diff --git a/CMakeLists.txt b/CMakeLists.txt index 3df1d82dbe09..3723a6e2e5ec 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -2,6 +2,16 @@ cmake_minimum_required(VERSION 3.14...3.28) # for add_link_options and implicit project("llama.cpp" C CXX) include(CheckIncludeFileCXX) +# C++26 where the compiler has it (amdclang 23, g++ 15); older toolchains keep C++17 +if("cxx_std_26" IN_LIST CMAKE_CXX_COMPILE_FEATURES) + set(CMAKE_CXX_STANDARD 26) +else() + set(CMAKE_CXX_STANDARD 17) +endif() +set(CMAKE_CXX_STANDARD_REQUIRED true) +# no C++20 module scanning: no module sources, and some toolchains lack clang-scan-deps +set(CMAKE_CXX_SCAN_FOR_MODULES OFF) + #set(CMAKE_WARN_DEPRECATED YES) set(CMAKE_WARN_UNUSED_CLI YES) diff --git a/benchmarks/NOTE-hrx-124-multipass-output-2026-09-27.md b/benchmarks/NOTE-hrx-124-multipass-output-2026-09-27.md new file mode 100644 index 000000000000..2b34b42c68bc --- /dev/null +++ b/benchmarks/NOTE-hrx-124-multipass-output-2026-09-27.md @@ -0,0 +1,91 @@ +# HRX decode-split multipass output pass: drop the 128x redundant expf and the scalar block loop (engine#124) + +Worktree: `~/1bit-engine-176/third_party/llama.cpp`, branch `1bit/hrx-124-output-pass` (base `00adc2b`). +Box: gfx1151 (Strix Halo), `Qwen3-Coder-30B-A3B-Instruct-Q4_K_M`, `-dev HRX0`. + +## What was wrong + +`@ggml.flash_attention.decode_split.reduce_completed.multipass` +(`ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/flash_attention_decode_split_f32_f16_wmma.loom`) +finished with a scalar output pass: `workitem -> one output channel`, then a serial +`scf.for %block = [%c0 to %active_block_count]` over **every** KV block, run by all 256 +workitems. Only 128 are live at `value_head_size = 128`, and each one recomputed +`scale = expf(partial_max[block] - maximum)` for every block, so `expf` ran once per +`(block, output element)` instead of once per `(block)`, and every element walked the +block dimension alone. That is the residual ~16% gap above capacity 2048 (issue #124). + +## The change (bit-exact) + +1. **LDS scale stage.** The lane-strided sum pass already computes + `scale = expf(partial_max[block] - maximum)` for every block; store it into a + per-row workgroup stage (`scale_stage_view[row, block]`) and also publish the row + `sum` (`sum_stage_view[row]`). The output pass then just reloads the exact f32 + value. Removes ~128x redundant `expf` and the redundant `partial_max` reloads. +2. **Vectorised all-rows output pass.** After the per-row max/sum phase, one pass + covers all query rows at once: workitem `w` owns a 4-channel *quad* + (`quad = tile*256 + w`, `row = quad / quads_per_row`, `channel = (quad % + quads_per_row)*4`), loads `vector<4xf16>` from `partial_output`, and accumulates + `vector<4xf32>` over blocks with `unroll(%c4) schedule(interleaved)`. For the + production GQA shape (8 query heads per KV head, `value_head_size = 128`) there are + exactly 256 quads, so the whole workgroup (all four subgroups) is busy and each + channel still accumulates its blocks in the same order as before. + +Because the per-channel accumulation order is unchanged and the vector ops are +elementwise, the reduce output is **bit-identical** to the previous multipass reducer +(verified end-to-end below). + +## Result + +`llama-bench -p 0 -n 8 -r 5 -d 1900,2000,2100,3000,4800`, 6 interleaved A/B rounds, +median of the 6 run means: + +| depth | capacity | blocks | base t/s | patched t/s | delta | path | +|---|---:|---:|---:|---:|---:|---| +| 1900 | 1920 | 30 | 70.72 | 71.48 | +1.1% | cooperative (untouched control) | +| 2000 | 2048 | 32 | 66.06 | 68.82 | +4.2% | cooperative (untouched control) | +| 2100 | 2112 | 33 | 58.19 | 67.15 | **+15.4%** | multipass | +| 3000 | 3008 | 47 | 50.84 | 60.89 | **+19.8%** | multipass | +| 4800 | 4864 | 76 | 41.83 | 51.02 | **+22.0%** | multipass | + +The boundary cliff is gone: `d2100/d2000` moves **0.881 -> 0.976** (and +`d4800/d2000` 0.633 -> 0.741). The two depths the reducer does not touch moved by ++1.1% / +4.2%, i.e. run-to-run noise. + +A throwaway diagnostic that deleted the output block loop entirely (keeping the +max/sum passes) measured **+21.6% / +50.0% / +30.5%** at d2100 / d3000 / d4800, which +is how the output pass was identified as the whole gap rather than the reduction math +(issue #124's diagnosis). + +## Evidence + +- **Bit-exact.** `llama-perplexity -c 2049 -b 1 -f <6600-token text> --save-all-logits` + (token-by-token decode, so the decode-split kernel is dispatched; `GGML_HRX_LOG_DISPATCH=1` + confirms `flash_attention_decode_split`). The 1,244,725,284-byte logits dumps for base + and patched are **byte-identical** (`cmp`), PPL 21.9760 +/- 0.81813 on both. +- **Correctness.** Buried code word `ZX-4718-QQ` at 4700 prompt tokens (capacity 4864, + multipass), `llama-server`, temperature 0, seed 42, `cache_prompt:false` -> answer + `ZX-4718-QQ`. PASS. +- **Faults.** 6 A/B rounds x 5-rep sweeps at all five depths on the patched build: + 0 `HSA_STATUS_ERROR_MEMORY_FAULT`, 0 `res = -3`. The cooperative/direct paths are + untouched, so <= 2048 (d1900/d2000) shows no tok/s or fault regression. + +## Reproduce + +``` +# build (engine worktree ~/1bit-engine-176, HRX llama.cpp build dir) +cmake --build build/hrx/llama --target llama-bench llama-perplexity -j 16 + +# dispatch needs libhsa from the HRX toolchain +export IREE_HAL_AMDGPU_LIBHSA_PATH=/opt/rocm-therock/lib/python3.14/site-packages/_rocm_sdk_core/lib/libhsa-runtime64.so.1 + +# performance +build/hrx/llama/bin/llama-bench -m models/Qwen3-Coder-30B-A3B-Instruct-Q4_K_M.gguf \ + -dev HRX0 -ngl 99 -p 0 -n 8 -r 5 -d 1900,2000,2100,3000,4800 + +# bit-exact reference vs the previous build +build/hrx/llama/bin/llama-perplexity -m models/Qwen3-Coder-30B-A3B-Instruct-Q4_K_M.gguf \ + -f ref-text.txt -c 2049 -b 1 -dev HRX0 -ngl 99 --save-all-logits logits.bin +``` + +Raw runs and the A/B driver used here live under the session scratch dir +(`ab5-*.json`, `ab5-summary.txt`, `diag-noloop.json`). diff --git a/benchmarks/NOTE-hrx-single-dispatch-verify-2026-09-26.md b/benchmarks/NOTE-hrx-single-dispatch-verify-2026-09-26.md new file mode 100644 index 000000000000..5acfbd5bba6d --- /dev/null +++ b/benchmarks/NOTE-hrx-single-dispatch-verify-2026-09-26.md @@ -0,0 +1,797 @@ +# HRX decode-split single-dispatch long reduce (engine#115) — build + first verification + +Worktree: `~/wt/hrx-fa-multipass3` (fork, base `fa1a456`). + +## Change +* `.loom` (`.../loom-libs/ops/flash_attention_decode_split_f32_f16_wmma.loom`): + both compose kernels' `index.assume %key_value_token_count` range widened + `1..2048` → `1..262144`; new **`reduce_completed.multipass`** (+ its + `template.decl`) and **`reduce_fused.multipass`** (`where capacity 2049..262144`), + both applied from inside the SAME `kernel.def launch` — same atomic completion + counter as `reduce_fused.cooperative`, so still ONE self-synchronising dispatch. + The multipass reducer is the `reduce_f32` math (lane-strided max pass; sum pass + that stores the per-block scale back into `partial_max`; per-channel output pass) + wrapped in a query-row loop, per KV head. +* `dispatch-flash-attention.cpp`: `kDecodeSplitMaxKeyValueTokenCapacity` + `2048` → `32768`. + +## Build +OK (after repairing the SDK wiring): configure with +`-Dhrx_DIR=$DEPS/libhrx/cmake/hrx -Dloomc_DIR=$DEPS/loom/binding/c/cmake/loomc +-DLLAMA_BUILD_TESTS=OFF` where +`DEPS=/home/bcloud/hrx-gfx1151/llama-b66/build/ggml/src/ggml-hrx/hrx/src/ggml-hrx-deps-build`, +plus `ln -s /tmp/hrx-main/libhrx/include $DEPS/libhrx/include` and the same for +`loom/binding/c/include`. `llama-server` and `llama-bench` link. + +## Verification so far (gfx1151, `-dev HRX0`) +* **4718-token repro** (`/tmp/repro_code_word.py 54`, capacity 4864): **PASS twice** + with the rigorous protocol (`temperature 0`, `seed 42`, `cache_prompt:false`) — + the model returns `ZX-4718-QQQ` (code word `ZX-4718-QQ`). The earlier + `47188QQ.ZX-4718QQ.`-style garbage was an artifact of KV-cache reuse + (`cache_prompt` default) with no seed, not the kernel. +* **≤2048** (54-section vs 20-section = 1742 tokens): the shipped cooperative path + alternates PASS/FAIL only on a one-character **case** slip (`ZX-4718-Qqq`), i.e. + the same model/reduction-order noise as the new path — not a regression. +* `llama-bench -p 0 -n 8`: `d1900 = 68.03`, `d2000 = 72.72` t/s. **`d2100` + (first capacity that selects `reduce_fused.multipass`) → `HSA_STATUS_ERROR_MEMORY_FAULT` + then `test_gen: failed to decode generation batch, res = -3`.** + +## Open blocker +GPU memory fault in `reduce_completed.multipass` (or a partials-buffer sizing +mismatch). Suspects: the reduce reading `active_block_count` blocks while the +produce wrote fewer / the view dims not matching the dispatch allocation; or an +out-of-range index in the lane-strided loops. Next: minimal repro at capacity 2112, +guard/dump the block indices, and diff the view shapes against `produce_partials`. + +Also note: `GGML_HRX_LOG_DISPATCH=1` logs each kernel once at load, so +decode-time selection above 2048 still needs a dedicated confirmation. + +## Isolation (added) +* `-d 2100` WITH `GGML_HRX_DISABLE_DISPATCH=flash_attention_decode_split` (fallback): + **works, 52.47 t/s**. WITHOUT (my multipass): **`HSA_STATUS_ERROR_MEMORY_FAULT` then + `res = -3`** — and there is no `all_rejected` / `select-templates` message, so this + is a genuine GPU page fault, not a dispatch-selection failure. +* Sizing is consistent: the dispatch allocates the partials from + `ceil_div(match.key_value_capacity, 64)` blocks and binds the loom config + `ggml.flash_attention.decode.key_value_token_capacity` to the same + `match.key_value_capacity` (dispatch-flash-attention.cpp:559, 612-613), so the launch's + producer count equals the partials' block count. Therefore the fault is **inside + `reduce_completed.multipass`**, not a buffer-size mismatch. +* Boundary: 32 blocks (capacity 2048, cooperative) OK; 33 blocks (capacity 2112, + multipass) faults. The ONLY things that change are the selected reducer variant and + the block count — so the bug is in the multipass reducer's 33+ block path. + +## Next debug step +Build a minimal repro at capacity 2112, then either (a) clamp the multipass reducer to +read only blocks the produce wrote (e.g. `active_block_count` vs the config-derived +`key_value_block_count`) and guard every index, or (b) bisect by replacing the +lane-strided loops with a single-block loop, rebuilding, and re-running +`llama-bench -d 2100`. Note `produce_partials`'s block guard is +`lt(workgroup_x, %key_value_block_count)` with `%key_value_block_count` from the CONFIG, +while `reduce_fused` receives `%producer_block_count` derived from +`%launch_key_value_token_capacity` — verify those two agree (off-by-one there would make +the reduce read one block past what the produce wrote). + +## Refined boundary (added) +`llama-bench -d 2048,2049` also fails (`res = -3`) with no `tg8 @ d` row. `-d 2000` +works. The reason: during the 8 generated tokens the KV count rises above the +requested depth, so `-d 2000` is the last depth whose whole generation stays at +`key_value_token_count <= 2048` (capacity 2048 = 32 blocks, cooperative), while +`-d 2048` and above hit 33+ blocks and select `reduce_fused.multipass`. +=> ANY decode step with `key_value_token_count > 2048` faults, i.e. the fault is +strictly in `reduce_completed.multipass`, and the very first >2048 block count +(33) is enough. + +## Bisect result — the fault is NOT in the reduce (added) +Two builds at `-d 2100`: +1. `reduce_completed.multipass` reduced to a minimal body (read block 0, store one + element): **still faults** (`res = -3`). +2. The `template.apply<...reduce_completed.multipass>` call removed entirely (produce + runs, no reduce at all): **still faults**. +=> The GPU page fault is **not** in `reduce_completed.multipass` (nor the reduce math, +query-row loop, or barriers). It is in the **produce / launch / allocation of the >2048 +compose path**. The compose kernel previously NEVER ran for `key_value_token_count > +2048` (its `index.assume` range was 1..2048, so the dispatch declined); widening it made +the kernel run there for the first time, exposing this. +Next: bisect the produce — check every view in `produce_partials` / +`produce_partials.active` (the partial buffers' block dim, and the key/value/mask +views) for a bound that does not scale (fixed 32, or the token count vs the buffer), +and check whether the kernel workload-parameter metadata for `key_value_token_count` +is still capped at 2048 in the manifest. + +## Control (c) inconclusive + design fork (added) +The launch-block-count clamp matched **3** sites (both compose kernels AND the +standalone `produce_partials_f32_f16_wmma` export), and clamping all three broke +llama-bench's warmup ("failed to run gen warmup") — so (c) must target only the compose +launch. Re-run it scoped. + +Design fork surfaced: the shipped compose kernel's inline `produce_partials` has never +run past 2048 and faults there for a still-unidentified reason (both bisects show it is +not the reduce). The alternative to widening `..._f32_f16_wmma` is a **dedicated +long-context compose kernel** that combines the standalone `produce_partials` (line ~976) +and `reduce_f32` (line ~999) templates into ONE `kernel.def launch` with the atomic +completion counter — i.e. exactly the maintainer's single-self-synchronising-dispatch bar, +but as a new kernel rather than a widened one. + +## Step (1) result (added) +Server at `-c 2400` + a ~2176-token prompt (N=25) FAILS: +`process_ubatch: failed to compute graph, compute status: -1` / `llama_decode: failed to +decode, ret = -3` — a GPU compute failure (not the bench's HSA memory-fault message). +So small capacities above 2048 (~2240) fail on the server too, while ~4736 (server with +`-c 8192`) passed. The fault is capacity-dependent: 2112 / 2240 fail, 4736 works. +Next: (a) confirm whether the 4718-token PASS actually used the split at all (compare +N=54 with vs without `GGML_HRX_DISABLE_DISPATCH=flash_attention_decode_split`); (b) find +what differs between a 33-38-block launch and a 74-block launch — e.g. transient/arena +sizing, or a workload-parameter range still declared `1..2048` in the compiled kernel +metadata (the .loom `index.assume` was widened, but the manifest / kernel-metadata range +was not checked). + +## CORRECTION (important - supersedes "Step (1) result" above) +The "Step (1) result" and the first sweep were run against a STALE binary: after control (c) +the `.loom` was restored and committed but never rebuilt, so the binary still contained the +hard-coded 32-workgroup clamp. Both results are void. + +After `ninja -C build-hrx llama-bench llama-server` on a clean GPU (no leftover servers): + +| depth | capacity | blocks | split ENABLED | split DISABLED | +|-------|----------|--------|---------------|----------------| +| 1900 | 1920 | 30 | OK 58.11 t/s | - | +| 2000 | 2048 | 32 | OK | - | +| 2049 | 2112 | 33 | FAULT (HSA_STATUS_ERROR_MEMORY_FAULT, surfaces on the prompt batch) | OK 54.05 | +| 2100 | 2112 | 33 | FAULT | OK 44.52 | +| 2500 | 2560 | 40 | FAULT | OK 39.75 | +| 3000 | 3008 | 47 | OK 44.33 | - | +| 4736 | 4736 | 74 | OK (server repro) | - | + +So: the split IS the cause (disabling it makes every faulting depth pass), the fault is NOT +monotonic in block count, and it is confined to a BAND: 33 and 40 blocks fault; 30, 32, 47 and +74 pass. Two bisects already excluded the reduce. Next experiment: sweep 41/44/46/48 to find the +exact upper edge of the fault band - that number identifies the residual constant (a launch or +transient sizing bound) in the newly-exercised >2048 compose path. The "capacity-dependent +2112/2240-fail-4736-works" conclusion above is wrong; the fault is a narrow block-count band. + +## CORRECTION 2 - the fault is NOT a llama-bench artifact, and the N=54 pass does not reproduce +Server at `-c 8192` with a ~2199-token prompt (N=25, capacity 2240) ALSO faults: + Warning: Queue error - HSA_STATUS_ERROR_MEMORY_FAULT + E graph_compute: wait for HRX graph replay commands failed ... AMDGPU memory access fault + at device address 0x00007f508ea94000 (reason mask 0x00000001) + E process_ubatch: failed to compute graph, compute status: -1 + E srv decode: Compute error. off = 0, n_batch = 2048, ret = -3 (n_tokens = 2199) +So the fault is real, reproduces on the server, and is NOT explained by +"padded capacity > key->ne[1]" (here capacity 2240 << KV length 8192), nor by llama-bench's +tight n_ctx. The earlier N=54 (4700-token, capacity 4736) PASS is NOT reproducible with the +current rebuilt binary -> the rebuild after the control-(c) clamp changed behaviour. DO NOT +trust the pre-clamp results; first re-establish a known-good baseline (clean checkout of the +multipass commit, or re-apply the multipass implementation) and re-test N=54 vs N=25. + +Failing set by block count: 33 and 40 fault; 30, 32, 47 and 74 pass (30/32/47 measured with the +rebuilt binary; 74 from the pre-clamp N=54 run - unverified). Reduce excluded by two bisects. +The fault address (0x7f508ea94000) and the create at n_batch=2048 show it is a real AMDGPU page +fault inside the graph, most likely a produce/partial OOB in the newly-exercised >2048 path. + +## Dispatch transient sizing is linear in block count (no static bound) +`match_flash_attention_decode_split_next_q8_dispatch`: + key_value_block_count = ceil_div(match.key_value_capacity, 64) + partial_scalar_count = kv_head_count * key_value_block_count * kDecodeRowCapacity(16) + partial_scalar_bytes = partial_scalar_count * 4 = 256 * blocks + partial_output_bytes = partial_scalar_count * value_head_size(128) * 2 = 16384 * blocks + q8_output_bytes = q8_1_x4_byte_count(query_token_count, output_hidden_size) + transients: partial_max/partial_sum -> partial_scalar_bytes; partial_output -> partial_output_bytes; + q8_output -> q8_output_bytes; completion_counter -> kv_head_count i32 (requests) + launch workload param: kernel.integer_parameters["key_value_token_count"] = match.key_value_token_count + (= mask->ne[0]); compile param key_value_token_capacity = match.key_value_capacity (padded) + bindings: key/value bound at the whole buffer byte_count; mask bound per-row at + row*mask->nb[1] with bytes = key_value_token_count*2 + +All sizes are strictly linear in `blocks` - there is NO residual constant/step in the transient +sizing, so the 33-46 band fault is NOT a static allocation bound. It is inside the kernel's own +memory access in the newly-exercised >2048 path. Prime suspects, in order: + (1) an `index.assume` promise violated at runtime (UB/miscompile) - loom lines 188 and 319, + `%score_key_end = assume [le(%score_key_end0, %bounded_key_value_token_count)]`, i.e. the + produce's per-tile end bound, and the block-end clamping for the last partial block; + (2) the produce's K/V tile load for the final block (tokens beyond the token count but within + the padded capacity) - only safe if the KV buffer is padded past the capacity; + (3) the per-row `mask` binding vs the mask view extent. +Next experiment: build a minimal produce that reads ONLY the last block (block_count-1) at 33 +blocks and see if it faults; then bisect the tile-end clamp. Baseline caveat: the rebuilt binary +no longer reproduces the pre-clamp N=54 pass, so re-establish a known-good baseline first. + +## BAND PINNED: capacity 2112..2560 (blocks 33..40) FAULTS; <=2048 and >=2624 are clean +Separate-process sweep, rebuilt binary (each depth its own llama-bench run): + d1900 cap1920 30 blocks -> OK + d2000 cap2048 32 blocks -> OK + d2049 cap2112 33 blocks -> FAULT (surfaced on test_prompt) + d2100 cap2112 33 blocks -> FAULT (test_gen) + d2500 cap2560 40 blocks -> FAULT + d2600 cap2624 41 blocks -> OK 43.39 <-- first working block count above the cap + d2800 cap2816 44 blocks -> OK 44.87 + d2900 cap2944 46 blocks -> OK 42.68 + d3000 cap3008 47 blocks -> OK 44.33 + d3060 cap3072 48 blocks -> OK 42.94 +=> the fault band is EXACTLY capacity 2112..2560 (blocks 33..40). 64*32=2048 (last PR#9 value) + works, 64*33..64*40 (2112..2560) fault, and 64*41=2624 upward works again. This starts exactly + where PR #9's cap ended, which is why it was never seen before. + +Correlate: partial_output transient = 16384*blocks bytes -> 528 KiB @33, 640 KiB @40, 656 KiB @41; +partial_max/sum = 256*blocks -> 8448 B @33, 10240 B @40, 10496 B @41. The 640 -> 656 KiB step is +the only discontinuity candidate near the band edge. Note the fault ALSO surfaced on the fallback +path for d2049 (test_prompt) and on the server prefill (n_batch=2048, N=25), so it may be a small +OOB whose address only leaves the allocation for those sizes, rather than a pure split-path bug - +but `GGML_HRX_DISABLE_DISPATCH=flash_attention_decode_split` made every faulting depth pass. + +Next: (1) instrument/log the transient arena offsets and the exact faulting address across the band +edges (2048 / 2112 / 2560 / 2624) to see which buffer the address sits past; (2) bisect the produce +with only the last block (block_count-1) active at 33 blocks; (3) re-check the K/V tile load for the +final partial block (tokens beyond mask->ne[0] but inside the padded capacity). + +## LEAD: two independent block counts in the same compose kernel (next thing to test) +The compose kernel derives its block count TWICE, from two different quantities: + loom:95-96 %active_key_value_block_count = div(pad(KEY_VALUE_TOKEN_COUNT), 64) -> %last_block_ordinal + (counter threshold for the "last partition" / who runs the reduce) + loom:1084 %key_value_block_count = div(pad(CONFIG CAPACITY), 64) + (the launch: workgroups(%key_value_block_count, kv_head_count, 1)) +The produce's partial views are sized by the SECOND (loom:494-496, +view<[kv_heads]x[%key_value_block_count]x16x...>), while the reduce_fused apply at +loom:1096/1122 passes `(%launch_key_value_token_capacity, %producer_block_count, %producer_block_count, ...)` +- i.e. partial_block_capacity == producer_block_count (a THIRD quantity, whose definition I have not +yet located; it is NOT %key_value_block_count textually). + +In the dispatch (dispatch-flash-attention.cpp:348-350) key_value_token_count = mask->ne[0] and +key_value_capacity = ceil_div(mask->ne[0],64)*64, so the two pad to the same value in the normal +case. IF they ever disagree (or if producer_block_count is taken from the wrong one of the pair), +the completion counter's threshold no longer matches the number of workgroups that actually +increment it: the reduce fires EARLY (before all producers wrote their partials) or NEVER. An early +fire has the last-arriving workgroup read partial rows whose block index exceeds +partial_block_capacity -> out-of-bounds read past partial_output -> exactly the observed +HSA_STATUS_ERROR_MEMORY_FAULT, and it would be band-sensitive because the mismatch magnitude is +(pad(capacity)/64 - pad(token_count)/64), which is 0 for most depths but non-zero for some. + +NEXT (concrete): (1) locate the definition of %producer_block_count in the loom (it is the +partial_block_capacity passed to reduce_fused - if it is derived from the token count while the +launch uses the capacity, that is the bug); (2) log all three at dispatch time for the band depths +and assert they are equal; (3) the two bisects excluded reduce_completed but BOTH kept the +reduce_fused wrapper (counter + pack_completed_q8), so the wrapper is still suspect. + +## Diagnostic that could crack the band in one run (added) +The fault depends on the CAPACITY VALUE (2112..2560), not on the token count: N=54 (KV starts at +4700) passes and N=25 (2199) fails on the SAME build with the SAME -c 8192. Two cheap diagnostics: + (1) ROUND UP: patch dispatch-flash-attention.cpp so + match.key_value_capacity = max(ceil_div(key_value_token_count,64)*64, 2624) + i.e. never select capacity 2112..2560. If every band depth then PASSES, the bug is a + capacity-VALUE-dependent codegen/allocation issue (the kernel is JIT-specialised per constant + capacity) rather than a function of the actual tokens. This is one edit + one rebuild. + (2) ARENA: the split's transients add ~0.6-1.2 MB to the graph arena; disabling the split removes + them and every fault disappears. Check whether the HRX graph/transient planner sizes the arena + correctly at those capacities - the sizes it is handed are exact and linear + (partial_scalar_bytes=256*blocks, partial_output_bytes=16384*blocks, alignment 256, + counter = kv_head_count i32). +Also unresolved: the loom checks exercise only 32 and 512 blocks, so 33..40 blocks is compiled but +verified by nothing. + +## n=1 CONTROLLED PROBE: the fault tracks `key_value_token_count % 64 == 0`, not the buffer length +Same capacity (2112), same block count (33), same JIT-specialised binary (compiled per constant +capacity), only the runtime workload parameter differs: + llama-bench -p 0 -n 1 -d 2110 -> kv=2111, tail=63 -> OK 17.06 t/s [tail branch] + llama-bench -p 0 -n 1 -d 2111 -> kv=2112, tail=0 -> FAULT res=-3 [full branch] + +This FALSIFIES both earlier theories: + - "padded capacity > key->ne[1]": at d2111 the padded cap (2112) EQUALS kv (2112) and it faults; + at d2110 the cap (2112) EXCEEDS kv (2111) and it passes. Exactly inverted. + - "partial last tile over-read": the PARTIAL-tile case is the one that PASSES. + +The only in-kernel difference is (loom:93-96) + %has_no_tail = (rem(bounded,64) == 0) + %is_full_block = %has_no_tail OR (block_ordinal != last_block_ordinal) +kv=2112 -> has_no_tail=true -> the LAST block (ordinal 32) takes the FULL 64-token branch. +kv=2111 -> has_no_tail=false -> that block takes the TAIL branch. +Same launch (33 workgroups), same buffers, same code. So the fault comes from the FULL branch +running on the last active block. + +CAVEAT (important): this is not universal. d3000 (-n 8) reaches kv=3008 (3008%64==0, 47 blocks, +full branch on the last block) and PASSES. So the trigger is the conjunction +(full branch on the last active block) x (block count in 33..40). The d2049/d2100/d2500 (-n 8) +failures surfaced with kv%64 != 0 (and d2049 failed on test_prompt, i.e. the fill/prefill rather +than the split decode), so those may be a second or compound phenomenon - re-measure every depth +with -n 1 to separate them. + +NEXT: (1) confirm the n=1 pair repeats; (2) probe -n 1 in pairs (64k-1, 64k) across and outside the +band - (2047,2048), (2111,2112), (2559,2560), (2623,2624), (3071,3072) - the pattern of which k +fault is the signature; (3) prime suspect is the produce's is_full_block / full-tile branch for the +last active block. + +## n=1 SWEEP: the fault is INTERMITTENT - "has_no_tail" and the clean 33..40 band are FALSIFIED +All seven kv==0 (mod 64) depths PASSED at -n 1, including d2111 which FAULTED at -n 1 an hour +earlier on the same binary: + k=32 kv=2048 d2047 OK | k=33 kv=2112 d2111 OK 17.56 (previously FAULT) | k=34 kv=2176 d2175 OK + k=40 kv=2560 d2559 OK | k=41 kv=2624 d2623 OK | k=47 kv=3008 d3007 OK | k=48 kv=3072 d3071 OK +So the fault is NOT a deterministic function of (kv%64, block count). It is INTERMITTENT. +Re-reading every measurement with that lens: + - EVERY -n 1 run has ever passed (1 decode step = 1 dispatch of the split). + - The -n 8 failures were all at capacity 2112 (d2049, d2100) and 2560 (d2500); every -n 8 run at + capacity >= 2624 passed. -n 8 = 8 decode steps = 8 split dispatches. +=> the fault needs MANY dispatches at those capacities, which points at state that persists ACROSS +dispatches rather than one bad access. PRIME SUSPECT: the completion counter transient +("common.decode.flash_attention.completion_counter", kv_head_count*4 bytes, allocated in the graph +arena). The protocol leaves it at 0 after a dispatch, so it is only correct if (a) it really is 0 at +graph start and (b) it is never aliased by another dispatch's transient in the arena. Either failure +makes old_counter match last_block_ordinal at the wrong moment -> the last-arriving workgroup runs the +reduce against unwritten partials -> garbage addresses -> the observed AMDGPU page fault. This ALSO +explains why bisect #2 (removing reduce_completed) did NOT stop the fault: both bisects kept the +reduce_fused wrapper and its counter protocol intact. +NEXT: (1) find whether the HRX dispatcher zero-initialises the completion_counter transient and +whether the arena planner can alias it across dispatches; (2) re-run the -n 8 band depths 5x each to +measure the per-run failure rate at capacity 2112/2560 vs 2624 (intermittency => a race, not a bound). + +## SMOKING GUN: the completion counter is NEVER zero-initialised +`dispatch/transient-allocator.cpp :: add_completion_counter_allocations` gives each completion counter +its own allocation appended at the TAIL of the arena: + allocation.alignment = 16; + allocation.arena_offset = align_up(plan.arena_size, allocation.alignment); + plan.arena_size = allocation.arena_offset + allocation.size; + completion_counters.byte_count = plan.arena_size - completion_counters.arena_offset; +so the counter region runs to the end of the arena. A grep for `memset`/`zero` over +transient-allocator.cpp and command-program.h returns NOTHING (only the string "has zero counters"). +=> nothing ever writes 0 into the completion counter. Correctness rests entirely on the kernel +protocol leaving it at 0 (each workgroup's release atomic add of -count at the end of the dispatch). +Anything that prevents that subtraction from completing - or any graph whose arena tail overlaps a +different graph's live data (the overlap check at transient-allocator.cpp:312 is per-graph only, and +the arena buffer is reused across graph computes) - leaves a stale non-zero counter. + +This consequence matches every observation: + - old_counter != 0 shifts WHICH workgroup sees old_counter == last_block_ordinal, so the reduce runs + against unwritten partials (garbage -> AMDGPU page fault) or never runs (wrong output); + - it is INTERMITTENT (depends on the arena's stale contents at that offset); + - it needs several dispatches to appear - exactly the -n 8 (fails at capacity 2112/2560) vs -n 1 + (always passed) split; + - it is arena-layout-sensitive, so only some capacities misbehaved; + - it explains why BOTH bisects failed to remove the fault: both kept the reduce_fused wrapper and + therefore the counter protocol. +NEXT: test by explicitly zeroing the counter (or pre-clearing the arena tail) before each dispatch and +re-running the -n 8 depths at capacity 2112/2560 several times to compare the failure rate before/after. + +## NEGATIVE RESULT: the multipass wrapper's counter protocol is IDENTICAL to the cooperative's +Compared loom ~1004-1030 (reduce_fused.cooperative) against ~1050-1076 (reduce_fused.multipass): +the increment (`view.atomic.rmw +1`, acq_rel/device, by workitem 0), the workgroup barrier, the +scratch load, `last_block_ordinal = producer_block_count - 1`, +`is_last_partition = (old_counter == last_block_ordinal)`, the global release barrier, and the single +`view.atomic.reduce %negative_key_value_block_count_i32` - all inside `scf.if %is_last_partition` +and guarded by `%workitem_is_zero` - are IDENTICAL. The ONLY difference between the two defs is which +`reduce_completed.*` template is applied inside that branch. +=> the completion-counter protocol is NOT the cause of the multipass-only fault. The never-zeroed +counter (previous section) is a real latent hazard but is NOT the trigger here, because the +cooperative path uses the exact same protocol and is correct at <=2048. +With the two bisects also excluding reduce_completed, this leaves an apparent contradiction: +produce == produce, wrapper == wrapper, reducer excluded - yet the path faults only above 2048. +Remaining explanations: + (a) the fault is not multipass-specific but CAPACITY-specific: the counter region is the TAIL of the + arena and its offset moves with the capacity, so a stale/aliased tail is capacity-dependent and + the <=2048 depths simply never landed on the bad layout; + (b) a Loom select-templates / JIT codegen defect for the multipass def at particular constant + capacities (the loom checks cover only 32 and 512 blocks). +NEXT: (1) re-run the -n 8 depths 5x each to measure the failure rate (intermittent vs deterministic); +(2) test (b) directly with a semantically-neutral change: widen reduce_fused.multipass's where from +range(2049, 262144) to e.g. range(2049, 1048576) - if the fault moves or vanishes, it is a +compiler/selection issue, not an access bug. + +## DECISIVE: the fault is a RACE (~4/5 failure at d2100 -n 8; run1 passed at 50.23 t/s) +Five separate `llama-bench -p 0 -n 8 -r 1 -d 2100` processes, same binary: + run1 PASS 50.23 t/s | run2 FAIL | run3 FAIL | run4 FAIL | run5 FAIL +=> NOT deterministic, and run1 proves the multipass path DOES work - at 50.23 t/s, essentially the +<=2048 speed, so the cliff really is removable. Combined with "every -n 1 run has ever passed", this is +the signature of a RACE needing several dispatches to appear. This FALSIFIES the deterministic +"33..40 block band" and "has_no_tail/is_full_block" readings (both were sampling artefacts). + +The completion counter is the only cross-dispatch state in this path: + - it is never zero-initialised (transient-allocator.cpp::add_completion_counter_allocations appends it + at the arena TAIL and nothing memsets it - see the earlier section), and + - the protocol assumes it is exactly 0 on entry and returns to exactly 0: + `last_block_ordinal = producer_block_count - 1` is only observed if the counter starts at 0. +If any dispatch leaves it non-zero (negative over-decrement, or a workgroup racing to an ordinal when +the counter starts negative), the "last partition" fires EARLY - before all producers wrote their +partials - and the reduce reads unwritten partial data. A persistently wrong counter also explains why +one bad dispatch POISONS the following decodes in the same process. + +NEXT: (1) make the counter start at 0 unconditionally (zero it before each dispatch, or pre-clear the +arena tail) and re-run -n 8 at d2100 five times - if all five pass, the race is confirmed and that is +the fix; (2) audit the decrement: every produce workgroup increments by 1 (N total) and only the +last-partition workgroup decrements by N; verify no workgroup can skip the increment (e.g. one masked +out by %block_has_attention, which skips the produce) while the count still says N. + +## FIX LOCATION: no fill/zero CommandKind, but ConstantInitialization exists +command-program.h: `enum class CommandKind { Invalid, Kernel }` - there is NO memset/fill command, so a +runtime "zero the counter before every replay" fix would need new plumbing. BUT CommandProgram has + std::vector constant_initializations; // {ValueId value; string name; size_t offset; vector data} +populated from `plan.constant_initializations` (dispatch-scheduler.cpp:253) and consumed in the resolver +(command-program.cpp:475-477, 521). `initialization_commands` likewise come only from +`plan.initialization_dispatches` (kernel dispatches). CommandProgram also already carries +`CompletionCounterPlan completion_counters {arena_offset, byte_count, count}`, so the runtime knows the +counter region exactly. + +=> CHEAPEST FIX: in match_flash_attention_decode_split_next_q8_dispatch (dispatch-flash-attention.cpp), +next to the existing + dispatch_match.completion_counter_requests.push_back({completion_counter, "...", kv_head_count}); +also register a constant initialization of kv_head_count*4 zero bytes for the completion_counter value, +so the region is zeroed once at program init. That is correct with no new command kinds IF the +N/-N +protocol is balanced (it returns the counter to its entry value, so entry==0 keeps it 0 forever). +Verify by re-running d2100 -n 8 five times (currently PASS,FAIL,FAIL,FAIL,FAIL). + +DECIDING QUESTION: re-derive whether each of the N per-head produce workgroups increments exactly once +and whether exactly one workgroup decrements by N. If balanced, once-at-init suffices; if not, the +decrement must be made unconditional (e.g. store 0 instead of add -N) so a stale counter cannot persist. + +## METHODOLOGICAL CORRECTION: both bisects are INVALID (single-run verdicts under an ~80% failure rate) +The fault is a RACE: 5 separate `llama-bench -p 0 -n 8 -d 2100` processes gave PASS,FAIL,FAIL,FAIL,FAIL. +A single run per bisect therefore CANNOT distinguish "the fault is gone" from "this run happened to +pass". Bisect #1 (minimal reduce body) and bisect #2 (reduce_completed apply removed) were each judged +on ONE run each, both of which faulted - but under an ~80% per-run failure rate, "it faulted once" is +the EXPECTED outcome even if that change fixed the fault completely (P(fault) = 0.8). +=> both bisects must be re-run with a race-aware protocol: N runs per variant, verdict = failure RATE. +=> reduce_completed.multipass is NOT excluded. It is the PRIME SUSPECT again, because it is the ONLY + difference between the correct <=2048 cooperative path and the failing >2048 path: the produce is + the same code, the counter wrapper is byte-identical (loom 1050-1076 vs 1004-1030), and only the + applied reduce_completed template differs. + +RACE-AWARE NEXT STEP (verdict = failure rate over 5 runs of `-p 0 -n 8 -d 2100`): + (a) current code -> expect ~4/5 fail (baseline) + (b) no-op reduce_completed.multipass body (insert `template.return` as its first statement) + (c) reduce_fused.multipass forced to call reduce_completed.cooperative (only valid <= 32 blocks, so + use it as a diagnostic: it will produce WRONG numbers but if the fault vanishes the race is in + my multipass reducer) +Keep /tmp/wmma.pre-bisect.loom before editing and rebuild with `ninja -C build-hrx llama-bench`. +Candidate race to look for while doing (b)/(c): in reduce_completed.multipass the SUM pass runs only in +`%is_first_subgroup` and writes the per-block scale back into partial_max_view, while the output pass +(all workitems, serial `scf.for %block = [%c0 to %active_block_count step %c1]`) reads it - verify the +workgroup barrier between them covers the scale stores, and that the reduce (which runs inside the +LAST-ARRIVING produce workgroup) is guaranteed to see every other workgroup's partial writes given the +release fence is `kernel.barrier scope(workgroup)` while the counter RMW is acq_rel/scope=device. + +## DECISIVE race-aware bisect: the reduce is EXCLUDED; the fault is in the PRODUCE +Protocol: 5 runs of `llama-bench -p 0 -n 8 -d 2100` per variant; verdict = FAILURE COUNT (not a single run). + (a) current code -> 1,0,1,1,1 = 4/5 FAULT (baseline) + (b) reduce_completed.multipass made a NO-OP (inserted `template.return` as its first body statement, + after loom line 832) -> 1,1,1,1,1 = 5/5 FAULT +Removing the ENTIRE multipass reduction does NOT reduce the fault rate (it rises). The reduce is +EXCLUDED with proper statistical power; the two earlier single-run "bisects" were invalid because the +fault is a race. +=> The fault is in the PRODUCE (or the launch/allocation) for >32 producer blocks. With the reducer a + no-op the remaining path in every workgroup is: produce -> counter increment -> (no reduce) -> + conditional decrement. So the race is in the produce, or in the increment/decrement sequence around it. + (The no-op variant writes garbage output, so this is a fault-RATE test only; llama-bench measures + speed, not correctness.) + +REMAINING HYPOTHESIS: memory ordering between the produce's partial writes and the counter signal. +Every workgroup's 256 workitems write partial_max/sum/output to GLOBAL memory, then workitem 0 does +`view.atomic.rmw` on completion_counter with {acq_rel, scope=device}, preceded only by +`kernel.barrier scope(workgroup) ordering(release)`. Because the reduce runs in a DIFFERENT +workgroup, that release/acquire chain must carry the partial writes across workgroups. If the release +fence's scope(workgroup) does not publish the other workitems' global writes to other workgroups, the +reduce reads unwritten partials - and, more importantly for a PAGE FAULT, the winner workgroup may +observe partial_output before it is allocated/written in that layout. This is identical in the +cooperative path, so the >32-block difference must come from the produce's own addressing/timing. +NEXT: (1) keep the reduce no-op'd and experiment with the ordering (e.g. device-scope release before +the counter RMW); (2) ALWAYS use 5 runs per variant - never single-run verdicts; (3) restore the reduce +once the race is found. +NOTE: /tmp/wmma.pre-bisect.loom is the pre-bisect .loom (restored below). + +## NEGATIVE: zeroing the completion counter does NOT fix the race +Implemented the counter init at dispatch-flash-attention.cpp:595 (patched, builds clean): + dispatch_match.constant_initializations.push_back({ + completion_counter, "common.decode.flash_attention.completion_counter", 0, + std::vector(kv_head_count * sizeof(int32_t), 0) }); +Result, 5 runs of `llama-bench -p 0 -n 8 -d 2100`: **1,1,1,1,1 = 5/5 FAULT** - no improvement over +the 4/5 baseline. => the stale-counter hypothesis is FALSIFIED (or the init never reaches the +transient - worth one check with the command-program dump if this thread is pursued). + +CURRENT EXCLUSION STATE (every verdict under the 5-run protocol): + - reduce_completed.multipass: EXCLUDED (no-op'd -> 5/5 fault, baseline 4/5). + - counter staleness: not the fix (5/5 with the counter zeroed at init). + - counter protocol SHAPE: identical to the working cooperative wrapper (loom 1050-1076 vs 1004-1030). + - static sizing: transient bytes are exactly 256*blocks and 16384*blocks (linear, no bound). +REMAINING SUSPECTS: + (1) the produce's own memory accesses at >32 blocks (it is the only code left running once the + reduce is no-op'd, besides the counter sequence); + (2) the arena/layout - the counter is appended at the arena TAIL and plan.arena_size is then + extended to cover it; verify the arena is really allocated at that final size and that the + counter region lies inside it (an out-of-arena counter is an OOB write); + (3) memory ordering between the produce's global partial writes and the counter RMW - the release + fence is scope(workgroup) while the reduce runs in a DIFFERENT workgroup. +PROTOCOL RULE (must be kept): the fault is a race with a ~80-100% per-(-n 8)-run rate, so ANY claim of +"fixed" or "not the cause" from a single bench run is invalid. Always use >= 5 runs and report rates. + +## CRITICAL CORRECTION: my counter test was INVALID - the ConstantInitialization patch broke everything +With the patch REVERTED and a clean rebuild: + d1900 (capacity 1920, <=2048 cooperative) 5 runs -> 0,0,0,0,0 = 0/5 FAULT (CLEAN, no regression) + d2100 (capacity 2112, multipass) 5 runs -> 1,1,1,1,1 = 5/5 FAULT +WITH the patch in the tree, d1900 ALSO gave 5/5 (and d4800 too). So that patch had a GLOBAL side +effect: registering the counter's `constant_initializations` entry corrupted/perturbed even the shipped +<=2048 path. Consequences: + 1. The earlier "counter init does not fix the race (5/5)" conclusion is INVALID - the measurement was + contaminated. The stale-counter hypothesis is NOT falsified; it was never actually tested. + Lesson: `constant_initializations` is NOT a safe way to write to a transient - do not use it again. + 2. <=2048 is confirmed clean at 5 runs (d1900 = 0/5), which satisfies the no-regression criterion. + 3. >2048 is ~100% failing at d2100 (5/5) with this build. + +COUNTER ARITHMETIC (re-derived, so the fix can be reasoned about): +each of the N per-head workgroups does `rmw addi +1` (counter goes 0->N); the workgroup whose returned +old_counter == N-1 fires the reduce and then adds -N, returning the counter to 0. This is CORRECT ONLY +IF the entry value is exactly 0. If the entry is NEGATIVE (C<0) then N-1 lies inside [C, C+N-1], so the +"last partition" fires EARLY - after only N-1-C increments, i.e. C workgroups short - and the reduce +(and pack_completed_q8) consume partials that were never written. If the entry is POSITIVE the reduce +never fires at all. Nothing in the runtime ever writes 0 here, so the entry value is whatever the arena +held at that offset. + +SAFE FIX TO TEST NEXT (cannot break other paths): in the two reduce_fused wrappers, replace the +last-partition reset + view.atomic.reduce %negative_key_value_block_count_i32, %completion_counter_view[%key_value_head] {ordering = release, scope = device} +with a plain release-ordered store of 0: + view.store %c0_i32, %completion_counter_view[%key_value_head] +The dispatch boundary serialises dispatches, so no other workgroup touches the counter at that point, and +the counter is then exactly 0 after EVERY dispatch regardless of its entry value (self-healing). If the +entry was ever negative, this removes the early fire after the first dispatch. Verify with >=5 runs at +d2100 (currently 5/5 fault) and re-check d1900 stays 0/5. + +## BREAKTHROUGH: the >2048 path WORKS except exactly capacity 2112 (33 blocks) +5-run protocol, rebuilt binary (ConstantInitialization patch reverted; the counter reset is now a plain +release-ordered store of 0 - semantically safe, d1900 stays 0/5): + d1900 cap1920 30 blocks -> 0/5 fault CLEAN (no regression at <=2048) + d2000 cap2048 32 blocks -> 0/5 fault CLEAN + d2100 cap2112 33 blocks -> 5/5 FAULT <-- the ONLY failing point found + d3000 cap3008 47 blocks -> 0/5 fault CLEAN + d4800 cap4864 76 blocks -> 0/5 fault CLEAN +AND the headline success case now runs clean: + llama-server -dev HRX0, N=54 -> prompt 4700 tokens, capacity 4864 + -> answer 'ZX-4718-Qqq' (buried word 'ZX-4718-QQ'), gpu_faults = 0 + i.e. the 4718-token case decodes with NO GPU fault and retrieves the code word with a single + CHARACTER-CASE slip - exactly the model/reduction-order noise the shipped <=2048 path already shows + (N=20 at 1742 tokens alternates PASS / 'ZX-4718-Qqq'). NOT a kernel error. + +=> the multipass implementation is FUNCTIONAL across most of the >2048 range. The sole remaining defect + is capacity 2112 = 33 producer blocks, which is EXACTLY the lower bound of my + `reduce_completed.multipass` where-clause `range(%partial_block_capacity0, 33, 4096)` and of + `reduce_fused.multipass`'s assume `range(%partial_block_capacity0, 33, 4096)` / + `range(%key_value_token_capacity, 2049, 262144)`. That boundary is now the prime suspect + (select-templates/assume bound at exactly 33), not the reduce math and not the counter. + +NEXT (highest value first): + (1) relax the boundary: raise `reduce_completed.cooperative`'s where/assume upper bound from 32 to 4096 + and lower `reduce_completed.multipass`'s lower bound from 33 to 1 (same for reduce_fused.multipass's + assume), rebuild, then run d2100 5x - if the fault disappears, the defect is the exact-33 bound. + (2) confirm -n 1 vs -n 2 vs -n 8 at d2100 (does the 2nd dispatch break?) to decide race vs bound. + (3) then run the full contract: 4718-token repro on HRX0 *and* the HRX0/Vulkan0 split, the ~3587-token + case, `llama-bench -p 0 -n 8 -d {1900,2000,2100,3000,4800}` (1900/2000/3000/4800 already 0/5), + and `GGML_HRX_LOG_DISPATCH=1` to confirm the decode-split (not the fallback) is selected above 2048. + +## DECISIVE LOCALIZATION: capacity 2112 (33 blocks) + a TAIL at block 32 => fault; needs only ONE dispatch +d2100 (cap 2112, 33 blocks) vs number of decode steps, current build: + -n 1: 1 | -n 2: 1 | -n 3: 1 | -n 4: 1 | -n 8: 1 (every one FAULTS) +So at capacity 2112 the fault needs only ONE dispatch - it is NOT a multi-dispatch race. (The <=2048, +3008 and 4864 paths are all 0/5 clean at -n 8, five runs each.) Re-reading every earlier measurement +with this: + d2100 kv=2101 cap2112 tail=53 -> FAULT (5/5 over -n 1..8, and 5/5 at -n 8 earlier) + d2111 kv=2112 cap2112 tail=0 -> PASSED twice, faulted once + d2000 kv=2008 cap2048 tail=24 -> 0/5 CLEAN + d3000 kv=3001 cap3008 tail=57 -> 0/5 CLEAN + d4800 kv=4801 cap4864 tail=1 -> 0/5 CLEAN +=> the trigger is the conjunction (capacity 2112 == 33 producer blocks) AND (the last block takes the + TAIL branch: has_no_tail == false, so %is_full_block is false for block_ordinal 32). The same capacity + with has_no_tail (kv == 2112 exactly) passes; a tail at 32, 47 or 76 blocks is fine. + 33 is EXACTLY the lower bound of `reduce_completed.multipass`'s where-range + `range(%partial_block_capacity0, 33, 4096)` and of the assume in `reduce_fused.multipass`. That + boundary is now the prime suspect - most plausibly the produce's TAIL branch at block_ordinal 32 / + the select-templates bound at exactly 33 - NOT the reduce math (no-op bisect) and NOT the counter + (reset-to-0 patch made no difference; both were measured, d1900 stayed 0/5 throughout). + +NEXT (highest value first): + (1) -n 1 at d2103..d2112 to confirm exactly which kv values fault at capacity 2112 (expect: the tail + ones fault, kv=2112 passes). + (2) Make the boundary non-coincident: raise `reduce_completed.cooperative`'s upper bound from 32 to + 4096 and lower `reduce_completed.multipass`'s lower bound from 33 to 1 (and the + `reduce_fused.multipass` assume likewise); rebuild; run d2100 5x. If the fault vanishes, the + exact-33 bound is the defect. + (3) Re-run the contract once fixed: 4718-token repro on HRX0 and the HRX0/Vulkan0 split (the HRX0 one + already gives 'ZX-4718-Qqq', gpu_faults=0), `llama-bench -p 0 -n 8 -d {1900,2000,2100,3000,4800}` + (1900/2000/3000/4800 already 0/5), and GGML_HRX_LOG_DISPATCH=1 to confirm the split - not the + fallback - is selected above 2048. + +## NEGATIVE: relaxing the exact-33 assume bound does NOT fix capacity 2112 +Replaced every `33, 4096` with `1, 4096` in the multipass where/assume clauses (def selection at 33 blocks +stays identical because reduce_completed.cooperative remains 1..32), rebuilt, ran the 5-run protocol: + d2100 (cap 2112) -> 1,1,1,1,1 = 5/5 FAULT (unchanged) + d1900 (cap 1920) -> 0,0,0,0,0 = 0/5 (still clean) +So the exact-33 lower bound is NOT the defect. Reverted. + +CONSOLIDATED STATE (every verdict 5 runs, current build): + CLEAN: d1900 (cap 1920, 30 blocks), d2000 (cap 2048, 32), d3000 (cap 3008, 47), d4800 (cap 4864, 76), + and the 4700-token server repro (cap 4864: 0 GPU faults, code word retrieved modulo one + character's case, the same noise the <=2048 path shows). + BROKEN: d2100 (cap 2112, 33 blocks) 5/5 - and it faults at -n 1,2,3,4 and 8, so it needs only ONE + dispatch; it is not a multi-dispatch race. +EXCLUDED for the 2112 case: reduce_completed (no-op -> 5/5), the completion counter (reset-to-0 -> 5/5, +plus the never-initialised finding), the counter protocol shape, transient sizing, arena sizing, and the +exact-33 assume bound. Trigger = (capacity 2112 == 33 blocks) AND (the last block takes the tail branch, +i.e. has_no_tail false -> %is_full_block false for block_ordinal 32). +=> the fault is in the PRODUCE's tail path for block_ordinal 32 at exactly 33 blocks, or in something + that only that configuration exercises. Note 33 = 32+1: the tail path at block 31 (cap 2048) and at + block 46 (cap 3008, also a tail) both work, so it is specifically block_ordinal 32 in the tail path. + +## Failure rate is a GRADIENT over block count (not a 33-only cliff) - decisive shape +5-run protocol, current build, `llama-bench -p 0 -n 8 -r 1 -d D`: + d2100 cap2112 33 blocks -> 1,1,1,1,1 = 5/5 FAULT (100%) + d2500 cap2560 40 blocks -> 1,0,1,0,0 = 2/5 FAULT (40%) + d2600 cap2624 41 blocks -> 0,0,0,0,1 = 1/5 FAULT (20%) + d1900 cap1920 30 blocks -> 0/5 | d2000 cap2048 32 -> 0/5 | d3000 cap3008 47 -> 0/5 | d4800 cap4864 76 -> 0/5 +The rate DECREASES smoothly with the block count above 33 (100% -> 40% -> 20% -> 0%). A smoothly varying +probability is NOT a deterministic bound bug and NOT a fixed "33-only" defect - it is the signature of a +LAYOUT/CODEGEN-dependent fault: the same logic works or faults depending on where the transients land in +the arena and/or on the JIT-specialised code emitted for that constant capacity. It also explains every +earlier contradictory single-run result (it is a per-run probability, not a law), and it means the +"33 blocks + tail" characterisation is a high-probability region rather than a precise trigger. + +NEXT EXPERIMENT (one line, in scope, and diagnostic): change the partial transients' alignment in +dispatch-flash-attention.cpp (currently alignment = 256 for partial_max / partial_sum / partial_output) +to e.g. 4096, rebuild, re-measure d2100 / d2500 / d2600 at 5 runs each. If the rates MOVE, the fault is +layout-sensitive and placement/alignment is the lever (and may be the fix); if they do not move at all, +the fault is in the JIT-specialised code and the produce must be restructured. + +DISPOSITION against the objective: criteria 1 (cap raised > 2048), 2 (no all_rejected) and 4 (no <=2048 +regression) are MET; criterion 3 is PARTIAL (4718-token HRX0 repro = 0 GPU faults and 'ZX-4718-QQ' +modulo one character's case; the ~3587-token case and the HRX0/Vulkan0 split are still unverified); +criterion 5 ('no sharp boundary cliff' at 2100/3000/4800) is NOT met - 3000 and 4800 are clean (0/5) but +2100 faults 5/5 and 2500 2/5. The residual defect is a layout/codegen-sensitive fault in the produce for +~33..40 producer blocks, which is kernel-codegen scope. + +## BREAKTHROUGH: aligning the partial transients to 4096 removes the GPU page fault +Changed the alignment of the partial_max / partial_sum / partial_output / q8_output transients in +dispatch-flash-attention.cpp from 256 to 4096 (sizes unchanged), rebuilt, contract depths at 5 runs: + d1900 cap1920 30 blocks -> 0,0,0,0,0 = 0/5 CLEAN + d2000 cap2048 32 blocks -> 0,0,0,0,0 = 0/5 CLEAN + d2100 cap2112 33 blocks -> 0,0,0,0,0 = 0/5 CLEAN (was 5/5 with alignment 256) + d3000 cap3008 47 blocks -> 0,1,0,0,0 = 1/5 (mostly clean) + d4800 cap4864 76 blocks -> 0,0,0,0,0 = 0/5 CLEAN +=> the cliff is essentially gone; objective criterion 5 is met at 1900/2000/2100/4800, with a small +residual at 3000. Rates are per-run probabilities, so single runs remain meaningless. + +MECHANISM: with a 256-byte alignment a transient can end exactly on a page boundary, so a SMALL overrun +past its end touches an unmapped page -> HSA_STATUS_ERROR_MEMORY_FAULT. With 4096-byte alignment the same +overrun lands inside the buffer's own (allocated) last page and is silent. Padding the SIZES by 64 KB while +aligning to 64 KB made it WORSE (d2100 4/5, d2500 4/5), so this is a placement effect, not simply "more +slack". => there IS a small out-of-bounds access past one of the partial transients; the alignment change +MASKS it rather than fixing it. It can silently corrupt, and the 1/5 at d3000 is its surviving signature. + +SUSPECTS IN THE PRODUCE (loom file): + - line 289: `%lane_output_channel = index.assume %lane_output_channel0 [range(%lane_output_channel0, 0, 636)]` + - a very loose upper bound (636) against a channel dimension of only value_head_size = 128, so the + compiler cannot prove the vector store in bounds. + - line 134: `%lane_has_output = index.cmp ult, %lane, %c32` - hard-codes 32 output lanes. + - line 282/284: `[range(0,480)]` / `[range(16,496)]` - similarly loose channel bounds. +Tightening those assumes to the true bounds (value_head_size) is the most likely ROOT fix, on top of +keeping the 4096 alignment as defence in depth. + +CONTRACT STATUS: criteria 1 (cap raised) and 2 (no all_rejected) MET; 4 (no <=2048 regression) MET +(d1900/d2000 = 0/5); 5 MET at 1900/2000/2100/4800 (and 1/5 at 3000); 3 PARTIAL - the 4718-token HRX0 repro +is clean, but the HRX0/Vulkan0 split repro and the ~3587-token case are still to run. Remaining for the +'land' task: run the split + 3587-token repros, confirm via GGML_HRX_LOG_DISPATCH that the split (not the +fallback) is selected above 2048, re-read issue #115, then open the PR. + +## AUDITOR REWORK: measured speed + the residual OOB is still live +Auditor rejected completion for: (1) the dispatch still FAILS the decode (res=-3) instead of declining +- the objective's hard constraint; (2) the headline speed recovery was never measured; (3) exact +retrieval on `-dev HRX0` alone is not met; (4) 1470 was not re-verified; (5) PRs #13/#121 are unmerged. + +SPEED MEASURED (llama-bench -p 0 -n 8 -r 3 -d D, -dev HRX0, align-4096 build): + d1470 cap1536 24 blocks -> 73.33 +- 13.09 + d1900 cap1920 30 blocks -> 64.65 +- 9.40 + d2000 cap2048 32 blocks -> 64.49 +- 9.76 + d2100 cap2112 33 blocks -> 53.88 +- 6.83 split selected (fallback 46.4) - +16%, but 2048 is 64.5 + d3000 cap3008 47 blocks -> FAULT (res = -3) on this -r 3 run + d4800 cap4864 76 blocks -> 37.04 +- 4.85 split selected (fallback 32.5) - +14% +=> <=2048 unchanged and clean (1470/1900/2000); the split IS re-selected above 2048 and beats the +fallback, BUT the multipass reducer only recovers about HALF the lost throughput (53.88 vs 64.5 = -16%, +against the fallback's -28%), and d3000 still fails the decode. So criterion 5 is only PARTIALLY met and +the "decline rather than fail" constraint is VIOLATED at d3000. + +OOB SEARCH NARROWED - the produce's channel stores ARE in bounds: + loom 486-501: output_tile_count = padded_value_head_size/128; lane_output_base = lane*4; + lane_has_output = lane < 32; lane_output_channel = output_tile*128 + lane*4; + valid = channel < value_head_size; published = lane_has_output AND valid. + For lane in 0..31 the channel is in {0,4,...,124} and the store is vector<4xf16>, so channel+4 <= 128 = + value_head_size -> IN BOUNDS. The cooperative reducer (loom 633+) uses lane*2 with the same < guard -> + also in bounds. => the "loose range(%lane_output_channel0,0,636)" theory is DEAD; that 636 bound covers + output_tile*128 for padded_value_head_size up to 512, not an out-of-range store. + K/V loads are element-guarded by bounded_key_value_token_count (loom 206/233/338) and the value stage is + bounded by value_width / output_stage_size (307/308/312). So the OOB is elsewhere - the next places to + instrument are the K/V global loads inside produce_partials.active (loom ~150-350) and the reduce's + partial views. + +NEXT (in priority order): + 1. FIND the OOB. Use the fault address from the HSA error plus the arena/transient offsets, or instrument + the produce's global loads. Until it is fixed the matcher must DECLINE rather than fail (objective + constraint) - that is the non-negotiable one. + 2. SPEED: the multipass output pass is a SERIAL `scf.for %block = [0 to active_block_count step 1]` over + all blocks per output element (loom ~911-920) - O(blocks) per element, executed by all 256 workitems. + That is the likely performance sink; make it lane-strided like the max/sum passes (the per-block scale + is already written into partial_max, so a lane-strided sum plus a subgroup reduction per channel would + do it). + 3. Re-run the -dev HRX0-ALONE 4718-token and ~3587-token repros and record the answers. + 4. Note the PRs are OPEN and unmerged - merging is the maintainer's call, not something the agent can do. + +## Wave-padding the block count does NOT fix it - hypothesis FALSIFIED; trigger is layout+token dependent +Patched the matcher to round the capacity up to a whole number of 64-block waves +(wave_block_count = ceil_div(ceil_div(kv,64),64)*64) so the reduce's lane-strided loops +(`scf.for %block = [%lane to %active_block_count step %c64]`) always have a uniform trip count across +lanes. Result (5 runs each, rebuilt): + d2100 cap4096 64 blocks -> 0,0,0,1,0 = 1/5 FAULT + d2500 cap4096 64 blocks -> 0,1,0,0,1 = 2/5 FAULT + d3000 cap4096 64 blocks -> 0,1,0,1,0 = 2/5 FAULT + d4800 cap8192 128 blocks -> 0/5 CLEAN + d1900 cap4096 64 blocks -> 0/5 CLEAN +=> the divergent-trip-count theory is FALSIFIED (a whole number of waves still faults) and 128 blocks is +clean. The three failing depths now share capacity 4096 AND block count 64 yet fault at DIFFERENT rates +(1/5, 2/5, 2/5) - the only remaining difference is the KV length (mask->ne[0] = 2108/2508/3008). So the +fault depends on the transient SIZE (block count), the TOKEN count, AND the arena layout (d2100: align 256 +-> 5/5, align 4096 -> 0/5, align 64K -> 4/5). + +REMAINING EXPLANATION, best fit: ALIASING/OVERLAP in the dispatch infrastructure. +transient-allocator.cpp has overlap logic but only against the COMPLETION COUNTER region +(`transient_allocation_overlaps_region(allocation, completion_counters.arena_offset, byte_count)`), and the +arena holds the graph tensors AND the transients. At capacity 4096 the partial_output transient is 1 MB +(4*64*16*128*2); at the objective's target capacity 32768 it would be 8 MB. If a transient may overlap a +graph tensor or another dispatch's transient, the layout-sensitive, intermittently-faulting behaviour +follows exactly. NEXT: audit TransientAllocator::allocate (transient-allocator.cpp:356-447) for +transient-vs-graph-tensor and transient-vs-transient overlap, and check whether graph tensors live in a +region disjoint from the transients. (Wave-padding reverted - it costs 2x compute for no benefit.) + +## DECISIVE LOCALIZATION: the fault is in `reduce_completed.multipass`; the cooperative reducer is CLEAN +Valid test: forced the multipass wrapper to apply `reduce_completed.cooperative` instead (widened its +`1, 32` bound sites to `1, 4096` - 4 text matches; swapped 1 apply site), rebuilt, 5 runs each: + d2100 (cap 2112, 33 blocks) -> 0,0,0,0,0 = 0/5 FAULT + d3000 (cap 3008, 47 blocks) -> 0,0,0,0,0 = 0/5 FAULT +=> `reduce_completed.multipass` IS the faulting code. This finally explains why the earlier bisects +misled me: the first two were SINGLE-RUN verdicts under an intermittent fault, and the later text-based +`template.return` insertion FAILED TO BUILD ("ninja: build stopped: subcommand failed"), so those 5/5 +numbers came from a stale binary. The reduce had never actually been excluded before. +NOTE: the cooperative is NOT a valid substitute above 32 blocks (its normalisation scratch is +`4x32x2xf32` and it assumes `lane < 32`), so this is a diagnostic, not the fix. + +THE DIFFERENCE between the two reducers - this is where the bug lives: + cooperative (CLEAN): max/sum passes are UNIFORM `scf.for %block = [%c0 to %active_block_count step %c1]` + (every lane iterates every block), scale stage in the 4x32x2 f32 workgroup scratch. + multipass (FAULTS): max/sum passes are LANE-DEPENDENT `scf.for %block = [%lane to %active_block_count + step %c64]` - a dynamic, DIVERGENT lower bound - plus a lane-strided global store of the per-block + scale back into partial_max. +=> prime suspect: the LANE-DEPENDENT (divergent) loop lower bound `%lane`. Not the access bounds (every +access is in bounds for block < active_block_count), not the trip count (wave-padding to a whole number of +64-block waves did not help; 64 blocks still faulted 1-2/5), and not the counter or the produce. + +NEXT: rewrite `reduce_completed.multipass`'s max/sum passes to avoid the divergent bound - e.g. +`scf.for %iteration = [%c0 to %iteration_count step %c1]` (uniform) with `%block = %lane + %iteration*64` +and a SELECTED value (plus a predicated store for the scale), so control flow is uniform across lanes. +Then re-run 5x at d2100/d3000/d4800 and re-check d1900/d2000 are still 0/5. + +## The fault is decisively IN `reduce_completed.multipass`; two structural fixes did NOT remove it +Localization (valid, 5 runs): forcing the multipass wrapper to apply `reduce_completed.cooperative` +(widening its `1, 32` bound sites to `1, 4096`) gave 0/5 faults at BOTH d2100 and d3000 - so the fault is +in my reducer, not in the produce, the wrapper, the counter, or the sizing. CAVEAT: the cooperative is NOT +valid above 32 blocks (its stage is 32-wide), so that build produced GARBAGE output - and therefore the +63.93 t/s "no cliff" speed number measured from that build is INVALID. + +Fix attempts - both reverted, both ~1/5, i.e. NO better than the PR #13 baseline: + (A) uniform loop bounds: `scf.for %block = [%lane to %active_block_count step %c64]` -> + `[%c0 to %active_block_count step %c1]` in the max and sum passes, plus a + `scf.select %workitem_is_zero` leader so the subgroup add-reduce does not multiply the sum by the + lane count -> 1/5, 1/5, 0/5 at d2100/d3000/d4800. + (B) uniform loops + publish `%lane_sum` directly to the reduction stage (dropping the subgroup add-reduce + entirely, since every lane then holds the same complete sum) -> 1/5, 1/5, 0/5. +=> the divergent lane-dependent loop bound is NOT the cause. + +STRUCTURAL DIFFERENCES between the fault-free cooperative and the faulting multipass, NOT yet addressed: + 1. the cooperative computes the sum AND the unnormalized output in ONE fused pass; the multipass uses + three passes (max, sum, output), and its output pass re-reads the per-block scale that the sum pass + wrote back into partial_max; + 2. the cooperative is PHASED (`scf.for %phase`) with a 32-lane stage; the multipass is single-shot; + 3. the cooperative's loops carry `unroll` / `schedule(interleaved)`; the multipass's do not. +NEXT: port the cooperative's reduction STRUCTURE (fused sum+output pass, uniform loops, and NO scale +written back into partial_max) into `reduce_completed.multipass` while keeping the scratch bounded, then +re-run d2100/d3000/d4800 at 5 runs each AND verify the code word is retrieved EXACTLY - a fault-free but +wrong reduction is not acceptable, as the forced-cooperative experiment proved. +Baseline for comparison (PR #13 loom, align-4096 cpp): d2100 1/5, d3000 1/5, d4800 0/5. diff --git a/common/arg.cpp b/common/arg.cpp index 79480e06f9d2..92728d34b55c 100644 --- a/common/arg.cpp +++ b/common/arg.cpp @@ -2586,6 +2586,19 @@ common_params_context common_params_parser_init(common_params & params, llama_ex else { throw std::invalid_argument("invalid value"); } } ).set_env("LLAMA_ARG_LOAD_MODE")); + add_opt(common_arg( + {"-lzm", "--lazy-mode"}, "MODE", + "on-demand reading of certain tensors, for example per-layer embeddings (default: auto)\n" + "- on: read the rows of such tensors from disk on demand instead of keeping them resident (requires mmap)\n" + "- auto: on, but only for tensors larger than 4 GiB\n" + "- off: always keep them resident", + [](common_params & params, const std::string & value) { + /**/ if (value == "on") { params.lazy_mode = LLAMA_LAZY_MODE_ON; } + else if (value == "auto") { params.lazy_mode = LLAMA_LAZY_MODE_AUTO; } + else if (value == "off") { params.lazy_mode = LLAMA_LAZY_MODE_OFF; } + else { throw std::invalid_argument("invalid value"); } + } + ).set_env("LLAMA_ARG_LAZY_MODE")); add_opt(common_arg( {"--numa"}, "TYPE", "attempt optimizations that help on some NUMA systems\n" diff --git a/common/chat.cpp b/common/chat.cpp index 0e06fa591f9d..98255ab4f2be 100644 --- a/common/chat.cpp +++ b/common/chat.cpp @@ -2365,7 +2365,7 @@ static common_chat_params common_chat_params_init_minimax_m3(const common_chat_t auto alternatives_of = [](const json & schema) -> std::optional { for (const auto * keyword : { "oneOf", "anyOf" }) { if (schema.contains(keyword) && schema.at(keyword).is_array() && !schema.at(keyword).empty()) { - return schema.at(keyword); + return std::optional(std::in_place, schema.at(keyword)); } } return std::nullopt; diff --git a/common/common.cpp b/common/common.cpp index ff27d392fb2e..c093133cb16b 100644 --- a/common/common.cpp +++ b/common/common.cpp @@ -1598,6 +1598,7 @@ struct llama_model_params common_model_params_to_llama(common_params & params) { mparams.main_gpu = params.main_gpu; mparams.split_mode = params.split_mode; mparams.load_mode = params.load_mode; + mparams.lazy_mode = params.lazy_mode; mparams.tensor_split = params.tensor_split; mparams.check_tensors = params.check_tensors; mparams.use_extra_bufts = !params.no_extra_bufts; diff --git a/common/common.h b/common/common.h index 919c0ea103a4..9572ab49a853 100644 --- a/common/common.h +++ b/common/common.h @@ -482,6 +482,8 @@ struct common_params { enum llama_split_mode split_mode = LLAMA_SPLIT_MODE_LAYER; // how to split the model across GPUs enum llama_load_mode load_mode = LLAMA_LOAD_MODE_MMAP; // how to load the model + enum llama_lazy_mode lazy_mode = LLAMA_LAZY_MODE_AUTO; // on-demand reading of tensors marked by the arch + common_cpu_params cpuparams; common_cpu_params cpuparams_batch; diff --git a/common/jinja/string.h b/common/jinja/string.h index c4963000adb8..930774f14505 100644 --- a/common/jinja/string.h +++ b/common/jinja/string.h @@ -1,5 +1,6 @@ #pragma once +#include #include #include #include @@ -31,7 +32,9 @@ struct string { parts.push_back({false, std::to_string(v)}); } string(double v) { - parts.push_back({false, std::to_string(v)}); + char buf[512]; + snprintf(buf, sizeof(buf), "%f", v); // std::to_string(double) format before C++26 + parts.push_back({false, buf}); } // mark all parts as user input diff --git a/common/jinja/value.h b/common/jinja/value.h index 5cf85e4f5443..c8ea48030294 100644 --- a/common/jinja/value.h +++ b/common/jinja/value.h @@ -5,6 +5,7 @@ #include #include +#include #include #include #include @@ -256,7 +257,9 @@ struct value_float_t : public value_t { virtual double as_float() const override { return val_flt; } virtual int64_t as_int() const override { return val_int; } virtual string as_string() const override { - std::string out = std::to_string(val_flt); + char buf[512]; + snprintf(buf, sizeof(buf), "%f", val_flt); // std::to_string(double) format before C++26 + std::string out = buf; out.erase(out.find_last_not_of('0') + 1, std::string::npos); // remove trailing zeros if (out.back() == '.') out.push_back('0'); // leave one zero if no decimals return out; diff --git a/common/sampling.cpp b/common/sampling.cpp index 256ac161e20f..55383fd84b17 100644 --- a/common/sampling.cpp +++ b/common/sampling.cpp @@ -11,6 +11,7 @@ #include #include #include +#include #include #include #include @@ -574,6 +575,54 @@ llama_token common_sampler_sample(struct common_sampler * gsmpl, struct llama_co gsmpl->set_logits(ctx, idx); + // engine#123: HRX can hand back all-NaN router logits when the driver migrates a page + // behind in-flight work (the engine#140-class nondeterminism). NaN is never a legitimate + // logit - unlike -inf, which masking legitimately uses - so refuse to sample from it + // instead of silently decoding a wrong-but-plausible token. + // + // engine#315: aborting on the FIRST NaN destroys the evidence along with the run, and an + // intermittent fault therefore ends a whole evaluation instead of being recorded. So + // measure the corruption first - how many entries, over what index range, and whether it + // is contiguous - which is what distinguishes "a migrated/torn page" from "one bad + // element". With GGML_HRX_NAN_CONTINUE set, substitute -inf and carry on: -inf makes those + // vocab entries unsampleable while leaving the rest of the distribution intact, so a long + // run survives. Off by default, so the guard's own default remains a hard abort. + { + float * logits = llama_get_logits_ith(ctx, idx); + if (logits != nullptr) { + const int32_t n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(llama_get_model(ctx))); + int32_t nan_count = 0; + int32_t nan_first = -1; + int32_t nan_last = -1; + for (int32_t i = 0; i < n_vocab; ++i) { + if (std::isnan(logits[i])) { + if (nan_first < 0) { + nan_first = i; + } + nan_last = i; + ++nan_count; + } + } + if (nan_count > 0) { + const bool contiguous = (nan_last - nan_first + 1) == nan_count; + const char * env = std::getenv("GGML_HRX_NAN_CONTINUE"); + const bool nan_continue = env != nullptr && env[0] != '\0' && env[0] != '0'; + LOG_ERR("%s: HRX returned NaN logits: count=%d of %d, index range [%d..%d], contiguous=%s, continuing=%s (engine#315)\n", + __func__, (int) nan_count, (int) n_vocab, (int) nan_first, (int) nan_last, + contiguous ? "yes" : "no", nan_continue ? "yes" : "no"); + if (nan_continue) { + for (int32_t i = nan_first; i <= nan_last; ++i) { + if (std::isnan(logits[i])) { + logits[i] = -INFINITY; + } + } + } else { + GGML_ABORT("HRX: NaN logits (engine#123)"); + } + } + } + } + // Check if a backend sampler has already sampled a token in which case we // return that token id directly. { diff --git a/common/speculative.cpp b/common/speculative.cpp index 5653a90b889c..a0e957ce33b3 100644 --- a/common/speculative.cpp +++ b/common/speculative.cpp @@ -1271,6 +1271,12 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { // call to pair with, so it's stashed here until that next call fires. std::vector> pending_h; // [n_seq][n_embd] + // Position that pending_h belongs to, or -1 when there is no carryover. + // The bridge is only valid for a batch that continues at pending_pos + 1; + // a fresh sequence starts at pos 0 while pending_h still holds the previous + // request's last row (engine #290). + std::vector pending_pos; + std::vector i_batch_beg; std::vector i_batch_end; @@ -1278,6 +1284,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { // Row 0 corresponds to the sampled token, row N to the Nth accepted draft token. std::vector> verify_h; std::vector verify_h_rows; + std::vector verify_pos_first; // [n_seq] — pos of verify_h[seq][0] std::vector i_last; std::vector> chain_h; @@ -1291,7 +1298,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { GGML_ASSERT(ctx_tgt && ctx_dft && "MTP requires ctx_tgt and ctx_dft to be set"); n_embd = llama_model_n_embd_out(llama_get_model(ctx_dft)); - GGML_ASSERT(n_embd == llama_model_n_embd(llama_get_model(ctx_tgt)) && + GGML_ASSERT(n_embd == llama_model_n_embd_out(llama_get_model(ctx_tgt)) && "MTP input row width must match the target h_nextn width"); n_mtp_layers = std::max(1, (int) llama_model_n_layer_nextn(llama_get_model(ctx_dft))); @@ -1339,7 +1346,9 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { llama_set_embeddings_nextn(ctx_tgt, true, /*masked*/ false); llama_set_embeddings_nextn(ctx_dft, true, /*masked*/ true); - is_mem_shared = llama_get_ctx_other(ctx_dft) == ctx_tgt; + char arch[64] = {0}; + llama_model_meta_val_str(llama_get_model(ctx_dft), "general.architecture", arch, sizeof(arch)); + is_mem_shared = llama_get_ctx_other(ctx_dft) == ctx_tgt && std::strcmp(arch, "gemma4-assistant") == 0; chain_heads = n_mtp_layers > 1 && !is_mem_shared; if (chain_heads) { @@ -1352,6 +1361,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { } pending_h.assign(n_seq, std::vector(n_embd, 0.0f)); + pending_pos.assign(n_seq, -1); i_last.assign(n_seq, -1); i_batch_beg.assign(n_seq, -1); @@ -1359,6 +1369,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { verify_h.assign(n_seq, {}); verify_h_rows.assign(n_seq, 0); + verify_pos_first.assign(n_seq, -1); } ~common_speculative_impl_draft_mtp() override { @@ -1461,7 +1472,19 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { continue; } - set_h(i_batch_beg[seq_id], pending_h[seq_id].data()); + // pending_h is the h-row that precedes the first token of this + // batch, so it only pairs with that token when the batch really + // continues from it. A fresh sequence starts at pos 0 while + // pending_h still holds the previous request's last row; seeding + // the MTP head with it there made the draft's first prediction + // depend on the previous request (engine #290). This is the same + // continuity guard eagle3 applies with pending_pos_last. + const llama_pos pos_beg = batch_in.pos[i_batch_beg[seq_id]]; + if (pending_pos[seq_id] >= 0 && pending_pos[seq_id] + 1 == pos_beg) { + set_h(i_batch_beg[seq_id], pending_h[seq_id].data()); + } else { + std::memset(batch.embd + (size_t) i_batch_beg[seq_id] * n_embd, 0, row_bytes); + } } auto * mem_dft = llama_get_memory(ctx_dft); @@ -1504,6 +1527,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { const int32_t n_rows = i_batch_end[seq_id] - i_batch_beg[seq_id] + 1; verify_h_rows[seq_id] = n_rows; verify_h[seq_id].resize((size_t) n_rows * n_embd); + verify_pos_first[seq_id] = batch_in.pos[i_batch_beg[seq_id]]; for (int32_t i = 0; i < n_rows; ++i) { const float * h = llama_get_embeddings_nextn_ith(ctx_tgt, i_batch_beg[seq_id] + i); @@ -1512,6 +1536,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { std::memcpy(pending_h[seq_id].data(), verify_h[seq_id].data() + (size_t) (n_rows - 1) * n_embd, row_bytes); + pending_pos[seq_id] = batch_in.pos[i_batch_end[seq_id]]; } return true; @@ -1681,6 +1706,37 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { const int32_t i_h = std::min(n_accepted, n_rows - 1); const size_t row_bytes = (size_t) n_embd * sizeof(float); std::memcpy(pending_h[seq_id].data(), verify_h[seq_id].data() + (size_t) i_h * n_embd, row_bytes); + pending_pos[seq_id] = verify_pos_first[seq_id] + i_h; + } + + bool get_state(llama_seq_id seq_id, std::vector & data) const override { + if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq || pending_pos[seq_id] < 0) { + return false; + } + + const llama_pos pos = pending_pos[seq_id]; + const std::vector & h = pending_h[seq_id]; + + data.resize(sizeof(llama_pos) + h.size() * sizeof(float)); + std::memcpy(data.data(), &pos, sizeof(llama_pos)); + std::memcpy(data.data() + sizeof(llama_pos), h.data(), h.size() * sizeof(float)); + return true; + } + + void set_state(llama_seq_id seq_id, const std::vector & data) override { + if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) { + return; + } + if (data.size() != sizeof(llama_pos) + (size_t) n_embd * sizeof(float)) { + return; + } + + llama_pos pos = -1; + std::memcpy(&pos, data.data(), sizeof(llama_pos)); + + pending_pos[seq_id] = pos; + pending_h[seq_id].resize(n_embd); + std::memcpy(pending_h[seq_id].data(), data.data() + sizeof(llama_pos), (size_t) n_embd * sizeof(float)); } bool need_embd() const override { @@ -2333,7 +2389,7 @@ common_speculative_init_result::common_speculative_init_result( model_path = params.speculative.draft.mparams.path; LOG_INF("%s: loading draft model '%s'\n", __func__, model_path.c_str()); - llama_model * model_dft = llama_model_load_from_file(params.model.path.c_str(), mparams); + llama_model * model_dft = llama_model_load_from_file(model_path.c_str(), mparams); if (model_dft == NULL) { LOG_ERR("%s: failed to load draft model, '%s'\n", __func__, model_path.c_str()); return; diff --git a/conversion/__init__.py b/conversion/__init__.py index 1a47b851a0e6..7dc95627ee49 100644 --- a/conversion/__init__.py +++ b/conversion/__init__.py @@ -39,11 +39,13 @@ "ChameleonForCausalLM": "chameleon", "ChameleonForConditionalGeneration": "chameleon", "ChatGLMForConditionalGeneration": "chatglm", + "CambrianQwenForCausalLM": "cambrian", "ChatGLMModel": "chatglm", "CodeShellForCausalLM": "codeshell", "CogVLMForCausalLM": "cogvlm", "Cohere2MoeForCausalLM": "command_r", "Cohere2ForCausalLM": "command_r", + "CodeGenForCausalLM": "codegen", "CohereForCausalLM": "command_r", "DbrxForCausalLM": "dbrx", "DeciLMForCausalLM": "deci", @@ -73,7 +75,9 @@ "FalconH1ForCausalLM": "falcon_h1", "FalconMambaForCausalLM": "mamba", "GPT2LMHeadModel": "gpt2", + "GPTJForCausalLM": "gptj", "GPTBigCodeForCausalLM": "starcoder", + "GPTNeoForCausalLM": "gptneo", "GPTNeoXForCausalLM": "gptneox", "GPTRefactForCausalLM": "refact", "Gemma2ForCausalLM": "gemma", @@ -180,6 +184,7 @@ "Olmo3ForCausalLM": "olmo", "OlmoForCausalLM": "olmo", "OlmoeForCausalLM": "olmo", + "OPTForCausalLM": "opt", "OpenELMForCausalLM": "openelm", "OrionForCausalLM": "orion", "PLMForCausalLM": "plm", @@ -190,6 +195,7 @@ "Phi3ForCausalLM": "phi", "Phi4ForCausalLMV": "phi", "PhiForCausalLM": "phi", + "PicoDecoderHF": "pico", "PhiMoEForCausalLM": "phi", "Plamo2ForCausalLM": "plamo", "Plamo3ForCausalLM": "plamo", @@ -215,6 +221,8 @@ "Qwen3_5ForConditionalGeneration": "qwen", "Qwen3_5MoeForCausalLM": "qwen", "Qwen3_5MoeForConditionalGeneration": "qwen", + "Qwen4ExpForCausalLM": "qwen4exp", + "Qwen4ExpForConditionalGeneration": "qwen4exp", "RND1": "qwen", "RWForCausalLM": "falcon", "RWKV6Qwen2ForCausalLM": "rwkv", @@ -253,6 +261,8 @@ "XverseForCausalLM": "xverse", "YoutuForCausalLM": "deepseek", "YoutuVLForConditionalGeneration": "deepseek", + "ZayaForCausalLM": "zaya", + "Zaya1VLForConditionalGeneration": "zaya", "modeling_grove_moe.GroveMoeForCausalLM": "grovemoe", "modeling_sarvam_moe.SarvamMoEForCausalLM": "bailingmoe", } @@ -300,6 +310,7 @@ "Qwen2VLForConditionalGeneration": "qwenvl", "Qwen2VLModel": "qwenvl", "Qwen2_5OmniModel": "qwenvl", + "Zaya1VLForConditionalGeneration": "zaya", "Qwen2_5_VLForConditionalGeneration": "qwenvl", "Qwen3ASRForConditionalGeneration": "qwen3vl", "Qwen3OmniMoeForConditionalGeneration": "qwen3vl", @@ -307,6 +318,7 @@ "Qwen3VLMoeForConditionalGeneration": "qwen3vl", "Qwen3_5ForConditionalGeneration": "qwen3vl", "Qwen3_5MoeForConditionalGeneration": "qwen3vl", + "Qwen4ExpForConditionalGeneration": "qwen4exp", "RADIOModel": "nemotron", "Sarashina2VisionForCausalLM": "sarashina2", "SmolVLMForConditionalGeneration": "smolvlm", diff --git a/conversion/base.py b/conversion/base.py index a7cd3fd904aa..af42497d1fb0 100644 --- a/conversion/base.py +++ b/conversion/base.py @@ -963,12 +963,16 @@ def load(): else: raise ValueError(f"Unknown file type: {self.ftype.name}") + # a chunked tensor quantizes as one chunk at a time, while it is written + quantize = data.quantize if isinstance(data, gguf.LazyChunkedTensor) else ( + lambda qtype, d=data: gguf.quants.quantize(d, qtype)) + try: - data = gguf.quants.quantize(data, data_qtype) + data = quantize(data_qtype) except gguf.QuantError as e: logger.warning("%s, %s", e, "falling back to F16") data_qtype = gguf.GGMLQuantizationType.F16 - data = gguf.quants.quantize(data, data_qtype) + data = quantize(data_qtype) shape = gguf.quant_shape_from_byte_shape(data.shape, data_qtype) if data.dtype == np.uint8 else data.shape diff --git a/conversion/bert.py b/conversion/bert.py index 0d25d0d62df5..1977be884ac9 100644 --- a/conversion/bert.py +++ b/conversion/bert.py @@ -262,7 +262,7 @@ def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Ca return super().filter_tensors((name, gen)) -@ModelBase.register("RobertaModel", "RobertaForSequenceClassification") +@ModelBase.register("RobertaModel", "RobertaForSequenceClassification", "RobertaForMaskedLM") class RobertaModel(BertModel): model_arch = gguf.MODEL_ARCH.BERT diff --git a/conversion/cambrian.py b/conversion/cambrian.py new file mode 100644 index 000000000000..c2a25653f041 --- /dev/null +++ b/conversion/cambrian.py @@ -0,0 +1,10 @@ +from __future__ import annotations + +from .base import ModelBase +from .qwen import Qwen2Model + + +@ModelBase.register("CambrianQwenForCausalLM") +class CambrianQwenModel(Qwen2Model): + """Cambrian's Qwen text tower: config model_type is `qwen2`, tensors and shapes are + Qwen2's. It runs on the qwen2 architecture; only the class name differs.""" diff --git a/conversion/chatglm.py b/conversion/chatglm.py index 801913075dbc..9bc195a8db52 100644 --- a/conversion/chatglm.py +++ b/conversion/chatglm.py @@ -8,7 +8,7 @@ from .base import ModelBase, SentencePieceTokenTypes, TextModel, gguf -@ModelBase.register("GlmForCausalLM", "ChatGLMModel", "ChatGLMForConditionalGeneration") +@ModelBase.register("GlmForCausalLM", "ChatGLMModel", "ChatGLMForConditionalGeneration", "ChatGlmForCausalLM") class ChatGLMModel(TextModel): model_arch = gguf.MODEL_ARCH.CHATGLM diff --git a/conversion/codegen.py b/conversion/codegen.py new file mode 100644 index 000000000000..8038da1a0679 --- /dev/null +++ b/conversion/codegen.py @@ -0,0 +1,46 @@ +from __future__ import annotations + +from typing import Iterable, TYPE_CHECKING + +import torch + +if TYPE_CHECKING: + from torch import Tensor + +from .base import ModelBase, TextModel, gguf + + +@ModelBase.register("CodeGenForCausalLM") +class CodeGenModel(TextModel): + model_arch = gguf.MODEL_ARCH.CODEGEN + + def set_gguf_parameters(self): + hparams = self.hparams + self.gguf_writer.add_block_count(self.block_count) + self.gguf_writer.add_context_length(hparams["n_positions"]) + self.gguf_writer.add_embedding_length(hparams["n_embd"]) + self.gguf_writer.add_feed_forward_length(4 * hparams["n_embd"]) + self.gguf_writer.add_head_count(hparams["n_head"]) + self.gguf_writer.add_layer_norm_eps(hparams["layer_norm_epsilon"]) + + n_embd_head = hparams["n_embd"] // hparams["n_head"] + n_rot = int(hparams.get("rotary_dim", n_embd_head)) + if n_rot < n_embd_head: + self.gguf_writer.add_rope_dimension_count(n_rot) + + self.gguf_writer.add_file_type(self.ftype) + + def set_vocab(self): + self._set_vocab_gpt2() + + def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]: + if name.endswith((".attn.bias", ".attn.masked_bias", ".attn.causal_mask")): # mask buffers + return + if name.endswith(".attn.qkv_proj.weight"): + # HF splits qkv_proj into mp_num = 4 blocks, each ordered (query, value, key) + # (CodeGenAttention.forward); the fused-QKV graph wants all of q, then k, then v + n_embd = self.hparams["n_embd"] + w = data_torch.reshape(4, 3, n_embd // 4, n_embd) + q, v, k = w[:, 0], w[:, 1], w[:, 2] + data_torch = torch.cat([t.reshape(n_embd, n_embd) for t in (q, k, v)], dim=0) + yield from super().modify_tensors(data_torch, name, bid) diff --git a/conversion/command_r.py b/conversion/command_r.py index 118565c66973..c6aed1559e7c 100644 --- a/conversion/command_r.py +++ b/conversion/command_r.py @@ -29,7 +29,7 @@ def set_gguf_parameters(self): self.gguf_writer.add_rope_scaling_type(gguf.RopeScalingType.NONE) -@ModelBase.register("Cohere2ForCausalLM") +@ModelBase.register("Cohere2ForCausalLM", "Cohere2Model") class Cohere2Model(TextModel): model_arch = gguf.MODEL_ARCH.COHERE2 diff --git a/conversion/deepseek.py b/conversion/deepseek.py index ea6ae23d58e7..cbc354358486 100644 --- a/conversion/deepseek.py +++ b/conversion/deepseek.py @@ -17,7 +17,7 @@ from .qwen import QwenModel -@ModelBase.register("DeepseekOCRForCausalLM", "UnlimitedOCRForCausalLM") +@ModelBase.register("DeepseekOCRForCausalLM", "UnlimitedOCRForCausalLM", "DeepseekOcrForConditionalGeneration") class DeepseekOCRVisionModel(MmprojModel): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) @@ -472,7 +472,7 @@ def set_gguf_parameters(self): self.gguf_writer.add_indexer_top_k(self.hparams["index_topk"]) -@ModelBase.register("DeepseekV4ForCausalLM") +@ModelBase.register("DeepseekV4ForCausalLM", "DeepSeekV4") class DeepseekV4Model(TextModel): model_arch = gguf.MODEL_ARCH.DEEPSEEK4 _skipped_mtp_tensors = 0 diff --git a/conversion/falcon.py b/conversion/falcon.py index 085fd4cd33ff..1c1037e85618 100644 --- a/conversion/falcon.py +++ b/conversion/falcon.py @@ -10,7 +10,7 @@ from .base import ModelBase, TextModel, gguf -@ModelBase.register("FalconForCausalLM", "RWForCausalLM") +@ModelBase.register("FalconForCausalLM", "RWForCausalLM", "RWModel") class FalconModel(TextModel): model_arch = gguf.MODEL_ARCH.FALCON diff --git a/conversion/gemma.py b/conversion/gemma.py index c552df732b0f..3c673adc89ce 100644 --- a/conversion/gemma.py +++ b/conversion/gemma.py @@ -13,7 +13,7 @@ from .base import MmprojModel, ModelBase, TextModel, gguf, logger -@ModelBase.register("GemmaForCausalLM") +@ModelBase.register("GemmaForCausalLM", "GemmaModel") class GemmaModel(TextModel): model_arch = gguf.MODEL_ARCH.GEMMA @@ -67,7 +67,7 @@ def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iter yield from super().modify_tensors(data_torch, name, bid) -@ModelBase.register("Gemma2ForCausalLM") +@ModelBase.register("Gemma2ForCausalLM", "Gemma2Model") class Gemma2Model(TextModel): model_arch = gguf.MODEL_ARCH.GEMMA2 @@ -117,7 +117,7 @@ def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iter yield from super().modify_tensors(data_torch, name, bid) -@ModelBase.register("Gemma3ForCausalLM", "Gemma3ForConditionalGeneration") +@ModelBase.register("Gemma3ForCausalLM", "Gemma3ForConditionalGeneration", "Gemma3Model") class Gemma3Model(TextModel): model_arch = gguf.MODEL_ARCH.GEMMA3 @@ -765,7 +765,7 @@ def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iter yield from super().modify_tensors(data_torch, name, bid) -@ModelBase.register("Gemma4UnifiedForConditionalGeneration") +@ModelBase.register("Gemma4UnifiedForConditionalGeneration", "Gemma4UnifiedForCausalLM") class Gemma4UnifiedModel(Gemma4Model): model_arch = gguf.MODEL_ARCH.GEMMA4 diff --git a/conversion/gpt2.py b/conversion/gpt2.py index 1cf06ae8b50c..bf328f913681 100644 --- a/conversion/gpt2.py +++ b/conversion/gpt2.py @@ -10,7 +10,7 @@ from .base import ModelBase, TextModel, gguf, logger -@ModelBase.register("GPT2LMHeadModel") +@ModelBase.register("GPT2LMHeadModel", "GPT2Model", "GPT2") class GPT2Model(TextModel): model_arch = gguf.MODEL_ARCH.GPT2 diff --git a/conversion/gpt_oss.py b/conversion/gpt_oss.py index d2c70c0bba56..88f49586e847 100644 --- a/conversion/gpt_oss.py +++ b/conversion/gpt_oss.py @@ -10,7 +10,7 @@ from .base import ModelBase, TextModel, gguf, logger -@ModelBase.register("GptOssForCausalLM") +@ModelBase.register("GptOssForCausalLM", "GPTOSSForCausalLM") class GptOssModel(TextModel): model_arch = gguf.MODEL_ARCH.GPT_OSS diff --git a/conversion/gptj.py b/conversion/gptj.py new file mode 100644 index 000000000000..64bf5b82e793 --- /dev/null +++ b/conversion/gptj.py @@ -0,0 +1,46 @@ +from __future__ import annotations + +from typing import Iterable, TYPE_CHECKING + + +if TYPE_CHECKING: + from torch import Tensor + +from .base import ModelBase, TextModel, gguf + + +@ModelBase.register("GPTJForCausalLM") +class GPTJModel(TextModel): + model_arch = gguf.MODEL_ARCH.GPTJ + + def set_gguf_parameters(self): + hparams = self.hparams + self.gguf_writer.add_block_count(self.block_count) + self.gguf_writer.add_context_length(hparams.get("n_positions", hparams.get("n_ctx", 2048))) + self.gguf_writer.add_embedding_length(hparams["n_embd"]) + self.gguf_writer.add_feed_forward_length(4 * hparams["n_embd"]) + self.gguf_writer.add_head_count(hparams["n_head"]) + self.gguf_writer.add_layer_norm_eps(hparams["layer_norm_epsilon"]) + + # GPT-J applies rotary to the first `rotary_dim` of each head (partial RoPE). + n_embd_head = hparams["n_embd"] // hparams["n_head"] + n_rot = int(hparams.get("rotary_dim", n_embd_head)) + if n_rot < n_embd_head: + self.gguf_writer.add_rope_dimension_count(n_rot) + + self.gguf_writer.add_file_type(self.ftype) + + def set_vocab(self): + self._set_vocab_gpt2() + + def get_vocab_base_pre(self, tokenizer) -> str: + # GPT-J's tokenizer is GPT-2's byte-level BPE; its checksum is just not in the list + return "gpt-2" + + def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]: + # HF keeps the projections as nn.Linear (no Conv1D permute); the tensor map + # already knows the GPT-J names, and biases ride the usual suffix path. + if name.endswith((".attn.bias", ".attn.masked_bias")): + return # causal-mask buffers older checkpoints carry, not weights + + yield from super().modify_tensors(data_torch, self.map_tensor_name(name), bid) diff --git a/conversion/gptneo.py b/conversion/gptneo.py new file mode 100644 index 000000000000..7e571bd1504f --- /dev/null +++ b/conversion/gptneo.py @@ -0,0 +1,34 @@ +from __future__ import annotations + +from typing import Iterable, TYPE_CHECKING + + +if TYPE_CHECKING: + from torch import Tensor + +from .base import ModelBase, TextModel, gguf + + +@ModelBase.register("GPTNeoForCausalLM") +class GPTNeoModel(TextModel): + model_arch = gguf.MODEL_ARCH.GPTNEO + + def set_gguf_parameters(self): + hparams = self.hparams + self.gguf_writer.add_block_count(self.block_count) + self.gguf_writer.add_context_length(hparams["max_position_embeddings"]) + self.gguf_writer.add_embedding_length(hparams["hidden_size"]) + self.gguf_writer.add_feed_forward_length(4 * hparams["hidden_size"]) + self.gguf_writer.add_head_count(hparams["num_heads"]) + self.gguf_writer.add_layer_norm_eps(hparams["layer_norm_epsilon"]) + self.gguf_writer.add_sliding_window(hparams["window_size"]) + self.gguf_writer.add_file_type(self.ftype) + + def set_vocab(self): + self._set_vocab_gpt2() + + def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]: + # attention bias buffers are rebuilt at runtime + if name.endswith((".attn.bias", ".attn.masked_bias", ".attn.attention.bias", ".attn.attention.masked_bias")): + return + yield from super().modify_tensors(data_torch, name, bid) diff --git a/conversion/gptneox.py b/conversion/gptneox.py index 6a42b12b15af..0caa4b6abd6f 100644 --- a/conversion/gptneox.py +++ b/conversion/gptneox.py @@ -12,7 +12,7 @@ from .base import ModelBase, TextModel, gguf, logger -@ModelBase.register("GPTNeoXForCausalLM") +@ModelBase.register("GPTNeoXForCausalLM", "GPTNeoXModel") class GPTNeoXModel(TextModel): model_arch = gguf.MODEL_ARCH.GPTNEOX diff --git a/conversion/lfm2.py b/conversion/lfm2.py index 70ce45658be5..5eb3f107cb40 100644 --- a/conversion/lfm2.py +++ b/conversion/lfm2.py @@ -64,7 +64,7 @@ def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iter yield from super().modify_tensors(data_torch, name, bid) -@ModelBase.register("Lfm2Model", "Lfm2BidirectionalModel") +@ModelBase.register("Lfm2Model", "Lfm2BidirectionalModel", "Lfm2BidirectionalForMaskedLM") class LFM2ColBertModel(LFM2Model): model_arch = gguf.MODEL_ARCH.LFM2 dense_tensor_name = "dense_2" @@ -92,7 +92,7 @@ def generate_extra_tensors(self) -> Iterable[tuple[str, Tensor]]: yield f"{self.dense_tensor_name}.weight", tensor.clone() -@ModelBase.register("Lfm2MoeForCausalLM") +@ModelBase.register("Lfm2MoeForCausalLM", "Lfm2MoEForCausalLM") class LFM2MoeModel(TextModel): model_arch = gguf.MODEL_ARCH.LFM2MOE diff --git a/conversion/llada.py b/conversion/llada.py index 98dc9de95b37..885514ed7b95 100644 --- a/conversion/llada.py +++ b/conversion/llada.py @@ -10,7 +10,7 @@ from .base import ModelBase, TextModel, gguf -@ModelBase.register("LLaDAModelLM") +@ModelBase.register("LLaDAModelLM", "LLaDAForCausalLM") class LLaDAModel(TextModel): model_arch = gguf.MODEL_ARCH.LLADA undo_permute = True diff --git a/conversion/llama.py b/conversion/llama.py index 9b3373f911e9..ea26e872a544 100644 --- a/conversion/llama.py +++ b/conversion/llama.py @@ -27,7 +27,15 @@ "Eagle3Speculator", "Eagle3DraftModel", "IQuestCoderForCausalLM", - "LlamaModel") + "LlamaModel", + "MistralModel", + "MixtralModel", + "LLaMA", + "LLAMA", + "LLaMAModel", + "LlaMAForCausalLM", + "llamaForCausalLM", + "LlamaForConditionalGeneration") class LlamaModel(TextModel): model_arch = gguf.MODEL_ARCH.LLAMA undo_permute = True diff --git a/conversion/mamba.py b/conversion/mamba.py index 43d559ffb0ae..f51956e8aa68 100644 --- a/conversion/mamba.py +++ b/conversion/mamba.py @@ -13,7 +13,7 @@ from .base import ModelBase, TextModel, gguf, logger -@ModelBase.register("MambaForCausalLM", "MambaLMHeadModel", "FalconMambaForCausalLM") +@ModelBase.register("MambaForCausalLM", "MambaLMHeadModel", "FalconMambaForCausalLM", "MambaModel") class MambaModel(TextModel): model_arch = gguf.MODEL_ARCH.MAMBA diff --git a/conversion/mpt.py b/conversion/mpt.py index 9557ab7fa642..0fb9ada2f706 100644 --- a/conversion/mpt.py +++ b/conversion/mpt.py @@ -8,7 +8,7 @@ from .base import ModelBase, TextModel, gguf -@ModelBase.register("MPTForCausalLM") +@ModelBase.register("MPTForCausalLM", "MptForCausalLM") class MPTModel(TextModel): model_arch = gguf.MODEL_ARCH.MPT diff --git a/conversion/olmo.py b/conversion/olmo.py index 1664c30e402e..0e9bab9a636a 100644 --- a/conversion/olmo.py +++ b/conversion/olmo.py @@ -44,7 +44,7 @@ class SeedOssModel(TextModel): @ModelBase.register("Olmo2ForCausalLM") -@ModelBase.register("Olmo3ForCausalLM") +@ModelBase.register("Olmo3ForCausalLM", "OLMo3ForCausalLM") class Olmo2Model(TextModel): model_arch = gguf.MODEL_ARCH.OLMO2 @@ -66,7 +66,7 @@ def set_gguf_parameters(self): self.gguf_writer.add_sliding_window_pattern(sliding_window_pattern) -@ModelBase.register("OlmoeForCausalLM") +@ModelBase.register("OlmoeForCausalLM", "OlmoeModel") class OlmoeModel(TextModel): model_arch = gguf.MODEL_ARCH.OLMOE diff --git a/conversion/opt.py b/conversion/opt.py new file mode 100644 index 000000000000..06eeae9edbfc --- /dev/null +++ b/conversion/opt.py @@ -0,0 +1,39 @@ +from __future__ import annotations + +from typing import Iterable, TYPE_CHECKING + + +if TYPE_CHECKING: + from torch import Tensor + +from .base import ModelBase, TextModel, gguf + + +@ModelBase.register("OPTForCausalLM") +class OPTModel(TextModel): + model_arch = gguf.MODEL_ARCH.OPT + + def set_gguf_parameters(self): + hparams = self.hparams + self.gguf_writer.add_block_count(self.block_count) + self.gguf_writer.add_context_length(hparams["max_position_embeddings"]) + self.gguf_writer.add_embedding_length(hparams["hidden_size"]) + self.gguf_writer.add_feed_forward_length(hparams["ffn_dim"]) + self.gguf_writer.add_head_count(hparams["num_attention_heads"]) + self.gguf_writer.add_layer_norm_eps(hparams.get("layer_norm_eps", 1e-5)) + self.gguf_writer.add_file_type(self.ftype) + + def set_vocab(self): + self._set_vocab_gpt2() + + def get_vocab_base_pre(self, tokenizer) -> str: + # OPT's tokenizer is GPT-2's byte-level BPE (same pre-tokenizer); its checksum is just + # not in the list + return "gpt-2" + + def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]: + # OPT's learned positions carry two padding rows (offset=2); drop them so the + # runtime's 0-based positions line up with the GGUF rows. + if name.endswith("embed_positions.weight"): + data_torch = data_torch[2:] + yield from super().modify_tensors(data_torch, name, bid) diff --git a/conversion/phi.py b/conversion/phi.py index df4bfe809af7..3b7074346710 100644 --- a/conversion/phi.py +++ b/conversion/phi.py @@ -13,7 +13,7 @@ from .base import MmprojModel, ModelBase, SentencePieceTokenTypes, TextModel, gguf, logger -@ModelBase.register("PhiForCausalLM") +@ModelBase.register("PhiForCausalLM", "Phi2Model") class Phi2Model(TextModel): model_arch = gguf.MODEL_ARCH.PHI2 @@ -35,7 +35,7 @@ def set_gguf_parameters(self): self.gguf_writer.add_add_bos_token(False) -@ModelBase.register("Phi3ForCausalLM", "Phi4ForCausalLMV") +@ModelBase.register("Phi3ForCausalLM", "Phi4ForCausalLMV", "Phi3Model") class Phi3MiniModel(TextModel): model_arch = gguf.MODEL_ARCH.PHI3 @@ -335,7 +335,7 @@ def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iter return -@ModelBase.register("PhiMoEForCausalLM") +@ModelBase.register("PhiMoEForCausalLM", "PhimoeForCausalLM") class PhiMoeModel(Phi3MiniModel): model_arch = gguf.MODEL_ARCH.PHIMOE diff --git a/conversion/pico.py b/conversion/pico.py new file mode 100644 index 000000000000..2ef3d77094a7 --- /dev/null +++ b/conversion/pico.py @@ -0,0 +1,38 @@ +from __future__ import annotations + +from typing import Iterable, TYPE_CHECKING + + +if TYPE_CHECKING: + from torch import Tensor + +from .base import ModelBase, TextModel, gguf + + +@ModelBase.register("PicoDecoderHF") +class PicoDecoderModel(TextModel): + """pico-lm's decoder: an ordinary Llama (RMSNorm, GQA, RoPE, SwiGLU) with its own + module names. It runs on the llama architecture; only the converter differs.""" + + model_arch = gguf.MODEL_ARCH.LLAMA + + def set_gguf_parameters(self): + c = self.hparams + self.gguf_writer.add_block_count(c["n_layers"]) + self.gguf_writer.add_context_length(c["max_seq_len"]) + self.gguf_writer.add_embedding_length(c["d_model"]) + self.gguf_writer.add_feed_forward_length(c["activation_hidden_dim"]) + self.gguf_writer.add_head_count(c["attention_n_heads"]) + self.gguf_writer.add_head_count_kv(c["attention_n_kv_heads"]) + self.gguf_writer.add_layer_norm_rms_eps(c["norm_eps"]) + self.gguf_writer.add_rope_freq_base(c.get("position_emb_theta", 10000.0)) + self.gguf_writer.add_file_type(self.ftype) + + def set_vocab(self): + self._set_vocab_gpt2() + + def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]: + # the HF wrapper nests the real module under `pico_decoder.` + if name.startswith("pico_decoder."): + name = name[len("pico_decoder."):] + yield from super().modify_tensors(data_torch, name, bid) diff --git a/conversion/qwen.py b/conversion/qwen.py index d1127f7431f6..fd384f477a4f 100644 --- a/conversion/qwen.py +++ b/conversion/qwen.py @@ -12,7 +12,7 @@ from .base import ModelBase, TextModel, gguf, logger -@ModelBase.register("QWenLMHeadModel") +@ModelBase.register("QWenLMHeadModel", "QwenForCausalLM", "QwenLMHeadModel", "QWenForCausalLM") class QwenModel(TextModel): model_arch = gguf.MODEL_ARCH.QWEN @@ -70,7 +70,7 @@ def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iter yield from super().modify_tensors(data_torch, name, bid) -@ModelBase.register("Qwen2MoeForCausalLM") +@ModelBase.register("Qwen2MoeForCausalLM", "Qwen2MoEForCausalLM") class Qwen2MoeModel(TextModel): model_arch = gguf.MODEL_ARCH.QWEN2MOE @@ -250,7 +250,7 @@ def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iter yield from super().modify_tensors(data_torch, name, bid) -@ModelBase.register("Qwen3MoeForCausalLM") +@ModelBase.register("Qwen3MoeForCausalLM", "Qwen3MoeModel") class Qwen3MoeModel(Qwen2MoeModel): model_arch = gguf.MODEL_ARCH.QWEN3MOE @@ -620,12 +620,12 @@ def prepare_metadata(self, vocab_only: bool): self.fname_out = self.fname_out.parent / f"mtp-{fname_default}.gguf" -@ModelBase.register("Qwen3_5ForConditionalGeneration", "Qwen3_5ForCausalLM") +@ModelBase.register("Qwen3_5ForConditionalGeneration", "Qwen3_5ForCausalLM", "Qwen3_5Model", "Qwen35ForCausalLM") class Qwen3_5TextModel(_Qwen35MtpMixin, _Qwen35MRopeMixin, _LinearAttentionVReorderBase): model_arch = gguf.MODEL_ARCH.QWEN35 -@ModelBase.register("Qwen3_5MoeForConditionalGeneration", "Qwen3_5MoeForCausalLM") +@ModelBase.register("Qwen3_5MoeForConditionalGeneration", "Qwen3_5MoeForCausalLM", "Qwen3_5MoEForCausalLM") class Qwen3_5MoeTextModel(_Qwen35MtpMixin, _Qwen35MRopeMixin, _LinearAttentionVReorderBase): model_arch = gguf.MODEL_ARCH.QWEN35MOE diff --git a/conversion/qwen4exp.py b/conversion/qwen4exp.py new file mode 100644 index 000000000000..168796d616b9 --- /dev/null +++ b/conversion/qwen4exp.py @@ -0,0 +1,195 @@ +from __future__ import annotations + +from typing import Iterable, cast + +import torch +from torch import Tensor + +import gguf +import numpy as np + +from .base import ModelBase +from .qwen import _LinearAttentionVReorderBase, _Qwen35MRopeMixin +from .qwen3vl import Qwen3VLVisionModel + + +@ModelBase.register("Qwen4ExpForConditionalGeneration", "Qwen4ExpForCausalLM") +@ModelBase.example("Qwen/Qwen3.8-Flash-Next") +class Qwen4ExpTextModel(_Qwen35MRopeMixin, _LinearAttentionVReorderBase): + """Qwen3.8-Flash-Next. + + Shares the Qwen3.5 gated delta net and interleaved mrope, and adds three things: + hyper-connections in place of every layer norm, QSA sparse attention on the full + attention layers, and PLE n-gram hash embeddings on a single layer. + """ + + model_arch = gguf.MODEL_ARCH.QWEN4EXP + + # the MTP block is a separate draft head; vLLM drops it too + supports_mtp_export = False + no_mtp = True + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + # only the shard names, so the table itself is never held + self._ple_shards: dict[int, str] = {} + self._ple_row_dim: int | None = None + + def _read_hash_constants(self, suffix: str) -> list[int]: + """Read an int64 PLE constant straight from the checkpoint. + + prepare_tensors() casts every non-float dtype to float32 before + modify_tensors() sees it (base.py), which would silently round these + 45-bit multipliers. Reading the lazy tensor here bypasses that. + """ + for name, gen in self.model_tensors.items(): + if name.endswith(suffix): + t = gen() + if t.dtype != torch.int64: + t = t.to(torch.int64) + return [int(x) for x in t.tolist()] + raise ValueError(f"PLE constant {suffix!r} missing from the checkpoint") + + def set_gguf_parameters(self): + super().set_gguf_parameters() + hp = self.hparams + + self.gguf_writer.add_hyper_connection_count(hp["hc_count"]) + self.gguf_writer.add_hyper_connection_low_rank(hp["hc_lowrank"]) + + n_layer = hp["num_hidden_layers"] + self.gguf_writer.add_indexer_head_count(hp["indexer_n_heads"]) + self.gguf_writer.add_indexer_key_length(hp["indexer_head_dim"]) + self.gguf_writer.add_indexer_top_k(hp["indexer_budget"]) + ratio = hp["indexer_compress_ratio"] + layer_types = hp["layer_types"] + self.gguf_writer.add_attention_compress_ratios( + [ratio if layer_types[i] == "full_attention" else 0 for i in range(n_layer)] + ) + + # ple_layer_ids is 1-based in the HF config; empty means no n-gram table, + # so emit no PLE keys rather than optional ones + ple_layers = [i - 1 for i in hp["ple_layer_ids"]] + if not ple_layers: + return + self.gguf_writer.add_ple_layers(ple_layers) + self.gguf_writer.add_ple_ngram_size(hp["ngram_size"]) + self.gguf_writer.add_ple_heads_per_ngram(hp["heads_per_ngram"]) + self.gguf_writer.add_ple_conv_kernel(hp["ple_conv_kernel_size"]) + self.gguf_writer.add_ple_eos_token_id(self._eos_token_id()) + # an image is decoded as an embeddings-only batch, so the graph has no placeholder + # ids to hash; carry the id and let it stand in for those positions + _img = self._image_token_id() + if _img is not None: + self.gguf_writer.add_ple_image_token_id(int(_img)) + if self._ple_row_dim is not None: + self.gguf_writer.add_embedding_length_per_layer_input(self._ple_row_dim) + + self.gguf_writer.add_ple_layer_multipliers( + self._read_hash_constants("ple_embedding.layer_multipliers")) + self.gguf_writer.add_ple_head_offsets( + self._read_hash_constants("ple_embedding.ngram_heads_offsets")) + self.gguf_writer.add_ple_head_vocab_sizes( + self._read_hash_constants("ple_embedding.ngram_heads_vocab_sizes")) + + def _image_token_id(self) -> int | None: + img = self.hparams.get("image_token_id") + return None if img is None else int(img) + + def _eos_token_id(self) -> int: + eos = self.hparams.get("eos_token_id") + if isinstance(eos, list): + # the PLE hash resets n-grams on the primary EOS + return int(eos[-1]) + if eos is None: + raise ValueError("eos_token_id is required: the PLE hash resets its n-grams on it") + return int(eos) + + def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]: + # int64 hash constants must stay exact; 1-D tensors force F32, so use KV + if name.endswith("ple_embedding.layer_multipliers"): + self._ple_multipliers = [int(x) for x in data_torch.tolist()] + return [] + if name.endswith("ple_embedding.ngram_heads_offsets"): + self._ple_head_offsets = [int(x) for x in data_torch.tolist()] + return [] + if name.endswith("ple_embedding.ngram_heads_vocab_sizes"): + self._ple_head_vocab_sizes = [int(x) for x in data_torch.tolist()] + return [] + + if ".ngram_embedding.shard_" in name: + return self._place_ple_shard(data_torch, name) + + # one projection feeds indexer q and k; split it, as minimax-m3 does + if ".indexer.index_qk_proj.weight" in name: + n_q = self.hparams["indexer_n_heads"] * self.hparams["indexer_head_dim"] + q = data_torch[:n_q] + k = data_torch[n_q:] + return [ + (self.format_tensor_name(gguf.MODEL_TENSOR.INDEXER_Q_PROJ, bid, ".weight"), q), + (self.format_tensor_name(gguf.MODEL_TENSOR.INDEXER_K_PROJ, bid, ".weight"), k), + ] + + # Gemma zero-centred gammas the inherited norm.weight rule misses + if name.endswith((".ple.norm_key.weight", ".ple.norm_query.weight", ".ple.norm_conv.weight", + ".indexer.q_layernorm.weight", ".indexer.k_layernorm.weight")): + return [(self.map_tensor_name(name), data_torch + 1)] + + if name.endswith(".ple.conv1d.weight"): + return [(self.map_tensor_name(name), data_torch.squeeze())] + + return super().modify_tensors(data_torch, name, bid) + + # the shards concatenate into a tensor of well over 100 GB + # use LazyChunkedTensor here, a single shard resident at a time + def _place_ple_shard(self, data_torch: Tensor, name: str) -> Iterable[tuple[str, Tensor]]: + + idx = int(name.rpartition(".shard_")[2].partition(".")[0]) + n_parts = self.hparams["split_ngram_parts"] + + self._ple_shards[idx] = name + self._ple_row_dim = int(data_torch.shape[-1]) + + if len(self._ple_shards) < n_parts: + return [] + + # the checkpoint may yield the shards in any order, the row order is by index + shards = [self._ple_shards[i] for i in sorted(self._ple_shards)] + rows = 0 + for shard in shards: + shape = self.model_tensors[shard]().shape + if int(shape[-1]) != self._ple_row_dim: + raise ValueError( + f"PLE shard {shard} has row dim {int(shape[-1])}, expected {self._ple_row_dim}") + rows += int(shape[0]) + + table = gguf.LazyChunkedTensor( + [self._load_ple_shard(shard) for shard in shards], + shape=(rows, self._ple_row_dim), + dtype=np.float32, + ) + gguf_name = gguf.TENSOR_NAMES[gguf.MODEL_TENSOR.PER_LAYER_TOKEN_EMBD] + return [(gguf_name + ".weight", cast(Tensor, table))] + + def _load_ple_shard(self, name: str): + def load() -> np.ndarray: + from .base import LazyTorchTensor + + # a fresh lazy tensor every call, or to_eager() memoizes every shard + eager = LazyTorchTensor.to_eager(self.model_tensors[name]()) + return eager.to(torch.float32).contiguous().numpy() + return load + + def prepare_tensors(self): + super().prepare_tensors() + n_parts = self.hparams.get("split_ngram_parts", 0) + if self._ple_shards and len(self._ple_shards) != n_parts: + raise ValueError( + f"got {len(self._ple_shards)} PLE embedding shards, expected {n_parts}" + ) + + +@ModelBase.register("Qwen4ExpForConditionalGeneration") +@ModelBase.example("Qwen/Qwen3.8-Flash-Next") +class Qwen4ExpVisionModel(Qwen3VLVisionModel): + """The vision tower is an unmodified Qwen3-VL ViT.""" diff --git a/conversion/rwkv.py b/conversion/rwkv.py index 2de0aa5346e9..9f0d34c95057 100644 --- a/conversion/rwkv.py +++ b/conversion/rwkv.py @@ -10,7 +10,7 @@ from .base import ModelBase, TextModel, gguf -@ModelBase.register("Rwkv6ForCausalLM") +@ModelBase.register("Rwkv6ForCausalLM", "RWKV6ForCausalLM") class Rwkv6Model(TextModel): model_arch = gguf.MODEL_ARCH.RWKV6 diff --git a/conversion/stablelm.py b/conversion/stablelm.py index 6e16378a031f..43ed6712096f 100644 --- a/conversion/stablelm.py +++ b/conversion/stablelm.py @@ -10,7 +10,7 @@ from .base import ModelBase, TextModel, gguf -@ModelBase.register("StableLmForCausalLM", "StableLMEpochForCausalLM", "LlavaStableLMEpochForCausalLM") +@ModelBase.register("StableLmForCausalLM", "StableLMEpochForCausalLM", "LlavaStableLMEpochForCausalLM", "StableLMForCausalLM") class StableLMModel(TextModel): model_arch = gguf.MODEL_ARCH.STABLELM diff --git a/conversion/starcoder.py b/conversion/starcoder.py index 0b4ffd84702a..d531ceb21e50 100644 --- a/conversion/starcoder.py +++ b/conversion/starcoder.py @@ -3,7 +3,7 @@ from .base import ModelBase, TextModel, gguf -@ModelBase.register("GPTBigCodeForCausalLM") +@ModelBase.register("GPTBigCodeForCausalLM", "GPTBigCodeModel", "GPTBigCodeLMHeadModel") class StarCoderModel(TextModel): model_arch = gguf.MODEL_ARCH.STARCODER @@ -18,6 +18,6 @@ def set_gguf_parameters(self): self.gguf_writer.add_file_type(self.ftype) -@ModelBase.register("Starcoder2ForCausalLM") +@ModelBase.register("Starcoder2ForCausalLM", "Starcoder2Model") class StarCoder2Model(TextModel): model_arch = gguf.MODEL_ARCH.STARCODER2 diff --git a/conversion/t5.py b/conversion/t5.py index 73dcfd1a2ced..7cbb4295f925 100644 --- a/conversion/t5.py +++ b/conversion/t5.py @@ -12,8 +12,8 @@ @ModelBase.register("T5WithLMHeadModel") -@ModelBase.register("T5ForConditionalGeneration") -@ModelBase.register("MT5ForConditionalGeneration") +@ModelBase.register("T5ForConditionalGeneration", "T5Model") +@ModelBase.register("MT5ForConditionalGeneration", "MT5Model") @ModelBase.register("UMT5ForConditionalGeneration") @ModelBase.register("UMT5Model") class T5Model(TextModel): diff --git a/conversion/zaya.py b/conversion/zaya.py new file mode 100644 index 000000000000..39b493bef09a --- /dev/null +++ b/conversion/zaya.py @@ -0,0 +1,459 @@ +# Copyright 2026 bong-water-water-bong +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +import re +import shutil +import tempfile + +from pathlib import Path +from typing import Callable, Iterable, TYPE_CHECKING + +if TYPE_CHECKING: + from torch import Tensor + +import torch + +from .base import ModelBase, TextModel, gguf, logger +from .qwenvl import Qwen2VLVisionModel + + + +def _special_tokens_as_control(tokenizer_dir, toktypes): + """Mark every token tokenizer.json lists as special: true as CONTROL. + + LlamaHfVocab leaves some of them NORMAL (ZAYA1's <|im_start|> and , which are also in the + base vocab), and llama.cpp matches special-token text in prompts only for CONTROL and + USER_DEFINED tokens, so the chat template's <|im_start|> reached the model spelled out as text.""" + import json + from pathlib import Path + path = Path(tokenizer_dir) / "tokenizer.json" + if not path.is_file(): + return toktypes + for added in json.loads(path.read_text(encoding="utf-8")).get("added_tokens", []): + tid = added.get("id") + if added.get("special") and isinstance(tid, int) and 0 <= tid < len(toktypes): + toktypes[tid] = gguf.TokenType.CONTROL + return toktypes + +@ModelBase.register("ZayaForCausalLM") +class ZayaModel(TextModel): + """Zyphra ZAYA1 (transformers naming): every layer is CCA attention + a top-1 MoE.""" + model_arch = gguf.MODEL_ARCH.ZAYA + + # checkpoint suffix -> (tensor, name suffix, squeeze dim 1) + _LAYER_MAP: dict[str, tuple[gguf.MODEL_TENSOR, str, bool]] = { + "input_layernorm.weight": (gguf.MODEL_TENSOR.ATTN_NORM, ".weight", False), + "post_attention_layernorm.weight": (gguf.MODEL_TENSOR.ATTN_POST_NORM, ".weight", False), + "self_attn.qkv_proj.q_proj.weight": (gguf.MODEL_TENSOR.ATTN_Q, ".weight", False), + "self_attn.qkv_proj.k_proj.weight": (gguf.MODEL_TENSOR.ATTN_K, ".weight", False), + "self_attn.qkv_proj.v_proj_current.weight": (gguf.MODEL_TENSOR.CCA_VAL_PROJ1, ".weight", False), + "self_attn.qkv_proj.v_proj_delayed.weight": (gguf.MODEL_TENSOR.CCA_VAL_PROJ2, ".weight", False), + "self_attn.qkv_proj.conv_qk_depthwise.weight": (gguf.MODEL_TENSOR.SSM_CONV1D, ".weight", True), + "self_attn.qkv_proj.conv_qk_depthwise.bias": (gguf.MODEL_TENSOR.SSM_CONV1D, ".bias", False), + "self_attn.qkv_proj.conv_qk_grouped.weight": (gguf.MODEL_TENSOR.CCA_CONV_GRP, ".weight", False), + "self_attn.qkv_proj.conv_qk_grouped.bias": (gguf.MODEL_TENSOR.CCA_CONV_GRP, ".bias", False), + "self_attn.qk_norm.temp": (gguf.MODEL_TENSOR.CCA_K_SCALE, ".weight", False), + "self_attn.o_proj.weight": (gguf.MODEL_TENSOR.ATTN_OUT, ".weight", False), + "mlp.gate.down_proj.weight": (gguf.MODEL_TENSOR.FFN_GATE_INP, ".weight", False), + "mlp.gate.down_proj.bias": (gguf.MODEL_TENSOR.FFN_GATE_INP, ".bias", False), + "mlp.gate.router_states_scale": (gguf.MODEL_TENSOR.ZAYA_ROUTER_EDA_SCALE, ".weight", False), + "mlp.gate.router_mlp.norm.weight": (gguf.MODEL_TENSOR.FFN_NORM, ".weight", False), + "mlp.gate.router_mlp.fc1.weight": (gguf.MODEL_TENSOR.FFN_GATE, ".weight", False), + "mlp.gate.router_mlp.fc1.bias": (gguf.MODEL_TENSOR.FFN_GATE, ".bias", False), + "mlp.gate.router_mlp.fc2.weight": (gguf.MODEL_TENSOR.ZAYA_ROUTER_MLP2, ".weight", False), + "mlp.gate.router_mlp.fc2.bias": (gguf.MODEL_TENSOR.ZAYA_ROUTER_MLP2, ".bias", False), + "mlp.gate.router_mlp.out_proj.weight": (gguf.MODEL_TENSOR.ZAYA_ROUTER_MLP4, ".weight", False), + "mlp.gate.balancing_biases": (gguf.MODEL_TENSOR.ZAYA_ROUTER_BIASES, ".weight", False), + "mlp.experts.gate_up_proj": (gguf.MODEL_TENSOR.FFN_GATE_UP_EXP, ".weight", False), + "mlp.experts.down_proj": (gguf.MODEL_TENSOR.FFN_DOWN_EXP, ".weight", False), + "post_attention_residual_scale.hidden_states_scale": (gguf.MODEL_TENSOR.RES_SCALE_HS, ".weight", False), + "post_attention_residual_scale.hidden_states_bias": (gguf.MODEL_TENSOR.RES_SCALE_HS, ".bias", False), + "post_attention_residual_scale.residual_scale": (gguf.MODEL_TENSOR.RES_SCALE_RES, ".weight", False), + "post_attention_residual_scale.residual_bias": (gguf.MODEL_TENSOR.RES_SCALE_RES, ".bias", False), + "post_mlp_residual_scale.hidden_states_scale": (gguf.MODEL_TENSOR.RES_SCALE_HS_MLP, ".weight", False), + "post_mlp_residual_scale.hidden_states_bias": (gguf.MODEL_TENSOR.RES_SCALE_HS_MLP, ".bias", False), + "post_mlp_residual_scale.residual_scale": (gguf.MODEL_TENSOR.RES_SCALE_RES_MLP, ".weight", False), + "post_mlp_residual_scale.residual_bias": (gguf.MODEL_TENSOR.RES_SCALE_RES_MLP, ".bias", False), + # ZAYA1-VL's vision-only LoRA (A: down to the rank, B: back up; experts stacked) + "vlora.q_a": (gguf.MODEL_TENSOR.ZAYA_VLORA_Q_A, ".weight", False), + "vlora.q_b": (gguf.MODEL_TENSOR.ZAYA_VLORA_Q_B, ".weight", False), + "vlora.k_a": (gguf.MODEL_TENSOR.ZAYA_VLORA_K_A, ".weight", False), + "vlora.k_b": (gguf.MODEL_TENSOR.ZAYA_VLORA_K_B, ".weight", False), + "vlora.v1_a": (gguf.MODEL_TENSOR.ZAYA_VLORA_V1_A, ".weight", False), + "vlora.v1_b": (gguf.MODEL_TENSOR.ZAYA_VLORA_V1_B, ".weight", False), + "vlora.v2_a": (gguf.MODEL_TENSOR.ZAYA_VLORA_V2_A, ".weight", False), + "vlora.v2_b": (gguf.MODEL_TENSOR.ZAYA_VLORA_V2_B, ".weight", False), + "vlora.o_a": (gguf.MODEL_TENSOR.ZAYA_VLORA_O_A, ".weight", False), + "vlora.o_b": (gguf.MODEL_TENSOR.ZAYA_VLORA_O_B, ".weight", False), + "vlora.gate_up_exps_a": (gguf.MODEL_TENSOR.ZAYA_VLORA_UP_EXPS_A, ".weight", False), + "vlora.gate_up_exps_b": (gguf.MODEL_TENSOR.ZAYA_VLORA_UP_EXPS_B, ".weight", False), + "vlora.down_exps_a": (gguf.MODEL_TENSOR.ZAYA_VLORA_DOWN_EXPS_A, ".weight", False), + "vlora.down_exps_b": (gguf.MODEL_TENSOR.ZAYA_VLORA_DOWN_EXPS_B, ".weight", False), + } + + def __init__(self, *args, **kwargs): + # ZAYA1-base, ZAYA1-reasoning-base and the *-legacy repos keep Zyphra's Megatron-style + # checkpoint: attention and MoE are separate "layers", with per-layer config lists. + # Normalize the config here, and the tensors in index_tensors, to the transformers layout. + hparams = kwargs.get("hparams") or ModelBase.load_hparams(args[0], False) + self._legacy = "zaya_layers" in hparams or "cca" in hparams + if self._legacy: + hparams = self._legacy_hparams(hparams) + kwargs["hparams"] = hparams + super().__init__(*args, **kwargs) + # tensors are prepared before set_vocab(), and the embedding is trimmed to this + self._n_vocab = gguf.LlamaHfVocab(self._tokenizer_dir()).vocab_size + + def _tokenizer_dir(self) -> Path: + # transformers' AutoTokenizer reads config.json, and its ZayaConfig rejects the legacy + # config (rope_scaling: false); load the tokenizer from its files alone + if not self._legacy: + return self.dir_model + if not hasattr(self, "_tok_tmp"): + self._tok_tmp = tempfile.TemporaryDirectory(prefix="zaya-tok-") + for f in ("tokenizer.json", "tokenizer_config.json", "special_tokens_map.json", "chat_template.jinja", "generation_config.json"): + if (self.dir_model / f).is_file(): + shutil.copy(self.dir_model / f, self._tok_tmp.name) + return Path(self._tok_tmp.name) + + @staticmethod + def _legacy_hparams(hp: dict) -> dict: + def first(v): # per-layer lists hold 0 on the layers the value does not apply to + return max(v) if isinstance(v, list) else v + for flag in ("zaya_use_eda", "zaya_use_mod", "scale_residual_merge", "cca", "gated_linear_unit"): + if hp.get(flag, True) is not True: + raise ValueError(f"zaya: legacy checkpoint with {flag}={hp[flag]} is not supported") + zl = hp.get("zaya_layers") + n_half = len(zl) if zl else hp["num_hidden_layers"] + if "vision_config" in hp: + n_half = 2 * hp["num_hidden_layers"] # ZAYA1-VL: layers.{i}.attn / layers.{i}.mlp + n_block = n_half // 2 + n_head = first(hp["cca_num_q_heads"]) if "cca_num_q_heads" in hp else hp["num_attention_heads"] + n_head_kv = first(hp["num_query_groups_list"]) if "num_query_groups_list" in hp else hp["num_query_groups"] + n_expert = max(x for x in zl if isinstance(x, int)) if zl else hp["num_experts"] + ffn = first(hp["ffn_hidden_size_list"]) if "ffn_hidden_size_list" in hp else hp["ffn_hidden_size"] + theta = hp.get("rope_theta", hp.get("rotary_base", 1e6)) + rotary = hp.get("partial_rotary_factor", hp.get("rope_pct", 0.5)) + out = dict(hp) + out.update({ + "num_hidden_layers": n_block, + "num_attention_heads": n_head, + "num_key_value_heads": n_head_kv, + "head_dim": hp.get("head_dim") or hp["kv_channels"], + "num_experts": n_expert, + "num_experts_per_tok": hp.get("moe_router_topk", 1), + "moe_intermediate_size": ffn // 2, # fc1 holds gate and up + "router_hidden_size": first(hp["zaya_mlp_expansion"]), + "rms_norm_eps": hp.get("norm_epsilon", 1e-5), + "layer_types": ["hybrid"] * n_block, + "rope_parameters": {"hybrid": {"rope_theta": theta, "partial_rotary_factor": rotary}}, + }) + swa = hp.get("swa_layers") + if swa and any(swa): + # per half-layer window on the attention layers; Megatron's window excludes the query + # position, transformers' sliding_window includes it (ZAYA1-74B-preview: 4096 -> 4097) + out["layer_types"] = ["hybrid_sliding" if swa[2 * i] else "hybrid" for i in range(n_block)] + out["sliding_window"] = max(swa) + 1 + out["rope_parameters"]["hybrid_sliding"] = {"rope_theta": hp.get("swa_rotary_base", theta), "partial_rotary_factor": rotary} + return out + + def index_tensors(self, remote_hf_model_id: str | None = None) -> dict[str, Callable[[], Tensor]]: + tensors = super().index_tensors(remote_hf_model_id=remote_hf_model_id) + if not self._legacy: + return tensors + n_block, n_expert = self.hparams["num_hidden_layers"], self.hparams["num_experts"] + vl: dict[str, Callable[[], Tensor]] = {} + if "vision_config" in self.hparams: + # ZAYA1-VL: blocks hold attn. and mlp. sublayers, which are the half-layers 2i and 2i+1; + # the vision tower goes to the mmproj + renamed: dict[str, Callable[[], Tensor]] = {} + for name, gen in tensors.items(): + if name.startswith("vision_tower."): + continue + m = re.match(r"model\.layers\.(\d+)\.(attn|mlp)\.(.*)", name) + if m is None: + renamed[name] = gen + continue + i, part, rest = int(m.group(1)), m.group(2), m.group(3) + lora = re.match(r"self_attn\.(?:qkv\.)?(lora_linear_q|lora_linear_k|lora_val_proj1|lora_val_proj2|lora_linear_o)\.([01])\.weight", rest) + if lora: + short = {"lora_linear_q": "q", "lora_linear_k": "k", "lora_val_proj1": "v1", + "lora_val_proj2": "v2", "lora_linear_o": "o"}[lora.group(1)] + vl[f"model.layers.{i}.vlora.{short}_{'ab'[int(lora.group(2))]}"] = gen + continue + if ".lora_fc" in rest: + continue # stacked below + renamed[f"model.layers.{2 * i + (part == 'mlp')}.{rest}"] = gen + for i in range(n_block): + pre = f"model.layers.{i}.mlp.zaya_block.experts.local_experts." + for src, dst in (("lora_fc1", "gate_up_exps"), ("lora_fc2", "down_exps")): + for seq, ab in (("0", "a"), ("1", "b")): + gens = [tensors[f"{pre}{e}.{src}.{seq}.weight"] for e in range(n_expert)] + vl[f"model.layers.{i}.vlora.{dst}_{ab}"] = lambda gens=gens: torch.stack([g() for g in gens]) + tensors = renamed + out: dict[str, Callable[[], Tensor]] = {} + rename = {"model.embed_tokens.weight": "model.embed_tokens.weight", "model.final_norm.weight": "model.norm.weight"} + for k in ("hidden_states_scale", "hidden_states_bias"): + rename[f"model.layers.0.res_scale.{k}"] = f"model.input_{k}" + for k in ("hidden_states_scale", "hidden_states_bias", "residual_scale", "residual_bias"): + rename[f"model.res_scale.{k}"] = f"model.layers.{n_block - 1}.post_mlp_residual_scale.{k}" + attn = { + "input_norm.weight": "input_layernorm.weight", + "self_attn.o_proj.weight": "self_attn.o_proj.weight", + "self_attn.qkv.temp": "self_attn.qk_norm.temp", + "self_attn.qkv.linear_q.weight": "self_attn.qkv_proj.q_proj.weight", + "self_attn.qkv.linear_k.weight": "self_attn.qkv_proj.k_proj.weight", + "self_attn.qkv.val_proj1.weight": "self_attn.qkv_proj.v_proj_current.weight", + "self_attn.qkv.val_proj2.weight": "self_attn.qkv_proj.v_proj_delayed.weight", + "self_attn.qkv.conv_qk.0.weight": "self_attn.qkv_proj.conv_qk_depthwise.weight", + "self_attn.qkv.conv_qk.0.bias": "self_attn.qkv_proj.conv_qk_depthwise.bias", + "self_attn.qkv.conv_qk.1.weight": "self_attn.qkv_proj.conv_qk_grouped.weight", + "self_attn.qkv.conv_qk.1.bias": "self_attn.qkv_proj.conv_qk_grouped.bias", + } + moe = { + "input_norm.weight": "post_attention_layernorm.weight", + "zaya_block.router.balancing_biases": "mlp.gate.balancing_biases", + "zaya_block.router.down_proj.weight": "mlp.gate.down_proj.weight", + "zaya_block.router.down_proj.bias": "mlp.gate.down_proj.bias", + "zaya_block.router.router_states_scale": "mlp.gate.router_states_scale", + "zaya_block.router.rmsnorm_eda.weight": "mlp.gate.router_mlp.norm.weight", + "zaya_block.router.router_mlp.0.weight": "mlp.gate.router_mlp.fc1.weight", + "zaya_block.router.router_mlp.0.bias": "mlp.gate.router_mlp.fc1.bias", + "zaya_block.router.router_mlp.2.weight": "mlp.gate.router_mlp.fc2.weight", + "zaya_block.router.router_mlp.2.bias": "mlp.gate.router_mlp.fc2.bias", + "zaya_block.router.router_mlp.4.weight": "mlp.gate.router_mlp.out_proj.weight", + } + res = ("hidden_states_scale", "hidden_states_bias", "residual_scale", "residual_bias") + for i in range(n_block): + a, m = f"model.layers.{2 * i}.", f"model.layers.{2 * i + 1}." + b = f"model.layers.{i}." + rename.update({a + k: b + v for k, v in attn.items()}) + rename.update({m + k: b + v for k, v in moe.items()}) + # a half-layer's res_scale merges the residual in front of it: the MoE half-layer's is + # the block's post-attention scale, the next attention half-layer's its post-MLP scale + rename.update({f"{m}res_scale.{k}": f"{b}post_attention_residual_scale.{k}" for k in res}) + if i + 1 < n_block: + rename.update({f"model.layers.{2 * i + 2}.res_scale.{k}": f"{b}post_mlp_residual_scale.{k}" for k in res}) + for name, gen in tensors.items(): + if ".local_experts." in name: + continue + if name == "lm_head.weight": + out[name] = gen + continue + if name not in rename: + raise ValueError(f"zaya: unmapped legacy tensor {name}") + out[rename[name]] = gen + for i in range(n_block): + pre = f"model.layers.{2 * i + 1}.zaya_block.experts.local_experts." + for src, dst in (("linear_fc1", "gate_up_proj"), ("linear_fc2", "down_proj")): + gens = [tensors[f"{pre}{e}.{src}.weight"] for e in range(n_expert)] + out[f"model.layers.{i}.mlp.experts.{dst}"] = lambda gens=gens: torch.stack([g() for g in gens]) + out.update(vl) + return out + + def set_vocab(self): + # the Gemma 3/4 BPE tokenizer; added tokens that are not special (such as "\n") stay + # USER_DEFINED so they are matched in prompts and rendered in output + vocab = gguf.LlamaHfVocab(self._tokenizer_dir()) + tokens, scores, toktypes = [], [], [] + for text, score, toktype in vocab.all_tokens(): + tokens.append(text) + scores.append(score) + toktypes.append(toktype) + + self.gguf_writer.add_tokenizer_model("gemma4") + self.gguf_writer.add_token_list(tokens) + self.gguf_writer.add_token_scores(scores) + self.gguf_writer.add_token_types(_special_tokens_as_control(self._tokenizer_dir(), toktypes)) + + special_vocab = gguf.SpecialVocab(self.dir_model, load_merges=True) + special_vocab.add_to_gguf(self.gguf_writer) + self.gguf_writer.add_add_space_prefix(False) + self.gguf_writer.add_add_bos_token(True) + + def set_gguf_parameters(self): + hp = self.hparams + head_dim = hp["head_dim"] + n_qk = (hp["num_attention_heads"] + hp["num_key_value_heads"]) * head_dim + rope = hp.get("rope_parameters", {}).get("hybrid", hp.get("rope_parameters", {})) + rotary = rope.get("partial_rotary_factor", hp.get("partial_rotary_factor", 0.5)) + + self.gguf_writer.add_block_count(self.block_count) + self.gguf_writer.add_context_length(hp["max_position_embeddings"]) + self.gguf_writer.add_embedding_length(hp["hidden_size"]) + self.gguf_writer.add_feed_forward_length(hp["moe_intermediate_size"]) + self.gguf_writer.add_head_count(hp["num_attention_heads"]) + self.gguf_writer.add_head_count_kv(hp["num_key_value_heads"]) + self.gguf_writer.add_key_length(head_dim) + self.gguf_writer.add_value_length(head_dim) + self.gguf_writer.add_layer_norm_rms_eps(hp.get("rms_norm_eps", 1e-5)) + self.gguf_writer.add_rope_dimension_count(int(rotary * head_dim)) + self.gguf_writer.add_rope_freq_base(float(rope.get("rope_theta", hp.get("rope_theta", 1e6)))) + self.gguf_writer.add_expert_count(hp["num_experts"]) + self.gguf_writer.add_expert_used_count(hp.get("num_experts_per_tok", 1)) + self.gguf_writer.add_expert_feed_forward_length(hp["router_hidden_size"]) + self.gguf_writer.add_ssm_conv_kernel(hp.get("cca_time0", 2)) + self.gguf_writer.add_ssm_state_size(2 * n_qk + hp["hidden_size"]) + self.gguf_writer.add_ssm_inner_size(1) + # ZAYA1-74B: layer_types alternates hybrid_sliding / hybrid; sliding layers attend to the + # last sliding_window positions and use their own rope theta (rope_parameters.hybrid_sliding) + layer_types = hp.get("layer_types") or [] + if hp.get("sliding_window") and "hybrid_sliding" in layer_types: + self.gguf_writer.add_sliding_window(int(hp["sliding_window"])) + self.gguf_writer.add_sliding_window_pattern([t == "hybrid_sliding" for t in layer_types]) + swa_rope = hp.get("rope_parameters", {}).get("hybrid_sliding", {}) + if "rope_theta" in swa_rope: + self.gguf_writer.add_rope_freq_base_swa(float(swa_rope["rope_theta"])) + self.gguf_writer.add_file_type(self.ftype) + logger.info(f"zaya: {self.block_count} layers, rope theta {rope.get('rope_theta')}, rotary {rotary}") + + def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]: + if name == "model.embed_tokens.weight": + # the embedding has padding rows past the tokenizer's vocab + yield self.format_tensor_name(gguf.MODEL_TENSOR.TOKEN_EMBD), data_torch[:self._n_vocab] + return + if name == "model.norm.weight": + yield self.format_tensor_name(gguf.MODEL_TENSOR.OUTPUT_NORM), data_torch + return + if name == "model.input_hidden_states_scale": + yield self.format_tensor_name(gguf.MODEL_TENSOR.INPUT_HIDDEN_STATES_SCALE), data_torch + return + if name == "model.input_hidden_states_bias": + yield self.format_tensor_name(gguf.MODEL_TENSOR.INPUT_HIDDEN_STATES_SCALE, suffix=".bias"), data_torch + return + if name == "lm_head.weight": + return # tied to the embedding + if bid is not None: + suffix = name.split(f"model.layers.{bid}.", 1)[-1] + if suffix in self._LAYER_MAP: + tensor, tsuffix, squeeze = self._LAYER_MAP[suffix] + if squeeze: + data_torch = data_torch.squeeze(1) + if tensor == gguf.MODEL_TENSOR.CCA_CONV_GRP and tsuffix == ".weight": + # (OC, IC_G, taps) -> tap-major (taps, OC, IC_G): each tap's weights are one + # contiguous [IC_G, OC] block, which the graph applies as one batched matmul + data_torch = data_torch.permute(2, 0, 1).contiguous() + yield self.format_tensor_name(tensor, bid, suffix=tsuffix), data_torch + return + raise ValueError(f"zaya: unmapped tensor {name}") + + +@ModelBase.register("Zaya1VLForConditionalGeneration") +class ZayaVLModel(ZayaModel): + """Zyphra ZAYA1-VL: the ZAYA1 language model (Megatron-style checkpoint) with vision-only LoRA + on CCA and on every expert, used on image tokens; the vision tower goes to the mmproj.""" + model_arch = gguf.MODEL_ARCH.ZAYA + + # Zyphra's template only takes content as a list of image and text parts. llama.cpp passes a + # message with media as a string holding a marker, <__media___> with an id random per server, + # where each image was; this one takes both, and puts the images (the markers) in front of the + # user turn, as Zyphra's does (mtmd wraps each image in <|vision_start|> ... <|vision_end|>). + _CHAT_TEMPLATE = ( + '{%- for message in messages -%}' + "{%- if message['content'] is string -%}" + "{%- set ns = namespace(text=message['content'], media='') -%}" + '{%- else -%}' + "{%- set ns = namespace(text=message['content'] | selectattr('type', 'equalto', 'text') | map(attribute='text') | join(''), media='') -%}" + "{%- for c in message['content'] | selectattr('type', 'equalto', 'image') -%}" + "{%- set ns.media = ns.media ~ '<|vision_start|><|vision_end|>\\n' -%}" + '{%- endfor -%}' + '{%- endif -%}' + "{%- set segs = ns.text.split('<__media_') -%}" + '{%- if segs | length > 1 -%}' + '{%- set ns.text = segs[0] -%}' + '{%- for seg in segs[1:] -%}' + "{%- set bits = seg.split('>') -%}" + "{%- set ns.media = ns.media ~ '<__media_' ~ bits[0] ~ '>\\n' -%}" + "{%- set ns.text = ns.text ~ (bits[1:] | join('>')) -%}" + '{%- endfor -%}' + '{%- set ns.text = ns.text | trim -%}' + '{%- endif -%}' + "{%- if message['role'] == 'user' -%}" + "{{ ns.media ~ '<|im_start|>user\\n' ~ ns.text ~ '<|im_end|>\\n' }}" + '{%- else -%}' + "{{ '<|im_start|>' ~ message['role'] ~ '\\n' ~ ns.text ~ '<|im_end|>' }}" + '{%- endif -%}' + '{%- endfor -%}' + '{%- if add_generation_prompt -%}' + "{{ '<|im_start|>assistant\\n' }}" + '{%- endif -%}' + ) + + def set_vocab(self): + vocab = gguf.LlamaHfVocab(self._tokenizer_dir()) + tokens, scores, toktypes = [], [], [] + for text, score, toktype in vocab.all_tokens(): + tokens.append(text) + scores.append(score) + toktypes.append(toktype) + + self.gguf_writer.add_tokenizer_model("gemma4") + self.gguf_writer.add_token_list(tokens) + self.gguf_writer.add_token_scores(scores) + self.gguf_writer.add_token_types(_special_tokens_as_control(self._tokenizer_dir(), toktypes)) + + special_vocab = gguf.SpecialVocab(self.dir_model, load_merges=True) + special_vocab.chat_template = self._CHAT_TEMPLATE + special_vocab.add_to_gguf(self.gguf_writer) + self.gguf_writer.add_add_space_prefix(False) + self.gguf_writer.add_add_bos_token(True) + + def set_gguf_parameters(self): + super().set_gguf_parameters() + if self.hparams.get("vision_lora"): + arch = self.gguf_writer.arch + self.gguf_writer.add_uint32(f"{arch}.vision_lora.attention_rank", int(self.hparams["vision_lora_rank_attn"])) + self.gguf_writer.add_uint32(f"{arch}.vision_lora.ffn_rank", int(self.hparams["vision_lora_rank_mlp"])) + + +@ModelBase.register("Zaya1VLForConditionalGeneration") +class ZayaVLVisionModel(Qwen2VLVisionModel): + """ZAYA1-VL's vision tower: Qwen2.5-VL's, under vision_tower.""" + + def __init__(self, *args, **kwargs): + # the checkpoint's vision_config lists only what differs from Qwen2.5-VL's defaults + # (no depth, heads, window pattern); fill it as transformers does + from transformers.models.qwen2_5_vl.configuration_qwen2_5_vl import Qwen2_5_VLVisionConfig + hparams = kwargs.get("hparams") or ModelBase.load_hparams(args[0], False) + vc = Qwen2_5_VLVisionConfig(**hparams["vision_config"]).to_dict() + kwargs["hparams"] = {**hparams, "vision_config": {**vc, **hparams["vision_config"]}} + super().__init__(*args, **kwargs) + assert self.hparams_vision is not None + # Qwen2VLVisionModel picks the projector from the top-level model_type (zaya1_vl here) + self.global_config["model_type"] = self.hparams_vision["model_type"] + + def set_gguf_parameters(self): + super().set_gguf_parameters() + # ZAYA1-VL attends to each image bidirectionally (the image tokens are not causal) + self.gguf_writer.add_vision_decode_non_causal(True) + + @classmethod + def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None: + name, gen = item + if not name.startswith("vision_tower."): + return None + return super().filter_tensors(("visual." + name.removeprefix("vision_tower."), gen)) + + def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]: + if "patch_embed.proj.weight" in name and data_torch.shape[2] == 1: + # temporal_patch_size 1: clip's Qwen2-VL graph sums two patch convs over the same + # still image (Qwen feeds the frame twice), so the second one is all zeros + w = data_torch[:, :, 0, ...] + yield (gguf.TENSOR_NAMES[gguf.MODEL_TENSOR.V_ENC_EMBD_PATCH] + ".weight", w) + yield (gguf.TENSOR_NAMES[gguf.MODEL_TENSOR.V_ENC_EMBD_PATCH] + ".weight.1", torch.zeros_like(w)) + return + yield from super().modify_tensors(data_torch, name, bid) diff --git a/convert_hf_to_gguf.py b/convert_hf_to_gguf.py index 2c5e62a16fbe..7a0d20c8de94 100755 --- a/convert_hf_to_gguf.py +++ b/convert_hf_to_gguf.py @@ -184,7 +184,7 @@ def main() -> None: if args.remote: hf_repo_id = args.model from huggingface_hub import snapshot_download - allowed_patterns = ["LICENSE", "*.json", "*.md", "*.txt", "tokenizer.model"] + allowed_patterns = ["LICENSE", "*.json", "*.jinja", "*.md", "*.txt", "tokenizer.model"] if args.sentence_transformers_dense_modules: # include sentence-transformers dense modules safetensors files allowed_patterns.append("*.safetensors") diff --git a/ggml/CMakeLists.txt b/ggml/CMakeLists.txt index 159da3afa0b0..c45eb3dc0772 100644 --- a/ggml/CMakeLists.txt +++ b/ggml/CMakeLists.txt @@ -213,6 +213,7 @@ set (GGML_CUDA_COMPRESSION_MODE "size" CACHE STRING set_property(CACHE GGML_CUDA_COMPRESSION_MODE PROPERTY STRINGS "none;speed;balance;size") option(GGML_HIP "ggml: use HIP" OFF) +option(GGML_HRX "ggml: use HRX" OFF) option(GGML_HIP_GRAPHS "ggml: use HIP graph" ON) option(GGML_HIP_RCCL "ggml: use ROCm Collective Comm. Library" OFF) option(GGML_HIP_NO_VMM "ggml: do not try to use HIP VMM" ON) @@ -285,8 +286,15 @@ option(GGML_BUILD_EXAMPLES "ggml: build examples" ${GGML_STANDALONE}) set(CMAKE_C_STANDARD 11) set(CMAKE_C_STANDARD_REQUIRED true) -set(CMAKE_CXX_STANDARD 17) +# C++26 where the compiler has it (amdclang 23, g++ 15); older toolchains keep C++17 +if("cxx_std_26" IN_LIST CMAKE_CXX_COMPILE_FEATURES) + set(CMAKE_CXX_STANDARD 26) +else() + set(CMAKE_CXX_STANDARD 17) +endif() set(CMAKE_CXX_STANDARD_REQUIRED true) +# no C++20 module scanning: no module sources, and some toolchains lack clang-scan-deps +set(CMAKE_CXX_SCAN_FOR_MODULES OFF) set(THREADS_PREFER_PTHREAD_FLAG ON) diff --git a/ggml/include/ggml-hrx.h b/ggml/include/ggml-hrx.h new file mode 100644 index 000000000000..d40c095b75bc --- /dev/null +++ b/ggml/include/ggml-hrx.h @@ -0,0 +1,26 @@ +#pragma once + +#include "ggml-backend.h" + +#ifdef __cplusplus +extern "C" { +#endif + +struct ggml_backend_hrx_cache_stats { + uint64_t graph_program_builds; + uint64_t graph_program_hits; + uint64_t prepared_program_builds; + uint64_t prepared_program_hits; +}; + +GGML_BACKEND_API ggml_backend_t ggml_backend_hrx_init(size_t device); +GGML_BACKEND_API bool ggml_backend_is_hrx(ggml_backend_t backend); +GGML_BACKEND_API int ggml_backend_hrx_get_device_count(void); +GGML_BACKEND_API ggml_backend_buffer_type_t ggml_backend_hrx_buffer_type(size_t device); +GGML_BACKEND_API bool ggml_backend_hrx_get_cache_stats(ggml_backend_t backend, + struct ggml_backend_hrx_cache_stats * stats); +GGML_BACKEND_API ggml_backend_reg_t ggml_backend_hrx_reg(void); + +#ifdef __cplusplus +} +#endif diff --git a/ggml/include/ggml.h b/ggml/include/ggml.h index 35f0c44ec421..00625779515d 100644 --- a/ggml/include/ggml.h +++ b/ggml/include/ggml.h @@ -430,7 +430,9 @@ extern "C" { GGML_TYPE_NVFP4 = 40, // NVFP4 (4 blocks, E4M3 scale) GGML_TYPE_Q1_0 = 41, GGML_TYPE_Q2_0 = 42, - GGML_TYPE_COUNT = 43, + GGML_TYPE_PQ2_0 = 142, // PrismML group-128 2-bit (ggml-prism.h); 43..141 unused + GGML_TYPE_PTQ1_0 = 143, // PrismML group-128 ternary (ggml-prism.h) + GGML_TYPE_COUNT = 144, }; // precision @@ -475,6 +477,8 @@ extern "C" { GGML_FTYPE_MOSTLY_NVFP4 = 26, // except 1d tensors GGML_FTYPE_MOSTLY_Q1_0 = 27, // except 1d tensors GGML_FTYPE_MOSTLY_Q2_0 = 28, // except 1d tensors + GGML_FTYPE_MOSTLY_PQ2_0 = 128, // except 1d tensors (PrismML) + GGML_FTYPE_MOSTLY_PTQ1_0 = 129, // except 1d tensors (PrismML) }; // available tensor operations: @@ -2627,11 +2631,21 @@ extern "C" { struct ggml_tensor * x, struct ggml_tensor * weights); + // hc_pre with a per-element gate (Qwen3.8-Flash-Next): gate [n_embd, hc, n_tokens] + // result[i, t] = scale*sum_h x[i, h, t]*sigmoid(gate[i, h, t]) + // + GGML_API struct ggml_tensor * ggml_dsv4_hc_pre_gated( + struct ggml_context * ctx, + struct ggml_tensor * x, + struct ggml_tensor * gate, + float scale); + // hc_post: x [n_embd, n_tokens], residual [n_embd, hc, n_tokens], // post [hc, n_tokens], comb [dst_hc, src_hc, n_tokens] // -> [n_embd, hc, n_tokens] // result[i, dst, t] = x[i, t]*post[dst, t] // + sum_src residual[i, src, t]*comb[dst, src, t] + // comb == NULL uses the identity: result[i, dst, t] = x[i, t]*post[dst, t] + residual[i, dst, t] // GGML_API struct ggml_tensor * ggml_dsv4_hc_post( struct ggml_context * ctx, diff --git a/ggml/src/CMakeLists.txt b/ggml/src/CMakeLists.txt index 82e9480c2f24..1716486ab845 100644 --- a/ggml/src/CMakeLists.txt +++ b/ggml/src/CMakeLists.txt @@ -206,6 +206,8 @@ add_library(ggml-base ggml-threading.h ggml-quants.c ggml-quants.h + ggml-prism-quants.c + ggml-prism.h gguf.cpp) set_target_properties(ggml-base PROPERTIES @@ -475,6 +477,7 @@ ggml_add_backend(CANN) ggml_add_backend(CUDA) ggml_add_backend(ET) ggml_add_backend(HIP) +ggml_add_backend(HRX) ggml_add_backend(METAL) ggml_add_backend(MUSA) ggml_add_backend(RPC) diff --git a/ggml/src/ggml-backend-reg.cpp b/ggml/src/ggml-backend-reg.cpp index e5959467071d..87a532b325c6 100644 --- a/ggml/src/ggml-backend-reg.cpp +++ b/ggml/src/ggml-backend-reg.cpp @@ -7,6 +7,7 @@ #include #include #include +#include #include #include #include @@ -34,6 +35,10 @@ #include "ggml-cuda.h" #endif +#ifdef GGML_USE_HRX +#include "ggml-hrx.h" +#endif + #ifdef GGML_USE_METAL #include "ggml-metal.h" #endif @@ -92,6 +97,16 @@ namespace fs = std::filesystem; +// UTF-8 C string -> path; fs::u8path is deprecated since C++20 +static fs::path path_from_u8(const char * str) { +#if defined(__cpp_lib_char8_t) + const std::string_view view(str); + return fs::path(std::u8string(view.begin(), view.end())); +#else + return fs::u8path(str); +#endif +} + static std::string path_str(const fs::path & path) { try { #if defined(__cpp_lib_char8_t) @@ -120,6 +135,9 @@ struct ggml_backend_registry { #ifdef GGML_USE_CUDA register_backend(ggml_backend_cuda_reg()); #endif +#ifdef GGML_USE_HRX + register_backend(ggml_backend_hrx_reg()); +#endif #ifdef GGML_USE_METAL register_backend(ggml_backend_metal_reg()); #endif @@ -463,36 +481,36 @@ static fs::path get_executable_path() { static fs::path backend_filename_prefix() { #ifdef _WIN32 - return fs::u8path("ggml-"); + return path_from_u8("ggml-"); #else - return fs::u8path("libggml-"); + return path_from_u8("libggml-"); #endif } static fs::path backend_filename_extension() { #ifdef _WIN32 - return fs::u8path(".dll"); + return path_from_u8(".dll"); #else - return fs::u8path(".so"); + return path_from_u8(".so"); #endif } static ggml_backend_reg_t ggml_backend_load_best(const char * name, bool silent, const char * user_search_path) { // enumerate all the files that match [lib]ggml-name-*.[so|dll] in the search paths - const fs::path name_path = fs::u8path(name); - const fs::path file_prefix = backend_filename_prefix().native() + name_path.native() + fs::u8path("-").native(); + const fs::path name_path = path_from_u8(name); + const fs::path file_prefix = backend_filename_prefix().native() + name_path.native() + path_from_u8("-").native(); const fs::path file_extension = backend_filename_extension(); std::vector search_paths; if (user_search_path == nullptr) { #ifdef GGML_BACKEND_DIR - search_paths.push_back(fs::u8path(GGML_BACKEND_DIR)); + search_paths.push_back(path_from_u8(GGML_BACKEND_DIR)); #endif // default search paths: executable directory, current directory search_paths.push_back(get_executable_path()); search_paths.push_back(fs::current_path()); } else { - search_paths.push_back(fs::u8path(user_search_path)); + search_paths.push_back(path_from_u8(user_search_path)); } int best_score = 0; diff --git a/ggml/src/ggml-cpu/CMakeLists.txt b/ggml/src/ggml-cpu/CMakeLists.txt index 836bae4d05a7..ca9d692e57b9 100644 --- a/ggml/src/ggml-cpu/CMakeLists.txt +++ b/ggml/src/ggml-cpu/CMakeLists.txt @@ -35,6 +35,8 @@ function(ggml_add_cpu_backend_variant_impl tag_name) ggml-cpu/hbm.h ggml-cpu/quants.c ggml-cpu/quants.h + ggml-cpu/prism-quants.c + ggml-cpu/prism-quants.h ggml-cpu/traits.cpp ggml-cpu/traits.h ggml-cpu/amx/amx.cpp diff --git a/ggml/src/ggml-cpu/ggml-cpu.c b/ggml/src/ggml-cpu/ggml-cpu.c index 491316f74912..c91f62f8edca 100644 --- a/ggml/src/ggml-cpu/ggml-cpu.c +++ b/ggml/src/ggml-cpu/ggml-cpu.c @@ -7,6 +7,7 @@ #include "ggml-cpu-impl.h" #include "ggml-impl.h" #include "quants.h" +#include "prism-quants.h" #include "ggml-threading.h" #include "unary-ops.h" #include "binary-ops.h" @@ -230,6 +231,18 @@ static const struct ggml_type_traits_cpu type_traits_cpu[GGML_TYPE_COUNT] = { .vec_dot_type = GGML_TYPE_Q8_0, .nrows = 1, }, + [GGML_TYPE_PQ2_0] = { + .from_float = quantize_row_pq2_0, + .vec_dot = ggml_vec_dot_pq2_0_q8_0, + .vec_dot_type = GGML_TYPE_Q8_0, + .nrows = 1, + }, + [GGML_TYPE_PTQ1_0] = { + .from_float = quantize_row_ptq1_0, + .vec_dot = ggml_vec_dot_ptq1_0_q8_0, + .vec_dot_type = GGML_TYPE_Q8_0, + .nrows = 1, + }, [GGML_TYPE_Q2_0] = { .from_float = quantize_row_q2_0, .vec_dot = ggml_vec_dot_q2_0_q8_0, diff --git a/ggml/src/ggml-cpu/ops.cpp b/ggml/src/ggml-cpu/ops.cpp index 42ec809ce521..6a1ee5e8d707 100644 --- a/ggml/src/ggml-cpu/ops.cpp +++ b/ggml/src/ggml-cpu/ops.cpp @@ -5047,6 +5047,8 @@ void ggml_compute_forward_get_rows( case GGML_TYPE_IQ4_XS: case GGML_TYPE_IQ3_S: case GGML_TYPE_IQ2_S: + case GGML_TYPE_PQ2_0: + case GGML_TYPE_PTQ1_0: { ggml_compute_forward_get_rows_q(params, dst); } break; @@ -5795,6 +5797,8 @@ void ggml_compute_forward_clamp( case GGML_TYPE_Q6_K: case GGML_TYPE_TQ1_0: case GGML_TYPE_TQ2_0: + case GGML_TYPE_PQ2_0: + case GGML_TYPE_PTQ1_0: case GGML_TYPE_IQ2_XXS: case GGML_TYPE_IQ2_XS: case GGML_TYPE_IQ3_XXS: @@ -11099,10 +11103,19 @@ static void ggml_compute_forward_dsv4_hc_pre_f32( const int64_t hc = x->ne[1]; const int64_t n_tokens = x->ne[2]; + const float scale = ggml_get_op_params_f32(dst, 0); + const bool gated = ggml_get_op_params_i32(dst, 1) != 0; + GGML_ASSERT(dst->ne[0] == n_embd); GGML_ASSERT(dst->ne[1] == n_tokens); - GGML_ASSERT(weights->ne[0] == hc); - GGML_ASSERT(weights->ne[1] == n_tokens); + if (gated) { + GGML_ASSERT(weights->ne[0] == n_embd); + GGML_ASSERT(weights->ne[1] == hc); + GGML_ASSERT(weights->ne[2] == n_tokens); + } else { + GGML_ASSERT(weights->ne[0] == hc); + GGML_ASSERT(weights->ne[1] == n_tokens); + } GGML_TENSOR_LOCALS(size_t, nbx, x, nb); GGML_TENSOR_LOCALS(size_t, nbw, weights, nb); @@ -11122,12 +11135,18 @@ static void ggml_compute_forward_dsv4_hc_pre_f32( float sum = 0.0f; for (int64_t ih = 0; ih < hc; ++ih) { - const float xv = *(const float *) ((const char *) x->data + i0*nbx0 + ih*nbx1 + it*nbx2); - const float wv = *(const float *) ((const char *) weights->data + ih*nbw0 + it*nbw1); + const float xv = *(const float *) ((const char *) x->data + i0*nbx0 + ih*nbx1 + it*nbx2); + float wv; + if (gated) { + const float gv = *(const float *) ((const char *) weights->data + i0*nbw0 + ih*nbw1 + it*nbw2); + wv = 1.0f / (1.0f + expf(-gv)); + } else { + wv = *(const float *) ((const char *) weights->data + ih*nbw0 + it*nbw1); + } sum += xv * wv; } - *(float *) ((char *) dst->data + i0*nbd0 + it*nbd1) = sum; + *(float *) ((char *) dst->data + i0*nbd0 + it*nbd1) = scale * sum; } } @@ -11161,7 +11180,6 @@ static void ggml_compute_forward_dsv4_hc_post_f32( GGML_ASSERT(x->type == GGML_TYPE_F32); GGML_ASSERT(residual->type == GGML_TYPE_F32); GGML_ASSERT(post->type == GGML_TYPE_F32); - GGML_ASSERT(comb->type == GGML_TYPE_F32); GGML_ASSERT(dst->type == GGML_TYPE_F32); const int64_t n_embd = x->ne[0]; @@ -11175,14 +11193,24 @@ static void ggml_compute_forward_dsv4_hc_post_f32( GGML_ASSERT(residual->ne[2] == n_tokens); GGML_ASSERT(post->ne[0] == hc); GGML_ASSERT(post->ne[1] == n_tokens); - GGML_ASSERT(comb->ne[0] == hc); - GGML_ASSERT(comb->ne[1] == hc); - GGML_ASSERT(comb->ne[2] == n_tokens); + + // comb == NULL: identity mixing, each stream keeps its own residual + size_t nbc0 = 0; + size_t nbc1 = 0; + size_t nbc2 = 0; + if (comb) { + GGML_ASSERT(comb->type == GGML_TYPE_F32); + GGML_ASSERT(comb->ne[0] == hc); + GGML_ASSERT(comb->ne[1] == hc); + GGML_ASSERT(comb->ne[2] == n_tokens); + nbc0 = comb->nb[0]; + nbc1 = comb->nb[1]; + nbc2 = comb->nb[2]; + } GGML_TENSOR_LOCALS(size_t, nbx, x, nb); GGML_TENSOR_LOCALS(size_t, nbr, residual, nb); GGML_TENSOR_LOCALS(size_t, nbp, post, nb); - GGML_TENSOR_LOCALS(size_t, nbc, comb, nb); GGML_TENSOR_LOCALS(size_t, nbd, dst, nb); const int ith = params->ith; @@ -11202,10 +11230,14 @@ static void ggml_compute_forward_dsv4_hc_post_f32( const float pv = *(const float *) ((const char *) post->data + idst*nbp0 + it*nbp1); float sum = xv * pv; - for (int64_t isrc = 0; isrc < hc; ++isrc) { - const float rv = *(const float *) ((const char *) residual->data + i0*nbr0 + isrc*nbr1 + it*nbr2); - const float cv = *(const float *) ((const char *) comb->data + idst*nbc0 + isrc*nbc1 + it*nbc2); - sum += rv * cv; + if (comb) { + for (int64_t isrc = 0; isrc < hc; ++isrc) { + const float rv = *(const float *) ((const char *) residual->data + i0*nbr0 + isrc*nbr1 + it*nbr2); + const float cv = *(const float *) ((const char *) comb->data + idst*nbc0 + isrc*nbc1 + it*nbc2); + sum += rv * cv; + } + } else { + sum += *(const float *) ((const char *) residual->data + i0*nbr0 + idst*nbr1 + it*nbr2); } *(float *) ((char *) dst->data + i0*nbd0 + idst*nbd1 + it*nbd2) = sum; diff --git a/ggml/src/ggml-cpu/prism-quants.c b/ggml/src/ggml-cpu/prism-quants.c new file mode 100644 index 000000000000..4776aada36e6 --- /dev/null +++ b/ggml/src/ggml-cpu/prism-quants.c @@ -0,0 +1,96 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// CPU dot products for PrismML's PQ2_0 and PTQ1_0 (ggml-prism.h) against Q8_0 activations, as PrismML's fork +// pairs them (four Q8_0 blocks per 128-value block). Plain C: the weights are decoded to int8 per block and +// summed per 32-value Q8_0 block, so the compiler can vectorise the inner loops. This is the reference and +// the fallback for these types; the GPU backends carry the fast paths. + +#include "prism-quants.h" + +#include "ggml-cpu-impl.h" +#include "simd-mappings.h" + +#include + +void quantize_row_pq2_0(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k) { + quantize_row_pq2_0_ref(x, (block_pq2_0 *) y, k); +} + +void quantize_row_ptq1_0(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k) { + quantize_row_ptq1_0_ref(x, (block_ptq1_0 *) y, k); +} + +static inline float prism_dot_q8_0_x4(const int8_t * GGML_RESTRICT q, const block_q8_0 * GGML_RESTRICT y) { + float sum = 0.0f; + for (int b = 0; b < 4; ++b) { + int sumi = 0; + for (int j = 0; j < QK8_0; ++j) { + sumi += (int) q[b * QK8_0 + j] * (int) y[b].qs[j]; + } + sum += GGML_CPU_FP16_TO_FP32(y[b].d) * (float) sumi; + } + return sum; +} + +void ggml_vec_dot_pq2_0_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, + const void * GGML_RESTRICT vy, size_t by, int nrc) { + assert(n % QK_PQ2_0 == 0); + assert(nrc == 1); + (void) bs; + (void) bx; + (void) by; + (void) nrc; + + const block_pq2_0 * GGML_RESTRICT x = (const block_pq2_0 *) vx; + const block_q8_0 * GGML_RESTRICT y = (const block_q8_0 *) vy; + const int nb = n / QK_PQ2_0; + + float sumf = 0.0f; + for (int i = 0; i < nb; ++i) { + int8_t q[QK_PQ2_0]; + for (int j = 0; j < QK_PQ2_0 / 4; ++j) { + const uint8_t b = x[i].qs[j]; + q[4 * j + 0] = (int8_t) (((b >> 0) & 3) - 1); + q[4 * j + 1] = (int8_t) (((b >> 2) & 3) - 1); + q[4 * j + 2] = (int8_t) (((b >> 4) & 3) - 1); + q[4 * j + 3] = (int8_t) (((b >> 6) & 3) - 1); + } + sumf += GGML_CPU_FP16_TO_FP32(x[i].d) * prism_dot_q8_0_x4(q, y + 4 * i); + } + *s = sumf; +} + +void ggml_vec_dot_ptq1_0_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, + const void * GGML_RESTRICT vy, size_t by, int nrc) { + assert(n % QK_PTQ1_0 == 0); + assert(nrc == 1); + (void) bs; + (void) bx; + (void) by; + (void) nrc; + + const block_ptq1_0 * GGML_RESTRICT x = (const block_ptq1_0 *) vx; + const block_q8_0 * GGML_RESTRICT y = (const block_q8_0 *) vy; + const int nb = n / QK_PTQ1_0; + + float sumf = 0.0f; + for (int i = 0; i < nb; ++i) { + int8_t q[QK_PTQ1_0]; + ggml_ptq1_0_trits(&x[i], q); + sumf += GGML_CPU_FP16_TO_FP32(x[i].d) * prism_dot_q8_0_x4(q, y + 4 * i); + } + *s = sumf; +} diff --git a/ggml/src/ggml-cpu/prism-quants.h b/ggml/src/ggml-cpu/prism-quants.h new file mode 100644 index 000000000000..a1fbc2997795 --- /dev/null +++ b/ggml/src/ggml-cpu/prism-quants.h @@ -0,0 +1,35 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// CPU backend entry points for PrismML's PQ2_0 and PTQ1_0 (ggml-prism.h). +#pragma once + +#include "ggml-prism.h" + +#ifdef __cplusplus +extern "C" { +#endif + +void quantize_row_pq2_0(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); +void quantize_row_ptq1_0(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); + +void ggml_vec_dot_pq2_0_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, + const void * GGML_RESTRICT vy, size_t by, int nrc); +void ggml_vec_dot_ptq1_0_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, + const void * GGML_RESTRICT vy, size_t by, int nrc); + +#ifdef __cplusplus +} +#endif diff --git a/ggml/src/ggml-cuda/dsv4-hc.cu b/ggml/src/ggml-cuda/dsv4-hc.cu index c4b19a787b0e..ca1d2dc8a482 100644 --- a/ggml/src/ggml-cuda/dsv4-hc.cu +++ b/ggml/src/ggml-cuda/dsv4-hc.cu @@ -100,6 +100,7 @@ static __global__ void dsv4_hc_comb_f32( } } +template static __global__ void dsv4_hc_pre_f32( const float * x, const float * weights, @@ -112,8 +113,10 @@ static __global__ void dsv4_hc_pre_f32( int64_t sx2, int64_t sw0, int64_t sw1, + int64_t sw2, int64_t sd0, - int64_t sd1) { + int64_t sd1, + float scale) { ggml_cuda_pdl_lc(); const int64_t ir = (int64_t) blockIdx.x * blockDim.x + threadIdx.x; const int64_t nr = n_embd * n_tokens; @@ -127,16 +130,22 @@ static __global__ void dsv4_hc_pre_f32( const int64_t i0 = ir % n_embd; const int64_t it = ir / n_embd; - float sum = x[i0*sx0 + it*sx2] * weights[it*sw1]; - for (int64_t ih = 1; ih < hc; ++ih) { + float sum = 0.0f; + for (int64_t ih = 0; ih < hc; ++ih) { const float xv = x[i0*sx0 + ih*sx1 + it*sx2]; - const float wv = weights[ih*sw0 + it*sw1]; + float wv; + if constexpr (gated) { + wv = 1.0f / (1.0f + expf(-weights[i0*sw0 + ih*sw1 + it*sw2])); + } else { + wv = weights[ih*sw0 + it*sw1]; + } sum += xv * wv; } - dst[i0*sd0 + it*sd1] = sum; + dst[i0*sd0 + it*sd1] = scale * sum; } +template static __global__ void dsv4_hc_post_f32( const float * x, const float * residual, @@ -174,8 +183,12 @@ static __global__ void dsv4_hc_post_f32( const int64_t it = ir / (n_embd * hc); float sum = x[i0*sx0 + it*sx1] * post[idst*sp0 + it*sp1]; - for (int64_t isrc = 0; isrc < hc; ++isrc) { - sum += residual[i0*sr0 + isrc*sr1 + it*sr2] * comb[idst*sc0 + isrc*sc1 + it*sc2]; + if constexpr (has_comb) { + for (int64_t isrc = 0; isrc < hc; ++isrc) { + sum += residual[i0*sr0 + isrc*sr1 + it*sr2] * comb[idst*sc0 + isrc*sc1 + it*sc2]; + } + } else { + sum += residual[i0*sr0 + idst*sr1 + it*sr2]; } dst[i0*sd0 + idst*sd1 + it*sd2] = sum; @@ -240,18 +253,23 @@ void ggml_cuda_op_dsv4_hc_pre(ggml_backend_cuda_context & ctx, ggml_tensor * dst const int64_t hc = x->ne[1]; const int64_t n_tokens = x->ne[2]; + const float scale = ggml_get_op_params_f32(dst, 0); + const bool gated = ggml_get_op_params_i32(dst, 1) != 0; + const int block_size = 256; const int64_t nr = n_embd * n_tokens; const dim3 block_dims(block_size, 1, 1); const dim3 grid_dims((nr + block_size - 1) / block_size, 1, 1); const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(grid_dims, block_dims, 0, ctx.stream()); - ggml_cuda_kernel_launch(dsv4_hc_pre_f32, launch_params, + auto kernel = gated ? dsv4_hc_pre_f32 : dsv4_hc_pre_f32; + ggml_cuda_kernel_launch(kernel, launch_params, (const float *) x->data, (const float *) weights->data, (float *) dst->data, n_embd, hc, n_tokens, nbx0 / sizeof(float), nbx1 / sizeof(float), nbx2 / sizeof(float), - nbw0 / sizeof(float), nbw1 / sizeof(float), - nbd0 / sizeof(float), nbd1 / sizeof(float)); + nbw0 / sizeof(float), nbw1 / sizeof(float), nbw2 / sizeof(float), + nbd0 / sizeof(float), nbd1 / sizeof(float), + scale); } void ggml_cuda_op_dsv4_hc_post(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { @@ -263,15 +281,18 @@ void ggml_cuda_op_dsv4_hc_post(ggml_backend_cuda_context & ctx, ggml_tensor * ds GGML_ASSERT(x->type == GGML_TYPE_F32); GGML_ASSERT(residual->type == GGML_TYPE_F32); GGML_ASSERT(post->type == GGML_TYPE_F32); - GGML_ASSERT(comb->type == GGML_TYPE_F32); + GGML_ASSERT(comb == nullptr || comb->type == GGML_TYPE_F32); GGML_ASSERT(dst->type == GGML_TYPE_F32); GGML_TENSOR_LOCALS(size_t, nbx, x, nb); GGML_TENSOR_LOCALS(size_t, nbr, residual, nb); GGML_TENSOR_LOCALS(size_t, nbp, post, nb); - GGML_TENSOR_LOCALS(size_t, nbc, comb, nb); GGML_TENSOR_LOCALS(size_t, nbd, dst, nb); + const size_t nbc0 = comb ? comb->nb[0] : 0; + const size_t nbc1 = comb ? comb->nb[1] : 0; + const size_t nbc2 = comb ? comb->nb[2] : 0; + const int64_t n_embd = x->ne[0]; const int64_t n_tokens = x->ne[1]; const int64_t hc = residual->ne[1]; @@ -282,9 +303,10 @@ void ggml_cuda_op_dsv4_hc_post(ggml_backend_cuda_context & ctx, ggml_tensor * ds const dim3 grid_dims((nr + block_size - 1) / block_size, 1, 1); const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(grid_dims, block_dims, 0, ctx.stream()); - ggml_cuda_kernel_launch(dsv4_hc_post_f32, launch_params, + auto kernel = comb ? dsv4_hc_post_f32 : dsv4_hc_post_f32; + ggml_cuda_kernel_launch(kernel, launch_params, (const float *) x->data, (const float *) residual->data, - (const float *) post->data, (const float *) comb->data, (float *) dst->data, + (const float *) post->data, comb ? (const float *) comb->data : nullptr, (float *) dst->data, n_embd, hc, n_tokens, nbx0 / sizeof(float), nbx1 / sizeof(float), nbr0 / sizeof(float), nbr1 / sizeof(float), nbr2 / sizeof(float), diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index 561ab7ac599f..34f14a4d2aca 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -5154,7 +5154,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g op->type == GGML_TYPE_F32; case GGML_OP_DSV4_HC_POST: return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 && - op->src[2]->type == GGML_TYPE_F32 && op->src[3]->type == GGML_TYPE_F32 && + op->src[2]->type == GGML_TYPE_F32 && (op->src[3] == nullptr || op->src[3]->type == GGML_TYPE_F32) && op->type == GGML_TYPE_F32; case GGML_OP_FLASH_ATTN_EXT: return ggml_cuda_flash_attn_ext_supported(dev_ctx->device, op); diff --git a/ggml/src/ggml-cuda/ssm-conv.cu b/ggml/src/ggml-cuda/ssm-conv.cu index 1463169cf78b..a9fe3377abb3 100644 --- a/ggml/src/ggml-cuda/ssm-conv.cu +++ b/ggml/src/ggml-cuda/ssm-conv.cu @@ -148,12 +148,13 @@ static void ssm_conv_f32_cuda(const float * src0, const float * src1, const floa }; switch (nc) { + case 2: launch_kernel(std::integral_constant{}); break; // zaya case 3: launch_kernel(std::integral_constant{}); break; case 4: launch_kernel(std::integral_constant{}); break; case 5: launch_kernel(std::integral_constant{}); break; case 9: launch_kernel(std::integral_constant{}); break; case 15: launch_kernel(std::integral_constant{}); break; - default: GGML_ABORT("Only support kernel sizes 3, 4, 5, 9, 15 right now."); + default: GGML_ABORT("Only support kernel sizes 2, 3, 4, 5, 9, 15 right now."); } } diff --git a/ggml/src/ggml-hrx/CMakeLists.txt b/ggml/src/ggml-hrx/CMakeLists.txt new file mode 100644 index 000000000000..58903a8c89cd --- /dev/null +++ b/ggml/src/ggml-hrx/CMakeLists.txt @@ -0,0 +1,349 @@ +set(HRX_SOURCE_DIR "" CACHE PATH "Optional HRX source tree to build instead of using installed hrx and loomc packages") + +if(HRX_SOURCE_DIR) + include(ExternalProject) + include(GNUInstallDirs) + + get_filename_component(HRX_SOURCE_DIR "${HRX_SOURCE_DIR}" ABSOLUTE) + if(NOT EXISTS "${HRX_SOURCE_DIR}/CMakeLists.txt") + message(FATAL_ERROR "HRX_SOURCE_DIR does not contain a CMakeLists.txt: ${HRX_SOURCE_DIR}") + endif() + + set(GGML_HRX_PREFIX "${CMAKE_CURRENT_BINARY_DIR}/hrx") + set(GGML_HRX_BUILD_DIR "${GGML_HRX_PREFIX}/src/ggml-hrx-deps-build") + set(GGML_HRX_LIB "${GGML_HRX_BUILD_DIR}/libhrx/src/libhrx/${CMAKE_SHARED_LIBRARY_PREFIX}hrx${CMAKE_SHARED_LIBRARY_SUFFIX}") + set(GGML_LOOMC_LIB "${GGML_HRX_BUILD_DIR}/loom/binding/c/${CMAKE_SHARED_LIBRARY_PREFIX}loomc${CMAKE_SHARED_LIBRARY_SUFFIX}") + set(GGML_HRX_LOOM_LINK "${GGML_HRX_BUILD_DIR}/loom/src/loom/tools/loom-link/loom-link${CMAKE_EXECUTABLE_SUFFIX}") + set(GGML_HRX_LOOM_FORMAT "${GGML_HRX_BUILD_DIR}/loom/src/loom/tools/loom-format/loom-format${CMAKE_EXECUTABLE_SUFFIX}") + set(GGML_HRX_IREE_BENCHMARK_LOOM "${GGML_HRX_BUILD_DIR}/loom/src/loom/tools/iree-benchmark-loom/iree-benchmark-loom${CMAKE_EXECUTABLE_SUFFIX}") + set(GGML_HRX_DEPS_TARGET ggml-hrx-deps) + + set(GGML_HRX_CMAKE_ARGS + -DCMAKE_BUILD_TYPE=${CMAKE_BUILD_TYPE} + -DCMAKE_C_COMPILER=${CMAKE_C_COMPILER} + -DCMAKE_CXX_COMPILER=${CMAKE_CXX_COMPILER} + -DIREE_BUILD_TESTS=OFF + -DIREE_BUILD_BENCHMARKS=OFF + -DIREE_HAL_DRIVER_DEFAULTS=OFF + -DIREE_HAL_DRIVER_AMDGPU=ON + -DIREE_HAL_DRIVER_LOCAL_SYNC=ON + -DIREE_HAL_DRIVER_LOCAL_TASK=ON + -DIREE_HAL_DRIVER_NULL=ON + -DLIBHRX_BUILD_CTS=OFF + ) + if(IREE_ROCM_PATH) + list(APPEND GGML_HRX_CMAKE_ARGS -DIREE_ROCM_PATH=${IREE_ROCM_PATH}) + endif() + if(FETCHCONTENT_BASE_DIR) + list(APPEND GGML_HRX_CMAKE_ARGS -DFETCHCONTENT_BASE_DIR=${FETCHCONTENT_BASE_DIR}) + endif() + + ExternalProject_Add(ggml-hrx-deps + SOURCE_DIR "${HRX_SOURCE_DIR}" + PREFIX "${GGML_HRX_PREFIX}" + CMAKE_ARGS ${GGML_HRX_CMAKE_ARGS} + BUILD_ALWAYS TRUE + BUILD_COMMAND ${CMAKE_COMMAND} --build . --target hrx loomc_shared loom_tools_loom-link_loom-link loom_tools_loom-format_loom-format loom_tools_iree-benchmark-loom_iree-benchmark-loom --config ${CMAKE_BUILD_TYPE} + INSTALL_COMMAND "" + BUILD_BYPRODUCTS "${GGML_HRX_LIB}" "${GGML_LOOMC_LIB}" "${GGML_HRX_LOOM_LINK}" "${GGML_HRX_LOOM_FORMAT}" "${GGML_HRX_IREE_BENCHMARK_LOOM}" + UPDATE_COMMAND "" + ) + + add_library(hrx::hrx SHARED IMPORTED GLOBAL) + set_target_properties(hrx::hrx PROPERTIES + IMPORTED_LOCATION "${GGML_HRX_LIB}" + INTERFACE_INCLUDE_DIRECTORIES "${HRX_SOURCE_DIR}/libhrx/include") + add_dependencies(hrx::hrx ggml-hrx-deps) + + add_library(loomc::loomc SHARED IMPORTED GLOBAL) + set_target_properties(loomc::loomc PROPERTIES + IMPORTED_LOCATION "${GGML_LOOMC_LIB}" + INTERFACE_INCLUDE_DIRECTORIES "${HRX_SOURCE_DIR}/loom/binding/c/include" + INTERFACE_COMPILE_DEFINITIONS LOOMC_USING_SHARED_LIBRARY) + add_dependencies(loomc::loomc ggml-hrx-deps) +else() + find_package(hrx CONFIG REQUIRED) + find_package(loomc CONFIG REQUIRED) + find_program(GGML_HRX_LOOM_LINK NAMES loom-link) + find_program(GGML_HRX_LOOM_FORMAT NAMES loom-format) + find_program(GGML_HRX_IREE_BENCHMARK_LOOM NAMES iree-benchmark-loom) +endif() + +find_package(Python3 REQUIRED COMPONENTS Interpreter) + +set(GGML_HRX_KERNEL_CORPUS_SOURCE_FORMAT "binary" CACHE STRING "Embedded Loom corpus source format: text or binary") +set_property(CACHE GGML_HRX_KERNEL_CORPUS_SOURCE_FORMAT PROPERTY STRINGS text binary) +if(GGML_HRX_KERNEL_CORPUS_SOURCE_FORMAT STREQUAL "binary") + if(NOT GGML_HRX_LOOM_LINK OR NOT GGML_HRX_LOOM_FORMAT) + message(FATAL_ERROR "GGML_HRX_KERNEL_CORPUS_SOURCE_FORMAT=binary requires loom-link and loom-format") + endif() + set(GGML_HRX_KERNEL_CORPUS_TOOL_ARGS + --loom-link "${GGML_HRX_LOOM_LINK}" + --loom-format "${GGML_HRX_LOOM_FORMAT}" + ) + set(GGML_HRX_KERNEL_CORPUS_TOOL_DEPENDS + "${GGML_HRX_LOOM_LINK}" + "${GGML_HRX_LOOM_FORMAT}" + ) +elseif(NOT GGML_HRX_KERNEL_CORPUS_SOURCE_FORMAT STREQUAL "text") + message(FATAL_ERROR "Unsupported GGML_HRX_KERNEL_CORPUS_SOURCE_FORMAT: ${GGML_HRX_KERNEL_CORPUS_SOURCE_FORMAT}") +endif() + +set(GGML_HRX_QWEN_MOE_KERNEL_CORPUS_DIR "${CMAKE_CURRENT_SOURCE_DIR}/kernel-corpus/kernels/qwen_moe") +set(GGML_HRX_QWEN_MOE_KERNEL_CORPUS_MANIFEST "${GGML_HRX_QWEN_MOE_KERNEL_CORPUS_DIR}/manifest.json") +set(GGML_HRX_HRX_KERNEL_CORPUS_DIR "${CMAKE_CURRENT_SOURCE_DIR}/kernel-corpus/kernels/hrx") +set(GGML_HRX_HRX_KERNEL_CORPUS_MANIFEST "${GGML_HRX_HRX_KERNEL_CORPUS_DIR}/manifest.json") +set(GGML_HRX_QWEN_KERNEL_CORPUS_DIR "${CMAKE_CURRENT_SOURCE_DIR}/kernel-corpus/kernels/qwen") +set(GGML_HRX_QWEN_KERNEL_CORPUS_MANIFEST "${GGML_HRX_QWEN_KERNEL_CORPUS_DIR}/manifest.json") +set(GGML_HRX_LOOM_LIBS_KERNEL_CORPUS_DIR "${CMAKE_CURRENT_SOURCE_DIR}/kernel-corpus/kernels/loom-libs") +set(GGML_HRX_LOOM_LIBS_KERNEL_CORPUS_MANIFEST "${GGML_HRX_LOOM_LIBS_KERNEL_CORPUS_DIR}/manifest.json") +set(GGML_HRX_KERNEL_CORPUS_MANIFESTS + "${GGML_HRX_QWEN_MOE_KERNEL_CORPUS_MANIFEST}" + "${GGML_HRX_HRX_KERNEL_CORPUS_MANIFEST}" + "${GGML_HRX_QWEN_KERNEL_CORPUS_MANIFEST}" + "${GGML_HRX_LOOM_LIBS_KERNEL_CORPUS_MANIFEST}" +) +set(GGML_HRX_KERNEL_CORPUS_MANIFEST_ARGS + --manifest "${GGML_HRX_QWEN_MOE_KERNEL_CORPUS_MANIFEST}" + --corpus-dir "${GGML_HRX_QWEN_MOE_KERNEL_CORPUS_DIR}" + --manifest "${GGML_HRX_HRX_KERNEL_CORPUS_MANIFEST}" + --corpus-dir "${GGML_HRX_HRX_KERNEL_CORPUS_DIR}" + --manifest "${GGML_HRX_QWEN_KERNEL_CORPUS_MANIFEST}" + --corpus-dir "${GGML_HRX_QWEN_KERNEL_CORPUS_DIR}" + --manifest "${GGML_HRX_LOOM_LIBS_KERNEL_CORPUS_MANIFEST}" + --corpus-dir "${GGML_HRX_LOOM_LIBS_KERNEL_CORPUS_DIR}" +) +set(GGML_HRX_KERNEL_CORPUS_SOURCES_INC "${CMAKE_CURRENT_BINARY_DIR}/kernel-corpus-sources.inc") +set(GGML_HRX_KERNEL_CORPUS_QWEN_INC "${CMAKE_CURRENT_BINARY_DIR}/kernel-corpus-qwen.inc") +set(GGML_HRX_KERNEL_CORPUS_CATALOG_INC "${CMAKE_CURRENT_BINARY_DIR}/kernel-corpus-catalog.inc") +set(GGML_HRX_KERNEL_CORPUS_DEPFILE "${CMAKE_CURRENT_BINARY_DIR}/kernel-corpus.d") + +add_custom_command( + OUTPUT + "${GGML_HRX_KERNEL_CORPUS_SOURCES_INC}" + "${GGML_HRX_KERNEL_CORPUS_QWEN_INC}" + "${GGML_HRX_KERNEL_CORPUS_CATALOG_INC}" + COMMAND ${Python3_EXECUTABLE} + "${CMAKE_CURRENT_SOURCE_DIR}/tools/generate_kernel_corpus.py" + --source-output "${GGML_HRX_KERNEL_CORPUS_SOURCES_INC}" + --corpus-output "${GGML_HRX_KERNEL_CORPUS_QWEN_INC}" + --catalog-output "${GGML_HRX_KERNEL_CORPUS_CATALOG_INC}" + ${GGML_HRX_KERNEL_CORPUS_MANIFEST_ARGS} + --source-format "${GGML_HRX_KERNEL_CORPUS_SOURCE_FORMAT}" + ${GGML_HRX_KERNEL_CORPUS_TOOL_ARGS} + --depfile "${GGML_HRX_KERNEL_CORPUS_DEPFILE}" + DEPENDS + "${CMAKE_CURRENT_SOURCE_DIR}/tools/generate_kernel_corpus.py" + ${GGML_HRX_KERNEL_CORPUS_MANIFESTS} + ${GGML_HRX_KERNEL_CORPUS_TOOL_DEPENDS} + DEPFILE "${GGML_HRX_KERNEL_CORPUS_DEPFILE}" + VERBATIM +) + +option(GGML_HRX_BUNDLE_RUNTIME_LIBS "Bundle HRX/ROCm runtime libraries next to the HRX backend" OFF) +set(GGML_HRX_BUNDLE_LIBRARY_DIRS "" CACHE STRING "Library directories to scan when GGML_HRX_BUNDLE_RUNTIME_LIBS=ON") + +add_library(ggml-hrx-kernel-corpus STATIC + status.h + kernel-corpus/kernel-corpus-json.cpp + kernel-corpus/kernel-corpus-json.h + kernel-corpus/kernel-corpus-catalog-verify.h + kernel-corpus/kernel-corpus-catalog.h + kernel-corpus/kernel-corpus.cpp + kernel-corpus/kernel-corpus.h + kernel-corpus/kernel-types.h + "${GGML_HRX_KERNEL_CORPUS_SOURCES_INC}" + "${GGML_HRX_KERNEL_CORPUS_QWEN_INC}" + "${GGML_HRX_KERNEL_CORPUS_CATALOG_INC}" +) +if(GGML_HRX_DEPS_TARGET) + add_dependencies(ggml-hrx-kernel-corpus ${GGML_HRX_DEPS_TARGET}) +endif() +target_include_directories(ggml-hrx-kernel-corpus PUBLIC . PRIVATE "${CMAKE_CURRENT_BINARY_DIR}" ../../../vendor) +target_compile_features(ggml-hrx-kernel-corpus PRIVATE cxx_std_17) +set_target_properties(ggml-hrx-kernel-corpus PROPERTIES POSITION_INDEPENDENT_CODE ON) + +ggml_add_backend_library(ggml-hrx + backend-buffer-binding.cpp + backend-buffer-binding.h + backend-context.h + dispatch/command-plan-metadata.cpp + dispatch/command-plan-metadata.h + dispatch/command-plan.h + dispatch/command-program-bindings.cpp + dispatch/command-program-bindings.h + dispatch/command-program-dump.cpp + dispatch/command-program-dump.h + dispatch/command-program-diagnostics.cpp + dispatch/command-program-diagnostics.h + dispatch/command-program.cpp + dispatch/command-program.h + dispatch/command-program-resolver.cpp + dispatch/command-program-resolver.h + dispatch/dispatch-scheduler.cpp + dispatch/dispatch-scheduler.h + dispatch/dispatch.h + dispatch/transient-allocator.cpp + dispatch/transient-allocator.h + dispatch/transient-reuse-guard.cpp + dispatch/transient-reuse-guard.h + dispatch_registration/common/dispatch-binary.cpp + dispatch_registration/common/dispatch-binary.h + dispatch_registration/common/dispatch-common.cpp + dispatch_registration/common/dispatch-common.h + dispatch_registration/common/dispatch-copy.cpp + dispatch_registration/common/dispatch-copy.h + dispatch_registration/common/dispatch-flash-attention.cpp + dispatch_registration/common/dispatch-flash-attention.h + dispatch_registration/common/dispatch-gated-mul-mat-id.cpp + dispatch_registration/common/dispatch-gated-mul-mat-id.h + dispatch_registration/common/dispatch-gated-mul-mat.cpp + dispatch_registration/common/dispatch-gated-mul-mat.h + dispatch_registration/common/dispatch-gather-add.cpp + dispatch_registration/common/dispatch-gather-add.h + dispatch_registration/common/dispatch-get-rows.cpp + dispatch_registration/common/dispatch-get-rows.h + dispatch_registration/common/dispatch-glu.cpp + dispatch_registration/common/dispatch-grouped-mul-mat.cpp + dispatch_registration/common/dispatch-grouped-mul-mat.h + dispatch_registration/common/dispatch-glu.h + dispatch_registration/common/dispatch-mul-mat-id.cpp + dispatch_registration/common/dispatch-mul-mat-id-common.h + dispatch_registration/common/dispatch-mul-mat-id.h + dispatch_registration/common/dispatch-moe-routing-layout.h + dispatch_registration/common/dispatch-mul-mat.cpp + dispatch_registration/common/dispatch-mul-mat-iq3-xxs.cpp + dispatch_registration/common/dispatch-mul-mat-iq3-xxs.h + dispatch_registration/common/dispatch-mul-mat-tail.cpp + dispatch_registration/common/dispatch-mul-mat-tail.h + dispatch_registration/common/dispatch-mul-mat-common.h + dispatch_registration/common/dispatch-mul-mat-weight-format.h + dispatch_registration/common/dispatch-mul-mat.h + dispatch_registration/common/dispatch-rope-set-rows.cpp + dispatch_registration/common/dispatch-rope-set-rows.h + dispatch_registration/common/dispatch-rope-utils.h + dispatch_registration/common/dispatch-rmsnorm.cpp + dispatch_registration/common/dispatch-rmsnorm.h + dispatch_registration/common/dispatch-scale.cpp + dispatch_registration/common/dispatch-small-rows.cpp + dispatch_registration/common/dispatch-small-rows.h + dispatch_registration/common/dispatch-mul-mat-id-decode.cpp + dispatch_registration/common/dispatch-mul-mat-id-decode.h + dispatch_registration/common/dispatch-res-scale-pair.cpp + dispatch_registration/common/dispatch-res-scale-pair.h + dispatch_registration/common/dispatch-hadamard.cpp + dispatch_registration/common/dispatch-hadamard.h + dispatch_registration/common/dispatch-zaya-cca-conv.cpp + dispatch_registration/common/dispatch-zaya-cca-conv.h + dispatch_registration/common/dispatch-zaya-cca-qk-norm.cpp + dispatch_registration/common/dispatch-zaya-cca-qk-norm.h + dispatch_registration/common/dispatch-kquant-decode.cpp + dispatch_registration/common/dispatch-kquant-decode.h + dispatch_registration/common/dispatch-add-id.cpp + dispatch_registration/common/dispatch-add-id.h + dispatch_registration/common/dispatch-swiglu-oai.cpp + dispatch_registration/common/dispatch-swiglu-oai.h + dispatch_registration/common/moe-placement-guard.cpp + dispatch_registration/common/moe-placement-guard.h + dispatch_registration/common/dispatch-attention-sink.cpp + dispatch_registration/common/dispatch-attention-sink.h + dispatch_registration/common/dispatch-softplus.cpp + dispatch_registration/common/dispatch-softplus.h + dispatch_registration/common/dispatch-scale.h + dispatch_registration/common/dispatch-unary.cpp + dispatch_registration/common/dispatch-unary.h + dispatch_registration/dispatch-registry.cpp + dispatch_registration/dispatch-registry.h + dispatch_registration/llm/dispatch-attention-qkv.cpp + dispatch_registration/llm/dispatch-attention-qkv.h + dispatch_registration/llm/dispatch-gated-delta-net.cpp + dispatch_registration/llm/dispatch-gated-delta-net.h + dispatch_registration/llm/dispatch-ssm-conv.cpp + dispatch_registration/llm/dispatch-ssm-conv.h + dispatch_registration/qwen/dispatch-llm-profiles.h + dispatch_registration/qwen/dispatch-llm-shapes.h + dispatch_registration/qwen/dispatch-moe-router.cpp + dispatch_registration/qwen/dispatch-moe-router.h + dispatch_registration/qwen/dispatch-qwen.cpp + dispatch_registration/qwen/dispatch-qwen.h + dispatch_registration/qwen/dispatch-qwen-attention-postprocess.cpp + dispatch_registration/qwen/dispatch-qwen-attention-postprocess.h + dispatch_registration/qwen/dispatch-qwen-matmul.cpp + dispatch_registration/qwen/dispatch-qwen-matmul.h + dispatch_registration/qwen/dispatch-qwen-rmsnorm.cpp + dispatch_registration/qwen/dispatch-qwen-rmsnorm.h + dispatch_registration/qwen/dispatch-routed-ffn.cpp + dispatch_registration/qwen/dispatch-routed-ffn.h + status.h + graph/graph.cpp + graph/graph-diagnostics.cpp + graph/graph-diagnostics.h + graph/graph.h + graph/graph-matcher.cpp + graph/graph-matcher.h + graph/graph-traversal.cpp + graph/graph-traversal.h + graph/op-params.cpp + graph/op-params.h + ggml-hrx.cpp + ggml-hrx-dmabuf.cpp + ggml-hrx-dmabuf.h + fused-context-claim.h + loom-jit.cpp + graph/value-map.cpp + graph/value-map.h + runtime/command-program-executor.cpp + runtime/hrx-sleeping-wait.cpp + runtime/hrx-sleeping-wait.h + runtime/command-program-executor.h + runtime/graph-executor.cpp + runtime/graph-executor.h + runtime/graph-program-cache.cpp + runtime/graph-program-cache-limit.cpp + runtime/graph-program-cache-limit.h + runtime/graph-record-order.cpp + runtime/graph-record-order.h + runtime/graph-program-cache.h + runtime/host-memory.cpp + runtime/host-memory.h + runtime/host-memory-ternary.cpp + runtime/host-memory-ternary.h + dispatch/ternary-q4-0.h + runtime/kernel-executable-cache.cpp + runtime/kernel-executable-cache.h + runtime/loom-kernel-jit.cpp + runtime/loom-kernel-jit.h + runtime/loom-jit-disk-cache.cpp + runtime/loom-jit-disk-cache.h + runtime/prepared-command-program-cache.cpp + runtime/prepared-command-program-cache.h + runtime/transient-arena.cpp + runtime/transient-arena.h + runtime/host-buffer-registry.cpp + runtime/host-buffer-registry.h +) +target_link_libraries(ggml-hrx PRIVATE ggml-hrx-kernel-corpus hrx::hrx loomc::loomc) +target_include_directories(ggml-hrx PRIVATE . "${CMAKE_CURRENT_BINARY_DIR}" ../../../vendor) +target_compile_definitions(ggml-hrx PRIVATE GGML_USE_HRX) +include("${CMAKE_CURRENT_SOURCE_DIR}/hip/ggml-hrx-hip.cmake") + +if (GGML_HRX_BUNDLE_RUNTIME_LIBS) + include("${CMAKE_CURRENT_SOURCE_DIR}/cmake/BundleRuntime.cmake") + ggml_hrx_bundle_runtime(ggml-hrx) +endif() + +add_executable(ggml-hrx-compile-kernel + tools/compile-kernel.cpp + tools/tool-utils.h + loom-jit.cpp +) +target_link_libraries(ggml-hrx-compile-kernel PRIVATE hrx::hrx loomc::loomc) +target_include_directories(ggml-hrx-compile-kernel PRIVATE .) +target_compile_features(ggml-hrx-compile-kernel PRIVATE cxx_std_17) + +add_executable(ggml-hrx-analyze-graph + tools/analyze-graph.cpp +) +target_link_libraries(ggml-hrx-analyze-graph PRIVATE ggml-hrx ggml ggml-hrx-kernel-corpus) +target_include_directories(ggml-hrx-analyze-graph PRIVATE . ../../../vendor) +target_compile_features(ggml-hrx-analyze-graph PRIVATE cxx_std_17) diff --git a/ggml/src/ggml-hrx/backend-buffer-binding.cpp b/ggml/src/ggml-hrx/backend-buffer-binding.cpp new file mode 100644 index 000000000000..bab5f12b1af7 --- /dev/null +++ b/ggml/src/ggml-hrx/backend-buffer-binding.cpp @@ -0,0 +1,86 @@ +#include "backend-buffer-binding.h" + +#include "ggml-backend-impl.h" +#include "ggml.h" + +ggml_backend_hrx_buffer_context * ggml_backend_hrx_buffer_context_from_buffer(ggml_backend_buffer_t buffer) { + return static_cast(buffer->context); +} + +size_t ggml_backend_hrx_tensor_offset(const ggml_backend_hrx_buffer_context * context, const ggml_tensor * tensor) { + return static_cast(static_cast(tensor->data) - context->base); +} + +void * ggml_backend_hrx_buffer_base(ggml_backend_buffer_t buffer) { + return ggml_backend_hrx_buffer_context_from_buffer(buffer)->base; +} + +bool ggml_backend_hrx_tensor_binding(const ggml_tensor * tensor, + ggml_backend_hrx_buffer_context ** out_context, + size_t * out_offset) { + if (tensor == nullptr) { + return false; + } + ggml_backend_buffer_t buffer = tensor->view_src != nullptr ? tensor->view_src->buffer : tensor->buffer; + if (buffer == nullptr || buffer->iface.get_base != ggml_backend_hrx_buffer_base) { + return false; + } + auto * context = ggml_backend_hrx_buffer_context_from_buffer(buffer); + const size_t offset = ggml_backend_hrx_tensor_offset(context, tensor); + if (context->buffer == nullptr || offset > buffer->size || ggml_nbytes(tensor) > buffer->size - offset) { + return false; + } + *out_context = context; + *out_offset = offset; + return true; +} + +bool ggml_backend_hrx_resolve_value_buffer(const ggml_tensor * tensor, ggml::hrx::ValueBufferBinding & binding) { + ggml_backend_hrx_buffer_context * context = nullptr; + size_t offset = 0; + if (!ggml_backend_hrx_tensor_binding(tensor, &context, &offset)) { + if (tensor == nullptr) { + return false; + } + const ggml_tensor * root = tensor->view_src != nullptr ? tensor->view_src : tensor; + ggml_backend_buffer_t buffer = root->buffer; + if (buffer == nullptr || !ggml_backend_buffer_is_host(buffer)) { + return false; + } + void * base = ggml_backend_buffer_get_base(buffer); + const size_t capacity = ggml_backend_buffer_get_size(buffer); + if (base == nullptr || tensor->data == nullptr || + static_cast(tensor->data) < static_cast(base)) { + return false; + } + const size_t host_offset = + static_cast(static_cast(tensor->data) - static_cast(base)); + if (host_offset > capacity || ggml_nbytes(tensor) > capacity - host_offset) { + return false; + } + const uint64_t buffer_address = static_cast(reinterpret_cast(buffer)); + const uint64_t base_address = static_cast(reinterpret_cast(base)); + binding.host_data = base; + binding.offset = host_offset; + binding.length = ggml_nbytes(tensor); + binding.identity = + buffer_address ^ (base_address + 0x9e3779b97f4a7c15ull + (buffer_address << 6) + (buffer_address >> 2)); + binding.generation = 1; + binding.capacity = capacity; + binding.weight = ggml_backend_buffer_get_usage(buffer) == GGML_BACKEND_BUFFER_USAGE_WEIGHTS; + return true; + } + ggml_backend_buffer_t buffer = tensor->view_src != nullptr ? tensor->view_src->buffer : tensor->buffer; + const bool directly_bindable = !ggml_backend_buffer_is_host(buffer) || context->direct_host_binding; + // Coherent HRX host allocations are directly device-addressable. Represent them with an HRX buffer handle so + // command-program preparation bypasses host materialization. Noncoherent host allocations remain host data. + binding.buffer = directly_bindable ? context->buffer : nullptr; + binding.host_data = directly_bindable ? nullptr : context->base; + binding.offset = offset; + binding.length = ggml_nbytes(tensor); + binding.identity = context->identity; + binding.generation = context->generation; + binding.capacity = buffer != nullptr ? buffer->size : 0; + binding.weight = buffer != nullptr && ggml_backend_buffer_get_usage(buffer) == GGML_BACKEND_BUFFER_USAGE_WEIGHTS; + return true; +} diff --git a/ggml/src/ggml-hrx/backend-buffer-binding.h b/ggml/src/ggml-hrx/backend-buffer-binding.h new file mode 100644 index 000000000000..59b925873f63 --- /dev/null +++ b/ggml/src/ggml-hrx/backend-buffer-binding.h @@ -0,0 +1,18 @@ +#pragma once + +#include "backend-context.h" +#include "ggml-backend.h" +#include "graph/value-map.h" + +#include + +struct ggml_tensor; + +ggml_backend_hrx_buffer_context * ggml_backend_hrx_buffer_context_from_buffer(ggml_backend_buffer_t buffer); +size_t ggml_backend_hrx_tensor_offset(const ggml_backend_hrx_buffer_context * context, const ggml_tensor * tensor); +void * ggml_backend_hrx_buffer_base(ggml_backend_buffer_t buffer); + +bool ggml_backend_hrx_tensor_binding(const ggml_tensor * tensor, + ggml_backend_hrx_buffer_context ** out_context, + size_t * out_offset); +bool ggml_backend_hrx_resolve_value_buffer(const ggml_tensor * tensor, ggml::hrx::ValueBufferBinding & binding); diff --git a/ggml/src/ggml-hrx/backend-context.h b/ggml/src/ggml-hrx/backend-context.h new file mode 100644 index 000000000000..3d510b310b59 --- /dev/null +++ b/ggml/src/ggml-hrx/backend-context.h @@ -0,0 +1,76 @@ +#pragma once + +#include "ggml-backend-impl.h" +#include "graph/value-map.h" +#include "runtime/graph-program-cache.h" +#include "runtime/host-buffer-registry.h" +#include "runtime/host-memory.h" +#include "runtime/kernel-executable-cache.h" +#include "runtime/prepared-command-program-cache.h" +#include "runtime/transient-arena.h" + +#include +#include +#include +#include +#include +#include +#include + +struct ggml_tensor; + +struct ggml_backend_hrx_device_context; + +struct ggml_backend_hrx_buffer_type_context { + ggml_backend_hrx_device_context * device; + std::string name; + bool host_visible = false; +}; + +struct ggml_backend_hrx_buffer_context { + ggml_backend_hrx_device_context * device; + hrx_buffer_t buffer; + uint8_t * base; + uint64_t identity; + uint64_t generation; + bool direct_host_binding; +}; + +struct ggml_backend_hrx_device_context { + hrx_device_t device = nullptr; + std::string name; + std::string description; + std::string architecture; + size_t memory_total = 0; + bool use_direct_host_bindings = false; + ggml_backend_buffer_type buft = {}; + ggml_backend_hrx_buffer_type_context buft_context = {}; + ggml_backend_buffer_type host_buft = {}; + ggml_backend_hrx_buffer_type_context host_buft_context = {}; + ggml::hrx::HostBufferRegistry host_buffers; + std::atomic synchronous_upload_fallbacks{ 0 }; + std::atomic synchronous_download_fallbacks{ 0 }; + std::mutex buffer_stream_mutex; + hrx_stream_t buffer_stream = nullptr; +}; + +struct ggml_backend_hrx_context { + ggml_backend_hrx_device_context * device; + hrx_stream_t stream; + ggml::hrx::KernelExecutableCache kernel_executables; + ggml::hrx::GraphProgramCache graph_programs; + ggml::hrx::PreparedCommandProgramCache prepared_programs; + ggml::hrx::TransientArena transient_arena; + ggml::hrx::HostTransferManager host_transfers; + ggml::hrx::HostWeightCache host_weights; + ggml::hrx::GraphReplayStreamState graph_replay_state; + std::string name; +}; + +struct ggml_backend_hrx_reg_context { + bool initialized = false; + std::vector> device_contexts; + std::vector devices; + + ~ggml_backend_hrx_reg_context(); +}; diff --git a/ggml/src/ggml-hrx/benchmarks/README.md b/ggml/src/ggml-hrx/benchmarks/README.md new file mode 100644 index 000000000000..58b8e83db874 --- /dev/null +++ b/ggml/src/ggml-hrx/benchmarks/README.md @@ -0,0 +1,199 @@ +# HRX Loom Benchmarks + +This directory contains model-scoped Loom benchmarks for HRX kernels. Benchmark sources stay separate from the production kernel corpus: they declare the kernels they need with `kernel.decl`, and the runner links those declarations against the production `.loom` files before benchmarking. + +Use these benchmarks to measure production kernels with `iree-benchmark-loom` while keeping model-shaped benchmark cases out of the embedded kernel catalog. Each model has one `.loom` file under `loom/`, and each workload shape has a sidecar manifest next to it, for example `llama32_3b_f16.pp512.json`. + +Generated model benchmarks are deduplicated by kernel shape. The sidecar manifest keeps a compact `dispatches` list with the generated benchmark name, kernel, compile/runtime parameters, source files, and `count`, so benchmark results can be weighted back to the full model without materializing one check case for every duplicate dispatch. + +Generated model benchmarks are performance fixtures, not numerical validators. They launch the model-shaped kernels with synthetic buffers and intentionally omit `check.expect.*` assertions; backend model smoke tests and the kernel corpus checks remain responsible for numerical validation. + +## Generating Llama Benchmarks + +First dump the HRX command programs for each model run: + +```sh +GGML_HRX_DUMP_COMMAND_PROGRAM_DIR=/tmp/hrx-llama32-pp256-dumps \ +/bin/llama-bench \ + --model /home/rsuderman/Downloads/gguf/llama-3.2/Llama-3.2-3B-Instruct-F16.gguf \ + --device HRX0 \ + --n-gpu-layers -1 \ + --batch-size 64 \ + --ubatch-size 64 \ + --repetitions 1 \ + --no-warmup \ + --output jsonl \ + --n-prompt 256 \ + --n-gen 0 \ + --n-depth 0 + +GGML_HRX_DUMP_COMMAND_PROGRAM_DIR=/tmp/hrx-llama32-pp512-dumps \ +/bin/llama-bench \ + --model /home/rsuderman/Downloads/gguf/llama-3.2/Llama-3.2-3B-Instruct-F16.gguf \ + --device HRX0 \ + --n-gpu-layers -1 \ + --batch-size 64 \ + --ubatch-size 64 \ + --repetitions 1 \ + --no-warmup \ + --output jsonl \ + --n-prompt 512 \ + --n-gen 0 \ + --n-depth 0 + +GGML_HRX_DUMP_COMMAND_PROGRAM_DIR=/tmp/hrx-llama32-tg8-dumps \ +/bin/llama-bench \ + --model /home/rsuderman/Downloads/gguf/llama-3.2/Llama-3.2-3B-Instruct-F16.gguf \ + --device HRX0 \ + --n-gpu-layers -1 \ + --batch-size 64 \ + --ubatch-size 64 \ + --repetitions 1 \ + --no-warmup \ + --output jsonl \ + --n-prompt 0 \ + --n-gen 8 \ + --n-depth 0 +``` + +For the Llama 3.2 1B Q4_K_XL decode benchmark, dump the tg32 command program with: + +```sh +GGML_HRX_DUMP_COMMAND_PROGRAM_DIR=/tmp/hrx-llama32-1b-tg32-dumps \ +/bin/llama-bench \ + --model /home/rsuderman/Downloads/gguf/lemonade/llamacpp-gguf-models/unsloth_Llama-3.2-1B-Instruct-GGUF/Llama-3.2-1B-Instruct-UD-Q4_K_XL.gguf \ + --device HRX0 \ + --n-gpu-layers -1 \ + --batch-size 512 \ + --ubatch-size 512 \ + --repetitions 1 \ + --output jsonl \ + --n-prompt 0 \ + --n-gen 32 \ + --n-depth 0 +``` + +Then generate the shared model benchmark file and per-scenario sidecars: + +```sh +ggml/src/ggml-hrx/tools/benchmarks/generate-model-benchmarks.py \ + --model llama32_3b_f16 \ + --scenario-dump pp256=/tmp/hrx-llama32-pp256-dumps \ + --scenario-dump pp512=/tmp/hrx-llama32-pp512-dumps \ + --scenario-dump tg8=/tmp/hrx-llama32-tg8-dumps + +ggml/src/ggml-hrx/tools/benchmarks/generate-model-benchmarks.py \ + --model llama32_1b_q4_k_xl \ + --scenario tg32 \ + --dump-dir /tmp/hrx-llama32-1b-tg32-dumps +``` + +Review the generated invoked-kernel set with: + +```sh +jq '.kernel_counts' ggml/src/ggml-hrx/benchmarks/loom/llama32_3b_f16..json + +ggml/src/ggml-hrx/tools/benchmarks/run-model-benchmarks.sh \ + --model llama32_3b_f16 \ + --scenario \ + --build-dir \ + --list +``` + +Run the generated benchmarks with: + +```sh +ggml/src/ggml-hrx/tools/benchmarks/run-model-benchmarks.sh \ + --model llama32_3b_f16 \ + --scenario \ + --build-dir \ + --output-dir /home/rsuderman/codex/project-workspaces/llama.cpp/gates/hrx-loom-benchmarks/llama32- +``` + +Use `--dry-run-only` to stop after planning, `--list` to inspect the generated benchmark set, or `--benchmark @name` to run a single invocation. Set `LOOM_LINK` or `IREE_BENCHMARK_LOOM` to override tool discovery. Set `DEVICE` to override the default `amdgpu` HAL device. + +The runner applies the same workload-argument specialization used by the HRX runtime after `loom-link`, and benchmarks the resulting `linked.runtime-specialized.loom` source by default. Use `--no-runtime-specialization` to benchmark the raw linked source. + +`loom-link` also receives each captured compile parameter as `--config==`. Some model kernels select templates from compile-time config values, so passing those configs during link keeps standalone benchmarks aligned with the HRX runtime path. + +Use `--continue-on-failure` to record failed or timed-out standalone cases and keep running the rest of the model benchmark set. Use `--benchmark-timeout-sec` to bound each link, dry-run, and benchmark command. Each benchmark directory keeps `loom-link`, dry-run, and benchmark command lines plus stdout/stderr so failures can be reproduced directly. + +Use `--profile-final-batch=false`, `--input-ring-count`, or repeated `--benchmark-extra-arg` flags when isolating `iree-benchmark-loom` behavior from production runtime behavior. After a run, summarize the weighted model time with: + +```sh +ggml/src/ggml-hrx/tools/benchmarks/summarize-model-benchmarks.py \ + /home/rsuderman/codex/project-workspaces/llama.cpp/gates/hrx-loom-benchmarks/llama32-pp512/results.jsonl +``` + +The summary uses `operation_timing_ns.p50 * count` by default and writes `summary.json` plus `summary.md` next to the runner results. + +To map likely fusion opportunities from a command-program dump and the weighted benchmark summary: + +```sh +ggml/src/ggml-hrx/tools/benchmarks/analyze-model-fusion-adjacency.py \ + --dump-dir /tmp/hrx-llama32-tg8-dumps \ + --scenario-manifest ggml/src/ggml-hrx/benchmarks/loom/llama32_3b_f16.tg8.json \ + --summary-json /home/rsuderman/codex/project-workspaces/llama.cpp/gates/hrx-loom-benchmarks/llama32-tg8/summary.json \ + --output-json /home/rsuderman/codex/project-workspaces/llama.cpp/gates/hrx-loom-benchmarks/llama32-tg8/fusion-adjacency.json \ + --output-md /home/rsuderman/codex/project-workspaces/llama.cpp/gates/hrx-loom-benchmarks/llama32-tg8/fusion-adjacency.md +``` + +The adjacency report uses transient value producer-consumer edges for true data dependencies and sequential cache-update windows for side-effect patterns such as RoPE, SET_ROWS, and decode flash attention. + +## Generating Qwen 30B Benchmarks + +Dump the HRX command programs with the Qwen 30B shard and the same pp256, pp512, and tg8 shapes used for model benchmarking: + +```sh +GGML_HRX_DUMP_COMMAND_PROGRAM_DIR=/tmp/hrx-qwen30b-pp256-dumps \ +/bin/llama-bench \ + --model /home/rsuderman/Downloads/gguf/qwen-30b/qwen3-30b-a3b-q4_k_m-00001-of-00020.gguf \ + --device HRX0 \ + --n-gpu-layers -1 \ + --batch-size 64 \ + --ubatch-size 64 \ + --repetitions 1 \ + --no-warmup \ + --output jsonl \ + --n-prompt 256 \ + --n-gen 0 \ + --n-depth 0 + +GGML_HRX_DUMP_COMMAND_PROGRAM_DIR=/tmp/hrx-qwen30b-pp512-dumps \ +/bin/llama-bench \ + --model /home/rsuderman/Downloads/gguf/qwen-30b/qwen3-30b-a3b-q4_k_m-00001-of-00020.gguf \ + --device HRX0 \ + --n-gpu-layers -1 \ + --batch-size 64 \ + --ubatch-size 64 \ + --repetitions 1 \ + --no-warmup \ + --output jsonl \ + --n-prompt 512 \ + --n-gen 0 \ + --n-depth 0 + +GGML_HRX_DUMP_COMMAND_PROGRAM_DIR=/tmp/hrx-qwen30b-tg8-dumps \ +/bin/llama-bench \ + --model /home/rsuderman/Downloads/gguf/qwen-30b/qwen3-30b-a3b-q4_k_m-00001-of-00020.gguf \ + --device HRX0 \ + --n-gpu-layers -1 \ + --batch-size 64 \ + --ubatch-size 64 \ + --repetitions 1 \ + --no-warmup \ + --output jsonl \ + --n-prompt 0 \ + --n-gen 8 \ + --n-depth 0 +``` + +Then regenerate the shared Qwen benchmark source and sidecars: + +```sh +ggml/src/ggml-hrx/tools/benchmarks/generate-model-benchmarks.py \ + --model qwen3_30b_a3b_q4_k_m \ + --scenario-dump pp256=/tmp/hrx-qwen30b-pp256-dumps \ + --scenario-dump pp512=/tmp/hrx-qwen30b-pp512-dumps \ + --scenario-dump tg8=/tmp/hrx-qwen30b-tg8-dumps +``` diff --git a/ggml/src/ggml-hrx/benchmarks/loom/llama32_3b_f16.loom b/ggml/src/ggml-hrx/benchmarks/loom/llama32_3b_f16.loom new file mode 100644 index 000000000000..3af4e58634f2 --- /dev/null +++ b/ggml/src/ggml-hrx/benchmarks/loom/llama32_3b_f16.loom @@ -0,0 +1,691 @@ +// Generated by tools/benchmarks/generate-model-benchmarks.py for llama32_3b_f16 scenarios: pp256, pp512, tg8. +// Regenerate from an HRX command program dump rather than editing by hand. + +target.decl @ggml_binary_f32_gfx11_wave64 + +target.decl @ggml_flash_attention_decode_split_gfx11_wave64 + +target.decl @ggml_flash_attention_gfx11_wave64 + +target.decl @ggml_gather_add_gfx11_wave64 + +target.decl @ggml_get_rows_f32_gfx11_wave64 + +target.decl @ggml_mul_mat_add_gfx11_wave64 + +target.decl @ggml_mul_mat_f32_f32_decode_gfx11_wave64 + +target.decl @ggml_mul_mat_gfx11_wave64 + +target.decl @ggml_mul_mat_swiglu_gfx11_wave64 + +target.decl @ggml_rmsnorm_binary_gfx11_wave32 + +target.decl @ggml_rope_f32_gfx11_wave32 + +target.decl @ggml_rope_set_rows_f32_gfx11_wave32 + +target.decl @ggml_set_rows_gfx11_wave64 + +target.decl @llm_attention_k_matmul_rope_set_rows_gfx11_wave64 + +target.decl @llm_attention_q_matmul_rope_gfx11_wave64 + +target.decl @llm_attention_v_matmul_set_rows_gfx11_wave64 + +kernel.decl target(@ggml_binary_f32_gfx11_wave64) @ggml_binary_f32(%element_count$0: index) launch(%element_count$1: index, %lhs: buffer, %rhs: buffer, %output: buffer) + +kernel.decl target(@ggml_flash_attention_decode_split_gfx11_wave64) @ggml_flash_attention_decode_split_f32_f16_wmma_next_q8(%key_value_token_count$5: index) launch(%key_value_token_count$6: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %completion_counter: buffer, %output: buffer, %next_q8_output: buffer) + +kernel.decl target(@ggml_flash_attention_gfx11_wave64) @ggml_flash_attention_f32_f16_wmma(%query_token_count$17: index, %key_value_token_count$18: index) launch(%query_token_count$19: index, %key_value_token_count$20: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %output: buffer) + +kernel.decl target(@ggml_gather_add_gfx11_wave64) @ggml_gather_add_f32(%source_token_count$26: index, %output_token_count$27: index, %hidden_size$28: index) launch(%source_token_count$29: index, %output_token_count$30: index, %hidden_size$31: index, %attention: buffer, %residual: buffer, %output_ids: buffer, %output: buffer) + +kernel.decl target(@ggml_get_rows_f32_gfx11_wave64) @ggml_get_rows_f32(%token_count$36: index, %row_count$37: index, %hidden_size$38: index) launch(%token_count$39: index, %row_count$40: index, %hidden_size$41: index, %token_ids: buffer, %weight: buffer, %output: buffer) + +kernel.decl target(@ggml_mul_mat_add_gfx11_wave64) @ggml_mul_mat_add_f32_f32_wmma(%token_count$45: index) launch(%token_count$46: index, %input: buffer, %weight: buffer, %residual_input: buffer, %residual_output: buffer) + +kernel.decl target(@ggml_mul_mat_f32_f32_decode_gfx11_wave64) @ggml_mul_mat_f32_f32_decode_wave64(%token_count$51: index, %input_size$52: index, %output_size$53: index) launch(%token_count$54: index, %input_size$55: index, %output_size$56: index, %input: buffer, %weight: buffer, %output: buffer) + +kernel.decl target(@ggml_mul_mat_gfx11_wave64) @ggml_mul_mat_f32_f32_wmma(%token_count$60: index) launch(%token_count$61: index, %input: buffer, %weight: buffer, %output: buffer) + +kernel.decl target(@ggml_mul_mat_swiglu_gfx11_wave64) @ggml_mul_mat_swiglu_f32_f32_wmma(%token_count$65: index) launch(%token_count$66: index, %input: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) + +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_rmsnorm_binary_f32(%token_count$71: index) launch(%token_count$72: index, %input: buffer, %rhs: buffer, %output: buffer) + +kernel.decl target(@ggml_rope_f32_gfx11_wave32) @ggml_rope_f32(%token_count$76: index) launch(%token_count$77: index, %positions: buffer, %input: buffer, %theta: buffer, %freq_factors: buffer, %output: buffer) + +kernel.decl target(@ggml_rope_set_rows_f32_gfx11_wave32) @ggml_rope_set_rows_f32(%token_count$83: index, %cache_row_count$84: index) launch(%token_count$85: index, %cache_row_count$86: index, %positions: buffer, %indices: buffer, %input: buffer, %theta: buffer, %freq_factors: buffer, %cache: buffer) + +kernel.decl target(@ggml_set_rows_gfx11_wave64) @ggml_set_rows(%token_count$93: index, %cache_row_count$94: index, %hidden_size$95: index) launch(%token_count$96: index, %cache_row_count$97: index, %hidden_size$98: index, %rows: buffer, %indices: buffer, %cache: buffer) + +kernel.decl target(@llm_attention_k_matmul_rope_set_rows_gfx11_wave64) @llm_attention_k_matmul_rope_set_rows_f32_f32_wmma(%token_count$102: index) launch(%token_count$103: index, %input: buffer, %weight: buffer, %positions: buffer, %indices: buffer, %theta: buffer, %freq_factors: buffer, %cache: buffer) + +kernel.decl target(@llm_attention_q_matmul_rope_gfx11_wave64) @llm_attention_q_matmul_rope_f32_f32_wmma(%token_count$111: index) launch(%token_count$112: index, %input: buffer, %weight: buffer, %positions: buffer, %theta: buffer, %freq_factors: buffer, %output: buffer) + +kernel.decl target(@llm_attention_v_matmul_set_rows_gfx11_wave64) @llm_attention_v_matmul_set_rows_f32_f32_wmma(%token_count$119: index) launch(%token_count$120: index, %input: buffer, %weight: buffer, %indices: buffer, %cache: buffer) + +// Scenario: pp256 +check.case public @llama32_3b_f16_pp256_000_ggml_get_rows_f32_case { + %token_count = check.literal value(64) : index + %row_count = check.literal value(128256) : index + %hidden_size = check.literal value(3072) : index + %token_ids = check.generate.fill value(0) : tensor<64xi32> + %weight = check.generate.fill value(0.0) : tensor<128256x3072xf16> + %output = check.generate.fill value(1.0) : tensor<64x3072xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<64xi32>, tensor<128256x3072xf16>, tensor<64x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp256_000_ggml_get_rows_f32_case> @llama32_3b_f16_pp256_000_ggml_get_rows_f32 + +check.case public @llama32_3b_f16_pp256_001_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(2.0) : tensor<64x3072xf32> + %rhs = check.generate.fill value(3.0) : tensor<3072xf32> + %output = check.generate.fill value(0.0) : tensor<64x3072xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<64x3072xf32>, tensor<3072xf32>, tensor<64x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp256_001_ggml_rmsnorm_binary_f32_case> @llama32_3b_f16_pp256_001_ggml_rmsnorm_binary_f32 + +check.case public @llama32_3b_f16_pp256_002_llm_attention_q_matmul_rope_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(1.0) : tensor<64x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<3072x3072xf16> + %positions = check.generate.fill value(0) : tensor<64xi32> + %theta = check.generate.fill value(0.0) : tensor<64xf32> + %freq_factors = check.generate.fill value(0.0) : tensor<64xf32> + %output = check.generate.fill value(0.0) : tensor<64x24x128xf32> + kernel.launch @llm_attention_q_matmul_rope_f32_f32_wmma[%token_count](%token_count, %input, %weight, %positions, %theta, %freq_factors, %output) : [index](index, tensor<64x3072xf32>, tensor<3072x3072xf16>, tensor<64xi32>, tensor<64xf32>, tensor<64xf32>, tensor<64x24x128xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp256_002_llm_attention_q_matmul_rope_f32_f32_wmma_case> @llama32_3b_f16_pp256_002_llm_attention_q_matmul_rope_f32_f32_wmma + +check.case public @llama32_3b_f16_pp256_003_llm_attention_v_matmul_set_rows_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(1.0) : tensor<64x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<1024x3072xf16> + %indices = check.generate.iota offset(0) step(1) period(256) : tensor<64xi64> + %cache = check.generate.fill value(0.0) : tensor<256x1024xf16> + kernel.launch @llm_attention_v_matmul_set_rows_f32_f32_wmma[%token_count](%token_count, %input, %weight, %indices, %cache) : [index](index, tensor<64x3072xf32>, tensor<1024x3072xf16>, tensor<64xi64>, tensor<256x1024xf16>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp256_003_llm_attention_v_matmul_set_rows_f32_f32_wmma_case> @llama32_3b_f16_pp256_003_llm_attention_v_matmul_set_rows_f32_f32_wmma + +check.case public @llama32_3b_f16_pp256_004_llm_attention_k_matmul_rope_set_rows_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(1.0) : tensor<64x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<1024x3072xf16> + %positions = check.generate.fill value(0) : tensor<64xi32> + %indices = check.generate.iota offset(0) step(1) period(256) : tensor<64xi64> + %theta = check.generate.fill value(0.0) : tensor<64xf32> + %freq_factors = check.generate.fill value(0.0) : tensor<64xf32> + %cache = check.generate.fill value(0.0) : tensor<256x8x128xf16> + kernel.launch @llm_attention_k_matmul_rope_set_rows_f32_f32_wmma[%token_count](%token_count, %input, %weight, %positions, %indices, %theta, %freq_factors, %cache) : [index](index, tensor<64x3072xf32>, tensor<1024x3072xf16>, tensor<64xi32>, tensor<64xi64>, tensor<64xf32>, tensor<64xf32>, tensor<256x8x128xf16>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp256_004_llm_attention_k_matmul_rope_set_rows_f32_f32_wmma_case> @llama32_3b_f16_pp256_004_llm_attention_k_matmul_rope_set_rows_f32_f32_wmma + +check.case public @llama32_3b_f16_pp256_005_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(64) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<64x24x128xf32> + %key = check.generate.fill value(0.0) : tensor<256x8x128xf16> + %value = check.generate.fill value(0.0) : tensor<256x8x128xf16> + %mask = check.generate.fill value(0.0) : tensor<64x256xf16> + %output = check.generate.fill value(1.0) : tensor<64x24x128xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<64x24x128xf32>, tensor<256x8x128xf16>, tensor<256x8x128xf16>, tensor<64x256xf16>, tensor<64x24x128xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp256_005_ggml_flash_attention_f32_f16_wmma_case> @llama32_3b_f16_pp256_005_ggml_flash_attention_f32_f16_wmma + +check.case public @llama32_3b_f16_pp256_006_ggml_mul_mat_add_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(1.0) : tensor<64x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<3072x3072xf16> + %residual_input = check.generate.fill value(0.25) : tensor<64x3072xf32> + %residual_output = check.generate.fill value(0.0) : tensor<64x3072xf32> + kernel.launch @ggml_mul_mat_add_f32_f32_wmma[%token_count](%token_count, %input, %weight, %residual_input, %residual_output) : [index](index, tensor<64x3072xf32>, tensor<3072x3072xf16>, tensor<64x3072xf32>, tensor<64x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp256_006_ggml_mul_mat_add_f32_f32_wmma_case> @llama32_3b_f16_pp256_006_ggml_mul_mat_add_f32_f32_wmma + +check.case public @llama32_3b_f16_pp256_007_ggml_mul_mat_swiglu_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0.0) : tensor<64x3072xf32> + %gate_weight = check.generate.fill value(1.0) : tensor<8192x3072xf16> + %up_weight = check.generate.fill value(1.0) : tensor<8192x3072xf16> + %output = check.generate.fill value(1.0) : tensor<64x8192xf32> + kernel.launch @ggml_mul_mat_swiglu_f32_f32_wmma[%token_count](%token_count, %input, %gate_weight, %up_weight, %output) : [index](index, tensor<64x3072xf32>, tensor<8192x3072xf16>, tensor<8192x3072xf16>, tensor<64x8192xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp256_007_ggml_mul_mat_swiglu_f32_f32_wmma_case> @llama32_3b_f16_pp256_007_ggml_mul_mat_swiglu_f32_f32_wmma + +check.case public @llama32_3b_f16_pp256_008_ggml_mul_mat_add_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(1.0) : tensor<64x8192xf32> + %weight = check.generate.fill value(1.0) : tensor<3072x8192xf16> + %residual_input = check.generate.fill value(0.25) : tensor<64x3072xf32> + %residual_output = check.generate.fill value(0.0) : tensor<64x3072xf32> + kernel.launch @ggml_mul_mat_add_f32_f32_wmma[%token_count](%token_count, %input, %weight, %residual_input, %residual_output) : [index](index, tensor<64x8192xf32>, tensor<3072x8192xf16>, tensor<64x3072xf32>, tensor<64x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp256_008_ggml_mul_mat_add_f32_f32_wmma_case> @llama32_3b_f16_pp256_008_ggml_mul_mat_add_f32_f32_wmma + +check.case public @llama32_3b_f16_pp256_009_ggml_mul_mat_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(1.0) : tensor<64x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<3072x3072xf16> + %output = check.generate.fill value(0.0) : tensor<64x3072xf32> + kernel.launch @ggml_mul_mat_f32_f32_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<64x3072xf32>, tensor<3072x3072xf16>, tensor<64x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp256_009_ggml_mul_mat_f32_f32_wmma_case> @llama32_3b_f16_pp256_009_ggml_mul_mat_f32_f32_wmma + +check.case public @llama32_3b_f16_pp256_010_ggml_gather_add_f32_case { + %source_token_count = check.literal value(64) : index + %output_token_count = check.literal value(1) : index + %hidden_size = check.literal value(3072) : index + %attention = check.generate.fill value(0.0) : tensor<64x3072xf32> + %residual = check.generate.fill value(0.0) : tensor<64x3072xf32> + %output_ids = check.generate.fill value(0) : tensor<1xi32> + %output = check.generate.fill value(1.0) : tensor<1x3072xf32> + kernel.launch @ggml_gather_add_f32[%source_token_count, %output_token_count, %hidden_size](%source_token_count, %output_token_count, %hidden_size, %attention, %residual, %output_ids, %output) : [index, index, index](index, index, index, tensor<64x3072xf32>, tensor<64x3072xf32>, tensor<1xi32>, tensor<1x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp256_010_ggml_gather_add_f32_case> @llama32_3b_f16_pp256_010_ggml_gather_add_f32 + +check.case public @llama32_3b_f16_pp256_011_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(2.0) : tensor<1x3072xf32> + %rhs = check.generate.fill value(3.0) : tensor<3072xf32> + %output = check.generate.fill value(0.0) : tensor<1x3072xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<1x3072xf32>, tensor<3072xf32>, tensor<1x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp256_011_ggml_rmsnorm_binary_f32_case> @llama32_3b_f16_pp256_011_ggml_rmsnorm_binary_f32 + +check.case public @llama32_3b_f16_pp256_012_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(3072) : index + %output_size = check.literal value(8192) : index + %input = check.generate.fill value(1.0) : tensor<1x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<8192x3072xf16> + %output = check.generate.fill value(0.0) : tensor<1x8192xf32> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<1x3072xf32>, tensor<8192x3072xf16>, tensor<1x8192xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp256_012_ggml_mul_mat_f32_f32_decode_wave64_case> @llama32_3b_f16_pp256_012_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @llama32_3b_f16_pp256_013_ggml_binary_f32_case { + %element_count = check.literal value(8192) : index + %lhs = check.generate.fill value(2.0) : tensor<8192xf32> + %rhs = check.generate.fill value(3.0) : tensor<8192xf32> + %output = check.generate.fill value(0.0) : tensor<8192xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<8192xf32>, tensor<8192xf32>, tensor<8192xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp256_013_ggml_binary_f32_case> @llama32_3b_f16_pp256_013_ggml_binary_f32 + +check.case public @llama32_3b_f16_pp256_014_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(8192) : index + %output_size = check.literal value(3072) : index + %input = check.generate.fill value(1.0) : tensor<1x8192xf32> + %weight = check.generate.fill value(1.0) : tensor<3072x8192xf16> + %output = check.generate.fill value(0.0) : tensor<1x3072xf32> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<1x8192xf32>, tensor<3072x8192xf16>, tensor<1x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp256_014_ggml_mul_mat_f32_f32_decode_wave64_case> @llama32_3b_f16_pp256_014_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @llama32_3b_f16_pp256_015_ggml_binary_f32_case { + %element_count = check.literal value(3072) : index + %lhs = check.generate.fill value(2.0) : tensor<3072xf32> + %rhs = check.generate.fill value(3.0) : tensor<3072xf32> + %output = check.generate.fill value(0.0) : tensor<3072xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<3072xf32>, tensor<3072xf32>, tensor<3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp256_015_ggml_binary_f32_case> @llama32_3b_f16_pp256_015_ggml_binary_f32 + +check.case public @llama32_3b_f16_pp256_016_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(3072) : index + %output_size = check.literal value(128256) : index + %input = check.generate.fill value(1.0) : tensor<1x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<128256x3072xf16> + %output = check.generate.fill value(0.0) : tensor<1x128256xf32> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<1x3072xf32>, tensor<128256x3072xf16>, tensor<1x128256xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp256_016_ggml_mul_mat_f32_f32_decode_wave64_case> @llama32_3b_f16_pp256_016_ggml_mul_mat_f32_f32_decode_wave64 + +// Scenario: pp512 +check.case public @llama32_3b_f16_pp512_000_ggml_get_rows_f32_case { + %token_count = check.literal value(64) : index + %row_count = check.literal value(128256) : index + %hidden_size = check.literal value(3072) : index + %token_ids = check.generate.fill value(0) : tensor<64xi32> + %weight = check.generate.fill value(0.0) : tensor<128256x3072xf16> + %output = check.generate.fill value(1.0) : tensor<64x3072xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<64xi32>, tensor<128256x3072xf16>, tensor<64x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_000_ggml_get_rows_f32_case> @llama32_3b_f16_pp512_000_ggml_get_rows_f32 + +check.case public @llama32_3b_f16_pp512_001_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(2.0) : tensor<64x3072xf32> + %rhs = check.generate.fill value(3.0) : tensor<3072xf32> + %output = check.generate.fill value(0.0) : tensor<64x3072xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<64x3072xf32>, tensor<3072xf32>, tensor<64x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_001_ggml_rmsnorm_binary_f32_case> @llama32_3b_f16_pp512_001_ggml_rmsnorm_binary_f32 + +check.case public @llama32_3b_f16_pp512_002_llm_attention_q_matmul_rope_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(1.0) : tensor<64x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<3072x3072xf16> + %positions = check.generate.fill value(0) : tensor<64xi32> + %theta = check.generate.fill value(0.0) : tensor<64xf32> + %freq_factors = check.generate.fill value(0.0) : tensor<64xf32> + %output = check.generate.fill value(0.0) : tensor<64x24x128xf32> + kernel.launch @llm_attention_q_matmul_rope_f32_f32_wmma[%token_count](%token_count, %input, %weight, %positions, %theta, %freq_factors, %output) : [index](index, tensor<64x3072xf32>, tensor<3072x3072xf16>, tensor<64xi32>, tensor<64xf32>, tensor<64xf32>, tensor<64x24x128xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_002_llm_attention_q_matmul_rope_f32_f32_wmma_case> @llama32_3b_f16_pp512_002_llm_attention_q_matmul_rope_f32_f32_wmma + +check.case public @llama32_3b_f16_pp512_003_llm_attention_v_matmul_set_rows_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(1.0) : tensor<64x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<1024x3072xf16> + %indices = check.generate.iota offset(0) step(1) period(512) : tensor<64xi64> + %cache = check.generate.fill value(0.0) : tensor<512x1024xf16> + kernel.launch @llm_attention_v_matmul_set_rows_f32_f32_wmma[%token_count](%token_count, %input, %weight, %indices, %cache) : [index](index, tensor<64x3072xf32>, tensor<1024x3072xf16>, tensor<64xi64>, tensor<512x1024xf16>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_003_llm_attention_v_matmul_set_rows_f32_f32_wmma_case> @llama32_3b_f16_pp512_003_llm_attention_v_matmul_set_rows_f32_f32_wmma + +check.case public @llama32_3b_f16_pp512_004_llm_attention_k_matmul_rope_set_rows_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(1.0) : tensor<64x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<1024x3072xf16> + %positions = check.generate.fill value(0) : tensor<64xi32> + %indices = check.generate.iota offset(0) step(1) period(512) : tensor<64xi64> + %theta = check.generate.fill value(0.0) : tensor<64xf32> + %freq_factors = check.generate.fill value(0.0) : tensor<64xf32> + %cache = check.generate.fill value(0.0) : tensor<512x8x128xf16> + kernel.launch @llm_attention_k_matmul_rope_set_rows_f32_f32_wmma[%token_count](%token_count, %input, %weight, %positions, %indices, %theta, %freq_factors, %cache) : [index](index, tensor<64x3072xf32>, tensor<1024x3072xf16>, tensor<64xi32>, tensor<64xi64>, tensor<64xf32>, tensor<64xf32>, tensor<512x8x128xf16>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_004_llm_attention_k_matmul_rope_set_rows_f32_f32_wmma_case> @llama32_3b_f16_pp512_004_llm_attention_k_matmul_rope_set_rows_f32_f32_wmma + +check.case public @llama32_3b_f16_pp512_005_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(64) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<64x24x128xf32> + %key = check.generate.fill value(0.0) : tensor<256x8x128xf16> + %value = check.generate.fill value(0.0) : tensor<256x8x128xf16> + %mask = check.generate.fill value(0.0) : tensor<64x256xf16> + %output = check.generate.fill value(1.0) : tensor<64x24x128xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<64x24x128xf32>, tensor<256x8x128xf16>, tensor<256x8x128xf16>, tensor<64x256xf16>, tensor<64x24x128xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_005_ggml_flash_attention_f32_f16_wmma_case> @llama32_3b_f16_pp512_005_ggml_flash_attention_f32_f16_wmma + +check.case public @llama32_3b_f16_pp512_006_ggml_mul_mat_add_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(1.0) : tensor<64x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<3072x3072xf16> + %residual_input = check.generate.fill value(0.25) : tensor<64x3072xf32> + %residual_output = check.generate.fill value(0.0) : tensor<64x3072xf32> + kernel.launch @ggml_mul_mat_add_f32_f32_wmma[%token_count](%token_count, %input, %weight, %residual_input, %residual_output) : [index](index, tensor<64x3072xf32>, tensor<3072x3072xf16>, tensor<64x3072xf32>, tensor<64x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_006_ggml_mul_mat_add_f32_f32_wmma_case> @llama32_3b_f16_pp512_006_ggml_mul_mat_add_f32_f32_wmma + +check.case public @llama32_3b_f16_pp512_007_ggml_mul_mat_swiglu_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0.0) : tensor<64x3072xf32> + %gate_weight = check.generate.fill value(1.0) : tensor<8192x3072xf16> + %up_weight = check.generate.fill value(1.0) : tensor<8192x3072xf16> + %output = check.generate.fill value(1.0) : tensor<64x8192xf32> + kernel.launch @ggml_mul_mat_swiglu_f32_f32_wmma[%token_count](%token_count, %input, %gate_weight, %up_weight, %output) : [index](index, tensor<64x3072xf32>, tensor<8192x3072xf16>, tensor<8192x3072xf16>, tensor<64x8192xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_007_ggml_mul_mat_swiglu_f32_f32_wmma_case> @llama32_3b_f16_pp512_007_ggml_mul_mat_swiglu_f32_f32_wmma + +check.case public @llama32_3b_f16_pp512_008_ggml_mul_mat_add_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(1.0) : tensor<64x8192xf32> + %weight = check.generate.fill value(1.0) : tensor<3072x8192xf16> + %residual_input = check.generate.fill value(0.25) : tensor<64x3072xf32> + %residual_output = check.generate.fill value(0.0) : tensor<64x3072xf32> + kernel.launch @ggml_mul_mat_add_f32_f32_wmma[%token_count](%token_count, %input, %weight, %residual_input, %residual_output) : [index](index, tensor<64x8192xf32>, tensor<3072x8192xf16>, tensor<64x3072xf32>, tensor<64x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_008_ggml_mul_mat_add_f32_f32_wmma_case> @llama32_3b_f16_pp512_008_ggml_mul_mat_add_f32_f32_wmma + +check.case public @llama32_3b_f16_pp512_009_ggml_mul_mat_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(1.0) : tensor<64x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<3072x3072xf16> + %output = check.generate.fill value(0.0) : tensor<64x3072xf32> + kernel.launch @ggml_mul_mat_f32_f32_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<64x3072xf32>, tensor<3072x3072xf16>, tensor<64x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_009_ggml_mul_mat_f32_f32_wmma_case> @llama32_3b_f16_pp512_009_ggml_mul_mat_f32_f32_wmma + +check.case public @llama32_3b_f16_pp512_010_ggml_gather_add_f32_case { + %source_token_count = check.literal value(64) : index + %output_token_count = check.literal value(1) : index + %hidden_size = check.literal value(3072) : index + %attention = check.generate.fill value(0.0) : tensor<64x3072xf32> + %residual = check.generate.fill value(0.0) : tensor<64x3072xf32> + %output_ids = check.generate.fill value(0) : tensor<1xi32> + %output = check.generate.fill value(1.0) : tensor<1x3072xf32> + kernel.launch @ggml_gather_add_f32[%source_token_count, %output_token_count, %hidden_size](%source_token_count, %output_token_count, %hidden_size, %attention, %residual, %output_ids, %output) : [index, index, index](index, index, index, tensor<64x3072xf32>, tensor<64x3072xf32>, tensor<1xi32>, tensor<1x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_010_ggml_gather_add_f32_case> @llama32_3b_f16_pp512_010_ggml_gather_add_f32 + +check.case public @llama32_3b_f16_pp512_011_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(2.0) : tensor<1x3072xf32> + %rhs = check.generate.fill value(3.0) : tensor<3072xf32> + %output = check.generate.fill value(0.0) : tensor<1x3072xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<1x3072xf32>, tensor<3072xf32>, tensor<1x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_011_ggml_rmsnorm_binary_f32_case> @llama32_3b_f16_pp512_011_ggml_rmsnorm_binary_f32 + +check.case public @llama32_3b_f16_pp512_012_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(3072) : index + %output_size = check.literal value(8192) : index + %input = check.generate.fill value(1.0) : tensor<1x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<8192x3072xf16> + %output = check.generate.fill value(0.0) : tensor<1x8192xf32> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<1x3072xf32>, tensor<8192x3072xf16>, tensor<1x8192xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_012_ggml_mul_mat_f32_f32_decode_wave64_case> @llama32_3b_f16_pp512_012_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @llama32_3b_f16_pp512_013_ggml_binary_f32_case { + %element_count = check.literal value(8192) : index + %lhs = check.generate.fill value(2.0) : tensor<8192xf32> + %rhs = check.generate.fill value(3.0) : tensor<8192xf32> + %output = check.generate.fill value(0.0) : tensor<8192xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<8192xf32>, tensor<8192xf32>, tensor<8192xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_013_ggml_binary_f32_case> @llama32_3b_f16_pp512_013_ggml_binary_f32 + +check.case public @llama32_3b_f16_pp512_014_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(8192) : index + %output_size = check.literal value(3072) : index + %input = check.generate.fill value(1.0) : tensor<1x8192xf32> + %weight = check.generate.fill value(1.0) : tensor<3072x8192xf16> + %output = check.generate.fill value(0.0) : tensor<1x3072xf32> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<1x8192xf32>, tensor<3072x8192xf16>, tensor<1x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_014_ggml_mul_mat_f32_f32_decode_wave64_case> @llama32_3b_f16_pp512_014_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @llama32_3b_f16_pp512_015_ggml_binary_f32_case { + %element_count = check.literal value(3072) : index + %lhs = check.generate.fill value(2.0) : tensor<3072xf32> + %rhs = check.generate.fill value(3.0) : tensor<3072xf32> + %output = check.generate.fill value(0.0) : tensor<3072xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<3072xf32>, tensor<3072xf32>, tensor<3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_015_ggml_binary_f32_case> @llama32_3b_f16_pp512_015_ggml_binary_f32 + +check.case public @llama32_3b_f16_pp512_016_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(3072) : index + %output_size = check.literal value(128256) : index + %input = check.generate.fill value(1.0) : tensor<1x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<128256x3072xf16> + %output = check.generate.fill value(0.0) : tensor<1x128256xf32> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<1x3072xf32>, tensor<128256x3072xf16>, tensor<1x128256xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_016_ggml_mul_mat_f32_f32_decode_wave64_case> @llama32_3b_f16_pp512_016_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @llama32_3b_f16_pp512_017_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(64) : index + %key_value_token_count = check.literal value(512) : index + %query = check.generate.fill value(0.0) : tensor<64x24x128xf32> + %key = check.generate.fill value(0.0) : tensor<512x8x128xf16> + %value = check.generate.fill value(0.0) : tensor<512x8x128xf16> + %mask = check.generate.fill value(0.0) : tensor<64x512xf16> + %output = check.generate.fill value(1.0) : tensor<64x24x128xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<64x24x128xf32>, tensor<512x8x128xf16>, tensor<512x8x128xf16>, tensor<64x512xf16>, tensor<64x24x128xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_pp512_017_ggml_flash_attention_f32_f16_wmma_case> @llama32_3b_f16_pp512_017_ggml_flash_attention_f32_f16_wmma + +// Scenario: tg8 +check.case public @llama32_3b_f16_tg8_000_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(128256) : index + %hidden_size = check.literal value(3072) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<128256x3072xf16> + %output = check.generate.fill value(1.0) : tensor<1x3072xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<128256x3072xf16>, tensor<1x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_tg8_000_ggml_get_rows_f32_case> @llama32_3b_f16_tg8_000_ggml_get_rows_f32 + +check.case public @llama32_3b_f16_tg8_001_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(2.0) : tensor<1x3072xf32> + %rhs = check.generate.fill value(3.0) : tensor<3072xf32> + %output = check.generate.fill value(0.0) : tensor<1x3072xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<1x3072xf32>, tensor<3072xf32>, tensor<1x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_tg8_001_ggml_rmsnorm_binary_f32_case> @llama32_3b_f16_tg8_001_ggml_rmsnorm_binary_f32 + +check.case public @llama32_3b_f16_tg8_002_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(3072) : index + %output_size = check.literal value(3072) : index + %input = check.generate.fill value(1.0) : tensor<1x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<3072x3072xf16> + %output = check.generate.fill value(0.0) : tensor<1x3072xf32> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<1x3072xf32>, tensor<3072x3072xf16>, tensor<1x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_tg8_002_ggml_mul_mat_f32_f32_decode_wave64_case> @llama32_3b_f16_tg8_002_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @llama32_3b_f16_tg8_003_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(3072) : index + %output_size = check.literal value(1024) : index + %input = check.generate.fill value(1.0) : tensor<1x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<1024x3072xf16> + %output = check.generate.fill value(0.0) : tensor<1x1024xf32> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<1x3072xf32>, tensor<1024x3072xf16>, tensor<1x1024xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_tg8_003_ggml_mul_mat_f32_f32_decode_wave64_case> @llama32_3b_f16_tg8_003_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @llama32_3b_f16_tg8_004_ggml_rope_f32_case { + %token_count = check.literal value(1) : index + %positions = check.generate.fill value(0) : tensor<1xi32> + %input = check.generate.fill value(0.0) : tensor<1x24x128xf32> + %theta = check.generate.fill value(0.0) : tensor<64xf32> + %freq_factors = check.generate.fill value(0.0) : tensor<64xf32> + %output = check.generate.fill value(1.0) : tensor<1x24x128xf32> + kernel.launch @ggml_rope_f32[%token_count](%token_count, %positions, %input, %theta, %freq_factors, %output) : [index](index, tensor<1xi32>, tensor<1x24x128xf32>, tensor<64xf32>, tensor<64xf32>, tensor<1x24x128xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_tg8_004_ggml_rope_f32_case> @llama32_3b_f16_tg8_004_ggml_rope_f32 + +check.case public @llama32_3b_f16_tg8_005_ggml_rope_set_rows_f32_case { + %token_count = check.literal value(1) : index + %cache_row_count = check.literal value(256) : index + %positions = check.generate.fill value(0) : tensor<1xi32> + %indices = check.generate.iota offset(0) step(1) period(256) : tensor<1xi64> + %input = check.generate.fill value(0.0) : tensor<1x8x128xf32> + %theta = check.generate.fill value(0.0) : tensor<64xf32> + %freq_factors = check.generate.fill value(0.0) : tensor<64xf32> + %cache = check.generate.fill value(0.0) : tensor<256x8x128xf16> + kernel.launch @ggml_rope_set_rows_f32[%token_count, %cache_row_count](%token_count, %cache_row_count, %positions, %indices, %input, %theta, %freq_factors, %cache) : [index, index](index, index, tensor<1xi32>, tensor<1xi64>, tensor<1x8x128xf32>, tensor<64xf32>, tensor<64xf32>, tensor<256x8x128xf16>) + check.return +} + +check.benchmark<@llama32_3b_f16_tg8_005_ggml_rope_set_rows_f32_case> @llama32_3b_f16_tg8_005_ggml_rope_set_rows_f32 + +check.case public @llama32_3b_f16_tg8_006_ggml_set_rows_case { + %token_count = check.literal value(1) : index + %cache_row_count = check.literal value(256) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<1x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(256) : tensor<1xi64> + %cache = check.generate.fill value(0.0) : tensor<256x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<1x1024xf32>, tensor<1xi64>, tensor<256x1024xf16>) + check.return +} + +check.benchmark<@llama32_3b_f16_tg8_006_ggml_set_rows_case> @llama32_3b_f16_tg8_006_ggml_set_rows + +check.case public @llama32_3b_f16_tg8_007_ggml_flash_attention_decode_split_f32_f16_wmma_next_q8_case { + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<24x128xf32> + %key = check.generate.fill value(0.0) : tensor<256x8x128xf16> + %value = check.generate.fill value(0.0) : tensor<256x8x128xf16> + %mask = check.generate.fill value(0.0) : tensor<256xf16> + %partial_max = check.generate.fill value(0.0) : tensor<8x4x16xf32> + %partial_sum = check.generate.fill value(0.0) : tensor<8x4x16xf32> + %partial_output = check.generate.fill value(0.0) : tensor<8x4x16x128xf16> + %completion_counter = check.generate.fill value(0) : tensor<8xi32> + %output = check.generate.fill value(1.0) : tensor<24x128xf32> + %next_q8_output = check.generate.fill value(0) : tensor<3456xi8> + kernel.launch @ggml_flash_attention_decode_split_f32_f16_wmma_next_q8[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output, %completion_counter, %output, %next_q8_output) : [index](index, tensor<24x128xf32>, tensor<256x8x128xf16>, tensor<256x8x128xf16>, tensor<256xf16>, tensor<8x4x16xf32>, tensor<8x4x16xf32>, tensor<8x4x16x128xf16>, tensor<8xi32>, tensor<24x128xf32>, tensor<3456xi8>) + check.return +} + +check.benchmark<@llama32_3b_f16_tg8_007_ggml_flash_attention_decode_split_f32_f16_wmma_next_q8_case> @llama32_3b_f16_tg8_007_ggml_flash_attention_decode_split_f32_f16_wmma_next_q8 + +check.case public @llama32_3b_f16_tg8_008_ggml_binary_f32_case { + %element_count = check.literal value(3072) : index + %lhs = check.generate.fill value(2.0) : tensor<3072xf32> + %rhs = check.generate.fill value(3.0) : tensor<3072xf32> + %output = check.generate.fill value(0.0) : tensor<3072xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<3072xf32>, tensor<3072xf32>, tensor<3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_tg8_008_ggml_binary_f32_case> @llama32_3b_f16_tg8_008_ggml_binary_f32 + +check.case public @llama32_3b_f16_tg8_009_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(3072) : index + %output_size = check.literal value(8192) : index + %input = check.generate.fill value(1.0) : tensor<1x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<8192x3072xf16> + %output = check.generate.fill value(0.0) : tensor<1x8192xf32> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<1x3072xf32>, tensor<8192x3072xf16>, tensor<1x8192xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_tg8_009_ggml_mul_mat_f32_f32_decode_wave64_case> @llama32_3b_f16_tg8_009_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @llama32_3b_f16_tg8_010_ggml_binary_f32_case { + %element_count = check.literal value(8192) : index + %lhs = check.generate.fill value(2.0) : tensor<8192xf32> + %rhs = check.generate.fill value(3.0) : tensor<8192xf32> + %output = check.generate.fill value(0.0) : tensor<8192xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<8192xf32>, tensor<8192xf32>, tensor<8192xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_tg8_010_ggml_binary_f32_case> @llama32_3b_f16_tg8_010_ggml_binary_f32 + +check.case public @llama32_3b_f16_tg8_011_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(8192) : index + %output_size = check.literal value(3072) : index + %input = check.generate.fill value(1.0) : tensor<1x8192xf32> + %weight = check.generate.fill value(1.0) : tensor<3072x8192xf16> + %output = check.generate.fill value(0.0) : tensor<1x3072xf32> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<1x8192xf32>, tensor<3072x8192xf16>, tensor<1x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_tg8_011_ggml_mul_mat_f32_f32_decode_wave64_case> @llama32_3b_f16_tg8_011_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @llama32_3b_f16_tg8_012_ggml_gather_add_f32_case { + %source_token_count = check.literal value(1) : index + %output_token_count = check.literal value(1) : index + %hidden_size = check.literal value(3072) : index + %attention = check.generate.fill value(0.0) : tensor<1x3072xf32> + %residual = check.generate.fill value(0.0) : tensor<1x3072xf32> + %output_ids = check.generate.fill value(0) : tensor<1xi32> + %output = check.generate.fill value(1.0) : tensor<1x3072xf32> + kernel.launch @ggml_gather_add_f32[%source_token_count, %output_token_count, %hidden_size](%source_token_count, %output_token_count, %hidden_size, %attention, %residual, %output_ids, %output) : [index, index, index](index, index, index, tensor<1x3072xf32>, tensor<1x3072xf32>, tensor<1xi32>, tensor<1x3072xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_tg8_012_ggml_gather_add_f32_case> @llama32_3b_f16_tg8_012_ggml_gather_add_f32 + +check.case public @llama32_3b_f16_tg8_013_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(3072) : index + %output_size = check.literal value(128256) : index + %input = check.generate.fill value(1.0) : tensor<1x3072xf32> + %weight = check.generate.fill value(1.0) : tensor<128256x3072xf16> + %output = check.generate.fill value(0.0) : tensor<1x128256xf32> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<1x3072xf32>, tensor<128256x3072xf16>, tensor<1x128256xf32>) + check.return +} + +check.benchmark<@llama32_3b_f16_tg8_013_ggml_mul_mat_f32_f32_decode_wave64_case> @llama32_3b_f16_tg8_013_ggml_mul_mat_f32_f32_decode_wave64 diff --git a/ggml/src/ggml-hrx/benchmarks/loom/llama32_3b_f16.pp256.json b/ggml/src/ggml-hrx/benchmarks/loom/llama32_3b_f16.pp256.json new file mode 100644 index 000000000000..1246ec73ebf8 --- /dev/null +++ b/ggml/src/ggml-hrx/benchmarks/loom/llama32_3b_f16.pp256.json @@ -0,0 +1,786 @@ +{ + "command_count": 259, + "dispatch_count": 17, + "dispatches": [ + { + "benchmark": "@llama32_3b_f16_pp256_000_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "3072", + "ggml.get_rows_f32.token_capacity": "64", + "ggml.get_rows_f32.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "hidden_size": 3072, + "row_count": 128256, + "token_count": 64 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp256_001_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "3072", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999975e-06" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 55, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp256_002_llm_attention_q_matmul_rope_f32_f32_wmma", + "compile_parameters": { + "ggml.workload.token_capacity": "64", + "llm.attention_qkv.head_count": "24", + "llm.attention_qkv.head_size": "128", + "llm.attention_qkv.input_size": "3072", + "llm.attention_qkv.output_size": "3072", + "llm.attention_qkv.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 28, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:llm_attention_q_matmul_rope_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/rope_f32.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/llm_attention_q_matmul_rope_f32_f32_wmma.loom" + ], + "sources": [ + "ops/llm_attention_q_matmul_rope_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/rope_f32.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "llm_attention_q_matmul_rope_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp256_003_llm_attention_v_matmul_set_rows_f32_f32_wmma", + "compile_parameters": { + "ggml.workload.token_capacity": "64", + "llm.attention_qkv.cache_output_format": "16", + "llm.attention_qkv.cache_row_count": "256", + "llm.attention_qkv.input_size": "3072", + "llm.attention_qkv.output_size": "1024", + "llm.attention_qkv.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 28, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:llm_attention_v_matmul_set_rows_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/llm_attention_v_matmul_set_rows_f32_f32_wmma.loom" + ], + "sources": [ + "ops/llm_attention_v_matmul_set_rows_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "llm_attention_v_matmul_set_rows_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp256_004_llm_attention_k_matmul_rope_set_rows_f32_f32_wmma", + "compile_parameters": { + "ggml.workload.token_capacity": "64", + "llm.attention_qkv.cache_output_format": "16", + "llm.attention_qkv.cache_row_count": "256", + "llm.attention_qkv.head_count": "8", + "llm.attention_qkv.head_size": "128", + "llm.attention_qkv.input_size": "3072", + "llm.attention_qkv.output_size": "1024", + "llm.attention_qkv.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 28, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:llm_attention_k_matmul_rope_set_rows_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/rope_f32.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/llm_attention_k_matmul_rope_set_rows_f32_f32_wmma.loom" + ], + "sources": [ + "ops/llm_attention_k_matmul_rope_set_rows_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/rope_f32.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "llm_attention_k_matmul_rope_set_rows_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp256_005_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.head_size": "128", + "ggml.flash_attention.key_value_head_count": "8", + "ggml.flash_attention.query_head_count": "24" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 28, + "integer_parameters": { + "key_value_token_count": 256, + "query_token_count": 64 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp256_006_ggml_mul_mat_add_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat_postops.input_size": "3072", + "ggml.mul_mat_postops.output_size": "3072", + "ggml.mul_mat_postops.weight_format": "16", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 27, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_add_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_add_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_add_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_add_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp256_007_ggml_mul_mat_swiglu_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat_swiglu.gate_weight_format": "16", + "ggml.mul_mat_swiglu.input_size": "3072", + "ggml.mul_mat_swiglu.op": "4", + "ggml.mul_mat_swiglu.output_size": "8192", + "ggml.mul_mat_swiglu.up_weight_format": "16", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 27, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_swiglu_f32_f32_wmma", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_swiglu_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_swiglu_f32_f32_wmma.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_swiglu_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp256_008_ggml_mul_mat_add_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat_postops.input_size": "8192", + "ggml.mul_mat_postops.output_size": "3072", + "ggml.mul_mat_postops.weight_format": "16", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 27, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_add_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_add_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_add_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_add_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp256_009_ggml_mul_mat_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat.input_size": "3072", + "ggml.mul_mat.output_accumulation": "0", + "ggml.mul_mat.output_size": "3072", + "ggml.mul_mat.output_unary_op": "23", + "ggml.mul_mat.weight_format": "16", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp256_010_ggml_gather_add_f32", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/hrx", + "count": 1, + "integer_parameters": { + "hidden_size": 3072, + "output_token_count": 1, + "source_token_count": 64 + }, + "kernel": "hrx:ggml_gather_add_f32", + "library_sources": [], + "primary_sources": [ + "gather_add_f32.loom" + ], + "sources": [ + "gather_add_f32.loom" + ], + "symbol": "ggml_gather_add_f32", + "workload_parameters": [ + { + "name": "source_token_count", + "type": "index" + }, + { + "name": "output_token_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp256_011_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "3072", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999975e-06" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "token_count": 1 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp256_012_ggml_mul_mat_f32_f32_decode_wave64", + "compile_parameters": { + "ggml.mul_mat_f32_f32_decode.output_capacity": "8192", + "ggml.mul_mat_f32_f32_decode.token_capacity": "1", + "ggml.mul_mat_f32_f32_decode.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "input_size": 3072, + "output_size": 8192, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp256_013_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "element_count": 8192 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp256_014_ggml_mul_mat_f32_f32_decode_wave64", + "compile_parameters": { + "ggml.mul_mat_f32_f32_decode.output_capacity": "3072", + "ggml.mul_mat_f32_f32_decode.token_capacity": "1", + "ggml.mul_mat_f32_f32_decode.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "input_size": 8192, + "output_size": 3072, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp256_015_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "element_count": 3072 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp256_016_ggml_mul_mat_f32_f32_decode_wave64", + "compile_parameters": { + "ggml.mul_mat_f32_f32_decode.output_capacity": "128256", + "ggml.mul_mat_f32_f32_decode.token_capacity": "1", + "ggml.mul_mat_f32_f32_decode.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "input_size": 3072, + "output_size": 128256, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + } + ], + "generated_count": 17, + "kernel_counts": { + "hrx:ggml_gather_add_f32": 1, + "loom_libs:ggml_binary_f32": 2, + "loom_libs:ggml_flash_attention_f32_f16_wmma": 28, + "loom_libs:ggml_get_rows_f32": 1, + "loom_libs:ggml_mul_mat_add_f32_f32_wmma": 54, + "loom_libs:ggml_mul_mat_f32_f32_decode_wave64": 4, + "loom_libs:ggml_mul_mat_f32_f32_wmma": 1, + "loom_libs:ggml_mul_mat_swiglu_f32_f32_wmma": 27, + "loom_libs:ggml_rmsnorm_binary_f32": 57, + "loom_libs:llm_attention_k_matmul_rope_set_rows_f32_f32_wmma": 28, + "loom_libs:llm_attention_q_matmul_rope_f32_f32_wmma": 28, + "loom_libs:llm_attention_v_matmul_set_rows_f32_f32_wmma": 28 + }, + "loom_source": "benchmarks/loom/llama32_3b_f16.loom", + "model": "llama32_3b_f16", + "scenario": "pp256", + "schema": "ggml-hrx-model-loom-benchmarks-v2" +} diff --git a/ggml/src/ggml-hrx/benchmarks/loom/llama32_3b_f16.pp512.json b/ggml/src/ggml-hrx/benchmarks/loom/llama32_3b_f16.pp512.json new file mode 100644 index 000000000000..d723d613a102 --- /dev/null +++ b/ggml/src/ggml-hrx/benchmarks/loom/llama32_3b_f16.pp512.json @@ -0,0 +1,819 @@ +{ + "command_count": 518, + "dispatch_count": 18, + "dispatches": [ + { + "benchmark": "@llama32_3b_f16_pp512_000_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "3072", + "ggml.get_rows_f32.token_capacity": "64", + "ggml.get_rows_f32.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "hidden_size": 3072, + "row_count": 128256, + "token_count": 64 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp512_001_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "3072", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999975e-06" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 110, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp512_002_llm_attention_q_matmul_rope_f32_f32_wmma", + "compile_parameters": { + "ggml.workload.token_capacity": "64", + "llm.attention_qkv.head_count": "24", + "llm.attention_qkv.head_size": "128", + "llm.attention_qkv.input_size": "3072", + "llm.attention_qkv.output_size": "3072", + "llm.attention_qkv.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 56, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:llm_attention_q_matmul_rope_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/rope_f32.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/llm_attention_q_matmul_rope_f32_f32_wmma.loom" + ], + "sources": [ + "ops/llm_attention_q_matmul_rope_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/rope_f32.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "llm_attention_q_matmul_rope_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp512_003_llm_attention_v_matmul_set_rows_f32_f32_wmma", + "compile_parameters": { + "ggml.workload.token_capacity": "64", + "llm.attention_qkv.cache_output_format": "16", + "llm.attention_qkv.cache_row_count": "512", + "llm.attention_qkv.input_size": "3072", + "llm.attention_qkv.output_size": "1024", + "llm.attention_qkv.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 56, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:llm_attention_v_matmul_set_rows_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/llm_attention_v_matmul_set_rows_f32_f32_wmma.loom" + ], + "sources": [ + "ops/llm_attention_v_matmul_set_rows_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "llm_attention_v_matmul_set_rows_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp512_004_llm_attention_k_matmul_rope_set_rows_f32_f32_wmma", + "compile_parameters": { + "ggml.workload.token_capacity": "64", + "llm.attention_qkv.cache_output_format": "16", + "llm.attention_qkv.cache_row_count": "512", + "llm.attention_qkv.head_count": "8", + "llm.attention_qkv.head_size": "128", + "llm.attention_qkv.input_size": "3072", + "llm.attention_qkv.output_size": "1024", + "llm.attention_qkv.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 56, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:llm_attention_k_matmul_rope_set_rows_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/rope_f32.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/llm_attention_k_matmul_rope_set_rows_f32_f32_wmma.loom" + ], + "sources": [ + "ops/llm_attention_k_matmul_rope_set_rows_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/rope_f32.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "llm_attention_k_matmul_rope_set_rows_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp512_005_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.head_size": "128", + "ggml.flash_attention.key_value_head_count": "8", + "ggml.flash_attention.query_head_count": "24" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 28, + "integer_parameters": { + "key_value_token_count": 256, + "query_token_count": 64 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp512_006_ggml_mul_mat_add_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat_postops.input_size": "3072", + "ggml.mul_mat_postops.output_size": "3072", + "ggml.mul_mat_postops.weight_format": "16", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 54, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_add_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_add_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_add_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_add_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp512_007_ggml_mul_mat_swiglu_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat_swiglu.gate_weight_format": "16", + "ggml.mul_mat_swiglu.input_size": "3072", + "ggml.mul_mat_swiglu.op": "4", + "ggml.mul_mat_swiglu.output_size": "8192", + "ggml.mul_mat_swiglu.up_weight_format": "16", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 54, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_swiglu_f32_f32_wmma", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_swiglu_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_swiglu_f32_f32_wmma.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_swiglu_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp512_008_ggml_mul_mat_add_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat_postops.input_size": "8192", + "ggml.mul_mat_postops.output_size": "3072", + "ggml.mul_mat_postops.weight_format": "16", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 54, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_add_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_add_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_add_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_add_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp512_009_ggml_mul_mat_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat.input_size": "3072", + "ggml.mul_mat.output_accumulation": "0", + "ggml.mul_mat.output_size": "3072", + "ggml.mul_mat.output_unary_op": "23", + "ggml.mul_mat.weight_format": "16", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp512_010_ggml_gather_add_f32", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/hrx", + "count": 2, + "integer_parameters": { + "hidden_size": 3072, + "output_token_count": 1, + "source_token_count": 64 + }, + "kernel": "hrx:ggml_gather_add_f32", + "library_sources": [], + "primary_sources": [ + "gather_add_f32.loom" + ], + "sources": [ + "gather_add_f32.loom" + ], + "symbol": "ggml_gather_add_f32", + "workload_parameters": [ + { + "name": "source_token_count", + "type": "index" + }, + { + "name": "output_token_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp512_011_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "3072", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999975e-06" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 4, + "integer_parameters": { + "token_count": 1 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp512_012_ggml_mul_mat_f32_f32_decode_wave64", + "compile_parameters": { + "ggml.mul_mat_f32_f32_decode.output_capacity": "8192", + "ggml.mul_mat_f32_f32_decode.token_capacity": "1", + "ggml.mul_mat_f32_f32_decode.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 4, + "integer_parameters": { + "input_size": 3072, + "output_size": 8192, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp512_013_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "element_count": 8192 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp512_014_ggml_mul_mat_f32_f32_decode_wave64", + "compile_parameters": { + "ggml.mul_mat_f32_f32_decode.output_capacity": "3072", + "ggml.mul_mat_f32_f32_decode.token_capacity": "1", + "ggml.mul_mat_f32_f32_decode.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "input_size": 8192, + "output_size": 3072, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp512_015_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "element_count": 3072 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp512_016_ggml_mul_mat_f32_f32_decode_wave64", + "compile_parameters": { + "ggml.mul_mat_f32_f32_decode.output_capacity": "128256", + "ggml.mul_mat_f32_f32_decode.token_capacity": "1", + "ggml.mul_mat_f32_f32_decode.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "input_size": 3072, + "output_size": 128256, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_pp512_017_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.head_size": "128", + "ggml.flash_attention.key_value_head_count": "8", + "ggml.flash_attention.query_head_count": "24" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 28, + "integer_parameters": { + "key_value_token_count": 512, + "query_token_count": 64 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + } + ], + "generated_count": 18, + "kernel_counts": { + "hrx:ggml_gather_add_f32": 2, + "loom_libs:ggml_binary_f32": 4, + "loom_libs:ggml_flash_attention_f32_f16_wmma": 56, + "loom_libs:ggml_get_rows_f32": 2, + "loom_libs:ggml_mul_mat_add_f32_f32_wmma": 108, + "loom_libs:ggml_mul_mat_f32_f32_decode_wave64": 8, + "loom_libs:ggml_mul_mat_f32_f32_wmma": 2, + "loom_libs:ggml_mul_mat_swiglu_f32_f32_wmma": 54, + "loom_libs:ggml_rmsnorm_binary_f32": 114, + "loom_libs:llm_attention_k_matmul_rope_set_rows_f32_f32_wmma": 56, + "loom_libs:llm_attention_q_matmul_rope_f32_f32_wmma": 56, + "loom_libs:llm_attention_v_matmul_set_rows_f32_f32_wmma": 56 + }, + "loom_source": "benchmarks/loom/llama32_3b_f16.loom", + "model": "llama32_3b_f16", + "scenario": "pp512", + "schema": "ggml-hrx-model-loom-benchmarks-v2" +} diff --git a/ggml/src/ggml-hrx/benchmarks/loom/llama32_3b_f16.tg8.json b/ggml/src/ggml-hrx/benchmarks/loom/llama32_3b_f16.tg8.json new file mode 100644 index 000000000000..c360cd1b0bb4 --- /dev/null +++ b/ggml/src/ggml-hrx/benchmarks/loom/llama32_3b_f16.tg8.json @@ -0,0 +1,606 @@ +{ + "command_count": 451, + "dispatch_count": 14, + "dispatches": [ + { + "benchmark": "@llama32_3b_f16_tg8_000_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "3072", + "ggml.get_rows_f32.token_capacity": "1", + "ggml.get_rows_f32.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "hidden_size": 3072, + "row_count": 128256, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_tg8_001_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "3072", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999975e-06" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 57, + "integer_parameters": { + "token_count": 1 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_tg8_002_ggml_mul_mat_f32_f32_decode_wave64", + "compile_parameters": { + "ggml.mul_mat_f32_f32_decode.output_capacity": "3072", + "ggml.mul_mat_f32_f32_decode.token_capacity": "1", + "ggml.mul_mat_f32_f32_decode.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 56, + "integer_parameters": { + "input_size": 3072, + "output_size": 3072, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_tg8_003_ggml_mul_mat_f32_f32_decode_wave64", + "compile_parameters": { + "ggml.mul_mat_f32_f32_decode.output_capacity": "1024", + "ggml.mul_mat_f32_f32_decode.token_capacity": "1", + "ggml.mul_mat_f32_f32_decode.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 56, + "integer_parameters": { + "input_size": 3072, + "output_size": 1024, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_tg8_004_ggml_rope_f32", + "compile_parameters": { + "ggml.rope_f32.head_count": "24", + "ggml.rope_f32.head_size": "128", + "ggml.rope_f32.mode": "0", + "ggml.rope_f32.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 28, + "integer_parameters": { + "token_count": 1 + }, + "kernel": "loom_libs:ggml_rope_f32", + "library_sources": [ + "motifs/rope_f32.loom" + ], + "primary_sources": [ + "ops/rope_f32.loom" + ], + "sources": [ + "ops/rope_f32.loom", + "motifs/rope_f32.loom" + ], + "symbol": "ggml_rope_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_tg8_005_ggml_rope_set_rows_f32", + "compile_parameters": { + "ggml.rope_set_rows_f32.head_count": "8", + "ggml.rope_set_rows_f32.head_size": "128", + "ggml.rope_set_rows_f32.mode": "0", + "ggml.rope_set_rows_f32.output_format": "16", + "ggml.rope_set_rows_f32.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 28, + "integer_parameters": { + "cache_row_count": 256, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_rope_set_rows_f32", + "library_sources": [ + "motifs/rope_f32.loom" + ], + "primary_sources": [ + "ops/rope_set_rows_f32.loom" + ], + "sources": [ + "ops/rope_set_rows_f32.loom", + "motifs/rope_f32.loom" + ], + "symbol": "ggml_rope_set_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_tg8_006_ggml_set_rows", + "compile_parameters": { + "ggml.set_rows.hidden_capacity": "1024", + "ggml.set_rows.input_format": "32", + "ggml.set_rows.output_format": "16", + "ggml.set_rows.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 28, + "integer_parameters": { + "cache_row_count": 256, + "hidden_size": 1024, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_set_rows", + "library_sources": [], + "primary_sources": [ + "ops/set_rows.loom" + ], + "sources": [ + "ops/set_rows.loom" + ], + "symbol": "ggml_set_rows", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_tg8_007_ggml_flash_attention_decode_split_f32_f16_wmma_next_q8", + "compile_parameters": { + "ggml.flash_attention.decode.key_value_token_capacity": "256", + "ggml.flash_attention.key_value_head_count": "8", + "ggml.flash_attention.query_head_count": "24" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 28, + "integer_parameters": { + "key_value_token_count": 256 + }, + "kernel": "loom_libs:ggml_flash_attention_decode_split_f32_f16_wmma_next_q8", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_decode_split_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_decode_split_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_decode_split_f32_f16_wmma_next_q8", + "workload_parameters": [ + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_tg8_008_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 55, + "integer_parameters": { + "element_count": 3072 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_tg8_009_ggml_mul_mat_f32_f32_decode_wave64", + "compile_parameters": { + "ggml.mul_mat_f32_f32_decode.output_capacity": "8192", + "ggml.mul_mat_f32_f32_decode.token_capacity": "1", + "ggml.mul_mat_f32_f32_decode.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 56, + "integer_parameters": { + "input_size": 3072, + "output_size": 8192, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_tg8_010_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 28, + "integer_parameters": { + "element_count": 8192 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_tg8_011_ggml_mul_mat_f32_f32_decode_wave64", + "compile_parameters": { + "ggml.mul_mat_f32_f32_decode.output_capacity": "3072", + "ggml.mul_mat_f32_f32_decode.token_capacity": "1", + "ggml.mul_mat_f32_f32_decode.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 28, + "integer_parameters": { + "input_size": 8192, + "output_size": 3072, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_tg8_012_ggml_gather_add_f32", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/hrx", + "count": 1, + "integer_parameters": { + "hidden_size": 3072, + "output_token_count": 1, + "source_token_count": 1 + }, + "kernel": "hrx:ggml_gather_add_f32", + "library_sources": [], + "primary_sources": [ + "gather_add_f32.loom" + ], + "sources": [ + "gather_add_f32.loom" + ], + "symbol": "ggml_gather_add_f32", + "workload_parameters": [ + { + "name": "source_token_count", + "type": "index" + }, + { + "name": "output_token_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@llama32_3b_f16_tg8_013_ggml_mul_mat_f32_f32_decode_wave64", + "compile_parameters": { + "ggml.mul_mat_f32_f32_decode.output_capacity": "128256", + "ggml.mul_mat_f32_f32_decode.token_capacity": "1", + "ggml.mul_mat_f32_f32_decode.weight_format": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "input_size": 3072, + "output_size": 128256, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + } + ], + "generated_count": 14, + "kernel_counts": { + "hrx:ggml_gather_add_f32": 1, + "loom_libs:ggml_binary_f32": 83, + "loom_libs:ggml_flash_attention_decode_split_f32_f16_wmma_next_q8": 28, + "loom_libs:ggml_get_rows_f32": 1, + "loom_libs:ggml_mul_mat_f32_f32_decode_wave64": 197, + "loom_libs:ggml_rmsnorm_binary_f32": 57, + "loom_libs:ggml_rope_f32": 28, + "loom_libs:ggml_rope_set_rows_f32": 28, + "loom_libs:ggml_set_rows": 28 + }, + "loom_source": "benchmarks/loom/llama32_3b_f16.loom", + "model": "llama32_3b_f16", + "scenario": "tg8", + "schema": "ggml-hrx-model-loom-benchmarks-v2" +} diff --git a/ggml/src/ggml-hrx/benchmarks/loom/qwen3_30b_a3b_q4_k_m.loom b/ggml/src/ggml-hrx/benchmarks/loom/qwen3_30b_a3b_q4_k_m.loom new file mode 100644 index 000000000000..e78d397230f1 --- /dev/null +++ b/ggml/src/ggml-hrx/benchmarks/loom/qwen3_30b_a3b_q4_k_m.loom @@ -0,0 +1,982 @@ +// Generated by tools/benchmarks/generate-model-benchmarks.py for qwen3_30b_a3b_q4_k_m scenarios: pp256, pp512, tg8. +// Regenerate from an HRX command program dump rather than editing by hand. + +target.decl @ggml_flash_attention_decode_split_gfx11_wave64 +target.decl @ggml_flash_attention_gfx11_wave64 +target.decl @ggml_gather_add_gfx11_wave64 +target.decl @ggml_get_rows_f32_gfx11_wave64 +target.decl @ggml_moe_routing_gfx11_wave32 +target.decl @ggml_mul_mat_add_gfx11_wave64 +target.decl @ggml_mul_mat_gfx11_wave64 +target.decl @ggml_mul_mat_id_f16_f16_gfx11_wave64 +target.decl @ggml_rmsnorm_binary_gfx11_wave32 +target.decl @qwen3_moe_attention_postprocess_gfx11_wave32 +target.decl @qwen3_moe_attention_prepare_gfx11_wave32 +target.decl @qwen3_moe_attention_qkv_postprocess_gfx11_wave32 +target.decl @qwen3_moe_gfx11_wave32 +target.decl @qwen3_moe_routed_down_next_norm_gfx11_wave32 +target.decl @qwen3_moe_routed_down_q6k_gfx11_wave64 +target.decl @qwen3_moe_router_fused_gfx11_wave64 +target.decl @qwen3_moe_router_gfx11_wave64 +target.decl @qwen3_moe_router_projection_gfx11_wave32 +target.decl @qwen_attention_metadata_gfx11_wave64 +target.decl @qwen_attention_state_gfx11_wave64 + +kernel.decl target(@ggml_flash_attention_decode_split_gfx11_wave64) @ggml_flash_attention_decode_split_f32_f16_wmma_next_q8(%key_value_token_count: index) launch(%key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %completion_counter: buffer, %output: buffer, %next_q8_output: buffer) +kernel.decl target(@ggml_flash_attention_gfx11_wave64) @ggml_flash_attention_f32_f16_wmma(%query_token_count: index, %key_value_token_count: index) launch(%query_token_count: index, %key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %output: buffer) +kernel.decl target(@ggml_gather_add_gfx11_wave64) @ggml_gather_add_f32(%source_token_count: index, %output_token_count: index, %hidden_size: index) launch(%source_token_count: index, %output_token_count: index, %hidden_size: index, %attention: buffer, %residual: buffer, %output_ids: buffer, %output: buffer) +kernel.decl target(@ggml_get_rows_f32_gfx11_wave64) @ggml_get_rows_f32(%token_count: index, %row_count: index, %hidden_size: index) launch(%token_count: index, %row_count: index, %hidden_size: index, %token_ids: buffer, %weight: buffer, %output: buffer) +kernel.decl @ggml_linear_q6k_q8_1_x4(%token_count: index, %input_size: index, %output_size: index) launch(%token_count: index, %input_size: index, %output_size: index, %q8_input: buffer, %weight: buffer, %output: buffer) +kernel.decl target(@ggml_moe_routing_gfx11_wave32) @ggml_moe_build_expert_partition_table(%token_count: index, %route_count: index, %expert_count: index) launch(%token_count: index, %route_count: index, %expert_count: index, %expert_table: buffer, %partition_table: buffer) +kernel.decl target(@ggml_moe_routing_gfx11_wave32) @ggml_moe_build_expert_table(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index) launch(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index, %route_ids: buffer, %expert_table: buffer) +kernel.decl target(@ggml_mul_mat_add_gfx11_wave64) @ggml_mul_mat_add_f32_f32_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %residual_input: buffer, %residual_output: buffer) +kernel.decl target(@ggml_mul_mat_gfx11_wave64) @ggml_mul_mat_f32_f32_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) +kernel.decl target(@ggml_mul_mat_id_f16_f16_gfx11_wave64) @ggml_mul_mat_id_f16_f16_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %expert_table: buffer, %weight: buffer, %output: buffer) +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_rmsnorm_binary_f32(%token_count: index) launch(%token_count: index, %input: buffer, %rhs: buffer, %output: buffer) +kernel.decl target(@qwen3_moe_attention_postprocess_gfx11_wave32) @qwen3_moe_attention_postprocess_f32_f16(%token_count: index, %cache_row_count: index) launch(%token_count: index, %cache_row_count: index, %positions: buffer, %key_cache_indices: buffer, %value_cache_indices: buffer, %query_input: buffer, %key_input: buffer, %value_input: buffer, %query_norm_weight: buffer, %key_norm_weight: buffer, %inverse_frequencies: buffer, %query_output: buffer, %key_cache: buffer, %value_cache: buffer) +kernel.decl target(@qwen3_moe_attention_qkv_postprocess_gfx11_wave32) @qwen3_moe_attention_qkv_postprocess_fused_decode(%token_count: index, %cache_row_count: index) launch(%token_count: index, %cache_row_count: index, %q8_input: buffer, %query_weight: buffer, %key_weight: buffer, %value_weight: buffer, %positions: buffer, %key_cache_indices: buffer, %value_cache_indices: buffer, %query_output_raw: buffer, %key_output_raw: buffer, %value_output_raw: buffer, %query_norm_weight: buffer, %key_norm_weight: buffer, %inverse_frequencies: buffer, %query_output: buffer, %key_cache: buffer, %value_cache: buffer, %completion_counters: buffer) +kernel.decl target(@qwen3_moe_attention_prepare_gfx11_wave32) @qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8(%token_count: index) launch(%token_count: index, %q8_input: buffer, %weight: buffer, %output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counter: buffer, %next_q8_output: buffer) +kernel.decl target(@qwen3_moe_attention_prepare_gfx11_wave32) @qwen3_moe_rmsnorm_f32_quantize_q8_1_x4(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %normalized_output: buffer, %q8_output: buffer) +kernel.decl target(@qwen3_moe_attention_prepare_gfx11_wave32) @qwen3_moe_routed_down_q4k_q8_1_x4_next_q8(%token_count: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index) launch(%token_count: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index, %q8_input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer, %norm_weight: buffer, %completion_counter: buffer, %next_q8_output: buffer) +kernel.decl target(@qwen3_moe_routed_down_q6k_gfx11_wave64) @qwen3_moe_routed_down_q6k_f32_wave64_next_q8(%token_count: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index) launch(%token_count: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index, %input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer, %norm_weight: buffer, %completion_counter: buffer, %next_q8_output: buffer) +kernel.decl target(@qwen3_moe_routed_down_next_norm_gfx11_wave32) @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32(%token_count: index) launch(%token_count: index, %route_weights: buffer, %routed_output: buffer, %hidden_state: buffer, %next_norm_weight: buffer, %next_projection_input: buffer) +kernel.decl target(@qwen3_moe_gfx11_wave32) @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %expert_table: buffer, %partition_table: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) +kernel.decl @qwen3_moe_routed_gate_up_swiglu_q4k_q8(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index, %output_size: index) launch(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index, %output_size: index, %q8_input: buffer, %route_ids: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) +kernel.decl @qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index, %output_size: index) launch(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index, %output_size: index, %q8_input: buffer, %route_ids: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer, %completion_counters: buffer, %next_q8_output: buffer) +kernel.decl target(@qwen3_moe_router_projection_gfx11_wave32) @qwen3_moe_router_projection_f32_four_row_wave32(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) +kernel.decl target(@qwen3_moe_router_fused_gfx11_wave64) @qwen3_moe_router_projection_top8_fused_decode_f32(%token_count: index, %route_id_stride: index) launch(%token_count: index, %route_id_stride: index, %input: buffer, %weight: buffer, %logits: buffer, %completion_counter: buffer, %route_ids: buffer, %route_weights: buffer) +kernel.decl target(@qwen3_moe_router_gfx11_wave64) @qwen3_moe_router_top8_f32(%token_count: index, %route_id_stride: index) launch(%token_count: index, %route_id_stride: index, %logits: buffer, %route_ids: buffer, %route_weights: buffer) +kernel.decl target(@qwen_attention_state_gfx11_wave64) @qwen_attention_context_base_capture() launch(%positions: buffer, %control: buffer) +kernel.decl target(@qwen_attention_metadata_gfx11_wave64) @qwen_attention_metadata(%token_count: index, %context_capacity: index) launch(%token_count: index, %context_capacity: index, %control: buffer, %positions: buffer, %key_cache_indices: buffer, %value_cache_indices: buffer, %attention_mask: buffer) + + +// Scenario: pp256 + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_000_qwen_attention_context_base_capture_case { + %positions = check.generate.fill value(0) : tensor<256xi8> + %control = check.generate.fill value(0) : tensor<4xi8> + kernel.launch @qwen_attention_context_base_capture[](%positions, %control) : [](tensor<256xi8>, tensor<4xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_000_qwen_attention_context_base_capture_case> @qwen3_30b_a3b_q4_k_m_pp256_000_qwen_attention_context_base_capture + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_001_qwen_attention_metadata_case { + %token_count = check.literal value(64) : index + %context_capacity = check.literal value(256) : index + %control = check.generate.fill value(0) : tensor<4xi8> + %positions = check.generate.fill value(0) : tensor<256xi8> + %key_cache_indices = check.generate.fill value(0) : tensor<512xi8> + %value_cache_indices = check.generate.fill value(0) : tensor<512xi8> + %attention_mask = check.generate.fill value(0) : tensor<32768xi8> + kernel.launch @qwen_attention_metadata[%token_count, %context_capacity](%token_count, %context_capacity, %control, %positions, %key_cache_indices, %value_cache_indices, %attention_mask) : [index, index](index, index, tensor<4xi8>, tensor<256xi8>, tensor<512xi8>, tensor<512xi8>, tensor<32768xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_001_qwen_attention_metadata_case> @qwen3_30b_a3b_q4_k_m_pp256_001_qwen_attention_metadata + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_002_ggml_get_rows_f32_case { + %token_count = check.literal value(64) : index + %row_count = check.literal value(151936) : index + %hidden_size = check.literal value(2048) : index + %token_ids = check.generate.fill value(0) : tensor<256xi8> + %weight = check.generate.fill value(0) : tensor<175030272xi8> + %output = check.generate.fill value(0) : tensor<524288xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<256xi8>, tensor<175030272xi8>, tensor<524288xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_002_ggml_get_rows_f32_case> @qwen3_30b_a3b_q4_k_m_pp256_002_ggml_get_rows_f32 + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_003_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(2.0) : tensor<64x2048xf32> + %rhs = check.generate.fill value(3.0) : tensor<2048xf32> + %output = check.generate.fill value(0.0) : tensor<64x2048xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<64x2048xf32>, tensor<2048xf32>, tensor<64x2048xf32>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_003_ggml_rmsnorm_binary_f32_case> @qwen3_30b_a3b_q4_k_m_pp256_003_ggml_rmsnorm_binary_f32 + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_004_ggml_mul_mat_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<524288xi8> + %weight = check.generate.fill value(0) : tensor<4718592xi8> + %output = check.generate.fill value(0) : tensor<1048576xi8> + kernel.launch @ggml_mul_mat_f32_f32_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<524288xi8>, tensor<4718592xi8>, tensor<1048576xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_004_ggml_mul_mat_f32_f32_wmma_case> @qwen3_30b_a3b_q4_k_m_pp256_004_ggml_mul_mat_f32_f32_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_005_ggml_mul_mat_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<524288xi8> + %weight = check.generate.fill value(0) : tensor<860160xi8> + %output = check.generate.fill value(0) : tensor<131072xi8> + kernel.launch @ggml_mul_mat_f32_f32_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<524288xi8>, tensor<860160xi8>, tensor<131072xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_005_ggml_mul_mat_f32_f32_wmma_case> @qwen3_30b_a3b_q4_k_m_pp256_005_ggml_mul_mat_f32_f32_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_006_ggml_mul_mat_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<524288xi8> + %weight = check.generate.fill value(0) : tensor<589824xi8> + %output = check.generate.fill value(0) : tensor<131072xi8> + kernel.launch @ggml_mul_mat_f32_f32_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<524288xi8>, tensor<589824xi8>, tensor<131072xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_006_ggml_mul_mat_f32_f32_wmma_case> @qwen3_30b_a3b_q4_k_m_pp256_006_ggml_mul_mat_f32_f32_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_007_qwen3_moe_attention_postprocess_f32_f16_case { + %token_count = check.literal value(64) : index + %cache_row_count = check.literal value(256) : index + %positions = check.generate.fill value(0) : tensor<256xi8> + %key_cache_indices = check.generate.fill value(0) : tensor<512xi8> + %value_cache_indices = check.generate.fill value(0) : tensor<512xi8> + %query_input = check.generate.fill value(0) : tensor<1048576xi8> + %key_input = check.generate.fill value(0) : tensor<131072xi8> + %value_input = check.generate.fill value(0) : tensor<131072xi8> + %query_norm_weight = check.generate.fill value(0) : tensor<512xi8> + %key_norm_weight = check.generate.fill value(0) : tensor<512xi8> + %inverse_frequencies = check.generate.fill value(0) : tensor<256xi8> + %query_output = check.generate.fill value(0) : tensor<1048576xi8> + %key_cache = check.generate.fill value(0) : tensor<262144xi8> + %value_cache = check.generate.fill value(0) : tensor<262144xi8> + kernel.launch @qwen3_moe_attention_postprocess_f32_f16[%token_count, %cache_row_count](%token_count, %cache_row_count, %positions, %key_cache_indices, %value_cache_indices, %query_input, %key_input, %value_input, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %query_output, %key_cache, %value_cache) : [index, index](index, index, tensor<256xi8>, tensor<512xi8>, tensor<512xi8>, tensor<1048576xi8>, tensor<131072xi8>, tensor<131072xi8>, tensor<512xi8>, tensor<512xi8>, tensor<256xi8>, tensor<1048576xi8>, tensor<262144xi8>, tensor<262144xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_007_qwen3_moe_attention_postprocess_f32_f16_case> @qwen3_30b_a3b_q4_k_m_pp256_007_qwen3_moe_attention_postprocess_f32_f16 + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_008_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(64) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<64x32x128xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x128xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<64x256xf16> + %output = check.generate.fill value(1.0) : tensor<64x32x128xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<64x32x128xf32>, tensor<256x4x128xf16>, tensor<256x4x128xf16>, tensor<64x256xf16>, tensor<64x32x128xf32>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_008_ggml_flash_attention_f32_f16_wmma_case> @qwen3_30b_a3b_q4_k_m_pp256_008_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_009_ggml_mul_mat_add_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<1048576xi8> + %weight = check.generate.fill value(0) : tensor<4718592xi8> + %residual_input = check.generate.fill value(0) : tensor<524288xi8> + %residual_output = check.generate.fill value(0) : tensor<524288xi8> + kernel.launch @ggml_mul_mat_add_f32_f32_wmma[%token_count](%token_count, %input, %weight, %residual_input, %residual_output) : [index](index, tensor<1048576xi8>, tensor<4718592xi8>, tensor<524288xi8>, tensor<524288xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_009_ggml_mul_mat_add_f32_f32_wmma_case> @qwen3_30b_a3b_q4_k_m_pp256_009_ggml_mul_mat_add_f32_f32_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_010_qwen3_moe_router_projection_f32_four_row_wave32_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<524288xi8> + %weight = check.generate.fill value(0) : tensor<1048576xi8> + %output = check.generate.fill value(0) : tensor<32768xi8> + kernel.launch @qwen3_moe_router_projection_f32_four_row_wave32[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<524288xi8>, tensor<1048576xi8>, tensor<32768xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_010_qwen3_moe_router_projection_f32_four_row_wave32_case> @qwen3_30b_a3b_q4_k_m_pp256_010_qwen3_moe_router_projection_f32_four_row_wave32 + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_011_qwen3_moe_router_top8_f32_case { + %token_count = check.literal value(64) : index + %route_id_stride = check.literal value(128) : index + %logits = check.generate.fill value(0) : tensor<32768xi8> + %route_ids = check.generate.fill value(0) : tensor<32768xi8> + %route_weights = check.generate.fill value(0) : tensor<2048xi8> + kernel.launch @qwen3_moe_router_top8_f32[%token_count, %route_id_stride](%token_count, %route_id_stride, %logits, %route_ids, %route_weights) : [index, index](index, index, tensor<32768xi8>, tensor<32768xi8>, tensor<2048xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_011_qwen3_moe_router_top8_f32_case> @qwen3_30b_a3b_q4_k_m_pp256_011_qwen3_moe_router_top8_f32 + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_012_ggml_moe_build_expert_table_case { + %token_count = check.literal value(64) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.fill value(0) : tensor<32768xi8> + %expert_table = check.generate.fill value(0) : tensor<33280xi8> + kernel.launch @ggml_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<32768xi8>, tensor<33280xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_012_ggml_moe_build_expert_table_case> @qwen3_30b_a3b_q4_k_m_pp256_012_ggml_moe_build_expert_table + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_013_ggml_moe_build_expert_partition_table_case { + %token_count = check.literal value(64) : index + %route_count = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %expert_table = check.generate.fill value(0) : tensor<33280xi8> + %partition_table = check.generate.fill value(0) : tensor<580xi8> + kernel.launch @ggml_moe_build_expert_partition_table[%token_count, %route_count, %expert_count](%token_count, %route_count, %expert_count, %expert_table, %partition_table) : [index, index, index](index, index, index, tensor<33280xi8>, tensor<580xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_013_ggml_moe_build_expert_partition_table_case> @qwen3_30b_a3b_q4_k_m_pp256_013_ggml_moe_build_expert_partition_table + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_014_qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<524288xi8> + %expert_table = check.generate.fill value(0) : tensor<33280xi8> + %partition_table = check.generate.fill value(0) : tensor<580xi8> + %gate_weight = check.generate.fill value(0) : tensor<113246208xi8> + %up_weight = check.generate.fill value(0) : tensor<113246208xi8> + %output = check.generate.fill value(0) : tensor<786432xi8> + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma[%token_count](%token_count, %input, %expert_table, %partition_table, %gate_weight, %up_weight, %output) : [index](index, tensor<524288xi8>, tensor<33280xi8>, tensor<580xi8>, tensor<113246208xi8>, tensor<113246208xi8>, tensor<786432xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_014_qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_case> @qwen3_30b_a3b_q4_k_m_pp256_014_qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_015_ggml_mul_mat_id_f16_f16_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<786432xi8> + %expert_table = check.generate.fill value(0) : tensor<33280xi8> + %weight = check.generate.fill value(0) : tensor<165150720xi8> + %output = check.generate.fill value(0) : tensor<2097152xi8> + kernel.launch @ggml_mul_mat_id_f16_f16_wmma[%token_count](%token_count, %input, %expert_table, %weight, %output) : [index](index, tensor<786432xi8>, tensor<33280xi8>, tensor<165150720xi8>, tensor<2097152xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_015_ggml_mul_mat_id_f16_f16_wmma_case> @qwen3_30b_a3b_q4_k_m_pp256_015_ggml_mul_mat_id_f16_f16_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_016_qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32_case { + %token_count = check.literal value(64) : index + %route_weights = check.generate.fill value(0) : tensor<2048xi8> + %routed_output = check.generate.fill value(0) : tensor<2097152xi8> + %hidden_state = check.generate.fill value(0) : tensor<524288xi8> + %next_norm_weight = check.generate.fill value(0) : tensor<8192xi8> + %next_projection_input = check.generate.fill value(0) : tensor<524288xi8> + kernel.launch @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32[%token_count](%token_count, %route_weights, %routed_output, %hidden_state, %next_norm_weight, %next_projection_input) : [index](index, tensor<2048xi8>, tensor<2097152xi8>, tensor<524288xi8>, tensor<8192xi8>, tensor<524288xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_016_qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32_case> @qwen3_30b_a3b_q4_k_m_pp256_016_qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32 + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_017_ggml_mul_mat_id_f16_f16_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<786432xi8> + %expert_table = check.generate.fill value(0) : tensor<33280xi8> + %weight = check.generate.fill value(0) : tensor<113246208xi8> + %output = check.generate.fill value(0) : tensor<2097152xi8> + kernel.launch @ggml_mul_mat_id_f16_f16_wmma[%token_count](%token_count, %input, %expert_table, %weight, %output) : [index](index, tensor<786432xi8>, tensor<33280xi8>, tensor<113246208xi8>, tensor<2097152xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_017_ggml_mul_mat_id_f16_f16_wmma_case> @qwen3_30b_a3b_q4_k_m_pp256_017_ggml_mul_mat_id_f16_f16_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_018_ggml_mul_mat_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<1048576xi8> + %weight = check.generate.fill value(0) : tensor<4718592xi8> + %output = check.generate.fill value(0) : tensor<524288xi8> + kernel.launch @ggml_mul_mat_f32_f32_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<1048576xi8>, tensor<4718592xi8>, tensor<524288xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_018_ggml_mul_mat_f32_f32_wmma_case> @qwen3_30b_a3b_q4_k_m_pp256_018_ggml_mul_mat_f32_f32_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_019_ggml_gather_add_f32_case { + %source_token_count = check.literal value(64) : index + %output_token_count = check.literal value(1) : index + %hidden_size = check.literal value(2048) : index + %attention = check.generate.fill value(0.0) : tensor<64x2048xf32> + %residual = check.generate.fill value(0.0) : tensor<64x2048xf32> + %output_ids = check.generate.fill value(0) : tensor<1xi32> + %output = check.generate.fill value(1.0) : tensor<1x2048xf32> + kernel.launch @ggml_gather_add_f32[%source_token_count, %output_token_count, %hidden_size](%source_token_count, %output_token_count, %hidden_size, %attention, %residual, %output_ids, %output) : [index, index, index](index, index, index, tensor<64x2048xf32>, tensor<64x2048xf32>, tensor<1xi32>, tensor<1x2048xf32>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_019_ggml_gather_add_f32_case> @qwen3_30b_a3b_q4_k_m_pp256_019_ggml_gather_add_f32 + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_020_qwen3_moe_rmsnorm_f32_quantize_q8_1_x4_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(0) : tensor<8192xi8> + %weight = check.generate.fill value(0) : tensor<8192xi8> + %normalized_output = check.generate.fill value(0) : tensor<8192xi8> + %q8_output = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @qwen3_moe_rmsnorm_f32_quantize_q8_1_x4[%token_count](%token_count, %input, %weight, %normalized_output, %q8_output) : [index](index, tensor<8192xi8>, tensor<8192xi8>, tensor<8192xi8>, tensor<2304xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_020_qwen3_moe_rmsnorm_f32_quantize_q8_1_x4_case> @qwen3_30b_a3b_q4_k_m_pp256_020_qwen3_moe_rmsnorm_f32_quantize_q8_1_x4 + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_021_qwen3_moe_router_projection_top8_fused_decode_f32_case { + %token_count = check.literal value(1) : index + %route_id_stride = check.literal value(128) : index + %input = check.generate.fill value(0) : tensor<8192xi8> + %weight = check.generate.fill value(0) : tensor<1048576xi8> + %logits = check.generate.fill value(0) : tensor<512xi8> + %completion_counter = check.generate.fill value(0) : tensor<4xi8> + %route_ids = check.generate.fill value(0) : tensor<512xi8> + %route_weights = check.generate.fill value(0) : tensor<32xi8> + kernel.launch @qwen3_moe_router_projection_top8_fused_decode_f32[%token_count, %route_id_stride](%token_count, %route_id_stride, %input, %weight, %logits, %completion_counter, %route_ids, %route_weights) : [index, index](index, index, tensor<8192xi8>, tensor<1048576xi8>, tensor<512xi8>, tensor<4xi8>, tensor<512xi8>, tensor<32xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_021_qwen3_moe_router_projection_top8_fused_decode_f32_case> @qwen3_30b_a3b_q4_k_m_pp256_021_qwen3_moe_router_projection_top8_fused_decode_f32 + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_022_qwen3_moe_routed_gate_up_swiglu_q4k_q8_case { + %token_count = check.literal value(1) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(768) : index + %q8_input = check.generate.fill value(0) : tensor<2304xi8> + %route_ids = check.generate.fill value(0) : tensor<512xi8> + %gate_weight = check.generate.fill value(0) : tensor<113246208xi8> + %up_weight = check.generate.fill value(0) : tensor<113246208xi8> + %output = check.generate.fill value(0) : tensor<24576xi8> + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8[%token_count, %route_count, %route_stride, %expert_count, %output_size](%token_count, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %output) : [index, index, index, index, index](index, index, index, index, index, tensor<2304xi8>, tensor<512xi8>, tensor<113246208xi8>, tensor<113246208xi8>, tensor<24576xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_022_qwen3_moe_routed_gate_up_swiglu_q4k_q8_case> @qwen3_30b_a3b_q4_k_m_pp256_022_qwen3_moe_routed_gate_up_swiglu_q4k_q8 + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_023_qwen3_moe_routed_down_q6k_f32_wave64_next_q8_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(8) : index + %route_id_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(2048) : index + %input = check.generate.fill value(0) : tensor<24576xi8> + %route_ids = check.generate.fill value(0) : tensor<512xi8> + %route_weights = check.generate.fill value(0) : tensor<32xi8> + %weight = check.generate.fill value(0) : tensor<165150720xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + %norm_weight = check.generate.fill value(0) : tensor<8192xi8> + %completion_counter = check.generate.fill value(0) : tensor<4xi8> + %next_q8_output = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @qwen3_moe_routed_down_q6k_f32_wave64_next_q8[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %input, %route_ids, %route_weights, %weight, %output, %norm_weight, %completion_counter, %next_q8_output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<24576xi8>, tensor<512xi8>, tensor<32xi8>, tensor<165150720xi8>, tensor<8192xi8>, tensor<8192xi8>, tensor<4xi8>, tensor<2304xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_023_qwen3_moe_routed_down_q6k_f32_wave64_next_q8_case> @qwen3_30b_a3b_q4_k_m_pp256_023_qwen3_moe_routed_down_q6k_f32_wave64_next_q8 + +check.case public @qwen3_30b_a3b_q4_k_m_pp256_024_ggml_linear_q6k_q8_1_x4_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(2048) : index + %output_size = check.literal value(151936) : index + %q8_input = check.generate.fill value(0) : tensor<2304xi8> + %weight = check.generate.fill value(0) : tensor<255252480xi8> + %output = check.generate.fill value(0) : tensor<607744xi8> + kernel.launch @ggml_linear_q6k_q8_1_x4[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %q8_input, %weight, %output) : [index, index, index](index, index, index, tensor<2304xi8>, tensor<255252480xi8>, tensor<607744xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp256_024_ggml_linear_q6k_q8_1_x4_case> @qwen3_30b_a3b_q4_k_m_pp256_024_ggml_linear_q6k_q8_1_x4 + +// Scenario: pp512 + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_000_qwen_attention_context_base_capture_case { + %positions = check.generate.fill value(0) : tensor<256xi8> + %control = check.generate.fill value(0) : tensor<4xi8> + kernel.launch @qwen_attention_context_base_capture[](%positions, %control) : [](tensor<256xi8>, tensor<4xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_000_qwen_attention_context_base_capture_case> @qwen3_30b_a3b_q4_k_m_pp512_000_qwen_attention_context_base_capture + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_001_qwen_attention_metadata_case { + %token_count = check.literal value(64) : index + %context_capacity = check.literal value(256) : index + %control = check.generate.fill value(0) : tensor<4xi8> + %positions = check.generate.fill value(0) : tensor<256xi8> + %key_cache_indices = check.generate.fill value(0) : tensor<512xi8> + %value_cache_indices = check.generate.fill value(0) : tensor<512xi8> + %attention_mask = check.generate.fill value(0) : tensor<32768xi8> + kernel.launch @qwen_attention_metadata[%token_count, %context_capacity](%token_count, %context_capacity, %control, %positions, %key_cache_indices, %value_cache_indices, %attention_mask) : [index, index](index, index, tensor<4xi8>, tensor<256xi8>, tensor<512xi8>, tensor<512xi8>, tensor<32768xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_001_qwen_attention_metadata_case> @qwen3_30b_a3b_q4_k_m_pp512_001_qwen_attention_metadata + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_002_ggml_get_rows_f32_case { + %token_count = check.literal value(64) : index + %row_count = check.literal value(151936) : index + %hidden_size = check.literal value(2048) : index + %token_ids = check.generate.fill value(0) : tensor<256xi8> + %weight = check.generate.fill value(0) : tensor<175030272xi8> + %output = check.generate.fill value(0) : tensor<524288xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<256xi8>, tensor<175030272xi8>, tensor<524288xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_002_ggml_get_rows_f32_case> @qwen3_30b_a3b_q4_k_m_pp512_002_ggml_get_rows_f32 + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_003_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(2.0) : tensor<64x2048xf32> + %rhs = check.generate.fill value(3.0) : tensor<2048xf32> + %output = check.generate.fill value(0.0) : tensor<64x2048xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<64x2048xf32>, tensor<2048xf32>, tensor<64x2048xf32>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_003_ggml_rmsnorm_binary_f32_case> @qwen3_30b_a3b_q4_k_m_pp512_003_ggml_rmsnorm_binary_f32 + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_004_ggml_mul_mat_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<524288xi8> + %weight = check.generate.fill value(0) : tensor<4718592xi8> + %output = check.generate.fill value(0) : tensor<1048576xi8> + kernel.launch @ggml_mul_mat_f32_f32_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<524288xi8>, tensor<4718592xi8>, tensor<1048576xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_004_ggml_mul_mat_f32_f32_wmma_case> @qwen3_30b_a3b_q4_k_m_pp512_004_ggml_mul_mat_f32_f32_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_005_ggml_mul_mat_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<524288xi8> + %weight = check.generate.fill value(0) : tensor<860160xi8> + %output = check.generate.fill value(0) : tensor<131072xi8> + kernel.launch @ggml_mul_mat_f32_f32_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<524288xi8>, tensor<860160xi8>, tensor<131072xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_005_ggml_mul_mat_f32_f32_wmma_case> @qwen3_30b_a3b_q4_k_m_pp512_005_ggml_mul_mat_f32_f32_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_006_ggml_mul_mat_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<524288xi8> + %weight = check.generate.fill value(0) : tensor<589824xi8> + %output = check.generate.fill value(0) : tensor<131072xi8> + kernel.launch @ggml_mul_mat_f32_f32_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<524288xi8>, tensor<589824xi8>, tensor<131072xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_006_ggml_mul_mat_f32_f32_wmma_case> @qwen3_30b_a3b_q4_k_m_pp512_006_ggml_mul_mat_f32_f32_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_007_qwen3_moe_attention_postprocess_f32_f16_case { + %token_count = check.literal value(64) : index + %cache_row_count = check.literal value(512) : index + %positions = check.generate.fill value(0) : tensor<256xi8> + %key_cache_indices = check.generate.fill value(0) : tensor<512xi8> + %value_cache_indices = check.generate.fill value(0) : tensor<512xi8> + %query_input = check.generate.fill value(0) : tensor<1048576xi8> + %key_input = check.generate.fill value(0) : tensor<131072xi8> + %value_input = check.generate.fill value(0) : tensor<131072xi8> + %query_norm_weight = check.generate.fill value(0) : tensor<512xi8> + %key_norm_weight = check.generate.fill value(0) : tensor<512xi8> + %inverse_frequencies = check.generate.fill value(0) : tensor<256xi8> + %query_output = check.generate.fill value(0) : tensor<1048576xi8> + %key_cache = check.generate.fill value(0) : tensor<524288xi8> + %value_cache = check.generate.fill value(0) : tensor<524288xi8> + kernel.launch @qwen3_moe_attention_postprocess_f32_f16[%token_count, %cache_row_count](%token_count, %cache_row_count, %positions, %key_cache_indices, %value_cache_indices, %query_input, %key_input, %value_input, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %query_output, %key_cache, %value_cache) : [index, index](index, index, tensor<256xi8>, tensor<512xi8>, tensor<512xi8>, tensor<1048576xi8>, tensor<131072xi8>, tensor<131072xi8>, tensor<512xi8>, tensor<512xi8>, tensor<256xi8>, tensor<1048576xi8>, tensor<524288xi8>, tensor<524288xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_007_qwen3_moe_attention_postprocess_f32_f16_case> @qwen3_30b_a3b_q4_k_m_pp512_007_qwen3_moe_attention_postprocess_f32_f16 + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_008_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(64) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<64x32x128xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x128xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<64x256xf16> + %output = check.generate.fill value(1.0) : tensor<64x32x128xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<64x32x128xf32>, tensor<256x4x128xf16>, tensor<256x4x128xf16>, tensor<64x256xf16>, tensor<64x32x128xf32>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_008_ggml_flash_attention_f32_f16_wmma_case> @qwen3_30b_a3b_q4_k_m_pp512_008_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_009_ggml_mul_mat_add_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<1048576xi8> + %weight = check.generate.fill value(0) : tensor<4718592xi8> + %residual_input = check.generate.fill value(0) : tensor<524288xi8> + %residual_output = check.generate.fill value(0) : tensor<524288xi8> + kernel.launch @ggml_mul_mat_add_f32_f32_wmma[%token_count](%token_count, %input, %weight, %residual_input, %residual_output) : [index](index, tensor<1048576xi8>, tensor<4718592xi8>, tensor<524288xi8>, tensor<524288xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_009_ggml_mul_mat_add_f32_f32_wmma_case> @qwen3_30b_a3b_q4_k_m_pp512_009_ggml_mul_mat_add_f32_f32_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_010_qwen3_moe_router_projection_f32_four_row_wave32_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<524288xi8> + %weight = check.generate.fill value(0) : tensor<1048576xi8> + %output = check.generate.fill value(0) : tensor<32768xi8> + kernel.launch @qwen3_moe_router_projection_f32_four_row_wave32[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<524288xi8>, tensor<1048576xi8>, tensor<32768xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_010_qwen3_moe_router_projection_f32_four_row_wave32_case> @qwen3_30b_a3b_q4_k_m_pp512_010_qwen3_moe_router_projection_f32_four_row_wave32 + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_011_qwen3_moe_router_top8_f32_case { + %token_count = check.literal value(64) : index + %route_id_stride = check.literal value(128) : index + %logits = check.generate.fill value(0) : tensor<32768xi8> + %route_ids = check.generate.fill value(0) : tensor<32768xi8> + %route_weights = check.generate.fill value(0) : tensor<2048xi8> + kernel.launch @qwen3_moe_router_top8_f32[%token_count, %route_id_stride](%token_count, %route_id_stride, %logits, %route_ids, %route_weights) : [index, index](index, index, tensor<32768xi8>, tensor<32768xi8>, tensor<2048xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_011_qwen3_moe_router_top8_f32_case> @qwen3_30b_a3b_q4_k_m_pp512_011_qwen3_moe_router_top8_f32 + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_012_ggml_moe_build_expert_table_case { + %token_count = check.literal value(64) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.fill value(0) : tensor<32768xi8> + %expert_table = check.generate.fill value(0) : tensor<33280xi8> + kernel.launch @ggml_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<32768xi8>, tensor<33280xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_012_ggml_moe_build_expert_table_case> @qwen3_30b_a3b_q4_k_m_pp512_012_ggml_moe_build_expert_table + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_013_ggml_moe_build_expert_partition_table_case { + %token_count = check.literal value(64) : index + %route_count = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %expert_table = check.generate.fill value(0) : tensor<33280xi8> + %partition_table = check.generate.fill value(0) : tensor<580xi8> + kernel.launch @ggml_moe_build_expert_partition_table[%token_count, %route_count, %expert_count](%token_count, %route_count, %expert_count, %expert_table, %partition_table) : [index, index, index](index, index, index, tensor<33280xi8>, tensor<580xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_013_ggml_moe_build_expert_partition_table_case> @qwen3_30b_a3b_q4_k_m_pp512_013_ggml_moe_build_expert_partition_table + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_014_qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<524288xi8> + %expert_table = check.generate.fill value(0) : tensor<33280xi8> + %partition_table = check.generate.fill value(0) : tensor<580xi8> + %gate_weight = check.generate.fill value(0) : tensor<113246208xi8> + %up_weight = check.generate.fill value(0) : tensor<113246208xi8> + %output = check.generate.fill value(0) : tensor<786432xi8> + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma[%token_count](%token_count, %input, %expert_table, %partition_table, %gate_weight, %up_weight, %output) : [index](index, tensor<524288xi8>, tensor<33280xi8>, tensor<580xi8>, tensor<113246208xi8>, tensor<113246208xi8>, tensor<786432xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_014_qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_case> @qwen3_30b_a3b_q4_k_m_pp512_014_qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_015_ggml_mul_mat_id_f16_f16_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<786432xi8> + %expert_table = check.generate.fill value(0) : tensor<33280xi8> + %weight = check.generate.fill value(0) : tensor<165150720xi8> + %output = check.generate.fill value(0) : tensor<2097152xi8> + kernel.launch @ggml_mul_mat_id_f16_f16_wmma[%token_count](%token_count, %input, %expert_table, %weight, %output) : [index](index, tensor<786432xi8>, tensor<33280xi8>, tensor<165150720xi8>, tensor<2097152xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_015_ggml_mul_mat_id_f16_f16_wmma_case> @qwen3_30b_a3b_q4_k_m_pp512_015_ggml_mul_mat_id_f16_f16_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_016_qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32_case { + %token_count = check.literal value(64) : index + %route_weights = check.generate.fill value(0) : tensor<2048xi8> + %routed_output = check.generate.fill value(0) : tensor<2097152xi8> + %hidden_state = check.generate.fill value(0) : tensor<524288xi8> + %next_norm_weight = check.generate.fill value(0) : tensor<8192xi8> + %next_projection_input = check.generate.fill value(0) : tensor<524288xi8> + kernel.launch @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32[%token_count](%token_count, %route_weights, %routed_output, %hidden_state, %next_norm_weight, %next_projection_input) : [index](index, tensor<2048xi8>, tensor<2097152xi8>, tensor<524288xi8>, tensor<8192xi8>, tensor<524288xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_016_qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32_case> @qwen3_30b_a3b_q4_k_m_pp512_016_qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32 + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_017_ggml_mul_mat_id_f16_f16_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<786432xi8> + %expert_table = check.generate.fill value(0) : tensor<33280xi8> + %weight = check.generate.fill value(0) : tensor<113246208xi8> + %output = check.generate.fill value(0) : tensor<2097152xi8> + kernel.launch @ggml_mul_mat_id_f16_f16_wmma[%token_count](%token_count, %input, %expert_table, %weight, %output) : [index](index, tensor<786432xi8>, tensor<33280xi8>, tensor<113246208xi8>, tensor<2097152xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_017_ggml_mul_mat_id_f16_f16_wmma_case> @qwen3_30b_a3b_q4_k_m_pp512_017_ggml_mul_mat_id_f16_f16_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_018_ggml_mul_mat_f32_f32_wmma_case { + %token_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<1048576xi8> + %weight = check.generate.fill value(0) : tensor<4718592xi8> + %output = check.generate.fill value(0) : tensor<524288xi8> + kernel.launch @ggml_mul_mat_f32_f32_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<1048576xi8>, tensor<4718592xi8>, tensor<524288xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_018_ggml_mul_mat_f32_f32_wmma_case> @qwen3_30b_a3b_q4_k_m_pp512_018_ggml_mul_mat_f32_f32_wmma + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_019_ggml_gather_add_f32_case { + %source_token_count = check.literal value(64) : index + %output_token_count = check.literal value(1) : index + %hidden_size = check.literal value(2048) : index + %attention = check.generate.fill value(0.0) : tensor<64x2048xf32> + %residual = check.generate.fill value(0.0) : tensor<64x2048xf32> + %output_ids = check.generate.fill value(0) : tensor<1xi32> + %output = check.generate.fill value(1.0) : tensor<1x2048xf32> + kernel.launch @ggml_gather_add_f32[%source_token_count, %output_token_count, %hidden_size](%source_token_count, %output_token_count, %hidden_size, %attention, %residual, %output_ids, %output) : [index, index, index](index, index, index, tensor<64x2048xf32>, tensor<64x2048xf32>, tensor<1xi32>, tensor<1x2048xf32>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_019_ggml_gather_add_f32_case> @qwen3_30b_a3b_q4_k_m_pp512_019_ggml_gather_add_f32 + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_020_qwen3_moe_rmsnorm_f32_quantize_q8_1_x4_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(0) : tensor<8192xi8> + %weight = check.generate.fill value(0) : tensor<8192xi8> + %normalized_output = check.generate.fill value(0) : tensor<8192xi8> + %q8_output = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @qwen3_moe_rmsnorm_f32_quantize_q8_1_x4[%token_count](%token_count, %input, %weight, %normalized_output, %q8_output) : [index](index, tensor<8192xi8>, tensor<8192xi8>, tensor<8192xi8>, tensor<2304xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_020_qwen3_moe_rmsnorm_f32_quantize_q8_1_x4_case> @qwen3_30b_a3b_q4_k_m_pp512_020_qwen3_moe_rmsnorm_f32_quantize_q8_1_x4 + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_021_qwen3_moe_router_projection_top8_fused_decode_f32_case { + %token_count = check.literal value(1) : index + %route_id_stride = check.literal value(128) : index + %input = check.generate.fill value(0) : tensor<8192xi8> + %weight = check.generate.fill value(0) : tensor<1048576xi8> + %logits = check.generate.fill value(0) : tensor<512xi8> + %completion_counter = check.generate.fill value(0) : tensor<4xi8> + %route_ids = check.generate.fill value(0) : tensor<512xi8> + %route_weights = check.generate.fill value(0) : tensor<32xi8> + kernel.launch @qwen3_moe_router_projection_top8_fused_decode_f32[%token_count, %route_id_stride](%token_count, %route_id_stride, %input, %weight, %logits, %completion_counter, %route_ids, %route_weights) : [index, index](index, index, tensor<8192xi8>, tensor<1048576xi8>, tensor<512xi8>, tensor<4xi8>, tensor<512xi8>, tensor<32xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_021_qwen3_moe_router_projection_top8_fused_decode_f32_case> @qwen3_30b_a3b_q4_k_m_pp512_021_qwen3_moe_router_projection_top8_fused_decode_f32 + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_022_qwen3_moe_routed_gate_up_swiglu_q4k_q8_case { + %token_count = check.literal value(1) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(768) : index + %q8_input = check.generate.fill value(0) : tensor<2304xi8> + %route_ids = check.generate.fill value(0) : tensor<512xi8> + %gate_weight = check.generate.fill value(0) : tensor<113246208xi8> + %up_weight = check.generate.fill value(0) : tensor<113246208xi8> + %output = check.generate.fill value(0) : tensor<24576xi8> + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8[%token_count, %route_count, %route_stride, %expert_count, %output_size](%token_count, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %output) : [index, index, index, index, index](index, index, index, index, index, tensor<2304xi8>, tensor<512xi8>, tensor<113246208xi8>, tensor<113246208xi8>, tensor<24576xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_022_qwen3_moe_routed_gate_up_swiglu_q4k_q8_case> @qwen3_30b_a3b_q4_k_m_pp512_022_qwen3_moe_routed_gate_up_swiglu_q4k_q8 + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_023_qwen3_moe_routed_down_q6k_f32_wave64_next_q8_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(8) : index + %route_id_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(2048) : index + %input = check.generate.fill value(0) : tensor<24576xi8> + %route_ids = check.generate.fill value(0) : tensor<512xi8> + %route_weights = check.generate.fill value(0) : tensor<32xi8> + %weight = check.generate.fill value(0) : tensor<165150720xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + %norm_weight = check.generate.fill value(0) : tensor<8192xi8> + %completion_counter = check.generate.fill value(0) : tensor<4xi8> + %next_q8_output = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @qwen3_moe_routed_down_q6k_f32_wave64_next_q8[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %input, %route_ids, %route_weights, %weight, %output, %norm_weight, %completion_counter, %next_q8_output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<24576xi8>, tensor<512xi8>, tensor<32xi8>, tensor<165150720xi8>, tensor<8192xi8>, tensor<8192xi8>, tensor<4xi8>, tensor<2304xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_023_qwen3_moe_routed_down_q6k_f32_wave64_next_q8_case> @qwen3_30b_a3b_q4_k_m_pp512_023_qwen3_moe_routed_down_q6k_f32_wave64_next_q8 + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_024_ggml_linear_q6k_q8_1_x4_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(2048) : index + %output_size = check.literal value(151936) : index + %q8_input = check.generate.fill value(0) : tensor<2304xi8> + %weight = check.generate.fill value(0) : tensor<255252480xi8> + %output = check.generate.fill value(0) : tensor<607744xi8> + kernel.launch @ggml_linear_q6k_q8_1_x4[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %q8_input, %weight, %output) : [index, index, index](index, index, index, tensor<2304xi8>, tensor<255252480xi8>, tensor<607744xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_024_ggml_linear_q6k_q8_1_x4_case> @qwen3_30b_a3b_q4_k_m_pp512_024_ggml_linear_q6k_q8_1_x4 + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_025_qwen_attention_metadata_case { + %token_count = check.literal value(64) : index + %context_capacity = check.literal value(512) : index + %control = check.generate.fill value(0) : tensor<4xi8> + %positions = check.generate.fill value(0) : tensor<256xi8> + %key_cache_indices = check.generate.fill value(0) : tensor<512xi8> + %value_cache_indices = check.generate.fill value(0) : tensor<512xi8> + %attention_mask = check.generate.fill value(0) : tensor<65536xi8> + kernel.launch @qwen_attention_metadata[%token_count, %context_capacity](%token_count, %context_capacity, %control, %positions, %key_cache_indices, %value_cache_indices, %attention_mask) : [index, index](index, index, tensor<4xi8>, tensor<256xi8>, tensor<512xi8>, tensor<512xi8>, tensor<65536xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_025_qwen_attention_metadata_case> @qwen3_30b_a3b_q4_k_m_pp512_025_qwen_attention_metadata + +check.case public @qwen3_30b_a3b_q4_k_m_pp512_026_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(64) : index + %key_value_token_count = check.literal value(512) : index + %query = check.generate.fill value(0.0) : tensor<64x32x128xf32> + %key = check.generate.fill value(0.0) : tensor<512x4x128xf16> + %value = check.generate.fill value(0.0) : tensor<512x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<64x512xf16> + %output = check.generate.fill value(1.0) : tensor<64x32x128xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<64x32x128xf32>, tensor<512x4x128xf16>, tensor<512x4x128xf16>, tensor<64x512xf16>, tensor<64x32x128xf32>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_pp512_026_ggml_flash_attention_f32_f16_wmma_case> @qwen3_30b_a3b_q4_k_m_pp512_026_ggml_flash_attention_f32_f16_wmma + +// Scenario: tg8 + +check.case public @qwen3_30b_a3b_q4_k_m_tg8_000_qwen_attention_context_base_capture_case { + %positions = check.generate.fill value(0) : tensor<4xi8> + %control = check.generate.fill value(0) : tensor<4xi8> + kernel.launch @qwen_attention_context_base_capture[](%positions, %control) : [](tensor<4xi8>, tensor<4xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_tg8_000_qwen_attention_context_base_capture_case> @qwen3_30b_a3b_q4_k_m_tg8_000_qwen_attention_context_base_capture + +check.case public @qwen3_30b_a3b_q4_k_m_tg8_001_qwen_attention_metadata_case { + %token_count = check.literal value(1) : index + %context_capacity = check.literal value(256) : index + %control = check.generate.fill value(0) : tensor<4xi8> + %positions = check.generate.fill value(0) : tensor<4xi8> + %key_cache_indices = check.generate.fill value(0) : tensor<8xi8> + %value_cache_indices = check.generate.fill value(0) : tensor<8xi8> + %attention_mask = check.generate.fill value(0) : tensor<512xi8> + kernel.launch @qwen_attention_metadata[%token_count, %context_capacity](%token_count, %context_capacity, %control, %positions, %key_cache_indices, %value_cache_indices, %attention_mask) : [index, index](index, index, tensor<4xi8>, tensor<4xi8>, tensor<8xi8>, tensor<8xi8>, tensor<512xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_tg8_001_qwen_attention_metadata_case> @qwen3_30b_a3b_q4_k_m_tg8_001_qwen_attention_metadata + +check.case public @qwen3_30b_a3b_q4_k_m_tg8_002_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(151936) : index + %hidden_size = check.literal value(2048) : index + %token_ids = check.generate.fill value(0) : tensor<4xi8> + %weight = check.generate.fill value(0) : tensor<175030272xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<4xi8>, tensor<175030272xi8>, tensor<8192xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_tg8_002_ggml_get_rows_f32_case> @qwen3_30b_a3b_q4_k_m_tg8_002_ggml_get_rows_f32 + +check.case public @qwen3_30b_a3b_q4_k_m_tg8_003_qwen3_moe_rmsnorm_f32_quantize_q8_1_x4_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(0) : tensor<8192xi8> + %weight = check.generate.fill value(0) : tensor<8192xi8> + %normalized_output = check.generate.fill value(0) : tensor<8192xi8> + %q8_output = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @qwen3_moe_rmsnorm_f32_quantize_q8_1_x4[%token_count](%token_count, %input, %weight, %normalized_output, %q8_output) : [index](index, tensor<8192xi8>, tensor<8192xi8>, tensor<8192xi8>, tensor<2304xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_tg8_003_qwen3_moe_rmsnorm_f32_quantize_q8_1_x4_case> @qwen3_30b_a3b_q4_k_m_tg8_003_qwen3_moe_rmsnorm_f32_quantize_q8_1_x4 + +check.case public @qwen3_30b_a3b_q4_k_m_tg8_004_qwen3_moe_attention_qkv_postprocess_fused_decode_case { + %token_count = check.literal value(1) : index + %cache_row_count = check.literal value(256) : index + %q8_input = check.generate.fill value(0) : tensor<2304xi8> + %query_weight = check.generate.fill value(0) : tensor<4718592xi8> + %key_weight = check.generate.fill value(0) : tensor<589824xi8> + %value_weight = check.generate.fill value(0) : tensor<860160xi8> + %positions = check.generate.fill value(0) : tensor<4xi8> + %key_cache_indices = check.generate.fill value(0) : tensor<8xi8> + %value_cache_indices = check.generate.fill value(0) : tensor<8xi8> + %query_output_raw = check.generate.fill value(0) : tensor<16384xi8> + %key_output_raw = check.generate.fill value(0) : tensor<2048xi8> + %value_output_raw = check.generate.fill value(0) : tensor<2048xi8> + %query_norm_weight = check.generate.fill value(0) : tensor<512xi8> + %key_norm_weight = check.generate.fill value(0) : tensor<512xi8> + %inverse_frequencies = check.generate.fill value(0) : tensor<256xi8> + %query_output = check.generate.fill value(0) : tensor<16384xi8> + %key_cache = check.generate.fill value(0) : tensor<262144xi8> + %value_cache = check.generate.fill value(0) : tensor<262144xi8> + %completion_counters = check.generate.fill value(0) : tensor<160xi8> + kernel.launch @qwen3_moe_attention_qkv_postprocess_fused_decode[%token_count, %cache_row_count](%token_count, %cache_row_count, %q8_input, %query_weight, %key_weight, %value_weight, %positions, %key_cache_indices, %value_cache_indices, %query_output_raw, %key_output_raw, %value_output_raw, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %query_output, %key_cache, %value_cache, %completion_counters) : [index, index](index, index, tensor<2304xi8>, tensor<4718592xi8>, tensor<589824xi8>, tensor<860160xi8>, tensor<4xi8>, tensor<8xi8>, tensor<8xi8>, tensor<16384xi8>, tensor<2048xi8>, tensor<2048xi8>, tensor<512xi8>, tensor<512xi8>, tensor<256xi8>, tensor<16384xi8>, tensor<262144xi8>, tensor<262144xi8>, tensor<160xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_tg8_004_qwen3_moe_attention_qkv_postprocess_fused_decode_case> @qwen3_30b_a3b_q4_k_m_tg8_004_qwen3_moe_attention_qkv_postprocess_fused_decode + +check.case public @qwen3_30b_a3b_q4_k_m_tg8_005_ggml_flash_attention_decode_split_f32_f16_wmma_next_q8_case { + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<32x128xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x128xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<256xf16> + %partial_max = check.generate.fill value(0.0) : tensor<4x4x16xf32> + %partial_sum = check.generate.fill value(0.0) : tensor<4x4x16xf32> + %partial_output = check.generate.fill value(0.0) : tensor<4x4x16x128xf16> + %completion_counter = check.generate.fill value(0) : tensor<4xi32> + %output = check.generate.fill value(1.0) : tensor<32x128xf32> + %next_q8_output = check.generate.fill value(0) : tensor<4608xi8> + kernel.launch @ggml_flash_attention_decode_split_f32_f16_wmma_next_q8[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output, %completion_counter, %output, %next_q8_output) : [index](index, tensor<32x128xf32>, tensor<256x4x128xf16>, tensor<256x4x128xf16>, tensor<256xf16>, tensor<4x4x16xf32>, tensor<4x4x16xf32>, tensor<4x4x16x128xf16>, tensor<4xi32>, tensor<32x128xf32>, tensor<4608xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_tg8_005_ggml_flash_attention_decode_split_f32_f16_wmma_next_q8_case> @qwen3_30b_a3b_q4_k_m_tg8_005_ggml_flash_attention_decode_split_f32_f16_wmma_next_q8 + +check.case public @qwen3_30b_a3b_q4_k_m_tg8_006_qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8_case { + %token_count = check.literal value(1) : index + %q8_input = check.generate.fill value(0) : tensor<4608xi8> + %weight = check.generate.fill value(0) : tensor<4718592xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + %norm_weight = check.generate.fill value(0) : tensor<8192xi8> + %normalized_output = check.generate.fill value(0) : tensor<8192xi8> + %completion_counter = check.generate.fill value(0) : tensor<4xi8> + %next_q8_output = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8[%token_count](%token_count, %q8_input, %weight, %output, %norm_weight, %normalized_output, %completion_counter, %next_q8_output) : [index](index, tensor<4608xi8>, tensor<4718592xi8>, tensor<8192xi8>, tensor<8192xi8>, tensor<8192xi8>, tensor<4xi8>, tensor<2304xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_tg8_006_qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8_case> @qwen3_30b_a3b_q4_k_m_tg8_006_qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8 + +check.case public @qwen3_30b_a3b_q4_k_m_tg8_007_qwen3_moe_router_projection_top8_fused_decode_f32_case { + %token_count = check.literal value(1) : index + %route_id_stride = check.literal value(128) : index + %input = check.generate.fill value(0) : tensor<8192xi8> + %weight = check.generate.fill value(0) : tensor<1048576xi8> + %logits = check.generate.fill value(0) : tensor<512xi8> + %completion_counter = check.generate.fill value(0) : tensor<4xi8> + %route_ids = check.generate.fill value(0) : tensor<512xi8> + %route_weights = check.generate.fill value(0) : tensor<32xi8> + kernel.launch @qwen3_moe_router_projection_top8_fused_decode_f32[%token_count, %route_id_stride](%token_count, %route_id_stride, %input, %weight, %logits, %completion_counter, %route_ids, %route_weights) : [index, index](index, index, tensor<8192xi8>, tensor<1048576xi8>, tensor<512xi8>, tensor<4xi8>, tensor<512xi8>, tensor<32xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_tg8_007_qwen3_moe_router_projection_top8_fused_decode_f32_case> @qwen3_30b_a3b_q4_k_m_tg8_007_qwen3_moe_router_projection_top8_fused_decode_f32 + +check.case public @qwen3_30b_a3b_q4_k_m_tg8_008_qwen3_moe_routed_gate_up_swiglu_q4k_q8_case { + %token_count = check.literal value(1) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(768) : index + %q8_input = check.generate.fill value(0) : tensor<2304xi8> + %route_ids = check.generate.fill value(0) : tensor<512xi8> + %gate_weight = check.generate.fill value(0) : tensor<113246208xi8> + %up_weight = check.generate.fill value(0) : tensor<113246208xi8> + %output = check.generate.fill value(0) : tensor<24576xi8> + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8[%token_count, %route_count, %route_stride, %expert_count, %output_size](%token_count, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %output) : [index, index, index, index, index](index, index, index, index, index, tensor<2304xi8>, tensor<512xi8>, tensor<113246208xi8>, tensor<113246208xi8>, tensor<24576xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_tg8_008_qwen3_moe_routed_gate_up_swiglu_q4k_q8_case> @qwen3_30b_a3b_q4_k_m_tg8_008_qwen3_moe_routed_gate_up_swiglu_q4k_q8 + +check.case public @qwen3_30b_a3b_q4_k_m_tg8_009_qwen3_moe_routed_down_q6k_f32_wave64_next_q8_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(8) : index + %route_id_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(2048) : index + %input = check.generate.fill value(0) : tensor<24576xi8> + %route_ids = check.generate.fill value(0) : tensor<512xi8> + %route_weights = check.generate.fill value(0) : tensor<32xi8> + %weight = check.generate.fill value(0) : tensor<165150720xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + %norm_weight = check.generate.fill value(0) : tensor<8192xi8> + %completion_counter = check.generate.fill value(0) : tensor<4xi8> + %next_q8_output = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @qwen3_moe_routed_down_q6k_f32_wave64_next_q8[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %input, %route_ids, %route_weights, %weight, %output, %norm_weight, %completion_counter, %next_q8_output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<24576xi8>, tensor<512xi8>, tensor<32xi8>, tensor<165150720xi8>, tensor<8192xi8>, tensor<8192xi8>, tensor<4xi8>, tensor<2304xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_tg8_009_qwen3_moe_routed_down_q6k_f32_wave64_next_q8_case> @qwen3_30b_a3b_q4_k_m_tg8_009_qwen3_moe_routed_down_q6k_f32_wave64_next_q8 + +check.case public @qwen3_30b_a3b_q4_k_m_tg8_010_qwen3_moe_attention_qkv_postprocess_fused_decode_case { + %token_count = check.literal value(1) : index + %cache_row_count = check.literal value(256) : index + %q8_input = check.generate.fill value(0) : tensor<2304xi8> + %query_weight = check.generate.fill value(0) : tensor<4718592xi8> + %key_weight = check.generate.fill value(0) : tensor<589824xi8> + %value_weight = check.generate.fill value(0) : tensor<589824xi8> + %positions = check.generate.fill value(0) : tensor<4xi8> + %key_cache_indices = check.generate.fill value(0) : tensor<8xi8> + %value_cache_indices = check.generate.fill value(0) : tensor<8xi8> + %query_output_raw = check.generate.fill value(0) : tensor<16384xi8> + %key_output_raw = check.generate.fill value(0) : tensor<2048xi8> + %value_output_raw = check.generate.fill value(0) : tensor<2048xi8> + %query_norm_weight = check.generate.fill value(0) : tensor<512xi8> + %key_norm_weight = check.generate.fill value(0) : tensor<512xi8> + %inverse_frequencies = check.generate.fill value(0) : tensor<256xi8> + %query_output = check.generate.fill value(0) : tensor<16384xi8> + %key_cache = check.generate.fill value(0) : tensor<262144xi8> + %value_cache = check.generate.fill value(0) : tensor<262144xi8> + %completion_counters = check.generate.fill value(0) : tensor<160xi8> + kernel.launch @qwen3_moe_attention_qkv_postprocess_fused_decode[%token_count, %cache_row_count](%token_count, %cache_row_count, %q8_input, %query_weight, %key_weight, %value_weight, %positions, %key_cache_indices, %value_cache_indices, %query_output_raw, %key_output_raw, %value_output_raw, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %query_output, %key_cache, %value_cache, %completion_counters) : [index, index](index, index, tensor<2304xi8>, tensor<4718592xi8>, tensor<589824xi8>, tensor<589824xi8>, tensor<4xi8>, tensor<8xi8>, tensor<8xi8>, tensor<16384xi8>, tensor<2048xi8>, tensor<2048xi8>, tensor<512xi8>, tensor<512xi8>, tensor<256xi8>, tensor<16384xi8>, tensor<262144xi8>, tensor<262144xi8>, tensor<160xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_tg8_010_qwen3_moe_attention_qkv_postprocess_fused_decode_case> @qwen3_30b_a3b_q4_k_m_tg8_010_qwen3_moe_attention_qkv_postprocess_fused_decode + +check.case public @qwen3_30b_a3b_q4_k_m_tg8_011_qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8_case { + %token_count = check.literal value(1) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(768) : index + %q8_input = check.generate.fill value(0) : tensor<2304xi8> + %route_ids = check.generate.fill value(0) : tensor<512xi8> + %gate_weight = check.generate.fill value(0) : tensor<113246208xi8> + %up_weight = check.generate.fill value(0) : tensor<113246208xi8> + %output = check.generate.fill value(0) : tensor<24576xi8> + %completion_counters = check.generate.fill value(0) : tensor<192xi8> + %next_q8_output = check.generate.fill value(0) : tensor<6912xi8> + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8[%token_count, %route_count, %route_stride, %expert_count, %output_size](%token_count, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %output, %completion_counters, %next_q8_output) : [index, index, index, index, index](index, index, index, index, index, tensor<2304xi8>, tensor<512xi8>, tensor<113246208xi8>, tensor<113246208xi8>, tensor<24576xi8>, tensor<192xi8>, tensor<6912xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_tg8_011_qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8_case> @qwen3_30b_a3b_q4_k_m_tg8_011_qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8 + +check.case public @qwen3_30b_a3b_q4_k_m_tg8_012_qwen3_moe_routed_down_q4k_q8_1_x4_next_q8_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(8) : index + %route_id_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(2048) : index + %q8_input = check.generate.fill value(0) : tensor<6912xi8> + %route_ids = check.generate.fill value(0) : tensor<512xi8> + %route_weights = check.generate.fill value(0) : tensor<32xi8> + %weight = check.generate.fill value(0) : tensor<113246208xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + %norm_weight = check.generate.fill value(0) : tensor<8192xi8> + %completion_counter = check.generate.fill value(0) : tensor<4xi8> + %next_q8_output = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @qwen3_moe_routed_down_q4k_q8_1_x4_next_q8[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %output, %norm_weight, %completion_counter, %next_q8_output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<6912xi8>, tensor<512xi8>, tensor<32xi8>, tensor<113246208xi8>, tensor<8192xi8>, tensor<8192xi8>, tensor<4xi8>, tensor<2304xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_tg8_012_qwen3_moe_routed_down_q4k_q8_1_x4_next_q8_case> @qwen3_30b_a3b_q4_k_m_tg8_012_qwen3_moe_routed_down_q4k_q8_1_x4_next_q8 + +check.case public @qwen3_30b_a3b_q4_k_m_tg8_013_ggml_linear_q6k_q8_1_x4_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(2048) : index + %output_size = check.literal value(151936) : index + %q8_input = check.generate.fill value(0) : tensor<2304xi8> + %weight = check.generate.fill value(0) : tensor<255252480xi8> + %output = check.generate.fill value(0) : tensor<607744xi8> + kernel.launch @ggml_linear_q6k_q8_1_x4[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %q8_input, %weight, %output) : [index, index, index](index, index, index, tensor<2304xi8>, tensor<255252480xi8>, tensor<607744xi8>) + check.return +} + +check.benchmark<@qwen3_30b_a3b_q4_k_m_tg8_013_ggml_linear_q6k_q8_1_x4_case> @qwen3_30b_a3b_q4_k_m_tg8_013_ggml_linear_q6k_q8_1_x4 diff --git a/ggml/src/ggml-hrx/benchmarks/loom/qwen3_30b_a3b_q4_k_m.pp256.json b/ggml/src/ggml-hrx/benchmarks/loom/qwen3_30b_a3b_q4_k_m.pp256.json new file mode 100644 index 000000000000..90a52c3f134c --- /dev/null +++ b/ggml/src/ggml-hrx/benchmarks/loom/qwen3_30b_a3b_q4_k_m.pp256.json @@ -0,0 +1,1073 @@ +{ + "command_count": 674, + "dispatch_count": 25, + "dispatches": [ + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_000_qwen_attention_context_base_capture", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/qwen", + "count": 1, + "integer_parameters": {}, + "kernel": "qwen:qwen_attention_context_base_capture", + "library_sources": [ + "motifs/mul_mat_id_f16_f16_wmma_core.loom" + ], + "primary_sources": [ + "attention_state_initialize.loom" + ], + "sources": [ + "attention_state_initialize.loom" + ], + "symbol": "qwen_attention_context_base_capture", + "workload_parameters": [] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_001_qwen_attention_metadata", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 1, + "integer_parameters": { + "context_capacity": 256, + "token_count": 64 + }, + "kernel": "qwen3_moe:qwen_attention_metadata", + "library_sources": [], + "primary_sources": [ + "../qwen/attention_metadata.loom" + ], + "sources": [ + "../qwen/attention_metadata.loom" + ], + "symbol": "qwen_attention_metadata", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "context_capacity", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_002_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "2048", + "ggml.get_rows_f32.token_capacity": "64", + "ggml.get_rows_f32.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "hidden_size": 2048, + "row_count": 151936, + "token_count": 64 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_003_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "2048", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_004_ggml_mul_mat_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat.input_size": "2048", + "ggml.mul_mat.output_accumulation": "0", + "ggml.mul_mat.output_size": "4096", + "ggml.mul_mat.output_unary_op": "23", + "ggml.mul_mat.weight_format": "4", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_005_ggml_mul_mat_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat.input_size": "2048", + "ggml.mul_mat.output_accumulation": "0", + "ggml.mul_mat.output_size": "512", + "ggml.mul_mat.output_unary_op": "23", + "ggml.mul_mat.weight_format": "6", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 24, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_006_ggml_mul_mat_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat.input_size": "2048", + "ggml.mul_mat.output_accumulation": "0", + "ggml.mul_mat.output_size": "512", + "ggml.mul_mat.output_unary_op": "23", + "ggml.mul_mat.weight_format": "4", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 72, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_007_qwen3_moe_attention_postprocess_f32_f16", + "compile_parameters": { + "qwen3_moe.attention.head_size": "128", + "qwen3_moe.attention.key_value_size": "512", + "qwen3_moe.attention.query_size": "4096", + "qwen3_moe.model.rms_epsilon": "0.000001", + "qwen3_moe.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 48, + "integer_parameters": { + "cache_row_count": 256, + "token_count": 64 + }, + "kernel": "qwen3_moe:qwen3_moe_attention_postprocess_f32_f16", + "library_sources": [ + "qwen3_moe/model_config.loom" + ], + "primary_sources": [ + "qwen3_moe/attention_postprocess_f32_f16.loom" + ], + "sources": [ + "qwen3_moe/attention_postprocess_f32_f16.loom", + "qwen3_moe/model_config.loom" + ], + "symbol": "qwen3_moe_attention_postprocess_f32_f16", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_008_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.head_size": "128", + "ggml.flash_attention.key_value_head_count": "4", + "ggml.flash_attention.query_head_count": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "key_value_token_count": 256, + "query_token_count": 64 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_009_ggml_mul_mat_add_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat_postops.input_size": "4096", + "ggml.mul_mat_postops.output_size": "2048", + "ggml.mul_mat_postops.weight_format": "4", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 47, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_add_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_add_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_add_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_add_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_010_qwen3_moe_router_projection_f32_four_row_wave32", + "compile_parameters": { + "qwen3_moe.model.hidden_size": "2048", + "qwen3_moe.router.expert_count": "128", + "qwen3_moe.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 47, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "qwen3_moe:qwen3_moe_router_projection_f32_four_row_wave32", + "library_sources": [ + "qwen3_moe/model_config.loom" + ], + "primary_sources": [ + "qwen3_moe/router_projection_f32.loom" + ], + "sources": [ + "qwen3_moe/router_projection_f32.loom", + "qwen3_moe/model_config.loom" + ], + "symbol": "qwen3_moe_router_projection_f32_four_row_wave32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_011_qwen3_moe_router_top8_f32", + "compile_parameters": { + "qwen3_moe.router.expert_count": "128", + "qwen3_moe.router.route_count": "8", + "qwen3_moe.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 47, + "integer_parameters": { + "route_id_stride": 128, + "token_count": 64 + }, + "kernel": "qwen3_moe:qwen3_moe_router_top8_f32", + "library_sources": [ + "qwen3_moe/model_config.loom" + ], + "primary_sources": [ + "qwen3_moe/router_top8_f32.loom" + ], + "sources": [ + "qwen3_moe/router_top8_f32.loom", + "qwen3_moe/model_config.loom" + ], + "symbol": "qwen3_moe_router_top8_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_012_ggml_moe_build_expert_table", + "compile_parameters": { + "ggml.moe_routing.descriptor_expert_mask": "127", + "ggml.moe_routing.descriptor_partition_shift": "7", + "ggml.moe_routing.descriptor_row_count_shift": "13", + "ggml.moe_routing.expert_count": "128", + "ggml.moe_routing.partition_workgroup_size": "128", + "ggml.moe_routing.route_count": "8", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 47, + "integer_parameters": { + "expert_count": 128, + "route_count": 8, + "route_stride": 128, + "token_count": 64 + }, + "kernel": "loom_libs:ggml_moe_build_expert_table", + "library_sources": [], + "primary_sources": [ + "ops/moe_routing_tables.loom" + ], + "sources": [ + "ops/moe_routing_tables.loom" + ], + "symbol": "ggml_moe_build_expert_table", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_013_ggml_moe_build_expert_partition_table", + "compile_parameters": { + "ggml.moe_routing.descriptor_expert_mask": "127", + "ggml.moe_routing.descriptor_partition_shift": "7", + "ggml.moe_routing.descriptor_row_count_shift": "13", + "ggml.moe_routing.expert_count": "128", + "ggml.moe_routing.partition_workgroup_size": "128", + "ggml.moe_routing.route_count": "8", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 47, + "integer_parameters": { + "expert_count": 128, + "route_count": 8, + "token_count": 64 + }, + "kernel": "loom_libs:ggml_moe_build_expert_partition_table", + "library_sources": [], + "primary_sources": [ + "ops/moe_routing_tables.loom" + ], + "sources": [ + "ops/moe_routing_tables.loom" + ], + "symbol": "ggml_moe_build_expert_partition_table", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_014_qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma", + "compile_parameters": { + "qwen3_moe.routed_gate_up.expert_count": "128", + "qwen3_moe.routed_gate_up.input_size": "2048", + "qwen3_moe.routed_gate_up.output_size": "768", + "qwen3_moe.routed_gate_up.route_count": "8", + "qwen3_moe.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 47, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "qwen3_moe:qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma", + "library_sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "qwen3_moe/routed_linear_q4k_f16_wmma.loom" + ], + "sources": [ + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_015_ggml_mul_mat_id_f16_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_id_f16_f16.expert_count": "128", + "ggml.mul_mat_id_f16_f16.input_size": "768", + "ggml.mul_mat_id_f16_f16.output_size": "2048", + "ggml.mul_mat_id_f16_f16.route_count": "8", + "ggml.mul_mat_id_f16_f16.weight_format": "6", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 23, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_id_f16_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_id_f16_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_id_f16_f16_wmma.loom", + "motifs/mul_mat_id_f16_f16_wmma_core.loom" + ], + "symbol": "ggml_mul_mat_id_f16_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_016_qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32", + "compile_parameters": { + "qwen3_moe.model.hidden_size": "2048", + "qwen3_moe.model.rms_epsilon": "0.000001", + "qwen3_moe.routed_down.expert_count": "128", + "qwen3_moe.routed_down.input_size": "768", + "qwen3_moe.routed_down.output_size": "2048", + "qwen3_moe.routed_down.route_count": "8", + "qwen3_moe.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 47, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "qwen3_moe:qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32", + "library_sources": [ + "qwen3_moe/routed_down_quantized_f16_wmma.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/routed_down_q6k.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_f32.loom" + ], + "sources": [ + "qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_f32.loom", + "qwen3_moe/routed_down_quantized_f16_wmma.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/routed_down_q6k.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_017_ggml_mul_mat_id_f16_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_id_f16_f16.expert_count": "128", + "ggml.mul_mat_id_f16_f16.input_size": "768", + "ggml.mul_mat_id_f16_f16.output_size": "2048", + "ggml.mul_mat_id_f16_f16.route_count": "8", + "ggml.mul_mat_id_f16_f16.weight_format": "4", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 24, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_id_f16_f16_wmma", + "library_sources": [ + "motifs/mul_mat_id_f16_f16_wmma_core.loom" + ], + "primary_sources": [ + "ops/mul_mat_id_f16_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_id_f16_f16_wmma.loom", + "motifs/mul_mat_id_f16_f16_wmma_core.loom" + ], + "symbol": "ggml_mul_mat_id_f16_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_018_ggml_mul_mat_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat.input_size": "4096", + "ggml.mul_mat.output_accumulation": "0", + "ggml.mul_mat.output_size": "2048", + "ggml.mul_mat.output_unary_op": "23", + "ggml.mul_mat.weight_format": "4", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_019_ggml_gather_add_f32", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/hrx", + "count": 1, + "integer_parameters": { + "hidden_size": 2048, + "output_token_count": 1, + "source_token_count": 64 + }, + "kernel": "hrx:ggml_gather_add_f32", + "library_sources": [], + "primary_sources": [ + "gather_add_f32.loom" + ], + "sources": [ + "gather_add_f32.loom" + ], + "symbol": "ggml_gather_add_f32", + "workload_parameters": [ + { + "name": "source_token_count", + "type": "index" + }, + { + "name": "output_token_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_020_qwen3_moe_rmsnorm_f32_quantize_q8_1_x4", + "compile_parameters": { + "ggml.quantize_q8_1_x4.group_capacity": "16", + "qwen3_moe.model.hidden_size": "2048", + "qwen3_moe.model.rms_epsilon": "0.000001", + "qwen3_moe.workload.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 1, + "integer_parameters": { + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen3_moe_rmsnorm_f32_quantize_q8_1_x4", + "library_sources": [ + "qwen3_moe/model_config.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "qwen3_moe/attention_prepare_quantized.loom" + ], + "sources": [ + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "qwen3_moe_rmsnorm_f32_quantize_q8_1_x4", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_021_qwen3_moe_router_projection_top8_fused_decode_f32", + "compile_parameters": { + "qwen3_moe.model.hidden_size": "2048", + "qwen3_moe.router.expert_count": "128", + "qwen3_moe.router.route_count": "8", + "qwen3_moe.workload.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 1, + "integer_parameters": { + "route_id_stride": 128, + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen3_moe_router_projection_top8_fused_decode_f32", + "library_sources": [ + "qwen3_moe/model_config.loom", + "qwen3_moe/router_projection_f32.loom", + "qwen3_moe/router_top8_f32.loom" + ], + "primary_sources": [ + "qwen3_moe/router_projection_top8_fused_f32.loom" + ], + "sources": [ + "qwen3_moe/router_projection_top8_fused_f32.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/router_projection_f32.loom", + "qwen3_moe/router_top8_f32.loom" + ], + "symbol": "qwen3_moe_router_projection_top8_fused_decode_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_022_qwen3_moe_routed_gate_up_swiglu_q4k_q8", + "compile_parameters": { + "qwen3_moe.routed_gate_up.expert_count": "128", + "qwen3_moe.routed_gate_up.input_size": "2048", + "qwen3_moe.routed_gate_up.output_size": "768", + "qwen3_moe.routed_gate_up.route_count": "8", + "qwen3_moe.workload.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 1, + "integer_parameters": { + "expert_count": 128, + "output_size": 768, + "route_count": 8, + "route_stride": 128, + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen3_moe_routed_gate_up_swiglu_q4k_q8", + "library_sources": [ + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom" + ], + "sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "qwen3_moe_routed_gate_up_swiglu_q4k_q8", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_023_qwen3_moe_routed_down_q6k_f32_wave64_next_q8", + "compile_parameters": { + "qwen3_moe.model.hidden_size": "2048", + "qwen3_moe.model.rms_epsilon": "0.000001", + "qwen3_moe.routed_down.expert_count": "128", + "qwen3_moe.routed_down.input_size": "768", + "qwen3_moe.routed_down.output_size": "2048", + "qwen3_moe.routed_down.route_count": "8", + "qwen3_moe.workload.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 1, + "integer_parameters": { + "expert_count": 128, + "input_size": 768, + "output_size": 2048, + "route_count": 8, + "route_id_stride": 128, + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen3_moe_routed_down_q6k_f32_wave64_next_q8", + "library_sources": [ + "ggml/linear_q6k_f32.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom" + ], + "primary_sources": [ + "qwen3_moe/routed_down_q6k.loom" + ], + "sources": [ + "qwen3_moe/routed_down_q6k.loom", + "ggml/linear_q6k_f32.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom" + ], + "symbol": "qwen3_moe_routed_down_q6k_f32_wave64_next_q8", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp256_024_ggml_linear_q6k_q8_1_x4", + "compile_parameters": { + "ggml.linear_q6k_q8_1_x4.output_capacity": "151936", + "ggml.linear_q6k_q8_1_x4.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 1, + "integer_parameters": { + "input_size": 2048, + "output_size": 151936, + "token_count": 1 + }, + "kernel": "qwen3_moe:ggml_linear_q6k_q8_1_x4", + "library_sources": [ + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "ggml/linear_q6k_q8_1_x4.loom" + ], + "sources": [ + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "ggml_linear_q6k_q8_1_x4", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + } + ], + "generated_count": 25, + "kernel_counts": { + "hrx:ggml_gather_add_f32": 1, + "loom_libs:ggml_flash_attention_f32_f16_wmma": 48, + "loom_libs:ggml_get_rows_f32": 1, + "loom_libs:ggml_moe_build_expert_partition_table": 47, + "loom_libs:ggml_moe_build_expert_table": 47, + "loom_libs:ggml_mul_mat_add_f32_f32_wmma": 47, + "loom_libs:ggml_mul_mat_f32_f32_wmma": 145, + "loom_libs:ggml_mul_mat_id_f16_f16_wmma": 47, + "loom_libs:ggml_rmsnorm_binary_f32": 48, + "qwen3_moe:ggml_linear_q6k_q8_1_x4": 1, + "qwen3_moe:qwen3_moe_attention_postprocess_f32_f16": 48, + "qwen3_moe:qwen3_moe_rmsnorm_f32_quantize_q8_1_x4": 1, + "qwen3_moe:qwen3_moe_routed_down_q6k_f32_wave64_next_q8": 1, + "qwen3_moe:qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32": 47, + "qwen3_moe:qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma": 47, + "qwen3_moe:qwen3_moe_routed_gate_up_swiglu_q4k_q8": 1, + "qwen3_moe:qwen3_moe_router_projection_f32_four_row_wave32": 47, + "qwen3_moe:qwen3_moe_router_projection_top8_fused_decode_f32": 1, + "qwen3_moe:qwen3_moe_router_top8_f32": 47, + "qwen3_moe:qwen_attention_metadata": 1, + "qwen:qwen_attention_context_base_capture": 1 + }, + "loom_source": "benchmarks/loom/qwen3_30b_a3b_q4_k_m.loom", + "model": "qwen3_30b_a3b_q4_k_m", + "scenario": "pp256", + "schema": "ggml-hrx-model-loom-benchmarks-v2" +} diff --git a/ggml/src/ggml-hrx/benchmarks/loom/qwen3_30b_a3b_q4_k_m.pp512.json b/ggml/src/ggml-hrx/benchmarks/loom/qwen3_30b_a3b_q4_k_m.pp512.json new file mode 100644 index 000000000000..692d376693e8 --- /dev/null +++ b/ggml/src/ggml-hrx/benchmarks/loom/qwen3_30b_a3b_q4_k_m.pp512.json @@ -0,0 +1,1135 @@ +{ + "command_count": 1348, + "dispatch_count": 27, + "dispatches": [ + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_000_qwen_attention_context_base_capture", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/qwen", + "count": 2, + "integer_parameters": {}, + "kernel": "qwen:qwen_attention_context_base_capture", + "library_sources": [ + "motifs/mul_mat_id_f16_f16_wmma_core.loom" + ], + "primary_sources": [ + "attention_state_initialize.loom" + ], + "sources": [ + "attention_state_initialize.loom" + ], + "symbol": "qwen_attention_context_base_capture", + "workload_parameters": [] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_001_qwen_attention_metadata", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 1, + "integer_parameters": { + "context_capacity": 256, + "token_count": 64 + }, + "kernel": "qwen3_moe:qwen_attention_metadata", + "library_sources": [], + "primary_sources": [ + "../qwen/attention_metadata.loom" + ], + "sources": [ + "../qwen/attention_metadata.loom" + ], + "symbol": "qwen_attention_metadata", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "context_capacity", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_002_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "2048", + "ggml.get_rows_f32.token_capacity": "64", + "ggml.get_rows_f32.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "hidden_size": 2048, + "row_count": 151936, + "token_count": 64 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_003_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "2048", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_004_ggml_mul_mat_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat.input_size": "2048", + "ggml.mul_mat.output_accumulation": "0", + "ggml.mul_mat.output_size": "4096", + "ggml.mul_mat.output_unary_op": "23", + "ggml.mul_mat.weight_format": "4", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_005_ggml_mul_mat_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat.input_size": "2048", + "ggml.mul_mat.output_accumulation": "0", + "ggml.mul_mat.output_size": "512", + "ggml.mul_mat.output_unary_op": "23", + "ggml.mul_mat.weight_format": "6", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_006_ggml_mul_mat_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat.input_size": "2048", + "ggml.mul_mat.output_accumulation": "0", + "ggml.mul_mat.output_size": "512", + "ggml.mul_mat.output_unary_op": "23", + "ggml.mul_mat.weight_format": "4", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 144, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_007_qwen3_moe_attention_postprocess_f32_f16", + "compile_parameters": { + "qwen3_moe.attention.head_size": "128", + "qwen3_moe.attention.key_value_size": "512", + "qwen3_moe.attention.query_size": "4096", + "qwen3_moe.model.rms_epsilon": "0.000001", + "qwen3_moe.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 96, + "integer_parameters": { + "cache_row_count": 512, + "token_count": 64 + }, + "kernel": "qwen3_moe:qwen3_moe_attention_postprocess_f32_f16", + "library_sources": [ + "qwen3_moe/model_config.loom" + ], + "primary_sources": [ + "qwen3_moe/attention_postprocess_f32_f16.loom" + ], + "sources": [ + "qwen3_moe/attention_postprocess_f32_f16.loom", + "qwen3_moe/model_config.loom" + ], + "symbol": "qwen3_moe_attention_postprocess_f32_f16", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_008_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.head_size": "128", + "ggml.flash_attention.key_value_head_count": "4", + "ggml.flash_attention.query_head_count": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "key_value_token_count": 256, + "query_token_count": 64 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_009_ggml_mul_mat_add_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat_postops.input_size": "4096", + "ggml.mul_mat_postops.output_size": "2048", + "ggml.mul_mat_postops.weight_format": "4", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 94, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_add_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_add_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_add_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_add_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_010_qwen3_moe_router_projection_f32_four_row_wave32", + "compile_parameters": { + "qwen3_moe.model.hidden_size": "2048", + "qwen3_moe.router.expert_count": "128", + "qwen3_moe.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 94, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "qwen3_moe:qwen3_moe_router_projection_f32_four_row_wave32", + "library_sources": [ + "qwen3_moe/model_config.loom" + ], + "primary_sources": [ + "qwen3_moe/router_projection_f32.loom" + ], + "sources": [ + "qwen3_moe/router_projection_f32.loom", + "qwen3_moe/model_config.loom" + ], + "symbol": "qwen3_moe_router_projection_f32_four_row_wave32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_011_qwen3_moe_router_top8_f32", + "compile_parameters": { + "qwen3_moe.router.expert_count": "128", + "qwen3_moe.router.route_count": "8", + "qwen3_moe.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 94, + "integer_parameters": { + "route_id_stride": 128, + "token_count": 64 + }, + "kernel": "qwen3_moe:qwen3_moe_router_top8_f32", + "library_sources": [ + "qwen3_moe/model_config.loom" + ], + "primary_sources": [ + "qwen3_moe/router_top8_f32.loom" + ], + "sources": [ + "qwen3_moe/router_top8_f32.loom", + "qwen3_moe/model_config.loom" + ], + "symbol": "qwen3_moe_router_top8_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_012_ggml_moe_build_expert_table", + "compile_parameters": { + "ggml.moe_routing.descriptor_expert_mask": "127", + "ggml.moe_routing.descriptor_partition_shift": "7", + "ggml.moe_routing.descriptor_row_count_shift": "13", + "ggml.moe_routing.expert_count": "128", + "ggml.moe_routing.partition_workgroup_size": "128", + "ggml.moe_routing.route_count": "8", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 94, + "integer_parameters": { + "expert_count": 128, + "route_count": 8, + "route_stride": 128, + "token_count": 64 + }, + "kernel": "loom_libs:ggml_moe_build_expert_table", + "library_sources": [], + "primary_sources": [ + "ops/moe_routing_tables.loom" + ], + "sources": [ + "ops/moe_routing_tables.loom" + ], + "symbol": "ggml_moe_build_expert_table", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_013_ggml_moe_build_expert_partition_table", + "compile_parameters": { + "ggml.moe_routing.descriptor_expert_mask": "127", + "ggml.moe_routing.descriptor_partition_shift": "7", + "ggml.moe_routing.descriptor_row_count_shift": "13", + "ggml.moe_routing.expert_count": "128", + "ggml.moe_routing.partition_workgroup_size": "128", + "ggml.moe_routing.route_count": "8", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 94, + "integer_parameters": { + "expert_count": 128, + "route_count": 8, + "token_count": 64 + }, + "kernel": "loom_libs:ggml_moe_build_expert_partition_table", + "library_sources": [], + "primary_sources": [ + "ops/moe_routing_tables.loom" + ], + "sources": [ + "ops/moe_routing_tables.loom" + ], + "symbol": "ggml_moe_build_expert_partition_table", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_014_qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma", + "compile_parameters": { + "qwen3_moe.routed_gate_up.expert_count": "128", + "qwen3_moe.routed_gate_up.input_size": "2048", + "qwen3_moe.routed_gate_up.output_size": "768", + "qwen3_moe.routed_gate_up.route_count": "8", + "qwen3_moe.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 94, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "qwen3_moe:qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma", + "library_sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "qwen3_moe/routed_linear_q4k_f16_wmma.loom" + ], + "sources": [ + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_015_ggml_mul_mat_id_f16_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_id_f16_f16.expert_count": "128", + "ggml.mul_mat_id_f16_f16.input_size": "768", + "ggml.mul_mat_id_f16_f16.output_size": "2048", + "ggml.mul_mat_id_f16_f16.route_count": "8", + "ggml.mul_mat_id_f16_f16.weight_format": "6", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 46, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_id_f16_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_id_f16_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_id_f16_f16_wmma.loom", + "motifs/mul_mat_id_f16_f16_wmma_core.loom" + ], + "symbol": "ggml_mul_mat_id_f16_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_016_qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32", + "compile_parameters": { + "qwen3_moe.model.hidden_size": "2048", + "qwen3_moe.model.rms_epsilon": "0.000001", + "qwen3_moe.routed_down.expert_count": "128", + "qwen3_moe.routed_down.input_size": "768", + "qwen3_moe.routed_down.output_size": "2048", + "qwen3_moe.routed_down.route_count": "8", + "qwen3_moe.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 94, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "qwen3_moe:qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32", + "library_sources": [ + "qwen3_moe/routed_down_quantized_f16_wmma.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/routed_down_q6k.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_f32.loom" + ], + "sources": [ + "qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_f32.loom", + "qwen3_moe/routed_down_quantized_f16_wmma.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/routed_down_q6k.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_017_ggml_mul_mat_id_f16_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_id_f16_f16.expert_count": "128", + "ggml.mul_mat_id_f16_f16.input_size": "768", + "ggml.mul_mat_id_f16_f16.output_size": "2048", + "ggml.mul_mat_id_f16_f16.route_count": "8", + "ggml.mul_mat_id_f16_f16.weight_format": "4", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_id_f16_f16_wmma", + "library_sources": [ + "motifs/mul_mat_id_f16_f16_wmma_core.loom" + ], + "primary_sources": [ + "ops/mul_mat_id_f16_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_id_f16_f16_wmma.loom", + "motifs/mul_mat_id_f16_f16_wmma_core.loom" + ], + "symbol": "ggml_mul_mat_id_f16_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_018_ggml_mul_mat_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat.input_size": "4096", + "ggml.mul_mat.output_accumulation": "0", + "ggml.mul_mat.output_size": "2048", + "ggml.mul_mat.output_unary_op": "23", + "ggml.mul_mat.weight_format": "4", + "ggml.workload.token_capacity": "64" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "token_count": 64 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_019_ggml_gather_add_f32", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/hrx", + "count": 2, + "integer_parameters": { + "hidden_size": 2048, + "output_token_count": 1, + "source_token_count": 64 + }, + "kernel": "hrx:ggml_gather_add_f32", + "library_sources": [], + "primary_sources": [ + "gather_add_f32.loom" + ], + "sources": [ + "gather_add_f32.loom" + ], + "symbol": "ggml_gather_add_f32", + "workload_parameters": [ + { + "name": "source_token_count", + "type": "index" + }, + { + "name": "output_token_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_020_qwen3_moe_rmsnorm_f32_quantize_q8_1_x4", + "compile_parameters": { + "ggml.quantize_q8_1_x4.group_capacity": "16", + "qwen3_moe.model.hidden_size": "2048", + "qwen3_moe.model.rms_epsilon": "0.000001", + "qwen3_moe.workload.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 2, + "integer_parameters": { + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen3_moe_rmsnorm_f32_quantize_q8_1_x4", + "library_sources": [ + "qwen3_moe/model_config.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "qwen3_moe/attention_prepare_quantized.loom" + ], + "sources": [ + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "qwen3_moe_rmsnorm_f32_quantize_q8_1_x4", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_021_qwen3_moe_router_projection_top8_fused_decode_f32", + "compile_parameters": { + "qwen3_moe.model.hidden_size": "2048", + "qwen3_moe.router.expert_count": "128", + "qwen3_moe.router.route_count": "8", + "qwen3_moe.workload.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 2, + "integer_parameters": { + "route_id_stride": 128, + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen3_moe_router_projection_top8_fused_decode_f32", + "library_sources": [ + "qwen3_moe/model_config.loom", + "qwen3_moe/router_projection_f32.loom", + "qwen3_moe/router_top8_f32.loom" + ], + "primary_sources": [ + "qwen3_moe/router_projection_top8_fused_f32.loom" + ], + "sources": [ + "qwen3_moe/router_projection_top8_fused_f32.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/router_projection_f32.loom", + "qwen3_moe/router_top8_f32.loom" + ], + "symbol": "qwen3_moe_router_projection_top8_fused_decode_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_022_qwen3_moe_routed_gate_up_swiglu_q4k_q8", + "compile_parameters": { + "qwen3_moe.routed_gate_up.expert_count": "128", + "qwen3_moe.routed_gate_up.input_size": "2048", + "qwen3_moe.routed_gate_up.output_size": "768", + "qwen3_moe.routed_gate_up.route_count": "8", + "qwen3_moe.workload.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 2, + "integer_parameters": { + "expert_count": 128, + "output_size": 768, + "route_count": 8, + "route_stride": 128, + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen3_moe_routed_gate_up_swiglu_q4k_q8", + "library_sources": [ + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom" + ], + "sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "qwen3_moe_routed_gate_up_swiglu_q4k_q8", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_023_qwen3_moe_routed_down_q6k_f32_wave64_next_q8", + "compile_parameters": { + "qwen3_moe.model.hidden_size": "2048", + "qwen3_moe.model.rms_epsilon": "0.000001", + "qwen3_moe.routed_down.expert_count": "128", + "qwen3_moe.routed_down.input_size": "768", + "qwen3_moe.routed_down.output_size": "2048", + "qwen3_moe.routed_down.route_count": "8", + "qwen3_moe.workload.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 2, + "integer_parameters": { + "expert_count": 128, + "input_size": 768, + "output_size": 2048, + "route_count": 8, + "route_id_stride": 128, + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen3_moe_routed_down_q6k_f32_wave64_next_q8", + "library_sources": [ + "ggml/linear_q6k_f32.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom" + ], + "primary_sources": [ + "qwen3_moe/routed_down_q6k.loom" + ], + "sources": [ + "qwen3_moe/routed_down_q6k.loom", + "ggml/linear_q6k_f32.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom" + ], + "symbol": "qwen3_moe_routed_down_q6k_f32_wave64_next_q8", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_024_ggml_linear_q6k_q8_1_x4", + "compile_parameters": { + "ggml.linear_q6k_q8_1_x4.output_capacity": "151936", + "ggml.linear_q6k_q8_1_x4.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 2, + "integer_parameters": { + "input_size": 2048, + "output_size": 151936, + "token_count": 1 + }, + "kernel": "qwen3_moe:ggml_linear_q6k_q8_1_x4", + "library_sources": [ + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "ggml/linear_q6k_q8_1_x4.loom" + ], + "sources": [ + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "ggml_linear_q6k_q8_1_x4", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_025_qwen_attention_metadata", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 1, + "integer_parameters": { + "context_capacity": 512, + "token_count": 64 + }, + "kernel": "qwen3_moe:qwen_attention_metadata", + "library_sources": [], + "primary_sources": [ + "../qwen/attention_metadata.loom" + ], + "sources": [ + "../qwen/attention_metadata.loom" + ], + "symbol": "qwen_attention_metadata", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "context_capacity", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_pp512_026_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.head_size": "128", + "ggml.flash_attention.key_value_head_count": "4", + "ggml.flash_attention.query_head_count": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "key_value_token_count": 512, + "query_token_count": 64 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + } + ], + "generated_count": 27, + "kernel_counts": { + "hrx:ggml_gather_add_f32": 2, + "loom_libs:ggml_flash_attention_f32_f16_wmma": 96, + "loom_libs:ggml_get_rows_f32": 2, + "loom_libs:ggml_moe_build_expert_partition_table": 94, + "loom_libs:ggml_moe_build_expert_table": 94, + "loom_libs:ggml_mul_mat_add_f32_f32_wmma": 94, + "loom_libs:ggml_mul_mat_f32_f32_wmma": 290, + "loom_libs:ggml_mul_mat_id_f16_f16_wmma": 94, + "loom_libs:ggml_rmsnorm_binary_f32": 96, + "qwen3_moe:ggml_linear_q6k_q8_1_x4": 2, + "qwen3_moe:qwen3_moe_attention_postprocess_f32_f16": 96, + "qwen3_moe:qwen3_moe_rmsnorm_f32_quantize_q8_1_x4": 2, + "qwen3_moe:qwen3_moe_routed_down_q6k_f32_wave64_next_q8": 2, + "qwen3_moe:qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32": 94, + "qwen3_moe:qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma": 94, + "qwen3_moe:qwen3_moe_routed_gate_up_swiglu_q4k_q8": 2, + "qwen3_moe:qwen3_moe_router_projection_f32_four_row_wave32": 94, + "qwen3_moe:qwen3_moe_router_projection_top8_fused_decode_f32": 2, + "qwen3_moe:qwen3_moe_router_top8_f32": 94, + "qwen3_moe:qwen_attention_metadata": 2, + "qwen:qwen_attention_context_base_capture": 2 + }, + "loom_source": "benchmarks/loom/qwen3_30b_a3b_q4_k_m.loom", + "model": "qwen3_30b_a3b_q4_k_m", + "scenario": "pp512", + "schema": "ggml-hrx-model-loom-benchmarks-v2" +} diff --git a/ggml/src/ggml-hrx/benchmarks/loom/qwen3_30b_a3b_q4_k_m.tg8.json b/ggml/src/ggml-hrx/benchmarks/loom/qwen3_30b_a3b_q4_k_m.tg8.json new file mode 100644 index 000000000000..35cbb8e814a9 --- /dev/null +++ b/ggml/src/ggml-hrx/benchmarks/loom/qwen3_30b_a3b_q4_k_m.tg8.json @@ -0,0 +1,672 @@ +{ + "command_count": 293, + "dispatch_count": 14, + "dispatches": [ + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_tg8_000_qwen_attention_context_base_capture", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/qwen", + "count": 1, + "integer_parameters": {}, + "kernel": "qwen:qwen_attention_context_base_capture", + "library_sources": [], + "primary_sources": [ + "attention_state_initialize.loom" + ], + "sources": [ + "attention_state_initialize.loom" + ], + "symbol": "qwen_attention_context_base_capture", + "workload_parameters": [] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_tg8_001_qwen_attention_metadata", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 1, + "integer_parameters": { + "context_capacity": 256, + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen_attention_metadata", + "library_sources": [], + "primary_sources": [ + "../qwen/attention_metadata.loom" + ], + "sources": [ + "../qwen/attention_metadata.loom" + ], + "symbol": "qwen_attention_metadata", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "context_capacity", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_tg8_002_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "2048", + "ggml.get_rows_f32.token_capacity": "1", + "ggml.get_rows_f32.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "hidden_size": 2048, + "row_count": 151936, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_tg8_003_qwen3_moe_rmsnorm_f32_quantize_q8_1_x4", + "compile_parameters": { + "ggml.quantize_q8_1_x4.group_capacity": "16", + "qwen3_moe.model.hidden_size": "2048", + "qwen3_moe.model.rms_epsilon": "0.000001", + "qwen3_moe.workload.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 1, + "integer_parameters": { + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen3_moe_rmsnorm_f32_quantize_q8_1_x4", + "library_sources": [ + "qwen3_moe/model_config.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "qwen3_moe/attention_prepare_quantized.loom" + ], + "sources": [ + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "qwen3_moe_rmsnorm_f32_quantize_q8_1_x4", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_tg8_004_qwen3_moe_attention_qkv_postprocess_fused_decode", + "compile_parameters": { + "qwen3_moe.attention.head_size": "128", + "qwen3_moe.attention.key_value_size": "512", + "qwen3_moe.attention.query_size": "4096", + "qwen3_moe.attention.value_uses_q6": "1", + "qwen3_moe.model.hidden_size": "2048", + "qwen3_moe.model.rms_epsilon": "0.000001", + "qwen3_moe.workload.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 24, + "integer_parameters": { + "cache_row_count": 256, + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen3_moe_attention_qkv_postprocess_fused_decode", + "library_sources": [ + "qwen3_moe/attention_postprocess_f32_f16.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/attention_qkv_quantized.loom", + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "qwen3_moe/attention_qkv_postprocess_fused.loom" + ], + "sources": [ + "qwen3_moe/attention_qkv_postprocess_fused.loom", + "qwen3_moe/attention_postprocess_f32_f16.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/attention_qkv_quantized.loom", + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "qwen3_moe_attention_qkv_postprocess_fused_decode", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_tg8_005_ggml_flash_attention_decode_split_f32_f16_wmma_next_q8", + "compile_parameters": { + "ggml.flash_attention.decode.key_value_token_capacity": "256", + "ggml.flash_attention.key_value_head_count": "4", + "ggml.flash_attention.query_head_count": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "key_value_token_count": 256 + }, + "kernel": "loom_libs:ggml_flash_attention_decode_split_f32_f16_wmma_next_q8", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_decode_split_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_decode_split_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_decode_split_f32_f16_wmma_next_q8", + "workload_parameters": [ + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_tg8_006_qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8", + "compile_parameters": { + "qwen3_moe.dense_quantized.input_size": "4096", + "qwen3_moe.dense_quantized.output_accumulation": "1", + "qwen3_moe.dense_quantized.output_size": "2048", + "qwen3_moe.model.hidden_size": "2048", + "qwen3_moe.model.rms_epsilon": "0.000001", + "qwen3_moe.workload.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 48, + "integer_parameters": { + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8", + "library_sources": [ + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom" + ], + "sources": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_tg8_007_qwen3_moe_router_projection_top8_fused_decode_f32", + "compile_parameters": { + "qwen3_moe.model.hidden_size": "2048", + "qwen3_moe.router.expert_count": "128", + "qwen3_moe.router.route_count": "8", + "qwen3_moe.workload.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 48, + "integer_parameters": { + "route_id_stride": 128, + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen3_moe_router_projection_top8_fused_decode_f32", + "library_sources": [ + "qwen3_moe/model_config.loom", + "qwen3_moe/router_projection_f32.loom", + "qwen3_moe/router_top8_f32.loom" + ], + "primary_sources": [ + "qwen3_moe/router_projection_top8_fused_f32.loom" + ], + "sources": [ + "qwen3_moe/router_projection_top8_fused_f32.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/router_projection_f32.loom", + "qwen3_moe/router_top8_f32.loom" + ], + "symbol": "qwen3_moe_router_projection_top8_fused_decode_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_tg8_008_qwen3_moe_routed_gate_up_swiglu_q4k_q8", + "compile_parameters": { + "qwen3_moe.routed_gate_up.expert_count": "128", + "qwen3_moe.routed_gate_up.input_size": "2048", + "qwen3_moe.routed_gate_up.output_size": "768", + "qwen3_moe.routed_gate_up.route_count": "8", + "qwen3_moe.workload.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 24, + "integer_parameters": { + "expert_count": 128, + "output_size": 768, + "route_count": 8, + "route_stride": 128, + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen3_moe_routed_gate_up_swiglu_q4k_q8", + "library_sources": [ + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom" + ], + "sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "qwen3_moe_routed_gate_up_swiglu_q4k_q8", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_tg8_009_qwen3_moe_routed_down_q6k_f32_wave64_next_q8", + "compile_parameters": { + "qwen3_moe.model.hidden_size": "2048", + "qwen3_moe.model.rms_epsilon": "0.000001", + "qwen3_moe.routed_down.expert_count": "128", + "qwen3_moe.routed_down.input_size": "768", + "qwen3_moe.routed_down.output_size": "2048", + "qwen3_moe.routed_down.route_count": "8", + "qwen3_moe.workload.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 24, + "integer_parameters": { + "expert_count": 128, + "input_size": 768, + "output_size": 2048, + "route_count": 8, + "route_id_stride": 128, + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen3_moe_routed_down_q6k_f32_wave64_next_q8", + "library_sources": [ + "ggml/linear_q6k_f32.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom" + ], + "primary_sources": [ + "qwen3_moe/routed_down_q6k.loom" + ], + "sources": [ + "qwen3_moe/routed_down_q6k.loom", + "ggml/linear_q6k_f32.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom" + ], + "symbol": "qwen3_moe_routed_down_q6k_f32_wave64_next_q8", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_tg8_010_qwen3_moe_attention_qkv_postprocess_fused_decode", + "compile_parameters": { + "qwen3_moe.attention.head_size": "128", + "qwen3_moe.attention.key_value_size": "512", + "qwen3_moe.attention.query_size": "4096", + "qwen3_moe.attention.value_uses_q6": "0", + "qwen3_moe.model.hidden_size": "2048", + "qwen3_moe.model.rms_epsilon": "0.000001", + "qwen3_moe.workload.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 24, + "integer_parameters": { + "cache_row_count": 256, + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen3_moe_attention_qkv_postprocess_fused_decode", + "library_sources": [ + "qwen3_moe/attention_postprocess_f32_f16.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/attention_qkv_quantized.loom", + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "qwen3_moe/attention_qkv_postprocess_fused.loom" + ], + "sources": [ + "qwen3_moe/attention_qkv_postprocess_fused.loom", + "qwen3_moe/attention_postprocess_f32_f16.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/attention_qkv_quantized.loom", + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "qwen3_moe_attention_qkv_postprocess_fused_decode", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_tg8_011_qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8", + "compile_parameters": { + "qwen3_moe.routed_gate_up.expert_count": "128", + "qwen3_moe.routed_gate_up.input_size": "2048", + "qwen3_moe.routed_gate_up.output_size": "768", + "qwen3_moe.routed_gate_up.route_count": "8", + "qwen3_moe.workload.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 24, + "integer_parameters": { + "expert_count": 128, + "output_size": 768, + "route_count": 8, + "route_stride": 128, + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8", + "library_sources": [ + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom" + ], + "sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_tg8_012_qwen3_moe_routed_down_q4k_q8_1_x4_next_q8", + "compile_parameters": { + "qwen3_moe.model.hidden_size": "2048", + "qwen3_moe.model.rms_epsilon": "0.000001", + "qwen3_moe.routed_down.expert_count": "128", + "qwen3_moe.routed_down.input_size": "768", + "qwen3_moe.routed_down.output_size": "2048", + "qwen3_moe.routed_down.route_count": "8", + "qwen3_moe.workload.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 24, + "integer_parameters": { + "expert_count": 128, + "input_size": 768, + "output_size": 2048, + "route_count": 8, + "route_id_stride": 128, + "token_count": 1 + }, + "kernel": "qwen3_moe:qwen3_moe_routed_down_q4k_q8_1_x4_next_q8", + "library_sources": [ + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "qwen3_moe/routed_down_q4k.loom" + ], + "sources": [ + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "qwen3_moe_routed_down_q4k_q8_1_x4_next_q8", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen3_30b_a3b_q4_k_m_tg8_013_ggml_linear_q6k_q8_1_x4", + "compile_parameters": { + "ggml.linear_q6k_q8_1_x4.output_capacity": "151936", + "ggml.linear_q6k_q8_1_x4.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/qwen_moe", + "count": 1, + "integer_parameters": { + "input_size": 2048, + "output_size": 151936, + "token_count": 1 + }, + "kernel": "qwen3_moe:ggml_linear_q6k_q8_1_x4", + "library_sources": [ + "ggml/quantize_q8_1_x4.loom" + ], + "primary_sources": [ + "ggml/linear_q6k_q8_1_x4.loom" + ], + "sources": [ + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom" + ], + "symbol": "ggml_linear_q6k_q8_1_x4", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + } + ], + "generated_count": 14, + "kernel_counts": { + "loom_libs:ggml_flash_attention_decode_split_f32_f16_wmma_next_q8": 48, + "loom_libs:ggml_get_rows_f32": 1, + "qwen3_moe:ggml_linear_q6k_q8_1_x4": 1, + "qwen3_moe:qwen3_moe_attention_qkv_postprocess_fused_decode": 48, + "qwen3_moe:qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8": 48, + "qwen3_moe:qwen3_moe_rmsnorm_f32_quantize_q8_1_x4": 1, + "qwen3_moe:qwen3_moe_routed_down_q4k_q8_1_x4_next_q8": 24, + "qwen3_moe:qwen3_moe_routed_down_q6k_f32_wave64_next_q8": 24, + "qwen3_moe:qwen3_moe_routed_gate_up_swiglu_q4k_q8": 24, + "qwen3_moe:qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8": 24, + "qwen3_moe:qwen3_moe_router_projection_top8_fused_decode_f32": 48, + "qwen3_moe:qwen_attention_metadata": 1, + "qwen:qwen_attention_context_base_capture": 1 + }, + "loom_source": "benchmarks/loom/qwen3_30b_a3b_q4_k_m.loom", + "model": "qwen3_30b_a3b_q4_k_m", + "scenario": "tg8", + "schema": "ggml-hrx-model-loom-benchmarks-v2" +} diff --git a/ggml/src/ggml-hrx/cmake/BundleRuntime.cmake b/ggml/src/ggml-hrx/cmake/BundleRuntime.cmake new file mode 100644 index 000000000000..91853b38ad18 --- /dev/null +++ b/ggml/src/ggml-hrx/cmake/BundleRuntime.cmake @@ -0,0 +1,204 @@ +# Capture the module directory while it is the active list file. On CMake 3.14, +# CMAKE_CURRENT_LIST_DIR inside a function refers to the function's call site. +set(_GGML_HRX_BUNDLE_RUNTIME_DIR "${CMAKE_CURRENT_LIST_DIR}") + +# Find every entry matching PATTERN in exactly one of SEARCH_DIRS. +# +# OUT_PATHS receives the sorted matching paths. Missing families are fatal. +function(_ggml_hrx_find_bundle_entries OUT_PATHS LABEL PATTERN SEARCH_DIRS) + set(GGML_HRX_SELECTED_PATHS) + set(GGML_HRX_SELECTED_DIR "") + foreach(GGML_HRX_SEARCH_DIR IN LISTS SEARCH_DIRS) + file(GLOB GGML_HRX_DIR_MATCHES + CONFIGURE_DEPENDS + LIST_DIRECTORIES FALSE + "${GGML_HRX_SEARCH_DIR}/${PATTERN}") + if (GGML_HRX_DIR_MATCHES) + if (NOT GGML_HRX_SELECTED_DIR STREQUAL "") + message(FATAL_ERROR "GGML_HRX_BUNDLE_RUNTIME_LIBS found ${LABEL} in multiple source directories: ${GGML_HRX_SELECTED_DIR};${GGML_HRX_SEARCH_DIR}") + endif() + set(GGML_HRX_SELECTED_DIR "${GGML_HRX_SEARCH_DIR}") + set(GGML_HRX_SELECTED_PATHS ${GGML_HRX_DIR_MATCHES}) + endif() + endforeach() + + if (NOT GGML_HRX_SELECTED_PATHS) + message(FATAL_ERROR "GGML_HRX_BUNDLE_RUNTIME_LIBS could not find required ${LABEL} matching ${PATTERN}. Searched: ${SEARCH_DIRS}") + endif() + + foreach(GGML_HRX_SELECTED_PATH IN LISTS GGML_HRX_SELECTED_PATHS) + if (IS_SYMLINK "${GGML_HRX_SELECTED_PATH}") + if (NOT EXISTS "${GGML_HRX_SELECTED_PATH}") + message(FATAL_ERROR "GGML_HRX_BUNDLE_RUNTIME_LIBS matched a broken symlink: ${GGML_HRX_SELECTED_PATH}") + endif() + elseif(NOT EXISTS "${GGML_HRX_SELECTED_PATH}") + message(FATAL_ERROR "GGML_HRX_BUNDLE_RUNTIME_LIBS matched a nonexistent file: ${GGML_HRX_SELECTED_PATH}") + endif() + endforeach() + list(LENGTH GGML_HRX_SELECTED_PATHS GGML_HRX_SELECTED_COUNT) + message(STATUS " ${LABEL}: ${GGML_HRX_SELECTED_DIR} (${GGML_HRX_SELECTED_COUNT} entries)") + + set(${OUT_PATHS} "${GGML_HRX_SELECTED_PATHS}" PARENT_SCOPE) +endfunction() + +# Keep the target's existing relative build and install RPATH entries and drop +# absolute entries. Append the paths needed by the adjacent runtime bundle. The +# resulting list is returned through OUT_VAR. +function(_ggml_hrx_collect_portable_rpath TARGET_NAME OUT_VAR) + set(GGML_HRX_PORTABLE_RPATH) + foreach(GGML_HRX_RPATH_PROPERTY BUILD_RPATH INSTALL_RPATH) + get_target_property(GGML_HRX_RPATH_ENTRIES "${TARGET_NAME}" "${GGML_HRX_RPATH_PROPERTY}") + if (NOT GGML_HRX_RPATH_ENTRIES) + continue() + endif() + foreach(GGML_HRX_RPATH_ENTRY IN LISTS GGML_HRX_RPATH_ENTRIES) + if (GGML_HRX_RPATH_ENTRY STREQUAL "") + continue() + endif() + if (IS_ABSOLUTE "${GGML_HRX_RPATH_ENTRY}") + continue() + endif() + list(APPEND GGML_HRX_PORTABLE_RPATH "${GGML_HRX_RPATH_ENTRY}") + endforeach() + endforeach() + list(APPEND GGML_HRX_PORTABLE_RPATH + "$ORIGIN" + "$ORIGIN/rocm_sysdeps/lib") + list(REMOVE_DUPLICATES GGML_HRX_PORTABLE_RPATH) + set(${OUT_VAR} "${GGML_HRX_PORTABLE_RPATH}" PARENT_SCOPE) +endfunction() + +# Add matching build-tree copy and install rules for SOURCES. +# RELATIVE_DESTINATION is appended below both backend destinations. +function(_ggml_hrx_add_bundle_rules TARGET_NAME SOURCES RELATIVE_DESTINATION INSTALL_BASE COPY_SCRIPT) + set(GGML_HRX_BUILD_DESTINATION "$") + set(GGML_HRX_INSTALL_DESTINATION "${INSTALL_BASE}") + if (NOT RELATIVE_DESTINATION STREQUAL "") + string(APPEND GGML_HRX_BUILD_DESTINATION "/${RELATIVE_DESTINATION}") + string(APPEND GGML_HRX_INSTALL_DESTINATION "/${RELATIVE_DESTINATION}") + endif() + + if (NOT RELATIVE_DESTINATION STREQUAL "") + add_custom_command(TARGET "${TARGET_NAME}" POST_BUILD + COMMAND "${CMAKE_COMMAND}" -E make_directory + "${GGML_HRX_BUILD_DESTINATION}" + VERBATIM) + endif() + add_custom_command(TARGET "${TARGET_NAME}" POST_BUILD + COMMAND "${CMAKE_COMMAND}" + "-DGGML_HRX_BUNDLE_SOURCES=${SOURCES}" + "-DGGML_HRX_BUNDLE_DESTINATION=${GGML_HRX_BUILD_DESTINATION}" + -P "${COPY_SCRIPT}" + VERBATIM) + install(FILES ${SOURCES} + DESTINATION "${GGML_HRX_INSTALL_DESTINATION}") +endfunction() + +# Discover, copy, and install the HRX runtime dependencies for TARGET_NAME. +function(ggml_hrx_bundle_runtime TARGET_NAME) + set(GGML_HRX_COPY_SCRIPT "${_GGML_HRX_BUNDLE_RUNTIME_DIR}/copy_bundle_entry.cmake") + if (NOT TARGET "${TARGET_NAME}") + message(FATAL_ERROR "ggml_hrx_bundle_runtime target does not exist: ${TARGET_NAME}") + endif() + if (NOT CMAKE_SYSTEM_NAME STREQUAL "Linux") + message(FATAL_ERROR "GGML_HRX_BUNDLE_RUNTIME_LIBS is currently implemented for Linux only") + endif() + if (NOT BUILD_SHARED_LIBS) + message(FATAL_ERROR "GGML_HRX_BUNDLE_RUNTIME_LIBS requires BUILD_SHARED_LIBS=ON") + endif() + + set(GGML_HRX_BUNDLE_SEARCH_DIRS) + foreach(GGML_HRX_BUNDLE_SEARCH_DIR IN LISTS GGML_HRX_BUNDLE_LIBRARY_DIRS) + if (GGML_HRX_BUNDLE_SEARCH_DIR STREQUAL "") + message(FATAL_ERROR "GGML_HRX_BUNDLE_LIBRARY_DIRS contains an empty directory entry") + endif() + get_filename_component(GGML_HRX_BUNDLE_SEARCH_DIR_ABSOLUTE "${GGML_HRX_BUNDLE_SEARCH_DIR}" ABSOLUTE BASE_DIR "${CMAKE_CURRENT_SOURCE_DIR}") + if (NOT IS_DIRECTORY "${GGML_HRX_BUNDLE_SEARCH_DIR_ABSOLUTE}") + message(FATAL_ERROR "GGML_HRX_BUNDLE_LIBRARY_DIRS contains a nonexistent directory: ${GGML_HRX_BUNDLE_SEARCH_DIR}") + endif() + get_filename_component(GGML_HRX_BUNDLE_SEARCH_DIR_CANONICAL "${GGML_HRX_BUNDLE_SEARCH_DIR_ABSOLUTE}" REALPATH) + list(APPEND GGML_HRX_BUNDLE_SEARCH_DIRS "${GGML_HRX_BUNDLE_SEARCH_DIR_CANONICAL}") + endforeach() + list(REMOVE_DUPLICATES GGML_HRX_BUNDLE_SEARCH_DIRS) + if (NOT GGML_HRX_BUNDLE_SEARCH_DIRS) + message(FATAL_ERROR "GGML_HRX_BUNDLE_LIBRARY_DIRS must list at least one directory when GGML_HRX_BUNDLE_RUNTIME_LIBS=ON") + endif() + + message(STATUS "HRX runtime bundling search directories: ${GGML_HRX_BUNDLE_SEARCH_DIRS}") + message(STATUS "HRX runtime bundle selection:") + _ggml_hrx_find_bundle_entries( + GGML_HRX_BUNDLE_HRX_LIBS + "libhrx" "libhrx.so*" "${GGML_HRX_BUNDLE_SEARCH_DIRS}") + _ggml_hrx_find_bundle_entries( + GGML_HRX_BUNDLE_LOOMC_LIBS + "libloomc" "libloomc.so*" "${GGML_HRX_BUNDLE_SEARCH_DIRS}") + _ggml_hrx_find_bundle_entries( + GGML_HRX_BUNDLE_HSA_RUNTIME_LIBS + "libhsa-runtime64" "libhsa-runtime64.so*" "${GGML_HRX_BUNDLE_SEARCH_DIRS}") + _ggml_hrx_find_bundle_entries( + GGML_HRX_BUNDLE_HSA_AQLPROFILE_LIBS + "libhsa-amd-aqlprofile64" "libhsa-amd-aqlprofile64.so*" "${GGML_HRX_BUNDLE_SEARCH_DIRS}") + _ggml_hrx_find_bundle_entries( + GGML_HRX_BUNDLE_ROCPROFILER_REGISTER_LIBS + "librocprofiler-register" "librocprofiler-register.so*" "${GGML_HRX_BUNDLE_SEARCH_DIRS}") + _ggml_hrx_find_bundle_entries( + GGML_HRX_BUNDLE_OMP_LIBS + "libomp" "libomp.so*" "${GGML_HRX_BUNDLE_SEARCH_DIRS}") + # HRX runtime bundles require the ROCm sysdeps overlay. + _ggml_hrx_find_bundle_entries( + GGML_HRX_BUNDLE_SYSDEP_LIBS + "rocm_sysdeps/lib overlay" "rocm_sysdeps/lib/*.so*" "${GGML_HRX_BUNDLE_SEARCH_DIRS}") + + set(GGML_HRX_BUNDLE_MAIN_LIBS + ${GGML_HRX_BUNDLE_HRX_LIBS} + ${GGML_HRX_BUNDLE_LOOMC_LIBS} + ${GGML_HRX_BUNDLE_HSA_RUNTIME_LIBS} + ${GGML_HRX_BUNDLE_HSA_AQLPROFILE_LIBS} + ${GGML_HRX_BUNDLE_ROCPROFILER_REGISTER_LIBS} + ${GGML_HRX_BUNDLE_OMP_LIBS} + ) + set_property(TARGET "${TARGET_NAME}" APPEND PROPERTY LINK_DEPENDS + ${GGML_HRX_BUNDLE_MAIN_LIBS} + ${GGML_HRX_BUNDLE_SYSDEP_LIBS} + "${GGML_HRX_COPY_SCRIPT}") + + # GGML_BACKEND_DL builds backends as runtime-loaded modules instead of + # normally linked libraries. Install dependencies next to the backend: + # modules use GGML_BACKEND_DIR or bin; linked backends use the standard + # library directory. + if (GGML_BACKEND_DL) + if (GGML_BACKEND_DIR) + set(GGML_HRX_BUNDLE_INSTALL_DIR "${GGML_BACKEND_DIR}") + else() + set(GGML_HRX_BUNDLE_INSTALL_DIR "${CMAKE_INSTALL_BINDIR}") + endif() + else() + set(GGML_HRX_BUNDLE_INSTALL_DIR "${CMAKE_INSTALL_LIBDIR}") + endif() + message(STATUS "HRX runtime bundle install directory: ${GGML_HRX_BUNDLE_INSTALL_DIR}") + + # Give build and install artifacts the same portable RUNPATH. Keep relative + # entries, add adjacent bundle directories, and prevent absolute HRX/ROCm + # link directories. + _ggml_hrx_collect_portable_rpath("${TARGET_NAME}" GGML_HRX_PORTABLE_RPATH) + set_target_properties("${TARGET_NAME}" PROPERTIES + BUILD_RPATH "${GGML_HRX_PORTABLE_RPATH}" + INSTALL_RPATH "${GGML_HRX_PORTABLE_RPATH}" + BUILD_WITH_INSTALL_RPATH TRUE + INSTALL_RPATH_USE_LINK_PATH FALSE + ) + message(STATUS "HRX backend RUNPATH: ${GGML_HRX_PORTABLE_RPATH}") + + _ggml_hrx_add_bundle_rules( + "${TARGET_NAME}" + "${GGML_HRX_BUNDLE_MAIN_LIBS}" + "" + "${GGML_HRX_BUNDLE_INSTALL_DIR}" + "${GGML_HRX_COPY_SCRIPT}") + _ggml_hrx_add_bundle_rules( + "${TARGET_NAME}" + "${GGML_HRX_BUNDLE_SYSDEP_LIBS}" + "rocm_sysdeps/lib" + "${GGML_HRX_BUNDLE_INSTALL_DIR}" + "${GGML_HRX_COPY_SCRIPT}") +endfunction() diff --git a/ggml/src/ggml-hrx/cmake/copy_bundle_entry.cmake b/ggml/src/ggml-hrx/cmake/copy_bundle_entry.cmake new file mode 100644 index 000000000000..70bb5bbe52ad --- /dev/null +++ b/ggml/src/ggml-hrx/cmake/copy_bundle_entry.cmake @@ -0,0 +1,33 @@ +if (NOT DEFINED GGML_HRX_BUNDLE_SOURCES OR NOT DEFINED GGML_HRX_BUNDLE_DESTINATION) + message(FATAL_ERROR "HRX bundle copy requires sources and destination") +endif() + +foreach(GGML_HRX_BUNDLE_SOURCE IN LISTS GGML_HRX_BUNDLE_SOURCES) + get_filename_component(GGML_HRX_BUNDLE_NAME "${GGML_HRX_BUNDLE_SOURCE}" NAME) + set(GGML_HRX_BUNDLE_DEST "${GGML_HRX_BUNDLE_DESTINATION}/${GGML_HRX_BUNDLE_NAME}") + if (IS_SYMLINK "${GGML_HRX_BUNDLE_SOURCE}") + if (NOT EXISTS "${GGML_HRX_BUNDLE_SOURCE}") + message(FATAL_ERROR "HRX bundle source is a broken symlink: ${GGML_HRX_BUNDLE_SOURCE}") + endif() + file(READ_SYMLINK "${GGML_HRX_BUNDLE_SOURCE}" GGML_HRX_BUNDLE_LINK_TARGET) + elseif (EXISTS "${GGML_HRX_BUNDLE_SOURCE}") + if (IS_DIRECTORY "${GGML_HRX_BUNDLE_SOURCE}") + message(FATAL_ERROR "HRX bundle source is not a regular file: ${GGML_HRX_BUNDLE_SOURCE}") + endif() + else() + message(FATAL_ERROR "HRX bundle source no longer exists: ${GGML_HRX_BUNDLE_SOURCE}") + endif() + + get_filename_component(GGML_HRX_BUNDLE_SOURCE_ABSOLUTE "${GGML_HRX_BUNDLE_SOURCE}" ABSOLUTE) + get_filename_component(GGML_HRX_BUNDLE_DEST_ABSOLUTE "${GGML_HRX_BUNDLE_DEST}" ABSOLUTE) + if ("${GGML_HRX_BUNDLE_SOURCE_ABSOLUTE}" STREQUAL "${GGML_HRX_BUNDLE_DEST_ABSOLUTE}") + continue() + endif() + + file(REMOVE "${GGML_HRX_BUNDLE_DEST}") + if (IS_SYMLINK "${GGML_HRX_BUNDLE_SOURCE}") + file(CREATE_LINK "${GGML_HRX_BUNDLE_LINK_TARGET}" "${GGML_HRX_BUNDLE_DEST}" SYMBOLIC) + else() + file(COPY "${GGML_HRX_BUNDLE_SOURCE}" DESTINATION "${GGML_HRX_BUNDLE_DESTINATION}") + endif() +endforeach() diff --git a/ggml/src/ggml-hrx/dispatch/command-plan-metadata.cpp b/ggml/src/ggml-hrx/dispatch/command-plan-metadata.cpp new file mode 100644 index 000000000000..89d9572d66f8 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/command-plan-metadata.cpp @@ -0,0 +1,146 @@ +#include "command-plan-metadata.h" + +#include +#include + +namespace ggml::hrx { +namespace { + +static bool metadata_matches(const CommandPlanResourceMetadata & lhs, const CommandPlanResourceMetadata & rhs) { + return lhs.kind == rhs.kind && lhs.size == rhs.size && + std::memcmp(lhs.bytes.data(), rhs.bytes.data(), lhs.size) == 0; +} + +static bool generated_resource_matches(const CommandPlanGeneratedResource & lhs, + const CommandPlanGeneratedResource & rhs) { + return lhs.source_value == rhs.source_value && lhs.role == rhs.role && lhs.generated_value == rhs.generated_value && + lhs.byte_count == rhs.byte_count && metadata_matches(lhs.metadata, rhs.metadata); +} + +static bool alternate_value_matches(const CommandPlanAlternateValue & lhs, const CommandPlanAlternateValue & rhs) { + return lhs.graph_value == rhs.graph_value && lhs.alternate_value == rhs.alternate_value && lhs.type == rhs.type && + lhs.byte_count == rhs.byte_count && lhs.name == rhs.name; +} + +static bool moe_routing_bundle_matches(const CommandPlanMoeRoutingBundle & lhs, + const CommandPlanMoeRoutingBundle & rhs) { + return lhs.route_ids == rhs.route_ids && lhs.route_weights == rhs.route_weights && + lhs.expert_table == rhs.expert_table && lhs.partition_table == rhs.partition_table && + lhs.expert_table_byte_count == rhs.expert_table_byte_count && + lhs.partition_table_byte_count == rhs.partition_table_byte_count && lhs.token_count == rhs.token_count && + lhs.route_count == rhs.route_count && lhs.route_stride == rhs.route_stride && + lhs.expert_count == rhs.expert_count; +} + +} // namespace + +void CommandPlanMetadata::clear() { + generated_resources_.clear(); + alternate_values_.clear(); + moe_routing_bundles_.clear(); +} + +bool CommandPlanMetadata::append(CommandPlanMetadata && other, Status & status) { + for (CommandPlanGeneratedResource & resource : other.generated_resources_) { + if (!append_generated_resource(std::move(resource), status)) { + return false; + } + } + for (CommandPlanAlternateValue & alternate : other.alternate_values_) { + if (!append_alternate_value(std::move(alternate), status)) { + return false; + } + } + for (CommandPlanMoeRoutingBundle & bundle : other.moe_routing_bundles_) { + if (!append_moe_routing_bundle(std::move(bundle), status)) { + return false; + } + } + return true; +} + +bool CommandPlanMetadata::append_generated_resource(CommandPlanGeneratedResource resource, Status & status) { + for (const CommandPlanGeneratedResource & existing : generated_resources_) { + if (existing.source_value == resource.source_value && existing.role == resource.role) { + if (generated_resource_matches(existing, resource)) { + return true; + } + status.log("conflicting generated resource for source value %d role %d", resource.source_value.value, + static_cast(resource.role)); + return false; + } + } + generated_resources_.push_back(std::move(resource)); + return true; +} + +bool CommandPlanMetadata::append_alternate_value(CommandPlanAlternateValue alternate, Status & status) { + for (const CommandPlanAlternateValue & existing : alternate_values_) { + if (existing.graph_value == alternate.graph_value && existing.type == alternate.type && + existing.byte_count == alternate.byte_count) { + if (alternate_value_matches(existing, alternate)) { + return true; + } + status.log("conflicting alternate value for graph value %d type %d byte_count %zu", + alternate.graph_value.value, static_cast(alternate.type), alternate.byte_count); + return false; + } + } + alternate_values_.push_back(std::move(alternate)); + return true; +} + +bool CommandPlanMetadata::append_moe_routing_bundle(CommandPlanMoeRoutingBundle bundle, Status & status) { + for (const CommandPlanMoeRoutingBundle & existing : moe_routing_bundles_) { + if (existing.route_ids == bundle.route_ids) { + if (moe_routing_bundle_matches(existing, bundle)) { + return true; + } + status.log("conflicting MoE routing bundle for route ids value %d", bundle.route_ids.value); + return false; + } + } + moe_routing_bundles_.push_back(std::move(bundle)); + return true; +} + +const CommandPlanGeneratedResource * CommandPlanMetadata::find_generated_resource(ValueId source_value, + GeneratedResourceRole role) const { + for (const CommandPlanGeneratedResource & resource : generated_resources_) { + if (resource.source_value == source_value && resource.role == role) { + return &resource; + } + } + return nullptr; +} + +const CommandPlanAlternateValue * CommandPlanMetadata::find_alternate_value(ValueId graph_value) const { + for (const CommandPlanAlternateValue & alternate : alternate_values_) { + if (alternate.graph_value == graph_value) { + return &alternate; + } + } + return nullptr; +} + +const CommandPlanAlternateValue * CommandPlanMetadata::find_alternate_value(ValueId graph_value, + ggml_type type, + size_t byte_count) const { + for (const CommandPlanAlternateValue & alternate : alternate_values_) { + if (alternate.graph_value == graph_value && alternate.type == type && alternate.byte_count == byte_count) { + return &alternate; + } + } + return nullptr; +} + +const CommandPlanMoeRoutingBundle * CommandPlanMetadata::find_moe_routing_bundle(ValueId route_ids) const { + for (const CommandPlanMoeRoutingBundle & bundle : moe_routing_bundles_) { + if (bundle.route_ids == route_ids) { + return &bundle; + } + } + return nullptr; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/command-plan-metadata.h b/ggml/src/ggml-hrx/dispatch/command-plan-metadata.h new file mode 100644 index 000000000000..0025b2b0fbbf --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/command-plan-metadata.h @@ -0,0 +1,137 @@ +#pragma once + +#include "dispatch.h" +#include "status.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { + +enum class GeneratedResourceRole { + MoeExpertTable, + MoePartitionTable, + F16K16Major, + Conv4Edges, +}; + +enum class CommandPlanResourceMetadataKind { + None, + MoeRoutingResource, +}; + +struct MoeRoutingResourceMetadata { + int64_t token_count = 0; + int64_t route_count = 0; + int64_t route_stride = 0; + int64_t expert_count = 0; +}; + +template constexpr CommandPlanResourceMetadataKind command_plan_resource_metadata_kind() { + static_assert(sizeof(T) == 0, "unsupported command plan resource metadata type"); + return CommandPlanResourceMetadataKind::None; +} + +template <> +constexpr CommandPlanResourceMetadataKind command_plan_resource_metadata_kind() { + return CommandPlanResourceMetadataKind::MoeRoutingResource; +} + +struct CommandPlanResourceMetadata { + static constexpr size_t kMaxBytes = 64; + + CommandPlanResourceMetadataKind kind = CommandPlanResourceMetadataKind::None; + size_t size = 0; + alignas(std::max_align_t) std::array bytes = {}; + + template bool read(T & value) const { + static_assert(std::is_trivially_copyable::value, "metadata payload must be trivially copyable"); + if (kind != command_plan_resource_metadata_kind() || size != sizeof(T)) { + return false; + } + std::memcpy(&value, bytes.data(), sizeof(T)); + return true; + } +}; + +template CommandPlanResourceMetadata make_command_plan_resource_metadata(const T & value) { + static_assert(std::is_trivially_copyable::value, "metadata payload must be trivially copyable"); + static_assert(sizeof(T) <= CommandPlanResourceMetadata::kMaxBytes, "metadata payload is too large"); + + CommandPlanResourceMetadata metadata; + metadata.kind = command_plan_resource_metadata_kind(); + metadata.size = sizeof(T); + std::memcpy(metadata.bytes.data(), &value, sizeof(T)); + return metadata; +} + +struct CommandPlanGeneratedResource { + ValueId source_value; + GeneratedResourceRole role = GeneratedResourceRole::MoeExpertTable; + ValueId generated_value; + size_t byte_count = 0; + CommandPlanResourceMetadata metadata; +}; + +struct CommandPlanAlternateValue { + ValueId graph_value; + ValueId alternate_value; + ggml_type type = GGML_TYPE_COUNT; + size_t byte_count = 0; + std::string name; +}; + +struct CommandPlanMoeRoutingBundle { + ValueId route_ids; + ValueId route_weights; + ValueId expert_table; + ValueId partition_table; + size_t expert_table_byte_count = 0; + size_t partition_table_byte_count = 0; + int64_t token_count = 0; + int64_t route_count = 0; + int64_t route_stride = 0; + int64_t expert_count = 0; +}; + +class CommandPlanMetadata { + public: + void clear(); + + bool append(CommandPlanMetadata && other, Status & status); + + bool append_generated_resource(CommandPlanGeneratedResource resource, Status & status); + + bool append_alternate_value(CommandPlanAlternateValue alternate, Status & status); + + bool append_moe_routing_bundle(CommandPlanMoeRoutingBundle bundle, Status & status); + + const CommandPlanGeneratedResource * find_generated_resource(ValueId source_value, + GeneratedResourceRole role) const; + + const CommandPlanAlternateValue * find_alternate_value(ValueId graph_value) const; + + const CommandPlanAlternateValue * find_alternate_value(ValueId graph_value, + ggml_type type, + size_t byte_count) const; + + const CommandPlanMoeRoutingBundle * find_moe_routing_bundle(ValueId route_ids) const; + + const std::vector & generated_resources() const { return generated_resources_; } + + const std::vector & alternate_values() const { return alternate_values_; } + + const std::vector & moe_routing_bundles() const { return moe_routing_bundles_; } + + private: + std::vector generated_resources_; + std::vector alternate_values_; + std::vector moe_routing_bundles_; +}; + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/command-plan.h b/ggml/src/ggml-hrx/dispatch/command-plan.h new file mode 100644 index 000000000000..a1f1abf62168 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/command-plan.h @@ -0,0 +1,115 @@ +#pragma once + +#include "command-plan-metadata.h" +#include "dispatch.h" +#include "graph/graph.h" +#include "status.h" + +#include +#include +#include +#include + +namespace ggml::hrx { + +struct CommandPlanTransient { + ValueId value; + std::string name; + size_t size = 0; + size_t alignment = 256; +}; + +struct CommandPlanConstantInitialization { + ValueId value; + std::string name; + size_t offset = 0; + std::vector data; +}; + +struct CommandPlanCompletionCounterRequest { + ValueId value; + std::string name; + uint32_t count = 0; +}; + +struct CommandPlan { + std::vector initialization_dispatches; + std::vector dispatches; + std::vector transients; + std::vector constant_initializations; + std::vector completion_counter_requests; + CommandPlanMetadata metadata; + Status status; + + bool valid() const { return status.success(); } +}; + +inline const CommandPlanAlternateValue * find_alternate_value(const CommandPlan & plan, ValueId graph_value) { + return plan.metadata.find_alternate_value(graph_value); +} + +inline const CommandPlanAlternateValue * find_alternate_value(const CommandPlan & plan, + ValueId graph_value, + ggml_type type, + size_t byte_count) { + return plan.metadata.find_alternate_value(graph_value, type, byte_count); +} + +inline bool same_full_value_range(const Value & lhs, const Value & rhs) { + return lhs.storage == rhs.storage && lhs.storage_offset == rhs.storage_offset && lhs.byte_count == rhs.byte_count; +} + +inline const CommandPlanAlternateValue * find_alternate_value(const Graph & graph, + const CommandPlan & plan, + ValueId graph_value, + ggml_type type, + size_t byte_count) { + const CommandPlanAlternateValue * exact = find_alternate_value(plan, graph_value, type, byte_count); + if (exact != nullptr) { + return exact; + } + + const Value * value = graph.values().find(graph_value); + if (value == nullptr) { + return nullptr; + } + + auto find_if_same_range = [&](ValueId candidate_id) -> const CommandPlanAlternateValue * { + const Value * candidate = graph.values().find(candidate_id); + if (candidate == nullptr || !same_full_value_range(*value, *candidate)) { + return nullptr; + } + return find_alternate_value(plan, candidate_id, type, byte_count); + }; + + ValueId alias = value->alias_source; + for (size_t i = 0; alias.value >= 0 && i < graph.values().size(); ++i) { + const CommandPlanAlternateValue * alternate = find_if_same_range(alias); + if (alternate != nullptr) { + return alternate; + } + const Value * alias_value = graph.values().find(alias); + if (alias_value == nullptr) { + break; + } + alias = alias_value->alias_source; + } + + const CommandPlanAlternateValue * root_alternate = find_if_same_range(value->storage_root); + if (root_alternate != nullptr) { + return root_alternate; + } + + for (const CommandPlanAlternateValue & alternate : plan.metadata.alternate_values()) { + if (alternate.type != type || alternate.byte_count != byte_count) { + continue; + } + const Value * alternate_value = graph.values().find(alternate.graph_value); + if (alternate_value != nullptr && same_full_value_range(*value, *alternate_value)) { + return &alternate; + } + } + return nullptr; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/command-program-bindings.cpp b/ggml/src/ggml-hrx/dispatch/command-program-bindings.cpp new file mode 100644 index 000000000000..741a41c9c17f --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/command-program-bindings.cpp @@ -0,0 +1,92 @@ +#include "command-program-bindings.h" + +#include +#include + +namespace ggml::hrx { +namespace { + +static void mix_hash(uint64_t & hash, uint64_t value) { + hash ^= value; + hash *= UINT64_C(1099511628211); +} + +} // namespace + +CommandProgramBindings CommandProgramBindings::from_value_map(const ValueMap & values) { + std::vector bindings; + CommandProgramBindings result; + for (const ValueId id : values.external_value_ids()) { + const Value * value = values.find(id); + if (value == nullptr) { + result.status.log("external value %d does not exist", id.value); + continue; + } + const std::optional buffer = values.resolve_buffer_binding(id); + if (!buffer.has_value()) { + result.status.log("external value %d is not bound", id.value); + continue; + } + bindings.push_back({ value->id, buffer->buffer, buffer->offset, buffer->length, buffer->identity, + buffer->generation, buffer->capacity, buffer->host_data, buffer->weight, + value->byte_count == 0 }); + } + return from_bindings(std::move(bindings), result.status); +} + +CommandProgramBindings CommandProgramBindings::from_bindings(std::vector bindings, + const Status & errors) { + CommandProgramBindings result; + result.status.append(errors); + result.bindings_ = std::move(bindings); + for (const CommandProgramBinding & binding : result.bindings_) { + if (binding.buffer == nullptr && binding.host_data == nullptr) { + result.status.log("external value %d has a null binding", binding.value.value); + } + if (binding.length == 0 && !binding.empty_value) { + result.status.log("external value %d has an empty binding", binding.value.value); + } + } + return result; +} + +const CommandProgramBinding * CommandProgramBindings::find(ValueId value) const { + for (const CommandProgramBinding & binding : bindings_) { + if (binding.value == value) { + return &binding; + } + } + return nullptr; +} + +CommandProgramBindingsHash command_program_bindings_hash(const CommandProgramBindings & bindings) { + uint64_t hash = UINT64_C(1469598103934665603); + mix_hash(hash, UINT64_C(0x6872782d62696e64)); + for (const CommandProgramBinding & binding : bindings.bindings()) { + mix_hash(hash, static_cast(static_cast(binding.value.value))); + mix_hash(hash, binding.host_data != nullptr ? 1 : 0); + mix_hash(hash, binding.identity); + mix_hash(hash, binding.generation); + mix_hash(hash, static_cast(binding.capacity)); + mix_hash(hash, static_cast(binding.offset)); + mix_hash(hash, static_cast(binding.length)); + mix_hash(hash, binding.weight ? 1 : 0); + mix_hash(hash, binding.empty_value ? 1 : 0); + } + mix_hash(hash, static_cast(bindings.bindings().size())); + return { hash }; +} + +CommandProgramBindingsFingerprint command_program_bindings_fingerprint(const CommandProgramBindings & bindings) { + std::ostringstream out; + out << "hrx-bindings-v1"; + for (const CommandProgramBinding & binding : bindings.bindings()) { + out << "|value=" << binding.value.value << "|kind=" << (binding.host_data != nullptr ? "host" : "device") + << "|identity=" << binding.identity << "|generation=" << binding.generation + << "|capacity=" << binding.capacity << "|offset=" << binding.offset << "|length=" << binding.length + << "|weight=" << (binding.weight ? 1 : 0) << "|empty=" << (binding.empty_value ? 1 : 0); + } + return { out.str() }; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/command-program-bindings.h b/ggml/src/ggml-hrx/dispatch/command-program-bindings.h new file mode 100644 index 000000000000..45ef75e96cd0 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/command-program-bindings.h @@ -0,0 +1,63 @@ +#pragma once + +#include "graph/value-map.h" +#include "status.h" + +#include +#include +#include +#include + +namespace ggml::hrx { + +struct CommandProgramBinding { + // A buffer is directly bindable by an HRX command program. Host data requires residency or staging before + // execution. These are alternate storage forms and should not both be populated. + ValueId value; + hrx_buffer_t buffer = nullptr; + size_t offset = 0; + size_t length = 0; + uint64_t identity = 0; + uint64_t generation = 0; + size_t capacity = 0; + void * host_data = nullptr; + bool weight = false; + bool empty_value = false; + // ggml graph input (GGML_TENSOR_FLAG_INPUT). A kernel may rewrite its device copy in place (the Qwen + // attention metadata kernel regenerates positions, cache indices and the mask), but that copy is never + // written back to the caller's host tensor: ggml reuses that memory for the next graph's inputs. + bool graph_input = false; + + bool requires_materialization() const { return host_data != nullptr; } +}; + +struct CommandProgramBindingsFingerprint { + std::string value; +}; + +struct CommandProgramBindingsHash { + uint64_t value = 0; +}; + +class CommandProgramBindings { + public: + static CommandProgramBindings from_value_map(const ValueMap & values); + static CommandProgramBindings from_bindings(std::vector bindings, + const Status & errors = {}); + + const CommandProgramBinding * find(ValueId value) const; + + const std::vector & bindings() const { return bindings_; } + + bool valid() const { return status.success(); } + + Status status; + + private: + std::vector bindings_; +}; + +CommandProgramBindingsHash command_program_bindings_hash(const CommandProgramBindings & bindings); +CommandProgramBindingsFingerprint command_program_bindings_fingerprint(const CommandProgramBindings & bindings); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/command-program-diagnostics.cpp b/ggml/src/ggml-hrx/dispatch/command-program-diagnostics.cpp new file mode 100644 index 000000000000..8fae265b6e4c --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/command-program-diagnostics.cpp @@ -0,0 +1,97 @@ +#include "command-program-diagnostics.h" + +#include + +namespace ggml::hrx { +namespace { + +template static std::string unknown_enum_name(Enum value) { + std::ostringstream out; + out << "Unknown(" << static_cast(value) << ")"; + return out.str(); +} + +static const char * binding_name(const CommandBinding & binding) { + return binding.name.empty() ? "" : binding.name.c_str(); +} + +} // namespace + +std::string command_kind_name(CommandKind kind) { + switch (kind) { + case CommandKind::Invalid: + return "Invalid"; + case CommandKind::Kernel: + return "Kernel"; + } + return unknown_enum_name(kind); +} + +std::string command_binding_origin_name(CommandBindingOrigin origin) { + switch (origin) { + case CommandBindingOrigin::GraphValue: + return "GraphValue"; + case CommandBindingOrigin::Transient: + return "Transient"; + case CommandBindingOrigin::ProgramConstant: + return "ProgramConstant"; + } + return unknown_enum_name(origin); +} + +std::string resource_access_name(ResourceAccess access) { + switch (access) { + case ResourceAccess::Read: + return "Read"; + case ResourceAccess::Write: + return "Write"; + case ResourceAccess::ReadWrite: + return "ReadWrite"; + } + return unknown_enum_name(access); +} + +std::string format_command_binding(const CommandBinding & binding) { + std::ostringstream out; + out << "binding " << binding_name(binding) << " value=" << binding.value.value + << " origin=" << command_binding_origin_name(binding.origin) + << " access=" << resource_access_name(binding.access) << " range=[" << binding.offset << ", " + << binding.offset + binding.length << ")"; + if (binding.layout != kNativeWeightLayout) { + out << " layout=" << binding.layout << " source_type=" << static_cast(binding.source_type) + << " input_size=" << binding.input_size << " output_size=" << binding.output_size + << " source_length=" << binding.source_length; + } + return out.str(); +} + +std::string format_command(const Command & command) { + std::ostringstream out; + out << "command " << command.ordinal << " kind=" << command_kind_name(command.kind) + << " kernel_id=" << command.kernel.kernel_id << " bindings=" << command.bindings.size() + << " deps=" << command.dependencies.size(); + return out.str(); +} + +std::string format_command_program(const CommandProgram & program) { + std::ostringstream out; + out << "command_program commands=" << program.commands.size() + << " transient_arena=" << program.transients.arena_size + << " transient_allocations=" << program.transients.allocations.size() + << " completion_counters=" << program.completion_counters.count << " completion_counter_range=[" + << program.completion_counters.arena_offset << ", " + << program.completion_counters.arena_offset + program.completion_counters.byte_count << ")"; + for (const Command & command : program.commands) { + out << '\n' << format_command(command); + for (const CommandBinding & binding : command.bindings) { + out << "\n " << format_command_binding(binding); + } + } + for (const TransientAllocation & allocation : program.transients.allocations) { + out << "\ntransient value=" << allocation.value.value << " range=[" << allocation.arena_offset << ", " + << allocation.arena_offset + allocation.size << ") alignment=" << allocation.alignment; + } + return out.str(); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/command-program-diagnostics.h b/ggml/src/ggml-hrx/dispatch/command-program-diagnostics.h new file mode 100644 index 000000000000..c961db905115 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/command-program-diagnostics.h @@ -0,0 +1,17 @@ +#pragma once + +#include "command-program.h" + +#include + +namespace ggml::hrx { + +std::string command_kind_name(CommandKind kind); +std::string command_binding_origin_name(CommandBindingOrigin origin); +std::string resource_access_name(ResourceAccess access); + +std::string format_command_binding(const CommandBinding & binding); +std::string format_command(const Command & command); +std::string format_command_program(const CommandProgram & program); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/command-program-dump.cpp b/ggml/src/ggml-hrx/dispatch/command-program-dump.cpp new file mode 100644 index 000000000000..24e793561970 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/command-program-dump.cpp @@ -0,0 +1,347 @@ +#include "command-program-dump.h" + +#include "command-program-diagnostics.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr const char * kDumpDirectoryEnvironment = "GGML_HRX_DUMP_COMMAND_PROGRAM_DIR"; + +static void mix_hash(uint64_t & hash, uint64_t value) { + hash ^= value; + hash *= UINT64_C(1099511628211); +} + +static uint64_t hash_string(const std::string & text) { + uint64_t hash = UINT64_C(1469598103934665603); + for (const char c : text) { + mix_hash(hash, static_cast(c)); + } + return hash; +} + +static std::string hex_u64(uint64_t value) { + std::ostringstream out; + out << std::hex << std::setw(16) << std::setfill('0') << value; + return out.str(); +} + +static std::string dot_escape(const std::string & text) { + std::string escaped; + escaped.reserve(text.size()); + for (const char c : text) { + if (c == '\\' || c == '"') { + escaped.push_back('\\'); + } + if (c == '\n') { + escaped += "\\n"; + } else { + escaped.push_back(c); + } + } + return escaped; +} + +static std::string json_escape(const std::string & text) { + std::string escaped; + escaped.reserve(text.size()); + for (const unsigned char c : text) { + switch (c) { + case '\\': + escaped += "\\\\"; + break; + case '"': + escaped += "\\\""; + break; + case '\b': + escaped += "\\b"; + break; + case '\f': + escaped += "\\f"; + break; + case '\n': + escaped += "\\n"; + break; + case '\r': + escaped += "\\r"; + break; + case '\t': + escaped += "\\t"; + break; + default: + if (c < 0x20) { + static constexpr char kHex[] = "0123456789abcdef"; + escaped += "\\u00"; + escaped.push_back(kHex[c >> 4]); + escaped.push_back(kHex[c & 0xf]); + } else { + escaped.push_back(static_cast(c)); + } + break; + } + } + return escaped; +} + +static std::string kernel_name(const KernelCorpus & corpus, const std::string & target, uint64_t kernel_id) { + const KernelResolveResult resolved = resolve_kernel_definition(corpus, target, kernel_id); + if (resolved.found()) { + return kernel_definition_name(*resolved.definition); + } + std::ostringstream out; + out << "kernel_id=" << kernel_id; + return out.str(); +} + +static void append_json_integer_map(std::ostringstream & out, const std::map & values) { + out << '{'; + size_t index = 0; + for (const auto & value : values) { + if (index++ > 0) { + out << ", "; + } + out << '"' << json_escape(value.first) << "\": " << value.second; + } + out << '}'; +} + +static void append_json_string_map(std::ostringstream & out, const std::map & values) { + out << '{'; + size_t index = 0; + for (const auto & value : values) { + if (index++ > 0) { + out << ", "; + } + out << '"' << json_escape(value.first) << "\": \"" << json_escape(value.second) << '"'; + } + out << '}'; +} + +static void append_kernel_list(std::ostringstream & out, + const std::vector & commands, + const char * phase, + const KernelCorpus & corpus, + const std::string & target) { + for (const Command & command : commands) { + out << phase << ' ' << command.ordinal << ' ' << kernel_name(corpus, target, command.kernel.kernel_id); + if (!command.dependencies.empty()) { + out << " deps="; + for (size_t i = 0; i < command.dependencies.size(); ++i) { + if (i > 0) { + out << ','; + } + out << command.dependencies[i]; + } + } + out << '\n'; + } +} + +static std::string format_kernel_list(const CommandProgram & program, + const KernelCorpus & corpus, + const std::string & target) { + std::ostringstream out; + append_kernel_list(out, program.initialization_commands, "init", corpus, target); + append_kernel_list(out, program.commands, "main", corpus, target); + return out.str(); +} + +static void append_dot_nodes(std::ostringstream & out, + const std::vector & commands, + const char * prefix, + const char * phase, + const KernelCorpus & corpus, + const std::string & target) { + for (const Command & command : commands) { + const std::string name = kernel_name(corpus, target, command.kernel.kernel_id); + out << " " << prefix << command.ordinal << " [label=\"" << phase << ' ' << command.ordinal << "\\n" + << dot_escape(name) << "\"];\n"; + } +} + +static std::string format_kernel_dot(const CommandProgram & program, + const KernelCorpus & corpus, + const std::string & target) { + std::ostringstream out; + out << "digraph hrx_kernel_invocations {\n"; + out << " rankdir=LR;\n"; + append_dot_nodes(out, program.initialization_commands, "init", "init", corpus, target); + append_dot_nodes(out, program.commands, "main", "main", corpus, target); + for (const Command & command : program.commands) { + for (const uint32_t dependency : command.dependencies) { + out << " main" << dependency << " -> main" << command.ordinal << ";\n"; + } + } + out << "}\n"; + return out.str(); +} + +static void append_json_command_list(std::ostringstream & out, + const std::vector & commands, + const char * phase, + const KernelCorpus & corpus, + const std::string & target, + const char * indent) { + for (size_t command_index = 0; command_index < commands.size(); ++command_index) { + const Command & command = commands[command_index]; + if (command_index > 0) { + out << ",\n"; + } + const std::string name = kernel_name(corpus, target, command.kernel.kernel_id); + out << indent << "{\n"; + out << indent << " \"phase\": \"" << phase << "\",\n"; + out << indent << " \"ordinal\": " << command.ordinal << ",\n"; + out << indent << " \"kind\": \"" << json_escape(command_kind_name(command.kind)) << "\",\n"; + out << indent << " \"kernel_id\": " << command.kernel.kernel_id << ",\n"; + out << indent << " \"kernel\": \"" << json_escape(name) << "\",\n"; + out << indent << " \"integer_parameters\": "; + append_json_integer_map(out, command.kernel.integer_parameters); + out << ",\n"; + out << indent << " \"compile_parameters\": "; + append_json_string_map(out, command.kernel.compile_parameters); + out << ",\n"; + out << indent << " \"dependencies\": ["; + for (size_t i = 0; i < command.dependencies.size(); ++i) { + if (i > 0) { + out << ", "; + } + out << command.dependencies[i]; + } + out << "],\n"; + out << indent << " \"bindings\": [\n"; + for (size_t binding_index = 0; binding_index < command.bindings.size(); ++binding_index) { + const CommandBinding & binding = command.bindings[binding_index]; + if (binding_index > 0) { + out << ",\n"; + } + out << indent << " {\n"; + out << indent << " \"index\": " << binding_index << ",\n"; + out << indent << " \"name\": \"" << json_escape(binding.name) << "\",\n"; + out << indent << " \"value\": " << binding.value.value << ",\n"; + out << indent << " \"origin\": \"" << json_escape(command_binding_origin_name(binding.origin)) << "\",\n"; + out << indent << " \"access\": \"" << json_escape(resource_access_name(binding.access)) << "\",\n"; + out << indent << " \"offset\": " << binding.offset << ",\n"; + out << indent << " \"length\": " << binding.length << "\n"; + out << indent << " }"; + } + out << '\n' << indent << " ]\n"; + out << indent << '}'; + } +} + +static std::string format_kernel_json(const CommandProgram & program, + const KernelCorpus & corpus, + const std::string & target, + const std::string & command_shape, + uint64_t dump_id, + const std::string & shape_hash) { + std::ostringstream out; + out << "{\n"; + out << " \"schema\": \"ggml-hrx-command-program-v1\",\n"; + out << " \"dump_id\": " << dump_id << ",\n"; + out << " \"shape_hash\": \"" << json_escape(shape_hash) << "\",\n"; + out << " \"target\": \"" << json_escape(target) << "\",\n"; + out << " \"command_shape\": \"" << json_escape(command_shape) << "\",\n"; + out << " \"transients\": {\n"; + out << " \"arena_size\": " << program.transients.arena_size << ",\n"; + out << " \"arena_alignment\": " << program.transients.arena_alignment << ",\n"; + out << " \"allocations\": [\n"; + for (size_t i = 0; i < program.transients.allocations.size(); ++i) { + const TransientAllocation & allocation = program.transients.allocations[i]; + if (i > 0) { + out << ",\n"; + } + out << " {\"value\": " << allocation.value.value << ", \"size\": " << allocation.size + << ", \"alignment\": " << allocation.alignment << ", \"arena_offset\": " << allocation.arena_offset + << '}'; + } + out << "\n ]\n"; + out << " },\n"; + out << " \"completion_counters\": {\n"; + out << " \"arena_offset\": " << program.completion_counters.arena_offset << ",\n"; + out << " \"byte_count\": " << program.completion_counters.byte_count << ",\n"; + out << " \"count\": " << program.completion_counters.count << "\n"; + out << " },\n"; + out << " \"commands\": [\n"; + bool wrote_any = false; + if (!program.initialization_commands.empty()) { + append_json_command_list(out, program.initialization_commands, "init", corpus, target, " "); + wrote_any = true; + } + if (!program.commands.empty()) { + if (wrote_any) { + out << ",\n"; + } + append_json_command_list(out, program.commands, "main", corpus, target, " "); + } + out << "\n ]\n"; + out << "}\n"; + return out.str(); +} + +static void write_file(const std::filesystem::path & path, const std::string & contents) { + std::ofstream output(path, std::ios::binary | std::ios::trunc); + if (!output) { + throw std::runtime_error("cannot create " + path.string()); + } + output << contents; + if (contents.empty() || contents.back() != '\n') { + output << '\n'; + } +} + +} // namespace + +Status dump_command_program_kernels_if_requested(const CommandProgram & program, + const KernelCorpus & corpus, + const std::string & target, + const std::string & command_shape) { + Status status; + const char * directory_value = std::getenv(kDumpDirectoryEnvironment); + if (directory_value == nullptr || directory_value[0] == '\0' || + (directory_value[0] == '0' && directory_value[1] == '\0')) { + return status; + } + + static std::mutex mutex; + static std::unordered_set dumped_shapes; + static std::atomic sequence{ 0 }; + + const std::filesystem::path root_directory(directory_value); + const std::string dump_key = root_directory.string() + '\n' + target + '\n' + command_shape; + uint64_t dump_id = 0; + { + std::lock_guard lock(mutex); + if (!dumped_shapes.insert(dump_key).second) { + return status; + } + dump_id = sequence.fetch_add(1); + } + + try { + const uint64_t shape_hash = hash_string(dump_key); + const std::string shape_hex = hex_u64(shape_hash); + std::ostringstream name; + name << "program-" << dump_id << "-shape-" << shape_hex; + const std::filesystem::path directory = root_directory / name.str(); + std::filesystem::create_directories(directory); + write_file(directory / "kernels.txt", format_kernel_list(program, corpus, target)); + write_file(directory / "kernels.dot", format_kernel_dot(program, corpus, target)); + write_file(directory / "program.json", format_kernel_json(program, corpus, target, command_shape, dump_id, shape_hex)); + } catch (const std::exception & error) { + status.log("failed to write HRX command program kernel dump: %s", error.what()); + } + return status; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/command-program-dump.h b/ggml/src/ggml-hrx/dispatch/command-program-dump.h new file mode 100644 index 000000000000..f0d911ebdd7f --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/command-program-dump.h @@ -0,0 +1,16 @@ +#pragma once + +#include "command-program.h" +#include "kernel-corpus/kernel-corpus.h" +#include "status.h" + +#include + +namespace ggml::hrx { + +Status dump_command_program_kernels_if_requested(const CommandProgram & program, + const KernelCorpus & corpus, + const std::string & target, + const std::string & command_shape); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/command-program-resolver.cpp b/ggml/src/ggml-hrx/dispatch/command-program-resolver.cpp new file mode 100644 index 000000000000..18805390c903 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/command-program-resolver.cpp @@ -0,0 +1,128 @@ +#include "command-program-resolver.h" + +#include "command-program-diagnostics.h" + +#include + +namespace ggml::hrx { +namespace { + +static Status resolve_command_binding(const Command & command, + const CommandProgram & program, + const CommandBinding & binding, + const CommandProgramBindings & bindings, + const TransientArenaAllocationRef * transient_arena, + ResolvedBufferRef & ref) { + Status status; + const std::string command_context = format_command(command); + const std::string binding_context = format_command_binding(binding); + if (binding.length == 0) { + status.log("%s %s has an empty range", command_context.c_str(), binding_context.c_str()); + return status; + } + switch (binding.origin) { + case CommandBindingOrigin::GraphValue: + { + const CommandProgramBinding * concrete = bindings.find(binding.value); + if (concrete == nullptr) { + status.log("%s %s is not bound", command_context.c_str(), binding_context.c_str()); + return status; + } + if (concrete->buffer == nullptr) { + status.log("%s %s has a null buffer", command_context.c_str(), binding_context.c_str()); + return status; + } + if (binding.offset > concrete->length || binding.length > concrete->length - binding.offset) { + status.log("%s %s is outside runtime binding length %zu", command_context.c_str(), + binding_context.c_str(), concrete->length); + return status; + } + ref = { concrete->buffer, concrete->offset + binding.offset, binding.length }; + return status; + } + case CommandBindingOrigin::Transient: + { + const TransientAllocation * allocation = find_transient_allocation(program.transients, binding.value); + if (allocation == nullptr) { + status.log("%s %s has no transient allocation", command_context.c_str(), binding_context.c_str()); + return status; + } + if (transient_arena == nullptr || transient_arena->buffer == nullptr) { + status.log("%s %s has no transient arena", command_context.c_str(), binding_context.c_str()); + return status; + } + if (transient_arena->allocation_id == kInvalidTransientArenaAllocationId) { + status.log("%s %s has no transient arena allocation id", command_context.c_str(), + binding_context.c_str()); + return status; + } + if (program.transients.arena_size > transient_arena->capacity) { + status.log("%s %s requires transient arena size %zu but only %zu bytes are available", + command_context.c_str(), binding_context.c_str(), program.transients.arena_size, + transient_arena->capacity); + return status; + } + if (binding.offset > allocation->size || binding.length > allocation->size - binding.offset) { + status.log("%s %s is outside transient allocation length %zu", command_context.c_str(), + binding_context.c_str(), allocation->size); + return status; + } + ref = { transient_arena->buffer, allocation->arena_offset + binding.offset, binding.length }; + return status; + } + case CommandBindingOrigin::ProgramConstant: + break; + } + status.log("%s %s has an unsupported binding origin", command_context.c_str(), binding_context.c_str()); + return status; +} + +} // namespace + +static void resolve_command_list(const CommandProgram & program, + const std::vector & commands, + const CommandProgramBindings & bindings, + const TransientArenaAllocationRef * transient_arena, + std::vector & resolved_commands, + Status & status) { + resolved_commands.reserve(commands.size()); + for (const Command & command : commands) { + ResolvedCommand resolved_command; + resolved_command.ordinal = command.ordinal; + resolved_command.kind = command.kind; + resolved_command.kernel = command.kernel; + resolved_command.bindings.reserve(command.bindings.size()); + + for (const CommandBinding & binding : command.bindings) { + ResolvedCommandBinding resolved_binding; + resolved_binding.binding = binding; + Status binding_status = + resolve_command_binding(command, program, binding, bindings, transient_arena, resolved_binding.ref); + if (binding_status.success()) { + resolved_command.bindings.push_back(resolved_binding); + } else { + status.append(binding_status); + } + } + resolved_commands.push_back(resolved_command); + } +} + +ResolvedCommandProgram resolve_command_program_bindings(const CommandProgram & program, + const CommandProgramBindings & bindings, + const TransientArenaAllocationRef * transient_arena) { + ResolvedCommandProgram result; + if (!program.valid()) { + result.status.append(program.status); + } + if (!bindings.valid()) { + result.status.append(bindings.status); + } + + resolve_command_list(program, program.initialization_commands, bindings, transient_arena, + result.initialization_commands, result.status); + resolve_command_list(program, program.commands, bindings, transient_arena, result.commands, result.status); + return result; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/command-program-resolver.h b/ggml/src/ggml-hrx/dispatch/command-program-resolver.h new file mode 100644 index 000000000000..b4a517fbb0eb --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/command-program-resolver.h @@ -0,0 +1,51 @@ +#pragma once + +#include "command-program-bindings.h" +#include "command-program.h" +#include "status.h" + +#include +#include +#include + +namespace ggml::hrx { + +struct ResolvedBufferRef { + hrx_buffer_t buffer = nullptr; + size_t offset = 0; + size_t length = 0; +}; + +static constexpr uint64_t kInvalidTransientArenaAllocationId = 0; + +struct TransientArenaAllocationRef { + hrx_buffer_t buffer = nullptr; + size_t capacity = 0; + uint64_t allocation_id = kInvalidTransientArenaAllocationId; +}; + +struct ResolvedCommandBinding { + CommandBinding binding; + ResolvedBufferRef ref; +}; + +struct ResolvedCommand { + uint32_t ordinal = 0; + CommandKind kind = CommandKind::Kernel; + KernelSpecialization kernel; + std::vector bindings; +}; + +struct ResolvedCommandProgram { + std::vector initialization_commands; + std::vector commands; + Status status; + + bool valid() const { return status.success(); } +}; + +ResolvedCommandProgram resolve_command_program_bindings(const CommandProgram & program, + const CommandProgramBindings & bindings, + const TransientArenaAllocationRef * transient_arena = nullptr); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/command-program.cpp b/ggml/src/ggml-hrx/dispatch/command-program.cpp new file mode 100644 index 000000000000..f789c5b09e62 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/command-program.cpp @@ -0,0 +1,540 @@ +#include "command-program.h" + +#include "command-program-diagnostics.h" +#include "transient-allocator.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static bool string_equal(const char * lhs, const char * rhs) { + return std::strcmp(lhs != nullptr ? lhs : "", rhs != nullptr ? rhs : "") == 0; +} + +static std::string string_value(const char * value) { + return value != nullptr ? value : ""; +} + +static CommandBindingOrigin command_binding_origin(const Graph & graph, const Value & value) { + if (value.kind == ValueKind::External) { + return CommandBindingOrigin::GraphValue; + } + const Value * root = graph.values().find(value.storage_root); + switch (root != nullptr ? root->kind : value.kind) { + case ValueKind::External: + return CommandBindingOrigin::GraphValue; + case ValueKind::Transient: + return CommandBindingOrigin::Transient; + } + return CommandBindingOrigin::GraphValue; +} + +struct StorageBindingTarget { + ValueId value; + size_t offset = 0; +}; + +static bool resource_access_writes(ResourceAccess access) { + return access == ResourceAccess::Write || access == ResourceAccess::ReadWrite; +} + +static bool byte_range_is_covered(size_t range_offset, size_t range_length, size_t cover_offset, size_t cover_length) { + if (range_offset < cover_offset) { + return false; + } + const size_t relative_offset = range_offset - cover_offset; + return relative_offset <= cover_length && range_length <= cover_length - relative_offset; +} + +static const Value * find_external_storage_binding_target(const Graph & graph, + const Value & source, + size_t binding_offset, + size_t binding_length) { + if (source.storage.value < 0 || binding_length > source.byte_count || + binding_offset > source.byte_count - binding_length) { + return nullptr; + } + + if (binding_offset > std::numeric_limits::max() - source.storage_offset) { + return nullptr; + } + const size_t range_offset = source.storage_offset + binding_offset; + const Value * best = nullptr; + for (const Value & candidate : graph.values().values()) { + if (candidate.kind != ValueKind::External || candidate.storage != source.storage) { + continue; + } + if (!byte_range_is_covered(range_offset, binding_length, candidate.storage_offset, candidate.byte_count)) { + continue; + } + if (best == nullptr || candidate.byte_count < best->byte_count) { + best = &candidate; + } + } + return best; +} + +static StorageBindingTarget storage_binding_target(const Graph & graph, + ValueId value, + size_t binding_offset, + size_t binding_length, + ResourceAccess access) { + StorageBindingTarget target; + target.value = value; + const Value * graph_value = graph.values().find(value); + if (graph_value == nullptr || graph_value->kind != ValueKind::Transient) { + return target; + } + if (resource_access_writes(access)) { + const Value * external_target = + find_external_storage_binding_target(graph, *graph_value, binding_offset, binding_length); + if (external_target != nullptr) { + target.value = external_target->id; + target.offset = graph_value->storage_offset + binding_offset - external_target->storage_offset; + return target; + } + } + const Value * root = graph.values().find(graph_value->storage_root); + if (root == nullptr) { + return target; + } + target.value = root->id; + target.offset = graph_value->storage_offset; + return target; +} + +static const CommandPlanTransient * find_plan_transient(const CommandPlan & plan, ValueId value) { + const auto found = std::find_if(plan.transients.begin(), plan.transients.end(), + [&](const CommandPlanTransient & transient) { return transient.value == value; }); + return found == plan.transients.end() ? nullptr : &*found; +} + +static const CommandPlanCompletionCounterRequest * find_plan_completion_counter_request(const CommandPlan & plan, + ValueId value) { + const auto found = + std::find_if(plan.completion_counter_requests.begin(), plan.completion_counter_requests.end(), + [&](const CommandPlanCompletionCounterRequest & request) { return request.value == value; }); + return found == plan.completion_counter_requests.end() ? nullptr : &*found; +} + +static void append_command(const Graph & graph, + const CommandPlan & plan, + const KernelCorpus & corpus, + const std::string & target, + const Dispatch & dispatch, + bool linear_dependency, + std::vector & commands, + Status & status) { + Command command; + command.ordinal = static_cast(commands.size()); + command.kind = CommandKind::Kernel; + command.kernel = dispatch.kernel; + // TODO: replace this linear ordinal dependency with real graph/resource dependency analysis. + if (linear_dependency && command.ordinal > 0) { + command.dependencies.push_back(command.ordinal - 1); + } + const KernelResolveResult resolved = resolve_kernel_definition(corpus, target, command.kernel.kernel_id); + const KernelDefinition * definition = resolved.definition; + if (!resolved.found()) { + status.log("%s", format_kernel_resolve_error(resolved, command.kernel.kernel_id).c_str()); + definition = nullptr; + } else if (dispatch.bindings.size() != definition->bindings.size()) { + status.log("command %u kernel %s has %zu bindings but its ABI requires %zu", command.ordinal, + kernel_definition_name(*definition).c_str(), dispatch.bindings.size(), definition->bindings.size()); + } + command.bindings.reserve(dispatch.bindings.size()); + for (size_t binding_index = 0; binding_index < dispatch.bindings.size(); ++binding_index) { + const DispatchBinding & binding = dispatch.bindings[binding_index]; + CommandBinding command_binding; + command_binding.value = binding.value; + command_binding.offset = binding.offset; + command_binding.length = binding.length; + command_binding.layout = binding.layout; + command_binding.source_type = binding.source_type; + command_binding.input_size = binding.input_size; + command_binding.output_size = binding.output_size; + command_binding.source_length = binding.source_length; + const Value * value = graph.values().find(command_binding.value); + const CommandPlanTransient * plan_transient = find_plan_transient(plan, command_binding.value); + const CommandPlanCompletionCounterRequest * completion_counter = + find_plan_completion_counter_request(plan, command_binding.value); + if (value == nullptr && plan_transient == nullptr && completion_counter == nullptr) { + status.log("command %u binding %zu references missing value %d", command.ordinal, binding_index, + command_binding.value.value); + } + if (definition != nullptr && binding_index < definition->bindings.size()) { + command_binding.name = string_value(definition->bindings[binding_index].name); + command_binding.access = definition->bindings[binding_index].access; + } + const StorageBindingTarget binding_target = storage_binding_target( + graph, command_binding.value, command_binding.offset, command_binding.length, command_binding.access); + command_binding.value = binding_target.value; + if (binding_target.offset > 0) { + if (binding_target.offset > std::numeric_limits::max() - command_binding.offset) { + status.log("command %u binding %zu storage alias offset overflows", command.ordinal, binding_index); + } else { + command_binding.offset += binding_target.offset; + } + } + const Value * target_value = graph.values().find(command_binding.value); + if (target_value != nullptr) { + command_binding.origin = command_binding_origin(graph, *target_value); + } else if (plan_transient != nullptr || completion_counter != nullptr) { + command_binding.origin = CommandBindingOrigin::Transient; + } + command.bindings.push_back(std::move(command_binding)); + } + commands.push_back(std::move(command)); +} + +static void verify_command_list(const std::vector & commands, + const TransientPlan & transients, + const KernelCorpus & corpus, + const std::string & target, + Status & status) { + for (size_t i = 0; i < commands.size(); ++i) { + const Command & command = commands[i]; + const std::string command_context = format_command(command); + if (command.ordinal != i) { + status.log("%s has non-contiguous ordinal at index %zu", command_context.c_str(), i); + } + if (command.kind != CommandKind::Kernel) { + status.log("%s is not a kernel command", command_context.c_str()); + } + KernelResolveResult resolved; + const KernelDefinition * definition = nullptr; + if (command.kind == CommandKind::Kernel) { + resolved = resolve_kernel_definition(corpus, target, command.kernel.kernel_id); + definition = resolved.definition; + } + if (command.kind == CommandKind::Kernel && !resolved.found()) { + status.log("%s: %s", command_context.c_str(), + format_kernel_resolve_error(resolved, command.kernel.kernel_id).c_str()); + } else if (definition != nullptr) { + if (command.bindings.size() != definition->bindings.size()) { + status.log("%s kernel %s has %zu bindings but its ABI requires %zu", command_context.c_str(), + kernel_definition_name(*definition).c_str(), command.bindings.size(), + definition->bindings.size()); + } + const size_t shared_count = std::min(command.bindings.size(), definition->bindings.size()); + for (size_t binding_index = 0; binding_index < shared_count; ++binding_index) { + const CommandBinding & binding = command.bindings[binding_index]; + const KernelBindingDefinition & abi = definition->bindings[binding_index]; + if (!string_equal(binding.name.c_str(), abi.name) || binding.access != abi.access) { + status.log("%s %s does not match ABI binding %zu", command_context.c_str(), + format_command_binding(binding).c_str(), binding_index); + } + } + } + if (command.bindings.empty()) { + status.log("%s has no bindings", command_context.c_str()); + } + for (uint32_t dependency : command.dependencies) { + if (dependency >= command.ordinal) { + status.log("%s has forward dependency %u", command_context.c_str(), dependency); + } + } + for (const CommandBinding & binding : command.bindings) { + const std::string binding_context = format_command_binding(binding); + if (binding.origin != CommandBindingOrigin::GraphValue && + binding.origin != CommandBindingOrigin::Transient) { + status.log("%s %s has an unsupported binding origin", command_context.c_str(), binding_context.c_str()); + } + if (binding.origin == CommandBindingOrigin::Transient) { + const TransientAllocation * allocation = find_transient_allocation(transients, binding.value); + if (allocation == nullptr) { + status.log("%s %s has no transient allocation", command_context.c_str(), binding_context.c_str()); + } else if (binding.offset > allocation->size || binding.length > allocation->size - binding.offset) { + status.log("%s %s is outside transient allocation length %zu", command_context.c_str(), + binding_context.c_str(), allocation->size); + } + } + if (binding.value.value < 0) { + status.log("%s %s has an invalid value id", command_context.c_str(), binding_context.c_str()); + } + if (binding.length == 0) { + status.log("%s %s has an empty binding", command_context.c_str(), binding_context.c_str()); + } + } + } +} + +struct VerifyTransientLifetime { + bool reserved = false; + bool has_lifetime = false; + uint32_t first_command = 0; + uint32_t last_command = 0; +}; + +static size_t saturated_range_end(size_t offset, size_t size) { + if (offset > std::numeric_limits::max() - size) { + return std::numeric_limits::max(); + } + return offset + size; +} + +static bool allocation_overlaps_region(const TransientAllocation & allocation, size_t offset, size_t size) { + if (size == 0) { + return false; + } + return allocation.arena_offset < saturated_range_end(offset, size) && + offset < saturated_range_end(allocation.arena_offset, allocation.size); +} + +static std::unordered_map collect_verify_transient_lifetimes( + const CommandProgram & program) { + std::unordered_map lifetimes; + lifetimes.reserve(program.transients.allocations.size()); + for (const TransientAllocation & allocation : program.transients.allocations) { + VerifyTransientLifetime & lifetime = lifetimes[allocation.value.value]; + if (allocation_overlaps_region(allocation, program.completion_counters.arena_offset, + program.completion_counters.byte_count)) { + lifetime.reserved = true; + } + } + for (const Command & command : program.initialization_commands) { + for (const CommandBinding & binding : command.bindings) { + if (binding.origin == CommandBindingOrigin::Transient) { + lifetimes[binding.value.value].reserved = true; + } + } + } + for (const ConstantInitialization & initialization : program.constant_initializations) { + lifetimes[initialization.value.value].reserved = true; + } + for (const Command & command : program.commands) { + for (const CommandBinding & binding : command.bindings) { + if (binding.origin != CommandBindingOrigin::Transient) { + continue; + } + VerifyTransientLifetime & lifetime = lifetimes[binding.value.value]; + if (lifetime.has_lifetime) { + lifetime.first_command = std::min(lifetime.first_command, command.ordinal); + lifetime.last_command = std::max(lifetime.last_command, command.ordinal); + } else { + lifetime.has_lifetime = true; + lifetime.first_command = command.ordinal; + lifetime.last_command = command.ordinal; + } + } + } + return lifetimes; +} + +static bool verify_transient_allocations_can_overlap( + const std::unordered_map & lifetimes, + const TransientAllocation & lhs, + const TransientAllocation & rhs) { + const auto lhs_lifetime = lifetimes.find(lhs.value.value); + const auto rhs_lifetime = lifetimes.find(rhs.value.value); + if (lhs_lifetime == lifetimes.end() || rhs_lifetime == lifetimes.end() || lhs_lifetime->second.reserved || + rhs_lifetime->second.reserved || !lhs_lifetime->second.has_lifetime || !rhs_lifetime->second.has_lifetime) { + return false; + } + return lhs_lifetime->second.last_command < rhs_lifetime->second.first_command || + rhs_lifetime->second.last_command < lhs_lifetime->second.first_command; +} + +static void verify_transient_allocations(const CommandProgram & program, Status & status) { + const std::unordered_map lifetimes = collect_verify_transient_lifetimes(program); + std::vector allocations_by_offset; + allocations_by_offset.reserve(program.transients.allocations.size()); + for (const TransientAllocation & allocation : program.transients.allocations) { + if (allocation.value.value < 0 || allocation.size == 0 || allocation.alignment == 0 || + allocation.arena_offset % allocation.alignment != 0 || + allocation.arena_offset > std::numeric_limits::max() - allocation.size || + allocation.arena_offset + allocation.size > program.transients.arena_size) { + status.log("invalid transient allocation for value %d", allocation.value.value); + } + allocations_by_offset.push_back(&allocation); + } + std::sort(allocations_by_offset.begin(), allocations_by_offset.end(), + [](const TransientAllocation * lhs, const TransientAllocation * rhs) { + if (lhs->arena_offset != rhs->arena_offset) { + return lhs->arena_offset < rhs->arena_offset; + } + return lhs->value.value < rhs->value.value; + }); + for (size_t i = 0; i < allocations_by_offset.size(); ++i) { + const TransientAllocation & allocation = *allocations_by_offset[i]; + const size_t end = saturated_range_end(allocation.arena_offset, allocation.size); + for (size_t j = i + 1; j < allocations_by_offset.size(); ++j) { + const TransientAllocation & other = *allocations_by_offset[j]; + if (other.arena_offset >= end) { + break; + } + if (!verify_transient_allocations_can_overlap(lifetimes, allocation, other)) { + status.log("transient allocations overlap"); + } + } + } +} + +static bool binding_requires_prior_write(const CommandBinding & binding) { + return binding.origin == CommandBindingOrigin::Transient && binding.access == ResourceAccess::Read; +} + +static bool binding_defines_transient_value(const CommandBinding & binding) { + return binding.origin == CommandBindingOrigin::Transient && resource_access_writes(binding.access); +} + +static int32_t find_later_transient_writer(const std::vector & commands, + size_t begin, + ValueId value) { + for (size_t i = begin; i < commands.size(); ++i) { + const Command & command = commands[i]; + for (const CommandBinding & binding : command.bindings) { + if (binding.value == value && binding_defines_transient_value(binding)) { + return static_cast(command.ordinal); + } + } + } + return -1; +} + +static void verify_transient_reads_are_defined(const CommandProgram & program, Status & status) { + std::unordered_set defined; + for (const Command & command : program.initialization_commands) { + for (const CommandBinding & binding : command.bindings) { + if (binding_defines_transient_value(binding)) { + defined.insert(binding.value.value); + } + } + } + for (const ConstantInitialization & initialization : program.constant_initializations) { + defined.insert(initialization.value.value); + } + if (program.completion_counters.byte_count > 0) { + for (const TransientAllocation & allocation : program.transients.allocations) { + if (allocation_overlaps_region(allocation, program.completion_counters.arena_offset, + program.completion_counters.byte_count)) { + defined.insert(allocation.value.value); + } + } + } + + for (size_t command_index = 0; command_index < program.commands.size(); ++command_index) { + const Command & command = program.commands[command_index]; + const std::string command_context = format_command(command); + for (size_t binding_index = 0; binding_index < command.bindings.size(); ++binding_index) { + const CommandBinding & binding = command.bindings[binding_index]; + if (!binding_requires_prior_write(binding) || defined.find(binding.value.value) != defined.end()) { + continue; + } + const int32_t later_writer = find_later_transient_writer(program.commands, command_index + 1, binding.value); + if (later_writer >= 0) { + status.log("%s %s reads transient value %d before write by command %d", command_context.c_str(), + format_command_binding(binding).c_str(), binding.value.value, later_writer); + } else { + status.log("%s %s reads transient value %d before write", command_context.c_str(), + format_command_binding(binding).c_str(), binding.value.value); + } + } + for (const CommandBinding & binding : command.bindings) { + if (binding_defines_transient_value(binding)) { + defined.insert(binding.value.value); + } + } + } +} + +} // namespace + +const TransientAllocation * find_transient_allocation(const TransientPlan & plan, ValueId value) { + const auto found = std::find_if(plan.allocations.begin(), plan.allocations.end(), + [&](const TransientAllocation & allocation) { return allocation.value == value; }); + return found == plan.allocations.end() ? nullptr : &*found; +} + +CommandProgram build_command_program(const Graph & graph, + const CommandPlan & plan, + const KernelCorpus & corpus, + const std::string & target) { + CommandProgram result; + if (!plan.valid()) { + result.status.append(plan.status); + return result; + } + + result.initialization_commands.reserve(plan.initialization_dispatches.size()); + for (const Dispatch & dispatch : plan.initialization_dispatches) { + append_command(graph, plan, corpus, target, dispatch, false, result.initialization_commands, result.status); + } + for (const Dispatch & dispatch : plan.dispatches) { + append_command(graph, plan, corpus, target, dispatch, true, result.commands, result.status); + } + result.transients = TransientAllocator::allocate(graph, plan, result.initialization_commands, result.commands, + result.completion_counters, result.status); + result.constant_initializations.reserve(plan.constant_initializations.size()); + for (const CommandPlanConstantInitialization & initialization : plan.constant_initializations) { + result.constant_initializations.push_back({ + initialization.value, + initialization.name, + initialization.offset, + initialization.data, + }); + } + return result; +} + +VerificationResult verify_command_program(const CommandProgram & program, + const KernelCorpus & corpus, + const std::string & target) { + VerificationResult result; + if (!program.valid()) { + result.status.append(program.status); + } + verify_command_list(program.initialization_commands, program.transients, corpus, target, result.status); + verify_command_list(program.commands, program.transients, corpus, target, result.status); + if (program.transients.arena_alignment == 0) { + result.status.log("transient arena has zero alignment"); + } + if (program.completion_counters.count == 0) { + if (program.completion_counters.byte_count != 0) { + result.status.log("completion counter region has bytes but no counters"); + } + } else { + if (program.completion_counters.byte_count == 0) { + result.status.log("completion counter region has counters but no bytes"); + } + if (program.completion_counters.byte_count % sizeof(int32_t) != 0) { + result.status.log("completion counter region byte count is not i32 aligned"); + } + if (program.completion_counters.arena_offset % 16 != 0) { + result.status.log("completion counter region is not 16-byte aligned"); + } + if (program.completion_counters.arena_offset > program.transients.arena_size || + program.completion_counters.byte_count > + program.transients.arena_size - program.completion_counters.arena_offset) { + result.status.log("completion counter region is outside transient arena"); + } + } + verify_transient_allocations(program, result.status); + verify_transient_reads_are_defined(program, result.status); + for (const ConstantInitialization & initialization : program.constant_initializations) { + const TransientAllocation * allocation = find_transient_allocation(program.transients, initialization.value); + if (allocation == nullptr) { + result.status.log("constant initialization %s references missing transient value %d", + initialization.name.c_str(), initialization.value.value); + continue; + } + if (initialization.data.empty()) { + result.status.log("constant initialization %s has no data", initialization.name.c_str()); + } + if (initialization.offset > allocation->size || + initialization.data.size() > allocation->size - initialization.offset) { + result.status.log("constant initialization %s is outside transient allocation length %zu", + initialization.name.c_str(), allocation->size); + } + } + return result; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/command-program.h b/ggml/src/ggml-hrx/dispatch/command-program.h new file mode 100644 index 000000000000..42796c30aa86 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/command-program.h @@ -0,0 +1,96 @@ +#pragma once + +#include "command-plan.h" +#include "graph/graph.h" +#include "kernel-corpus/kernel-corpus.h" +#include "kernel-corpus/kernel-types.h" +#include "status.h" + +#include +#include +#include +#include + +namespace ggml::hrx { + +enum class CommandKind : uint8_t { + Invalid, + Kernel, +}; + +enum class CommandBindingOrigin : uint8_t { + GraphValue, + Transient, + ProgramConstant, +}; + +struct CommandBinding { + std::string name; + ValueId value; + CommandBindingOrigin origin = CommandBindingOrigin::GraphValue; + size_t offset = 0; + size_t length = 0; + ResourceAccess access = ResourceAccess::Read; + std::string layout = kNativeWeightLayout; + ggml_type source_type = GGML_TYPE_COUNT; + int64_t input_size = 0; + int64_t output_size = 0; + size_t source_length = 0; +}; + +struct Command { + uint32_t ordinal = 0; + CommandKind kind = CommandKind::Kernel; + KernelSpecialization kernel; + std::vector bindings; + std::vector dependencies; +}; + +struct TransientAllocation { + ValueId value; + size_t size = 0; + size_t alignment = 1; + size_t arena_offset = 0; +}; + +struct TransientPlan { + size_t arena_size = 0; + size_t arena_alignment = 1; + std::vector allocations; +}; + +struct ConstantInitialization { + ValueId value; + std::string name; + size_t offset = 0; + std::vector data; +}; + +struct CompletionCounterPlan { + size_t arena_offset = 0; + size_t byte_count = 0; + uint32_t count = 0; +}; + +struct CommandProgram { + std::vector initialization_commands; + std::vector commands; + TransientPlan transients; + CompletionCounterPlan completion_counters; + std::vector constant_initializations; + Status status; + + bool valid() const { return status.success(); } +}; + +const TransientAllocation * find_transient_allocation(const TransientPlan & plan, ValueId value); + +CommandProgram build_command_program(const Graph & graph, + const CommandPlan & plan, + const KernelCorpus & corpus, + const std::string & target); +VerificationResult verify_command_program(const CommandProgram & program, + const KernelCorpus & corpus, + const std::string & target); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/dispatch-scheduler.cpp b/ggml/src/ggml-hrx/dispatch/dispatch-scheduler.cpp new file mode 100644 index 000000000000..d51ec80db79e --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/dispatch-scheduler.cpp @@ -0,0 +1,321 @@ +#include "dispatch-scheduler.h" + +#include "ggml.h" +#include "graph/graph-traversal.h" + +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static bool match_covers_root(const DispatchMatch & match, size_t root_index) { + return std::find(match.covered_nodes.begin(), match.covered_nodes.end(), root_index) != match.covered_nodes.end(); +} + +static bool match_overlaps_covered_nodes(const DispatchMatch & match, const std::vector & covered_nodes) { + for (const size_t node_index : match.covered_nodes) { + if (node_index >= covered_nodes.size() || covered_nodes[node_index]) { + return true; + } + } + return false; +} + +static bool try_match_registration(const Graph & graph, + const GraphNode * node, + size_t node_index, + const std::vector & covered_nodes, + const CommandPlan & plan, + const DispatchRegistry & registry, + ValueId next_plan_value, + DispatchMatch & match, + DispatchMatchDiagnostics * diagnostics) { + const DispatchMatchContext context = { + graph, node, node_index, covered_nodes, plan, next_plan_value, + }; + return registry.match(context, match, diagnostics); +} + +static void clear_plan_results(CommandPlan & plan) { + plan.initialization_dispatches.clear(); + plan.dispatches.clear(); + plan.transients.clear(); + plan.constant_initializations.clear(); + plan.completion_counter_requests.clear(); + plan.metadata.clear(); +} + +static void append_value_summary(std::ostringstream & stream, const Graph & graph, ValueId value_id) { + const Value * value = graph.values().find(value_id); + if (value == nullptr) { + stream << value_id.value << ":missing"; + return; + } + stream << value_id.value << ":" << ggml_type_name(value->type) << "[" << value->ne[0] << "," << value->ne[1] << "," + << value->ne[2] << "," << value->ne[3] << "]"; + const GraphNode * producer = graph.index().producer(value_id); + if (producer != nullptr) { + stream << "<-" << ggml_op_name(producer->op); + } +} + +static void append_node_summary(std::ostringstream & stream, const Graph & graph, const GraphNode * node) { + if (node == nullptr) { + stream << "null"; + return; + } + size_t node_index = 0; + if (graph.index().node_index(node, node_index)) { + stream << node_index << ":"; + } + stream << ggml_op_name(node->op); +} + +static std::string unsupported_node_message(const Graph & graph, size_t index, const GraphNode & node) { + std::ostringstream stream; + stream << "unsupported HRX node " << index << ": " << ggml_op_name(node.op) << " output="; + append_value_summary(stream, graph, node.output); + stream << " inputs=["; + for (size_t i = 0; i < node.inputs.size(); ++i) { + if (i > 0) { + stream << ", "; + } + append_value_summary(stream, graph, node.inputs[i]); + } + stream << "]"; + stream << " consumers=["; + const std::vector & consumers = graph.index().consumers(node.output); + for (size_t i = 0; i < consumers.size(); ++i) { + if (i > 0) { + stream << ", "; + } + append_node_summary(stream, graph, consumers[i]); + } + stream << "]"; + return stream.str(); +} + +static bool value_is_available(const Graph & graph, ValueId value, const std::vector & covered_nodes) { + const GraphNode * producer = graph.index().producer(value); + if (producer == nullptr) { + return true; + } + size_t producer_index = 0; + return graph.index().node_index(producer, producer_index) && producer_index < covered_nodes.size() && + covered_nodes[producer_index]; +} + +static bool can_elide_layout_alias_node(const Graph & graph, + const GraphNode & node, + const std::vector & covered_nodes) { + return is_layout_alias_node(graph, node) && value_is_available(graph, node.inputs[0], covered_nodes); +} + +static bool node_inputs_are_available(const Graph & graph, + const GraphNode & node, + const std::vector & covered_nodes) { + for (const ValueId input : node.inputs) { + if (!value_is_available(graph, input, covered_nodes)) { + return false; + } + } + return true; +} + +static bool can_elide_zero_output_node(const Graph & graph, + const GraphNode & node, + const std::vector & covered_nodes) { + const Value * output = graph.values().find(node.output); + return output != nullptr && (output->element_count == 0 || output->byte_count == 0) && + node_inputs_are_available(graph, node, covered_nodes); +} + +static bool can_elide_node(const Graph & graph, const GraphNode & node, const std::vector & covered_nodes) { + return can_elide_layout_alias_node(graph, node, covered_nodes) || + can_elide_zero_output_node(graph, node, covered_nodes); +} + +static bool apply_value_aliases(Graph & graph, const DispatchMatch & match, Status & status) { + for (const DispatchValueAliasRequest & alias : match.value_aliases) { + Status alias_status = graph.values().alias_storage(alias.target_value, alias.source_value); + if (!alias_status.success()) { + status.append(alias_status); + return false; + } + } + return true; +} + +} // namespace + +bool DispatchScheduler::schedule_graph(Graph & graph, const DispatchTarget & target) { + return this->schedule_graph(graph, target, nullptr); +} + +bool DispatchScheduler::schedule_graph(Graph & graph, + const DispatchTarget & target, + DispatchScheduleDiagnostics * diagnostics) { + plan_ = {}; + if (diagnostics != nullptr) { + *diagnostics = {}; + } + const DispatchRegistry * registry = find_dispatch_registry(target); + if (registry == nullptr) { + plan_.status.log("no HRX dispatch registry for target %s", target.architecture.c_str()); + return false; + } + const std::vector & nodes = graph.nodes(); + if (!graph.has_index()) { + plan_.status.log("HRX graph is missing graph index"); + return false; + } + std::vector covered_nodes(nodes.size(), false); + Status pending_diagnostics; + const GraphTraversalOrder traversal = GraphTraversalOrder::build(graph); + for (const GraphNode * node : traversal.nodes()) { + size_t i = 0; + if (node == nullptr || !graph.index().node_index(node, i)) { + plan_.status.log("HRX traversal references a node outside the graph"); + clear_plan_results(plan_); + return false; + } + if (covered_nodes[i]) { + continue; + } + // leaf nodes (ggml_build_forward_expand emits them) have no producer command: their storage is bound from outside + if (node->op == GGML_OP_NONE) { + covered_nodes[i] = true; + continue; + } + if (can_elide_zero_output_node(graph, *node, covered_nodes)) { + covered_nodes[i] = true; + continue; + } + DispatchMatch match; + const ValueId next_plan_value(static_cast(graph.values().size() + plan_.transients.size() + + plan_.completion_counter_requests.size())); + DispatchMatchDiagnostics match_diagnostics; + if (!try_match_registration(graph, node, i, covered_nodes, plan_, *registry, next_plan_value, match, + &match_diagnostics)) { + if (can_elide_node(graph, *node, covered_nodes)) { + pending_diagnostics.append(match.status); + covered_nodes[i] = true; + continue; + } + plan_.status.append(pending_diagnostics); + plan_.status.append(match.status); + const std::string message = unsupported_node_message(graph, i, *node); + plan_.status.log("%s", message.c_str()); + if (diagnostics != nullptr) { + diagnostics->unsupported_node_index = i; + diagnostics->unsupported_node = node; + diagnostics->unsupported_message = message; + diagnostics->match = std::move(match_diagnostics); + } + clear_plan_results(plan_); + return false; + } + if (match.covered_nodes.empty() || match.dispatches.empty() || !match_covers_root(match, i) || + match_overlaps_covered_nodes(match, covered_nodes)) { + plan_.status.log("invalid HRX dispatch match for node %zu: %s", i, ggml_op_name(node->op)); + if (diagnostics != nullptr) { + diagnostics->unsupported_node_index = i; + diagnostics->unsupported_node = node; + diagnostics->unsupported_message = "invalid HRX dispatch match"; + diagnostics->match = std::move(match_diagnostics); + } + clear_plan_results(plan_); + return false; + } + if (!apply_value_aliases(graph, match, plan_.status)) { + if (diagnostics != nullptr) { + diagnostics->unsupported_node_index = i; + diagnostics->unsupported_node = node; + diagnostics->unsupported_message = "invalid HRX value alias"; + diagnostics->match = std::move(match_diagnostics); + } + clear_plan_results(plan_); + return false; + } + for (Dispatch & dispatch : match.initialization_dispatches) { + plan_.initialization_dispatches.push_back(std::move(dispatch)); + } + for (Dispatch & dispatch : match.dispatches) { + plan_.dispatches.push_back(std::move(dispatch)); + } + for (CommandPlanTransient & transient : match.transients) { + plan_.transients.push_back(std::move(transient)); + } + for (CommandPlanConstantInitialization & initialization : match.constant_initializations) { + plan_.constant_initializations.push_back(std::move(initialization)); + } + for (CommandPlanCompletionCounterRequest & request : match.completion_counter_requests) { + plan_.completion_counter_requests.push_back(std::move(request)); + } + if (!plan_.metadata.append(std::move(match.metadata), plan_.status)) { + clear_plan_results(plan_); + return false; + } + for (const size_t covered_node : match.covered_nodes) { + covered_nodes[covered_node] = true; + } + } + for (size_t i = 0; i < nodes.size(); ++i) { + if (!covered_nodes[i]) { + if (can_elide_node(graph, nodes[i], covered_nodes)) { + covered_nodes[i] = true; + continue; + } + plan_.status.append(pending_diagnostics); + const std::string message = unsupported_node_message(graph, i, nodes[i]); + plan_.status.log("%s", message.c_str()); + if (diagnostics != nullptr) { + DispatchMatch match; + const ValueId next_plan_value(static_cast(graph.values().size() + plan_.transients.size() + + plan_.completion_counter_requests.size())); + DispatchMatchDiagnostics match_diagnostics; + try_match_registration(graph, &nodes[i], i, covered_nodes, plan_, *registry, next_plan_value, match, + &match_diagnostics); + diagnostics->unsupported_node_index = i; + diagnostics->unsupported_node = &nodes[i]; + diagnostics->unsupported_message = message; + diagnostics->match = std::move(match_diagnostics); + } + clear_plan_results(plan_); + return false; + } + } + return true; +} + +bool DispatchScheduler::supports_node(const Graph & graph, const GraphNode * node, const DispatchTarget & target) { + const DispatchRegistry * registry = find_dispatch_registry(target); + if (registry == nullptr) { + return false; + } + if (node == nullptr || !graph.has_index()) { + return false; + } + size_t node_index = 0; + if (!graph.index().node_index(node, node_index)) { + return false; + } + const std::vector covered_nodes(graph.nodes().size(), false); + DispatchMatch match; + const ValueId next_plan_value(static_cast(graph.values().size())); + const CommandPlan plan; + return try_match_registration(graph, node, node_index, covered_nodes, plan, *registry, next_plan_value, match, + nullptr); +} + +bool DispatchScheduler::can_schedule_graph(const Graph & graph, const DispatchTarget & target) { + Graph graph_copy = graph; + DispatchScheduler scheduler; + return scheduler.schedule_graph(graph_copy, target); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/dispatch-scheduler.h b/ggml/src/ggml-hrx/dispatch/dispatch-scheduler.h new file mode 100644 index 000000000000..b533c0ed78d2 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/dispatch-scheduler.h @@ -0,0 +1,40 @@ +#pragma once + +#include "command-plan.h" +#include "dispatch_registration/dispatch-registry.h" +#include "graph/graph.h" + +#include +#include + +namespace ggml::hrx { + +struct DispatchScheduleDiagnostics { + size_t unsupported_node_index = 0; + const GraphNode * unsupported_node = nullptr; + std::string unsupported_message; + DispatchMatchDiagnostics match; +}; + +class DispatchScheduler { + public: + bool schedule_graph(Graph & graph, const DispatchTarget & target); + bool schedule_graph(Graph & graph, const DispatchTarget & target, DispatchScheduleDiagnostics * diagnostics); + + const CommandPlan & plan() const { return plan_; } + + const std::vector & dispatches() const { return plan_.dispatches; } + + const std::string & error() const { + static const std::string empty; + return plan_.status.errors().empty() ? empty : plan_.status.errors().front(); + } + + static bool supports_node(const Graph & graph, const GraphNode * node, const DispatchTarget & target); + static bool can_schedule_graph(const Graph & graph, const DispatchTarget & target); + + private: + CommandPlan plan_; +}; + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/dispatch.h b/ggml/src/ggml-hrx/dispatch/dispatch.h new file mode 100644 index 000000000000..0e5ea1b0cfa9 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/dispatch.h @@ -0,0 +1,59 @@ +#pragma once + +#include "graph/value-map.h" +#include "kernel-corpus/kernel-corpus-catalog.h" +#include "kernel-corpus/kernel-types.h" + +#include +#include +#include +#include +#include + +namespace ggml::hrx { + +inline constexpr const char kNativeWeightLayout[] = "ggml-native"; +inline constexpr const char kQ4KPackedK256Row64Layout[] = "q4k-packed-k256-row64"; +inline constexpr const char kSymmetricI4K32Row64Layout[] = "symi4-k32-row64"; +inline constexpr const char kSymmetricI2K32EightGroupsShared4Layout[] = "symi2-k32-eightgroups-shared4-payload-first"; +inline constexpr const char kSymmetricI4K32EightGroupsShared4Layout[] = "symi4-k32-eightgroups-shared4-payload-first"; +inline constexpr const char kSymmetricI4K32EightGroupsShared4MultistartLayout[] = + "symi4-k32-eightgroups-shared4-multistart-payload-first"; +inline constexpr const char kQ6KSymmetricI2PackedK256Row64ScaleRowLayout[] = + "q6k-symi2-k32-eightgroups-shared4-plus-packed-k256-row64-scalerow"; +inline constexpr const char kSymmetricI4K64Row64Layout[] = "symi4-k64-row64"; +inline constexpr const char kQ5KSymmetricI5K32Layout[] = "q5k-symi5-k32-native-footprint"; +inline constexpr const char kQ5KSymmetricI8K256Row64Layout[] = "q5k-symi8-k256-row64"; +inline constexpr const char kQ6KI8K32Row64Layout[] = "q6k-i8-k32-row64"; +inline constexpr const char kQ6KPackedK256Row64ScaleRowLayout[] = "q6k-packed-k256-row64-scalerow"; + +struct KernelSpecialization { + uint64_t kernel_id = kUncatalogedKernelId; + std::map integer_parameters; + std::map compile_parameters; +}; + +inline KernelSpecialization make_kernel_specialization(KernelCatalogRef ref) { + KernelSpecialization kernel; + kernel.kernel_id = ref.id; + return kernel; +} + +struct DispatchBinding { + ValueId value; + size_t offset = 0; + size_t length = 0; + // Non-native layouts request a backend-resident materialization; the graph value remains unchanged. + std::string layout = kNativeWeightLayout; + ggml_type source_type = GGML_TYPE_COUNT; + int64_t input_size = 0; + int64_t output_size = 0; + size_t source_length = 0; +}; + +struct Dispatch { + KernelSpecialization kernel; + std::vector bindings; +}; + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/ternary-q4-0.h b/ggml/src/ggml-hrx/dispatch/ternary-q4-0.h new file mode 100644 index 000000000000..cb35db0d8931 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/ternary-q4-0.h @@ -0,0 +1,48 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Packed ternary weights for the K-quant decode kernels. A ternary model written as exact Q4_0 (each +// 128-value group of trits t in {-1, 0, 1} with one fp16 scale d becomes four Q4_0 blocks with nibbles +// t + 8 and the same d; engine tools/ternary_to_q4_0.py) is repacked at upload into 68 bytes per 256 +// values: d0, d1 (the two group scales), then 64 bytes of 2-bit codes t + 1. Decode then reads 2.125 +// instead of 4.5 bits per weight. Opt-in with GGML_HRX_TERNARY_Q4_0=1 (1bit serve sets it for files +// stamped onebit.ternary_q4_0); the upload checks every block and fails rather than change a weight. +#pragma once + +#include +#include +#include + +namespace ggml::hrx { + +inline constexpr const char kTernaryQ40K128Layout[] = "ternary-q4_0-k128-t2"; + +// The K-quant decode kernels' weight format value for this layout (ops/kquant_decode_f32.loom). +inline constexpr int64_t kTernaryQ40K128FormatValue = 90; + +inline bool ternary_q4_0_enabled() { + static const bool enabled = [] { + const char * value = std::getenv("GGML_HRX_TERNARY_Q4_0"); + return value != nullptr && value[0] != '\0' && value[0] != '0'; + }(); + return enabled; +} + +// Bytes of the packed layout for a K x rows weight (K a multiple of 256). +inline size_t ternary_q4_0_k128_bytes(int64_t input_size, int64_t rows) { + return static_cast(input_size / 256) * 68u * static_cast(rows); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/transient-allocator.cpp b/ggml/src/ggml-hrx/dispatch/transient-allocator.cpp new file mode 100644 index 000000000000..fde1d182484b --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/transient-allocator.cpp @@ -0,0 +1,464 @@ +#include "transient-allocator.h" +#include "transient-reuse-guard.h" + +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static size_t align_up(size_t value, size_t alignment) { + return alignment == 0 ? value : (value + alignment - 1) / alignment * alignment; +} + +struct TransientAllocationRequest { + ValueId value; + size_t required_size = 0; +}; + +struct TransientAllocationInterval { + TransientAllocation allocation; + uint32_t first_use = 0; + uint32_t last_use = 0; +}; + +struct FreeTransientBlock { + size_t offset = 0; + size_t size = 0; +}; + +struct ActiveTransientAllocation { + size_t offset = 0; + size_t size = 0; + uint32_t last_use = 0; +}; + +static const CommandPlanTransient * find_plan_transient(const CommandPlan & plan, ValueId value) { + const auto found = std::find_if(plan.transients.begin(), plan.transients.end(), + [&](const CommandPlanTransient & transient) { return transient.value == value; }); + return found == plan.transients.end() ? nullptr : &*found; +} + +static const CommandPlanCompletionCounterRequest * find_plan_completion_counter_request(const CommandPlan & plan, + ValueId value) { + const auto found = + std::find_if(plan.completion_counter_requests.begin(), plan.completion_counter_requests.end(), + [&](const CommandPlanCompletionCounterRequest & request) { return request.value == value; }); + return found == plan.completion_counter_requests.end() ? nullptr : &*found; +} + +static void add_transient_allocation_request(std::vector & requests, + ValueId value, + size_t required_size) { + for (TransientAllocationRequest & request : requests) { + if (request.value == value) { + request.required_size = std::max(request.required_size, required_size); + return; + } + } + requests.push_back({ value, required_size }); +} + +static const TransientAllocationRequest * find_transient_allocation_request( + const std::vector & requests, + ValueId value) { + const auto found = std::find_if(requests.begin(), requests.end(), + [&](const TransientAllocationRequest & request) { return request.value == value; }); + return found == requests.end() ? nullptr : &*found; +} + +static bool contains_value(const std::vector & values, ValueId value) { + return std::find(values.begin(), values.end(), value) != values.end(); +} + +static bool make_transient_allocation(const Graph & graph, + const CommandPlan & command_plan, + const TransientAllocationRequest & request, + TransientAllocation & allocation, + Status & errors) { + const Value * graph_value = graph.values().find(request.value); + const CommandPlanTransient * plan_transient = find_plan_transient(command_plan, request.value); + if (graph_value == nullptr && plan_transient == nullptr) { + errors.log("transient value %d is missing from graph values and command plan transients", request.value.value); + return false; + } + if (graph_value != nullptr && graph_value->kind != ValueKind::Transient) { + errors.log("transient value %d aliases a non-transient graph value", request.value.value); + return false; + } + + allocation.value = request.value; + allocation.size = + std::max(graph_value != nullptr ? graph_value->byte_count : plan_transient->size, request.required_size); + allocation.alignment = plan_transient != nullptr ? plan_transient->alignment : 256; + allocation.arena_offset = 0; + return true; +} + +static void add_transient_allocation(const Graph & graph, + const CommandPlan & command_plan, + const TransientAllocationRequest & request, + TransientPlan & plan, + Status & errors) { + TransientAllocation allocation; + if (!make_transient_allocation(graph, command_plan, request, allocation, errors)) { + return; + } + allocation.arena_offset = align_up(plan.arena_size, allocation.alignment); + plan.arena_size = allocation.arena_offset + allocation.size; + plan.allocations.push_back(allocation); +} + +static void add_transient_interval(std::vector & intervals, + TransientAllocation allocation, + uint32_t command_ordinal) { + for (TransientAllocationInterval & interval : intervals) { + if (interval.allocation.value == allocation.value) { + interval.allocation.size = std::max(interval.allocation.size, allocation.size); + interval.allocation.alignment = std::max(interval.allocation.alignment, allocation.alignment); + interval.first_use = std::min(interval.first_use, command_ordinal); + interval.last_use = std::max(interval.last_use, command_ordinal); + return; + } + } + intervals.push_back({ allocation, command_ordinal, command_ordinal }); +} + +static void add_free_transient_block(std::vector & free_blocks, size_t offset, size_t size) { + if (size > 0) { + free_blocks.push_back({ offset, size }); + } +} + +static void coalesce_free_transient_blocks(std::vector & free_blocks) { + std::sort(free_blocks.begin(), free_blocks.end(), + [](const FreeTransientBlock & lhs, const FreeTransientBlock & rhs) { return lhs.offset < rhs.offset; }); + size_t write_index = 0; + for (const FreeTransientBlock & block : free_blocks) { + if (block.size == 0) { + continue; + } + if (write_index > 0) { + FreeTransientBlock & previous = free_blocks[write_index - 1]; + const size_t previous_end = previous.offset + previous.size; + if (previous_end == block.offset) { + previous.size += block.size; + continue; + } + } + free_blocks[write_index++] = block; + } + free_blocks.resize(write_index); +} + +static void release_completed_transient_allocations(std::vector & active, + std::vector & free_blocks, + uint32_t first_use) { + size_t write_index = 0; + for (size_t read_index = 0; read_index < active.size(); ++read_index) { + if (active[read_index].last_use < first_use) { + add_free_transient_block(free_blocks, active[read_index].offset, active[read_index].size); + continue; + } + if (write_index != read_index) { + active[write_index] = active[read_index]; + } + ++write_index; + } + active.resize(write_index); +} + +static bool try_allocate_from_free_blocks(std::vector & free_blocks, + size_t size, + size_t alignment, + size_t & offset) { + coalesce_free_transient_blocks(free_blocks); + for (size_t i = 0; i < free_blocks.size(); ++i) { + const size_t block_begin = free_blocks[i].offset; + const size_t block_end = free_blocks[i].offset + free_blocks[i].size; + const size_t aligned_offset = align_up(block_begin, alignment); + const bool aligned_in_block = aligned_offset >= block_begin && aligned_offset <= block_end; + if (!aligned_in_block || size > block_end - aligned_offset) { + continue; + } + + offset = aligned_offset; + const FreeTransientBlock block = free_blocks[i]; + free_blocks.erase(free_blocks.begin() + static_cast(i)); + add_free_transient_block(free_blocks, block.offset, aligned_offset - block.offset); + add_free_transient_block(free_blocks, aligned_offset + size, block_end - (aligned_offset + size)); + return true; + } + return false; +} + +static void pack_transient_intervals(std::vector & intervals, TransientPlan & plan) { + std::sort(intervals.begin(), intervals.end(), + [](const TransientAllocationInterval & lhs, const TransientAllocationInterval & rhs) { + if (lhs.first_use != rhs.first_use) { + return lhs.first_use < rhs.first_use; + } + if (lhs.last_use != rhs.last_use) { + return lhs.last_use < rhs.last_use; + } + return lhs.allocation.value.value < rhs.allocation.value.value; + }); + + std::vector active; + std::vector free_blocks; + for (TransientAllocationInterval & interval : intervals) { + release_completed_transient_allocations(active, free_blocks, interval.first_use); + + size_t offset = 0; + if (!try_allocate_from_free_blocks(free_blocks, interval.allocation.size, interval.allocation.alignment, + offset)) { + offset = align_up(plan.arena_size, interval.allocation.alignment); + plan.arena_size = offset + interval.allocation.size; + } + + interval.allocation.arena_offset = offset; + active.push_back({ offset, interval.allocation.size, interval.last_use }); + plan.allocations.push_back(interval.allocation); + } +} + +static void add_completion_counter_allocations(const CommandPlan & command_plan, + const std::vector & binding_requests, + TransientPlan & plan, + CompletionCounterPlan & completion_counters, + Status & errors) { + for (size_t i = 0; i < command_plan.completion_counter_requests.size(); ++i) { + const CommandPlanCompletionCounterRequest & request = command_plan.completion_counter_requests[i]; + if (request.value.value < 0) { + errors.log("completion counter request %s has invalid value %d", request.name.c_str(), request.value.value); + continue; + } + if (request.count == 0) { + errors.log("completion counter request %s has zero counters", request.name.c_str()); + continue; + } + for (size_t j = i + 1; j < command_plan.completion_counter_requests.size(); ++j) { + if (request.value == command_plan.completion_counter_requests[j].value) { + errors.log("duplicate completion counter request value %d", request.value.value); + } + } + const size_t byte_count = static_cast(request.count) * sizeof(int32_t); + const TransientAllocationRequest * binding_request = + find_transient_allocation_request(binding_requests, request.value); + if (binding_request != nullptr && binding_request->required_size > byte_count) { + errors.log("completion counter request %s requires %zu bytes but binding uses %zu bytes", + request.name.c_str(), byte_count, binding_request->required_size); + continue; + } + if (completion_counters.count > std::numeric_limits::max() - request.count) { + errors.log("completion counter count overflows"); + continue; + } + + TransientAllocation allocation; + allocation.value = request.value; + allocation.size = byte_count; + allocation.alignment = 16; + allocation.arena_offset = align_up(plan.arena_size, allocation.alignment); + if (completion_counters.byte_count == 0) { + completion_counters.arena_offset = allocation.arena_offset; + } + plan.arena_size = allocation.arena_offset + allocation.size; + completion_counters.byte_count = + plan.arena_size > completion_counters.arena_offset ? plan.arena_size - completion_counters.arena_offset : 0; + completion_counters.count += request.count; + plan.allocations.push_back(allocation); + } +} + +static void record_transient_lifetime(std::vector & lifetimes, + const CommandBinding & binding, + uint32_t command_ordinal) { + for (TransientAllocationInterval & lifetime : lifetimes) { + if (lifetime.allocation.value == binding.value) { + lifetime.first_use = std::min(lifetime.first_use, command_ordinal); + lifetime.last_use = std::max(lifetime.last_use, command_ordinal); + return; + } + } + TransientAllocation allocation; + allocation.value = binding.value; + lifetimes.push_back({ allocation, command_ordinal, command_ordinal }); +} + +static const TransientAllocationInterval * find_transient_lifetime( + const std::vector & lifetimes, + ValueId value) { + const auto found = + std::find_if(lifetimes.begin(), lifetimes.end(), + [&](const TransientAllocationInterval & lifetime) { return lifetime.allocation.value == value; }); + return found == lifetimes.end() ? nullptr : &*found; +} + +static bool transient_allocation_overlaps_region(const TransientAllocation & allocation, + size_t region_offset, + size_t region_size) { + if (region_size == 0) { + return false; + } + return allocation.arena_offset < region_offset + region_size && + region_offset < allocation.arena_offset + allocation.size; +} + +static bool transient_allocation_has_reserved_lifetime(const CommandProgram & program, + const TransientAllocation & allocation) { + if (transient_allocation_overlaps_region(allocation, program.completion_counters.arena_offset, + program.completion_counters.byte_count)) { + return true; + } + for (const Command & command : program.initialization_commands) { + for (const CommandBinding & binding : command.bindings) { + if (binding.origin == CommandBindingOrigin::Transient && binding.value == allocation.value) { + return true; + } + } + } + for (const ConstantInitialization & initialization : program.constant_initializations) { + if (initialization.value == allocation.value) { + return true; + } + } + return false; +} + +static std::vector collect_main_transient_lifetimes(const CommandProgram & program) { + std::vector lifetimes; + for (const Command & command : program.commands) { + for (const CommandBinding & binding : command.bindings) { + if (binding.origin == CommandBindingOrigin::Transient) { + record_transient_lifetime(lifetimes, binding, command.ordinal); + } + } + } + return lifetimes; +} + +static bool transient_lifetimes_disjoint(const std::vector & lifetimes, + const TransientAllocation & lhs, + const TransientAllocation & rhs) { + const TransientAllocationInterval * lhs_lifetime = find_transient_lifetime(lifetimes, lhs.value); + const TransientAllocationInterval * rhs_lifetime = find_transient_lifetime(lifetimes, rhs.value); + if (lhs_lifetime == nullptr || rhs_lifetime == nullptr) { + return false; + } + return lhs_lifetime->last_use < rhs_lifetime->first_use || rhs_lifetime->last_use < lhs_lifetime->first_use; +} + +} // namespace + +TransientPlan TransientAllocator::allocate(const Graph & graph, + const CommandPlan & command_plan, + const std::vector & initialization_commands, + const std::vector & commands, + CompletionCounterPlan & completion_counters, + Status & errors) { + TransientPlan plan; + plan.arena_alignment = 256; + std::vector reserved_values; + std::vector reserved_requests; + std::vector packable_requests; + std::vector completion_counter_binding_requests; + + for (const Command & command : initialization_commands) { + for (const CommandBinding & binding : command.bindings) { + if (binding.origin == CommandBindingOrigin::Transient && + find_plan_completion_counter_request(command_plan, binding.value) == nullptr && + !contains_value(reserved_values, binding.value)) { + reserved_values.push_back(binding.value); + } + } + } + for (const CommandPlanConstantInitialization & initialization : command_plan.constant_initializations) { + const bool completion_counter = + find_plan_completion_counter_request(command_plan, initialization.value) != nullptr; + if (!completion_counter && !contains_value(reserved_values, initialization.value)) { + reserved_values.push_back(initialization.value); + } + if (initialization.offset > std::numeric_limits::max() - initialization.data.size()) { + errors.log("constant initialization %s range overflows", initialization.name.c_str()); + continue; + } + if (completion_counter) { + add_transient_allocation_request(completion_counter_binding_requests, initialization.value, + initialization.offset + initialization.data.size()); + } else { + add_transient_allocation_request(reserved_requests, initialization.value, + initialization.offset + initialization.data.size()); + } + } + + auto append_command_bindings = [&](const std::vector & command_list, bool packable) { + for (const Command & command : command_list) { + for (const CommandBinding & binding : command.bindings) { + if (binding.origin != CommandBindingOrigin::Transient) { + continue; + } + if (binding.offset > std::numeric_limits::max() - binding.length) { + errors.log("transient value %d binding range overflows", binding.value.value); + continue; + } + if (find_plan_completion_counter_request(command_plan, binding.value) != nullptr) { + add_transient_allocation_request(completion_counter_binding_requests, binding.value, + binding.offset + binding.length); + } else if (!packable || contains_value(reserved_values, binding.value)) { + add_transient_allocation_request(reserved_requests, binding.value, binding.offset + binding.length); + } else { + add_transient_allocation_request(packable_requests, binding.value, binding.offset + binding.length); + } + } + } + }; + append_command_bindings(initialization_commands, false); + append_command_bindings(commands, true); + add_completion_counter_allocations(command_plan, completion_counter_binding_requests, plan, completion_counters, + errors); + for (const TransientAllocationRequest & request : reserved_requests) { + add_transient_allocation(graph, command_plan, request, plan, errors); + } + + std::vector intervals; + for (const Command & command : commands) { + for (const CommandBinding & binding : command.bindings) { + if (binding.origin != CommandBindingOrigin::Transient || + find_plan_completion_counter_request(command_plan, binding.value) != nullptr || + contains_value(reserved_values, binding.value)) { + continue; + } + const TransientAllocationRequest * request = + find_transient_allocation_request(packable_requests, binding.value); + if (request == nullptr) { + continue; + } + TransientAllocation allocation; + if (make_transient_allocation(graph, command_plan, *request, allocation, errors)) { + add_transient_interval(intervals, allocation, command.ordinal); + } + } + } + std::vector lifetimes; // transient-reuse-guard.h + for (const TransientAllocationInterval & interval : intervals) lifetimes.push_back({ interval.allocation, interval.first_use, interval.last_use }); + if (!pack_transients_without_false_dependencies(commands, lifetimes, plan)) pack_transient_intervals(intervals, plan); + plan.arena_size = align_up(plan.arena_size, plan.arena_alignment); + return plan; +} + +bool TransientAllocator::allocations_can_overlap(const CommandProgram & program, + const TransientAllocation & lhs, + const TransientAllocation & rhs) { + if (transient_allocation_has_reserved_lifetime(program, lhs) || + transient_allocation_has_reserved_lifetime(program, rhs)) { + return false; + } + const std::vector lifetimes = collect_main_transient_lifetimes(program); + return transient_lifetimes_disjoint(lifetimes, lhs, rhs); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/transient-allocator.h b/ggml/src/ggml-hrx/dispatch/transient-allocator.h new file mode 100644 index 000000000000..ad059b7de60f --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/transient-allocator.h @@ -0,0 +1,23 @@ +#pragma once + +#include "command-program.h" + +#include + +namespace ggml::hrx { + +class TransientAllocator { + public: + static TransientPlan allocate(const Graph & graph, + const CommandPlan & command_plan, + const std::vector & initialization_commands, + const std::vector & commands, + CompletionCounterPlan & completion_counters, + Status & errors); + + static bool allocations_can_overlap(const CommandProgram & program, + const TransientAllocation & lhs, + const TransientAllocation & rhs); +}; + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/transient-reuse-guard.cpp b/ggml/src/ggml-hrx/dispatch/transient-reuse-guard.cpp new file mode 100644 index 000000000000..a52180da0b61 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/transient-reuse-guard.cpp @@ -0,0 +1,327 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "transient-reuse-guard.h" + +#include "ggml-impl.h" + +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +// Ancestor bitsets cost commands^2 / 8 bytes; 8192 commands = 8 MiB, built once per program. +constexpr size_t kMaxCommands = 8192; + +// The guarded packing is used while its arena stays within 2x the unguarded one plus this much. +constexpr size_t kMaxGuardGrowth = size_t{ 8 } << 20; + +// Programs whose unguarded arena exceeds this are prompt programs (decode measured 0.1-2.4 MB unguarded). +constexpr size_t kMaxUnguardedArena = size_t{ 16 } << 20; + +bool reuse_guard_enabled() { + static const bool enabled = [] { + const char * value = std::getenv("GGML_HRX_TRANSIENT_REUSE"); + return value == nullptr || (std::strcmp(value, "legacy") != 0 && std::strcmp(value, "0") != 0); + }(); + return enabled; +} + +size_t align_up(size_t value, size_t alignment) { + return alignment == 0 ? value : (value + alignment - 1) / alignment * alignment; +} + +bool writes(ResourceAccess access) { + return access == ResourceAccess::Write || access == ResourceAccess::ReadWrite; +} + +bool reads(ResourceAccess access) { + return access == ResourceAccess::Read || access == ResourceAccess::ReadWrite; +} + +// ancestors[c] holds bit d when command d must finish before command c, following read/write hazards on +// the same value (transients by value id, before arena placement; graph values and program constants by +// value id). Different graph values that alias one buffer are not linked here; a missing link only makes +// the packer reuse less, never more. +class CommandAncestors { + public: + explicit CommandAncestors(const std::vector & commands) : + words_((commands.size() + 63) / 64), + bits_(commands.size() * words_, 0) { + struct Access { + int32_t last_writer = -1; + std::vector readers; + }; + std::map, Access> accesses; + std::vector deps; + for (size_t c = 0; c < commands.size(); ++c) { + deps.clear(); + for (const CommandBinding & binding : commands[c].bindings) { + const auto key = std::make_pair(static_cast(binding.origin), binding.value.value); + auto found = accesses.find(key); + if (found == accesses.end()) { + continue; + } + if (found->second.last_writer >= 0) { + deps.push_back(static_cast(found->second.last_writer)); + } + if (writes(binding.access)) { + deps.insert(deps.end(), found->second.readers.begin(), found->second.readers.end()); + } + } + uint64_t * row = &bits_[c * words_]; + for (uint32_t d : deps) { + if (d == c) { + continue; + } + const uint64_t * dep_row = &bits_[static_cast(d) * words_]; + for (size_t w = 0; w < words_; ++w) { + row[w] |= dep_row[w]; + } + row[d / 64] |= uint64_t{ 1 } << (d % 64); + } + for (const CommandBinding & binding : commands[c].bindings) { + Access & access = accesses[std::make_pair(static_cast(binding.origin), binding.value.value)]; + if (writes(binding.access)) { + access.last_writer = static_cast(c); + access.readers.clear(); + } + if (reads(binding.access)) { + access.readers.push_back(static_cast(c)); + } + } + } + } + + bool precedes(uint32_t before, uint32_t command) const { + return (bits_[static_cast(command) * words_ + before / 64] >> (before % 64)) & 1; + } + + private: + size_t words_; + std::vector bits_; +}; + +struct FreeRange { + size_t offset = 0; + size_t size = 0; + std::vector users; // every command that touched a value previously placed here +}; + +struct LiveRange { + size_t offset = 0; + size_t size = 0; + uint32_t last_use = 0; + std::vector users; +}; + +bool log_enabled() { + static const bool enabled = [] { + const char * value = std::getenv("GGML_HRX_LOG_TRANSIENT_REUSE"); + return value != nullptr && value[0] != '\0' && std::strcmp(value, "0") != 0; + }(); + return enabled; +} + +bool pack(const std::vector & commands, + const CommandAncestors * ancestors, + std::vector & lifetimes, + TransientPlan & plan); + +} // namespace + +bool pack_transients_without_false_dependencies(const std::vector & commands, + std::vector & lifetimes, + TransientPlan & plan) { + if (!reuse_guard_enabled() || commands.empty() || commands.size() > kMaxCommands) { + return false; + } + for (size_t i = 0; i < commands.size(); ++i) { + if (commands[i].ordinal != i) { + return false; + } + } + // The guard costs arena space: every value whose range cannot be reused gets fresh space. Decode + // programs grow by a few MiB at most (measured ZAYA1-8B 2.2 -> 2.9 MB, GLM-4.7-Flash 0.2 -> 0.9 MB); + // 512-token prompt programs grew 3-15x (ZAYA1-8B 19 -> 299 MB), and there the barriers are a small + // share of the long kernels, so such programs keep the stock packing. + const CommandAncestors ancestors(commands); + std::vector guarded_lifetimes = lifetimes; + TransientPlan guarded = plan; + std::vector unguarded_lifetimes = lifetimes; + TransientPlan unguarded = plan; + if (!pack(commands, &ancestors, guarded_lifetimes, guarded) || + !pack(commands, nullptr, unguarded_lifetimes, unguarded)) { + return false; + } + // Prompt programs (unguarded arena over 16 MiB) keep the stock packing as well: with the guard, pp512 + // read 2.2% lower on Qwen3-8B (larger arena, more activation traffic). + const size_t limit = 2 * unguarded.arena_size + kMaxGuardGrowth; + const bool use = guarded.arena_size <= limit && unguarded.arena_size <= kMaxUnguardedArena; + if (log_enabled()) { + GGML_LOG_WARN("ggml-hrx transient reuse: %zu commands, %zu values, arena %zu bytes guarded, %zu unguarded%s\n", + commands.size(), lifetimes.size(), guarded.arena_size, unguarded.arena_size, + use ? "" : " (prompt-sized or over the growth limit: stock packing)"); + } + if (!use) { + return false; + } + lifetimes = std::move(guarded_lifetimes); + plan = std::move(guarded); + return true; +} + +namespace { + +// |ancestors| == nullptr packs without the guard (any freed range is reusable); used to bound the +// guarded arena's growth. +bool pack(const std::vector & commands, + const CommandAncestors * ancestors, + std::vector & lifetimes, + TransientPlan & plan) { + + std::map> users_by_value; + for (const Command & command : commands) { + for (const CommandBinding & binding : command.bindings) { + if (binding.origin == CommandBindingOrigin::Transient) { + std::vector & users = users_by_value[binding.value.value]; + if (users.empty() || users.back() != command.ordinal) { + users.push_back(command.ordinal); + } + } + } + } + + std::sort(lifetimes.begin(), lifetimes.end(), [](const TransientReuseLifetime & lhs, const TransientReuseLifetime & rhs) { + if (lhs.first_use != rhs.first_use) { + return lhs.first_use < rhs.first_use; + } + if (lhs.last_use != rhs.last_use) { + return lhs.last_use < rhs.last_use; + } + return lhs.allocation.value.value < rhs.allocation.value.value; + }); + + TransientPlan packed = plan; + std::vector live; + std::vector free_ranges; + for (TransientReuseLifetime & lifetime : lifetimes) { + const uint32_t writer = lifetime.first_use; + if (writer >= commands.size()) { + return false; + } + // Release values whose last use is before this value's first command. + for (size_t i = 0; i < live.size();) { + if (live[i].last_use < writer) { + free_ranges.push_back({ live[i].offset, live[i].size, std::move(live[i].users) }); + live[i] = std::move(live.back()); + live.pop_back(); + } else { + ++i; + } + } + std::sort(free_ranges.begin(), free_ranges.end(), + [](const FreeRange & lhs, const FreeRange & rhs) { return lhs.offset < rhs.offset; }); + + auto reusable = [&](const FreeRange & range) { + if (ancestors == nullptr) { + return true; + } + for (uint32_t user : range.users) { + if (user >= writer || !ancestors->precedes(user, writer)) { + return false; + } + } + return true; + }; + + // Best fit over runs of adjacent reusable ranges. + const size_t size = lifetime.allocation.size; + const size_t alignment = std::max(lifetime.allocation.alignment, 1); + size_t best_offset = 0, best_waste = SIZE_MAX, best_first = 0, best_last = 0; + bool found = false; + for (size_t first = 0; first < free_ranges.size(); ++first) { + if (!reusable(free_ranges[first]) || + align_up(free_ranges[first].offset, alignment) >= free_ranges[first].offset + free_ranges[first].size) { + continue; + } + size_t end = free_ranges[first].offset + free_ranges[first].size; + size_t last = first; + while (true) { + const size_t offset = align_up(free_ranges[first].offset, alignment); + if (offset + size <= end) { + const size_t waste = end - free_ranges[first].offset - size; + if (waste < best_waste) { + best_waste = waste, best_offset = offset, best_first = first, best_last = last, found = true; + } + break; + } + if (last + 1 >= free_ranges.size() || free_ranges[last + 1].offset != end || + !reusable(free_ranges[last + 1])) { + break; + } + ++last; + end = free_ranges[last].offset + free_ranges[last].size; + } + } + + size_t offset = 0; + if (found) { + offset = best_offset; + // Keep the parts of the run outside [offset, offset + size) free, with their own users. + std::vector kept; + const FreeRange & head = free_ranges[best_first]; + if (offset > head.offset) { + kept.push_back({ head.offset, offset - head.offset, head.users }); + } + const FreeRange & tail = free_ranges[best_last]; + const size_t tail_end = tail.offset + tail.size; + if (offset + size < tail_end) { + const size_t begin = std::max(offset + size, tail.offset); + kept.push_back({ begin, tail_end - begin, tail.users }); + } + free_ranges.erase(free_ranges.begin() + static_cast(best_first), + free_ranges.begin() + static_cast(best_last) + 1); + free_ranges.insert(free_ranges.end(), kept.begin(), kept.end()); + } else { + offset = align_up(packed.arena_size, alignment); + packed.arena_size = offset + size; + } + + lifetime.allocation.arena_offset = offset; + LiveRange range; + range.offset = offset; + range.size = size; + range.last_use = lifetime.last_use; + const auto users = users_by_value.find(lifetime.allocation.value.value); + if (users != users_by_value.end()) { + range.users = users->second; + } else { + range.users = { lifetime.first_use, lifetime.last_use }; + } + live.push_back(std::move(range)); + packed.allocations.push_back(lifetime.allocation); + } + + plan = std::move(packed); + return true; +} + +} // namespace +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch/transient-reuse-guard.h b/ggml/src/ggml-hrx/dispatch/transient-reuse-guard.h new file mode 100644 index 000000000000..ea798ea7a38e --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch/transient-reuse-guard.h @@ -0,0 +1,52 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Transient arena packing that does not add execution barriers. +// +// The stock packer hands a freed arena range to the next value that fits, often the value written by the +// very next command. The graph recorder tracks arena ranges, so that reuse becomes a write-after-read +// dependency on the command that last read the range, and HRX puts an execution barrier before the new +// writer even when the two commands share no data (two projections of one input, for example). +// +// This packer reuses a freed range only for a value whose first command already depends, through the +// program's own data dependencies, on every command that used the range before. Such a reuse adds no +// dependency the recorder would not already have. Otherwise the value gets fresh arena space. A program +// whose unguarded arena exceeds 16 MiB, or whose guarded arena would exceed twice the unguarded one plus +// 8 MiB (512-token prompt programs), keeps the stock packing. GGML_HRX_TRANSIENT_REUSE=legacy restores the stock packer everywhere, and +// GGML_HRX_LOG_TRANSIENT_REUSE=1 logs both arena sizes for every program. + +#pragma once + +#include "command-program.h" + +#include +#include + +namespace ggml::hrx { + +struct TransientReuseLifetime { + TransientAllocation allocation; + uint32_t first_use = 0; + uint32_t last_use = 0; +}; + +// Assigns arena offsets to |lifetimes| (main-program transients) and appends them to |plan|, growing +// plan.arena_size as needed. Returns false, leaving |plan| untouched, when the dependency-aware packer is +// disabled or the program is too large for it; the caller then uses the stock packer. +bool pack_transients_without_false_dependencies(const std::vector & commands, + std::vector & lifetimes, + TransientPlan & plan); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-add-id.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-add-id.cpp new file mode 100644 index 000000000000..2d4fc7f14664 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-add-id.cpp @@ -0,0 +1,95 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// GGML_OP_ADD_ID on F32 (ops/add_id_f32.loom): gpt-oss adds a per-expert bias to each MUL_MAT_ID output row, +// output[t][r] = input[t][r] + bias[ids[t][r]]. Without this every MoE layer leaves HRX three times. + +#include "dispatch-add-id.h" + +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include + +namespace ggml::hrx { + +namespace { + +static constexpr KernelCatalogRef kAddIdF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_add_id_f32"); + +bool flat_3d(const Value & value) { + return value.ne[3] == 1; +} + +bool match_add_id_f32(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->inputs.size() != 3) { + return false; + } + const Value * input = context.graph.values().find(node->inputs[0]); + const Value * bias = context.graph.values().find(node->inputs[1]); + const Value * ids = context.graph.values().find(node->inputs[2]); + const Value * output = context.graph.values().find(node->output); + if (input == nullptr || bias == nullptr || ids == nullptr || output == nullptr || input->type != GGML_TYPE_F32 || + bias->type != GGML_TYPE_F32 || ids->type != GGML_TYPE_I32 || output->type != GGML_TYPE_F32 || + !input->contiguous || !bias->contiguous || !output->contiguous || !flat_3d(*input) || !flat_3d(*output)) { + return false; + } + const int64_t width = input->ne[0]; + const int64_t rows = input->ne[1]; + const int64_t tokens = input->ne[2]; + const int64_t expert_count = bias->ne[1]; + if (output->ne[0] != width || output->ne[1] != rows || output->ne[2] != tokens || bias->ne[0] != width || + bias->ne[2] != 1 || bias->ne[3] != 1 || ids->ne[0] != rows || ids->ne[1] != tokens || ids->ne[2] != 1 || + ids->ne[3] != 1 || ids->nb[0] != sizeof(int32_t) || ids->nb[1] % sizeof(int32_t) != 0) { + return false; + } + const int64_t ids_stride = tokens > 1 ? static_cast(ids->nb[1] / sizeof(int32_t)) : rows; + // the kernel's launch ranges + if (width < 1 || width > 1048576 || rows < 1 || rows > 4096 || tokens < 1 || tokens > 65536 || + ids_stride < rows || ids_stride > 4096 || expert_count < 1 || expert_count > 4096) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kAddIdF32Kernel); + dispatch.kernel.integer_parameters.emplace("width", width); + dispatch.kernel.integer_parameters.emplace("rows", rows); + dispatch.kernel.integer_parameters.emplace("tokens", tokens); + dispatch.kernel.integer_parameters.emplace("ids_stride", ids_stride); + dispatch.kernel.integer_parameters.emplace("expert_count", expert_count); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ bias->id, 0, bias->byte_count }); + dispatch.bindings.push_back({ ids->id, 0, ids->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_add_id_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "extra.add_id_f32", + GGML_OP_ADD_ID, + DispatchMatchKind::SingleOp, + 1, + DispatchSource::Common, + match_add_id_f32, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-add-id.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-add-id.h new file mode 100644 index 000000000000..bc30ba4cda32 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-add-id.h @@ -0,0 +1,24 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +namespace ggml::hrx { + +void register_add_id_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-attention-sink.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-attention-sink.cpp new file mode 100644 index 000000000000..17fba12dceaa --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-attention-sink.cpp @@ -0,0 +1,135 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Attention sinks for FLASH_ATTN_EXT (gpt-oss): the FlashAttention dispatch runs without the sink and +// ops/attention_sink_f32.loom then rescales its output in place, row by row, by S / (S + exp(sink - M)), which is +// exact (see the kernel). dispatch-flash-attention.cpp accepts the 5-input node when attention_sinks_supported +// holds and appends this dispatch after its own. + +#include "dispatch-attention-sink.h" + +#include "dispatch-mul-mat-common.h" +#include "graph/op-params.h" +#include "hip/hip-dispatches.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include + +namespace ggml::hrx { + +namespace { + +static constexpr KernelCatalogRef kAttentionSinkF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_attention_sink_f32"); + +std::string index_config(int64_t value) { + return std::to_string(value); +} + +// True when the two values share storage and their byte ranges intersect. +bool overlaps(const Graph & graph, const Value & lhs, const Value & rhs) { + if (lhs.id == rhs.id) { + return true; + } + if (!graph.values().same_storage(lhs.id, rhs.id)) { + return false; + } + const size_t lhs_end = lhs.storage_offset + lhs.byte_count; + const size_t rhs_end = rhs.storage_offset + rhs.byte_count; + return lhs.storage_offset < rhs_end && rhs.storage_offset < lhs_end; +} + +} // namespace + +bool attention_sinks_supported(const Graph & graph, const GraphNode & node) { + if (node.op != GGML_OP_FLASH_ATTN_EXT || node.inputs.size() != 5) { + return false; + } + const Value * query = graph.values().find(node.inputs[0]); + const Value * sinks = graph.values().find(node.inputs[4]); + if (query == nullptr || sinks == nullptr || sinks->type != GGML_TYPE_F32 || !sinks->contiguous) { + return false; + } + const int64_t query_head_count = query->ne[2]; + return sinks->ne[0] == query_head_count && sinks->ne[1] == 1 && sinks->ne[2] == 1 && sinks->ne[3] == 1; +} + +bool append_attention_sink_dispatch(const Graph & graph, const GraphNode & node, DispatchMatch & dispatch_match) { + if (!attention_sinks_supported(graph, node)) { + return false; + } + const Value * query = graph.values().find(node.inputs[0]); + const Value * key = graph.values().find(node.inputs[1]); + const Value * mask = graph.values().find(node.inputs[3]); + const Value * sinks = graph.values().find(node.inputs[4]); + const Value * output = graph.values().find(node.output); + const FlashAttnExtParams * params = op_params_as(node.params); + if (key == nullptr || mask == nullptr || output == nullptr || params == nullptr) { + return false; + } + // The FlashAttention matchers have checked these layouts (query [tokens][heads][d] f32, key + // [capacity][kv_heads][d] f16, output [tokens][heads][dv] f32); the mask rows must be key_count apart. + const int64_t tokens = query->ne[1]; + const int64_t query_head_count = query->ne[2]; + const int64_t qk_head_size = query->ne[0]; + const int64_t key_capacity = key->ne[1]; + const int64_t key_value_head_count = key->ne[2]; + const int64_t value_head_size = output->ne[0]; + const int64_t key_count = mask->ne[0]; + if (mask->type != GGML_TYPE_F16 || mask->nb[0] != sizeof(ggml_fp16_t) || + mask->nb[1] != static_cast(key_count) * sizeof(ggml_fp16_t) || key_count < 1 || + key_count > key_capacity || query_head_count > 256 || key_value_head_count > 256 || qk_head_size % 16 != 0 || + value_head_size % 16 != 0 || qk_head_size > 576 || value_head_size > 576) { + return false; + } + // The output is rewritten in place after FlashAttention wrote it; no input may overlap its bytes (views of one + // allocation share a storage root but are distinct values, so compare storage, not value ids). + for (const Value * input : { query, key, mask, sinks }) { + if (overlaps(graph, *input, *output)) { + return false; + } + } + + // A HIP kernel add-on may take the rescale (hip/hip-dispatches.h). + if (Dispatch hip_dispatch; hip_attention_sink_dispatch({ query, key, mask, sinks, output, tokens, key_count, + query_head_count, key_value_head_count, qk_head_size, + value_head_size, params->scale }, + hip_dispatch)) { + dispatch_match.dispatches.push_back(std::move(hip_dispatch)); + return true; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kAttentionSinkF32Kernel); + dispatch.kernel.integer_parameters.emplace("tokens", tokens); + dispatch.kernel.integer_parameters.emplace("key_count", key_count); + dispatch.kernel.integer_parameters.emplace("key_capacity", key_capacity); + dispatch.kernel.compile_parameters.emplace("ggml.attention_sink.query_head_count", index_config(query_head_count)); + dispatch.kernel.compile_parameters.emplace("ggml.attention_sink.key_value_head_count", + index_config(key_value_head_count)); + dispatch.kernel.compile_parameters.emplace("ggml.attention_sink.qk_head_size", index_config(qk_head_size)); + dispatch.kernel.compile_parameters.emplace("ggml.attention_sink.value_head_size", index_config(value_head_size)); + dispatch.kernel.compile_parameters.emplace("ggml.attention_sink.scale", common_to_config_value(params->scale)); + dispatch.bindings.push_back({ query->id, 0, query->byte_count }); + dispatch.bindings.push_back({ key->id, 0, key->byte_count }); + dispatch.bindings.push_back({ mask->id, 0, mask->byte_count }); + dispatch.bindings.push_back({ sinks->id, 0, sinks->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-attention-sink.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-attention-sink.h new file mode 100644 index 000000000000..75c91aa2fdb1 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-attention-sink.h @@ -0,0 +1,29 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include "dispatch_registration/dispatch-registry.h" +#include "graph/graph.h" + +namespace ggml::hrx { + +// A 5-input FLASH_ATTN_EXT whose fifth input is F32 sinks, one per query head. +bool attention_sinks_supported(const Graph & graph, const GraphNode & node); + +// Appends the in-place sink rescale of the node's output (run after the FlashAttention dispatch). +bool append_attention_sink_dispatch(const Graph & graph, const GraphNode & node, DispatchMatch & dispatch_match); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-binary.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-binary.cpp new file mode 100644 index 000000000000..84f34aa28c79 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-binary.cpp @@ -0,0 +1,348 @@ +#include "dispatch-binary.h" + +#include "dispatch-layout-utils.h" +#include "dispatch-mul-mat-common.h" +#include "ggml.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kBinaryF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_binary_f32"); +static constexpr KernelCatalogRef kBinaryBcF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_binary_bc_f32"); +static constexpr KernelCatalogRef kBinarySwiGluSymmetricI4K32Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_binary_swiglu_symmetric_i4_k32"); + +static bool same_shape(const Value & lhs, const Value & rhs) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (lhs.ne[i] != rhs.ne[i]) { + return false; + } + } + return true; +} + +static bool positive_shape(const Value & value) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (value.ne[i] <= 0) { + return false; + } + } + return value.element_count > 0; +} + +static bool packed_f32_layout(const Value & value) { + size_t expected_stride = sizeof(float); + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (value.nb[i] != expected_stride) { + return false; + } + expected_stride *= static_cast(value.ne[i]); + } + return true; +} + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static bool storage_ranges_disjoint(const Value & lhs, size_t lhs_byte_count, const Value & rhs, size_t rhs_byte_count) { + if (lhs.storage != rhs.storage) { + return true; + } + if (lhs.storage_offset > std::numeric_limits::max() - lhs_byte_count || + rhs.storage_offset > std::numeric_limits::max() - rhs_byte_count) { + return false; + } + return lhs.storage_offset + lhs_byte_count <= rhs.storage_offset || + rhs.storage_offset + rhs_byte_count <= lhs.storage_offset; +} + +static bool storage_ranges_disjoint(const Value & lhs, const Value & rhs) { + return storage_ranges_disjoint(lhs, lhs.byte_count, rhs, rhs.byte_count); +} + +static bool binary_output_storage_is_safe(const Value & lhs, const Value & rhs, const Value & output) { + return storage_ranges_disjoint(lhs, output) && storage_ranges_disjoint(rhs, output); +} + +static bool binary_noalias_storage_is_safe(const Value & lhs, const Value & rhs, const Value & output) { + return storage_ranges_disjoint(lhs, rhs) && binary_output_storage_is_safe(lhs, rhs, output); +} + +static bool binary_output_storage_is_safe(const Value & lhs, + size_t lhs_byte_count, + const Value & rhs, + size_t rhs_byte_count, + const Value & output) { + return storage_ranges_disjoint(lhs, lhs_byte_count, output, output.byte_count) && + storage_ranges_disjoint(rhs, rhs_byte_count, output, output.byte_count); +} + +static bool supported_source_layout(const Graph & graph, const Value & value) { + if (value.alias_source.value < 0) { + return true; + } + if (value.storage_offset != 0) { + return false; + } + const GraphNode * producer = graph.index().producer(value.id); + return producer != nullptr && producer->op == GGML_OP_RESHAPE; +} + +static bool broadcastable_to(const Value & source, const Value & output) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (source.ne[i] != output.ne[i] && source.ne[i] != 1) { + return false; + } + } + return true; +} + +static bool binary_kind_allows_broadcast(BinaryKind kind, const Value & lhs, const Value & rhs, const Value & output) { + const bool lhs_full = same_shape(lhs, output); + const bool rhs_full = same_shape(rhs, output); + if (!broadcastable_to(lhs, output) || !broadcastable_to(rhs, output) || (!lhs_full && !rhs_full)) { + return false; + } + + switch (kind) { + case BinaryKind::Add: + case BinaryKind::Mul: + return true; + case BinaryKind::Sub: + case BinaryKind::Div: + return lhs_full; + case BinaryKind::SwiGLU: + case BinaryKind::GeGLU: + case BinaryKind::RegLU: + case BinaryKind::GeGLUErf: + case BinaryKind::GeGLUQuick: + return lhs_full && rhs_full; + } + return false; +} + +static uint32_t broadcast_dim_flag(const Value & source, const Value & output, int dim) { + return source.ne[dim] == 1 && output.ne[dim] != 1 ? 1 : 0; +} + +static void add_binary_shape_parameters(Dispatch & dispatch, + const Value & lhs, + const Value & rhs, + const Value & output) { + dispatch.kernel.integer_parameters.emplace("element_count", output.element_count); + dispatch.kernel.integer_parameters.emplace("ne0", output.ne[0]); + dispatch.kernel.integer_parameters.emplace("ne1", output.ne[1]); + dispatch.kernel.integer_parameters.emplace("ne2", output.ne[2]); + dispatch.kernel.integer_parameters.emplace("ne3", output.ne[3]); + dispatch.kernel.integer_parameters.emplace("src0_element_count", lhs.element_count); + dispatch.kernel.integer_parameters.emplace("src1_element_count", rhs.element_count); +} + +static void add_broadcast_config(Dispatch & dispatch, + const char * config_prefix, + const char * source_prefix, + const Value & source, + const Value & output) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + dispatch.kernel.compile_parameters.emplace( + std::string(config_prefix) + source_prefix + "_broadcast_dim" + std::to_string(i), + std::to_string(broadcast_dim_flag(source, output, i))); + } +} + +static void bind_binary_source_buffer(Dispatch & dispatch, const Value & value, size_t byte_count) { + dispatch.bindings.push_back({ value.storage_root, value.storage_offset, byte_count }); +} + +static void bind_binary_buffers(Dispatch & dispatch, + const Value & lhs, + size_t lhs_byte_count, + const Value & rhs, + size_t rhs_byte_count, + const Value & output) { + bind_binary_source_buffer(dispatch, lhs, lhs_byte_count); + bind_binary_source_buffer(dispatch, rhs, rhs_byte_count); + dispatch.bindings.push_back({ output.id, 0, output.byte_count }); +} + +static void add_binary_strided_parameters(Dispatch & dispatch, + const Value & lhs, + const Value & rhs, + const Value & output, + size_t lhs_byte_count, + size_t rhs_byte_count) { + dispatch.kernel.integer_parameters.emplace("element_count", output.element_count); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.ne0", std::to_string(output.ne[0])); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.ne1", std::to_string(output.ne[1])); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.ne2", std::to_string(output.ne[2])); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src0_stride1", + std::to_string(lhs.nb[1] / sizeof(float))); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src0_stride2", + std::to_string(lhs.nb[2] / sizeof(float))); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src0_stride3", + std::to_string(lhs.nb[3] / sizeof(float))); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src1_stride1", + std::to_string(rhs.nb[1] / sizeof(float))); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src1_stride2", + std::to_string(rhs.nb[2] / sizeof(float))); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src1_stride3", + std::to_string(rhs.nb[3] / sizeof(float))); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src0_span", + std::to_string(lhs_byte_count / sizeof(float))); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src1_span", + std::to_string(rhs_byte_count / sizeof(float))); +} + +static bool match_binary_swiglu_symmetric_i4_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_GLU || node->inputs.size() != 2 || + !common_is_swiglu_params(node->params)) { + return false; + } + + const Value * lhs = graph_value(context.graph, node->inputs[0]); + const Value * rhs = graph_value(context.graph, node->inputs[1]); + const Value * output = graph_value(context.graph, node->output); + if (lhs == nullptr || rhs == nullptr || output == nullptr || lhs->type != GGML_TYPE_F32 || + rhs->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || !same_shape(*lhs, *output) || + !same_shape(*rhs, *output) || !packed_f32_layout(*lhs) || !packed_f32_layout(*rhs) || + !packed_f32_layout(*output) || output->alias_source.value >= 0 || + !supported_source_layout(context.graph, *lhs) || !supported_source_layout(context.graph, *rhs) || + !binary_noalias_storage_is_safe(*lhs, *rhs, *output) || output->ne[0] < 256 || output->ne[0] > 32768 || + output->ne[0] % 64 != 0 || output->element_count <= 0 || output->element_count % output->ne[0] != 0 || + !common_has_symmetric_i4_lowrow_consumer(context.graph, *output)) { + return false; + } + + const int64_t input_size = output->ne[0]; + const int64_t token_count = output->element_count / input_size; + if (token_count < 1 || token_count > 16) { + return false; + } + const CommonSymmetricI4ActivationLayout activation_layout = + common_symmetric_i4_activation_layout(input_size, token_count); + if (activation_layout.total_bytes == 0) { + return false; + } + const ValueId activation = context.next_plan_value; + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kBinarySwiGluSymmetricI4K32Kernel); + dispatch.kernel.compile_parameters.emplace("ggml.binary_swiglu_symmetric_i4.input_size", + std::to_string(input_size)); + dispatch.kernel.compile_parameters.emplace("ggml.binary_swiglu_symmetric_i4.token_count", + std::to_string(token_count)); + bind_binary_buffers(dispatch, *lhs, lhs->byte_count, *rhs, rhs->byte_count, *output); + dispatch.bindings.push_back({ activation, 0, activation_layout.payload_bytes }); + dispatch.bindings.push_back({ activation, activation_layout.scales_offset, activation_layout.metadata_bytes }); + dispatch.bindings.push_back({ activation, activation_layout.sums_offset, activation_layout.metadata_bytes }); + + Status metadata_status; + if (!match.metadata.append_alternate_value({ output->id, activation, GGML_TYPE_COUNT, activation_layout.total_bytes, + kCommonSymmetricI4K32ActivationAlternateName }, + metadata_status)) { + match.status.append(metadata_status); + return false; + } + + match.covered_nodes.push_back(context.root_index); + match.dispatches.push_back(std::move(dispatch)); + match.transients.push_back( + { activation, kCommonSymmetricI4K32ActivationAlternateName, activation_layout.total_bytes, 256 }); + return match.status.success(); +} + +static bool match_binary_f32_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->inputs.size() != 2) { + return false; + } + + const BinaryParams * params = op_params_as(node->params); + if (params == nullptr || !binary_kind_supported(params->op)) { + return false; + } + + const Value * output = graph_value(context.graph, node->output); + const Value * lhs = graph_value(context.graph, node->inputs[0]); + const Value * rhs = graph_value(context.graph, node->inputs[1]); + if (output == nullptr || lhs == nullptr || rhs == nullptr) { + return false; + } + + size_t lhs_byte_count = 0; + size_t rhs_byte_count = 0; + if (output->type != GGML_TYPE_F32 || lhs->type != GGML_TYPE_F32 || rhs->type != GGML_TYPE_F32 || + !positive_shape(*output) || !output->contiguous || !packed_f32_layout(*output) || + !strided_f32_storage_span_bytes(*lhs, lhs_byte_count) || + !strided_f32_storage_span_bytes(*rhs, rhs_byte_count) || + output->alias_source.value >= 0 || + !binary_output_storage_is_safe(*lhs, lhs_byte_count, *rhs, rhs_byte_count, *output) || + static_cast(output->element_count) > std::numeric_limits::max()) { + return false; + } + + Dispatch dispatch; + if (same_shape(*lhs, *output) && same_shape(*rhs, *output)) { + dispatch.kernel = make_kernel_specialization(kBinaryF32Kernel); + add_binary_strided_parameters(dispatch, *lhs, *rhs, *output, lhs_byte_count, rhs_byte_count); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.op", + std::to_string(binary_kind_config_value(params->op))); + } else { + if (!lhs->contiguous || !rhs->contiguous || !packed_f32_layout(*lhs) || !packed_f32_layout(*rhs)) { + return false; + } + if (!binary_kind_allows_broadcast(params->op, *lhs, *rhs, *output)) { + return false; + } + dispatch.kernel = make_kernel_specialization(kBinaryBcF32Kernel); + add_binary_shape_parameters(dispatch, *lhs, *rhs, *output); + dispatch.kernel.compile_parameters.emplace("ggml.binary_bc_f32.op", + std::to_string(binary_kind_config_value(params->op))); + add_broadcast_config(dispatch, "ggml.binary_bc_f32.", "src0", *lhs, *output); + add_broadcast_config(dispatch, "ggml.binary_bc_f32.", "src1", *rhs, *output); + } + bind_binary_buffers(dispatch, *lhs, lhs_byte_count, *rhs, rhs_byte_count, *output); + + match.covered_nodes.push_back(context.root_index); + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static void register_binary_dispatch_for(DispatchRegistryBuilder & registry, ggml_op root_op) { + registry.add({ + "common.binary_f32", + root_op, + DispatchMatchKind::SingleOp, + 0, + DispatchSource::Common, + match_binary_f32_dispatch, + }); +} + +} // namespace + +void register_binary_dispatch(DispatchRegistryBuilder & registry) { + registry.add({ + "common.binary_swiglu_symmetric_i4_k32", + GGML_OP_GLU, + DispatchMatchKind::Fused, + 100, + DispatchSource::Common, + match_binary_swiglu_symmetric_i4_dispatch, + }); + register_binary_dispatch_for(registry, GGML_OP_ADD); + register_binary_dispatch_for(registry, GGML_OP_SUB); + register_binary_dispatch_for(registry, GGML_OP_MUL); + register_binary_dispatch_for(registry, GGML_OP_DIV); + register_binary_dispatch_for(registry, GGML_OP_GLU); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-binary.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-binary.h new file mode 100644 index 000000000000..6445daded097 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-binary.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_binary_dispatch(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-common.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-common.cpp new file mode 100644 index 000000000000..4d66217e6f12 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-common.cpp @@ -0,0 +1,61 @@ +#include "dispatch-common.h" + +#include "dispatch-binary.h" +#include "dispatch-copy.h" +#include "dispatch-flash-attention.h" +#include "dispatch-gated-mul-mat-id.h" +#include "dispatch-gated-mul-mat.h" +#include "dispatch-gather-add.h" +#include "dispatch-get-rows.h" +#include "dispatch-glu.h" +#include "dispatch-mul-mat-id.h" +#include "dispatch-mul-mat.h" +#include "dispatch-rmsnorm.h" +#include "dispatch-rope-set-rows.h" +#include "dispatch-scale.h" +#include "dispatch-grouped-mul-mat.h" +#include "dispatch-small-rows.h" +#include "dispatch-mul-mat-id-decode.h" +#include "dispatch-res-scale-pair.h" +#include "dispatch-hadamard.h" +#include "dispatch-zaya-cca-conv.h" +#include "dispatch-zaya-cca-qk-norm.h" +#include "dispatch-kquant-decode.h" +#include "dispatch-add-id.h" +#include "dispatch-softplus.h" +#include "dispatch-swiglu-oai.h" +#include "dispatch-unary.h" +#include "hip/hip-dispatches.h" + +namespace ggml::hrx { + +void register_common_dispatches(DispatchRegistryBuilder & registry) { + register_binary_dispatch(registry); + register_copy_dispatch(registry); + register_flash_attention_dispatches(registry); + register_gated_mul_mat_id_dispatches(registry); + register_gated_mul_mat_dispatches(registry); + register_gather_add_dispatch(registry); + register_grouped_mul_mat_dispatch(registry); + register_get_rows_dispatches(registry); + register_glu_dispatches(registry); + register_mul_mat_id_dispatches(registry); + register_mul_mat_dispatches(registry); + register_rope_set_rows_dispatches(registry); + register_scale_dispatch(registry); + register_small_rows_dispatches(registry); + register_mul_mat_id_decode_dispatches(registry); + register_res_scale_pair_dispatches(registry); + register_hadamard_dispatches(registry); + register_zaya_cca_conv_dispatches(registry); + register_zaya_cca_qk_norm_dispatches(registry); + register_kquant_decode_dispatches(registry); + register_softplus_dispatches(registry); + register_add_id_dispatches(registry); + register_swiglu_oai_dispatches(registry); + register_unary_dispatch(registry); + register_rmsnorm_dispatches(registry); + register_hip_dispatches(registry); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-common.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-common.h new file mode 100644 index 000000000000..d7fe2aed8246 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-common.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_common_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-copy.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-copy.cpp new file mode 100644 index 000000000000..cb531ccd1cfe --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-copy.cpp @@ -0,0 +1,227 @@ +#include "dispatch-copy.h" + +#include "dispatch-layout-utils.h" +#include "ggml.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr uint64_t kMaximumCopyElements = uint64_t{ 1 } << 30; +static constexpr KernelCatalogRef kCopyF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_copy_f32"); +static constexpr KernelCatalogRef kCopyStridedSourceF32Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_copy_strided_source_f32"); +static constexpr KernelCatalogRef kConcatDim0F32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_concat_dim0_f32"); +static constexpr KernelCatalogRef kConcatDim0StridedSourceF32Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_concat_dim0_strided_source_f32"); + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static bool packed_f32_layout(const Value & value) { + size_t expected_stride = sizeof(float); + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (value.nb[i] != expected_stride) { + return false; + } + expected_stride *= static_cast(value.ne[i]); + } + return true; +} + +static bool f32_element_strides(const Value & value) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (value.ne[i] <= 0 || value.nb[i] <= 0 || value.nb[i] % sizeof(float) != 0) { + return false; + } + } + return value.byte_count % sizeof(float) == 0; +} + +static void add_concat_strided_source_parameters(KernelSpecialization & kernel, const Value & lhs, const Value & rhs) { + static constexpr const char * kPrefix = "ggml.concat_dim0_strided_source_f32."; + kernel.integer_parameters.emplace("lhs_span", lhs.byte_count / sizeof(float)); + kernel.integer_parameters.emplace("rhs_span", rhs.byte_count / sizeof(float)); + kernel.compile_parameters.emplace(std::string(kPrefix) + "lhs_width", std::to_string(lhs.ne[0])); + kernel.compile_parameters.emplace(std::string(kPrefix) + "rhs_width", std::to_string(rhs.ne[0])); + kernel.compile_parameters.emplace(std::string(kPrefix) + "row_count", + std::to_string(lhs.ne[1] * lhs.ne[2] * lhs.ne[3])); + kernel.compile_parameters.emplace(std::string(kPrefix) + "row_ne1", std::to_string(lhs.ne[1])); + kernel.compile_parameters.emplace(std::string(kPrefix) + "row_ne2", std::to_string(lhs.ne[2])); + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + kernel.compile_parameters.emplace(std::string(kPrefix) + "lhs_stride" + std::to_string(i), + std::to_string(lhs.nb[i] / sizeof(float))); + kernel.compile_parameters.emplace(std::string(kPrefix) + "rhs_stride" + std::to_string(i), + std::to_string(rhs.nb[i] / sizeof(float))); + } +} + +static bool make_copy_f32_dispatch(const Value & source, const Value & output, Dispatch & dispatch) { + if (source.type != GGML_TYPE_F32 || output.type != GGML_TYPE_F32 || !output.contiguous || + !packed_f32_layout(output) || !f32_element_strides(source) || source.element_count <= 0 || + source.element_count != output.element_count || + static_cast(source.element_count) > kMaximumCopyElements || source.storage == output.storage) { + return false; + } + + size_t source_span = source.byte_count; + if (packed_f32_layout(source)) { + dispatch.kernel = make_kernel_specialization(kCopyF32Kernel); + } else { + if (!strided_f32_storage_span_bytes(source, source_span)) { + return false; + } + const size_t source_span_elements = source_span / sizeof(float); + if (source_span_elements > kMaximumCopyElements) { + return false; + } + dispatch.kernel = make_kernel_specialization(kCopyStridedSourceF32Kernel); + dispatch.kernel.integer_parameters.emplace("source_span", source_span_elements); + dispatch.kernel.compile_parameters.emplace("ggml.copy_strided_source_f32.ne0", std::to_string(source.ne[0])); + dispatch.kernel.compile_parameters.emplace("ggml.copy_strided_source_f32.ne1", std::to_string(source.ne[1])); + dispatch.kernel.compile_parameters.emplace("ggml.copy_strided_source_f32.ne2", std::to_string(source.ne[2])); + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + dispatch.kernel.compile_parameters.emplace("ggml.copy_strided_source_f32.stride" + std::to_string(i), + std::to_string(source.nb[i] / sizeof(float))); + } + } + dispatch.kernel.integer_parameters.emplace("element_count", source.element_count); + dispatch.bindings.push_back({ source.storage_root, source.storage_offset, source_span }); + dispatch.bindings.push_back({ output.id, 0, output.byte_count }); + return true; +} + +static bool match_copy_f32_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_CPY || node->inputs.size() != 2) { + return false; + } + + const Value * source = graph_value(context.graph, node->inputs[0]); + const Value * target = graph_value(context.graph, node->inputs[1]); + const Value * output = graph_value(context.graph, node->output); + if (source == nullptr || target == nullptr || output == nullptr || target->type != GGML_TYPE_F32 || + !target->contiguous || !packed_f32_layout(*target) || target->element_count != output->element_count || + target->byte_count != output->byte_count || source->storage == target->storage) { + return false; + } + + Dispatch dispatch; + if (!make_copy_f32_dispatch(*source, *output, dispatch)) { + return false; + } + + match.covered_nodes.push_back(context.root_index); + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_cont_f32_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_CONT || node->inputs.size() != 1) { + return false; + } + + const Value * source = graph_value(context.graph, node->inputs[0]); + const Value * output = graph_value(context.graph, node->output); + if (source == nullptr || output == nullptr) { + return false; + } + + Dispatch dispatch; + if (!make_copy_f32_dispatch(*source, *output, dispatch)) { + return false; + } + + match.covered_nodes.push_back(context.root_index); + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_concat_dim0_f32_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_CONCAT || node->inputs.size() != 2) { + return false; + } + + const Value * lhs = graph_value(context.graph, node->inputs[0]); + const Value * rhs = graph_value(context.graph, node->inputs[1]); + const Value * output = graph_value(context.graph, node->output); + if (lhs == nullptr || rhs == nullptr || output == nullptr || lhs->type != GGML_TYPE_F32 || + rhs->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || !output->contiguous || + !packed_f32_layout(*output) || !f32_element_strides(*lhs) || !f32_element_strides(*rhs) || lhs->ne[0] <= 0 || + lhs->ne[0] > 65536 || rhs->ne[0] <= 0 || rhs->ne[0] > 65536 || output->ne[0] != lhs->ne[0] + rhs->ne[0] || + output->element_count <= 0 || static_cast(output->element_count) > kMaximumCopyElements || + lhs->storage == rhs->storage || lhs->storage == output->storage || rhs->storage == output->storage) { + return false; + } + for (int dim = 1; dim < GGML_MAX_DIMS; ++dim) { + if (lhs->ne[dim] != rhs->ne[dim] || lhs->ne[dim] != output->ne[dim]) { + return false; + } + } + + const int64_t row_count = output->element_count / output->ne[0]; + if (row_count < 1 || row_count > 1048576) { + return false; + } + + Dispatch dispatch; + if (packed_f32_layout(*lhs) && packed_f32_layout(*rhs)) { + dispatch.kernel = make_kernel_specialization(kConcatDim0F32Kernel); + dispatch.kernel.compile_parameters.emplace("ggml.concat_dim0_f32.lhs_width", std::to_string(lhs->ne[0])); + dispatch.kernel.compile_parameters.emplace("ggml.concat_dim0_f32.rhs_width", std::to_string(rhs->ne[0])); + dispatch.kernel.compile_parameters.emplace("ggml.concat_dim0_f32.row_count", std::to_string(row_count)); + } else { + const size_t lhs_span = lhs->byte_count / sizeof(float); + const size_t rhs_span = rhs->byte_count / sizeof(float); + if (lhs_span > kMaximumCopyElements || rhs_span > kMaximumCopyElements) { + return false; + } + dispatch.kernel = make_kernel_specialization(kConcatDim0StridedSourceF32Kernel); + add_concat_strided_source_parameters(dispatch.kernel, *lhs, *rhs); + } + dispatch.bindings.push_back({ lhs->id, 0, lhs->byte_count }); + dispatch.bindings.push_back({ rhs->id, 0, rhs->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + + match.covered_nodes.push_back(context.root_index); + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_copy_dispatch(DispatchRegistryBuilder & registry) { + registry.add({ + "common.copy_f32", + GGML_OP_CPY, + DispatchMatchKind::SingleOp, + 0, + DispatchSource::Common, + match_copy_f32_dispatch, + }); + registry.add({ + "common.cont_f32", + GGML_OP_CONT, + DispatchMatchKind::SingleOp, + 0, + DispatchSource::Common, + match_cont_f32_dispatch, + }); + registry.add({ + "common.concat_dim0_f32", + GGML_OP_CONCAT, + DispatchMatchKind::SingleOp, + 0, + DispatchSource::Common, + match_concat_dim0_f32_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-copy.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-copy.h new file mode 100644 index 000000000000..6228aac1946b --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-copy.h @@ -0,0 +1,9 @@ +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +namespace ggml::hrx { + +void register_copy_dispatch(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-flash-attention.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-flash-attention.cpp new file mode 100644 index 000000000000..20e9150450d5 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-flash-attention.cpp @@ -0,0 +1,713 @@ +#include "dispatch-attention-sink.h" +#include "dispatch-flash-attention.h" + +#include "ggml.h" +#include "graph/graph-matcher.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +// Test knob for the decode-split partial transients' alignment. Production is 4096; +// the engine#123/#140 rig oracle used 256, a layout that exposes the divergence. +// Defaults to 4096 so production behaviour is unchanged unless explicitly requested. +static size_t decode_split_partial_alignment() { + static const size_t value = []() -> size_t { + const char * env = std::getenv("GGML_HRX_FA_PARTIAL_ALIGN"); + if (env == nullptr || *env == '\0') { + return 4096u; + } + char * end = nullptr; + const long parsed = std::strtol(env, &end, 10); + if (end == env || parsed <= 0 || parsed > (1 << 20)) { + return 4096u; + } + return static_cast(parsed); + }(); + return value; +} + +static constexpr KernelCatalogRef kFlashAttentionF32F16WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_flash_attention_f32_f16_wmma"); +static constexpr KernelCatalogRef kFlashAttentionDecodeSplitNextQ8Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_flash_attention_decode_split_f32_f16_wmma_next_q8"); +static constexpr KernelCatalogRef kCopyTransposeF16Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_copy_transpose_f16"); +static constexpr int64_t kDecodeRowCapacity = 16; +static constexpr int64_t kDecodeKvTileSize = 64; +// The decode-split reduce kernel (ggml.flash_attention.decode_split.reduce_fused, in the loom-libs +// kernel corpus) has exactly two template.def implementations: one for key_value_token_capacity +// 64-256, one for 257-2048. Nothing covers above 2048 - its cooperative reducer gives each +// subgroup lane one partial KV block (capped at 32) and its per-workgroup scratch buffer is sized +// for exactly 32 blocks x 64 tokens. Offering this dispatch above the ceiling makes the kernel +// selector reject every candidate ("all_rejected") and the whole decode fail; match this bound so +// the scheduler falls through to the general flash_attention_f32_f16_wmma dispatch instead (lower +// priority, still correct here, just not split-parallelized for very long decode contexts). +static constexpr int64_t kDecodeSplitMaxKeyValueTokenCapacity = 32768; +// ggml.copy_transpose_f16 declares row_count and column_count in [32, 32768] (copy_f32.loom). +// Past that the JIT refuses the specialization ("violates constraint 'range'") and the whole +// prompt batch fails, so a longer context keeps V in the row-major cache layout instead. +static constexpr int64_t kCopyTransposeF16MaxExtent = 32768; +static constexpr int64_t kPrefillQkHeadSizeBlock = 16; +static constexpr int64_t kPrefillValueHeadSizeBlock = 64; +static constexpr int64_t kPrefillMinQkHeadSize = kPrefillQkHeadSizeBlock; +static constexpr int64_t kPrefillMinValueHeadSize = kPrefillValueHeadSizeBlock; +// Current cap keeps full-head Q/K staging within the kernel's fixed LDS budget +// (576 = DeepSeek-V2/V3 / GLM-4.7-Flash MLA key_length; 592 is the LDS ceiling). +static constexpr int64_t kPrefillMaxHeadSize = 576; + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static bool nearly_equal(float lhs, float rhs) { + return std::fabs(lhs - rhs) <= 1.0e-6f; +} + +static bool is_supported_token_count(int64_t token_count) { + return token_count >= 1 && token_count <= 2048; +} + +static bool is_supported_key_value_token_count(int64_t token_count) { + return token_count >= 1 && token_count <= 262144; +} + +static bool is_supported_decode_key_value_token_count(int64_t token_count) { + return token_count >= 1 && token_count <= 262144; +} + +static bool is_supported_decode_query_length(int64_t query_length) { + return query_length >= 1 && query_length < kDecodeRowCapacity; +} + +static bool is_supported_head_count(int64_t head_count) { + return head_count >= 1 && head_count <= 64; +} + +static bool is_supported_qk_head_size(int64_t head_size) { + return head_size >= kPrefillMinQkHeadSize && head_size <= kPrefillMaxHeadSize && + head_size % kPrefillQkHeadSizeBlock == 0; +} + +static bool is_supported_value_head_size(int64_t head_size) { + return head_size >= kPrefillMinValueHeadSize && head_size <= kPrefillMaxHeadSize && + head_size % kPrefillValueHeadSizeBlock == 0; +} + +static bool has_query_layout(const Value & value, int64_t query_head_count, int64_t head_size) { + const size_t element_size = sizeof(float); + return value.nb[0] == element_size && + value.nb[1] == static_cast(query_head_count * head_size) * element_size && + (value.ne[2] == 1 || value.nb[2] == static_cast(head_size) * element_size); +} + +static bool has_key_value_layout(const Value & value, int64_t key_value_head_count, int64_t head_size) { + const size_t element_size = sizeof(ggml_fp16_t); + return value.nb[0] == element_size && + value.nb[1] == static_cast(key_value_head_count * head_size) * element_size && + (value.ne[2] == 1 || value.nb[2] == static_cast(head_size) * element_size); +} + +static bool has_mask_layout(const Value & value, int64_t key_value_token_count) { + const size_t element_size = sizeof(ggml_fp16_t); + return value.nb[0] == element_size && value.nb[1] == static_cast(key_value_token_count) * element_size; +} + +static bool has_output_layout(const Value & value, int64_t query_head_count, int64_t head_size) { + const size_t element_size = sizeof(float); + return value.nb[0] == element_size && value.nb[1] == static_cast(head_size) * element_size && + value.nb[2] == static_cast(query_head_count * head_size) * element_size; +} + +static bool has_flash_attention_params(const GraphNode & node) { + const FlashAttnExtParams * params = op_params_as(node.params); + if (params == nullptr) { + return false; + } + return nearly_equal(params->max_bias, 0.0f) && nearly_equal(params->logit_softcap, 0.0f) && + (params->prec == GGML_PREC_DEFAULT || params->prec == GGML_PREC_F32); +} + +static size_t attention_mask_byte_count(int64_t query_token_count, int64_t key_value_token_count) { + if (query_token_count <= 0 || key_value_token_count <= 0) { + return 0; + } + return static_cast(query_token_count) * static_cast(key_value_token_count) * sizeof(ggml_fp16_t); +} + +static std::string to_config_value(int64_t value) { + return std::to_string(value); +} + +static size_t q8_1_x4_byte_count(int64_t row_count, int64_t hidden_size) { + if (row_count <= 0 || hidden_size <= 0) { + return 0; + } + // Packed Q8 stores four 32-element blocks in each physical group. + const int64_t padded_hidden_size = (hidden_size + 127) / 128 * 128; + return static_cast(row_count) * ggml_row_size(GGML_TYPE_Q8_1, padded_hidden_size); +} + +static int64_t ceil_div(int64_t value, int64_t divisor) { + return (value + divisor - 1) / divisor; +} + +static ValueId match_value(const DispatchMatchContext & context, const DispatchMatch & dispatch_match, int32_t offset) { + return ValueId(context.next_plan_value.value + static_cast(dispatch_match.transients.size()) + + static_cast(dispatch_match.completion_counter_requests.size()) + offset); +} + +struct FlashAttentionMatch { + const Value * query = nullptr; + const Value * key = nullptr; + const Value * value = nullptr; + const Value * mask = nullptr; + const Value * output = nullptr; + const GraphNode * output_layout = nullptr; + ValueId mask_binding_value; + size_t mask_binding_bytes = 0; + int64_t query_token_count = 0; + int64_t key_value_token_count = 0; + int64_t query_head_count = 0; + int64_t key_value_head_count = 0; + int64_t qk_head_size = 0; + int64_t value_head_size = 0; + float attention_scale = 0.0f; + + bool matched() const { + return query != nullptr && key != nullptr && value != nullptr && mask != nullptr && output != nullptr; + } +}; + +struct DecodeSplitFlashAttentionMatch { + const Value * query = nullptr; + const Value * key = nullptr; + const Value * value = nullptr; + const Value * mask = nullptr; + const Value * output = nullptr; + const GraphNode * output_layout = nullptr; + int64_t query_token_count = 0; + int64_t key_value_token_count = 0; + int64_t key_value_capacity = 0; + int64_t query_head_count = 0; + int64_t key_value_head_count = 0; + int64_t qk_head_size = 0; + int64_t value_head_size = 0; + float attention_scale = 0.0f; + + bool matched() const { + return query != nullptr && key != nullptr && value != nullptr && mask != nullptr && output != nullptr; + } +}; + +static FlashAttentionMatch match_flash_attention_f32_f16(const Graph & graph, + const CommandPlan & plan, + const GraphNode * node) { + FlashAttentionMatch match; + if (node == nullptr || node->op != GGML_OP_FLASH_ATTN_EXT || + (node->inputs.size() != 4 && !attention_sinks_supported(graph, *node))) { + return match; + } + + const Value * query = graph_value(graph, node->inputs[0]); + const Value * key = graph_value(graph, node->inputs[1]); + const Value * value = graph_value(graph, node->inputs[2]); + const Value * mask = graph_value(graph, node->inputs[3]); + const Value * output = graph_value(graph, node->output); + if (query == nullptr || key == nullptr || value == nullptr || mask == nullptr || output == nullptr) { + return {}; + } + if (query->type != GGML_TYPE_F32 || key->type != GGML_TYPE_F16 || value->type != GGML_TYPE_F16 || + mask->type != GGML_TYPE_F16 || output->type != GGML_TYPE_F32) { + return {}; + } + const int64_t qk_head_size = query->ne[0]; + const int64_t value_head_size = value->ne[0]; + if (!is_supported_qk_head_size(qk_head_size) || !is_supported_value_head_size(value_head_size) || + !has_flash_attention_params(*node)) { + return {}; + } + if (key->ne[0] != qk_head_size || output->ne[0] != value_head_size) { + return {}; + } + if (query->ne[3] != 1 || key->ne[3] != 1 || value->ne[3] != 1 || output->ne[3] != 1 || mask->ne[2] != 1 || + mask->ne[3] != 1) { + return {}; + } + + const int64_t query_token_count = query->ne[1]; + const int64_t query_head_count = query->ne[2]; + const int64_t key_value_capacity = key->ne[1]; + const int64_t key_value_head_count = key->ne[2]; + if (!is_supported_token_count(query_token_count) || key_value_capacity < query_token_count || + !is_supported_head_count(query_head_count) || !is_supported_head_count(key_value_head_count) || + query_head_count % key_value_head_count != 0) { + return {}; + } + if (value->ne[1] != key_value_capacity || value->ne[2] != key_value_head_count || + mask->ne[0] > key_value_capacity || mask->ne[1] != query_token_count || output->ne[1] != query_head_count || + output->ne[2] != query_token_count) { + return {}; + } + + int64_t key_value_token_count = key_value_capacity; + ValueId mask_binding_value = mask->id; + size_t mask_binding_bytes = mask->byte_count; + bool mask_binding_is_alternate = false; + const size_t compact_mask_bytes = attention_mask_byte_count(query_token_count, query_token_count); + const auto * compact_mask = find_alternate_value(plan, mask->id, GGML_TYPE_F16, compact_mask_bytes); + const bool mask_is_compact = mask->ne[0] == query_token_count; + const bool mask_is_capacity = mask->ne[0] == key_value_capacity; + if (compact_mask != nullptr && mask_is_capacity && mask->ne[0] > query_token_count) { + key_value_token_count = query_token_count; + mask_binding_value = compact_mask->alternate_value; + mask_binding_bytes = compact_mask->byte_count; + mask_binding_is_alternate = true; + } else if (!mask_is_compact && !mask_is_capacity) { + return {}; + } + if (!is_supported_key_value_token_count(key_value_token_count)) { + return {}; + } + + if (!has_query_layout(*query, query_head_count, qk_head_size) || + !has_key_value_layout(*key, key_value_head_count, qk_head_size) || + // MLA (deepseek2/GLM) stores the value with the QK head-size stride (the KV cache row + // is padded to key_length), not the value head-size stride, so match the key stride (#95). + !has_key_value_layout(*value, key_value_head_count, qk_head_size) || + !has_output_layout(*output, query_head_count, value_head_size)) { + return {}; + } + if (!mask_binding_is_alternate && !has_mask_layout(*mask, key_value_token_count)) { + return {}; + } + + match.query = query; + match.key = key; + match.value = value; + match.mask = mask; + match.output = output; + match.output_layout = find_single_layout_alias_consumer(graph, output->id); + match.mask_binding_value = mask_binding_value; + match.mask_binding_bytes = mask_binding_bytes; + match.query_token_count = query_token_count; + match.key_value_token_count = key_value_token_count; + match.query_head_count = query_head_count; + match.key_value_head_count = key_value_head_count; + match.qk_head_size = qk_head_size; + match.value_head_size = value_head_size; + match.attention_scale = op_params_as(node->params)->scale; + return match; +} + +static DecodeSplitFlashAttentionMatch match_decode_split_flash_attention_f32_f16(const Graph & graph, + const GraphNode * node) { + DecodeSplitFlashAttentionMatch match; + if (node == nullptr || node->op != GGML_OP_FLASH_ATTN_EXT || node->inputs.size() != 4) { + return match; + } + + const Value * query = graph_value(graph, node->inputs[0]); + const Value * key = graph_value(graph, node->inputs[1]); + const Value * value = graph_value(graph, node->inputs[2]); + const Value * mask = graph_value(graph, node->inputs[3]); + const Value * output = graph_value(graph, node->output); + if (query == nullptr || key == nullptr || value == nullptr || mask == nullptr || output == nullptr) { + return {}; + } + if (query->type != GGML_TYPE_F32 || key->type != GGML_TYPE_F16 || value->type != GGML_TYPE_F16 || + mask->type != GGML_TYPE_F16 || output->type != GGML_TYPE_F32) { + return {}; + } + + const int64_t qk_head_size = query->ne[0]; + const int64_t value_head_size = value->ne[0]; + if (!is_supported_qk_head_size(qk_head_size) || !is_supported_value_head_size(value_head_size) || + !has_flash_attention_params(*node)) { + return {}; + } + if (key->ne[0] != qk_head_size || output->ne[0] != value_head_size) { + return {}; + } + if (query->ne[3] != 1 || key->ne[3] != 1 || value->ne[3] != 1 || output->ne[3] != 1 || mask->ne[2] != 1 || + mask->ne[3] != 1) { + return {}; + } + + const int64_t query_token_count = query->ne[1]; + const int64_t query_head_count = query->ne[2]; + const int64_t key_value_capacity = key->ne[1]; + const int64_t key_value_head_count = key->ne[2]; + const int64_t key_value_token_count = mask->ne[0]; + if (!is_supported_decode_query_length(query_token_count) || + !is_supported_decode_key_value_token_count(key_value_token_count) || + key_value_capacity < key_value_token_count || !is_supported_head_count(query_head_count) || + !is_supported_head_count(key_value_head_count) || query_head_count % key_value_head_count != 0) { + return {}; + } + if (value->ne[1] != key_value_capacity || value->ne[2] != key_value_head_count || + mask->ne[1] != query_token_count || output->ne[1] != query_head_count || output->ne[2] != query_token_count) { + return {}; + } + if (!has_query_layout(*query, query_head_count, qk_head_size) || + !has_key_value_layout(*key, key_value_head_count, qk_head_size) || + !has_key_value_layout(*value, key_value_head_count, value_head_size) || + !has_mask_layout(*mask, key_value_token_count) || + !has_output_layout(*output, query_head_count, value_head_size)) { + return {}; + } + + match.query = query; + match.key = key; + match.value = value; + match.mask = mask; + match.output = output; + match.output_layout = find_single_layout_alias_consumer(graph, output->id); + match.query_token_count = query_token_count; + match.key_value_token_count = key_value_token_count; + match.key_value_capacity = ceil_div(key_value_token_count, kDecodeKvTileSize) * kDecodeKvTileSize; + if (match.key_value_capacity > kDecodeSplitMaxKeyValueTokenCapacity) { + return {}; + } + + // engine#123 step 6: the decode-split pack handoff is only PROVEN with the production + // partial-transient alignment. Any other layout is a test rig whose safety is not + // established, so DECLINE (fall back to flash_attention_f32_f16_wmma) rather than run an + // unproven layout and risk a probabilistic fault. Production behaviour is unchanged. + if (decode_split_partial_alignment() != 4096u) { + return {}; + } + match.query_head_count = query_head_count; + match.key_value_head_count = key_value_head_count; + match.qk_head_size = qk_head_size; + match.value_head_size = value_head_size; + match.attention_scale = op_params_as(node->params)->scale; + return match; +} + +static std::string to_config_value(float value) { + std::ostringstream out; + out.precision(9); + out << value; + return out.str(); +} + +static void add_flash_attention_decode_compile_parameters(KernelSpecialization & kernel, + int64_t query_head_count, + int64_t key_value_head_count, + int64_t qk_head_size, + int64_t value_head_size, + float attention_scale) { + kernel.compile_parameters.emplace("ggml.flash_attention.query_head_count", to_config_value(query_head_count)); + kernel.compile_parameters.emplace("ggml.flash_attention.key_value_head_count", + to_config_value(key_value_head_count)); + kernel.compile_parameters.emplace("ggml.flash_attention.qk_head_size", to_config_value(qk_head_size)); + kernel.compile_parameters.emplace("ggml.flash_attention.value_head_size", to_config_value(value_head_size)); + kernel.compile_parameters.emplace("ggml.flash_attention.attention_scale", to_config_value(attention_scale)); +} + +static void add_flash_attention_compile_parameters(KernelSpecialization & kernel, + int64_t query_head_count, + int64_t key_value_head_count, + int64_t qk_head_size, + int64_t value_head_size, + float attention_scale, + bool apply_gate, + int64_t gate_stride_head, + int64_t gate_stride_token, + int64_t value_stride) { + kernel.compile_parameters.emplace("ggml.flash_attention.query_head_count", to_config_value(query_head_count)); + kernel.compile_parameters.emplace("ggml.flash_attention.key_value_head_count", + to_config_value(key_value_head_count)); + kernel.compile_parameters.emplace("ggml.flash_attention.qk_head_size", to_config_value(qk_head_size)); + kernel.compile_parameters.emplace("ggml.flash_attention.value_head_size", to_config_value(value_head_size)); + kernel.compile_parameters.emplace("ggml.flash_attention.value_stride", to_config_value(value_stride)); + kernel.compile_parameters.emplace("ggml.flash_attention.attention_scale", to_config_value(attention_scale)); + kernel.compile_parameters.emplace("ggml.flash_attention.apply_gate", apply_gate ? "1" : "0"); + kernel.compile_parameters.emplace("ggml.flash_attention.gate_stride_head", to_config_value(gate_stride_head)); + kernel.compile_parameters.emplace("ggml.flash_attention.gate_stride_token", to_config_value(gate_stride_token)); +} + +static const GraphNode * single_consumer_with_op(const Graph & graph, ValueId value, ggml_op op) { + const std::vector & consumers = graph.index().consumers(value); + return consumers.size() == 1 && consumers.front() != nullptr && consumers.front()->op == op ? consumers.front() : + nullptr; +} + +static DispatchBinding prepare_flash_attention_value(const DispatchMatchContext & context, + const FlashAttentionMatch & match, + DispatchMatch & dispatch_match, + KernelSpecialization & attention) { + const int64_t columns = match.key_value_head_count * match.value_head_size; + // ggml_copy_transpose_f16 reads `columns`-wide contiguous rows. MLA (deepseek2/GLM) stores V with the QK + // head-size stride (the matcher accepts that layout, #95), so its rows are wider than `columns`: transposing + // would read misaligned rows. Such a V keeps the row-major layout, whose kernel takes the real row stride. + const bool contiguous_rows = match.value->nb[1] == static_cast(columns) * sizeof(ggml_fp16_t); + if (!contiguous_rows || match.query_token_count < 256 || match.key_value_token_count < 512 || + match.key_value_token_count % 32 != 0 || match.key_value_token_count > kCopyTransposeF16MaxExtent || + columns % 32 != 0 || columns > kCopyTransposeF16MaxExtent) { + return { match.value->id, 0, match.value->byte_count }; + } + + const size_t bytes = static_cast(match.key_value_token_count * columns) * sizeof(ggml_fp16_t); + const ValueId transposed = match_value(context, dispatch_match, 0); + dispatch_match.transients.push_back({ transposed, "common.flash_attention.transposed_value", bytes, 256 }); + + Dispatch copy; + copy.kernel = make_kernel_specialization(kCopyTransposeF16Kernel); + copy.kernel.compile_parameters.emplace("ggml.copy_transpose_f16.row_count", to_config_value(match.key_value_token_count)); + copy.kernel.compile_parameters.emplace("ggml.copy_transpose_f16.column_count", to_config_value(columns)); + copy.bindings.push_back({ match.value->id, 0, bytes }); + copy.bindings.push_back({ transposed, 0, bytes }); + dispatch_match.dispatches.push_back(std::move(copy)); + attention.compile_parameters.emplace("ggml.flash_attention.value_layout", "1"); + return { transposed, 0, bytes }; +} + +static bool match_flash_attention_gate_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const FlashAttentionMatch match = match_flash_attention_f32_f16(context.graph, context.plan, context.root_node); + if (!match.matched() || context.root_node->inputs.size() != 4 || match.output_layout == nullptr || + match.output_layout->op != GGML_OP_RESHAPE) { + return false; + } + + const GraphNode * reshape = match.output_layout; + const Value * reshaped = graph_value(context.graph, reshape->output); + const GraphNode * mul = + reshaped != nullptr ? single_consumer_with_op(context.graph, reshaped->id, GGML_OP_MUL) : nullptr; + if (reshape->inputs.size() != 1 || reshaped == nullptr || mul == nullptr || mul->inputs.size() != 2 || + reshaped->type != GGML_TYPE_F32 || !reshaped->contiguous || + reshaped->ne[0] != match.value_head_size * match.query_head_count || + reshaped->ne[1] != match.query_token_count || reshaped->ne[2] != 1 || reshaped->ne[3] != 1) { + return false; + } + + ValueId gate_value_id; + if (mul->inputs[0] == reshape->output) { + gate_value_id = mul->inputs[1]; + } else if (mul->inputs[1] == reshape->output) { + gate_value_id = mul->inputs[0]; + } else { + return false; + } + + const GraphNode * sigmoid = context.graph.index().producer(gate_value_id); + const UnaryParams * sigmoid_params = sigmoid != nullptr ? op_params_as(sigmoid->params) : nullptr; + if (sigmoid == nullptr || sigmoid->op != GGML_OP_UNARY || sigmoid->inputs.size() != 1 || + sigmoid_params == nullptr || sigmoid_params->op != UnaryKind::Sigmoid || + single_consumer_with_op(context.graph, sigmoid->output, GGML_OP_MUL) != mul) { + return false; + } + + const GraphNode * cont = context.graph.index().producer(sigmoid->inputs[0]); + if (cont == nullptr || cont->op != GGML_OP_CONT || cont->inputs.size() != 1 || + single_consumer_with_op(context.graph, cont->output, GGML_OP_UNARY) != sigmoid) { + return false; + } + + const GraphNode * gate_view = context.graph.index().producer(cont->inputs[0]); + const Value * raw_gate = gate_view != nullptr ? graph_value(context.graph, gate_view->output) : nullptr; + const Value * output = graph_value(context.graph, mul->output); + if (gate_view == nullptr || gate_view->op != GGML_OP_VIEW || gate_view->inputs.size() != 1 || raw_gate == nullptr || + output == nullptr || raw_gate->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || + raw_gate->ne[0] != match.value_head_size || raw_gate->ne[1] != match.query_head_count || + raw_gate->ne[2] != match.query_token_count || raw_gate->ne[3] != 1 || raw_gate->nb[0] != sizeof(float) || + output->ne != reshaped->ne || !output->contiguous || + single_consumer_with_op(context.graph, gate_view->output, GGML_OP_CONT) != cont) { + return false; + } + if (output->storage_root == match.query->storage_root || output->storage_root == match.key->storage_root || + output->storage_root == match.value->storage_root || output->storage_root == match.mask->storage_root || + output->storage_root == raw_gate->storage_root) { + return false; + } + + dispatch_match.covered_nodes.push_back(context.root_index); + for (const GraphNode * covered : { reshape, gate_view, cont, sigmoid, mul }) { + if (!append_covered_node_index_once(context.graph, context.covered_nodes, covered, + dispatch_match.covered_nodes)) { + return false; + } + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kFlashAttentionF32F16WmmaKernel); + dispatch.kernel.integer_parameters.emplace("query_token_count", match.query_token_count); + dispatch.kernel.integer_parameters.emplace("key_value_token_count", match.key_value_token_count); + add_flash_attention_compile_parameters(dispatch.kernel, match.query_head_count, match.key_value_head_count, + match.qk_head_size, match.value_head_size, match.attention_scale, true, + static_cast(raw_gate->nb[1] / sizeof(float)), + static_cast(raw_gate->nb[2] / sizeof(float)), + static_cast(match.value->nb[1] / sizeof(ggml_fp16_t))); + dispatch.bindings.push_back({ match.query->id, 0, match.query->byte_count }); + dispatch.bindings.push_back({ match.key->id, 0, match.key->byte_count }); + dispatch.bindings.push_back(prepare_flash_attention_value(context, match, dispatch_match, dispatch.kernel)); + dispatch.bindings.push_back({ match.mask_binding_value, 0, match.mask_binding_bytes }); + dispatch.bindings.push_back({ raw_gate->id, 0, raw_gate->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_flash_attention_f32_f16_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const FlashAttentionMatch match = match_flash_attention_f32_f16(context.graph, context.plan, context.root_node); + if (!match.matched()) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kFlashAttentionF32F16WmmaKernel); + dispatch.kernel.integer_parameters.emplace("query_token_count", match.query_token_count); + dispatch.kernel.integer_parameters.emplace("key_value_token_count", match.key_value_token_count); + add_flash_attention_compile_parameters(dispatch.kernel, match.query_head_count, match.key_value_head_count, + match.qk_head_size, match.value_head_size, match.attention_scale, false, 1, + 1, static_cast(match.value->nb[1] / sizeof(ggml_fp16_t))); + dispatch.bindings.push_back({ match.query->id, 0, match.query->byte_count }); + dispatch.bindings.push_back({ match.key->id, 0, match.key->byte_count }); + dispatch.bindings.push_back(prepare_flash_attention_value(context, match, dispatch_match, dispatch.kernel)); + dispatch.bindings.push_back({ match.mask_binding_value, 0, match.mask_binding_bytes }); + dispatch.bindings.push_back({ match.query->id, 0, match.query->byte_count }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + + dispatch_match.covered_nodes.push_back(context.root_index); + if (match.output_layout != nullptr) { + if (!append_covered_node_index_once(context.graph, context.covered_nodes, match.output_layout, + dispatch_match.covered_nodes)) { + return false; + } + } + dispatch_match.dispatches.push_back(std::move(dispatch)); + if (context.root_node->inputs.size() == 5 && + !append_attention_sink_dispatch(context.graph, *context.root_node, dispatch_match)) { + return false; + } + return true; +} + +static bool match_flash_attention_decode_split_next_q8_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const DecodeSplitFlashAttentionMatch match = + match_decode_split_flash_attention_f32_f16(context.graph, context.root_node); + if (!match.matched()) { + return false; + } + + const int64_t key_value_block_count = ceil_div(match.key_value_capacity, kDecodeKvTileSize); + const size_t partial_scalar_count = static_cast(match.key_value_head_count) * + static_cast(key_value_block_count) * + static_cast(kDecodeRowCapacity); + const size_t partial_value_count = partial_scalar_count * static_cast(match.value_head_size); + const size_t partial_scalar_bytes = partial_scalar_count * sizeof(float); + const size_t partial_output_bytes = partial_value_count * sizeof(ggml_fp16_t); + const int64_t query_hidden_size = match.query_head_count * match.qk_head_size; + const int64_t output_hidden_size = match.query_head_count * match.value_head_size; + const size_t q8_row_bytes = q8_1_x4_byte_count(1, output_hidden_size); + const size_t q8_output_bytes = q8_1_x4_byte_count(match.query_token_count, output_hidden_size); + if (partial_scalar_bytes == 0 || partial_output_bytes == 0 || q8_row_bytes == 0 || q8_output_bytes == 0) { + return false; + } + + const ValueId partial_max = match_value(context, dispatch_match, 0); + const ValueId partial_sum = match_value(context, dispatch_match, 1); + const ValueId partial_output = match_value(context, dispatch_match, 2); + const ValueId completion_counter = match_value(context, dispatch_match, 3); + const ValueId q8_output = match_value(context, dispatch_match, 4); + + dispatch_match.transients.push_back( + { partial_max, "common.decode.flash_attention.partial_max", partial_scalar_bytes, decode_split_partial_alignment() }); + dispatch_match.transients.push_back( + { partial_sum, "common.decode.flash_attention.partial_sum", partial_scalar_bytes, decode_split_partial_alignment() }); + dispatch_match.transients.push_back( + { partial_output, "common.decode.flash_attention.partial_output", partial_output_bytes, decode_split_partial_alignment() }); + dispatch_match.transients.push_back( + { q8_output, "common.decode.flash_attention.next_q8_output", q8_output_bytes, 4096 }); + dispatch_match.completion_counter_requests.push_back({ + completion_counter, + "common.decode.flash_attention.completion_counter", + static_cast(match.key_value_head_count), + }); + + Status metadata_status; + if (!dispatch_match.metadata.append_alternate_value({ match.output->id, q8_output, GGML_TYPE_Q8_1, q8_output_bytes, + "common.decode.flash_attention.next_q8_output" }, + metadata_status)) { + dispatch_match.status.append(metadata_status); + return false; + } + + const size_t query_row_bytes = static_cast(query_hidden_size) * sizeof(float); + const size_t mask_row_bytes = static_cast(match.key_value_token_count) * sizeof(ggml_fp16_t); + const size_t output_row_bytes = static_cast(output_hidden_size) * sizeof(float); + for (int64_t row = 0; row < match.query_token_count; ++row) { + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kFlashAttentionDecodeSplitNextQ8Kernel); + dispatch.kernel.integer_parameters.emplace("key_value_token_count", match.key_value_token_count); + add_flash_attention_decode_compile_parameters(dispatch.kernel, match.query_head_count, + match.key_value_head_count, match.qk_head_size, + match.value_head_size, match.attention_scale); + dispatch.kernel.compile_parameters.emplace("ggml.flash_attention.decode.key_value_token_capacity", + to_config_value(match.key_value_capacity)); + dispatch.bindings.push_back( + { match.query->id, static_cast(row) * match.query->nb[1], query_row_bytes }); + dispatch.bindings.push_back({ match.key->id, 0, match.key->byte_count }); + dispatch.bindings.push_back({ match.value->id, 0, match.value->byte_count }); + dispatch.bindings.push_back({ match.mask->id, static_cast(row) * match.mask->nb[1], mask_row_bytes }); + dispatch.bindings.push_back({ partial_max, 0, partial_scalar_bytes }); + dispatch.bindings.push_back({ partial_sum, 0, partial_scalar_bytes }); + dispatch.bindings.push_back({ partial_output, 0, partial_output_bytes }); + dispatch.bindings.push_back( + { completion_counter, 0, static_cast(match.key_value_head_count) * sizeof(int32_t) }); + dispatch.bindings.push_back( + { match.output->id, static_cast(row) * match.output->nb[2], output_row_bytes }); + dispatch.bindings.push_back({ q8_output, static_cast(row) * q8_row_bytes, q8_row_bytes }); + dispatch_match.dispatches.push_back(std::move(dispatch)); + } + + dispatch_match.covered_nodes.push_back(context.root_index); + if (match.output_layout != nullptr) { + if (!append_covered_node_index_once(context.graph, context.covered_nodes, match.output_layout, + dispatch_match.covered_nodes)) { + return false; + } + } + return true; +} + +} // namespace + +void register_flash_attention_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "common.flash_attention_f32_f16_gate", + GGML_OP_FLASH_ATTN_EXT, + DispatchMatchKind::Fused, + 100, + DispatchSource::Common, + match_flash_attention_gate_dispatch, + }); + registry.add({ + "common.flash_attention_decode_split_next_q8", + GGML_OP_FLASH_ATTN_EXT, + DispatchMatchKind::SingleOp, + 75, + DispatchSource::Common, + match_flash_attention_decode_split_next_q8_dispatch, + }); + registry.add({ + "common.flash_attention_f32_f16_wmma", + GGML_OP_FLASH_ATTN_EXT, + DispatchMatchKind::SingleOp, + 50, + DispatchSource::Common, + match_flash_attention_f32_f16_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-flash-attention.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-flash-attention.h new file mode 100644 index 000000000000..125a0beeffce --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-flash-attention.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_flash_attention_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gated-mul-mat-id.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gated-mul-mat-id.cpp new file mode 100644 index 000000000000..38bbda68a0cc --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gated-mul-mat-id.cpp @@ -0,0 +1,217 @@ +#include "dispatch-gated-mul-mat-id.h" + +#include "dispatch-mul-mat-id-common.h" +#include "ggml.h" +#include "graph/graph-matcher.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kMulMatIdSwiGLUF32F32WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_id_swiglu_f32_f32_wmma"); + +struct MulMatIdSwiGLUMatch { + const Value * input = nullptr; + const Value * route_ids = nullptr; + const Value * gate_weight = nullptr; + const Value * up_weight = nullptr; + const Value * gate_output = nullptr; + const Value * up_output = nullptr; + const Value * output = nullptr; + const GraphNode * gate_node = nullptr; + const GraphNode * up_node = nullptr; + const GraphNode * glu_node = nullptr; + CommandPlanMoeRoutingBundle routing_bundle; + bool has_routing_bundle = false; + int64_t input_size = 0; + int64_t output_size = 0; + int64_t token_count = 0; + int64_t route_count = 0; + int64_t route_stride = 0; + int64_t input_route_count = 0; + int64_t expert_count = 0; + CommonMulMatWeightFormat gate_format = CommonMulMatWeightFormat::Q4K; + CommonMulMatWeightFormat up_format = CommonMulMatWeightFormat::Q4K; + + bool matched() const { + return input != nullptr && route_ids != nullptr && gate_weight != nullptr && up_weight != nullptr && + gate_output != nullptr && up_output != nullptr && output != nullptr && gate_node != nullptr && + up_node != nullptr && glu_node != nullptr; + } +}; + +static MulMatIdSwiGLUMatch match_mul_mat_id_swiglu(const DispatchMatchContext & context) { + MulMatIdSwiGLUMatch match; + const CommonMulMatIdMatch root = common_match_mul_mat_id_any_format(context.graph, context.root_node, context.plan); + if (!root.matched() || !context.graph.has_index()) { + return match; + } + + const std::vector & root_consumers = context.graph.index().consumers(context.root_node->output); + if (root_consumers.size() != 1 || root_consumers.front() == nullptr || root_consumers.front()->op != GGML_OP_GLU) { + return {}; + } + + const GraphNode * glu_node = root_consumers.front(); + if (glu_node->inputs.size() != 2 || !common_mul_mat_id_is_swiglu_params(glu_node->params)) { + return {}; + } + + size_t glu_index = 0; + if (!context.graph.index().node_index(glu_node, glu_index) || glu_index >= context.covered_nodes.size() || + context.covered_nodes[glu_index]) { + return {}; + } + + const bool root_is_gate = glu_node->inputs[0] == context.root_node->output; + const bool root_is_up = glu_node->inputs[1] == context.root_node->output; + if (!root_is_gate && !root_is_up) { + return {}; + } + + const ValueId peer_output_id = root_is_gate ? glu_node->inputs[1] : glu_node->inputs[0]; + const GraphNode * peer_node = context.graph.index().producer(peer_output_id); + if (peer_node == nullptr || peer_node == context.root_node || peer_node->op != GGML_OP_MUL_MAT_ID) { + return {}; + } + + size_t peer_index = 0; + if (!context.graph.index().node_index(peer_node, peer_index) || peer_index >= context.covered_nodes.size() || + context.covered_nodes[peer_index]) { + return {}; + } + + const std::vector & peer_consumers = context.graph.index().consumers(peer_output_id); + if (peer_consumers.size() != 1 || peer_consumers.front() != glu_node) { + return {}; + } + + const CommonMulMatIdMatch peer = common_match_mul_mat_id_any_format(context.graph, peer_node, context.plan); + if (!peer.matched() || peer.input->id != root.input->id || peer.route_ids->id != root.route_ids->id || + peer.input_size != root.input_size || peer.output_size != root.output_size || + peer.token_count != root.token_count || peer.route_count != root.route_count || + peer.input_route_count != root.input_route_count || peer.expert_count != root.expert_count || + !common_mul_mat_id_same_shape(*root.output, *peer.output)) { + return {}; + } + + const Value * output = common_mul_mat_id_graph_value(context.graph, glu_node->output); + if (output == nullptr || output->type != GGML_TYPE_F32 || !output->contiguous || + !common_mul_mat_id_same_shape(*output, *root.output)) { + return {}; + } + + match.input = root.input; + match.route_ids = root.route_ids; + match.gate_weight = root_is_gate ? root.weight : peer.weight; + match.up_weight = root_is_gate ? peer.weight : root.weight; + match.gate_output = root_is_gate ? root.output : peer.output; + match.up_output = root_is_gate ? peer.output : root.output; + match.output = output; + match.gate_node = root_is_gate ? context.root_node : peer_node; + match.up_node = root_is_gate ? peer_node : context.root_node; + match.glu_node = glu_node; + match.routing_bundle = root.routing_bundle; + match.has_routing_bundle = root.has_routing_bundle; + match.input_size = root.input_size; + match.output_size = root.output_size; + match.token_count = root.token_count; + match.route_count = root.route_count; + match.route_stride = root.route_stride; + match.input_route_count = root.input_route_count; + match.expert_count = root.expert_count; + match.gate_format = root_is_gate ? root.weight_format : peer.weight_format; + match.up_format = root_is_gate ? peer.weight_format : root.weight_format; + return match; +} + +static CommonMulMatIdMatch routing_match_for_swiglu(const MulMatIdSwiGLUMatch & match) { + CommonMulMatIdMatch routed; + routed.input = match.input; + routed.weight = match.gate_weight; + routed.output = match.gate_output; + routed.route_ids = match.route_ids; + routed.routing_bundle = match.routing_bundle; + routed.has_routing_bundle = match.has_routing_bundle; + routed.input_size = match.input_size; + routed.output_size = match.output_size; + routed.token_count = match.token_count; + routed.route_count = match.route_count; + routed.route_stride = match.route_stride; + routed.input_route_count = match.input_route_count; + routed.expert_count = match.expert_count; + routed.weight_format = match.gate_format; + return routed; +} + +static bool match_mul_mat_id_swiglu_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const MulMatIdSwiGLUMatch match = match_mul_mat_id_swiglu(context); + if (!match.matched()) { + return false; + } + + const CommonMulMatIdMatch routed = routing_match_for_swiglu(match); + CommandPlanMoeRoutingBundle routing_bundle; + if (!common_mul_mat_id_ensure_moe_routing_bundle(context, routed, dispatch_match, routing_bundle)) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kMulMatIdSwiGLUF32F32WmmaKernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.workload.token_capacity", + common_mul_mat_id_to_config_value(match.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_swiglu.input_size", + common_mul_mat_id_to_config_value(match.input_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_swiglu.output_size", + common_mul_mat_id_to_config_value(match.output_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_swiglu.expert_count", + common_mul_mat_id_to_config_value(match.expert_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_swiglu.route_count", + common_mul_mat_id_to_config_value(match.route_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_swiglu.input_route_count", + common_mul_mat_id_to_config_value(match.input_route_count)); + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat_id_swiglu.gate_weight_format", + common_mul_mat_id_to_config_value(common_mul_mat_format_config_value(match.gate_format))); + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat_id_swiglu.up_weight_format", + common_mul_mat_id_to_config_value(common_mul_mat_format_config_value(match.up_format))); + dispatch.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + dispatch.bindings.push_back({ routing_bundle.expert_table, 0, routing_bundle.expert_table_byte_count }); + dispatch.bindings.push_back({ routing_bundle.partition_table, 0, routing_bundle.partition_table_byte_count }); + dispatch.bindings.push_back({ match.gate_weight->id, 0, match.gate_weight->byte_count }); + dispatch.bindings.push_back({ match.up_weight->id, 0, match.up_weight->byte_count }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + + if (!append_covered_node_index_once(context.graph, context.covered_nodes, match.gate_node, + dispatch_match.covered_nodes) || + !append_covered_node_index_once(context.graph, context.covered_nodes, match.up_node, + dispatch_match.covered_nodes) || + !append_covered_node_index_once(context.graph, context.covered_nodes, match.glu_node, + dispatch_match.covered_nodes)) { + return false; + } + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_gated_mul_mat_id_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "common.mul_mat_id_swiglu.f32_f32_wmma", + GGML_OP_MUL_MAT_ID, + DispatchMatchKind::Fused, + 200, + DispatchSource::Common, + match_mul_mat_id_swiglu_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gated-mul-mat-id.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gated-mul-mat-id.h new file mode 100644 index 000000000000..37cfb4dc9ada --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gated-mul-mat-id.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_gated_mul_mat_id_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gated-mul-mat.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gated-mul-mat.cpp new file mode 100644 index 000000000000..779a2ff53e1b --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gated-mul-mat.cpp @@ -0,0 +1,862 @@ +#include "dispatch-gated-mul-mat.h" + +#include "dispatch-mul-mat-common.h" +#include "ggml.h" +#include "graph/graph-matcher.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kMulMatF32F32WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_f32_f32_wmma"); +static constexpr KernelCatalogRef kMulMatSwiGLUF32F32WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_swiglu_f32_f32_wmma"); +static constexpr KernelCatalogRef kMulMatSwiGLUF32F32DecodeWave64Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_swiglu_f32_f32_decode_wave64"); +static constexpr KernelCatalogRef kMulMatSwiGLUF32F32LowTokenDotKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_swiglu_f32_f32_lowtoken_dot"); +static constexpr KernelCatalogRef kMulMatSwiGLUQ4Q8LowTokenDotKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_swiglu_q4_q8_1_x4_lowtoken_dot"); +static constexpr KernelCatalogRef kMulMatSwiGLUQ4Q8OutputKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_swiglu_q4_q8_1_x4_lowtoken_dot_q8_output"); +static constexpr KernelCatalogRef kMulMatSwiGLUQ4Q8PrefillKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_swiglu_q4_k_f16_wmma_prefill_wave32"); +static constexpr KernelCatalogRef kQuantizeF32SymmetricI4K32Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_quantize_f32_symmetric_i4_k32"); +static constexpr KernelCatalogRef kMulMatSymmetricI4LowRowAdjacentDualWmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma"); +static constexpr KernelCatalogRef kMulMatSymmetricI4LowRowAdjacentDualDirectDotKernels[] = { + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c1"), + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c2"), + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c3"), + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c4"), + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c5"), +}; +static constexpr KernelCatalogRef kMulMatSwiGLUSymmetricI4WmmaQ8PlaneKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_swiglu_symmetric_i4_wmma_q8_plane"); +static constexpr KernelCatalogRef kMulMatQ5KQ8PlaneWmmaToken256Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_q5_k_q8_plane_wmmai8_token256"); + +struct MulMatSwiGLUMatch { + const Value * input = nullptr; + const Value * gate_weight = nullptr; + const Value * up_weight = nullptr; + const Value * gate_output = nullptr; + const Value * up_output = nullptr; + const Value * output = nullptr; + const GraphNode * gate_node = nullptr; + const GraphNode * up_node = nullptr; + const GraphNode * glu_node = nullptr; + CommonMulMatWeightFormat gate_format = CommonMulMatWeightFormat::Q4K; + CommonMulMatWeightFormat up_format = CommonMulMatWeightFormat::Q4K; + BinaryKind op = BinaryKind::SwiGLU; + int64_t input_size = 0; + int64_t output_size = 0; + int64_t token_count = 0; + + bool topology_matched() const { + return input != nullptr && gate_weight != nullptr && up_weight != nullptr && gate_output != nullptr && + up_output != nullptr && output != nullptr && gate_node != nullptr && up_node != nullptr && + glu_node != nullptr; + } + + bool matched() const { return topology_matched() && token_count >= 1; } + + bool decode_matched() const { return topology_matched() && token_count == 1; } +}; + +struct MulMatSwiGLUProjectionMatch { + MulMatSwiGLUMatch gate_up; + const GraphNode * projection_node = nullptr; + const Value * projection_weight = nullptr; + const Value * projection_output = nullptr; + int64_t projection_size = 0; + + bool matched() const { + return gate_up.matched() && projection_node != nullptr && projection_weight != nullptr && + projection_output != nullptr; + } +}; + +struct PackedMulMatGluMatch { + const Value * input = nullptr; + const Value * weight = nullptr; + const Value * packed_output = nullptr; + const Value * output = nullptr; + const GraphNode * matmul_node = nullptr; + const GraphNode * glu_node = nullptr; + BinaryKind op = BinaryKind::SwiGLU; + int64_t input_size = 0; + int64_t output_size = 0; + int64_t token_count = 0; + size_t gate_weight_offset = 0; + size_t up_weight_offset = 0; + size_t weight_half_bytes = 0; + CommonMulMatWeightFormat weight_format = CommonMulMatWeightFormat::Q4K; + + bool matched() const { + return input != nullptr && weight != nullptr && packed_output != nullptr && output != nullptr && + matmul_node != nullptr && glu_node != nullptr && token_count > 1; + } +}; + +static bool supported_symmetric_i4_pair(ggml_type gate_type, ggml_type up_type) { + return (gate_type == GGML_TYPE_Q4_K && up_type == GGML_TYPE_Q4_K) || + (gate_type == GGML_TYPE_Q5_K && up_type == GGML_TYPE_IQ4_XS) || + (gate_type == GGML_TYPE_IQ4_XS && up_type == GGML_TYPE_Q5_K); +} + +static bool supported_mixed_symmetric_i4_pair(ggml_type gate_type, ggml_type up_type) { + return (gate_type == GGML_TYPE_Q5_K && up_type == GGML_TYPE_IQ4_XS) || + (gate_type == GGML_TYPE_IQ4_XS && up_type == GGML_TYPE_Q5_K); +} + +static bool distinct_storage(const Graph & graph, const Value & lhs, const Value & rhs) { + return !graph.values().same_storage(lhs.id, rhs.id); +} + +static bool checked_mul_size(size_t lhs, size_t rhs, size_t & result) { + if (rhs != 0 && lhs > std::numeric_limits::max() / rhs) { + return false; + } + result = lhs * rhs; + return true; +} + +static size_t symmetric_i4_weight_byte_count(int64_t input_size, int64_t output_size) { + return static_cast(output_size) * static_cast(input_size / 256) * size_t{ 144 }; +} + +static DispatchBinding symmetric_i4_weight_binding(const Value & weight, int64_t input_size, int64_t output_size) { + DispatchBinding binding; + binding.value = weight.id; + binding.length = symmetric_i4_weight_byte_count(input_size, output_size); + binding.layout = kSymmetricI4K32Row64Layout; + binding.source_type = weight.type; + binding.input_size = input_size; + binding.output_size = output_size; + binding.source_length = weight.byte_count; + return binding; +} + +static DispatchBinding symmetric_i5_weight_binding(const Value & weight, int64_t input_size, int64_t output_size) { + DispatchBinding binding; + binding.value = weight.id; + binding.length = weight.byte_count; + binding.layout = kQ5KSymmetricI5K32Layout; + binding.source_type = weight.type; + binding.input_size = input_size; + binding.output_size = output_size; + binding.source_length = weight.byte_count; + return binding; +} + +static MulMatSwiGLUMatch match_mul_mat_swiglu(const DispatchMatchContext & context) { + MulMatSwiGLUMatch match; + CommonMulMatMatch root = + common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatF32F32WmmaKernel, false); + if (!root.matched()) { + root = common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatF32F32WmmaKernel, true); + } + if (!root.matched() || !context.graph.has_index()) { + return match; + } + + const std::vector & root_consumers = context.graph.index().consumers(context.root_node->output); + if (root_consumers.size() != 1 || root_consumers.front() == nullptr) { + return {}; + } + + const GraphNode * glu_node = root_consumers.front(); + BinaryKind binary_op; + if (glu_node->inputs.size() != 2 || !common_fused_binary_kind_from_params(glu_node->params, binary_op)) { + return {}; + } + + size_t glu_index = 0; + if (!context.graph.index().node_index(glu_node, glu_index) || glu_index >= context.covered_nodes.size() || + context.covered_nodes[glu_index]) { + return {}; + } + + const bool root_is_gate = glu_node->inputs[0] == context.root_node->output; + const bool root_is_up = glu_node->inputs[1] == context.root_node->output; + if (!root_is_gate && !root_is_up) { + return {}; + } + + const ValueId peer_output_id = root_is_gate ? glu_node->inputs[1] : glu_node->inputs[0]; + const GraphNode * peer_node = context.graph.index().producer(peer_output_id); + if (peer_node == nullptr || peer_node == context.root_node || peer_node->op != GGML_OP_MUL_MAT) { + return {}; + } + + size_t peer_index = 0; + if (!context.graph.index().node_index(peer_node, peer_index) || peer_index >= context.covered_nodes.size() || + context.covered_nodes[peer_index]) { + return {}; + } + + const std::vector & peer_consumers = context.graph.index().consumers(peer_output_id); + if (peer_consumers.size() != 1 || peer_consumers.front() != glu_node) { + return {}; + } + + CommonMulMatMatch peer = common_match_mul_mat_any_format(context.graph, peer_node, kMulMatF32F32WmmaKernel, false); + if (!peer.matched()) { + peer = common_match_mul_mat_any_format(context.graph, peer_node, kMulMatF32F32WmmaKernel, true); + } + if (!peer.matched() || peer.input->id != root.input->id || peer.input_size != root.input_size || + peer.output_size != root.output_size || peer.token_count != root.token_count || + !common_same_shape(*root.output, *peer.output)) { + return {}; + } + + const Value * output = common_graph_value(context.graph, glu_node->output); + if (output == nullptr || output->type != GGML_TYPE_F32 || !output->contiguous || + !common_same_shape(*output, *root.output)) { + return {}; + } + + match.input = root.input; + match.gate_weight = root_is_gate ? root.weight : peer.weight; + match.up_weight = root_is_gate ? peer.weight : root.weight; + match.gate_output = root_is_gate ? root.output : peer.output; + match.up_output = root_is_gate ? peer.output : root.output; + match.output = output; + match.gate_node = root_is_gate ? context.root_node : peer_node; + match.up_node = root_is_gate ? peer_node : context.root_node; + match.glu_node = glu_node; + match.gate_format = root_is_gate ? root.weight_format : peer.weight_format; + match.up_format = root_is_gate ? peer.weight_format : root.weight_format; + match.op = binary_op; + match.input_size = root.input_size; + match.output_size = root.output_size; + match.token_count = root.token_count; + return match; +} + +// Fuses a packed projection followed by one-input GLU. The packed weight is split into +// gate/up halves and lowered through the split-weight gated matmul kernel. +static PackedMulMatGluMatch match_packed_mul_mat_glu(const DispatchMatchContext & context) { + PackedMulMatGluMatch match; + const CommonMulMatMatch root = + common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatSwiGLUF32F32WmmaKernel, false); + if (!root.matched() || !context.graph.has_index() || root.weight->alias_source.value >= 0) { + return {}; + } + + const std::vector & root_consumers = context.graph.index().consumers(context.root_node->output); + if (root_consumers.size() != 1 || root_consumers.front() == nullptr || root_consumers.front()->op != GGML_OP_GLU) { + return {}; + } + + const GraphNode * glu_node = root_consumers.front(); + BinaryKind op; + if (glu_node->inputs.size() != 1 || glu_node->inputs[0] != context.root_node->output || + !common_fused_binary_kind_from_params(glu_node->params, op)) { + return {}; + } + + const GluParams * glu_params = op_params_as(glu_node->params); + if (glu_params == nullptr) { + return {}; + } + + size_t glu_index = 0; + if (!context.graph.index().node_index(glu_node, glu_index) || glu_index >= context.covered_nodes.size() || + context.covered_nodes[glu_index]) { + return {}; + } + + const Value * output = common_graph_value(context.graph, glu_node->output); + if (output == nullptr || output->type != GGML_TYPE_F32 || !output->contiguous || root.output_size % 2 != 0 || + output->ne[0] != root.output_size / 2 || output->ne[1] != root.token_count || output->ne[2] != 1 || + output->ne[3] != 1 || output->alias_source.value >= 0) { + return {}; + } + + const size_t row_bytes = ggml_row_size(root.weight->type, root.input_size); + size_t half_bytes = 0; + if (row_bytes == 0 || !checked_mul_size(row_bytes, static_cast(output->ne[0]), half_bytes) || + half_bytes > root.weight->byte_count || half_bytes > std::numeric_limits::max() - half_bytes || + 2 * half_bytes > root.weight->byte_count) { + return {}; + } + + match.input = root.input; + match.weight = root.weight; + match.packed_output = root.output; + match.output = output; + match.matmul_node = context.root_node; + match.glu_node = glu_node; + match.op = op; + match.input_size = root.input_size; + match.output_size = output->ne[0]; + match.token_count = root.token_count; + match.gate_weight_offset = glu_params->swapped ? half_bytes : 0; + match.up_weight_offset = glu_params->swapped ? 0 : half_bytes; + match.weight_half_bytes = half_bytes; + match.weight_format = root.weight_format; + return match; +} + +static bool match_packed_mul_mat_glu_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const PackedMulMatGluMatch match = match_packed_mul_mat_glu(context); + if (!match.matched()) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kMulMatSwiGLUF32F32WmmaKernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.workload.token_capacity", + common_to_config_value(match.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_swiglu.input_size", + common_to_config_value(match.input_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_swiglu.output_size", + common_to_config_value(match.output_size)); + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat_swiglu.gate_weight_format", + common_to_config_value(common_mul_mat_format_config_value(match.weight_format))); + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat_swiglu.up_weight_format", + common_to_config_value(common_mul_mat_format_config_value(match.weight_format))); + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat_swiglu.op", common_to_config_value(static_cast(binary_kind_config_value(match.op)))); + dispatch.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + dispatch.bindings.push_back({ match.weight->id, match.gate_weight_offset, match.weight_half_bytes }); + dispatch.bindings.push_back({ match.weight->id, match.up_weight_offset, match.weight_half_bytes }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + + if (!append_covered_node_index_once(context.graph, context.covered_nodes, match.matmul_node, + dispatch_match.covered_nodes) || + !append_covered_node_index_once(context.graph, context.covered_nodes, match.glu_node, + dispatch_match.covered_nodes)) { + return false; + } + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_mul_mat_swiglu_symmetric_i4_lowrow_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const MulMatSwiGLUMatch match = match_mul_mat_swiglu(context); + if (!match.topology_matched() || + !supported_mixed_symmetric_i4_pair(match.gate_weight->type, match.up_weight->type) || + match.gate_weight->alias_source.value >= 0 || match.up_weight->alias_source.value >= 0 || + match.token_count < 1 || match.token_count > 16 || match.input_size % 64 != 0 || match.output_size % 64 != 0 || + !distinct_storage(context.graph, *match.gate_weight, *match.up_weight) || + !distinct_storage(context.graph, *match.gate_output, *match.up_output)) { + return false; + } + + const CommonSymmetricI4ActivationLayout activation_layout = + common_symmetric_i4_activation_layout(match.input_size, match.token_count); + const CommandPlanAlternateValue * alternate = find_alternate_value(context.graph, context.plan, match.input->id, + GGML_TYPE_COUNT, activation_layout.total_bytes); + if (alternate != nullptr && alternate->name != kCommonSymmetricI4K32ActivationAlternateName) { + alternate = nullptr; + } + + const ValueId activation = alternate != nullptr ? alternate->alternate_value : context.next_plan_value; + if (alternate == nullptr) { + dispatch_match.transients.push_back( + { activation, kCommonSymmetricI4K32ActivationAlternateName, activation_layout.total_bytes, 256 }); + + Dispatch quantize; + quantize.kernel = make_kernel_specialization(kQuantizeF32SymmetricI4K32Kernel); + quantize.kernel.compile_parameters.emplace("ggml.quantize_symmetric_i4_k32.input_size", + common_to_config_value(match.input_size)); + quantize.kernel.compile_parameters.emplace("ggml.quantize_symmetric_i4_k32.token_count", + common_to_config_value(match.token_count)); + quantize.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + quantize.bindings.push_back({ activation, 0, activation_layout.payload_bytes }); + quantize.bindings.push_back({ activation, activation_layout.scales_offset, activation_layout.metadata_bytes }); + quantize.bindings.push_back({ activation, activation_layout.sums_offset, activation_layout.metadata_bytes }); + dispatch_match.dispatches.push_back(std::move(quantize)); + + Status metadata_status; + if (!dispatch_match.metadata.append_alternate_value( + { match.input->id, activation, GGML_TYPE_COUNT, activation_layout.total_bytes, + kCommonSymmetricI4K32ActivationAlternateName }, + metadata_status)) { + dispatch_match.status.append(metadata_status); + return false; + } + } + + const bool use_direct_dot = + match.token_count <= 5 && match.input_size % 256 == 0 && match.output_size >= 2 * match.input_size; + + Dispatch gate_up; + gate_up.kernel = make_kernel_specialization( + use_direct_dot ? kMulMatSymmetricI4LowRowAdjacentDualDirectDotKernels[match.token_count - 1] : + kMulMatSymmetricI4LowRowAdjacentDualWmmaKernel); + gate_up.kernel.compile_parameters.emplace("ggml.mul_mat.symmetric_i4.lowrow.input_size", + common_to_config_value(match.input_size)); + gate_up.kernel.compile_parameters.emplace("ggml.mul_mat.symmetric_i4.lowrow.output_size", + common_to_config_value(match.output_size)); + gate_up.kernel.compile_parameters.emplace("ggml.mul_mat.symmetric_i4.lowrow.token_count", + common_to_config_value(match.token_count)); + gate_up.kernel.compile_parameters.emplace( + "ggml.mul_mat.symmetric_i4.lowrow.row_group_size", + common_to_config_value( + static_cast(common_symmetric_shared4_row_group_size(match.input_size, match.output_size, 4)))); + gate_up.bindings.push_back( + match.gate_weight->type == GGML_TYPE_Q5_K && match.token_count == 1 ? + common_symmetric_i4_shared4_multistart_weight_binding(*match.gate_weight, match.input_size, + match.output_size) : + common_symmetric_i4_shared4_weight_binding(*match.gate_weight, match.input_size, match.output_size)); + gate_up.bindings.push_back( + match.up_weight->type == GGML_TYPE_Q5_K && match.token_count == 1 ? + common_symmetric_i4_shared4_multistart_weight_binding(*match.up_weight, match.input_size, + match.output_size) : + common_symmetric_i4_shared4_weight_binding(*match.up_weight, match.input_size, match.output_size)); + gate_up.bindings.push_back({ match.gate_output->id, 0, match.gate_output->byte_count }); + gate_up.bindings.push_back({ match.up_output->id, 0, match.up_output->byte_count }); + gate_up.bindings.push_back({ activation, 0, activation_layout.payload_bytes }); + gate_up.bindings.push_back({ activation, activation_layout.scales_offset, activation_layout.metadata_bytes }); + gate_up.bindings.push_back({ activation, activation_layout.sums_offset, activation_layout.metadata_bytes }); + + if (!append_covered_node_index_once(context.graph, context.covered_nodes, match.gate_node, + dispatch_match.covered_nodes) || + !append_covered_node_index_once(context.graph, context.covered_nodes, match.up_node, + dispatch_match.covered_nodes)) { + return false; + } + dispatch_match.dispatches.push_back(std::move(gate_up)); + return dispatch_match.status.success(); +} + +static MulMatSwiGLUProjectionMatch match_mul_mat_swiglu_q5_projection(const DispatchMatchContext & context) { + MulMatSwiGLUProjectionMatch match; + match.gate_up = match_mul_mat_swiglu(context); + if (!match.gate_up.matched() || match.gate_up.op != BinaryKind::SwiGLU || + !supported_symmetric_i4_pair(match.gate_up.gate_weight->type, match.gate_up.up_weight->type) || + match.gate_up.gate_weight->alias_source.value >= 0 || match.gate_up.up_weight->alias_source.value >= 0 || + match.gate_up.token_count < 256 || match.gate_up.token_count % 256 != 0 || + match.gate_up.output_size % 256 != 0 || + !distinct_storage(context.graph, *match.gate_up.gate_weight, *match.gate_up.up_weight)) { + return {}; + } + + const GraphNode * projection_node = + common_find_only_consumer_with_op(context.graph, match.gate_up.output->id, GGML_OP_MUL_MAT); + if (projection_node == nullptr || projection_node->inputs.size() != 2 || + projection_node->inputs[1] != match.gate_up.output->id) { + return {}; + } + + size_t projection_index = 0; + if (!context.graph.index().node_index(projection_node, projection_index) || + projection_index >= context.covered_nodes.size() || context.covered_nodes[projection_index]) { + return {}; + } + + const CommonMulMatMatch projection = + common_match_mul_mat_any_format(context.graph, projection_node, kMulMatQ5KQ8PlaneWmmaToken256Kernel, false); + if (!projection.matched() || projection.weight->type != GGML_TYPE_Q5_K || + projection.weight->alias_source.value >= 0 || projection.input->id != match.gate_up.output->id || + projection.input_size != match.gate_up.output_size || projection.token_count != match.gate_up.token_count || + projection.output_size % 64 != 0 || !distinct_storage(context.graph, *projection.weight, *projection.output)) { + return {}; + } + + match.projection_node = projection_node; + match.projection_weight = projection.weight; + match.projection_output = projection.output; + match.projection_size = projection.output_size; + return match; +} + +static bool match_mul_mat_swiglu_q5_projection_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const MulMatSwiGLUProjectionMatch match = match_mul_mat_swiglu_q5_projection(context); + if (!match.matched()) { + return false; + } + + const size_t input_elements = + static_cast(match.gate_up.token_count) * static_cast(match.gate_up.input_size); + const size_t i4_payload_bytes = input_elements / 2; + const size_t i4_metadata_bytes = input_elements / 8; + const size_t q8_output_bytes = + static_cast(match.gate_up.token_count) * ggml_row_size(GGML_TYPE_Q8_1, match.gate_up.output_size); + const ValueId i4_payload = context.next_plan_value; + const ValueId i4_scales(context.next_plan_value.value + 1); + const ValueId i4_sums(context.next_plan_value.value + 2); + const ValueId q8_output(context.next_plan_value.value + 3); + + dispatch_match.transients.push_back( + { i4_payload, "common.mul_mat_swiglu.symmetric_i4.payload", i4_payload_bytes, 256 }); + dispatch_match.transients.push_back( + { i4_scales, "common.mul_mat_swiglu.symmetric_i4.scales", i4_metadata_bytes, 256 }); + dispatch_match.transients.push_back({ i4_sums, "common.mul_mat_swiglu.symmetric_i4.sums", i4_metadata_bytes, 256 }); + dispatch_match.transients.push_back({ q8_output, "common.mul_mat_swiglu.q8_plane", q8_output_bytes, 256 }); + + Dispatch quantize; + quantize.kernel = make_kernel_specialization(kQuantizeF32SymmetricI4K32Kernel); + quantize.kernel.compile_parameters.emplace("ggml.quantize_symmetric_i4_k32.input_size", + common_to_config_value(match.gate_up.input_size)); + quantize.kernel.compile_parameters.emplace("ggml.quantize_symmetric_i4_k32.token_count", + common_to_config_value(match.gate_up.token_count)); + quantize.bindings.push_back({ match.gate_up.input->id, 0, match.gate_up.input->byte_count }); + quantize.bindings.push_back({ i4_payload, 0, i4_payload_bytes }); + quantize.bindings.push_back({ i4_scales, 0, i4_metadata_bytes }); + quantize.bindings.push_back({ i4_sums, 0, i4_metadata_bytes }); + + Dispatch gate_up; + gate_up.kernel = make_kernel_specialization(kMulMatSwiGLUSymmetricI4WmmaQ8PlaneKernel); + gate_up.kernel.compile_parameters.emplace("ggml.mul_mat_swiglu.symmetric_i4.input_size", + common_to_config_value(match.gate_up.input_size)); + gate_up.kernel.compile_parameters.emplace("ggml.mul_mat_swiglu.symmetric_i4.output_size", + common_to_config_value(match.gate_up.output_size)); + gate_up.kernel.compile_parameters.emplace("ggml.mul_mat_swiglu.symmetric_i4.token_count", + common_to_config_value(match.gate_up.token_count)); + gate_up.bindings.push_back( + symmetric_i4_weight_binding(*match.gate_up.gate_weight, match.gate_up.input_size, match.gate_up.output_size)); + gate_up.bindings.push_back( + symmetric_i4_weight_binding(*match.gate_up.up_weight, match.gate_up.input_size, match.gate_up.output_size)); + gate_up.bindings.push_back({ match.gate_up.input->id, 0, match.gate_up.input->byte_count }); + // The q8-plane variant publishes only q8_output (publish_f32 is false), so its output binding is a + // placeholder: bind the already-written input, not the SwiGLU value this path never materializes + // (the executor rejects reading a transient before any write). + gate_up.bindings.push_back({ match.gate_up.input->id, 0, match.gate_up.input->byte_count }); + gate_up.bindings.push_back({ i4_payload, 0, i4_payload_bytes }); + gate_up.bindings.push_back({ i4_scales, 0, i4_metadata_bytes }); + gate_up.bindings.push_back({ i4_sums, 0, i4_metadata_bytes }); + gate_up.bindings.push_back({ q8_output, 0, q8_output_bytes }); + + Dispatch projection; + projection.kernel = make_kernel_specialization(kMulMatQ5KQ8PlaneWmmaToken256Kernel); + projection.kernel.integer_parameters.emplace("token_count", match.gate_up.token_count); + projection.kernel.compile_parameters.emplace("ggml.mul_mat_q5_k_q8_plane.input_size", + common_to_config_value(match.gate_up.output_size)); + projection.kernel.compile_parameters.emplace("ggml.mul_mat_q5_k_q8_plane.output_size", + common_to_config_value(match.projection_size)); + projection.kernel.compile_parameters.emplace("ggml.mul_mat_q5_k_q8_plane.token_capacity", + common_to_config_value(match.gate_up.token_count)); + projection.bindings.push_back({ q8_output, 0, q8_output_bytes }); + projection.bindings.push_back( + symmetric_i5_weight_binding(*match.projection_weight, match.gate_up.output_size, match.projection_size)); + projection.bindings.push_back({ match.projection_output->id, 0, match.projection_output->byte_count }); + + if (!append_covered_node_index_once(context.graph, context.covered_nodes, match.gate_up.gate_node, + dispatch_match.covered_nodes) || + !append_covered_node_index_once(context.graph, context.covered_nodes, match.gate_up.up_node, + dispatch_match.covered_nodes) || + !append_covered_node_index_once(context.graph, context.covered_nodes, match.gate_up.glu_node, + dispatch_match.covered_nodes) || + !append_covered_node_index_once(context.graph, context.covered_nodes, match.projection_node, + dispatch_match.covered_nodes)) { + return false; + } + + dispatch_match.dispatches.push_back(std::move(quantize)); + dispatch_match.dispatches.push_back(std::move(gate_up)); + dispatch_match.dispatches.push_back(std::move(projection)); + return true; +} + +static bool match_mul_mat_swiglu_q4_q8_prefill_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const MulMatSwiGLUMatch match = match_mul_mat_swiglu(context); + if (!match.matched() || match.gate_format != CommonMulMatWeightFormat::Q4K || + match.up_format != CommonMulMatWeightFormat::Q4K || match.token_count < 256 || + match.token_count > 2048 || match.token_count % 256 != 0 || match.input_size % 256 != 0 || + match.output_size % 64 != 0 || match.gate_weight->alias_source.value >= 0 || + match.up_weight->alias_source.value >= 0 || + !distinct_storage(context.graph, *match.gate_weight, *match.up_weight) || + !distinct_storage(context.graph, *match.gate_output, *match.up_output)) { + return false; + } + + // Keep in sync with ggml_swiglu_use_packed_f16 in the shared operation. + const bool packed_input = match.token_count % 512 == 0 && + (match.token_count / 512) * (match.output_size / 64) >= 64; + DispatchBinding activation; + const bool prepared = packed_input ? + common_prepare_k16_major_f16_input(context, *match.input, match.input_size, match.token_count, + dispatch_match, activation) : + common_prepare_f16_input(context, *match.input, match.input_size, match.token_count, + dispatch_match, activation); + if (!prepared) { + return false; + } + + const GraphNode * consumer = common_find_only_consumer_with_op(context.graph, match.output->id, GGML_OP_MUL_MAT); + const CommonMulMatMatch projection = + common_match_mul_mat_any_format(context.graph, consumer, kMulMatSwiGLUQ4Q8PrefillKernel, false); + const bool packed_output = projection.matched() && projection.input->id == match.output->id && + projection.weight->alias_source.value < 0 && + common_mul_mat_uses_k16_major_f16(projection.weight_format, projection.input_size, + projection.output_size, projection.token_count, true); + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kMulMatSwiGLUQ4Q8PrefillKernel); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_swiglu.f16_output_layout", packed_output ? "1" : "0"); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_swiglu.op", + common_to_config_value(static_cast(binary_kind_config_value(match.op)))); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_swiglu.input_size", + common_to_config_value(match.input_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_swiglu.output_size", + common_to_config_value(match.output_size)); + dispatch.kernel.compile_parameters.emplace("ggml.workload.token_capacity", + common_to_config_value(match.token_count)); + dispatch.bindings.push_back(activation); + for (const Value * weight : { match.gate_weight, match.up_weight }) { + dispatch.bindings.push_back({ weight->id, 0, weight->byte_count, kQ4KPackedK256Row64Layout, weight->type, + match.input_size, match.output_size, weight->byte_count }); + } + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + const size_t f16_bytes = match.output->byte_count / 2; + const ValueId f16_output(context.next_plan_value.value + static_cast(dispatch_match.transients.size())); + const char * f16_name = packed_output ? "common.mul_mat_swiglu.k16_major_f16" : "common.mul_mat_swiglu.f16"; + dispatch_match.transients.push_back({ f16_output, f16_name, f16_bytes, 256 }); + dispatch.bindings.push_back({ f16_output, 0, f16_bytes }); + Status status; + const bool recorded = packed_output ? + dispatch_match.metadata.append_generated_resource( + { match.output->id, GeneratedResourceRole::F16K16Major, f16_output, f16_bytes, {} }, status) : + dispatch_match.metadata.append_alternate_value( + { match.output->id, f16_output, GGML_TYPE_F16, f16_bytes, f16_name }, status); + if (!recorded) { + dispatch_match.status.append(status); + return false; + } + for (const GraphNode * node : { match.gate_node, match.up_node, match.glu_node }) { + if (!append_covered_node_index_once(context.graph, context.covered_nodes, node, dispatch_match.covered_nodes)) { + return false; + } + } + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool supports_swiglu_q8_output_shape(const MulMatSwiGLUMatch & match) { + return match.op == BinaryKind::SwiGLU && match.token_count >= 1 && match.token_count <= 5 && + common_is_supported_dense_input_size(CommonMulMatWeightFormat::Q4K, match.input_size) && + common_is_supported_dense_output_size(match.output_size) && match.output_size % 128 == 0; +} + +static bool has_qualified_swiglu_q8_consumer(const Graph & graph, const MulMatSwiGLUMatch & match) { + if (!supports_swiglu_q8_output_shape(match) || + !distinct_storage(graph, *match.input, *match.gate_weight) || + !distinct_storage(graph, *match.input, *match.up_weight) || + !distinct_storage(graph, *match.input, *match.output) || + !distinct_storage(graph, *match.gate_weight, *match.up_weight) || + !distinct_storage(graph, *match.gate_weight, *match.output) || + !distinct_storage(graph, *match.up_weight, *match.output)) { + return false; + } + for (const GraphNode * consumer : graph.index().consumers(match.output->id)) { + CommonMulMatMatch projection = + common_match_mul_mat_any_format(graph, consumer, kMulMatSwiGLUQ4Q8OutputKernel, true); + if (!projection.matched()) { + projection = common_match_mul_mat_any_format(graph, consumer, kMulMatSwiGLUQ4Q8OutputKernel, false); + } + if (projection.matched() && projection.input->id == match.output->id && + projection.weight->alias_source.value < 0 && projection.output_size >= 64 && + projection.output_size % 64 == 0 && + (projection.weight_format == CommonMulMatWeightFormat::Q4K || + projection.weight_format == CommonMulMatWeightFormat::Q6K)) { + return true; + } + } + return false; +} + +static bool match_mul_mat_swiglu_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const MulMatSwiGLUMatch match = match_mul_mat_swiglu(context); + if (!match.matched()) { + return false; + } + + // Prompt chunks whose gate and up weights the q8_1 x4 prefill kernel takes are left to it: + // two q8_1 x4 matmuls plus a separate GLU beat this f32 WMMA body (Qwen3.8-27B UD-Q4_K_XL + // pp512 143 -> 363 tok/s, KLD vs BF16 0.00720 -> 0.00711). + const auto q8_x4_format = [](CommonMulMatWeightFormat f) { + return f == CommonMulMatWeightFormat::Q4K || f == CommonMulMatWeightFormat::Q5K || + f == CommonMulMatWeightFormat::IQ4_XS; + }; + // A chunk with a remainder is split by the q8_1 x4 matcher only for Q5_K/IQ4_XS weights. + const int64_t tail = match.token_count % 256; + const bool tail_ok = tail == 0 || (tail >= 2 && match.gate_format != CommonMulMatWeightFormat::Q4K && + match.up_format != CommonMulMatWeightFormat::Q4K); + if (common_q8_prefill_relaxed() && q8_x4_format(match.gate_format) && q8_x4_format(match.up_format) && + match.token_count >= 256 && match.token_count <= 2048 && tail_ok && + match.input_size % 256 == 0 && match.output_size % 64 == 0) { + return false; + } + + const bool pack_q4 = match.token_count <= 5 && match.output_size % 64 == 0 && + match.gate_weight->alias_source.value < 0 && match.up_weight->alias_source.value < 0 && + match.gate_format == CommonMulMatWeightFormat::Q4K && + match.up_format == CommonMulMatWeightFormat::Q4K; + const bool use_q8 = pack_q4; + const bool publish_q8 = use_q8 && has_qualified_swiglu_q8_consumer(context.graph, match); + DispatchBinding activation = { match.input->id, 0, match.input->byte_count }; + if (use_q8 && !common_prepare_q8_1_x4_input(context, *match.input, match.input_size, match.token_count, + dispatch_match, activation, + CommonQ8ActivationPolicy::AllowStandaloneQuantize)) { + return false; + } + + Dispatch dispatch; + const bool use_direct_dot = match.token_count <= 5 && match.input_size % 256 == 0; + const bool use_wmma = !use_q8 && !use_direct_dot; + if (use_wmma && match.token_count < 2) { + return false; + } + dispatch.kernel = make_kernel_specialization(publish_q8 ? kMulMatSwiGLUQ4Q8OutputKernel : + use_q8 ? kMulMatSwiGLUQ4Q8LowTokenDotKernel : + (use_direct_dot ? kMulMatSwiGLUF32F32LowTokenDotKernel : kMulMatSwiGLUF32F32WmmaKernel)); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.workload.token_capacity", + common_to_config_value(match.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_swiglu.input_size", + common_to_config_value(match.input_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_swiglu.output_size", + common_to_config_value(match.output_size)); + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat_swiglu.gate_weight_format", + common_to_config_value(common_mul_mat_format_config_value(pack_q4 ? CommonMulMatWeightFormat::Q4KRow64 : + match.gate_format))); + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat_swiglu.up_weight_format", + common_to_config_value(common_mul_mat_format_config_value(pack_q4 ? CommonMulMatWeightFormat::Q4KRow64 : + match.up_format))); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_swiglu.op", + common_to_config_value(static_cast(binary_kind_config_value(match.op)))); + dispatch.bindings.push_back(activation); + if (pack_q4) { + for (const Value * weight : { match.gate_weight, match.up_weight }) { + dispatch.bindings.push_back({ weight->id, 0, weight->byte_count, kQ4KPackedK256Row64Layout, weight->type, + match.input_size, match.output_size, weight->byte_count }); + } + } else { + dispatch.bindings.push_back({ match.gate_weight->id, 0, match.gate_weight->byte_count }); + dispatch.bindings.push_back({ match.up_weight->id, 0, match.up_weight->byte_count }); + } + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + + if (publish_q8) { + const ValueId q8_output(context.next_plan_value.value + static_cast(dispatch_match.transients.size())); + const size_t bytes = static_cast(match.token_count) * ggml_row_size(GGML_TYPE_Q8_1, match.output_size); + constexpr const char * name = "common.mul_mat_swiglu.q8_1_x4"; + Status status; + if (!dispatch_match.metadata.append_alternate_value( + { match.output->id, q8_output, GGML_TYPE_Q8_1, bytes, name }, status)) { + dispatch_match.status.append(status); + return false; + } + dispatch.bindings.push_back({ q8_output, 0, bytes }); + dispatch_match.transients.push_back({ q8_output, name, bytes, 256 }); + } + + if (!append_covered_node_index_once(context.graph, context.covered_nodes, match.gate_node, + dispatch_match.covered_nodes) || + !append_covered_node_index_once(context.graph, context.covered_nodes, match.up_node, + dispatch_match.covered_nodes) || + !append_covered_node_index_once(context.graph, context.covered_nodes, match.glu_node, + dispatch_match.covered_nodes)) { + return false; + } + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_decode_mul_mat_swiglu_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const MulMatSwiGLUMatch match = match_mul_mat_swiglu(context); + if (!match.decode_matched()) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kMulMatSwiGLUF32F32DecodeWave64Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.workload.token_capacity", + common_to_config_value(match.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_swiglu.input_size", + common_to_config_value(match.input_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_swiglu.output_size", + common_to_config_value(match.output_size)); + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat_swiglu.gate_weight_format", + common_to_config_value(common_mul_mat_format_config_value(match.gate_format))); + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat_swiglu.up_weight_format", + common_to_config_value(common_mul_mat_format_config_value(match.up_format))); + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat_swiglu.op", common_to_config_value(static_cast(binary_kind_config_value(match.op)))); + dispatch.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + dispatch.bindings.push_back({ match.gate_weight->id, 0, match.gate_weight->byte_count }); + dispatch.bindings.push_back({ match.up_weight->id, 0, match.up_weight->byte_count }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + + if (!append_covered_node_index_once(context.graph, context.covered_nodes, match.gate_node, + dispatch_match.covered_nodes) || + !append_covered_node_index_once(context.graph, context.covered_nodes, match.up_node, + dispatch_match.covered_nodes) || + !append_covered_node_index_once(context.graph, context.covered_nodes, match.glu_node, + dispatch_match.covered_nodes)) { + return false; + } + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_gated_mul_mat_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "common.mul_mat_swiglu.q4_k_q8_1_x4_prefill", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 310, + DispatchSource::Common, + match_mul_mat_swiglu_q4_q8_prefill_dispatch, + }); + registry.add({ + "common.mul_mat_swiglu.symmetric_i4_lowrow_adjacent_dual", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 310, + DispatchSource::Common, + match_mul_mat_swiglu_symmetric_i4_lowrow_dispatch, + }); + registry.add({ + "common.mul_mat_swiglu_q5_projection.symmetric_i4_q8_plane", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 300, + DispatchSource::Common, + match_mul_mat_swiglu_q5_projection_dispatch, + }); + registry.add({ + "common.mul_mat_swiglu.f32_f32_wmma", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 290, + DispatchSource::Common, + match_mul_mat_swiglu_dispatch, + }); + registry.add({ + "common.packed_mul_mat_glu.f32_f32_wmma", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 285, + DispatchSource::Common, + match_packed_mul_mat_glu_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gated-mul-mat.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gated-mul-mat.h new file mode 100644 index 000000000000..07e1477be1f8 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gated-mul-mat.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_gated_mul_mat_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gather-add.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gather-add.cpp new file mode 100644 index 000000000000..a1992ba6eb29 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gather-add.cpp @@ -0,0 +1,214 @@ +#include "dispatch-gather-add.h" + +#include "ggml.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kGatherAddF32Kernel = GGML_HRX_KERNEL_REF("hrx", "ggml_gather_add_f32"); + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static bool same_shape(const Value & lhs, const Value & rhs) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (lhs.ne[i] != rhs.ne[i]) { + return false; + } + } + return true; +} + +static bool is_2d_f32(const Value & value) { + return value.type == GGML_TYPE_F32 && value.contiguous && value.ne[0] > 0 && value.ne[1] > 0 && value.ne[2] == 1 && + value.ne[3] == 1; +} + +static bool is_row_id_tensor(const Value & value, int64_t output_token_count) { + return value.type == GGML_TYPE_I32 && value.contiguous && value.element_count == output_token_count && + output_token_count > 0; +} + +static bool is_supported_hidden_size(int64_t hidden_size) { + return hidden_size >= 128 && hidden_size <= 32768 && hidden_size % 128 == 0; +} + +static bool is_supported_token_count(int64_t token_count) { + return token_count >= 1 && token_count <= 2048; +} + +static bool value_is_available(const Graph & graph, ValueId value, const std::vector & covered_nodes) { + const GraphNode * producer = graph.index().producer(value); + if (producer == nullptr) { + return true; + } + size_t producer_index = 0; + return graph.index().node_index(producer, producer_index) && producer_index < covered_nodes.size() && + covered_nodes[producer_index]; +} + +struct GatherAddMatch { + const GraphNode * first_get_rows = nullptr; + const GraphNode * second_get_rows = nullptr; + const GraphNode * add_node = nullptr; + const Value * first_source = nullptr; + const Value * second_source = nullptr; + const Value * row_ids = nullptr; + const Value * output = nullptr; + size_t first_get_rows_index = 0; + size_t second_get_rows_index = 0; + size_t add_node_index = 0; + int64_t source_token_count = 0; + int64_t output_token_count = 0; + int64_t hidden_size = 0; + + bool matched() const { + return first_get_rows != nullptr && second_get_rows != nullptr && add_node != nullptr && + first_source != nullptr && second_source != nullptr && row_ids != nullptr && output != nullptr; + } +}; + +static const GraphNode * find_single_add_consumer(const Graph & graph, const GraphNode & get_rows) { + const std::vector & consumers = graph.index().consumers(get_rows.output); + if (consumers.size() != 1) { + return nullptr; + } + const GraphNode * add = consumers.front(); + return add != nullptr && add->op == GGML_OP_ADD && add->inputs.size() == 2 ? add : nullptr; +} + +static const GraphNode * peer_get_rows_input(const Graph & graph, + const GraphNode & add, + const GraphNode & root_get_rows) { + ValueId peer_output; + if (add.inputs[0] == root_get_rows.output) { + peer_output = add.inputs[1]; + } else if (add.inputs[1] == root_get_rows.output) { + peer_output = add.inputs[0]; + } else { + return nullptr; + } + + const GraphNode * peer = graph.index().producer(peer_output); + return peer != nullptr && peer->op == GGML_OP_GET_ROWS && peer->inputs.size() == 2 ? peer : nullptr; +} + +static GatherAddMatch match_gather_add_f32(const Graph & graph, const GraphNode * node, size_t node_index) { + GatherAddMatch match; + if (node == nullptr || node->op != GGML_OP_GET_ROWS || node->inputs.size() != 2 || !graph.has_index()) { + return match; + } + + const GraphNode * add_node = find_single_add_consumer(graph, *node); + const GraphNode * peer = add_node != nullptr ? peer_get_rows_input(graph, *add_node, *node) : nullptr; + if (add_node == nullptr || peer == nullptr) { + return {}; + } + + size_t peer_index = 0; + size_t add_index = 0; + if (!graph.index().node_index(peer, peer_index) || !graph.index().node_index(add_node, add_index)) { + return {}; + } + + const Value * first_source = graph_value(graph, node->inputs[0]); + const Value * first_ids = graph_value(graph, node->inputs[1]); + const Value * first_output = graph_value(graph, node->output); + const Value * second_source = graph_value(graph, peer->inputs[0]); + const Value * second_ids = graph_value(graph, peer->inputs[1]); + const Value * second_output = graph_value(graph, peer->output); + const Value * output = graph_value(graph, add_node->output); + if (first_source == nullptr || first_ids == nullptr || first_output == nullptr || second_source == nullptr || + second_ids == nullptr || second_output == nullptr || output == nullptr) { + return {}; + } + if (node->inputs[1] != peer->inputs[1]) { + return {}; + } + if (!is_2d_f32(*first_source) || !is_2d_f32(*second_source) || !is_2d_f32(*first_output) || + !is_2d_f32(*second_output) || !is_2d_f32(*output)) { + return {}; + } + if (!same_shape(*first_source, *second_source) || !same_shape(*first_output, *second_output) || + !same_shape(*first_output, *output)) { + return {}; + } + + const int64_t hidden_size = first_source->ne[0]; + const int64_t source_token_count = first_source->ne[1]; + const int64_t output_token_count = first_output->ne[1]; + if (first_output->ne[0] != hidden_size || !is_row_id_tensor(*first_ids, output_token_count) || + !is_row_id_tensor(*second_ids, output_token_count)) { + return {}; + } + if (!is_supported_hidden_size(hidden_size) || !is_supported_token_count(source_token_count) || + !is_supported_token_count(output_token_count)) { + return {}; + } + + match.first_get_rows = node; + match.second_get_rows = peer; + match.add_node = add_node; + match.first_source = first_source; + match.second_source = second_source; + match.row_ids = first_ids; + match.output = output; + match.first_get_rows_index = node_index; + match.second_get_rows_index = peer_index; + match.add_node_index = add_index; + match.source_token_count = source_token_count; + match.output_token_count = output_token_count; + match.hidden_size = hidden_size; + return match; +} + +} // namespace + +static bool match_gather_add_f32_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + const GatherAddMatch gather_add = match_gather_add_f32(context.graph, context.root_node, context.root_index); + if (!gather_add.matched() || gather_add.first_get_rows_index >= context.covered_nodes.size() || + gather_add.second_get_rows_index >= context.covered_nodes.size() || + gather_add.add_node_index >= context.covered_nodes.size() || + context.covered_nodes[gather_add.first_get_rows_index] || + context.covered_nodes[gather_add.second_get_rows_index] || context.covered_nodes[gather_add.add_node_index]) { + return false; + } + if (!value_is_available(context.graph, gather_add.first_source->id, context.covered_nodes) || + !value_is_available(context.graph, gather_add.second_source->id, context.covered_nodes)) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kGatherAddF32Kernel); + dispatch.kernel.integer_parameters.emplace("source_token_count", gather_add.source_token_count); + dispatch.kernel.integer_parameters.emplace("output_token_count", gather_add.output_token_count); + dispatch.kernel.integer_parameters.emplace("hidden_size", gather_add.hidden_size); + dispatch.bindings.push_back({ gather_add.first_source->id, 0, gather_add.first_source->byte_count }); + dispatch.bindings.push_back({ gather_add.second_source->id, 0, gather_add.second_source->byte_count }); + dispatch.bindings.push_back({ gather_add.row_ids->id, 0, gather_add.row_ids->byte_count }); + dispatch.bindings.push_back({ gather_add.output->id, 0, gather_add.output->byte_count }); + + match.covered_nodes.push_back(gather_add.first_get_rows_index); + match.covered_nodes.push_back(gather_add.second_get_rows_index); + match.covered_nodes.push_back(gather_add.add_node_index); + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +void register_gather_add_dispatch(DispatchRegistryBuilder & registry) { + registry.add({ + "common.gather_add_f32", + GGML_OP_GET_ROWS, + DispatchMatchKind::Fused, + 1000, + DispatchSource::Common, + match_gather_add_f32_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gather-add.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gather-add.h new file mode 100644 index 000000000000..9f8d6e94a4d0 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-gather-add.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_gather_add_dispatch(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-get-rows.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-get-rows.cpp new file mode 100644 index 000000000000..821261dcbaad --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-get-rows.cpp @@ -0,0 +1,322 @@ +#include "dispatch-get-rows.h" + +#include "../qwen/dispatch-llm-profiles.h" +#include "dispatch-mul-mat-weight-format.h" +#include "ggml.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kGetRowsF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_get_rows_f32"); +static constexpr KernelCatalogRef kGetRowsF32NextKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_get_rows_f32_next"); +static constexpr int64_t kMaximumHiddenElements = int64_t{ 1 } << 30; +static constexpr int64_t kQ1_0GetRowsFormat = 10; +static constexpr int64_t kQwenHiddenSize = kQwen30BMoeDispatchProfile.hidden_size; +static constexpr int64_t kQwenVocabularyCount = 151936; +static constexpr int64_t kMaxGetRowsRowCount = 262208; + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static bool is_1d_or_2d_column(const Value & value) { + return value.ne[0] > 0 && value.ne[1] == 1 && value.ne[2] == 1 && value.ne[3] == 1; +} + +static bool is_2d(const Value & value) { + return value.ne[0] > 0 && value.ne[1] > 0 && value.ne[2] == 1 && value.ne[3] == 1; +} + +static bool is_supported_hidden_size(ggml_type type, int64_t hidden_size) { + if (hidden_size < 4 || hidden_size > kMaximumHiddenElements || hidden_size % 4 != 0) { + return false; + } + + const int64_t block_size = ggml_blck_size(type); + return block_size > 0 && hidden_size % block_size == 0; +} + +static bool is_supported_token_count(int64_t token_count) { + return token_count >= 1 && token_count <= 2048; +} + +static bool is_supported_row_count(int64_t row_count) { + return row_count >= 1 && row_count <= kMaxGetRowsRowCount; +} + +static std::string to_config_value(int64_t value) { + return std::to_string(value); +} + +static size_t row_byte_count(ggml_type type, int64_t token_count, int64_t hidden_size) { + if (token_count <= 0 || hidden_size <= 0) { + return 0; + } + return static_cast(token_count) * ggml_row_size(type, hidden_size); +} + +static bool is_qwen_q6k_q8_consumer(const Graph & graph, const GraphNode * consumer, const Value & input) { + if (consumer == nullptr || consumer->op != GGML_OP_MUL_MAT || consumer->inputs.size() != 2 || + consumer->inputs[1] != input.id) { + return false; + } + + const Value * weight = graph_value(graph, consumer->inputs[0]); + const Value * output = graph_value(graph, consumer->output); + if (weight == nullptr || output == nullptr || weight->type != GGML_TYPE_Q6_K || output->type != GGML_TYPE_F32 || + !weight->contiguous || !output->contiguous) { + return false; + } + + return input.ne[0] == kQwenHiddenSize && input.ne[1] == 1 && input.ne[2] == 1 && input.ne[3] == 1 && + weight->ne[0] == kQwenHiddenSize && weight->ne[1] == kQwenVocabularyCount && weight->ne[2] == 1 && + weight->ne[3] == 1 && output->ne[0] == kQwenVocabularyCount && output->ne[1] == 1 && output->ne[2] == 1 && + output->ne[3] == 1; +} + +static void append_unique_demand(std::vector & demands, ggml_type type) { + if (std::find(demands.begin(), demands.end(), type) == demands.end()) { + demands.push_back(type); + } +} + +static std::vector collect_alternate_demands(const Graph & graph, const Value & value) { + std::vector demands; + if (!graph.has_index()) { + return demands; + } + + const std::vector & consumers = graph.index().consumers(value.id); + for (const GraphNode * consumer : consumers) { + if (is_qwen_q6k_q8_consumer(graph, consumer, value)) { + append_unique_demand(demands, GGML_TYPE_Q8_1); + } + } + return demands; +} + +static const char * alternate_name(ggml_type type) { + switch (type) { + case GGML_TYPE_Q8_1: + return "common.get_rows.q8_1_x4"; + case GGML_TYPE_F16: + return "common.get_rows.f16"; + case GGML_TYPE_F32: + return "common.get_rows.f32"; + default: + return "common.get_rows.next"; + } +} + +static bool get_rows_format_for_type(ggml_type type, int64_t & format) { + if (type == GGML_TYPE_Q1_0) { + format = kQ1_0GetRowsFormat; + return true; + } + + CommonMulMatWeightFormat common_format; + if (!common_mul_mat_format_for_type(type, common_format)) { + return false; + } + + format = common_mul_mat_format_config_value(common_format); + return true; +} + +struct GetRowsMatch { + const Value * ids = nullptr; + const Value * weight = nullptr; + const Value * output = nullptr; + int64_t weight_format_value = -1; + int64_t token_count = 0; + int64_t row_count = 0; + int64_t hidden_size = 0; + + bool matched() const { + return ids != nullptr && weight != nullptr && output != nullptr && weight_format_value >= 0; + } +}; + +static void add_common_compile_parameters(Dispatch & dispatch, const GetRowsMatch & match) { + dispatch.kernel.compile_parameters.emplace("ggml.get_rows_f32.token_capacity", to_config_value(match.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.get_rows_f32.hidden_capacity", to_config_value(match.hidden_size)); + dispatch.kernel.compile_parameters.emplace("ggml.get_rows_f32.weight_format", + to_config_value(match.weight_format_value)); +} + +static void add_common_integer_parameters(Dispatch & dispatch, const GetRowsMatch & match) { + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.integer_parameters.emplace("row_count", match.row_count); + dispatch.kernel.integer_parameters.emplace("hidden_size", match.hidden_size); +} + +static void add_primary_bindings(Dispatch & dispatch, const GetRowsMatch & match) { + dispatch.bindings.push_back({ match.ids->id, 0, match.ids->byte_count }); + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); +} + +static GetRowsMatch match_get_rows_f32(const Graph & graph, const GraphNode * node) { + GetRowsMatch match; + if (node == nullptr || node->op != GGML_OP_GET_ROWS || node->inputs.size() != 2) { + return match; + } + + const Value * weight = graph_value(graph, node->inputs[0]); + const Value * ids = graph_value(graph, node->inputs[1]); + const Value * output = graph_value(graph, node->output); + if (weight == nullptr || ids == nullptr || output == nullptr || weight->type == GGML_TYPE_IQ3_S || + weight->type == GGML_TYPE_IQ4_NL || ids->type != GGML_TYPE_I32 || output->type != GGML_TYPE_F32 || + !weight->contiguous || !ids->contiguous || !output->contiguous || !is_2d(*weight) || + !is_1d_or_2d_column(*ids) || !is_2d(*output)) { + return {}; + } + + int64_t weight_format_value = 0; + if (!get_rows_format_for_type(weight->type, weight_format_value)) { + return {}; + } + + const int64_t hidden_size = weight->ne[0]; + const int64_t row_count = weight->ne[1]; + const int64_t token_count = ids->ne[0]; + if (output->ne[0] != hidden_size || output->ne[1] != token_count || + !is_supported_hidden_size(weight->type, hidden_size) || !is_supported_row_count(row_count) || + !is_supported_token_count(token_count)) { + return {}; + } + + match.ids = ids; + match.weight = weight; + match.output = output; + match.weight_format_value = weight_format_value; + match.token_count = token_count; + match.row_count = row_count; + match.hidden_size = hidden_size; + return match; +} + +static Dispatch make_get_rows_dispatch(const GetRowsMatch & match) { + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kGetRowsF32Kernel); + add_common_integer_parameters(dispatch, match); + add_common_compile_parameters(dispatch, match); + add_primary_bindings(dispatch, match); + return dispatch; +} + +static Dispatch make_get_rows_next_dispatch(const GetRowsMatch & match, + CommonMulMatWeightFormat next_format, + ValueId next_value, + size_t next_byte_count) { + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kGetRowsF32NextKernel); + add_common_integer_parameters(dispatch, match); + add_common_compile_parameters(dispatch, match); + dispatch.kernel.compile_parameters.emplace("ggml.get_rows_f32.next_format", + to_config_value(common_mul_mat_format_config_value(next_format))); + add_primary_bindings(dispatch, match); + dispatch.bindings.push_back({ next_value, 0, next_byte_count }); + return dispatch; +} + +static bool append_alternate_metadata(DispatchMatch & dispatch_match, + const Value & output, + ValueId alternate_value, + ggml_type type, + size_t byte_count) { + Status metadata_status; + if (!dispatch_match.metadata.append_alternate_value( + { output.id, alternate_value, type, byte_count, alternate_name(type) }, metadata_status)) { + dispatch_match.status.append(metadata_status); + return false; + } + return true; +} + +static bool match_get_rows_f32_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const GetRowsMatch match = match_get_rows_f32(context.graph, context.root_node); + if (!match.matched()) { + return false; + } + + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(make_get_rows_dispatch(match)); + return true; +} + +static bool match_get_rows_f32_next_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const GetRowsMatch match = match_get_rows_f32(context.graph, context.root_node); + if (!match.matched()) { + return false; + } + + const std::vector demands = collect_alternate_demands(context.graph, *match.output); + if (demands.empty()) { + return false; + } + + for (ggml_type type : demands) { + CommonMulMatWeightFormat next_format; + if (!common_mul_mat_alternate_format_for_type(type, next_format)) { + return false; + } + + const size_t next_byte_count = row_byte_count(type, match.token_count, match.hidden_size); + if (next_byte_count == 0) { + return false; + } + + if (type == GGML_TYPE_F32 && next_byte_count == match.output->byte_count) { + if (!append_alternate_metadata(dispatch_match, *match.output, match.output->id, type, next_byte_count)) { + return false; + } + continue; + } + + const ValueId next_value(context.next_plan_value.value + + static_cast(dispatch_match.transients.size())); + dispatch_match.dispatches.push_back( + make_get_rows_next_dispatch(match, next_format, next_value, next_byte_count)); + dispatch_match.transients.push_back({ next_value, alternate_name(type), next_byte_count, 256 }); + if (!append_alternate_metadata(dispatch_match, *match.output, next_value, type, next_byte_count)) { + return false; + } + } + + if (dispatch_match.dispatches.empty()) { + dispatch_match.dispatches.push_back(make_get_rows_dispatch(match)); + } + dispatch_match.covered_nodes.push_back(context.root_index); + return dispatch_match.status.success(); +} + +} // namespace + +void register_get_rows_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "common.get_rows.f32_next", + GGML_OP_GET_ROWS, + DispatchMatchKind::Fused, + 150, + DispatchSource::Common, + match_get_rows_f32_next_dispatch, + }); + registry.add({ + "common.get_rows.f32", + GGML_OP_GET_ROWS, + DispatchMatchKind::SingleOp, + 100, + DispatchSource::Common, + match_get_rows_f32_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-get-rows.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-get-rows.h new file mode 100644 index 000000000000..579f725c3ed0 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-get-rows.h @@ -0,0 +1,9 @@ +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +namespace ggml::hrx { + +void register_get_rows_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-glu.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-glu.cpp new file mode 100644 index 000000000000..66f13832a181 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-glu.cpp @@ -0,0 +1,200 @@ +#include "dispatch-glu.h" + +#include "dispatch-layout-utils.h" +#include "dispatch-mul-mat-common.h" +#include "ggml.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kBinaryF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_binary_f32"); + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static bool positive_shape(const Value & value) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (value.ne[i] <= 0) { + return false; + } + } + return value.element_count > 0; +} + +static bool packed_f32_layout(const Value & value) { + size_t expected_stride = sizeof(float); + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (value.nb[i] != expected_stride) { + return false; + } + expected_stride *= static_cast(value.ne[i]); + } + return true; +} + +static bool storage_ranges_disjoint(const Value & lhs, + size_t lhs_offset, + size_t lhs_byte_count, + const Value & rhs, + size_t rhs_offset, + size_t rhs_byte_count) { + if (lhs.storage != rhs.storage) { + return true; + } + if (lhs.storage_offset > std::numeric_limits::max() - lhs_offset || + rhs.storage_offset > std::numeric_limits::max() - rhs_offset) { + return false; + } + const size_t lhs_start = lhs.storage_offset + lhs_offset; + const size_t rhs_start = rhs.storage_offset + rhs_offset; + if (lhs_start > std::numeric_limits::max() - lhs_byte_count || + rhs_start > std::numeric_limits::max() - rhs_byte_count) { + return false; + } + return lhs_start + lhs_byte_count <= rhs_start || rhs_start + rhs_byte_count <= lhs_start; +} + +static bool packed_glu_half_span_bytes(const Value & input, + const Value & output, + size_t half_offset, + size_t & byte_count) { + if (input.type != GGML_TYPE_F32 || input.nb[0] != sizeof(float)) { + return false; + } + for (int i = 1; i < GGML_MAX_DIMS; ++i) { + if (input.nb[i] % sizeof(float) != 0) { + return false; + } + } + + size_t max_offset = 0; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (output.ne[i] <= 0) { + return false; + } + const size_t extent = static_cast(output.ne[i] - 1); + if (extent != 0 && input.nb[i] > std::numeric_limits::max() / extent) { + return false; + } + const size_t dim_offset = extent * input.nb[i]; + if (max_offset > std::numeric_limits::max() - dim_offset) { + return false; + } + max_offset += dim_offset; + } + if (max_offset > std::numeric_limits::max() - sizeof(float)) { + return false; + } + byte_count = max_offset + sizeof(float); + if (input.storage_offset > std::numeric_limits::max() - half_offset || + input.storage_offset + half_offset > std::numeric_limits::max() - byte_count) { + return false; + } + return input.storage_offset + half_offset + byte_count <= input.storage_byte_count; +} + +static void add_packed_glu_parameters(Dispatch & dispatch, + const Value & input, + const Value & output, + size_t half_span) { + dispatch.kernel.integer_parameters.emplace("element_count", output.element_count); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.ne0", std::to_string(output.ne[0])); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.ne1", std::to_string(output.ne[1])); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.ne2", std::to_string(output.ne[2])); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src0_stride1", + std::to_string(input.nb[1] / sizeof(float))); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src0_stride2", + std::to_string(input.nb[2] / sizeof(float))); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src0_stride3", + std::to_string(input.nb[3] / sizeof(float))); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src1_stride1", + std::to_string(input.nb[1] / sizeof(float))); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src1_stride2", + std::to_string(input.nb[2] / sizeof(float))); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src1_stride3", + std::to_string(input.nb[3] / sizeof(float))); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src0_span", std::to_string(half_span / sizeof(float))); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src1_span", std::to_string(half_span / sizeof(float))); +} + +static void bind_packed_glu_source(Dispatch & dispatch, const Value & input, size_t offset, size_t byte_count) { + dispatch.bindings.push_back({ input.storage_root, input.storage_offset + offset, byte_count }); +} + +static bool match_packed_glu_f32_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_GLU || node->inputs.size() != 1) { + return false; + } + + BinaryKind op; + if (!common_fused_binary_kind_from_params(node->params, op)) { + return false; + } + + const GluParams * params = op_params_as(node->params); + if (params == nullptr) { + return false; + } + + const Value * input = graph_value(context.graph, node->inputs[0]); + const Value * output = graph_value(context.graph, node->output); + if (input == nullptr || output == nullptr || input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || + !positive_shape(*output) || input->ne[0] != 2 * output->ne[0] || output->alias_source.value >= 0 || + !output->contiguous || !packed_f32_layout(*output) || + static_cast(output->element_count) > std::numeric_limits::max()) { + return false; + } + for (int i = 1; i < GGML_MAX_DIMS; ++i) { + if (input->ne[i] != output->ne[i]) { + return false; + } + } + + const size_t half_offset = static_cast(output->ne[0]) * sizeof(float); + size_t half_span = 0; + size_t input_span = 0; + if (!packed_glu_half_span_bytes(*input, *output, 0, half_span) || + !packed_glu_half_span_bytes(*input, *output, half_offset, half_span) || + !strided_f32_storage_span_bytes(*input, input_span) || + !storage_ranges_disjoint(*input, 0, input_span, *output, 0, output->byte_count)) { + return false; + } + + const size_t gate_offset = params->swapped ? half_offset : 0; + const size_t up_offset = params->swapped ? 0 : half_offset; + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kBinaryF32Kernel); + add_packed_glu_parameters(dispatch, *input, *output, half_span); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.op", std::to_string(binary_kind_config_value(op))); + bind_packed_glu_source(dispatch, *input, gate_offset, half_span); + bind_packed_glu_source(dispatch, *input, up_offset, half_span); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + + match.covered_nodes.push_back(context.root_index); + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_glu_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "common.packed_glu_f32", + GGML_OP_GLU, + DispatchMatchKind::SingleOp, + 10, + DispatchSource::Common, + match_packed_glu_f32_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-glu.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-glu.h new file mode 100644 index 000000000000..fec5fa69ed07 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-glu.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_glu_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-grouped-mul-mat.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-grouped-mul-mat.cpp new file mode 100644 index 000000000000..f3955e098b26 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-grouped-mul-mat.cpp @@ -0,0 +1,94 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +#include "dispatch-grouped-mul-mat.h" + +#include "ggml.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kGroupedMulMatF16F32Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_grouped_mul_mat_f16_f32"); + +// ne[0..2] as given, ne[3] == 1, and densely packed with the given element size +static bool packed_3d(const Value & value, int64_t ne0, int64_t ne1, int64_t ne2, size_t element_size) { + if (value.ne[0] != ne0 || value.ne[1] != ne1 || value.ne[2] != ne2 || value.ne[3] != 1) { + return false; + } + return value.nb[0] == element_size && value.nb[1] == element_size * static_cast(ne0) && + value.nb[2] == value.nb[1] * static_cast(ne1); +} + +static bool match_grouped_mul_mat_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_MUL_MAT || node->inputs.size() != 2) { + return false; + } + const Value * weight = context.graph.values().find(node->inputs[0]); + const Value * input = context.graph.values().find(node->inputs[1]); + const Value * output = context.graph.values().find(node->output); + if (weight == nullptr || input == nullptr || output == nullptr) { + return false; + } + if (weight->type != GGML_TYPE_F16 || input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32) { + return false; + } + const int64_t k = weight->ne[0]; + const int64_t n = weight->ne[1]; + const int64_t g = weight->ne[2]; + const int64_t m = input->ne[1]; + // only the batched case: 2-D weights belong to the regular MUL_MAT kernels + if (g < 2 || g > 4096 || k < 1 || k > 65536 || n < 1 || n > 65536 || m < 1 || m > 65535) { + return false; + } + if (!packed_3d(*weight, k, n, g, sizeof(uint16_t)) || !packed_3d(*input, k, m, g, sizeof(float)) || + !packed_3d(*output, n, m, g, sizeof(float)) || output->alias_source.value >= 0 || + input->storage == output->storage || weight->storage == output->storage) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kGroupedMulMatF16F32Kernel); + dispatch.kernel.integer_parameters.emplace("input_size", k); + dispatch.kernel.integer_parameters.emplace("output_size", n); + dispatch.kernel.integer_parameters.emplace("token_count", m); + dispatch.kernel.integer_parameters.emplace("group_count", g); + dispatch.bindings.push_back({ weight->id, 0, weight->byte_count }); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + + match.covered_nodes.push_back(context.root_index); + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_grouped_mul_mat_dispatch(DispatchRegistryBuilder & registry) { + registry.add({ + "common.grouped_mul_mat_f16_f32", + GGML_OP_MUL_MAT, + DispatchMatchKind::SingleOp, + 0, + DispatchSource::Common, + match_grouped_mul_mat_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-grouped-mul-mat.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-grouped-mul-mat.h new file mode 100644 index 000000000000..39c65942472f --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-grouped-mul-mat.h @@ -0,0 +1,25 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +namespace ggml::hrx { + +// MUL_MAT with a batched F16 weight (one matrix per group, no broadcast), such as ZAYA's +// grouped convolution: ggml_grouped_mul_mat_f16_f32. +void register_grouped_mul_mat_dispatch(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-hadamard.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-hadamard.cpp new file mode 100644 index 000000000000..22f14eb632fe --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-hadamard.cpp @@ -0,0 +1,109 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// A MUL_MAT that llama marks GGML_HINT_SRC0_IS_HADAMARD (src0 is the normalized n x n Sylvester +// matrix) as one Walsh-Hadamard transform per row (ops/hadamard_f32.loom), as the CPU, Vulkan, +// CUDA and Metal backends do for that hint. The dense matmul matchers stop at 2048 rows, so +// Bonsai's prompt-time rotations (512 tokens x 17 blocks of 1024) used to run on the CPU in every +// layer: pp512 13 tok/s. The matrix is not read; only its size picks the transform. + +#include "dispatch-hadamard.h" + +#include "dispatch-mul-mat-common.h" +#include "graph/graph-matcher.h" + +#include + +namespace ggml::hrx { + +namespace { + +static constexpr KernelCatalogRef kHadamardKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_hadamard_f32"); +constexpr int32_t kHintSrc0IsHadamard = 1; // GGML_HINT_SRC0_IS_HADAMARD + +bool packed_f32_rows(const Value * value, int64_t n) { + return value != nullptr && value->type == GGML_TYPE_F32 && value->contiguous && value->ne[0] == n && + value->ne[2] == 1 && value->ne[3] == 1 && value->element_count == value->ne[0] * value->ne[1]; +} + +int log2_exact(int64_t n) { + int k = 0; + while ((int64_t{ 1 } << k) < n) { + ++k; + } + return (int64_t{ 1 } << k) == n ? k : -1; +} + +// Each workgroup reads its whole row before it writes, so the output may be the input itself, +// but not a shifted overlap of it. +bool in_place_or_disjoint(const Value & input, const Value & output) { + if (input.storage != output.storage || input.storage_offset == output.storage_offset) { + return true; + } + return input.storage_offset + input.byte_count <= output.storage_offset || + output.storage_offset + output.byte_count <= input.storage_offset; +} + +bool match_hadamard(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_MUL_MAT || node->inputs.size() != 2) { + return false; + } + const MulMatParams * params = op_params_as(node->params); + if (params == nullptr || params->hint != kHintSrc0IsHadamard) { + return false; + } + const Value * rotation = common_graph_value(context.graph, node->inputs[0]); + const Value * input = common_graph_value(context.graph, node->inputs[1]); + const Value * output = common_graph_value(context.graph, node->output); + if (rotation == nullptr || rotation->ne[0] != rotation->ne[1] || rotation->ne[2] != 1 || rotation->ne[3] != 1) { + return false; + } + const int64_t n = rotation->ne[0]; + const int k = log2_exact(n); + if (k < 6 || k > 12 || !packed_f32_rows(input, n) || !packed_f32_rows(output, n) || + output->ne[1] != input->ne[1] || output->ne[1] < 1 || !in_place_or_disjoint(*input, *output)) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kHadamardKernel); + dispatch.kernel.integer_parameters.emplace("row_count", input->ne[1]); + dispatch.kernel.compile_parameters.emplace("ggml.hadamard_f32.log2_block", std::to_string(k)); + dispatch.bindings.push_back({ input->storage_root, input->storage_offset, input->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_hadamard_dispatches(DispatchRegistryBuilder & registry) { + // Fused, not SingleOp: the registry tries every Fused matcher of an op before any SingleOp + // one, and the dense MUL_MAT matchers are Fused, so as SingleOp a hinted MUL_MAT of up to + // 2048 rows (Bonsai decode: 17) went to a dense WMMA kernel (1e-2 off the exact product). + // Priority 400 is above every dense MUL_MAT matcher (the highest is 315). + registry.add({ + "common.hadamard_f32", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 400, + DispatchSource::Common, + match_hadamard, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-hadamard.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-hadamard.h new file mode 100644 index 000000000000..317389782d59 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-hadamard.h @@ -0,0 +1,24 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +namespace ggml::hrx { + +void register_hadamard_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-kquant-decode.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-kquant-decode.cpp new file mode 100644 index 000000000000..c56f79febee0 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-kquant-decode.cpp @@ -0,0 +1,357 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Decode projections (1 token, or 2-8 for MTP / speculative verify batches) on Q2_K, Q3_K, Q4_K, Q5_K, +// Q6_K, IQ1_S, IQ1_M, IQ2_XXS, IQ2_XS, IQ2_S, IQ3_XXS, IQ3_S, IQ4_NL, IQ4_XS, Q8_0, TQ1_0, TQ2_0 and MXFP4 weights, read in +// their GGUF block layout (and exact-ternary Q4_0 repacked to 2 bits, see kquant_ternary; Q1_0 and PrismML's PQ2_0 / +// PTQ1_0 in their group-128 GGUF layout) +// (ops/kquant_decode_f32.loom): FFN gate/up pairs fused with SwiGLU, and plain projections with an +// optional following residual ADD. Mixed-quant models (Unsloth UD-Q4_K_XL and similar) pair these +// types freely per layer; without this they take the generic dequantize-4-values-at-a-time +// kernels. + +#include "dispatch-kquant-decode.h" + +#include "dispatch-mul-mat-common.h" +#include "dispatch/ternary-q4-0.h" +#include "graph/graph-matcher.h" + +#include +#include +#include + +namespace ggml::hrx { + +namespace { + +static constexpr KernelCatalogRef kKQuantSwiGLUDecodeKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_kquant_swiglu_decode_f32"); +static constexpr KernelCatalogRef kKQuantMulMatDecodeKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_kquant_mul_mat_decode_f32"); +// 2..8 tokens (MTP / speculative verify batches): weights dequantized once per lane for all tokens. +static constexpr KernelCatalogRef kKQuantSwiGLUDecodeTokensKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_kquant_swiglu_decode_tokens_f32"); +static constexpr KernelCatalogRef kKQuantMulMatDecodeTokensKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_kquant_mul_mat_decode_tokens_f32"); + +bool kquant_format(CommonMulMatWeightFormat format) { + switch (format) { + case CommonMulMatWeightFormat::Q2K: + case CommonMulMatWeightFormat::IQ1_S: + case CommonMulMatWeightFormat::IQ1_M: + case CommonMulMatWeightFormat::IQ2_XXS: + case CommonMulMatWeightFormat::IQ2_XS: + case CommonMulMatWeightFormat::IQ3_XXS: + case CommonMulMatWeightFormat::IQ2_S: + case CommonMulMatWeightFormat::Q3K: + case CommonMulMatWeightFormat::Q4K: + case CommonMulMatWeightFormat::Q5K: + case CommonMulMatWeightFormat::Q6K: + case CommonMulMatWeightFormat::IQ3_S: + case CommonMulMatWeightFormat::IQ4_NL: + case CommonMulMatWeightFormat::IQ4_XS: + case CommonMulMatWeightFormat::Q8_0: + case CommonMulMatWeightFormat::Q1_0: + case CommonMulMatWeightFormat::PQ2_0: + case CommonMulMatWeightFormat::PTQ1_0: + case CommonMulMatWeightFormat::TQ1_0: + case CommonMulMatWeightFormat::TQ2_0: + case CommonMulMatWeightFormat::MXFP4: + return true; + default: + return false; + } +} + +// Q4_0 joins as packed ternary (format 90, dispatch/ternary-q4-0.h) when GGML_HRX_TERNARY_Q4_0 is set: the +// weight binding asks for the repacked layout, which the upload verifies value by value. +bool kquant_ternary(const CommonMulMatMatch & match) { + return match.weight_format == CommonMulMatWeightFormat::Q4_0 && ternary_q4_0_enabled() && + match.input_size % 256 == 0; +} + +bool kquant_supported(const CommonMulMatMatch & match) { + return kquant_format(match.weight_format) || kquant_ternary(match); +} + +int64_t kquant_format_value(const CommonMulMatMatch & match) { + return kquant_ternary(match) ? kTernaryQ40K128FormatValue : common_mul_mat_format_config_value(match.weight_format); +} + +DispatchBinding kquant_weight_binding(const CommonMulMatMatch & match) { + if (kquant_ternary(match)) { + return { match.weight->id, + 0, + ternary_q4_0_k128_bytes(match.input_size, match.output_size), + kTernaryQ40K128Layout, + match.weight->type, + match.input_size, + match.output_size, + match.weight->byte_count }; + } + return { match.weight->id, 0, match.weight->byte_count }; +} + +// The kernels' config ranges (ops/kquant_decode_f32.loom). +constexpr int64_t kMaxInputSize = 65536; +constexpr int64_t kMaxOutputSize = 1048576; +constexpr int64_t kMaxTokens = 8; + +// A MUL_MAT with 1..kMaxTokens tokens: the common matcher's decode form admits exactly one token and +// its prefill form two or more. +CommonMulMatMatch match_few_token_mul_mat(const Graph & graph, const GraphNode * node, KernelCatalogRef kernel) { + CommonMulMatMatch match = common_match_mul_mat_any_format(graph, node, kernel, true); + if (!match.matched()) { + match = common_match_mul_mat_any_format(graph, node, kernel, false); + } + if (!match.matched() || match.token_count > kMaxTokens) { + return {}; + } + return match; +} + +// A value is ready at this root when its producer (looking through layout aliases) is a graph input +// or an already-claimed node. A dispatch is emitted at its root's position, so reading a value that +// a later-rooted fusion will produce fails with "reads transient value before write". +bool value_ready(const DispatchMatchContext & context, ValueId id) { + const Graph & graph = context.graph; + for (int depth = 0; depth < 16; ++depth) { + const GraphNode * producer = graph.index().producer(id); + if (producer == nullptr) { + return true; + } + size_t index = 0; + if (!graph.index().node_index(producer, index) || index >= context.covered_nodes.size()) { + return false; + } + if (context.covered_nodes[index]) { + return true; + } + if (!is_layout_alias_node(graph, *producer) || producer->inputs.empty()) { + return false; + } + id = producer->inputs[0]; + } + return false; +} + +// Some fused producers (e.g. qwen3_moe's routed down + next-layer norm) publish the next input only +// as a Q8_1 alternate and never write the F32 value when every consumer reads the alternate. +// These kernels read F32, so leave such inputs to the q8 consumers. +bool has_q8_alternate(const DispatchMatchContext & context, const Value & input, int64_t input_size, + int64_t token_count) { + const size_t bytes = static_cast(token_count) * ggml_row_size(GGML_TYPE_Q8_1, input_size); + return find_alternate_value(context.graph, context.plan, input.id, GGML_TYPE_Q8_1, bytes) != nullptr; +} + +bool uncovered(const DispatchMatchContext & context, const GraphNode * node) { + size_t index = 0; + return context.graph.index().node_index(node, index) && index < context.covered_nodes.size() && + !context.covered_nodes[index]; +} + +bool match_kquant_swiglu_decode(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const Graph & graph = context.graph; + if (!graph.has_index()) { + return false; + } + const CommonMulMatMatch root = + match_few_token_mul_mat(graph, context.root_node, kKQuantSwiGLUDecodeKernel); + if (!root.matched() || root.token_count < 1 || root.token_count > kMaxTokens || !kquant_supported(root) || + root.input_size % 256 != 0 || root.input_size > kMaxInputSize || root.output_size > kMaxOutputSize || + !value_ready(context, root.input->id) || + has_q8_alternate(context, *root.input, root.input_size, root.token_count)) { + return false; + } + + // The root's only consumer is a SwiGLU whose other operand is a matching MUL_MAT of the same input. + const std::vector & root_consumers = graph.index().consumers(context.root_node->output); + const GraphNode * glu = root_consumers.size() == 1 ? root_consumers.front() : nullptr; + BinaryKind op; + if (glu == nullptr || glu->inputs.size() != 2 || !common_fused_binary_kind_from_params(glu->params, op) || + op != BinaryKind::SwiGLU || !uncovered(context, glu)) { + return false; + } + const bool root_is_gate = glu->inputs[0] == context.root_node->output; + if (!root_is_gate && glu->inputs[1] != context.root_node->output) { + return false; + } + const ValueId peer_id = root_is_gate ? glu->inputs[1] : glu->inputs[0]; + const GraphNode * peer = graph.index().producer(peer_id); + if (peer == nullptr || peer == context.root_node || !uncovered(context, peer) || + graph.index().consumers(peer_id).size() != 1) { + return false; + } + const CommonMulMatMatch other = match_few_token_mul_mat(graph, peer, kKQuantSwiGLUDecodeKernel); + if (!other.matched() || !kquant_supported(other) || other.input->id != root.input->id || + other.input_size != root.input_size || other.output_size != root.output_size || + other.token_count != root.token_count) { + return false; + } + const Value * output = common_graph_value(graph, glu->output); + if (output == nullptr || output->type != GGML_TYPE_F32 || !output->contiguous || + !common_same_shape(*output, *root.output)) { + return false; + } + + const CommonMulMatMatch & gate = root_is_gate ? root : other; + const CommonMulMatMatch & up = root_is_gate ? other : root; + for (const GraphNode * node : { context.root_node, peer, glu }) { + if (!append_covered_node_index_once(graph, context.covered_nodes, node, dispatch_match.covered_nodes)) { + dispatch_match.covered_nodes.clear(); + return false; + } + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(root.token_count == 1 ? kKQuantSwiGLUDecodeKernel : + kKQuantSwiGLUDecodeTokensKernel); + if (root.token_count > 1) { + dispatch.kernel.compile_parameters.emplace("ggml.kquant_decode.token_count", + std::to_string(root.token_count)); + } + dispatch.kernel.compile_parameters.emplace("ggml.kquant_swiglu_decode.input_size", + std::to_string(root.input_size)); + dispatch.kernel.compile_parameters.emplace("ggml.kquant_swiglu_decode.output_size", + std::to_string(root.output_size)); + dispatch.kernel.compile_parameters.emplace("ggml.kquant_swiglu_decode.gate_weight_format", + std::to_string(kquant_format_value(gate))); + dispatch.kernel.compile_parameters.emplace("ggml.kquant_swiglu_decode.up_weight_format", + std::to_string(kquant_format_value(up))); + dispatch.bindings.push_back({ root.input->id, 0, root.input->byte_count }); + dispatch.bindings.push_back(kquant_weight_binding(gate)); + dispatch.bindings.push_back(kquant_weight_binding(up)); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + + +bool match_kquant_mul_mat_decode_impl(const DispatchMatchContext & context, DispatchMatch & dispatch_match, + bool require_add) { + const Graph & graph = context.graph; + if (!graph.has_index()) { + return false; + } + const CommonMulMatMatch root = + match_few_token_mul_mat(graph, context.root_node, kKQuantMulMatDecodeKernel); + if (!root.matched() || root.token_count < 1 || root.token_count > kMaxTokens || !kquant_supported(root) || + root.input_size % 256 != 0 || root.input_size > kMaxInputSize || root.output_size > kMaxOutputSize || + !value_ready(context, root.input->id) || + has_q8_alternate(context, *root.input, root.input_size, root.token_count)) { + return false; + } + + // Fold a following residual ADD (the projection's only consumer) into the store. + const Value * addend = nullptr; + const Value * output = root.output; + const GraphNode * add = common_find_only_consumer_with_op(graph, root.output->id, GGML_OP_ADD); + if (add != nullptr && common_binary_node_is_add(*add) && uncovered(context, add)) { + const bool root_is_lhs = add->inputs[0] == root.output->id; + const Value * other = common_graph_value(graph, root_is_lhs ? add->inputs[1] : add->inputs[0]); + const Value * sum = common_graph_value(graph, add->output); + if ((root_is_lhs || add->inputs[1] == root.output->id) && other != nullptr && sum != nullptr && + other->type == GGML_TYPE_F32 && sum->type == GGML_TYPE_F32 && other->contiguous && sum->contiguous && + common_same_shape(*root.output, *other) && common_same_shape(*root.output, *sum) && + value_ready(context, other->id)) { + addend = other; + output = sum; + } else { + add = nullptr; + } + } else { + add = nullptr; + } + if (require_add && add == nullptr) { + return false; + } + + if (!append_covered_node_index_once(graph, context.covered_nodes, context.root_node, + dispatch_match.covered_nodes) || + (add != nullptr && + !append_covered_node_index_once(graph, context.covered_nodes, add, dispatch_match.covered_nodes))) { + dispatch_match.covered_nodes.clear(); + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(root.token_count == 1 ? kKQuantMulMatDecodeKernel : + kKQuantMulMatDecodeTokensKernel); + if (root.token_count > 1) { + dispatch.kernel.compile_parameters.emplace("ggml.kquant_decode.token_count", + std::to_string(root.token_count)); + } + dispatch.kernel.compile_parameters.emplace("ggml.kquant_mul_mat_decode.input_size", + std::to_string(root.input_size)); + dispatch.kernel.compile_parameters.emplace("ggml.kquant_mul_mat_decode.output_size", + std::to_string(root.output_size)); + dispatch.kernel.compile_parameters.emplace("ggml.kquant_mul_mat_decode.weight_format", + std::to_string(kquant_format_value(root))); + dispatch.kernel.compile_parameters.emplace("ggml.kquant_mul_mat_decode.add", addend != nullptr ? "1" : "0"); + dispatch.bindings.push_back({ root.input->id, 0, root.input->byte_count }); + dispatch.bindings.push_back(kquant_weight_binding(root)); + // Without an ADD the addend is never read; bind the input in its place. + const Value * addend_binding = addend != nullptr ? addend : root.input; + dispatch.bindings.push_back({ addend_binding->id, 0, addend_binding->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +bool match_kquant_mul_mat_decode(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + return match_kquant_mul_mat_decode_impl(context, dispatch_match, false); +} + +bool match_kquant_mul_mat_add_decode(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + return match_kquant_mul_mat_decode_impl(context, dispatch_match, true); +} + +} // namespace + +void register_kquant_decode_dispatches(DispatchRegistryBuilder & registry) { + // Above common.mul_mat_swiglu.symmetric_i4_lowrow_adjacent_dual (310): that path repacks the + // weights to int4 at first use and quantizes activations to int4, and under llama-server it + // made 27B decode 10.7 tok/s against 11.9 with this kernel. + registry.add({ + "kquant.swiglu.decode_f32", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 315, + DispatchSource::Common, + match_kquant_swiglu_decode, + }); + // Projection + residual ADD above common.mul_mat_postops.f32_f32_wmma (180), which otherwise takes + // the 2-8 token (MTP verify) projections through the prefill WMMA kernel. + registry.add({ + "kquant.mul_mat_add.decode_f32", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 185, + DispatchSource::Common, + match_kquant_mul_mat_add_decode, + }); + // Plain projections above the generic common.mul_mat f32 matchers (80/70/60), below every + // specialized one. + registry.add({ + "kquant.mul_mat.decode_f32", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 85, + DispatchSource::Common, + match_kquant_mul_mat_decode, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-kquant-decode.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-kquant-decode.h new file mode 100644 index 000000000000..e4d74dc9acbb --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-kquant-decode.h @@ -0,0 +1,24 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +namespace ggml::hrx { + +void register_kquant_decode_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-layout-utils.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-layout-utils.h new file mode 100644 index 000000000000..e0c6d594732a --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-layout-utils.h @@ -0,0 +1,56 @@ +#pragma once + +#include "ggml.h" +#include "graph/value-map.h" + +#include +#include + +namespace ggml::hrx { + +inline bool strided_f32_storage_span_bytes(const Value & value, size_t & byte_count) { + if (value.type != GGML_TYPE_F32 || value.nb[0] != sizeof(float)) { + return false; + } + for (int i = 1; i < GGML_MAX_DIMS; ++i) { + if (value.nb[i] % sizeof(float) != 0) { + return false; + } + } + + size_t max_offset = 0; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (value.ne[i] <= 0) { + return false; + } + const size_t extent = static_cast(value.ne[i] - 1); + if (extent != 0 && value.nb[i] > std::numeric_limits::max() / extent) { + return false; + } + const size_t dim_offset = extent * value.nb[i]; + if (max_offset > std::numeric_limits::max() - dim_offset) { + return false; + } + max_offset += dim_offset; + } + if (max_offset > std::numeric_limits::max() - sizeof(float)) { + return false; + } + + byte_count = max_offset + sizeof(float); + if (value.storage_offset > std::numeric_limits::max() - byte_count) { + return false; + } + return value.storage_offset + byte_count <= value.storage_byte_count; +} + +inline bool strided_f32_storage_span_elements(const Value & value, size_t & element_count) { + size_t byte_count = 0; + if (!strided_f32_storage_span_bytes(value, byte_count) || byte_count % sizeof(float) != 0) { + return false; + } + element_count = byte_count / sizeof(float); + return true; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-moe-routing-layout.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-moe-routing-layout.h new file mode 100644 index 000000000000..b09cbcd1ddb6 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-moe-routing-layout.h @@ -0,0 +1,59 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +// Descriptor layout of the MoE partition table that the MoE router dispatch builds. +// +// ggml_moe_build_expert_partition_table packs each 32-row partition as +// expert | partition << partition_shift | (row_count - 1) << row_count_shift, and it +// runs one workgroup with one lane per expert. Two layouts are in use: +// +// - Qwen layout (7-bit expert: mask 127, shifts 7 and 13, 128 lanes). The qwen3_moe +// routed kernels decode this layout with hard-coded constants, and they only run +// for 128 experts. +// - Common layout (9-bit expert: mask 511, shifts 9 and 15, 512 lanes). The common +// mul_mat_id kernels (ggml_moe_unpack_expert_partition_descriptor) decode this one +// with hard-coded constants, for up to 512 experts. +// +// The router used the Qwen layout for every expert count. With more than 128 experts +// (Qwen3.6-35B-A3B has 256) the table builder masked expert ids to 7 bits, gave no +// partition to experts 128 and up (128 lanes), and the common mul_mat_id kernels that +// consume the table decoded row counts as partition ordinals. They then read assignment +// ordinals past the expert table and faulted (HSA_STATUS_ERROR_MEMORY_FAULT) on any +// prompt batch in which one expert received two or more tokens. No Qwen-layout +// consumer matches above 128 experts, so the router emits the common layout there. + +#include + +namespace ggml::hrx { + +struct MoeRoutingDescriptorLayout { + const char * expert_mask; + const char * partition_shift; + const char * row_count_shift; + const char * partition_workgroup_size; +}; + +inline constexpr int64_t kMoeRoutingQwenLayoutMaxExpertCount = 128; + +inline MoeRoutingDescriptorLayout moe_router_descriptor_layout(int64_t expert_count) { + if (expert_count <= kMoeRoutingQwenLayoutMaxExpertCount) { + return { "127", "7", "13", "128" }; + } + return { "511", "9", "15", "512" }; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-common.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-common.h new file mode 100644 index 000000000000..b2ffcd1335f7 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-common.h @@ -0,0 +1,607 @@ +#pragma once + +#include "../dispatch-registry.h" +#include "dispatch-mul-mat-weight-format.h" +#include "ggml.h" +#include "graph/graph-matcher.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include +#include +#include +#include + +// Q5_K/IQ4_XS prefill matmuls take the q8_1 x4 kernel like Q4_K (default on; +// GGML_HRX_Q8_PREFILL_RELAX=0 restores the previous policy). Defined in dispatch-mul-mat.cpp. +namespace ggml::hrx { +bool common_q8_prefill_relaxed(); +} // namespace ggml::hrx + +namespace ggml::hrx { + +struct CommonMulMatMatch { + const Value * input = nullptr; + const Value * weight = nullptr; + const Value * output = nullptr; + KernelCatalogRef kernel = {}; + int64_t input_size = 0; + int64_t output_size = 0; + int64_t token_count = 0; + CommonMulMatWeightFormat weight_format = CommonMulMatWeightFormat::Q4K; + UnaryKind output_unary_op = UnaryKind::Identity; + size_t unary_node_index = 0; + bool has_fused_unary = false; + + bool matched() const { + return input != nullptr && weight != nullptr && output != nullptr && kernel.id != kUncatalogedKernelId; + } +}; + +struct CommonSymmetricI4ActivationLayout { + size_t payload_bytes = 0; + size_t scales_offset = 0; + size_t metadata_bytes = 0; + size_t sums_offset = 0; + size_t total_bytes = 0; +}; + +inline constexpr const char kCommonSymmetricI4K32ActivationAlternateName[] = + "common.mul_mat.symmetric_i4_k32.activation"; + +enum class CommonQ8ActivationPolicy { + ExistingAlternateOnly, + AllowStandaloneQuantize, +}; + +inline size_t common_align_up(size_t value, size_t alignment) { + return (value + alignment - 1) / alignment * alignment; +} + +inline CommonSymmetricI4ActivationLayout common_symmetric_i4_activation_layout(int64_t input_size, + int64_t token_count) { + const size_t element_count = static_cast(input_size) * static_cast(token_count); + const size_t payload_bytes = element_count / 2; + const size_t metadata_bytes = element_count / 8; + const size_t scales_offset = common_align_up(payload_bytes, 256); + const size_t sums_offset = common_align_up(scales_offset + metadata_bytes, 256); + return { + payload_bytes, scales_offset, metadata_bytes, sums_offset, common_align_up(sums_offset + metadata_bytes, 256), + }; +} + +inline size_t common_symmetric_shared4_row_group_size(int64_t input_size, int64_t output_size, int64_t quant_bits) { + const size_t logical_output_size = static_cast(output_size); + const size_t materialized_row_bytes = 4 + 8 * 32 * static_cast(quant_bits) / 8; + const size_t materialized_bytes = + logical_output_size * static_cast(input_size / 256) * materialized_row_bytes; + if (quant_bits == 4 && materialized_bytes < size_t{ 16 } * 1024 * 1024) { + return 32; + } + if (quant_bits == 4 && materialized_bytes <= size_t{ 32 } * 1024 * 1024 && output_size > input_size) { + return 96; + } + return ((logical_output_size + 255) / 256) * 32; +} + +inline size_t common_symmetric_shared4_weight_byte_count(int64_t input_size, int64_t output_size, int64_t quant_bits) { + const size_t logical_output_size = static_cast(output_size); + const size_t row_group_size = common_symmetric_shared4_row_group_size(input_size, output_size, quant_bits); + const size_t padded_output_size = (logical_output_size + row_group_size - 1) / row_group_size * row_group_size; + const size_t materialized_row_bytes = 4 + 8 * 32 * static_cast(quant_bits) / 8; + return padded_output_size * static_cast(input_size / 256) * materialized_row_bytes; +} + +inline size_t common_symmetric_i4_shared4_weight_byte_count(int64_t input_size, int64_t output_size) { + return common_symmetric_shared4_weight_byte_count(input_size, output_size, 4); +} + +inline DispatchBinding common_symmetric_i4_shared4_weight_binding(const Value & weight, + int64_t input_size, + int64_t output_size) { + DispatchBinding binding; + binding.value = weight.id; + binding.length = common_symmetric_i4_shared4_weight_byte_count(input_size, output_size); + binding.layout = kSymmetricI4K32EightGroupsShared4Layout; + binding.source_type = weight.type; + binding.input_size = input_size; + binding.output_size = output_size; + binding.source_length = weight.byte_count; + return binding; +} + +inline DispatchBinding common_symmetric_i4_shared4_multistart_weight_binding(const Value & weight, + int64_t input_size, + int64_t output_size) { + DispatchBinding binding = common_symmetric_i4_shared4_weight_binding(weight, input_size, output_size); + binding.layout = kSymmetricI4K32EightGroupsShared4MultistartLayout; + return binding; +} + +inline const Value * common_graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +inline bool common_same_shape(const Value & lhs, const Value & rhs) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (lhs.ne[i] != rhs.ne[i]) { + return false; + } + } + return true; +} + +inline bool common_is_supported_dense_input_size(int64_t input_size) { + return input_size >= 256 && input_size <= 32768 && input_size % 256 == 0; +} + +inline bool common_is_supported_dense_input_size(CommonMulMatWeightFormat format, int64_t input_size) { + return common_mul_mat_supported_dense_input_size(format, input_size); +} + +inline bool common_is_supported_dense_output_size(int64_t output_size) { + return output_size >= 1 && output_size <= 262144; +} + +// Select the packed prefill schedule. Long Q4 contractions only benefit when +// their producer supplies the layout without a separate conversion pass. +inline bool common_mul_mat_uses_k16_major_f16(CommonMulMatWeightFormat format, + int64_t input_size, + int64_t output_size, + int64_t token_count, + bool packed_producer = false) { + if (token_count < 512 || token_count > 2048 || token_count % 512 != 0 || + input_size % 256 != 0 || output_size % 64 != 0) { + return false; + } + if (format == CommonMulMatWeightFormat::Q4K) { + return token_count == 512 && output_size / 64 >= 64 && + (input_size <= 2 * output_size || (packed_producer && input_size <= 4 * output_size)); + } + if (format == CommonMulMatWeightFormat::Q6K) { + const int64_t tile = input_size > output_size ? 64 : 128; + return output_size % tile == 0 && (token_count / 512) * (output_size / tile) >= 32; + } + return false; +} + +inline ggml_type common_mul_mat_format_type(CommonMulMatWeightFormat format) { + switch (format) { + case CommonMulMatWeightFormat::Q1_0: + return GGML_TYPE_Q1_0; + case CommonMulMatWeightFormat::PQ2_0: + return GGML_TYPE_PQ2_0; + case CommonMulMatWeightFormat::PTQ1_0: + return GGML_TYPE_PTQ1_0; + case CommonMulMatWeightFormat::Q2K: + return GGML_TYPE_Q2_K; + case CommonMulMatWeightFormat::Q3K: + return GGML_TYPE_Q3_K; + case CommonMulMatWeightFormat::Q4K: + case CommonMulMatWeightFormat::Q4KRow64: + return GGML_TYPE_Q4_K; + case CommonMulMatWeightFormat::Q5K: + return GGML_TYPE_Q5_K; + case CommonMulMatWeightFormat::Q6K: + case CommonMulMatWeightFormat::Q6KRow64: + return GGML_TYPE_Q6_K; + case CommonMulMatWeightFormat::Q4_0: + return GGML_TYPE_Q4_0; + case CommonMulMatWeightFormat::Q4_1: + return GGML_TYPE_Q4_1; + case CommonMulMatWeightFormat::Q5_0: + return GGML_TYPE_Q5_0; + case CommonMulMatWeightFormat::Q5_1: + return GGML_TYPE_Q5_1; + case CommonMulMatWeightFormat::IQ3_XXS: + return GGML_TYPE_IQ3_XXS; + case CommonMulMatWeightFormat::IQ1_S: + return GGML_TYPE_IQ1_S; + case CommonMulMatWeightFormat::IQ1_M: + return GGML_TYPE_IQ1_M; + case CommonMulMatWeightFormat::IQ2_XXS: + return GGML_TYPE_IQ2_XXS; + case CommonMulMatWeightFormat::IQ2_XS: + return GGML_TYPE_IQ2_XS; + case CommonMulMatWeightFormat::IQ2_S: + return GGML_TYPE_IQ2_S; + case CommonMulMatWeightFormat::IQ3_S: + return GGML_TYPE_IQ3_S; + case CommonMulMatWeightFormat::IQ4_NL: + return GGML_TYPE_IQ4_NL; + case CommonMulMatWeightFormat::MXFP4: + return GGML_TYPE_MXFP4; + case CommonMulMatWeightFormat::TQ1_0: + return GGML_TYPE_TQ1_0; + case CommonMulMatWeightFormat::TQ2_0: + return GGML_TYPE_TQ2_0; + case CommonMulMatWeightFormat::IQ4_XS: + return GGML_TYPE_IQ4_XS; + case CommonMulMatWeightFormat::Q8_0: + return GGML_TYPE_Q8_0; + case CommonMulMatWeightFormat::Q8_1: + return GGML_TYPE_Q8_1; + case CommonMulMatWeightFormat::F16: + return GGML_TYPE_F16; + case CommonMulMatWeightFormat::BF16: + return GGML_TYPE_BF16; + case CommonMulMatWeightFormat::F32: + return GGML_TYPE_F32; + } + return GGML_TYPE_COUNT; +} + +inline bool common_is_supported_dense_decode_output_size(CommonMulMatWeightFormat format, + int64_t input_size, + int64_t output_size) { + static constexpr uint64_t kMaxDenseDecodeOutputSize = 1048576; + static constexpr uint64_t kAmdgpuAddressableByteRangeSize = uint64_t{ 1 } << 32; + + if (output_size < 1) { + return false; + } + + const ggml_type type = common_mul_mat_format_type(format); + if (type == GGML_TYPE_COUNT) { + return false; + } + + const size_t weight_row_size = ggml_row_size(type, input_size); + if (weight_row_size == 0) { + return false; + } + + const uint64_t addressable_output_size = kAmdgpuAddressableByteRangeSize / static_cast(weight_row_size); + const uint64_t max_output_size = std::min(kMaxDenseDecodeOutputSize, addressable_output_size); + return static_cast(output_size) <= max_output_size; +} + +inline bool common_is_supported_rmsnorm_hidden_size(int64_t hidden_size) { + return hidden_size >= 128 && hidden_size <= 32768 && hidden_size % 128 == 0; +} + +inline bool common_is_supported_prefill_token_count(int64_t token_count) { + return token_count > 1 && token_count <= 2048; +} + +inline bool common_is_supported_decode_token_count(int64_t token_count) { + return token_count == 1; +} + +inline bool common_is_2d(const Value & value) { + return value.ne[0] > 0 && value.ne[1] > 0 && value.ne[2] == 1 && value.ne[3] == 1; +} + +inline std::string common_to_config_value(int64_t value) { + return std::to_string(value); +} + +inline std::string common_to_config_value(float value) { + std::ostringstream out; + out.precision(9); + out << value; + return out.str(); +} + +inline int64_t common_ceil_div(int64_t value, int64_t divisor) { + return (value + divisor - 1) / divisor; +} + +inline bool common_is_weight_shape(const Value & weight, int64_t hidden_size) { + if (weight.ne[0] != hidden_size) { + return false; + } + for (int i = 1; i < GGML_MAX_DIMS; ++i) { + if (weight.ne[i] != 1) { + return false; + } + } + return true; +} + +inline const GraphNode * common_find_single_consumer_with_op(const Graph & graph, ValueId value, ggml_op op) { + if (!graph.has_index()) { + return nullptr; + } + const GraphNode * match = nullptr; + for (const GraphNode * consumer : graph.index().consumers(value)) { + if (consumer == nullptr || consumer->op != op) { + continue; + } + if (match != nullptr) { + return nullptr; + } + match = consumer; + } + return match; +} + +inline const GraphNode * common_find_only_consumer_with_op(const Graph & graph, ValueId value, ggml_op op) { + if (!graph.has_index()) { + return nullptr; + } + const std::vector & consumers = graph.index().consumers(value); + if (consumers.size() != 1 || consumers.front() == nullptr || consumers.front()->op != op) { + return nullptr; + } + return consumers.front(); +} + +inline bool common_has_direct_symmetric_i4_lowrow_consumer(const Graph & graph, const Value & value) { + if (!graph.has_index() || value.type != GGML_TYPE_F32 || !value.contiguous || value.ne[0] < 256 || + value.ne[0] > 32768 || value.ne[0] % 64 != 0 || value.element_count <= 0 || + value.element_count % value.ne[0] != 0) { + return false; + } + + const int64_t token_count = value.element_count / value.ne[0]; + if (token_count < 1 || token_count > 16) { + return false; + } + + for (const GraphNode * consumer : graph.index().consumers(value.id)) { + if (consumer == nullptr || consumer->op != GGML_OP_MUL_MAT || consumer->inputs.size() != 2 || + consumer->inputs[1] != value.id) { + continue; + } + + const Value * weight = common_graph_value(graph, consumer->inputs[0]); + const Value * output = common_graph_value(graph, consumer->output); + if (weight == nullptr || output == nullptr || + (weight->type != GGML_TYPE_Q5_K && weight->type != GGML_TYPE_IQ4_XS) || weight->alias_source.value >= 0 || + !weight->contiguous || !output->contiguous || output->type != GGML_TYPE_F32 || + weight->ne[0] != value.ne[0] || weight->ne[1] != output->ne[0] || output->ne[0] % 64 != 0 || + output->ne[1] != token_count || output->ne[2] != 1 || output->ne[3] != 1) { + continue; + } + return true; + } + return false; +} + +inline bool common_has_symmetric_i4_lowrow_consumer(const Graph & graph, const Value & value) { + if (common_has_direct_symmetric_i4_lowrow_consumer(graph, value)) { + return true; + } + + for (const GraphNode * consumer : graph.index().consumers(value.id)) { + if (consumer == nullptr || !is_layout_alias_node(graph, *consumer) || consumer->inputs.size() != 1 || + consumer->inputs.front() != value.id) { + continue; + } + + const Value * reshaped = common_graph_value(graph, consumer->output); + if (reshaped != nullptr && reshaped->type == GGML_TYPE_F32 && reshaped->contiguous && + reshaped->storage_root == value.storage_root && reshaped->element_count == value.element_count && + reshaped->byte_count == value.byte_count && + common_has_direct_symmetric_i4_lowrow_consumer(graph, *reshaped)) { + return true; + } + } + return false; +} + +inline bool common_binary_node_is_mul(const GraphNode & node) { + if (node.op != GGML_OP_MUL || node.inputs.size() != 2) { + return false; + } + const BinaryParams * binary_params = op_params_as(node.params); + return binary_params == nullptr || binary_params->op == BinaryKind::Mul; +} + +inline bool common_binary_node_is_add(const GraphNode & node) { + if (node.op != GGML_OP_ADD || node.inputs.size() != 2) { + return false; + } + const BinaryParams * binary_params = op_params_as(node.params); + return binary_params == nullptr || binary_params->op == BinaryKind::Add; +} + +inline bool common_is_swiglu_params(const OpParams & params) { + const BinaryParams * binary_params = op_params_as(params); + if (binary_params != nullptr) { + return binary_params->op == BinaryKind::SwiGLU; + } + + const GluParams * glu_params = op_params_as(params); + return glu_params != nullptr && glu_params->op == GGML_GLU_OP_SWIGLU; +} + +inline bool common_fused_binary_kind_from_params(const OpParams & params, BinaryKind & kind) { + const BinaryParams * binary_params = op_params_as(params); + if (binary_params != nullptr && binary_kind_supported(binary_params->op)) { + kind = binary_params->op; + return true; + } + + const GluParams * glu_params = op_params_as(params); + if (glu_params == nullptr) { + return false; + } + + switch (glu_params->op) { + case GGML_GLU_OP_REGLU: + kind = BinaryKind::RegLU; + return true; + case GGML_GLU_OP_SWIGLU: + kind = BinaryKind::SwiGLU; + return true; + case GGML_GLU_OP_GEGLU: + kind = BinaryKind::GeGLU; + return true; + case GGML_GLU_OP_GEGLU_ERF: + kind = BinaryKind::GeGLUErf; + return true; + case GGML_GLU_OP_GEGLU_QUICK: + kind = BinaryKind::GeGLUQuick; + return true; + default: + return false; + } +} + +inline CommonMulMatMatch common_match_mul_mat_any_format(const Graph & graph, + const GraphNode * node, + KernelCatalogRef kernel, + bool decode) { + CommonMulMatMatch match; + if (node == nullptr || node->op != GGML_OP_MUL_MAT || node->inputs.size() != 2) { + return match; + } + + const Value * weight = common_graph_value(graph, node->inputs[0]); + const Value * input = common_graph_value(graph, node->inputs[1]); + const Value * output = common_graph_value(graph, node->output); + if (weight == nullptr || input == nullptr || output == nullptr || !common_is_2d(*weight) || !common_is_2d(*input) || + !common_is_2d(*output) || !weight->contiguous || !input->contiguous || !output->contiguous || + input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32) { + return {}; + } + + CommonMulMatWeightFormat format = CommonMulMatWeightFormat::Q4K; + if (!common_mul_mat_format_for_type(weight->type, format)) { + return {}; + } + + const int64_t input_size = weight->ne[0]; + const int64_t output_size = weight->ne[1]; + const int64_t token_count = input->ne[1]; + const bool token_count_supported = decode ? common_is_supported_decode_token_count(token_count) : + common_is_supported_prefill_token_count(token_count); + if (input->ne[0] != input_size || output->ne[0] != output_size || output->ne[1] != token_count || + !token_count_supported || !common_is_supported_dense_input_size(format, input_size) || + !(decode ? common_is_supported_dense_decode_output_size(format, input_size, output_size) : + common_is_supported_dense_output_size(output_size))) { + return {}; + } + + match.input = input; + match.weight = weight; + match.output = output; + match.kernel = kernel; + match.input_size = input_size; + match.output_size = output_size; + match.token_count = token_count; + match.weight_format = format; + return match; +} + +// Share the conversion used by F16-operand matmuls. +inline bool common_prepare_f16_input(const DispatchMatchContext & context, + const Value & input, + int64_t input_size, + int64_t token_count, + DispatchMatch & match, + DispatchBinding & binding) { + const size_t bytes = static_cast(token_count * input_size) * sizeof(ggml_fp16_t); + const CommandPlanAlternateValue * alternate = + find_alternate_value(context.graph, context.plan, input.id, GGML_TYPE_F16, bytes); + if (alternate != nullptr) { + binding = { alternate->alternate_value, 0, bytes }; + return true; + } + + constexpr KernelCatalogRef kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_copy_f32_f16"); + constexpr const char * name = "common.mul_mat.f16"; + const ValueId activation = context.next_plan_value; + match.transients.push_back({ activation, name, bytes, 256 }); + Dispatch convert; + convert.kernel = make_kernel_specialization(kernel); + convert.kernel.integer_parameters.emplace("element_count", token_count * input_size); + convert.bindings.push_back({ input.id, 0, input.byte_count }); + convert.bindings.push_back({ activation, 0, bytes }); + match.dispatches.push_back(std::move(convert)); + Status status; + if (!match.metadata.append_alternate_value({ input.id, activation, GGML_TYPE_F16, bytes, name }, status)) { + match.status.append(status); + return false; + } + binding = { activation, 0, bytes }; + return true; +} + +// Private K16-major copy; it is not published as an ordinary F16 alternate. +inline bool common_prepare_k16_major_f16_input(const DispatchMatchContext & context, + const Value & input, + int64_t input_size, + int64_t token_count, + DispatchMatch & match, + DispatchBinding & binding) { + constexpr KernelCatalogRef kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_copy_f16_k16_major"); + constexpr const char * name = "common.mul_mat.k16_major_f16"; + const size_t bytes = static_cast(token_count * input_size) * sizeof(ggml_fp16_t); + const CommandPlanGeneratedResource * generated = + context.plan.metadata.find_generated_resource(input.id, GeneratedResourceRole::F16K16Major); + if (generated != nullptr && generated->byte_count == bytes) { + binding = { generated->generated_value, 0, bytes }; + return true; + } + const CommandPlanAlternateValue * alternate = + find_alternate_value(context.graph, context.plan, input.id, GGML_TYPE_F16, bytes); + const ValueId packed(context.next_plan_value.value + static_cast(match.transients.size())); + match.transients.push_back({ packed, name, bytes, 256 }); + Dispatch copy; + copy.kernel = make_kernel_specialization(kernel); + copy.kernel.integer_parameters.emplace("token_count", token_count); + copy.kernel.compile_parameters.emplace("ggml.copy_f16_k16_major.input_size", common_to_config_value(input_size)); + copy.kernel.compile_parameters.emplace("ggml.copy_f16_k16_major.token_count", common_to_config_value(token_count)); + copy.kernel.compile_parameters.emplace("ggml.copy_f16_k16_major.input_is_f16", alternate != nullptr ? "1" : "0"); + copy.bindings.push_back(alternate != nullptr ? DispatchBinding{ alternate->alternate_value, 0, bytes } : + DispatchBinding{ input.id, 0, input.byte_count }); + copy.bindings.push_back({ packed, 0, bytes }); + match.dispatches.push_back(std::move(copy)); + Status status; + if (!match.metadata.append_generated_resource( + { input.id, GeneratedResourceRole::F16K16Major, packed, bytes, {} }, status)) { + match.status.append(status); + return false; + } + binding = { packed, 0, bytes }; + return true; +} + +// Share one activation producer between ordinary and fused matmuls. +inline bool common_prepare_q8_1_x4_input(const DispatchMatchContext & context, + const Value & input, + int64_t input_size, + int64_t token_count, + DispatchMatch & match, + DispatchBinding & binding, + CommonQ8ActivationPolicy policy) { + const size_t bytes = static_cast(token_count) * ggml_row_size(GGML_TYPE_Q8_1, input_size); + const CommandPlanAlternateValue * alternate = + find_alternate_value(context.graph, context.plan, input.id, GGML_TYPE_Q8_1, bytes); + if (alternate != nullptr) { + binding = { alternate->alternate_value, 0, bytes }; + return true; + } + if (policy != CommonQ8ActivationPolicy::AllowStandaloneQuantize) { + return false; + } + + constexpr KernelCatalogRef kernel = GGML_HRX_KERNEL_REF("qwen3_moe", "ggml_quantize_q8_1_x4_f32"); + constexpr const char * name = "common.mul_mat.q8_1_x4"; + const ValueId activation = context.next_plan_value; + match.transients.push_back({ activation, name, bytes, 256 }); + Dispatch quantize; + quantize.kernel = make_kernel_specialization(kernel); + quantize.kernel.integer_parameters.emplace("token_count", token_count); + quantize.kernel.integer_parameters.emplace("input_size", input_size); + quantize.kernel.compile_parameters.emplace("ggml.quantize_q8_1_x4.group_capacity", + common_to_config_value(token_count * input_size / 128)); + quantize.bindings.push_back({ input.id, 0, input.byte_count }); + quantize.bindings.push_back({ activation, 0, bytes }); + match.dispatches.push_back(std::move(quantize)); + + Status status; + if (!match.metadata.append_alternate_value({ input.id, activation, GGML_TYPE_Q8_1, bytes, name }, status)) { + match.status.append(status); + return false; + } + binding = { activation, 0, bytes }; + return true; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id-common.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id-common.h new file mode 100644 index 000000000000..467fda4eea73 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id-common.h @@ -0,0 +1,364 @@ +#pragma once + +#include "../dispatch-registry.h" +#include "dispatch-mul-mat-weight-format.h" +#include "ggml.h" +#include "graph/graph-matcher.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { + +struct CommonMulMatIdMatch { + const Value * input = nullptr; + const Value * weight = nullptr; + const Value * output = nullptr; + const Value * route_ids = nullptr; + CommandPlanMoeRoutingBundle routing_bundle; + bool has_routing_bundle = false; + int64_t input_size = 0; + int64_t output_size = 0; + int64_t token_count = 0; + int64_t route_count = 0; + int64_t route_stride = 0; // I32 elements between tokens' route ids + int64_t input_route_count = 0; + int64_t expert_count = 0; + CommonMulMatWeightFormat weight_format = CommonMulMatWeightFormat::Q4K; + + bool matched() const { return input != nullptr && weight != nullptr && output != nullptr && route_ids != nullptr; } +}; + +inline const Value * common_mul_mat_id_graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +inline bool common_mul_mat_id_is_shape(const Value & value, int64_t ne0, int64_t ne1, int64_t ne2, int64_t ne3) { + return value.ne[0] == ne0 && value.ne[1] == ne1 && value.ne[2] == ne2 && value.ne[3] == ne3; +} + +inline bool common_mul_mat_id_same_shape(const Value & lhs, const Value & rhs) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (lhs.ne[i] != rhs.ne[i]) { + return false; + } + } + return true; +} + +inline bool common_mul_mat_id_supported_dense_input_size(CommonMulMatWeightFormat format, int64_t input_size) { + return common_mul_mat_supported_dense_input_size(format, input_size); +} + +inline bool common_mul_mat_id_supported_dense_output_size(int64_t output_size) { + return output_size >= 1 && output_size <= 262144; +} + +inline bool common_mul_mat_id_supported_rmsnorm_hidden_size(int64_t hidden_size) { + return hidden_size >= 128 && hidden_size <= 32768 && hidden_size % 128 == 0; +} + +inline bool common_mul_mat_id_supported_token_count(int64_t token_count) { + return token_count >= 1 && token_count <= 2048; +} + +inline bool common_mul_mat_id_supported_route_count(int64_t route_count) { + return route_count >= 1 && route_count <= 32; +} + +inline bool common_mul_mat_id_supported_expert_count(int64_t expert_count) { + return expert_count >= 1 && expert_count <= 512; +} + +inline std::string common_mul_mat_id_to_config_value(int64_t value) { + return std::to_string(value); +} + +inline std::string common_mul_mat_id_to_config_value(float value) { + std::ostringstream out; + out.precision(9); + out << value; + return out.str(); +} + +inline int64_t common_mul_mat_id_ceil_div(int64_t value, int64_t divisor) { + return (value + divisor - 1) / divisor; +} + +inline int64_t common_mul_mat_id_partition_descriptor_capacity(int64_t token_count, + int64_t route_count, + int64_t expert_count) { + return common_mul_mat_id_ceil_div(token_count * route_count, 32) + expert_count; +} + +inline const GraphNode * common_mul_mat_id_find_single_consumer_with_op(const Graph & graph, + ValueId value, + ggml_op op) { + if (!graph.has_index()) { + return nullptr; + } + const GraphNode * match = nullptr; + for (const GraphNode * consumer : graph.index().consumers(value)) { + if (consumer == nullptr || consumer->op != op) { + continue; + } + if (match != nullptr) { + return nullptr; + } + match = consumer; + } + return match; +} + +inline const GraphNode * common_mul_mat_id_find_only_consumer_with_op(const Graph & graph, ValueId value, ggml_op op) { + if (!graph.has_index()) { + return nullptr; + } + const std::vector & consumers = graph.index().consumers(value); + if (consumers.size() != 1 || consumers.front() == nullptr || consumers.front()->op != op) { + return nullptr; + } + return consumers.front(); +} + +inline bool common_mul_mat_id_binary_node_is_mul(const GraphNode & node) { + if (node.op != GGML_OP_MUL || node.inputs.size() != 2) { + return false; + } + const BinaryParams * binary_params = op_params_as(node.params); + return binary_params == nullptr || binary_params->op == BinaryKind::Mul; +} + +inline bool common_mul_mat_id_binary_node_is_add(const GraphNode & node) { + if (node.op != GGML_OP_ADD || node.inputs.size() != 2) { + return false; + } + const BinaryParams * binary_params = op_params_as(node.params); + return binary_params == nullptr || binary_params->op == BinaryKind::Add; +} + +inline bool common_mul_mat_id_is_swiglu_params(const OpParams & params) { + const BinaryParams * binary_params = op_params_as(params); + if (binary_params != nullptr) { + return binary_params->op == BinaryKind::SwiGLU; + } + + const GluParams * glu_params = op_params_as(params); + return glu_params != nullptr && glu_params->op == GGML_GLU_OP_SWIGLU; +} + +inline bool common_mul_mat_id_is_weight_shape(const Value & weight, int64_t hidden_size) { + if (weight.ne[0] != hidden_size) { + return false; + } + for (int i = 1; i < GGML_MAX_DIMS; ++i) { + if (weight.ne[i] != 1) { + return false; + } + } + return true; +} + +inline size_t common_mul_mat_id_expert_table_size(int64_t token_count, int64_t expert_count) { + return static_cast(expert_count + expert_count * token_count) * sizeof(int32_t); +} + +inline size_t common_mul_mat_id_partition_table_size(int64_t token_count, int64_t route_count, int64_t expert_count) { + const int64_t assignment_count = token_count * route_count; + const int64_t assignment_partition_count = (assignment_count + 31) / 32; + return static_cast(1 + assignment_partition_count + expert_count) * sizeof(int32_t); +} + +inline bool common_mul_mat_id_bundle_matches(const CommandPlanMoeRoutingBundle & bundle, + ValueId route_ids, + int64_t token_count, + int64_t route_count, + int64_t expert_count) { + return bundle.route_ids == route_ids && bundle.expert_table.value >= 0 && bundle.partition_table.value >= 0 && + bundle.expert_table_byte_count == common_mul_mat_id_expert_table_size(token_count, expert_count) && + bundle.partition_table_byte_count == + common_mul_mat_id_partition_table_size(token_count, route_count, expert_count) && + bundle.token_count == token_count && bundle.route_count == route_count && + bundle.expert_count == expert_count && bundle.route_stride >= route_count; +} + +inline CommonMulMatIdMatch common_match_mul_mat_id_any_format(const Graph & graph, + const GraphNode * node, + const CommandPlan & plan) { + CommonMulMatIdMatch match; + if (node == nullptr || node->op != GGML_OP_MUL_MAT_ID || node->inputs.size() != 3 || !graph.has_index()) { + return match; + } + + const Value * weight = common_mul_mat_id_graph_value(graph, node->inputs[0]); + const Value * input = common_mul_mat_id_graph_value(graph, node->inputs[1]); + const Value * route_ids = common_mul_mat_id_graph_value(graph, node->inputs[2]); + const Value * output = common_mul_mat_id_graph_value(graph, node->output); + if (weight == nullptr || input == nullptr || route_ids == nullptr || output == nullptr || !weight->contiguous || + !input->contiguous || !output->contiguous || input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || + route_ids->type != GGML_TYPE_I32) { + return {}; + } + + CommonMulMatWeightFormat format = CommonMulMatWeightFormat::Q4K; + if (!common_mul_mat_format_for_type(weight->type, format)) { + return {}; + } + + const int64_t input_size = weight->ne[0]; + const int64_t output_size = weight->ne[1]; + const int64_t expert_count = weight->ne[2]; + const int64_t input_route_count = input->ne[1]; + const int64_t token_count = input->ne[2]; + const int64_t route_count = route_ids->ne[0]; + if (!common_mul_mat_id_is_shape(*weight, input_size, output_size, expert_count, 1) || + !common_mul_mat_id_is_shape(*input, input_size, input_route_count, token_count, 1) || + !common_mul_mat_id_is_shape(*route_ids, route_count, token_count, 1, 1) || + !common_mul_mat_id_is_shape(*output, output_size, route_count, token_count, 1) || + !common_mul_mat_id_supported_dense_input_size(format, input_size) || + !common_mul_mat_id_supported_dense_output_size(output_size) || + !common_mul_mat_id_supported_token_count(token_count) || + !common_mul_mat_id_supported_route_count(route_count) || + !common_mul_mat_id_supported_expert_count(expert_count) || input_route_count <= 0 || + input_route_count > route_count || route_count % input_route_count != 0) { + return {}; + } + // llama.cpp's route ids are a [n_expert_used, n_tokens] view of the [n_expert, n_tokens] argsort, + // so rows are n_expert apart; the routing tables must read them at that stride + if (route_ids->nb[0] != sizeof(int32_t) || route_ids->nb[1] % sizeof(int32_t) != 0) { + return {}; + } + const int64_t route_stride = static_cast(route_ids->nb[1] / sizeof(int32_t)); + if (token_count > 1 && (route_stride < route_count || route_stride > expert_count)) { + return {}; + } + + match.input = input; + match.weight = weight; + match.output = output; + match.route_ids = route_ids; + match.input_size = input_size; + match.output_size = output_size; + match.token_count = token_count; + match.route_count = route_count; + match.route_stride = token_count > 1 ? route_stride : route_count; + match.input_route_count = input_route_count; + match.expert_count = expert_count; + match.weight_format = format; + + const CommandPlanMoeRoutingBundle * routing_bundle = plan.metadata.find_moe_routing_bundle(route_ids->id); + if (routing_bundle != nullptr && + common_mul_mat_id_bundle_matches(*routing_bundle, route_ids->id, token_count, route_count, expert_count)) { + match.routing_bundle = *routing_bundle; + match.has_routing_bundle = true; + } + return match; +} + +inline void common_mul_mat_id_add_moe_routing_compile_parameters(Dispatch & dispatch, + const CommonMulMatIdMatch & match) { + dispatch.kernel.compile_parameters.emplace("ggml.workload.token_capacity", + common_mul_mat_id_to_config_value(match.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.moe_routing.route_count", + common_mul_mat_id_to_config_value(match.route_count)); + dispatch.kernel.compile_parameters.emplace("ggml.moe_routing.expert_count", + common_mul_mat_id_to_config_value(match.expert_count)); + dispatch.kernel.compile_parameters.emplace("ggml.moe_routing.descriptor_expert_mask", "511"); + dispatch.kernel.compile_parameters.emplace("ggml.moe_routing.descriptor_partition_shift", "9"); + dispatch.kernel.compile_parameters.emplace("ggml.moe_routing.descriptor_row_count_shift", "15"); + dispatch.kernel.compile_parameters.emplace("ggml.moe_routing.partition_workgroup_size", "512"); +} + +inline bool common_mul_mat_id_ensure_moe_routing_bundle(const DispatchMatchContext & context, + const CommonMulMatIdMatch & match, + DispatchMatch & dispatch_match, + CommandPlanMoeRoutingBundle & routing_bundle) { + static constexpr KernelCatalogRef kMoeBuildExpertTableKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_moe_build_expert_table"); + static constexpr KernelCatalogRef kMoeBuildExpertPartitionTableKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_moe_build_expert_partition_table"); + + if (match.has_routing_bundle) { + routing_bundle = match.routing_bundle; + return true; + } + + const ValueId expert_table_value(context.next_plan_value.value + + static_cast(dispatch_match.transients.size())); + const ValueId partition_table_value(expert_table_value.value + 1); + const size_t expert_table_bytes = common_mul_mat_id_expert_table_size(match.token_count, match.expert_count); + const size_t partition_table_bytes = + common_mul_mat_id_partition_table_size(match.token_count, match.route_count, match.expert_count); + const int64_t route_stride = match.route_stride; + const size_t route_ids_length = + static_cast((match.token_count - 1) * route_stride + match.route_count) * sizeof(int32_t); + + dispatch_match.transients.push_back( + { expert_table_value, "common.moe_routing.expert_table", expert_table_bytes, 256 }); + dispatch_match.transients.push_back( + { partition_table_value, "common.moe_routing.partition_table", partition_table_bytes, 256 }); + + const CommandPlanResourceMetadata routing_metadata = make_command_plan_resource_metadata(MoeRoutingResourceMetadata{ + match.token_count, + match.route_count, + route_stride, + match.expert_count, + }); + routing_bundle = { + match.route_ids->id, ValueId(), expert_table_value, partition_table_value, expert_table_bytes, + partition_table_bytes, match.token_count, match.route_count, route_stride, match.expert_count, + }; + + Status metadata_status; + if (!dispatch_match.metadata.append_generated_resource( + { + match.route_ids->id, + GeneratedResourceRole::MoeExpertTable, + expert_table_value, + expert_table_bytes, + routing_metadata, + }, + metadata_status) || + !dispatch_match.metadata.append_generated_resource( + { + match.route_ids->id, + GeneratedResourceRole::MoePartitionTable, + partition_table_value, + partition_table_bytes, + routing_metadata, + }, + metadata_status) || + !dispatch_match.metadata.append_moe_routing_bundle(routing_bundle, metadata_status)) { + dispatch_match.status.append(metadata_status); + return false; + } + + Dispatch expert_table_dispatch; + expert_table_dispatch.kernel = make_kernel_specialization(kMoeBuildExpertTableKernel); + expert_table_dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + expert_table_dispatch.kernel.integer_parameters.emplace("route_count", match.route_count); + expert_table_dispatch.kernel.integer_parameters.emplace("route_stride", route_stride); + expert_table_dispatch.kernel.integer_parameters.emplace("expert_count", match.expert_count); + common_mul_mat_id_add_moe_routing_compile_parameters(expert_table_dispatch, match); + expert_table_dispatch.bindings.push_back({ match.route_ids->id, 0, route_ids_length }); + expert_table_dispatch.bindings.push_back({ expert_table_value, 0, expert_table_bytes }); + dispatch_match.dispatches.push_back(std::move(expert_table_dispatch)); + + Dispatch partition_table_dispatch; + partition_table_dispatch.kernel = make_kernel_specialization(kMoeBuildExpertPartitionTableKernel); + partition_table_dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + partition_table_dispatch.kernel.integer_parameters.emplace("route_count", match.route_count); + partition_table_dispatch.kernel.integer_parameters.emplace("expert_count", match.expert_count); + common_mul_mat_id_add_moe_routing_compile_parameters(partition_table_dispatch, match); + partition_table_dispatch.bindings.push_back({ expert_table_value, 0, expert_table_bytes }); + partition_table_dispatch.bindings.push_back({ partition_table_value, 0, partition_table_bytes }); + dispatch_match.dispatches.push_back(std::move(partition_table_dispatch)); + return true; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id-decode.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id-decode.cpp new file mode 100644 index 000000000000..868cd3ad920d --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id-decode.cpp @@ -0,0 +1,84 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// MUL_MAT_ID for a few tokens (decode) as a GEMV over the selected experts' rows +// (ops/mul_mat_id_decode_f32.loom), ahead of the WMMA mul_mat_id kernel, which is built for +// batches and reads a single token's expert at ~45 GB/s. Larger batches still take the WMMA path. + +#include "dispatch-mul-mat-id-decode.h" + +#include "dispatch-mul-mat-id-common.h" + +namespace ggml::hrx { + +namespace { + +static constexpr KernelCatalogRef kMulMatIdDecodeKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_id_decode_f32_wave64"); + +// The GEMV pays off while token * slot rows are few; past this the WMMA kernel's reuse wins. +constexpr int64_t kMaximumDecodeTokens = 4; +constexpr int64_t kMaximumDecodeRows = 64; + +bool match_mul_mat_id_decode_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const CommonMulMatIdMatch match = common_match_mul_mat_id_any_format(context.graph, context.root_node, context.plan); + if (!match.matched() || match.token_count < 1 || match.token_count > kMaximumDecodeTokens || + match.token_count * match.route_count > kMaximumDecodeRows || match.route_stride < match.route_count || + match.input_size < 256 || match.input_size > 32768 || match.input_size % 32 != 0 || + match.output_size < 1 || match.expert_count < 1 || + (match.input_route_count != 1 && match.input_route_count != match.route_count)) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kMulMatIdDecodeKernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.integer_parameters.emplace("slot_count", match.route_count); + dispatch.kernel.integer_parameters.emplace("input_rows", match.input_route_count); + dispatch.kernel.integer_parameters.emplace("input_size", match.input_size); + dispatch.kernel.integer_parameters.emplace("output_size", match.output_size); + dispatch.kernel.integer_parameters.emplace("expert_count", match.expert_count); + dispatch.kernel.integer_parameters.emplace("route_stride", match.route_stride); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_decode.row_capacity", + common_mul_mat_id_to_config_value(match.token_count * match.route_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_decode.output_capacity", + common_mul_mat_id_to_config_value(match.output_size)); + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat_id_decode.weight_format", + common_mul_mat_id_to_config_value(common_mul_mat_format_config_value(match.weight_format))); + dispatch.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count }); + dispatch.bindings.push_back({ match.route_ids->id, 0, match.route_ids->byte_count }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_mul_mat_id_decode_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "common.mul_mat_id.decode_f32_wave64", + GGML_OP_MUL_MAT_ID, + DispatchMatchKind::SingleOp, + 150, + DispatchSource::Common, + match_mul_mat_id_decode_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id-decode.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id-decode.h new file mode 100644 index 000000000000..f5d1f71f1624 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id-decode.h @@ -0,0 +1,24 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +namespace ggml::hrx { + +void register_mul_mat_id_decode_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id.cpp new file mode 100644 index 000000000000..1bee27ec8a95 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id.cpp @@ -0,0 +1,310 @@ +#include "dispatch-mul-mat-id.h" + +#include "dispatch-mul-mat-id-common.h" +#include "ggml.h" +#include "graph/graph-matcher.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kMulMatIdF32F32WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_id_f32_f32_wmma"); +static constexpr KernelCatalogRef kMulMatIdPostOpsF32F32WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_id_postops_f32_f32_wmma"); +static constexpr KernelCatalogRef kMulMatIdPostOpsNextRmsNormF32F32WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_id_postops_next_rmsnorm_f32_f32_wmma"); + +struct MulMatIdPostOpsMatch { + CommonMulMatIdMatch root; + const Value * bias = nullptr; + const Value * residual_input = nullptr; + const Value * residual_output = nullptr; + const Value * norm_weight = nullptr; + const Value * normalized_output = nullptr; + std::vector add_nodes; + const GraphNode * rms_node = nullptr; + const GraphNode * mul_node = nullptr; + KernelCatalogRef kernel = {}; + float epsilon = 0.0f; + bool has_bias = false; + bool has_residual = false; + bool has_rmsnorm = false; + + bool matched() const { + return root.matched() && residual_output != nullptr && kernel.id != kUncatalogedKernelId && + (has_bias || has_residual); + } +}; + +static bool is_bias_shape(const Value & value, int64_t output_size) { + if (value.type != GGML_TYPE_F32 || !value.contiguous || value.ne[0] != output_size) { + return false; + } + for (int i = 1; i < GGML_MAX_DIMS; ++i) { + if (value.ne[i] != 1) { + return false; + } + } + return true; +} + +static MulMatIdPostOpsMatch match_mul_mat_id_postops(const DispatchMatchContext & context) { + MulMatIdPostOpsMatch match; + const CommonMulMatIdMatch root = common_match_mul_mat_id_any_format(context.graph, context.root_node, context.plan); + if (!root.matched() || !context.graph.has_index()) { + return match; + } + + const Value * current = root.output; + for (int add_index = 0; add_index < 2; ++add_index) { + const GraphNode * add_node = + common_mul_mat_id_find_only_consumer_with_op(context.graph, current->id, GGML_OP_ADD); + if (add_node == nullptr || !common_mul_mat_id_binary_node_is_add(*add_node)) { + break; + } + + const bool current_is_lhs = add_node->inputs[0] == current->id; + const bool current_is_rhs = add_node->inputs[1] == current->id; + if (!current_is_lhs && !current_is_rhs) { + return {}; + } + + const ValueId other_id = current_is_lhs ? add_node->inputs[1] : add_node->inputs[0]; + const Value * other = common_mul_mat_id_graph_value(context.graph, other_id); + const Value * output = common_mul_mat_id_graph_value(context.graph, add_node->output); + if (other == nullptr || output == nullptr || output->type != GGML_TYPE_F32 || !output->contiguous || + !common_mul_mat_id_same_shape(*root.output, *output)) { + return {}; + } + + if (!match.has_bias && is_bias_shape(*other, root.output_size)) { + match.bias = other; + match.add_nodes.push_back(add_node); + match.has_bias = true; + current = output; + continue; + } + + if (!match.has_residual && other->type == GGML_TYPE_F32 && other->contiguous && + common_mul_mat_id_same_shape(*root.output, *other)) { + match.residual_input = other; + match.add_nodes.push_back(add_node); + match.has_residual = true; + current = output; + continue; + } + + break; + } + + if (!match.has_bias && !match.has_residual) { + return {}; + } + + match.residual_output = current; + + const bool may_match_next_rmsnorm = + match.has_residual && common_mul_mat_id_supported_rmsnorm_hidden_size(root.output_size); + if (may_match_next_rmsnorm) { + const GraphNode * rms_node = + common_mul_mat_id_find_single_consumer_with_op(context.graph, current->id, GGML_OP_RMS_NORM); + if (rms_node != nullptr && rms_node->inputs.size() == 1) { + const RmsNormParams * rms_params = op_params_as(rms_node->params); + const Value * rms_output = common_mul_mat_id_graph_value(context.graph, rms_node->output); + if (rms_params != nullptr && std::isfinite(rms_params->eps) && rms_params->eps > 0.0f && + rms_output != nullptr && rms_output->type == GGML_TYPE_F32 && rms_output->contiguous && + common_mul_mat_id_same_shape(*current, *rms_output)) { + const GraphNode * mul_node = + common_mul_mat_id_find_single_consumer_with_op(context.graph, rms_node->output, GGML_OP_MUL); + if (mul_node != nullptr && common_mul_mat_id_binary_node_is_mul(*mul_node)) { + const bool rms_is_lhs = mul_node->inputs[0] == rms_node->output; + const bool rms_is_rhs = mul_node->inputs[1] == rms_node->output; + if (rms_is_lhs || rms_is_rhs) { + const ValueId norm_weight_id = rms_is_lhs ? mul_node->inputs[1] : mul_node->inputs[0]; + const Value * norm_weight = common_mul_mat_id_graph_value(context.graph, norm_weight_id); + const Value * normalized_output = + common_mul_mat_id_graph_value(context.graph, mul_node->output); + if (norm_weight != nullptr && normalized_output != nullptr && + norm_weight->type == GGML_TYPE_F32 && normalized_output->type == GGML_TYPE_F32 && + norm_weight->contiguous && normalized_output->contiguous && + common_mul_mat_id_is_weight_shape(*norm_weight, root.output_size) && + common_mul_mat_id_same_shape(*current, *normalized_output)) { + match.norm_weight = norm_weight; + match.normalized_output = normalized_output; + match.rms_node = rms_node; + match.mul_node = mul_node; + match.epsilon = rms_params->eps; + match.has_rmsnorm = true; + } + } + } + } + } + } + + match.root = root; + match.kernel = match.has_rmsnorm ? kMulMatIdPostOpsNextRmsNormF32F32WmmaKernel : kMulMatIdPostOpsF32F32WmmaKernel; + return match; +} + +} // namespace + +static bool match_mul_mat_id_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const CommonMulMatIdMatch match = + common_match_mul_mat_id_any_format(context.graph, context.root_node, context.plan); + if (!match.matched()) { + return false; + } + + CommandPlanMoeRoutingBundle routing_bundle; + if (!common_mul_mat_id_ensure_moe_routing_bundle(context, match, dispatch_match, routing_bundle)) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kMulMatIdF32F32WmmaKernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.workload.token_capacity", + common_mul_mat_id_to_config_value(match.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id.input_size", + common_mul_mat_id_to_config_value(match.input_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id.output_size", + common_mul_mat_id_to_config_value(match.output_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id.expert_count", + common_mul_mat_id_to_config_value(match.expert_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id.route_count", + common_mul_mat_id_to_config_value(match.route_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id.input_route_count", + common_mul_mat_id_to_config_value(match.input_route_count)); + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat_id.weight_format", + common_mul_mat_id_to_config_value(common_mul_mat_format_config_value(match.weight_format))); + dispatch.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + dispatch.bindings.push_back({ routing_bundle.expert_table, 0, routing_bundle.expert_table_byte_count }); + dispatch.bindings.push_back({ routing_bundle.partition_table, 0, routing_bundle.partition_table_byte_count }); + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_mul_mat_id_postops_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const MulMatIdPostOpsMatch match = match_mul_mat_id_postops(context); + if (!match.matched()) { + return false; + } + + CommandPlanMoeRoutingBundle routing_bundle; + if (!common_mul_mat_id_ensure_moe_routing_bundle(context, match.root, dispatch_match, routing_bundle)) { + return false; + } + + const int64_t completion_counter_count = + match.has_rmsnorm ? common_mul_mat_id_partition_descriptor_capacity( + match.root.token_count, match.root.route_count, match.root.expert_count) : + 0; + const ValueId completion_counters(context.next_plan_value.value + + static_cast(dispatch_match.transients.size())); + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(match.kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.root.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.workload.token_capacity", + common_mul_mat_id_to_config_value(match.root.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_postops.input_size", + common_mul_mat_id_to_config_value(match.root.input_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_postops.output_size", + common_mul_mat_id_to_config_value(match.root.output_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_postops.expert_count", + common_mul_mat_id_to_config_value(match.root.expert_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_postops.route_count", + common_mul_mat_id_to_config_value(match.root.route_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_postops.input_route_count", + common_mul_mat_id_to_config_value(match.root.input_route_count)); + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat_id_postops.weight_format", + common_mul_mat_id_to_config_value(common_mul_mat_format_config_value(match.root.weight_format))); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_postops.has_bias", match.has_bias ? "1" : "0"); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_postops.has_residual", match.has_residual ? "1" : "0"); + if (match.has_rmsnorm) { + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_postops.rms_epsilon", + common_mul_mat_id_to_config_value(match.epsilon)); + } + + dispatch.bindings.push_back({ match.root.input->id, 0, match.root.input->byte_count }); + dispatch.bindings.push_back({ routing_bundle.expert_table, 0, routing_bundle.expert_table_byte_count }); + dispatch.bindings.push_back({ routing_bundle.partition_table, 0, routing_bundle.partition_table_byte_count }); + dispatch.bindings.push_back({ match.root.weight->id, 0, match.root.weight->byte_count }); + if (match.has_bias) { + dispatch.bindings.push_back({ match.bias->id, 0, match.bias->byte_count }); + } else { + dispatch.bindings.push_back({ match.residual_output->id, 0, match.residual_output->byte_count }); + } + if (match.has_residual) { + dispatch.bindings.push_back({ match.residual_input->id, 0, match.residual_input->byte_count }); + } else { + dispatch.bindings.push_back({ match.residual_output->id, 0, match.residual_output->byte_count }); + } + dispatch.bindings.push_back({ match.residual_output->id, 0, match.residual_output->byte_count }); + if (match.has_rmsnorm) { + dispatch.bindings.push_back({ match.norm_weight->id, 0, match.norm_weight->byte_count }); + dispatch.bindings.push_back({ match.normalized_output->id, 0, match.normalized_output->byte_count }); + dispatch.bindings.push_back( + { completion_counters, 0, static_cast(completion_counter_count) * sizeof(int32_t) }); + } + + if (!append_covered_node_index_once(context.graph, context.covered_nodes, context.root_node, + dispatch_match.covered_nodes)) { + return false; + } + for (const GraphNode * add_node : match.add_nodes) { + if (!append_covered_node_index_once(context.graph, context.covered_nodes, add_node, + dispatch_match.covered_nodes)) { + return false; + } + } + if (match.has_rmsnorm && (!append_covered_node_index_once(context.graph, context.covered_nodes, match.rms_node, + dispatch_match.covered_nodes) || + !append_covered_node_index_once(context.graph, context.covered_nodes, match.mul_node, + dispatch_match.covered_nodes))) { + return false; + } + + if (match.has_rmsnorm) { + dispatch_match.completion_counter_requests.push_back({ completion_counters, + "common.mul_mat_id_postops.completion_counters", + static_cast(completion_counter_count) }); + } + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +void register_mul_mat_id_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "common.mul_mat_id_postops.f32_f32_wmma", + GGML_OP_MUL_MAT_ID, + DispatchMatchKind::Fused, + 180, + DispatchSource::Common, + match_mul_mat_id_postops_dispatch, + }); + registry.add({ + "common.mul_mat_id.f32_f32_wmma", + GGML_OP_MUL_MAT_ID, + DispatchMatchKind::SingleOp, + 100, + DispatchSource::Common, + match_mul_mat_id_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id.h new file mode 100644 index 000000000000..931158c11c48 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_mul_mat_id_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-iq3-xxs.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-iq3-xxs.cpp new file mode 100644 index 000000000000..a1788e35b08a --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-iq3-xxs.cpp @@ -0,0 +1,194 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "dispatch-mul-mat-iq3-xxs.h" + +#include "dispatch-mul-mat-common.h" +#include "ggml.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kMulMatVecIq3XxsF32Kernel = + GGML_HRX_KERNEL_REF("hrx", "ggml_mul_mat_vec_iq3xxs_f32"); + +// IQ3_XXS codebook: 256 x u32 grid (LE) | 128 sign bytes | 8 sign-mask bytes +static const uint8_t kIq3XxsTables[1168] = { + 0x04, 0x04, 0x04, 0x04, 0x14, 0x04, 0x04, 0x04, 0x24, 0x04, 0x04, 0x04, + 0x0c, 0x0c, 0x04, 0x04, 0x1c, 0x0c, 0x04, 0x04, 0x3e, 0x0c, 0x04, 0x04, + 0x04, 0x14, 0x04, 0x04, 0x14, 0x14, 0x04, 0x04, 0x0c, 0x1c, 0x04, 0x04, + 0x14, 0x24, 0x04, 0x04, 0x1c, 0x3e, 0x04, 0x04, 0x2c, 0x3e, 0x04, 0x04, + 0x0c, 0x04, 0x0c, 0x04, 0x1c, 0x04, 0x0c, 0x04, 0x04, 0x0c, 0x0c, 0x04, + 0x14, 0x0c, 0x0c, 0x04, 0x0c, 0x14, 0x0c, 0x04, 0x2c, 0x14, 0x0c, 0x04, + 0x04, 0x1c, 0x0c, 0x04, 0x14, 0x1c, 0x0c, 0x04, 0x0c, 0x24, 0x0c, 0x04, + 0x24, 0x2c, 0x0c, 0x04, 0x04, 0x3e, 0x0c, 0x04, 0x04, 0x04, 0x14, 0x04, + 0x14, 0x04, 0x14, 0x04, 0x24, 0x04, 0x14, 0x04, 0x0c, 0x0c, 0x14, 0x04, + 0x04, 0x14, 0x14, 0x04, 0x14, 0x14, 0x14, 0x04, 0x0c, 0x1c, 0x14, 0x04, + 0x1c, 0x1c, 0x14, 0x04, 0x3e, 0x1c, 0x14, 0x04, 0x0c, 0x2c, 0x14, 0x04, + 0x3e, 0x2c, 0x14, 0x04, 0x2c, 0x3e, 0x14, 0x04, 0x0c, 0x04, 0x1c, 0x04, + 0x3e, 0x04, 0x1c, 0x04, 0x04, 0x0c, 0x1c, 0x04, 0x14, 0x0c, 0x1c, 0x04, + 0x2c, 0x14, 0x1c, 0x04, 0x04, 0x3e, 0x1c, 0x04, 0x1c, 0x0c, 0x24, 0x04, + 0x3e, 0x1c, 0x24, 0x04, 0x24, 0x24, 0x24, 0x04, 0x3e, 0x2c, 0x24, 0x04, + 0x1c, 0x3e, 0x24, 0x04, 0x2c, 0x3e, 0x24, 0x04, 0x0c, 0x04, 0x2c, 0x04, + 0x3e, 0x04, 0x2c, 0x04, 0x14, 0x1c, 0x2c, 0x04, 0x14, 0x2c, 0x2c, 0x04, + 0x2c, 0x1c, 0x34, 0x04, 0x24, 0x34, 0x34, 0x04, 0x04, 0x0c, 0x3e, 0x04, + 0x24, 0x0c, 0x3e, 0x04, 0x34, 0x0c, 0x3e, 0x04, 0x1c, 0x24, 0x3e, 0x04, + 0x0c, 0x34, 0x3e, 0x04, 0x0c, 0x04, 0x04, 0x0c, 0x1c, 0x04, 0x04, 0x0c, + 0x04, 0x0c, 0x04, 0x0c, 0x14, 0x0c, 0x04, 0x0c, 0x0c, 0x14, 0x04, 0x0c, + 0x1c, 0x14, 0x04, 0x0c, 0x04, 0x1c, 0x04, 0x0c, 0x14, 0x1c, 0x04, 0x0c, + 0x24, 0x1c, 0x04, 0x0c, 0x3e, 0x24, 0x04, 0x0c, 0x04, 0x2c, 0x04, 0x0c, + 0x04, 0x04, 0x0c, 0x0c, 0x14, 0x04, 0x0c, 0x0c, 0x0c, 0x0c, 0x0c, 0x0c, + 0x04, 0x14, 0x0c, 0x0c, 0x14, 0x14, 0x0c, 0x0c, 0x0c, 0x04, 0x14, 0x0c, + 0x1c, 0x04, 0x14, 0x0c, 0x04, 0x0c, 0x14, 0x0c, 0x14, 0x0c, 0x14, 0x0c, + 0x0c, 0x14, 0x14, 0x0c, 0x04, 0x1c, 0x14, 0x0c, 0x14, 0x3e, 0x14, 0x0c, + 0x04, 0x04, 0x1c, 0x0c, 0x14, 0x04, 0x1c, 0x0c, 0x04, 0x14, 0x1c, 0x0c, + 0x0c, 0x1c, 0x1c, 0x0c, 0x34, 0x24, 0x1c, 0x0c, 0x34, 0x34, 0x1c, 0x0c, + 0x0c, 0x04, 0x24, 0x0c, 0x2c, 0x04, 0x24, 0x0c, 0x04, 0x2c, 0x24, 0x0c, + 0x04, 0x14, 0x2c, 0x0c, 0x24, 0x14, 0x2c, 0x0c, 0x34, 0x24, 0x2c, 0x0c, + 0x0c, 0x3e, 0x2c, 0x0c, 0x2c, 0x04, 0x34, 0x0c, 0x14, 0x14, 0x3e, 0x0c, + 0x04, 0x24, 0x3e, 0x0c, 0x04, 0x04, 0x04, 0x14, 0x14, 0x04, 0x04, 0x14, + 0x0c, 0x0c, 0x04, 0x14, 0x1c, 0x0c, 0x04, 0x14, 0x04, 0x14, 0x04, 0x14, + 0x14, 0x14, 0x04, 0x14, 0x34, 0x14, 0x04, 0x14, 0x0c, 0x1c, 0x04, 0x14, + 0x14, 0x24, 0x04, 0x14, 0x0c, 0x04, 0x0c, 0x14, 0x1c, 0x04, 0x0c, 0x14, + 0x2c, 0x04, 0x0c, 0x14, 0x04, 0x0c, 0x0c, 0x14, 0x14, 0x0c, 0x0c, 0x14, + 0x0c, 0x14, 0x0c, 0x14, 0x04, 0x1c, 0x0c, 0x14, 0x1c, 0x34, 0x0c, 0x14, + 0x3e, 0x34, 0x0c, 0x14, 0x04, 0x3e, 0x0c, 0x14, 0x04, 0x04, 0x14, 0x14, + 0x14, 0x04, 0x14, 0x14, 0x0c, 0x0c, 0x14, 0x14, 0x3e, 0x0c, 0x14, 0x14, + 0x04, 0x14, 0x14, 0x14, 0x14, 0x14, 0x14, 0x14, 0x3e, 0x1c, 0x14, 0x14, + 0x04, 0x24, 0x14, 0x14, 0x2c, 0x2c, 0x14, 0x14, 0x0c, 0x04, 0x1c, 0x14, + 0x04, 0x0c, 0x1c, 0x14, 0x24, 0x0c, 0x1c, 0x14, 0x04, 0x3e, 0x1c, 0x14, + 0x24, 0x3e, 0x1c, 0x14, 0x2c, 0x1c, 0x24, 0x14, 0x1c, 0x2c, 0x24, 0x14, + 0x1c, 0x04, 0x2c, 0x14, 0x3e, 0x14, 0x2c, 0x14, 0x0c, 0x24, 0x2c, 0x14, + 0x24, 0x3e, 0x2c, 0x14, 0x0c, 0x04, 0x3e, 0x14, 0x1c, 0x04, 0x3e, 0x14, + 0x34, 0x0c, 0x3e, 0x14, 0x2c, 0x24, 0x3e, 0x14, 0x0c, 0x04, 0x04, 0x1c, + 0x04, 0x0c, 0x04, 0x1c, 0x14, 0x0c, 0x04, 0x1c, 0x0c, 0x14, 0x04, 0x1c, + 0x1c, 0x14, 0x04, 0x1c, 0x04, 0x2c, 0x04, 0x1c, 0x2c, 0x34, 0x04, 0x1c, + 0x14, 0x3e, 0x04, 0x1c, 0x04, 0x04, 0x0c, 0x1c, 0x14, 0x04, 0x0c, 0x1c, + 0x04, 0x14, 0x0c, 0x1c, 0x0c, 0x1c, 0x0c, 0x1c, 0x24, 0x24, 0x0c, 0x1c, + 0x34, 0x24, 0x0c, 0x1c, 0x0c, 0x04, 0x14, 0x1c, 0x1c, 0x04, 0x14, 0x1c, + 0x04, 0x0c, 0x14, 0x1c, 0x2c, 0x14, 0x14, 0x1c, 0x14, 0x2c, 0x14, 0x1c, + 0x14, 0x3e, 0x14, 0x1c, 0x0c, 0x0c, 0x1c, 0x1c, 0x1c, 0x1c, 0x1c, 0x1c, + 0x04, 0x1c, 0x24, 0x1c, 0x3e, 0x24, 0x24, 0x1c, 0x14, 0x3e, 0x24, 0x1c, + 0x04, 0x04, 0x2c, 0x1c, 0x34, 0x04, 0x2c, 0x1c, 0x14, 0x14, 0x2c, 0x1c, + 0x2c, 0x2c, 0x2c, 0x1c, 0x24, 0x0c, 0x34, 0x1c, 0x34, 0x1c, 0x34, 0x1c, + 0x1c, 0x34, 0x34, 0x1c, 0x1c, 0x1c, 0x3e, 0x1c, 0x04, 0x34, 0x3e, 0x1c, + 0x24, 0x04, 0x04, 0x24, 0x3e, 0x0c, 0x04, 0x24, 0x2c, 0x1c, 0x04, 0x24, + 0x3e, 0x1c, 0x04, 0x24, 0x1c, 0x2c, 0x04, 0x24, 0x3e, 0x2c, 0x04, 0x24, + 0x24, 0x3e, 0x0c, 0x24, 0x04, 0x14, 0x14, 0x24, 0x3e, 0x1c, 0x14, 0x24, + 0x04, 0x24, 0x14, 0x24, 0x04, 0x34, 0x14, 0x24, 0x34, 0x34, 0x14, 0x24, + 0x3e, 0x04, 0x1c, 0x24, 0x2c, 0x24, 0x1c, 0x24, 0x24, 0x04, 0x24, 0x24, + 0x0c, 0x2c, 0x24, 0x24, 0x24, 0x34, 0x24, 0x24, 0x2c, 0x14, 0x2c, 0x24, + 0x1c, 0x24, 0x2c, 0x24, 0x04, 0x3e, 0x2c, 0x24, 0x2c, 0x04, 0x3e, 0x24, + 0x04, 0x0c, 0x3e, 0x24, 0x14, 0x0c, 0x3e, 0x24, 0x04, 0x1c, 0x3e, 0x24, + 0x14, 0x0c, 0x04, 0x2c, 0x0c, 0x24, 0x04, 0x2c, 0x04, 0x3e, 0x04, 0x2c, + 0x04, 0x04, 0x0c, 0x2c, 0x34, 0x04, 0x0c, 0x2c, 0x34, 0x14, 0x0c, 0x2c, + 0x2c, 0x2c, 0x0c, 0x2c, 0x24, 0x0c, 0x14, 0x2c, 0x14, 0x1c, 0x14, 0x2c, + 0x14, 0x3e, 0x14, 0x2c, 0x14, 0x04, 0x1c, 0x2c, 0x1c, 0x2c, 0x1c, 0x2c, + 0x04, 0x0c, 0x24, 0x2c, 0x1c, 0x14, 0x24, 0x2c, 0x3e, 0x14, 0x24, 0x2c, + 0x14, 0x3e, 0x24, 0x2c, 0x14, 0x04, 0x2c, 0x2c, 0x0c, 0x1c, 0x2c, 0x2c, + 0x04, 0x2c, 0x34, 0x2c, 0x24, 0x14, 0x3e, 0x2c, 0x14, 0x24, 0x3e, 0x2c, + 0x24, 0x14, 0x04, 0x34, 0x24, 0x24, 0x04, 0x34, 0x34, 0x24, 0x04, 0x34, + 0x24, 0x34, 0x04, 0x34, 0x0c, 0x14, 0x0c, 0x34, 0x0c, 0x34, 0x0c, 0x34, + 0x3e, 0x0c, 0x14, 0x34, 0x24, 0x34, 0x14, 0x34, 0x04, 0x1c, 0x1c, 0x34, + 0x34, 0x1c, 0x1c, 0x34, 0x24, 0x24, 0x24, 0x34, 0x2c, 0x04, 0x2c, 0x34, + 0x14, 0x2c, 0x2c, 0x34, 0x1c, 0x1c, 0x34, 0x34, 0x1c, 0x04, 0x3e, 0x34, + 0x0c, 0x14, 0x3e, 0x34, 0x1c, 0x04, 0x04, 0x3e, 0x2c, 0x04, 0x04, 0x3e, + 0x3e, 0x04, 0x04, 0x3e, 0x04, 0x0c, 0x04, 0x3e, 0x14, 0x1c, 0x04, 0x3e, + 0x14, 0x2c, 0x04, 0x3e, 0x34, 0x14, 0x0c, 0x3e, 0x04, 0x24, 0x0c, 0x3e, + 0x14, 0x0c, 0x14, 0x3e, 0x2c, 0x24, 0x14, 0x3e, 0x14, 0x2c, 0x14, 0x3e, + 0x04, 0x04, 0x1c, 0x3e, 0x2c, 0x0c, 0x1c, 0x3e, 0x1c, 0x1c, 0x1c, 0x3e, + 0x04, 0x34, 0x1c, 0x3e, 0x0c, 0x14, 0x24, 0x3e, 0x0c, 0x24, 0x24, 0x3e, + 0x04, 0x04, 0x2c, 0x3e, 0x14, 0x04, 0x2c, 0x3e, 0x24, 0x14, 0x2c, 0x3e, + 0x04, 0x1c, 0x34, 0x3e, 0x00, 0x81, 0x82, 0x03, 0x84, 0x05, 0x06, 0x87, + 0x88, 0x09, 0x0a, 0x8b, 0x0c, 0x8d, 0x8e, 0x0f, 0x90, 0x11, 0x12, 0x93, + 0x14, 0x95, 0x96, 0x17, 0x18, 0x99, 0x9a, 0x1b, 0x9c, 0x1d, 0x1e, 0x9f, + 0xa0, 0x21, 0x22, 0xa3, 0x24, 0xa5, 0xa6, 0x27, 0x28, 0xa9, 0xaa, 0x2b, + 0xac, 0x2d, 0x2e, 0xaf, 0x30, 0xb1, 0xb2, 0x33, 0xb4, 0x35, 0x36, 0xb7, + 0xb8, 0x39, 0x3a, 0xbb, 0x3c, 0xbd, 0xbe, 0x3f, 0xc0, 0x41, 0x42, 0xc3, + 0x44, 0xc5, 0xc6, 0x47, 0x48, 0xc9, 0xca, 0x4b, 0xcc, 0x4d, 0x4e, 0xcf, + 0x50, 0xd1, 0xd2, 0x53, 0xd4, 0x55, 0x56, 0xd7, 0xd8, 0x59, 0x5a, 0xdb, + 0x5c, 0xdd, 0xde, 0x5f, 0x60, 0xe1, 0xe2, 0x63, 0xe4, 0x65, 0x66, 0xe7, + 0xe8, 0x69, 0x6a, 0xeb, 0x6c, 0xed, 0xee, 0x6f, 0xf0, 0x71, 0x72, 0xf3, + 0x74, 0xf5, 0xf6, 0x77, 0x78, 0xf9, 0xfa, 0x7b, 0xfc, 0x7d, 0x7e, 0xff, + 0x01, 0x02, 0x04, 0x08, 0x10, 0x20, 0x40, 0x80, 0x00, 0x00, 0x00, 0x00, + 0x00, 0x00, 0x00, 0x00, +}; + +// IQ3_XXS weight x f32 activation matvec, one workitem per output row (Unsloth UD GGUFs) +static bool match_iq3_xxs_matvec_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_MUL_MAT || node->inputs.size() != 2) { + return false; + } + const Value * weight = common_graph_value(context.graph, node->inputs[0]); + const Value * input = common_graph_value(context.graph, node->inputs[1]); + const Value * output = common_graph_value(context.graph, node->output); + if (weight == nullptr || input == nullptr || output == nullptr) { + return false; + } + if (weight->type != GGML_TYPE_IQ3_XXS || input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32) { + return false; + } + if (!common_is_2d(*weight) || !common_is_2d(*input) || !common_is_2d(*output) || !weight->contiguous || !input->contiguous || + !output->contiguous) { + return false; + } + const int64_t input_size = weight->ne[0]; + const int64_t output_size = weight->ne[1]; + if (input_size <= 0 || input_size % 256 != 0 || input->ne[0] != input_size || output->ne[0] != output_size) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kMulMatVecIq3XxsF32Kernel); + dispatch.kernel.integer_parameters.emplace("input_size", input_size); + dispatch.kernel.integer_parameters.emplace("output_size", output_size); + dispatch.kernel.integer_parameters.emplace("token_count", input->ne[1]); + const ValueId tables = ValueId(context.next_plan_value.value + + static_cast(dispatch_match.transients.size()) + + static_cast(dispatch_match.completion_counter_requests.size())); + dispatch_match.transients.push_back({ tables, "common.mul_mat.iq3_xxs.tables", sizeof(kIq3XxsTables), 16 }); + dispatch_match.constant_initializations.push_back({ tables, "common.mul_mat.iq3_xxs.tables", 0, + std::vector(kIq3XxsTables, + kIq3XxsTables + sizeof(kIq3XxsTables)) }); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ weight->id, 0, weight->byte_count }); + dispatch.bindings.push_back({ tables, 0, sizeof(kIq3XxsTables) }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_mul_mat_iq3_xxs_dispatch(DispatchRegistryBuilder & registry) { + registry.add({ + "common.mul_mat.iq3_xxs_matvec_f32", + GGML_OP_MUL_MAT, + DispatchMatchKind::SingleOp, + // below kquant.mul_mat.decode_f32 (85), which reads IQ3_XXS through the shared lane functions; + // above the generic common.mul_mat f32 matchers (80/70/60) for shapes the K-quant kernels decline + 84, + DispatchSource::Common, + match_iq3_xxs_matvec_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-iq3-xxs.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-iq3-xxs.h new file mode 100644 index 000000000000..fd6c7e4a83ec --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-iq3-xxs.h @@ -0,0 +1,25 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +// IQ3_XXS weight MUL_MAT (Unsloth UD GGUFs): registered by register_mul_mat_dispatches +void register_mul_mat_iq3_xxs_dispatch(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-tail.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-tail.cpp new file mode 100644 index 000000000000..39c14003b855 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-tail.cpp @@ -0,0 +1,62 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "dispatch-mul-mat-tail.h" + +#include "graph/op-params.h" + +#include +#include + +namespace ggml::hrx { + +namespace { + +constexpr KernelCatalogRef kMulMatTailKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_f32_f32_wmma"); + +} // namespace + +bool common_append_mul_mat_token_tail(const CommonMulMatMatch & match, + int64_t head_tokens, + DispatchMatch & dispatch_match) { + const int64_t tail_tokens = match.token_count - head_tokens; + if (head_tokens <= 0 || tail_tokens < 2) { + return false; + } + const size_t input_row = static_cast(match.input_size) * sizeof(float); + const size_t output_row = static_cast(match.output_size) * sizeof(float); + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kMulMatTailKernel); + dispatch.kernel.integer_parameters.emplace("token_count", tail_tokens); + dispatch.kernel.compile_parameters.emplace("ggml.workload.token_capacity", common_to_config_value(tail_tokens)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat.input_size", common_to_config_value(match.input_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat.output_size", common_to_config_value(match.output_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat.output_accumulation", "0"); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat.output_unary_op", + std::to_string(unary_kind_config_value(UnaryKind::Identity))); + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat.weight_format", + common_to_config_value(common_mul_mat_format_config_value(match.weight_format))); + dispatch.bindings.push_back({ match.input->id, static_cast(head_tokens) * input_row, + static_cast(tail_tokens) * input_row }); + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count }); + dispatch.bindings.push_back({ match.output->id, static_cast(head_tokens) * output_row, + static_cast(tail_tokens) * output_row }); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-tail.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-tail.h new file mode 100644 index 000000000000..23f798829011 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-tail.h @@ -0,0 +1,32 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include "dispatch-mul-mat-common.h" + +namespace ggml::hrx { + +// The q8_1 x4 prefill kernel takes token counts that are multiples of 256. For +// a prompt chunk with a remainder, it runs the 256-aligned head and this +// appends one generic F32 WMMA matmul for tokens [head_tokens, token_count): +// input and output are token-major, so the tail is the same buffers bound at +// the head's byte offset. Returns false (appending nothing) for an empty head +// or a tail shorter than 2 tokens. +bool common_append_mul_mat_token_tail(const CommonMulMatMatch & match, + int64_t head_tokens, + DispatchMatch & dispatch_match); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-weight-format.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-weight-format.h new file mode 100644 index 000000000000..c76178e4da43 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-weight-format.h @@ -0,0 +1,241 @@ +#pragma once + +#include "ggml.h" + +#include + +namespace ggml::hrx { + +enum class CommonMulMatWeightFormat { + Q1_0, + PQ2_0, // PrismML group-128 2-bit (ggml-prism.h) + PTQ1_0, // PrismML group-128 ternary + Q2K, + Q3K, + Q4K, + Q4KRow64, + Q5K, + Q6K, + Q6KRow64, + Q4_0, + Q4_1, + Q5_0, + Q5_1, + IQ1_S, + IQ1_M, + IQ2_XXS, + IQ2_XS, + IQ3_XXS, + IQ2_S, + IQ3_S, + IQ4_NL, + MXFP4, + TQ1_0, + TQ2_0, + IQ4_XS, + Q8_0, + Q8_1, + F16, + BF16, + F32, +}; + +inline bool common_mul_mat_format_for_type(ggml_type type, CommonMulMatWeightFormat & format) { + switch (type) { + case GGML_TYPE_Q1_0: + format = CommonMulMatWeightFormat::Q1_0; + return true; + case GGML_TYPE_PQ2_0: + format = CommonMulMatWeightFormat::PQ2_0; + return true; + case GGML_TYPE_PTQ1_0: + format = CommonMulMatWeightFormat::PTQ1_0; + return true; + case GGML_TYPE_Q2_K: + format = CommonMulMatWeightFormat::Q2K; + return true; + case GGML_TYPE_Q3_K: + format = CommonMulMatWeightFormat::Q3K; + return true; + case GGML_TYPE_Q4_K: + format = CommonMulMatWeightFormat::Q4K; + return true; + case GGML_TYPE_Q5_K: + format = CommonMulMatWeightFormat::Q5K; + return true; + case GGML_TYPE_Q6_K: + format = CommonMulMatWeightFormat::Q6K; + return true; + case GGML_TYPE_Q4_0: + format = CommonMulMatWeightFormat::Q4_0; + return true; + case GGML_TYPE_Q4_1: + format = CommonMulMatWeightFormat::Q4_1; + return true; + case GGML_TYPE_Q5_0: + format = CommonMulMatWeightFormat::Q5_0; + return true; + case GGML_TYPE_Q5_1: + format = CommonMulMatWeightFormat::Q5_1; + return true; + case GGML_TYPE_IQ3_XXS: + format = CommonMulMatWeightFormat::IQ3_XXS; + return true; + case GGML_TYPE_IQ1_S: + format = CommonMulMatWeightFormat::IQ1_S; + return true; + case GGML_TYPE_IQ1_M: + format = CommonMulMatWeightFormat::IQ1_M; + return true; + case GGML_TYPE_IQ2_XXS: + format = CommonMulMatWeightFormat::IQ2_XXS; + return true; + case GGML_TYPE_IQ2_XS: + format = CommonMulMatWeightFormat::IQ2_XS; + return true; + case GGML_TYPE_IQ2_S: + format = CommonMulMatWeightFormat::IQ2_S; + return true; + case GGML_TYPE_IQ3_S: + format = CommonMulMatWeightFormat::IQ3_S; + return true; + case GGML_TYPE_IQ4_NL: + format = CommonMulMatWeightFormat::IQ4_NL; + return true; + case GGML_TYPE_MXFP4: + format = CommonMulMatWeightFormat::MXFP4; + return true; + case GGML_TYPE_TQ1_0: + format = CommonMulMatWeightFormat::TQ1_0; + return true; + case GGML_TYPE_TQ2_0: + format = CommonMulMatWeightFormat::TQ2_0; + return true; + case GGML_TYPE_IQ4_XS: + format = CommonMulMatWeightFormat::IQ4_XS; + return true; + case GGML_TYPE_Q8_0: + format = CommonMulMatWeightFormat::Q8_0; + return true; + case GGML_TYPE_Q8_1: + format = CommonMulMatWeightFormat::Q8_1; + return true; + case GGML_TYPE_F16: + format = CommonMulMatWeightFormat::F16; + return true; + case GGML_TYPE_BF16: + format = CommonMulMatWeightFormat::BF16; + return true; + case GGML_TYPE_F32: + format = CommonMulMatWeightFormat::F32; + return true; + default: + return false; + } +} + +inline int64_t common_mul_mat_format_config_value(CommonMulMatWeightFormat format) { + switch (format) { + case CommonMulMatWeightFormat::Q1_0: + return 10; + case CommonMulMatWeightFormat::PQ2_0: + return 72; + case CommonMulMatWeightFormat::PTQ1_0: + return 73; + case CommonMulMatWeightFormat::Q2K: + return 12; + case CommonMulMatWeightFormat::Q3K: + return 11; + case CommonMulMatWeightFormat::Q4K: + return 4; + case CommonMulMatWeightFormat::Q4KRow64: + return 44; + case CommonMulMatWeightFormat::Q5K: + return 5; + case CommonMulMatWeightFormat::Q6K: + return 6; + case CommonMulMatWeightFormat::Q6KRow64: + return 46; + case CommonMulMatWeightFormat::Q4_0: + return 40; + case CommonMulMatWeightFormat::Q4_1: + return 41; + case CommonMulMatWeightFormat::Q5_0: + return 50; + case CommonMulMatWeightFormat::Q5_1: + return 51; + case CommonMulMatWeightFormat::IQ3_XXS: + return 28; + case CommonMulMatWeightFormat::IQ1_S: + return 26; + case CommonMulMatWeightFormat::IQ1_M: + return 27; + case CommonMulMatWeightFormat::IQ2_XXS: + return 24; + case CommonMulMatWeightFormat::IQ2_XS: + return 25; + case CommonMulMatWeightFormat::IQ2_S: + return 22; + case CommonMulMatWeightFormat::IQ3_S: + return 21; + case CommonMulMatWeightFormat::IQ4_NL: + return 20; + case CommonMulMatWeightFormat::MXFP4: + return 39; + case CommonMulMatWeightFormat::TQ1_0: + return 34; + case CommonMulMatWeightFormat::TQ2_0: + return 35; + case CommonMulMatWeightFormat::IQ4_XS: + return 23; + case CommonMulMatWeightFormat::Q8_0: + return 80; + case CommonMulMatWeightFormat::Q8_1: + return 81; + case CommonMulMatWeightFormat::F16: + return 16; + case CommonMulMatWeightFormat::BF16: + return 30; + case CommonMulMatWeightFormat::F32: + return 32; + } + return 0; +} + +inline bool common_mul_mat_dense_float_format(CommonMulMatWeightFormat format) { + return format == CommonMulMatWeightFormat::F16 || format == CommonMulMatWeightFormat::BF16 || + format == CommonMulMatWeightFormat::F32; +} + +inline bool common_mul_mat_supported_dense_input_size(CommonMulMatWeightFormat format, int64_t input_size) { + switch (format) { + case CommonMulMatWeightFormat::Q1_0: + case CommonMulMatWeightFormat::PQ2_0: + case CommonMulMatWeightFormat::PTQ1_0: + return input_size >= 256 && input_size <= 32768 && input_size % 128 == 0; + case CommonMulMatWeightFormat::Q4_0: + case CommonMulMatWeightFormat::Q4_1: + case CommonMulMatWeightFormat::Q5_0: + case CommonMulMatWeightFormat::Q5_1: + case CommonMulMatWeightFormat::IQ4_NL: + case CommonMulMatWeightFormat::MXFP4: + case CommonMulMatWeightFormat::Q8_0: + case CommonMulMatWeightFormat::Q8_1: + return input_size >= 256 && input_size <= 32768 && input_size % 32 == 0; + default: + return input_size >= 256 && input_size <= 32768 && input_size % 256 == 0; + } +} + +inline bool common_mul_mat_alternate_format_for_type(ggml_type type, CommonMulMatWeightFormat & format) { + switch (type) { + case GGML_TYPE_Q8_1: + case GGML_TYPE_F16: + case GGML_TYPE_F32: + return common_mul_mat_format_for_type(type, format); + default: + return false; + } +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat.cpp new file mode 100644 index 000000000000..08ca1190510c --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat.cpp @@ -0,0 +1,1482 @@ +#include "dispatch-mul-mat.h" + +#include + +#include "dispatch-mul-mat-common.h" +#include "dispatch-mul-mat-iq3-xxs.h" +#include "dispatch-mul-mat-tail.h" +#include "ggml.h" +#include "graph/graph-matcher.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kMulMatF32F32WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_f32_f32_wmma"); +static constexpr KernelCatalogRef kMulMatF32F32NarrowSplitK4Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_f32_f32_narrow_split_k4"); +static constexpr KernelCatalogRef kMulMatF32F32DecodeWave64Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_f32_f32_decode_wave64"); +static constexpr KernelCatalogRef kMulMatAddF32F32DecodeWave64Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_add_f32_f32_decode_wave64"); +static constexpr KernelCatalogRef kMulMatBiasF32F32WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_bias_f32_f32_wmma"); +static constexpr KernelCatalogRef kMulMatAddF32F32WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_add_f32_f32_wmma"); +static constexpr KernelCatalogRef kMulMatBiasAddF32F32WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_bias_add_f32_f32_wmma"); +static constexpr KernelCatalogRef kQuantizeF32SymmetricI4K64PlaneKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_quantize_f32_symmetric_i4_k64_plane"); +static constexpr KernelCatalogRef kMulMatSymmetricI4WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_symmetric_i4_wmma"); +static constexpr KernelCatalogRef kQuantizeF32SymmetricI4K32Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_quantize_f32_symmetric_i4_k32"); +static constexpr KernelCatalogRef kMulMatSymmetricI4LowRowWmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_symmetric_i4_lowrow_wmma"); +static constexpr KernelCatalogRef kMulMatSymmetricI4LowRowSplitK2WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma"); +static constexpr KernelCatalogRef kMulMatSymmetricI4LowRowSplitK2DirectDotKernels[] = { + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c1"), + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c2"), + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c3"), + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c4"), + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c5"), +}; +static constexpr KernelCatalogRef kQuantizeF32SymmetricI8K256Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_quantize_f32_symmetric_i8_k256"); +static constexpr KernelCatalogRef kMulMatQ5KSymmetricI8WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_q5_k_symmetric_i8_wmma"); +static constexpr KernelCatalogRef kMulMatQ5KIQ4XSQ8_1X4WmmaToken256Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_q5_k_iq4_xs_q8_1_x4_wmma_token256"); +static constexpr KernelCatalogRef kMulMatQ4KF16WmmaPrefillWave32Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_q4_k_f16_wmma_prefill_wave32"); +static constexpr KernelCatalogRef kMulMatQ6KPackedToken1F16WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_q6_k_packed_token1_f16_wmma"); +static constexpr int64_t kMulMatQ6KPackedMaxOutputSize = 262144; +static constexpr KernelCatalogRef kMulMatQ6KI8PrepackedF16WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma"); +static constexpr KernelCatalogRef kMulMatQ6KF32WmmaPrefillWave32Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_q6_k_f32_wmma_prefill_wave32"); +static constexpr KernelCatalogRef kMulMatQ6KF16WmmaPrefillWave32Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_q6_k_f16_wmma_prefill_wave32"); +static constexpr KernelCatalogRef kSelectSymmetricI4K32GroupsKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_select_symmetric_i4_k32_groups"); +static constexpr KernelCatalogRef kMulMatQ6KSymmetricI2ScanToken1Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_q6_k_symmetric_i2_scan_token1"); +static constexpr KernelCatalogRef kTopK8F32PartitionsRegisterKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_top_k8_f32_partitions_register"); +static constexpr KernelCatalogRef kTopK128F32ReduceGatherRegisterKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_top_k128_f32_reduce_gather_register"); +static constexpr KernelCatalogRef kFillNegativeF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_fill_negative_f32"); +static constexpr KernelCatalogRef kMulMatQ6KPackedSelectedRefineToken1Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_q6_k_packed_selected_refine_token1"); + +struct MulMatPostOpsMatch { + const Value * input = nullptr; + const Value * weight = nullptr; + const Value * projection_output = nullptr; + const Value * bias = nullptr; + const Value * residual_input = nullptr; + const Value * residual_output = nullptr; + const GraphNode * bias_add_node = nullptr; + const GraphNode * residual_add_node = nullptr; + const GraphNode * layout_node = nullptr; + std::vector add_nodes; + KernelCatalogRef kernel = {}; + CommonMulMatWeightFormat weight_format = CommonMulMatWeightFormat::Q4K; + int64_t input_size = 0; + int64_t output_size = 0; + int64_t token_count = 0; + bool has_bias = false; + bool has_residual = false; + bool requires_q8_activation = false; + + bool matched() const { + return input != nullptr && weight != nullptr && projection_output != nullptr && residual_output != nullptr && + kernel.id != kUncatalogedKernelId && (has_bias || has_residual); + } +}; + +struct DecodeMulMatAddMatch { + CommonMulMatMatch root; + const Value * residual_input = nullptr; + const Value * residual_output = nullptr; + const GraphNode * add_node = nullptr; + + bool matched() const { + return root.matched() && residual_input != nullptr && residual_output != nullptr && add_node != nullptr; + } +}; + +static bool is_bias_shape(const Value & value, int64_t output_size) { + if (value.kind != ValueKind::External || value.type != GGML_TYPE_F32 || !value.contiguous || + value.ne[0] != output_size) { + return false; + } + for (int i = 1; i < GGML_MAX_DIMS; ++i) { + if (value.ne[i] != 1) { + return false; + } + } + return true; +} + +static bool has_normalized_rope_consumer(const Graph & graph, const Value & value) { + if (!graph.has_index()) { + return false; + } + for (const GraphNode * layout : graph.index().consumers(value.id)) { + if (layout == nullptr || (layout->op != GGML_OP_RESHAPE && layout->op != GGML_OP_VIEW)) { + continue; + } + const GraphNode * rms = common_find_single_consumer_with_op(graph, layout->output, GGML_OP_RMS_NORM); + const GraphNode * mul = + rms != nullptr ? common_find_single_consumer_with_op(graph, rms->output, GGML_OP_MUL) : nullptr; + if (mul != nullptr && common_find_single_consumer_with_op(graph, mul->output, GGML_OP_ROPE) != nullptr) { + return true; + } + } + return false; +} + +static bool has_reshaped_unary_mul_consumer(const Graph & graph, const Value & value) { + const GraphNode * reshape = common_find_single_consumer_with_op(graph, value.id, GGML_OP_RESHAPE); + const GraphNode * unary = + reshape != nullptr ? common_find_single_consumer_with_op(graph, reshape->output, GGML_OP_UNARY) : nullptr; + return unary != nullptr && common_find_single_consumer_with_op(graph, unary->output, GGML_OP_MUL) != nullptr; +} + +static bool has_reshaped_transpose_concat_consumer(const Graph & graph, const Value & value) { + const GraphNode * reshape = common_find_single_consumer_with_op(graph, value.id, GGML_OP_RESHAPE); + const GraphNode * transpose = + reshape != nullptr ? common_find_single_consumer_with_op(graph, reshape->output, GGML_OP_TRANSPOSE) : nullptr; + return transpose != nullptr && + common_find_single_consumer_with_op(graph, transpose->output, GGML_OP_CONCAT) != nullptr; +} + +static const Value * match_scaled_normalized_branch(const Graph & graph, ValueId branch_id, int64_t hidden_size) { + const Value * branch = common_graph_value(graph, branch_id); + const GraphNode * mul = graph.index().producer(branch_id); + if (branch == nullptr || mul == nullptr || !common_binary_node_is_mul(*mul) || branch->type != GGML_TYPE_F32 || + !common_is_2d(*branch) || branch->ne[0] != hidden_size || branch->ne[1] != 1) { + return nullptr; + } + + const GraphNode * rms = nullptr; + const Value * scale = nullptr; + for (ValueId input_id : mul->inputs) { + const GraphNode * producer = graph.index().producer(input_id); + if (producer != nullptr && producer->op == GGML_OP_RMS_NORM) { + if (rms != nullptr) { + return nullptr; + } + rms = producer; + } else { + scale = common_graph_value(graph, input_id); + } + } + if (rms == nullptr || rms->inputs.size() != 1 || scale == nullptr || scale->type != GGML_TYPE_F32 || + !common_is_2d(*scale) || scale->ne[0] != hidden_size || scale->ne[1] != 1) { + return nullptr; + } + + const Value * normalized = common_graph_value(graph, rms->output); + const Value * source = common_graph_value(graph, rms->inputs[0]); + if (normalized == nullptr || source == nullptr || normalized->type != GGML_TYPE_F32 || + source->type != GGML_TYPE_F32 || !common_is_2d(*normalized) || !common_is_2d(*source) || + normalized->ne[0] != hidden_size || normalized->ne[1] != 1 || source->ne[0] != hidden_size || + source->ne[1] != 1) { + return nullptr; + } + return source; +} + +static bool matches_embedded_external_concat_projection(const Graph & graph, + const GraphNode & projection, + int64_t hidden_size) { + if (projection.op != GGML_OP_MUL_MAT || projection.inputs.size() != 2) { + return false; + } + + const Value * weight = common_graph_value(graph, projection.inputs[0]); + const Value * input = common_graph_value(graph, projection.inputs[1]); + const Value * output = common_graph_value(graph, projection.output); + if (weight == nullptr || input == nullptr || output == nullptr || !common_is_2d(*weight) || !common_is_2d(*input) || + !common_is_2d(*output) || !weight->contiguous || !input->contiguous || !output->contiguous || + input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || weight->ne[0] != 2 * hidden_size || + weight->ne[1] != hidden_size || input->ne[0] != 2 * hidden_size || input->ne[1] != 1 || + output->ne[0] != hidden_size || output->ne[1] != 1) { + return false; + } + + const GraphNode * concat = graph.index().producer(input->id); + if (concat == nullptr || concat->op != GGML_OP_CONCAT || concat->inputs.size() != 2) { + return false; + } + const Value * first_source = match_scaled_normalized_branch(graph, concat->inputs[0], hidden_size); + const Value * second_source = match_scaled_normalized_branch(graph, concat->inputs[1], hidden_size); + if (first_source == nullptr || second_source == nullptr) { + return false; + } + + const GraphNode * first_producer = graph.index().producer(first_source->id); + const GraphNode * second_producer = graph.index().producer(second_source->id); + const bool first_is_embedding = + first_producer != nullptr && first_producer->op == GGML_OP_GET_ROWS && first_producer->inputs.size() == 2; + const bool second_is_embedding = + second_producer != nullptr && second_producer->op == GGML_OP_GET_ROWS && second_producer->inputs.size() == 2; + const bool first_is_external = first_producer == nullptr; + const bool second_is_external = second_producer == nullptr; + return (first_is_embedding && second_is_external) || (second_is_embedding && first_is_external); +} + +static bool has_embedded_external_concat_projection_ancestor(const Graph & graph, + ValueId endpoint_input, + int64_t hidden_size) { + if (!graph.has_index() || endpoint_input.value < 0 || hidden_size <= 0) { + return false; + } + + std::vector visited(graph.values().size(), false); + std::vector pending = { endpoint_input }; + while (!pending.empty()) { + const ValueId current = pending.back(); + pending.pop_back(); + if (current.value < 0 || static_cast(current.value) >= visited.size() || visited[current.value]) { + continue; + } + visited[current.value] = true; + const GraphNode * producer = graph.index().producer(current); + if (producer == nullptr) { + continue; + } + if (matches_embedded_external_concat_projection(graph, *producer, hidden_size)) { + return true; + } + pending.insert(pending.end(), producer->inputs.begin(), producer->inputs.end()); + } + return false; +} + +static bool match_symmetric_i4_low_row_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + // Disabled: this symmetric matmul route causes a substantial numeric performance regression. + (void) context; + (void) dispatch_match; + return false; + + CommonMulMatMatch match = + common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatSymmetricI4LowRowWmmaKernel, false); + if (!match.matched()) { + match = + common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatSymmetricI4LowRowWmmaKernel, true); + } + if (!match.matched() || (match.weight->type != GGML_TYPE_Q5_K && match.weight->type != GGML_TYPE_IQ4_XS) || + match.weight->alias_source.value >= 0 || match.token_count < 1 || match.token_count > 16 || + match.input_size % 64 != 0 || match.output_size % 64 != 0) { + return false; + } + const CommonSymmetricI4ActivationLayout activation_layout = + common_symmetric_i4_activation_layout(match.input_size, match.token_count); + const CommandPlanAlternateValue * alternate = find_alternate_value(context.graph, context.plan, match.input->id, + GGML_TYPE_COUNT, activation_layout.total_bytes); + if (alternate != nullptr && alternate->name != kCommonSymmetricI4K32ActivationAlternateName) { + alternate = nullptr; + } + + const ValueId activation = alternate != nullptr ? alternate->alternate_value : context.next_plan_value; + const int32_t next_value = context.next_plan_value.value + (alternate == nullptr ? 1 : 0); + + if (alternate == nullptr) { + dispatch_match.transients.push_back( + { activation, kCommonSymmetricI4K32ActivationAlternateName, activation_layout.total_bytes, 256 }); + + Dispatch quantize; + quantize.kernel = make_kernel_specialization(kQuantizeF32SymmetricI4K32Kernel); + quantize.kernel.compile_parameters.emplace("ggml.quantize_symmetric_i4_k32.input_size", + common_to_config_value(match.input_size)); + quantize.kernel.compile_parameters.emplace("ggml.quantize_symmetric_i4_k32.token_count", + common_to_config_value(match.token_count)); + quantize.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + quantize.bindings.push_back({ activation, 0, activation_layout.payload_bytes }); + quantize.bindings.push_back({ activation, activation_layout.scales_offset, activation_layout.metadata_bytes }); + quantize.bindings.push_back({ activation, activation_layout.sums_offset, activation_layout.metadata_bytes }); + dispatch_match.dispatches.push_back(std::move(quantize)); + + Status metadata_status; + if (!dispatch_match.metadata.append_alternate_value( + { match.input->id, activation, GGML_TYPE_COUNT, activation_layout.total_bytes, + kCommonSymmetricI4K32ActivationAlternateName }, + metadata_status)) { + dispatch_match.status.append(metadata_status); + return false; + } + } + + const size_t materialized_weight_bytes = + common_symmetric_i4_shared4_weight_byte_count(match.input_size, match.output_size); + const bool split_k = match.input_size % 128 == 0 && materialized_weight_bytes >= size_t{ 16 } * 1024 * 1024 && + match.output_size <= 2 * match.input_size; + + const ValueId partial = split_k ? ValueId(next_value) : ValueId{}; + const ValueId completion_counters = split_k ? ValueId(next_value + 1) : ValueId{}; + const uint32_t completion_counter_count = + static_cast(common_ceil_div(match.output_size, 16) * common_ceil_div(match.token_count, 16)); + if (split_k) { + dispatch_match.transients.push_back( + { partial, "common.mul_mat.symmetric_i4_lowrow.split_k_partial", match.output->byte_count, 256 }); + dispatch_match.completion_counter_requests.push_back({ + completion_counters, + "common.mul_mat.symmetric_i4_lowrow.split_k_completion_counters", + completion_counter_count, + }); + } + + const bool use_split_direct_dot = split_k && match.token_count <= 5 && match.input_size % 256 == 0; + + Dispatch contraction; + contraction.kernel = make_kernel_specialization( + use_split_direct_dot ? kMulMatSymmetricI4LowRowSplitK2DirectDotKernels[match.token_count - 1] : + split_k ? kMulMatSymmetricI4LowRowSplitK2WmmaKernel : + kMulMatSymmetricI4LowRowWmmaKernel); + contraction.kernel.compile_parameters.emplace("ggml.mul_mat.symmetric_i4.lowrow.input_size", + common_to_config_value(match.input_size)); + contraction.kernel.compile_parameters.emplace("ggml.mul_mat.symmetric_i4.lowrow.output_size", + common_to_config_value(match.output_size)); + contraction.kernel.compile_parameters.emplace("ggml.mul_mat.symmetric_i4.lowrow.token_count", + common_to_config_value(match.token_count)); + contraction.kernel.compile_parameters.emplace( + "ggml.mul_mat.symmetric_i4.lowrow.row_group_size", + common_to_config_value( + static_cast(common_symmetric_shared4_row_group_size(match.input_size, match.output_size, 4)))); + contraction.bindings.push_back( + match.weight->type == GGML_TYPE_Q5_K && match.token_count == 1 ? + common_symmetric_i4_shared4_multistart_weight_binding(*match.weight, match.input_size, match.output_size) : + common_symmetric_i4_shared4_weight_binding(*match.weight, match.input_size, match.output_size)); + contraction.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + contraction.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + contraction.bindings.push_back({ activation, 0, activation_layout.payload_bytes }); + contraction.bindings.push_back({ activation, activation_layout.scales_offset, activation_layout.metadata_bytes }); + contraction.bindings.push_back({ activation, activation_layout.sums_offset, activation_layout.metadata_bytes }); + if (split_k) { + contraction.bindings.push_back({ partial, 0, match.output->byte_count }); + contraction.bindings.push_back({ completion_counters, 0, completion_counter_count * sizeof(int32_t) }); + } + + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(std::move(contraction)); + return dispatch_match.status.success(); +} + +static bool match_q6_k_token1_shortlist_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + constexpr int64_t selected_group_count = 96; + constexpr int64_t candidate_count = 128; + constexpr int64_t refine_tile_size = 64; + constexpr size_t partition_entry_count = 1024; + + const CommonMulMatMatch match = + common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatQ6KSymmetricI2ScanToken1Kernel, true); + if (!match.matched() || !context.graph.has_index() || match.weight->type != GGML_TYPE_Q6_K || + match.token_count != 1 || match.weight->alias_source.value >= 0 || match.input_size % 256 != 0 || + match.input_size > 8192 || match.output_size < 65536 || match.output_size % 64 != 0 || + match.output_size < 16 * match.input_size || !context.graph.index().consumers(match.output->id).empty() || + !has_embedded_external_concat_projection_ancestor(context.graph, match.input->id, match.input_size)) { + return false; + } + + const CommonSymmetricI4ActivationLayout activation_layout = + common_symmetric_i4_activation_layout(match.input_size, match.token_count); + const CommandPlanAlternateValue * alternate = find_alternate_value(context.graph, context.plan, match.input->id, + GGML_TYPE_COUNT, activation_layout.total_bytes); + if (alternate != nullptr && alternate->name != kCommonSymmetricI4K32ActivationAlternateName) { + alternate = nullptr; + } + + const ValueId activation = alternate != nullptr ? alternate->alternate_value : context.next_plan_value; + const int32_t transient_base = context.next_plan_value.value + (alternate == nullptr ? 1 : 0); + const ValueId selected_groups(transient_base); + const ValueId partial_values(transient_base + 1); + const ValueId partial_ids(transient_base + 2); + const ValueId candidates(transient_base + 3); + const ValueId candidate_values(transient_base + 4); + + if (alternate == nullptr) { + dispatch_match.transients.push_back( + { activation, kCommonSymmetricI4K32ActivationAlternateName, activation_layout.total_bytes, 256 }); + + Dispatch quantize; + quantize.kernel = make_kernel_specialization(kQuantizeF32SymmetricI4K32Kernel); + quantize.kernel.compile_parameters.emplace("ggml.quantize_symmetric_i4_k32.input_size", + common_to_config_value(match.input_size)); + quantize.kernel.compile_parameters.emplace("ggml.quantize_symmetric_i4_k32.token_count", + common_to_config_value(match.token_count)); + quantize.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + quantize.bindings.push_back({ activation, 0, activation_layout.payload_bytes }); + quantize.bindings.push_back({ activation, activation_layout.scales_offset, activation_layout.metadata_bytes }); + quantize.bindings.push_back({ activation, activation_layout.sums_offset, activation_layout.metadata_bytes }); + dispatch_match.dispatches.push_back(std::move(quantize)); + + Status metadata_status; + if (!dispatch_match.metadata.append_alternate_value( + { match.input->id, activation, GGML_TYPE_COUNT, activation_layout.total_bytes, + kCommonSymmetricI4K32ActivationAlternateName }, + metadata_status)) { + dispatch_match.status.append(metadata_status); + return false; + } + } + + dispatch_match.transients.push_back({ selected_groups, "common.mul_mat.q6_k_shortlist.selected_groups", + static_cast(selected_group_count) * sizeof(int32_t), 256 }); + dispatch_match.transients.push_back( + { partial_values, "common.mul_mat.q6_k_shortlist.partial_values", partition_entry_count * sizeof(float), 256 }); + dispatch_match.transients.push_back( + { partial_ids, "common.mul_mat.q6_k_shortlist.partial_ids", partition_entry_count * sizeof(int32_t), 256 }); + dispatch_match.transients.push_back( + { candidates, "common.mul_mat.q6_k_shortlist.candidates", candidate_count * sizeof(int32_t), 256 }); + dispatch_match.transients.push_back( + { candidate_values, "common.mul_mat.q6_k_shortlist.candidate_values", candidate_count * sizeof(float), 256 }); + + Dispatch select_groups; + select_groups.kernel = make_kernel_specialization(kSelectSymmetricI4K32GroupsKernel); + select_groups.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_shortlist.input_size", + common_to_config_value(match.input_size)); + select_groups.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_shortlist.selected_group_count", + common_to_config_value(selected_group_count)); + select_groups.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + select_groups.bindings.push_back( + { selected_groups, 0, static_cast(selected_group_count) * sizeof(int32_t) }); + dispatch_match.dispatches.push_back(std::move(select_groups)); + + const size_t symmetric_i2_bytes = + common_symmetric_shared4_weight_byte_count(match.input_size, match.output_size, 2); + if (symmetric_i2_bytes > std::numeric_limits::max() - match.weight->byte_count) { + return false; + } + const size_t materialized_weight_bytes = symmetric_i2_bytes + match.weight->byte_count; + Dispatch scan; + scan.kernel = make_kernel_specialization(kMulMatQ6KSymmetricI2ScanToken1Kernel); + scan.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_shortlist.input_size", + common_to_config_value(match.input_size)); + scan.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_shortlist.output_size", + common_to_config_value(match.output_size)); + scan.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_shortlist.selected_group_count", + common_to_config_value(selected_group_count)); + scan.bindings.push_back({ match.weight->id, 0, materialized_weight_bytes, + kQ6KSymmetricI2PackedK256Row64ScaleRowLayout, match.weight->type, match.input_size, + match.output_size, match.weight->byte_count }); + scan.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + scan.bindings.push_back({ activation, 0, activation_layout.payload_bytes }); + scan.bindings.push_back({ activation, activation_layout.scales_offset, activation_layout.metadata_bytes }); + scan.bindings.push_back({ selected_groups, 0, static_cast(selected_group_count) * sizeof(int32_t) }); + dispatch_match.dispatches.push_back(std::move(scan)); + + Dispatch partition_top_k; + partition_top_k.kernel = make_kernel_specialization(kTopK8F32PartitionsRegisterKernel); + partition_top_k.kernel.integer_parameters.emplace("element_count", match.output_size); + partition_top_k.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + partition_top_k.bindings.push_back({ partial_values, 0, partition_entry_count * sizeof(float) }); + partition_top_k.bindings.push_back({ partial_ids, 0, partition_entry_count * sizeof(int32_t) }); + dispatch_match.dispatches.push_back(std::move(partition_top_k)); + + Dispatch reduce_top_k; + reduce_top_k.kernel = make_kernel_specialization(kTopK128F32ReduceGatherRegisterKernel); + reduce_top_k.kernel.integer_parameters.emplace("element_count", match.output_size); + reduce_top_k.bindings.push_back({ partial_values, 0, partition_entry_count * sizeof(float) }); + reduce_top_k.bindings.push_back({ partial_ids, 0, partition_entry_count * sizeof(int32_t) }); + reduce_top_k.bindings.push_back({ candidates, 0, candidate_count * sizeof(int32_t) }); + reduce_top_k.bindings.push_back({ candidate_values, 0, candidate_count * sizeof(float) }); + dispatch_match.dispatches.push_back(std::move(reduce_top_k)); + + Dispatch fill; + fill.kernel = make_kernel_specialization(kFillNegativeF32Kernel); + fill.kernel.integer_parameters.emplace("element_count", match.output_size); + fill.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + dispatch_match.dispatches.push_back(std::move(fill)); + + for (int64_t candidate_offset = 0; candidate_offset < candidate_count; candidate_offset += refine_tile_size) { + Dispatch refine; + refine.kernel = make_kernel_specialization(kMulMatQ6KPackedSelectedRefineToken1Kernel); + refine.kernel.integer_parameters.emplace("token_count", match.token_count); + refine.kernel.integer_parameters.emplace("candidate_count", refine_tile_size); + refine.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_packed.input_size", + common_to_config_value(match.input_size)); + refine.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_packed.output_size", + common_to_config_value(match.output_size)); + refine.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_packed.output_accumulation", "0"); + refine.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_packed.weight_offset", + common_to_config_value(static_cast(symmetric_i2_bytes))); + refine.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + refine.bindings.push_back({ match.weight->id, 0, materialized_weight_bytes, + kQ6KSymmetricI2PackedK256Row64ScaleRowLayout, match.weight->type, match.input_size, + match.output_size, match.weight->byte_count }); + refine.bindings.push_back({ candidates, static_cast(candidate_offset) * sizeof(int32_t), + static_cast(refine_tile_size) * sizeof(int32_t) }); + refine.bindings.push_back({ candidate_values, static_cast(candidate_offset) * sizeof(float), + static_cast(refine_tile_size) * sizeof(float) }); + refine.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + dispatch_match.dispatches.push_back(std::move(refine)); + } + + dispatch_match.covered_nodes.push_back(context.root_index); + return dispatch_match.status.success(); +} + +static bool build_q6_k_aligned_skinny_dispatch(const CommonMulMatMatch & match, + size_t root_index, + KernelCatalogRef kernel, + DispatchMatch & dispatch_match) { + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_packed.input_size", + common_to_config_value(match.input_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_packed.output_size", + common_to_config_value(match.output_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_packed.output_accumulation", "0"); + dispatch.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count, kQ6KPackedK256Row64ScaleRowLayout, + match.weight->type, match.input_size, match.output_size, match.weight->byte_count }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + + dispatch_match.covered_nodes.push_back(root_index); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_q6_k_aligned_skinny_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + CommonMulMatMatch match = + common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatQ6KPackedToken1F16WmmaKernel, false); + if (!match.matched()) { + match = common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatQ6KPackedToken1F16WmmaKernel, + true); + } + if (!match.matched() || !context.graph.has_index() || match.weight->type != GGML_TYPE_Q6_K || + match.weight->alias_source.value >= 0 || match.token_count > 16 || match.input_size % 256 != 0 || + match.output_size % 64 != 0 || match.output_size > kMulMatQ6KPackedMaxOutputSize || + (match.token_count > 1 && match.output_size < 8 * match.input_size) || + !context.graph.index().consumers(match.output->id).empty()) { + return false; + } + + return build_q6_k_aligned_skinny_dispatch(match, context.root_index, kMulMatQ6KPackedToken1F16WmmaKernel, + dispatch_match); +} + +static bool match_q6_k_i8_prepacked_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const CommonMulMatMatch match = + common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatQ6KI8PrepackedF16WmmaKernel, false); + if (!match.matched() || match.weight->type != GGML_TYPE_Q6_K || match.weight->alias_source.value >= 0 || + match.token_count <= 5 || match.token_count > 16 || match.input_size % 256 != 0 || match.output_size % 64 != 0 || + (match.output_size > match.input_size / 2 && match.token_count > 5)) { + return false; + } + + const size_t block_count = static_cast(match.input_size / ggml_blck_size(GGML_TYPE_Q6_K)); + size_t materialized_weight_bytes = static_cast(match.output_size) * block_count; + if (materialized_weight_bytes > std::numeric_limits::max() / size_t{ 274 }) { + return false; + } + materialized_weight_bytes *= size_t{ 274 }; + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kMulMatQ6KI8PrepackedF16WmmaKernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_packed.input_size", + common_to_config_value(match.input_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_packed.output_size", + common_to_config_value(match.output_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_packed.output_accumulation", "0"); + dispatch.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + dispatch.bindings.push_back({ match.weight->id, 0, materialized_weight_bytes, kQ6KI8K32Row64Layout, + match.weight->type, match.input_size, match.output_size, match.weight->byte_count }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool build_q6_k_prefill_wave32_dispatch(const CommonMulMatMatch & match, + ValueId input, + size_t input_bytes, + KernelCatalogRef kernel, + size_t root_index, + DispatchMatch & dispatch_match) { + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_packed.input_size", + common_to_config_value(match.input_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_packed.output_size", + common_to_config_value(match.output_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_packed.output_accumulation", "0"); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_q6_k_packed.token_capacity", + common_to_config_value(match.token_count)); + dispatch.bindings.push_back({ input, 0, input_bytes }); + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + + dispatch_match.covered_nodes.push_back(root_index); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_q6_k_prefill_wave32_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const CommonMulMatMatch match = common_match_mul_mat_any_format( + context.graph, context.root_node, kMulMatQ6KF32WmmaPrefillWave32Kernel, false); + if (!match.matched() || match.weight->type != GGML_TYPE_Q6_K || match.weight->alias_source.value >= 0 || + match.token_count < 128 || match.token_count > 2048 || match.token_count % 128 != 0 || + match.input_size % 256 != 0 || match.output_size % 64 != 0) { + return false; + } + + const bool packed_input = common_mul_mat_uses_k16_major_f16( + match.weight_format, match.input_size, match.output_size, match.token_count); + DispatchBinding activation; + const bool prepared = packed_input ? + common_prepare_k16_major_f16_input(context, *match.input, match.input_size, match.token_count, + dispatch_match, activation) : + common_prepare_f16_input(context, *match.input, match.input_size, match.token_count, + dispatch_match, activation); + if (!prepared) { + return false; + } + return build_q6_k_prefill_wave32_dispatch(match, activation.value, activation.length, + kMulMatQ6KF16WmmaPrefillWave32Kernel, context.root_index, + dispatch_match); +} + +static bool match_symmetric_i4_prefill_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + // Disabled: this symmetric matmul route causes a substantial numeric performance regression. + (void) context; + (void) dispatch_match; + return false; + + const CommonMulMatMatch match = + common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatSymmetricI4WmmaKernel, false); + if (!match.matched() || !context.graph.has_index() || match.weight->type != GGML_TYPE_Q5_K || + match.weight->alias_source.value >= 0 || match.token_count < 128 || match.token_count % 128 != 0 || + match.input_size % 64 != 0 || match.output_size < match.input_size || match.output_size % 128 != 0) { + return false; + } + + const GraphNode * input_producer = context.graph.index().producer(match.input->id); + if (input_producer != nullptr && input_producer->op == GGML_OP_GLU) { + return false; + } + if (!has_normalized_rope_consumer(context.graph, *match.output) && + !has_reshaped_unary_mul_consumer(context.graph, *match.output) && + !has_reshaped_transpose_concat_consumer(context.graph, *match.output)) { + return false; + } + + const CommonSymmetricI4ActivationLayout activation_layout = + common_symmetric_i4_activation_layout(match.input_size, match.token_count); + const ValueId activation = context.next_plan_value; + dispatch_match.transients.push_back( + { activation, "common.mul_mat.symmetric_i4_k64.activation", activation_layout.total_bytes, 256 }); + + Dispatch quantize; + quantize.kernel = make_kernel_specialization(kQuantizeF32SymmetricI4K64PlaneKernel); + quantize.kernel.compile_parameters.emplace("ggml.quantize_symmetric_i4_k64.input_size", + common_to_config_value(match.input_size)); + quantize.kernel.compile_parameters.emplace("ggml.quantize_symmetric_i4_k64.token_count", + common_to_config_value(match.token_count)); + quantize.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + quantize.bindings.push_back({ activation, 0, activation_layout.payload_bytes }); + quantize.bindings.push_back({ activation, activation_layout.scales_offset, activation_layout.metadata_bytes }); + quantize.bindings.push_back({ activation, activation_layout.sums_offset, activation_layout.metadata_bytes }); + + const size_t materialized_weight_bytes = + static_cast(match.output_size) * static_cast(match.input_size / 256) * size_t{ 144 }; + Dispatch contraction; + contraction.kernel = make_kernel_specialization(kMulMatSymmetricI4WmmaKernel); + contraction.kernel.compile_parameters.emplace("ggml.mul_mat.symmetric_i4.input_size", + common_to_config_value(match.input_size)); + contraction.kernel.compile_parameters.emplace("ggml.mul_mat.symmetric_i4.output_size", + common_to_config_value(match.output_size)); + contraction.kernel.compile_parameters.emplace("ggml.mul_mat.symmetric_i4.token_count", + common_to_config_value(match.token_count)); + contraction.bindings.push_back({ match.weight->id, 0, materialized_weight_bytes, kSymmetricI4K64Row64Layout, + match.weight->type, match.input_size, match.output_size, + match.weight->byte_count }); + contraction.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + contraction.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + contraction.bindings.push_back({ activation, 0, activation_layout.payload_bytes }); + contraction.bindings.push_back({ activation, activation_layout.scales_offset, activation_layout.metadata_bytes }); + contraction.bindings.push_back({ activation, activation_layout.sums_offset, activation_layout.metadata_bytes }); + + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(std::move(quantize)); + dispatch_match.dispatches.push_back(std::move(contraction)); + return true; +} + +static bool match_q5_k_symmetric_i8_prefill_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + // Disabled: this symmetric matmul route causes a substantial numeric performance regression. + (void) context; + (void) dispatch_match; + return false; + + const CommonMulMatMatch match = + common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatQ5KSymmetricI8WmmaKernel, false); + if (!match.matched() || !context.graph.has_index() || match.weight->type != GGML_TYPE_Q5_K || + match.weight->alias_source.value >= 0 || match.token_count < 256 || match.token_count % 256 != 0 || + match.input_size % 256 != 0 || match.output_size % 64 != 0 || match.output_size > match.input_size) { + return false; + } + + const GraphNode * input_producer = context.graph.index().producer(match.input->id); + if (input_producer != nullptr && input_producer->op == GGML_OP_GLU) { + return false; + } + + const size_t element_count = static_cast(match.token_count) * static_cast(match.input_size); + const size_t metadata_bytes = + static_cast(match.token_count) * static_cast(match.input_size / 256) * sizeof(int32_t); + const size_t quantized_bytes = element_count + metadata_bytes; + const ValueId quantized = context.next_plan_value; + dispatch_match.transients.push_back( + { quantized, "common.mul_mat.symmetric_i8_k256.activation", quantized_bytes, 256 }); + + Dispatch quantize; + quantize.kernel = make_kernel_specialization(kQuantizeF32SymmetricI8K256Kernel); + quantize.kernel.integer_parameters.emplace("token_count", match.token_count); + quantize.kernel.integer_parameters.emplace("input_size_arg", match.input_size); + quantize.kernel.compile_parameters.emplace("ggml.quantize_symmetric_i8_k256.token_count", + common_to_config_value(match.token_count)); + quantize.kernel.compile_parameters.emplace("ggml.quantize_symmetric_i8_k256.input_size", + common_to_config_value(match.input_size)); + quantize.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + quantize.bindings.push_back({ quantized, 0, quantized_bytes }); + + const size_t materialized_weight_bytes = + static_cast(match.output_size) * static_cast(match.input_size / 256) * size_t{ 258 }; + Dispatch contraction; + contraction.kernel = make_kernel_specialization(kMulMatQ5KSymmetricI8WmmaKernel); + contraction.kernel.integer_parameters.emplace("token_count", match.token_count); + contraction.kernel.compile_parameters.emplace("ggml.mul_mat.symmetric_i8.input_size", + common_to_config_value(match.input_size)); + contraction.kernel.compile_parameters.emplace("ggml.mul_mat.symmetric_i8.output_size", + common_to_config_value(match.output_size)); + contraction.kernel.compile_parameters.emplace("ggml.mul_mat.symmetric_i8.token_count", + common_to_config_value(match.token_count)); + contraction.bindings.push_back({ quantized, 0, quantized_bytes }); + contraction.bindings.push_back({ match.weight->id, 0, materialized_weight_bytes, kQ5KSymmetricI8K256Row64Layout, + match.weight->type, match.input_size, match.output_size, + match.weight->byte_count }); + contraction.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(std::move(quantize)); + dispatch_match.dispatches.push_back(std::move(contraction)); + return true; +} + + +static bool match_packed_q8_1_x4_prefill_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const CommonMulMatMatch match = common_match_mul_mat_any_format(context.graph, context.root_node, + kMulMatQ5KIQ4XSQ8_1X4WmmaToken256Kernel, false); + // A chunk with a remainder runs its 256-aligned head here and the rest on the generic + // kernel (common_append_mul_mat_token_tail), when the relaxed policy is on. Q4_K is left + // out: its head binds the weight in the packed Row64 layout, and one weight cannot be + // resident in two layouts ("conflicting resident layout requests"). + const int64_t tail_tokens = match.matched() ? match.token_count % 256 : 0; + const bool split_tail = common_q8_prefill_relaxed() && tail_tokens >= 2 && + match.matched() && match.weight->type != GGML_TYPE_Q4_K; + if (!match.matched() || !context.graph.has_index() || + (match.weight->type != GGML_TYPE_Q4_K && match.weight->type != GGML_TYPE_Q5_K && + match.weight->type != GGML_TYPE_IQ4_XS) || + match.token_count < 256 || match.token_count > 2048 || + (match.token_count % 256 != 0 && !split_tail) || + match.input_size % 256 != 0 || match.output_size % 64 != 0) { + return false; + } + const int64_t q8_tokens = match.token_count - (split_tail ? tail_tokens : 0); + + for (const GraphNode * consumer : context.graph.index().consumers(match.output->id)) { + if (!common_q8_prefill_relaxed() && match.weight->type != GGML_TYPE_Q4_K && consumer != nullptr && + consumer->op == GGML_OP_GLU) { + return false; + } + } + + const bool use_f16 = !split_tail && match.weight->type == GGML_TYPE_Q4_K && + match.weight->alias_source.value < 0 && match.output_size >= match.input_size / 4; + DispatchBinding activation; + const bool packed_input = use_f16 && common_mul_mat_uses_k16_major_f16( + match.weight_format, match.input_size, match.output_size, match.token_count, + context.plan.metadata.find_generated_resource(match.input->id, GeneratedResourceRole::F16K16Major) != nullptr); + if (use_f16) { + const bool prepared = packed_input ? + common_prepare_k16_major_f16_input(context, *match.input, match.input_size, match.token_count, + dispatch_match, activation) : + common_prepare_f16_input(context, *match.input, match.input_size, match.token_count, + dispatch_match, activation); + if (!prepared) { + return false; + } + } else if (!common_prepare_q8_1_x4_input(context, *match.input, match.input_size, q8_tokens, + dispatch_match, activation, + match.weight->type == GGML_TYPE_Q4_K || common_q8_prefill_relaxed() ? + CommonQ8ActivationPolicy::AllowStandaloneQuantize : + CommonQ8ActivationPolicy::ExistingAlternateOnly)) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(use_f16 ? kMulMatQ4KF16WmmaPrefillWave32Kernel : + kMulMatQ5KIQ4XSQ8_1X4WmmaToken256Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", q8_tokens); + dispatch.kernel.compile_parameters.emplace(use_f16 ? "ggml.mul_mat.input_size" : + "ggml.mul_mat_q8_1_x4.input_size", + common_to_config_value(match.input_size)); + dispatch.kernel.compile_parameters.emplace(use_f16 ? "ggml.mul_mat.output_size" : + "ggml.mul_mat_q8_1_x4.output_size", + common_to_config_value(match.output_size)); + dispatch.kernel.compile_parameters.emplace(use_f16 ? "ggml.workload.token_capacity" : + "ggml.mul_mat_q8_1_x4.token_capacity", + common_to_config_value(q8_tokens)); + const bool pack_q4 = match.weight_format == CommonMulMatWeightFormat::Q4K && + match.weight->alias_source.value < 0; + if (use_f16) { + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat.f16_input_layout", packed_input ? "1" : "0"); + } else { + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat_q8_1_x4.weight_format", + common_to_config_value(common_mul_mat_format_config_value(pack_q4 ? CommonMulMatWeightFormat::Q4KRow64 : + match.weight_format))); + } + dispatch.bindings.push_back(activation); + if (pack_q4) { + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count, kQ4KPackedK256Row64Layout, + match.weight->type, match.input_size, match.output_size, + match.weight->byte_count }); + } else { + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count }); + } + dispatch.bindings.push_back( + { match.output->id, 0, + split_tail ? static_cast(q8_tokens) * static_cast(match.output_size) * sizeof(float) : + match.output->byte_count }); + + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(std::move(dispatch)); + if (split_tail && !common_append_mul_mat_token_tail(match, q8_tokens, dispatch_match)) { + return false; + } + return true; +} + +static bool match_q4_k_q8_1_x4_prefill_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + if (context.root_node == nullptr || context.root_node->inputs.empty()) { + return false; + } + const Value * weight = common_graph_value(context.graph, context.root_node->inputs[0]); + return weight != nullptr && weight->type == GGML_TYPE_Q4_K && + match_packed_q8_1_x4_prefill_dispatch(context, dispatch_match); +} + +static MulMatPostOpsMatch match_mul_mat_postops(const DispatchMatchContext & context) { + MulMatPostOpsMatch match; + CommonMulMatMatch root = + common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatF32F32WmmaKernel, false); + if (!root.matched()) { + root = common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatF32F32WmmaKernel, true); + } + if (!root.matched() || !context.graph.has_index()) { + return match; + } + + const bool lowtoken_contraction = root.token_count <= 5 && root.weight->alias_source.value < 0 && + (root.weight_format == CommonMulMatWeightFormat::Q4K || + root.weight_format == CommonMulMatWeightFormat::Q6K) && + root.input_size >= 4096 && root.output_size >= 4096 && + root.output_size <= root.input_size && root.output_size % 64 == 0; + const bool token1_requires_q8_activation = root.token_count == 1 && !lowtoken_contraction; + + const Value * current = root.output; + if (lowtoken_contraction) { + const GraphNode * reshape = common_find_only_consumer_with_op(context.graph, current->id, GGML_OP_RESHAPE); + const Value * reshaped = reshape != nullptr ? common_graph_value(context.graph, reshape->output) : nullptr; + if (reshaped != nullptr && is_layout_alias_node(context.graph, *reshape) && reshaped->contiguous && + same_full_value_range(*current, *reshaped) && common_same_shape(*current, *reshaped)) { + match.layout_node = reshape; + current = reshaped; + } + } + const int add_limit = root.token_count == 1 || match.layout_node != nullptr ? 1 : 2; + for (int add_index = 0; add_index < add_limit; ++add_index) { + const GraphNode * add_node = common_find_only_consumer_with_op(context.graph, current->id, GGML_OP_ADD); + if (add_node == nullptr || !common_binary_node_is_add(*add_node)) { + break; + } + + const bool current_is_lhs = add_node->inputs[0] == current->id; + const bool current_is_rhs = add_node->inputs[1] == current->id; + if (!current_is_lhs && !current_is_rhs) { + return {}; + } + + const ValueId other_id = current_is_lhs ? add_node->inputs[1] : add_node->inputs[0]; + const Value * other = common_graph_value(context.graph, other_id); + const Value * output = common_graph_value(context.graph, add_node->output); + if (other == nullptr || output == nullptr || output->type != GGML_TYPE_F32 || !output->contiguous || + !common_same_shape(*root.output, *output)) { + return {}; + } + + if (!match.has_bias && is_bias_shape(*other, root.output_size)) { + match.bias = other; + match.bias_add_node = add_node; + match.add_nodes.push_back(add_node); + match.has_bias = true; + current = output; + continue; + } + + if (!match.has_residual && other->type == GGML_TYPE_F32 && other->contiguous && + common_same_shape(*root.output, *other)) { + match.residual_input = other; + match.residual_output = output; + match.residual_add_node = add_node; + match.add_nodes.push_back(add_node); + match.has_residual = true; + current = output; + continue; + } + + break; + } + + if (!match.has_bias && !match.has_residual) { + return {}; + } + + match.residual_output = current; + + if (match.has_bias && match.has_residual) { + match.kernel = kMulMatBiasAddF32F32WmmaKernel; + } else if (match.has_residual) { + match.kernel = kMulMatAddF32F32WmmaKernel; + } else if (match.has_bias) { + match.kernel = kMulMatBiasF32F32WmmaKernel; + } else { + return {}; + } + + match.input = root.input; + match.weight = root.weight; + match.projection_output = root.output; + match.weight_format = root.weight_format; + match.input_size = root.input_size; + match.output_size = root.output_size; + match.token_count = root.token_count; + match.requires_q8_activation = token1_requires_q8_activation; + return match; +} + +static bool try_match_fused_unary(const DispatchMatchContext & context, CommonMulMatMatch & match) { + if (!match.matched()) { + return false; + } + + const std::vector & consumers = context.graph.index().consumers(context.root_node->output); + if (consumers.size() != 1 || consumers.front() == nullptr) { + return false; + } + + const GraphNode * unary = consumers.front(); + size_t unary_index = 0; + if (!context.graph.index().node_index(unary, unary_index) || unary_index >= context.covered_nodes.size() || + context.covered_nodes[unary_index] || unary->inputs.size() != 1) { + return false; + } + + const UnaryParams * params = op_params_as(unary->params); + if (params == nullptr || !unary_kind_supported(params->op)) { + return false; + } + + const Value * unary_input = common_graph_value(context.graph, unary->inputs[0]); + const Value * unary_output = common_graph_value(context.graph, unary->output); + if (unary_input == nullptr || unary_output == nullptr || unary_input->id != match.output->id || + unary_output->type != GGML_TYPE_F32 || !unary_output->contiguous || + !common_same_shape(*match.output, *unary_output)) { + return false; + } + + match.output = unary_output; + match.output_unary_op = params->op; + match.unary_node_index = unary_index; + match.has_fused_unary = true; + return true; +} + +} // namespace + +// Q5_K/IQ4_XS prefill matmuls quantize their own q8_1 activations and may feed a +// GLU, like Q4_K. Qwen3.8-27B UD-Q4_K_XL on gfx1151: pp512 99 -> 141 tok/s, KLD vs BF16 +// unchanged (0.00717 / 0.00720). GGML_HRX_Q8_PREFILL_RELAX=0 restores the previous policy. +bool common_q8_prefill_relaxed() { + static const bool relaxed = [] { + const char * env = std::getenv("GGML_HRX_Q8_PREFILL_RELAX"); + return env == nullptr || std::atoi(env) != 0; + }(); + return relaxed; +} + +static bool build_mul_mat_dispatch(const DispatchMatchContext & context, + const CommonMulMatMatch & match, + DispatchMatch & dispatch_match, + CommonQ8ActivationPolicy q8_policy = CommonQ8ActivationPolicy::ExistingAlternateOnly) { + const bool split_k = match.weight_format == CommonMulMatWeightFormat::Q4K && !match.has_fused_unary && + match.input_size >= 4096 && match.input_size % 1024 == 0 && + match.output_size <= 64 && match.output_size % 4 == 0 && + match.token_count >= 128 && match.token_count % 32 == 0 && + common_ceil_div(match.token_count, 32) < 32; + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(split_k ? kMulMatF32F32NarrowSplitK4Kernel : match.kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.workload.token_capacity", + common_to_config_value(match.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat.input_size", common_to_config_value(match.input_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat.output_size", common_to_config_value(match.output_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat.output_accumulation", "0"); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat.output_unary_op", + std::to_string(unary_kind_config_value(match.output_unary_op))); + const bool pack_rows = match.token_count <= 5 && match.output_size % 64 == 0 && + match.weight->alias_source.value < 0; + const bool can_pack_q4 = pack_rows && match.weight_format == CommonMulMatWeightFormat::Q4K; + const bool can_pack_q6 = pack_rows && match.weight_format == CommonMulMatWeightFormat::Q6K; + bool pack_q4 = false; + bool pack_q6 = false; + DispatchBinding activation = { match.input->id, 0, match.input->byte_count }; + if (can_pack_q4 || can_pack_q6) { + if (common_prepare_q8_1_x4_input(context, *match.input, match.input_size, match.token_count, + dispatch_match, activation, q8_policy)) { + pack_q4 = can_pack_q4; + pack_q6 = can_pack_q6; + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat.activation_format", + std::to_string(GGML_TYPE_Q8_1)); + } else if (q8_policy == CommonQ8ActivationPolicy::AllowStandaloneQuantize) { + return false; + } + } + const CommonMulMatWeightFormat format = pack_q4 ? CommonMulMatWeightFormat::Q4KRow64 : + pack_q6 ? CommonMulMatWeightFormat::Q6KRow64 : match.weight_format; + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat.weight_format", + common_to_config_value(common_mul_mat_format_config_value(format))); + dispatch.bindings.push_back(activation); + if (pack_q4 || pack_q6) { + const char * layout = pack_q4 ? kQ4KPackedK256Row64Layout : kQ6KPackedK256Row64ScaleRowLayout; + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count, layout, + match.weight->type, match.input_size, match.output_size, + match.weight->byte_count }); + } else { + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count }); + } + if (split_k) { + const ValueId partial = context.next_plan_value; + const ValueId counters(partial.value + 1); + const size_t partial_bytes = 4 * match.output->byte_count; + const uint32_t counter_count = static_cast(match.token_count / 32); + dispatch_match.transients.push_back( + { partial, "common.mul_mat.narrow.split_k_partial", partial_bytes, 256 }); + dispatch_match.completion_counter_requests.push_back( + { counters, "common.mul_mat.narrow.split_k_counters", counter_count }); + dispatch.bindings.push_back({ partial, 0, partial_bytes }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + dispatch.bindings.push_back({ counters, 0, counter_count * sizeof(int32_t) }); + } else { + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + } + + dispatch_match.covered_nodes.push_back(context.root_index); + if (match.has_fused_unary) { + dispatch_match.covered_nodes.push_back(match.unary_node_index); + } + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_q6_k_token1_final_projection_q8_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + CommonMulMatMatch match = + common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatF32F32WmmaKernel, true); + if (!match.matched()) { + match = common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatF32F32WmmaKernel, false); + } + if (!match.matched() || !context.graph.has_index() || match.weight->type != GGML_TYPE_Q6_K || + match.token_count < 1 || match.token_count > 5 || match.weight->alias_source.value >= 0 || + match.input_size % 256 != 0 || match.output_size % 64 != 0 || + match.output_size > kMulMatQ6KPackedMaxOutputSize || + !context.graph.index().consumers(match.output->id).empty()) { + return false; + } + + if (match.token_count == 1 && match.output_size >= 16 * match.input_size) { + return build_q6_k_aligned_skinny_dispatch(match, context.root_index, kMulMatQ6KPackedToken1F16WmmaKernel, + dispatch_match); + } + + return build_mul_mat_dispatch(context, match, dispatch_match, + CommonQ8ActivationPolicy::AllowStandaloneQuantize); +} + +static bool match_mul_mat_postops_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const MulMatPostOpsMatch match = match_mul_mat_postops(context); + if (!match.matched() || match.input_size % 256 != 0) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(match.kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.workload.token_capacity", + common_to_config_value(match.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_postops.input_size", + common_to_config_value(match.input_size)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_postops.output_size", + common_to_config_value(match.output_size)); + const bool pack_rows = match.token_count <= 5 && match.output_size % 64 == 0 && + match.weight->alias_source.value < 0; + const bool can_pack_q4 = pack_rows && match.weight_format == CommonMulMatWeightFormat::Q4K; + const bool can_pack_q6 = pack_rows && match.weight_format == CommonMulMatWeightFormat::Q6K; + bool pack_q4 = false; + bool pack_q6 = false; + DispatchBinding activation = { match.input->id, 0, match.input->byte_count }; + if (match.requires_q8_activation && !can_pack_q4 && !can_pack_q6) { + return false; + } + if (can_pack_q4 || can_pack_q6) { + if (common_prepare_q8_1_x4_input(context, *match.input, match.input_size, match.token_count, + dispatch_match, activation, + CommonQ8ActivationPolicy::ExistingAlternateOnly)) { + pack_q4 = can_pack_q4; + pack_q6 = can_pack_q6; + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat.activation_format", + std::to_string(GGML_TYPE_Q8_1)); + } + } + if (match.requires_q8_activation && !pack_q4 && !pack_q6) { + return false; + } + const CommonMulMatWeightFormat format = pack_q4 ? CommonMulMatWeightFormat::Q4KRow64 : + pack_q6 ? CommonMulMatWeightFormat::Q6KRow64 : match.weight_format; + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat_postops.weight_format", + common_to_config_value(common_mul_mat_format_config_value(format))); + dispatch.bindings.push_back(activation); + if (pack_q4 || pack_q6) { + const char * layout = pack_q4 ? kQ4KPackedK256Row64Layout : kQ6KPackedK256Row64ScaleRowLayout; + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count, layout, + match.weight->type, match.input_size, match.output_size, + match.weight->byte_count }); + } else { + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count }); + } + if (match.has_bias) { + dispatch.bindings.push_back({ match.bias->id, 0, match.bias->byte_count }); + } + if (match.has_residual) { + dispatch.bindings.push_back({ match.residual_input->id, 0, match.residual_input->byte_count }); + } + dispatch.bindings.push_back({ match.residual_output->id, 0, match.residual_output->byte_count }); + + if (!append_covered_node_index_once(context.graph, context.covered_nodes, context.root_node, + dispatch_match.covered_nodes)) { + return false; + } + if (match.layout_node != nullptr && + !append_covered_node_index_once(context.graph, context.covered_nodes, match.layout_node, + dispatch_match.covered_nodes)) { + return false; + } + for (const GraphNode * add_node : match.add_nodes) { + if (!append_covered_node_index_once(context.graph, context.covered_nodes, add_node, + dispatch_match.covered_nodes)) { + return false; + } + } + + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static DecodeMulMatAddMatch match_decode_mul_mat_add(const DispatchMatchContext & context) { + DecodeMulMatAddMatch match; + CommonMulMatMatch root = + common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatAddF32F32DecodeWave64Kernel, true); + if (!root.matched() || !context.graph.has_index() || root.token_count != 1) { + return match; + } + + const GraphNode * add_node = common_find_only_consumer_with_op(context.graph, root.output->id, GGML_OP_ADD); + if (add_node == nullptr || !common_binary_node_is_add(*add_node)) { + return match; + } + + const bool root_is_lhs = add_node->inputs[0] == root.output->id; + const bool root_is_rhs = add_node->inputs[1] == root.output->id; + if (!root_is_lhs && !root_is_rhs) { + return match; + } + + const ValueId residual_id = root_is_lhs ? add_node->inputs[1] : add_node->inputs[0]; + const Value * residual = common_graph_value(context.graph, residual_id); + const Value * output = common_graph_value(context.graph, add_node->output); + if (residual == nullptr || output == nullptr || residual->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || + !residual->contiguous || !output->contiguous || !common_same_shape(*root.output, *residual) || + !common_same_shape(*root.output, *output)) { + return match; + } + + size_t add_index = 0; + if (!context.graph.index().node_index(add_node, add_index) || add_index >= context.covered_nodes.size() || + context.covered_nodes[add_index]) { + return match; + } + + root.output = output; + match.root = root; + match.residual_input = residual; + match.residual_output = output; + match.add_node = add_node; + return match; +} + +static bool build_decode_mul_mat_add_dispatch(const DecodeMulMatAddMatch & match, + const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kMulMatAddF32F32DecodeWave64Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.root.token_count); + dispatch.kernel.integer_parameters.emplace("input_size", match.root.input_size); + dispatch.kernel.integer_parameters.emplace("output_size", match.root.output_size); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_f32_f32_decode.token_capacity", + common_to_config_value(match.root.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_f32_f32_decode.output_capacity", + common_to_config_value(match.root.output_size)); + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat_f32_f32_decode.weight_format", + common_to_config_value(common_mul_mat_format_config_value(match.root.weight_format))); + dispatch.bindings.push_back({ match.root.input->id, 0, match.root.input->byte_count }); + dispatch.bindings.push_back({ match.root.weight->id, 0, match.root.weight->byte_count }); + dispatch.bindings.push_back({ match.residual_input->id, 0, match.residual_input->byte_count }); + dispatch.bindings.push_back({ match.residual_output->id, 0, match.residual_output->byte_count }); + + if (!append_covered_node_index_once(context.graph, context.covered_nodes, context.root_node, + dispatch_match.covered_nodes) || + !append_covered_node_index_once(context.graph, context.covered_nodes, match.add_node, + dispatch_match.covered_nodes)) { + return false; + } + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static void build_decode_mul_mat_dispatch(const CommonMulMatMatch & match, + DispatchMatch & dispatch_match, + size_t root_index) { + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(match.kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.integer_parameters.emplace("input_size", match.input_size); + dispatch.kernel.integer_parameters.emplace("output_size", match.output_size); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_f32_f32_decode.token_capacity", + common_to_config_value(match.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_f32_f32_decode.output_capacity", + common_to_config_value(match.output_size)); + dispatch.kernel.compile_parameters.emplace( + "ggml.mul_mat_f32_f32_decode.weight_format", + common_to_config_value(common_mul_mat_format_config_value(match.weight_format))); + dispatch.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + + dispatch_match.covered_nodes.push_back(root_index); + dispatch_match.dispatches.push_back(std::move(dispatch)); +} + +static bool match_mul_mat_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + CommonMulMatMatch match = + common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatF32F32WmmaKernel, false); + if (!match.matched()) { + return false; + } + try_match_fused_unary(context, match); + return build_mul_mat_dispatch(context, match, dispatch_match); +} + +static bool match_mul_mat_unary_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + CommonMulMatMatch match = + common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatF32F32WmmaKernel, false); + if (!match.matched() || !try_match_fused_unary(context, match)) { + return false; + } + return build_mul_mat_dispatch(context, match, dispatch_match); +} + +static bool match_decode_mul_mat_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + CommonMulMatMatch match = + common_match_mul_mat_any_format(context.graph, context.root_node, kMulMatF32F32DecodeWave64Kernel, true); + if (!match.matched()) { + return false; + } + if ((match.weight_format == CommonMulMatWeightFormat::Q4K || + match.weight_format == CommonMulMatWeightFormat::Q6K) && match.output_size % 64 == 0 && + common_is_supported_dense_output_size(match.output_size) && match.weight->alias_source.value < 0) { + match.kernel = kMulMatF32F32WmmaKernel; + return build_mul_mat_dispatch(context, match, dispatch_match); + } else { + build_decode_mul_mat_dispatch(match, dispatch_match, context.root_index); + } + return true; +} + +static bool match_decode_mul_mat_add_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const DecodeMulMatAddMatch match = match_decode_mul_mat_add(context); + if (!match.matched()) { + return false; + } + return build_decode_mul_mat_add_dispatch(match, context, dispatch_match); +} + + +void register_mul_mat_dispatches(DispatchRegistryBuilder & registry) { + register_mul_mat_iq3_xxs_dispatch(registry); + registry.add({ + "common.mul_mat.q6_k_token1_shortlist", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 310, + DispatchSource::Common, + match_q6_k_token1_shortlist_dispatch, + }); + registry.add({ + "common.mul_mat.q6_k_prefill_wave32", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 295, + DispatchSource::Common, + match_q6_k_prefill_wave32_dispatch, + }); + registry.add({ + "common.mul_mat.q6_k_token1_final_projection_q8", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 305, + DispatchSource::Common, + match_q6_k_token1_final_projection_q8_dispatch, + }); + registry.add({ + "common.mul_mat.q6_k_aligned_skinny", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 300, + DispatchSource::Common, + match_q6_k_aligned_skinny_dispatch, + }); + registry.add({ + "common.mul_mat.q5_k_iq4_xs_symmetric_i4_lowrow", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 285, + DispatchSource::Common, + match_symmetric_i4_low_row_dispatch, + }); + registry.add({ + "common.mul_mat.q6_k_i8_prepacked", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 285, + DispatchSource::Common, + match_q6_k_i8_prepacked_dispatch, + }); + registry.add({ + "common.mul_mat.q5_k_iq4_xs_q8_1_x4_prefill", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 260, + DispatchSource::Common, + match_packed_q8_1_x4_prefill_dispatch, + }); + registry.add({ + "common.mul_mat.q4_k_q8_1_x4_prefill", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 305, + DispatchSource::Common, + match_q4_k_q8_1_x4_prefill_dispatch, + }); + registry.add({ + "common.mul_mat.q5_k_symmetric_i4_prefill", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 270, + DispatchSource::Common, + match_symmetric_i4_prefill_dispatch, + }); + registry.add({ + "common.mul_mat.q5_k_symmetric_i8_prefill", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 280, + DispatchSource::Common, + match_q5_k_symmetric_i8_prefill_dispatch, + }); + registry.add({ + "common.mul_mat_unary.f32_f32_wmma", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 290, + DispatchSource::Common, + match_mul_mat_unary_dispatch, + }); + registry.add({ + "common.mul_mat_postops.f32_f32_wmma", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 180, + DispatchSource::Common, + match_mul_mat_postops_dispatch, + }); + registry.add({ + "common.mul_mat.f32_f32_wmma", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 80, + DispatchSource::Common, + match_mul_mat_dispatch, + }); + registry.add({ + "common.mul_mat_add.f32_f32_decode", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 70, + DispatchSource::Common, + match_decode_mul_mat_add_dispatch, + }); + registry.add({ + "common.mul_mat.f32_f32_decode", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 60, + DispatchSource::Common, + match_decode_mul_mat_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat.h new file mode 100644 index 000000000000..33a822434eaa --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_mul_mat_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-res-scale-pair.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-res-scale-pair.cpp new file mode 100644 index 000000000000..427b9c34500f --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-res-scale-pair.cpp @@ -0,0 +1,186 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// ZAYA's residual scale, residual = (x + bx) * sx + (r + br) * sr, in two dispatches +// (ops/res_scale_pair_f32.loom) instead of up to five ADD / MUL dispatches, twice per layer. At +// decode each of those is ~1.3 us of work plus a ~1.8 us gap. The graph computes the r side +// long before x exists, so one dispatch cannot cover both: each side becomes (a + b) * s where +// it starts, and the side that comes last also takes the final ADD, reading the other side's +// result as an addend. A bias ADD an earlier fusion already took (the output projection's +// bias) stays there. Every removed intermediate has exactly one consumer. + +#include "dispatch-res-scale-pair.h" + +#include "dispatch-mul-mat-common.h" +#include "graph/graph-matcher.h" + +#include + +namespace ggml::hrx { + +namespace { + +static constexpr KernelCatalogRef kResScalePairKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_res_scale_pair_f32"); + +// A full activation: packed f32, [row_size, rows] with nothing beyond dim 1. +bool is_activation(const Graph & graph, const Value * value) { + if (value == nullptr || value->type != GGML_TYPE_F32 || !value->contiguous || value->ne[0] < 1 || + value->ne[2] != 1 || value->ne[3] != 1 || value->element_count != value->ne[0] * value->ne[1]) { + return false; + } + if (value->alias_source.value < 0) { + return true; + } + const GraphNode * producer = graph.index().producer(value->id); + return value->storage_offset == 0 && producer != nullptr && producer->op == GGML_OP_RESHAPE; +} + +// The output may reuse x's or r's storage (each element is read and written by one thread at +// one index), but only exactly in place: a shifted overlap would race. +bool in_place_or_disjoint(const Value & source, const Value & output) { + if (source.storage != output.storage || source.storage_offset == output.storage_offset) { + return true; + } + return source.storage_offset + source.byte_count <= output.storage_offset || + output.storage_offset + output.byte_count <= source.storage_offset; +} + +// A per-channel parameter: packed f32 [row_size], broadcast over rows. +bool is_channel_parameter(const Value * value, int64_t row_size) { + return value != nullptr && value->type == GGML_TYPE_F32 && value->contiguous && value->ne[0] == row_size && + value->element_count == row_size && value->alias_source.value < 0; +} + +// node = BINARY(activation, parameter) with the activation as input 0 (the order ZAYA builds). +bool split_scaled(const Graph & graph, const GraphNode & node, int64_t row_size, const Value *& activation, + const Value *& parameter) { + activation = common_graph_value(graph, node.inputs[0]); + parameter = common_graph_value(graph, node.inputs[1]); + return is_activation(graph, activation) && activation->ne[0] == row_size && + is_channel_parameter(parameter, row_size); +} + +bool covered(const DispatchMatchContext & context, const GraphNode * node) { + size_t index = 0; + return node != nullptr && context.graph.index().node_index(node, index) && index < context.covered_nodes.size() && + context.covered_nodes[index]; +} + +// Rooted at a bias ADD (a + b, feeding only a MUL) or at the MUL (a * s) itself. +bool match_res_scale_pair(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const Graph & graph = context.graph; + const GraphNode * root = context.root_node; + if (root == nullptr || !graph.has_index()) { + return false; + } + + std::vector nodes; + const Value * a = nullptr; + const Value * bias = nullptr; + const Value * s = nullptr; + const GraphNode * mul = root; + if (common_binary_node_is_add(*root)) { + const Value * out = common_graph_value(graph, root->output); + mul = common_find_only_consumer_with_op(graph, root->output, GGML_OP_MUL); + if (out == nullptr || mul == nullptr || mul->inputs[0] != root->output || + !split_scaled(graph, *root, out->ne[0], a, bias)) { + return false; + } + nodes.push_back(root); + } + const Value * mul_out = common_graph_value(graph, mul->output); + const Value * scaled = nullptr; + if (mul_out == nullptr || !common_binary_node_is_mul(*mul) || !split_scaled(graph, *mul, mul_out->ne[0], scaled, s)) { + return false; + } + if (a == nullptr) { + a = scaled; + } + nodes.push_back(mul); + + // Take the final ADD when the other side is already computed (its producer covered, or a leaf). + const Value * output = mul_out; + const Value * addend = nullptr; + const GraphNode * sum = common_find_only_consumer_with_op(graph, mul->output, GGML_OP_ADD); + if (sum != nullptr && common_binary_node_is_add(*sum)) { + const ValueId other_id = sum->inputs[0] == mul->output ? sum->inputs[1] : sum->inputs[0]; + const Value * other = common_graph_value(graph, other_id); + const Value * sum_out = common_graph_value(graph, sum->output); + const GraphNode * other_producer = graph.index().producer(other_id); + if (other != nullptr && sum_out != nullptr && other_id != mul->output && + (other_producer == nullptr || covered(context, other_producer)) && is_activation(graph, other) && + is_activation(graph, sum_out) && sum_out->alias_source.value < 0 && common_same_shape(*other, *sum_out) && + common_same_shape(*sum_out, *a) && in_place_or_disjoint(*other, *sum_out)) { + addend = other; + output = sum_out; + nodes.push_back(sum); + } + } + if (!is_activation(graph, output) || output->alias_source.value >= 0 || !common_same_shape(*output, *a) || + !in_place_or_disjoint(*a, *output)) { + return false; + } + // Nothing to gain from a lone MUL. + if (nodes.size() < 2) { + return false; + } + for (const GraphNode * node : nodes) { + if (!append_covered_node_index_once(graph, context.covered_nodes, node, dispatch_match.covered_nodes)) { + dispatch_match.covered_nodes.clear(); + return false; + } + } + + const auto source = [](const Value * value) { + return DispatchBinding{ value->storage_root, value->storage_offset, value->byte_count }; + }; + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kResScalePairKernel); + dispatch.kernel.integer_parameters.emplace("element_count", output->element_count); + dispatch.kernel.integer_parameters.emplace("row_size", output->ne[0]); + dispatch.kernel.integer_parameters.emplace("row_count", output->ne[1]); + dispatch.kernel.compile_parameters.emplace("ggml.res_scale_pair_f32.has_bias", bias != nullptr ? "1" : "0"); + dispatch.kernel.compile_parameters.emplace("ggml.res_scale_pair_f32.has_addend", addend != nullptr ? "1" : "0"); + dispatch.bindings.push_back(source(a)); + dispatch.bindings.push_back(source(bias != nullptr ? bias : s)); // unread without has_bias + dispatch.bindings.push_back(source(s)); + dispatch.bindings.push_back(source(addend != nullptr ? addend : a)); // unread without has_addend + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_res_scale_pair_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "common.res_scale_pair.add_root_f32", + GGML_OP_ADD, + DispatchMatchKind::Fused, + 250, + DispatchSource::Common, + match_res_scale_pair, + }); + registry.add({ + "common.res_scale_pair.mul_root_f32", + GGML_OP_MUL, + DispatchMatchKind::Fused, + 250, + DispatchSource::Common, + match_res_scale_pair, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-res-scale-pair.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-res-scale-pair.h new file mode 100644 index 000000000000..f9a7d704faf6 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-res-scale-pair.h @@ -0,0 +1,24 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +namespace ggml::hrx { + +void register_res_scale_pair_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-rmsnorm.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-rmsnorm.cpp new file mode 100644 index 000000000000..6641689e6d19 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-rmsnorm.cpp @@ -0,0 +1,1069 @@ +#include "dispatch-rmsnorm.h" + +#include "dispatch-layout-utils.h" +#include "dispatch-mul-mat-common.h" +#include "ggml.h" +#include "graph/graph-matcher.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kRmsNormBinaryF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_rmsnorm_binary_f32"); +static constexpr KernelCatalogRef kRmsNormBinaryF32K16Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_rmsnorm_binary_f32_k16"); +static constexpr KernelCatalogRef kRmsNormBinaryQ8_1X4Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_rmsnorm_binary_q8_1_x4"); +static constexpr KernelCatalogRef kRmsNormF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_rmsnorm_f32"); +static constexpr KernelCatalogRef kRmsNormMulRopeF32Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_rmsnorm_mul_rope_f32"); +static constexpr KernelCatalogRef kAddRmsNormBinarySymmetricI4K32Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_add_rmsnorm_binary_symmetric_i4_k32"); +static constexpr KernelCatalogRef kRmsNormBinarySymmetricI4K32Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_rmsnorm_binary_symmetric_i4_k32"); +static constexpr KernelCatalogRef kRmsNormGateSiluMulSymmetricI4K32Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32"); +static constexpr KernelCatalogRef kRmsNormGateF32F16Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_rmsnorm_gate_f32_f16"); +static constexpr KernelCatalogRef kRmsNormGateF32Q8_1X4Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_rmsnorm_gate_f32_q8_1_x4"); + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static bool same_shape(const Value & lhs, const Value & rhs) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (lhs.ne[i] != rhs.ne[i]) { + return false; + } + } + return true; +} + +static bool is_supported_f32_hidden_size(int64_t hidden_size) { + return hidden_size >= 64 && hidden_size <= 32768 && hidden_size % 64 == 0; +} + +static bool is_supported_binary_f32_hidden_size(int64_t hidden_size) { + return hidden_size >= 128 && hidden_size <= 32768 && hidden_size % 128 == 0; +} + +static bool is_supported_symmetric_i4_hidden_size(int64_t hidden_size) { + return hidden_size >= 128 && hidden_size <= 32768 && hidden_size % 128 == 0; +} + +static bool is_supported_token_count(int64_t token_count) { + return token_count >= 1 && token_count <= (1 << 20); +} + +static bool supported_rmsnorm_input_layout(const Value & input, + int64_t hidden_size, + int64_t token_count, + int64_t & input_stride, + size_t & input_span_bytes) { + if (input.type != GGML_TYPE_F32 || input.ne[0] != hidden_size || input.nb[0] != sizeof(float)) { + return false; + } + + if (input.contiguous) { + input_stride = hidden_size; + input_span_bytes = input.byte_count; + return true; + } + + if (input.ne[1] != token_count || input.ne[2] != 1 || input.ne[3] != 1 || + input.nb[1] % sizeof(float) != 0) { + return false; + } + input_stride = static_cast(input.nb[1] / sizeof(float)); + if (input_stride < hidden_size || input_stride > 1048576) { + return false; + } + + return strided_f32_storage_span_bytes(input, input_span_bytes); +} + +static bool is_binary_op(ggml_op op) { + return op == GGML_OP_ADD || op == GGML_OP_SUB || op == GGML_OP_MUL || op == GGML_OP_DIV; +} + +static bool binary_kind_requires_order(BinaryKind kind) { + return kind == BinaryKind::Sub || kind == BinaryKind::Div; +} + +static bool is_packed_q8_consumer(const Graph & graph, const GraphNode * consumer, const Value & input) { + if (consumer == nullptr || consumer->op != GGML_OP_MUL_MAT || consumer->inputs.size() != 2 || + consumer->inputs[1] != input.id) { + return false; + } + const Value * weight = graph_value(graph, consumer->inputs[0]); + const Value * output = graph_value(graph, consumer->output); + if (weight == nullptr || output == nullptr || input.type != GGML_TYPE_F32 || + output->type != GGML_TYPE_F32 || !input.contiguous || !weight->contiguous || !output->contiguous) { + return false; + } + const bool prefill = (weight->type == GGML_TYPE_Q5_K || weight->type == GGML_TYPE_IQ4_XS) && + input.ne[1] >= 256 && input.ne[1] <= 2048 && input.ne[1] % 256 == 0; + const bool decode = input.ne[1] >= 1 && input.ne[1] <= 5 && weight->alias_source.value < 0 && + (weight->type == GGML_TYPE_Q4_K || + (weight->type == GGML_TYPE_Q6_K && !graph.index().consumers(output->id).empty())); + return (prefill || decode) && input.ne[0] >= 256 && input.ne[0] <= 32768 && input.ne[0] % 256 == 0 && + input.ne[2] == 1 && input.ne[3] == 1 && + weight->ne[0] == input.ne[0] && weight->ne[1] >= 64 && weight->ne[1] <= 262144 && weight->ne[1] % 64 == 0 && + weight->ne[2] == 1 && weight->ne[3] == 1 && output->ne[0] == weight->ne[1] && output->ne[1] == input.ne[1] && + output->ne[2] == 1 && output->ne[3] == 1; +} + +static bool has_packed_q8_consumer(const Graph & graph, const Value & value) { + if (!graph.has_index()) { + return false; + } + for (const GraphNode * consumer : graph.index().consumers(value.id)) { + if (is_packed_q8_consumer(graph, consumer, value)) { + return true; + } + } + return false; +} + +static size_t q8_1_x4_byte_count(int64_t token_count, int64_t hidden_size) { + if (token_count <= 0 || hidden_size <= 0) { + return 0; + } + return static_cast(token_count) * ggml_row_size(GGML_TYPE_Q8_1, hidden_size); +} + +static bool is_weight_shape(const Value & weight, int64_t hidden_size) { + if (weight.ne[0] != hidden_size) { + return false; + } + for (int i = 1; i < GGML_MAX_DIMS; ++i) { + if (weight.ne[i] != 1) { + return false; + } + } + return true; +} + +struct RmsNormBinaryMatch { + const GraphNode * rms_node = nullptr; + const GraphNode * binary_node = nullptr; + const Value * input = nullptr; + const Value * rhs = nullptr; + const Value * output = nullptr; + size_t rms_node_index = 0; + size_t binary_node_index = 0; + int64_t hidden_size = 0; + int64_t token_count = 0; + BinaryKind op = BinaryKind::Add; + float epsilon = 0.0f; + + bool matched() const { + return rms_node != nullptr && binary_node != nullptr && input != nullptr && rhs != nullptr && output != nullptr; + } +}; + +struct RmsNormMatch { + const GraphNode * rms_node = nullptr; + const Value * input = nullptr; + const Value * output = nullptr; + size_t input_span = 0; + size_t rms_node_index = 0; + int64_t input_stride = 0; + int64_t hidden_size = 0; + int64_t token_count = 0; + float epsilon = 0.0f; + + bool matched() const { return rms_node != nullptr && input != nullptr && output != nullptr; } +}; + +struct RmsNormMulRopeMatch { + const GraphNode * rms = nullptr; + const GraphNode * binary = nullptr; + const GraphNode * rope = nullptr; + const Value * input = nullptr; + const Value * weight = nullptr; + const Value * positions = nullptr; + const Value * output = nullptr; + const RmsNormParams * rms_params = nullptr; + const RopeParams * rope_params = nullptr; + + bool matched() const { + return rms != nullptr && binary != nullptr && rope != nullptr && input != nullptr && weight != nullptr && + positions != nullptr && output != nullptr && rms_params != nullptr && rope_params != nullptr; + } +}; + +struct AddRmsNormBinarySymmetricI4Match { + const GraphNode * add_node = nullptr; + const Value * lhs = nullptr; + const Value * rhs = nullptr; + const Value * residual = nullptr; + RmsNormBinaryMatch rms_binary; + size_t add_node_index = 0; + + bool matched() const { + return add_node != nullptr && lhs != nullptr && rhs != nullptr && residual != nullptr && rms_binary.matched(); + } +}; + +struct RmsNormGateMatch { + std::vector covered; + const Value * input = nullptr; + const Value * weight = nullptr; + const Value * raw_gate = nullptr; + const Value * output = nullptr; + int64_t hidden_size = 0; + int64_t token_count = 0; + float epsilon = 0.0f; + UnaryKind gate_op = UnaryKind::Silu; + + bool matched() const { + return !covered.empty() && input != nullptr && weight != nullptr && raw_gate != nullptr && output != nullptr; + } +}; + +template static bool pairwise_distinct_storage_roots(const std::array & values) { + for (size_t lhs = 0; lhs < values.size(); ++lhs) { + for (size_t rhs = lhs + 1; rhs < values.size(); ++rhs) { + if (values[lhs]->storage_root == values[rhs]->storage_root) { + return false; + } + } + } + return true; +} + +static RmsNormMulRopeMatch match_rmsnorm_mul_rope_f32(const Graph & graph, const GraphNode * rms) { + RmsNormMulRopeMatch match; + if (rms == nullptr || rms->op != GGML_OP_RMS_NORM || rms->inputs.size() != 1 || !graph.has_index()) { + return match; + } + + const std::vector & rms_consumers = graph.index().consumers(rms->output); + if (rms_consumers.size() != 1 || rms_consumers.front() == nullptr || rms_consumers.front()->op != GGML_OP_MUL || + rms_consumers.front()->inputs.size() != 2) { + return {}; + } + const GraphNode * binary = rms_consumers.front(); + const std::vector & binary_consumers = graph.index().consumers(binary->output); + if (binary_consumers.size() != 1 || binary_consumers.front() == nullptr || + binary_consumers.front()->op != GGML_OP_ROPE || binary_consumers.front()->inputs.size() != 2 || + binary_consumers.front()->inputs[0] != binary->output) { + return {}; + } + const GraphNode * rope = binary_consumers.front(); + + const Value * input = graph_value(graph, rms->inputs[0]); + const Value * rms_output = graph_value(graph, rms->output); + const Value * binary_output = graph_value(graph, binary->output); + const Value * positions = graph_value(graph, rope->inputs[1]); + const Value * output = graph_value(graph, rope->output); + const Value * weight = nullptr; + if (binary->inputs[0] == rms->output) { + weight = graph_value(graph, binary->inputs[1]); + } else if (binary->inputs[1] == rms->output) { + weight = graph_value(graph, binary->inputs[0]); + } + const RmsNormParams * rms_params = op_params_as(rms->params); + const RopeParams * rope_params = op_params_as(rope->params); + if (input == nullptr || rms_output == nullptr || binary_output == nullptr || positions == nullptr || + output == nullptr || weight == nullptr || rms_params == nullptr || rope_params == nullptr) { + return {}; + } + + if (input->type != GGML_TYPE_F32 || rms_output->type != GGML_TYPE_F32 || binary_output->type != GGML_TYPE_F32 || + output->type != GGML_TYPE_F32 || weight->type != GGML_TYPE_F32 || positions->type != GGML_TYPE_I32 || + input->ne != rms_output->ne || input->ne != binary_output->ne || input->ne != output->ne || + !rms_output->contiguous || !binary_output->contiguous || !output->contiguous || !weight->contiguous || + !positions->contiguous || weight->ne[0] != input->ne[0] || weight->ne[1] != 1 || weight->ne[2] != 1 || + weight->ne[3] != 1 || input->ne[0] < 2 || input->ne[0] > 512 || input->ne[0] % 2 != 0 || input->ne[1] < 1 || + input->ne[1] > 1048576 || input->ne[2] < 1 || input->ne[2] > 1048576 || input->ne[3] < 1 || + input->ne[3] > 1024 || input->nb[0] != sizeof(float) || output->nb[0] != sizeof(float) || + input->ne[2] > std::numeric_limits::max() / 4 || positions->element_count < 4 * input->ne[2]) { + return {}; + } + for (int i = 1; i < GGML_MAX_DIMS; ++i) { + if (input->nb[i] % sizeof(float) != 0 || output->nb[i] % sizeof(float) != 0) { + return {}; + } + } + + const bool sectioned_mode = rope_params->mode == GGML_ROPE_TYPE_MROPE || rope_params->mode == GGML_ROPE_TYPE_IMROPE; + int64_t section_count = 0; + for (int section : rope_params->sections) { + if (section < 0) { + return {}; + } + section_count += section; + } + if (!sectioned_mode || rope_params->n_dims < 2 || rope_params->n_dims > input->ne[0] || + rope_params->n_dims % 2 != 0 || section_count != rope_params->n_dims / 2 || !std::isfinite(rms_params->eps) || + rms_params->eps <= 0.0f || !std::isfinite(rope_params->freq_base) || rope_params->freq_base <= 0.0f || + !std::isfinite(rope_params->freq_scale) || rope_params->freq_scale <= 0.0f || + !std::isfinite(rope_params->attn_factor) || rope_params->attn_factor <= 0.0f || + rope_params->ext_factor != 0.0f || output->storage_root == input->storage_root || + output->storage_root == weight->storage_root || output->storage_root == positions->storage_root) { + return {}; + } + + match.rms = rms; + match.binary = binary; + match.rope = rope; + match.input = input; + match.weight = weight; + match.positions = positions; + match.output = output; + match.rms_params = rms_params; + match.rope_params = rope_params; + return match; +} + +static RmsNormMatch match_rmsnorm_f32(const Graph & graph, const GraphNode * node, size_t node_index) { + RmsNormMatch match; + if (node == nullptr || node->op != GGML_OP_RMS_NORM || node->inputs.size() != 1) { + return match; + } + + const RmsNormParams * rms_params = op_params_as(node->params); + if (rms_params == nullptr || !std::isfinite(rms_params->eps) || rms_params->eps <= 0.0f) { + return {}; + } + + const Value * input = graph_value(graph, node->inputs[0]); + const Value * output = graph_value(graph, node->output); + if (input == nullptr || output == nullptr) { + return {}; + } + if (input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || !output->contiguous || + !same_shape(*input, *output)) { + return {}; + } + if (input->storage_root == output->storage_root) { + return {}; + } + + const int64_t hidden_size = output->ne[0]; + if (!is_supported_f32_hidden_size(hidden_size)) { + return {}; + } + if (hidden_size == 0 || output->element_count <= 0 || output->element_count % hidden_size != 0) { + return {}; + } + const int64_t token_count = output->element_count / hidden_size; + if (!is_supported_token_count(token_count)) { + return {}; + } + int64_t input_stride = 0; + size_t input_span = 0; + if (!supported_rmsnorm_input_layout(*input, hidden_size, token_count, input_stride, input_span)) { + return {}; + } + + match.rms_node = node; + match.input = input; + match.output = output; + match.input_span = input_span; + match.rms_node_index = node_index; + match.input_stride = input_stride; + match.hidden_size = hidden_size; + match.token_count = token_count; + match.epsilon = rms_params->eps; + return match; +} + +static RmsNormBinaryMatch match_rmsnorm_binary_f32(const Graph & graph, const GraphNode * node, size_t node_index) { + RmsNormBinaryMatch match; + if (node == nullptr || node->op != GGML_OP_RMS_NORM || node->inputs.size() != 1 || !graph.has_index()) { + return match; + } + + const RmsNormParams * rms_params = op_params_as(node->params); + if (rms_params == nullptr || !std::isfinite(rms_params->eps) || rms_params->eps <= 0.0f) { + return {}; + } + const std::vector & consumers = graph.index().consumers(node->output); + if (consumers.size() != 1) { + return {}; + } + const GraphNode * binary_node = consumers.front(); + size_t binary_node_index; + if (binary_node == nullptr || !is_binary_op(binary_node->op) || binary_node->inputs.size() != 2 || + !graph.index().node_index(binary_node, binary_node_index)) { + return {}; + } + const BinaryParams * binary_params = op_params_as(binary_node->params); + if (binary_params == nullptr || !binary_kind_supported(binary_params->op)) { + return {}; + } + + const bool rms_is_lhs = binary_node->inputs[0] == node->output; + const bool rms_is_rhs = binary_node->inputs[1] == node->output; + if (!rms_is_lhs && !rms_is_rhs) { + return {}; + } + if (rms_is_rhs && binary_kind_requires_order(binary_params->op)) { + return {}; + } + + const ValueId rhs_id = rms_is_lhs ? binary_node->inputs[1] : binary_node->inputs[0]; + const Value * input = graph_value(graph, node->inputs[0]); + const Value * rms = graph_value(graph, node->output); + const Value * rhs = graph_value(graph, rhs_id); + const Value * output = graph_value(graph, binary_node->output); + if (input == nullptr || rms == nullptr || rhs == nullptr || output == nullptr) { + return {}; + } + if (input->type != GGML_TYPE_F32 || rms->type != GGML_TYPE_F32 || rhs->type != GGML_TYPE_F32 || + output->type != GGML_TYPE_F32) { + return {}; + } + if (!input->contiguous || !rms->contiguous || !rhs->contiguous || !output->contiguous) { + return {}; + } + if (!same_shape(*input, *rms) || !same_shape(*input, *output)) { + return {}; + } + + const int64_t hidden_size = output->ne[0]; + if (!is_supported_binary_f32_hidden_size(hidden_size) || !is_weight_shape(*rhs, hidden_size)) { + return {}; + } + if (hidden_size == 0 || output->element_count <= 0 || output->element_count % hidden_size != 0) { + return {}; + } + const int64_t token_count = output->element_count / hidden_size; + if (!is_supported_token_count(token_count)) { + return {}; + } + + match.rms_node = node; + match.binary_node = binary_node; + match.input = input; + match.rhs = rhs; + match.output = output; + match.rms_node_index = node_index; + match.binary_node_index = binary_node_index; + match.hidden_size = hidden_size; + match.token_count = token_count; + match.op = binary_params->op; + match.epsilon = rms_params->eps; + return match; +} + +static AddRmsNormBinarySymmetricI4Match match_add_rmsnorm_binary_symmetric_i4(const Graph & graph, + const GraphNode * node, + size_t node_index) { + AddRmsNormBinarySymmetricI4Match match; + if (node == nullptr || node->op != GGML_OP_ADD || node->inputs.size() != 2 || !graph.has_index()) { + return match; + } + + const Value * lhs = graph_value(graph, node->inputs[0]); + const Value * rhs = graph_value(graph, node->inputs[1]); + const Value * residual = graph_value(graph, node->output); + if (lhs == nullptr || rhs == nullptr || residual == nullptr || lhs->type != GGML_TYPE_F32 || + rhs->type != GGML_TYPE_F32 || residual->type != GGML_TYPE_F32 || !lhs->contiguous || !rhs->contiguous || + !residual->contiguous || !same_shape(*lhs, *residual) || !same_shape(*rhs, *residual)) { + return {}; + } + + const GraphNode * rms_node = nullptr; + for (const GraphNode * consumer : graph.index().consumers(residual->id)) { + if (consumer == nullptr || consumer->op != GGML_OP_RMS_NORM) { + continue; + } + if (rms_node != nullptr) { + return {}; + } + rms_node = consumer; + } + size_t rms_node_index = 0; + if (rms_node == nullptr || !graph.index().node_index(rms_node, rms_node_index)) { + return {}; + } + + RmsNormBinaryMatch rms_binary = match_rmsnorm_binary_f32(graph, rms_node, rms_node_index); + if (!rms_binary.matched() || rms_binary.input->id != residual->id || rms_binary.op != BinaryKind::Mul || + !is_supported_symmetric_i4_hidden_size(rms_binary.hidden_size) || + !common_has_symmetric_i4_lowrow_consumer(graph, *rms_binary.output) || + !pairwise_distinct_storage_roots( + std::array{ lhs, rhs, residual, rms_binary.rhs, rms_binary.output })) { + return {}; + } + + match.add_node = node; + match.lhs = lhs; + match.rhs = rhs; + match.residual = residual; + match.rms_binary = rms_binary; + match.add_node_index = node_index; + return match; +} + +static RmsNormGateMatch match_rmsnorm_gate(const Graph & graph, + const GraphNode * node, + size_t node_index) { + RmsNormGateMatch match; + const RmsNormBinaryMatch rms_binary = match_rmsnorm_binary_f32(graph, node, node_index); + if (!rms_binary.matched() || rms_binary.op != BinaryKind::Mul || rms_binary.hidden_size % 64 != 0) { + return match; + } + + const GraphNode * terminal = common_find_only_consumer_with_op(graph, rms_binary.output->id, GGML_OP_MUL); + if (terminal == nullptr || !common_binary_node_is_mul(*terminal)) { + return {}; + } + const ValueId activated_id = + terminal->inputs[0] == rms_binary.output->id ? terminal->inputs[1] : terminal->inputs[0]; + const Value * activated = graph_value(graph, activated_id); + const GraphNode * unary = activated != nullptr ? graph.index().producer(activated->id) : nullptr; + const UnaryParams * unary_params = unary != nullptr ? op_params_as(unary->params) : nullptr; + if (unary == nullptr || unary->op != GGML_OP_UNARY || unary->inputs.size() != 1 || unary_params == nullptr || + !unary_kind_supported(unary_params->op) || + common_find_only_consumer_with_op(graph, activated->id, GGML_OP_MUL) != terminal) { + return {}; + } + + const Value * gate_input = graph_value(graph, unary->inputs[0]); + const GraphNode * gate_reshape = gate_input != nullptr ? graph.index().producer(gate_input->id) : nullptr; + const bool has_gate_reshape = + gate_reshape != nullptr && gate_reshape->op == GGML_OP_RESHAPE && gate_reshape->inputs.size() == 1; + const Value * raw_gate = has_gate_reshape ? graph_value(graph, gate_reshape->inputs[0]) : gate_input; + const Value * output = graph_value(graph, terminal->output); + if (gate_input == nullptr || raw_gate == nullptr || output == nullptr || gate_input->type != GGML_TYPE_F32 || + activated->type != GGML_TYPE_F32 || raw_gate->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || + !gate_input->contiguous || !activated->contiguous || !raw_gate->contiguous || !output->contiguous || + !same_shape(*rms_binary.input, *gate_input) || !same_shape(*rms_binary.input, *activated) || + !same_shape(*rms_binary.input, *output) || raw_gate->element_count != output->element_count || + graph.index().consumers(gate_input->id).size() != 1 || + !pairwise_distinct_storage_roots( + std::array{ rms_binary.input, rms_binary.rhs, raw_gate, output })) { + return {}; + } + + match.covered = { rms_binary.rms_node, rms_binary.binary_node }; + if (has_gate_reshape) { + match.covered.push_back(gate_reshape); + } + match.covered.push_back(unary); + match.covered.push_back(terminal); + match.input = rms_binary.input; + match.weight = rms_binary.rhs; + match.raw_gate = raw_gate; + match.output = output; + match.hidden_size = rms_binary.hidden_size; + match.token_count = rms_binary.token_count; + match.epsilon = rms_binary.epsilon; + match.gate_op = unary_params->op; + return match; +} + +static std::string to_config_value(int64_t value) { + return std::to_string(value); +} + +static std::string to_config_value(float value) { + std::ostringstream out; + out.precision(9); + out << value; + return out.str(); +} + +} // namespace + +static bool match_rmsnorm_mul_rope_f32_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + const RmsNormMulRopeMatch fused = match_rmsnorm_mul_rope_f32(context.graph, context.root_node); + if (!fused.matched()) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kRmsNormMulRopeF32Kernel); + auto & config = dispatch.kernel.compile_parameters; + config.emplace("ggml.rmsnorm_mul_rope.hidden_size", to_config_value(fused.input->ne[0])); + config.emplace("ggml.rmsnorm_mul_rope.ne1", to_config_value(fused.input->ne[1])); + config.emplace("ggml.rmsnorm_mul_rope.ne2", to_config_value(fused.input->ne[2])); + config.emplace("ggml.rmsnorm_mul_rope.ne3", to_config_value(fused.input->ne[3])); + config.emplace("ggml.rmsnorm_mul_rope.input_stride1", + to_config_value(static_cast(fused.input->nb[1] / sizeof(float)))); + config.emplace("ggml.rmsnorm_mul_rope.input_stride2", + to_config_value(static_cast(fused.input->nb[2] / sizeof(float)))); + config.emplace("ggml.rmsnorm_mul_rope.input_stride3", + to_config_value(static_cast(fused.input->nb[3] / sizeof(float)))); + config.emplace("ggml.rmsnorm_mul_rope.output_stride1", + to_config_value(static_cast(fused.output->nb[1] / sizeof(float)))); + config.emplace("ggml.rmsnorm_mul_rope.output_stride2", + to_config_value(static_cast(fused.output->nb[2] / sizeof(float)))); + config.emplace("ggml.rmsnorm_mul_rope.output_stride3", + to_config_value(static_cast(fused.output->nb[3] / sizeof(float)))); + config.emplace("ggml.rmsnorm_mul_rope.n_dims", to_config_value(static_cast(fused.rope_params->n_dims))); + config.emplace("ggml.rmsnorm_mul_rope.section0", + to_config_value(static_cast(fused.rope_params->sections[0]))); + config.emplace("ggml.rmsnorm_mul_rope.section1", + to_config_value(static_cast(fused.rope_params->sections[1]))); + config.emplace("ggml.rmsnorm_mul_rope.section2", + to_config_value(static_cast(fused.rope_params->sections[2]))); + config.emplace("ggml.rmsnorm_mul_rope.section3", + to_config_value(static_cast(fused.rope_params->sections[3]))); + config.emplace("ggml.rmsnorm_mul_rope.mode", to_config_value(static_cast(fused.rope_params->mode))); + config.emplace("ggml.rmsnorm_mul_rope.workgroup_size", "256"); + config.emplace("ggml.rmsnorm_mul_rope.epsilon", to_config_value(fused.rms_params->eps)); + config.emplace("ggml.rmsnorm_mul_rope.freq_base", to_config_value(fused.rope_params->freq_base)); + config.emplace("ggml.rmsnorm_mul_rope.freq_scale", to_config_value(fused.rope_params->freq_scale)); + config.emplace("ggml.rmsnorm_mul_rope.attn_factor", to_config_value(fused.rope_params->attn_factor)); + dispatch.bindings.push_back({ fused.input->id, 0, fused.input->byte_count }); + dispatch.bindings.push_back({ fused.weight->id, 0, fused.weight->byte_count }); + dispatch.bindings.push_back({ fused.positions->id, 0, fused.positions->byte_count }); + dispatch.bindings.push_back({ fused.output->id, 0, fused.output->byte_count }); + + match.covered_nodes.push_back(context.root_index); + if (!append_covered_node_index_once(context.graph, context.covered_nodes, fused.binary, match.covered_nodes) || + !append_covered_node_index_once(context.graph, context.covered_nodes, fused.rope, match.covered_nodes)) { + return false; + } + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_add_rmsnorm_binary_symmetric_i4_dispatch(const DispatchMatchContext & context, + DispatchMatch & match) { + const AddRmsNormBinarySymmetricI4Match fused = + match_add_rmsnorm_binary_symmetric_i4(context.graph, context.root_node, context.root_index); + if (!fused.matched() || fused.add_node_index >= context.covered_nodes.size() || + fused.rms_binary.rms_node_index >= context.covered_nodes.size() || + fused.rms_binary.binary_node_index >= context.covered_nodes.size() || + context.covered_nodes[fused.add_node_index] || context.covered_nodes[fused.rms_binary.rms_node_index] || + context.covered_nodes[fused.rms_binary.binary_node_index]) { + return false; + } + + const CommonSymmetricI4ActivationLayout activation_layout = + common_symmetric_i4_activation_layout(fused.rms_binary.hidden_size, fused.rms_binary.token_count); + if (activation_layout.total_bytes == 0) { + return false; + } + const ValueId activation = context.next_plan_value; + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kAddRmsNormBinarySymmetricI4K32Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", fused.rms_binary.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.add_rmsnorm_binary_symmetric_i4.hidden_size", + to_config_value(fused.rms_binary.hidden_size)); + dispatch.kernel.compile_parameters.emplace("ggml.add_rmsnorm_binary_symmetric_i4.rms_epsilon", + to_config_value(fused.rms_binary.epsilon)); + dispatch.bindings.push_back({ fused.lhs->id, 0, fused.lhs->byte_count }); + dispatch.bindings.push_back({ fused.rhs->id, 0, fused.rhs->byte_count }); + dispatch.bindings.push_back({ fused.residual->id, 0, fused.residual->byte_count }); + dispatch.bindings.push_back({ fused.rms_binary.rhs->id, 0, fused.rms_binary.rhs->byte_count }); + dispatch.bindings.push_back({ fused.rms_binary.output->id, 0, fused.rms_binary.output->byte_count }); + dispatch.bindings.push_back({ activation, 0, activation_layout.payload_bytes }); + dispatch.bindings.push_back({ activation, activation_layout.scales_offset, activation_layout.metadata_bytes }); + dispatch.bindings.push_back({ activation, activation_layout.sums_offset, activation_layout.metadata_bytes }); + + Status metadata_status; + if (!match.metadata.append_alternate_value( + { fused.rms_binary.output->id, activation, GGML_TYPE_COUNT, activation_layout.total_bytes, + kCommonSymmetricI4K32ActivationAlternateName }, + metadata_status)) { + match.status.append(metadata_status); + return false; + } + + match.covered_nodes.push_back(fused.add_node_index); + match.covered_nodes.push_back(fused.rms_binary.rms_node_index); + match.covered_nodes.push_back(fused.rms_binary.binary_node_index); + match.dispatches.push_back(std::move(dispatch)); + match.transients.push_back( + { activation, kCommonSymmetricI4K32ActivationAlternateName, activation_layout.total_bytes, 256 }); + return match.status.success(); +} + +static bool match_rmsnorm_binary_symmetric_i4_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + const std::vector & nodes = context.graph.nodes(); + if (context.root_index >= nodes.size()) { + return false; + } + const RmsNormBinaryMatch fused = + match_rmsnorm_binary_f32(context.graph, &nodes[context.root_index], context.root_index); + if (!fused.matched() || fused.op != BinaryKind::Mul || + !is_supported_symmetric_i4_hidden_size(fused.hidden_size) || fused.token_count > 16 || + fused.rms_node_index >= context.covered_nodes.size() || + fused.binary_node_index >= context.covered_nodes.size() || context.covered_nodes[fused.rms_node_index] || + context.covered_nodes[fused.binary_node_index] || + !common_has_symmetric_i4_lowrow_consumer(context.graph, *fused.output) || + !pairwise_distinct_storage_roots(std::array{ fused.input, fused.rhs, fused.output })) { + return false; + } + + const CommonSymmetricI4ActivationLayout activation_layout = + common_symmetric_i4_activation_layout(fused.hidden_size, fused.token_count); + if (activation_layout.total_bytes == 0) { + return false; + } + const ValueId activation = context.next_plan_value; + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kRmsNormBinarySymmetricI4K32Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", fused.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.rmsnorm_binary_symmetric_i4.hidden_size", + to_config_value(fused.hidden_size)); + dispatch.kernel.compile_parameters.emplace("ggml.rmsnorm_binary_symmetric_i4.rms_epsilon", + to_config_value(fused.epsilon)); + dispatch.bindings.push_back({ fused.input->id, 0, fused.input->byte_count }); + dispatch.bindings.push_back({ fused.rhs->id, 0, fused.rhs->byte_count }); + dispatch.bindings.push_back({ fused.output->id, 0, fused.output->byte_count }); + dispatch.bindings.push_back({ activation, 0, activation_layout.payload_bytes }); + dispatch.bindings.push_back({ activation, activation_layout.scales_offset, activation_layout.metadata_bytes }); + dispatch.bindings.push_back({ activation, activation_layout.sums_offset, activation_layout.metadata_bytes }); + + Status metadata_status; + if (!match.metadata.append_alternate_value( + { fused.output->id, activation, GGML_TYPE_COUNT, activation_layout.total_bytes, + kCommonSymmetricI4K32ActivationAlternateName }, + metadata_status)) { + match.status.append(metadata_status); + return false; + } + + match.covered_nodes.push_back(fused.rms_node_index); + match.covered_nodes.push_back(fused.binary_node_index); + match.dispatches.push_back(std::move(dispatch)); + match.transients.push_back( + { activation, kCommonSymmetricI4K32ActivationAlternateName, activation_layout.total_bytes, 256 }); + return match.status.success(); +} + +static bool match_rmsnorm_gate_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + const RmsNormGateMatch fused = match_rmsnorm_gate(context.graph, context.root_node, context.root_index); + if (!fused.matched()) { + return false; + } + + const bool use_i4 = fused.gate_op == UnaryKind::Silu && + common_has_symmetric_i4_lowrow_consumer(context.graph, *fused.output); + if (!use_i4) { + if (fused.hidden_size > 1024) { + return false; + } + const Value * packed_input = fused.output; + const GraphNode * consumer = common_find_only_consumer_with_op(context.graph, packed_input->id, GGML_OP_MUL_MAT); + if (consumer == nullptr) { + const GraphNode * reshape = common_find_only_consumer_with_op(context.graph, packed_input->id, GGML_OP_RESHAPE); + const Value * reshaped = reshape != nullptr ? graph_value(context.graph, reshape->output) : nullptr; + if (reshaped != nullptr && is_layout_alias_node(context.graph, *reshape) && reshaped->contiguous && + same_full_value_range(*packed_input, *reshaped)) { + packed_input = reshaped; + consumer = common_find_only_consumer_with_op(context.graph, packed_input->id, GGML_OP_MUL_MAT); + } + } + const CommonMulMatMatch projection = + common_match_mul_mat_any_format(context.graph, consumer, kRmsNormGateF32F16Kernel, false); + const bool packed_output = projection.matched() && projection.input->id == packed_input->id && + projection.weight->alias_source.value < 0 && + common_mul_mat_uses_k16_major_f16(projection.weight_format, projection.input_size, + projection.output_size, projection.token_count); + const bool q8_output = !packed_output && packed_input->ne[1] <= 5 && + is_packed_q8_consumer(context.graph, consumer, *packed_input); + const size_t activation_bytes = q8_output ? q8_1_x4_byte_count(fused.token_count, fused.hidden_size) : + fused.output->byte_count / 2; + const ValueId activation = context.next_plan_value; + const char * activation_name = q8_output ? "common.rmsnorm_gate.q8_1_x4" : + packed_output ? "common.rmsnorm_gate.k16_major_f16" : "common.rmsnorm_gate.f16"; + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(q8_output ? kRmsNormGateF32Q8_1X4Kernel : kRmsNormGateF32F16Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", fused.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.rmsnorm_gate_f32.hidden_size", + to_config_value(fused.hidden_size)); + dispatch.kernel.compile_parameters.emplace("ggml.rmsnorm_gate_f32.rms_epsilon", + to_config_value(fused.epsilon)); + dispatch.kernel.compile_parameters.emplace("ggml.rmsnorm_gate_f32.gate_op", + std::to_string(unary_kind_config_value(fused.gate_op))); + dispatch.kernel.compile_parameters.emplace("ggml.rmsnorm_gate_f32.f16_output_row_width", + to_config_value(packed_output ? projection.input_size : 0)); + dispatch.bindings.push_back({ fused.input->id, 0, fused.input->byte_count }); + dispatch.bindings.push_back({ fused.weight->id, 0, fused.weight->byte_count }); + dispatch.bindings.push_back({ fused.raw_gate->id, 0, fused.raw_gate->byte_count }); + dispatch.bindings.push_back({ fused.output->id, 0, fused.output->byte_count }); + dispatch.bindings.push_back({ activation, 0, activation_bytes }); + Status status; + const bool recorded = packed_output ? + match.metadata.append_generated_resource( + { packed_input->id, GeneratedResourceRole::F16K16Major, activation, activation_bytes, {} }, status) : + match.metadata.append_alternate_value( + { fused.output->id, activation, q8_output ? GGML_TYPE_Q8_1 : GGML_TYPE_F16, + activation_bytes, activation_name }, status); + if (!recorded) { + match.status.append(status); + return false; + } + for (const GraphNode * covered : fused.covered) { + if (!append_covered_node_index_once(context.graph, context.covered_nodes, covered, match.covered_nodes)) { + return false; + } + } + match.transients.push_back({ activation, activation_name, activation_bytes, 256 }); + match.dispatches.push_back(std::move(dispatch)); + return true; + } + + const CommonSymmetricI4ActivationLayout activation_layout = + common_symmetric_i4_activation_layout(fused.hidden_size, fused.token_count); + if (activation_layout.total_bytes == 0) { + return false; + } + const ValueId activation = context.next_plan_value; + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kRmsNormGateSiluMulSymmetricI4K32Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", fused.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.rmsnorm_gate_silu_mul_symmetric_i4.hidden_size", + to_config_value(fused.hidden_size)); + dispatch.kernel.compile_parameters.emplace("ggml.rmsnorm_gate_silu_mul_symmetric_i4.rms_epsilon", + to_config_value(fused.epsilon)); + dispatch.bindings.push_back({ fused.input->id, 0, fused.input->byte_count }); + dispatch.bindings.push_back({ fused.weight->id, 0, fused.weight->byte_count }); + dispatch.bindings.push_back({ fused.raw_gate->id, 0, fused.raw_gate->byte_count }); + dispatch.bindings.push_back({ fused.output->id, 0, fused.output->byte_count }); + dispatch.bindings.push_back({ activation, 0, activation_layout.payload_bytes }); + dispatch.bindings.push_back({ activation, activation_layout.scales_offset, activation_layout.metadata_bytes }); + dispatch.bindings.push_back({ activation, activation_layout.sums_offset, activation_layout.metadata_bytes }); + + Status metadata_status; + if (!match.metadata.append_alternate_value( + { fused.output->id, activation, GGML_TYPE_COUNT, activation_layout.total_bytes, + kCommonSymmetricI4K32ActivationAlternateName }, + metadata_status)) { + match.status.append(metadata_status); + return false; + } + for (const GraphNode * covered : fused.covered) { + if (!append_covered_node_index_once(context.graph, context.covered_nodes, covered, match.covered_nodes)) { + return false; + } + } + match.dispatches.push_back(std::move(dispatch)); + match.transients.push_back( + { activation, kCommonSymmetricI4K32ActivationAlternateName, activation_layout.total_bytes, 256 }); + return match.status.success(); +} + +static bool has_qualified_rmsnorm_k16_consumer(const Graph & graph, const RmsNormBinaryMatch & rms) { + const bool qualified_shape = + (rms.token_count == 512 && (rms.hidden_size == 4096 || rms.hidden_size == 5120 || + rms.hidden_size == 6144 || rms.hidden_size == 8192)) || + (rms.token_count == 1024 && rms.hidden_size == 5120); + if (!qualified_shape || rms.op != BinaryKind::Mul || rms.output->ne[1] != rms.token_count || + rms.output->ne[2] != 1 || rms.output->ne[3] != 1 || + !pairwise_distinct_storage_roots(std::array{ rms.input, rms.rhs, rms.output })) { + return false; + } + for (const GraphNode * consumer : graph.index().consumers(rms.output->id)) { + const CommonMulMatMatch projection = + common_match_mul_mat_any_format(graph, consumer, kRmsNormBinaryF32K16Kernel, false); + if (projection.matched() && projection.input->id == rms.output->id && + projection.weight->alias_source.value < 0 && + common_mul_mat_uses_k16_major_f16(projection.weight_format, projection.input_size, + projection.output_size, projection.token_count)) { + return true; + } + } + return false; +} + +static bool match_rmsnorm_binary_f32_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + const std::vector & nodes = context.graph.nodes(); + if (context.root_index >= nodes.size()) { + return false; + } + const RmsNormBinaryMatch rms_match = + match_rmsnorm_binary_f32(context.graph, &nodes[context.root_index], context.root_index); + if (!rms_match.matched() || rms_match.rms_node_index >= context.covered_nodes.size() || + rms_match.binary_node_index >= context.covered_nodes.size() || + context.covered_nodes[rms_match.rms_node_index] || context.covered_nodes[rms_match.binary_node_index]) { + return false; + } + + const bool packed_output = has_qualified_rmsnorm_k16_consumer(context.graph, rms_match); + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(packed_output ? kRmsNormBinaryF32K16Kernel : kRmsNormBinaryF32Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", rms_match.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.rmsnorm_binary_f32.hidden_size", + to_config_value(rms_match.hidden_size)); + dispatch.kernel.compile_parameters.emplace("ggml.rmsnorm_binary_f32.rms_epsilon", + to_config_value(rms_match.epsilon)); + dispatch.kernel.compile_parameters.emplace("ggml.rmsnorm_binary_f32.op", + std::to_string(binary_kind_config_value(rms_match.op))); + dispatch.bindings.push_back({ rms_match.input->id, 0, rms_match.input->byte_count }); + dispatch.bindings.push_back({ rms_match.rhs->id, 0, rms_match.rhs->byte_count }); + dispatch.bindings.push_back({ rms_match.output->id, 0, rms_match.output->byte_count }); + + if (packed_output) { + const ValueId packed = context.next_plan_value; + const size_t bytes = rms_match.output->byte_count / 2; + Status status; + if (!match.metadata.append_generated_resource( + { rms_match.output->id, GeneratedResourceRole::F16K16Major, packed, bytes, {} }, status)) { + match.status.append(status); + return false; + } + dispatch.bindings.push_back({ packed, 0, bytes }); + match.transients.push_back({ packed, "common.rmsnorm_binary.k16_major_f16", bytes, 256 }); + } + + match.covered_nodes.push_back(rms_match.rms_node_index); + match.covered_nodes.push_back(rms_match.binary_node_index); + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_rmsnorm_binary_q8_1_x4_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + const std::vector & nodes = context.graph.nodes(); + if (context.root_index >= nodes.size()) { + return false; + } + const RmsNormBinaryMatch rms_match = + match_rmsnorm_binary_f32(context.graph, &nodes[context.root_index], context.root_index); + if (!rms_match.matched() || rms_match.rms_node_index >= context.covered_nodes.size() || + rms_match.binary_node_index >= context.covered_nodes.size() || + context.covered_nodes[rms_match.rms_node_index] || context.covered_nodes[rms_match.binary_node_index] || + !has_packed_q8_consumer(context.graph, *rms_match.output)) { + return false; + } + + const size_t q8_byte_count = q8_1_x4_byte_count(rms_match.token_count, rms_match.hidden_size); + if (q8_byte_count == 0) { + return false; + } + + const ValueId q8_value = context.next_plan_value; + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kRmsNormBinaryQ8_1X4Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", rms_match.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.rmsnorm_binary_q8_1_x4.hidden_size", + to_config_value(rms_match.hidden_size)); + dispatch.kernel.compile_parameters.emplace("ggml.rmsnorm_binary_q8_1_x4.rms_epsilon", + to_config_value(rms_match.epsilon)); + dispatch.kernel.compile_parameters.emplace("ggml.rmsnorm_binary_q8_1_x4.op", + std::to_string(binary_kind_config_value(rms_match.op))); + dispatch.bindings.push_back({ rms_match.input->id, 0, rms_match.input->byte_count }); + dispatch.bindings.push_back({ rms_match.rhs->id, 0, rms_match.rhs->byte_count }); + dispatch.bindings.push_back({ rms_match.output->id, 0, rms_match.output->byte_count }); + dispatch.bindings.push_back({ q8_value, 0, q8_byte_count }); + + Status metadata_status; + if (!match.metadata.append_alternate_value( + { rms_match.output->id, q8_value, GGML_TYPE_Q8_1, q8_byte_count, "common.rmsnorm_binary.q8_1_x4" }, + metadata_status)) { + match.status.append(metadata_status); + return false; + } + + match.covered_nodes.push_back(rms_match.rms_node_index); + match.covered_nodes.push_back(rms_match.binary_node_index); + match.dispatches.push_back(std::move(dispatch)); + match.transients.push_back({ q8_value, "common.rmsnorm_binary.q8_1_x4", q8_byte_count, 256 }); + return match.status.success(); +} + +static bool match_rmsnorm_f32_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + const std::vector & nodes = context.graph.nodes(); + if (context.root_index >= nodes.size()) { + return false; + } + const RmsNormMatch rms_match = match_rmsnorm_f32(context.graph, &nodes[context.root_index], context.root_index); + if (!rms_match.matched() || rms_match.rms_node_index >= context.covered_nodes.size() || + context.covered_nodes[rms_match.rms_node_index]) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kRmsNormF32Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", rms_match.token_count); + dispatch.kernel.compile_parameters.emplace("ggml.rmsnorm_f32.hidden_size", to_config_value(rms_match.hidden_size)); + dispatch.kernel.compile_parameters.emplace("ggml.rmsnorm_f32.rms_epsilon", to_config_value(rms_match.epsilon)); + dispatch.kernel.compile_parameters.emplace("ggml.rmsnorm_f32.input_stride", + to_config_value(rms_match.input_stride)); + dispatch.bindings.push_back( + { rms_match.input->storage_root, rms_match.input->storage_offset, rms_match.input_span }); + dispatch.bindings.push_back({ rms_match.output->id, 0, rms_match.output->byte_count }); + + match.covered_nodes.push_back(rms_match.rms_node_index); + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +bool common_match_rmsnorm_gate_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + return match_rmsnorm_gate_dispatch(context, match); +} + +void register_rmsnorm_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "common.add_rmsnorm_binary_symmetric_i4_k32", + GGML_OP_ADD, + DispatchMatchKind::Fused, + 400, + DispatchSource::Common, + match_add_rmsnorm_binary_symmetric_i4_dispatch, + }); + registry.add({ + "common.rmsnorm_gate", + GGML_OP_RMS_NORM, + DispatchMatchKind::Fused, + 400, + DispatchSource::Common, + match_rmsnorm_gate_dispatch, + }); + registry.add({ + "common.rmsnorm_binary_symmetric_i4_k32", + GGML_OP_RMS_NORM, + DispatchMatchKind::Fused, + 350, + DispatchSource::Common, + match_rmsnorm_binary_symmetric_i4_dispatch, + }); + registry.add({ + "common.rmsnorm_mul_rope_f32", + GGML_OP_RMS_NORM, + DispatchMatchKind::Fused, + 300, + DispatchSource::Common, + match_rmsnorm_mul_rope_f32_dispatch, + }); + registry.add({ + "common.rmsnorm_binary_q8_1_x4", + GGML_OP_RMS_NORM, + DispatchMatchKind::Fused, + 150, + DispatchSource::Common, + match_rmsnorm_binary_q8_1_x4_dispatch, + }); + registry.add({ + "common.rmsnorm_binary_f32", + GGML_OP_RMS_NORM, + DispatchMatchKind::Fused, + 100, + DispatchSource::Common, + match_rmsnorm_binary_f32_dispatch, + }); + registry.add({ + "common.rmsnorm_f32", + GGML_OP_RMS_NORM, + DispatchMatchKind::SingleOp, + 0, + DispatchSource::Common, + match_rmsnorm_f32_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-rmsnorm.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-rmsnorm.h new file mode 100644 index 000000000000..12d043813715 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-rmsnorm.h @@ -0,0 +1,11 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +bool common_match_rmsnorm_gate_dispatch(const DispatchMatchContext & context, DispatchMatch & match); + +void register_rmsnorm_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-rope-set-rows.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-rope-set-rows.cpp new file mode 100644 index 000000000000..cf003f1ba073 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-rope-set-rows.cpp @@ -0,0 +1,714 @@ +#include "dispatch-rope-set-rows.h" + +#include "dispatch-layout-utils.h" +#include "dispatch-rope-utils.h" +#include "ggml.h" +#include "graph/graph-matcher.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kRopeF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_rope_f32"); +static constexpr KernelCatalogRef kSetRowsKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_set_rows"); +static constexpr KernelCatalogRef kSetRowsScatterKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_set_rows_scatter"); +static constexpr KernelCatalogRef kRopeSetRowsF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_rope_set_rows_f32"); + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static std::string to_config_value(int64_t value) { + return std::to_string(value); +} + +static bool same_shape(const Value & lhs, const Value & rhs) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (lhs.ne[i] != rhs.ne[i]) { + return false; + } + } + return true; +} + +static bool is_1d_shape(const Value & value, int64_t ne0) { + return value.ne[0] == ne0 && value.ne[1] == 1 && value.ne[2] == 1 && value.ne[3] == 1; +} + +static bool is_2d_shape(const Value & value, int64_t ne0, int64_t ne1) { + return value.ne[0] == ne0 && value.ne[1] == ne1 && value.ne[2] == 1 && value.ne[3] == 1; +} + +static std::string value_layout_for_log(const Value * value) { + if (value == nullptr) { + return "null"; + } + std::ostringstream out; + out << ggml_type_name(value->type) << " ne=[" << value->ne[0] << "," << value->ne[1] << "," << value->ne[2] << "," + << value->ne[3] << "] nb=[" << value->nb[0] << "," << value->nb[1] << "," << value->nb[2] << "," << value->nb[3] + << "] contiguous=" << (value->contiguous ? 1 : 0); + if (value->alias_source.value >= 0) { + out << " alias=" << value->alias_source.value << " storage_offset=" << value->storage_offset; + } + return out.str(); +} + +static void log_rope_reject(Status * status, + const char * reason, + const Value * input, + const Value * positions, + const Value * output, + const Value * freq_factors = nullptr) { + if (status == nullptr) { + return; + } + status->log("ROPE matcher rejected node: %s input=%s positions=%s output=%s freq_factors=%s", reason, + value_layout_for_log(input).c_str(), value_layout_for_log(positions).c_str(), + value_layout_for_log(output).c_str(), value_layout_for_log(freq_factors).c_str()); +} + +static void log_set_rows_reject(Status * status, + const char * reason, + const Value * rows, + const Value * indices, + const Value * cache, + const Value * output) { + if (status == nullptr) { + return; + } + status->log("SET_ROWS matcher rejected node: %s rows=%s indices=%s cache=%s output=%s", reason, + value_layout_for_log(rows).c_str(), value_layout_for_log(indices).c_str(), + value_layout_for_log(cache).c_str(), value_layout_for_log(output).c_str()); +} + +static bool is_rope_shape(const Value & value) { + return value.ne[0] >= 4 && value.ne[0] <= 1024 && value.ne[0] % 4 == 0 && value.ne[1] >= 1 && value.ne[1] <= 64 && + value.ne[2] >= 1 && value.ne[2] <= 2048 && value.ne[3] == 1; +} + +static bool is_packed_f32_rope_layout(const Value & value) { + if (value.type != GGML_TYPE_F32 || !is_rope_shape(value) || value.nb[0] != sizeof(float)) { + return false; + } + return value.nb[1] == static_cast(value.ne[0]) * sizeof(float) && + value.nb[2] == static_cast(value.ne[1]) * value.nb[1] && + value.nb[3] == static_cast(value.ne[2]) * value.nb[2]; +} + +static bool is_supported_rope_input_layout(const Value & value, size_t & span_elements) { + if (!is_rope_shape(value) || value.nb[1] % sizeof(float) != 0 || value.nb[2] % sizeof(float) != 0 || + !strided_f32_storage_span_elements(value, span_elements)) { + return false; + } + + const size_t stride1 = value.nb[1] / sizeof(float); + const size_t stride2 = value.nb[2] / sizeof(float); + if (stride1 < static_cast(value.ne[0]) || stride2 < static_cast(value.ne[1]) * stride1) { + return false; + } + + const size_t token_span = static_cast(value.ne[2] - 1) * stride2; + const size_t head_span = static_cast(value.ne[1] - 1) * stride1; + const size_t required = token_span + head_span + static_cast(value.ne[0]); + return required <= span_elements; +} + +static bool is_supported_rope_params(const RopeParams & params, int64_t head_size) { + const bool supported_mode = params.mode == GGML_ROPE_TYPE_NORMAL || params.mode == GGML_ROPE_TYPE_NEOX; + return params.n_dims >= 4 && params.n_dims <= head_size && params.n_dims % 4 == 0 && supported_mode && + std::isfinite(params.freq_base) && params.freq_base > 0.0f && std::isfinite(params.freq_scale) && + params.freq_scale > 0.0f && std::isfinite(params.ext_factor) && std::isfinite(params.attn_factor); +} + +static bool build_rope_theta_table(const GraphNode & rope, + int64_t n_dims, + std::vector & data, + float & mscale) { + const RopeParams * params = op_params_as(rope.params); + if (params == nullptr) { + return false; + } + + RopeFrequencyTable table; + if (!build_rope_frequency_table(*params, n_dims, table)) { + return false; + } + data = std::move(table.data); + mscale = table.mscale; + return true; +} + +static void build_unit_frequency_factors(int64_t n_dims, std::vector & data) { + data.resize(static_cast(n_dims / 2) * sizeof(float)); + const float one = 1.0f; + for (int64_t i = 0; i < n_dims / 2; ++i) { + std::memcpy(data.data() + static_cast(i) * sizeof(float), &one, sizeof(one)); + } +} + +static bool format_value(ggml_type type, int64_t & value) { + switch (type) { + case GGML_TYPE_F16: + value = 16; + return true; + case GGML_TYPE_F32: + value = 32; + return true; + default: + return false; + } +} + +static bool supported_token_count(int64_t token_count) { + return token_count >= 1 && token_count <= 2048; +} + +static bool supported_cache_row_count(int64_t row_count) { + return row_count >= 1 && row_count <= 1048576; +} + +static bool supported_hidden_size(int64_t hidden_size) { + return hidden_size >= 4 && hidden_size <= 32768 && hidden_size % 4 == 0; +} + +static bool set_rows_supported_row_layout(const Value & rows, + int64_t hidden_size, + int64_t token_count, + int64_t & input_stride) { + if (!is_2d_shape(rows, hidden_size, token_count)) { + return false; + } + + const size_t element_size = ggml_type_size(rows.type); + if (rows.nb[0] != element_size || rows.nb[1] % element_size != 0) { + return false; + } + + input_stride = static_cast(rows.nb[1] / element_size); + return input_stride >= hidden_size && input_stride <= 1048576; +} + +static bool supported_set_rows_input_layout(const Value & rows, + int64_t hidden_size, + int64_t token_count, + int64_t & input_stride, + size_t & rows_span_bytes) { + if (!set_rows_supported_row_layout(rows, hidden_size, token_count, input_stride)) { + return false; + } + + if (rows.contiguous) { + rows_span_bytes = rows.byte_count; + return true; + } + + if (rows.type != GGML_TYPE_F32) { + return false; + } + + return strided_f32_storage_span_bytes(rows, rows_span_bytes); +} + +static ValueId next_match_transient_value(const DispatchMatchContext & context, const DispatchMatch & dispatch_match) { + return ValueId(context.next_plan_value.value + static_cast(dispatch_match.transients.size()) + + static_cast(dispatch_match.completion_counter_requests.size())); +} + +struct RopeMatch { + const GraphNode * node = nullptr; + const Value * input = nullptr; + const Value * positions = nullptr; + const Value * freq_factors = nullptr; + const Value * output = nullptr; + int64_t token_count = 0; + int64_t head_count = 0; + int64_t head_size = 0; + int64_t n_dims = 0; + int64_t input_stride1 = 0; + int64_t input_stride2 = 0; + size_t input_span = 0; + size_t input_span_bytes = 0; + float rope_mscale = 1.0f; + int64_t mode = 0; + size_t theta_bytes = 0; + size_t freq_factors_bytes = 0; + std::vector theta_data; + std::vector freq_factors_data; + + bool matched() const { return node != nullptr && input != nullptr && positions != nullptr && output != nullptr; } +}; + +struct SetRowsMatch { + const GraphNode * node = nullptr; + const Value * rows = nullptr; + const Value * indices = nullptr; + const Value * cache = nullptr; + const Value * output = nullptr; + size_t rows_span_bytes = 0; + int64_t input_stride = 0; + int64_t row_format = 0; + int64_t output_format = 0; + int64_t token_count = 0; + int64_t cache_row_count = 0; + int64_t hidden_size = 0; + + bool matched() const { return node != nullptr && rows != nullptr && indices != nullptr && output != nullptr; } +}; + +struct RopeSetRowsMatch { + RopeMatch rope; + SetRowsMatch set_rows; + const GraphNode * layout = nullptr; + const Value * cache_rows = nullptr; + size_t set_rows_idx = 0; + size_t layout_idx = 0; + + bool matched() const { return rope.matched() && set_rows.matched() && cache_rows != nullptr; } +}; + +static RopeMatch match_rope_f32(const Graph & graph, const GraphNode * node, Status * status = nullptr) { + RopeMatch match; + if (node == nullptr || node->op != GGML_OP_ROPE || node->inputs.size() < 2 || node->inputs.size() > 3) { + log_rope_reject(status, "root is not a supported ROPE arity", nullptr, nullptr, nullptr); + return match; + } + + const Value * input = graph_value(graph, node->inputs[0]); + const Value * positions = graph_value(graph, node->inputs[1]); + const Value * output = graph_value(graph, node->output); + size_t input_span = 0; + if (input == nullptr || positions == nullptr || output == nullptr || input->type != GGML_TYPE_F32 || + output->type != GGML_TYPE_F32 || positions->type != GGML_TYPE_I32 || + !is_supported_rope_input_layout(*input, input_span) || !is_packed_f32_rope_layout(*output) || + !positions->contiguous || !same_shape(*input, *output)) { + log_rope_reject(status, "input/output/positions shape or layout is unsupported", input, positions, output); + return {}; + } + + const int64_t head_size = input->ne[0]; + const int64_t head_count = input->ne[1]; + const int64_t token_count = input->ne[2]; + const RopeParams * params = op_params_as(node->params); + if (params == nullptr || !is_supported_rope_params(*params, head_size) || !is_1d_shape(*positions, token_count)) { + std::string reason = "ROPE params or positions shape is unsupported"; + if (params == nullptr) { + reason += " params=missing"; + } else { + std::ostringstream out; + out << reason << " params={n_dims=" << params->n_dims << ", mode=" << params->mode + << ", n_ctx_orig=" << params->n_ctx_orig << ", freq_base=" << params->freq_base + << ", freq_scale=" << params->freq_scale << ", ext_factor=" << params->ext_factor + << ", attn_factor=" << params->attn_factor << ", beta_fast=" << params->beta_fast + << ", beta_slow=" << params->beta_slow << "}"; + reason = out.str(); + } + log_rope_reject(status, reason.c_str(), input, positions, output); + return {}; + } + const int64_t n_dims = params->n_dims; + + std::vector theta_data; + float rope_mscale = 1.0f; + if (!build_rope_theta_table(*node, n_dims, theta_data, rope_mscale)) { + log_rope_reject(status, "failed to build supported ROPE theta table", input, positions, output); + return {}; + } + + const Value * freq_factors = nullptr; + size_t freq_factors_bytes = 0; + std::vector freq_factors_data; + if (node->inputs.size() == 3) { + freq_factors = graph_value(graph, node->inputs[2]); + if (freq_factors == nullptr || freq_factors->type != GGML_TYPE_F32 || !freq_factors->contiguous || + !is_1d_shape(*freq_factors, n_dims / 2)) { + log_rope_reject(status, "explicit frequency factors shape or layout is unsupported", input, positions, + output, freq_factors); + return {}; + } + freq_factors_bytes = freq_factors->byte_count; + } else { + build_unit_frequency_factors(n_dims, freq_factors_data); + freq_factors_bytes = freq_factors_data.size(); + } + + match.node = node; + match.input = input; + match.positions = positions; + match.freq_factors = freq_factors; + match.output = output; + match.token_count = token_count; + match.head_count = head_count; + match.head_size = head_size; + match.n_dims = n_dims; + match.input_stride1 = static_cast(input->nb[1] / sizeof(float)); + match.input_stride2 = static_cast(input->nb[2] / sizeof(float)); + match.input_span = input_span; + match.input_span_bytes = input_span * sizeof(float); + match.rope_mscale = rope_mscale; + match.mode = params->mode; + match.theta_bytes = theta_data.size(); + match.freq_factors_bytes = freq_factors_bytes; + match.theta_data = std::move(theta_data); + match.freq_factors_data = std::move(freq_factors_data); + return match; +} + +static SetRowsMatch match_set_rows_2d(const Graph & graph, const GraphNode * node, Status * status = nullptr) { + SetRowsMatch match; + if (node == nullptr || node->op != GGML_OP_SET_ROWS || node->inputs.size() != 3) { + log_set_rows_reject(status, "root is not a supported SET_ROWS arity", nullptr, nullptr, nullptr, nullptr); + return match; + } + + const Value * rows = graph_value(graph, node->inputs[0]); + const Value * indices = graph_value(graph, node->inputs[1]); + const Value * cache = graph_value(graph, node->inputs[2]); + const Value * output = graph_value(graph, node->output); + if (rows == nullptr || indices == nullptr || cache == nullptr || output == nullptr || !indices->contiguous || + !cache->contiguous || indices->type != GGML_TYPE_I64 || output->type != cache->type || + !same_shape(*cache, *output) || !graph.values().same_storage(cache->id, output->id)) { + log_set_rows_reject(status, "input/output shape or storage is unsupported", rows, indices, cache, output); + return {}; + } + + int64_t row_format = 0; + int64_t output_format = 0; + if (!format_value(rows->type, row_format) || !format_value(output->type, output_format)) { + log_set_rows_reject(status, "input or output format is unsupported", rows, indices, cache, output); + return {}; + } + if (rows->type == GGML_TYPE_F16 && output->type != GGML_TYPE_F16) { + log_set_rows_reject(status, "f16 input requires f16 output", rows, indices, cache, output); + return {}; + } + + const int64_t hidden_size = rows->ne[0]; + const int64_t token_count = rows->ne[1]; + const int64_t cache_row_count = cache->ne[1]; + int64_t input_stride = 0; + size_t rows_span_bytes = 0; + if (!supported_set_rows_input_layout(*rows, hidden_size, token_count, input_stride, rows_span_bytes) || + !is_1d_shape(*indices, token_count) || !is_2d_shape(*cache, hidden_size, cache_row_count) || + !supported_hidden_size(hidden_size) || !supported_token_count(token_count) || + !supported_cache_row_count(cache_row_count)) { + log_set_rows_reject(status, "shape or rows layout is unsupported", rows, indices, cache, output); + return {}; + } + + match.node = node; + match.rows = rows; + match.indices = indices; + match.cache = cache; + match.output = output; + match.rows_span_bytes = rows_span_bytes; + match.input_stride = input_stride; + match.row_format = row_format; + match.output_format = output_format; + match.token_count = token_count; + match.cache_row_count = cache_row_count; + match.hidden_size = hidden_size; + return match; +} + +static void add_rope_compile_parameters(Dispatch & dispatch, const RopeMatch & match) { + dispatch.kernel.compile_parameters.emplace("ggml.rope_f32.head_size", to_config_value(match.head_size)); + dispatch.kernel.compile_parameters.emplace("ggml.rope_f32.n_dims", to_config_value(match.n_dims)); + dispatch.kernel.compile_parameters.emplace("ggml.rope_f32.head_count", to_config_value(match.head_count)); + dispatch.kernel.compile_parameters.emplace("ggml.rope_f32.token_capacity", to_config_value(match.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.rope_f32.input_stride1", to_config_value(match.input_stride1)); + dispatch.kernel.compile_parameters.emplace("ggml.rope_f32.input_stride2", to_config_value(match.input_stride2)); + dispatch.kernel.compile_parameters.emplace("ggml.rope_f32.mscale", rope_mscale_config_value(match.rope_mscale)); + dispatch.kernel.compile_parameters.emplace("ggml.rope_f32.mode", to_config_value(match.mode)); +} + +static void add_set_rows_compile_parameters(Dispatch & dispatch, const SetRowsMatch & match) { + dispatch.kernel.compile_parameters.emplace("ggml.set_rows.token_capacity", to_config_value(match.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.set_rows.hidden_capacity", to_config_value(match.hidden_size)); + dispatch.kernel.compile_parameters.emplace("ggml.set_rows.input_format", to_config_value(match.row_format)); + dispatch.kernel.compile_parameters.emplace("ggml.set_rows.output_format", to_config_value(match.output_format)); + dispatch.kernel.compile_parameters.emplace("ggml.set_rows.input_stride", to_config_value(match.input_stride)); +} + +static ValueId add_constant_binding(const DispatchMatchContext & context, + DispatchMatch & dispatch_match, + const char * name, + size_t byte_count, + const std::vector & data) { + const ValueId value = next_match_transient_value(context, dispatch_match); + dispatch_match.transients.push_back({ value, name, byte_count, 256 }); + dispatch_match.constant_initializations.push_back({ + value, + name, + 0, + data, + }); + return value; +} + +static std::pair add_rope_frequency_bindings(const DispatchMatchContext & context, + const RopeMatch & match, + DispatchMatch & dispatch_match, + const char * prefix) { + const std::string theta_name = std::string(prefix) + ".theta"; + const ValueId theta = + add_constant_binding(context, dispatch_match, theta_name.c_str(), match.theta_bytes, match.theta_data); + if (match.freq_factors != nullptr) { + return { theta, match.freq_factors->id }; + } + + const std::string factors_name = std::string(prefix) + ".freq_factors"; + const ValueId freq_factors = add_constant_binding(context, dispatch_match, factors_name.c_str(), + match.freq_factors_bytes, match.freq_factors_data); + return { theta, freq_factors }; +} + +static Dispatch make_rope_dispatch(const DispatchMatchContext & context, + const RopeMatch & match, + DispatchMatch & dispatch_match) { + const auto [theta, freq_factors] = add_rope_frequency_bindings(context, match, dispatch_match, "common.rope"); + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kRopeF32Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + add_rope_compile_parameters(dispatch, match); + dispatch.kernel.compile_parameters.emplace("ggml.rope_f32.input_span", to_config_value(match.input_span)); + dispatch.bindings.push_back({ match.positions->id, 0, match.positions->byte_count }); + dispatch.bindings.push_back({ match.input->storage_root, match.input->storage_offset, match.input_span_bytes }); + dispatch.bindings.push_back({ theta, 0, match.theta_bytes }); + dispatch.bindings.push_back({ freq_factors, 0, match.freq_factors_bytes }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + return dispatch; +} + +static Dispatch make_set_rows_dispatch(const SetRowsMatch & match) { + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kSetRowsKernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.integer_parameters.emplace("cache_row_count", match.cache_row_count); + dispatch.kernel.integer_parameters.emplace("hidden_size", match.hidden_size); + add_set_rows_compile_parameters(dispatch, match); + dispatch.bindings.push_back({ match.rows->storage_root, match.rows->storage_offset, match.rows_span_bytes }); + dispatch.bindings.push_back({ match.indices->id, 0, match.indices->byte_count }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + return dispatch; +} + +static Dispatch make_rope_set_rows_dispatch(const DispatchMatchContext & context, + const RopeSetRowsMatch & match, + DispatchMatch & dispatch_match) { + const auto [theta, freq_factors] = + add_rope_frequency_bindings(context, match.rope, dispatch_match, "common.rope_set_rows"); + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kRopeSetRowsF32Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.rope.token_count); + dispatch.kernel.integer_parameters.emplace("cache_row_count", match.set_rows.cache_row_count); + dispatch.kernel.compile_parameters.emplace("ggml.rope_set_rows_f32.head_size", + to_config_value(match.rope.head_size)); + dispatch.kernel.compile_parameters.emplace("ggml.rope_set_rows_f32.n_dims", to_config_value(match.rope.n_dims)); + dispatch.kernel.compile_parameters.emplace("ggml.rope_set_rows_f32.head_count", + to_config_value(match.rope.head_count)); + dispatch.kernel.compile_parameters.emplace("ggml.rope_set_rows_f32.token_capacity", + to_config_value(match.rope.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.rope_set_rows_f32.input_stride1", + to_config_value(match.rope.input_stride1)); + dispatch.kernel.compile_parameters.emplace("ggml.rope_set_rows_f32.input_stride2", + to_config_value(match.rope.input_stride2)); + dispatch.kernel.compile_parameters.emplace("ggml.rope_set_rows_f32.input_span", + to_config_value(match.rope.input_span)); + dispatch.kernel.compile_parameters.emplace("ggml.rope_set_rows_f32.mscale", + rope_mscale_config_value(match.rope.rope_mscale)); + dispatch.kernel.compile_parameters.emplace("ggml.rope_set_rows_f32.output_format", + to_config_value(match.set_rows.output_format)); + dispatch.kernel.compile_parameters.emplace("ggml.rope_set_rows_f32.mode", to_config_value(match.rope.mode)); + dispatch.bindings.push_back({ match.rope.positions->id, 0, match.rope.positions->byte_count }); + dispatch.bindings.push_back({ match.set_rows.indices->id, 0, match.set_rows.indices->byte_count }); + dispatch.bindings.push_back( + { match.rope.input->storage_root, match.rope.input->storage_offset, match.rope.input_span_bytes }); + dispatch.bindings.push_back({ theta, 0, match.rope.theta_bytes }); + dispatch.bindings.push_back({ freq_factors, 0, match.rope.freq_factors_bytes }); + dispatch.bindings.push_back({ match.set_rows.output->id, 0, match.set_rows.output->byte_count }); + return dispatch; +} + +static bool find_rope_set_rows_consumer(const Graph & graph, const GraphNode & rope, RopeSetRowsMatch & match) { + if (!graph.has_index() || !graph.index().has_single_consumer(rope.output)) { + return false; + } + + const GraphNode * consumer = graph.index().consumers(rope.output).front(); + if (consumer == nullptr) { + return false; + } + + if (consumer->op == GGML_OP_SET_ROWS) { + match.set_rows_idx = 0; + match.layout = nullptr; + return graph.index().node_index(consumer, match.set_rows_idx) && + (match.set_rows = match_set_rows_2d(graph, consumer)).matched(); + } + + if (!is_layout_alias_node(graph, *consumer) || !graph.index().has_single_consumer(consumer->output)) { + return false; + } + + const GraphNode * set_rows = graph.index().consumers(consumer->output).front(); + if (set_rows == nullptr || set_rows->op != GGML_OP_SET_ROWS) { + return false; + } + match.layout = consumer; + return graph.index().node_index(consumer, match.layout_idx) && + graph.index().node_index(set_rows, match.set_rows_idx) && + (match.set_rows = match_set_rows_2d(graph, set_rows)).matched(); +} + +static RopeSetRowsMatch match_rope_set_rows_f32(const Graph & graph, const GraphNode * node) { + RopeSetRowsMatch match; + match.rope = match_rope_f32(graph, node); + if (!match.rope.matched() || !find_rope_set_rows_consumer(graph, *node, match)) { + return {}; + } + + match.cache_rows = match.set_rows.rows; + if (match.set_rows.row_format != 32 || match.set_rows.hidden_size != match.rope.head_size * match.rope.head_count || + match.set_rows.token_count != match.rope.token_count || match.cache_rows->type != GGML_TYPE_F32) { + return {}; + } + return match; +} + +static bool match_rope_set_rows_f32_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const RopeSetRowsMatch match = match_rope_set_rows_f32(context.graph, context.root_node); + if (!match.matched()) { + return false; + } + + if (!append_covered_node_index_once(context.graph, context.covered_nodes, context.root_node, + dispatch_match.covered_nodes) || + (match.layout != nullptr && !append_covered_node_index_once(context.graph, context.covered_nodes, match.layout, + dispatch_match.covered_nodes)) || + !append_covered_node_index_once(context.graph, context.covered_nodes, match.set_rows.node, + dispatch_match.covered_nodes)) { + return false; + } + + dispatch_match.dispatches.push_back(make_rope_set_rows_dispatch(context, match, dispatch_match)); + return true; +} + +static bool match_rope_f32_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const RopeMatch match = match_rope_f32(context.graph, context.root_node, &dispatch_match.status); + if (!match.matched()) { + return false; + } + + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(make_rope_dispatch(context, match, dispatch_match)); + return true; +} + +static SetRowsMatch match_set_rows_scatter(const Graph & graph, const GraphNode * node) { + SetRowsMatch match; + if (node == nullptr || node->op != GGML_OP_SET_ROWS || node->inputs.size() != 3) { + return match; + } + const Value * rows = graph_value(graph, node->inputs[0]); + const Value * indices = graph_value(graph, node->inputs[1]); + const Value * cache = graph_value(graph, node->inputs[2]); + const Value * output = graph_value(graph, node->output); + // Element-wise scatter: the non-FA path flattens the K/V rows, so hidden_size == 1 and each + // value is written to its own cache row. + if (rows == nullptr || indices == nullptr || cache == nullptr || output == nullptr || !rows->contiguous || + !indices->contiguous || !cache->contiguous || rows->type != GGML_TYPE_F32 || + output->type != GGML_TYPE_F16 || indices->type != GGML_TYPE_I64 || rows->ne[0] != 1 || cache->ne[0] != 1 || + !same_shape(*cache, *output) || !graph.values().same_storage(cache->id, output->id) || + rows->ne[1] != indices->ne[0] || rows->ne[1] < 1 || rows->ne[1] > 1048576 || cache->ne[1] < 1 || + cache->ne[1] > 134217728) { + return {}; + } + match.node = node; + match.rows = rows; + match.indices = indices; + match.cache = cache; + match.output = output; + match.rows_span_bytes = rows->byte_count; + match.token_count = rows->ne[1]; + match.cache_row_count = cache->ne[1]; + match.hidden_size = 1; + return match; +} + +static bool match_set_rows_scatter_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const SetRowsMatch match = match_set_rows_scatter(context.graph, context.root_node); + if (!match.matched()) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kSetRowsScatterKernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.integer_parameters.emplace("cache_row_count", match.cache_row_count); + dispatch.kernel.compile_parameters.emplace("ggml.set_rows_scatter.token_capacity", + to_config_value(match.token_count)); + dispatch.bindings.push_back({ match.rows->storage_root, match.rows->storage_offset, match.rows_span_bytes }); + dispatch.bindings.push_back({ match.indices->id, 0, match.indices->byte_count }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_set_rows_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const SetRowsMatch match = match_set_rows_2d(context.graph, context.root_node, &dispatch_match.status); + if (!match.matched()) { + return false; + } + + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(make_set_rows_dispatch(match)); + return true; +} + +} // namespace + +void register_rope_set_rows_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "common.rope_set_rows.f32", + GGML_OP_ROPE, + DispatchMatchKind::Fused, + 200, + DispatchSource::Common, + match_rope_set_rows_f32_dispatch, + }); + registry.add({ + "common.rope.f32", + GGML_OP_ROPE, + DispatchMatchKind::SingleOp, + 100, + DispatchSource::Common, + match_rope_f32_dispatch, + }); + registry.add({ + "common.set_rows_scatter", + GGML_OP_SET_ROWS, + DispatchMatchKind::SingleOp, + 200, + DispatchSource::Common, + match_set_rows_scatter_dispatch, + }); + registry.add({ + "common.set_rows", + GGML_OP_SET_ROWS, + DispatchMatchKind::SingleOp, + 100, + DispatchSource::Common, + match_set_rows_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-rope-set-rows.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-rope-set-rows.h new file mode 100644 index 000000000000..8e9497a6fe68 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-rope-set-rows.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_rope_set_rows_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-rope-utils.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-rope-utils.h new file mode 100644 index 000000000000..7981ccd1c03e --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-rope-utils.h @@ -0,0 +1,72 @@ +#pragma once + +#include "ggml.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { + +struct RopeFrequencyTable { + std::vector data; + float mscale = 1.0f; +}; + +inline float rope_yarn_ramp(float low, float high, int64_t i0) { + const float y = (static_cast(i0 / 2) - low) / std::max(0.001f, high - low); + return 1.0f - std::min(1.0f, std::max(0.0f, y)); +} + +inline float rope_yarn_corr_dim(int n_dims, int n_ctx_orig, float n_rot, float base) { + return static_cast(n_dims) * std::log(static_cast(n_ctx_orig) / (n_rot * 2.0f * static_cast(M_PI))) / + (2.0f * std::log(base)); +} + +inline void rope_yarn_corr_dims(int n_dims, int n_ctx_orig, float freq_base, float beta_fast, float beta_slow, float dims[2]) { + const float start = std::floor(rope_yarn_corr_dim(n_dims, n_ctx_orig, beta_fast, freq_base)); + const float end = std::ceil(rope_yarn_corr_dim(n_dims, n_ctx_orig, beta_slow, freq_base)); + dims[0] = std::max(0.0f, start); + dims[1] = std::min(static_cast(n_dims - 1), end); +} + +inline bool build_rope_frequency_table(const RopeParams & params, int64_t n_dims, RopeFrequencyTable & table) { + if (n_dims <= 0 || n_dims % 2 != 0 || params.freq_base <= 0.0f || !std::isfinite(params.freq_base) || + params.freq_scale <= 0.0f || !std::isfinite(params.freq_scale) || !std::isfinite(params.ext_factor) || + !std::isfinite(params.attn_factor)) { + return false; + } + + float corr_dims[2] = { 0.0f, 0.0f }; + rope_yarn_corr_dims(static_cast(n_dims), params.n_ctx_orig, params.freq_base, params.beta_fast, params.beta_slow, + corr_dims); + + table.mscale = params.attn_factor; + if (params.ext_factor != 0.0f) { + table.mscale *= 1.0f + 0.1f * std::log(1.0f / params.freq_scale); + } + + table.data.resize(static_cast(n_dims / 2) * sizeof(float)); + const float theta_scale = std::pow(params.freq_base, -2.0f / static_cast(n_dims)); + float theta = 1.0f; + for (int64_t i = 0; i < n_dims / 2; ++i) { + const int64_t dim = 2 * i; + const float ramp_mix = rope_yarn_ramp(corr_dims[0], corr_dims[1], dim) * params.ext_factor; + const float frequency = params.freq_scale * theta * (1.0f - ramp_mix) + theta * ramp_mix; + std::memcpy(table.data.data() + static_cast(i) * sizeof(float), &frequency, sizeof(frequency)); + theta *= theta_scale; + } + return true; +} + +inline std::string rope_mscale_config_value(float value) { + char buffer[64]; + std::snprintf(buffer, sizeof(buffer), "%.9g", static_cast(value)); + return std::string(buffer); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-scale.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-scale.cpp new file mode 100644 index 000000000000..5c23046162a1 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-scale.cpp @@ -0,0 +1,123 @@ +#include "dispatch-scale.h" + +#include "ggml.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kScaleF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_scale_f32"); + +static bool same_shape(const Value & lhs, const Value & rhs) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (lhs.ne[i] != rhs.ne[i]) { + return false; + } + } + return true; +} + +static bool positive_shape(const Value & value) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (value.ne[i] <= 0) { + return false; + } + } + return value.element_count > 0; +} + +static bool packed_f32_layout(const Value & value) { + size_t expected_stride = sizeof(float); + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (value.nb[i] != expected_stride) { + return false; + } + expected_stride *= static_cast(value.ne[i]); + } + return true; +} + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static bool distinct_storage(const Value & input, const Value & output) { + return input.storage != output.storage; +} + +static bool supported_source_layout(const Graph & graph, const Value & value) { + if (value.alias_source.value < 0) { + return true; + } + if (value.storage_offset != 0) { + return false; + } + const GraphNode * producer = graph.index().producer(value.id); + return producer != nullptr && (producer->op == GGML_OP_RESHAPE || producer->op == GGML_OP_CLAMP); +} + +static std::string format_float_config(float value) { + std::ostringstream out; + out.precision(9); + out << value; + return out.str(); +} + +static bool match_scale_f32_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_SCALE || node->inputs.size() != 1) { + return false; + } + + const ScaleParams * params = op_params_as(node->params); + if (params == nullptr) { + return false; + } + + const Value * output = graph_value(context.graph, node->output); + const Value * input = graph_value(context.graph, node->inputs[0]); + if (output == nullptr || input == nullptr) { + return false; + } + + if (output->type != GGML_TYPE_F32 || input->type != GGML_TYPE_F32 || !same_shape(*output, *input) || + !positive_shape(*output) || !output->contiguous || !input->contiguous || !packed_f32_layout(*output) || + !packed_f32_layout(*input) || output->alias_source.value >= 0 || + !supported_source_layout(context.graph, *input) || !distinct_storage(*input, *output) || + static_cast(output->element_count) > std::numeric_limits::max()) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kScaleF32Kernel); + dispatch.kernel.integer_parameters.emplace("element_count", output->element_count); + dispatch.kernel.compile_parameters.emplace("ggml.scale_f32.scale", format_float_config(params->scale)); + dispatch.kernel.compile_parameters.emplace("ggml.scale_f32.bias", format_float_config(params->bias)); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + + match.covered_nodes.push_back(context.root_index); + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_scale_dispatch(DispatchRegistryBuilder & registry) { + registry.add({ + "common.scale_f32", + GGML_OP_SCALE, + DispatchMatchKind::SingleOp, + 0, + DispatchSource::Common, + match_scale_f32_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-scale.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-scale.h new file mode 100644 index 000000000000..7c962a0062fd --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-scale.h @@ -0,0 +1,9 @@ +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +namespace ggml::hrx { + +void register_scale_dispatch(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.cpp new file mode 100644 index 000000000000..cfda7ae003a0 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.cpp @@ -0,0 +1,732 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +#include "dispatch-small-rows.h" + +#include "ggml.h" +#include "graph/graph-matcher.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kSoftmaxRowsKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_softmax_rows_f32"); +static constexpr KernelCatalogRef kSumRowsKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_sum_rows_f32"); +static constexpr KernelCatalogRef kArgsortRowsKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_argsort_rows_f32"); +static constexpr KernelCatalogRef kGetRowsSmallKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_get_rows_small_f32"); +static constexpr KernelCatalogRef kCopyStridedKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_copy_strided_f32"); +static constexpr KernelCatalogRef kNormRowsKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_norm_rows_f32"); +static constexpr KernelCatalogRef kBinaryStridedKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_binary_strided_f32"); +static constexpr KernelCatalogRef kClampKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_clamp_f32"); +static constexpr KernelCatalogRef kClampInplaceKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_clamp_inplace_f32"); +static constexpr KernelCatalogRef kCopyF32F16Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_copy_strided_f32_f16"); +static constexpr KernelCatalogRef kAttentionStridedKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_attention_strided_f32_f16"); +static constexpr KernelCatalogRef kAttentionRowsKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_attention_rows_f32_f16"); +static constexpr KernelCatalogRef kRopeRotateHalfKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_rope_rotate_half_f32"); +static constexpr KernelCatalogRef kGegluStridedKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_geglu_strided_f32"); +static constexpr KernelCatalogRef kMulMatSmallF16Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_small_f16_f32"); +static constexpr KernelCatalogRef kMulMatSmallF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_small_f32_f32"); +static constexpr KernelCatalogRef kMulMatSmallQ8Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_small_q8_0_f32"); +static constexpr KernelCatalogRef kMulMatRowsQ8Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_rows_q8_0_f32"); + +static bool packed(const Value & value, size_t element_size) { + size_t stride = element_size; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (value.ne[i] <= 0 || value.nb[i] != stride) { + return false; + } + stride *= static_cast(value.ne[i]); + } + return true; +} + +static int64_t rows_of(const Value & value) { + return value.ne[1] * value.ne[2] * value.ne[3]; +} + +static bool same_shape(const Value & a, const Value & b) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (a.ne[i] != b.ne[i]) { + return false; + } + } + return true; +} + +// the single input and the output of a row op, both packed F32 (the output may be I32) +static bool row_op_values(const DispatchMatchContext & context, ggml_op op, ggml_type output_type, + const Value *& input, const Value *& output) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != op || node->inputs.size() != 1) { + return false; + } + input = context.graph.values().find(node->inputs[0]); + output = context.graph.values().find(node->output); + return input != nullptr && output != nullptr && input->type == GGML_TYPE_F32 && output->type == output_type && + packed(*input, sizeof(float)) && packed(*output, ggml_type_size(output_type)) && + output->alias_source.value < 0 && input->storage != output->storage && rows_of(*input) <= 16777216; +} + +static void finish(const DispatchMatchContext & context, DispatchMatch & match, Dispatch && dispatch) { + match.covered_nodes.push_back(context.root_index); + match.dispatches.push_back(std::move(dispatch)); +} + +static bool match_softmax_rows(const DispatchMatchContext & context, DispatchMatch & match) { + const Value * input = nullptr; + const Value * output = nullptr; + if (!row_op_values(context, GGML_OP_SOFT_MAX, GGML_TYPE_F32, input, output) || !same_shape(*input, *output) || + input->ne[0] > 4096) { + return false; + } + const SoftMaxParams * params = op_params_as(context.root_node->params); + if (params == nullptr || params->scale != 1.0f || params->max_bias != 0.0f) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kSoftmaxRowsKernel); + dispatch.kernel.integer_parameters.emplace("column_count", input->ne[0]); + dispatch.kernel.integer_parameters.emplace("row_count", rows_of(*input)); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + finish(context, match, std::move(dispatch)); + return true; +} + +static bool match_sum_rows(const DispatchMatchContext & context, DispatchMatch & match) { + const Value * input = nullptr; + const Value * output = nullptr; + if (!row_op_values(context, GGML_OP_SUM_ROWS, GGML_TYPE_F32, input, output) || input->ne[0] > 4096 || + output->ne[0] != 1 || output->ne[1] != input->ne[1] || output->ne[2] != input->ne[2] || + output->ne[3] != input->ne[3]) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kSumRowsKernel); + dispatch.kernel.integer_parameters.emplace("column_count", input->ne[0]); + dispatch.kernel.integer_parameters.emplace("row_count", rows_of(*input)); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + finish(context, match, std::move(dispatch)); + return true; +} + +static bool match_argsort_rows(const DispatchMatchContext & context, DispatchMatch & match) { + const Value * input = nullptr; + const Value * output = nullptr; + if (!row_op_values(context, GGML_OP_ARGSORT, GGML_TYPE_I32, input, output) || !same_shape(*input, *output) || + input->ne[0] > 1024) { + return false; + } + const ArgsortParams * params = op_params_as(context.root_node->params); + if (params == nullptr) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kArgsortRowsKernel); + dispatch.kernel.integer_parameters.emplace("column_count", input->ne[0]); + dispatch.kernel.integer_parameters.emplace("row_count", rows_of(*input)); + dispatch.kernel.integer_parameters.emplace("descending", params->order == GGML_SORT_ORDER_DESC ? 1 : 0); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + finish(context, match, std::move(dispatch)); + return true; +} + +// NORM (LayerNorm without affine; its weight and bias are separate MUL/ADD nodes) on packed F32 rows +static bool match_norm_rows(const DispatchMatchContext & context, DispatchMatch & match) { + const Value * input = nullptr; + const Value * output = nullptr; + if (!row_op_values(context, GGML_OP_NORM, GGML_TYPE_F32, input, output) || !same_shape(*input, *output) || + input->ne[0] > 65536) { + return false; + } + const RmsNormParams * params = op_params_as(context.root_node->params); + if (params == nullptr) { + return false; + } + std::ostringstream eps; + eps.precision(9); + eps << params->eps; + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kNormRowsKernel); + dispatch.kernel.integer_parameters.emplace("column_count", input->ne[0]); + dispatch.kernel.integer_parameters.emplace("row_count", rows_of(*input)); + dispatch.kernel.compile_parameters.emplace("ggml.norm_rows_f32.epsilon", eps.str()); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + finish(context, match, std::move(dispatch)); + return true; +} + +// GET_ROWS within batches: source [W, S, B], ids [R, B] (I32, rows may be strided, as the first k +// columns of an ARGSORT are), output [W, R, B]. Only rows the +// regular get_rows kernel does not take (narrower than 4 floats, or not a multiple of 4). +static bool match_get_rows_small(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_GET_ROWS || node->inputs.size() != 2) { + return false; + } + const Value * source = context.graph.values().find(node->inputs[0]); + const Value * ids = context.graph.values().find(node->inputs[1]); + const Value * output = context.graph.values().find(node->output); + if (source == nullptr || ids == nullptr || output == nullptr || source->type != GGML_TYPE_F32 || + ids->type != GGML_TYPE_I32 || output->type != GGML_TYPE_F32 || !packed(*source, sizeof(float)) || + !packed(*output, sizeof(float)) || output->alias_source.value >= 0 || ids->nb[0] != sizeof(int32_t) || + ids->nb[1] % sizeof(int32_t) != 0 || static_cast(ids->nb[1] / sizeof(int32_t)) < ids->ne[0]) { + return false; + } + const int64_t w = source->ne[0]; + const int64_t s = source->ne[1]; + const int64_t b = source->ne[2]; + const int64_t r = ids->ne[0]; + if ((w >= 4 && w % 4 == 0) || w > 65536 || s > 65536 || b > 65536 || r > 65536 || source->ne[3] != 1 || + ids->ne[1] != b || ids->ne[2] != 1 || ids->ne[3] != 1 || output->ne[0] != w || output->ne[1] != r || + output->ne[2] != b || output->ne[3] != 1) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kGetRowsSmallKernel); + dispatch.kernel.integer_parameters.emplace("width", w); + dispatch.kernel.integer_parameters.emplace("id_count", r); + dispatch.kernel.integer_parameters.emplace("batch_count", b); + dispatch.kernel.integer_parameters.emplace("source_rows", s); + dispatch.kernel.integer_parameters.emplace("id_stride", static_cast(ids->nb[1] / sizeof(int32_t))); + dispatch.bindings.push_back({ source->id, 0, source->byte_count }); + dispatch.bindings.push_back({ ids->id, 0, ids->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + finish(context, match, std::move(dispatch)); + return true; +} + + +// CONT of a strided F32 view (permuted, transposed or sliced) into a packed output. Only sources +// that are not packed: packed copies belong to the regular copy kernel. +static bool match_copy_strided(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_CONT || node->inputs.size() != 1) { + return false; + } + const Value * input = context.graph.values().find(node->inputs[0]); + const Value * output = context.graph.values().find(node->output); + if (input == nullptr || output == nullptr || input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || + packed(*input, sizeof(float)) || !packed(*output, sizeof(float)) || output->alias_source.value >= 0 || + input->storage == output->storage || input->element_count != output->element_count) { + return false; + } + int64_t strides[GGML_MAX_DIMS]; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (input->nb[i] % sizeof(float) != 0 || input->ne[i] != output->ne[i]) { + return false; + } + strides[i] = static_cast(input->nb[i] / sizeof(float)); + } + const int64_t extent = static_cast(input->byte_count / sizeof(float)); + if (extent < 1 || extent > 268435456 || output->element_count > 268435456) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kCopyStridedKernel); + static const char * const ne_names[] = { "ne0", "ne1", "ne2", "ne3" }; + static const char * const s_names[] = { "s0", "s1", "s2", "s3" }; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + dispatch.kernel.integer_parameters.emplace(ne_names[i], output->ne[i]); + dispatch.kernel.integer_parameters.emplace(s_names[i], strides[i]); + } + dispatch.kernel.integer_parameters.emplace("source_extent", extent); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + finish(context, match, std::move(dispatch)); + return true; +} + + +// REPEAT that only broadcasts (every input dim is 1 or the output's): the strided copy with a zero +// stride on the broadcast dims. Tiling repeats (output a multiple of a larger input) are not claimed. +static bool match_repeat_broadcast(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_REPEAT || node->inputs.size() != 1) { + return false; + } + const Value * input = context.graph.values().find(node->inputs[0]); + const Value * output = context.graph.values().find(node->output); + if (input == nullptr || output == nullptr || input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || + !packed(*output, sizeof(float)) || output->alias_source.value >= 0 || input->storage == output->storage || + output->element_count > 268435456) { + return false; + } + int64_t strides[GGML_MAX_DIMS]; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (input->nb[i] % sizeof(float) != 0 || (input->ne[i] != 1 && input->ne[i] != output->ne[i])) { + return false; + } + strides[i] = input->ne[i] == 1 ? 0 : static_cast(input->nb[i] / sizeof(float)); + } + const int64_t extent = static_cast(input->byte_count / sizeof(float)); + if (extent < 1 || extent > 268435456) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kCopyStridedKernel); + static const char * const ne_names[] = { "ne0", "ne1", "ne2", "ne3" }; + static const char * const s_names[] = { "s0", "s1", "s2", "s3" }; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + dispatch.kernel.integer_parameters.emplace(ne_names[i], output->ne[i]); + dispatch.kernel.integer_parameters.emplace(s_names[i], strides[i]); + } + dispatch.kernel.integer_parameters.emplace("source_extent", extent); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + finish(context, match, std::move(dispatch)); + return true; +} + +// ADD / SUB / MUL / DIV of F32 values with any element strides, broadcasting either input (ggml's +// rule: an input dim of 1 against a larger output dim), into a packed output. Registered below the +// packed binary kernels (priority -10), so it only takes what they refuse: strided views mostly. +static bool match_binary_strided(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->inputs.size() != 2) { + return false; + } + int64_t op = -1; + switch (node->op) { + case GGML_OP_ADD: op = 0; break; + case GGML_OP_SUB: op = 1; break; + case GGML_OP_MUL: op = 2; break; + case GGML_OP_DIV: op = 3; break; + default: return false; + } + const Value * lhs = context.graph.values().find(node->inputs[0]); + const Value * rhs = context.graph.values().find(node->inputs[1]); + const Value * output = context.graph.values().find(node->output); + if (lhs == nullptr || rhs == nullptr || output == nullptr || lhs->type != GGML_TYPE_F32 || rhs->type != GGML_TYPE_F32 || + output->type != GGML_TYPE_F32 || !packed(*output, sizeof(float)) || output->alias_source.value >= 0 || + lhs->storage == output->storage || rhs->storage == output->storage || output->element_count > 268435456) { + return false; + } + int64_t a[GGML_MAX_DIMS], b[GGML_MAX_DIMS]; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + for (const Value * v : { lhs, rhs }) { + if (v->nb[i] % sizeof(float) != 0 || (v->ne[i] != 1 && v->ne[i] != output->ne[i])) { + return false; + } + } + a[i] = lhs->ne[i] == 1 ? 0 : static_cast(lhs->nb[i] / sizeof(float)); + b[i] = rhs->ne[i] == 1 ? 0 : static_cast(rhs->nb[i] / sizeof(float)); + } + const int64_t a_extent = static_cast(lhs->byte_count / sizeof(float)); + const int64_t b_extent = static_cast(rhs->byte_count / sizeof(float)); + if (a_extent < 1 || b_extent < 1 || a_extent > 268435456 || b_extent > 268435456) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kBinaryStridedKernel); + static const char * const ne_names[] = { "ne0", "ne1", "ne2", "ne3" }; + static const char * const a_names[] = { "a0", "a1", "a2", "a3" }; + static const char * const b_names[] = { "b0", "b1", "b2", "b3" }; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + dispatch.kernel.integer_parameters.emplace(ne_names[i], output->ne[i]); + dispatch.kernel.integer_parameters.emplace(a_names[i], a[i]); + dispatch.kernel.integer_parameters.emplace(b_names[i], b[i]); + } + dispatch.kernel.integer_parameters.emplace("a_extent", a_extent); + dispatch.kernel.integer_parameters.emplace("b_extent", b_extent); + dispatch.kernel.integer_parameters.emplace("op", op); + dispatch.bindings.push_back({ lhs->id, 0, lhs->byte_count }); + dispatch.bindings.push_back({ rhs->id, 0, rhs->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + finish(context, match, std::move(dispatch)); + return true; +} + +static std::string f32_config(float value) { + std::ostringstream out; + out.precision(9); + out << value; + return out.str(); +} + +// CLAMP of a packed F32 tensor on its own (fused router clamps match first, at their own priority) +static bool match_clamp(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_CLAMP || node->inputs.size() != 1) { + return false; + } + const Value * input = context.graph.values().find(node->inputs[0]); + const Value * output = context.graph.values().find(node->output); + const ClampParams * params = op_params_as(node->params); + if (input == nullptr || output == nullptr || params == nullptr || input->type != GGML_TYPE_F32 || + output->type != GGML_TYPE_F32 || !packed(*input, sizeof(float)) || !packed(*output, sizeof(float)) || + !same_shape(*input, *output) || output->element_count > 268435456) { + return false; + } + // ggml_clamp is in place: the output is a view of the input, same layout + const bool in_place = output->alias_source.value >= 0; + if (in_place ? (output->storage != input->storage || output->storage_offset != input->storage_offset) + : input->storage == output->storage) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(in_place ? kClampInplaceKernel : kClampKernel); + dispatch.kernel.integer_parameters.emplace("element_count", output->element_count); + dispatch.kernel.compile_parameters.emplace("ggml.clamp_f32.min", f32_config(params->min)); + dispatch.kernel.compile_parameters.emplace("ggml.clamp_f32.max", f32_config(params->max)); + if (!in_place) { + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + } + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + finish(context, match, std::move(dispatch)); + return true; +} + +// CPY F32 (any element strides) -> packed F16: ggml_cpy(src, dst) returns a view of dst, so the +// output aliases the second input, which only gives the destination's layout +static bool match_copy_f32_f16(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_CPY || node->inputs.empty() || node->inputs.size() > 2) { + return false; + } + const Value * input = context.graph.values().find(node->inputs[0]); + const Value * output = context.graph.values().find(node->output); + if (input == nullptr || output == nullptr || input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F16 || + !packed(*output, ggml_type_size(GGML_TYPE_F16)) || input->storage == output->storage || + input->element_count != output->element_count || output->element_count > 268435456) { + return false; + } + int64_t strides[GGML_MAX_DIMS]; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (input->nb[i] % sizeof(float) != 0 || input->ne[i] != output->ne[i]) { + return false; + } + strides[i] = static_cast(input->nb[i] / sizeof(float)); + } + const int64_t extent = static_cast(input->byte_count / sizeof(float)); + if (extent < 1 || extent > 268435456) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kCopyF32F16Kernel); + static const char * const ne_names[] = { "ne0", "ne1", "ne2", "ne3" }; + static const char * const s_names[] = { "s0", "s1", "s2", "s3" }; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + dispatch.kernel.integer_parameters.emplace(ne_names[i], output->ne[i]); + dispatch.kernel.integer_parameters.emplace(s_names[i], strides[i]); + } + dispatch.kernel.integer_parameters.emplace("source_extent", extent); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + finish(context, match, std::move(dispatch)); + return true; +} + +// FLASH_ATTN_EXT with the layouts the flash-attention kernels refuse (one contiguous block per head, +// as encoders lay them out): F32 query [d, n_q, h], F16 key/value [d, n_kv, h_kv], F16 mask +// [n_kv, >= n_q], output [dv, h, n_q] packed; no ALiBi, no softcap. Registered below them. +static bool match_attention_strided(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_FLASH_ATTN_EXT || node->inputs.size() != 4) { + return false; + } + const Value * q = context.graph.values().find(node->inputs[0]); + const Value * k = context.graph.values().find(node->inputs[1]); + const Value * v = context.graph.values().find(node->inputs[2]); + const Value * mask = context.graph.values().find(node->inputs[3]); + const Value * out = context.graph.values().find(node->output); + const FlashAttnExtParams * params = op_params_as(node->params); + if (q == nullptr || k == nullptr || v == nullptr || mask == nullptr || out == nullptr || params == nullptr || + q->type != GGML_TYPE_F32 || k->type != GGML_TYPE_F16 || v->type != GGML_TYPE_F16 || mask->type != GGML_TYPE_F16 || + out->type != GGML_TYPE_F32 || params->max_bias != 0.0f || params->logit_softcap != 0.0f) { + return false; + } + const int64_t d = q->ne[0], dv = v->ne[0], nq = q->ne[1], nkv = k->ne[1], nh = q->ne[2], nhkv = k->ne[2]; + if (k->ne[0] != d || v->ne[1] != nkv || v->ne[2] != nhkv || nhkv < 1 || nh % nhkv != 0 || d > 1024 || dv > 1024 || + dv > 1024 || q->ne[3] != 1 || k->ne[3] != 1 || v->ne[3] != 1 || mask->ne[0] < nkv || mask->ne[1] < nq || + mask->ne[2] != 1 || mask->ne[3] != 1 || out->ne[0] != dv || out->ne[1] != nh || out->ne[2] != nq || + out->ne[3] != 1 || !packed(*out, sizeof(float)) || q->nb[0] != sizeof(float) || k->nb[0] != 2 || v->nb[0] != 2 || + mask->nb[0] != 2 || q->nb[1] % 4 || q->nb[2] % 4 || k->nb[1] % 2 || k->nb[2] % 2 || v->nb[1] % 2 || v->nb[2] % 2 || + mask->nb[1] % 2 || nq > 65536 || nkv > 65536 || nh > 1024 || dv < 1 || dv > 1024 || (dv % 32) != 0) { + return false; + } + Dispatch dispatch; + // scores once per (query, head) in workgroup memory when they fit (ONEBIT_HRX_ATTN_PER_LANE=1: the old way) + const bool rows = nkv <= 2048 && d <= 256 && dv <= 256 && std::getenv("ONEBIT_HRX_ATTN_PER_LANE") == nullptr; + dispatch.kernel = make_kernel_specialization(rows ? kAttentionRowsKernel : kAttentionStridedKernel); + auto & ip = dispatch.kernel.integer_parameters; + ip.emplace("qk_size", d); + ip.emplace("v_size", dv); + ip.emplace("q_count", nq); + ip.emplace("kv_count", nkv); + ip.emplace("head_count", nh); + ip.emplace("kv_head_count", nhkv); + ip.emplace("q_s1", static_cast(q->nb[1] / 4)); + ip.emplace("q_s2", static_cast(q->nb[2] / 4)); + ip.emplace("k_s1", static_cast(k->nb[1] / 2)); + ip.emplace("k_s2", static_cast(k->nb[2] / 2)); + ip.emplace("v_s1", static_cast(v->nb[1] / 2)); + ip.emplace("v_s2", static_cast(v->nb[2] / 2)); + ip.emplace("m_s1", static_cast(mask->nb[1] / 2)); + ip.emplace("q_extent", static_cast(q->byte_count / 4)); + ip.emplace("k_extent", static_cast(k->byte_count / 2)); + ip.emplace("v_extent", static_cast(v->byte_count / 2)); + ip.emplace("m_extent", static_cast(mask->byte_count / 2)); + dispatch.kernel.compile_parameters.emplace("ggml.attention_strided.scale", f32_config(params->scale)); + for (const Value * b : { q, k, v, mask, out }) { + dispatch.bindings.push_back({ b->id, 0, b->byte_count }); + } + finish(context, match, std::move(dispatch)); + return true; +} + +// MUL_MAT of a small F16/F32 weight [K, N] (rows packed, any K) or Q8_0 weight (K a multiple of 32) with F32 columns [K, T] into a +// packed [N, T]: heads and projections the tiled matmul kernels refuse (K not a multiple of 256). +// Registered below them; one workitem per output, so it is for small N x T only. +static bool match_mul_mat_small(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_MUL_MAT || node->inputs.size() != 2) { + return false; + } + const Value * w = context.graph.values().find(node->inputs[0]); + const Value * x = context.graph.values().find(node->inputs[1]); + const Value * out = context.graph.values().find(node->output); + const bool q8 = w != nullptr && w->type == GGML_TYPE_Q8_0; + if (w == nullptr || x == nullptr || out == nullptr || + (w->type != GGML_TYPE_F16 && w->type != GGML_TYPE_F32 && !q8) || x->type != GGML_TYPE_F32 || + out->type != GGML_TYPE_F32 || !packed(*out, sizeof(float))) { + return false; + } + // element size, or for Q8_0 the block size (w_s1 then counts blocks) + const size_t wsz = q8 ? ggml_type_size(GGML_TYPE_Q8_0) : ggml_type_size(w->type); + const int64_t k = w->ne[0], n = w->ne[1], t = x->ne[1]; + if (x->ne[0] != k || w->ne[2] != 1 || w->ne[3] != 1 || x->ne[2] != 1 || x->ne[3] != 1 || out->ne[0] != n || + out->ne[1] != t || out->ne[2] != 1 || out->ne[3] != 1 || (!q8 && w->nb[0] != wsz) || w->nb[1] % wsz != 0 || + (q8 && k % 32 != 0) || x->nb[0] != sizeof(float) || x->nb[1] % sizeof(float) != 0 || n * t > (1 << 22) || + k > 1048576) { + return false; + } + Dispatch dispatch; + // Q8_0 with enough K: a workgroup per output, lanes over the blocks (ONEBIT_HRX_Q8_PER_OUTPUT=1: the old way) + const bool rows = q8 && k >= 32 * 64 && std::getenv("ONEBIT_HRX_Q8_PER_OUTPUT") == nullptr; + dispatch.kernel = make_kernel_specialization(rows ? kMulMatRowsQ8Kernel + : q8 ? kMulMatSmallQ8Kernel + : w->type == GGML_TYPE_F16 ? kMulMatSmallF16Kernel : kMulMatSmallF32Kernel); + auto & ip = dispatch.kernel.integer_parameters; + ip.emplace("k_size", k); + ip.emplace("n_size", n); + ip.emplace("t_count", t); + ip.emplace("w_s1", static_cast(w->nb[1] / wsz)); + ip.emplace("x_s1", static_cast(x->nb[1] / sizeof(float))); + ip.emplace("w_extent", static_cast(w->byte_count / wsz)); + ip.emplace("x_extent", static_cast(x->byte_count / sizeof(float))); + for (const Value * b : { w, x, out }) { + dispatch.bindings.push_back({ b->id, 0, b->byte_count }); + } + finish(context, match, std::move(dispatch)); + return true; +} + +// the node producing `value` when it is `op` and `value` has no other consumer +static const GraphNode * sole_producer(const Graph & graph, ValueId value, ggml_op op) { + const GraphNode * node = graph.index().producer(value); + return node != nullptr && node->op == op && graph.index().has_single_consumer(value) ? node : nullptr; +} + +// Rotate-half RoPE lowered to eight nodes (ModernBERT through ggmlc): +// out = ADD(MUL(x, cos), MUL(CONT(CONCAT(NEG(CONT(x[half:])), CONT(x[:half]))), sin)) +// with x packed F32 [d, T, H] and cos/sin [d, T] broadcast over heads: one kernel for all eight. +static bool match_rope_rotate_half(const DispatchMatchContext & context, DispatchMatch & match) { + // rooted at x * cos, the chain's first node in graph order (a fused match covers later nodes only) + const Graph & graph = context.graph; + const GraphNode * mul_cos = context.root_node; + if (mul_cos == nullptr || mul_cos->op != GGML_OP_MUL || mul_cos->inputs.size() != 2 || !graph.has_index() || + !graph.index().has_single_consumer(mul_cos->output)) { + return false; + } + const GraphNode * add = graph.index().consumers(mul_cos->output).front(); + if (add == nullptr || add->op != GGML_OP_ADD || add->inputs.size() != 2) { + return false; + } + for (int order = 0; order < 2; ++order) { + if (add->inputs[order] != mul_cos->output) { + continue; + } + const GraphNode * mul_sin = sole_producer(graph, add->inputs[1 - order], GGML_OP_MUL); + if (mul_sin == nullptr || mul_sin->inputs.size() != 2) { + continue; + } + const GraphNode * cat_cont = sole_producer(graph, mul_sin->inputs[0], GGML_OP_CONT); + if (cat_cont == nullptr) { + continue; + } + const GraphNode * concat = sole_producer(graph, cat_cont->inputs[0], GGML_OP_CONCAT); + if (concat == nullptr || concat->inputs.size() != 2) { + continue; + } + const GraphNode * neg = sole_producer(graph, concat->inputs[0], GGML_OP_UNARY); + const GraphNode * lo_cont = sole_producer(graph, concat->inputs[1], GGML_OP_CONT); + if (neg == nullptr || lo_cont == nullptr || neg->inputs.size() != 1) { + continue; + } + const UnaryParams * neg_params = op_params_as(neg->params); + const GraphNode * hi_cont = sole_producer(graph, neg->inputs[0], GGML_OP_CONT); + if (neg_params == nullptr || neg_params->op != UnaryKind::Neg || hi_cont == nullptr) { + continue; + } + const Value * x = graph.values().find(mul_cos->inputs[0]); + const Value * cs = graph.values().find(mul_cos->inputs[1]); + const Value * sn = graph.values().find(mul_sin->inputs[1]); + const Value * hi = graph.values().find(hi_cont->inputs[0]); + const Value * lo = graph.values().find(lo_cont->inputs[0]); + const Value * out = graph.values().find(add->output); + if (x == nullptr || cs == nullptr || sn == nullptr || hi == nullptr || lo == nullptr || out == nullptr) { + continue; + } + const int64_t d = x->ne[0], t = x->ne[1], h = x->ne[2], half = d / 2; + bool ok = x->type == GGML_TYPE_F32 && cs->type == GGML_TYPE_F32 && sn->type == GGML_TYPE_F32 && + out->type == GGML_TYPE_F32 && packed(*x, sizeof(float)) && packed(*out, sizeof(float)) && + same_shape(*x, *out) && x->ne[3] == 1 && d % 2 == 0 && d <= 4096 && h <= 4096 && + out->storage != x->storage && out->alias_source.value < 0; + // the two halves are views of x at its offset and half a row further, with x's strides + ok = ok && hi->storage == x->storage && lo->storage == x->storage && hi->ne[0] == half && lo->ne[0] == half && + lo->storage_offset == x->storage_offset && hi->storage_offset == x->storage_offset + half * sizeof(float) && + hi->nb == x->nb && lo->nb == x->nb && hi->ne[1] == t && hi->ne[2] == h && lo->ne[1] == t && lo->ne[2] == h; + // cos and sin: [d, T], rows possibly strided, broadcast over heads + for (const Value * v : { cs, sn }) { + ok = ok && v->ne[0] == d && v->ne[1] == t && v->ne[2] == 1 && v->ne[3] == 1 && v->nb[0] == sizeof(float) && + v->nb[1] % sizeof(float) == 0 && v->storage != out->storage; + } + if (!ok) { + continue; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kRopeRotateHalfKernel); + auto & ip = dispatch.kernel.integer_parameters; + ip.emplace("ne0", d); + ip.emplace("ne1", t); + ip.emplace("ne2", h); + ip.emplace("cos_s1", static_cast(cs->nb[1] / sizeof(float))); + ip.emplace("sin_s1", static_cast(sn->nb[1] / sizeof(float))); + ip.emplace("cos_extent", static_cast(cs->byte_count / sizeof(float))); + ip.emplace("sin_extent", static_cast(sn->byte_count / sizeof(float))); + for (const Value * b : { x, cs, sn, out }) { + dispatch.bindings.push_back({ b->id, 0, b->byte_count }); + } + match.covered_nodes.push_back(context.root_index); + for (const GraphNode * n : { hi_cont, neg, lo_cont, concat, cat_cont, mul_sin, add }) { + if (!append_covered_node_index_once(graph, context.covered_nodes, n, match.covered_nodes)) { + return false; + } + } + match.dispatches.push_back(std::move(dispatch)); + return true; + } + return false; +} + +// GEGLU lowered to CONT(gate view) -> GELU -> MUL(., up view), rooted at the CONT: one kernel +// reads both strided halves and writes gelu(gate) * up. +static bool match_geglu_strided(const DispatchMatchContext & context, DispatchMatch & match) { + const Graph & graph = context.graph; + const GraphNode * cont = context.root_node; + if (cont == nullptr || cont->op != GGML_OP_CONT || cont->inputs.size() != 1 || !graph.has_index() || + !graph.index().has_single_consumer(cont->output)) { + return false; + } + const GraphNode * gelu = graph.index().consumers(cont->output).front(); + const UnaryParams * gp = gelu != nullptr && gelu->op == GGML_OP_UNARY ? op_params_as(gelu->params) : nullptr; + if (gp == nullptr || gp->op != UnaryKind::Gelu || !graph.index().has_single_consumer(gelu->output)) { + return false; + } + const GraphNode * mul = graph.index().consumers(gelu->output).front(); + if (mul == nullptr || mul->op != GGML_OP_MUL || mul->inputs.size() != 2 || mul->inputs[0] != gelu->output) { + return false; + } + const Value * a = graph.values().find(cont->inputs[0]); + const Value * b = graph.values().find(mul->inputs[1]); + const Value * out = graph.values().find(mul->output); + if (a == nullptr || b == nullptr || out == nullptr || a->type != GGML_TYPE_F32 || b->type != GGML_TYPE_F32 || + out->type != GGML_TYPE_F32 || !packed(*out, sizeof(float)) || !same_shape(*a, *out) || !same_shape(*b, *out) || + out->ne[2] != 1 || out->ne[3] != 1 || a->nb[0] != sizeof(float) || b->nb[0] != sizeof(float) || + a->nb[1] % sizeof(float) != 0 || b->nb[1] % sizeof(float) != 0 || out->storage == a->storage || + out->storage == b->storage || out->alias_source.value >= 0) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kGegluStridedKernel); + auto & ip = dispatch.kernel.integer_parameters; + ip.emplace("n_size", out->ne[0]); + ip.emplace("t_count", out->ne[1]); + ip.emplace("a_s1", static_cast(a->nb[1] / sizeof(float))); + ip.emplace("b_s1", static_cast(b->nb[1] / sizeof(float))); + ip.emplace("a_extent", static_cast(a->byte_count / sizeof(float))); + ip.emplace("b_extent", static_cast(b->byte_count / sizeof(float))); + for (const Value * v : { a, b, out }) { + dispatch.bindings.push_back({ v->id, 0, v->byte_count }); + } + match.covered_nodes.push_back(context.root_index); + for (const GraphNode * n : { gelu, mul }) { + if (!append_covered_node_index_once(graph, context.covered_nodes, n, match.covered_nodes)) { + return false; + } + } + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_small_rows_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ "common.softmax_rows_f32", GGML_OP_SOFT_MAX, DispatchMatchKind::SingleOp, 0, DispatchSource::Common, + match_softmax_rows }); + registry.add({ "common.binary_strided_f32.add", GGML_OP_ADD, DispatchMatchKind::SingleOp, -10, DispatchSource::Common, + match_binary_strided }); + registry.add({ "common.binary_strided_f32.sub", GGML_OP_SUB, DispatchMatchKind::SingleOp, -10, DispatchSource::Common, + match_binary_strided }); + registry.add({ "common.binary_strided_f32.mul", GGML_OP_MUL, DispatchMatchKind::SingleOp, -10, DispatchSource::Common, + match_binary_strided }); + registry.add({ "common.binary_strided_f32.div", GGML_OP_DIV, DispatchMatchKind::SingleOp, -10, DispatchSource::Common, + match_binary_strided }); + registry.add({ "common.geglu_strided_f32", GGML_OP_CONT, DispatchMatchKind::Fused, 300, DispatchSource::Common, + match_geglu_strided }); + registry.add({ "common.rope_rotate_half_f32", GGML_OP_MUL, DispatchMatchKind::Fused, 300, DispatchSource::Common, + match_rope_rotate_half }); + registry.add({ "common.mul_mat_small_f32", GGML_OP_MUL_MAT, DispatchMatchKind::SingleOp, -10, DispatchSource::Common, + match_mul_mat_small }); + registry.add({ "common.attention_strided_f32_f16", GGML_OP_FLASH_ATTN_EXT, DispatchMatchKind::SingleOp, -10, + DispatchSource::Common, match_attention_strided }); + registry.add({ "common.copy_strided_f32_f16", GGML_OP_CPY, DispatchMatchKind::SingleOp, -10, DispatchSource::Common, + match_copy_f32_f16 }); + registry.add({ "common.clamp_f32", GGML_OP_CLAMP, DispatchMatchKind::SingleOp, -10, DispatchSource::Common, + match_clamp }); + registry.add({ "common.norm_rows_f32", GGML_OP_NORM, DispatchMatchKind::SingleOp, 0, DispatchSource::Common, + match_norm_rows }); + registry.add({ "common.sum_rows_f32", GGML_OP_SUM_ROWS, DispatchMatchKind::SingleOp, 0, DispatchSource::Common, + match_sum_rows }); + registry.add({ "common.argsort_rows_f32", GGML_OP_ARGSORT, DispatchMatchKind::SingleOp, 0, DispatchSource::Common, + match_argsort_rows }); + registry.add({ "common.get_rows_small_f32", GGML_OP_GET_ROWS, DispatchMatchKind::SingleOp, 0, + DispatchSource::Common, match_get_rows_small }); + registry.add({ "common.copy_strided_f32", GGML_OP_CONT, DispatchMatchKind::SingleOp, 0, DispatchSource::Common, + match_copy_strided }); + registry.add({ "common.repeat_broadcast_f32", GGML_OP_REPEAT, DispatchMatchKind::SingleOp, 0, DispatchSource::Common, + match_repeat_broadcast }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.h new file mode 100644 index 000000000000..76b6464e177e --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.h @@ -0,0 +1,26 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +namespace ggml::hrx { + +// SOFT_MAX (no mask, scale 1), SUM_ROWS, NORM, ARGSORT and narrow GET_ROWS on short F32 rows, such +// as a MoE router's, CONT of strided F32 views and broadcast REPEAT: ggml_softmax_rows_f32, ggml_sum_rows_f32, +// ggml_argsort_rows_f32, ggml_norm_rows_f32, ggml_get_rows_small_f32 and ggml_copy_strided_f32. +void register_small_rows_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-softplus.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-softplus.cpp new file mode 100644 index 000000000000..ca0d4cd5aef4 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-softplus.cpp @@ -0,0 +1,74 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// UNARY softplus on contiguous F32 values (ops/softplus_f32.loom). common.unary_f32 covers the +// exact-math unary kinds only; without this, a standalone softplus (Qwen3.5/3.8 delta-net gates in +// multi-sequence batches) is claimed by HRX and then fails as an unsupported node. + +#include "dispatch-softplus.h" + +#include "graph/op-params.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include + +namespace ggml::hrx { + +namespace { + +static constexpr KernelCatalogRef kSoftplusF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_softplus_f32"); + +bool match_softplus_f32(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->inputs.size() != 1) { + return false; + } + const UnaryParams * params = op_params_as(node->params); + if (params == nullptr || params->op != UnaryKind::SoftPlus) { + return false; + } + const Value * input = context.graph.values().find(node->inputs[0]); + const Value * output = context.graph.values().find(node->output); + if (input == nullptr || output == nullptr || input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || + !input->contiguous || !output->contiguous || input->element_count != output->element_count || + output->element_count <= 0 || static_cast(output->element_count) > (uint64_t{ 1 } << 27)) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kSoftplusF32Kernel); + dispatch.kernel.integer_parameters.emplace("element_count", output->element_count); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_softplus_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "extra.softplus_f32", + GGML_OP_UNARY, + DispatchMatchKind::SingleOp, + 1, + DispatchSource::Common, + match_softplus_f32, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-softplus.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-softplus.h new file mode 100644 index 000000000000..6867f23a4e9f --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-softplus.h @@ -0,0 +1,24 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +namespace ggml::hrx { + +void register_softplus_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-swiglu-oai.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-swiglu-oai.cpp new file mode 100644 index 000000000000..e89a7f21114e --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-swiglu-oai.cpp @@ -0,0 +1,87 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// GGML_GLU_OP_SWIGLU_OAI on F32 with gate and up as two contiguous tensors of the same shape +// (ops/swiglu_oai_f32.loom): gpt-oss's clamped SwiGLU between the expert up/gate and down projections. + +#include "dispatch-swiglu-oai.h" + +#include "dispatch-mul-mat-common.h" +#include "graph/op-params.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include + +namespace ggml::hrx { + +namespace { + +static constexpr KernelCatalogRef kSwigluOaiF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_swiglu_oai_f32"); + +bool same_shape(const Value & a, const Value & b) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (a.ne[i] != b.ne[i]) { + return false; + } + } + return true; +} + +bool match_swiglu_oai_f32(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->inputs.size() != 2) { + return false; + } + const GluParams * params = op_params_as(node->params); + if (params == nullptr || params->op != GGML_GLU_OP_SWIGLU_OAI) { + return false; + } + const Value * gate = context.graph.values().find(node->inputs[0]); + const Value * up = context.graph.values().find(node->inputs[1]); + const Value * output = context.graph.values().find(node->output); + if (gate == nullptr || up == nullptr || output == nullptr || gate->type != GGML_TYPE_F32 || + up->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || !gate->contiguous || !up->contiguous || + !output->contiguous || !same_shape(*gate, *up) || !same_shape(*gate, *output) || output->element_count <= 0 || + static_cast(output->element_count) > (uint64_t{ 1 } << 27)) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kSwigluOaiF32Kernel); + dispatch.kernel.integer_parameters.emplace("element_count", output->element_count); + dispatch.kernel.compile_parameters.emplace("ggml.swiglu_oai.alpha", common_to_config_value(params->alpha)); + dispatch.kernel.compile_parameters.emplace("ggml.swiglu_oai.limit", common_to_config_value(params->limit)); + dispatch.bindings.push_back({ gate->id, 0, gate->byte_count }); + dispatch.bindings.push_back({ up->id, 0, up->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_swiglu_oai_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "extra.swiglu_oai_f32", + GGML_OP_GLU, + DispatchMatchKind::SingleOp, + 1, + DispatchSource::Common, + match_swiglu_oai_f32, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-swiglu-oai.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-swiglu-oai.h new file mode 100644 index 000000000000..6c3cb783990b --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-swiglu-oai.h @@ -0,0 +1,24 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +namespace ggml::hrx { + +void register_swiglu_oai_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-unary.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-unary.cpp new file mode 100644 index 000000000000..4622b109141e --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-unary.cpp @@ -0,0 +1,150 @@ +#include "dispatch-unary.h" + +#include "ggml.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kUnaryF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_unary_f32"); +static constexpr KernelCatalogRef kScaleBiasF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_scale_bias_f32"); + +static bool same_shape(const Value & lhs, const Value & rhs) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (lhs.ne[i] != rhs.ne[i]) { + return false; + } + } + return true; +} + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static std::string float_config_value(float value) { + std::ostringstream stream; + stream << std::setprecision(std::numeric_limits::max_digits10) << value; + return stream.str(); +} + +static bool scale_storage_is_safe(const Value & input, const Value & output) { + if (input.storage_root != output.storage_root) { + return true; + } + if (input.storage_offset == output.storage_offset && input.byte_count == output.byte_count) { + return true; + } + if (input.storage_offset > std::numeric_limits::max() - input.byte_count || + output.storage_offset > std::numeric_limits::max() - output.byte_count) { + return false; + } + return input.storage_offset + input.byte_count <= output.storage_offset || + output.storage_offset + output.byte_count <= input.storage_offset; +} + +static bool match_scale_bias_f32_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->inputs.size() != 1) { + return false; + } + + const ScaleParams * params = op_params_as(node->params); + const Value * input = graph_value(context.graph, node->inputs[0]); + const Value * output = graph_value(context.graph, node->output); + if (params == nullptr || input == nullptr || output == nullptr || !std::isfinite(params->scale) || + !std::isfinite(params->bias) || input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || + !input->contiguous || !output->contiguous || !same_shape(*input, *output) || + !scale_storage_is_safe(*input, *output) || output->element_count <= 0 || + static_cast(output->element_count) > std::numeric_limits::max()) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kScaleBiasF32Kernel); + dispatch.kernel.integer_parameters.emplace("element_count", output->element_count); + dispatch.kernel.compile_parameters.emplace("ggml.scale.scale", float_config_value(params->scale)); + dispatch.kernel.compile_parameters.emplace("ggml.scale.bias", float_config_value(params->bias)); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + + match.covered_nodes.push_back(context.root_index); + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_unary_f32_dispatch(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->inputs.size() != 1) { + return false; + } + + const UnaryParams * params = op_params_as(node->params); + if (params == nullptr || !unary_kind_supported(params->op)) { + return false; + } + + const Value * output = graph_value(context.graph, node->output); + const Value * input = graph_value(context.graph, node->inputs[0]); + if (output == nullptr || input == nullptr) { + return false; + } + + if (output->type != GGML_TYPE_F32 || input->type != GGML_TYPE_F32 || !same_shape(*output, *input) || + !output->contiguous || !input->contiguous || output->element_count <= 0 || + static_cast(output->element_count) > std::numeric_limits::max()) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kUnaryF32Kernel); + dispatch.kernel.integer_parameters.emplace("element_count", output->element_count); + dispatch.kernel.compile_parameters.emplace("ggml.unary_f32.op", + std::to_string(unary_kind_config_value(params->op))); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + + match.covered_nodes.push_back(context.root_index); + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static void register_unary_dispatch_for(DispatchRegistryBuilder & registry, ggml_op root_op) { + registry.add({ + "common.unary_f32", + root_op, + DispatchMatchKind::SingleOp, + 0, + DispatchSource::Common, + match_unary_f32_dispatch, + }); +} + +} // namespace + +void register_unary_dispatch(DispatchRegistryBuilder & registry) { + registry.add({ + "common.scale_bias_f32", + GGML_OP_SCALE, + DispatchMatchKind::SingleOp, + 0, + DispatchSource::Common, + match_scale_bias_f32_dispatch, + }); + register_unary_dispatch_for(registry, GGML_OP_UNARY); + register_unary_dispatch_for(registry, GGML_OP_SQR); + register_unary_dispatch_for(registry, GGML_OP_SQRT); + register_unary_dispatch_for(registry, GGML_OP_LOG); + register_unary_dispatch_for(registry, GGML_OP_SIN); + register_unary_dispatch_for(registry, GGML_OP_COS); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-unary.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-unary.h new file mode 100644 index 000000000000..c42d89d69f0b --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-unary.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_unary_dispatch(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-zaya-cca-conv.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-zaya-cca-conv.cpp new file mode 100644 index 000000000000..f28907c3220f --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-zaya-cca-conv.cpp @@ -0,0 +1,287 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// ZAYA's CCA convolution at decode (one token, one sequence) as one dispatch +// (ops/zaya_cca_conv_decode_f32.loom) instead of the 13 that its ~35 graph nodes take. Matched from +// the Q/K concat, following the exact decode structure of src/models/zaya.cpp: +// +// QKraw = CONCAT(Qraw, Kraw) -> reshape(s) -> conv_input = CONCAT(conv_state, .) [3, C] +// conv_input -> VIEW (steps 1..2) -> CONT -> RESHAPE -> CPY into the conv-state cache +// conv_input -> SSM_CONV(dw) -> ADD(dw bias) = QK_dw [C, 2] +// per tap t: QK_dw -> VIEW(step t) -> PERMUTE -> CONT -> RESHAPE -> MUL_MAT(W_t, .) [128, 1, G] +// ADD(tap 0, tap 1) -> RESHAPE -> PERMUTE -> CONT -> RESHAPE -> ADD(grp bias) = QK_grp [C] +// +// Anything else (prefill shapes, several sequences, other layouts) is left to the generic path. + +#include "dispatch-zaya-cca-conv.h" + +#include "dispatch-mul-mat-common.h" +#include "graph/graph-matcher.h" + +#include +#include + +namespace ggml::hrx { + +namespace { + +static constexpr KernelCatalogRef kZayaCcaConvKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_zaya_cca_conv_decode_f32"); + +constexpr int64_t kGroupSize = 128; + +struct Walk { + const Graph & graph; + std::vector nodes; + + const Value * value(ValueId id) const { return common_graph_value(graph, id); } + + // The only consumer of |id|, when there is exactly one. + const GraphNode * only_consumer(ValueId id) const { + const std::vector & consumers = graph.index().consumers(id); + return consumers.size() == 1 ? consumers.front() : nullptr; + } + + // Follow single-consumer layout aliases (reshape/view/permute) from |id|; returns the last value. + ValueId skip_aliases(ValueId id) { + for (;;) { + const GraphNode * next = only_consumer(id); + if (next == nullptr || !is_layout_alias_node(graph, *next)) { + return id; + } + nodes.push_back(next); + id = next->output; + } + } + + // |id|'s only consumer, which must be |op| with |id| as input |slot|. + const GraphNode * expect(ValueId id, ggml_op op, size_t slot) { + const GraphNode * next = only_consumer(id); + if (next == nullptr || next->op != op || next->inputs.size() <= slot || next->inputs[slot] != id) { + return nullptr; + } + nodes.push_back(next); + return next; + } +}; + +bool is_f32(const Value * v, int64_t ne0, int64_t ne1) { + return v != nullptr && v->type == GGML_TYPE_F32 && v->contiguous && v->ne[0] == ne0 && v->ne[1] == ne1 && + v->ne[2] == 1 && v->ne[3] == 1; +} + +// One tap: QK_dw -> VIEW -> PERMUTE -> CONT -> RESHAPE -> MUL_MAT(W view, .). Returns the MUL_MAT +// and its weight view, and the step the view reads (its byte offset over QK_dw's row stride). +const GraphNode * match_tap(Walk & walk, const GraphNode * view, const Value & qk_dw, int64_t channels, + int64_t & step, const Value *& weight_view) { + const Value * v = walk.value(view->output); + if (v == nullptr || !is_layout_alias_node(walk.graph, *view) || v->ne[0] != kGroupSize || + v->ne[1] != channels / kGroupSize || v->storage_offset < qk_dw.storage_offset || + (v->storage_offset - qk_dw.storage_offset) % qk_dw.nb[1] != 0) { + return nullptr; + } + step = static_cast((v->storage_offset - qk_dw.storage_offset) / qk_dw.nb[1]); + walk.nodes.push_back(view); + const GraphNode * permute = walk.expect(view->output, GGML_OP_PERMUTE, 0); + const GraphNode * cont = permute != nullptr ? walk.expect(permute->output, GGML_OP_CONT, 0) : nullptr; + if (cont == nullptr) { + return nullptr; + } + const ValueId x = walk.skip_aliases(cont->output); + const GraphNode * mul = walk.expect(x, GGML_OP_MUL_MAT, 1); + if (mul == nullptr) { + return nullptr; + } + weight_view = walk.value(mul->inputs[0]); + return mul; +} + +bool match_zaya_cca_conv(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const Graph & graph = context.graph; + const GraphNode * root = context.root_node; + if (root == nullptr || !graph.has_index() || root->op != GGML_OP_CONCAT || root->inputs.size() != 2) { + return false; + } + Walk walk{ graph, { root } }; + const Value * q = walk.value(root->inputs[0]); + const Value * k = walk.value(root->inputs[1]); + const Value * qk = walk.value(root->output); + if (q == nullptr || k == nullptr || !is_f32(q, q->ne[0], 1) || !is_f32(k, k->ne[0], 1)) { + return false; + } + const int64_t channels = q->ne[0] + k->ne[0]; + if (!is_f32(qk, channels, 1) || channels % kGroupSize != 0 || channels < kGroupSize || channels > 8192) { + return false; + } + + // conv_input = CONCAT(conv_state, QKraw reshaped to [1, C]): [3, C]. + const ValueId qk_col = walk.skip_aliases(root->output); + const GraphNode * cat = walk.expect(qk_col, GGML_OP_CONCAT, 1); + if (cat == nullptr) { + return false; + } + const Value * state = walk.value(cat->inputs[0]); + const Value * input = walk.value(cat->output); + if (state == nullptr || state->type != GGML_TYPE_F32 || !state->contiguous || state->ne[0] != 2 || + state->ne[1] != channels || state->element_count != 2 * channels || input == nullptr || + input->ne[0] != 3 || input->ne[1] != channels || input->ne[2] != 1 || input->ne[3] != 1) { + return false; + } + + // Its two consumers: the state-update view and the depthwise conv. + const std::vector & input_users = graph.index().consumers(cat->output); + if (input_users.size() != 2) { + return false; + } + const GraphNode * state_view = nullptr; + const GraphNode * conv = nullptr; + for (const GraphNode * user : input_users) { + if (user->op == GGML_OP_SSM_CONV && user->inputs.size() == 2 && user->inputs[0] == cat->output) { + conv = user; + } else if (is_layout_alias_node(graph, *user)) { + state_view = user; + } + } + if (conv == nullptr || state_view == nullptr) { + return false; + } + + // State update: VIEW (steps 1..2) -> CONT -> RESHAPE -> CPY into the cache. + const Value * last_states = walk.value(state_view->output); + if (last_states == nullptr || last_states->ne[0] != 2 || last_states->ne[1] != channels || + last_states->storage_offset != input->storage_offset + input->nb[0]) { + return false; + } + walk.nodes.push_back(state_view); + const GraphNode * state_cont = walk.expect(state_view->output, GGML_OP_CONT, 0); + if (state_cont == nullptr) { + return false; + } + const GraphNode * state_copy = walk.expect(walk.skip_aliases(state_cont->output), GGML_OP_CPY, 0); + const Value * new_state = state_copy != nullptr ? walk.value(state_copy->output) : nullptr; + if (new_state == nullptr || new_state->type != GGML_TYPE_F32 || !new_state->contiguous || + new_state->element_count != 2 * channels) { + return false; + } + + // Depthwise conv + bias. + walk.nodes.push_back(conv); + const Value * dw = walk.value(conv->inputs[1]); + const GraphNode * dw_add = walk.expect(conv->output, GGML_OP_ADD, 0); + if (dw == nullptr || dw->type != GGML_TYPE_F32 || !dw->contiguous || dw->ne[0] != 2 || dw->ne[1] != channels || + dw_add == nullptr || !common_binary_node_is_add(*dw_add)) { + return false; + } + const Value * dw_bias = walk.value(dw_add->inputs[1]); + const Value * qk_dw = walk.value(dw_add->output); + if (dw_bias == nullptr || dw_bias->type != GGML_TYPE_F32 || !dw_bias->contiguous || + dw_bias->element_count != channels || !is_f32(qk_dw, channels, 2)) { + return false; + } + + // Two taps, one per step, each into one grouped matmul. + const std::vector & tap_views = graph.index().consumers(dw_add->output); + if (tap_views.size() != 2) { + return false; + } + const GraphNode * tap_mul[2] = { nullptr, nullptr }; + const Value * tap_weight[2] = { nullptr, nullptr }; + for (const GraphNode * view : tap_views) { + int64_t step = -1; + const Value * weight_view = nullptr; + const GraphNode * mul = match_tap(walk, view, *qk_dw, channels, step, weight_view); + if (mul == nullptr || step < 0 || step > 1 || tap_mul[step] != nullptr) { + return false; + } + tap_mul[step] = mul; + tap_weight[step] = weight_view; + } + + // Both weight views are slices of one F16 [128, C, 2] tensor: tap t at t * nb[2]. + const Value * weight = tap_weight[0] != nullptr ? walk.value(tap_weight[0]->storage_root) : nullptr; + if (weight == nullptr || tap_weight[1] == nullptr || weight->type != GGML_TYPE_F16 || !weight->contiguous || + weight->ne[0] != kGroupSize || weight->ne[1] != channels || weight->ne[2] != 2 || weight->ne[3] != 1 || + tap_weight[1]->storage_root != tap_weight[0]->storage_root || + tap_weight[0]->storage_offset != weight->storage_offset || + tap_weight[1]->storage_offset != weight->storage_offset + weight->nb[2] || + tap_weight[0]->ne[0] != kGroupSize || tap_weight[0]->ne[1] != kGroupSize || + tap_weight[0]->ne[2] != channels / kGroupSize) { + return false; + } + + // ADD(tap 0, tap 1) -> RESHAPE -> PERMUTE -> CONT -> RESHAPE -> ADD(grp bias). + const GraphNode * sum = walk.only_consumer(tap_mul[0]->output); + if (sum == nullptr || !common_binary_node_is_add(*sum) || sum->inputs[0] != tap_mul[0]->output || + sum->inputs[1] != tap_mul[1]->output || walk.only_consumer(tap_mul[1]->output) != sum) { + return false; + } + walk.nodes.push_back(sum); + const ValueId sum_view = walk.skip_aliases(sum->output); + const GraphNode * sum_cont = walk.expect(sum_view, GGML_OP_CONT, 0); + if (sum_cont == nullptr) { + return false; + } + const GraphNode * bias_add = walk.expect(walk.skip_aliases(sum_cont->output), GGML_OP_ADD, 0); + if (bias_add == nullptr || !common_binary_node_is_add(*bias_add)) { + return false; + } + const Value * grp_bias = walk.value(bias_add->inputs[1]); + const Value * output = walk.value(bias_add->output); + if (grp_bias == nullptr || grp_bias->type != GGML_TYPE_F32 || !grp_bias->contiguous || + grp_bias->element_count != channels || !is_f32(output, channels, 1) || output->alias_source.value >= 0) { + return false; + } + + for (const GraphNode * node : walk.nodes) { + if (!append_covered_node_index_once(graph, context.covered_nodes, node, dispatch_match.covered_nodes)) { + dispatch_match.covered_nodes.clear(); + return false; + } + } + + const auto source = [](const Value * v) { + return DispatchBinding{ v->storage_root, v->storage_offset, v->byte_count }; + }; + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kZayaCcaConvKernel); + dispatch.kernel.compile_parameters.emplace("ggml.zaya_cca_conv.channels", std::to_string(channels)); + dispatch.kernel.compile_parameters.emplace("ggml.zaya_cca_conv.q_size", std::to_string(q->ne[0])); + dispatch.bindings.push_back(source(q)); + dispatch.bindings.push_back(source(k)); + dispatch.bindings.push_back(source(state)); + dispatch.bindings.push_back(source(dw)); + dispatch.bindings.push_back(source(dw_bias)); + dispatch.bindings.push_back({ weight->storage_root, weight->storage_offset, weight->byte_count }); + dispatch.bindings.push_back(source(grp_bias)); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + dispatch.bindings.push_back({ new_state->id, 0, new_state->byte_count }); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_zaya_cca_conv_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "zaya.cca_conv.decode_f32", + GGML_OP_CONCAT, + DispatchMatchKind::Fused, + 500, + DispatchSource::Common, + match_zaya_cca_conv, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-zaya-cca-conv.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-zaya-cca-conv.h new file mode 100644 index 000000000000..abc04b165d37 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-zaya-cca-conv.h @@ -0,0 +1,24 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +namespace ggml::hrx { + +void register_zaya_cca_conv_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-zaya-cca-qk-norm.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-zaya-cca-qk-norm.cpp new file mode 100644 index 000000000000..890335b41cab --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-zaya-cca-qk-norm.cpp @@ -0,0 +1,270 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// ZAYA's CCA query/key mixing and normalization at decode (one token) as one dispatch +// (ops/zaya_cca_qk_norm_decode_f32.loom), between the CCA convolution and RoPE. Matched from the +// copy of the convolution output's query slice, following src/models/zaya.cpp: +// +// Qcur_pre_rope = RMS_NORM(C_q + SCALE(Qpre + REPEAT(Kpre), 0.5)) +// Kcur_pre_rope = RMS_NORM(C_k + SCALE(SCALE(SUM_ROWS(CONT(PERMUTE(Qgroup))), 1/gqa) + Kpre, 0.5)) +// * k_scale +// +// with C_q, C_k the query and key slices of the convolution output and Qpre, Kpre the plain Q/K +// projections. Every node of both chains comes after the convolution in graph order, so the +// dispatch sits at its first node with all inputs ready and both outputs still ahead of RoPE. + +#include "dispatch-zaya-cca-qk-norm.h" + +#include "dispatch-mul-mat-common.h" +#include "graph/graph-matcher.h" +#include "graph/op-params.h" + +#include +#include +#include +#include +#include + +namespace ggml::hrx { + +namespace { + +static constexpr KernelCatalogRef kZayaCcaQkNormKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_zaya_cca_qk_norm_decode_f32"); + +struct Chain { + const Graph & graph; + std::vector nodes; + + const Value * value(ValueId id) const { return common_graph_value(graph, id); } + + const GraphNode * only_consumer(ValueId id) const { + const std::vector & consumers = graph.index().consumers(id); + return consumers.size() == 1 ? consumers.front() : nullptr; + } + + // Forward through single-consumer layout aliases. + ValueId skip_aliases(ValueId id) { + for (;;) { + const GraphNode * next = only_consumer(id); + if (next == nullptr || !is_layout_alias_node(graph, *next)) { + return id; + } + nodes.push_back(next); + id = next->output; + } + } + + // Backward through layout aliases to the value they view. + ValueId source(ValueId id) { + for (;;) { + const GraphNode * producer = graph.index().producer(id); + if (producer == nullptr || !is_layout_alias_node(graph, *producer)) { + return id; + } + nodes.push_back(producer); + id = producer->inputs[0]; + } + } + + // The non-alias producer of |id| (after skipping aliases), which must be |op| and consumed only + // along this chain. + const GraphNode * producer(ValueId id, ggml_op op) { + const ValueId base = source(id); + const GraphNode * p = graph.index().producer(base); + if (p == nullptr || p->op != op || only_consumer(base) == nullptr) { + return nullptr; + } + nodes.push_back(p); + return p; + } +}; + +bool scale_is(const GraphNode * node, float value) { + const ScaleParams * params = node != nullptr ? op_params_as(node->params) : nullptr; + return params != nullptr && params->bias == 0.0f && std::fabs(params->scale - value) <= 1e-7f * std::fabs(value); +} + +bool is_rows(const Value * v, int64_t ne0, int64_t ne1) { + return v != nullptr && v->type == GGML_TYPE_F32 && v->contiguous && v->ne[0] == ne0 && v->ne[1] == ne1 && + v->ne[2] == 1 && v->ne[3] == 1; +} + +// |add| = ADD(conv slice, SCALE(mean, 0.5)); returns the SCALE's input (the mean's sum). +const GraphNode * mean_sum(Chain & chain, const GraphNode * add, ValueId conv_slice) { + if (add == nullptr || !common_binary_node_is_add(*add) || add->inputs[0] != conv_slice) { + return nullptr; + } + const GraphNode * half = chain.producer(add->inputs[1], GGML_OP_SCALE); + if (!scale_is(half, 0.5f)) { + return nullptr; + } + const GraphNode * sum = chain.producer(half->inputs[0], GGML_OP_ADD); + return sum != nullptr && common_binary_node_is_add(*sum) ? sum : nullptr; +} + +bool match_zaya_cca_qk_norm(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const Graph & graph = context.graph; + const GraphNode * root = context.root_node; + if (root == nullptr || !graph.has_index() || root->op != GGML_OP_CONT || root->inputs.size() != 1) { + return false; + } + Chain chain{ graph, { root } }; + + // The convolution output C and its two slices: query [0, q_size), key [q_size, q_size + k_size). + const GraphNode * q_slice = graph.index().producer(root->inputs[0]); + if (q_slice == nullptr || !is_layout_alias_node(graph, *q_slice)) { + return false; + } + const ValueId conv_id = q_slice->inputs[0]; + const Value * conv = chain.value(conv_id); + const Value * q_view = chain.value(root->inputs[0]); + if (conv == nullptr || q_view == nullptr || conv->type != GGML_TYPE_F32 || !conv->contiguous || + conv->ne[1] != 1 || conv->ne[2] != 1 || conv->ne[3] != 1 || q_view->storage_offset != conv->storage_offset) { + return false; + } + const std::vector & slices = graph.index().consumers(conv_id); + if (slices.size() != 2) { + return false; + } + const GraphNode * k_slice = slices[0] == q_slice ? slices[1] : slices[0]; + const Value * k_view = chain.value(k_slice->output); + if (!is_layout_alias_node(graph, *k_slice) || k_view == nullptr || + k_view->storage_offset != conv->storage_offset + q_view->byte_count || + q_view->element_count + k_view->element_count != conv->element_count) { + return false; + } + + // Query side: CONT -> reshape [D, n_head] -> ADD(., mean) -> RMS_NORM. + const ValueId cq = chain.skip_aliases(root->output); + const Value * cq_val = chain.value(cq); + const GraphNode * add_q = chain.only_consumer(cq); + if (cq_val == nullptr || cq_val->ne[2] != 1 || cq_val->ne[3] != 1 || add_q == nullptr) { + return false; + } + const int64_t head_dim = cq_val->ne[0]; + const int64_t n_head = cq_val->ne[1]; + chain.nodes.push_back(add_q); + const GraphNode * q_sum = mean_sum(chain, add_q, cq); + const GraphNode * norm_q = chain.only_consumer(add_q->output); + if (q_sum == nullptr || norm_q == nullptr || norm_q->op != GGML_OP_RMS_NORM) { + return false; + } + chain.nodes.push_back(norm_q); + // Qpre + reshape(REPEAT(Kpre)). + const ValueId qraw_id = chain.source(q_sum->inputs[0]); + const GraphNode * repeat = chain.producer(q_sum->inputs[1], GGML_OP_REPEAT); + const ValueId kraw_id = repeat != nullptr ? chain.source(repeat->inputs[0]) : ValueId{}; + const Value * qraw = chain.value(qraw_id); + const Value * kraw = chain.value(kraw_id); + if (repeat == nullptr || !is_rows(qraw, head_dim * n_head, 1) || kraw == nullptr || kraw->ne[0] % head_dim != 0 || + !is_rows(kraw, kraw->ne[0], 1)) { + return false; + } + const int64_t n_head_kv = kraw->ne[0] / head_dim; + if (n_head_kv < 1 || n_head % n_head_kv != 0 || k_view->element_count != kraw->ne[0]) { + return false; + } + const int64_t gqa = n_head / n_head_kv; + + // Key side: slice -> CONT -> reshape [D, n_head_kv] -> ADD(., mean) -> RMS_NORM -> MUL(k_scale). + chain.nodes.push_back(k_slice); + const GraphNode * cont_k = chain.only_consumer(k_slice->output); + if (cont_k == nullptr || cont_k->op != GGML_OP_CONT) { + return false; + } + chain.nodes.push_back(cont_k); + const ValueId ck = chain.skip_aliases(cont_k->output); + const GraphNode * add_k = chain.only_consumer(ck); + if (add_k == nullptr) { + return false; + } + chain.nodes.push_back(add_k); + const GraphNode * k_sum = mean_sum(chain, add_k, ck); + const GraphNode * norm_k = chain.only_consumer(add_k->output); + const GraphNode * mul_k = norm_k != nullptr ? chain.only_consumer(norm_k->output) : nullptr; + if (k_sum == nullptr || norm_k == nullptr || norm_k->op != GGML_OP_RMS_NORM || mul_k == nullptr || + !common_binary_node_is_mul(*mul_k) || mul_k->inputs[0] != norm_k->output || + chain.source(k_sum->inputs[1]) != kraw_id) { + return false; + } + chain.nodes.push_back(norm_k); + chain.nodes.push_back(mul_k); + // SCALE(SUM_ROWS(CONT(PERMUTE(reshape(Qpre)))), 1/gqa). + const GraphNode * inv_gqa = chain.producer(k_sum->inputs[0], GGML_OP_SCALE); + const GraphNode * sum_rows = inv_gqa != nullptr ? chain.producer(inv_gqa->inputs[0], GGML_OP_SUM_ROWS) : nullptr; + const GraphNode * cont_q = sum_rows != nullptr ? chain.producer(sum_rows->inputs[0], GGML_OP_CONT) : nullptr; + if (!scale_is(inv_gqa, 1.0f / static_cast(gqa)) || cont_q == nullptr || + chain.source(cont_q->inputs[0]) != qraw_id) { + return false; + } + + const RmsNormParams * eps_q = op_params_as(norm_q->params); + const RmsNormParams * eps_k = op_params_as(norm_k->params); + const Value * scale = chain.value(chain.source(mul_k->inputs[1])); + const Value * q_out = chain.value(norm_q->output); + const Value * k_out = chain.value(mul_k->output); + if (eps_q == nullptr || eps_k == nullptr || eps_q->eps != eps_k->eps || scale == nullptr || + scale->type != GGML_TYPE_F32 || !scale->contiguous || scale->element_count != n_head_kv || + q_out == nullptr || k_out == nullptr || q_out->element_count != head_dim * n_head || + k_out->element_count != head_dim * n_head_kv || !q_out->contiguous || !k_out->contiguous || + q_out->alias_source.value >= 0 || k_out->alias_source.value >= 0 || head_dim % 64 != 0 || head_dim > 256 || + n_head > 128 || n_head_kv > 128) { + return false; + } + + for (const GraphNode * node : chain.nodes) { + if (!append_covered_node_index_once(graph, context.covered_nodes, node, dispatch_match.covered_nodes)) { + dispatch_match.covered_nodes.clear(); + return false; + } + } + + const auto input = [](const Value * v) { + return DispatchBinding{ v->storage_root, v->storage_offset, v->byte_count }; + }; + char eps[32]; + std::snprintf(eps, sizeof(eps), "%.9g", static_cast(eps_q->eps)); + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kZayaCcaQkNormKernel); + dispatch.kernel.compile_parameters.emplace("ggml.zaya_cca_qk_norm.head_dim", std::to_string(head_dim)); + dispatch.kernel.compile_parameters.emplace("ggml.zaya_cca_qk_norm.n_head", std::to_string(n_head)); + dispatch.kernel.compile_parameters.emplace("ggml.zaya_cca_qk_norm.n_head_kv", std::to_string(n_head_kv)); + dispatch.kernel.compile_parameters.emplace("ggml.zaya_cca_qk_norm.gqa", std::to_string(gqa)); + dispatch.kernel.compile_parameters.emplace("ggml.zaya_cca_qk_norm.rms_epsilon", eps); + dispatch.bindings.push_back(input(conv)); + dispatch.bindings.push_back(input(qraw)); + dispatch.bindings.push_back(input(kraw)); + dispatch.bindings.push_back(input(scale)); + dispatch.bindings.push_back({ q_out->id, 0, q_out->byte_count }); + dispatch.bindings.push_back({ k_out->id, 0, k_out->byte_count }); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_zaya_cca_qk_norm_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "zaya.cca_qk_norm.decode_f32", + GGML_OP_CONT, + DispatchMatchKind::Fused, + 500, + DispatchSource::Common, + match_zaya_cca_qk_norm, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-zaya-cca-qk-norm.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-zaya-cca-qk-norm.h new file mode 100644 index 000000000000..b363585f5bf8 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-zaya-cca-qk-norm.h @@ -0,0 +1,24 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +namespace ggml::hrx { + +void register_zaya_cca_qk_norm_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/moe-placement-guard.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/moe-placement-guard.cpp new file mode 100644 index 000000000000..d330afeb99fb --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/moe-placement-guard.cpp @@ -0,0 +1,83 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Placement guard for gpt-oss's MoE tail ops (ADD_ID, SWIGLU_OAI). When the expert MUL_MAT_IDs they follow run on +// the CPU, running SWIGLU_OAI on HRX between those CPU splits gives wrong results in the full model, although every +// op is right on its own and the same block is right in isolation (root cause open: engine #286). Until that is +// found, these ops are only claimed when their MUL_MAT_ID will run on HRX too. +// +// "Will run on HRX" takes two checks, because the scheduler places a MUL_MAT_ID with its expert weights: +// - HRX can execute the MUL_MAT_ID (a batch HRX declines, or a disabled dispatch, sends it to the CPU), and +// - the expert weights sit in a buffer HRX can read. Experts in a CPU-only buffer (CPU_REPACK, where llama.cpp +// puts them when HRX declined the MUL_MAT_ID at load time, or where an override puts them) keep every +// MUL_MAT_ID on the CPU, decode included, even when HRX could execute the decode shape. + +#include "moe-placement-guard.h" + +#include "ggml-backend.h" +#include "ggml.h" + +namespace ggml::hrx { + +namespace { + +// The MUL_MAT_ID an ADD_ID / SWIGLU_OAI input comes from, directly or through ADD_ID, or nullptr. +const ggml_tensor * expert_source(const ggml_tensor * tensor) { + if (tensor == nullptr) { + return nullptr; + } + if (tensor->op == GGML_OP_MUL_MAT_ID) { + return tensor; + } + if (tensor->op == GGML_OP_ADD_ID && tensor->src[0] != nullptr && tensor->src[0]->op == GGML_OP_MUL_MAT_ID) { + return tensor->src[0]; + } + return nullptr; +} + +// False when the expert weights are already in a buffer this device cannot use; true when it can, or before +// allocation (no buffer yet). +bool experts_readable(ggml_backend_dev_t device, const ggml_tensor * experts) { + const ggml_tensor * weights = experts->src[0]; + if (weights == nullptr) { + return true; + } + const ggml_tensor * storage = weights->view_src != nullptr ? weights->view_src : weights; + ggml_backend_buffer_t buffer = storage->buffer; + return buffer == nullptr || ggml_backend_dev_supports_buft(device, ggml_backend_buffer_get_type(buffer)); +} + +} // namespace + +bool moe_tail_claimable(ggml_backend_dev_t device, const ggml_tensor * op) { + if (op == nullptr) { + return true; + } + const bool add_id = op->op == GGML_OP_ADD_ID; + const bool swiglu_oai = op->op == GGML_OP_GLU && ggml_get_glu_op(op) == GGML_GLU_OP_SWIGLU_OAI; + if (!add_id && !swiglu_oai) { + return true; + } + for (int i = 0; i < 2; ++i) { + const ggml_tensor * experts = expert_source(op->src[i]); + if (experts != nullptr && + (!experts_readable(device, experts) || !ggml_backend_dev_supports_op(device, experts))) { + return false; + } + } + return true; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/moe-placement-guard.h b/ggml/src/ggml-hrx/dispatch_registration/common/moe-placement-guard.h new file mode 100644 index 000000000000..b0f8aee78a26 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/moe-placement-guard.h @@ -0,0 +1,29 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include "ggml-backend.h" + +struct ggml_tensor; + +namespace ggml::hrx { + +// False for an ADD_ID / SWIGLU_OAI whose expert MUL_MAT_ID (directly, or through ADD_ID) will not run on this +// device: the device declines the op, or the expert weights sit in a buffer it cannot read (see +// moe-placement-guard.cpp). True for everything else. +bool moe_tail_claimable(ggml_backend_dev_t device, const ggml_tensor * op); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/dispatch-registry.cpp b/ggml/src/ggml-hrx/dispatch_registration/dispatch-registry.cpp new file mode 100644 index 000000000000..41de31611db4 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/dispatch-registry.cpp @@ -0,0 +1,214 @@ +#include +#include +#include +#include +#include +#include "dispatch-registry.h" + +#include "common/dispatch-common.h" +#include "llm/dispatch-attention-qkv.h" +#include "llm/dispatch-gated-delta-net.h" +#include "llm/dispatch-ssm-conv.h" +#include "qwen/dispatch-qwen.h" + +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr size_t kOpCount = static_cast(GGML_OP_COUNT); + +static const std::vector kEmptyRegistrations; + +static bool valid_root_op(ggml_op op) { + return op >= 0 && static_cast(op) < kOpCount; +} + +static void sort_registrations(std::vector & registrations) { + std::stable_sort( + registrations.begin(), registrations.end(), + [](const DispatchRegistration & lhs, const DispatchRegistration & rhs) { return lhs.priority > rhs.priority; }); +} + +static bool qwen_dispatch_disabled_from_environment() { + const char * value = std::getenv("GGML_HRX_DISABLE_QWEN_DISPATCH"); + if (value == nullptr || value[0] == '\0') { + return false; + } + return std::strcmp(value, "0") != 0 && std::strcmp(value, "false") != 0 && std::strcmp(value, "FALSE") != 0 && + std::strcmp(value, "off") != 0 && std::strcmp(value, "OFF") != 0; +} + +static DispatchRegistry build_registry(bool include_qwen) { + DispatchRegistryBuilder builder; + register_common_dispatches(builder); + register_llm_attention_qkv_dispatches(builder); + register_llm_gated_delta_net_dispatch(builder); + register_llm_ssm_conv_dispatch(builder); + if (include_qwen) { + register_qwen_dispatches(builder); + } + return builder.build(); +} + +} // namespace + +// Debug switches. GGML_HRX_DISABLE_DISPATCH: comma-separated substrings of registration names to +// skip ("fused" skips every fused registration). GGML_HRX_LOG_DISPATCH=1: print each +// registration name the first time it matches. +static bool dispatch_disabled(const DispatchRegistration & registration) { + static const std::vector patterns = [] { + std::vector out; + const char * env = std::getenv("GGML_HRX_DISABLE_DISPATCH"); + std::string list = env != nullptr ? env : ""; + size_t start = 0; + while (start <= list.size()) { + size_t end = list.find(',', start); + if (end == std::string::npos) { + end = list.size(); + } + if (end > start) { + out.push_back(list.substr(start, end - start)); + } + start = end + 1; + } + return out; + }(); + const std::string name = registration.name != nullptr ? registration.name : ""; + for (const std::string & p : patterns) { + if ((p == "fused" && registration.kind == DispatchMatchKind::Fused) || name.find(p) != std::string::npos) { + return true; + } + } + return false; +} + +static void dispatch_log_match(const DispatchRegistration & registration) { + static const bool enabled = std::getenv("GGML_HRX_LOG_DISPATCH") != nullptr; + if (!enabled) { + return; + } + static std::mutex mutex; + static std::set seen; + const std::string name = registration.name != nullptr ? registration.name : ""; + std::lock_guard lock(mutex); + if (seen.insert(name).second) { + std::fprintf(stderr, "ggml_hrx dispatch: %s (%s)\n", name.c_str(), + registration.kind == DispatchMatchKind::Fused ? "fused" : "single"); + } +} + +bool DispatchRegistry::match(const DispatchMatchContext & context, DispatchMatch & match) const { + return this->match(context, match, nullptr); +} + +bool DispatchRegistry::match(const DispatchMatchContext & context, + DispatchMatch & match, + DispatchMatchDiagnostics * diagnostics) const { + if (context.root_node == nullptr || !valid_root_op(context.root_node->op) || registrations_by_root_.empty()) { + return false; + } + if (diagnostics != nullptr) { + diagnostics->root_op = context.root_node->op; + diagnostics->attempts.clear(); + } + const std::vector & registrations = + registrations_by_root_[static_cast(context.root_node->op)].ordered; + for (const DispatchRegistration & registration : registrations) { + if (dispatch_disabled(registration)) { + continue; + } + DispatchMatch candidate; + if (registration.matcher != nullptr && registration.matcher(context, candidate)) { + dispatch_log_match(registration); + if (diagnostics != nullptr) { + diagnostics->attempts.push_back({ + registration.name != nullptr ? registration.name : "", + registration.root_op, + registration.kind, + registration.priority, + registration.source, + true, + candidate.covered_nodes, + candidate.status.errors(), + }); + } + match = std::move(candidate); + return true; + } + if (diagnostics != nullptr) { + diagnostics->attempts.push_back({ + registration.name != nullptr ? registration.name : "", + registration.root_op, + registration.kind, + registration.priority, + registration.source, + false, + candidate.covered_nodes, + candidate.status.errors(), + }); + } + match.status.append(candidate.status); + } + return false; +} + +const std::vector & DispatchRegistry::registrations_for_root(ggml_op root_op) const { + if (!valid_root_op(root_op) || registrations_by_root_.empty()) { + return kEmptyRegistrations; + } + return registrations_by_root_[static_cast(root_op)].ordered; +} + +void DispatchRegistryBuilder::add(DispatchRegistration registration) { + if (!valid_root_op(registration.root_op) || registration.matcher == nullptr) { + return; + } + if (registry_.registrations_by_root_.empty()) { + registry_.registrations_by_root_.resize(kOpCount); + } + DispatchRegistry::RegistrationGroup & group = + registry_.registrations_by_root_[static_cast(registration.root_op)]; + if (registration.kind == DispatchMatchKind::Fused) { + group.fused.push_back(std::move(registration)); + } else { + group.single_op.push_back(registration); + registry_.single_op_registrations_.push_back(std::move(registration)); + } +} + +DispatchRegistry DispatchRegistryBuilder::build() { + if (registry_.registrations_by_root_.empty()) { + registry_.registrations_by_root_.resize(kOpCount); + } + for (DispatchRegistry::RegistrationGroup & group : registry_.registrations_by_root_) { + sort_registrations(group.fused); + sort_registrations(group.single_op); + group.ordered = group.fused; + group.ordered.insert(group.ordered.end(), group.single_op.begin(), group.single_op.end()); + } + sort_registrations(registry_.single_op_registrations_); + return std::move(registry_); +} + +const DispatchRegistry * find_dispatch_registry(const DispatchTarget & target) { + static const DispatchRegistry gfx1100_registry = build_registry(true); + static const DispatchRegistry gfx1100_generic_only_registry = build_registry(false); + static const DispatchRegistry gfx1151_registry = build_registry(true); + static const DispatchRegistry gfx1151_generic_only_registry = build_registry(false); + + const bool qwen_disabled = qwen_dispatch_disabled_from_environment(); + + if (target.architecture == "gfx1100") { + return qwen_disabled ? &gfx1100_generic_only_registry : &gfx1100_registry; + } + if (target.architecture == "gfx1151") { + return qwen_disabled ? &gfx1151_generic_only_registry : &gfx1151_registry; + } + return nullptr; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/dispatch-registry.h b/ggml/src/ggml-hrx/dispatch_registration/dispatch-registry.h new file mode 100644 index 000000000000..bf92f9f82e1e --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/dispatch-registry.h @@ -0,0 +1,117 @@ +#pragma once + +#include "dispatch/command-plan.h" +#include "ggml.h" +#include "graph/graph.h" +#include "status.h" + +#include +#include +#include + +namespace ggml::hrx { + +struct DispatchTarget { + std::string architecture; +}; + +enum class DispatchMatchKind { + Fused, + SingleOp, +}; + +enum class DispatchSource { + Common, + Llm, + Qwen, +}; + +struct DispatchMatchContext { + const Graph & graph; + const GraphNode * root_node = nullptr; + size_t root_index = 0; + const std::vector & covered_nodes; + const CommandPlan & plan; + ValueId next_plan_value; +}; + +struct DispatchValueAliasRequest { + ValueId source_value; + ValueId target_value; +}; + +struct DispatchMatch { + std::vector initialization_dispatches; + std::vector covered_nodes; + std::vector dispatches; + std::vector transients; + std::vector constant_initializations; + std::vector completion_counter_requests; + std::vector value_aliases; + CommandPlanMetadata metadata; + Status status; +}; + +using DispatchMatcher = bool (*)(const DispatchMatchContext & context, DispatchMatch & match); + +struct DispatchRegistration { + const char * name = ""; + ggml_op root_op = GGML_OP_NONE; + DispatchMatchKind kind = DispatchMatchKind::SingleOp; + int priority = 0; + DispatchSource source = DispatchSource::Common; + DispatchMatcher matcher = nullptr; +}; + +struct DispatchRegistrationAttempt { + std::string name; + ggml_op root_op = GGML_OP_NONE; + DispatchMatchKind kind = DispatchMatchKind::SingleOp; + int priority = 0; + DispatchSource source = DispatchSource::Common; + bool matched = false; + std::vector covered_nodes; + std::vector errors; +}; + +struct DispatchMatchDiagnostics { + ggml_op root_op = GGML_OP_NONE; + std::vector attempts; +}; + +class DispatchRegistry { + public: + bool match(const DispatchMatchContext & context, DispatchMatch & match) const; + bool match(const DispatchMatchContext & context, + DispatchMatch & match, + DispatchMatchDiagnostics * diagnostics) const; + + const std::vector & registrations_for_root(ggml_op root_op) const; + + const std::vector & single_op_registrations() const { return single_op_registrations_; } + + private: + friend class DispatchRegistryBuilder; + + struct RegistrationGroup { + std::vector fused; + std::vector single_op; + std::vector ordered; + }; + + std::vector registrations_by_root_; + std::vector single_op_registrations_; +}; + +class DispatchRegistryBuilder { + public: + void add(DispatchRegistration registration); + DispatchRegistry build(); + + private: + DispatchRegistry registry_; +}; + +const DispatchRegistry * find_dispatch_registry(const DispatchTarget & target); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-attention-qkv.cpp b/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-attention-qkv.cpp new file mode 100644 index 000000000000..ea81bda320f3 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-attention-qkv.cpp @@ -0,0 +1,677 @@ +#include "dispatch-attention-qkv.h" + +#include "../common/dispatch-mul-mat-common.h" +#include "ggml.h" +#include "graph/graph-matcher.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kAttentionQMatMulRopeF32F32WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_attention_q_matmul_rope_f32_f32_wmma"); +static constexpr KernelCatalogRef kAttentionKMatMulRopeSetRowsF32F32WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_attention_k_matmul_rope_set_rows_f32_f32_wmma"); +static constexpr KernelCatalogRef kAttentionVMatMulSetRowsF32F32WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_attention_v_matmul_set_rows_f32_f32_wmma"); +static constexpr KernelCatalogRef kAttentionQMatMulRopeF32F32DecodeKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_attention_q_matmul_rope_decode_f32_f32"); +static constexpr KernelCatalogRef kAttentionKMatMulRopeSetRowsF32F32DecodeKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_attention_k_matmul_rope_set_rows_decode_f32_f32"); +static constexpr KernelCatalogRef kAttentionVMatMulSetRowsF32F32DecodeKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_attention_v_matmul_set_rows_decode_f32_f32"); + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static bool is_1d_shape(const Value & value, int64_t ne0) { + return value.ne[0] == ne0 && value.ne[1] == 1 && value.ne[2] == 1 && value.ne[3] == 1; +} + +static bool is_2d_shape(const Value & value, int64_t ne0, int64_t ne1) { + return value.ne[0] == ne0 && value.ne[1] == ne1 && value.ne[2] == 1 && value.ne[3] == 1; +} + +static bool is_2d(const Value & value) { + return value.ne[0] > 0 && value.ne[1] > 0 && value.ne[2] == 1 && value.ne[3] == 1; +} + +static bool is_attention_rope_shape(const Value & value) { + return value.ne[0] >= 4 && value.ne[0] <= 1024 && value.ne[0] % 4 == 0 && value.ne[1] >= 1 && value.ne[1] <= 64 && + value.ne[2] >= 1 && value.ne[2] <= 2048 && value.ne[3] == 1; +} + +static bool is_supported_dense_input_size(int64_t input_size) { + return input_size >= 256 && input_size <= 32768 && input_size % 256 == 0; +} + +static bool is_supported_decode_input_size(int64_t input_size) { + return input_size >= 256 && input_size <= 32768 && input_size % 32 == 0; +} + +static bool is_supported_dense_output_size(int64_t output_size) { + return output_size >= 1 && output_size <= 262144; +} + +static bool is_supported_decode_output_size(int64_t output_size) { + return output_size >= 1 && output_size <= 32768; +} + +static bool is_supported_prefill_token_count(int64_t token_count) { + return token_count > 1 && token_count <= 2048; +} + +static bool is_supported_decode_token_count(int64_t token_count) { + return token_count == 1; +} + +static bool is_supported_attention_token_count(int64_t token_count) { + return is_supported_decode_token_count(token_count) || is_supported_prefill_token_count(token_count); +} + +static bool is_supported_attention_input_size(int64_t input_size, int64_t token_count) { + return is_supported_decode_token_count(token_count) ? is_supported_decode_input_size(input_size) : + is_supported_dense_input_size(input_size); +} + +static bool is_supported_attention_output_size(int64_t output_size, int64_t token_count) { + return is_supported_decode_token_count(token_count) ? is_supported_decode_output_size(output_size) : + is_supported_dense_output_size(output_size); +} + +static bool is_supported_attention_cache_row_count(int64_t row_count) { + return row_count >= 1 && row_count <= 1048576; +} + +static bool same_shape(const Value & lhs, const Value & rhs) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (lhs.ne[i] != rhs.ne[i]) { + return false; + } + } + return true; +} + +static std::string to_config_value(int64_t value) { + return std::to_string(value); +} + +static bool is_supported_attention_rope_params(const RopeParams & params, int64_t head_size) { + return params.n_dims == head_size && params.mode == GGML_ROPE_TYPE_NORMAL && std::isfinite(params.freq_base) && + params.freq_base > 0.0f && std::isfinite(params.freq_scale) && params.freq_scale > 0.0f && + params.ext_factor == 0.0f && params.attn_factor == 1.0f; +} + +static bool build_rope_theta_table(const GraphNode & rope, int64_t head_size, std::vector & data) { + const RopeParams * params = op_params_as(rope.params); + if (params == nullptr || !is_supported_attention_rope_params(*params, head_size)) { + return false; + } + + data.resize(static_cast(head_size / 2) * sizeof(float)); + const float theta_scale = std::pow(params->freq_base, -2.0f / static_cast(head_size)); + float theta = params->freq_scale; + for (int64_t i = 0; i < head_size / 2; ++i) { + std::memcpy(data.data() + static_cast(i) * sizeof(float), &theta, sizeof(theta)); + theta *= theta_scale; + } + return true; +} + +static void build_unit_frequency_factors(int64_t head_size, std::vector & data) { + data.resize(static_cast(head_size / 2) * sizeof(float)); + const float one = 1.0f; + for (int64_t i = 0; i < head_size / 2; ++i) { + std::memcpy(data.data() + static_cast(i) * sizeof(float), &one, sizeof(one)); + } +} + +static bool cache_output_format_value(ggml_type type, int64_t & value) { + switch (type) { + case GGML_TYPE_F16: + value = 16; + return true; + case GGML_TYPE_F32: + value = 32; + return true; + default: + return false; + } +} + +static ValueId next_match_transient_value(const DispatchMatchContext & context, const DispatchMatch & dispatch_match) { + return ValueId(context.next_plan_value.value + static_cast(dispatch_match.transients.size()) + + static_cast(dispatch_match.completion_counter_requests.size())); +} + +static bool node_is_available(const DispatchMatchContext & context, const GraphNode * node) { + size_t node_index = 0; + return node != nullptr && context.graph.index().node_index(node, node_index) && + node_index < context.covered_nodes.size() && !context.covered_nodes[node_index]; +} + +static const GraphNode * find_only_available_consumer(const DispatchMatchContext & context, ValueId value) { + if (!context.graph.has_index()) { + return nullptr; + } + const std::vector & consumers = context.graph.index().consumers(value); + if (consumers.size() != 1 || !node_is_available(context, consumers.front())) { + return nullptr; + } + return consumers.front(); +} + +static const GraphNode * find_only_consumer_after_layout_aliases(const DispatchMatchContext & context, + const Value * start, + std::vector & layouts, + const Value *& current, + size_t max_layouts) { + current = start; + for (;;) { + if (current == nullptr) { + return nullptr; + } + const GraphNode * consumer = find_only_available_consumer(context, current->id); + if (consumer == nullptr) { + return nullptr; + } + if (!is_layout_alias_node(context.graph, *consumer)) { + return consumer; + } + if (layouts.size() >= max_layouts) { + return nullptr; + } + const Value * output = graph_value(context.graph, consumer->output); + if (output == nullptr || output->type != current->type || !output->contiguous) { + return nullptr; + } + layouts.push_back(consumer); + current = output; + } +} + +struct AttentionMatMulMatch { + const Value * input = nullptr; + const Value * weight = nullptr; + const Value * output = nullptr; + int64_t input_size = 0; + int64_t output_size = 0; + int64_t token_count = 0; + CommonMulMatWeightFormat weight_format = CommonMulMatWeightFormat::Q4K; + + bool matched() const { return input != nullptr && weight != nullptr && output != nullptr; } +}; + +struct AttentionRopeMatch { + const GraphNode * node = nullptr; + const Value * input = nullptr; + const Value * positions = nullptr; + const Value * freq_factors = nullptr; + const Value * output = nullptr; + int64_t token_count = 0; + int64_t head_count = 0; + int64_t head_size = 0; + size_t theta_bytes = 0; + size_t freq_factors_bytes = 0; + std::vector theta_data; + std::vector freq_factors_data; + + bool matched() const { return node != nullptr && input != nullptr && positions != nullptr && output != nullptr; } +}; + +struct AttentionSetRowsMatch { + const GraphNode * node = nullptr; + const Value * rows = nullptr; + const Value * indices = nullptr; + const Value * output = nullptr; + int64_t output_format = 0; + int64_t token_count = 0; + int64_t cache_row_count = 0; + int64_t output_size = 0; + + bool matched() const { return node != nullptr && rows != nullptr && indices != nullptr && output != nullptr; } +}; + +enum class AttentionQkvProjectionKind { + Query, + Key, + Value, +}; + +struct AttentionQkvMatch { + AttentionMatMulMatch root; + AttentionRopeMatch rope; + AttentionSetRowsMatch set_rows; + std::vector layout_nodes; + KernelCatalogRef kernel = {}; + AttentionQkvProjectionKind kind = AttentionQkvProjectionKind::Query; + + bool matched() const { return root.matched() && kernel.id != kUncatalogedKernelId; } +}; + +static bool can_use_packed_lowtoken_attention_matmul(const AttentionMatMulMatch & root) { + if (root.token_count < 1 || root.token_count > 5 || root.output_size % 64 != 0 || + root.weight->alias_source.value >= 0) { + return false; + } + if (root.weight_format != CommonMulMatWeightFormat::Q4K && + root.weight_format != CommonMulMatWeightFormat::Q6K) { + return false; + } + return common_is_supported_dense_input_size(root.weight_format, root.input_size); +} + +static CommonMulMatWeightFormat packed_attention_lowtoken_format(CommonMulMatWeightFormat format) { + return format == CommonMulMatWeightFormat::Q4K ? CommonMulMatWeightFormat::Q4KRow64 : + CommonMulMatWeightFormat::Q6KRow64; +} + +static AttentionMatMulMatch match_attention_matmul_any_format(const Graph & graph, const GraphNode * node) { + AttentionMatMulMatch match; + if (node == nullptr || node->op != GGML_OP_MUL_MAT || node->inputs.size() != 2) { + return match; + } + + const Value * weight = graph_value(graph, node->inputs[0]); + const Value * input = graph_value(graph, node->inputs[1]); + const Value * output = graph_value(graph, node->output); + if (weight == nullptr || input == nullptr || output == nullptr || !is_2d(*weight) || !is_2d(*input) || + !is_2d(*output) || !weight->contiguous || !input->contiguous || !output->contiguous || + input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32) { + return {}; + } + + CommonMulMatWeightFormat format = CommonMulMatWeightFormat::Q4K; + if (!common_mul_mat_format_for_type(weight->type, format)) { + return {}; + } + + const int64_t input_size = weight->ne[0]; + const int64_t output_size = weight->ne[1]; + const int64_t token_count = input->ne[1]; + if (input->ne[0] != input_size || output->ne[0] != output_size || output->ne[1] != token_count || + !is_supported_attention_token_count(token_count) || + !is_supported_attention_input_size(input_size, token_count) || + !is_supported_attention_output_size(output_size, token_count)) { + return {}; + } + + match.input = input; + match.weight = weight; + match.output = output; + match.input_size = input_size; + match.output_size = output_size; + match.token_count = token_count; + match.weight_format = format; + return match; +} + +static AttentionRopeMatch match_attention_rope(const Graph & graph, + const GraphNode * node, + const Value * expected_input) { + AttentionRopeMatch match; + if (node == nullptr || node->op != GGML_OP_ROPE || node->inputs.size() < 2 || node->inputs.size() > 3) { + return match; + } + + const Value * input = graph_value(graph, node->inputs[0]); + const Value * positions = graph_value(graph, node->inputs[1]); + const Value * output = graph_value(graph, node->output); + if (input == nullptr || positions == nullptr || output == nullptr || expected_input == nullptr || + input->id != expected_input->id || input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || + positions->type != GGML_TYPE_I32 || !input->contiguous || !output->contiguous || !positions->contiguous || + !is_attention_rope_shape(*input) || !same_shape(*input, *output)) { + return {}; + } + + const int64_t head_size = input->ne[0]; + const int64_t head_count = input->ne[1]; + const int64_t token_count = input->ne[2]; + const RopeParams * params = op_params_as(node->params); + if (params == nullptr || !is_supported_attention_rope_params(*params, head_size) || + !is_1d_shape(*positions, token_count)) { + return {}; + } + + std::vector theta_data; + if (!build_rope_theta_table(*node, head_size, theta_data)) { + return {}; + } + + const Value * freq_factors = nullptr; + size_t freq_factors_bytes = 0; + std::vector freq_factors_data; + if (node->inputs.size() == 3) { + freq_factors = graph_value(graph, node->inputs[2]); + if (freq_factors == nullptr || freq_factors->type != GGML_TYPE_F32 || !freq_factors->contiguous || + !is_1d_shape(*freq_factors, head_size / 2)) { + return {}; + } + freq_factors_bytes = freq_factors->byte_count; + } else { + build_unit_frequency_factors(head_size, freq_factors_data); + freq_factors_bytes = freq_factors_data.size(); + } + + match.node = node; + match.input = input; + match.positions = positions; + match.freq_factors = freq_factors; + match.output = output; + match.token_count = token_count; + match.head_count = head_count; + match.head_size = head_size; + match.theta_bytes = theta_data.size(); + match.freq_factors_bytes = freq_factors_bytes; + match.theta_data = std::move(theta_data); + match.freq_factors_data = std::move(freq_factors_data); + return match; +} + +static AttentionSetRowsMatch match_attention_set_rows(const Graph & graph, + const GraphNode * node, + const Value * expected_rows) { + AttentionSetRowsMatch match; + if (node == nullptr || node->op != GGML_OP_SET_ROWS || node->inputs.size() != 3) { + return match; + } + + const Value * rows = graph_value(graph, node->inputs[0]); + const Value * indices = graph_value(graph, node->inputs[1]); + const Value * cache = graph_value(graph, node->inputs[2]); + const Value * output = graph_value(graph, node->output); + if (rows == nullptr || indices == nullptr || cache == nullptr || output == nullptr || expected_rows == nullptr || + rows->id != expected_rows->id || rows->type != GGML_TYPE_F32 || indices->type != GGML_TYPE_I64 || + output->type != cache->type || !rows->contiguous || !indices->contiguous || !cache->contiguous || + !same_shape(*cache, *output) || !graph.values().same_storage(cache->id, output->id)) { + return {}; + } + + int64_t output_format = 0; + if (!cache_output_format_value(output->type, output_format)) { + return {}; + } + + const int64_t output_size = rows->ne[0]; + const int64_t token_count = rows->ne[1]; + const int64_t cache_row_count = cache->ne[1]; + if (!is_2d_shape(*rows, output_size, token_count) || !is_1d_shape(*indices, token_count) || + !is_2d_shape(*cache, output_size, cache_row_count) || + !is_supported_attention_output_size(output_size, token_count) || + !is_supported_attention_token_count(token_count) || !is_supported_attention_cache_row_count(cache_row_count)) { + return {}; + } + + match.node = node; + match.rows = rows; + match.indices = indices; + match.output = output; + match.output_format = output_format; + match.token_count = token_count; + match.cache_row_count = cache_row_count; + match.output_size = output_size; + return match; +} + +static AttentionQkvMatch match_attention_qkv_projection(const DispatchMatchContext & context) { + AttentionQkvMatch match; + AttentionMatMulMatch root = match_attention_matmul_any_format(context.graph, context.root_node); + if (!root.matched() || !context.graph.has_index()) { + return match; + } + + const Value * after_projection = nullptr; + const GraphNode * consumer = + find_only_consumer_after_layout_aliases(context, root.output, match.layout_nodes, after_projection, 2); + if (consumer == nullptr) { + return {}; + } + + if (consumer->op == GGML_OP_ROPE) { + match.rope = match_attention_rope(context.graph, consumer, after_projection); + if (!match.rope.matched() || match.rope.token_count != root.token_count || + match.rope.head_size * match.rope.head_count != root.output_size) { + return {}; + } + + const Value * after_rope = nullptr; + std::vector set_rows_layouts; + const GraphNode * rope_consumer = + find_only_consumer_after_layout_aliases(context, match.rope.output, set_rows_layouts, after_rope, 2); + if (rope_consumer != nullptr && rope_consumer->op == GGML_OP_SET_ROWS) { + match.set_rows = match_attention_set_rows(context.graph, rope_consumer, after_rope); + if (!match.set_rows.matched() || match.set_rows.token_count != root.token_count || + match.set_rows.output_size != root.output_size) { + return {}; + } + match.layout_nodes.insert(match.layout_nodes.end(), set_rows_layouts.begin(), set_rows_layouts.end()); + match.root = root; + match.kernel = is_supported_decode_token_count(root.token_count) ? + kAttentionKMatMulRopeSetRowsF32F32DecodeKernel : + kAttentionKMatMulRopeSetRowsF32F32WmmaKernel; + if (can_use_packed_lowtoken_attention_matmul(root)) { + match.root.weight_format = packed_attention_lowtoken_format(root.weight_format); + match.kernel = kAttentionKMatMulRopeSetRowsF32F32WmmaKernel; + } + match.kind = AttentionQkvProjectionKind::Key; + return match; + } + + match.root = root; + match.kernel = is_supported_decode_token_count(root.token_count) ? kAttentionQMatMulRopeF32F32DecodeKernel : + kAttentionQMatMulRopeF32F32WmmaKernel; + if (can_use_packed_lowtoken_attention_matmul(root)) { + match.root.weight_format = packed_attention_lowtoken_format(root.weight_format); + match.kernel = kAttentionQMatMulRopeF32F32WmmaKernel; + } + match.kind = AttentionQkvProjectionKind::Query; + return match; + } + + if (consumer->op == GGML_OP_SET_ROWS) { + if (!is_supported_decode_token_count(root.token_count) && + !common_mul_mat_dense_float_format(root.weight_format)) { + return {}; + } + match.set_rows = match_attention_set_rows(context.graph, consumer, after_projection); + if (!match.set_rows.matched() || match.set_rows.token_count != root.token_count || + match.set_rows.output_size != root.output_size) { + return {}; + } + match.root = root; + match.kernel = is_supported_decode_token_count(root.token_count) ? kAttentionVMatMulSetRowsF32F32DecodeKernel : + kAttentionVMatMulSetRowsF32F32WmmaKernel; + if (can_use_packed_lowtoken_attention_matmul(root)) { + match.root.weight_format = packed_attention_lowtoken_format(root.weight_format); + match.kernel = kAttentionVMatMulSetRowsF32F32WmmaKernel; + } + match.kind = AttentionQkvProjectionKind::Value; + return match; + } + + return {}; +} + +static ValueId add_attention_constant_binding(const DispatchMatchContext & context, + DispatchMatch & dispatch_match, + const char * name, + size_t byte_count, + const std::vector & data) { + const ValueId value = next_match_transient_value(context, dispatch_match); + dispatch_match.transients.push_back({ value, name, byte_count, 256 }); + dispatch_match.constant_initializations.push_back({ + value, + name, + 0, + data, + }); + return value; +} + +static std::pair add_attention_rope_frequency_bindings(const DispatchMatchContext & context, + const AttentionRopeMatch & match, + DispatchMatch & dispatch_match) { + const ValueId theta = add_attention_constant_binding(context, dispatch_match, "llm.attention_qkv.theta", + match.theta_bytes, match.theta_data); + if (match.freq_factors != nullptr) { + return { theta, match.freq_factors->id }; + } + + const ValueId freq_factors = add_attention_constant_binding( + context, dispatch_match, "llm.attention_qkv.freq_factors", match.freq_factors_bytes, match.freq_factors_data); + return { theta, freq_factors }; +} + +static void add_attention_qkv_compile_parameters(Dispatch & dispatch, const AttentionQkvMatch & match) { + dispatch.kernel.compile_parameters.emplace("ggml.workload.token_capacity", to_config_value(match.root.token_count)); + dispatch.kernel.compile_parameters.emplace("llm.attention_qkv.input_size", to_config_value(match.root.input_size)); + dispatch.kernel.compile_parameters.emplace("llm.attention_qkv.output_size", + to_config_value(match.root.output_size)); + if (is_supported_decode_token_count(match.root.token_count)) { + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_f32_f32_decode.token_capacity", + to_config_value(match.root.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_f32_f32_decode.output_capacity", + to_config_value(match.root.output_size)); + } + dispatch.kernel.compile_parameters.emplace( + "llm.attention_qkv.weight_format", + to_config_value(common_mul_mat_format_config_value(match.root.weight_format))); + if (match.kind == AttentionQkvProjectionKind::Query || match.kind == AttentionQkvProjectionKind::Key) { + dispatch.kernel.compile_parameters.emplace("llm.attention_qkv.head_size", + to_config_value(match.rope.head_size)); + dispatch.kernel.compile_parameters.emplace("llm.attention_qkv.head_count", + to_config_value(match.rope.head_count)); + } + if (match.kind == AttentionQkvProjectionKind::Key || match.kind == AttentionQkvProjectionKind::Value) { + dispatch.kernel.compile_parameters.emplace("llm.attention_qkv.cache_row_count", + to_config_value(match.set_rows.cache_row_count)); + dispatch.kernel.compile_parameters.emplace("llm.attention_qkv.cache_output_format", + to_config_value(match.set_rows.output_format)); + } +} + +static bool cover_attention_qkv_nodes(const DispatchMatchContext & context, + const AttentionQkvMatch & match, + DispatchMatch & dispatch_match) { + if (!append_covered_node_index_once(context.graph, context.covered_nodes, context.root_node, + dispatch_match.covered_nodes)) { + return false; + } + for (const GraphNode * layout : match.layout_nodes) { + if (!append_covered_node_index_once(context.graph, context.covered_nodes, layout, + dispatch_match.covered_nodes)) { + return false; + } + } + if ((match.kind == AttentionQkvProjectionKind::Query || match.kind == AttentionQkvProjectionKind::Key) && + !append_covered_node_index_once(context.graph, context.covered_nodes, match.rope.node, + dispatch_match.covered_nodes)) { + return false; + } + if ((match.kind == AttentionQkvProjectionKind::Key || match.kind == AttentionQkvProjectionKind::Value) && + !append_covered_node_index_once(context.graph, context.covered_nodes, match.set_rows.node, + dispatch_match.covered_nodes)) { + return false; + } + return true; +} + +static bool match_attention_qkv_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + AttentionQkvMatch match = match_attention_qkv_projection(context); + if (!match.matched() || !cover_attention_qkv_nodes(context, match, dispatch_match)) { + return false; + } + + Dispatch dispatch; + bool pack_q4 = match.root.weight_format == CommonMulMatWeightFormat::Q4KRow64; + bool pack_q6 = match.root.weight_format == CommonMulMatWeightFormat::Q6KRow64; + bool use_q8_activation = false; + DispatchBinding activation = { match.root.input->id, 0, match.root.input->byte_count }; + if (pack_q4 || pack_q6) { + if (common_prepare_q8_1_x4_input(context, *match.root.input, match.root.input_size, match.root.token_count, + dispatch_match, activation, + CommonQ8ActivationPolicy::ExistingAlternateOnly)) { + use_q8_activation = true; + } else { + CommonMulMatWeightFormat fallback_format = CommonMulMatWeightFormat::Q4K; + if (!common_mul_mat_format_for_type(match.root.weight->type, fallback_format)) { + return false; + } + match.root.weight_format = fallback_format; + pack_q4 = false; + pack_q6 = false; + if (is_supported_decode_token_count(match.root.token_count)) { + switch (match.kind) { + case AttentionQkvProjectionKind::Query: + match.kernel = kAttentionQMatMulRopeF32F32DecodeKernel; + break; + case AttentionQkvProjectionKind::Key: + match.kernel = kAttentionKMatMulRopeSetRowsF32F32DecodeKernel; + break; + case AttentionQkvProjectionKind::Value: + match.kernel = kAttentionVMatMulSetRowsF32F32DecodeKernel; + break; + } + } + } + } + dispatch.kernel = make_kernel_specialization(match.kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.root.token_count); + add_attention_qkv_compile_parameters(dispatch, match); + if (use_q8_activation) { + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat.activation_format", std::to_string(GGML_TYPE_Q8_1)); + } + dispatch.bindings.push_back(activation); + if (pack_q4 || pack_q6) { + const char * layout = pack_q4 ? kQ4KPackedK256Row64Layout : kQ6KPackedK256Row64ScaleRowLayout; + dispatch.bindings.push_back({ match.root.weight->id, 0, match.root.weight->byte_count, layout, + match.root.weight->type, match.root.input_size, match.root.output_size, + match.root.weight->byte_count }); + } else { + dispatch.bindings.push_back({ match.root.weight->id, 0, match.root.weight->byte_count }); + } + + if (match.kind == AttentionQkvProjectionKind::Query) { + const auto [theta, freq_factors] = add_attention_rope_frequency_bindings(context, match.rope, dispatch_match); + dispatch.bindings.push_back({ match.rope.positions->id, 0, match.rope.positions->byte_count }); + dispatch.bindings.push_back({ theta, 0, match.rope.theta_bytes }); + dispatch.bindings.push_back({ freq_factors, 0, match.rope.freq_factors_bytes }); + dispatch.bindings.push_back({ match.rope.output->id, 0, match.rope.output->byte_count }); + } else if (match.kind == AttentionQkvProjectionKind::Key) { + const auto [theta, freq_factors] = add_attention_rope_frequency_bindings(context, match.rope, dispatch_match); + dispatch.bindings.push_back({ match.rope.positions->id, 0, match.rope.positions->byte_count }); + dispatch.bindings.push_back({ match.set_rows.indices->id, 0, match.set_rows.indices->byte_count }); + dispatch.bindings.push_back({ theta, 0, match.rope.theta_bytes }); + dispatch.bindings.push_back({ freq_factors, 0, match.rope.freq_factors_bytes }); + dispatch.bindings.push_back({ match.set_rows.output->id, 0, match.set_rows.output->byte_count }); + } else { + dispatch.bindings.push_back({ match.set_rows.indices->id, 0, match.set_rows.indices->byte_count }); + dispatch.bindings.push_back({ match.set_rows.output->id, 0, match.set_rows.output->byte_count }); + } + + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_llm_attention_qkv_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "llm.attention_qkv_matmul_postprocess.f32_f32_wmma", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 210, + DispatchSource::Llm, + match_attention_qkv_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-attention-qkv.h b/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-attention-qkv.h new file mode 100644 index 000000000000..27884ce1468a --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-attention-qkv.h @@ -0,0 +1,9 @@ +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +namespace ggml::hrx { + +void register_llm_attention_qkv_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-gated-delta-net.cpp b/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-gated-delta-net.cpp new file mode 100644 index 000000000000..d82613c864e1 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-gated-delta-net.cpp @@ -0,0 +1,1232 @@ +#include "dispatch-gated-delta-net.h" + +#include "../common/dispatch-mul-mat-common.h" +#include "../common/dispatch-rmsnorm.h" +#include "ggml.h" +#include "graph/graph-matcher.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kGatedDeltaNetProjectionEpilogueKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_gated_delta_net_projection_epilogue_f32"); +static constexpr KernelCatalogRef kGatedDeltaNetPrefillKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_gated_delta_net_f32_wmma_head128"); +static constexpr KernelCatalogRef kGatedDeltaNetRmsNormGateKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_gated_delta_net_f32_wmma_head128_rmsnorm_gate"); +static constexpr KernelCatalogRef kRmsNormGateKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_rmsnorm_gate_f32_f16"); +static constexpr KernelCatalogRef kGatedDeltaNetPrefillProjectionEpilogueKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_gated_delta_net_f32_wmma_head128_projection_epilogue"); +static constexpr KernelCatalogRef kGatedDeltaNetInplaceKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_gated_delta_net_f32_wmma_head128_inplace"); +static constexpr KernelCatalogRef kGatedDeltaNetInplaceProjectionEpilogueKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_gated_delta_net_f32_wmma_head128_inplace_projection_epilogue"); +static constexpr KernelCatalogRef kGatedDeltaNetSnapshotKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_gated_delta_net_f32_wmma_head128_snapshot"); +static constexpr KernelCatalogRef kGatedDeltaNetSnapshotProjectionEpilogueKernel = GGML_HRX_KERNEL_REF( + "loom_libs", "llm_gated_delta_net_f32_wmma_head128_snapshot_projection_epilogue"); +static constexpr KernelCatalogRef kGatedDeltaNetSelectedSnapshotKernel = GGML_HRX_KERNEL_REF( + "loom_libs", "llm_gated_delta_net_f32_wmma_head128_selected_snapshot_projection_epilogue"); +static constexpr KernelCatalogRef kGatedDeltaNetSelectedRmsQ8Kernel = GGML_HRX_KERNEL_REF( + "loom_libs", "llm_gated_delta_net_f32_wmma_head128_selected_snapshot_projection_rms_gate_q8"); +static constexpr KernelCatalogRef kRmsNormGateQ8Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_rmsnorm_gate_f32_q8_1_x4"); +static constexpr KernelCatalogRef kCopyF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_copy_f32"); +static constexpr KernelCatalogRef kMulMatSymmetricI4LowRowAdjacentDualWmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma"); +static constexpr KernelCatalogRef kMulMatDualQ4F32DecodeKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_dual_q4_f32_decode"); + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static bool is_shape(const Value & value, int64_t ne0, int64_t ne1, int64_t ne2, int64_t ne3) { + return value.ne[0] == ne0 && value.ne[1] == ne1 && value.ne[2] == ne2 && value.ne[3] == ne3; +} + +static bool is_f32(const Value * value) { + return value != nullptr && value->type == GGML_TYPE_F32; +} + +static bool distinct_storage(const Value & lhs, const Value & rhs) { + return lhs.storage != rhs.storage; +} + +static const GraphNode * producer_with_op(const Graph & graph, ValueId value, ggml_op op) { + const GraphNode * producer = graph.index().producer(value); + return producer != nullptr && producer->op == op ? producer : nullptr; +} + +static bool has_unary_op(const GraphNode * node, UnaryKind op) { + const UnaryParams * params = node != nullptr ? op_params_as(node->params) : nullptr; + return node != nullptr && node->op == GGML_OP_UNARY && params != nullptr && params->op == op; +} + +static bool append_covered_node(const DispatchMatchContext & context, const GraphNode * node, DispatchMatch & match) { + return append_covered_node_index_once(context.graph, context.covered_nodes, node, match.covered_nodes); +} + +struct GatedDeltaNetMatch { + std::vector covered; + const Value * alpha_raw = nullptr; + const Value * beta_raw = nullptr; + const Value * bias = nullptr; + const Value * a_scale = nullptr; + const Value * gate = nullptr; + const Value * gate_flat = nullptr; + const Value * beta = nullptr; + const Value * raw_q = nullptr; + const Value * raw_k = nullptr; + const Value * v = nullptr; + const Value * state = nullptr; + const Value * gdn_output = nullptr; + const Value * attention = nullptr; + const Value * new_state = nullptr; + const Value * cache = nullptr; + int64_t width = 0; + int64_t q_head_count = 0; + int64_t head_count = 0; + int64_t token_count = 0; + int64_t sequence_count = 0; + int64_t snapshot_count = 0; + float l2_epsilon = 0.0f; + + bool matched() const { + return !covered.empty() && gate != nullptr && beta != nullptr && raw_q != nullptr && raw_k != nullptr && + v != nullptr && state != nullptr && gdn_output != nullptr && new_state != nullptr && cache != nullptr; + } + + bool has_projection_epilogue() const { + return alpha_raw != nullptr && beta_raw != nullptr && bias != nullptr && a_scale != nullptr && + gate_flat != nullptr; + } +}; + +struct GatedDeltaNetProjectionPairMatch { + const GraphNode * first_node = nullptr; + const GraphNode * second_node = nullptr; + const Value * input = nullptr; + const Value * first_weight = nullptr; + const Value * second_weight = nullptr; + const Value * first_output = nullptr; + const Value * second_output = nullptr; + ValueId activation; + int64_t input_size = 0; + int64_t output_size = 0; + int64_t token_count = 0; + bool native_f32 = false; + + bool matched() const { + return first_node != nullptr && second_node != nullptr && input != nullptr && first_weight != nullptr && + second_weight != nullptr && first_output != nullptr && second_output != nullptr; + } +}; + +static GatedDeltaNetMatch match_gated_delta_net(const Graph & graph, const GraphNode * node) { + GatedDeltaNetMatch match; + if (node == nullptr || node->op != GGML_OP_L2_NORM || node->inputs.size() != 1 || !graph.has_index()) { + return match; + } + + const Value * raw_q = graph_value(graph, node->inputs[0]); + const Value * q_norm = graph_value(graph, node->output); + if (!is_f32(raw_q) || !is_f32(q_norm)) { + return {}; + } + const std::vector & q_consumers = graph.index().consumers(q_norm->id); + if (q_consumers.size() != 1 || q_consumers.front() == nullptr || + q_consumers.front()->op != GGML_OP_GATED_DELTA_NET || q_consumers.front()->inputs.size() != 6 || + q_consumers.front()->inputs[0] != q_norm->id) { + return {}; + } + const GraphNode * gdn = q_consumers.front(); + + const Value * k_norm = graph_value(graph, gdn->inputs[1]); + const Value * v = graph_value(graph, gdn->inputs[2]); + const Value * gate = graph_value(graph, gdn->inputs[3]); + const Value * beta = graph_value(graph, gdn->inputs[4]); + const Value * state = graph_value(graph, gdn->inputs[5]); + const Value * gdn_output = graph_value(graph, gdn->output); + const GraphNode * k_norm_node = k_norm != nullptr ? producer_with_op(graph, k_norm->id, GGML_OP_L2_NORM) : nullptr; + const Value * raw_k = k_norm_node != nullptr && k_norm_node->inputs.size() == 1 ? + graph_value(graph, k_norm_node->inputs[0]) : + nullptr; + if (!is_f32(k_norm) || !is_f32(raw_k) || !is_f32(v) || !is_f32(gate) || !is_f32(beta) || !is_f32(state) || + !is_f32(gdn_output)) { + return {}; + } + const RmsNormParams * l2_params = op_params_as(node->params); + if (l2_params == nullptr || !std::isfinite(l2_params->eps) || l2_params->eps <= 0.0f || + !op_params_equivalent(GGML_OP_L2_NORM, node->params, k_norm_node->params)) { + return {}; + } + + const int64_t width = raw_q->ne[0]; + const int64_t q_head_count = raw_q->ne[1]; + const int64_t token_count = raw_q->ne[2]; + const int64_t sequence_count = raw_q->ne[3]; + const int64_t head_count = v->ne[1]; + if (width != 128 || q_head_count <= 0 || q_head_count > 4096 || head_count <= 0 || head_count > 4096 || + token_count < 1 || token_count > 512 || sequence_count < 1 || sequence_count > 3) { + return {}; + } + const int64_t hidden_size = width * (2 * q_head_count + head_count); + if (!is_shape(*raw_k, width, q_head_count, token_count, sequence_count) || + !is_shape(*q_norm, width, q_head_count, token_count, sequence_count) || + !is_shape(*k_norm, width, q_head_count, token_count, sequence_count) || + !is_shape(*v, width, head_count, token_count, sequence_count) || + !is_shape(*gate, 1, head_count, token_count, sequence_count) || + !is_shape(*beta, 1, head_count, token_count, sequence_count) || + !is_shape(*state, width, width, head_count, sequence_count) || gdn_output->ne[0] != width * head_count || + gdn_output->ne[2] != 1 || gdn_output->ne[3] != 1) { + return {}; + } + const int64_t attention_rows = token_count * sequence_count; + const int64_t snapshot_rows = width * sequence_count; + if (gdn_output->ne[1] <= attention_rows || (gdn_output->ne[1] - attention_rows) % snapshot_rows != 0) { + return {}; + } + const int64_t snapshot_count = (gdn_output->ne[1] - attention_rows) / snapshot_rows; + if (snapshot_count < 1 || snapshot_count > 5) { + return {}; + } + if (!q_norm->contiguous || !k_norm->contiguous || !state->contiguous || !gdn_output->contiguous || + raw_q->storage != raw_k->storage || raw_q->storage != v->storage || raw_q->storage_offset != 0 || + raw_k->storage_offset != static_cast(width * q_head_count) * sizeof(float) || + v->storage_offset != static_cast(2 * width * q_head_count) * sizeof(float) || + raw_q->nb[0] != sizeof(float) || raw_q->nb[1] != static_cast(width) * sizeof(float) || + raw_q->nb[2] != static_cast(hidden_size) * sizeof(float) || raw_k->nb != raw_q->nb || + v->nb[0] != sizeof(float) || v->nb[1] != static_cast(width) * sizeof(float) || + v->nb[2] != static_cast(hidden_size) * sizeof(float)) { + return {}; + } + + const GraphNode * gate_reshape = producer_with_op(graph, gate->id, GGML_OP_RESHAPE); + const Value * gate_flat = gate_reshape != nullptr && gate_reshape->inputs.size() == 1 ? + graph_value(graph, gate_reshape->inputs[0]) : + nullptr; + const GraphNode * gate_mul = gate_flat != nullptr ? producer_with_op(graph, gate_flat->id, GGML_OP_MUL) : nullptr; + if (gate_mul == nullptr || gate_mul->inputs.size() != 2) { + return {}; + } + const Value * alpha_softplus = graph_value(graph, gate_mul->inputs[0]); + const Value * a_scale = graph_value(graph, gate_mul->inputs[1]); + const GraphNode * softplus = + alpha_softplus != nullptr ? producer_with_op(graph, alpha_softplus->id, GGML_OP_UNARY) : nullptr; + if (!has_unary_op(softplus, UnaryKind::SoftPlus) || softplus->inputs.size() != 1) { + return {}; + } + const Value * alpha_biased = graph_value(graph, softplus->inputs[0]); + const GraphNode * add_bias = + alpha_biased != nullptr ? producer_with_op(graph, alpha_biased->id, GGML_OP_ADD) : nullptr; + if (add_bias == nullptr || add_bias->inputs.size() != 2) { + return {}; + } + const Value * alpha = graph_value(graph, add_bias->inputs[0]); + const Value * bias = graph_value(graph, add_bias->inputs[1]); + const GraphNode * alpha_reshape = alpha != nullptr ? producer_with_op(graph, alpha->id, GGML_OP_RESHAPE) : nullptr; + const Value * alpha_raw = alpha_reshape != nullptr && alpha_reshape->inputs.size() == 1 ? + graph_value(graph, alpha_reshape->inputs[0]) : + nullptr; + + const GraphNode * sigmoid = producer_with_op(graph, beta->id, GGML_OP_UNARY); + if (!has_unary_op(sigmoid, UnaryKind::Sigmoid) || sigmoid->inputs.size() != 1) { + return {}; + } + const Value * beta_pre = graph_value(graph, sigmoid->inputs[0]); + const GraphNode * beta_reshape = + beta_pre != nullptr ? producer_with_op(graph, beta_pre->id, GGML_OP_RESHAPE) : nullptr; + const Value * beta_raw = beta_reshape != nullptr && beta_reshape->inputs.size() == 1 ? + graph_value(graph, beta_reshape->inputs[0]) : + nullptr; + if (!is_f32(alpha_raw) || !is_f32(beta_raw) || !is_f32(alpha) || !is_f32(alpha_biased) || !is_f32(alpha_softplus) || + !is_f32(a_scale) || !is_f32(bias) || !is_f32(gate_flat) || !is_f32(beta_pre) || + !is_shape(*alpha_raw, head_count, token_count * sequence_count, 1, 1) || + !is_shape(*beta_raw, head_count, token_count * sequence_count, 1, 1) || + !is_shape(*gate_flat, head_count, token_count, sequence_count, 1) || !is_shape(*bias, head_count, 1, 1, 1) || + !is_shape(*a_scale, head_count, 1, 1, 1) || !alpha_raw->contiguous || !beta_raw->contiguous || + !bias->contiguous || !a_scale->contiguous || !gate_flat->contiguous || !beta->contiguous) { + return {}; + } + + const std::array epilogue_values = { alpha_raw, beta_raw, bias, a_scale, gate_flat, beta }; + for (size_t lhs = 0; lhs < epilogue_values.size(); ++lhs) { + for (size_t rhs = lhs + 1; rhs < epilogue_values.size(); ++rhs) { + if (!distinct_storage(*epilogue_values[lhs], *epilogue_values[rhs])) { + return {}; + } + } + } + + const GraphNode * state_reshape = producer_with_op(graph, state->id, GGML_OP_RESHAPE); + if (state_reshape == nullptr || state_reshape->inputs.size() != 1) { + return {}; + } + const std::vector & gdn_consumers = graph.index().consumers(gdn_output->id); + if (gdn_consumers.size() != 2) { + return {}; + } + const GraphNode * new_state_view = nullptr; + const GraphNode * attention_view = nullptr; + const size_t attention_bytes = + static_cast(width * head_count * token_count * sequence_count) * sizeof(float); + for (const GraphNode * consumer : gdn_consumers) { + const Value * output = + consumer != nullptr && consumer->op == GGML_OP_VIEW ? graph_value(graph, consumer->output) : nullptr; + if (output != nullptr && output->storage == gdn_output->storage && + output->storage_offset == gdn_output->storage_offset + attention_bytes) { + new_state_view = consumer; + } else if (output != nullptr && output->storage == gdn_output->storage && + output->storage_offset == gdn_output->storage_offset && + is_shape(*output, width, head_count, token_count, sequence_count)) { + attention_view = consumer; + } else { + return {}; + } + } + const Value * new_state = new_state_view != nullptr ? graph_value(graph, new_state_view->output) : nullptr; + if (new_state == nullptr || attention_view == nullptr) { + return {}; + } + const int64_t written_snapshot_count = std::min(token_count, snapshot_count); + const size_t state_bytes = static_cast(width * width * head_count) * sizeof(float); + const size_t state_plane_bytes = state_bytes * static_cast(sequence_count); + const bool single_snapshot_state = + snapshot_count == 1 && is_shape(*new_state, width, width, head_count, sequence_count); + const bool rollback_snapshot_state = + is_shape(*new_state, width * width * head_count, sequence_count, written_snapshot_count, 1); + if ((!single_snapshot_state && !rollback_snapshot_state) || !new_state->contiguous || + new_state->storage != gdn_output->storage || + new_state->storage_offset != gdn_output->storage_offset + attention_bytes || + new_state->byte_count != state_plane_bytes * static_cast(written_snapshot_count)) { + return {}; + } + const std::vector & state_consumers = graph.index().consumers(new_state->id); + if (state_consumers.size() != 1 || state_consumers.front() == nullptr || + state_consumers.front()->op != GGML_OP_CPY || state_consumers.front()->inputs.size() != 2 || + state_consumers.front()->inputs[0] != new_state->id) { + return {}; + } + const GraphNode * cache_copy = state_consumers.front(); + const Value * cache_target = graph_value(graph, cache_copy->inputs[1]); + const Value * cache = graph_value(graph, cache_copy->output); + const GraphNode * cache_view = cache_target != nullptr ? graph.index().producer(cache_target->id) : nullptr; + if (!is_f32(cache_target) || !is_f32(cache) || cache_view == nullptr) { + return {}; + } + const size_t snapshot_stride_count = static_cast(written_snapshot_count - 1); + if (snapshot_stride_count != 0 && + cache->nb[2] > (std::numeric_limits::max() - state_bytes) / snapshot_stride_count) { + return {}; + } + const size_t required_cache_bytes = state_plane_bytes + snapshot_stride_count * cache->nb[2]; + if (cache_view->op != GGML_OP_VIEW || cache_view->inputs.size() != 1 || + !same_full_value_range(*cache_target, *cache) || cache->ne[0] != width * width * head_count || + cache->ne[1] != sequence_count || cache->ne[2] != written_snapshot_count || cache->ne[3] != 1 || + cache->nb[0] != sizeof(float) || cache->nb[1] != state_bytes || + (written_snapshot_count > 1 && cache->nb[2] < state_plane_bytes) || cache->byte_count < required_cache_bytes || + !distinct_storage(*new_state, *cache)) { + return {}; + } + + if (!distinct_storage(*gdn_output, *state) || !distinct_storage(*gdn_output, *gate_flat) || + !distinct_storage(*gdn_output, *beta)) { + return {}; + } + + match.covered = { alpha_reshape, add_bias, softplus, gate_mul, gate_reshape, + beta_reshape, sigmoid, node, k_norm_node, state_reshape, + gdn, new_state_view, cache_view, cache_copy, attention_view }; + match.alpha_raw = alpha_raw; + match.beta_raw = beta_raw; + match.bias = bias; + match.a_scale = a_scale; + match.gate = gate; + match.gate_flat = gate_flat; + match.beta = beta; + match.raw_q = raw_q; + match.raw_k = raw_k; + match.v = v; + match.state = state; + match.gdn_output = gdn_output; + match.attention = graph_value(graph, attention_view->output); + match.new_state = new_state; + match.cache = cache; + match.width = width; + match.q_head_count = q_head_count; + match.head_count = head_count; + match.token_count = token_count; + match.sequence_count = sequence_count; + match.snapshot_count = snapshot_count; + match.l2_epsilon = l2_params->eps; + return match; +} + +static GatedDeltaNetProjectionPairMatch match_gated_delta_net_projection_pair(const DispatchMatchContext & context) { + GatedDeltaNetProjectionPairMatch match; + const GraphNode * first_node = context.root_node; + if (first_node == nullptr || first_node->op != GGML_OP_MUL_MAT || first_node->inputs.size() != 2 || + !context.graph.has_index()) { + return match; + } + + const GraphNode * reshape = common_find_only_consumer_with_op(context.graph, first_node->output, GGML_OP_RESHAPE); + const GraphNode * add = + reshape != nullptr ? common_find_only_consumer_with_op(context.graph, reshape->output, GGML_OP_ADD) : nullptr; + const GraphNode * softplus = + add != nullptr ? common_find_only_consumer_with_op(context.graph, add->output, GGML_OP_UNARY) : nullptr; + const GraphNode * mul = has_unary_op(softplus, UnaryKind::SoftPlus) ? + common_find_only_consumer_with_op(context.graph, softplus->output, GGML_OP_MUL) : + nullptr; + const GraphNode * gate_reshape = + mul != nullptr ? common_find_only_consumer_with_op(context.graph, mul->output, GGML_OP_RESHAPE) : nullptr; + const GraphNode * gdn = + gate_reshape != nullptr ? + common_find_only_consumer_with_op(context.graph, gate_reshape->output, GGML_OP_GATED_DELTA_NET) : + nullptr; + if (gdn == nullptr || gdn->inputs.size() != 6) { + return {}; + } + + const GraphNode * q_norm = producer_with_op(context.graph, gdn->inputs[0], GGML_OP_L2_NORM); + const GatedDeltaNetMatch gdn_match = match_gated_delta_net(context.graph, q_norm); + if (!gdn_match.matched() || !gdn_match.has_projection_epilogue() || gdn_match.alpha_raw->id != first_node->output) { + return {}; + } + + const GraphNode * second_node = producer_with_op(context.graph, gdn_match.beta_raw->id, GGML_OP_MUL_MAT); + if (second_node == nullptr || second_node == first_node || second_node->inputs.size() != 2 || + second_node->inputs[1] != first_node->inputs[1]) { + return {}; + } + + CommonMulMatMatch first = common_match_mul_mat_any_format( + context.graph, first_node, kMulMatSymmetricI4LowRowAdjacentDualWmmaKernel, false); + CommonMulMatMatch second = common_match_mul_mat_any_format( + context.graph, second_node, kMulMatSymmetricI4LowRowAdjacentDualWmmaKernel, false); + if (!first.matched()) { + first = common_match_mul_mat_any_format( + context.graph, first_node, kMulMatDualQ4F32DecodeKernel, true); + } + if (!second.matched()) { + second = common_match_mul_mat_any_format( + context.graph, second_node, kMulMatDualQ4F32DecodeKernel, true); + } + if (!first.matched() || !second.matched() || first.weight->type != GGML_TYPE_Q4_K || + second.weight->type != GGML_TYPE_Q4_K || first.weight->alias_source.value >= 0 || + second.weight->alias_source.value >= 0 || first.input->id != second.input->id || + first.input_size != second.input_size || first.output_size != second.output_size || + first.token_count != second.token_count || first.output_size != gdn_match.head_count || + first.token_count != gdn_match.token_count * gdn_match.sequence_count || + !distinct_storage(*first.weight, *second.weight) || !distinct_storage(*first.output, *second.output)) { + return {}; + } + + const CommonSymmetricI4ActivationLayout activation_layout = + common_symmetric_i4_activation_layout(first.input_size, first.token_count); + const CommandPlanAlternateValue * alternate = find_alternate_value(context.graph, context.plan, first.input->id, + GGML_TYPE_COUNT, activation_layout.total_bytes); + const bool symmetric_input = first.token_count > 1 && + alternate != nullptr && alternate->name == kCommonSymmetricI4K32ActivationAlternateName; + const bool native_f32 = is_f32(first.input) && first.input->contiguous && + first.token_count == 1 && first.output_size >= 24 && first.output_size < 64 && + first.input_size >= 4096 && first.input_size <= 32768 && first.input_size % 1024 == 0; + if (!symmetric_input && !native_f32) { + return {}; + } + if (!symmetric_input) { + size_t second_index = 0; + if (!context.graph.index().node_index(second_node, second_index) || second_index <= context.root_index) { + return {}; + } + for (const Value * input : { first.input, second.weight }) { + const GraphNode * producer = context.graph.index().producer(input->storage_root); + size_t index = 0; + if (producer != nullptr && (!context.graph.index().node_index(producer, index) || + index >= context.covered_nodes.size() || !context.covered_nodes[index])) { + return {}; + } + } + const auto & nodes = context.graph.nodes(); + for (size_t i = context.root_index + 1; i < second_index; ++i) { + const GraphNode & node = nodes[i]; + if (node.op == GGML_OP_NONE || node.op == GGML_OP_VIEW || node.op == GGML_OP_RESHAPE || + node.op == GGML_OP_PERMUTE || node.op == GGML_OP_TRANSPOSE) { + continue; + } + const Value * output = graph_value(context.graph, node.output); + if (output != nullptr && (output->storage == first.input->storage || + output->storage == second.weight->storage)) { + return {}; + } + } + } + + match.first_node = first_node; + match.second_node = second_node; + match.input = first.input; + match.first_weight = first.weight; + match.second_weight = second.weight; + match.first_output = first.output; + match.second_output = second.output; + match.activation = symmetric_input ? alternate->alternate_value : ValueId{}; + match.native_f32 = !symmetric_input; + match.input_size = first.input_size; + match.output_size = first.output_size; + match.token_count = first.token_count; + return match; +} + +static GatedDeltaNetMatch match_direct_gated_delta_net(const Graph & graph, const GraphNode * gdn) { + GatedDeltaNetMatch match; + if (gdn == nullptr || gdn->op != GGML_OP_GATED_DELTA_NET || gdn->inputs.size() != 6 || !graph.has_index()) { + return match; + } + + const Value * q_norm = graph_value(graph, gdn->inputs[0]); + const Value * k_norm = graph_value(graph, gdn->inputs[1]); + const Value * v = graph_value(graph, gdn->inputs[2]); + const Value * gate = graph_value(graph, gdn->inputs[3]); + const Value * beta = graph_value(graph, gdn->inputs[4]); + const Value * state = graph_value(graph, gdn->inputs[5]); + const Value * gdn_output = graph_value(graph, gdn->output); + if (!is_f32(q_norm) || !is_f32(k_norm) || !is_f32(v) || !is_f32(gate) || !is_f32(beta) || !is_f32(state) || + !is_f32(gdn_output)) { + return {}; + } + + const int64_t width = q_norm->ne[0]; + const int64_t q_head_count = q_norm->ne[1]; + const int64_t token_count = q_norm->ne[2]; + const int64_t sequence_count = q_norm->ne[3]; + const int64_t head_count = v->ne[1]; + if (width != 128 || q_head_count <= 0 || q_head_count > 4096 || head_count <= 0 || head_count > 4096 || + token_count < 1 || token_count > 512 || sequence_count < 1 || sequence_count > 3) { + return {}; + } + + if (!is_shape(*q_norm, width, q_head_count, token_count, sequence_count) || + !is_shape(*k_norm, width, q_head_count, token_count, sequence_count) || + !is_shape(*v, width, head_count, token_count, sequence_count) || + !is_shape(*gate, 1, head_count, token_count, sequence_count) || + !is_shape(*beta, 1, head_count, token_count, sequence_count) || + !is_shape(*state, width, width, head_count, sequence_count) || gdn_output->ne[0] != width * head_count || + gdn_output->ne[2] != 1 || gdn_output->ne[3] != 1 || !q_norm->contiguous || !k_norm->contiguous || + !gate->contiguous || !beta->contiguous || !state->contiguous || !gdn_output->contiguous || + q_norm->nb[0] != sizeof(float) || q_norm->nb[1] != static_cast(width) * sizeof(float) || + q_norm->nb[2] != static_cast(width * q_head_count) * sizeof(float) || k_norm->nb != q_norm->nb || + v->nb[0] != sizeof(float) || v->nb[1] != static_cast(width) * sizeof(float) || + v->nb[2] < static_cast(width * head_count) * sizeof(float) || v->nb[2] % sizeof(float) != 0 || + v->nb[3] % sizeof(float) != 0) { + return {}; + } + + const int64_t attention_rows = token_count * sequence_count; + const int64_t snapshot_rows = width * sequence_count; + if (gdn_output->ne[1] <= attention_rows || (gdn_output->ne[1] - attention_rows) % snapshot_rows != 0) { + return {}; + } + const int64_t snapshot_count = (gdn_output->ne[1] - attention_rows) / snapshot_rows; + if (snapshot_count < 1 || snapshot_count > 5) { + return {}; + } + + const std::vector & gdn_consumers = graph.index().consumers(gdn_output->id); + if (gdn_consumers.size() != 2) { + return {}; + } + const GraphNode * attention_view = nullptr; + const GraphNode * new_state_view = nullptr; + const size_t attention_bytes = + static_cast(width * head_count * token_count * sequence_count) * sizeof(float); + for (const GraphNode * consumer : gdn_consumers) { + const Value * output = + consumer != nullptr && consumer->op == GGML_OP_VIEW ? graph_value(graph, consumer->output) : nullptr; + if (output != nullptr && output->storage == gdn_output->storage && + output->storage_offset == gdn_output->storage_offset && + is_shape(*output, width, head_count, token_count, sequence_count)) { + attention_view = consumer; + } else if (output != nullptr && output->storage == gdn_output->storage && + output->storage_offset == gdn_output->storage_offset + attention_bytes) { + new_state_view = consumer; + } else { + return {}; + } + } + + const Value * new_state = new_state_view != nullptr ? graph_value(graph, new_state_view->output) : nullptr; + if (attention_view == nullptr || new_state == nullptr) { + return {}; + } + const int64_t written_snapshot_count = std::min(token_count, snapshot_count); + const size_t state_bytes = static_cast(width * width * head_count) * sizeof(float); + const size_t state_plane_bytes = state_bytes * static_cast(sequence_count); + const bool single_snapshot_state = + snapshot_count == 1 && is_shape(*new_state, width, width, head_count, sequence_count); + const bool rollback_snapshot_state = + is_shape(*new_state, width * width * head_count, sequence_count, written_snapshot_count, 1); + if ((!single_snapshot_state && !rollback_snapshot_state) || !new_state->contiguous || + new_state->storage != gdn_output->storage || + new_state->storage_offset != gdn_output->storage_offset + attention_bytes || + new_state->byte_count != state_plane_bytes * static_cast(written_snapshot_count)) { + return {}; + } + + const std::vector & state_consumers = graph.index().consumers(new_state->id); + if (state_consumers.size() != 1 || state_consumers.front() == nullptr || + state_consumers.front()->op != GGML_OP_CPY || state_consumers.front()->inputs.size() != 2 || + state_consumers.front()->inputs[0] != new_state->id) { + return {}; + } + const GraphNode * cache_copy = state_consumers.front(); + const Value * cache_target = graph_value(graph, cache_copy->inputs[1]); + const Value * cache = graph_value(graph, cache_copy->output); + const GraphNode * cache_view = cache_target != nullptr ? graph.index().producer(cache_target->id) : nullptr; + if (!is_f32(cache_target) || !is_f32(cache) || + (cache_view != nullptr && (cache_view->op != GGML_OP_VIEW || cache_view->inputs.size() != 1)) || + !same_full_value_range(*cache_target, *cache) || cache->ne[0] != width * width * head_count || + cache->ne[1] != sequence_count || cache->ne[2] != written_snapshot_count || cache->ne[3] != 1 || + cache->nb[0] != sizeof(float) || cache->nb[1] != state_bytes || + (written_snapshot_count > 1 && cache->nb[2] < state_plane_bytes) || !distinct_storage(*new_state, *cache) || + !distinct_storage(*gdn_output, *state) || !distinct_storage(*gdn_output, *gate) || + !distinct_storage(*gdn_output, *beta)) { + return {}; + } + const size_t snapshot_stride_count = static_cast(written_snapshot_count - 1); + if (snapshot_stride_count != 0 && + cache->nb[2] > (std::numeric_limits::max() - state_bytes) / snapshot_stride_count) { + return {}; + } + const size_t required_cache_bytes = state_plane_bytes + snapshot_stride_count * cache->nb[2]; + if (cache->byte_count < required_cache_bytes) { + return {}; + } + + match.covered = { gdn, attention_view, new_state_view, cache_copy }; + match.gate = gate; + match.beta = beta; + match.raw_q = q_norm; + match.raw_k = k_norm; + match.v = v; + match.state = state; + match.gdn_output = gdn_output; + match.attention = graph_value(graph, attention_view->output); + match.new_state = new_state; + match.cache = cache; + match.width = width; + match.q_head_count = q_head_count; + match.head_count = head_count; + match.token_count = token_count; + match.sequence_count = sequence_count; + match.snapshot_count = snapshot_count; + match.l2_epsilon = 1.0e-6f; + return match; +} + +static void set_compile_parameter(KernelSpecialization & kernel, const char * name, int64_t value) { + kernel.compile_parameters.emplace(name, std::to_string(value)); +} + +static void set_float_compile_parameter(KernelSpecialization & kernel, const char * name, float value) { + std::ostringstream stream; + stream << std::setprecision(std::numeric_limits::max_digits10) << value; + kernel.compile_parameters.emplace(name, stream.str()); +} + +static void configure_gated_delta_net_kernel(Dispatch & dispatch, + const GatedDeltaNetMatch & match, + int64_t token_count) { + set_compile_parameter(dispatch.kernel, "llm.gated_delta_net.head_width", match.width); + set_compile_parameter(dispatch.kernel, "llm.gated_delta_net.head_count", match.head_count); + set_compile_parameter(dispatch.kernel, "llm.gated_delta_net.token_count", token_count); + set_compile_parameter(dispatch.kernel, "llm.gated_delta_net.sequence_count", match.sequence_count); + set_compile_parameter(dispatch.kernel, "llm.gated_delta_net.qk_stride1", match.raw_q->nb[1] / sizeof(float)); + set_compile_parameter(dispatch.kernel, "llm.gated_delta_net.qk_stride2", match.raw_q->nb[2] / sizeof(float)); + set_compile_parameter(dispatch.kernel, "llm.gated_delta_net.qk_stride3", match.raw_q->nb[3] / sizeof(float)); + set_compile_parameter(dispatch.kernel, "llm.gated_delta_net.value_stride1", match.v->nb[1] / sizeof(float)); + set_compile_parameter(dispatch.kernel, "llm.gated_delta_net.value_stride2", match.v->nb[2] / sizeof(float)); + set_compile_parameter(dispatch.kernel, "llm.gated_delta_net.value_stride3", match.v->nb[3] / sizeof(float)); + set_compile_parameter(dispatch.kernel, "llm.gated_delta_net.scalar_stride1", match.beta->nb[1] / sizeof(float)); + set_compile_parameter(dispatch.kernel, "llm.gated_delta_net.scalar_stride2", match.beta->nb[2] / sizeof(float)); + set_compile_parameter(dispatch.kernel, "llm.gated_delta_net.scalar_stride3", match.beta->nb[3] / sizeof(float)); + set_compile_parameter(dispatch.kernel, "llm.gated_delta_net.query_head_count", match.q_head_count); + set_compile_parameter(dispatch.kernel, "llm.gated_delta_net.query_sequence_ratio", 1); + set_float_compile_parameter(dispatch.kernel, "llm.gated_delta_net.l2_epsilon", match.l2_epsilon); + set_compile_parameter(dispatch.kernel, "llm.gated_delta_net.workgroup_size", 256); +} + +static bool match_gated_delta_net_projection_pair_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const GatedDeltaNetProjectionPairMatch match = match_gated_delta_net_projection_pair(context); + if (!match.matched() || !append_covered_node(context, match.first_node, dispatch_match) || + !append_covered_node(context, match.second_node, dispatch_match)) { + return false; + } + + if (match.native_f32) { + Dispatch projections; + projections.kernel = make_kernel_specialization(kMulMatDualQ4F32DecodeKernel); + set_compile_parameter(projections.kernel, "ggml.mul_mat_dual_q4_f32_c1.input_size", match.input_size); + set_compile_parameter(projections.kernel, "ggml.mul_mat_dual_q4_f32_c1.output_size", match.output_size); + projections.bindings = { + { match.input->id, 0, match.input->byte_count }, + { match.first_weight->id, 0, match.first_weight->byte_count }, + { match.second_weight->id, 0, match.second_weight->byte_count }, + { match.first_output->id, 0, match.first_output->byte_count }, + { match.second_output->id, 0, match.second_output->byte_count }, + }; + dispatch_match.dispatches.push_back(std::move(projections)); + return true; + } + + const CommonSymmetricI4ActivationLayout activation_layout = + common_symmetric_i4_activation_layout(match.input_size, match.token_count); + + Dispatch projections; + projections.kernel = make_kernel_specialization(kMulMatSymmetricI4LowRowAdjacentDualWmmaKernel); + projections.kernel.compile_parameters.emplace("ggml.mul_mat.symmetric_i4.lowrow.input_size", + common_to_config_value(match.input_size)); + projections.kernel.compile_parameters.emplace("ggml.mul_mat.symmetric_i4.lowrow.output_size", + common_to_config_value(match.output_size)); + projections.kernel.compile_parameters.emplace("ggml.mul_mat.symmetric_i4.lowrow.token_count", + common_to_config_value(match.token_count)); + projections.kernel.compile_parameters.emplace( + "ggml.mul_mat.symmetric_i4.lowrow.row_group_size", + common_to_config_value(static_cast( + common_symmetric_shared4_row_group_size(match.input_size, match.output_size, 4)))); + projections.bindings.push_back( + common_symmetric_i4_shared4_weight_binding(*match.first_weight, match.input_size, match.output_size)); + projections.bindings.push_back( + common_symmetric_i4_shared4_weight_binding(*match.second_weight, match.input_size, match.output_size)); + projections.bindings.push_back({ match.first_output->id, 0, match.first_output->byte_count }); + projections.bindings.push_back({ match.second_output->id, 0, match.second_output->byte_count }); + projections.bindings.push_back({ match.activation, 0, activation_layout.payload_bytes }); + projections.bindings.push_back( + { match.activation, activation_layout.scales_offset, activation_layout.metadata_bytes }); + projections.bindings.push_back( + { match.activation, activation_layout.sums_offset, activation_layout.metadata_bytes }); + dispatch_match.dispatches.push_back(std::move(projections)); + return true; +} + +static bool match_gated_delta_net_rmsnorm_gate(const DispatchMatchContext & context, + const GatedDeltaNetMatch & gdn, + DispatchMatch & match) { + if (gdn.attention->kind != ValueKind::Transient || gdn.gdn_output->kind != ValueKind::Transient) { + return false; + } + const GraphNode * rms = common_find_only_consumer_with_op(context.graph, gdn.attention->id, GGML_OP_RMS_NORM); + DispatchMatchContext rms_context = context; + rms_context.root_node = rms; + if (rms == nullptr || !context.graph.index().node_index(rms, rms_context.root_index) || + !common_match_rmsnorm_gate_dispatch(rms_context, match) || match.dispatches.size() != 1) { + return false; + } + const Dispatch & norm = match.dispatches.front(); + if (norm.kernel.kernel_id != kRmsNormGateKernel.id || norm.bindings.size() != 5 || + norm.kernel.compile_parameters.at("ggml.rmsnorm_gate_f32.hidden_size") != std::to_string(gdn.width) || + norm.kernel.compile_parameters.at("ggml.rmsnorm_gate_f32.f16_output_row_width") != + std::to_string(gdn.width * gdn.head_count)) { + return false; + } + // These inputs remain live while the fused operation publishes its rows. + const Value * output = graph_value(context.graph, norm.bindings[3].value); + for (const Value * input : {gdn.raw_q, gdn.raw_k, gdn.v, gdn.state, gdn.gate, gdn.beta, gdn.cache}) { + if (!distinct_storage(*output, *input)) { + return false; + } + } + return true; +} + +static bool match_gated_delta_net_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const GatedDeltaNetMatch match = context.root_node != nullptr && context.root_node->op == GGML_OP_GATED_DELTA_NET ? + match_direct_gated_delta_net(context.graph, context.root_node) : + match_gated_delta_net(context.graph, context.root_node); + if (!match.matched()) { + return false; + } + + const bool can_fuse_projection_epilogue = match.has_projection_epilogue() && match.token_count <= 16; + DispatchMatch rms_match; + const bool fuse_rmsnorm_gate = match.snapshot_count == 1 && !can_fuse_projection_epilogue && + match_gated_delta_net_rmsnorm_gate(context, match, rms_match); + Dispatch rms_gate; + if (fuse_rmsnorm_gate) { + rms_gate = std::move(rms_match.dispatches.front()); + rms_match.dispatches.clear(); + dispatch_match = std::move(rms_match); + } + + for (const GraphNode * node : match.covered) { + if (!append_covered_node(context, node, dispatch_match)) { + return false; + } + } + + const int64_t written_snapshot_count = std::min(match.token_count, match.snapshot_count); + const int64_t prefix_token_count = match.token_count - written_snapshot_count; + const size_t cache_stride = match.cache->nb[2]; + const bool fuse_snapshot_projection_epilogue = + can_fuse_projection_epilogue && match.snapshot_count != 1 && prefix_token_count == 0 && + written_snapshot_count == match.token_count && (match.token_count >= 2 || match.sequence_count > 1) && + match.token_count <= 5 && cache_stride % sizeof(float) == 0; + + if (match.has_projection_epilogue() && !can_fuse_projection_epilogue) { + Dispatch epilogue; + epilogue.kernel = make_kernel_specialization(kGatedDeltaNetProjectionEpilogueKernel); + set_compile_parameter(epilogue.kernel, "llm.gated_delta_net.epilogue_head_count", match.head_count); + set_compile_parameter(epilogue.kernel, "llm.gated_delta_net.epilogue_element_count", + match.head_count * match.token_count * match.sequence_count); + set_compile_parameter(epilogue.kernel, "llm.gated_delta_net.epilogue_workgroup_size", 256); + epilogue.bindings.push_back({ match.alpha_raw->id, 0, match.alpha_raw->byte_count }); + epilogue.bindings.push_back({ match.beta_raw->id, 0, match.beta_raw->byte_count }); + epilogue.bindings.push_back({ match.bias->id, 0, match.bias->byte_count }); + epilogue.bindings.push_back({ match.a_scale->id, 0, match.a_scale->byte_count }); + epilogue.bindings.push_back({ match.gate_flat->id, 0, match.gate_flat->byte_count }); + epilogue.bindings.push_back({ match.beta->id, 0, match.beta->byte_count }); + dispatch_match.dispatches.push_back(std::move(epilogue)); + } + + if (match.snapshot_count == 1) { + Dispatch gdn; + gdn.kernel = make_kernel_specialization(fuse_rmsnorm_gate ? kGatedDeltaNetRmsNormGateKernel : + can_fuse_projection_epilogue ? + kGatedDeltaNetPrefillProjectionEpilogueKernel : + kGatedDeltaNetPrefillKernel); + configure_gated_delta_net_kernel(gdn, match, match.token_count); + gdn.bindings.push_back({ match.raw_q->id, 0, match.raw_q->byte_count }); + gdn.bindings.push_back({ match.raw_k->id, 0, match.raw_k->byte_count }); + gdn.bindings.push_back({ match.v->id, 0, match.v->byte_count }); + if (can_fuse_projection_epilogue) { + gdn.bindings.push_back({ match.alpha_raw->id, 0, match.alpha_raw->byte_count }); + gdn.bindings.push_back({ match.beta_raw->id, 0, match.beta_raw->byte_count }); + gdn.bindings.push_back({ match.bias->id, 0, match.bias->byte_count }); + gdn.bindings.push_back({ match.a_scale->id, 0, match.a_scale->byte_count }); + } else { + gdn.bindings.push_back({ match.gate->id, 0, match.gate->byte_count }); + gdn.bindings.push_back({ match.beta->id, 0, match.beta->byte_count }); + } + gdn.bindings.push_back({ match.state->id, 0, match.state->byte_count }); + gdn.bindings.push_back({ match.gdn_output->id, 0, match.gdn_output->byte_count }); + if (fuse_rmsnorm_gate) { + for (const char * parameter : {"ggml.rmsnorm_gate_f32.rms_epsilon", "ggml.rmsnorm_gate_f32.gate_op"}) { + gdn.kernel.compile_parameters.emplace(parameter, rms_gate.kernel.compile_parameters.at(parameter)); + } + gdn.bindings.insert(gdn.bindings.end(), rms_gate.bindings.begin() + 1, rms_gate.bindings.end()); + } + + Dispatch cache_copy; + cache_copy.kernel = make_kernel_specialization(kCopyF32Kernel); + cache_copy.kernel.integer_parameters.emplace("element_count", match.new_state->element_count); + cache_copy.bindings.push_back({ match.new_state->id, 0, match.new_state->byte_count }); + cache_copy.bindings.push_back({ match.cache->id, 0, match.cache->byte_count }); + + dispatch_match.dispatches.push_back(std::move(gdn)); + dispatch_match.dispatches.push_back(std::move(cache_copy)); + return true; + } + + const size_t q_token_bytes = static_cast(match.width * match.q_head_count) * sizeof(float); + const size_t v_token_bytes = static_cast(match.width * match.head_count) * sizeof(float); + const size_t gate_token_bytes = static_cast(match.head_count) * sizeof(float); + const size_t attention_token_bytes = v_token_bytes; + const size_t state_bytes = static_cast(match.width * match.width * match.head_count) * sizeof(float); + const size_t state_plane_bytes = state_bytes * static_cast(match.sequence_count); + + if (prefix_token_count == 0 && written_snapshot_count == match.token_count && + (match.token_count >= 2 || match.sequence_count > 1) && match.token_count <= 5 && + cache_stride % sizeof(float) == 0) { + Dispatch gdn; + gdn.kernel = make_kernel_specialization(fuse_snapshot_projection_epilogue ? + kGatedDeltaNetSnapshotProjectionEpilogueKernel : + kGatedDeltaNetSnapshotKernel); + configure_gated_delta_net_kernel(gdn, match, match.token_count); + set_compile_parameter(gdn.kernel, "llm.gated_delta_net.snapshot_stride", cache_stride / sizeof(float)); + gdn.bindings.push_back({ match.raw_q->id, 0, match.raw_q->byte_count }); + gdn.bindings.push_back({ match.raw_k->id, 0, match.raw_k->byte_count }); + gdn.bindings.push_back({ match.v->id, 0, match.v->byte_count }); + if (fuse_snapshot_projection_epilogue) { + gdn.bindings.push_back({ match.alpha_raw->id, 0, match.alpha_raw->byte_count }); + gdn.bindings.push_back({ match.beta_raw->id, 0, match.beta_raw->byte_count }); + gdn.bindings.push_back({ match.bias->id, 0, match.bias->byte_count }); + gdn.bindings.push_back({ match.a_scale->id, 0, match.a_scale->byte_count }); + } else { + gdn.bindings.push_back({ match.gate->id, 0, match.gate->byte_count }); + gdn.bindings.push_back({ match.beta->id, 0, match.beta->byte_count }); + } + gdn.bindings.push_back({ match.state->id, 0, state_plane_bytes }); + gdn.bindings.push_back( + { match.cache->id, 0, state_plane_bytes + static_cast(written_snapshot_count - 1) * cache_stride }); + gdn.bindings.push_back( + { match.gdn_output->id, 0, + static_cast(match.token_count * match.sequence_count) * attention_token_bytes }); + dispatch_match.dispatches.push_back(std::move(gdn)); + return true; + } + + auto append_state_copy = [&](ValueId source, size_t source_offset, ValueId target, size_t target_offset) { + Dispatch copy; + copy.kernel = make_kernel_specialization(kCopyF32Kernel); + copy.kernel.integer_parameters.emplace("element_count", state_bytes / sizeof(float)); + copy.bindings.push_back({ source, source_offset, state_bytes }); + copy.bindings.push_back({ target, target_offset, state_bytes }); + dispatch_match.dispatches.push_back(std::move(copy)); + }; + + if (prefix_token_count > 0) { + const size_t q_span = static_cast(prefix_token_count - 1) * match.raw_q->nb[2] + q_token_bytes; + const size_t k_span = static_cast(prefix_token_count - 1) * match.raw_k->nb[2] + q_token_bytes; + const size_t v_span = static_cast(prefix_token_count - 1) * match.v->nb[2] + v_token_bytes; + const size_t gate_span = static_cast(prefix_token_count - 1) * match.gate->nb[2] + gate_token_bytes; + const size_t beta_span = static_cast(prefix_token_count - 1) * match.beta->nb[2] + gate_token_bytes; + const size_t prefix_attention_bytes = static_cast(prefix_token_count) * attention_token_bytes; + + Dispatch prefix; + prefix.kernel = make_kernel_specialization(can_fuse_projection_epilogue ? + kGatedDeltaNetPrefillProjectionEpilogueKernel : + kGatedDeltaNetPrefillKernel); + configure_gated_delta_net_kernel(prefix, match, prefix_token_count); + prefix.bindings.push_back({ match.raw_q->id, 0, q_span }); + prefix.bindings.push_back({ match.raw_k->id, 0, k_span }); + prefix.bindings.push_back({ match.v->id, 0, v_span }); + if (can_fuse_projection_epilogue) { + const size_t alpha_span = static_cast(prefix_token_count - 1) * match.alpha_raw->nb[1] + + gate_token_bytes; + const size_t beta_raw_span = static_cast(prefix_token_count - 1) * match.beta_raw->nb[1] + + gate_token_bytes; + prefix.bindings.push_back({ match.alpha_raw->id, 0, alpha_span }); + prefix.bindings.push_back({ match.beta_raw->id, 0, beta_raw_span }); + prefix.bindings.push_back({ match.bias->id, 0, match.bias->byte_count }); + prefix.bindings.push_back({ match.a_scale->id, 0, match.a_scale->byte_count }); + } else { + prefix.bindings.push_back({ match.gate->id, 0, gate_span }); + prefix.bindings.push_back({ match.beta->id, 0, beta_span }); + } + prefix.bindings.push_back({ match.state->id, 0, state_bytes }); + prefix.bindings.push_back({ match.gdn_output->id, 0, prefix_attention_bytes + state_bytes }); + dispatch_match.dispatches.push_back(std::move(prefix)); + append_state_copy(match.gdn_output->id, prefix_attention_bytes, match.cache->id, + static_cast(written_snapshot_count - 1) * cache_stride); + } else { + append_state_copy(match.state->id, 0, match.cache->id, + static_cast(written_snapshot_count - 1) * cache_stride); + } + + for (int64_t snapshot = 0; snapshot < written_snapshot_count; ++snapshot) { + const int64_t token = prefix_token_count + snapshot; + const int64_t slot = written_snapshot_count - 1 - snapshot; + if (snapshot > 0) { + append_state_copy(match.cache->id, static_cast(slot + 1) * cache_stride, match.cache->id, + static_cast(slot) * cache_stride); + } + + Dispatch gdn; + gdn.kernel = make_kernel_specialization(can_fuse_projection_epilogue ? + kGatedDeltaNetInplaceProjectionEpilogueKernel : + kGatedDeltaNetInplaceKernel); + configure_gated_delta_net_kernel(gdn, match, 1); + gdn.bindings.push_back({ match.raw_q->id, static_cast(token) * match.raw_q->nb[2], q_token_bytes }); + gdn.bindings.push_back({ match.raw_k->id, static_cast(token) * match.raw_k->nb[2], q_token_bytes }); + gdn.bindings.push_back({ match.v->id, static_cast(token) * match.v->nb[2], v_token_bytes }); + if (can_fuse_projection_epilogue) { + gdn.bindings.push_back( + { match.alpha_raw->id, static_cast(token) * match.alpha_raw->nb[1], gate_token_bytes }); + gdn.bindings.push_back( + { match.beta_raw->id, static_cast(token) * match.beta_raw->nb[1], gate_token_bytes }); + gdn.bindings.push_back({ match.bias->id, 0, match.bias->byte_count }); + gdn.bindings.push_back({ match.a_scale->id, 0, match.a_scale->byte_count }); + } else { + gdn.bindings.push_back( + { match.gate->id, static_cast(token) * match.gate->nb[2], gate_token_bytes }); + gdn.bindings.push_back( + { match.beta->id, static_cast(token) * match.beta->nb[2], gate_token_bytes }); + } + gdn.bindings.push_back({ match.cache->id, static_cast(slot) * cache_stride, state_bytes }); + gdn.bindings.push_back( + { match.gdn_output->id, static_cast(token) * attention_token_bytes, attention_token_bytes }); + dispatch_match.dispatches.push_back(std::move(gdn)); + } + + return true; +} +static bool match_gated_delta_net_rmsnorm_q8(const DispatchMatchContext & context, + const GatedDeltaNetMatch & gdn, + const Value & state_cache, + DispatchMatch & result) { + if (gdn.head_count < 24 || gdn.attention->kind != ValueKind::Transient || !gdn.attention->contiguous) { + return false; + } + const Graph & graph = context.graph; + const Value * input = gdn.attention; + const GraphNode * reshape = nullptr; + const GraphNode * rms = common_find_only_consumer_with_op(graph, input->id, GGML_OP_RMS_NORM); + if (rms == nullptr) { + reshape = common_find_only_consumer_with_op(graph, input->id, GGML_OP_RESHAPE); + const Value * reshaped = reshape != nullptr ? graph_value(graph, reshape->output) : nullptr; + if (reshaped == nullptr || !is_layout_alias_node(graph, *reshape) || !reshaped->contiguous || + !same_full_value_range(*input, *reshaped)) { + return false; + } + input = reshaped; + rms = common_find_only_consumer_with_op(graph, input->id, GGML_OP_RMS_NORM); + } + DispatchMatchContext rms_context = context; + rms_context.root_node = rms; + DispatchMatch match; + if (rms == nullptr || !graph.index().node_index(rms, rms_context.root_index) || + !common_match_rmsnorm_gate_dispatch(rms_context, match) || match.dispatches.size() != 1) { + return false; + } + const Dispatch & norm = match.dispatches.front(); + if (norm.kernel.kernel_id != kRmsNormGateQ8Kernel.id || norm.bindings.size() != 5 || + norm.bindings[0].value != input->id || + norm.kernel.integer_parameters.at("token_count") != gdn.head_count * gdn.token_count || + norm.kernel.compile_parameters.at("ggml.rmsnorm_gate_f32.hidden_size") != "128" || + norm.kernel.compile_parameters.at("ggml.rmsnorm_gate_f32.gate_op") != "15" || + norm.kernel.compile_parameters.at("ggml.rmsnorm_gate_f32.f16_output_row_width") != "0") { + return false; + } + const Value * output = graph_value(graph, norm.bindings[3].value); + if (output == nullptr || output->kind != ValueKind::Transient || !distinct_storage(*output, state_cache)) { + return false; + } + for (const Value * value : {gdn.raw_q, gdn.raw_k, gdn.v, gdn.alpha_raw, gdn.beta_raw, + gdn.bias, gdn.a_scale, gdn.state, gdn.cache}) { + if (!distinct_storage(*output, *value)) { + return false; + } + } + for (size_t binding : {size_t{1}, size_t{2}}) { + const Value * value = graph_value(graph, norm.bindings[binding].value); + if (value == nullptr || !distinct_storage(*value, *gdn.cache) || !distinct_storage(*value, *output)) { + return false; + } + const GraphNode * producer = graph.index().producer(value->storage_root); + size_t index = 0; + if (producer != nullptr && (!graph.index().node_index(producer, index) || + index >= context.covered_nodes.size() || !context.covered_nodes[index])) { + return false; + } + } + if (reshape != nullptr && !append_covered_node(context, reshape, match)) { + return false; + } + result = std::move(match); + return true; +} + +static bool match_selected_gated_delta_net_dispatch(const DispatchMatchContext & context, DispatchMatch & result) { + const Graph & graph = context.graph; + const GraphNode * gather = context.root_node; + if (gather == nullptr || gather->op != GGML_OP_GET_ROWS || gather->inputs.size() != 2 || !graph.has_index()) { + return false; + } + const Value * cache = graph_value(graph, gather->inputs[0]); + const Value * ids = graph_value(graph, gather->inputs[1]); + const Value * gathered = graph_value(graph, gather->output); + if (!is_f32(cache) || !is_f32(gathered) || ids == nullptr || ids->type != GGML_TYPE_I32 || + !cache->contiguous || !ids->contiguous || !gathered->contiguous || + cache->ne[0] <= 0 || cache->ne[1] < 1 || cache->ne[1] > 262208 || + cache->ne[2] != 1 || cache->ne[3] != 1 || cache->element_count > (int64_t{1} << 30) || + !is_shape(*ids, 1, 1, 1, 1) || !is_shape(*gathered, cache->ne[0], 1, 1, 1) || + gathered->kind != ValueKind::Transient || !graph.index().has_single_consumer(gathered->id)) { + return false; + } + const GraphNode * reshape = graph.index().consumers(gathered->id).front(); + if (reshape == nullptr || reshape->op != GGML_OP_RESHAPE || reshape->inputs.size() != 1 || + !graph.index().has_single_consumer(reshape->output)) { + return false; + } + const GraphNode * gdn = graph.index().consumers(reshape->output).front(); + if (gdn == nullptr || gdn->op != GGML_OP_GATED_DELTA_NET || gdn->inputs.size() != 6 || + gdn->inputs[5] != reshape->output) { + return false; + } + const GraphNode * q_norm = producer_with_op(graph, gdn->inputs[0], GGML_OP_L2_NORM); + const GatedDeltaNetMatch match = match_gated_delta_net(graph, q_norm); + if (!match.matched() || !match.has_projection_epilogue() || match.sequence_count != 1 || + match.token_count < 1 || match.token_count > 5 || match.snapshot_count < match.token_count || + (match.token_count == 1 && match.snapshot_count == 1) || + cache->ne[0] != match.width * match.width * match.head_count || + !same_full_value_range(*gathered, *match.state) || match.gdn_output->kind != ValueKind::Transient) { + return false; + } + const size_t state_bytes = static_cast(cache->ne[0]) * sizeof(float); + if (match.cache->nb[2] % state_bytes != 0) { + return false; + } + if (cache->storage == match.cache->storage && + (match.cache->storage_offset < cache->storage_offset || + (match.cache->storage_offset - cache->storage_offset) % state_bytes != 0 || + match.cache->storage_offset - cache->storage_offset > cache->byte_count || + match.cache->byte_count > cache->byte_count - (match.cache->storage_offset - cache->storage_offset))) { + return false; + } + const auto internal = [&](const GraphNode * node) { + return std::find(match.covered.begin(), match.covered.end(), node) != match.covered.end(); + }; + const std::array inputs = { + match.raw_q, match.raw_k, match.v, match.alpha_raw, match.beta_raw, match.bias, match.a_scale, cache, ids + }; + for (const Value * input : inputs) { + const GraphNode * producer = graph.index().producer(input->storage_root); + size_t index = 0; + if (producer != nullptr && (!graph.index().node_index(producer, index) || + index >= context.covered_nodes.size() || !context.covered_nodes[index])) { + return false; + } + if (!distinct_storage(*input, *match.gdn_output) || + (input != cache && !distinct_storage(*input, *match.cache))) { + return false; + } + } + size_t copy_index = 0; + for (const GraphNode * node : match.covered) { + const Value * output = graph_value(graph, node->output); + if (node->op == GGML_OP_CPY) { + if (!graph.index().node_index(node, copy_index)) { + return false; + } + } + if (output->id == match.cache->id || output->id == match.attention->id) { + continue; + } + if (!is_layout_alias_node(graph, *node) && output->kind != ValueKind::Transient) { + return false; + } + for (const GraphNode * consumer : graph.index().consumers(output->id)) { + if (!internal(consumer)) { + return false; + } + } + } + if (copy_index <= context.root_index) { + return false; + } + // The fused dispatch publishes snapshots at the gather's position. + for (size_t index = context.root_index + 1; index < copy_index; ++index) { + const GraphNode & node = graph.nodes()[index]; + const Value * output = graph_value(graph, node.output); + if (internal(&node) || is_layout_alias_node(graph, node) || output->byte_count == 0) { + continue; + } + const auto touches_cache = [&](const Value * value) { + return value != nullptr && value->byte_count != 0 && value->storage == match.cache->storage; + }; + if (touches_cache(output)) { + return false; + } + for (ValueId input : node.inputs) { + if (touches_cache(graph_value(graph, input))) { + return false; + } + } + } + DispatchMatchContext gdn_context = context; + gdn_context.root_node = q_norm; + if (!graph.index().node_index(q_norm, gdn_context.root_index)) { + return false; + } + DispatchMatch selected; + DispatchMatch rms_match; + const bool fuse_rms_q8 = match_gated_delta_net_rmsnorm_q8(context, match, *cache, rms_match); + if ((match.token_count == 1 && !fuse_rms_q8) || + !match_gated_delta_net_dispatch(gdn_context, selected) || + selected.dispatches.size() != (match.token_count == 1 ? 2 : 1) || + selected.dispatches.back().bindings.size() != (match.token_count == 1 ? 9 : 10) || + !append_covered_node(context, gather, selected)) { + return false; + } + if (fuse_rms_q8) { + Dispatch dispatch = std::move(selected.dispatches.back()); + const Dispatch & norm = rms_match.dispatches.front(); + const auto snapshot = dispatch.bindings[match.token_count == 1 ? 7 : 8]; + auto kernel = make_kernel_specialization(kGatedDeltaNetSelectedRmsQ8Kernel); + kernel.compile_parameters = std::move(dispatch.kernel.compile_parameters); + set_compile_parameter(kernel, "llm.gated_delta_net.state_row_count", cache->ne[1]); + set_compile_parameter(kernel, "llm.gated_delta_net.snapshot_stride", match.cache->nb[2] / sizeof(float)); + for (const auto & parameter : norm.kernel.compile_parameters) { + kernel.compile_parameters.emplace(parameter); + } + dispatch.kernel = std::move(kernel); + dispatch.bindings.resize(7); + dispatch.bindings.push_back({ cache->id, 0, cache->byte_count }); + dispatch.bindings.push_back(snapshot); + dispatch.bindings.push_back(norm.bindings[3]); + dispatch.bindings.push_back({ ids->id, 0, ids->byte_count }); + dispatch.bindings.push_back(norm.bindings[1]); + dispatch.bindings.push_back(norm.bindings[2]); + dispatch.bindings.push_back(norm.bindings[4]); + selected.dispatches.clear(); + selected.dispatches.push_back(std::move(dispatch)); + selected.covered_nodes.insert(selected.covered_nodes.end(), rms_match.covered_nodes.begin(), rms_match.covered_nodes.end()); + selected.transients = std::move(rms_match.transients); + if (!selected.metadata.append(std::move(rms_match.metadata), selected.status)) { + return false; + } + result = std::move(selected); + return true; + } + Dispatch & dispatch = selected.dispatches.front(); + auto kernel = make_kernel_specialization(kGatedDeltaNetSelectedSnapshotKernel); + kernel.compile_parameters = std::move(dispatch.kernel.compile_parameters); + dispatch.kernel = std::move(kernel); + set_compile_parameter(dispatch.kernel, "llm.gated_delta_net.state_row_count", cache->ne[1]); + dispatch.bindings[7] = { cache->id, 0, cache->byte_count }; + dispatch.bindings.push_back({ ids->id, 0, ids->byte_count }); + result = std::move(selected); + return true; +} +} // namespace + +void register_llm_gated_delta_net_dispatch(DispatchRegistryBuilder & registry) { + registry.add({ + "llm.gated_delta_net.selected_snapshot.f32_wmma_head128", + GGML_OP_GET_ROWS, + DispatchMatchKind::Fused, + 250, + DispatchSource::Llm, + match_selected_gated_delta_net_dispatch, + }); + registry.add({ + "llm.gated_delta_net.projection_pair.q4_k", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 300, + DispatchSource::Llm, + match_gated_delta_net_projection_pair_dispatch, + }); + registry.add({ + "llm.gated_delta_net.direct.f32_wmma_head128", + GGML_OP_GATED_DELTA_NET, + DispatchMatchKind::Fused, + 200, + DispatchSource::Llm, + match_gated_delta_net_dispatch, + }); + registry.add({ + "llm.gated_delta_net.f32_wmma_head128", + GGML_OP_L2_NORM, + DispatchMatchKind::Fused, + 200, + DispatchSource::Llm, + match_gated_delta_net_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-gated-delta-net.h b/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-gated-delta-net.h new file mode 100644 index 000000000000..a02c2d2c1489 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-gated-delta-net.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_llm_gated_delta_net_dispatch(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-ssm-conv.cpp b/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-ssm-conv.cpp new file mode 100644 index 000000000000..15ea9323451f --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-ssm-conv.cpp @@ -0,0 +1,726 @@ +#include "dispatch-ssm-conv.h" + +#include "../common/dispatch-mul-mat-common.h" + +#include "ggml.h" +#include "graph/graph-matcher.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kMulMatConv4Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_quantized_f16_wmma_prefill_conv4"); +static constexpr KernelCatalogRef kMulMatConv4InteriorKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_quantized_f16_wmma_prefill_conv4_interior"); +static constexpr KernelCatalogRef kSsmConvFinishKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_ssm_conv_dconv4_silu_prefill_finish_f32"); +static constexpr KernelCatalogRef kSsmConvSnapshotKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_ssm_conv_snapshot_window_tail_f32"); +static constexpr KernelCatalogRef kSsmConvPrefillKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_ssm_conv_dconv4_silu_prefill_512_wg1024"); +static constexpr KernelCatalogRef kSsmConvDecodeKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_ssm_conv_dconv4_silu_decode_f32"); +static constexpr KernelCatalogRef kSsmConvRollbackKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_ssm_conv_dconv4_silu_rollback_f32"); +static constexpr KernelCatalogRef kSsmConvGenericKernel = GGML_HRX_KERNEL_REF("loom_libs", "llm_ssm_conv_f32"); +static constexpr KernelCatalogRef kSsmConvGenericBinaryKernel = + GGML_HRX_KERNEL_REF("loom_libs", "llm_ssm_conv_binary_f32"); + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static bool is_f32(const Value * value) { + return value != nullptr && value->type == GGML_TYPE_F32; +} + +static bool is_shape(const Value & value, int64_t ne0, int64_t ne1, int64_t ne2, int64_t ne3) { + return value.ne[0] == ne0 && value.ne[1] == ne1 && value.ne[2] == ne2 && value.ne[3] == ne3; +} + +static bool same_shape(const Value & lhs, const Value & rhs) { + for (int dim = 0; dim < GGML_MAX_DIMS; ++dim) { + if (lhs.ne[dim] != rhs.ne[dim]) { + return false; + } + } + return true; +} + +static bool packed_f32_layout(const Value & value) { + size_t expected_stride = sizeof(float); + for (int dim = 0; dim < GGML_MAX_DIMS; ++dim) { + if (value.ne[dim] <= 0 || value.nb[dim] != expected_stride) { + return false; + } + expected_stride *= static_cast(value.ne[dim]); + } + return true; +} + +static bool supported_hidden_size(int64_t hidden_size) { + return hidden_size >= 32 && hidden_size <= 65536 && hidden_size % 32 == 0; +} + +static const Value * root_alias_value(const Graph & graph, const Value * value) { + for (size_t depth = 0; value != nullptr && value->alias_source.value >= 0 && depth < graph.values().size(); + ++depth) { + value = graph_value(graph, value->alias_source); + } + return value; +} + +static bool distinct_storage(const Value & lhs, const Value & rhs) { + return lhs.storage != rhs.storage; +} + +static bool ranges_overlap(const Value & lhs, const Value & rhs) { + if (lhs.storage != rhs.storage || lhs.byte_count == 0 || rhs.byte_count == 0) { + return false; + } + if (lhs.storage_offset <= rhs.storage_offset) { + return rhs.storage_offset - lhs.storage_offset < lhs.byte_count; + } + return lhs.storage_offset - rhs.storage_offset < rhs.byte_count; +} + +static bool append_covered_node(const DispatchMatchContext & context, const GraphNode * node, DispatchMatch & match) { + return append_covered_node_index_once(context.graph, context.covered_nodes, node, match.covered_nodes); +} + +static void set_compile_parameter(KernelSpecialization & kernel, const char * name, int64_t value) { + kernel.compile_parameters.emplace(name, std::to_string(value)); +} + +struct SsmConvCacheUpdate { + const GraphNode * state_tail_view = nullptr; + const GraphNode * cache_view = nullptr; + const GraphNode * cache_copy = nullptr; + const Value * state_tail = nullptr; + const Value * cache = nullptr; + int64_t source_row = 0; +}; + +struct SsmConvPrefillMatch { + const GraphNode * concat = nullptr; + const GraphNode * ssm = nullptr; + const GraphNode * silu = nullptr; + const Value * state = nullptr; + const Value * x = nullptr; + const Value * filter = nullptr; + const Value * output = nullptr; + std::vector cache_updates; + int64_t hidden_size = 0; + int64_t token_count = 0; + int64_t sequence_count = 0; + + bool matched() const { + return concat != nullptr && ssm != nullptr && silu != nullptr && state != nullptr && x != nullptr && + filter != nullptr && output != nullptr && !cache_updates.empty(); + } +}; + +struct SsmConvCoreMatch { + const GraphNode * ssm = nullptr; + const Value * window = nullptr; + const Value * filter = nullptr; + const Value * conv_output = nullptr; + const GraphNode * unary = nullptr; + const GraphNode * binary = nullptr; + const Value * binary_operand = nullptr; + const Value * output = nullptr; + int64_t d_conv = 0; + int64_t d_inner = 0; + int64_t token_count = 0; + int64_t sequence_count = 0; + UnaryKind unary_op = UnaryKind::Identity; + BinaryKind binary_op = BinaryKind::Mul; + bool binary_lhs = true; + + bool matched() const { + return ssm != nullptr && window != nullptr && filter != nullptr && conv_output != nullptr && output != nullptr; + } + + bool has_unary_fusion() const { return unary != nullptr; } + + bool has_binary_fusion() const { return binary != nullptr; } +}; + +static SsmConvPrefillMatch match_ssm_conv_prefill(const Graph & graph, const GraphNode * node) { + SsmConvPrefillMatch match; + if (node == nullptr || node->op != GGML_OP_CONCAT || node->inputs.size() != 2 || !graph.has_index()) { + return match; + } + + const Value * state = graph_value(graph, node->inputs[0]); + const Value * x_transposed = graph_value(graph, node->inputs[1]); + const Value * window = graph_value(graph, node->output); + if (!is_f32(state) || !is_f32(x_transposed) || !is_f32(window)) { + return {}; + } + + const GraphNode * transpose = graph.index().producer(x_transposed->id); + if (transpose == nullptr || transpose->op != GGML_OP_TRANSPOSE || transpose->inputs.size() != 1) { + return {}; + } + const Value * x_layout = graph_value(graph, transpose->inputs[0]); + const Value * x = root_alias_value(graph, x_layout); + if (!is_f32(x_layout) || !x_layout->contiguous || !is_f32(x) || !x->contiguous) { + return {}; + } + + const int64_t hidden_size = x_layout->ne[0]; + const int64_t token_count = x_layout->ne[1]; + const int64_t sequence_count = x_layout->ne[2]; + if (!supported_hidden_size(hidden_size) || token_count < 1 || token_count > 512 || sequence_count < 1 || + sequence_count > 3 || !is_shape(*x_layout, hidden_size, token_count, sequence_count, 1) || + !is_shape(*x_transposed, token_count, hidden_size, sequence_count, 1) || + !is_shape(*state, 3, hidden_size, sequence_count, 1) || + !is_shape(*window, token_count + 3, hidden_size, sequence_count, 1)) { + return {}; + } + if (x_layout->nb[0] != sizeof(float) || x_layout->nb[1] != static_cast(hidden_size) * sizeof(float) || + x_layout->nb[2] != static_cast(hidden_size * token_count) * sizeof(float) || + x_transposed->nb[0] != static_cast(hidden_size) * sizeof(float) || + x_transposed->nb[1] != sizeof(float) || + x_transposed->nb[2] != static_cast(hidden_size * token_count) * sizeof(float) || + state->nb[0] != sizeof(float) || state->nb[1] != 3 * sizeof(float) || + state->nb[2] != static_cast(3 * hidden_size) * sizeof(float) || window->nb[0] != sizeof(float) || + window->nb[1] != static_cast(token_count + 3) * sizeof(float) || + window->nb[2] != static_cast((token_count + 3) * hidden_size) * sizeof(float)) { + return {}; + } + + const std::vector & window_consumers = graph.index().consumers(window->id); + if (window_consumers.size() < 2 || window_consumers.size() > 6) { + return {}; + } + std::vector state_tail_views; + const GraphNode * ssm = nullptr; + for (const GraphNode * consumer : window_consumers) { + if (consumer != nullptr && consumer->op == GGML_OP_VIEW) { + state_tail_views.push_back(consumer); + } else if (consumer != nullptr && consumer->op == GGML_OP_SSM_CONV) { + if (ssm != nullptr) { + return {}; + } + ssm = consumer; + } else { + return {}; + } + } + if (state_tail_views.empty() || state_tail_views.size() > 5 || ssm == nullptr || ssm->inputs.size() != 2 || + ssm->inputs[0] != window->id) { + return {}; + } + + const Value * filter = graph_value(graph, ssm->inputs[1]); + const Value * ssm_output = graph_value(graph, ssm->output); + if (!is_f32(filter) || !is_f32(ssm_output) || !filter->contiguous || !ssm_output->contiguous || + !is_shape(*filter, 4, hidden_size, 1, 1) || + !is_shape(*ssm_output, hidden_size, token_count, sequence_count, 1) || filter->nb[0] != sizeof(float) || + filter->nb[1] != 4 * sizeof(float)) { + return {}; + } + + std::vector cache_updates; + for (const GraphNode * state_tail_view : state_tail_views) { + if (state_tail_view == nullptr || state_tail_view->inputs.size() != 1) { + return {}; + } + const Value * state_tail = graph_value(graph, state_tail_view->output); + if (!is_f32(state_tail) || state_tail->alias_source != window->id || + state_tail->storage_offset < window->storage_offset || + (state_tail->storage_offset - window->storage_offset) % sizeof(float) != 0 || + !is_shape(*state_tail, 3, hidden_size, sequence_count, 1)) { + return {}; + } + const int64_t source_row = + static_cast((state_tail->storage_offset - window->storage_offset) / sizeof(float)); + if (source_row < 0 || source_row > token_count) { + return {}; + } + const std::vector & tail_consumers = graph.index().consumers(state_tail->id); + if (tail_consumers.size() != 1 || tail_consumers.front() == nullptr || + tail_consumers.front()->op != GGML_OP_CPY || tail_consumers.front()->inputs.size() != 2 || + tail_consumers.front()->inputs[0] != state_tail->id) { + return {}; + } + const GraphNode * cache_copy = tail_consumers.front(); + const Value * cache_target = graph_value(graph, cache_copy->inputs[1]); + const Value * cache = graph_value(graph, cache_copy->output); + const GraphNode * cache_view = cache_target != nullptr ? graph.index().producer(cache_target->id) : nullptr; + if (!is_f32(cache_target) || !is_f32(cache) || cache_view == nullptr || cache_view->op != GGML_OP_VIEW || + cache_view->inputs.size() != 1 || !same_full_value_range(*cache_target, *cache) || + cache->byte_count != static_cast(3 * hidden_size * sequence_count) * sizeof(float)) { + return {}; + } + cache_updates.push_back({ state_tail_view, cache_view, cache_copy, state_tail, cache, source_row }); + } + + std::sort(cache_updates.begin(), cache_updates.end(), + [](const auto & lhs, const auto & rhs) { return lhs.cache->storage_offset < rhs.cache->storage_offset; }); + for (size_t slot = 0; slot < cache_updates.size(); ++slot) { + const int64_t expected_source_row = std::max(0, token_count - static_cast(slot)); + if (cache_updates[slot].source_row != expected_source_row || + cache_updates[slot].cache->storage != cache_updates.front().cache->storage) { + return {}; + } + for (size_t prior = 0; prior < slot; ++prior) { + if (ranges_overlap(*cache_updates[slot].cache, *cache_updates[prior].cache)) { + return {}; + } + } + } + + const std::vector & ssm_consumers = graph.index().consumers(ssm_output->id); + if (ssm_consumers.size() != 1 || ssm_consumers.front() == nullptr || ssm_consumers.front()->op != GGML_OP_UNARY || + ssm_consumers.front()->inputs.size() != 1) { + return {}; + } + const GraphNode * silu = ssm_consumers.front(); + const UnaryParams * unary = op_params_as(silu->params); + const Value * output = graph_value(graph, silu->output); + if (unary == nullptr || unary->op != UnaryKind::Silu || !is_f32(output) || !output->contiguous || + !is_shape(*output, hidden_size, token_count, sequence_count, 1)) { + return {}; + } + + if (!distinct_storage(*state, *x) || !distinct_storage(*state, *filter) || !distinct_storage(*x, *filter) || + (!distinct_storage(*output, *x) && !same_full_value_range(*output, *x)) || !distinct_storage(*output, *state) || + !distinct_storage(*output, *filter)) { + return {}; + } + for (const SsmConvCacheUpdate & update : cache_updates) { + if (!distinct_storage(*state, *update.cache) || !distinct_storage(*x, *update.cache) || + !distinct_storage(*filter, *update.cache) || !distinct_storage(*output, *update.cache)) { + return {}; + } + } + + match.concat = node; + match.ssm = ssm; + match.silu = silu; + match.state = state; + match.x = x; + match.filter = filter; + match.output = output; + match.cache_updates = std::move(cache_updates); + match.hidden_size = hidden_size; + match.token_count = token_count; + match.sequence_count = sequence_count; + return match; +} + +static SsmConvCoreMatch match_ssm_conv_core(const Graph & graph, const GraphNode * node) { + SsmConvCoreMatch match; + if (node == nullptr || node->op != GGML_OP_SSM_CONV || node->inputs.size() != 2) { + return match; + } + + const Value * window = graph_value(graph, node->inputs[0]); + const Value * filter = graph_value(graph, node->inputs[1]); + const Value * conv_output = graph_value(graph, node->output); + if (!is_f32(window) || !is_f32(filter) || !is_f32(conv_output) || !packed_f32_layout(*window) || + !packed_f32_layout(*filter) || !packed_f32_layout(*conv_output)) { + return {}; + } + + const int64_t d_conv = filter->ne[0]; + const int64_t d_inner = filter->ne[1]; + const int64_t token_count = window->ne[0] - d_conv + 1; + const int64_t sequence_count = window->ne[2]; + if (d_conv < 1 || d_conv > 16 || !supported_hidden_size(d_inner) || token_count < 1 || token_count > 512 || + sequence_count < 1 || sequence_count > 4 || window->ne[1] != d_inner || window->ne[3] != 1 || + filter->ne[2] != 1 || filter->ne[3] != 1 || !is_shape(*conv_output, d_inner, token_count, sequence_count, 1) || + ranges_overlap(*window, *conv_output) || ranges_overlap(*filter, *conv_output)) { + return {}; + } + + match.ssm = node; + match.window = window; + match.filter = filter; + match.conv_output = conv_output; + match.output = conv_output; + match.d_conv = d_conv; + match.d_inner = d_inner; + match.token_count = token_count; + match.sequence_count = sequence_count; + return match; +} + +static bool try_match_unary_fusion(const Graph & graph, SsmConvCoreMatch & match) { + if (!graph.has_index()) { + return false; + } + const std::vector & consumers = graph.index().consumers(match.conv_output->id); + if (consumers.size() != 1 || consumers.front() == nullptr || consumers.front()->op != GGML_OP_UNARY || + consumers.front()->inputs.size() != 1) { + return false; + } + + const GraphNode * unary = consumers.front(); + const UnaryParams * params = op_params_as(unary->params); + const Value * output = graph_value(graph, unary->output); + if (params == nullptr || !unary_kind_supported(params->op) || output == nullptr || output->type != GGML_TYPE_F32 || + !packed_f32_layout(*output) || !same_shape(*output, *match.conv_output) || + ranges_overlap(*match.window, *output) || ranges_overlap(*match.filter, *output)) { + return false; + } + + match.unary = unary; + match.output = output; + match.unary_op = params->op; + return true; +} + +static bool try_match_binary_fusion(const Graph & graph, SsmConvCoreMatch & match) { + if (!graph.has_index()) { + return false; + } + const std::vector & consumers = graph.index().consumers(match.conv_output->id); + if (consumers.size() != 1 || consumers.front() == nullptr || consumers.front()->inputs.size() != 2) { + return false; + } + + const GraphNode * binary = consumers.front(); + const BinaryParams * params = op_params_as(binary->params); + if (params == nullptr || params->op != BinaryKind::Mul) { + return false; + } + + const bool conv_is_lhs = binary->inputs[0] == match.conv_output->id; + const bool conv_is_rhs = binary->inputs[1] == match.conv_output->id; + if (conv_is_lhs == conv_is_rhs) { + return false; + } + + const Value * operand = graph_value(graph, conv_is_lhs ? binary->inputs[1] : binary->inputs[0]); + const Value * output = graph_value(graph, binary->output); + if (!is_f32(operand) || output == nullptr || output->type != GGML_TYPE_F32 || !packed_f32_layout(*operand) || + !packed_f32_layout(*output) || !same_shape(*operand, *match.conv_output) || + !same_shape(*output, *match.conv_output) || ranges_overlap(*match.window, *output) || + ranges_overlap(*match.filter, *output) || ranges_overlap(*operand, *output)) { + return false; + } + + match.binary = binary; + match.binary_operand = operand; + match.output = output; + match.binary_op = params->op; + match.binary_lhs = conv_is_lhs; + return true; +} + +static void set_generic_ssm_conv_parameters(KernelSpecialization & kernel, const SsmConvCoreMatch & match) { + set_compile_parameter(kernel, "llm.ssm_conv.generic.d_conv", match.d_conv); + set_compile_parameter(kernel, "llm.ssm_conv.generic.d_inner", match.d_inner); + set_compile_parameter(kernel, "llm.ssm_conv.generic.n_t", match.token_count); + set_compile_parameter(kernel, "llm.ssm_conv.generic.n_s", match.sequence_count); + set_compile_parameter(kernel, "llm.ssm_conv.generic.unary_op", unary_kind_config_value(match.unary_op)); + set_compile_parameter(kernel, "llm.ssm_conv.generic.workgroup_size", 256); +} + +static bool append_ssm_conv_covered_nodes(const DispatchMatchContext & context, + const SsmConvPrefillMatch & match, + DispatchMatch & dispatch_match) { + if (!append_covered_node(context, match.concat, dispatch_match) || + !append_covered_node(context, match.ssm, dispatch_match) || + !append_covered_node(context, match.silu, dispatch_match)) { + return false; + } + for (const SsmConvCacheUpdate & update : match.cache_updates) { + if (!append_covered_node(context, update.state_tail_view, dispatch_match) || + !append_covered_node(context, update.cache_view, dispatch_match) || + !append_covered_node(context, update.cache_copy, dispatch_match)) { + return false; + } + } + return true; +} + +static bool match_ssm_conv_prefill_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const SsmConvPrefillMatch match = match_ssm_conv_prefill(context.graph, context.root_node); + if (!match.matched()) { + return false; + } + + if (!append_ssm_conv_covered_nodes(context, match, dispatch_match)) { + return false; + } + + const CommandPlanGeneratedResource * edges = + context.plan.metadata.find_generated_resource(match.x->id, GeneratedResourceRole::Conv4Edges); + if (edges != nullptr) { + Dispatch finish; + finish.kernel = make_kernel_specialization(kSsmConvFinishKernel); + set_compile_parameter(finish.kernel, "llm.ssm_conv.generic.d_inner", match.hidden_size); + finish.bindings.push_back({ match.state->id, 0, match.state->byte_count }); + finish.bindings.push_back({ match.filter->id, 0, match.filter->byte_count }); + finish.bindings.push_back({ edges->generated_value, 0, edges->byte_count }); + finish.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + const Value * cache = match.cache_updates.front().cache; + finish.bindings.push_back({ cache->id, 0, cache->byte_count }); + dispatch_match.dispatches.push_back(std::move(finish)); + return true; + } + + const bool optimized_prefill = + match.token_count == 512 && match.sequence_count == 1 && match.cache_updates.size() == 1 && + match.hidden_size >= 8192 && match.hidden_size <= 10240; + if (!optimized_prefill) { + Dispatch ssm_dispatch; + const bool decode = match.token_count == 1 && match.sequence_count == 1 && match.cache_updates.size() == 1; + if (decode) { + ssm_dispatch.kernel = make_kernel_specialization(kSsmConvDecodeKernel); + set_compile_parameter(ssm_dispatch.kernel, "llm.ssm_conv.decode.d_inner", match.hidden_size); + set_compile_parameter(ssm_dispatch.kernel, "llm.ssm_conv.decode.workgroup_size", 256); + } else { + ssm_dispatch.kernel = make_kernel_specialization(kSsmConvRollbackKernel); + set_compile_parameter(ssm_dispatch.kernel, "llm.ssm_conv.rollback.d_inner", match.hidden_size); + set_compile_parameter(ssm_dispatch.kernel, "llm.ssm_conv.rollback.n_t", match.token_count); + set_compile_parameter(ssm_dispatch.kernel, "llm.ssm_conv.rollback.n_s", match.sequence_count); + set_compile_parameter(ssm_dispatch.kernel, "llm.ssm_conv.rollback.cache_count", match.cache_updates.size()); + set_compile_parameter(ssm_dispatch.kernel, "llm.ssm_conv.rollback.workgroup_size", 256); + } + ssm_dispatch.bindings.push_back({ match.state->id, 0, match.state->byte_count }); + ssm_dispatch.bindings.push_back({ match.x->id, 0, match.x->byte_count }); + ssm_dispatch.bindings.push_back({ match.filter->id, 0, match.filter->byte_count }); + ssm_dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + if (decode) { + const Value * cache = match.cache_updates.front().cache; + ssm_dispatch.bindings.push_back({ cache->id, 0, cache->byte_count }); + } else { + for (size_t slot = 0; slot < 5; ++slot) { + const Value * cache = match.cache_updates[std::min(slot, match.cache_updates.size() - 1)].cache; + ssm_dispatch.bindings.push_back({ cache->id, 0, cache->byte_count }); + } + } + dispatch_match.dispatches.push_back(std::move(ssm_dispatch)); + return true; + } + + const size_t snapshot_bytes = static_cast(64 * match.hidden_size) * sizeof(float); + const ValueId snapshot = context.next_plan_value; + dispatch_match.transients.push_back({ snapshot, "llm.ssm_conv.window_snapshot", snapshot_bytes, 256 }); + + Dispatch snapshot_dispatch; + snapshot_dispatch.kernel = make_kernel_specialization(kSsmConvSnapshotKernel); + set_compile_parameter(snapshot_dispatch.kernel, "llm.ssm_conv.snapshot.d_conv", 4); + set_compile_parameter(snapshot_dispatch.kernel, "llm.ssm_conv.snapshot.d_inner", match.hidden_size); + set_compile_parameter(snapshot_dispatch.kernel, "llm.ssm_conv.snapshot.n_t", match.token_count); + set_compile_parameter(snapshot_dispatch.kernel, "llm.ssm_conv.snapshot.n_s", match.sequence_count); + set_compile_parameter(snapshot_dispatch.kernel, "llm.ssm_conv.snapshot.state_row_stride", 1); + set_compile_parameter(snapshot_dispatch.kernel, "llm.ssm_conv.snapshot.state_channel_stride", 3); + set_compile_parameter(snapshot_dispatch.kernel, "llm.ssm_conv.snapshot.x_row_stride", match.hidden_size); + set_compile_parameter(snapshot_dispatch.kernel, "llm.ssm_conv.snapshot.dst_row_stride", match.hidden_size); + set_compile_parameter(snapshot_dispatch.kernel, "llm.ssm_conv.snapshot.cache_row_stride", 1); + set_compile_parameter(snapshot_dispatch.kernel, "llm.ssm_conv.snapshot.cache_channel_stride", 3); + set_compile_parameter(snapshot_dispatch.kernel, "llm.ssm_conv.snapshot.workgroup_size", 256); + snapshot_dispatch.bindings.push_back({ match.state->id, 0, match.state->byte_count }); + snapshot_dispatch.bindings.push_back({ match.x->id, 0, match.x->byte_count }); + snapshot_dispatch.bindings.push_back({ snapshot, 0, snapshot_bytes }); + const Value * cache = match.cache_updates.front().cache; + snapshot_dispatch.bindings.push_back({ cache->id, 0, cache->byte_count }); + + Dispatch ssm_dispatch; + ssm_dispatch.kernel = make_kernel_specialization(kSsmConvPrefillKernel); + set_compile_parameter(ssm_dispatch.kernel, "llm.ssm_conv.prefill.d_conv", 4); + set_compile_parameter(ssm_dispatch.kernel, "llm.ssm_conv.prefill.d_inner", match.hidden_size); + set_compile_parameter(ssm_dispatch.kernel, "llm.ssm_conv.prefill.n_t", match.token_count); + set_compile_parameter(ssm_dispatch.kernel, "llm.ssm_conv.prefill.n_s", match.sequence_count); + set_compile_parameter(ssm_dispatch.kernel, "llm.ssm_conv.prefill.state_row_stride", match.hidden_size); + set_compile_parameter(ssm_dispatch.kernel, "llm.ssm_conv.prefill.x_row_stride", match.hidden_size); + ssm_dispatch.bindings.push_back({ snapshot, 0, snapshot_bytes }); + ssm_dispatch.bindings.push_back({ match.x->id, 0, match.x->byte_count }); + ssm_dispatch.bindings.push_back({ match.filter->id, 0, match.filter->byte_count }); + ssm_dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + + dispatch_match.dispatches.push_back(std::move(snapshot_dispatch)); + dispatch_match.dispatches.push_back(std::move(ssm_dispatch)); + return true; +} + +static bool match_ssm_conv_generic_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + SsmConvCoreMatch match = match_ssm_conv_core(context.graph, context.root_node); + if (!match.matched()) { + return false; + } + + try_match_binary_fusion(context.graph, match) || try_match_unary_fusion(context.graph, match); + + if (!append_covered_node(context, match.ssm, dispatch_match)) { + return false; + } + if (match.has_binary_fusion() && !append_covered_node(context, match.binary, dispatch_match)) { + return false; + } + if (match.has_unary_fusion() && !append_covered_node(context, match.unary, dispatch_match)) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = + make_kernel_specialization(match.has_binary_fusion() ? kSsmConvGenericBinaryKernel : kSsmConvGenericKernel); + set_generic_ssm_conv_parameters(dispatch.kernel, match); + dispatch.bindings.push_back({ match.window->id, 0, match.window->byte_count }); + dispatch.bindings.push_back({ match.filter->id, 0, match.filter->byte_count }); + if (match.has_binary_fusion()) { + dispatch.kernel.compile_parameters.emplace("llm.ssm_conv.generic.binary_op", + std::to_string(binary_kind_config_value(match.binary_op))); + dispatch.kernel.compile_parameters.emplace("llm.ssm_conv.generic.binary_lhs", match.binary_lhs ? "1" : "0"); + dispatch.bindings.push_back({ match.binary_operand->id, 0, match.binary_operand->byte_count }); + } + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_mul_mat_conv4_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const CommonMulMatMatch mat = common_match_mul_mat_any_format( + context.graph, context.root_node, kMulMatConv4Kernel, false); + if (!mat.matched() || !context.graph.has_index() || mat.token_count != 512 || + (mat.weight->type != GGML_TYPE_Q4_K && mat.weight->type != GGML_TYPE_Q6_K) || + mat.weight->alias_source.value >= 0 || mat.input_size % 256 != 0 || + mat.output_size % 64 != 0 || mat.output_size / 64 < 32 || + !common_mul_mat_uses_k16_major_f16(mat.weight_format, mat.input_size, mat.output_size, mat.token_count)) { + return false; + } + + const Value * value = mat.output; + std::vector layouts; + const GraphNode * concat = nullptr; + while (value != nullptr && value->kind == ValueKind::Transient) { + const auto & consumers = context.graph.index().consumers(value->id); + if (consumers.size() != 1 || consumers.front() == nullptr) { + return false; + } + const GraphNode * consumer = consumers.front(); + if (consumer->op == GGML_OP_CONCAT) { + concat = consumer; + break; + } + if (!is_layout_alias_node(context.graph, *consumer)) { + return false; + } + layouts.push_back(consumer); + value = graph_value(context.graph, consumer->output); + } + const SsmConvPrefillMatch conv = match_ssm_conv_prefill(context.graph, concat); + if (!conv.matched() || conv.x->id != mat.output->id || conv.sequence_count != 1 || + conv.token_count != 512 || conv.cache_updates.size() != 1) { + return false; + } + const Value * window = graph_value(context.graph, conv.concat->output); + const Value * intermediate = graph_value(context.graph, conv.ssm->output); + if (window->kind != ValueKind::Transient || intermediate->kind != ValueKind::Transient) { + return false; + } + bool deferred_state = false; + for (const Value * input : {conv.state, conv.filter}) { + const GraphNode * producer = context.graph.index().producer(input->id); + size_t index = 0; + if (producer != nullptr && (!context.graph.index().node_index(producer, index) || !context.covered_nodes[index])) { + if (input == conv.filter) { + return false; + } + deferred_state = true; + } + } + if (deferred_state && (conv.hidden_size < 8192 || conv.hidden_size > 10240)) { + return false; + } + const Value * cache = conv.cache_updates.front().cache; + for (const Value * output : {conv.output, cache}) { + if (!distinct_storage(*output, *mat.input) || !distinct_storage(*output, *mat.weight)) { + return false; + } + } + if (!append_covered_node(context, context.root_node, dispatch_match) || + (!deferred_state && !append_ssm_conv_covered_nodes(context, conv, dispatch_match))) { + return false; + } + for (const GraphNode * layout : layouts) { + if (!append_covered_node(context, layout, dispatch_match)) { + return false; + } + } + + DispatchBinding activation; + if (!common_prepare_k16_major_f16_input(context, *mat.input, mat.input_size, mat.token_count, + dispatch_match, activation)) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(deferred_state ? kMulMatConv4InteriorKernel : kMulMatConv4Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", mat.token_count); + set_compile_parameter(dispatch.kernel, "ggml.mul_mat.input_size", mat.input_size); + set_compile_parameter(dispatch.kernel, "ggml.mul_mat.output_size", mat.output_size); + set_compile_parameter(dispatch.kernel, "ggml.mul_mat.weight_format", mat.weight->type == GGML_TYPE_Q4_K ? 4 : 6); + dispatch.bindings.push_back(activation); + if (mat.weight->type == GGML_TYPE_Q4_K) { + dispatch.bindings.push_back({mat.weight->id, 0, mat.weight->byte_count, kQ4KPackedK256Row64Layout, + mat.weight->type, mat.input_size, mat.output_size, mat.weight->byte_count}); + } else { + dispatch.bindings.push_back({mat.weight->id, 0, mat.weight->byte_count}); + } + if (!deferred_state) { + dispatch.bindings.push_back({conv.state->id, 0, conv.state->byte_count}); + } + dispatch.bindings.push_back({conv.filter->id, 0, conv.filter->byte_count}); + dispatch.bindings.push_back({conv.output->id, 0, conv.output->byte_count}); + if (deferred_state) { + const size_t edge_bytes = static_cast(6 * conv.hidden_size) * sizeof(float); + const ValueId edges(context.next_plan_value.value + static_cast(dispatch_match.transients.size())); + dispatch_match.transients.push_back({ edges, "llm.ssm_conv.projection_edges", edge_bytes, 256 }); + Status status; + if (!dispatch_match.metadata.append_generated_resource( + { mat.output->id, GeneratedResourceRole::Conv4Edges, edges, edge_bytes, {} }, status)) { + dispatch_match.status.append(status); + return false; + } + dispatch.bindings.push_back({ edges, 0, edge_bytes }); + } else { + dispatch.bindings.push_back({cache->id, 0, cache->byte_count}); + } + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_llm_ssm_conv_dispatch(DispatchRegistryBuilder & registry) { + registry.add({ + "llm.ssm_conv.quantized_prefill", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 306, + DispatchSource::Llm, + match_mul_mat_conv4_dispatch, + }); + registry.add({ + "llm.ssm_conv.dconv4_silu", + GGML_OP_CONCAT, + DispatchMatchKind::Fused, + 200, + DispatchSource::Llm, + match_ssm_conv_prefill_dispatch, + }); + registry.add({ + "llm.ssm_conv.generic_f32", + GGML_OP_SSM_CONV, + DispatchMatchKind::Fused, + 50, + DispatchSource::Llm, + match_ssm_conv_generic_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-ssm-conv.h b/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-ssm-conv.h new file mode 100644 index 000000000000..a17e37af73d7 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/llm/dispatch-ssm-conv.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_llm_ssm_conv_dispatch(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-llm-profiles.h b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-llm-profiles.h new file mode 100644 index 000000000000..ef9a06f9f8b0 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-llm-profiles.h @@ -0,0 +1,36 @@ +#pragma once + +#include + +namespace ggml::hrx { + +struct LlmMoeDispatchProfile { + const char * name = ""; + int64_t hidden_size = 0; + int64_t expert_hidden_size = 0; + int64_t expert_count = 0; + int64_t route_count = 0; + int64_t max_token_count = 0; + float rms_norm_epsilon = 0.0f; +}; + +constexpr LlmMoeDispatchProfile kLlmMoeQwen30BDispatchProfile = { + "qwen30b", 2048, 768, 128, 8, 2048, 0.000001f, +}; + +static constexpr const LlmMoeDispatchProfile & kActiveLlmMoeDispatchProfile = kLlmMoeQwen30BDispatchProfile; +static constexpr const LlmMoeDispatchProfile & kQwen30BMoeDispatchProfile = kLlmMoeQwen30BDispatchProfile; + +constexpr bool is_llm_supported_query_length(const LlmMoeDispatchProfile & profile, int64_t query_length) { + return query_length >= 1 && query_length <= profile.max_token_count; +} + +constexpr bool is_llm_decode_query_length(int64_t query_length) { + return query_length == 1; +} + +constexpr bool is_llm_prefill_query_length(const LlmMoeDispatchProfile & profile, int64_t query_length) { + return query_length > 1 && query_length <= profile.max_token_count; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-llm-shapes.h b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-llm-shapes.h new file mode 100644 index 000000000000..d24637a2c37b --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-llm-shapes.h @@ -0,0 +1,29 @@ +#pragma once + +#include "dispatch-llm-profiles.h" + +#include + +namespace ggml::hrx { + +constexpr bool is_llm_prefill_512_query_length(const LlmMoeDispatchProfile & profile, int64_t query_length) { + return is_llm_prefill_query_length(profile, query_length) && query_length == 512; +} + +constexpr bool is_qwen_supported_query_length(int64_t query_length) { + return is_llm_supported_query_length(kQwen30BMoeDispatchProfile, query_length); +} + +constexpr bool is_qwen_decode_query_length(int64_t query_length) { + return is_llm_decode_query_length(query_length); +} + +constexpr bool is_qwen_prefill_query_length(int64_t query_length) { + return is_llm_prefill_query_length(kQwen30BMoeDispatchProfile, query_length); +} + +constexpr bool is_qwen_prefill_512_query_length(int64_t query_length) { + return is_llm_prefill_512_query_length(kQwen30BMoeDispatchProfile, query_length); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-moe-router.cpp b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-moe-router.cpp new file mode 100644 index 000000000000..b79ea4bf1408 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-moe-router.cpp @@ -0,0 +1,723 @@ +#include "dispatch-moe-router.h" + +#include "../common/dispatch-moe-routing-layout.h" + +#include "dispatch-llm-shapes.h" +#include "ggml.h" +#include "graph/graph-matcher.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kQwenRouterTop8F32Kernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "qwen3_moe_router_top8_f32"); +static constexpr KernelCatalogRef kQwenRouterProjectionTop8FusedDecodeF32Kernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "qwen3_moe_router_projection_top8_fused_decode_f32"); +static constexpr KernelCatalogRef kQwenRouterProjectionF32FourRowWave32Kernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "qwen3_moe_router_projection_f32_four_row_wave32"); +static constexpr KernelCatalogRef kMoeBuildExpertTableKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_moe_build_expert_table"); +static constexpr KernelCatalogRef kMoeBuildExpertPartitionTableKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_moe_build_expert_partition_table"); +static constexpr KernelCatalogRef kQwenBuildExpertTablePartitionPrefill512Kernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "qwen3_moe_build_expert_table_partition_prefill_512"); + +static constexpr const LlmMoeDispatchProfile & kMoeRouterProfile = kActiveLlmMoeDispatchProfile; +static constexpr size_t kMoeRouterPlanTransientAlignment = 256; + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static bool nearly_equal(float lhs, float rhs) { + return std::fabs(lhs - rhs) <= 1.0e-12f; +} + +static const GraphNode * find_consumer_with_op(const Graph & graph, ValueId value, ggml_op op) { + const std::vector consumers = consumers_with_op_through_layout_aliases(graph, value, op); + return consumers.empty() ? nullptr : consumers.front(); +} + +static const GraphNode * find_single_consumer_with_op(const Graph & graph, ValueId value, ggml_op op) { + const std::vector consumers = consumers_with_op_through_layout_aliases(graph, value, op); + return consumers.size() == 1 ? consumers.front() : nullptr; +} + +static const GraphNode * find_consumer_with_op_and_input(const Graph & graph, + ValueId value, + ggml_op op, + ValueId input) { + for (const GraphNode * consumer : consumers_with_op_through_layout_aliases(graph, value, op)) { + if (consumer != nullptr && node_has_input_or_alias(graph, *consumer, input)) { + return consumer; + } + } + return nullptr; +} + +static bool is_shape(const Value & value, int64_t ne0, int64_t ne1, int64_t ne2, int64_t ne3) { + return value.ne[0] == ne0 && value.ne[1] == ne1 && value.ne[2] == ne2 && value.ne[3] == ne3; +} + +static bool is_2d(const Value & value) { + return value.ne[0] > 0 && value.ne[1] > 0 && value.ne[2] == 1 && value.ne[3] == 1; +} + +static bool same_shape(const Value & lhs, const Value & rhs) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (lhs.ne[i] != rhs.ne[i]) { + return false; + } + } + return true; +} + +static bool is_supported_expert_count(int64_t expert_count) { + return expert_count >= 32 && expert_count <= 512 && expert_count % 32 == 0; +} + +static bool is_supported_route_count(int64_t route_count, int64_t expert_count) { + return route_count >= 1 && route_count <= 32 && route_count <= expert_count; +} + +static bool is_supported_route_stride(int64_t route_stride, int64_t route_count, int64_t expert_count) { + return route_stride >= route_count && route_stride <= expert_count; +} + +static bool is_default_scale_softmax(const GraphNode & node) { + const SoftMaxParams * params = op_params_as(node.params); + return params != nullptr && nearly_equal(params->scale, 1.0f) && nearly_equal(params->max_bias, 0.0f); +} + +static bool is_descending_argsort(const GraphNode & node) { + const ArgsortParams * params = op_params_as(node.params); + return params != nullptr && params->order == GGML_SORT_ORDER_DESC; +} + +static bool is_topk_normalization_clamp(const GraphNode & node) { + const ClampParams * params = op_params_as(node.params); + return params != nullptr && params->min >= 0.0f && params->min <= 1.0e-4f && std::isinf(params->max) && + params->max > 0.0f; +} + +static std::string to_config_value(int64_t value) { + return std::to_string(value); +} + +static size_t expert_table_size(int64_t token_count, int64_t expert_count) { + return static_cast(expert_count + expert_count * token_count) * sizeof(int32_t); +} + +static size_t partition_table_size(int64_t token_count, int64_t route_count, int64_t expert_count) { + const int64_t assignment_count = token_count * route_count; + const int64_t assignment_partition_count = (assignment_count + 31) / 32; + return static_cast(1 + assignment_partition_count + expert_count) * sizeof(int32_t); +} + +static std::string value_summary(const Graph & graph, const Value * value) { + if (value == nullptr) { + return "missing"; + } + std::ostringstream stream; + stream << value->id.value << ":" << ggml_type_name(value->type) << "[" << value->ne[0] << "," << value->ne[1] << "," + << value->ne[2] << "," << value->ne[3] << "] nb=[" << value->nb[0] << "," << value->nb[1] << "," + << value->nb[2] << "," << value->nb[3] << "]"; + const GraphNode * producer = graph.index().producer(value->id); + if (producer != nullptr) { + stream << "<-" << ggml_op_name(producer->op); + } + return stream.str(); +} + +static bool is_moe_router_candidate_root(const Graph & graph, const GraphNode * softmax_node) { + if (softmax_node == nullptr || softmax_node->op != GGML_OP_SOFT_MAX || softmax_node->inputs.size() != 1 || + !graph.has_index() || !is_default_scale_softmax(*softmax_node)) { + return false; + } + const Value * logits = graph_value(graph, softmax_node->inputs[0]); + const Value * probs = graph_value(graph, softmax_node->output); + if (logits == nullptr || probs == nullptr || logits->type != GGML_TYPE_F32 || probs->type != GGML_TYPE_F32) { + return false; + } + return find_consumer_with_op(graph, softmax_node->output, GGML_OP_RESHAPE) != nullptr && + find_consumer_with_op(graph, softmax_node->output, GGML_OP_ARGSORT) != nullptr; +} + +static void log_router_reject(Status * status, + const Graph & graph, + const GraphNode * node, + const std::string & reason) { + if (status == nullptr || !is_moe_router_candidate_root(graph, node)) { + return; + } + const Value * logits = node == nullptr || node->inputs.empty() ? nullptr : graph_value(graph, node->inputs[0]); + const Value * probs = node == nullptr ? nullptr : graph_value(graph, node->output); + status->log("MoE router top-k matcher rejected node: %s logits=%s probs=%s", reason.c_str(), + value_summary(graph, logits).c_str(), value_summary(graph, probs).c_str()); +} + +static bool append_covered_node(const DispatchMatchContext & context, const GraphNode * node, DispatchMatch & match) { + return append_covered_node_index_once(context.graph, context.covered_nodes, node, match.covered_nodes); +} + +struct RouterTop8Match { + const Value * logits = nullptr; + const Value * route_ids = nullptr; + const Value * route_weights = nullptr; + int64_t token_count = 0; + int64_t expert_count = 0; + int64_t route_count = 0; + int64_t route_stride = 0; + + bool matched() const { + return logits != nullptr && route_ids != nullptr && route_weights != nullptr && token_count > 0 && + is_supported_expert_count(expert_count) && is_supported_route_count(route_count, expert_count) && + is_supported_route_stride(route_stride, route_count, expert_count); + } +}; + +static bool supports_fused_prefill_expert_table_partition(const RouterTop8Match & router_match) { + // Matches the reference prefill recipe gate; q=1 uses decode routing paths. + return is_llm_prefill_512_query_length(kMoeRouterProfile, router_match.token_count) && + router_match.route_count == kMoeRouterProfile.route_count && + router_match.route_stride == router_match.route_count && + router_match.expert_count == kMoeRouterProfile.expert_count; +} + +static RouterTop8Match match_moe_router_top8(const Graph & graph, const GraphNode * softmax_node, Status * status) { + RouterTop8Match match; + if (softmax_node == nullptr || softmax_node->op != GGML_OP_SOFT_MAX || softmax_node->inputs.size() != 1 || + !graph.has_index() || !is_default_scale_softmax(*softmax_node)) { + return match; + } + + const Value * logits = graph_value(graph, softmax_node->inputs[0]); + const Value * probs = graph_value(graph, softmax_node->output); + if (logits == nullptr || probs == nullptr || logits->type != GGML_TYPE_F32 || probs->type != GGML_TYPE_F32 || + !same_shape(*logits, *probs) || !logits->contiguous || !probs->contiguous) { + log_router_reject(status, graph, softmax_node, "logits/probs must be same contiguous F32 shape"); + return {}; + } + const int64_t expert_count = logits->ne[0]; + const int64_t token_count = logits->ne[1]; + if (!is_shape(*logits, expert_count, token_count, 1, 1) || !is_supported_expert_count(expert_count) || + !is_llm_supported_query_length(kMoeRouterProfile, token_count)) { + log_router_reject(status, graph, softmax_node, "unsupported logits expert/token shape"); + return {}; + } + + const GraphNode * probs_reshape = find_consumer_with_op(graph, softmax_node->output, GGML_OP_RESHAPE); + const GraphNode * argsort = find_consumer_with_op(graph, softmax_node->output, GGML_OP_ARGSORT); + if (probs_reshape == nullptr || argsort == nullptr || argsort->inputs.size() != 1 || + !is_descending_argsort(*argsort)) { + log_router_reject(status, graph, softmax_node, "missing probability reshape or descending argsort"); + return {}; + } + + const Value * probs_reshaped = graph_value(graph, probs_reshape->output); + const Value * argsort_output = graph_value(graph, argsort->output); + if (probs_reshaped == nullptr || argsort_output == nullptr || probs_reshaped->type != GGML_TYPE_F32 || + argsort_output->type != GGML_TYPE_I32 || !is_shape(*probs_reshaped, 1, expert_count, token_count, 1) || + !is_shape(*argsort_output, expert_count, token_count, 1, 1)) { + log_router_reject(status, graph, softmax_node, "probability reshape or argsort output shape is incompatible"); + return {}; + } + + const GraphNode * topk_view = find_consumer_with_op(graph, argsort->output, GGML_OP_VIEW); + if (topk_view == nullptr || topk_view->inputs.size() != 1) { + return {}; + } + const Value * route_ids = graph_value(graph, topk_view->output); + const int64_t route_count = route_ids == nullptr ? 0 : route_ids->ne[0]; + if (route_ids == nullptr || route_ids->type != GGML_TYPE_I32 || + !is_shape(*route_ids, route_count, token_count, 1, 1) || !is_supported_route_count(route_count, expert_count) || + route_ids->nb[0] != sizeof(int32_t) || route_ids->nb[1] % sizeof(int32_t) != 0) { + log_router_reject(status, graph, softmax_node, "top-k route id view shape or stride is incompatible"); + return {}; + } + const int64_t route_stride = static_cast(route_ids->nb[1] / sizeof(int32_t)); + if (!is_supported_route_stride(route_stride, route_count, expert_count)) { + log_router_reject(status, graph, softmax_node, "top-k route id stride is outside supported bounds"); + return {}; + } + + const GraphNode * get_rows = + find_consumer_with_op_and_input(graph, probs_reshape->output, GGML_OP_GET_ROWS, topk_view->output); + if (get_rows == nullptr || get_rows->inputs.size() != 2) { + log_router_reject(status, graph, softmax_node, "missing GET_ROWS from reshaped probabilities and top-k ids"); + return {}; + } + const Value * selected_weights = graph_value(graph, get_rows->output); + if (selected_weights == nullptr || selected_weights->type != GGML_TYPE_F32 || + !is_shape(*selected_weights, 1, route_count, token_count, 1)) { + log_router_reject(status, graph, softmax_node, "selected route weight shape is incompatible"); + return {}; + } + + const GraphNode * weights_reshape = find_consumer_with_op(graph, get_rows->output, GGML_OP_RESHAPE); + const Value * weights_flat = weights_reshape == nullptr ? nullptr : graph_value(graph, weights_reshape->output); + if (weights_flat == nullptr || weights_flat->type != GGML_TYPE_F32 || + !is_shape(*weights_flat, route_count, token_count, 1, 1)) { + log_router_reject(status, graph, softmax_node, "flattened route weight shape is incompatible"); + return {}; + } + + const GraphNode * sum_rows = find_consumer_with_op(graph, weights_reshape->output, GGML_OP_SUM_ROWS); + const Value * sum = sum_rows == nullptr ? nullptr : graph_value(graph, sum_rows->output); + if (sum == nullptr || sum->type != GGML_TYPE_F32 || !is_shape(*sum, 1, token_count, 1, 1)) { + log_router_reject(status, graph, softmax_node, "missing SUM_ROWS over selected route weights"); + return {}; + } + + const GraphNode * clamp = find_consumer_with_op(graph, sum_rows->output, GGML_OP_CLAMP); + const Value * clamped_sum = clamp == nullptr ? nullptr : graph_value(graph, clamp->output); + if (clamped_sum == nullptr || clamped_sum->type != GGML_TYPE_F32 || !is_shape(*clamped_sum, 1, token_count, 1, 1) || + !is_topk_normalization_clamp(*clamp)) { + log_router_reject(status, graph, softmax_node, "missing supported CLAMP on selected route weight sum"); + return {}; + } + + const GraphNode * div = find_consumer_with_op_and_input(graph, weights_reshape->output, GGML_OP_DIV, clamp->output); + const Value * normalized = div == nullptr ? nullptr : graph_value(graph, div->output); + if (normalized == nullptr || normalized->type != GGML_TYPE_F32 || + !is_shape(*normalized, route_count, token_count, 1, 1)) { + log_router_reject(status, graph, softmax_node, "missing DIV normalization for selected route weights"); + return {}; + } + + const GraphNode * output_reshape = find_consumer_with_op(graph, div->output, GGML_OP_RESHAPE); + const Value * route_weights = output_reshape == nullptr ? nullptr : graph_value(graph, output_reshape->output); + if (route_weights == nullptr || route_weights->type != GGML_TYPE_F32 || + !is_shape(*route_weights, 1, route_count, token_count, 1) || !route_weights->contiguous) { + log_router_reject(status, graph, softmax_node, "route weight output shape is incompatible"); + return {}; + } + + match.logits = logits; + match.route_ids = route_ids; + match.route_weights = route_weights; + match.token_count = token_count; + match.expert_count = expert_count; + match.route_count = route_count; + match.route_stride = route_stride; + return match; +} + +static bool append_moe_router_top8_coverage(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * softmax = context.root_node; + const GraphNode * probs_reshape = find_consumer_with_op(context.graph, softmax->output, GGML_OP_RESHAPE); + const GraphNode * argsort = find_consumer_with_op(context.graph, softmax->output, GGML_OP_ARGSORT); + const GraphNode * topk_view = + argsort == nullptr ? nullptr : find_consumer_with_op(context.graph, argsort->output, GGML_OP_VIEW); + const GraphNode * get_rows = + probs_reshape == nullptr || topk_view == nullptr ? + nullptr : + find_consumer_with_op_and_input(context.graph, probs_reshape->output, GGML_OP_GET_ROWS, topk_view->output); + const GraphNode * weights_reshape = + get_rows == nullptr ? nullptr : find_consumer_with_op(context.graph, get_rows->output, GGML_OP_RESHAPE); + const GraphNode * sum_rows = weights_reshape == nullptr ? + nullptr : + find_consumer_with_op(context.graph, weights_reshape->output, GGML_OP_SUM_ROWS); + const GraphNode * clamp = + sum_rows == nullptr ? nullptr : find_consumer_with_op(context.graph, sum_rows->output, GGML_OP_CLAMP); + const GraphNode * div = + weights_reshape == nullptr || clamp == nullptr ? + nullptr : + find_consumer_with_op_and_input(context.graph, weights_reshape->output, GGML_OP_DIV, clamp->output); + const GraphNode * output_reshape = + div == nullptr ? nullptr : find_consumer_with_op(context.graph, div->output, GGML_OP_RESHAPE); + + return append_covered_node(context, softmax, match) && append_covered_node(context, probs_reshape, match) && + append_covered_node(context, argsort, match) && append_covered_node(context, topk_view, match) && + append_covered_node(context, get_rows, match) && append_covered_node(context, weights_reshape, match) && + append_covered_node(context, sum_rows, match) && append_covered_node(context, clamp, match) && + append_covered_node(context, div, match) && append_covered_node(context, output_reshape, match); +} + +static bool append_moe_router_top8_coverage_from_softmax(const DispatchMatchContext & context, + const GraphNode * softmax, + DispatchMatch & match) { + if (softmax == nullptr) { + return false; + } + DispatchMatchContext softmax_context = context; + softmax_context.root_node = softmax; + if (!context.graph.index().node_index(softmax, softmax_context.root_index)) { + return false; + } + return append_moe_router_top8_coverage(softmax_context, match); +} + +static void add_routed_gate_up_compile_parameters(Dispatch & dispatch, const RouterTop8Match & router_match) { + dispatch.kernel.compile_parameters.emplace("qwen3_moe.routed_gate_up.expert_count", + to_config_value(router_match.expert_count)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.routed_gate_up.route_count", + to_config_value(router_match.route_count)); +} + +static void add_moe_routing_compile_parameters(Dispatch & dispatch, const RouterTop8Match & router_match) { + dispatch.kernel.compile_parameters.emplace("ggml.workload.token_capacity", + to_config_value(router_match.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.moe_routing.route_count", + to_config_value(router_match.route_count)); + dispatch.kernel.compile_parameters.emplace("ggml.moe_routing.expert_count", + to_config_value(router_match.expert_count)); + // the descriptor layout depends on the expert count (dispatch-moe-routing-layout.h) + const MoeRoutingDescriptorLayout layout = moe_router_descriptor_layout(router_match.expert_count); + dispatch.kernel.compile_parameters.emplace("ggml.moe_routing.descriptor_expert_mask", layout.expert_mask); + dispatch.kernel.compile_parameters.emplace("ggml.moe_routing.descriptor_partition_shift", layout.partition_shift); + dispatch.kernel.compile_parameters.emplace("ggml.moe_routing.descriptor_row_count_shift", layout.row_count_shift); + dispatch.kernel.compile_parameters.emplace("ggml.moe_routing.partition_workgroup_size", + layout.partition_workgroup_size); +} + +struct RouterProjectionTop8Match { + const GraphNode * projection = nullptr; + const GraphNode * softmax = nullptr; + const Value * input = nullptr; + const Value * weight = nullptr; + const Value * logits = nullptr; + RouterTop8Match top8; + + bool matched() const { + return projection != nullptr && softmax != nullptr && input != nullptr && weight != nullptr && + logits != nullptr && top8.matched(); + } +}; + +struct RouterProjectionMatch { + const Value * input = nullptr; + const Value * weight = nullptr; + const Value * output = nullptr; + int64_t token_count = 0; + + bool matched() const { return input != nullptr && weight != nullptr && output != nullptr && token_count > 0; } +}; + +static RouterProjectionTop8Match match_moe_router_projection_top8_decode(const DispatchMatchContext & context, + Status * status) { + RouterProjectionTop8Match match; + const GraphNode * projection = context.root_node; + if (projection == nullptr || projection->op != GGML_OP_MUL_MAT || projection->inputs.size() != 2 || + !context.graph.has_index()) { + return match; + } + + const Value * weight = graph_value(context.graph, projection->inputs[0]); + const Value * input = graph_value(context.graph, projection->inputs[1]); + const Value * logits = graph_value(context.graph, projection->output); + if (weight == nullptr || input == nullptr || logits == nullptr || weight->type != GGML_TYPE_F32 || + input->type != GGML_TYPE_F32 || logits->type != GGML_TYPE_F32 || !weight->contiguous || !input->contiguous || + !logits->contiguous || !is_shape(*input, kMoeRouterProfile.hidden_size, 1, 1, 1) || + !is_shape(*weight, kMoeRouterProfile.hidden_size, kMoeRouterProfile.expert_count, 1, 1) || + !is_shape(*logits, kMoeRouterProfile.expert_count, 1, 1, 1)) { + return {}; + } + + const GraphNode * softmax = find_single_consumer_with_op(context.graph, projection->output, GGML_OP_SOFT_MAX); + RouterTop8Match top8 = match_moe_router_top8(context.graph, softmax, status); + if (!top8.matched() || top8.token_count != 1) { + return {}; + } + + match.projection = projection; + match.softmax = softmax; + match.input = input; + match.weight = weight; + match.logits = logits; + match.top8 = top8; + return match; +} + +static RouterProjectionMatch match_moe_router_projection_f32(const DispatchMatchContext & context) { + RouterProjectionMatch match; + const GraphNode * projection = context.root_node; + if (projection == nullptr || projection->op != GGML_OP_MUL_MAT || projection->inputs.size() != 2) { + return match; + } + + const Value * weight = graph_value(context.graph, projection->inputs[0]); + const Value * input = graph_value(context.graph, projection->inputs[1]); + const Value * output = graph_value(context.graph, projection->output); + if (weight == nullptr || input == nullptr || output == nullptr || !is_2d(*weight) || !is_2d(*input) || + !is_2d(*output) || !weight->contiguous || !input->contiguous || !output->contiguous || + weight->type != GGML_TYPE_F32 || input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32) { + return {}; + } + + const int64_t input_size = weight->ne[0]; + const int64_t output_size = weight->ne[1]; + const int64_t token_count = input->ne[1]; + if (input_size != kMoeRouterProfile.hidden_size || output_size != kMoeRouterProfile.expert_count || + input->ne[0] != input_size || output->ne[0] != output_size || output->ne[1] != token_count || + !is_llm_supported_query_length(kMoeRouterProfile, token_count)) { + return {}; + } + + match.input = input; + match.weight = weight; + match.output = output; + match.token_count = token_count; + return match; +} + +static bool has_common_matmul_fusible_unary_consumer(const DispatchMatchContext & context, const Value & output) { + const std::vector & consumers = context.graph.index().consumers(output.id); + if (consumers.size() != 1 || consumers.front() == nullptr) { + return false; + } + + const GraphNode * unary = consumers.front(); + size_t unary_index = 0; + if (!context.graph.index().node_index(unary, unary_index) || unary_index >= context.covered_nodes.size() || + context.covered_nodes[unary_index] || unary->inputs.size() != 1 || unary->inputs[0] != output.id) { + return false; + } + + const UnaryParams * params = op_params_as(unary->params); + const Value * result = graph_value(context.graph, unary->output); + return params != nullptr && unary_kind_supported(params->op) && result != nullptr && + result->type == GGML_TYPE_F32 && result->contiguous && same_shape(output, *result); +} + +} // namespace + +static bool match_moe_router_projection_f32_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const RouterProjectionMatch match = match_moe_router_projection_f32(context); + if (!match.matched()) { + return false; + } + if (has_common_matmul_fusible_unary_consumer(context, *match.output)) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kQwenRouterProjectionF32FourRowWave32Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.workload.token_capacity", to_config_value(match.token_count)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.model.hidden_size", + to_config_value(kMoeRouterProfile.hidden_size)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.router.expert_count", + to_config_value(kMoeRouterProfile.expert_count)); + dispatch.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_moe_router_projection_top8_fused_decode_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const RouterProjectionTop8Match match = match_moe_router_projection_top8_decode(context, &dispatch_match.status); + if (!match.matched()) { + return false; + } + + const ValueId completion_counter_value = context.next_plan_value; + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kQwenRouterProjectionTop8FusedDecodeF32Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.top8.token_count); + dispatch.kernel.integer_parameters.emplace("route_id_stride", match.top8.route_stride); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.model.hidden_size", + to_config_value(kMoeRouterProfile.hidden_size)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.router.expert_count", + to_config_value(match.top8.expert_count)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.router.route_count", to_config_value(match.top8.route_count)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.workload.token_capacity", + to_config_value(match.top8.token_count)); + + // The ids view is a graph output when the experts are on another device, and its buffer ends at + // (token_count - 1) * route_stride + route_count ids, not token_count * route_stride (#95 uses the same bound + // for MUL_MAT_ID's route ids: dispatch-mul-mat-id-common.h). + const size_t route_id_length = + static_cast((match.top8.token_count - 1) * match.top8.route_stride + match.top8.route_count) * + sizeof(int32_t); + dispatch.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count }); + dispatch.bindings.push_back({ match.logits->id, 0, match.logits->byte_count }); + dispatch.bindings.push_back({ completion_counter_value, 0, sizeof(int32_t) }); + dispatch.bindings.push_back({ match.top8.route_ids->id, 0, route_id_length }); + dispatch.bindings.push_back({ match.top8.route_weights->id, 0, match.top8.route_weights->byte_count }); + + dispatch_match.completion_counter_requests.push_back({ + completion_counter_value, + "qwen.router.decode_projection_top8_completion_counter", + 1, + }); + if (!append_covered_node(context, match.projection, dispatch_match) || + !append_moe_router_top8_coverage_from_softmax(context, match.softmax, dispatch_match)) { + return false; + } + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_moe_router_top8_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const RouterTop8Match router_match = + match_moe_router_top8(context.graph, context.root_node, &dispatch_match.status); + if (!router_match.matched()) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kQwenRouterTop8F32Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", router_match.token_count); + dispatch.kernel.integer_parameters.emplace("route_id_stride", router_match.route_stride); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.router.expert_count", + to_config_value(router_match.expert_count)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.router.route_count", + to_config_value(router_match.route_count)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.workload.token_capacity", + to_config_value(router_match.token_count)); + + // Same bound as above: a graph-output ids view's buffer ends at (token_count - 1) * route_stride + + // route_count, not token_count * route_stride. + const size_t route_id_length = + static_cast((router_match.token_count - 1) * router_match.route_stride + router_match.route_count) * + sizeof(int32_t); + dispatch.bindings.push_back({ router_match.logits->id, 0, router_match.logits->byte_count }); + dispatch.bindings.push_back({ router_match.route_ids->id, 0, route_id_length }); + dispatch.bindings.push_back({ router_match.route_weights->id, 0, router_match.route_weights->byte_count }); + + if (!append_moe_router_top8_coverage(context, dispatch_match)) { + return false; + } + dispatch_match.dispatches.push_back(std::move(dispatch)); + + const ValueId expert_table_value(context.next_plan_value.value); + const ValueId partition_table_value(context.next_plan_value.value + 1); + const ValueId completion_counter_value(context.next_plan_value.value + 2); + const size_t expert_table_bytes = expert_table_size(router_match.token_count, router_match.expert_count); + const size_t partition_table_bytes = + partition_table_size(router_match.token_count, router_match.route_count, router_match.expert_count); + const bool use_fused_prefill_expert_table_partition = supports_fused_prefill_expert_table_partition(router_match); + dispatch_match.transients.push_back( + { expert_table_value, "qwen.router.expert_table", expert_table_bytes, kMoeRouterPlanTransientAlignment }); + dispatch_match.transients.push_back({ partition_table_value, "qwen.router.partition_table", partition_table_bytes, + kMoeRouterPlanTransientAlignment }); + if (use_fused_prefill_expert_table_partition) { + dispatch_match.completion_counter_requests.push_back({ + completion_counter_value, + "qwen.router.prefill_expert_table_partition_completion_counter", + 1, + }); + } + const CommandPlanResourceMetadata routing_metadata = make_command_plan_resource_metadata(MoeRoutingResourceMetadata{ + router_match.token_count, + router_match.route_count, + router_match.route_stride, + router_match.expert_count, + }); + Status metadata_status; + if (!dispatch_match.metadata.append_generated_resource( + { + router_match.route_ids->id, + GeneratedResourceRole::MoeExpertTable, + expert_table_value, + expert_table_bytes, + routing_metadata, + }, + metadata_status) || + !dispatch_match.metadata.append_generated_resource( + { + router_match.route_ids->id, + GeneratedResourceRole::MoePartitionTable, + partition_table_value, + partition_table_bytes, + routing_metadata, + }, + metadata_status) || + !dispatch_match.metadata.append_moe_routing_bundle( + { + router_match.route_ids->id, + router_match.route_weights->id, + expert_table_value, + partition_table_value, + expert_table_bytes, + partition_table_bytes, + router_match.token_count, + router_match.route_count, + router_match.route_stride, + router_match.expert_count, + }, + metadata_status)) { + return false; + } + + if (use_fused_prefill_expert_table_partition) { + Dispatch expert_table_partition_dispatch; + expert_table_partition_dispatch.kernel = + make_kernel_specialization(kQwenBuildExpertTablePartitionPrefill512Kernel); + expert_table_partition_dispatch.kernel.integer_parameters.emplace("token_count", router_match.token_count); + expert_table_partition_dispatch.kernel.integer_parameters.emplace("route_count", router_match.route_count); + expert_table_partition_dispatch.kernel.integer_parameters.emplace("route_stride", router_match.route_stride); + expert_table_partition_dispatch.kernel.integer_parameters.emplace("expert_count", router_match.expert_count); + expert_table_partition_dispatch.bindings.push_back({ router_match.route_ids->id, 0, route_id_length }); + expert_table_partition_dispatch.bindings.push_back({ expert_table_value, 0, expert_table_bytes }); + expert_table_partition_dispatch.bindings.push_back({ partition_table_value, 0, partition_table_bytes }); + expert_table_partition_dispatch.bindings.push_back({ completion_counter_value, 0, sizeof(int32_t) }); + dispatch_match.dispatches.push_back(std::move(expert_table_partition_dispatch)); + } else { + Dispatch expert_table_dispatch; + expert_table_dispatch.kernel = make_kernel_specialization(kMoeBuildExpertTableKernel); + expert_table_dispatch.kernel.integer_parameters.emplace("token_count", router_match.token_count); + expert_table_dispatch.kernel.integer_parameters.emplace("route_count", router_match.route_count); + expert_table_dispatch.kernel.integer_parameters.emplace("route_stride", router_match.route_stride); + expert_table_dispatch.kernel.integer_parameters.emplace("expert_count", router_match.expert_count); + add_moe_routing_compile_parameters(expert_table_dispatch, router_match); + expert_table_dispatch.bindings.push_back({ router_match.route_ids->id, 0, route_id_length }); + expert_table_dispatch.bindings.push_back({ expert_table_value, 0, expert_table_bytes }); + dispatch_match.dispatches.push_back(std::move(expert_table_dispatch)); + + Dispatch partition_table_dispatch; + partition_table_dispatch.kernel = make_kernel_specialization(kMoeBuildExpertPartitionTableKernel); + partition_table_dispatch.kernel.integer_parameters.emplace("token_count", router_match.token_count); + partition_table_dispatch.kernel.integer_parameters.emplace("route_count", router_match.route_count); + partition_table_dispatch.kernel.integer_parameters.emplace("expert_count", router_match.expert_count); + add_moe_routing_compile_parameters(partition_table_dispatch, router_match); + partition_table_dispatch.bindings.push_back({ expert_table_value, 0, expert_table_bytes }); + partition_table_dispatch.bindings.push_back({ partition_table_value, 0, partition_table_bytes }); + dispatch_match.dispatches.push_back(std::move(partition_table_dispatch)); + } + return true; +} + +void register_moe_router_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "llm.moe_router.projection_top8_fused_decode", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 1200, + DispatchSource::Llm, + match_moe_router_projection_top8_fused_decode_dispatch, + }); + registry.add({ + "llm.moe_router.top8_f32", + GGML_OP_SOFT_MAX, + DispatchMatchKind::Fused, + 1000, + DispatchSource::Llm, + match_moe_router_top8_dispatch, + }); + registry.add({ + "llm.moe_router.projection_f32_four_row_wave32", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 90, + DispatchSource::Llm, + match_moe_router_projection_f32_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-moe-router.h b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-moe-router.h new file mode 100644 index 000000000000..ff174cff5d12 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-moe-router.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_moe_router_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-attention-postprocess.cpp b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-attention-postprocess.cpp new file mode 100644 index 000000000000..9b0dfd1cff88 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-attention-postprocess.cpp @@ -0,0 +1,1025 @@ +#include "dispatch-qwen-attention-postprocess.h" + +#include "../common/dispatch-rope-utils.h" +#include "dispatch-llm-shapes.h" +#include "ggml.h" +#include "graph/graph-matcher.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kQwenAttentionPostprocessF32F16Kernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "qwen3_moe_attention_postprocess_f32_f16"); +static constexpr KernelCatalogRef kQwenAttentionQkvPostprocessFusedDecodeKernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "qwen3_moe_attention_qkv_postprocess_fused_decode"); +static constexpr KernelCatalogRef kQwenAttentionContextBaseCaptureKernel = + GGML_HRX_KERNEL_REF("qwen", "qwen_attention_context_base_capture"); +static constexpr KernelCatalogRef kQwenAttentionMetadataKernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "qwen_attention_metadata"); +static constexpr int64_t kQwenAttentionHeadSize = 128; + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static bool is_shape(const Value & value, int64_t ne0, int64_t ne1, int64_t ne2, int64_t ne3) { + return value.ne[0] == ne0 && value.ne[1] == ne1 && value.ne[2] == ne2 && value.ne[3] == ne3; +} + +static bool is_2d(const Value & value) { + return value.ne[0] > 0 && value.ne[1] > 0 && value.ne[2] == 1 && value.ne[3] == 1; +} + +static size_t q8_1_x4_byte_count(int64_t token_count, int64_t hidden_size) { + if (token_count <= 0 || hidden_size <= 0) { + return 0; + } + return static_cast(token_count) * ggml_row_size(GGML_TYPE_Q8_1, hidden_size); +} + +static bool is_supported_token_count(int64_t token_count) { + return token_count >= 1 && token_count <= 2048; +} + +static bool is_supported_head_count(int64_t head_count) { + return head_count >= 1 && head_count <= 64; +} + +static bool is_qwen_rms_norm_epsilon(float eps) { + return eps >= 0.0000009f && eps <= 0.0000011f; +} + +static bool is_qwen_implicit_rope_contract(const RopeParams & params) { + return params.n_dims == kQwenAttentionHeadSize && params.mode == GGML_ROPE_TYPE_NEOX && + std::isfinite(params.freq_base) && params.freq_base > 0.0f && std::isfinite(params.freq_scale) && + params.freq_scale > 0.0f && std::isfinite(params.ext_factor) && std::isfinite(params.attn_factor); +} + +static bool build_inverse_frequency_table(const GraphNode & rope, std::vector & data, float & rope_mscale) { + const RopeParams * params = op_params_as(rope.params); + if (params == nullptr || !is_qwen_implicit_rope_contract(*params)) { + return false; + } + + RopeFrequencyTable table; + if (!build_rope_frequency_table(*params, params->n_dims, table)) { + return false; + } + data = std::move(table.data); + rope_mscale = table.mscale; + return true; +} + +static bool is_supported_cache_index_type(ggml_type type) { + return type == GGML_TYPE_I64; +} + +static const GraphNode * find_single_consumer_with_op(const Graph & graph, ValueId value, ggml_op op) { + const GraphNode * match = nullptr; + for (const GraphNode * consumer : graph.index().consumers(value)) { + if (consumer == nullptr || consumer->op != op) { + continue; + } + if (match != nullptr) { + return nullptr; + } + match = consumer; + } + return match; +} + +static const GraphNode * producer_with_op(const Graph & graph, ValueId value, ggml_op op) { + const GraphNode * producer = graph.index().producer(value); + return producer != nullptr && producer->op == op ? producer : nullptr; +} + +static bool append_covered_node(const DispatchMatchContext & context, const GraphNode * node, DispatchMatch & match) { + return append_covered_node_index_once(context.graph, context.covered_nodes, node, match.covered_nodes); +} + +static std::string to_config_value(int64_t value) { + return std::to_string(value); +} + +static std::string value_summary(const Graph & graph, const Value * value) { + if (value == nullptr) { + return "missing"; + } + std::ostringstream stream; + stream << value->id.value << ":" << ggml_type_name(value->type) << "[" << value->ne[0] << "," << value->ne[1] << "," + << value->ne[2] << "," << value->ne[3] << "] nb=[" << value->nb[0] << "," << value->nb[1] << "," + << value->nb[2] << "," << value->nb[3] << "]"; + const GraphNode * producer = graph.index().producer(value->id); + if (producer != nullptr) { + stream << "<-" << ggml_op_name(producer->op); + } + if (value->alias_source.value >= 0) { + stream << " alias=" << value->alias_source.value << " storage_root=" << value->storage_root.value + << " storage_offset=" << value->storage_offset; + } + return stream.str(); +} + +static std::string node_summary(const Graph & graph, const GraphNode * node) { + if (node == nullptr) { + return "missing"; + } + std::ostringstream stream; + size_t index = 0; + if (graph.index().node_index(node, index)) { + stream << index << ":"; + } + stream << ggml_op_name(node->op) << " output=" << value_summary(graph, graph_value(graph, node->output)); + return stream.str(); +} + +static std::string rope_params_summary(const GraphNode & node) { + const RopeParams * params = op_params_as(node.params); + if (params == nullptr) { + return "missing"; + } + std::ostringstream stream; + stream << "n_dims=" << params->n_dims << " mode=" << params->mode << " n_ctx_orig=" << params->n_ctx_orig + << " freq_base=" << params->freq_base << " freq_scale=" << params->freq_scale + << " ext_factor=" << params->ext_factor << " attn_factor=" << params->attn_factor + << " beta_fast=" << params->beta_fast << " beta_slow=" << params->beta_slow; + return stream.str(); +} + +static bool is_attention_postprocess_candidate_root(const Graph & graph, const GraphNode * root) { + if (root == nullptr || root->op != GGML_OP_RESHAPE || root->inputs.size() != 1 || !graph.has_index()) { + return false; + } + const GraphNode * projection = producer_with_op(graph, root->inputs[0], GGML_OP_MUL_MAT); + const Value * raw_input = graph_value(graph, root->inputs[0]); + const Value * reshaped = graph_value(graph, root->output); + return projection != nullptr && raw_input != nullptr && reshaped != nullptr && raw_input->type == GGML_TYPE_F32 && + reshaped->type == GGML_TYPE_F32 && raw_input->ne[1] > 0 && reshaped->ne[0] == kQwenAttentionHeadSize && + reshaped->ne[2] == raw_input->ne[1]; +} + +static void log_attention_reject(Status * status, + const Graph & graph, + const GraphNode * root, + const std::string & reason) { + if (status == nullptr || !is_attention_postprocess_candidate_root(graph, root)) { + return; + } + const Value * input = root->inputs.empty() ? nullptr : graph_value(graph, root->inputs[0]); + const Value * output = graph_value(graph, root->output); + status->log("qwen attention postprocess matcher rejected node: %s root=%s input=%s output=%s", reason.c_str(), + node_summary(graph, root).c_str(), value_summary(graph, input).c_str(), + value_summary(graph, output).c_str()); +} + +static bool append_postprocess_node(const DispatchMatchContext & context, + const GraphNode * root, + const char * role, + const GraphNode * node, + DispatchMatch & dispatch_match, + Status * status) { + if (append_covered_node(context, node, dispatch_match)) { + return true; + } + std::string reason = std::string("cannot cover ") + role + " node " + node_summary(context.graph, node); + log_attention_reject(status, context.graph, root, reason); + return false; +} + +struct NormRopeChain { + const GraphNode * projection_node = nullptr; + const GraphNode * reshape_node = nullptr; + const GraphNode * rms_node = nullptr; + const GraphNode * mul_node = nullptr; + const GraphNode * rope_node = nullptr; + const Value * projection_input = nullptr; + const Value * raw_input = nullptr; + const Value * reshaped = nullptr; + const Value * norm_weight = nullptr; + const Value * positions = nullptr; + const Value * inverse_freqs = nullptr; + size_t inverse_freqs_byte_count = 0; + std::vector inverse_freqs_data; + float rope_mscale = 1.0f; + const Value * output = nullptr; + int64_t token_count = 0; + int64_t head_count = 0; + + bool matched() const { + return projection_node != nullptr && reshape_node != nullptr && rms_node != nullptr && mul_node != nullptr && + rope_node != nullptr && projection_input != nullptr && raw_input != nullptr && reshaped != nullptr && + norm_weight != nullptr && positions != nullptr && + (inverse_freqs != nullptr || !inverse_freqs_data.empty()) && inverse_freqs_byte_count > 0 && + output != nullptr && token_count > 0 && head_count > 0; + } +}; + +struct CachePublishChain { + NormRopeChain key; + const GraphNode * layout_node = nullptr; + const GraphNode * set_rows_node = nullptr; + const Value * cache_indices = nullptr; + const Value * cache = nullptr; + int64_t cache_row_count = 0; + bool key_publish_path = false; + + bool matched_key() const { + return key_publish_path && key.matched() && layout_node != nullptr && set_rows_node != nullptr && + cache_indices != nullptr && cache != nullptr && cache_row_count > 0; + } +}; + +struct ValuePublishChain { + const GraphNode * projection_node = nullptr; + const GraphNode * reshape_node = nullptr; + const GraphNode * layout_node = nullptr; + const GraphNode * set_rows_node = nullptr; + const Value * projection_input = nullptr; + const Value * raw_input = nullptr; + const Value * cache_indices = nullptr; + const Value * cache = nullptr; + int64_t token_count = 0; + int64_t head_count = 0; + int64_t cache_row_count = 0; + + bool matched() const { + return projection_node != nullptr && reshape_node != nullptr && layout_node != nullptr && + set_rows_node != nullptr && projection_input != nullptr && raw_input != nullptr && + cache_indices != nullptr && cache != nullptr && token_count > 0 && head_count > 0 && cache_row_count > 0; + } +}; + +struct FlashInputLayoutChain { + const GraphNode * query_layout = nullptr; + const GraphNode * query_permute = nullptr; + const GraphNode * key_layout = nullptr; + const GraphNode * key_permute = nullptr; + const GraphNode * value_layout = nullptr; + const GraphNode * value_permute = nullptr; + const GraphNode * flash = nullptr; + + bool matched() const { + return query_layout != nullptr && query_permute != nullptr && key_layout != nullptr && key_permute != nullptr && + value_layout != nullptr && value_permute != nullptr && flash != nullptr; + } +}; + +struct AttentionPostprocessMatch { + NormRopeChain query; + CachePublishChain key; + ValuePublishChain value; + FlashInputLayoutChain flash_layouts; + + bool matched() const { return query.matched() && key.matched_key() && value.matched(); } +}; + +static bool matching_inverse_frequencies(const NormRopeChain & lhs, const NormRopeChain & rhs) { + if (lhs.inverse_freqs != nullptr || rhs.inverse_freqs != nullptr) { + return lhs.inverse_freqs != nullptr && rhs.inverse_freqs != nullptr && + lhs.inverse_freqs->id == rhs.inverse_freqs->id && lhs.rope_mscale == rhs.rope_mscale; + } + return lhs.inverse_freqs_data == rhs.inverse_freqs_data && lhs.rope_mscale == rhs.rope_mscale; +} + +static bool has_qwen_rope_params(const GraphNode & node) { + const RopeParams * params = op_params_as(node.params); + return params != nullptr && params->n_dims == kQwenAttentionHeadSize && params->mode == GGML_ROPE_TYPE_NEOX; +} + +static bool has_qwen_rms_params(const GraphNode & node) { + const RmsNormParams * params = op_params_as(node.params); + return params != nullptr && is_qwen_rms_norm_epsilon(params->eps); +} + +static bool is_norm_weight(const Value & value) { + return value.type == GGML_TYPE_F32 && is_shape(value, kQwenAttentionHeadSize, 1, 1, 1); +} + +static bool is_inverse_frequency_table(const Value & value) { + return value.type == GGML_TYPE_F32 && is_shape(value, kQwenAttentionHeadSize / 2, 1, 1, 1); +} + +static bool match_projection_reshape(const Graph & graph, + const GraphNode * reshape, + NormRopeChain & chain, + Status * status, + const std::string & label) { + if (reshape == nullptr || reshape->op != GGML_OP_RESHAPE || reshape->inputs.size() != 1) { + log_attention_reject(status, graph, reshape, label + " projection reshape is not a single-input RESHAPE"); + return false; + } + + const GraphNode * projection = producer_with_op(graph, reshape->inputs[0], GGML_OP_MUL_MAT); + const Value * raw_input = graph_value(graph, reshape->inputs[0]); + const Value * reshaped = graph_value(graph, reshape->output); + const Value * projection_input = + projection == nullptr || projection->inputs.size() != 2 ? nullptr : graph_value(graph, projection->inputs[1]); + if (projection == nullptr || projection_input == nullptr || raw_input == nullptr || reshaped == nullptr) { + log_attention_reject(status, graph, reshape, label + " projection producer or values are missing"); + return false; + } + if (raw_input->type != GGML_TYPE_F32 || reshaped->type != GGML_TYPE_F32 || !is_2d(*raw_input) || + reshaped->ne[0] != kQwenAttentionHeadSize || reshaped->ne[3] != 1) { + log_attention_reject(status, graph, reshape, + label + " projection reshape has incompatible type, rank, or head size"); + return false; + } + + const int64_t head_count = reshaped->ne[1]; + const int64_t token_count = reshaped->ne[2]; + if (!is_supported_head_count(head_count) || !is_supported_token_count(token_count) || + raw_input->ne[0] != head_count * kQwenAttentionHeadSize || raw_input->ne[1] != token_count) { + log_attention_reject(status, graph, reshape, label + " projection reshape has unsupported head/token shape"); + return false; + } + + chain.projection_node = projection; + chain.reshape_node = reshape; + chain.projection_input = projection_input; + chain.raw_input = raw_input; + chain.reshaped = reshaped; + chain.token_count = token_count; + chain.head_count = head_count; + return true; +} + +static bool match_norm_rope_chain_from_reshape(const Graph & graph, + const GraphNode * reshape, + NormRopeChain & chain, + Status * status = nullptr, + const std::string & label = "attention") { + if (!match_projection_reshape(graph, reshape, chain, status, label)) { + return false; + } + + const GraphNode * rms = find_single_consumer_with_op(graph, chain.reshaped->id, GGML_OP_RMS_NORM); + if (rms == nullptr || rms->inputs.size() != 1 || !has_qwen_rms_params(*rms)) { + log_attention_reject(status, graph, reshape, label + " chain is missing supported RMS_NORM"); + return false; + } + + const GraphNode * mul = find_single_consumer_with_op(graph, rms->output, GGML_OP_MUL); + if (mul == nullptr || mul->inputs.size() != 2) { + log_attention_reject(status, graph, reshape, label + " chain is missing norm-weight MUL"); + return false; + } + ValueId weight_id; + if (mul->inputs[0] == rms->output) { + weight_id = mul->inputs[1]; + } else if (mul->inputs[1] == rms->output) { + weight_id = mul->inputs[0]; + } else { + log_attention_reject(status, graph, reshape, label + " norm-weight MUL does not consume RMS output"); + return false; + } + const Value * norm_weight = graph_value(graph, weight_id); + if (norm_weight == nullptr || !is_norm_weight(*norm_weight)) { + log_attention_reject(status, graph, reshape, label + " norm weight shape is incompatible"); + return false; + } + + const GraphNode * rope = find_single_consumer_with_op(graph, mul->output, GGML_OP_ROPE); + if (rope == nullptr || rope->inputs.size() < 2 || rope->inputs.size() > 3 || rope->inputs[0] != mul->output || + !has_qwen_rope_params(*rope)) { + std::string reason = label + " chain is missing supported ROPE"; + if (rope != nullptr) { + reason += " params=" + rope_params_summary(*rope); + } + log_attention_reject(status, graph, reshape, reason); + return false; + } + const Value * positions = graph_value(graph, rope->inputs[1]); + const Value * output = graph_value(graph, rope->output); + size_t inverse_freqs_byte_count = 0; + const Value * inverse_freqs = nullptr; + std::vector inverse_freqs_data; + float rope_mscale = 1.0f; + if (rope->inputs.size() == 3) { + inverse_freqs = graph_value(graph, rope->inputs[2]); + if (inverse_freqs == nullptr || !is_inverse_frequency_table(*inverse_freqs)) { + log_attention_reject(status, graph, reshape, label + " explicit inverse-frequency table is incompatible"); + return false; + } + const RopeParams * params = op_params_as(rope->params); + if (params == nullptr || !is_qwen_implicit_rope_contract(*params)) { + log_attention_reject(status, graph, reshape, + label + " explicit inverse-frequency table has unsupported ROPE params=" + + rope_params_summary(*rope)); + return false; + } + RopeFrequencyTable table; + if (!build_rope_frequency_table(*params, params->n_dims, table)) { + return false; + } + rope_mscale = table.mscale; + inverse_freqs_byte_count = inverse_freqs->byte_count; + } else if (build_inverse_frequency_table(*rope, inverse_freqs_data, rope_mscale)) { + inverse_freqs_byte_count = inverse_freqs_data.size(); + } else { + log_attention_reject(status, graph, reshape, + label + " implicit inverse-frequency table cannot be derived from ROPE params=" + + rope_params_summary(*rope)); + return false; + } + if (positions == nullptr || output == nullptr || positions->type != GGML_TYPE_I32 || + !is_shape(*positions, chain.token_count, 1, 1, 1) || output->type != GGML_TYPE_F32 || + !is_shape(*output, kQwenAttentionHeadSize, chain.head_count, chain.token_count, 1)) { + log_attention_reject(status, graph, reshape, label + " positions or ROPE output shape is incompatible"); + return false; + } + + chain.rms_node = rms; + chain.mul_node = mul; + chain.rope_node = rope; + chain.norm_weight = norm_weight; + chain.positions = positions; + chain.inverse_freqs = inverse_freqs; + chain.inverse_freqs_byte_count = inverse_freqs_byte_count; + chain.inverse_freqs_data = std::move(inverse_freqs_data); + chain.rope_mscale = rope_mscale; + chain.output = output; + return true; +} + +static const GraphNode * find_cache_read_layout(const Graph & graph, const Value & cache, int64_t head_count) { + const GraphNode * match = nullptr; + for (const GraphNode * consumer : layout_alias_consumers(graph, cache.id)) { + const Value * output = graph_value(graph, consumer->output); + if (output == nullptr || output->type != GGML_TYPE_F16 || output->ne[0] != kQwenAttentionHeadSize || + output->ne[1] != head_count || output->ne[3] != 1) { + continue; + } + if (match != nullptr) { + return nullptr; + } + match = consumer; + } + return match; +} + +static int64_t cache_row_count_for_value(const Value & cache, int64_t head_count) { + if (cache.type != GGML_TYPE_F16 || head_count <= 0) { + return 0; + } + if (cache.ne[0] == kQwenAttentionHeadSize * head_count && cache.ne[2] == 1 && cache.ne[3] == 1) { + return cache.ne[1]; + } + if (cache.ne[0] == kQwenAttentionHeadSize && cache.ne[2] == head_count && cache.ne[3] == 1) { + return cache.ne[1]; + } + return 0; +} + +static bool match_key_publish_chain(const Graph & graph, const GraphNode * set_rows, CachePublishChain & chain) { + if (set_rows == nullptr || set_rows->op != GGML_OP_SET_ROWS || set_rows->inputs.size() != 3) { + return false; + } + const GraphNode * layout = graph.index().producer(set_rows->inputs[0]); + if (layout == nullptr || !is_layout_alias_node(graph, *layout) || layout->inputs.size() != 1) { + return false; + } + + const GraphNode * rope = producer_with_op(graph, layout->inputs[0], GGML_OP_ROPE); + if (rope == nullptr) { + return false; + } + const GraphNode * mul = producer_with_op(graph, rope->inputs.empty() ? ValueId() : rope->inputs[0], GGML_OP_MUL); + if (mul == nullptr || mul->inputs.size() != 2) { + return false; + } + const GraphNode * rms = nullptr; + if (mul->inputs[0] != rope->inputs[0]) { + rms = producer_with_op(graph, mul->inputs[0], GGML_OP_RMS_NORM); + } + if (rms == nullptr && mul->inputs[1] != rope->inputs[0]) { + rms = producer_with_op(graph, mul->inputs[1], GGML_OP_RMS_NORM); + } + if (rms == nullptr || rms->inputs.size() != 1) { + return false; + } + const GraphNode * reshape = producer_with_op(graph, rms->inputs[0], GGML_OP_RESHAPE); + NormRopeChain key_chain; + if (!match_norm_rope_chain_from_reshape(graph, reshape, key_chain) || key_chain.rope_node != rope) { + return false; + } + + const Value * cache_indices = graph_value(graph, set_rows->inputs[1]); + const Value * cache = graph_value(graph, set_rows->inputs[2]); + if (cache_indices == nullptr || cache == nullptr || !is_supported_cache_index_type(cache_indices->type) || + !is_shape(*cache_indices, key_chain.token_count, 1, 1, 1)) { + return false; + } + const int64_t cache_row_count = cache_row_count_for_value(*cache, key_chain.head_count); + if (cache_row_count <= 0) { + return false; + } + + chain.key = key_chain; + chain.layout_node = layout; + chain.set_rows_node = set_rows; + chain.cache_indices = cache_indices; + chain.cache = cache; + chain.cache_row_count = cache_row_count; + chain.key_publish_path = true; + return true; +} + +static bool match_value_publish_chain(const Graph & graph, const GraphNode * set_rows, ValuePublishChain & chain) { + if (set_rows == nullptr || set_rows->op != GGML_OP_SET_ROWS || set_rows->inputs.size() != 3) { + return false; + } + const GraphNode * layout = graph.index().producer(set_rows->inputs[0]); + if (layout == nullptr || !is_layout_alias_node(graph, *layout) || layout->inputs.size() != 1) { + return false; + } + + const GraphNode * reshape = producer_with_op(graph, layout->inputs[0], GGML_OP_RESHAPE); + if (reshape == nullptr) { + reshape = layout; + } + if (reshape == nullptr || reshape->op != GGML_OP_RESHAPE || reshape->inputs.size() != 1) { + return false; + } + + NormRopeChain projection_shape; + if (!match_projection_reshape(graph, reshape, projection_shape, nullptr, "value")) { + return false; + } + + const Value * cache_indices = graph_value(graph, set_rows->inputs[1]); + const Value * cache = graph_value(graph, set_rows->inputs[2]); + if (cache_indices == nullptr || cache == nullptr || !is_supported_cache_index_type(cache_indices->type) || + !is_shape(*cache_indices, projection_shape.token_count, 1, 1, 1)) { + return false; + } + const int64_t cache_row_count = cache_row_count_for_value(*cache, projection_shape.head_count); + if (cache_row_count <= 0) { + return false; + } + + chain.projection_node = projection_shape.projection_node; + chain.reshape_node = projection_shape.reshape_node; + chain.layout_node = layout; + chain.set_rows_node = set_rows; + chain.projection_input = projection_shape.projection_input; + chain.raw_input = projection_shape.raw_input; + chain.cache_indices = cache_indices; + chain.cache = cache; + chain.token_count = projection_shape.token_count; + chain.head_count = projection_shape.head_count; + chain.cache_row_count = cache_row_count; + return true; +} + +static bool same_projection_input(const NormRopeChain & lhs, const NormRopeChain & rhs) { + return lhs.projection_input != nullptr && rhs.projection_input != nullptr && + lhs.projection_input->id == rhs.projection_input->id; +} + +static bool same_projection_input(const NormRopeChain & lhs, const ValuePublishChain & rhs) { + return lhs.projection_input != nullptr && rhs.projection_input != nullptr && + lhs.projection_input->id == rhs.projection_input->id; +} + +static FlashInputLayoutChain match_flash_input_layouts(const Graph & graph, const AttentionPostprocessMatch & match) { + FlashInputLayoutChain layouts; + const GraphNode * query_layout = find_single_layout_alias_consumer(graph, match.query.output->id); + if (query_layout == nullptr) { + return layouts; + } + const GraphNode * query_permute = find_single_consumer_with_op(graph, query_layout->output, GGML_OP_PERMUTE); + if (query_permute == nullptr) { + return {}; + } + + const GraphNode * key_layout = find_cache_read_layout(graph, *match.key.cache, match.key.key.head_count); + if (key_layout == nullptr) { + return {}; + } + const GraphNode * key_permute = find_single_consumer_with_op(graph, key_layout->output, GGML_OP_PERMUTE); + if (key_permute == nullptr) { + return {}; + } + + const GraphNode * value_layout = find_cache_read_layout(graph, *match.value.cache, match.value.head_count); + if (value_layout == nullptr) { + return {}; + } + const GraphNode * value_permute = find_single_consumer_with_op(graph, value_layout->output, GGML_OP_PERMUTE); + if (value_permute == nullptr) { + return {}; + } + + const GraphNode * flash = find_single_consumer_with_op(graph, query_permute->output, GGML_OP_FLASH_ATTN_EXT); + if (flash == nullptr || flash->inputs.size() != 4 || flash->inputs[0] != query_permute->output || + flash->inputs[1] != key_permute->output || flash->inputs[2] != value_permute->output) { + return {}; + } + + layouts.query_layout = query_layout; + layouts.query_permute = query_permute; + layouts.key_layout = key_layout; + layouts.key_permute = key_permute; + layouts.value_layout = value_layout; + layouts.value_permute = value_permute; + layouts.flash = flash; + return layouts; +} + +static AttentionPostprocessMatch match_qwen_attention_postprocess(const Graph & graph, + const GraphNode * root, + Status * status) { + AttentionPostprocessMatch match; + if (root == nullptr || root->op != GGML_OP_RESHAPE || !graph.has_index()) { + return match; + } + if (!match_norm_rope_chain_from_reshape(graph, root, match.query, status, "query")) { + return {}; + } + + for (const GraphNode & node : graph.nodes()) { + if (node.op != GGML_OP_SET_ROWS) { + continue; + } + CachePublishChain key; + if (!match.key.matched_key() && match_key_publish_chain(graph, &node, key) && + same_projection_input(match.query, key.key)) { + match.key = key; + continue; + } + ValuePublishChain value; + if (!match.value.matched() && match_value_publish_chain(graph, &node, value) && + same_projection_input(match.query, value)) { + match.value = value; + } + } + + if (!match.matched()) { + if (!match.key.matched_key()) { + log_attention_reject(status, graph, root, "no matching key ROPE cache publish chain found"); + } + if (!match.value.matched()) { + log_attention_reject(status, graph, root, "no matching value cache publish chain found"); + } + return {}; + } + if (match.query.token_count != match.key.key.token_count || match.query.token_count != match.value.token_count || + match.key.key.head_count != match.value.head_count || + match.query.positions->id != match.key.key.positions->id || + !matching_inverse_frequencies(match.query, match.key.key) || + match.key.cache_row_count != match.value.cache_row_count) { + log_attention_reject(status, graph, root, "query/key/value postprocess invariants are incompatible"); + return {}; + } + match.flash_layouts = match_flash_input_layouts(graph, match); + return match; +} + +static bool append_postprocess_covered_nodes(const DispatchMatchContext & context, + const AttentionPostprocessMatch & postprocess, + DispatchMatch & dispatch_match, + Status * status) { + // TODO: move fused matcher coverage into a shared builder that records GraphNode pointers during matching and + // materializes scheduler indices once. This is constant-size today, but the explicit list will not scale well as + // Qwen fused patterns grow. + const GraphNode * root = postprocess.query.reshape_node; + if (!append_postprocess_node(context, root, "query reshape", postprocess.query.reshape_node, dispatch_match, + status) || + !append_postprocess_node(context, root, "query rms", postprocess.query.rms_node, dispatch_match, status) || + !append_postprocess_node(context, root, "query mul", postprocess.query.mul_node, dispatch_match, status) || + !append_postprocess_node(context, root, "query rope", postprocess.query.rope_node, dispatch_match, status) || + !append_postprocess_node(context, root, "key reshape", postprocess.key.key.reshape_node, dispatch_match, + status) || + !append_postprocess_node(context, root, "key rms", postprocess.key.key.rms_node, dispatch_match, status) || + !append_postprocess_node(context, root, "key mul", postprocess.key.key.mul_node, dispatch_match, status) || + !append_postprocess_node(context, root, "key rope", postprocess.key.key.rope_node, dispatch_match, status) || + !append_postprocess_node(context, root, "key layout", postprocess.key.layout_node, dispatch_match, status) || + !append_postprocess_node(context, root, "key set rows", postprocess.key.set_rows_node, dispatch_match, + status) || + !append_postprocess_node(context, root, "value reshape", postprocess.value.reshape_node, dispatch_match, + status) || + !append_postprocess_node(context, root, "value layout", postprocess.value.layout_node, dispatch_match, + status) || + !append_postprocess_node(context, root, "value set rows", postprocess.value.set_rows_node, dispatch_match, + status)) { + return false; + } + if (postprocess.flash_layouts.matched() && + (!append_postprocess_node(context, root, "flash query layout", postprocess.flash_layouts.query_layout, + dispatch_match, status) || + !append_postprocess_node(context, root, "flash query permute", postprocess.flash_layouts.query_permute, + dispatch_match, status) || + !append_postprocess_node(context, root, "flash key layout", postprocess.flash_layouts.key_layout, + dispatch_match, status) || + !append_postprocess_node(context, root, "flash key permute", postprocess.flash_layouts.key_permute, + dispatch_match, status) || + !append_postprocess_node(context, root, "flash value layout", postprocess.flash_layouts.value_layout, + dispatch_match, status) || + !append_postprocess_node(context, root, "flash value permute", postprocess.flash_layouts.value_permute, + dispatch_match, status))) { + return false; + } + return true; +} + +static bool has_attention_metadata_initialization(const CommandPlan & plan) { + for (const Dispatch & dispatch : plan.initialization_dispatches) { + if (dispatch.kernel.kernel_id == kQwenAttentionMetadataKernel.id) { + return true; + } + } + return false; +} + +static ValueId next_match_transient_value(const DispatchMatchContext & context, const DispatchMatch & dispatch_match) { + return ValueId(context.next_plan_value.value + static_cast(dispatch_match.transients.size()) + + static_cast(dispatch_match.completion_counter_requests.size())); +} + +static bool append_attention_metadata_initialization(const DispatchMatchContext & context, + const AttentionPostprocessMatch & match, + DispatchMatch & dispatch_match) { + if (has_attention_metadata_initialization(context.plan)) { + return true; + } + if (!match.flash_layouts.matched() || match.flash_layouts.flash->inputs.size() != 4) { + return true; + } + + const Value * mask = graph_value(context.graph, match.flash_layouts.flash->inputs[3]); + if (mask == nullptr || mask->type != GGML_TYPE_F16 || mask->ne[0] <= 0 || mask->ne[1] != match.query.token_count) { + return false; + } + int64_t context_capacity = mask->ne[0]; + size_t mask_byte_count = mask->byte_count; + if (context_capacity > 32768 && match.query.token_count <= 32768) { + context_capacity = match.query.token_count; + mask_byte_count = static_cast(match.query.token_count) * static_cast(match.query.token_count) * + sizeof(ggml_fp16_t); + Status metadata_status; + if (!dispatch_match.metadata.append_alternate_value( + { mask->id, mask->id, GGML_TYPE_F16, mask_byte_count, "qwen.attention.compact_mask" }, + metadata_status)) { + dispatch_match.status.append(metadata_status); + return false; + } + } + + const ValueId control = next_match_transient_value(context, dispatch_match); + dispatch_match.transients.push_back({ control, "qwen.attention.control", sizeof(int32_t), 16 }); + + Dispatch context_capture; + context_capture.kernel = make_kernel_specialization(kQwenAttentionContextBaseCaptureKernel); + context_capture.bindings.push_back({ match.query.positions->id, 0, match.query.positions->byte_count }); + context_capture.bindings.push_back({ control, 0, sizeof(int32_t) }); + dispatch_match.initialization_dispatches.push_back(std::move(context_capture)); + + Dispatch metadata; + metadata.kernel = make_kernel_specialization(kQwenAttentionMetadataKernel); + metadata.kernel.integer_parameters.emplace("token_count", match.query.token_count); + metadata.kernel.integer_parameters.emplace("context_capacity", context_capacity); + metadata.bindings.push_back({ control, 0, sizeof(int32_t) }); + metadata.bindings.push_back({ match.query.positions->id, 0, match.query.positions->byte_count }); + metadata.bindings.push_back({ match.key.cache_indices->id, 0, match.key.cache_indices->byte_count }); + metadata.bindings.push_back({ match.value.cache_indices->id, 0, match.value.cache_indices->byte_count }); + metadata.bindings.push_back({ mask->id, 0, mask_byte_count }); + dispatch_match.initialization_dispatches.push_back(std::move(metadata)); + return true; +} + +static const Value * projection_weight(const Graph & graph, const GraphNode * projection) { + return projection == nullptr || projection->inputs.size() != 2 ? nullptr : + graph_value(graph, projection->inputs[0]); +} + +static bool is_qwen_attention_projection_weight(const Value & weight, int64_t input_size, int64_t output_size) { + return (weight.type == GGML_TYPE_Q4_K || weight.type == GGML_TYPE_Q6_K) && weight.contiguous && + is_shape(weight, input_size, output_size, 1, 1); +} + +static uint32_t attention_qkv_completion_counter_count(const AttentionPostprocessMatch & match) { + return static_cast(match.query.head_count + 2 * match.key.key.head_count); +} + +} // namespace + +static bool match_qwen_attention_qkv_postprocess_fused_decode_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const GraphNode * root = context.root_node; + if (root != nullptr && root->op == GGML_OP_MUL_MAT) { + root = find_single_consumer_with_op(context.graph, root->output, GGML_OP_RESHAPE); + } + const AttentionPostprocessMatch match = + match_qwen_attention_postprocess(context.graph, root, &dispatch_match.status); + if (!match.matched() || !is_qwen_decode_query_length(match.query.token_count)) { + return false; + } + if (match.query.projection_input->id != match.key.key.projection_input->id || + match.query.projection_input->id != match.value.projection_input->id) { + return false; + } + + const int64_t hidden_size = match.query.projection_input->ne[0]; + const int64_t query_size = match.query.head_count * kQwenAttentionHeadSize; + const int64_t key_value_size = match.key.key.head_count * kQwenAttentionHeadSize; + if (hidden_size != 2048 || match.query.head_count != 32 || match.key.key.head_count != 4 || + match.value.head_count != match.key.key.head_count) { + return false; + } + + const Value * query_weight = projection_weight(context.graph, match.query.projection_node); + const Value * key_weight = projection_weight(context.graph, match.key.key.projection_node); + const Value * value_weight = projection_weight(context.graph, match.value.projection_node); + if (query_weight == nullptr || key_weight == nullptr || value_weight == nullptr || + query_weight->type != GGML_TYPE_Q4_K || key_weight->type != GGML_TYPE_Q4_K || + (value_weight->type != GGML_TYPE_Q4_K && value_weight->type != GGML_TYPE_Q6_K) || + !is_qwen_attention_projection_weight(*query_weight, hidden_size, query_size) || + !is_qwen_attention_projection_weight(*key_weight, hidden_size, key_value_size) || + !is_qwen_attention_projection_weight(*value_weight, hidden_size, key_value_size)) { + return false; + } + + const size_t q8_input_bytes = q8_1_x4_byte_count(match.query.token_count, hidden_size); + const CommandPlanAlternateValue * q8_input = + find_alternate_value(context.plan, match.query.projection_input->id, GGML_TYPE_Q8_1, q8_input_bytes); + if (q8_input == nullptr) { + return false; + } + if (!append_postprocess_covered_nodes(context, match, dispatch_match, &dispatch_match.status)) { + return false; + } + if (!append_postprocess_node(context, root, "query projection", match.query.projection_node, dispatch_match, + &dispatch_match.status) || + !append_postprocess_node(context, root, "key projection", match.key.key.projection_node, dispatch_match, + &dispatch_match.status) || + !append_postprocess_node(context, root, "value projection", match.value.projection_node, dispatch_match, + &dispatch_match.status)) { + return false; + } + if (!append_attention_metadata_initialization(context, match, dispatch_match)) { + return false; + } + + const bool synthetic_inverse_frequencies = match.query.inverse_freqs == nullptr; + const ValueId inverse_frequencies_value = synthetic_inverse_frequencies ? + next_match_transient_value(context, dispatch_match) : + match.query.inverse_freqs->id; + const size_t inverse_frequencies_size = match.query.inverse_freqs_byte_count; + if (synthetic_inverse_frequencies) { + dispatch_match.transients.push_back({ inverse_frequencies_value, + "qwen.attention_qkv_decode.inverse_frequencies", inverse_frequencies_size, + 256 }); + dispatch_match.constant_initializations.push_back({ + inverse_frequencies_value, + "qwen.attention_qkv_decode.inverse_frequencies", + 0, + match.query.inverse_freqs_data, + }); + } + + const ValueId completion_counters = next_match_transient_value(context, dispatch_match); + const uint32_t completion_counter_count = attention_qkv_completion_counter_count(match); + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kQwenAttentionQkvPostprocessFusedDecodeKernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.query.token_count); + dispatch.kernel.integer_parameters.emplace("cache_row_count", match.key.cache_row_count); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.model.hidden_size", to_config_value(hidden_size)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.attention.query_size", to_config_value(query_size)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.attention.key_value_size", to_config_value(key_value_size)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.attention.value_uses_q6", + value_weight->type == GGML_TYPE_Q6_K ? "1" : "0"); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.attention.head_size", + to_config_value(kQwenAttentionHeadSize)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.attention.rope_mscale", + rope_mscale_config_value(match.query.rope_mscale)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.model.rms_epsilon", "0.000001"); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.workload.token_capacity", + to_config_value(match.query.token_count)); + + dispatch.bindings.push_back({ q8_input->alternate_value, 0, q8_input->byte_count }); + dispatch.bindings.push_back({ query_weight->id, 0, query_weight->byte_count }); + dispatch.bindings.push_back({ key_weight->id, 0, key_weight->byte_count }); + dispatch.bindings.push_back({ value_weight->id, 0, value_weight->byte_count }); + dispatch.bindings.push_back({ match.query.positions->id, 0, match.query.positions->byte_count }); + dispatch.bindings.push_back({ match.key.cache_indices->id, 0, match.key.cache_indices->byte_count }); + dispatch.bindings.push_back({ match.value.cache_indices->id, 0, match.value.cache_indices->byte_count }); + dispatch.bindings.push_back({ match.query.raw_input->id, 0, match.query.raw_input->byte_count }); + dispatch.bindings.push_back({ match.key.key.raw_input->id, 0, match.key.key.raw_input->byte_count }); + dispatch.bindings.push_back({ match.value.raw_input->id, 0, match.value.raw_input->byte_count }); + dispatch.bindings.push_back({ match.query.norm_weight->id, 0, match.query.norm_weight->byte_count }); + dispatch.bindings.push_back({ match.key.key.norm_weight->id, 0, match.key.key.norm_weight->byte_count }); + dispatch.bindings.push_back({ inverse_frequencies_value, 0, inverse_frequencies_size }); + dispatch.bindings.push_back({ match.query.output->id, 0, match.query.output->byte_count }); + dispatch.bindings.push_back({ match.key.cache->id, 0, match.key.cache->byte_count }); + dispatch.bindings.push_back({ match.value.cache->id, 0, match.value.cache->byte_count }); + dispatch.bindings.push_back({ completion_counters, 0, completion_counter_count * sizeof(int32_t) }); + + dispatch_match.completion_counter_requests.push_back({ + completion_counters, + "qwen.attention_qkv_decode.completion_counters", + completion_counter_count, + }); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_qwen_attention_postprocess_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const AttentionPostprocessMatch match = + match_qwen_attention_postprocess(context.graph, context.root_node, &dispatch_match.status); + if (!match.matched()) { + return false; + } + if (!append_postprocess_covered_nodes(context, match, dispatch_match, &dispatch_match.status)) { + return false; + } + if (!append_attention_metadata_initialization(context, match, dispatch_match)) { + return false; + } + + Dispatch dispatch; + const bool synthetic_inverse_frequencies = match.query.inverse_freqs == nullptr; + const ValueId inverse_frequencies_value = synthetic_inverse_frequencies ? + next_match_transient_value(context, dispatch_match) : + match.query.inverse_freqs->id; + const size_t inverse_frequencies_size = match.query.inverse_freqs_byte_count; + dispatch.kernel = make_kernel_specialization(kQwenAttentionPostprocessF32F16Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.query.token_count); + dispatch.kernel.integer_parameters.emplace("cache_row_count", match.key.cache_row_count); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.model.rms_epsilon", "0.000001"); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.attention.head_size", + to_config_value(kQwenAttentionHeadSize)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.attention.query_size", + to_config_value(match.query.head_count * kQwenAttentionHeadSize)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.attention.key_value_size", + to_config_value(match.key.key.head_count * kQwenAttentionHeadSize)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.attention.rope_mscale", + rope_mscale_config_value(match.query.rope_mscale)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.workload.token_capacity", + to_config_value(match.query.token_count)); + dispatch.bindings.push_back({ match.query.positions->id, 0, match.query.positions->byte_count }); + dispatch.bindings.push_back({ match.key.cache_indices->id, 0, match.key.cache_indices->byte_count }); + dispatch.bindings.push_back({ match.value.cache_indices->id, 0, match.value.cache_indices->byte_count }); + dispatch.bindings.push_back({ match.query.raw_input->id, 0, match.query.raw_input->byte_count }); + dispatch.bindings.push_back({ match.key.key.raw_input->id, 0, match.key.key.raw_input->byte_count }); + dispatch.bindings.push_back({ match.value.raw_input->id, 0, match.value.raw_input->byte_count }); + dispatch.bindings.push_back({ match.query.norm_weight->id, 0, match.query.norm_weight->byte_count }); + dispatch.bindings.push_back({ match.key.key.norm_weight->id, 0, match.key.key.norm_weight->byte_count }); + dispatch.bindings.push_back({ inverse_frequencies_value, 0, inverse_frequencies_size }); + dispatch.bindings.push_back({ match.query.output->id, 0, match.query.output->byte_count }); + dispatch.bindings.push_back({ match.key.cache->id, 0, match.key.cache->byte_count }); + dispatch.bindings.push_back({ match.value.cache->id, 0, match.value.cache->byte_count }); + + dispatch_match.dispatches.push_back(std::move(dispatch)); + if (synthetic_inverse_frequencies) { + dispatch_match.transients.push_back({ inverse_frequencies_value, + "qwen.attention_postprocess.inverse_frequencies", + inverse_frequencies_size, 256 }); + dispatch_match.constant_initializations.push_back({ + inverse_frequencies_value, + "qwen.attention_postprocess.inverse_frequencies", + 0, + match.query.inverse_freqs_data, + }); + } + return true; +} + +void register_qwen_attention_postprocess_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "qwen.attention_qkv_postprocess_fused_decode", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 300, + DispatchSource::Qwen, + match_qwen_attention_qkv_postprocess_fused_decode_dispatch, + }); + registry.add({ + "qwen.attention_qkv_postprocess_fused_decode", + GGML_OP_RESHAPE, + DispatchMatchKind::Fused, + 200, + DispatchSource::Qwen, + match_qwen_attention_qkv_postprocess_fused_decode_dispatch, + }); + registry.add({ + "qwen.attention_postprocess_f32_f16", + GGML_OP_RESHAPE, + DispatchMatchKind::Fused, + 100, + DispatchSource::Qwen, + match_qwen_attention_postprocess_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-attention-postprocess.h b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-attention-postprocess.h new file mode 100644 index 000000000000..66cb74342732 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-attention-postprocess.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_qwen_attention_postprocess_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-matmul.cpp b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-matmul.cpp new file mode 100644 index 000000000000..202932340635 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-matmul.cpp @@ -0,0 +1,513 @@ +#include "dispatch-qwen-matmul.h" + +#include "dispatch-llm-shapes.h" +#include "ggml.h" +#include "graph/graph-matcher.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kQwenDenseLinearQ6KF16WmmaKernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "qwen3_moe_dense_linear_q6k_f16_wmma"); +static constexpr KernelCatalogRef kQwenDenseLinearQ4KQ8NextQ8Kernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8"); +static constexpr KernelCatalogRef kGgmlLinearQ6KQ8_1X4Kernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "ggml_linear_q6k_q8_1_x4"); + +static constexpr int64_t kQwenHiddenSize = kQwen30BMoeDispatchProfile.hidden_size; +static constexpr int64_t kQwenVocabularyCount = 151936; + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static bool is_2d(const Value & value) { + return value.ne[0] > 0 && value.ne[1] > 0 && value.ne[2] == 1 && value.ne[3] == 1; +} + +static bool is_supported_dense_input_size(int64_t input_size) { + return input_size >= 256 && input_size <= 32768 && input_size % 256 == 0; +} + +static bool is_supported_dense_output_size(int64_t output_size) { + return output_size >= 1 && output_size <= 262144; +} + +static bool is_qwen_endpoint_projection(int64_t input_size, int64_t output_size) { + return input_size == kQwenHiddenSize && output_size == kQwenVocabularyCount; +} + +static std::string to_config_value(int64_t value) { + return std::to_string(value); +} + +struct QwenMatmulMatch { + const Value * input = nullptr; + const Value * weight = nullptr; + const Value * output = nullptr; + KernelCatalogRef kernel = {}; + ValueId input_value = {}; + size_t input_bytes = 0; + int64_t input_size = 0; + int64_t output_size = 0; + int64_t token_count = 0; + bool dense = false; + + bool matched() const { + return input != nullptr && weight != nullptr && output != nullptr && kernel.id != kUncatalogedKernelId; + } +}; + +struct QwenAttentionOutputNextQ8Match { + const Value * input = nullptr; + const CommandPlanAlternateValue * input_alternate = nullptr; + const Value * weight = nullptr; + const Value * projection_output = nullptr; + const Value * residual_input = nullptr; + const Value * residual_output = nullptr; + const Value * norm_weight = nullptr; + const Value * normalized_output = nullptr; + const GraphNode * projection_get_rows = nullptr; + const GraphNode * residual_get_rows = nullptr; + const GraphNode * add_node = nullptr; + const GraphNode * rms_node = nullptr; + const GraphNode * mul_node = nullptr; + int64_t input_size = 0; + int64_t output_size = 0; + int64_t token_count = 0; + + bool matched() const { + return input != nullptr && input_alternate != nullptr && weight != nullptr && projection_output != nullptr && + residual_input != nullptr && residual_output != nullptr && norm_weight != nullptr && + normalized_output != nullptr && add_node != nullptr && rms_node != nullptr && mul_node != nullptr; + } +}; + +static size_t q8_1_x4_byte_count(int64_t token_count, int64_t input_size) { + if (token_count <= 0 || input_size <= 0) { + return 0; + } + return static_cast(token_count) * ggml_row_size(GGML_TYPE_Q8_1, input_size); +} + +static const GraphNode * find_single_consumer_with_op(const Graph & graph, ValueId value, ggml_op op) { + const GraphNode * match = nullptr; + for (const GraphNode * consumer : graph.index().consumers(value)) { + if (consumer == nullptr || consumer->op != op) { + continue; + } + if (match != nullptr) { + return nullptr; + } + match = consumer; + } + return match; +} + +static bool value_has_no_uncovered_consumers_except(const DispatchMatchContext & context, + ValueId value, + const GraphNode * expected_consumer) { + for (const GraphNode * consumer : context.graph.index().consumers(value)) { + if (consumer == expected_consumer) { + continue; + } + size_t consumer_index = 0; + if (!context.graph.index().node_index(consumer, consumer_index) || + consumer_index >= context.covered_nodes.size() || !context.covered_nodes[consumer_index]) { + return false; + } + } + return true; +} + +static bool same_value_layout(const Value & lhs, const Value & rhs) { + if (lhs.type != rhs.type || lhs.byte_count != rhs.byte_count || lhs.element_count != rhs.element_count || + lhs.contiguous != rhs.contiguous) { + return false; + } + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (lhs.ne[i] != rhs.ne[i] || lhs.nb[i] != rhs.nb[i]) { + return false; + } + } + return true; +} + +static const Value * get_rows_source_with_same_layout(const Graph & graph, const GraphNode * node) { + if (node == nullptr || node->op != GGML_OP_GET_ROWS || node->inputs.size() != 2) { + return nullptr; + } + const Value * source = graph_value(graph, node->inputs[0]); + const Value * output = graph_value(graph, node->output); + if (source == nullptr || output == nullptr || !same_value_layout(*source, *output)) { + return nullptr; + } + return source; +} + +static QwenMatmulMatch match_qwen_q6k_q8_matmul(const Graph & graph, const GraphNode * node, const CommandPlan & plan) { + QwenMatmulMatch match; + if (node == nullptr || node->op != GGML_OP_MUL_MAT || node->inputs.size() != 2) { + return match; + } + + const Value * weight = graph_value(graph, node->inputs[0]); + const Value * input = graph_value(graph, node->inputs[1]); + const Value * output = graph_value(graph, node->output); + if (weight == nullptr || input == nullptr || output == nullptr) { + return {}; + } + if (!is_2d(*weight) || !is_2d(*input) || !is_2d(*output)) { + return {}; + } + if (!weight->contiguous || !input->contiguous || !output->contiguous) { + return {}; + } + if (weight->type != GGML_TYPE_Q6_K || input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32) { + return {}; + } + + const int64_t input_size = weight->ne[0]; + const int64_t output_size = weight->ne[1]; + const int64_t token_count = input->ne[1]; + if (input->ne[0] != input_size || output->ne[0] != output_size || output->ne[1] != token_count) { + return {}; + } + if (!is_qwen_decode_query_length(token_count) || input_size != kQwenHiddenSize || + output_size != kQwenVocabularyCount || !is_supported_dense_input_size(input_size) || + !is_supported_dense_output_size(output_size)) { + return {}; + } + + const size_t q8_byte_count = q8_1_x4_byte_count(token_count, input_size); + const CommandPlanAlternateValue * alternate = + find_alternate_value(graph, plan, input->id, GGML_TYPE_Q8_1, q8_byte_count); + if (alternate == nullptr) { + return {}; + } + + match.input = input; + match.weight = weight; + match.output = output; + match.input_value = alternate->alternate_value; + match.input_bytes = alternate->byte_count; + match.kernel = kGgmlLinearQ6KQ8_1X4Kernel; + match.input_size = input_size; + match.output_size = output_size; + match.token_count = token_count; + return match; +} + +static QwenMatmulMatch match_qwen_decode_endpoint_q6k_matmul(const Graph & graph, const GraphNode * node) { + QwenMatmulMatch match; + if (node == nullptr || node->op != GGML_OP_MUL_MAT || node->inputs.size() != 2) { + return match; + } + + const Value * weight = graph_value(graph, node->inputs[0]); + const Value * input = graph_value(graph, node->inputs[1]); + const Value * output = graph_value(graph, node->output); + if (weight == nullptr || input == nullptr || output == nullptr || !is_2d(*weight) || !is_2d(*input) || + !is_2d(*output) || !weight->contiguous || !input->contiguous || !output->contiguous || + weight->type != GGML_TYPE_Q6_K || input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32) { + return {}; + } + + const int64_t input_size = weight->ne[0]; + const int64_t output_size = weight->ne[1]; + const int64_t token_count = input->ne[1]; + if (input->ne[0] != input_size || output->ne[0] != output_size || output->ne[1] != token_count || + !is_qwen_decode_query_length(token_count) || !is_qwen_endpoint_projection(input_size, output_size) || + !is_supported_dense_input_size(input_size) || !is_supported_dense_output_size(output_size)) { + return {}; + } + + match.input = input; + match.weight = weight; + match.output = output; + match.input_value = input->id; + match.input_bytes = input->byte_count; + match.kernel = kQwenDenseLinearQ6KF16WmmaKernel; + match.input_size = input_size; + match.output_size = output_size; + match.token_count = token_count; + match.dense = true; + return match; +} + +static QwenAttentionOutputNextQ8Match match_qwen_attention_output_next_q8(const DispatchMatchContext & context) { + QwenAttentionOutputNextQ8Match match; + const Graph & graph = context.graph; + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_MUL_MAT || node->inputs.size() != 2 || !graph.has_index()) { + return match; + } + + const Value * weight = graph_value(graph, node->inputs[0]); + const Value * input = graph_value(graph, node->inputs[1]); + const Value * projection_output = graph_value(graph, node->output); + if (weight == nullptr || input == nullptr || projection_output == nullptr || !is_2d(*weight) || !is_2d(*input) || + !is_2d(*projection_output)) { + return {}; + } + if (weight->type != GGML_TYPE_Q4_K || input->type != GGML_TYPE_F32 || projection_output->type != GGML_TYPE_F32 || + !weight->contiguous || !input->contiguous || !projection_output->contiguous) { + return {}; + } + + const int64_t input_size = weight->ne[0]; + const int64_t output_size = weight->ne[1]; + const int64_t token_count = input->ne[1]; + if (token_count != 1 || input->ne[0] != input_size || projection_output->ne[0] != output_size || + projection_output->ne[1] != token_count || input_size != 4096 || output_size != kQwenHiddenSize) { + return {}; + } + + const CommandPlanAlternateValue * input_alternate = find_alternate_value( + graph, context.plan, input->id, GGML_TYPE_Q8_1, q8_1_x4_byte_count(token_count, input_size)); + if (input_alternate == nullptr) { + return {}; + } + + const Value * selected_projection = projection_output; + const GraphNode * projection_get_rows = nullptr; + const GraphNode * add_node = find_single_consumer_with_op(graph, projection_output->id, GGML_OP_ADD); + if (add_node == nullptr) { + projection_get_rows = find_single_consumer_with_op(graph, projection_output->id, GGML_OP_GET_ROWS); + if (get_rows_source_with_same_layout(graph, projection_get_rows) != projection_output) { + return {}; + } + selected_projection = graph_value(graph, projection_get_rows->output); + add_node = find_single_consumer_with_op(graph, projection_get_rows->output, GGML_OP_ADD); + } + if (add_node == nullptr || add_node->inputs.size() != 2) { + return {}; + } + const Value * residual_input = nullptr; + const Value * selected_residual = nullptr; + const GraphNode * residual_get_rows = nullptr; + const GraphNode * residual_consumer = add_node; + if (add_node->inputs[0] == selected_projection->id) { + selected_residual = graph_value(graph, add_node->inputs[1]); + } else if (add_node->inputs[1] == selected_projection->id) { + selected_residual = graph_value(graph, add_node->inputs[0]); + } + if (selected_residual == nullptr) { + return {}; + } + const GraphNode * selected_residual_producer = graph.index().producer(selected_residual->id); + residual_input = get_rows_source_with_same_layout(graph, selected_residual_producer); + if (residual_input != nullptr) { + residual_get_rows = selected_residual_producer; + residual_consumer = residual_get_rows; + } else { + residual_input = selected_residual; + } + const Value * residual_output = graph_value(graph, add_node->output); + if (residual_input == nullptr || residual_output == nullptr || residual_input->type != GGML_TYPE_F32 || + residual_output->type != GGML_TYPE_F32 || !same_value_layout(*selected_projection, *selected_residual) || + !same_value_layout(*selected_projection, *residual_input) || + !same_value_layout(*selected_projection, *residual_output) || + // The in-place residual add reuses residual_input's storage for residual_output, which is only + // safe when both are transient (an external value shares a persistent buffer) (#95). + residual_input->kind != ValueKind::Transient || residual_output->kind != ValueKind::Transient || + !value_has_no_uncovered_consumers_except(context, residual_input->id, residual_consumer)) { + return {}; + } + + const GraphNode * rms_node = find_single_consumer_with_op(graph, residual_output->id, GGML_OP_RMS_NORM); + if (rms_node == nullptr || rms_node->inputs.size() != 1) { + return {}; + } + const Value * rms_output = graph_value(graph, rms_node->output); + const GraphNode * mul_node = find_single_consumer_with_op(graph, rms_node->output, GGML_OP_MUL); + if (rms_output == nullptr || mul_node == nullptr || mul_node->inputs.size() != 2) { + return {}; + } + + const Value * norm_weight = nullptr; + if (mul_node->inputs[0] == rms_node->output) { + norm_weight = graph_value(graph, mul_node->inputs[1]); + } else if (mul_node->inputs[1] == rms_node->output) { + norm_weight = graph_value(graph, mul_node->inputs[0]); + } + const Value * normalized_output = graph_value(graph, mul_node->output); + if (norm_weight == nullptr || normalized_output == nullptr || norm_weight->type != GGML_TYPE_F32 || + normalized_output->type != GGML_TYPE_F32 || !norm_weight->contiguous || !normalized_output->contiguous || + norm_weight->ne[0] != output_size || normalized_output->ne[0] != output_size || + normalized_output->ne[1] != token_count) { + return {}; + } + + match.input = input; + match.input_alternate = input_alternate; + match.weight = weight; + match.projection_output = projection_output; + match.residual_input = residual_input; + match.residual_output = residual_output; + match.norm_weight = norm_weight; + match.normalized_output = normalized_output; + match.projection_get_rows = projection_get_rows; + match.residual_get_rows = residual_get_rows; + match.add_node = add_node; + match.rms_node = rms_node; + match.mul_node = mul_node; + match.input_size = input_size; + match.output_size = output_size; + match.token_count = token_count; + return match; +} + +} // namespace + +static void build_qwen_matmul_dispatch(const QwenMatmulMatch & match, + DispatchMatch & dispatch_match, + size_t root_index) { + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(match.kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + if (match.kernel.id == kGgmlLinearQ6KQ8_1X4Kernel.id) { + dispatch.kernel.integer_parameters.emplace("input_size", match.input_size); + dispatch.kernel.integer_parameters.emplace("output_size", match.output_size); + dispatch.kernel.compile_parameters.emplace("ggml.linear_q6k_q8_1_x4.token_capacity", + to_config_value(match.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.linear_q6k_q8_1_x4.output_capacity", + to_config_value(match.output_size)); + } else { + dispatch.kernel.compile_parameters.emplace("qwen3_moe.workload.token_capacity", + to_config_value(match.token_count)); + } + if (match.dense) { + dispatch.kernel.compile_parameters.emplace("qwen3_moe.dense_quantized.input_size", + to_config_value(match.input_size)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.dense_quantized.output_size", + to_config_value(match.output_size)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.dense_quantized.output_accumulation", "0"); + } + dispatch.bindings.push_back({ match.input_value, 0, match.input_bytes }); + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + + dispatch_match.covered_nodes.push_back(root_index); + dispatch_match.dispatches.push_back(std::move(dispatch)); +} + +static bool match_qwen_q6k_q8_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { + const QwenMatmulMatch match = match_qwen_q6k_q8_matmul(context.graph, context.root_node, context.plan); + if (!match.matched()) { + return false; + } + build_qwen_matmul_dispatch(match, dispatch_match, context.root_index); + return true; +} + +static bool match_qwen_decode_endpoint_q6k_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const QwenMatmulMatch match = match_qwen_decode_endpoint_q6k_matmul(context.graph, context.root_node); + if (!match.matched()) { + return false; + } + build_qwen_matmul_dispatch(match, dispatch_match, context.root_index); + return true; +} + +static bool match_qwen_attention_output_next_q8_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const QwenAttentionOutputNextQ8Match match = match_qwen_attention_output_next_q8(context); + if (!match.matched()) { + return false; + } + + const ValueId completion_counter(context.next_plan_value.value); + const ValueId q8_output(context.next_plan_value.value + 1); + const size_t q8_output_bytes = q8_1_x4_byte_count(match.token_count, match.output_size); + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kQwenDenseLinearQ4KQ8NextQ8Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.dense_quantized.input_size", + to_config_value(match.input_size)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.dense_quantized.output_size", + to_config_value(match.output_size)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.dense_quantized.output_accumulation", "1"); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.model.hidden_size", to_config_value(match.output_size)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.model.rms_epsilon", "0.000001"); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.workload.token_capacity", to_config_value(match.token_count)); + dispatch.bindings.push_back({ match.input_alternate->alternate_value, 0, match.input_alternate->byte_count }); + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count }); + dispatch.bindings.push_back({ match.residual_output->id, 0, match.residual_output->byte_count }); + dispatch.bindings.push_back({ match.norm_weight->id, 0, match.norm_weight->byte_count }); + dispatch.bindings.push_back({ match.normalized_output->id, 0, match.normalized_output->byte_count }); + dispatch.bindings.push_back({ completion_counter, 0, sizeof(int32_t) }); + dispatch.bindings.push_back({ q8_output, 0, q8_output_bytes }); + + dispatch_match.value_aliases.push_back({ match.residual_input->id, match.residual_output->id }); + dispatch_match.completion_counter_requests.push_back({ + completion_counter, + "qwen.decode.attention_output.completion_counter", + 1, + }); + dispatch_match.transients.push_back( + { q8_output, "qwen.decode.attention_output.next_q8_output", q8_output_bytes, 256 }); + Status metadata_status; + if (!dispatch_match.metadata.append_alternate_value( + { match.normalized_output->id, q8_output, GGML_TYPE_Q8_1, q8_output_bytes, + "qwen.decode.attention_output.next_q8_output" }, + metadata_status)) { + dispatch_match.status.append(metadata_status); + return false; + } + + if (!append_covered_node_index_once(context.graph, context.covered_nodes, context.root_node, + dispatch_match.covered_nodes) || + (match.projection_get_rows != nullptr && + !append_covered_node_index_once(context.graph, context.covered_nodes, match.projection_get_rows, + dispatch_match.covered_nodes)) || + (match.residual_get_rows != nullptr && + !append_covered_node_index_once(context.graph, context.covered_nodes, match.residual_get_rows, + dispatch_match.covered_nodes)) || + !append_covered_node_index_once(context.graph, context.covered_nodes, match.add_node, + dispatch_match.covered_nodes) || + !append_covered_node_index_once(context.graph, context.covered_nodes, match.rms_node, + dispatch_match.covered_nodes) || + !append_covered_node_index_once(context.graph, context.covered_nodes, match.mul_node, + dispatch_match.covered_nodes)) { + return false; + } + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +void register_qwen_matmul_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "qwen.matmul.attention_output_q4k_q8_1_x4_next_q8", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 300, + DispatchSource::Qwen, + match_qwen_attention_output_next_q8_dispatch, + }); + registry.add({ + "qwen.matmul.q6k_q8_1_x4", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 305, + DispatchSource::Qwen, + match_qwen_q6k_q8_dispatch, + }); + registry.add({ + "qwen.matmul.decode_endpoint_q6k_f16_wmma", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 100, + DispatchSource::Qwen, + match_qwen_decode_endpoint_q6k_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-matmul.h b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-matmul.h new file mode 100644 index 000000000000..59eb88679ff3 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-matmul.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_qwen_matmul_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-rmsnorm.cpp b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-rmsnorm.cpp new file mode 100644 index 000000000000..db6b45794dbd --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-rmsnorm.cpp @@ -0,0 +1,400 @@ +#include "dispatch-qwen-rmsnorm.h" + +#include "dispatch-llm-profiles.h" +#include "ggml.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kQwenRmsNormF32QuantizeQ8_1X4Kernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "qwen3_moe_rmsnorm_f32_quantize_q8_1_x4"); +static constexpr KernelCatalogRef kGgmlLinearQ6KQ8_1X4Kernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "ggml_linear_q6k_q8_1_x4"); +static constexpr float kQwenRmsNormEpsilon = kQwen30BMoeDispatchProfile.rms_norm_epsilon; +static constexpr int64_t kQwenHiddenSize = kQwen30BMoeDispatchProfile.hidden_size; +static constexpr int64_t kQwenVocabularyCount = 151936; + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static bool same_shape(const Value & lhs, const Value & rhs) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (lhs.ne[i] != rhs.ne[i]) { + return false; + } + } + return true; +} + +static bool is_qwen_rms_norm_epsilon(float eps) { + return std::fabs(eps - kQwenRmsNormEpsilon) <= 1.0e-12f; +} + +static bool is_supported_hidden_size(int64_t hidden_size) { + return hidden_size >= 128 && hidden_size <= 32768 && hidden_size % 128 == 0; +} + +static bool is_supported_token_count(int64_t token_count) { + return is_llm_supported_query_length(kQwen30BMoeDispatchProfile, token_count); +} + +static bool has_decode_q8_consumer(const Graph & graph, ValueId value) { + if (!graph.has_index()) { + return false; + } + for (const GraphNode * consumer : graph.index().consumers(value)) { + if (consumer != nullptr && (consumer->op == GGML_OP_MUL_MAT || consumer->op == GGML_OP_MUL_MAT_ID)) { + return true; + } + } + return false; +} + +static size_t q8_1_x4_byte_count(int64_t token_count, int64_t hidden_size) { + if (token_count <= 0 || hidden_size <= 0) { + return 0; + } + return static_cast(token_count) * ggml_row_size(GGML_TYPE_Q8_1, hidden_size); +} + +static bool is_weight_shape(const Value & weight, int64_t hidden_size) { + if (weight.ne[0] != hidden_size) { + return false; + } + for (int i = 1; i < GGML_MAX_DIMS; ++i) { + if (weight.ne[i] != 1) { + return false; + } + } + return true; +} + +struct RmsNormMatch { + const GraphNode * rms_node = nullptr; + const GraphNode * mul_node = nullptr; + const Value * input = nullptr; + const Value * weight = nullptr; + const Value * output = nullptr; + size_t rms_node_index = 0; + size_t mul_node_index = 0; + int64_t hidden_size = 0; + int64_t token_count = 0; + int64_t q8_group_count = 0; + + bool matched() const { + return rms_node != nullptr && mul_node != nullptr && input != nullptr && weight != nullptr && output != nullptr; + } +}; + +static RmsNormMatch match_qwen_rmsnorm_f32(const Graph & graph, const GraphNode * node, size_t node_index) { + RmsNormMatch match; + if (node == nullptr || node->op != GGML_OP_RMS_NORM || node->inputs.size() != 1 || !graph.has_index()) { + return match; + } + + const RmsNormParams * rms_params = op_params_as(node->params); + if (rms_params == nullptr || !is_qwen_rms_norm_epsilon(rms_params->eps)) { + return {}; + } + const std::vector & consumers = graph.index().consumers(node->output); + if (consumers.size() != 1) { + return {}; + } + const GraphNode * mul_node = consumers.front(); + size_t mul_node_index; + if (mul_node == nullptr || mul_node->op != GGML_OP_MUL || mul_node->inputs.size() != 2 || + !graph.index().node_index(mul_node, mul_node_index)) { + return {}; + } + + const Value * weight = nullptr; + for (ValueId input : mul_node->inputs) { + if (input != node->output) { + weight = graph_value(graph, input); + } + } + const Value * input = graph_value(graph, node->inputs[0]); + const Value * rms = graph_value(graph, node->output); + const Value * output = graph_value(graph, mul_node->output); + if (input == nullptr || rms == nullptr || weight == nullptr || output == nullptr) { + return {}; + } + if (input->type != GGML_TYPE_F32 || rms->type != GGML_TYPE_F32 || weight->type != GGML_TYPE_F32 || + output->type != GGML_TYPE_F32) { + return {}; + } + if (!input->contiguous || !rms->contiguous || !weight->contiguous || !output->contiguous) { + return {}; + } + if (!same_shape(*input, *rms) || !same_shape(*input, *output)) { + return {}; + } + + const int64_t hidden_size = output->ne[0]; + if (!is_supported_hidden_size(hidden_size) || !is_weight_shape(*weight, hidden_size)) { + return {}; + } + if (hidden_size == 0 || output->element_count <= 0 || output->element_count % hidden_size != 0) { + return {}; + } + const int64_t token_count = output->element_count / hidden_size; + if (!is_supported_token_count(token_count)) { + return {}; + } + + match.rms_node = node; + match.mul_node = mul_node; + match.input = input; + match.weight = weight; + match.output = output; + match.rms_node_index = node_index; + match.mul_node_index = mul_node_index; + match.hidden_size = hidden_size; + match.token_count = token_count; + match.q8_group_count = token_count * ((hidden_size + 127) / 128); + return match; +} + +struct QwenEndpointProjectionMatch { + const GraphNode * projection_node = nullptr; + const Value * weight = nullptr; + const Value * output = nullptr; + size_t projection_node_index = 0; + + bool matched() const { return projection_node != nullptr && weight != nullptr && output != nullptr; } +}; + +static QwenEndpointProjectionMatch match_qwen_endpoint_projection(const Graph & graph, const RmsNormMatch & match) { + QwenEndpointProjectionMatch projection_match; + if (!graph.has_index()) { + return projection_match; + } + const std::vector & consumers = graph.index().consumers(match.output->id); + if (consumers.size() != 1) { + return {}; + } + const GraphNode * consumer = consumers.front(); + if (consumer == nullptr || consumer->op != GGML_OP_MUL_MAT || consumer->inputs.size() != 2) { + return {}; + } + size_t consumer_index = 0; + if (!graph.index().node_index(consumer, consumer_index)) { + return {}; + } + const Value * weight = graph_value(graph, consumer->inputs[0]); + const Value * input = graph_value(graph, consumer->inputs[1]); + const Value * output = graph_value(graph, consumer->output); + if (weight == nullptr || input == nullptr || output == nullptr) { + return {}; + } + if (input->id != match.output->id || weight->type != GGML_TYPE_Q6_K || output->type != GGML_TYPE_F32 || + !weight->contiguous || !input->contiguous || !output->contiguous || match.hidden_size != kQwenHiddenSize || + !is_supported_token_count(match.token_count) || weight->ne[0] != match.hidden_size || + weight->ne[1] != kQwenVocabularyCount || output->ne[0] != kQwenVocabularyCount || + output->ne[1] != match.token_count || output->ne[2] != 1 || output->ne[3] != 1) { + return {}; + } + + projection_match.projection_node = consumer; + projection_match.projection_node_index = consumer_index; + projection_match.weight = weight; + projection_match.output = output; + return projection_match; +} + +static RmsNormMatch match_qwen_endpoint_rmsnorm_from_projection(const Graph & graph, const GraphNode * projection) { + if (projection == nullptr || projection->op != GGML_OP_MUL_MAT || projection->inputs.size() != 2 || + !graph.has_index()) { + return {}; + } + + const Value * input = graph_value(graph, projection->inputs[1]); + if (input == nullptr) { + return {}; + } + const GraphNode * mul_node = graph.index().producer(input->id); + size_t mul_node_index; + if (mul_node == nullptr || mul_node->op != GGML_OP_MUL || mul_node->inputs.size() != 2 || + !graph.index().node_index(mul_node, mul_node_index)) { + return {}; + } + + for (const ValueId mul_input : mul_node->inputs) { + const GraphNode * rms_node = graph.index().producer(mul_input); + size_t rms_node_index; + if (rms_node != nullptr && rms_node->op == GGML_OP_RMS_NORM && + graph.index().node_index(rms_node, rms_node_index)) { + return match_qwen_rmsnorm_f32(graph, rms_node, rms_node_index); + } + } + return {}; +} + +static std::string to_config_value(int64_t value) { + return std::to_string(value); +} + +static std::string to_config_value(float value) { + std::ostringstream out; + out.precision(9); + out << value; + return out.str(); +} + +} // namespace + +static bool match_qwen_rmsnorm_f32_quantize_q8_1_x4_dispatch(const DispatchMatchContext & context, + DispatchMatch & match) { + const std::vector & nodes = context.graph.nodes(); + if (context.root_index >= nodes.size()) { + return false; + } + const RmsNormMatch rms_match = + context.root_node->op == GGML_OP_MUL_MAT ? + match_qwen_endpoint_rmsnorm_from_projection(context.graph, context.root_node) : + match_qwen_rmsnorm_f32(context.graph, &nodes[context.root_index], context.root_index); + const QwenEndpointProjectionMatch projection_match = + rms_match.matched() ? match_qwen_endpoint_projection(context.graph, rms_match) : QwenEndpointProjectionMatch{}; + if (!rms_match.matched() || rms_match.rms_node_index >= context.covered_nodes.size() || + rms_match.mul_node_index >= context.covered_nodes.size() || context.covered_nodes[rms_match.rms_node_index] || + context.covered_nodes[rms_match.mul_node_index] || !projection_match.matched() || + projection_match.projection_node_index >= context.covered_nodes.size() || + context.covered_nodes[projection_match.projection_node_index] || + (context.root_node->op == GGML_OP_MUL_MAT && projection_match.projection_node != context.root_node)) { + return false; + } + + const size_t q8_byte_count = q8_1_x4_byte_count(rms_match.token_count, rms_match.hidden_size); + if (q8_byte_count == 0) { + return false; + } + + const ValueId q8_value = context.next_plan_value; + + Dispatch rms_dispatch; + rms_dispatch.kernel = make_kernel_specialization(kQwenRmsNormF32QuantizeQ8_1X4Kernel); + rms_dispatch.kernel.integer_parameters.emplace("token_count", rms_match.token_count); + rms_dispatch.kernel.compile_parameters.emplace("qwen3_moe.model.hidden_size", + to_config_value(rms_match.hidden_size)); + rms_dispatch.kernel.compile_parameters.emplace("qwen3_moe.model.rms_epsilon", "0.000001"); + rms_dispatch.kernel.compile_parameters.emplace("qwen3_moe.workload.token_capacity", + to_config_value(rms_match.token_count)); + rms_dispatch.kernel.compile_parameters.emplace("ggml.quantize_q8_1_x4.group_capacity", + to_config_value(rms_match.q8_group_count)); + rms_dispatch.bindings.push_back({ rms_match.input->id, 0, rms_match.input->byte_count }); + rms_dispatch.bindings.push_back({ rms_match.weight->id, 0, rms_match.weight->byte_count }); + rms_dispatch.bindings.push_back({ rms_match.output->id, 0, rms_match.output->byte_count }); + rms_dispatch.bindings.push_back({ q8_value, 0, q8_byte_count }); + + Dispatch projection_dispatch; + projection_dispatch.kernel = make_kernel_specialization(kGgmlLinearQ6KQ8_1X4Kernel); + projection_dispatch.kernel.integer_parameters.emplace("token_count", rms_match.token_count); + projection_dispatch.kernel.integer_parameters.emplace("input_size", rms_match.hidden_size); + projection_dispatch.kernel.integer_parameters.emplace("output_size", kQwenVocabularyCount); + projection_dispatch.kernel.compile_parameters.emplace("ggml.linear_q6k_q8_1_x4.token_capacity", + to_config_value(rms_match.token_count)); + projection_dispatch.kernel.compile_parameters.emplace("ggml.linear_q6k_q8_1_x4.output_capacity", + to_config_value(kQwenVocabularyCount)); + projection_dispatch.bindings.push_back({ q8_value, 0, q8_byte_count }); + projection_dispatch.bindings.push_back({ projection_match.weight->id, 0, projection_match.weight->byte_count }); + projection_dispatch.bindings.push_back({ projection_match.output->id, 0, projection_match.output->byte_count }); + + match.covered_nodes.push_back(rms_match.rms_node_index); + match.covered_nodes.push_back(rms_match.mul_node_index); + match.covered_nodes.push_back(projection_match.projection_node_index); + match.dispatches.push_back(std::move(rms_dispatch)); + match.dispatches.push_back(std::move(projection_dispatch)); + match.transients.push_back({ q8_value, "qwen.rmsnorm.q8_1_x4", q8_byte_count, 256 }); + return match.status.success(); +} + +static bool match_qwen_decode_rmsnorm_f32_quantize_q8_1_x4_dispatch(const DispatchMatchContext & context, + DispatchMatch & match) { + const std::vector & nodes = context.graph.nodes(); + if (context.root_index >= nodes.size()) { + return false; + } + const RmsNormMatch rms_match = + match_qwen_rmsnorm_f32(context.graph, &nodes[context.root_index], context.root_index); + if (!rms_match.matched() || rms_match.token_count != 1 || rms_match.hidden_size != kQwenHiddenSize || + rms_match.rms_node_index >= context.covered_nodes.size() || + rms_match.mul_node_index >= context.covered_nodes.size() || context.covered_nodes[rms_match.rms_node_index] || + context.covered_nodes[rms_match.mul_node_index]) { + return false; + } + if (!has_decode_q8_consumer(context.graph, rms_match.output->id)) { + return false; + } + + const size_t q8_byte_count = q8_1_x4_byte_count(rms_match.token_count, rms_match.hidden_size); + if (q8_byte_count == 0) { + return false; + } + + const ValueId q8_value = context.next_plan_value; + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kQwenRmsNormF32QuantizeQ8_1X4Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", rms_match.token_count); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.model.hidden_size", to_config_value(rms_match.hidden_size)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.model.rms_epsilon", "0.000001"); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.workload.token_capacity", + to_config_value(rms_match.token_count)); + dispatch.kernel.compile_parameters.emplace("ggml.quantize_q8_1_x4.group_capacity", + to_config_value(rms_match.q8_group_count)); + dispatch.bindings.push_back({ rms_match.input->id, 0, rms_match.input->byte_count }); + dispatch.bindings.push_back({ rms_match.weight->id, 0, rms_match.weight->byte_count }); + dispatch.bindings.push_back({ rms_match.output->id, 0, rms_match.output->byte_count }); + dispatch.bindings.push_back({ q8_value, 0, q8_byte_count }); + + Status metadata_status; + if (!match.metadata.append_alternate_value( + { rms_match.output->id, q8_value, GGML_TYPE_Q8_1, q8_byte_count, "qwen.decode.q8_hidden" }, + metadata_status)) { + match.status.append(metadata_status); + return false; + } + + match.covered_nodes.push_back(rms_match.rms_node_index); + match.covered_nodes.push_back(rms_match.mul_node_index); + match.dispatches.push_back(std::move(dispatch)); + match.transients.push_back({ q8_value, "qwen.decode.q8_hidden", q8_byte_count, 256 }); + return match.status.success(); +} + +void register_qwen_rmsnorm_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "qwen.endpoint_rmsnorm_q6k_q8_1_x4", + GGML_OP_MUL_MAT, + DispatchMatchKind::Fused, + 1200, + DispatchSource::Qwen, + match_qwen_rmsnorm_f32_quantize_q8_1_x4_dispatch, + }); + registry.add({ + "qwen.decode_rmsnorm_f32_quantize_q8_1_x4", + GGML_OP_RMS_NORM, + DispatchMatchKind::Fused, + 1150, + DispatchSource::Qwen, + match_qwen_decode_rmsnorm_f32_quantize_q8_1_x4_dispatch, + }); + registry.add({ + "qwen.rmsnorm_f32_quantize_q8_1_x4", + GGML_OP_RMS_NORM, + DispatchMatchKind::Fused, + 1100, + DispatchSource::Qwen, + match_qwen_rmsnorm_f32_quantize_q8_1_x4_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-rmsnorm.h b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-rmsnorm.h new file mode 100644 index 000000000000..e728f90a63dc --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen-rmsnorm.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_qwen_rmsnorm_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen.cpp b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen.cpp new file mode 100644 index 000000000000..28dad1bfff41 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen.cpp @@ -0,0 +1,19 @@ +#include "dispatch-qwen.h" + +#include "dispatch-moe-router.h" +#include "dispatch-qwen-attention-postprocess.h" +#include "dispatch-qwen-matmul.h" +#include "dispatch-qwen-rmsnorm.h" +#include "dispatch-routed-ffn.h" + +namespace ggml::hrx { + +void register_qwen_dispatches(DispatchRegistryBuilder & registry) { + register_qwen_attention_postprocess_dispatches(registry); + register_qwen_matmul_dispatches(registry); + register_routed_ffn_dispatches(registry); + register_qwen_rmsnorm_dispatches(registry); + register_moe_router_dispatches(registry); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen.h b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen.h new file mode 100644 index 000000000000..377477b7f7a5 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-qwen.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_qwen_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-routed-ffn.cpp b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-routed-ffn.cpp new file mode 100644 index 000000000000..e9e97ce38247 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-routed-ffn.cpp @@ -0,0 +1,1321 @@ +#include "dispatch-routed-ffn.h" + +#include "dispatch-llm-shapes.h" +#include "dispatch_registration/common/dispatch-mul-mat-weight-format.h" +#include "ggml.h" +#include "graph/graph-matcher.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kCommonRoutedGateUpSwiGLUF16WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_id_swiglu_f16_f16_wmma"); +static constexpr KernelCatalogRef kQwenRoutedGateUpSwiGLUQ4KQ8Kernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "qwen3_moe_routed_gate_up_swiglu_q4k_q8"); +static constexpr KernelCatalogRef kQwenRoutedGateUpSwiGLUQ4KQ8NextQ8Kernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8"); +static constexpr KernelCatalogRef kCommonMulMatIdF16F16WmmaKernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_id_f16_f16_wmma"); +static constexpr KernelCatalogRef kQwenRoutedDownQ4KQ8NextQ8Kernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "qwen3_moe_routed_down_q4k_q8_1_x4_next_q8"); +static constexpr KernelCatalogRef kQwenRoutedDownQ6KF32Wave64NextQ8Kernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "qwen3_moe_routed_down_q6k_f32_wave64_next_q8"); +static constexpr KernelCatalogRef kQwenRoutedDownWeightedReduceF16F32Kernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "qwen3_moe_routed_down_weighted_reduce_f16_f32"); +static constexpr KernelCatalogRef kQwenRoutedDownWeightedReduceNextRmsNormF32Kernel = + GGML_HRX_KERNEL_REF("qwen3_moe", "qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32"); + +static constexpr const LlmMoeDispatchProfile & kRoutedFfnProfile = kActiveLlmMoeDispatchProfile; +static constexpr int64_t kRoutedFfnInputSize = kRoutedFfnProfile.hidden_size; +static constexpr int64_t kRoutedFfnExpertHiddenSize = kRoutedFfnProfile.expert_hidden_size; +static constexpr int64_t kRoutedFfnExpertCount = kRoutedFfnProfile.expert_count; +static constexpr int64_t kRoutedFfnRouteCount = kRoutedFfnProfile.route_count; +static constexpr size_t kRoutedFfnPlanTransientAlignment = 256; +static constexpr const char * kRoutedFfnF16GateUpOutputName = "qwen.moe.gate_up_swiglu_f16"; +static constexpr const char * kRoutedFfnF16RoutedDownOutputName = "qwen.moe.routed_down_f16"; +static constexpr const char * kRoutedFfnQ8GateUpOutputName = "qwen.decode.moe.gate_up_swiglu_q8"; +static constexpr const char * kRoutedFfnQ8HiddenOutputName = "qwen.decode.moe.hidden_q8"; + +static const Value * graph_value(const Graph & graph, ValueId id) { + return graph.values().find(id); +} + +static bool is_shape(const Value & value, int64_t ne0, int64_t ne1, int64_t ne2, int64_t ne3) { + return value.ne[0] == ne0 && value.ne[1] == ne1 && value.ne[2] == ne2 && value.ne[3] == ne3; +} + +static bool same_shape(const Value & lhs, const Value & rhs) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (lhs.ne[i] != rhs.ne[i]) { + return false; + } + } + return true; +} + +static bool is_profile_rms_norm_epsilon(float eps) { + const float expected = kRoutedFfnProfile.rms_norm_epsilon; + return eps >= expected * 0.9f && eps <= expected * 1.1f; +} + +static bool is_routed_ffn_gate_up_weight(const Value & value) { + return (value.type == GGML_TYPE_Q4_K || value.type == GGML_TYPE_Q6_K) && value.contiguous && + is_shape(value, kRoutedFfnInputSize, kRoutedFfnExpertHiddenSize, kRoutedFfnExpertCount, 1); +} + +static bool is_fast_routed_ffn_gate_up_weight(const Value & value) { + const bool supported_type = value.type == GGML_TYPE_Q4_K || value.type == GGML_TYPE_Q6_K || + value.type == GGML_TYPE_Q8_0 || value.type == GGML_TYPE_Q8_1 || + value.type == GGML_TYPE_F16; + return supported_type && value.contiguous && + is_shape(value, kRoutedFfnInputSize, kRoutedFfnExpertHiddenSize, kRoutedFfnExpertCount, 1); +} + +static bool is_routed_ffn_down_weight(const Value & value) { + return (value.type == GGML_TYPE_Q4_K || value.type == GGML_TYPE_Q6_K) && value.contiguous && + is_shape(value, kRoutedFfnExpertHiddenSize, kRoutedFfnInputSize, kRoutedFfnExpertCount, 1); +} + +static bool is_routed_ffn_projection_output(const Value & value, int64_t token_count) { + return value.type == GGML_TYPE_F32 && value.contiguous && + is_shape(value, kRoutedFfnExpertHiddenSize, kRoutedFfnRouteCount, token_count, 1); +} + +static bool is_routed_ffn_down_output(const Value & value, int64_t token_count) { + return value.type == GGML_TYPE_F32 && value.contiguous && + is_shape(value, kRoutedFfnInputSize, kRoutedFfnRouteCount, token_count, 1); +} + +static const GraphNode * find_consumer_with_op(const Graph & graph, ValueId value, ggml_op op) { + for (const GraphNode * consumer : graph.index().consumers(value)) { + if (consumer != nullptr && consumer->op == op) { + return consumer; + } + } + return nullptr; +} + +static std::vector find_consumers_with_op(const Graph & graph, ValueId value, ggml_op op) { + std::vector matches; + for (const GraphNode * consumer : graph.index().consumers(value)) { + if (consumer != nullptr && consumer->op == op) { + matches.push_back(consumer); + } + } + return matches; +} + +static const GraphNode * find_single_consumer_with_op(const Graph & graph, ValueId value, ggml_op op) { + const std::vector consumers = find_consumers_with_op(graph, value, op); + return consumers.size() == 1 ? consumers.front() : nullptr; +} + +static const GraphNode * producer_with_op(const Graph & graph, ValueId value, ggml_op op) { + const GraphNode * producer = graph.index().producer(value); + return producer != nullptr && producer->op == op ? producer : nullptr; +} + +static bool is_swiglu_params(const OpParams & params) { + const BinaryParams * binary_params = op_params_as(params); + if (binary_params != nullptr) { + return binary_params->op == BinaryKind::SwiGLU; + } + + const GluParams * glu_params = op_params_as(params); + return glu_params != nullptr && glu_params->op == GGML_GLU_OP_SWIGLU; +} + +static bool append_covered_node(const DispatchMatchContext & context, const GraphNode * node, DispatchMatch & match) { + return append_covered_node_index_once(context.graph, context.covered_nodes, node, match.covered_nodes); +} + +static std::string to_config_value(int64_t value) { + return std::to_string(value); +} + +static size_t expert_table_size(int64_t token_count) { + return static_cast(kRoutedFfnExpertCount + kRoutedFfnExpertCount * token_count) * sizeof(int32_t); +} + +static size_t partition_table_size(int64_t token_count) { + const int64_t assignment_count = token_count * kRoutedFfnRouteCount; + const int64_t assignment_partition_count = (assignment_count + 31) / 32; + return static_cast(1 + assignment_partition_count + kRoutedFfnExpertCount) * sizeof(int32_t); +} + +static size_t f16_gate_up_output_size(int64_t token_count) { + return static_cast(token_count * kRoutedFfnRouteCount * kRoutedFfnExpertHiddenSize) * sizeof(ggml_fp16_t); +} + +static size_t f16_routed_down_output_size(int64_t token_count) { + return static_cast(token_count * kRoutedFfnRouteCount * kRoutedFfnInputSize) * sizeof(ggml_fp16_t); +} + +static size_t q8_1_x4_byte_count(int64_t row_count, int64_t input_size) { + if (row_count <= 0 || input_size <= 0) { + return 0; + } + return static_cast(row_count) * ggml_row_size(GGML_TYPE_Q8_1, input_size); +} + +static uint32_t gate_up_completion_counter_count(int64_t token_count) { + const int64_t physical_group_count = (kRoutedFfnExpertHiddenSize + 127) / 128; + return static_cast(token_count * kRoutedFfnRouteCount * physical_group_count); +} + +static void add_routed_down_compile_parameters(Dispatch & dispatch, int64_t token_count) { + dispatch.kernel.compile_parameters.emplace("qwen3_moe.routed_down.input_size", + to_config_value(kRoutedFfnExpertHiddenSize)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.routed_down.route_count", + to_config_value(kRoutedFfnRouteCount)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.routed_down.expert_count", + to_config_value(kRoutedFfnExpertCount)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.routed_down.output_size", + to_config_value(kRoutedFfnInputSize)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.workload.token_capacity", to_config_value(token_count)); +} + +static void add_common_routed_down_compile_parameters(Dispatch & dispatch, int64_t token_count, ggml_type weight_type) { + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_f16_f16.input_size", + to_config_value(kRoutedFfnExpertHiddenSize)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_f16_f16.route_count", + to_config_value(kRoutedFfnRouteCount)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_f16_f16.expert_count", + to_config_value(kRoutedFfnExpertCount)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_f16_f16.output_size", + to_config_value(kRoutedFfnInputSize)); + dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_id_f16_f16.weight_format", + weight_type == GGML_TYPE_Q4_K ? "4" : "6"); + dispatch.kernel.compile_parameters.emplace("ggml.workload.token_capacity", to_config_value(token_count)); +} + +struct RoutedGateUpMatch { + const Value * gate_weight = nullptr; + const Value * up_weight = nullptr; + const Value * input = nullptr; + const Value * route_ids = nullptr; + const Value * gate_output = nullptr; + const Value * up_output = nullptr; + const Value * glu_output = nullptr; + const CommandPlanMoeRoutingBundle * routing_bundle = nullptr; + const GraphNode * gate_node = nullptr; + const GraphNode * up_node = nullptr; + const GraphNode * glu_node = nullptr; + int64_t token_count = 0; + + bool matched() const { + return gate_weight != nullptr && up_weight != nullptr && input != nullptr && route_ids != nullptr && + gate_output != nullptr && up_output != nullptr && glu_output != nullptr && routing_bundle != nullptr && + gate_node != nullptr && up_node != nullptr && glu_node != nullptr && token_count > 0; + } +}; + +struct DecodeRoutedGateUpMatch { + const Value * gate_weight = nullptr; + const Value * up_weight = nullptr; + const Value * input = nullptr; + const Value * route_ids = nullptr; + const Value * glu_output = nullptr; + const CommandPlanAlternateValue * input_alternate = nullptr; + const GraphNode * down_node = nullptr; + const Value * down_weight = nullptr; + const GraphNode * gate_node = nullptr; + const GraphNode * up_node = nullptr; + const GraphNode * glu_node = nullptr; + int64_t token_count = 0; + int64_t route_stride = 0; + bool publish_q8 = false; + + bool matched() const { + return gate_weight != nullptr && up_weight != nullptr && input != nullptr && route_ids != nullptr && + glu_output != nullptr && input_alternate != nullptr && down_node != nullptr && down_weight != nullptr && + gate_node != nullptr && up_node != nullptr && glu_node != nullptr && token_count == 1 && + route_stride >= kRoutedFfnRouteCount; + } +}; + +struct RoutedDownMatch { + const Value * input_graph_value = nullptr; + const CommandPlanAlternateValue * input_alternate = nullptr; + const Value * weight = nullptr; + const Value * output = nullptr; + const Value * route_ids = nullptr; + const CommandPlanMoeRoutingBundle * routing_bundle = nullptr; + KernelCatalogRef kernel = {}; + int64_t token_count = 0; + + bool matched() const { + return input_graph_value != nullptr && input_alternate != nullptr && weight != nullptr && output != nullptr && + route_ids != nullptr && routing_bundle != nullptr && kernel.id != kUncatalogedKernelId && + token_count > 0; + } +}; + +struct WeightedReduceNextRmsNormMatch { + const GraphNode * rms_node = nullptr; + const GraphNode * mul_node = nullptr; + const Value * norm_weight = nullptr; + const Value * output = nullptr; + + bool matched() const { + return rms_node != nullptr && mul_node != nullptr && norm_weight != nullptr && output != nullptr; + } +}; + +struct WeightedReduceMatch { + const Value * route_weights = nullptr; + const Value * routed_output = nullptr; + const CommandPlanAlternateValue * routed_alternate = nullptr; + const Value * residual_input = nullptr; + const Value * output = nullptr; + const GraphNode * weighted_node = nullptr; + std::vector views; + std::vector reductions; + const GraphNode * residual = nullptr; + WeightedReduceNextRmsNormMatch next_rmsnorm; + int64_t token_count = 0; + + bool topology_matched() const { + return route_weights != nullptr && routed_output != nullptr && residual_input != nullptr && output != nullptr && + weighted_node != nullptr && !views.empty() && residual != nullptr && token_count > 0; + } + + bool matched() const { return topology_matched() && routed_alternate != nullptr; } +}; + +struct DecodeRoutedDownMatch { + const Value * input_graph_value = nullptr; + const CommandPlanAlternateValue * input_alternate = nullptr; + const Value * weight = nullptr; + const Value * output = nullptr; + const Value * route_ids = nullptr; + WeightedReduceMatch reduce; + KernelCatalogRef kernel = {}; + int64_t token_count = 0; + int64_t route_stride = 0; + bool input_is_q8 = false; + + bool matched() const { + return input_graph_value != nullptr && weight != nullptr && output != nullptr && route_ids != nullptr && + reduce.topology_matched() && reduce.next_rmsnorm.matched() && kernel.id != kUncatalogedKernelId && + token_count == 1 && route_stride >= kRoutedFfnRouteCount && (!input_is_q8 || input_alternate != nullptr); + } +}; + +static bool bundle_matches_moe_routing(const CommandPlanMoeRoutingBundle & bundle, + ValueId route_ids, + int64_t token_count) { + return bundle.route_ids == route_ids && bundle.route_weights.value >= 0 && bundle.expert_table.value >= 0 && + bundle.partition_table.value >= 0 && bundle.expert_table_byte_count == expert_table_size(token_count) && + bundle.partition_table_byte_count == partition_table_size(token_count) && + bundle.token_count == token_count && bundle.route_count == kRoutedFfnRouteCount && + bundle.expert_count == kRoutedFfnExpertCount && bundle.route_stride >= kRoutedFfnRouteCount; +} + +static bool match_same_route_projection(const Graph & graph, + const GraphNode & node, + ValueId expected_input, + ValueId expected_route_ids, + int64_t token_count, + const Value *& weight, + const Value *& output) { + if (node.op != GGML_OP_MUL_MAT_ID || node.inputs.size() != 3 || node.inputs[1] != expected_input || + node.inputs[2] != expected_route_ids) { + return false; + } + weight = graph_value(graph, node.inputs[0]); + output = graph_value(graph, node.output); + return weight != nullptr && output != nullptr && is_routed_ffn_gate_up_weight(*weight) && + is_routed_ffn_projection_output(*output, token_count); +} + +static bool match_same_fast_route_projection(const Graph & graph, + const GraphNode & node, + ValueId expected_input, + ValueId expected_route_ids, + int64_t token_count, + const Value *& weight, + const Value *& output) { + if (node.op != GGML_OP_MUL_MAT_ID || node.inputs.size() != 3 || node.inputs[1] != expected_input || + node.inputs[2] != expected_route_ids) { + return false; + } + weight = graph_value(graph, node.inputs[0]); + output = graph_value(graph, node.output); + return weight != nullptr && output != nullptr && is_fast_routed_ffn_gate_up_weight(*weight) && + is_routed_ffn_projection_output(*output, token_count); +} + +static RoutedDownMatch match_routed_ffn_down_grouped(const DispatchMatchContext & context) { + RoutedDownMatch match; + const GraphNode * root = context.root_node; + if (root == nullptr || root->op != GGML_OP_MUL_MAT_ID || root->inputs.size() != 3 || !context.graph.has_index()) { + return match; + } + + const Value * weight = graph_value(context.graph, root->inputs[0]); + const Value * input = graph_value(context.graph, root->inputs[1]); + const Value * route_ids = graph_value(context.graph, root->inputs[2]); + const Value * root_output = graph_value(context.graph, root->output); + if (weight == nullptr || input == nullptr || route_ids == nullptr || root_output == nullptr || + !is_routed_ffn_down_weight(*weight) || !is_routed_ffn_projection_output(*input, input->ne[2]) || + route_ids->type != GGML_TYPE_I32 || !is_shape(*route_ids, kRoutedFfnRouteCount, input->ne[2], 1, 1)) { + return {}; + } + + const int64_t token_count = input->ne[2]; + if (!is_llm_supported_query_length(kRoutedFfnProfile, token_count) || + !is_routed_ffn_down_output(*root_output, token_count)) { + return {}; + } + + const CommandPlanMoeRoutingBundle * routing_bundle = context.plan.metadata.find_moe_routing_bundle(route_ids->id); + if (routing_bundle == nullptr || !bundle_matches_moe_routing(*routing_bundle, route_ids->id, token_count)) { + return {}; + } + + const CommandPlanAlternateValue * input_alternate = + find_alternate_value(context.plan, input->id, GGML_TYPE_F16, f16_gate_up_output_size(token_count)); + if (input_alternate == nullptr) { + return {}; + } + + match.input_graph_value = input; + match.input_alternate = input_alternate; + match.weight = weight; + match.output = root_output; + match.route_ids = route_ids; + match.routing_bundle = routing_bundle; + match.kernel = kCommonMulMatIdF16F16WmmaKernel; + match.token_count = token_count; + return match; +} + +static RoutedGateUpMatch match_routed_ffn_gate_up_swiglu(const DispatchMatchContext & context) { + RoutedGateUpMatch match; + const GraphNode * root = context.root_node; + if (root == nullptr || root->op != GGML_OP_MUL_MAT_ID || root->inputs.size() != 3 || !context.graph.has_index()) { + return match; + } + + const Value * root_weight = graph_value(context.graph, root->inputs[0]); + const Value * input = graph_value(context.graph, root->inputs[1]); + const Value * route_ids = graph_value(context.graph, root->inputs[2]); + const Value * root_output = graph_value(context.graph, root->output); + if (root_weight == nullptr || input == nullptr || route_ids == nullptr || root_output == nullptr || + !is_fast_routed_ffn_gate_up_weight(*root_weight) || input->type != GGML_TYPE_F32 || !input->contiguous || + !is_shape(*input, kRoutedFfnInputSize, 1, input->ne[2], 1) || route_ids->type != GGML_TYPE_I32 || + !is_shape(*route_ids, kRoutedFfnRouteCount, input->ne[2], 1, 1)) { + return {}; + } + + const int64_t token_count = input->ne[2]; + if (!is_llm_supported_query_length(kRoutedFfnProfile, token_count) || + !is_routed_ffn_projection_output(*root_output, token_count)) { + return {}; + } + + const CommandPlanMoeRoutingBundle * routing_bundle = context.plan.metadata.find_moe_routing_bundle(route_ids->id); + if (routing_bundle == nullptr || !bundle_matches_moe_routing(*routing_bundle, route_ids->id, token_count)) { + return {}; + } + + const GraphNode * glu_node = find_consumer_with_op(context.graph, root->output, GGML_OP_GLU); + if (glu_node == nullptr || glu_node->inputs.size() != 2) { + return {}; + } + if (!is_swiglu_params(glu_node->params)) { + return {}; + } + + const GraphNode * gate_node = producer_with_op(context.graph, glu_node->inputs[0], GGML_OP_MUL_MAT_ID); + const GraphNode * up_node = producer_with_op(context.graph, glu_node->inputs[1], GGML_OP_MUL_MAT_ID); + if (gate_node == nullptr || up_node == nullptr || gate_node == up_node || (gate_node != root && up_node != root)) { + return {}; + } + + const Value * gate_weight = nullptr; + const Value * gate_output = nullptr; + const Value * up_weight = nullptr; + const Value * up_output = nullptr; + if (!match_same_fast_route_projection(context.graph, *gate_node, input->id, route_ids->id, token_count, gate_weight, + gate_output) || + !match_same_fast_route_projection(context.graph, *up_node, input->id, route_ids->id, token_count, up_weight, + up_output)) { + return {}; + } + if (!same_shape(*gate_output, *up_output)) { + return {}; + } + + const Value * glu_output = graph_value(context.graph, glu_node->output); + if (glu_output == nullptr || glu_output->kind != ValueKind::Transient || + !is_routed_ffn_projection_output(*glu_output, token_count)) { + return {}; + } + + match.gate_weight = gate_weight; + match.up_weight = up_weight; + match.input = input; + match.route_ids = route_ids; + match.gate_output = gate_output; + match.up_output = up_output; + match.glu_output = glu_output; + match.routing_bundle = routing_bundle; + match.gate_node = gate_node; + match.up_node = up_node; + match.glu_node = glu_node; + match.token_count = token_count; + return match; +} + +static DecodeRoutedGateUpMatch match_decode_routed_ffn_gate_up_swiglu(const DispatchMatchContext & context) { + DecodeRoutedGateUpMatch match; + const GraphNode * root = context.root_node; + if (root == nullptr || root->op != GGML_OP_MUL_MAT_ID || root->inputs.size() != 3 || !context.graph.has_index()) { + return match; + } + + const Value * root_weight = graph_value(context.graph, root->inputs[0]); + const Value * input = graph_value(context.graph, root->inputs[1]); + const Value * route_ids = graph_value(context.graph, root->inputs[2]); + const Value * root_output = graph_value(context.graph, root->output); + if (root_weight == nullptr || input == nullptr || route_ids == nullptr || root_output == nullptr || + !is_routed_ffn_gate_up_weight(*root_weight) || input->type != GGML_TYPE_F32 || !input->contiguous || + !is_shape(*input, kRoutedFfnInputSize, 1, 1, 1) || route_ids->type != GGML_TYPE_I32 || + !is_shape(*route_ids, kRoutedFfnRouteCount, 1, 1, 1) || !is_routed_ffn_projection_output(*root_output, 1)) { + return {}; + } + + const int64_t route_stride = static_cast(route_ids->nb[1] / sizeof(int32_t)); + const CommandPlanAlternateValue * input_alternate = find_alternate_value( + context.graph, context.plan, input->id, GGML_TYPE_Q8_1, q8_1_x4_byte_count(1, kRoutedFfnInputSize)); + if (input_alternate == nullptr) { + return {}; + } + + const GraphNode * glu_node = find_consumer_with_op(context.graph, root->output, GGML_OP_GLU); + if (glu_node == nullptr || glu_node->inputs.size() != 2) { + return {}; + } + if (!is_swiglu_params(glu_node->params)) { + return {}; + } + + const GraphNode * gate_node = producer_with_op(context.graph, glu_node->inputs[0], GGML_OP_MUL_MAT_ID); + const GraphNode * up_node = producer_with_op(context.graph, glu_node->inputs[1], GGML_OP_MUL_MAT_ID); + if (gate_node == nullptr || up_node == nullptr || gate_node == up_node || (gate_node != root && up_node != root)) { + return {}; + } + + const Value * gate_weight = nullptr; + const Value * gate_output = nullptr; + const Value * up_weight = nullptr; + const Value * up_output = nullptr; + if (!match_same_route_projection(context.graph, *gate_node, input->id, route_ids->id, 1, gate_weight, + gate_output) || + !match_same_route_projection(context.graph, *up_node, input->id, route_ids->id, 1, up_weight, up_output) || + !same_shape(*gate_output, *up_output)) { + return {}; + } + + const Value * glu_output = graph_value(context.graph, glu_node->output); + if (glu_output == nullptr || glu_output->kind != ValueKind::Transient || + !is_routed_ffn_projection_output(*glu_output, 1)) { + return {}; + } + + const GraphNode * down_node = find_single_consumer_with_op(context.graph, glu_output->id, GGML_OP_MUL_MAT_ID); + const Value * down_weight = down_node == nullptr || down_node->inputs.size() != 3 ? + nullptr : + graph_value(context.graph, down_node->inputs[0]); + if (down_node == nullptr || down_weight == nullptr || !is_routed_ffn_down_weight(*down_weight)) { + return {}; + } + + match.gate_weight = gate_weight; + match.up_weight = up_weight; + match.input = input; + match.route_ids = route_ids; + match.glu_output = glu_output; + match.input_alternate = input_alternate; + match.down_node = down_node; + match.down_weight = down_weight; + match.gate_node = gate_node; + match.up_node = up_node; + match.glu_node = glu_node; + match.token_count = 1; + match.route_stride = route_stride; + match.publish_q8 = down_weight->type == GGML_TYPE_Q4_K; + return match; +} + +static WeightedReduceNextRmsNormMatch match_qwen_weighted_reduce_next_rmsnorm(const DispatchMatchContext & context, + const Value & residual) { + WeightedReduceNextRmsNormMatch match; + const GraphNode * rms_node = find_single_consumer_with_op(context.graph, residual.id, GGML_OP_RMS_NORM); + if (rms_node == nullptr || rms_node->inputs.size() != 1) { + return match; + } + const RmsNormParams * rms_params = op_params_as(rms_node->params); + if (rms_params == nullptr || !is_profile_rms_norm_epsilon(rms_params->eps)) { + return {}; + } + + const Value * rms = graph_value(context.graph, rms_node->output); + if (rms == nullptr || rms->type != GGML_TYPE_F32 || !same_shape(*rms, residual)) { + return {}; + } + + const GraphNode * mul_node = find_single_consumer_with_op(context.graph, rms_node->output, GGML_OP_MUL); + if (mul_node == nullptr || mul_node->inputs.size() != 2) { + return {}; + } + + const Value * norm_weight = nullptr; + for (ValueId input : mul_node->inputs) { + if (input != rms_node->output) { + norm_weight = graph_value(context.graph, input); + } + } + const Value * output = graph_value(context.graph, mul_node->output); + if (norm_weight == nullptr || output == nullptr || norm_weight->type != GGML_TYPE_F32 || + output->type != GGML_TYPE_F32 || !norm_weight->contiguous || !output->contiguous || + !is_shape(*norm_weight, kRoutedFfnInputSize, 1, 1, 1) || !same_shape(*output, residual)) { + return {}; + } + + match.rms_node = rms_node; + match.mul_node = mul_node; + match.norm_weight = norm_weight; + match.output = output; + return match; +} + +static bool append_node_if_uncovered(const DispatchMatchContext & context, + const GraphNode * node, + std::vector & nodes) { + size_t index = 0; + if (node == nullptr || !context.graph.index().node_index(node, index) || index >= context.covered_nodes.size() || + context.covered_nodes[index]) { + return false; + } + for (const GraphNode * existing : nodes) { + if (existing == node) { + return true; + } + } + nodes.push_back(node); + return true; +} + +static bool node_is_covered(const DispatchMatchContext & context, const GraphNode * node) { + size_t index = 0; + return node != nullptr && context.graph.index().node_index(node, index) && index < context.covered_nodes.size() && + context.covered_nodes[index]; +} + +static bool has_uncovered_q6_k_final_projection_consumer(const DispatchMatchContext & context, ValueId input) { + for (const GraphNode * consumer : consumers_with_op_through_layout_aliases(context.graph, input, GGML_OP_MUL_MAT)) { + if (consumer == nullptr || node_is_covered(context, consumer) || consumer->inputs.size() < 2 || + !node_has_input_or_alias(context.graph, *consumer, input)) { + continue; + } + + const Value * weight = graph_value(context.graph, consumer->inputs[0]); + const Value * rhs = graph_value(context.graph, consumer->inputs[1]); + const Value * output = graph_value(context.graph, consumer->output); + if (weight == nullptr || rhs == nullptr || output == nullptr || weight->type != GGML_TYPE_Q6_K || + rhs->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || weight->alias_source.value >= 0 || + rhs->ne[0] <= 0 || output->ne[0] <= 0) { + continue; + } + + const int64_t input_size = rhs->ne[0]; + const int64_t output_size = output->ne[0]; + if (input_size % 256 == 0 && output_size % 64 == 0 && output_size >= 16 * input_size) { + return true; + } + } + return false; +} + +static bool residual_input_is_safe_for_in_place(const DispatchMatchContext & context, + const WeightedReduceMatch & match) { + if (match.residual_input == nullptr || match.residual == nullptr) { + return false; + } + // In-place reuse overwrites the residual input's storage, which is only safe for a + // transient (temporary) value; an external value shares a persistent buffer (#95). + if (match.residual_input->kind != ValueKind::Transient) { + return false; + } + for (const GraphNode * consumer : context.graph.index().consumers(match.residual_input->id)) { + if (consumer == match.residual || node_is_covered(context, consumer)) { + continue; + } + return false; + } + return true; +} + +static const Value * find_qwen_route_weights_for_route_ids(const Graph & graph, + ValueId route_ids, + int64_t token_count) { + const GraphNode * get_rows = find_single_consumer_with_op(graph, route_ids, GGML_OP_GET_ROWS); + if (get_rows == nullptr || get_rows->inputs.size() != 2) { + return nullptr; + } + const Value * selected = graph_value(graph, get_rows->output); + if (selected == nullptr || selected->type != GGML_TYPE_F32 || + !is_shape(*selected, 1, kRoutedFfnRouteCount, token_count, 1)) { + return nullptr; + } + const GraphNode * reshape = find_single_consumer_with_op(graph, get_rows->output, GGML_OP_RESHAPE); + const Value * flat_weights = reshape == nullptr ? nullptr : graph_value(graph, reshape->output); + if (flat_weights == nullptr || flat_weights->type != GGML_TYPE_F32 || + !is_shape(*flat_weights, kRoutedFfnRouteCount, token_count, 1, 1)) { + return nullptr; + } + const GraphNode * sum_rows = find_single_consumer_with_op(graph, reshape->output, GGML_OP_SUM_ROWS); + const GraphNode * clamp = + sum_rows == nullptr ? nullptr : find_single_consumer_with_op(graph, sum_rows->output, GGML_OP_CLAMP); + const GraphNode * div = + clamp == nullptr ? nullptr : find_single_consumer_with_op(graph, clamp->output, GGML_OP_DIV); + const GraphNode * output_reshape = + div == nullptr ? nullptr : find_single_consumer_with_op(graph, div->output, GGML_OP_RESHAPE); + const Value * weights = output_reshape == nullptr ? nullptr : graph_value(graph, output_reshape->output); + if (weights == nullptr || weights->type != GGML_TYPE_F32 || + !is_shape(*weights, 1, kRoutedFfnRouteCount, token_count, 1)) { + return nullptr; + } + return weights; +} + +static WeightedReduceMatch match_routed_ffn_down_weighted_reduce_topology(const DispatchMatchContext & context, + const GraphNode * weighted, + const Value * routed_output, + const Value * route_weights) { + WeightedReduceMatch match; + if (weighted == nullptr || weighted->op != GGML_OP_MUL || weighted->inputs.size() != 2 || + routed_output == nullptr || route_weights == nullptr || !context.graph.has_index()) { + return match; + } + if (!node_has_input_or_alias(context.graph, *weighted, routed_output->id) || + !node_has_input_or_alias(context.graph, *weighted, route_weights->id)) { + return {}; + } + const Value * weighted_output = graph_value(context.graph, weighted->output); + if (!is_routed_ffn_down_output(*routed_output, routed_output->ne[2]) || route_weights->type != GGML_TYPE_F32 || + !route_weights->contiguous || !is_shape(*route_weights, 1, kRoutedFfnRouteCount, routed_output->ne[2], 1) || + weighted_output == nullptr || !same_shape(*weighted_output, *routed_output)) { + return {}; + } + + const int64_t token_count = routed_output->ne[2]; + if (!is_llm_supported_query_length(kRoutedFfnProfile, token_count)) { + return {}; + } + + std::vector views = + layout_alias_consumers_with_op(context.graph, weighted->output, GGML_OP_VIEW); + if (views.size() != kRoutedFfnRouteCount) { + return {}; + } + + std::set routed_values; + std::vector owned_views; + for (const GraphNode * view : views) { + const Value * value = view == nullptr ? nullptr : graph_value(context.graph, view->output); + if (value == nullptr || value->type != GGML_TYPE_F32 || + !is_shape(*value, kRoutedFfnInputSize, token_count, 1, 1) || + !append_node_if_uncovered(context, view, owned_views)) { + return {}; + } + routed_values.insert(view->output.value); + } + + std::vector reductions; + bool changed = true; + while (changed) { + changed = false; + const std::set values = routed_values; + for (int32_t value : values) { + for (const GraphNode * add : find_consumers_with_op(context.graph, ValueId(value), GGML_OP_ADD)) { + if (add == nullptr || add->inputs.size() != 2) { + continue; + } + bool already_owned = false; + for (const GraphNode * reduction : reductions) { + if (reduction == add) { + already_owned = true; + break; + } + } + if (already_owned) { + continue; + } + bool all_routed = true; + for (ValueId input : add->inputs) { + all_routed = all_routed && routed_values.count(input.value) != 0; + } + if (!all_routed) { + continue; + } + const Value * output = graph_value(context.graph, add->output); + if (output == nullptr || output->type != GGML_TYPE_F32 || + !is_shape(*output, kRoutedFfnInputSize, token_count, 1, 1) || + !append_node_if_uncovered(context, add, reductions)) { + return {}; + } + routed_values.insert(add->output.value); + changed = true; + } + } + } + + const GraphNode * residual = nullptr; + const Value * residual_input = nullptr; + for (int32_t value : routed_values) { + for (const GraphNode * add : find_consumers_with_op(context.graph, ValueId(value), GGML_OP_ADD)) { + if (add == nullptr || add->inputs.size() != 2) { + continue; + } + bool is_reduction = false; + for (const GraphNode * reduction : reductions) { + if (reduction == add) { + is_reduction = true; + break; + } + } + if (is_reduction) { + continue; + } + int routed_input_count = 0; + const Value * non_routed_input = nullptr; + for (ValueId input : add->inputs) { + if (routed_values.count(input.value) != 0) { + ++routed_input_count; + } else { + non_routed_input = graph_value(context.graph, input); + } + } + if (routed_input_count != 1 || non_routed_input == nullptr || residual != nullptr) { + return {}; + } + residual = add; + residual_input = non_routed_input; + } + } + if (residual == nullptr || reductions.size() + 1 != views.size()) { + return {}; + } + const Value * output = graph_value(context.graph, residual->output); + if (output == nullptr || residual_input == nullptr || output->type != GGML_TYPE_F32 || + residual_input->type != GGML_TYPE_F32 || !is_shape(*output, kRoutedFfnInputSize, token_count, 1, 1) || + !same_shape(*output, *residual_input) || output->byte_count != residual_input->byte_count || + !output->contiguous || !residual_input->contiguous) { + return {}; + } + + match.route_weights = route_weights; + match.routed_output = routed_output; + match.residual_input = residual_input; + match.output = output; + match.weighted_node = weighted; + match.views = std::move(owned_views); + match.reductions = std::move(reductions); + match.residual = residual; + match.next_rmsnorm = match_qwen_weighted_reduce_next_rmsnorm(context, *output); + match.token_count = token_count; + return match; +} + +static WeightedReduceMatch match_routed_ffn_down_weighted_reduce(const DispatchMatchContext & context) { + WeightedReduceMatch match; + const GraphNode * weighted = context.root_node; + if (weighted == nullptr || weighted->op != GGML_OP_MUL || weighted->inputs.size() != 2 || + !context.graph.has_index()) { + return match; + } + + const Value * routed_output = nullptr; + const Value * route_weights = nullptr; + for (ValueId input : weighted->inputs) { + const Value * value = graph_value(context.graph, input); + if (value == nullptr) { + return {}; + } + if (is_routed_ffn_down_output(*value, value->ne[2])) { + routed_output = value; + } else if (value->type == GGML_TYPE_F32 && value->contiguous && + is_shape(*value, 1, kRoutedFfnRouteCount, value->ne[2], 1)) { + route_weights = value; + } + } + match = match_routed_ffn_down_weighted_reduce_topology(context, weighted, routed_output, route_weights); + if (!match.topology_matched()) { + return {}; + } + const int64_t token_count = match.token_count; + const CommandPlanAlternateValue * routed_alternate = + find_alternate_value(context.plan, routed_output->id, GGML_TYPE_F16, f16_routed_down_output_size(token_count)); + if (routed_alternate == nullptr) { + return {}; + } + + bool known_route_weights = false; + for (const CommandPlanMoeRoutingBundle & bundle : context.plan.metadata.moe_routing_bundles()) { + if (bundle.route_weights == route_weights->id && + bundle_matches_moe_routing(bundle, bundle.route_ids, token_count)) { + known_route_weights = true; + break; + } + } + if (!known_route_weights) { + return {}; + } + + match.routed_alternate = routed_alternate; + return match; +} + +static bool match_routed_ffn_gate_up_swiglu_f16_wmma_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const RoutedGateUpMatch match = match_routed_ffn_gate_up_swiglu(context); + if (!match.matched()) { + return false; + } + if (match.gate_weight->type != match.up_weight->type) { + return false; + } + + CommonMulMatWeightFormat weight_format; + if (!common_mul_mat_format_for_type(match.gate_weight->type, weight_format)) { + return false; + } + const std::string weight_format_config = to_config_value(common_mul_mat_format_config_value(weight_format)); + const std::string config_prefix = "ggml.mul_mat_id_swiglu_f16_f16"; + + const ValueId f16_output(context.next_plan_value.value); + const size_t f16_output_bytes = f16_gate_up_output_size(match.token_count); + dispatch_match.transients.push_back( + { f16_output, kRoutedFfnF16GateUpOutputName, f16_output_bytes, kRoutedFfnPlanTransientAlignment }); + Status metadata_status; + if (!dispatch_match.metadata.append_alternate_value( + { match.glu_output->id, f16_output, GGML_TYPE_F16, f16_output_bytes, kRoutedFfnF16GateUpOutputName }, + metadata_status)) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kCommonRoutedGateUpSwiGLUF16WmmaKernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.compile_parameters.emplace(config_prefix + ".input_size", to_config_value(kRoutedFfnInputSize)); + dispatch.kernel.compile_parameters.emplace(config_prefix + ".expert_count", to_config_value(kRoutedFfnExpertCount)); + dispatch.kernel.compile_parameters.emplace(config_prefix + ".route_count", to_config_value(kRoutedFfnRouteCount)); + dispatch.kernel.compile_parameters.emplace(config_prefix + ".output_size", + to_config_value(kRoutedFfnExpertHiddenSize)); + dispatch.kernel.compile_parameters.emplace(config_prefix + ".gate_weight_format", weight_format_config); + dispatch.kernel.compile_parameters.emplace(config_prefix + ".up_weight_format", weight_format_config); + dispatch.kernel.compile_parameters.emplace(config_prefix + ".descriptor_expert_mask", "127"); + dispatch.kernel.compile_parameters.emplace(config_prefix + ".descriptor_partition_shift", "7"); + dispatch.kernel.compile_parameters.emplace(config_prefix + ".descriptor_row_count_shift", "13"); + + dispatch.bindings.push_back({ match.input->id, 0, match.input->byte_count }); + dispatch.bindings.push_back( + { match.routing_bundle->expert_table, 0, match.routing_bundle->expert_table_byte_count }); + dispatch.bindings.push_back( + { match.routing_bundle->partition_table, 0, match.routing_bundle->partition_table_byte_count }); + dispatch.bindings.push_back({ match.gate_weight->id, 0, match.gate_weight->byte_count }); + dispatch.bindings.push_back({ match.up_weight->id, 0, match.up_weight->byte_count }); + dispatch.bindings.push_back({ f16_output, 0, f16_output_bytes }); + + if (!append_covered_node(context, match.gate_node, dispatch_match) || + !append_covered_node(context, match.up_node, dispatch_match) || + !append_covered_node(context, match.glu_node, dispatch_match)) { + return false; + } + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_decode_routed_ffn_gate_up_swiglu_q4k_q8_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const DecodeRoutedGateUpMatch match = match_decode_routed_ffn_gate_up_swiglu(context); + if (!match.matched()) { + return false; + } + + const size_t q8_output_bytes = + q8_1_x4_byte_count(match.token_count * kRoutedFfnRouteCount, kRoutedFfnExpertHiddenSize); + const ValueId q8_output = match.publish_q8 ? ValueId(context.next_plan_value.value) : ValueId(); + const ValueId completion_counters = match.publish_q8 ? ValueId(context.next_plan_value.value + 1) : ValueId(); + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(match.publish_q8 ? kQwenRoutedGateUpSwiGLUQ4KQ8NextQ8Kernel : + kQwenRoutedGateUpSwiGLUQ4KQ8Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.integer_parameters.emplace("route_count", kRoutedFfnRouteCount); + dispatch.kernel.integer_parameters.emplace("route_stride", match.route_stride); + dispatch.kernel.integer_parameters.emplace("expert_count", kRoutedFfnExpertCount); + dispatch.kernel.integer_parameters.emplace("output_size", kRoutedFfnExpertHiddenSize); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.routed_gate_up.input_size", + to_config_value(kRoutedFfnInputSize)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.routed_gate_up.expert_count", + to_config_value(kRoutedFfnExpertCount)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.routed_gate_up.route_count", + to_config_value(kRoutedFfnRouteCount)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.routed_gate_up.output_size", + to_config_value(kRoutedFfnExpertHiddenSize)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.workload.token_capacity", to_config_value(match.token_count)); + + const size_t route_id_length = static_cast(match.token_count * match.route_stride) * sizeof(int32_t); + dispatch.bindings.push_back({ match.input_alternate->alternate_value, 0, match.input_alternate->byte_count }); + dispatch.bindings.push_back({ match.route_ids->id, 0, route_id_length }); + dispatch.bindings.push_back({ match.gate_weight->id, 0, match.gate_weight->byte_count }); + dispatch.bindings.push_back({ match.up_weight->id, 0, match.up_weight->byte_count }); + dispatch.bindings.push_back({ match.glu_output->id, 0, match.glu_output->byte_count }); + if (match.publish_q8) { + dispatch.bindings.push_back( + { completion_counters, 0, gate_up_completion_counter_count(match.token_count) * sizeof(int32_t) }); + dispatch.bindings.push_back({ q8_output, 0, q8_output_bytes }); + dispatch_match.completion_counter_requests.push_back({ + completion_counters, + "qwen.decode.moe.gate_up_completion_counters", + gate_up_completion_counter_count(match.token_count), + }); + dispatch_match.transients.push_back( + { q8_output, kRoutedFfnQ8GateUpOutputName, q8_output_bytes, kRoutedFfnPlanTransientAlignment }); + Status metadata_status; + if (!dispatch_match.metadata.append_alternate_value( + { match.glu_output->id, q8_output, GGML_TYPE_Q8_1, q8_output_bytes, kRoutedFfnQ8GateUpOutputName }, + metadata_status)) { + dispatch_match.status.append(metadata_status); + return false; + } + } + + if (!append_covered_node(context, match.gate_node, dispatch_match) || + !append_covered_node(context, match.up_node, dispatch_match) || + !append_covered_node(context, match.glu_node, dispatch_match)) { + return false; + } + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static DecodeRoutedDownMatch match_decode_routed_ffn_down_next_q8(const DispatchMatchContext & context) { + DecodeRoutedDownMatch match; + const GraphNode * root = context.root_node; + if (root == nullptr || root->op != GGML_OP_MUL_MAT_ID || root->inputs.size() != 3 || !context.graph.has_index()) { + return match; + } + + const Value * weight = graph_value(context.graph, root->inputs[0]); + const Value * input = graph_value(context.graph, root->inputs[1]); + const Value * route_ids = graph_value(context.graph, root->inputs[2]); + const Value * root_output = graph_value(context.graph, root->output); + if (weight == nullptr || input == nullptr || route_ids == nullptr || root_output == nullptr || + !is_routed_ffn_down_weight(*weight) || !is_routed_ffn_projection_output(*input, 1) || + !is_routed_ffn_down_output(*root_output, 1) || route_ids->type != GGML_TYPE_I32 || + route_ids->nb[0] != sizeof(int32_t) || route_ids->nb[1] % sizeof(int32_t) != 0 || + !is_shape(*route_ids, kRoutedFfnRouteCount, 1, 1, 1)) { + return {}; + } + + const int64_t route_stride = static_cast(route_ids->nb[1] / sizeof(int32_t)); + const GraphNode * glu_node = producer_with_op(context.graph, input->id, GGML_OP_GLU); + const GraphNode * gate_node = glu_node == nullptr || glu_node->inputs.size() != 2 ? + nullptr : + producer_with_op(context.graph, glu_node->inputs[0], GGML_OP_MUL_MAT_ID); + const Value * gate_input = gate_node == nullptr || gate_node->inputs.size() != 3 ? + nullptr : + graph_value(context.graph, gate_node->inputs[1]); + const CommandPlanAlternateValue * gate_input_q8 = + gate_input == nullptr ? nullptr : + find_alternate_value(context.graph, context.plan, gate_input->id, GGML_TYPE_Q8_1, + q8_1_x4_byte_count(1, kRoutedFfnInputSize)); + if (gate_input_q8 == nullptr) { + return {}; + } + + const Value * route_weights = find_qwen_route_weights_for_route_ids(context.graph, route_ids->id, 1); + const GraphNode * weighted = + find_single_consumer_with_op_through_layout_aliases(context.graph, root_output->id, GGML_OP_MUL); + WeightedReduceMatch reduce = + match_routed_ffn_down_weighted_reduce_topology(context, weighted, root_output, route_weights); + if (!reduce.topology_matched() || !reduce.next_rmsnorm.matched() || + // The in-place reduce reuses the residual input's storage for the output, which is only safe + // when both are transient (an external value shares a persistent buffer) (#95). + reduce.output->kind != ValueKind::Transient || + !residual_input_is_safe_for_in_place(context, reduce)) { + return {}; + } + if (has_uncovered_q6_k_final_projection_consumer(context, reduce.next_rmsnorm.output->id)) { + return {}; + } + + const bool input_is_q8 = weight->type == GGML_TYPE_Q4_K; + const CommandPlanAlternateValue * input_alternate = + input_is_q8 ? find_alternate_value(context.graph, context.plan, input->id, GGML_TYPE_Q8_1, + q8_1_x4_byte_count(kRoutedFfnRouteCount, kRoutedFfnExpertHiddenSize)) : + nullptr; + if (input_is_q8 && input_alternate == nullptr) { + return {}; + } + + match.input_graph_value = input; + match.input_alternate = input_alternate; + match.weight = weight; + match.output = reduce.output; + match.route_ids = route_ids; + match.reduce = std::move(reduce); + match.kernel = + weight->type == GGML_TYPE_Q4_K ? kQwenRoutedDownQ4KQ8NextQ8Kernel : kQwenRoutedDownQ6KF32Wave64NextQ8Kernel; + match.token_count = 1; + match.route_stride = route_stride; + match.input_is_q8 = input_is_q8; + return match; +} + +static bool match_decode_routed_ffn_down_next_q8_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const DecodeRoutedDownMatch match = match_decode_routed_ffn_down_next_q8(context); + if (!match.matched()) { + return false; + } + + const ValueId completion_counter_value(context.next_plan_value.value); + const ValueId q8_output(context.next_plan_value.value + 1); + const size_t q8_output_bytes = q8_1_x4_byte_count(match.token_count, kRoutedFfnInputSize); + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(match.kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + dispatch.kernel.integer_parameters.emplace("input_size", kRoutedFfnExpertHiddenSize); + dispatch.kernel.integer_parameters.emplace("route_count", kRoutedFfnRouteCount); + dispatch.kernel.integer_parameters.emplace("route_id_stride", match.route_stride); + dispatch.kernel.integer_parameters.emplace("expert_count", kRoutedFfnExpertCount); + dispatch.kernel.integer_parameters.emplace("output_size", kRoutedFfnInputSize); + add_routed_down_compile_parameters(dispatch, match.token_count); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.model.hidden_size", to_config_value(kRoutedFfnInputSize)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.model.rms_epsilon", "0.000001"); + + const size_t route_id_length = static_cast(match.token_count * match.route_stride) * sizeof(int32_t); + if (match.input_is_q8) { + dispatch.bindings.push_back({ match.input_alternate->alternate_value, 0, match.input_alternate->byte_count }); + } else { + dispatch.bindings.push_back({ match.input_graph_value->id, 0, match.input_graph_value->byte_count }); + } + dispatch.bindings.push_back({ match.route_ids->id, 0, route_id_length }); + dispatch.bindings.push_back({ match.reduce.route_weights->id, 0, match.reduce.route_weights->byte_count }); + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + dispatch.bindings.push_back( + { match.reduce.next_rmsnorm.norm_weight->id, 0, match.reduce.next_rmsnorm.norm_weight->byte_count }); + dispatch.bindings.push_back({ completion_counter_value, 0, sizeof(int32_t) }); + dispatch.bindings.push_back({ q8_output, 0, q8_output_bytes }); + + dispatch_match.value_aliases.push_back({ match.reduce.residual_input->id, match.output->id }); + dispatch_match.completion_counter_requests.push_back({ + completion_counter_value, + "qwen.decode.moe.routed_down_completion_counter", + 1, + }); + dispatch_match.transients.push_back( + { q8_output, kRoutedFfnQ8HiddenOutputName, q8_output_bytes, kRoutedFfnPlanTransientAlignment }); + Status metadata_status; + if (!dispatch_match.metadata.append_alternate_value( + { match.reduce.next_rmsnorm.output->id, q8_output, GGML_TYPE_Q8_1, q8_output_bytes, + kRoutedFfnQ8HiddenOutputName }, + metadata_status)) { + dispatch_match.status.append(metadata_status); + return false; + } + + if (!append_covered_node(context, context.root_node, dispatch_match) || + !append_covered_node(context, match.reduce.weighted_node, dispatch_match)) { + return false; + } + for (const GraphNode * view : match.reduce.views) { + if (!append_covered_node(context, view, dispatch_match)) { + return false; + } + } + for (const GraphNode * reduction : match.reduce.reductions) { + if (!append_covered_node(context, reduction, dispatch_match)) { + return false; + } + } + if (!append_covered_node(context, match.reduce.residual, dispatch_match) || + !append_covered_node(context, match.reduce.next_rmsnorm.rms_node, dispatch_match) || + !append_covered_node(context, match.reduce.next_rmsnorm.mul_node, dispatch_match)) { + return false; + } + + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool build_routed_ffn_down_grouped_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match, + ggml_type expected_weight_type) { + const RoutedDownMatch match = match_routed_ffn_down_grouped(context); + if (!match.matched() || match.weight->type != expected_weight_type) { + return false; + } + + const ValueId f16_output(context.next_plan_value.value); + const size_t f16_output_bytes = f16_routed_down_output_size(match.token_count); + dispatch_match.transients.push_back( + { f16_output, kRoutedFfnF16RoutedDownOutputName, f16_output_bytes, kRoutedFfnPlanTransientAlignment }); + Status metadata_status; + if (!dispatch_match.metadata.append_alternate_value( + { match.output->id, f16_output, GGML_TYPE_F16, f16_output_bytes, kRoutedFfnF16RoutedDownOutputName }, + metadata_status)) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(match.kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + add_common_routed_down_compile_parameters(dispatch, match.token_count, match.weight->type); + dispatch.bindings.push_back({ match.input_alternate->alternate_value, 0, match.input_alternate->byte_count }); + dispatch.bindings.push_back( + { match.routing_bundle->expert_table, 0, match.routing_bundle->expert_table_byte_count }); + dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count }); + dispatch.bindings.push_back({ f16_output, 0, f16_output_bytes }); + + dispatch_match.covered_nodes.push_back(context.root_index); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_routed_ffn_down_weighted_reduce_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + const WeightedReduceMatch match = match_routed_ffn_down_weighted_reduce(context); + if (!match.matched()) { + return false; + } + const bool use_next_rmsnorm = match.next_rmsnorm.matched() && match.output->kind == ValueKind::Transient && + residual_input_is_safe_for_in_place(context, match); + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(use_next_rmsnorm ? kQwenRoutedDownWeightedReduceNextRmsNormF32Kernel : + kQwenRoutedDownWeightedReduceF16F32Kernel); + dispatch.kernel.integer_parameters.emplace("token_count", match.token_count); + add_routed_down_compile_parameters(dispatch, match.token_count); + if (use_next_rmsnorm) { + dispatch_match.value_aliases.push_back({ match.residual_input->id, match.output->id }); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.model.hidden_size", to_config_value(kRoutedFfnInputSize)); + dispatch.kernel.compile_parameters.emplace("qwen3_moe.model.rms_epsilon", "0.000001"); + dispatch.bindings.push_back({ match.route_weights->id, 0, match.route_weights->byte_count }); + dispatch.bindings.push_back({ match.routed_alternate->alternate_value, 0, match.routed_alternate->byte_count }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + dispatch.bindings.push_back( + { match.next_rmsnorm.norm_weight->id, 0, match.next_rmsnorm.norm_weight->byte_count }); + dispatch.bindings.push_back({ match.next_rmsnorm.output->id, 0, match.next_rmsnorm.output->byte_count }); + } else { + dispatch.bindings.push_back({ match.route_weights->id, 0, match.route_weights->byte_count }); + dispatch.bindings.push_back({ match.routed_alternate->alternate_value, 0, match.routed_alternate->byte_count }); + dispatch.bindings.push_back({ match.residual_input->id, 0, match.residual_input->byte_count }); + dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count }); + } + + if (!append_covered_node(context, match.weighted_node, dispatch_match)) { + return false; + } + for (const GraphNode * view : match.views) { + if (!append_covered_node(context, view, dispatch_match)) { + return false; + } + } + for (const GraphNode * reduction : match.reductions) { + if (!append_covered_node(context, reduction, dispatch_match)) { + return false; + } + } + if (!append_covered_node(context, match.residual, dispatch_match)) { + return false; + } + if (use_next_rmsnorm && (!append_covered_node(context, match.next_rmsnorm.rms_node, dispatch_match) || + !append_covered_node(context, match.next_rmsnorm.mul_node, dispatch_match))) { + return false; + } + + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_routed_ffn_down_q4k_f16_wmma_grouped_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + return build_routed_ffn_down_grouped_dispatch(context, dispatch_match, GGML_TYPE_Q4_K); +} + +static bool match_routed_ffn_down_q6k_f16_wmma_grouped_dispatch(const DispatchMatchContext & context, + DispatchMatch & dispatch_match) { + return build_routed_ffn_down_grouped_dispatch(context, dispatch_match, GGML_TYPE_Q6_K); +} + +} // namespace + +void register_routed_ffn_dispatches(DispatchRegistryBuilder & registry) { + registry.add({ + "llm.routed_ffn.decode_gate_up_swiglu_q4k_q8", + GGML_OP_MUL_MAT_ID, + DispatchMatchKind::Fused, + 1200, + DispatchSource::Llm, + match_decode_routed_ffn_gate_up_swiglu_q4k_q8_dispatch, + }); + registry.add({ + "llm.routed_ffn.decode_down_next_q8", + GGML_OP_MUL_MAT_ID, + DispatchMatchKind::Fused, + 1150, + DispatchSource::Llm, + match_decode_routed_ffn_down_next_q8_dispatch, + }); + registry.add({ + "llm.routed_ffn.gate_up_swiglu_f16_wmma", + GGML_OP_MUL_MAT_ID, + DispatchMatchKind::Fused, + 1000, + DispatchSource::Llm, + match_routed_ffn_gate_up_swiglu_f16_wmma_dispatch, + }); + registry.add({ + "llm.routed_ffn.down_q4k_f16_wmma_grouped", + GGML_OP_MUL_MAT_ID, + DispatchMatchKind::Fused, + 900, + DispatchSource::Llm, + match_routed_ffn_down_q4k_f16_wmma_grouped_dispatch, + }); + registry.add({ + "llm.routed_ffn.down_q6k_f16_wmma_grouped", + GGML_OP_MUL_MAT_ID, + DispatchMatchKind::Fused, + 900, + DispatchSource::Llm, + match_routed_ffn_down_q6k_f16_wmma_grouped_dispatch, + }); + registry.add({ + "llm.routed_ffn.down_weighted_reduce", + GGML_OP_MUL, + DispatchMatchKind::Fused, + 800, + DispatchSource::Llm, + match_routed_ffn_down_weighted_reduce_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-routed-ffn.h b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-routed-ffn.h new file mode 100644 index 000000000000..54589e193b73 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/qwen/dispatch-routed-ffn.h @@ -0,0 +1,9 @@ +#pragma once + +#include "../dispatch-registry.h" + +namespace ggml::hrx { + +void register_routed_ffn_dispatches(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/fused-context-claim.h b/ggml/src/ggml-hrx/fused-context-claim.h new file mode 100644 index 000000000000..16357796bf30 --- /dev/null +++ b/ggml/src/ggml-hrx/fused-context-claim.h @@ -0,0 +1,125 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +// Claim for nodes that ggml-hrx executes only inside a fused dispatch. +// +// device_supports_op claims a node when can_execute_standalone_op_as_graph says the +// dispatcher can run it alone, so that everything else falls back to another backend. +// Some nodes of the Qwen3 MoE and Qwen3.5/3.6 graphs have no standalone dispatch and +// run only inside a fused pattern: +// +// - the MoE router chain SOFT_MAX -> ARGSORT -> GET_ROWS -> SUM_ROWS -> CLAMP +// (llm.moe_router.top8_f32 and the decode projection fusion); +// - the gated-delta-net chain L2_NORM, SOFTPLUS and GATED_DELTA_NET; +// - the per-head RMS_NORM and ROPE of the attention fusions (head_dim 64..256). +// +// With the standalone test alone these nodes went to the CPU, which split every MoE +// and DeltaNet layer: Qwen3.6-35B-A3B prefill fell to 67 tok/s at pp512 and +// Qwen3-30B-A3B prompt decode failed ("value alias target is not transient"). +// +// A node keeps the per-op claim when it has the producers of those patterns, which +// only a model graph gives it. The single-op and small multi-op graphs of +// test-backend-ops feed these ops from leaves, from other producers or with shapes the +// fusions do not take (4-D, rows wider than a head), so they still go through the +// standalone test. + +#include "ggml.h" + +namespace ggml::hrx { + +// the op that produced a tensor's storage, looking through views and reshapes; GGML_OP_NONE for a leaf +inline ggml_op fused_context_producer_op(const ggml_tensor * tensor) { + if (tensor == nullptr) { + return GGML_OP_NONE; + } + while (tensor->view_src != nullptr) { + tensor = tensor->view_src; + } + return tensor->op; +} + +// the tensor that owns a tensor's storage, looking through views and reshapes +inline const ggml_tensor * fused_context_root(const ggml_tensor * tensor) { + while (tensor != nullptr && tensor->view_src != nullptr) { + tensor = tensor->view_src; + } + return tensor; +} + +// true for the router's GET_ROWS: softmax probabilities gathered at the argsort's top-k ids, +// the chain the top-8 router dispatch starts from (a sigmoid router, as in GLM-4.7-Flash, has none) +inline bool fused_context_router_get_rows(const ggml_tensor * op) { + return op != nullptr && op->op == GGML_OP_GET_ROWS && + fused_context_producer_op(op->src[0]) == GGML_OP_SOFT_MAX && + fused_context_producer_op(op->src[1]) == GGML_OP_ARGSORT; +} + +inline bool fused_context_fed_by_op(const ggml_tensor * op) { + for (const ggml_tensor * source : op->src) { + if (source != nullptr && fused_context_producer_op(source) != GGML_OP_NONE) { + return true; + } + } + return false; +} + +inline bool fused_context_is_3d(const ggml_tensor * op) { + return op->ne[3] == 1; +} + +inline bool fused_context_is_head_row(const ggml_tensor * op) { + return op->ne[0] == 64 || op->ne[0] == 128 || op->ne[0] == 256; +} + +// true when op is a node that only a fused dispatch executes, with the producers of that fused pattern +inline bool fused_context_claim(const ggml_tensor * op) { + if (op == nullptr || !fused_context_is_3d(op)) { + return false; + } + switch (op->op) { + // MoE router: softmax over [n_expert, n_tokens] router logits, then top-k + case GGML_OP_SOFT_MAX: + return op->src[1] == nullptr && fused_context_producer_op(op->src[0]) == GGML_OP_MUL_MAT && + op->ne[2] == 1 && op->ne[0] >= 32 && op->ne[0] <= 512 && op->ne[0] % 32 == 0; + case GGML_OP_ARGSORT: + return fused_context_producer_op(op->src[0]) == GGML_OP_SOFT_MAX; + case GGML_OP_GET_ROWS: + return fused_context_producer_op(op->src[0]) == GGML_OP_SOFT_MAX && + fused_context_producer_op(op->src[1]) == GGML_OP_ARGSORT; + case GGML_OP_SUM_ROWS: + return fused_context_router_get_rows(fused_context_root(op->src[0])); + case GGML_OP_CLAMP: { + const ggml_tensor * sum = fused_context_root(op->src[0]); + return sum != nullptr && sum->op == GGML_OP_SUM_ROWS && + fused_context_router_get_rows(fused_context_root(sum->src[0])); + } + // gated delta net + case GGML_OP_L2_NORM: + case GGML_OP_GATED_DELTA_NET: + return fused_context_fed_by_op(op); + case GGML_OP_UNARY: + return ggml_get_unary_op(op) == GGML_UNARY_OP_SOFTPLUS && fused_context_fed_by_op(op); + // per-head attention normalisation and rotation + case GGML_OP_RMS_NORM: + case GGML_OP_ROPE: + return fused_context_is_head_row(op) && fused_context_fed_by_op(op); + default: + return false; + } +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/ggml-hrx-dmabuf.cpp b/ggml/src/ggml-hrx/ggml-hrx-dmabuf.cpp new file mode 100644 index 000000000000..499b39f7493d --- /dev/null +++ b/ggml/src/ggml-hrx/ggml-hrx-dmabuf.cpp @@ -0,0 +1,40 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "ggml-hrx-dmabuf.h" + +#include +#include +#include + +namespace ggml::hrx { + +bool export_dmabuf(void * device_ptr, size_t size, int * fd, size_t * offset) { + using export_fn = int (*)(const void *, size_t, int *, uint64_t *); + // libhsa is already loaded by the HRX runtime; look the symbol up in that copy + static export_fn hsa_export = []() -> export_fn { + const char * path = getenv("IREE_HAL_AMDGPU_LIBHSA_PATH"); + void * lib = dlopen(path != nullptr ? path : "libhsa-runtime64.so.1", RTLD_NOW | RTLD_NOLOAD); + return lib != nullptr ? (export_fn) dlsym(lib, "hsa_amd_portable_export_dmabuf") : nullptr; + }(); + uint64_t off = 0; + if (hsa_export == nullptr || device_ptr == nullptr || hsa_export(device_ptr, size, fd, &off) != 0) { + return false; + } + *offset = (size_t) off; + return true; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/ggml-hrx-dmabuf.h b/ggml/src/ggml-hrx/ggml-hrx-dmabuf.h new file mode 100644 index 000000000000..34c6ee6e4c4b --- /dev/null +++ b/ggml/src/ggml-hrx/ggml-hrx-dmabuf.h @@ -0,0 +1,25 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include + +namespace ggml::hrx { + +// Export an HRX device allocation as a dma-buf through HSA, so another API (Vulkan) can map it with no copy. +bool export_dmabuf(void * device_ptr, size_t size, int * fd, size_t * offset); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/ggml-hrx.cpp b/ggml/src/ggml-hrx/ggml-hrx.cpp new file mode 100644 index 000000000000..8090e3bd593e --- /dev/null +++ b/ggml/src/ggml-hrx/ggml-hrx.cpp @@ -0,0 +1,1264 @@ +#include "dispatch_registration/common/moe-placement-guard.h" +#include "ggml-hrx.h" +#include "ggml-hrx-dmabuf.h" + +#include "backend-buffer-binding.h" +#include "backend-context.h" +#include "fused-context-claim.h" +#include "ggml-backend-impl.h" +#include "ggml-impl.h" +#include "graph/op-params.h" +#include "hip/hip-capabilities.h" +#include "hrx_runtime.h" +#include "kernel-corpus/kernel-corpus.h" +#include "loom-jit.h" +#include "runtime/hrx-sleeping-wait.h" +#include "runtime/graph-executor.h" +#include "runtime/graph-program-cache.h" +#include "runtime/kernel-executable-cache.h" +#include "runtime/prepared-command-program-cache.h" +#include "runtime/transient-arena.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace { + +static constexpr size_t GGML_HRX_ALIGNMENT = 256; +// GGML represents tensor locations as host pointers and derives view/arena offsets with ordinary pointer arithmetic. +// Device-local HRX buffers have no host address to return, so expose a non-null sentinel base as an offset coordinate. +static constexpr uintptr_t GGML_HRX_FAKE_PTR_BASE = 0x1000; +static std::atomic g_allocation_generation{ 1 }; + +static bool environment_flag_enabled(const char * name) { + const char * value = std::getenv(name); + return value != nullptr && value[0] != '\0' && value[0] != '0'; +} + +static void log_hrx_device_event(void *, const hrx_device_event_t * event) { + if (event == nullptr || event->type != HRX_DEVICE_EVENT_TYPE_ASAN_REPORT || event->payload.data == nullptr || + event->payload.data_length < sizeof(hrx_device_asan_report_t)) { + return; + } + hrx_device_asan_report_t report; + std::memcpy(&report, event->payload.data, sizeof(report)); + GGML_LOG_ERROR("HRX ASAN: executable=%" PRIu64 " export=%u site=%" PRIu64 " access=%u address=0x%016" PRIx64 + " length=%" PRIu64 " workgroup=(%u,%u,%u) workitem=(%u,%u,%u) shadow=0x%016" PRIx64 + " value=0x%016" PRIx64 " dispatch=0x%016" PRIx64 "\n", + event->source.executable_id, event->source.export_ordinal, report.site_id, report.access_kind, + report.fault_address, report.access_length, report.workgroup_id[0], report.workgroup_id[1], + report.workgroup_id[2], report.workitem_id[0], report.workitem_id[1], report.workitem_id[2], + report.shadow_address, report.shadow_value, report.source_dispatch_ptr); +} + +static bool hrx_check(hrx_status_t status, const char * expression, const char * file, int line) { + if (hrx_status_is_ok(status)) { + return true; + } + char * message = nullptr; + size_t length = 0; + hrx_status_to_string(status, &message, &length); + GGML_LOG_ERROR("%s:%d: %s failed: %s\n", file, line, expression, + message != nullptr ? message : "unknown HRX error"); + hrx_status_free_message(message); + hrx_status_ignore(status); + return false; +} + +#define HRX_CHECK(expression) hrx_check((expression), #expression, __FILE__, __LINE__) + +} // namespace + +ggml_backend_hrx_reg_context::~ggml_backend_hrx_reg_context() { + for (auto & context : device_contexts) { + if (context->buffer_stream != nullptr) { + hrx_stream_release(context->buffer_stream); + } + if (context->device != nullptr) { + hrx_device_release(context->device); + } + } + if (initialized) { + hrx_status_t status = hrx_gpu_shutdown(); + if (!hrx_status_is_ok(status)) { + hrx_status_ignore(status); + } + } +} + +namespace { + +static std::optional device_string_property(hrx_device_t device, + hrx_device_property_t property, + const char * property_name) { + std::vector buffer(64); + while (buffer.size() <= 4096) { + hrx_status_t status = hrx_device_get_property(device, property, buffer.data(), buffer.size()); + if (hrx_status_is_ok(status)) { + return std::string(buffer.data()); + } + if (hrx_status_code(status) != HRX_STATUS_OUT_OF_RANGE) { + hrx_check(status, property_name, __FILE__, __LINE__); + return std::nullopt; + } + hrx_status_ignore(status); + buffer.resize(buffer.size() * 2); + } + GGML_LOG_ERROR("%s exceeds the maximum supported property string length\n", property_name); + return std::nullopt; +} + +static ggml_guid_t ggml_backend_hrx_guid() { + static ggml_guid guid = { + 0xd2, 0x3d, 0x72, 0x83, 0xb2, 0x82, 0x4d, 0xe0, 0x8a, 0x3e, 0x21, 0x1d, 0x68, 0x87, 0x2f, 0x4b, + }; + return &guid; +} + +static ggml_backend_hrx_device_context * device_context(ggml_backend_dev_t device) { + return static_cast(device->context); +} + +static ggml_backend_hrx_buffer_context * buffer_context(ggml_backend_buffer_t buffer) { + return ggml_backend_hrx_buffer_context_from_buffer(buffer); +} + +static size_t tensor_offset(const ggml_backend_hrx_buffer_context * context, const ggml_tensor * tensor) { + return ggml_backend_hrx_tensor_offset(context, tensor); +} + +static const char * buffer_type_name(ggml_backend_buffer_type_t buft) { + return static_cast(buft->context)->name.c_str(); +} + +static bool buffer_type_is_host(ggml_backend_buffer_type_t buft) { + return static_cast(buft->context)->host_visible; +} + +static bool buffer_submit_and_wait(ggml_backend_hrx_device_context * device, + hrx_status_t (*submit)(hrx_stream_t, void *), + void * user_data) { + std::lock_guard lock(device->buffer_stream_mutex); + if (!HRX_CHECK(submit(device->buffer_stream, user_data))) { + return false; + } + return HRX_CHECK(hrx_stream_synchronize(device->buffer_stream)); +} + +struct FillBufferArgs { + hrx_buffer_t buffer; + size_t offset; + size_t size; + uint8_t value; +}; + +static hrx_status_t submit_fill_buffer(hrx_stream_t stream, void * user_data) { + auto * args = static_cast(user_data); + return hrx_stream_fill_buffer(stream, args->buffer, args->offset, args->size, &args->value, sizeof(args->value)); +} + +struct CopyBufferArgs { + hrx_buffer_t source; + size_t source_offset; + hrx_buffer_t destination; + size_t destination_offset; + size_t size; +}; + +static hrx_status_t submit_copy_buffer(hrx_stream_t stream, void * user_data) { + auto * args = static_cast(user_data); + return hrx_stream_copy_buffer(stream, args->source, args->source_offset, args->destination, + args->destination_offset, args->size); +} + +static void buffer_free(ggml_backend_buffer_t buffer) { + auto * context = buffer_context(buffer); + if (context->base != reinterpret_cast(GGML_HRX_FAKE_PTR_BASE)) { + context->device->host_buffers.remove(context->buffer); + } + if (context->buffer != nullptr) { + hrx_buffer_release(context->buffer); + } + delete context; +} + +static void buffer_memset(ggml_backend_buffer_t buffer, + ggml_tensor * tensor, + uint8_t value, + size_t offset, + size_t size) { + if (size == 0) { + return; + } + auto * context = buffer_context(buffer); + const size_t destination_offset = tensor_offset(context, tensor) + offset; + GGML_ASSERT(destination_offset <= buffer->size && size <= buffer->size - destination_offset); + if (context->base != reinterpret_cast(GGML_HRX_FAKE_PTR_BASE)) { + std::memset(context->base + destination_offset, value, size); + return; + } + FillBufferArgs args{ context->buffer, destination_offset, size, value }; + if (!buffer_submit_and_wait(context->device, submit_fill_buffer, &args)) { + GGML_LOG_ERROR("%s: HRX buffer fill failed\n", __func__); + } +} + +static void buffer_set(ggml_backend_buffer_t buffer, + ggml_tensor * tensor, + const void * data, + size_t offset, + size_t size) { + if (size == 0) { + return; + } + auto * context = buffer_context(buffer); + const size_t destination_offset = tensor_offset(context, tensor) + offset; + GGML_ASSERT(destination_offset <= buffer->size && size <= buffer->size - destination_offset); + if (context->base != reinterpret_cast(GGML_HRX_FAKE_PTR_BASE)) { + std::memcpy(context->base + destination_offset, data, size); + return; + } + if (!HRX_CHECK(hrx_synchronous_h2d(context->device->device, data, context->buffer, destination_offset, size))) { + GGML_LOG_ERROR("%s: HRX buffer upload failed\n", __func__); + } +} + +static void buffer_get(ggml_backend_buffer_t buffer, + const ggml_tensor * tensor, + void * data, + size_t offset, + size_t size) { + if (size == 0) { + return; + } + auto * context = buffer_context(buffer); + const size_t source_offset = tensor_offset(context, tensor) + offset; + GGML_ASSERT(source_offset <= buffer->size && size <= buffer->size - source_offset); + if (context->base != reinterpret_cast(GGML_HRX_FAKE_PTR_BASE)) { + std::memcpy(data, context->base + source_offset, size); + return; + } + if (!HRX_CHECK(hrx_synchronous_d2h(context->device->device, context->buffer, source_offset, data, size))) { + GGML_LOG_ERROR("%s: HRX buffer download failed\n", __func__); + } +} + +static bool buffer_copy(ggml_backend_buffer_t buffer, const ggml_tensor * source, ggml_tensor * destination) { + ggml_backend_buffer_t source_buffer = source->view_src != nullptr ? source->view_src->buffer : source->buffer; + if (source_buffer == nullptr || source_buffer->iface.get_base != ggml_backend_hrx_buffer_base) { + return false; + } + auto * source_context = buffer_context(source_buffer); + auto * destination_context = buffer_context(buffer); + if (source_context->device != destination_context->device) { + return false; + } + const size_t source_offset = tensor_offset(source_context, source); + const size_t destination_offset = tensor_offset(destination_context, destination); + const size_t size = ggml_nbytes(source); + if (source_offset > source_buffer->size || size > source_buffer->size - source_offset || + destination_offset > buffer->size || size > buffer->size - destination_offset) { + return false; + } + CopyBufferArgs args{ source_context->buffer, source_offset, destination_context->buffer, destination_offset, size }; + return buffer_submit_and_wait(destination_context->device, submit_copy_buffer, &args); +} + +static void buffer_clear(ggml_backend_buffer_t buffer, uint8_t value) { + if (buffer->size == 0) { + return; + } + auto * context = buffer_context(buffer); + if (context->base != reinterpret_cast(GGML_HRX_FAKE_PTR_BASE)) { + std::memset(context->base, value, buffer->size); + return; + } + FillBufferArgs args{ context->buffer, 0, buffer->size, value }; + if (!buffer_submit_and_wait(context->device, submit_fill_buffer, &args)) { + GGML_LOG_ERROR("%s: HRX buffer clear failed\n", __func__); + } +} + +static const ggml_backend_buffer_i buffer_i = { + buffer_free, ggml_backend_hrx_buffer_base, + nullptr, buffer_memset, + buffer_set, buffer_get, + nullptr, nullptr, + buffer_copy, buffer_clear, + nullptr, +}; + +static ggml_backend_buffer_t buffer_alloc(ggml_backend_buffer_type_t buft, size_t size) { + auto * type_context = static_cast(buft->context); + const bool host_visible = type_context->host_visible; + const bool direct_host_binding = host_visible && type_context->device->use_direct_host_bindings; + hrx_memory_type_t memory_type = HRX_MEMORY_TYPE_DEVICE_LOCAL; + // Direct command-program bindings require coherent CPU/GPU visibility. Otherwise HRX host buffers are pinned + // transfer memory: DEVICE_VISIBLE permits handle-based stream copies without implying direct device access. + if (host_visible) { + memory_type = HRX_MEMORY_TYPE_HOST_LOCAL | HRX_MEMORY_TYPE_DEVICE_VISIBLE; + if (direct_host_binding) { + memory_type |= HRX_MEMORY_TYPE_HOST_COHERENT; + } + } + hrx_buffer_params_t params = { + memory_type, + HRX_MEMORY_ACCESS_ALL, + host_visible ? + HRX_BUFFER_USAGE_DEFAULT | HRX_BUFFER_USAGE_MAPPING_SCOPED | HRX_BUFFER_USAGE_MAPPING_PERSISTENT : + HRX_BUFFER_USAGE_DEFAULT, + 0, + }; + hrx_buffer_t allocation = nullptr; + if (size > 0 && !HRX_CHECK(hrx_allocator_allocate_buffer(hrx_device_allocator(type_context->device->device), params, + size, &allocation))) { + return nullptr; + } + uint8_t * base = reinterpret_cast(GGML_HRX_FAKE_PTR_BASE); + if (host_visible && size > 0) { + void * mapped = nullptr; + if (!HRX_CHECK(hrx_buffer_map(allocation, HRX_MAP_READ | HRX_MAP_WRITE, 0, size, &mapped))) { + hrx_buffer_release(allocation); + return nullptr; + } + base = static_cast(mapped); + } + const uint64_t generation = g_allocation_generation.fetch_add(1); + auto * context = new (std::nothrow) ggml_backend_hrx_buffer_context{ + type_context->device, allocation, base, generation, generation, direct_host_binding, + }; + if (context == nullptr) { + if (allocation != nullptr) { + hrx_buffer_release(allocation); + } + return nullptr; + } + if (host_visible && allocation != nullptr) { + type_context->device->host_buffers.add(allocation, base, size); + } + return ggml_backend_buffer_init(buft, buffer_i, context, size); +} + +static size_t buffer_alignment(ggml_backend_buffer_type_t buft) { + GGML_UNUSED(buft); + return GGML_HRX_ALIGNMENT; +} + +static size_t buffer_max_size(ggml_backend_buffer_type_t buft) { + return static_cast(buft->context)->device->memory_total; +} + +static const ggml_backend_buffer_type_i buffer_type_i = { + buffer_type_name, buffer_alloc, buffer_alignment, buffer_max_size, nullptr, buffer_type_is_host, +}; + +static const char * backend_name(ggml_backend_t backend) { + return static_cast(backend->context)->name.c_str(); +} + +static void backend_synchronize(ggml_backend_t backend); + +static void backend_free(ggml_backend_t backend) { + auto * context = static_cast(backend->context); + backend_synchronize(backend); + context->prepared_programs.clear(); + context->graph_programs.clear(); + context->kernel_executables.clear(); + context->transient_arena.clear(); + context->host_weights.clear(); + context->host_transfers.clear(); + hrx_stream_release(context->stream); + delete context; + delete backend; +} + +static bool synchronous_upload_fallback(ggml_backend_hrx_context * backend, + const void * source, + hrx_buffer_t destination, + size_t destination_offset, + size_t size) { + const uint64_t fallback = backend->device->synchronous_upload_fallbacks.fetch_add(1, std::memory_order_relaxed); + if (fallback == 0) { + GGML_LOG_WARN( + "ggml_hrx: synchronous upload fallback for an unregistered host pointer; use the HRX host " + "buffer type for asynchronous transfers\n"); + } + // Compatibility path for arbitrary GGML pointers. Keep the synchronization explicit until a bounded staging ring + // with transfer retirement is available. + const ggml::hrx::Status status = + backend->host_transfers.upload_synchronous(backend->stream, source, destination, destination_offset, size); + if (!status.success()) { + GGML_LOG_ERROR("%s: %s\n", __func__, status.errors().front().c_str()); + return false; + } + return true; +} + +static bool synchronous_download_fallback(ggml_backend_hrx_context * backend, + hrx_buffer_t source, + size_t source_offset, + void * destination, + size_t size) { + const uint64_t fallback = backend->device->synchronous_download_fallbacks.fetch_add(1, std::memory_order_relaxed); + if (fallback == 0) { + GGML_LOG_WARN( + "ggml_hrx: synchronous download fallback for an unregistered host pointer; use the HRX host " + "buffer type for asynchronous transfers\n"); + } + // Compatibility path for arbitrary GGML pointers. Keep the synchronization explicit until a bounded staging ring + // with transfer retirement is available. + const ggml::hrx::Status status = + backend->host_transfers.download_synchronous(backend->stream, source, source_offset, destination, size); + if (!status.success()) { + GGML_LOG_ERROR("%s: %s\n", __func__, status.errors().front().c_str()); + return false; + } + return true; +} + +static void backend_set_tensor_async(ggml_backend_t backend, + ggml_tensor * tensor, + const void * data, + size_t offset, + size_t size) { + auto * backend_context = static_cast(backend->context); + ggml_backend_hrx_buffer_context * context = nullptr; + size_t tensor_base = 0; + if (!ggml_backend_hrx_tensor_binding(tensor, &context, &tensor_base) || offset > ggml_nbytes(tensor) || + size > ggml_nbytes(tensor) - offset) { + GGML_LOG_ERROR("%s: invalid HRX tensor upload\n", __func__); + return; + } + ggml::hrx::HostBufferRef source = backend_context->device->host_buffers.find(data, size); + if (source.valid()) { + // Registered host buffers can participate directly in the stream command buffer. + HRX_CHECK(hrx_stream_copy_buffer(backend_context->stream, source.buffer(), source.offset(), context->buffer, + tensor_base + offset, size)); + } else { + synchronous_upload_fallback(backend_context, data, context->buffer, tensor_base + offset, size); + } +} + +static void backend_get_tensor_async(ggml_backend_t backend, + const ggml_tensor * tensor, + void * data, + size_t offset, + size_t size) { + auto * backend_context = static_cast(backend->context); + ggml_backend_hrx_buffer_context * context = nullptr; + size_t tensor_base = 0; + if (!ggml_backend_hrx_tensor_binding(tensor, &context, &tensor_base) || offset > ggml_nbytes(tensor) || + size > ggml_nbytes(tensor) - offset) { + GGML_LOG_ERROR("%s: invalid HRX tensor download\n", __func__); + return; + } + ggml::hrx::HostBufferRef destination = backend_context->device->host_buffers.find(data, size); + if (destination.valid()) { + // Registered host buffers can participate directly in the stream command buffer. + HRX_CHECK(hrx_stream_copy_buffer(backend_context->stream, context->buffer, tensor_base + offset, + destination.buffer(), destination.offset(), size)); + } else { + synchronous_download_fallback(backend_context, context->buffer, tensor_base + offset, data, size); + } +} + +static bool backend_copy_tensor_async(ggml_backend_t backend_src, + ggml_backend_t backend_dst, + const ggml_tensor * source, + ggml_tensor * destination) { + GGML_UNUSED(backend_src); + auto * destination_backend = static_cast(backend_dst->context); + ggml_backend_hrx_buffer_context * destination_context = nullptr; + size_t destination_offset = 0; + if (!ggml_backend_hrx_tensor_binding(destination, &destination_context, &destination_offset)) { + return false; + } + ggml_backend_hrx_buffer_context * source_context = nullptr; + size_t source_offset = 0; + const size_t size = ggml_nbytes(source); + if (ggml_backend_hrx_tensor_binding(source, &source_context, &source_offset)) { + if (source_context->device != destination_context->device) { + return false; + } + return HRX_CHECK(hrx_stream_copy_buffer(destination_backend->stream, source_context->buffer, source_offset, + destination_context->buffer, destination_offset, size)); + } + ggml_backend_buffer_t source_buffer = source->view_src != nullptr ? source->view_src->buffer : source->buffer; + if (source_buffer != nullptr && ggml_backend_buffer_is_host(source_buffer)) { + return synchronous_upload_fallback(destination_backend, source->data, destination_context->buffer, + destination_offset, size); + } + return false; +} + +static void backend_synchronize(ggml_backend_t backend) { + auto * context = static_cast(backend->context); + static thread_local ggml::hrx::WaitHistory synchronize_wait_history; + if (HRX_CHECK(ggml::hrx::stream_synchronize_sleeping(context->stream, synchronize_wait_history))) { + context->graph_replay_state.mark_stream_synchronized(); + } +} + +static const char * status_first_error(const ggml::hrx::Status & status) { + return status.errors().empty() ? "" : status.errors().front().c_str(); +} + +static enum ggml_status graph_compute(ggml_backend_t backend, ggml_cgraph * graph) { + auto * context = static_cast(backend->context); + const ggml::hrx::GraphExecutor executor = ggml::hrx::GraphExecutor(*context); + const ggml::hrx::GraphExecutionResult result = executor.execute(*graph); + if (!result.success()) { + GGML_LOG_ERROR("%s: %s\n", __func__, status_first_error(result.status)); + } + return result.code; +} + +static const ggml_backend_i backend_i = { + backend_name, + backend_free, + backend_set_tensor_async, + backend_get_tensor_async, + nullptr, + nullptr, + backend_copy_tensor_async, + backend_synchronize, + nullptr, + nullptr, + nullptr, + nullptr, + graph_compute, + nullptr, + nullptr, + nullptr, +}; + +static const char * device_name(ggml_backend_dev_t device) { + return device_context(device)->name.c_str(); +} + +static const char * device_description(ggml_backend_dev_t device) { + return device_context(device)->description.c_str(); +} + +static void device_memory(ggml_backend_dev_t device, size_t * free, size_t * total) { + *free = device_context(device)->memory_total; + *total = device_context(device)->memory_total; +} + +static enum ggml_backend_dev_type device_type(ggml_backend_dev_t device) { + GGML_UNUSED(device); + return GGML_BACKEND_DEVICE_TYPE_GPU; +} + +static void device_props(ggml_backend_dev_t device, ggml_backend_dev_props * props) { + props->name = device_name(device); + props->description = device_description(device); + device_memory(device, &props->memory_free, &props->memory_total); + props->type = GGML_BACKEND_DEVICE_TYPE_GPU; + props->device_id = nullptr; + props->caps = { true, true, false, false }; +} + +static ggml_backend_t device_init(ggml_backend_dev_t device, const char * parameters) { + GGML_UNUSED(parameters); + auto * device_ctx = device_context(device); + hrx_stream_t stream = nullptr; + if (!HRX_CHECK(hrx_stream_create(device_ctx->device, 0, &stream))) { + return nullptr; + } + auto * context = new (std::nothrow) ggml_backend_hrx_context; + if (context != nullptr) { + context->device = device_ctx; + context->stream = stream; + context->name = device_ctx->name; + } + auto * backend = context != nullptr ? new (std::nothrow) + ggml_backend{ ggml_backend_hrx_guid(), backend_i, device, context } : + nullptr; + if (backend == nullptr) { + delete context; + hrx_stream_release(stream); + } + return backend; +} + +static ggml_backend_buffer_type_t device_buffer_type(ggml_backend_dev_t device) { + return &device_context(device)->buft; +} + +static ggml_backend_buffer_type_t device_host_buffer_type(ggml_backend_dev_t device) { + return &device_context(device)->host_buft; +} + +static bool eager_capability_declared(enum ggml_op op) { + switch (op) { + // The scheduler probes preallocated weight tensors as NONE operations when deciding whether their buffer type is + // usable by this backend. Fused ops are declared here so graph-claim can validate the full dispatch pattern. + // TODO: split this into placement capability and exact graph execution capability once graph claiming owns the + // full decision. + case GGML_OP_NONE: + case GGML_OP_ADD: + case GGML_OP_ADD_ID: + case GGML_OP_ARGSORT: + case GGML_OP_CLAMP: + case GGML_OP_CONCAT: + case GGML_OP_CONT: + case GGML_OP_CPY: + case GGML_OP_DIV: + case GGML_OP_FLASH_ATTN_EXT: + case GGML_OP_GATED_DELTA_NET: + case GGML_OP_GET_ROWS: + case GGML_OP_GLU: + case GGML_OP_L2_NORM: + case GGML_OP_MUL: + case GGML_OP_MUL_MAT: + case GGML_OP_MUL_MAT_ID: + case GGML_OP_NORM: + case GGML_OP_PERMUTE: + case GGML_OP_REPEAT: // broadcast-only, through the strided copy (dispatch-small-rows.cpp) + case GGML_OP_RESHAPE: + case GGML_OP_RMS_NORM: + case GGML_OP_ROPE: + case GGML_OP_SCALE: + case GGML_OP_SET_ROWS: + case GGML_OP_SOFT_MAX: + case GGML_OP_SSM_CONV: + case GGML_OP_SUM_ROWS: + case GGML_OP_TRANSPOSE: + case GGML_OP_UNARY: + case GGML_OP_VIEW: + return true; + default: + return ggml::hrx::hip_eager_op_declared(op); // ops a HIP matcher declared (hip/hip-capabilities.h) + } +} + +static bool zero_output_elision_safe_op(enum ggml_op op) { + switch (op) { + case GGML_OP_DUP: + case GGML_OP_ADD: + case GGML_OP_ADD_ID: + case GGML_OP_ADD1: + case GGML_OP_SUB: + case GGML_OP_MUL: + case GGML_OP_DIV: + case GGML_OP_SQR: + case GGML_OP_SQRT: + case GGML_OP_LOG: + case GGML_OP_SIN: + case GGML_OP_COS: + case GGML_OP_SUM: + case GGML_OP_SUM_ROWS: + case GGML_OP_CUMSUM: + case GGML_OP_MEAN: + case GGML_OP_ARGMAX: + case GGML_OP_COUNT_EQUAL: + case GGML_OP_REPEAT: + case GGML_OP_REPEAT_BACK: + case GGML_OP_CONCAT: + case GGML_OP_SILU_BACK: + case GGML_OP_NORM: + case GGML_OP_RMS_NORM: + case GGML_OP_RMS_NORM_BACK: + case GGML_OP_GROUP_NORM: + case GGML_OP_L2_NORM: + case GGML_OP_MUL_MAT: + case GGML_OP_MUL_MAT_ID: + case GGML_OP_OUT_PROD: + case GGML_OP_SCALE: + case GGML_OP_CONT: + case GGML_OP_GET_ROWS: + case GGML_OP_GET_ROWS_BACK: + case GGML_OP_DIAG: + case GGML_OP_DIAG_MASK_INF: + case GGML_OP_DIAG_MASK_ZERO: + case GGML_OP_SOFT_MAX: + case GGML_OP_SOFT_MAX_BACK: + case GGML_OP_ROPE: + case GGML_OP_ROPE_BACK: + case GGML_OP_CLAMP: + case GGML_OP_CONV_TRANSPOSE_1D: + case GGML_OP_IM2COL: + case GGML_OP_IM2COL_BACK: + case GGML_OP_IM2COL_3D: + case GGML_OP_COL2IM_1D: + case GGML_OP_CONV_2D: + case GGML_OP_CONV_3D: + case GGML_OP_CONV_2D_DW: + case GGML_OP_CONV_TRANSPOSE_2D: + case GGML_OP_POOL_1D: + case GGML_OP_POOL_2D: + case GGML_OP_POOL_2D_BACK: + case GGML_OP_UPSCALE: + case GGML_OP_PAD: + case GGML_OP_PAD_REFLECT_1D: + case GGML_OP_ROLL: + case GGML_OP_ARANGE: + case GGML_OP_TIMESTEP_EMBEDDING: + case GGML_OP_ARGSORT: + case GGML_OP_TOP_K: + case GGML_OP_LEAKY_RELU: + case GGML_OP_TRI: + case GGML_OP_FILL: + case GGML_OP_FLASH_ATTN_EXT: + case GGML_OP_FLASH_ATTN_BACK: + case GGML_OP_SSM_CONV: + case GGML_OP_SSM_SCAN: + case GGML_OP_WIN_PART: + case GGML_OP_WIN_UNPART: + case GGML_OP_GET_REL_POS: + case GGML_OP_ADD_REL_POS: + case GGML_OP_RWKV_WKV6: + case GGML_OP_GATED_LINEAR_ATTN: + case GGML_OP_RWKV_WKV7: + case GGML_OP_SOLVE_TRI: + case GGML_OP_GATED_DELTA_NET: + case GGML_OP_LIGHTNING_INDEXER: + case GGML_OP_DSV4_HC_COMB: + case GGML_OP_DSV4_HC_PRE: + case GGML_OP_DSV4_HC_POST: + case GGML_OP_UNARY: + case GGML_OP_CROSS_ENTROPY_LOSS: + case GGML_OP_CROSS_ENTROPY_LOSS_BACK: + case GGML_OP_GLU: + return true; + default: + return false; + } +} + +static bool zero_output_elision_supported(const ggml_tensor * op) { + return op != nullptr && (ggml_nelements(op) == 0 || ggml_nbytes(op) == 0) && zero_output_elision_safe_op(op->op); +} + +struct TensorStorageRange { + const ggml_tensor * root = nullptr; + size_t offset = 0; + size_t size = 0; + bool valid = false; +}; + +static TensorStorageRange tensor_storage_range(const ggml_tensor * tensor) { + TensorStorageRange range; + if (tensor == nullptr) { + return range; + } + range.size = ggml_nbytes(tensor); + while (tensor->view_src != nullptr) { + if (range.offset > std::numeric_limits::max() - tensor->view_offs) { + return {}; + } + range.offset += tensor->view_offs; + tensor = tensor->view_src; + } + range.root = tensor; + range.valid = true; + return range; +} + +static TensorStorageRange tensor_storage_range(const ggml_tensor * tensor, size_t size) { + TensorStorageRange range = tensor_storage_range(tensor); + if (range.valid) { + range.size = size; + } + return range; +} + +static bool tensor_storage_exactly_overlaps(const TensorStorageRange & lhs, const TensorStorageRange & rhs) { + return lhs.root == rhs.root && lhs.offset == rhs.offset && lhs.size == rhs.size; +} + +static bool tensor_storage_ranges_disjoint(const TensorStorageRange & lhs, const TensorStorageRange & rhs) { + if (lhs.root != rhs.root) { + return true; + } + if (lhs.offset > std::numeric_limits::max() - lhs.size || + rhs.offset > std::numeric_limits::max() - rhs.size) { + return false; + } + return lhs.offset + lhs.size <= rhs.offset || rhs.offset + rhs.size <= lhs.offset; +} + +static bool scale_storage_is_safe(const ggml_tensor * input, const ggml_tensor * output) { + const TensorStorageRange input_range = tensor_storage_range(input); + const TensorStorageRange output_range = tensor_storage_range(output); + if (!input_range.valid || !output_range.valid) { + return false; + } + return tensor_storage_exactly_overlaps(input_range, output_range) || + tensor_storage_ranges_disjoint(input_range, output_range); +} + +static bool binary_output_storage_is_safe(const ggml_tensor * lhs, const ggml_tensor * rhs, const ggml_tensor * output) { + const TensorStorageRange lhs_range = tensor_storage_range(lhs); + const TensorStorageRange rhs_range = tensor_storage_range(rhs); + const TensorStorageRange output_range = tensor_storage_range(output); + if (!lhs_range.valid || !rhs_range.valid || !output_range.valid) { + return false; + } + return tensor_storage_ranges_disjoint(lhs_range, output_range) && + tensor_storage_ranges_disjoint(rhs_range, output_range); +} + +static bool binary_output_storage_is_safe(const ggml_tensor * lhs, + size_t lhs_byte_count, + const ggml_tensor * rhs, + size_t rhs_byte_count, + const ggml_tensor * output) { + const TensorStorageRange lhs_range = tensor_storage_range(lhs, lhs_byte_count); + const TensorStorageRange rhs_range = tensor_storage_range(rhs, rhs_byte_count); + const TensorStorageRange output_range = tensor_storage_range(output); + if (!lhs_range.valid || !rhs_range.valid || !output_range.valid) { + return false; + } + return tensor_storage_ranges_disjoint(lhs_range, output_range) && + tensor_storage_ranges_disjoint(rhs_range, output_range); +} + +static bool tensor_has_positive_shape(const ggml_tensor * tensor) { + if (tensor == nullptr || ggml_nelements(tensor) <= 0) { + return false; + } + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (tensor->ne[i] <= 0) { + return false; + } + } + return true; +} + +static bool tensor_has_packed_f32_layout(const ggml_tensor * tensor) { + if (tensor == nullptr) { + return false; + } + size_t expected_stride = sizeof(float); + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (tensor->nb[i] != expected_stride) { + return false; + } + expected_stride *= static_cast(tensor->ne[i]); + } + return true; +} + +static bool tensor_has_strided_f32_layout(const ggml_tensor * tensor) { + if (tensor == nullptr || tensor->nb[0] != sizeof(float)) { + return false; + } + for (int i = 1; i < GGML_MAX_DIMS; ++i) { + if (tensor->nb[i] % sizeof(float) != 0) { + return false; + } + } + return true; +} + +static bool tensor_storage_span_bytes(const ggml_tensor * tensor, size_t & byte_count) { + if (!tensor_has_strided_f32_layout(tensor)) { + return false; + } + size_t max_offset = 0; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (tensor->ne[i] <= 0) { + return false; + } + const size_t extent = static_cast(tensor->ne[i] - 1); + if (extent != 0 && tensor->nb[i] > std::numeric_limits::max() / extent) { + return false; + } + const size_t dim_offset = extent * tensor->nb[i]; + if (max_offset > std::numeric_limits::max() - dim_offset) { + return false; + } + max_offset += dim_offset; + } + if (max_offset > std::numeric_limits::max() - sizeof(float)) { + return false; + } + byte_count = max_offset + sizeof(float); + return true; +} + +static bool tensor_broadcastable_to(const ggml_tensor * source, const ggml_tensor * output) { + if (source == nullptr || output == nullptr) { + return false; + } + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (source->ne[i] != output->ne[i] && source->ne[i] != 1) { + return false; + } + } + return true; +} + +static bool binary_kind_allows_broadcast(ggml::hrx::BinaryKind kind, + const ggml_tensor * lhs, + const ggml_tensor * rhs, + const ggml_tensor * output) { + const bool lhs_full = ggml_are_same_shape(lhs, output); + const bool rhs_full = ggml_are_same_shape(rhs, output); + if (!tensor_broadcastable_to(lhs, output) || !tensor_broadcastable_to(rhs, output) || (!lhs_full && !rhs_full)) { + return false; + } + + switch (kind) { + case ggml::hrx::BinaryKind::Add: + case ggml::hrx::BinaryKind::Mul: + return true; + case ggml::hrx::BinaryKind::Sub: + case ggml::hrx::BinaryKind::Div: + return lhs_full; + case ggml::hrx::BinaryKind::SwiGLU: + case ggml::hrx::BinaryKind::GeGLU: + case ggml::hrx::BinaryKind::RegLU: + case ggml::hrx::BinaryKind::GeGLUErf: + case ggml::hrx::BinaryKind::GeGLUQuick: + return lhs_full && rhs_full; + } + return false; +} + +static bool supported_binary_f32_tensor(const ggml_tensor * op) { + if (op == nullptr || op->src[0] == nullptr || op->src[1] == nullptr) { + return false; + } + if (op->type != GGML_TYPE_F32 || op->src[0]->type != GGML_TYPE_F32 || op->src[1]->type != GGML_TYPE_F32) { + return false; + } + if (!tensor_has_positive_shape(op)) { + return false; + } + if (!ggml_is_contiguous(op) || !tensor_has_packed_f32_layout(op) || op->view_src != nullptr) { + return false; + } + size_t lhs_byte_count = 0; + size_t rhs_byte_count = 0; + if (!tensor_storage_span_bytes(op->src[0], lhs_byte_count) || + !tensor_storage_span_bytes(op->src[1], rhs_byte_count)) { + return false; + } + if (!binary_output_storage_is_safe(op->src[0], lhs_byte_count, op->src[1], rhs_byte_count, op)) { + return false; + } + + ggml::hrx::BinaryKind binary_kind; + if (!ggml::hrx::import_binary_kind(*op, binary_kind)) { + return false; + } + + if (!ggml::hrx::binary_kind_supported(binary_kind)) { + return false; + } + + if ((!ggml_are_same_shape(op->src[0], op) || !ggml_are_same_shape(op->src[1], op)) && + (!ggml_is_contiguous(op->src[0]) || !ggml_is_contiguous(op->src[1]) || + !tensor_has_packed_f32_layout(op->src[0]) || !tensor_has_packed_f32_layout(op->src[1]))) { + return false; + } + + if (!binary_kind_allows_broadcast(binary_kind, op->src[0], op->src[1], op)) { + return false; + } + + return true; +} + +static bool supported_scale_f32_tensor(const ggml_tensor * op) { + return op != nullptr && op->op == GGML_OP_SCALE && op->src[0] != nullptr && op->type == GGML_TYPE_F32 && + op->src[0]->type == GGML_TYPE_F32 && ggml_are_same_shape(op, op->src[0]) && tensor_has_positive_shape(op) && + ggml_is_contiguous(op) && ggml_is_contiguous(op->src[0]) && tensor_has_packed_f32_layout(op) && + tensor_has_packed_f32_layout(op->src[0]) && scale_storage_is_safe(op->src[0], op) && + static_cast(ggml_nelements(op)) <= std::numeric_limits::max(); +} + +static bool supported_qwen_attention_projection_get_rows_tensor(const ggml_tensor * op) { + if (op == nullptr || op->op != GGML_OP_GET_ROWS || op->src[0] == nullptr || op->src[1] == nullptr || + op->src[0]->op != GGML_OP_MUL_MAT || op->type != GGML_TYPE_F32 || op->src[0]->type != GGML_TYPE_F32 || + op->src[1]->type != GGML_TYPE_I32 || !ggml_is_contiguous(op) || !ggml_is_contiguous(op->src[0]) || + !ggml_is_contiguous(op->src[1]) || op->ne[0] != op->src[0]->ne[0] || op->ne[1] <= 0 || + op->src[1]->ne[0] != op->ne[1] || op->ne[2] != 1 || op->ne[3] != 1) { + return false; + } + return true; +} + +static bool supported_qwen_attention_residual_add_tensor(const ggml_tensor * op) { + if (op == nullptr || op->op != GGML_OP_ADD || op->src[0] == nullptr || op->src[1] == nullptr || + op->type != GGML_TYPE_F32 || op->src[0]->type != GGML_TYPE_F32 || op->src[1]->type != GGML_TYPE_F32 || + !ggml_is_contiguous(op) || !ggml_is_contiguous(op->src[0]) || !ggml_is_contiguous(op->src[1]) || + !ggml_are_same_shape(op, op->src[0]) || !ggml_are_same_shape(op, op->src[1])) { + return false; + } + + return op->src[0]->op == GGML_OP_GET_ROWS || op->src[1]->op == GGML_OP_GET_ROWS || + op->src[0]->op == GGML_OP_MUL_MAT || op->src[1]->op == GGML_OP_MUL_MAT; +} + +static bool supported_qwen_routed_ffn_reduce_add_tensor(const ggml_tensor * op) { + if (op == nullptr || op->op != GGML_OP_ADD || op->src[0] == nullptr || op->src[1] == nullptr || + op->type != GGML_TYPE_F32 || op->src[0]->type != GGML_TYPE_F32 || op->src[1]->type != GGML_TYPE_F32 || + !ggml_is_contiguous(op) || !ggml_are_same_shape(op, op->src[0]) || !ggml_are_same_shape(op, op->src[1]) || + op->ne[0] != 2048 || op->ne[3] != 1 || + !((op->ne[1] > 0 && op->ne[2] == 1) || (op->ne[1] == 1 && op->ne[2] > 0))) { + return false; + } + + return op->src[0]->op == GGML_OP_VIEW || op->src[1]->op == GGML_OP_VIEW || op->src[0]->op == GGML_OP_ADD || + op->src[1]->op == GGML_OP_ADD; +} + +static bool supported_unary_f32_tensor(const ggml_tensor * op) { + if (op == nullptr || op->src[0] == nullptr || op->type != GGML_TYPE_F32 || op->src[0]->type != GGML_TYPE_F32 || + !ggml_are_same_shape(op, op->src[0]) || !ggml_is_contiguous(op) || !ggml_is_contiguous(op->src[0])) { + return false; + } + + ggml::hrx::UnaryKind unary_kind; + if (!ggml::hrx::import_unary_kind(*op, unary_kind)) { + return false; + } + + return ggml::hrx::unary_kind_supported(unary_kind); +} + +static bool device_supports_op(ggml_backend_dev_t device, const ggml_tensor * op) { + if (op == nullptr) { + return false; + } + if (!ggml::hrx::moe_tail_claimable(device, op)) { + return false; + } + // GET_ROWS gives wrong rows for batched IQ3_S sources (test-backend-ops GET_ROWS) + if (op->op == GGML_OP_GET_ROWS && op->src[0] != nullptr && + op->src[0]->type == GGML_TYPE_IQ3_S && op->src[0]->ne[2] * op->src[0]->ne[3] > 1) { + return false; + } + if (zero_output_elision_supported(op)) { + return true; + } + const bool supported_binary = supported_binary_f32_tensor(op); + if (supported_binary) { + return true; + } + if (supported_scale_f32_tensor(op)) { + return true; + } + if (supported_qwen_attention_projection_get_rows_tensor(op)) { + return true; + } + if (supported_qwen_attention_residual_add_tensor(op)) { + return true; + } + if (supported_qwen_routed_ffn_reduce_add_tensor(op)) { + return true; + } + const bool supported_unary = supported_unary_f32_tensor(op); + if (supported_unary) { + return true; + } + if (op->op == GGML_OP_GET_ROWS && op->src[0] != nullptr && + (op->src[0]->type == GGML_TYPE_IQ3_S || op->src[0]->type == GGML_TYPE_IQ4_NL)) { + return false; + } + if (op->op == GGML_OP_NONE) { + return true; + } + // a node already placed in HRX device memory (for example a KV cache view) cannot move to another backend + const ggml_tensor * placed = op->view_src != nullptr ? op->view_src : op; + if (placed->buffer != nullptr && ggml_backend_buffer_get_type(placed->buffer) == &device_context(device)->buft) { + return eager_capability_declared(op->op); + } + // otherwise claim only nodes the dispatcher can execute, so the rest falls back to another backend instead of failing the graph + // a node that only a fused dispatch executes keeps the per-op claim inside a model graph (fused-context-claim.h) + return eager_capability_declared(op->op) && + (ggml::hrx::fused_context_claim(op) || + ggml::hrx::can_execute_standalone_op_as_graph(op, device_context(device)->architecture)); +} + +static bool device_supports_buffer_type(ggml_backend_dev_t device, ggml_backend_buffer_type_t buft) { + auto * context = device_context(device); + return buft == &context->buft || buft == &context->host_buft || ggml_backend_buft_is_host(buft); +} + +static const ggml_backend_device_i device_i = { + device_name, + device_description, + device_memory, + device_type, + device_props, + device_init, + device_buffer_type, + device_host_buffer_type, + nullptr, + device_supports_op, + device_supports_buffer_type, + nullptr, + nullptr, + nullptr, + nullptr, +}; + +static const char * registry_name(ggml_backend_reg_t registry) { + GGML_UNUSED(registry); + return "HRX"; +} + +static size_t registry_device_count(ggml_backend_reg_t registry) { + return static_cast(registry->context)->devices.size(); +} + +static ggml_backend_dev_t registry_device(ggml_backend_reg_t registry, size_t index) { + auto * context = static_cast(registry->context); + GGML_ASSERT(index < context->devices.size()); + return &context->devices[index]; +} + +// [1bit] zero-copy sharing: proc "ggml_backend_buffer_export_dmabuf" (ggml-hrx-dmabuf.cpp) +static bool ggml_backend_hrx_buffer_export_dmabuf(ggml_backend_buffer_t buffer, int * fd, size_t * offset) { + void * device_ptr = nullptr; + return buffer != nullptr && buffer->iface.get_base == ggml_backend_hrx_buffer_base && + HRX_CHECK(hrx_buffer_get_device_ptr(buffer_context(buffer)->buffer, &device_ptr)) && + ggml::hrx::export_dmabuf(device_ptr, buffer->size, fd, offset); +} + +static void * registry_proc(ggml_backend_reg_t registry, const char * name) { + if (std::strcmp(name, "ggml_backend_buffer_export_dmabuf") == 0) { + return (void *) ggml_backend_hrx_buffer_export_dmabuf; + } + GGML_UNUSED(registry); + GGML_UNUSED(name); + return nullptr; +} + +static const ggml_backend_reg_i registry_i = { registry_name, registry_device_count, registry_device, registry_proc }; + +static std::unique_ptr create_registry_context() { + auto context = std::make_unique(); + if (environment_flag_enabled("GGML_HRX_LOG_DEVICE_EVENTS")) { + hrx_device_event_sink_t sink = { log_hrx_device_event, nullptr }; + if (!HRX_CHECK(hrx_runtime_set_device_event_sink(sink))) { + return context; + } + } + hrx_status_t status = hrx_gpu_initialize(0); + if (hrx_status_is_ok(status)) { + context->initialized = true; + } else if (hrx_status_code(status) == HRX_STATUS_ALREADY_EXISTS) { + hrx_status_ignore(status); + } else { + hrx_status_ignore(status); + return context; + } + int count = 0; + if (!HRX_CHECK(hrx_gpu_device_count(&count))) { + return context; + } + context->device_contexts.reserve(count); + context->devices.reserve(count); + for (int i = 0; i < count; ++i) { + hrx_device_t hrx_device = nullptr; + if (!HRX_CHECK(hrx_gpu_device_get(i, &hrx_device)) || hrx_device == nullptr) { + continue; + } + hrx_device_retain(hrx_device); + auto device_ctx = std::make_unique(); + device_ctx->device = hrx_device; + device_ctx->name = "HRX" + std::to_string(i); + device_ctx->use_direct_host_bindings = environment_flag_enabled("GGML_HRX_USE_UNIFIED_MEMORY"); + if (device_ctx->use_direct_host_bindings) { + GGML_LOG_INFO("ggml_hrx: direct coherent host bindings enabled by GGML_HRX_USE_UNIFIED_MEMORY\n"); + } + const std::optional name = + device_string_property(hrx_device, HRX_DEVICE_PROPERTY_NAME, "query HRX device name"); + const std::optional architecture = + device_string_property(hrx_device, HRX_DEVICE_PROPERTY_ARCHITECTURE, "query HRX device architecture"); + if (!name || !architecture) { + hrx_device_release(hrx_device); + continue; + } + uint64_t memory = 0; + if (!HRX_CHECK( + hrx_device_get_property(hrx_device, HRX_DEVICE_PROPERTY_TOTAL_MEMORY, &memory, sizeof(memory)))) { + hrx_device_release(hrx_device); + continue; + } + if (!HRX_CHECK(hrx_stream_create(hrx_device, 0, &device_ctx->buffer_stream))) { + hrx_device_release(hrx_device); + continue; + } + device_ctx->memory_total = static_cast(memory); + device_ctx->description = *name + " (" + *architecture + ")"; + device_ctx->architecture = *architecture; + device_ctx->buft_context = { device_ctx.get(), device_ctx->name, false }; + device_ctx->buft = { buffer_type_i, nullptr, &device_ctx->buft_context }; + device_ctx->host_buft_context = { device_ctx.get(), device_ctx->name + "_HOST", true }; + device_ctx->host_buft = { buffer_type_i, nullptr, &device_ctx->host_buft_context }; + context->device_contexts.emplace_back(std::move(device_ctx)); + context->devices.push_back({ device_i, nullptr, context->device_contexts.back().get() }); + context->device_contexts.back()->buft.device = &context->devices.back(); + context->device_contexts.back()->host_buft.device = &context->devices.back(); + } + return context; +} + +} // namespace + +ggml_backend_reg_t ggml_backend_hrx_reg() { + static std::unique_ptr context = create_registry_context(); + static ggml_backend_reg registry = { GGML_BACKEND_API_VERSION, registry_i, context.get() }; + for (auto & device : context->devices) { + device.reg = ®istry; + } + return ®istry; +} + +ggml_backend_t ggml_backend_hrx_init(size_t device) { + ggml_backend_reg_t registry = ggml_backend_hrx_reg(); + if (device >= ggml_backend_reg_dev_count(registry)) { + return nullptr; + } + return ggml_backend_dev_init(ggml_backend_reg_dev_get(registry, device), nullptr); +} + +bool ggml_backend_is_hrx(ggml_backend_t backend) { + return backend != nullptr && ggml_guid_matches(backend->guid, ggml_backend_hrx_guid()); +} + +bool ggml_backend_hrx_get_cache_stats(ggml_backend_t backend, ggml_backend_hrx_cache_stats * stats) { + if (!ggml_backend_is_hrx(backend) || stats == nullptr) { + return false; + } + auto * context = static_cast(backend->context); + const ggml::hrx::GraphProgramCacheStats graph_stats = context->graph_programs.stats(); + const ggml::hrx::PreparedCommandProgramCacheStats prepared_stats = context->prepared_programs.stats(); + stats->graph_program_builds = graph_stats.builds; + stats->graph_program_hits = graph_stats.hits; + stats->prepared_program_builds = graph_stats.prepared_program_builds + prepared_stats.builds; + stats->prepared_program_hits = graph_stats.prepared_program_hits + prepared_stats.hits; + return true; +} + +int ggml_backend_hrx_get_device_count() { + return static_cast(ggml_backend_reg_dev_count(ggml_backend_hrx_reg())); +} + +ggml_backend_buffer_type_t ggml_backend_hrx_buffer_type(size_t device) { + ggml_backend_reg_t registry = ggml_backend_hrx_reg(); + return device < ggml_backend_reg_dev_count(registry) ? + ggml_backend_dev_buffer_type(ggml_backend_reg_dev_get(registry, device)) : + nullptr; +} + +GGML_BACKEND_DL_IMPL(ggml_backend_hrx_reg) diff --git a/ggml/src/ggml-hrx/graph/graph-diagnostics.cpp b/ggml/src/ggml-hrx/graph/graph-diagnostics.cpp new file mode 100644 index 000000000000..8c0427ffaaa1 --- /dev/null +++ b/ggml/src/ggml-hrx/graph/graph-diagnostics.cpp @@ -0,0 +1,556 @@ +#include "graph-diagnostics.h" + +#include "ggml-impl.h" +#include "ggml.h" + +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +using json = nlohmann::ordered_json; + +const char * value_kind_name(ValueKind kind) { + switch (kind) { + case ValueKind::External: + return "external"; + case ValueKind::Transient: + return "transient"; + } + return "unknown"; +} + +bool parse_value_kind(const std::string & name, ValueKind & kind) { + if (name == "external") { + kind = ValueKind::External; + return true; + } + if (name == "transient") { + kind = ValueKind::Transient; + return true; + } + return false; +} + +const char * match_kind_name(DispatchMatchKind kind) { + switch (kind) { + case DispatchMatchKind::Fused: + return "fused"; + case DispatchMatchKind::SingleOp: + return "single_op"; + } + return "unknown"; +} + +const char * dispatch_source_name(DispatchSource source) { + switch (source) { + case DispatchSource::Common: + return "common"; + case DispatchSource::Llm: + return "llm"; + case DispatchSource::Qwen: + return "qwen"; + } + return "unknown"; +} + +json dims_json(const std::array & values) { + json result = json::array(); + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + result.push_back(values[i]); + } + return result; +} + +json strides_json(const std::array & values) { + json result = json::array(); + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + result.push_back(values[i]); + } + return result; +} + +Status read_dims(const json & item, const char * name, std::array & values) { + Status status; + if (!item.contains(name) || !item[name].is_array() || item[name].size() != GGML_MAX_DIMS) { + status.log("snapshot array %s must contain %d values", name, GGML_MAX_DIMS); + return status; + } + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + values[i] = item[name][i].get(); + } + return status; +} + +Status read_strides(const json & item, const char * name, std::array & values) { + Status status; + if (!item.contains(name) || !item[name].is_array() || item[name].size() != GGML_MAX_DIMS) { + status.log("snapshot array %s must contain %d values", name, GGML_MAX_DIMS); + return status; + } + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + values[i] = item[name][i].get(); + } + return status; +} + +json op_params_json(const OpParams & params) { + return std::visit( + [](const auto & value) -> json { + using T = std::decay_t; + if constexpr (std::is_same_v) { + return { + { "kind", "none" } + }; + } else if constexpr (std::is_same_v) { + return { + { "kind", "rms_norm" }, + { "eps", value.eps } + }; + } else if constexpr (std::is_same_v) { + return { + { "kind", "flash_attn_ext" }, + { "scale", value.scale }, + { "max_bias", value.max_bias }, + { "logit_softcap", value.logit_softcap }, + { "prec", static_cast(value.prec) }, + }; + } else if constexpr (std::is_same_v) { + return { + { "kind", "soft_max" }, + { "scale", value.scale }, + { "max_bias", value.max_bias } + }; + } else if constexpr (std::is_same_v) { + return { + { "kind", "argsort" }, + { "order", static_cast(value.order) } + }; + } else if constexpr (std::is_same_v) { + return { + { "kind", "clamp" }, + { "min", value.min }, + { "max", value.max } + }; + } else if constexpr (std::is_same_v) { + return { + { "kind", "glu" }, + { "op", static_cast(value.op) }, + { "swapped", value.swapped }, + { "alpha", value.alpha }, + { "limit", value.limit } + }; + } else if constexpr (std::is_same_v) { + return { + { "kind", "scale" }, + { "scale", value.scale }, + { "bias", value.bias } + }; + } else if constexpr (std::is_same_v) { + return { + { "kind", "binary" }, + { "op", static_cast(value.op) } + }; + } else if constexpr (std::is_same_v) { + return { + { "kind", "unary" }, + { "op", static_cast(value.op) } + }; + } else if constexpr (std::is_same_v) { + return { + { "kind", "mul_mat" }, + { "hint", value.hint } + }; + } else if constexpr (std::is_same_v) { + return { + { "kind", "rope" }, + { "n_dims", value.n_dims }, + { "mode", value.mode }, + { "n_ctx_orig", value.n_ctx_orig }, + { "freq_base", value.freq_base }, + { "freq_scale", value.freq_scale }, + { "ext_factor", value.ext_factor }, + { "attn_factor", value.attn_factor }, + { "beta_fast", value.beta_fast }, + { "beta_slow", value.beta_slow }, + { "sections", value.sections }, + }; + } + }, + params); +} + +OpParams parse_op_params(const json & item) { + const std::string kind = item.value("kind", "none"); + if (kind == "rms_norm") { + return RmsNormParams{ item.value("eps", 0.0f) }; + } + if (kind == "flash_attn_ext") { + return FlashAttnExtParams{ + item.value("scale", 0.0f), + item.value("max_bias", 0.0f), + item.value("logit_softcap", 0.0f), + static_cast(item.value("prec", static_cast(GGML_PREC_DEFAULT))), + }; + } + if (kind == "soft_max") { + return SoftMaxParams{ item.value("scale", 0.0f), item.value("max_bias", 0.0f) }; + } + if (kind == "argsort") { + return ArgsortParams{ static_cast( + item.value("order", static_cast(GGML_SORT_ORDER_ASC))) }; + } + if (kind == "clamp") { + const float min = item.contains("min") && !item["min"].is_null() ? item["min"].get() : 0.0f; + const float max = item.contains("max") && !item["max"].is_null() ? item["max"].get() : + std::numeric_limits::infinity(); + return ClampParams{ min, max }; + } + if (kind == "glu") { + return GluParams{ static_cast(item.value("op", static_cast(GGML_GLU_OP_REGLU))), + item.value("swapped", false), item.value("alpha", 0.0f), item.value("limit", 0.0f) }; + } + if (kind == "mul_mat") { + return MulMatParams{ item.value("hint", 0) }; + } + if (kind == "scale") { + return ScaleParams{ item.value("scale", 0.0f), item.value("bias", 0.0f) }; + } + if (kind == "binary") { + return BinaryParams{ static_cast(item.value("op", static_cast(BinaryKind::Add))) }; + } + if (kind == "unary") { + return UnaryParams{ static_cast(item.value("op", static_cast(UnaryKind::Abs))) }; + } + if (kind == "rope") { + std::array sections = {}; + if (item.contains("sections")) { + sections = item["sections"].get>(); + } + return RopeParams{ + item.value("n_dims", 0), item.value("mode", 0), + item.value("n_ctx_orig", 0), item.value("freq_base", 0.0f), + item.value("freq_scale", 0.0f), item.value("ext_factor", 0.0f), + item.value("attn_factor", 0.0f), item.value("beta_fast", 0.0f), + item.value("beta_slow", 0.0f), sections, + }; + } + return std::monostate{}; +} + +json value_json(const Value & value) { + return { + { "id", value.id.value }, + { "kind", value_kind_name(value.kind) }, + { "storage", value.storage.value }, + { "storage_root", value.storage_root.value }, + { "alias_source", value.alias_source.value }, + { "storage_offset", value.storage_offset }, + { "storage_byte_count", value.storage_byte_count }, + { "type", static_cast(value.type) }, + { "type_name", ggml_type_name(value.type) }, + { "ne", dims_json(value.ne) }, + { "nb", strides_json(value.nb) }, + { "element_count", value.element_count }, + { "byte_count", value.byte_count }, + { "contiguous", value.contiguous }, + }; +} + +json storage_json(const ValueStorage & storage) { + return { + { "id", storage.id.value }, + { "root", storage.root.value }, + { "byte_count", storage.byte_count }, + }; +} + +json node_json(const GraphNode & node) { + json inputs = json::array(); + for (ValueId input : node.inputs) { + inputs.push_back(input.value); + } + return { + { "op", static_cast(node.op) }, + { "op_name", ggml_op_name(node.op) }, + { "output", node.output.value }, + { "inputs", std::move(inputs) }, + { "params", op_params_json(node.params) }, + }; +} + +std::string value_summary(const Graph & graph, ValueId id) { + std::ostringstream out; + const Value * value = graph.values().find(id); + if (value == nullptr) { + out << id.value << ":missing"; + return out.str(); + } + out << id.value << ":" << ggml_type_name(value->type) << "[" << value->ne[0] << "," << value->ne[1] << "," + << value->ne[2] << "," << value->ne[3] << "] " << value_kind_name(value->kind); + if (value->alias_source.value >= 0) { + out << " alias=" << value->alias_source.value << " storage_offset=" << value->storage_offset; + } + return out.str(); +} + +json attempts_json(const DispatchMatchDiagnostics & diagnostics) { + json attempts = json::array(); + for (const DispatchRegistrationAttempt & attempt : diagnostics.attempts) { + json covered = json::array(); + for (size_t node : attempt.covered_nodes) { + covered.push_back(node); + } + attempts.push_back({ + { "name", attempt.name }, + { "root_op", static_cast(attempt.root_op) }, + { "root_op_name", ggml_op_name(attempt.root_op) }, + { "kind", match_kind_name(attempt.kind) }, + { "priority", attempt.priority }, + { "source", dispatch_source_name(attempt.source) }, + { "matched", attempt.matched }, + { "covered_nodes", std::move(covered) }, + { "errors", attempt.errors }, + }); + } + return attempts; +} + +void write_file(const std::filesystem::path & path, const std::string & contents) { + std::filesystem::create_directories(path.parent_path()); + std::ofstream output(path, std::ios::binary | std::ios::trunc); + if (!output) { + throw std::runtime_error("cannot create " + path.string()); + } + output << contents; + if (contents.empty() || contents.back() != '\n') { + output << '\n'; + } +} + +} // namespace + +std::string serialize_graph_snapshot_json(const Graph & graph, const std::string & target, uint64_t uid) { + json values = json::array(); + for (const Value & value : graph.values().values()) { + values.push_back(value_json(value)); + } + json storages = json::array(); + for (const ValueStorage & storage : graph.values().storages()) { + storages.push_back(storage_json(storage)); + } + json nodes = json::array(); + for (const GraphNode & node : graph.nodes()) { + nodes.push_back(node_json(node)); + } + json root = { + { "schema", "ggml-hrx-graph-snapshot-v1" }, + { "uid", uid }, + { "target", target }, + { "storages", std::move(storages) }, + { "values", std::move(values) }, + { "nodes", std::move(nodes) }, + }; + return root.dump(2); +} + +std::string format_graph_snapshot_text(const Graph & graph, const std::string & target, uint64_t uid) { + std::ostringstream out; + out << "schema=ggml-hrx-graph-snapshot-v1\n"; + out << "uid=" << uid << "\ntarget=" << target << "\nvalues=" << graph.values().size() + << "\nnodes=" << graph.nodes().size() << '\n'; + for (size_t i = 0; i < graph.nodes().size(); ++i) { + const GraphNode & node = graph.nodes()[i]; + out << "node " << i << " " << ggml_op_name(node.op) << " output=" << value_summary(graph, node.output) + << " inputs=["; + for (size_t j = 0; j < node.inputs.size(); ++j) { + if (j > 0) { + out << ", "; + } + out << value_summary(graph, node.inputs[j]); + } + out << "] consumers=["; + const std::vector & consumers = graph.index().consumers(node.output); + for (size_t j = 0; j < consumers.size(); ++j) { + size_t consumer_index = 0; + if (j > 0) { + out << ", "; + } + if (graph.index().node_index(consumers[j], consumer_index)) { + out << consumer_index << ":" << ggml_op_name(consumers[j]->op); + } + } + out << "]\n"; + } + return out.str(); +} + +GraphSnapshotLoadResult load_graph_snapshot_json(const std::string & contents) { + GraphSnapshotLoadResult result; + try { + const json root = json::parse(contents); + if (root.value("schema", "") != "ggml-hrx-graph-snapshot-v1") { + result.status.log("unsupported graph snapshot schema"); + return result; + } + result.uid = root.value("uid", 0ULL); + result.target = root.value("target", ""); + + ValueMap & values = result.graph.values(); + for (const json & item : root.at("storages")) { + Status status = values.add_snapshot_storage({ + ValueStorageId(item.at("id").get()), + ValueId(item.at("root").get()), + item.at("byte_count").get(), + }); + if (!status.success()) { + result.status.append(status); + return result; + } + } + for (const json & item : root.at("values")) { + Value value = { + ValueId(item.at("id").get()), + ValueKind::External, + ValueStorageId(item.at("storage").get()), + ValueId(item.at("storage_root").get()), + ValueId(item.at("alias_source").get()), + item.at("storage_offset").get(), + item.at("storage_byte_count").get(), + static_cast(item.at("type").get()), + {}, + {}, + item.at("element_count").get(), + item.at("byte_count").get(), + item.at("contiguous").get(), + nullptr, + std::nullopt, + }; + if (!parse_value_kind(item.at("kind").get(), value.kind)) { + result.status.log("snapshot value %d has unknown kind", value.id.value); + return result; + } + Status dims_status = read_dims(item, "ne", value.ne); + if (!dims_status.success()) { + result.status.append(dims_status); + return result; + } + Status strides_status = read_strides(item, "nb", value.nb); + if (!strides_status.success()) { + result.status.append(strides_status); + return result; + } + Status status = values.add_snapshot_value(std::move(value)); + if (!status.success()) { + result.status.append(status); + return result; + } + } + for (const json & item : root.at("nodes")) { + std::vector inputs; + for (const json & input : item.at("inputs")) { + inputs.push_back(ValueId(input.get())); + } + GraphNode & node = result.graph.add_node(static_cast(item.at("op").get()), + ValueId(item.at("output").get()), std::move(inputs)); + node.params = parse_op_params(item.at("params")); + } + result.status.append(result.graph.build_index()); + } catch (const std::exception & error) { + result.status.log("failed to load graph snapshot: %s", error.what()); + } + return result; +} + +Status write_graph_snapshot(const std::filesystem::path & directory, + const Graph & graph, + const std::string & target, + uint64_t uid) { + Status status; + try { + static std::atomic sequence{ 0 }; + const uint64_t id = sequence.fetch_add(1); + std::ostringstream name; + name << "graph-" << id << "-uid-" << uid << "-" << target << "-" << graph.nodes().size() << "-nodes"; + const std::filesystem::path base = directory / name.str(); + write_file(base.string() + ".json", serialize_graph_snapshot_json(graph, target, uid)); + write_file(base.string() + ".txt", format_graph_snapshot_text(graph, target, uid)); + } catch (const std::exception & error) { + status.log("failed to write HRX graph snapshot: %s", error.what()); + } + return status; +} + +std::string format_schedule_diagnostics_text(const Graph & graph, + const CommandPlan & plan, + const DispatchScheduleDiagnostics & diagnostics) { + std::ostringstream out; + out << "valid=" << (plan.valid() ? "true" : "false") << '\n'; + for (const std::string & error : plan.status.errors()) { + out << "error=" << error << '\n'; + } + if (diagnostics.unsupported_node != nullptr) { + out << "unsupported_node=" << diagnostics.unsupported_node_index << ":" + << ggml_op_name(diagnostics.unsupported_node->op) << '\n'; + out << "unsupported_message=" << diagnostics.unsupported_message << '\n'; + out << "output=" << value_summary(graph, diagnostics.unsupported_node->output) << '\n'; + for (size_t i = 0; i < diagnostics.unsupported_node->inputs.size(); ++i) { + out << "input" << i << "=" << value_summary(graph, diagnostics.unsupported_node->inputs[i]) << '\n'; + } + } + out << "matcher_attempts=" << diagnostics.match.attempts.size() << '\n'; + for (const DispatchRegistrationAttempt & attempt : diagnostics.match.attempts) { + out << "attempt name=" << attempt.name << " kind=" << match_kind_name(attempt.kind) + << " source=" << dispatch_source_name(attempt.source) << " priority=" << attempt.priority + << " matched=" << (attempt.matched ? "true" : "false") << " covered=["; + for (size_t i = 0; i < attempt.covered_nodes.size(); ++i) { + if (i > 0) { + out << ","; + } + out << attempt.covered_nodes[i]; + } + out << "]\n"; + for (const std::string & error : attempt.errors) { + out << " error=" << error << '\n'; + } + } + return out.str(); +} + +std::string serialize_schedule_diagnostics_json(const Graph & graph, + const CommandPlan & plan, + const DispatchScheduleDiagnostics & diagnostics) { + json errors = json::array(); + for (const std::string & error : plan.status.errors()) { + errors.push_back(error); + } + json root = { + { "schema", "ggml-hrx-schedule-diagnostics-v1" }, + { "valid", plan.valid() }, + { "errors", std::move(errors) }, + { "matcher_attempts", attempts_json(diagnostics.match) }, + }; + if (diagnostics.unsupported_node != nullptr) { + json inputs = json::array(); + for (ValueId input : diagnostics.unsupported_node->inputs) { + inputs.push_back(input.value); + } + root["unsupported_node"] = { + { "index", diagnostics.unsupported_node_index }, + { "op", static_cast(diagnostics.unsupported_node->op) }, + { "op_name", ggml_op_name(diagnostics.unsupported_node->op) }, + { "output", diagnostics.unsupported_node->output.value }, + { "inputs", std::move(inputs) }, + { "message", diagnostics.unsupported_message }, + }; + } + GGML_UNUSED(graph); + return root.dump(2); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/graph/graph-diagnostics.h b/ggml/src/ggml-hrx/graph/graph-diagnostics.h new file mode 100644 index 000000000000..9956301ca9a8 --- /dev/null +++ b/ggml/src/ggml-hrx/graph/graph-diagnostics.h @@ -0,0 +1,39 @@ +#pragma once + +#include "dispatch/dispatch-scheduler.h" +#include "graph.h" +#include "status.h" + +#include +#include +#include + +namespace ggml::hrx { + +struct GraphSnapshotLoadResult { + uint64_t uid = 0; + std::string target; + Graph graph; + Status status; + + bool valid() const { return status.success(); } +}; + +std::string serialize_graph_snapshot_json(const Graph & graph, const std::string & target, uint64_t uid); +std::string format_graph_snapshot_text(const Graph & graph, const std::string & target, uint64_t uid); + +GraphSnapshotLoadResult load_graph_snapshot_json(const std::string & contents); + +Status write_graph_snapshot(const std::filesystem::path & directory, + const Graph & graph, + const std::string & target, + uint64_t uid); + +std::string format_schedule_diagnostics_text(const Graph & graph, + const CommandPlan & plan, + const DispatchScheduleDiagnostics & diagnostics); +std::string serialize_schedule_diagnostics_json(const Graph & graph, + const CommandPlan & plan, + const DispatchScheduleDiagnostics & diagnostics); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/graph/graph-matcher.cpp b/ggml/src/ggml-hrx/graph/graph-matcher.cpp new file mode 100644 index 000000000000..0d466f8d5e0f --- /dev/null +++ b/ggml/src/ggml-hrx/graph/graph-matcher.cpp @@ -0,0 +1,116 @@ +#include "graph-matcher.h" + +namespace ggml::hrx { +namespace { + +static bool node_in_list(const GraphNode * node, const std::vector & nodes) { + for (const GraphNode * candidate : nodes) { + if (candidate == node) { + return true; + } + } + return false; +} + +static void append_unique_node(std::vector & nodes, const GraphNode * node) { + if (node != nullptr && !node_in_list(node, nodes)) { + nodes.push_back(node); + } +} + +} // namespace + +std::vector layout_alias_consumers(const Graph & graph, ValueId value) { + std::vector matches; + for (const GraphNode * consumer : graph.index().consumers(value)) { + if (consumer != nullptr && is_layout_alias_node(graph, *consumer)) { + matches.push_back(consumer); + } + } + return matches; +} + +std::vector layout_alias_consumers_with_op(const Graph & graph, ValueId value, ggml_op op) { + std::vector matches; + for (const GraphNode * consumer : layout_alias_consumers(graph, value)) { + if (consumer->op == op) { + matches.push_back(consumer); + } + } + return matches; +} + +const GraphNode * find_single_layout_alias_consumer(const Graph & graph, ValueId value) { + const std::vector matches = layout_alias_consumers(graph, value); + return matches.size() == 1 ? matches.front() : nullptr; +} + +const GraphNode * find_single_layout_alias_consumer_with_op(const Graph & graph, ValueId value, ggml_op op) { + const std::vector matches = layout_alias_consumers_with_op(graph, value, op); + return matches.size() == 1 ? matches.front() : nullptr; +} + +std::vector consumers_with_op_through_layout_aliases(const Graph & graph, + ValueId value, + ggml_op op) { + std::vector matches; + for (const GraphNode * consumer : graph.index().consumers(value)) { + if (consumer == nullptr) { + continue; + } + if (consumer->op == op) { + append_unique_node(matches, consumer); + } + if (!is_layout_alias_node(graph, *consumer)) { + continue; + } + for (const GraphNode * alias_consumer : graph.index().consumers(consumer->output)) { + if (alias_consumer != nullptr && alias_consumer->op == op) { + append_unique_node(matches, alias_consumer); + } + } + } + return matches; +} + +const GraphNode * find_single_consumer_with_op_through_layout_aliases(const Graph & graph, ValueId value, ggml_op op) { + const std::vector matches = consumers_with_op_through_layout_aliases(graph, value, op); + return matches.size() == 1 ? matches.front() : nullptr; +} + +bool node_has_input_or_alias(const Graph & graph, const GraphNode & node, ValueId input) { + const Value * input_value = graph.values().find(input); + for (ValueId candidate : node.inputs) { + if (candidate == input) { + return true; + } + const Value * candidate_value = graph.values().find(candidate); + if (candidate_value != nullptr && candidate_value->alias_source == input) { + return true; + } + if (input_value != nullptr && input_value->alias_source == candidate) { + return true; + } + } + return false; +} + +bool append_covered_node_index_once(const Graph & graph, + const std::vector & covered_nodes, + const GraphNode * node, + std::vector & covered_indices) { + size_t index = 0; + if (node == nullptr || !graph.index().node_index(node, index) || index >= covered_nodes.size() || + covered_nodes[index]) { + return false; + } + for (const size_t covered : covered_indices) { + if (covered == index) { + return true; + } + } + covered_indices.push_back(index); + return true; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/graph/graph-matcher.h b/ggml/src/ggml-hrx/graph/graph-matcher.h new file mode 100644 index 000000000000..507165d45eff --- /dev/null +++ b/ggml/src/ggml-hrx/graph/graph-matcher.h @@ -0,0 +1,25 @@ +#pragma once + +#include "graph.h" + +#include +#include + +namespace ggml::hrx { + +std::vector layout_alias_consumers(const Graph & graph, ValueId value); +std::vector layout_alias_consumers_with_op(const Graph & graph, ValueId value, ggml_op op); + +const GraphNode * find_single_layout_alias_consumer(const Graph & graph, ValueId value); +const GraphNode * find_single_layout_alias_consumer_with_op(const Graph & graph, ValueId value, ggml_op op); + +std::vector consumers_with_op_through_layout_aliases(const Graph & graph, ValueId value, ggml_op op); +const GraphNode * find_single_consumer_with_op_through_layout_aliases(const Graph & graph, ValueId value, ggml_op op); + +bool node_has_input_or_alias(const Graph & graph, const GraphNode & node, ValueId input); +bool append_covered_node_index_once(const Graph & graph, + const std::vector & covered_nodes, + const GraphNode * node, + std::vector & covered_indices); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/graph/graph-traversal.cpp b/ggml/src/ggml-hrx/graph/graph-traversal.cpp new file mode 100644 index 000000000000..22edc6bba011 --- /dev/null +++ b/ggml/src/ggml-hrx/graph/graph-traversal.cpp @@ -0,0 +1,163 @@ +#include "graph-traversal.h" + +#include +#include + +namespace ggml::hrx { +namespace { + +static bool is_root_op(ggml_op op) { + switch (op) { + case GGML_OP_MUL_MAT: + case GGML_OP_MUL_MAT_ID: + case GGML_OP_FLASH_ATTN_EXT: + case GGML_OP_CONV_TRANSPOSE_1D: + case GGML_OP_CONV_2D: + case GGML_OP_CONV_3D: + case GGML_OP_CONV_2D_DW: + case GGML_OP_CONV_TRANSPOSE_2D: + case GGML_OP_SSM_CONV: + return true; + default: + return false; + } +} + +static bool is_fusable_followup_op(ggml_op op) { + switch (op) { + case GGML_OP_ADD: + case GGML_OP_MUL: + return true; + default: + return false; + } +} + +static void erase_ready(size_t node_index, std::set & root_queue, std::set & regular_queue) { + root_queue.erase(node_index); + regular_queue.erase(node_index); +} + +static void add_ready_node(const GraphNode & node, + size_t node_index, + const std::vector & selected, + std::set & root_queue, + std::set & regular_queue) { + if (node_index >= selected.size() || selected[node_index]) { + return; + } + if (is_root_op(node.op)) { + root_queue.insert(node_index); + } else { + regular_queue.insert(node_index); + } +} + +static bool select_merge_candidate(const Graph & graph, + size_t selected_node, + const std::vector & pending_inputs, + const std::vector & selected, + size_t & next_node) { + const std::vector & nodes = graph.nodes(); + if (selected_node >= nodes.size()) { + return false; + } + + bool found = false; + size_t best_index = 0; + for (const GraphNode * consumer : graph.index().consumers(nodes[selected_node].output)) { + size_t consumer_index = 0; + if (consumer == nullptr || !graph.index().node_index(consumer, consumer_index) || + consumer_index >= selected.size() || selected[consumer_index] || pending_inputs[consumer_index] != 0 || + !is_fusable_followup_op(consumer->op)) { + continue; + } + if (!found || consumer_index < best_index) { + found = true; + best_index = consumer_index; + } + } + if (!found) { + return false; + } + next_node = best_index; + return true; +} + +} // namespace + +GraphTraversalOrder GraphTraversalOrder::build(const Graph & graph) { + GraphTraversalOrder result; + const std::vector & nodes = graph.nodes(); + result.nodes_.reserve(nodes.size()); + if (!graph.has_index()) { + for (const GraphNode & node : nodes) { + result.nodes_.push_back(&node); + } + return result; + } + + std::vector pending_inputs(nodes.size(), 0); + for (size_t i = 0; i < nodes.size(); ++i) { + for (ValueId input : nodes[i].inputs) { + if (graph.index().producer(input) != nullptr) { + ++pending_inputs[i]; + } + } + } + + std::set root_queue; + std::set regular_queue; + std::vector selected(nodes.size(), false); + for (size_t i = 0; i < nodes.size(); ++i) { + if (pending_inputs[i] == 0) { + add_ready_node(nodes[i], i, selected, root_queue, regular_queue); + } + } + + bool has_previous = false; + size_t previous_node = 0; + while (result.nodes_.size() < nodes.size()) { + size_t selected_index = 0; + if (has_previous && select_merge_candidate(graph, previous_node, pending_inputs, selected, selected_index)) { + erase_ready(selected_index, root_queue, regular_queue); + } else if (!root_queue.empty()) { + selected_index = *root_queue.begin(); + root_queue.erase(root_queue.begin()); + } else if (!regular_queue.empty()) { + selected_index = *regular_queue.begin(); + regular_queue.erase(regular_queue.begin()); + } else { + break; + } + + if (selected[selected_index]) { + continue; + } + selected[selected_index] = true; + result.nodes_.push_back(&nodes[selected_index]); + has_previous = true; + previous_node = selected_index; + + for (const GraphNode * consumer : graph.index().consumers(nodes[selected_index].output)) { + size_t consumer_index = 0; + if (consumer == nullptr || !graph.index().node_index(consumer, consumer_index) || + consumer_index >= pending_inputs.size() || selected[consumer_index]) { + continue; + } + --pending_inputs[consumer_index]; + if (pending_inputs[consumer_index] == 0) { + add_ready_node(*consumer, consumer_index, selected, root_queue, regular_queue); + } + } + } + + for (size_t i = 0; i < nodes.size(); ++i) { + if (!selected[i]) { + result.nodes_.push_back(&nodes[i]); + } + } + return result; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/graph/graph-traversal.h b/ggml/src/ggml-hrx/graph/graph-traversal.h new file mode 100644 index 000000000000..c5cb1f27a751 --- /dev/null +++ b/ggml/src/ggml-hrx/graph/graph-traversal.h @@ -0,0 +1,21 @@ +#pragma once + +#include "graph.h" + +#include + +namespace ggml::hrx { + +class GraphTraversalOrder { + public: + GraphTraversalOrder() = default; + + static GraphTraversalOrder build(const Graph & graph); + + const std::vector & nodes() const { return nodes_; } + + private: + std::vector nodes_; +}; + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/graph/graph.cpp b/ggml/src/ggml-hrx/graph/graph.cpp new file mode 100644 index 000000000000..f88f34faba02 --- /dev/null +++ b/ggml/src/ggml-hrx/graph/graph.cpp @@ -0,0 +1,201 @@ +#include "graph.h" + +#include "ggml-impl.h" + +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static const ggml_tensor * tensor_storage_root(const ggml_tensor * tensor) { + while (tensor != nullptr && tensor->view_src != nullptr) { + tensor = tensor->view_src; + } + return tensor; +} + +static int32_t graph_tensor_use_count(const ggml_cgraph & graph, const ggml_tensor * tensor) { + if (tensor == nullptr || graph.use_counts == nullptr || graph.visited_hash_set.keys == nullptr) { + return -1; + } + const size_t hash_pos = ggml_hash_find(&graph.visited_hash_set, tensor); + if (hash_pos == GGML_HASHSET_FULL || !ggml_bitset_get(graph.visited_hash_set.used, hash_pos)) { + return -1; + } + return graph.use_counts[hash_pos]; +} + +static bool tensor_is_external(const ggml_tensor * tensor, + const std::unordered_map & use_counts, + const std::unordered_set & graph_nodes, + const ggml_cgraph & graph) { + const ggml_tensor * root = tensor_storage_root(tensor); + if (root == nullptr || root->op == GGML_OP_NONE || graph_nodes.find(root) == graph_nodes.end()) { + return true; + } + if (tensor->flags & GGML_TENSOR_FLAG_OUTPUT) { + return true; + } + const auto found = use_counts.find(tensor); + const int local_uses = found != use_counts.end() ? found->second : 0; + const int32_t graph_uses = graph_tensor_use_count(graph, tensor); + if (graph_uses > local_uses) { + return true; + } + return local_uses == 0; +} + +} // namespace + +GraphIndex GraphIndex::build(const Graph & graph) { + GraphIndex index; + const std::vector & nodes = graph.nodes(); + for (size_t i = 0; i < nodes.size(); ++i) { + const GraphNode & node = nodes[i]; + index.node_indices_.emplace(&node, i); + index.producers_.emplace(node.output.value, &node); + for (ValueId input : node.inputs) { + index.consumers_[input.value].push_back(&node); + } + } + return index; +} + +const GraphNode * GraphIndex::producer(ValueId value) const { + const auto found = producers_.find(value.value); + return found == producers_.end() ? nullptr : found->second; +} + +const std::vector & GraphIndex::consumers(ValueId value) const { + static const std::vector empty; + const auto found = consumers_.find(value.value); + return found == consumers_.end() ? empty : found->second; +} + +bool GraphIndex::has_single_consumer(ValueId value) const { + return consumers(value).size() == 1; +} + +bool GraphIndex::node_index(const GraphNode * node, size_t & index) const { + const auto found = node_indices_.find(node); + if (found == node_indices_.end()) { + return false; + } + index = found->second; + return true; +} + +Graph::Graph(const Graph & other) : values_(other.values_), nodes_(other.nodes_) { + if (other.has_index()) { + index_ = GraphIndex::build(*this); + } +} + +Graph & Graph::operator=(const Graph & other) { + if (this == &other) { + return *this; + } + values_ = other.values_; + nodes_ = other.nodes_; + index_.reset(); + if (other.has_index()) { + index_ = GraphIndex::build(*this); + } + return *this; +} + +Graph::Graph(Graph && other) : values_(std::move(other.values_)), nodes_(std::move(other.nodes_)) { + if (other.has_index()) { + index_ = GraphIndex::build(*this); + } +} + +Graph & Graph::operator=(Graph && other) { + if (this == &other) { + return *this; + } + values_ = std::move(other.values_); + nodes_ = std::move(other.nodes_); + index_.reset(); + if (other.has_index()) { + index_ = GraphIndex::build(*this); + } + return *this; +} + +GraphNode & Graph::add_node(ggml_op op, ValueId output, std::vector inputs) { + index_.reset(); + GraphNode node; + node.op = op; + node.output = output; + node.inputs = std::move(inputs); + nodes_.push_back(std::move(node)); + return nodes_.back(); +} + +Status Graph::build_index() { + index_ = GraphIndex::build(*this); + return {}; +} + +const GraphIndex & Graph::index() const { + assert(index_.has_value()); + return *index_; +} + +GraphImportResult import_ggml_graph(const ggml_cgraph & graph) { + GraphImportResult result; + std::unordered_map use_counts; + std::unordered_set graph_nodes; + graph_nodes.reserve(static_cast(graph.n_nodes)); + for (int i = 0; i < graph.n_nodes; ++i) { + const ggml_tensor * node = graph.nodes[i]; + if (node == nullptr) { + result.status.log("ggml graph contains a null node"); + return result; + } + graph_nodes.insert(node); + for (const ggml_tensor * source : node->src) { + if (source != nullptr) { + ++use_counts[source]; + } + } + } + + ValueMap & values = result.graph.values(); + for (int i = 0; i < graph.n_nodes; ++i) { + const ggml_tensor * node = graph.nodes[i]; + std::vector inputs; + for (const ggml_tensor * source : node->src) { + if (source == nullptr) { + continue; + } + const ValueKind kind = + tensor_is_external(source, use_counts, graph_nodes, graph) ? ValueKind::External : ValueKind::Transient; + inputs.push_back(values.get_or_add_tensor_value(source, kind)); + } + + const ValueKind output_kind = + tensor_is_external(node, use_counts, graph_nodes, graph) ? ValueKind::External : ValueKind::Transient; + const ValueId output = values.get_or_add_tensor_value(node, output_kind); + GraphNode & graph_node = result.graph.add_node(node->op, output, std::move(inputs)); + graph_node.params = import_op_params(*node); + } + + result.status.append(result.graph.build_index()); + return result; +} + +bool is_layout_alias_op(ggml_op op) { + return op == GGML_OP_VIEW || op == GGML_OP_RESHAPE || op == GGML_OP_PERMUTE || op == GGML_OP_TRANSPOSE; +} + +bool is_layout_alias_node(const Graph & graph, const GraphNode & node) { + return is_layout_alias_op(node.op) && node.inputs.size() == 1 && + graph.values().same_storage(node.output, node.inputs[0]); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/graph/graph.h b/ggml/src/ggml-hrx/graph/graph.h new file mode 100644 index 000000000000..60b26a59daa7 --- /dev/null +++ b/ggml/src/ggml-hrx/graph/graph.h @@ -0,0 +1,85 @@ +#pragma once + +#include "ggml.h" +#include "op-params.h" +#include "status.h" +#include "value-map.h" + +#include +#include +#include +#include +#include + +struct ggml_cgraph; +struct ggml_tensor; + +namespace ggml::hrx { + +struct GraphNode { + ggml_op op; + ValueId output; + std::vector inputs; + OpParams params; +}; + +class Graph; + +class GraphIndex { + public: + GraphIndex() = default; + + static GraphIndex build(const Graph & graph); + + const GraphNode * producer(ValueId value) const; + const std::vector & consumers(ValueId value) const; + bool has_single_consumer(ValueId value) const; + bool node_index(const GraphNode * node, size_t & index) const; + + private: + std::unordered_map producers_; + std::unordered_map> consumers_; + std::unordered_map node_indices_; +}; + +class Graph { + public: + Graph() = default; + Graph(const Graph & other); + Graph & operator=(const Graph & other); + Graph(Graph && other); + Graph & operator=(Graph && other); + + GraphNode & add_node(ggml_op op, ValueId output, std::vector inputs); + + Status build_index(); + + bool has_index() const { return index_.has_value(); } + + const GraphIndex & index() const; + + const std::vector & nodes() const { return nodes_; } + + const ValueMap & values() const { return values_; } + + ValueMap & values() { return values_; } + + private: + ValueMap values_; + std::vector nodes_; + std::optional index_; +}; + +struct GraphImportResult { + Graph graph; + Status status; + + bool valid() const { return status.success(); } +}; + +GraphImportResult import_ggml_graph(const ggml_cgraph & graph); + +bool is_layout_alias_op(ggml_op op); +bool is_layout_alias_node(const Graph & graph, const GraphNode & node); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/graph/op-params.cpp b/ggml/src/ggml-hrx/graph/op-params.cpp new file mode 100644 index 000000000000..efdbeca02e79 --- /dev/null +++ b/ggml/src/ggml-hrx/graph/op-params.cpp @@ -0,0 +1,360 @@ +#include "op-params.h" + +#include "ggml-impl.h" + +#include + +namespace ggml::hrx { +namespace { + +static bool nearly_equal(float lhs, float rhs) { + if (lhs == rhs) { + return true; + } + return std::fabs(lhs - rhs) <= 1.0e-12f; +} + +static bool rms_norm_params_equivalent(const OpParams & lhs, const OpParams & rhs) { + const RmsNormParams * lhs_params = op_params_as(lhs); + const RmsNormParams * rhs_params = op_params_as(rhs); + return lhs_params != nullptr && rhs_params != nullptr && nearly_equal(lhs_params->eps, rhs_params->eps); +} + +static bool flash_attn_ext_params_equivalent(const OpParams & lhs, const OpParams & rhs) { + const FlashAttnExtParams * lhs_params = op_params_as(lhs); + const FlashAttnExtParams * rhs_params = op_params_as(rhs); + return lhs_params != nullptr && rhs_params != nullptr && nearly_equal(lhs_params->scale, rhs_params->scale) && + nearly_equal(lhs_params->max_bias, rhs_params->max_bias) && + nearly_equal(lhs_params->logit_softcap, rhs_params->logit_softcap) && lhs_params->prec == rhs_params->prec; +} + +static bool soft_max_params_equivalent(const OpParams & lhs, const OpParams & rhs) { + const SoftMaxParams * lhs_params = op_params_as(lhs); + const SoftMaxParams * rhs_params = op_params_as(rhs); + return lhs_params != nullptr && rhs_params != nullptr && nearly_equal(lhs_params->scale, rhs_params->scale) && + nearly_equal(lhs_params->max_bias, rhs_params->max_bias); +} + +static bool argsort_params_equivalent(const OpParams & lhs, const OpParams & rhs) { + const ArgsortParams * lhs_params = op_params_as(lhs); + const ArgsortParams * rhs_params = op_params_as(rhs); + return lhs_params != nullptr && rhs_params != nullptr && lhs_params->order == rhs_params->order; +} + +static bool clamp_params_equivalent(const OpParams & lhs, const OpParams & rhs) { + const ClampParams * lhs_params = op_params_as(lhs); + const ClampParams * rhs_params = op_params_as(rhs); + return lhs_params != nullptr && rhs_params != nullptr && nearly_equal(lhs_params->min, rhs_params->min) && + nearly_equal(lhs_params->max, rhs_params->max); +} + +static bool glu_params_equivalent(const OpParams & lhs, const OpParams & rhs) { + const GluParams * lhs_params = op_params_as(lhs); + const GluParams * rhs_params = op_params_as(rhs); + return lhs_params != nullptr && rhs_params != nullptr && lhs_params->op == rhs_params->op && + lhs_params->swapped == rhs_params->swapped && lhs_params->alpha == rhs_params->alpha && + lhs_params->limit == rhs_params->limit; +} + +static bool scale_params_equivalent(const OpParams & lhs, const OpParams & rhs) { + const ScaleParams * lhs_params = op_params_as(lhs); + const ScaleParams * rhs_params = op_params_as(rhs); + return lhs_params != nullptr && rhs_params != nullptr && nearly_equal(lhs_params->scale, rhs_params->scale) && + nearly_equal(lhs_params->bias, rhs_params->bias); +} + +static bool binary_params_equivalent(const OpParams & lhs, const OpParams & rhs) { + const BinaryParams * lhs_params = op_params_as(lhs); + const BinaryParams * rhs_params = op_params_as(rhs); + return lhs_params != nullptr && rhs_params != nullptr && lhs_params->op == rhs_params->op; +} + +static bool unary_params_equivalent(const OpParams & lhs, const OpParams & rhs) { + const UnaryParams * lhs_params = op_params_as(lhs); + const UnaryParams * rhs_params = op_params_as(rhs); + return lhs_params != nullptr && rhs_params != nullptr && lhs_params->op == rhs_params->op; +} + +static bool rope_params_equivalent(const OpParams & lhs, const OpParams & rhs) { + const RopeParams * lhs_params = op_params_as(lhs); + const RopeParams * rhs_params = op_params_as(rhs); + return lhs_params != nullptr && rhs_params != nullptr && lhs_params->n_dims == rhs_params->n_dims && + lhs_params->mode == rhs_params->mode && lhs_params->n_ctx_orig == rhs_params->n_ctx_orig && + nearly_equal(lhs_params->freq_base, rhs_params->freq_base) && + nearly_equal(lhs_params->freq_scale, rhs_params->freq_scale) && + nearly_equal(lhs_params->ext_factor, rhs_params->ext_factor) && + nearly_equal(lhs_params->attn_factor, rhs_params->attn_factor) && + nearly_equal(lhs_params->beta_fast, rhs_params->beta_fast) && + nearly_equal(lhs_params->beta_slow, rhs_params->beta_slow) && lhs_params->sections == rhs_params->sections; +} + +} // namespace + +bool import_binary_kind(const ggml_tensor & tensor, BinaryKind & kind) { + switch (tensor.op) { + case GGML_OP_ADD: + kind = BinaryKind::Add; + return true; + case GGML_OP_SUB: + kind = BinaryKind::Sub; + return true; + case GGML_OP_MUL: + kind = BinaryKind::Mul; + return true; + case GGML_OP_DIV: + kind = BinaryKind::Div; + return true; + case GGML_OP_GLU: + if (tensor.src[1] == nullptr) { + return false; + } + switch (ggml_get_glu_op(&tensor)) { + case GGML_GLU_OP_REGLU: + kind = BinaryKind::RegLU; + return true; + case GGML_GLU_OP_SWIGLU: + kind = BinaryKind::SwiGLU; + return true; + case GGML_GLU_OP_GEGLU: + kind = BinaryKind::GeGLU; + return true; + case GGML_GLU_OP_GEGLU_ERF: + kind = BinaryKind::GeGLUErf; + return true; + case GGML_GLU_OP_GEGLU_QUICK: + kind = BinaryKind::GeGLUQuick; + return true; + default: + return false; + } + default: + return false; + } +} + +bool binary_kind_supported(BinaryKind kind) { + return static_cast(kind) <= static_cast(BinaryKind::GeGLUQuick); +} + +uint32_t binary_kind_config_value(BinaryKind kind) { + return static_cast(kind); +} + +static bool unary_kind_from_ggml_unary_op(ggml_unary_op op, UnaryKind & kind) { + switch (op) { + case GGML_UNARY_OP_ABS: + kind = UnaryKind::Abs; + return true; + case GGML_UNARY_OP_SGN: + kind = UnaryKind::Sgn; + return true; + case GGML_UNARY_OP_NEG: + kind = UnaryKind::Neg; + return true; + case GGML_UNARY_OP_STEP: + kind = UnaryKind::Step; + return true; + case GGML_UNARY_OP_TANH: + kind = UnaryKind::Tanh; + return true; + case GGML_UNARY_OP_ELU: + kind = UnaryKind::Elu; + return true; + case GGML_UNARY_OP_RELU: + kind = UnaryKind::Relu; + return true; + case GGML_UNARY_OP_SIGMOID: + kind = UnaryKind::Sigmoid; + return true; + case GGML_UNARY_OP_GELU: + kind = UnaryKind::Gelu; + return true; + case GGML_UNARY_OP_GELU_QUICK: + kind = UnaryKind::GeluQuick; + return true; + case GGML_UNARY_OP_SILU: + kind = UnaryKind::Silu; + return true; + case GGML_UNARY_OP_HARDSWISH: + kind = UnaryKind::HardSwish; + return true; + case GGML_UNARY_OP_HARDSIGMOID: + kind = UnaryKind::HardSigmoid; + return true; + case GGML_UNARY_OP_EXP: + kind = UnaryKind::Exp; + return true; + case GGML_UNARY_OP_EXPM1: + kind = UnaryKind::Expm1; + return true; + case GGML_UNARY_OP_SOFTPLUS: + kind = UnaryKind::SoftPlus; + return true; + case GGML_UNARY_OP_GELU_ERF: + kind = UnaryKind::GeluErf; + return true; + case GGML_UNARY_OP_XIELU: + kind = UnaryKind::Xielu; + return true; + case GGML_UNARY_OP_FLOOR: + kind = UnaryKind::Floor; + return true; + case GGML_UNARY_OP_CEIL: + kind = UnaryKind::Ceil; + return true; + case GGML_UNARY_OP_ROUND: + kind = UnaryKind::Round; + return true; + case GGML_UNARY_OP_TRUNC: + kind = UnaryKind::Trunc; + return true; + default: + return false; + } +} + +bool import_unary_kind(const ggml_tensor & tensor, UnaryKind & kind) { + switch (tensor.op) { + case GGML_OP_UNARY: + return unary_kind_from_ggml_unary_op(ggml_get_unary_op(&tensor), kind); + case GGML_OP_SQR: + kind = UnaryKind::Sqr; + return true; + case GGML_OP_SQRT: + kind = UnaryKind::Sqrt; + return true; + case GGML_OP_LOG: + kind = UnaryKind::Log; + return true; + case GGML_OP_SIN: + kind = UnaryKind::Sin; + return true; + case GGML_OP_COS: + kind = UnaryKind::Cos; + return true; + default: + return false; + } +} + +bool unary_kind_supported(UnaryKind kind) { + return static_cast(kind) <= static_cast(UnaryKind::Identity); +} + +uint32_t unary_kind_config_value(UnaryKind kind) { + return static_cast(kind); +} + +OpParams import_op_params(const ggml_tensor & tensor) { + BinaryKind binary_kind; + if (import_binary_kind(tensor, binary_kind)) { + return BinaryParams{ binary_kind }; + } + + UnaryKind unary_kind; + if (import_unary_kind(tensor, unary_kind)) { + return UnaryParams{ unary_kind }; + } + + switch (tensor.op) { + case GGML_OP_RMS_NORM: + case GGML_OP_L2_NORM: + case GGML_OP_NORM: + return RmsNormParams{ ggml_get_op_params_f32(&tensor, 0) }; + case GGML_OP_SOFT_MAX: + return SoftMaxParams{ + ggml_get_op_params_f32(&tensor, 0), + ggml_get_op_params_f32(&tensor, 1), + }; + case GGML_OP_FLASH_ATTN_EXT: + return FlashAttnExtParams{ + ggml_get_op_params_f32(&tensor, 0), + ggml_get_op_params_f32(&tensor, 1), + ggml_get_op_params_f32(&tensor, 2), + ggml_flash_attn_ext_get_prec(&tensor), + }; + case GGML_OP_ARGSORT: + return ArgsortParams{ static_cast(ggml_get_op_params_i32(&tensor, 0)) }; + case GGML_OP_CLAMP: + return ClampParams{ + ggml_get_op_params_f32(&tensor, 0), + ggml_get_op_params_f32(&tensor, 1), + }; + case GGML_OP_GLU: + return GluParams{ ggml_get_glu_op(&tensor), ggml_get_op_params_i32(&tensor, 1) != 0, + ggml_get_op_params_f32(&tensor, 2), ggml_get_op_params_f32(&tensor, 3) }; + case GGML_OP_MUL_MAT: + if (ggml_get_op_params_i32(&tensor, 1) != 0) { + return MulMatParams{ ggml_get_op_params_i32(&tensor, 1) }; + } + return {}; + case GGML_OP_SCALE: + return ScaleParams{ + ggml_get_op_params_f32(&tensor, 0), + ggml_get_op_params_f32(&tensor, 1), + }; + case GGML_OP_ROPE: + return RopeParams{ + ggml_get_op_params_i32(&tensor, 1), + ggml_get_op_params_i32(&tensor, 2), + ggml_get_op_params_i32(&tensor, 4), + ggml_get_op_params_f32(&tensor, 5), + ggml_get_op_params_f32(&tensor, 6), + ggml_get_op_params_f32(&tensor, 7), + ggml_get_op_params_f32(&tensor, 8), + ggml_get_op_params_f32(&tensor, 9), + ggml_get_op_params_f32(&tensor, 10), + { + ggml_get_op_params_i32(&tensor, 11), + ggml_get_op_params_i32(&tensor, 12), + ggml_get_op_params_i32(&tensor, 13), + ggml_get_op_params_i32(&tensor, 14), + }, + }; + default: + return std::monostate{}; + } +} + +bool op_params_equivalent(ggml_op op, const OpParams & lhs, const OpParams & rhs) { + switch (op) { + case GGML_OP_RMS_NORM: + case GGML_OP_L2_NORM: + case GGML_OP_NORM: + return rms_norm_params_equivalent(lhs, rhs); + case GGML_OP_SOFT_MAX: + return soft_max_params_equivalent(lhs, rhs); + case GGML_OP_FLASH_ATTN_EXT: + return flash_attn_ext_params_equivalent(lhs, rhs); + case GGML_OP_ARGSORT: + return argsort_params_equivalent(lhs, rhs); + case GGML_OP_CLAMP: + return clamp_params_equivalent(lhs, rhs); + case GGML_OP_GLU: + return binary_params_equivalent(lhs, rhs) || glu_params_equivalent(lhs, rhs); + case GGML_OP_SCALE: + return scale_params_equivalent(lhs, rhs); + case GGML_OP_ADD: + case GGML_OP_SUB: + case GGML_OP_MUL: + case GGML_OP_DIV: + return binary_params_equivalent(lhs, rhs); + case GGML_OP_UNARY: + case GGML_OP_SQR: + case GGML_OP_SQRT: + case GGML_OP_LOG: + case GGML_OP_SIN: + case GGML_OP_COS: + return unary_params_equivalent(lhs, rhs); + case GGML_OP_ROPE: + return rope_params_equivalent(lhs, rhs); + default: + return lhs.index() == rhs.index(); + } +} + +bool op_params_equivalent(ggml_op op, const OpParams & lhs, const ggml_tensor & rhs) { + return op_params_equivalent(op, lhs, import_op_params(rhs)); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/graph/op-params.h b/ggml/src/ggml-hrx/graph/op-params.h new file mode 100644 index 000000000000..c3e8bc820fab --- /dev/null +++ b/ggml/src/ggml-hrx/graph/op-params.h @@ -0,0 +1,150 @@ +#pragma once + +#include "ggml.h" + +#include +#include +#include + +struct ggml_tensor; + +namespace ggml::hrx { + +struct RmsNormParams { + float eps = 0.0f; +}; + +struct FlashAttnExtParams { + float scale = 0.0f; + float max_bias = 0.0f; + float logit_softcap = 0.0f; + ggml_prec prec = GGML_PREC_DEFAULT; +}; + +struct SoftMaxParams { + float scale = 0.0f; + float max_bias = 0.0f; +}; + +struct ArgsortParams { + ggml_sort_order order = GGML_SORT_ORDER_ASC; +}; + +struct ClampParams { + float min = 0.0f; + float max = 0.0f; +}; + +struct GluParams { + ggml_glu_op op = GGML_GLU_OP_REGLU; + bool swapped = false; + float alpha = 0.0f; // GGML_GLU_OP_SWIGLU_OAI + float limit = 0.0f; // GGML_GLU_OP_SWIGLU_OAI +}; + +struct ScaleParams { + float scale = 0.0f; + float bias = 0.0f; +}; + +enum class BinaryKind : uint32_t { + Add = 0, + Sub = 1, + Mul = 2, + Div = 3, + SwiGLU = 4, + GeGLU = 5, + RegLU = 6, + GeGLUErf = 7, + GeGLUQuick = 8, +}; + +struct BinaryParams { + BinaryKind op = BinaryKind::Add; +}; + +enum class UnaryKind : uint32_t { + Neg = 0, + Abs = 1, + Relu = 2, + Step = 3, + Sqr = 4, + Sgn = 5, + Floor = 6, + Ceil = 7, + Round = 8, + Trunc = 9, + Tanh = 10, + Elu = 11, + Sigmoid = 12, + Gelu = 13, + GeluQuick = 14, + Silu = 15, + HardSwish = 16, + HardSigmoid = 17, + Exp = 18, + Expm1 = 19, + GeluErf = 20, + Sqrt = 21, + Log = 22, + Identity = 23, + + // Parameterized or remaining exact-math unary routes need separate lowering work. + SoftPlus = 100, + Xielu, + Sin, + Cos, +}; + +struct UnaryParams { + UnaryKind op = UnaryKind::Abs; +}; + +struct MulMatParams { + int32_t hint = 0; // ggml_mul_mat_set_hint: GGML_HINT_SRC0_IS_HADAMARD +}; + +struct RopeParams { + int n_dims = 0; + int mode = 0; + int n_ctx_orig = 0; + float freq_base = 0.0f; + float freq_scale = 0.0f; + float ext_factor = 0.0f; + float attn_factor = 0.0f; + float beta_fast = 0.0f; + float beta_slow = 0.0f; + std::array sections = {}; +}; + +// clang-format off +using OpParams = std::variant< + std::monostate, + RmsNormParams, + FlashAttnExtParams, + SoftMaxParams, + ArgsortParams, + ClampParams, + GluParams, + ScaleParams, + BinaryParams, + UnaryParams, + RopeParams, + MulMatParams>; +// clang-format on + +template const T * op_params_as(const OpParams & params) { + return std::get_if(¶ms); +} + +OpParams import_op_params(const ggml_tensor & tensor); +bool op_params_equivalent(ggml_op op, const OpParams & lhs, const OpParams & rhs); +bool op_params_equivalent(ggml_op op, const OpParams & lhs, const ggml_tensor & rhs); +bool import_binary_kind(const ggml_tensor & tensor, BinaryKind & kind); +bool binary_kind_supported(BinaryKind kind); +uint32_t binary_kind_config_value(BinaryKind kind); +bool import_unary_kind(const ggml_tensor & tensor, UnaryKind & kind); +bool unary_kind_supported(UnaryKind kind); +uint32_t unary_kind_config_value(UnaryKind kind); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/graph/value-map.cpp b/ggml/src/ggml-hrx/graph/value-map.cpp new file mode 100644 index 000000000000..8f37c6d3abc9 --- /dev/null +++ b/ggml/src/ggml-hrx/graph/value-map.cpp @@ -0,0 +1,313 @@ +#include "value-map.h" + +#include "ggml-impl.h" + +#include + +namespace ggml::hrx { +namespace { + +static const ggml_tensor * tensor_storage_root(const ggml_tensor * tensor) { + while (tensor != nullptr && tensor->view_src != nullptr) { + tensor = tensor->view_src; + } + return tensor; +} + +static size_t tensor_storage_offset(const ggml_tensor * tensor) { + return tensor != nullptr && tensor->view_src != nullptr ? tensor->view_offs : 0; +} + +static bool tensor_storage_relative_offset(const ggml_tensor * source, const ggml_tensor * tensor, size_t & offset) { + if (source == nullptr || tensor == nullptr || tensor_storage_root(source) != tensor_storage_root(tensor)) { + return false; + } + const size_t source_offset = tensor_storage_offset(source); + const size_t tensor_offset = tensor_storage_offset(tensor); + if (tensor_offset < source_offset) { + return false; + } + const size_t relative_offset = tensor_offset - source_offset; + if (ggml_nbytes(source) == 0) { + // A zero-byte (empty) source can only contain a zero-byte tensor; the byte + // offset of a degenerate view into an empty tensor is meaningless (#95). + offset = relative_offset; + return ggml_nbytes(tensor) == 0; + } + if (relative_offset > ggml_nbytes(source) || ggml_nbytes(tensor) > ggml_nbytes(source) - relative_offset) { + return false; + } + offset = relative_offset; + return true; +} + +} // namespace + +const Value * ValueMap::find_alias_source(const ggml_tensor * tensor) const { + if (tensor == nullptr || tensor->view_src == nullptr) { + return nullptr; + } + const Value * exact_source = find_tensor(tensor->view_src); + if (exact_source != nullptr) { + return exact_source; + } + const Value * best_source = nullptr; + size_t best_source_size = 0; + for (const Value & value : values_) { + size_t relative_offset = 0; + if (value.tensor == nullptr || !tensor_storage_relative_offset(value.tensor, tensor, relative_offset)) { + continue; + } + if (best_source == nullptr || ggml_nbytes(value.tensor) < best_source_size) { + best_source = &value; + best_source_size = ggml_nbytes(value.tensor); + } + } + return best_source; +} + +void ValueMap::promote_storage_root_external(ValueId storage_root) { + if (storage_root.value < 0 || static_cast(storage_root.value) >= values_.size()) { + return; + } + values_[static_cast(storage_root.value)].kind = ValueKind::External; +} + +ValueId ValueMap::get_or_add_tensor_value(const ggml_tensor * tensor, ValueKind kind) { + const auto found = tensor_values_.find(tensor); + if (found != tensor_values_.end()) { + Value & value = values_[found->second]; + if (kind == ValueKind::External) { + value.kind = ValueKind::External; + promote_storage_root_external(value.storage_root); + } + return value.id; + } + + const ValueId id(static_cast(values_.size())); + const Value * alias_source = find_alias_source(tensor); + ValueStorageId storage; + ValueId storage_root; + ValueId alias_source_id; + size_t storage_offset = 0; + size_t storage_byte_count = ggml_nbytes(tensor); + if (alias_source != nullptr) { + size_t relative_offset = 0; + if (!tensor_storage_relative_offset(alias_source->tensor, tensor, relative_offset)) { + relative_offset = tensor->view_offs; + } + storage = alias_source->storage; + storage_root = alias_source->storage_root; + alias_source_id = alias_source->id; + storage_offset = alias_source->storage_offset + relative_offset; + storage_byte_count = alias_source->storage_byte_count; + const Value * root = find(storage_root); + if (root != nullptr && root->kind == ValueKind::External) { + kind = ValueKind::External; + } + if (kind == ValueKind::External) { + promote_storage_root_external(storage_root); + } + } else { + storage = ValueStorageId(static_cast(storages_.size())); + storage_root = id; + storages_.push_back({ storage, storage_root, storage_byte_count }); + } + + Value value = { + id, + kind, + storage, + storage_root, + alias_source_id, + storage_offset, + storage_byte_count, + tensor->type, + {}, + {}, + ggml_nelements(tensor), + ggml_nbytes(tensor), + ggml_is_contiguous(tensor), + tensor, + std::nullopt, + }; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + value.ne[i] = tensor->ne[i]; + value.nb[i] = tensor->nb[i]; + } + + values_.push_back(std::move(value)); + tensor_values_.emplace(tensor, values_.size() - 1); + return values_.back().id; +} + +const Value * ValueMap::find(ValueId id) const { + if (id.value < 0 || static_cast(id.value) >= values_.size()) { + return nullptr; + } + return &values_[static_cast(id.value)]; +} + +const ValueStorage * ValueMap::find_storage(ValueStorageId id) const { + if (id.value < 0 || static_cast(id.value) >= storages_.size()) { + return nullptr; + } + return &storages_[static_cast(id.value)]; +} + +bool ValueMap::bind_buffer(ValueId id, ValueBufferBinding binding) { + if (id.value < 0 || static_cast(id.value) >= values_.size()) { + return false; + } + Value & value = values_[static_cast(id.value)]; + if (value.kind != ValueKind::External) { + return false; + } + value.buffer = std::move(binding); + return true; +} + +std::optional ValueMap::resolve_buffer_binding(ValueId id) const { + const Value * value = find(id); + if (value == nullptr) { + return std::nullopt; + } + if (value->buffer.has_value()) { + return value->buffer; + } + if (value->storage_root == value->id) { + return std::nullopt; + } + const Value * root = find(value->storage_root); + if (root == nullptr || !root->buffer.has_value()) { + return std::nullopt; + } + ValueBufferBinding binding = *root->buffer; + if (value->storage_offset > binding.length) { + return std::nullopt; + } + if (value->byte_count > binding.length - value->storage_offset) { + return std::nullopt; + } + binding.offset += value->storage_offset; + binding.length = value->byte_count; + return binding; +} + +std::vector ValueMap::external_value_ids() const { + std::vector ids; + for (const Value & value : values_) { + if (value.kind == ValueKind::External) { + ids.push_back(value.id); + } + } + return ids; +} + +Status ValueMap::alias_storage(ValueId target, ValueId source) { + Status status; + if (target.value < 0 || static_cast(target.value) >= values_.size()) { + status.log("value alias target %d does not exist", target.value); + return status; + } + if (source.value < 0 || static_cast(source.value) >= values_.size()) { + status.log("value alias source %d does not exist", source.value); + return status; + } + if (target == source) { + status.log("value alias target %d aliases itself", target.value); + return status; + } + + Value & target_value = values_[static_cast(target.value)]; + const Value & source_value = values_[static_cast(source.value)]; + if (target_value.storage == source_value.storage) { + return status; + } + if (target_value.kind != ValueKind::Transient) { + status.log("value alias target %d is not transient", target.value); + return status; + } + if (target_value.type != source_value.type || target_value.byte_count != source_value.byte_count || + target_value.element_count != source_value.element_count || + target_value.contiguous != source_value.contiguous) { + status.log("value alias target %d is incompatible with source %d", target.value, source.value); + return status; + } + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (target_value.ne[i] != source_value.ne[i] || target_value.nb[i] != source_value.nb[i]) { + status.log("value alias target %d has a different layout than source %d", target.value, source.value); + return status; + } + } + if (target_value.alias_source.value >= 0 && target_value.alias_source != source) { + status.log("value alias target %d already aliases source %d", target.value, target_value.alias_source.value); + return status; + } + + target_value.storage = source_value.storage; + target_value.storage_root = source_value.storage_root; + target_value.alias_source = source; + target_value.storage_offset = source_value.storage_offset; + target_value.storage_byte_count = source_value.storage_byte_count; + return status; +} + +ValueId ValueMap::storage_root(ValueId id) const { + const Value * value = find(id); + return value == nullptr ? ValueId() : value->storage_root; +} + +bool ValueMap::same_storage(ValueId lhs, ValueId rhs) const { + const Value * lhs_value = find(lhs); + const Value * rhs_value = find(rhs); + return lhs_value != nullptr && rhs_value != nullptr && lhs_value->storage == rhs_value->storage; +} + +Status ValueMap::add_snapshot_storage(ValueStorage storage) { + Status status; + if (storage.id.value < 0 || static_cast(storage.id.value) != storages_.size()) { + status.log("snapshot storage id %d is not the next storage id %zu", storage.id.value, storages_.size()); + return status; + } + if (storage.root.value < 0) { + status.log("snapshot storage %d has invalid root value %d", storage.id.value, storage.root.value); + return status; + } + storages_.push_back(storage); + return status; +} + +Status ValueMap::add_snapshot_value(Value value) { + Status status; + if (value.id.value < 0 || static_cast(value.id.value) != values_.size()) { + status.log("snapshot value id %d is not the next value id %zu", value.id.value, values_.size()); + return status; + } + if (value.storage.value < 0 || static_cast(value.storage.value) >= storages_.size()) { + status.log("snapshot value %d references missing storage %d", value.id.value, value.storage.value); + return status; + } + if (value.storage_root.value < 0 || static_cast(value.storage_root.value) > values_.size()) { + status.log("snapshot value %d references invalid storage root %d", value.id.value, value.storage_root.value); + return status; + } + if (value.alias_source.value >= 0 && static_cast(value.alias_source.value) >= values_.size()) { + status.log("snapshot value %d references missing alias source %d", value.id.value, value.alias_source.value); + return status; + } + value.tensor = nullptr; + value.buffer.reset(); + values_.push_back(std::move(value)); + return status; +} + +const Value * ValueMap::find_tensor(const ggml_tensor * tensor) const { + const auto found = tensor_values_.find(tensor); + if (found == tensor_values_.end()) { + return nullptr; + } + return &values_[found->second]; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/graph/value-map.h b/ggml/src/ggml-hrx/graph/value-map.h new file mode 100644 index 000000000000..b026de65fde6 --- /dev/null +++ b/ggml/src/ggml-hrx/graph/value-map.h @@ -0,0 +1,128 @@ +#pragma once + +#include "ggml.h" +#include "status.h" + +#include +#include +#include +#include +#include +#include + +struct ggml_tensor; +typedef struct hrx_buffer_s * hrx_buffer_t; + +namespace ggml::hrx { + +struct ValueId { + ValueId() : value(-1) {} + + explicit ValueId(int32_t value) : value(value) {} + + int32_t value; +}; + +struct ValueStorageId { + ValueStorageId() : value(-1) {} + + explicit ValueStorageId(int32_t value) : value(value) {} + + int32_t value; +}; + +inline bool operator==(ValueId lhs, ValueId rhs) { + return lhs.value == rhs.value; +} + +inline bool operator!=(ValueId lhs, ValueId rhs) { + return !(lhs == rhs); +} + +inline bool operator==(ValueStorageId lhs, ValueStorageId rhs) { + return lhs.value == rhs.value; +} + +inline bool operator!=(ValueStorageId lhs, ValueStorageId rhs) { + return !(lhs == rhs); +} + +enum class ValueKind : uint8_t { + External, + Transient, +}; + +struct ValueBufferBinding { + // A buffer is directly bindable by an HRX command program. Host data requires residency or staging before + // execution. These are alternate storage forms and should not both be populated. + hrx_buffer_t buffer = nullptr; + size_t offset = 0; + size_t length = 0; + uint64_t identity = 0; + uint64_t generation = 0; + size_t capacity = 0; + void * host_data = nullptr; + bool weight = false; + + bool requires_materialization() const { return host_data != nullptr; } +}; + +struct Value { + ValueId id; + ValueKind kind; + ValueStorageId storage; + ValueId storage_root; + ValueId alias_source; + size_t storage_offset = 0; + size_t storage_byte_count = 0; + ggml_type type; + std::array ne; + std::array nb; + int64_t element_count = 0; + size_t byte_count = 0; + bool contiguous = false; + const ggml_tensor * tensor = nullptr; + std::optional buffer; +}; + +struct ValueStorage { + ValueStorageId id; + ValueId root; + size_t byte_count = 0; +}; + +class ValueMap { + public: + ValueMap() = default; + + ValueId get_or_add_tensor_value(const ggml_tensor * tensor, ValueKind kind); + + const Value * find(ValueId id) const; + const Value * find_tensor(const ggml_tensor * tensor) const; + const ValueStorage * find_storage(ValueStorageId id) const; + bool bind_buffer(ValueId id, ValueBufferBinding binding); + std::optional resolve_buffer_binding(ValueId id) const; + std::vector external_value_ids() const; + Status alias_storage(ValueId target, ValueId source); + ValueId storage_root(ValueId id) const; + bool same_storage(ValueId lhs, ValueId rhs) const; + + const std::vector & values() const { return values_; } + + const std::vector & storages() const { return storages_; } + + size_t size() const { return values_.size(); } + + Status add_snapshot_storage(ValueStorage storage); + Status add_snapshot_value(Value value); + + private: + const Value * find_alias_source(const ggml_tensor * tensor) const; + void promote_storage_root_external(ValueId storage_root); + + std::vector values_; + std::vector storages_; + std::unordered_map tensor_values_; +}; + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/hip/README.md b/ggml/src/ggml-hrx/hip/README.md new file mode 100644 index 000000000000..9f458314469f --- /dev/null +++ b/ggml/src/ggml-hrx/hip/README.md @@ -0,0 +1,108 @@ + + +# HIP kernels on the HRX backend + +HRX can run code objects built by hipcc/amdclang++. The HRX runtime loads them +(`hrx_executable_load_data`, the same call the Loom JIT result goes through) and dispatches +them. The HIP runtime and ggml-hip are not involved. This directory lets a ggml-hrx matcher +dispatch a HIP kernel the same way it dispatches a Loom kernel. + +## Layout + +| file | role | +|---|---| +| `kernels/*.hip` | kernel sources; every file here is built for each target in `GGML_HRX_HIP_TARGETS` | +| `ggml-hrx-hip.cmake` | compiles each `.hip` to a raw code object (`--cuda-device-only --no-gpu-bundle-output`), embeds the code objects, adds the sources below | +| `embed_hip_code_objects.py` | writes `ggml-hrx-hip-code-objects.inc` (bytes + digest per stem/target) | +| `hip-code-objects.{h,cpp}` | lookup of embedded code objects by (stem, target) | +| `hip-kernel-registry.{h,cpp}` | `HipKernel` description -> `KernelDefinition` (family `hip`), found by `resolve_kernel_definition` | +| `hip-kernel-loader.{h,cpp}` | gives the kernel executable cache a finished "compile" (embedded code object + launch geometry) in place of a Loom JIT | +| `hip-dispatches.{h,cpp}` | **the one place HIP matchers are registered** | +| `dispatch-hip-scale.cpp`, `kernels/hip_scale_f32.hip` | worked example (GGML_OP_SCALE, opt-in `GGML_HRX_HIP_EXAMPLE_SCALE=1`) | +| `hip-smoke.cpp`, `kernels/hip_smoke.hip` | `ggml-hrx-hip-smoke`: libhrx-only load/dispatch/check of a code object | +| `GGML_HRX_HIP_ADDON_DIR` | optional out-of-tree kernel add-on, built in the same way (see "Kernel add-ons") | + +Hooks in upstream ggml-hrx files are one line each: the `include()` in `CMakeLists.txt`, the +HIP lookup in `resolve_kernel_definition` (kernel-corpus.cpp), the HIP branch in +`KernelExecutableCache::get_or_compile`, the `friend struct HipCodeObjectLoader` in +loom-kernel-jit.h and `register_hip_dispatches` in dispatch-common.cpp. + +## Adding a kernel + +1. **Kernel**: `kernels/.hip`, with the Apache header. The arguments are all pointers first + (each one is an HRX binding, in order), then `uint32_t` by-value arguments (the launch + parameters, in order). Pass floats as `uint32_t` bit patterns (`__builtin_bit_cast`). Use + `extern "C" __global__` with `__launch_bounds__`. `blockDim`/`gridDim` work, because HRX fills + the hidden kernargs. gfx11 is wave32: use `__shfl_xor(v, o, 32)` and + `__builtin_amdgcn_wmma_*_w32`. +2. **Matcher**: `dispatch-hip-.cpp` (our header). In its `register_hip__dispatch`: + - `register_hip_kernel({ "", "", bindings, launch_parameters, workload_parameters, launch_fn })`; + - `registry.add({ "hip.", GGML_OP_..., kind, priority, DispatchSource::Common, matcher })`. + Name registrations `hip.*` so that `GGML_HRX_DISABLE_DISPATCH=hip.` turns every HIP matcher off. + The matcher builds `Dispatch` as for Loom: `make_kernel_specialization(hip_kernel_ref(""))`, + `integer_parameters` for every launch and workload parameter, and bindings in pointer order. + `launch_fn(params, geometry)` sets the workgroup count and size. The executable cache key + holds only the workload parameters, so list the parameters that change the grid there and + nothing else. +3. **Register**: add one line to `register_hip_dispatches` in `hip-dispatches.cpp` and the + matcher source to `target_sources(ggml-hrx ...)` in `ggml-hrx-hip.cmake`. (In an add-on, + the line goes in its `ggml_hrx_hip_addon_register` and the source is picked up by the glob.) +4. **Check** (loom worker rules): `test-backend-ops -o -b HRX0`, run several times, with + `GGML_HRX_LOG_DISPATCH=1` to confirm `hip.` matched. Add a known-answer probe, a model + KLD or teacher-forced comparison, and repeated identical requests. Keep a HIP kernel only + if an interleaved A/B shows it at least as fast as the path it replaces. + +**Trap: fused registrations always come first.** For a root op the registry tries every +`DispatchMatchKind::Fused` registration before any `SingleOp` one; `priority` only orders +registrations within the same kind. A `SingleOp` HIP matcher with a high priority still loses +to any fused Loom matcher that takes the node, and a fused HIP matcher wins over every single-op +matcher whatever its priority. Check `GGML_HRX_LOG_DISPATCH=1` rather than reasoning from priorities. + +At load, a missing (stem, target) code object fails with `no code object ''`. +A kernel ABI that differs from the registration (binding count or constant bytes) fails with +`compiled ABI does not match manifest`. + +## Build + +`GGML_HRX_HIP_COMPILER` (default: `amdclang++` from `CMAKE_CXX_COMPILER` or /opt/rocm-therock/bin), +`GGML_HRX_HIP_TARGETS` (default `gfx1151`, semicolon list), `GGML_HRX_HIP_FLAGS` (default `-O3`). +Look at the ISA with +`llvm-objdump -d --mcpu=gfx1151 build/ggml/src/ggml-hrx/hip-code-objects/.gfx1151.hsaco`. + +## Kernel add-ons + +Kernels and matchers can live outside this tree. Configure with +`-DGGML_HRX_HIP_ADDON_DIR=` (absolute, or relative to the top source directory): + +| add-on path | what the build does with it | +|---|---| +| `/kernels/*.hip` | compiled and embedded with `kernels/*.hip` (same targets, flags and lookup by stem; `-I` the kernel's own directory) | +| `/*.cpp`, `/*.h` | compiled into ggml-hrx, with `GGML_HRX_HIP_ADDON` defined | +| `/addon.cmake` | included at the end of `ggml-hrx-hip.cmake` if present (extra tools; relative source paths resolve against `ggml/src/ggml-hrx`) | + +The add-on defines `ggml::hrx::ggml_hrx_hip_addon_register(DispatchRegistryBuilder &)` +(declared in `hip-dispatches.h`); `register_hip_dispatches` calls it after the matchers in this +directory, once per registry build. Add-on sources include `hip/hip-dispatches.h` and +`hip/hip-kernel-registry.h` like the example does. Kernel stems must not repeat a stem from +`kernels/` (configure fails). An add-on can also take the attention-sink rescale that follows +FlashAttention through `set_hip_attention_sink_hook`. A matcher for an op that ggml-hrx's +`supports_op` does not list by itself declares the op from its `register_*` function with +`hip_declare_eager_op(op)` (`hip/hip-capabilities.h`); `eager_capability_declared` consults that +set for every op outside its built-in list, so the node reaches the dispatcher instead of another +backend. Without the option nothing changes: only the kernels and matchers in this directory are +built, and the declared-op set stays empty. diff --git a/ggml/src/ggml-hrx/hip/dispatch-hip-scale.cpp b/ggml/src/ggml-hrx/hip/dispatch-hip-scale.cpp new file mode 100644 index 000000000000..8a12124e9a82 --- /dev/null +++ b/ggml/src/ggml-hrx/hip/dispatch-hip-scale.cpp @@ -0,0 +1,108 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Example matcher for the HRX HIP recipe: GGML_OP_SCALE (packed f32) -> hip_scale_f32 +// (hip/kernels/hip_scale_f32.hip). Opt-in with GGML_HRX_HIP_EXAMPLE_SCALE=1 so it never replaces +// the Loom scale kernel in normal runs; it exists to prove and document the matcher path. + +#include "hip/hip-dispatches.h" +#include "hip/hip-kernel-registry.h" + +#include "ggml.h" + +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +constexpr KernelCatalogRef kHipScaleKernel = hip_kernel_ref("hip_scale_f32"); + +bool packed_f32(const Value & value) { + size_t stride = sizeof(float); + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (value.ne[i] <= 0 || value.nb[i] != stride) { + return false; + } + stride *= static_cast(value.ne[i]); + } + return value.type == GGML_TYPE_F32 && value.contiguous; +} + +bool launch_hip_scale(const std::map & parameters, HipLaunchGeometry & geometry) { + const auto n = parameters.find("element_count"); + if (n == parameters.end() || n->second <= 0) { + return false; + } + const int64_t groups = (n->second + 255) / 256; + geometry.workgroup_count = { static_cast(groups < 4096 ? groups : 4096), 1, 1 }; + geometry.workgroup_size = { 256, 1, 1 }; + return true; +} + +uint32_t float_bits(float value) { + uint32_t bits = 0; + std::memcpy(&bits, &value, sizeof(bits)); + return bits; +} + +bool match_hip_scale(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_SCALE || node->inputs.size() != 1) { + return false; + } + const ScaleParams * params = op_params_as(node->params); + const Value * output = context.graph.values().find(node->output); + const Value * input = context.graph.values().find(node->inputs[0]); + if (params == nullptr || output == nullptr || input == nullptr || !packed_f32(*output) || !packed_f32(*input) || + input->element_count != output->element_count || output->alias_source.value >= 0 || + input->alias_source.value >= 0 || input->storage == output->storage || + static_cast(output->element_count) > std::numeric_limits::max()) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kHipScaleKernel); + dispatch.kernel.integer_parameters.emplace("element_count", output->element_count); + dispatch.kernel.integer_parameters.emplace("scale_bits", float_bits(params->scale)); + dispatch.kernel.integer_parameters.emplace("bias_bits", float_bits(params->bias)); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + match.covered_nodes.push_back(context.root_index); + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_hip_scale_dispatch(DispatchRegistryBuilder & registry) { + const char * enabled = std::getenv("GGML_HRX_HIP_EXAMPLE_SCALE"); + if (enabled == nullptr || std::strcmp(enabled, "1") != 0 || !hip_code_object_embedded("hip_scale_f32")) { + return; + } + register_hip_kernel({ + "hip_scale_f32", + "hip_scale_f32", + { { "input", ResourceAccess::Read }, { "output", ResourceAccess::Write } }, + { "element_count", "scale_bits", "bias_bits" }, + { "element_count" }, + launch_hip_scale, + }); + registry.add({ "hip.scale_f32", GGML_OP_SCALE, DispatchMatchKind::SingleOp, 10, DispatchSource::Common, + match_hip_scale }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/hip/embed_hip_code_objects.py b/ggml/src/ggml-hrx/hip/embed_hip_code_objects.py new file mode 100644 index 000000000000..6da0652cb3ad --- /dev/null +++ b/ggml/src/ggml-hrx/hip/embed_hip_code_objects.py @@ -0,0 +1,57 @@ +#!/usr/bin/env python3 +# Copyright 2026 bong-water-water-bong +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Embeds HIP code objects (raw AMDGPU ELF, one per .hip stem and target) as a C++ table. + +usage: embed_hip_code_objects.py OUTPUT.inc STEM:TARGET:PATH ... +""" + +import hashlib +import sys + + +def main() -> int: + out_path, entries = sys.argv[1], sys.argv[2:] + lines = ["// Generated by embed_hip_code_objects.py. Do not edit.", ""] + table = [] + for index, entry in enumerate(entries): + stem, target, path = entry.split(":", 2) + data = open(path, "rb").read() + if data[:4] != b"\x7fELF": + print(f"{path}: not a raw ELF code object (offload bundle?)", file=sys.stderr) + return 1 + digest = hashlib.sha256(data).hexdigest()[:16] + lines.append(f"static const unsigned char kHipCodeObject{index}[] = {{") + for offset in range(0, len(data), 24): + lines.append(" " + ",".join(str(b) for b in data[offset:offset + 24]) + ",") + lines.append("};") + table.append(f' {{ "{stem}", "{target}", kHipCodeObject{index}, sizeof(kHipCodeObject{index}), "{digest}" }},') + lines.append("static const HipEmbeddedCodeObject kHipEmbeddedCodeObjects[] = {") + lines.extend(table) + lines.append(" { nullptr, nullptr, nullptr, 0, nullptr },") + lines.append("};") + text = "\n".join(lines) + "\n" + try: + if open(out_path).read() == text: + return 0 + except OSError: + pass + open(out_path, "w").write(text) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/ggml/src/ggml-hrx/hip/ggml-hrx-hip.cmake b/ggml/src/ggml-hrx/hip/ggml-hrx-hip.cmake new file mode 100644 index 000000000000..997309821fce --- /dev/null +++ b/ggml/src/ggml-hrx/hip/ggml-hrx-hip.cmake @@ -0,0 +1,134 @@ +# Copyright 2026 bong-water-water-bong +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# HIP kernels for the HRX backend: every hip/kernels/*.hip is compiled by amdclang++ into one raw +# code object per target in GGML_HRX_HIP_TARGETS and embedded in libggml-hrx. Nothing here links +# the HIP runtime: HRX loads and dispatches the code objects. See hip/README.md. + +set(GGML_HRX_HIP_TARGETS "gfx1151" CACHE STRING "GPU targets for the HRX HIP kernels (semicolon list)") +set(GGML_HRX_HIP_FLAGS "-O3" CACHE STRING "Extra amdclang++ flags for the HRX HIP kernels") +set(GGML_HRX_HIP_ADDON_DIR "" CACHE PATH + "Optional HIP kernel add-on directory: /kernels/*.hip join the code objects, /*.cpp join ggml-hrx, /addon.cmake is included if present (hip/README.md)") + +set(_ggml_hrx_hip_default_compiler "") +if (CMAKE_CXX_COMPILER MATCHES "amdclang\\+\\+$") + set(_ggml_hrx_hip_default_compiler "${CMAKE_CXX_COMPILER}") +endif() +find_program(GGML_HRX_HIP_COMPILER NAMES amdclang++ + HINTS ${_ggml_hrx_hip_default_compiler} /opt/rocm-therock/bin $ENV{ROCM_PATH}/bin /opt/rocm/bin + DOC "amdclang++ used to compile the HRX HIP kernels") +if (_ggml_hrx_hip_default_compiler AND NOT GGML_HRX_HIP_COMPILER) + set(GGML_HRX_HIP_COMPILER "${_ggml_hrx_hip_default_compiler}") +endif() +if (NOT GGML_HRX_HIP_COMPILER) + message(FATAL_ERROR "GGML_HRX needs amdclang++ for its HIP kernels; set GGML_HRX_HIP_COMPILER") +endif() + +file(GLOB GGML_HRX_HIP_KERNEL_SOURCES CONFIGURE_DEPENDS "${CMAKE_CURRENT_SOURCE_DIR}/hip/kernels/*.hip") + +# Add-on: its kernels are built and embedded exactly like the ones above, its matchers are compiled into +# ggml-hrx, and register_hip_dispatches calls the ggml_hrx_hip_addon_register it defines. +set(GGML_HRX_HIP_ADDON_SOURCES) +if (GGML_HRX_HIP_ADDON_DIR) + get_filename_component(GGML_HRX_HIP_ADDON_PATH "${GGML_HRX_HIP_ADDON_DIR}" ABSOLUTE BASE_DIR "${CMAKE_SOURCE_DIR}") + if (NOT IS_DIRECTORY "${GGML_HRX_HIP_ADDON_PATH}/kernels") + message(FATAL_ERROR "GGML_HRX_HIP_ADDON_DIR=${GGML_HRX_HIP_ADDON_DIR}: no kernels/ directory") + endif() + file(GLOB _ggml_hrx_hip_addon_kernels CONFIGURE_DEPENDS "${GGML_HRX_HIP_ADDON_PATH}/kernels/*.hip") + file(GLOB GGML_HRX_HIP_ADDON_SOURCES CONFIGURE_DEPENDS "${GGML_HRX_HIP_ADDON_PATH}/*.cpp" "${GGML_HRX_HIP_ADDON_PATH}/*.h") + list(APPEND GGML_HRX_HIP_KERNEL_SOURCES ${_ggml_hrx_hip_addon_kernels}) + list(LENGTH _ggml_hrx_hip_addon_kernels _ggml_hrx_hip_addon_kernel_count) + message(STATUS "GGML_HRX: HIP kernel add-on ${GGML_HRX_HIP_ADDON_PATH} (${_ggml_hrx_hip_addon_kernel_count} kernels)") +endif() + +# Code objects are looked up by file stem, so stems must be unique across this directory and the add-on. +set(_ggml_hrx_hip_stems) +foreach(_source ${GGML_HRX_HIP_KERNEL_SOURCES}) + get_filename_component(_stem "${_source}" NAME_WE) + if (_stem IN_LIST _ggml_hrx_hip_stems) + message(FATAL_ERROR "GGML_HRX: two HIP kernels named '${_stem}' (${_source})") + endif() + list(APPEND _ggml_hrx_hip_stems "${_stem}") +endforeach() + +separate_arguments(_ggml_hrx_hip_flags UNIX_COMMAND "${GGML_HRX_HIP_FLAGS}") +set(_ggml_hrx_hip_dir "${CMAKE_CURRENT_BINARY_DIR}/hip-code-objects") +file(MAKE_DIRECTORY "${_ggml_hrx_hip_dir}") + +set(_ggml_hrx_hip_objects) +set(_ggml_hrx_hip_entries) +foreach(_source ${GGML_HRX_HIP_KERNEL_SOURCES}) + get_filename_component(_stem "${_source}" NAME_WE) + get_filename_component(_source_dir "${_source}" DIRECTORY) + foreach(_target ${GGML_HRX_HIP_TARGETS}) + set(_object "${_ggml_hrx_hip_dir}/${_stem}.${_target}.hsaco") + add_custom_command( + OUTPUT "${_object}" + COMMAND "${GGML_HRX_HIP_COMPILER}" -x hip -std=c++17 --offload-arch=${_target} --cuda-device-only + --no-gpu-bundle-output ${_ggml_hrx_hip_flags} + -I "${_source_dir}" -I "${CMAKE_CURRENT_SOURCE_DIR}/hip/kernels" + -MD -MF "${_object}.d" -o "${_object}" "${_source}" + DEPENDS "${_source}" + DEPFILE "${_object}.d" + COMMENT "HIP code object ${_stem} (${_target})" + VERBATIM) + list(APPEND _ggml_hrx_hip_objects "${_object}") + list(APPEND _ggml_hrx_hip_entries "${_stem}:${_target}:${_object}") + endforeach() +endforeach() + +set(GGML_HRX_HIP_CODE_OBJECTS_INC "${CMAKE_CURRENT_BINARY_DIR}/ggml-hrx-hip-code-objects.inc") +add_custom_command( + OUTPUT "${GGML_HRX_HIP_CODE_OBJECTS_INC}" + COMMAND ${Python3_EXECUTABLE} "${CMAKE_CURRENT_SOURCE_DIR}/hip/embed_hip_code_objects.py" + "${GGML_HRX_HIP_CODE_OBJECTS_INC}" ${_ggml_hrx_hip_entries} + DEPENDS "${CMAKE_CURRENT_SOURCE_DIR}/hip/embed_hip_code_objects.py" ${_ggml_hrx_hip_objects} + COMMENT "Embedding HRX HIP code objects" + VERBATIM) +add_custom_target(ggml-hrx-hip-code-objects DEPENDS "${GGML_HRX_HIP_CODE_OBJECTS_INC}") + +# Code objects + kernel registry live with the kernel corpus (resolve_kernel_definition finds HIP +# kernels there); the loader that feeds the executable cache lives in ggml-hrx. +target_sources(ggml-hrx-kernel-corpus PRIVATE + hip/hip-code-objects.cpp + hip/hip-code-objects.h + hip/hip-kernel-registry.cpp + hip/hip-kernel-registry.h + "${GGML_HRX_HIP_CODE_OBJECTS_INC}") +add_dependencies(ggml-hrx-kernel-corpus ggml-hrx-hip-code-objects) +target_sources(ggml-hrx PRIVATE + hip/hip-kernel-loader.cpp + hip/hip-kernel-loader.h + hip/hip-capabilities.cpp + hip/hip-capabilities.h + hip/hip-dispatches.cpp + hip/hip-dispatches.h + hip/dispatch-hip-scale.cpp + ${GGML_HRX_HIP_ADDON_SOURCES}) +if (GGML_HRX_HIP_ADDON_DIR) + target_compile_definitions(ggml-hrx PRIVATE GGML_HRX_HIP_ADDON) +endif() + +# Standalone check that HRX loads and runs a hipcc code object (hip/kernels/hip_smoke.hip). +add_executable(ggml-hrx-hip-smoke hip/hip-smoke.cpp hip/hip-code-objects.cpp) +target_link_libraries(ggml-hrx-hip-smoke PRIVATE hrx::hrx) +target_include_directories(ggml-hrx-hip-smoke PRIVATE . "${CMAKE_CURRENT_BINARY_DIR}") +target_compile_features(ggml-hrx-hip-smoke PRIVATE cxx_std_17) +add_dependencies(ggml-hrx-hip-smoke ggml-hrx-hip-code-objects) + +# Add-on extras (standalone harnesses, benches); paths in it resolve as in this file. +if (GGML_HRX_HIP_ADDON_DIR) + include("${GGML_HRX_HIP_ADDON_PATH}/addon.cmake" OPTIONAL) +endif() diff --git a/ggml/src/ggml-hrx/hip/hip-capabilities.cpp b/ggml/src/ggml-hrx/hip/hip-capabilities.cpp new file mode 100644 index 000000000000..3eda21f530f0 --- /dev/null +++ b/ggml/src/ggml-hrx/hip/hip-capabilities.cpp @@ -0,0 +1,51 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "hip/hip-capabilities.h" + +#include "dispatch_registration/dispatch-registry.h" + +#include +#include + +namespace ggml::hrx { +namespace { + +std::array, GGML_OP_COUNT> g_declared_ops = {}; + +} // namespace + +void hip_declare_eager_op(enum ggml_op op) { + if (op >= 0 && op < GGML_OP_COUNT) { + g_declared_ops[op].store(true, std::memory_order_release); + } +} + +bool hip_eager_op_declared(enum ggml_op op) { + // the registries are function-local statics: the first lookup builds them (and runs every + // register_* function, which is where declarations happen) exactly once + static const bool registries_built = [] { + for (const char * architecture : { "gfx1151", "gfx1100" }) { + DispatchTarget target; + target.architecture = architecture; + find_dispatch_registry(target); + } + return true; + }(); + (void) registries_built; + return op >= 0 && op < GGML_OP_COUNT && g_declared_ops[op].load(std::memory_order_acquire); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/hip/hip-capabilities.h b/ggml/src/ggml-hrx/hip/hip-capabilities.h new file mode 100644 index 000000000000..f67229c0ce03 --- /dev/null +++ b/ggml/src/ggml-hrx/hip/hip-capabilities.h @@ -0,0 +1,36 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Extra eager-capability ops declared by HIP matchers. ggml-hrx's supports_op only claims a node +// whose op is in its built-in capability list; a HIP matcher for an op outside that list (built in +// here or from an add-on, see README.md "Kernel add-ons") declares the op while it registers, and +// eager_capability_declared consults this set for every op it does not know. Empty without such a +// declaration, so a build without matchers for extra ops behaves as before. + +#pragma once + +#include "ggml.h" + +namespace ggml::hrx { + +// Called from a matcher's register_* function (so from register_hip_dispatches / the add-on's +// ggml_hrx_hip_addon_register) while the dispatch registry is being built. +void hip_declare_eager_op(enum ggml_op op); + +// True when a matcher declared op. Builds the dispatch registries on first use so the declarations +// exist before the scheduler probes supports_op. +bool hip_eager_op_declared(enum ggml_op op); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/hip/hip-code-objects.cpp b/ggml/src/ggml-hrx/hip/hip-code-objects.cpp new file mode 100644 index 000000000000..32bf7ade9c02 --- /dev/null +++ b/ggml/src/ggml-hrx/hip/hip-code-objects.cpp @@ -0,0 +1,65 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "hip/hip-code-objects.h" + +#include + +namespace ggml::hrx { +namespace { + +struct HipEmbeddedCodeObject { + const char * stem; + const char * target; + const unsigned char * data; + size_t size; + const char * digest; +}; + +// Generated by hip/embed_hip_code_objects.py: kHipEmbeddedCodeObjects, terminated by a null stem. +#include "ggml-hrx-hip-code-objects.inc" + +const HipEmbeddedCodeObject * find_code_object(const char * stem, const char * target) { + if (stem == nullptr) { + return nullptr; + } + for (const HipEmbeddedCodeObject * object = kHipEmbeddedCodeObjects; object->stem != nullptr; ++object) { + if (std::strcmp(object->stem, stem) == 0 && (target == nullptr || std::strcmp(object->target, target) == 0)) { + return object; + } + } + return nullptr; +} + +} // namespace + +const void * hip_code_object_data(const char * stem, const char * target, size_t * size) { + const HipEmbeddedCodeObject * object = find_code_object(stem, target); + if (size != nullptr) { + *size = object != nullptr ? object->size : 0; + } + return object != nullptr ? object->data : nullptr; +} + +bool hip_code_object_embedded(const char * stem) { + return find_code_object(stem, nullptr) != nullptr; +} + +const char * hip_code_object_digest(const char * stem) { + const HipEmbeddedCodeObject * object = find_code_object(stem, nullptr); + return object != nullptr ? object->digest : nullptr; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/hip/hip-code-objects.h b/ggml/src/ggml-hrx/hip/hip-code-objects.h new file mode 100644 index 000000000000..cfa041c93d36 --- /dev/null +++ b/ggml/src/ggml-hrx/hip/hip-code-objects.h @@ -0,0 +1,34 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// HIP code objects embedded at build time (ggml-hrx-hip.cmake): one raw AMDGPU ELF per +// hip/kernels/.hip and target. No HRX or Loom dependencies, so tools can link it alone. + +#pragma once + +#include + +namespace ggml::hrx { + +// Code object bytes for (stem, target), or nullptr; target nullptr = any target. +const void * hip_code_object_data(const char * stem, const char * target, size_t * size); + +// True when the build embedded this stem for at least one target. +bool hip_code_object_embedded(const char * stem); + +// Short content digest of the stem's first embedded code object, or nullptr. +const char * hip_code_object_digest(const char * stem); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/hip/hip-dispatches.cpp b/ggml/src/ggml-hrx/hip/hip-dispatches.cpp new file mode 100644 index 000000000000..4bf58ba09086 --- /dev/null +++ b/ggml/src/ggml-hrx/hip/hip-dispatches.cpp @@ -0,0 +1,40 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "hip/hip-dispatches.h" + +namespace ggml::hrx { +namespace { + +HipAttentionSinkHook g_attention_sink_hook = nullptr; + +} // namespace + +void register_hip_dispatches(DispatchRegistryBuilder & registry) { + register_hip_scale_dispatch(registry); +#ifdef GGML_HRX_HIP_ADDON + ggml_hrx_hip_addon_register(registry); +#endif +} + +void set_hip_attention_sink_hook(HipAttentionSinkHook hook) { + g_attention_sink_hook = hook; +} + +bool hip_attention_sink_dispatch(const HipAttentionSinkArgs & args, Dispatch & dispatch) { + return g_attention_sink_hook != nullptr && g_attention_sink_hook(args, dispatch); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/hip/hip-dispatches.h b/ggml/src/ggml-hrx/hip/hip-dispatches.h new file mode 100644 index 000000000000..0859f2d893f5 --- /dev/null +++ b/ggml/src/ggml-hrx/hip/hip-dispatches.h @@ -0,0 +1,58 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// The one place HIP matchers are registered (hooked from register_common_dispatches), plus the +// entry points a HIP kernel add-on (GGML_HRX_HIP_ADDON_DIR, see README.md) plugs into. + +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +#include + +namespace ggml::hrx { + +void register_hip_dispatches(DispatchRegistryBuilder & registry); + +// Matchers built from this directory (one line each in hip-dispatches.cpp). +void register_hip_scale_dispatch(DispatchRegistryBuilder & registry); + +// Defined by the add-on when GGML_HRX_HIP_ADDON_DIR is set (the build then defines GGML_HRX_HIP_ADDON); +// registers the add-on's kernels and matchers. Called after the matchers above, once per registry build. +void ggml_hrx_hip_addon_register(DispatchRegistryBuilder & registry); + +// Attention-sink rescale hook. dispatch-attention-sink.cpp appends the rescale after FlashAttention; it +// has checked the node (query [tokens][heads][d] f32, key [capacity][kv_heads][d] f16, f16 mask rows +// key_count apart, f32 sinks, output [tokens][heads][dv] f32, no input overlapping the output). A hook +// that returns true has filled dispatch and replaces the Loom kernel; false keeps the Loom kernel. +struct HipAttentionSinkArgs { + const Value * query; + const Value * key; + const Value * mask; + const Value * sinks; + const Value * output; + int64_t tokens; + int64_t key_count; + int64_t heads; + int64_t kv_heads; + int64_t qk_head_size; + int64_t value_head_size; + float scale; +}; +using HipAttentionSinkHook = bool (*)(const HipAttentionSinkArgs & args, Dispatch & dispatch); +void set_hip_attention_sink_hook(HipAttentionSinkHook hook); +bool hip_attention_sink_dispatch(const HipAttentionSinkArgs & args, Dispatch & dispatch); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/hip/hip-kernel-loader.cpp b/ggml/src/ggml-hrx/hip/hip-kernel-loader.cpp new file mode 100644 index 000000000000..43d84b516dbd --- /dev/null +++ b/ggml/src/ggml-hrx/hip/hip-kernel-loader.cpp @@ -0,0 +1,69 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "hip/hip-kernel-loader.h" + +#include "ggml-impl.h" +#include "hip/hip-code-objects.h" +#include "hip/hip-kernel-registry.h" +#include "hrx_runtime.h" + +#include + +namespace ggml::hrx { + +LoomCompiledKernelRef make_hip_compiled_kernel(const std::string & key, + const KernelDefinition & definition, + const Dispatch & dispatch, + const char * target) { + auto compiled_ref = std::make_shared(key, LoomKernelCompileRequest{}); + ggml_hrx_loom_jit_compile_result result; + std::string error; + + const HipLaunchFunction launch = find_hip_kernel_launch(definition.id); + size_t size = 0; + const void * data = hip_code_object_data(definition.source, target, &size); + HipLaunchGeometry geometry; + if (launch == nullptr) { + error = "HIP kernel " + kernel_definition_name(definition) + " is not registered"; + } else if (data == nullptr) { + error = std::string("no ") + (target != nullptr ? target : "?") + " code object '" + definition.source + + "' embedded for HIP kernel " + kernel_definition_name(definition); + } else if (!launch(dispatch.kernel.integer_parameters, geometry)) { + error = "HIP kernel " + kernel_definition_name(definition) + " rejected its launch parameters"; + } else { + void * copy = nullptr; + hrx_status_t status = hrx_host_allocator_malloc_uninitialized(hrx_host_allocator_system(), size, ©); + if (!hrx_status_is_ok(status)) { + hrx_status_ignore(status); + error = "HIP code object allocation failed"; + } else { + std::memcpy(copy, data, size); + result.hsaco_data = copy; + result.hsaco_size = size; + result.launch_config.workgroup_count = geometry.workgroup_count; + result.launch_config.workgroup_size = geometry.workgroup_size; + result.launch_config.subgroup_size = (target != nullptr && std::strncmp(target, "gfx9", 4) == 0) ? 64 : 32; + } + } + if (!error.empty()) { + GGML_LOG_ERROR("%s: %s\n", __func__, error.c_str()); + } + const bool ok = error.empty(); + HipCodeObjectLoader::complete(*compiled_ref, std::move(result), ok, std::move(error)); + return compiled_ref; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/hip/hip-kernel-loader.h b/ggml/src/ggml-hrx/hip/hip-kernel-loader.h new file mode 100644 index 000000000000..967c835d9239 --- /dev/null +++ b/ggml/src/ggml-hrx/hip/hip-kernel-loader.h @@ -0,0 +1,43 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Turns a registered HIP kernel (hip-kernel-registry.h) into the compile result the HRX kernel +// executable cache loads: the embedded code object for the device target plus the launch +// geometry from the kernel's launch function. Hooked into KernelExecutableCache::get_or_compile. + +#pragma once + +#include "dispatch/dispatch.h" +#include "kernel-corpus/kernel-corpus.h" +#include "runtime/loom-kernel-jit.h" + +#include + +namespace ggml::hrx { + +LoomCompiledKernelRef make_hip_compiled_kernel(const std::string & key, + const KernelDefinition & definition, + const Dispatch & dispatch, + const char * target); + +// Lets make_hip_compiled_kernel complete a LoomCompiledKernel without a Loom compile. +struct HipCodeObjectLoader { + static void complete(LoomCompiledKernel & kernel, ggml_hrx_loom_jit_compile_result compiled, bool success, + std::string error) { + kernel.complete(std::move(compiled), success, std::move(error)); + } +}; + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/hip/hip-kernel-registry.cpp b/ggml/src/ggml-hrx/hip/hip-kernel-registry.cpp new file mode 100644 index 000000000000..3f264ce1f35b --- /dev/null +++ b/ggml/src/ggml-hrx/hip/hip-kernel-registry.cpp @@ -0,0 +1,103 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "hip/hip-kernel-registry.h" + +#include "hip/hip-code-objects.h" + +#include +#include +#include + +namespace ggml::hrx { +namespace { + +struct RegisteredHipKernel { + std::string name; + std::string code_object; + std::string digest; + std::vector bindings; + std::vector launch_parameters; + std::vector workload_parameters; + HipLaunchFunction launch = nullptr; + KernelDefinition definition; +}; + +std::mutex g_mutex; +std::deque g_kernels; // deque: definitions keep their addresses + +const RegisteredHipKernel * find_registered(uint64_t kernel_id) { + for (const RegisteredHipKernel & kernel : g_kernels) { + if (kernel.definition.id == kernel_id) { + return &kernel; + } + } + return nullptr; +} + +} // namespace + +void register_hip_kernel(const HipKernel & kernel) { + std::lock_guard lock(g_mutex); + const uint64_t id = hip_kernel_ref(kernel.name).id; + if (find_registered(id) != nullptr) { + return; + } + RegisteredHipKernel & entry = g_kernels.emplace_back(); + entry.name = kernel.name; + entry.code_object = kernel.code_object; + // The digest of any target's code object for this stem identifies the build in cache keys. + const char * digest = hip_code_object_digest(kernel.code_object); + entry.digest = digest != nullptr ? digest : "missing"; + entry.bindings = kernel.bindings; + for (const char * name : kernel.launch_parameters) { + entry.launch_parameters.push_back({ name, "index" }); + } + for (const char * name : kernel.workload_parameters) { + entry.workload_parameters.push_back({ name, "index" }); + } + entry.launch = kernel.launch; + + KernelDefinition & definition = entry.definition; + definition.family = kHipKernelFamily; + definition.name = entry.name.c_str(); + definition.id = id; + definition.source = entry.code_object.c_str(); + definition.symbol = entry.name.c_str(); + definition.backend = "amdgpu"; + definition.source_digest = entry.digest.c_str(); + definition.bindings = { entry.bindings.data(), entry.bindings.size() }; + definition.launch_parameters = { entry.launch_parameters.data(), entry.launch_parameters.size() }; + definition.workload_parameters = { entry.workload_parameters.data(), entry.workload_parameters.size() }; + definition.compile_recipe.mode = "hip"; +} + +const KernelDefinition * find_hip_kernel_definition(uint64_t kernel_id) { + std::lock_guard lock(g_mutex); + const RegisteredHipKernel * kernel = find_registered(kernel_id); + return kernel != nullptr ? &kernel->definition : nullptr; +} + +bool is_hip_kernel_definition(const KernelDefinition & definition) { + return definition.compile_recipe.mode != nullptr && std::strcmp(definition.compile_recipe.mode, "hip") == 0; +} + +HipLaunchFunction find_hip_kernel_launch(uint64_t kernel_id) { + std::lock_guard lock(g_mutex); + const RegisteredHipKernel * kernel = find_registered(kernel_id); + return kernel != nullptr ? kernel->launch : nullptr; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/hip/hip-kernel-registry.h b/ggml/src/ggml-hrx/hip/hip-kernel-registry.h new file mode 100644 index 000000000000..08ac2cc70a34 --- /dev/null +++ b/ggml/src/ggml-hrx/hip/hip-kernel-registry.h @@ -0,0 +1,76 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// HIP kernels dispatched through HRX. A .hip file under hip/kernels is compiled by amdclang++ at +// build time (one raw code object per GPU target, see ggml-hrx-hip.cmake), embedded in +// libggml-hrx, and loaded with hrx_executable_load_data like a Loom JIT result. A matcher refers +// to a HIP kernel with hip_kernel_ref("name") and fills Dispatch exactly as for a Loom kernel: +// +// - bindings -> the kernel's pointer arguments, in order; +// - launch_parameters -> uint32_t by-value arguments after the pointers (packed as constants); +// - workload_parameters -> integer parameters the launch function reads to size the grid. +// They are part of the executable cache key, so keep them to what changes the geometry. +// +// See hip/README.md for the full recipe. + +#pragma once + +#include "hip/hip-code-objects.h" +#include "kernel-corpus/kernel-corpus.h" + +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { + +inline constexpr const char kHipKernelFamily[] = "hip"; + +constexpr KernelCatalogRef hip_kernel_ref(const char * name) { + return kernel_catalog_ref(kHipKernelFamily, name); +} + +struct HipLaunchGeometry { + std::array workgroup_count = { 1, 1, 1 }; + std::array workgroup_size = { 1, 1, 1 }; +}; + +// Computes the grid from the dispatch's integer parameters; false rejects the dispatch. +using HipLaunchFunction = bool (*)(const std::map & parameters, HipLaunchGeometry & geometry); + +struct HipKernel { + const char * name = ""; // extern "C" __global__ symbol + const char * code_object = ""; // stem of the .hip file that defines it + std::vector bindings; + std::vector launch_parameters; // uint32_t arguments after the pointers + std::vector workload_parameters; // read by launch + HipLaunchFunction launch = nullptr; +}; + +// Registers a kernel (idempotent per name). Call it from the matcher's register_* function. +void register_hip_kernel(const HipKernel & kernel); + +// The registered kernel for a catalog id, or nullptr (hook in resolve_kernel_definition). +const KernelDefinition * find_hip_kernel_definition(uint64_t kernel_id); + +bool is_hip_kernel_definition(const KernelDefinition & definition); + +// The launch function of a registered kernel, or nullptr. +HipLaunchFunction find_hip_kernel_launch(uint64_t kernel_id); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/hip/hip-smoke.cpp b/ggml/src/ggml-hrx/hip/hip-smoke.cpp new file mode 100644 index 000000000000..2d0a235ad023 --- /dev/null +++ b/ggml/src/ggml-hrx/hip/hip-smoke.cpp @@ -0,0 +1,126 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// ggml-hrx-hip-smoke: loads the embedded hip_smoke code object through HRX (no HIP runtime), +// dispatches both kernels and checks the results on the host. Exit 0 = pass. + +#include "hip/hip-code-objects.h" +#include "hrx_runtime.h" + +#include +#include +#include +#include +#include +#include + +#define CHECK(expr) \ + do { \ + hrx_status_t _s = (expr); \ + if (!hrx_status_is_ok(_s)) { \ + char * _m = nullptr; \ + size_t _n = 0; \ + hrx_status_to_string(_s, &_m, &_n); \ + std::fprintf(stderr, "FAIL %s: %.*s\n", #expr, (int) _n, _m ? _m : ""); \ + return 1; \ + } \ + } while (0) + +int main() { + CHECK(hrx_gpu_initialize(0)); + hrx_device_t device = nullptr; + CHECK(hrx_gpu_device_get(0, &device)); + char arch[64] = {}; + CHECK(hrx_device_get_property(device, HRX_DEVICE_PROPERTY_ARCHITECTURE, arch, sizeof(arch))); + size_t size = 0; + const void * data = ggml::hrx::hip_code_object_data("hip_smoke", arch, &size); + if (data == nullptr) { + std::fprintf(stderr, "FAIL no hip_smoke code object for %s\n", arch); + return 1; + } + hrx_executable_t executable = nullptr; + CHECK(hrx_executable_load_data(device, data, size, "amdgpu", arch, &executable)); + + uint32_t axpy = 0, block_sum = 0; + CHECK(hrx_executable_lookup_export_by_name(executable, "hip_smoke_axpy", &axpy)); + CHECK(hrx_executable_lookup_export_by_name(executable, "hip_smoke_block_sum", &block_sum)); + for (uint32_t ordinal : { axpy, block_sum }) { + hrx_executable_export_info_t info = {}; + CHECK(hrx_executable_export_info(executable, ordinal, &info)); + std::printf("export %s: bindings=%u constants=%uB parameters=%u\n", info.name, info.binding_count, + info.constant_byte_length, info.parameter_count); + } + + hrx_stream_t stream = nullptr; + CHECK(hrx_stream_create(device, 0, &stream)); + const uint32_t n = 1000003, blocks = (n + 255) / 256; + hrx_buffer_t x = nullptr, y = nullptr, out = nullptr; + const hrx_memory_type_t mem = HRX_MEMORY_TYPE_HOST_LOCAL | HRX_MEMORY_TYPE_DEVICE_VISIBLE; + CHECK(hrx_buffer_allocate(stream, n * sizeof(float), mem, HRX_BUFFER_USAGE_DEFAULT | HRX_BUFFER_USAGE_MAPPING_SCOPED, &x)); + CHECK(hrx_buffer_allocate(stream, n * sizeof(float), mem, HRX_BUFFER_USAGE_DEFAULT | HRX_BUFFER_USAGE_MAPPING_SCOPED, &y)); + CHECK(hrx_buffer_allocate(stream, blocks * sizeof(float), mem, HRX_BUFFER_USAGE_DEFAULT | HRX_BUFFER_USAGE_MAPPING_SCOPED, &out)); + CHECK(hrx_stream_synchronize(stream)); + float * px = nullptr; + float * py = nullptr; + CHECK(hrx_buffer_map(x, HRX_MAP_WRITE, 0, n * sizeof(float), (void **) &px)); + CHECK(hrx_buffer_map(y, HRX_MAP_WRITE, 0, n * sizeof(float), (void **) &py)); + for (uint32_t i = 0; i < n; ++i) { + px[i] = (float) (i % 97) * 0.25f; + py[i] = (float) (i % 13); + } + CHECK(hrx_buffer_unmap(x)); + CHECK(hrx_buffer_unmap(y)); + + struct { + uint32_t n; + float a; + } axpy_constants = { n, 3.0f }; + hrx_buffer_ref_t axpy_refs[2] = { { x, 0, n * sizeof(float) }, { y, 0, n * sizeof(float) } }; + hrx_dispatch_config_t config = { { blocks, 1, 1 }, { 256, 1, 1 }, 32 }; + CHECK(hrx_stream_dispatch(stream, executable, axpy, &config, &axpy_constants, sizeof(axpy_constants), axpy_refs, 2, 0)); + uint32_t sum_constants = n; + hrx_buffer_ref_t sum_refs[2] = { { y, 0, n * sizeof(float) }, { out, 0, blocks * sizeof(float) } }; + CHECK(hrx_stream_dispatch(stream, executable, block_sum, &config, &sum_constants, sizeof(sum_constants), sum_refs, 2, 0)); + CHECK(hrx_stream_synchronize(stream)); + + CHECK(hrx_buffer_map(y, HRX_MAP_READ, 0, n * sizeof(float), (void **) &py)); + float * pout = nullptr; + CHECK(hrx_buffer_map(out, HRX_MAP_READ, 0, blocks * sizeof(float), (void **) &pout)); + size_t bad = 0; + for (uint32_t i = 0; i < n; ++i) { + const float want = 3.0f * ((float) (i % 97) * 0.25f) + (float) (i % 13); + bad += py[i] != want; + } + double max_err = 0; + for (uint32_t b = 0; b < blocks; ++b) { + double want = 0; + for (uint32_t i = b * 256; i < std::min(n, (b + 1) * 256); ++i) { + want += 3.0 * ((i % 97) * 0.25) + (i % 13); + } + max_err = std::fmax(max_err, std::fabs(want - pout[b]) / std::fmax(1.0, std::fabs(want))); + } + CHECK(hrx_buffer_unmap(y)); + CHECK(hrx_buffer_unmap(out)); + std::printf("arch %s: axpy mismatches %zu/%u, block_sum max rel err %.3g over %u blocks\n", arch, bad, n, max_err, + blocks); + hrx_buffer_release(x); + hrx_buffer_release(y); + hrx_buffer_release(out); + hrx_stream_release(stream); + hrx_executable_release(executable); + const bool pass = bad == 0 && max_err < 1e-5; + std::printf("%s\n", pass ? "PASS" : "FAIL"); + return pass ? 0 : 1; +} diff --git a/ggml/src/ggml-hrx/hip/kernels/hip_scale_f32.hip b/ggml/src/ggml-hrx/hip/kernels/hip_scale_f32.hip new file mode 100644 index 000000000000..6d491d3854e6 --- /dev/null +++ b/ggml/src/ggml-hrx/hip/kernels/hip_scale_f32.hip @@ -0,0 +1,32 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Example HIP kernel for the HRX HIP recipe (hip/README.md): GGML_OP_SCALE on packed f32, +// y = x * scale + bias. Floats travel as uint32 bit patterns because HRX launch parameters are +// uint32. Matcher: hip/dispatch-hip-scale.cpp (opt-in, GGML_HRX_HIP_EXAMPLE_SCALE=1). + +#include +#include + +extern "C" __global__ void __launch_bounds__(256) +hip_scale_f32(const float * __restrict__ x, float * __restrict__ y, uint32_t n, uint32_t scale_bits, + uint32_t bias_bits) { + const float scale = __builtin_bit_cast(float, scale_bits); + const float bias = __builtin_bit_cast(float, bias_bits); + const uint32_t stride = gridDim.x * blockDim.x; + for (uint32_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n; i += stride) { + y[i] = x[i] * scale + bias; + } +} diff --git a/ggml/src/ggml-hrx/hip/kernels/hip_smoke.hip b/ggml/src/ggml-hrx/hip/kernels/hip_smoke.hip new file mode 100644 index 000000000000..f2d69b72d843 --- /dev/null +++ b/ggml/src/ggml-hrx/hip/kernels/hip_smoke.hip @@ -0,0 +1,50 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Smoke kernels for the HRX HIP path (ggml-hrx-hip-smoke): pointer + by-value arguments, +// blockDim/gridDim (hidden kernargs), LDS and a wave32 shuffle reduction. + +#include +#include + +extern "C" __global__ void __launch_bounds__(256) +hip_smoke_axpy(const float * __restrict__ x, float * __restrict__ y, uint32_t n, float a) { + const uint32_t i = blockIdx.x * blockDim.x + threadIdx.x; + if (i < n) { + y[i] = a * x[i] + y[i]; + } +} + +// out[b] = sum of x over block b (256 threads per block, grid = gridDim.x blocks). +extern "C" __global__ void __launch_bounds__(256) +hip_smoke_block_sum(const float * __restrict__ x, float * __restrict__ out, uint32_t n) { + __shared__ float partial[256 / 32]; + const uint32_t i = blockIdx.x * blockDim.x + threadIdx.x; + float v = i < n ? x[i] : 0.0f; + for (int offset = 16; offset > 0; offset >>= 1) { + v += __shfl_xor(v, offset, 32); + } + if ((threadIdx.x & 31) == 0) { + partial[threadIdx.x / 32] = v; + } + __syncthreads(); + if (threadIdx.x == 0) { + float s = 0.0f; + for (uint32_t w = 0; w < blockDim.x / 32; ++w) { + s += partial[w]; + } + out[blockIdx.x] = s + (gridDim.x == 0 ? 1.0f : 0.0f); + } +} diff --git a/ggml/src/ggml-hrx/hrx-interop-utils.h b/ggml/src/ggml-hrx/hrx-interop-utils.h new file mode 100644 index 000000000000..08d3c2190639 --- /dev/null +++ b/ggml/src/ggml-hrx/hrx-interop-utils.h @@ -0,0 +1,31 @@ +#pragma once + +#include "hrx_runtime.h" + +#include +#include + +namespace ggml::hrx { + +// Success has no payload; failure carries the diagnostic produced by HRX or +// by the caller. Keeping this distinct from an empty string makes status tests +// explicit at API boundaries. +using ErrorResult = std::optional; + +inline ErrorResult take_status(hrx_status_t status) { + if (hrx_status_is_ok(status)) { + return std::nullopt; + } + char * message = nullptr; + size_t length = 0; + hrx_status_t format_status = hrx_status_to_string(status, &message, &length); + if (!hrx_status_is_ok(format_status)) { + hrx_status_ignore(format_status); + } + std::string result = message != nullptr ? std::string(message, length) : std::string("unknown HRX error"); + hrx_status_free_message(message); + hrx_status_ignore(status); + return result; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus-catalog-verify.h b/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus-catalog-verify.h new file mode 100644 index 000000000000..15af8e0bf4f8 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus-catalog-verify.h @@ -0,0 +1,16 @@ +#pragma once + +#include "kernel-corpus-catalog.h" + +namespace ggml::hrx { + +#include "kernel-corpus-catalog.inc" + +} // namespace ggml::hrx + +#define GGML_HRX_KERNEL_REF(family_literal, name_literal) \ + ([] { \ + static_assert(::ggml::hrx::kernel_catalog_entry_exists(family_literal, name_literal), \ + "unknown HRX kernel catalog entry"); \ + return ::ggml::hrx::kernel_catalog_ref(family_literal, name_literal); \ + }()) diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus-catalog.h b/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus-catalog.h new file mode 100644 index 000000000000..11e480071261 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus-catalog.h @@ -0,0 +1,56 @@ +#pragma once + +#include +#include + +namespace ggml::hrx { + +static constexpr uint64_t kUncatalogedKernelId = 0; + +constexpr bool kernel_catalog_name_equal(const char * lhs, const char * rhs) { + while (*lhs != 0 && *rhs != 0) { + if (*lhs != *rhs) { + return false; + } + ++lhs; + ++rhs; + } + return *lhs == *rhs; +} + +constexpr uint64_t kernel_catalog_id(const char * family, const char * name) { + uint64_t hash = UINT64_C(1469598103934665603); + while (*family != 0) { + hash ^= static_cast(*family); + hash *= UINT64_C(1099511628211); + ++family; + } + hash ^= 0; + hash *= UINT64_C(1099511628211); + while (*name != 0) { + hash ^= static_cast(*name); + hash *= UINT64_C(1099511628211); + ++name; + } + return hash; +} + +struct KernelCatalogRef { + const char * family = ""; + const char * name = ""; + uint64_t id = kUncatalogedKernelId; + + constexpr bool valid() const { + return id != kUncatalogedKernelId && family != nullptr && family[0] != 0 && name != nullptr && name[0] != 0; + } +}; + +constexpr KernelCatalogRef kernel_catalog_ref(const char * family, const char * name) { + return { + family, + name, + family != nullptr && name != nullptr ? kernel_catalog_id(family, name) : kUncatalogedKernelId, + }; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus-json.cpp b/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus-json.cpp new file mode 100644 index 000000000000..222ee9e259d4 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus-json.cpp @@ -0,0 +1,102 @@ +#include "kernel-corpus-json.h" + +#include + +namespace ggml::hrx { +namespace { + +const char * kernel_resource_access_name(ResourceAccess access) { + switch (access) { + case ResourceAccess::Read: + return "read"; + case ResourceAccess::Write: + return "write"; + case ResourceAccess::ReadWrite: + return "read_write"; + } + return "unknown"; +} + +nlohmann::ordered_json string_span_json(KernelSpan values) { + nlohmann::ordered_json result = nlohmann::ordered_json::array(); + for (const char * value : values) { + result.push_back(value != nullptr ? value : ""); + } + return result; +} + +nlohmann::ordered_json source_ref_span_json(KernelSpan values) { + nlohmann::ordered_json result = nlohmann::ordered_json::array(); + for (const KernelSourceRef & value : values) { + result.push_back(value.path != nullptr ? value.path : ""); + } + return result; +} + +nlohmann::ordered_json compile_config_json(KernelSpan values) { + nlohmann::ordered_json result = nlohmann::ordered_json::object(); + for (const KernelCompileConfig & value : values) { + result[value.key != nullptr ? value.key : ""] = value.value != nullptr ? value.value : ""; + } + return result; +} + +} // namespace + +std::string serialize_kernel_corpus_json(const KernelCorpus & corpus) { + nlohmann::ordered_json root = { + { "schema", corpus.schema }, + { "upstream_revision", corpus.upstream_revision }, + { "corpus_digest", corpus.corpus_digest }, + { "recipe_digest", corpus.recipe_digest }, + { "plan_case_count", corpus.plan_case_count }, + { "kernels", nlohmann::ordered_json::array() }, + }; + for (const KernelDefinition & kernel : corpus.kernels) { + nlohmann::ordered_json item = { + { "family", kernel.family }, + { "name", kernel.name }, + { "id", kernel.id }, + { "source", kernel.source }, + { "dependencies", string_span_json(kernel.dependencies) }, + { "symbol", kernel.symbol }, + { "backend", kernel.backend }, + { "target_selector", kernel.target_selector }, + { "compile_config", compile_config_json(kernel.compile_config) }, + { "scalar_parameters", string_span_json(kernel.scalar_parameters) }, + { "source_digest", kernel.source_digest }, + { "compile_recipe", + { + { "mode", kernel.compile_recipe.mode }, + { "link_module", kernel.compile_recipe.link_module }, + { "primary_sources", source_ref_span_json(kernel.compile_recipe.primary_sources) }, + { "library_sources", source_ref_span_json(kernel.compile_recipe.library_sources) }, + } }, + { "workload_parameters", nlohmann::ordered_json::array() }, + { "launch_parameters", nlohmann::ordered_json::array() }, + { "bindings", nlohmann::ordered_json::array() }, + }; + for (const KernelScalarDefinition & parameter : kernel.workload_parameters) { + item["workload_parameters"].push_back({ + { "name", parameter.name }, + { "type", parameter.type } + }); + } + for (const KernelScalarDefinition & parameter : kernel.launch_parameters) { + item["launch_parameters"].push_back({ + { "name", parameter.name }, + { "type", parameter.type } + }); + } + for (const KernelBindingDefinition & binding : kernel.bindings) { + item["bindings"].push_back({ + { "name", binding.name }, + { "access", kernel_resource_access_name(binding.access) } + }); + } + root["kernels"].push_back(std::move(item)); + } + return root.dump(); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus-json.h b/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus-json.h new file mode 100644 index 000000000000..2edf760b5976 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus-json.h @@ -0,0 +1,11 @@ +#pragma once + +#include "kernel-corpus.h" + +#include + +namespace ggml::hrx { + +std::string serialize_kernel_corpus_json(const KernelCorpus & corpus); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus.cpp b/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus.cpp new file mode 100644 index 000000000000..7a4bf065124f --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus.cpp @@ -0,0 +1,291 @@ +#include "kernel-corpus.h" +#include "hip/hip-kernel-registry.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +const char * kernel_resource_access_name(ResourceAccess access) { + switch (access) { + case ResourceAccess::Read: + return "read"; + case ResourceAccess::Write: + return "write"; + case ResourceAccess::ReadWrite: + return "read_write"; + } + return "unknown"; +} + +struct KernelSourceRecordEntry { + const char * source_path; + const KernelSource * source; +}; + +static bool string_equal(const char * lhs, const char * rhs) { + return std::strcmp(lhs != nullptr ? lhs : "", rhs != nullptr ? rhs : "") == 0; +} + +static bool string_empty(const char * value) { + return value == nullptr || value[0] == 0; +} + +static bool contains_source_ref(KernelSpan values, const char * path) { + return std::find_if(values.begin(), values.end(), + [&](const KernelSourceRef & item) { return string_equal(item.path, path); }) != values.end(); +} + +static bool string_span_equal(KernelSpan lhs, KernelSpan rhs) { + return lhs.size() == rhs.size() && std::equal(lhs.begin(), lhs.end(), rhs.begin(), + [](const char * a, const char * b) { return string_equal(a, b); }); +} + +static bool scalar_span_equal(KernelSpan lhs, KernelSpan rhs) { + return lhs.size() == rhs.size() && + std::equal(lhs.begin(), lhs.end(), rhs.begin(), + [](const KernelScalarDefinition & a, const KernelScalarDefinition & b) { + return string_equal(a.name, b.name) && string_equal(a.type, b.type); + }); +} + +static bool binding_span_equal(KernelSpan lhs, KernelSpan rhs) { + return lhs.size() == rhs.size() && + std::equal(lhs.begin(), lhs.end(), rhs.begin(), + [](const KernelBindingDefinition & a, const KernelBindingDefinition & b) { + return string_equal(a.name, b.name) && a.access == b.access; + }); +} + +static bool kernel_variant_contract_equal(const KernelDefinition & lhs, const KernelDefinition & rhs) { + return string_equal(lhs.backend, rhs.backend) && string_span_equal(lhs.scalar_parameters, rhs.scalar_parameters) && + scalar_span_equal(lhs.workload_parameters, rhs.workload_parameters) && + scalar_span_equal(lhs.launch_parameters, rhs.launch_parameters) && + binding_span_equal(lhs.bindings, rhs.bindings); +} + +// clang-format off +#include "kernel-corpus-sources.inc" +#include "kernel-corpus-qwen.inc" +// clang-format on + +} // namespace + +const KernelSource * get_kernel_source(const char * source_path) { + if (source_path == nullptr) { + return nullptr; + } + for (const KernelSourceRecordEntry & entry : kKernelSourceRecords) { + if (std::strcmp(source_path, entry.source_path) == 0) { + return entry.source; + } + } + return nullptr; +} + +const KernelCorpus & get_qwen_kernel_corpus() { + return kQwenKernelCorpus; +} + +KernelResolveResult resolve_kernel_definition(const KernelCorpus & corpus, + const std::string & target, + uint64_t kernel_id) { + if (kernel_id == kUncatalogedKernelId) { + return { KernelResolveStatus::UncatalogedKernel, nullptr }; + } + if (const KernelDefinition * hip = find_hip_kernel_definition(kernel_id)) { + return { KernelResolveStatus::Found, hip }; + } + const KernelDefinition * first_match = nullptr; + const KernelDefinition * default_variant = nullptr; + bool target_mismatch = false; + for (const KernelDefinition & kernel : corpus.kernels) { + if (kernel.id != kernel_id) { + continue; + } + if (first_match == nullptr) { + first_match = &kernel; + } else if (!string_equal(first_match->family, kernel.family) || !string_equal(first_match->name, kernel.name)) { + return { KernelResolveStatus::HashCollision, nullptr }; + } + if (string_equal(kernel.target_selector, target.c_str())) { + return { KernelResolveStatus::Found, &kernel }; + } + if (string_empty(kernel.target_selector)) { + default_variant = &kernel; + } else { + target_mismatch = true; + } + } + if (default_variant != nullptr) { + return { KernelResolveStatus::Found, default_variant }; + } + if (target_mismatch) { + return { KernelResolveStatus::UnsupportedTarget, first_match }; + } + return { KernelResolveStatus::MissingActiveCorpusEntry, nullptr }; +} + +const char * kernel_resolve_status_name(KernelResolveStatus status) { + switch (status) { + case KernelResolveStatus::Found: + return "found"; + case KernelResolveStatus::UncatalogedKernel: + return "uncataloged_kernel"; + case KernelResolveStatus::MissingActiveCorpusEntry: + return "missing_active_corpus_entry"; + case KernelResolveStatus::HashCollision: + return "hash_collision"; + case KernelResolveStatus::UnsupportedTarget: + return "unsupported_target"; + } + return "unknown"; +} + +std::string kernel_definition_name(const KernelDefinition & definition) { + return std::string(definition.family != nullptr ? definition.family : "") + ":" + + (definition.name != nullptr ? definition.name : ""); +} + +std::string kernel_definition_name_or_id(const KernelDefinition * definition, uint64_t kernel_id) { + if (definition != nullptr) { + return kernel_definition_name(*definition); + } + return "kernel_id=" + std::to_string(kernel_id); +} + +std::string format_kernel_resolve_error(const KernelResolveResult & result, uint64_t kernel_id) { + const std::string label = kernel_definition_name_or_id(result.definition, kernel_id); + switch (result.status) { + case KernelResolveStatus::Found: + return ""; + case KernelResolveStatus::UncatalogedKernel: + return "uncataloged kernel " + label; + case KernelResolveStatus::MissingActiveCorpusEntry: + return "cataloged kernel " + label + " is not available in the active corpus"; + case KernelResolveStatus::HashCollision: + return "kernel catalog id collision while resolving " + label; + case KernelResolveStatus::UnsupportedTarget: + return "cataloged kernel " + label + " has no implementation for the requested target"; + } + return "unknown kernel resolution failure for " + label; +} + +VerificationResult verify_kernel_corpus(const KernelCorpus & corpus) { + VerificationResult result; + if (!string_equal(corpus.schema, "ggml-hrx-kernel-corpus-v2")) { + result.status.log("unsupported kernel corpus schema"); + } + if (string_empty(corpus.upstream_revision)) { + result.status.log("kernel corpus has no upstream revision"); + } + if (string_empty(corpus.corpus_digest)) { + result.status.log("kernel corpus has no digest"); + } + if (string_empty(corpus.recipe_digest)) { + result.status.log("kernel corpus has no BUILD.bazel recipe digest"); + } + if (corpus.plan_case_count == 0) { + result.status.log("kernel corpus has no compile plan cases"); + } + std::set variants; + std::map contracts; + for (const KernelDefinition & kernel : corpus.kernels) { + if (string_empty(kernel.family) || string_empty(kernel.name) || string_empty(kernel.source) || + string_empty(kernel.symbol) || string_empty(kernel.backend) || string_empty(kernel.source_digest)) { + result.status.log("kernel definition is incomplete"); + } + if (kernel.id != kernel_catalog_id(kernel.family != nullptr ? kernel.family : "", + kernel.name != nullptr ? kernel.name : "")) { + result.status.log("kernel %s has an invalid catalog id", kernel.name != nullptr ? kernel.name : ""); + } + const bool source_is_primary = contains_source_ref(kernel.compile_recipe.primary_sources, kernel.source); + const bool source_is_library = contains_source_ref(kernel.compile_recipe.library_sources, kernel.source); + if ((!string_equal(kernel.compile_recipe.mode, "direct") && + !string_equal(kernel.compile_recipe.mode, "archive")) || + kernel.compile_recipe.primary_sources.empty() || (!source_is_primary && !source_is_library) || + (string_equal(kernel.compile_recipe.mode, "archive") && string_empty(kernel.compile_recipe.link_module))) { + result.status.log("kernel %s has an invalid BUILD compile recipe", + kernel.name != nullptr ? kernel.name : ""); + } + for (const KernelSourceRef & source : kernel.compile_recipe.primary_sources) { + if (string_empty(source.path) || source.contents == nullptr) { + result.status.log("kernel %s has an invalid embedded primary source reference", + kernel.name != nullptr ? kernel.name : ""); + } + } + for (const KernelSourceRef & source : kernel.compile_recipe.library_sources) { + if (string_empty(source.path) || source.contents == nullptr) { + result.status.log("kernel %s has an invalid embedded library source reference", + kernel.name != nullptr ? kernel.name : ""); + } + } + const std::string full_name = std::string(kernel.family != nullptr ? kernel.family : "") + ":" + + std::string(kernel.name != nullptr ? kernel.name : ""); + const std::string target_selector = kernel.target_selector != nullptr ? kernel.target_selector : ""; + if (!variants.insert(full_name + "@" + target_selector).second) { + result.status.log("kernel corpus repeats target variant %s@%s", full_name.c_str(), + target_selector.empty() ? "default" : target_selector.c_str()); + } + const auto contract = contracts.emplace(full_name, &kernel); + if (!contract.second && !kernel_variant_contract_equal(*contract.first->second, kernel)) { + result.status.log("kernel target variants disagree on ABI for %s", full_name.c_str()); + } + std::set binding_names; + for (const KernelBindingDefinition & binding : kernel.bindings) { + if (string_empty(binding.name) || + !binding_names.insert(binding.name != nullptr ? binding.name : "").second) { + result.status.log("kernel %s has invalid binding names", kernel.name != nullptr ? kernel.name : ""); + } + } + if (kernel.bindings.size() == 0) { + result.status.log("kernel %s has no binding ABI", kernel.name != nullptr ? kernel.name : ""); + } + } + return result; +} + +std::string format_kernel_corpus(const KernelCorpus & corpus) { + std::ostringstream out; + out << "kernel-corpus " << corpus.schema << " revision=" << corpus.upstream_revision + << " digest=" << corpus.corpus_digest << " recipe=" << corpus.recipe_digest + << " kernels=" << corpus.kernels.size() << " plan_cases=" << corpus.plan_case_count << '\n'; + for (const KernelDefinition & kernel : corpus.kernels) { + out << " kernel " << kernel.family << ':' << kernel.name << " id=0x" << std::hex << kernel.id << std::dec + << " backend=" << kernel.backend + << " target=" << (string_empty(kernel.target_selector) ? "default" : kernel.target_selector) << " symbol=@" + << kernel.symbol << " source=" << kernel.source << " sha256=" << kernel.source_digest << '\n'; + out << " recipe " << kernel.compile_recipe.mode; + if (!string_empty(kernel.compile_recipe.link_module)) { + out << " module=" << kernel.compile_recipe.link_module; + } + out << " primary="; + for (size_t i = 0; i < kernel.compile_recipe.primary_sources.size(); ++i) { + out << (i ? "," : "") << kernel.compile_recipe.primary_sources[i].path; + } + out << " libraries="; + for (size_t i = 0; i < kernel.compile_recipe.library_sources.size(); ++i) { + out << (i ? "," : "") << kernel.compile_recipe.library_sources[i].path; + } + out << '\n'; + for (size_t i = 0; i < kernel.bindings.size(); ++i) { + out << " binding[" << i << "] " << kernel.bindings[i].name << ' ' + << kernel_resource_access_name(kernel.bindings[i].access) << '\n'; + } + for (const KernelScalarDefinition & parameter : kernel.workload_parameters) { + out << " workload " << parameter.name << ' ' << parameter.type << '\n'; + } + for (const KernelScalarDefinition & parameter : kernel.launch_parameters) { + out << " launch " << parameter.name << ' ' << parameter.type << '\n'; + } + } + return out.str(); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus.h b/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus.h new file mode 100644 index 000000000000..501eae8c19fa --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernel-corpus.h @@ -0,0 +1,129 @@ +#pragma once + +#include "kernel-corpus-catalog.h" +#include "kernel-types.h" + +#include +#include +#include + +namespace ggml::hrx { + +template struct KernelSpan { + const T * items = nullptr; + size_t count = 0; + + const T * begin() const { return items; } + + const T * end() const { return items == nullptr ? nullptr : items + count; } + + const T * data() const { return items; } + + size_t size() const { return count; } + + bool empty() const { return count == 0; } + + const T & operator[](size_t index) const { return items[index]; } + + const T & front() const { return items[0]; } +}; + +struct KernelCompileConfig { + const char * key = ""; + const char * value = ""; +}; + +struct KernelBindingDefinition { + const char * name = ""; + ResourceAccess access = ResourceAccess::Read; +}; + +struct KernelScalarDefinition { + const char * name = ""; + const char * type = ""; +}; + +enum KernelSourceFormat { + KERNEL_SOURCE_FORMAT_TEXT, + KERNEL_SOURCE_FORMAT_BINARY, +}; + +struct KernelSourceSpan { + const char * data; + size_t length; + KernelSourceFormat format; +}; + +struct KernelSource { + KernelSourceSpan source; + const KernelSourceSpan * dependencies; + size_t dependency_count; +}; + +struct KernelSourceRef { + const char * path = ""; + const KernelSource * contents = nullptr; +}; + +struct KernelCompileRecipe { + const char * mode = ""; + const char * link_module = ""; + KernelSpan primary_sources; + KernelSpan library_sources; +}; + +struct KernelDefinition { + const char * family = ""; + const char * name = ""; + uint64_t id = kUncatalogedKernelId; + const char * source = ""; + KernelSpan dependencies; + const char * symbol = ""; + const char * backend = ""; + const char * target_selector = ""; + KernelSpan compile_config; + KernelSpan scalar_parameters; + KernelSpan bindings; + const char * source_digest = ""; + KernelSpan workload_parameters; + KernelSpan launch_parameters; + KernelCompileRecipe compile_recipe; +}; + +struct KernelCorpus { + const char * schema = "ggml-hrx-kernel-corpus-v2"; + const char * upstream_revision = ""; + const char * corpus_digest = ""; + const char * recipe_digest = ""; + size_t plan_case_count = 0; + KernelSpan kernels; +}; + +enum class KernelResolveStatus : uint8_t { + Found, + UncatalogedKernel, + MissingActiveCorpusEntry, + HashCollision, + UnsupportedTarget, +}; + +struct KernelResolveResult { + KernelResolveStatus status = KernelResolveStatus::MissingActiveCorpusEntry; + const KernelDefinition * definition = nullptr; + + bool found() const { return status == KernelResolveStatus::Found && definition != nullptr; } +}; + +const KernelSource * get_kernel_source(const char * source_path); +const KernelCorpus & get_qwen_kernel_corpus(); +KernelResolveResult resolve_kernel_definition(const KernelCorpus & corpus, + const std::string & target, + uint64_t kernel_id); +const char * kernel_resolve_status_name(KernelResolveStatus status); +std::string kernel_definition_name(const KernelDefinition & definition); +std::string kernel_definition_name_or_id(const KernelDefinition * definition, uint64_t kernel_id); +std::string format_kernel_resolve_error(const KernelResolveResult & result, uint64_t kernel_id); +VerificationResult verify_kernel_corpus(const KernelCorpus & corpus); +std::string format_kernel_corpus(const KernelCorpus & corpus); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernel-types.h b/ggml/src/ggml-hrx/kernel-corpus/kernel-types.h new file mode 100644 index 000000000000..eb63504cb07f --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernel-types.h @@ -0,0 +1,22 @@ +#pragma once + +#include "status.h" + +#include +#include + +namespace ggml::hrx { + +enum class ResourceAccess : uint8_t { + Read, + Write, + ReadWrite, +}; + +struct VerificationResult { + Status status; + + bool valid() const { return status.success(); } +}; + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/hrx/gather_add_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/hrx/gather_add_f32.loom new file mode 100644 index 000000000000..f645dedd9d6d --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/hrx/gather_add_f32.loom @@ -0,0 +1,66 @@ +// Copyright 2026 The IREE Authors +// +// Licensed under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception + +// Publishes an indexed residual without assuming that the requested rows are +// contiguous or at the end of the source. This is the physical form of two +// GGML GET_ROWS operations followed by ADD. +amdgpu.target @ggml_gather_add_gfx11_wave64 {subgroup_size = 64} + +kernel.def target(@ggml_gather_add_gfx11_wave64) export("ggml_gather_add_f32") @ggml_gather_add_f32(%source_token_count: index, %output_token_count: index, %hidden_size: index) { + %one = index.constant 1 : index + %twofiftysix = index.constant 256 : index + %rounding = index.constant 255 : index + %rounded_width = index.add %hidden_size, %rounding : index + %column_workgroup_count = index.div %rounded_width, %twofiftysix : index + kernel.launch.config workgroups(%column_workgroup_count, %output_token_count, %one) workgroup_size(%twofiftysix, %one, %one) : index +} launch(%source_token_count: index, %output_token_count: index, %hidden_size: index, %attention: buffer, %residual: buffer, %output_ids: buffer, %output: buffer) { + %source_count = index.assume %source_token_count [range(%source_token_count, 1, 2048)] : index + %output_count = index.assume %output_token_count [range(%output_token_count, 1, 2048)] : index + %width = index.assume %hidden_size [range(%hidden_size, 1, 65536)] : index + %column_workgroup = kernel.workgroup.id : index + %output_row0 = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %twofiftysix = index.constant 256 : index + %output_row, %launch_output_count = index.assume %output_row0, %output_count [lt(%output_row0, %output_count)] : index, index + %base0 = index.mul %column_workgroup, %twofiftysix : index + %column0 = index.add %base0, %workitem : index + %column = index.assume %column0 [range(%column0, 0, 65535)] : index + %in_bounds = index.cmp ult, %column, %width : index + %zero_offset = index.constant 0 : offset + %attention_noalias, %residual_noalias, %ids_noalias, %output_noalias = buffer.assume.noalias %attention, %residual, %output_ids, %output : buffer, buffer, buffer, buffer + %attention_view = buffer.view %attention_noalias[%zero_offset] : buffer -> view<[%source_count]x[%width]xf32> + %residual_view = buffer.view %residual_noalias[%zero_offset] : buffer -> view<[%source_count]x[%width]xf32> + %ids_view = buffer.view %ids_noalias[%zero_offset] : buffer -> view<[%launch_output_count]xi32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%launch_output_count]x[%width]xf32> + scf.if %in_bounds { + %safe_column = index.assume %column [lt(%column, %width)] : index + %source_row_i32 = view.load %ids_view[%output_row] : view<[%launch_output_count]xi32> -> i32 + %source_row0 = index.cast %source_row_i32 : i32 to index + %source_row = index.assume %source_row0 [range(%source_row0, 0, 2047), lt(%source_row0, %source_count)] : index + %attention_value = view.load %attention_view[%source_row, %safe_column] : view<[%source_count]x[%width]xf32> -> f32 + %residual_value = view.load %residual_view[%source_row, %safe_column] : view<[%source_count]x[%width]xf32> -> f32 + %sum = scalar.addf %attention_value, %residual_value : f32 + view.store %sum, %output_view[%output_row, %safe_column] : f32, view<[%launch_output_count]x[%width]xf32> + } + kernel.return +} + +// Select source row 2 rather than a positional tail inferred by the host. +check.case public @ggml_gather_add_noncontiguous_case { + %one = check.literal value(1) : index + %three = check.literal value(3) : index + %four = check.literal value(4) : index + %attention = check.generate.iota offset(0.0) step(1.0) : tensor<3x4xf32> + %residual = check.generate.fill value(100.0) : tensor<3x4xf32> + %output_ids = check.generate.fill value(2) : tensor<1xi32> + %output = check.generate.fill value(0.0) : tensor<1x4xf32> + %expected = check.generate.iota offset(108.0) step(1.0) : tensor<1x4xf32> + kernel.launch @ggml_gather_add_f32[%three, %one, %four](%three, %one, %four, %attention, %residual, %output_ids, %output) : [index, index, index](index, index, index, tensor<3x4xf32>, tensor<3x4xf32>, tensor<1xi32>, tensor<1x4xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<1x4xf32> + check.return +} + +check.benchmark<@ggml_gather_add_noncontiguous_case> @ggml_gather_add_noncontiguous diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/hrx/manifest.json b/ggml/src/ggml-hrx/kernel-corpus/kernels/hrx/manifest.json new file mode 100644 index 000000000000..328dfd28126c --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/hrx/manifest.json @@ -0,0 +1,126 @@ +{ + "schema": "ggml-hrx-kernel-corpus-v1", + "upstream_revision": "local", + "files": [ + { + "path": "gather_add_f32.loom" + }, + { + "path": "mul_mat_vec_iq3xxs_f32.loom" + } + ], + "exports": [ + { + "name": "ggml_gather_add_f32", + "symbol": "ggml_gather_add_f32", + "source": "gather_add_f32.loom", + "workload_parameters": [ + { + "name": "source_token_count", + "type": "index" + }, + { + "name": "output_token_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "source_token_count", + "type": "index" + }, + { + "name": "output_token_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ], + "bindings": [ + "attention", + "residual", + "output_ids", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "gather_add_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [], + "family": "hrx" + }, + { + "name": "ggml_mul_mat_vec_iq3xxs_f32", + "symbol": "ggml_mul_mat_vec_iq3xxs_f32", + "source": "mul_mat_vec_iq3xxs_f32.loom", + "workload_parameters": [ + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + }, + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + }, + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "activation", + "weight", + "tables", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "mul_mat_vec_iq3xxs_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [], + "family": "hrx" + } + ], + "link_modules": [], + "plan_cases": [] +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/hrx/mul_mat_vec_iq3xxs_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/hrx/mul_mat_vec_iq3xxs_f32.loom new file mode 100644 index 000000000000..8eb5dced7f90 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/hrx/mul_mat_vec_iq3xxs_f32.loom @@ -0,0 +1,169 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// IQ3_XXS matrix-vector product: out[row] = dot(weight[row,:], activation[:]). +// +// One workitem per output row; each row loops over its 256-value IQ3_XXS blocks +// and decodes them with the GGML grid/sign codebook (see dequantize_row_iq3_xxs). +amdgpu.target @ggml_mmviq3_gfx11_wave64 {subgroup_size = 64} + +kernel.def target(@ggml_mmviq3_gfx11_wave64) export("ggml_mul_mat_vec_iq3xxs_f32") @ggml_mul_mat_vec_iq3xxs_f32(%input_size: index, %output_size: index, %token_count: index) { + %c1 = index.constant 1 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %rows_plus = index.add %output_size, %c63 : index + %wgx = index.div %rows_plus, %c64 : index + kernel.launch.config workgroups(%wgx, %token_count, %c1) workgroup_size(%c64, %c1, %c1) : index +} launch(%input_size: index, %output_size: index, %token_count: index, %activation: buffer, %weight: buffer, %tables: buffer, %output: buffer) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c7 = index.constant 7 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c98 = index.constant 98 : index + %c256 = index.constant 256 : index + %o1 = index.constant 1 : offset + %o2 = index.constant 2 : offset + %o4 = index.constant 4 : offset + %o66 = index.constant 66 : offset + %o98 = index.constant 98 : offset + %o1024 = index.constant 1024 : offset + %o1152 = index.constant 1152 : offset + %c255 = index.constant 255 : index + %zero = scalar.constant 0.0 : f32 + %halfc = scalar.constant 0.5 : f32 + %negone = scalar.constant -1.0 : f32 + %posone = scalar.constant 1.0 : f32 + %wg = kernel.workgroup.id : index + %tok = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %rowi = index.mul %wg, %c64 : index + %row = index.add %rowi, %lane : index + %row_ok = index.cmp ult, %row, %output_size : index + %activation_na, %weight_na, %output_na, %tables_na = buffer.assume.noalias %activation, %weight, %output, %tables : buffer, buffer, buffer, buffer + %blocks = index.div %input_size, %c256 : index + %row_bytes_i = index.mul %blocks, %c98 : index + %row_bytes = index.scale %row_bytes_i, %o1 : index, offset -> offset + %row_off = index.scale %row, %row_bytes : index, offset -> offset + scf.if %row_ok { + %acc = scf.for %kb = [%c0 to %blocks step %c1](%acc0 = %zero : f32) -> (f32) { + %kbo = index.scale %kb, %o98 : index, offset -> offset + %blk = index.add %row_off, %kbo : offset + %dview = buffer.view %weight_na[%blk] : buffer -> view<1xf16> + %dv = view.load %dview[%c0] : view<1xf16> -> f16 + %d = scalar.extf %dv : f16 to f32 + %acc1 = scf.for %v = [%c0 to %c256 step %c1](%accv = %acc0 : f32) -> (f32) { + // --- decode weight value %v of block %kb --- + %ib32 = index.div %v, %c32 : index + %remv = index.rem %v, %c32 : index + %lv = index.div %remv, %c8 : index + %pv = index.rem %remv, %c8 : index + %halfv = index.div %pv, %c4 : index + %jjv = index.rem %pv, %c4 : index + %qsa = index.mul %ib32, %c8 : index + %qsb = index.mul %lv, %c2 : index + %qsc = index.add %qsa, %qsb : index + %qsiv = index.add %qsc, %halfv : index + %qsov = index.scale %qsiv, %o1 : index, offset -> offset + %qsbase = index.add %blk, %o2 : offset + %qspv = index.add %qsbase, %qsov : offset + %qsview = buffer.view %weight_na[%qspv] : buffer -> view<1xi8> + %gidx_i8 = view.load %qsview[%c0] : view<1xi8> -> i8 + %gidx_s = index.cast %gidx_i8 : i8 to index + %gidx0 = index.andi %gidx_s, %c255 : index + %gidx = index.assume %gidx0 [range(%gidx0, 0, 255)] : index + %ga = index.mul %gidx, %c4 : index + %goi = index.add %ga, %jjv : index + %go = index.scale %goi, %o1 : index, offset -> offset + %gridview = buffer.view %tables_na[%go] : buffer -> view<1xi8> + %gbyte_i8 = view.load %gridview[%c0] : view<1xi8> -> i8 + %gbyte = scalar.uitofp %gbyte_i8 : i8 to f32 + %ssa = index.mul %ib32, %c4 : index + %sso = index.scale %ssa, %o1 : index, offset -> offset + %ssb = index.add %blk, %o66 : offset + %ssp = index.add %ssb, %sso : offset + %auxview = buffer.view %weight_na[%ssp] : buffer -> view<1xi32> + %aux32 = vector.load %auxview[%c0] : view<1xi32> -> vector<1xi32> + %c28v = vector.constant 28 : vector<1xi32> + %snib = vector.shrui %aux32, %c28v : vector<1xi32> + %sf_i32 = vector.extract %snib[0] : vector<1xi32> -> i32 + %sf = scalar.uitofp %sf_i32 : i32 to f32 + %scp = scalar.addf %halfc, %sf : f32 + %dhalf = scalar.mulf %d, %halfc : f32 + %dbv = scalar.mulf %dhalf, %scp : f32 + %sh7 = index.mul %lv, %c7 : index + %sh7_i32 = index.cast %sh7 : index to i32 + %shv = vector.splat %sh7_i32 : vector<1xi32> + %sgnsh = vector.shrui %aux32, %shv : vector<1xi32> + %c127v = vector.constant 127 : vector<1xi32> + %sgnmk = vector.andi %sgnsh, %c127v : vector<1xi32> + %sgnidx_i32 = vector.extract %sgnmk[0] : vector<1xi32> -> i32 + %sgnidx_s = index.cast %sgnidx_i32 : i32 to index + %sgnidx = index.assume %sgnidx_s [range(%sgnidx_s, 0, 127)] : index + %sgno = index.scale %sgnidx, %o1 : index, offset -> offset + %ksabs = index.add %o1024, %sgno : offset + %ksview = buffer.view %tables_na[%ksabs] : buffer -> view<1xi8> + %sgni8 = view.load %ksview[%c0] : view<1xi8> -> i8 + %sgni32 = index.cast %sgni8 : i8 to index + %kmo = index.scale %pv, %o1 : index, offset -> offset + %kmabs = index.add %o1152, %kmo : offset + %kmview = buffer.view %tables_na[%kmabs] : buffer -> view<1xi8> + %kmi8 = view.load %kmview[%c0] : view<1xi8> -> i8 + %kmi32 = index.cast %kmi8 : i8 to index + %andb = index.andi %sgni32, %kmi32 : index + %isneg = index.cmp ne, %andb, %c0 : index + %signf = scf.select %isneg, %negone, %posone : f32 + %dbg = scalar.mulf %dbv, %gbyte : f32 + %wval = scalar.mulf %dbg, %signf : f32 + // --- activation value x[kb*256 + v] --- + %xtok = index.mul %tok, %input_size : index + %xai = index.mul %kb, %c256 : index + %xai1 = index.add %xtok, %xai : index + %xai2 = index.add %xai1, %v : index + %xao = index.scale %xai2, %o4 : index, offset -> offset + %xview = buffer.view %activation_na[%xao] : buffer -> view<1xf32> + %xval = view.load %xview[%c0] : view<1xf32> -> f32 + %prod = scalar.mulf %wval, %xval : f32 + %acc2 = scalar.addf %accv, %prod : f32 + scf.yield %acc2 : f32 + } + scf.yield %acc1 : f32 + } + %otok = index.mul %tok, %output_size : index + %orow = index.add %otok, %row : index + %outo = index.scale %orow, %o4 : index, offset -> offset + %outv = buffer.view %output_na[%outo] : buffer -> view<1xf32> + view.store %acc, %outv[%c0] : f32, view<1xf32> + } + kernel.return +} + +check.case public @ggml_mul_mat_vec_iq3xxs_f32_coverage { + %in = check.literal value(256) : index + %out = check.literal value(64) : index + %tok = check.literal value(1) : index + %activation = check.generate.fill value(0.0) : tensor<256xf32> + %weight = check.generate.fill value(0) : tensor<64x98xi8> + %tables = check.generate.fill value(0) : tensor<1168xi8> + %output = check.generate.fill value(0.0) : tensor<64xf32> + kernel.launch @ggml_mul_mat_vec_iq3xxs_f32[%in, %out, %tok](%in, %out, %tok, %activation, %weight, %tables, %output) : [index, index, index](index, index, index, tensor<256xf32>, tensor<64x98xi8>, tensor<1168xi8>, tensor<64xf32>) + check.expect.close actual(%output) expected(%output) atol(0.0) rtol(0.0) nan(same) : tensor<64xf32> + check.return +} + +check.benchmark<@ggml_mul_mat_vec_iq3xxs_f32_coverage> @ggml_mul_mat_vec_iq3xxs_f32_benchmark diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/manifest.json b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/manifest.json new file mode 100644 index 000000000000..15030b241b02 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/manifest.json @@ -0,0 +1,8138 @@ +{ + "schema": "ggml-hrx-kernel-corpus-v1", + "upstream_revision": "local", + "files": [ + { + "path": "ops/unary_f32.loom" + }, + { + "path": "ops/scale_f32.loom" + }, + { + "path": "ops/scale_bias_f32.loom" + }, + { + "path": "ops/copy_f32.loom" + }, + { + "path": "ops/binary_f32.loom" + }, + { + "path": "ops/binary_bc_f32.loom" + }, + { + "path": "ops/rmsnorm_f32.loom" + }, + { + "path": "ops/rmsnorm_binary_f32.loom" + }, + { + "path": "ops/rmsnorm_binary_q8_1_x4.loom" + }, + { + "path": "ops/mul_mat_f32_f32_wmma.loom" + }, + { + "path": "ops/llm_attention_q_matmul_rope_f32_f32_wmma.loom" + }, + { + "path": "ops/llm_attention_k_matmul_rope_set_rows_f32_f32_wmma.loom" + }, + { + "path": "ops/llm_attention_v_matmul_set_rows_f32_f32_wmma.loom" + }, + { + "path": "ops/llm_attention_q_matmul_rope_decode_f32_f32.loom" + }, + { + "path": "ops/llm_attention_k_matmul_rope_set_rows_decode_f32_f32.loom" + }, + { + "path": "ops/llm_attention_v_matmul_set_rows_decode_f32_f32.loom" + }, + { + "path": "ops/mul_mat_id_f32_f32_wmma.loom" + }, + { + "path": "ops/mul_mat_id_f16_f16_wmma.loom" + }, + { + "path": "ops/mul_mat_id_swiglu_f32_f32_wmma.loom" + }, + { + "path": "ops/mul_mat_id_swiglu_f16_f16_wmma.loom" + }, + { + "path": "motifs/mul_mat_id_swiglu_f16_f16_wmma_body.loom" + }, + { + "path": "motifs/mul_mat_id_swiglu_f16_f16_accumulate.loom" + }, + { + "path": "ops/mul_mat_id_postops_f32_f32_wmma.loom" + }, + { + "path": "ops/mul_mat_id_postops_next_rmsnorm_f32_f32_wmma.loom" + }, + { + "path": "ops/moe_routing_tables.loom" + }, + { + "path": "ops/mul_mat_bias_f32_f32_wmma.loom" + }, + { + "path": "ops/mul_mat_add_f32_f32_wmma.loom" + }, + { + "path": "ops/mul_mat_bias_add_f32_f32_wmma.loom" + }, + { + "path": "ops/mul_mat_add_next_rmsnorm_f32_f32_wmma.loom" + }, + { + "path": "ops/mul_mat_bias_add_next_rmsnorm_f32_f32_wmma.loom" + }, + { + "path": "ops/mul_mat_swiglu_f32_f32_wmma.loom" + }, + { + "path": "ops/mul_mat_swiglu_f32_f32_decode.loom" + }, + { + "path": "ops/mul_mat_add_f32_f32_decode.loom" + }, + { + "path": "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + }, + { + "path": "ops/mul_mat_symmetric_i4_wmma.loom" + }, + { + "path": "ops/mul_mat_symmetric_i8_wmma.loom" + }, + { + "path": "ops/mul_mat_q5_k_q8_plane_wmma.loom" + }, + { + "path": "ops/mul_mat_q6_k_packed_f16_wmma.loom" + }, + { + "path": "ops/mul_mat_dual_q4_f32_decode.loom" + }, + { + "path": "ops/mul_mat_f32_f32_decode.loom" + }, + { + "path": "ops/get_rows_f32.loom" + }, + { + "path": "ops/gated_delta_net_f32_wmma.loom" + }, + { + "path": "ops/ssm_conv_f32.loom" + }, + { + "path": "ops/ssm_conv_generic_f32.loom" + }, + { + "path": "ops/rope_f32.loom" + }, + { + "path": "ops/set_rows.loom" + }, + { + "path": "ops/rope_set_rows_f32.loom" + }, + { + "path": "ops/flash_attention_f32_f16_wmma.loom" + }, + { + "path": "ops/flash_attention_decode_split_f32_f16_wmma.loom" + }, + { + "path": "motifs/unary_f32_apply.loom" + }, + { + "path": "motifs/binary_f32_apply.loom" + }, + { + "path": "motifs/rmsnorm_f32.loom" + }, + { + "path": "motifs/quantize_q8_1_x4.loom" + }, + { + "path": "motifs/publish_f32.loom" + }, + { + "path": "motifs/rope_f32.loom" + }, + { + "path": "motifs/dequant.loom" + }, + { + "path": "motifs/dequant_prism.loom" + }, + { + "path": "motifs/dequant_1bit.loom" + }, + { + "path": "motifs/mul_mat_f32_f32_wmma_core.loom" + }, + { + "path": "motifs/mul_mat_quantized_f16_prefill.loom" + }, + { + "path": "motifs/mul_mat_id_f32_f32_wmma_core.loom" + }, + { + "path": "motifs/mul_mat_id_f16_f16_wmma_core.loom" + }, + { + "path": "motifs/mul_mat_id_q4k_f16_wmma_projection.loom" + }, + { + "path": "motifs/mul_mat_id_q5k_f16_wmma_projection.loom" + }, + { + "path": "motifs/mul_mat_id_q6k_f16_wmma_projection.loom" + }, + { + "path": "motifs/mul_mat_id_f32_f32_postops.loom" + }, + { + "path": "motifs/mul_mat_f32_f32_postops.loom" + }, + { + "path": "motifs/llm_attention_qkv_matmul_postops.loom" + }, + { + "path": "motifs/q4_k_f16.loom" + }, + { + "path": "motifs/q6_k_f16.loom" + }, + { + "path": "motifs/q8_0_f16.loom" + }, + { + "path": "motifs/q8_1_f16.loom" + }, + { + "path": "motifs/f16_f16.loom" + }, + { + "path": "motifs/bf16_f16.loom" + }, + { + "path": "motifs/f32_f16.loom" + }, + { + "path": "ops/grouped_mul_mat_f16_f32.loom" + }, + { + "path": "ops/small_rows_f32.loom" + }, + { + "path": "ops/mul_mat_id_decode_f32.loom" + }, + { + "path": "ops/res_scale_pair_f32.loom" + }, + { + "path": "ops/zaya_cca_conv_decode_f32.loom" + }, + { + "path": "ops/zaya_cca_qk_norm_decode_f32.loom" + }, + { + "path": "ops/kquant_decode_f32.loom" + }, + { + "path": "ops/softplus_f32.loom" + }, + { + "path": "ops/hadamard_f32.loom" + }, + { + "path": "ops/add_id_f32.loom" + }, + { + "path": "ops/swiglu_oai_f32.loom" + }, + { + "path": "ops/attention_sink_f32.loom" + } + ], + "exports": [ + { + "name": "ggml_copy_transpose_f16", + "family": "loom_libs", + "symbol": "ggml_copy_transpose_f16", + "source": "ops/copy_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/copy_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_copy_f16_k16_major", + "family": "loom_libs", + "symbol": "ggml_copy_f16_k16_major", + "source": "ops/copy_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/copy_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_copy_f32_f16", + "family": "loom_libs", + "symbol": "ggml_copy_f32_f16", + "source": "ops/copy_f32.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "bindings": [ + "source", + "output" + ], + "binding_access": [ + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/copy_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_copy_f32", + "family": "loom_libs", + "symbol": "ggml_copy_f32", + "source": "ops/copy_f32.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "bindings": [ + "source", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/copy_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_copy_strided_source_f32", + "family": "loom_libs", + "symbol": "ggml_copy_strided_source_f32", + "source": "ops/copy_f32.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + }, + { + "name": "source_span", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + }, + { + "name": "source_span", + "type": "index" + } + ], + "bindings": [ + "source", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/copy_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_concat_dim0_f32", + "family": "loom_libs", + "symbol": "ggml_concat_dim0_f32", + "source": "ops/copy_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "lhs", + "rhs", + "output" + ], + "binding_access": [ + "read", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/copy_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_concat_dim0_strided_source_f32", + "family": "loom_libs", + "symbol": "ggml_concat_dim0_strided_source_f32", + "source": "ops/copy_f32.loom", + "workload_parameters": [ + { + "name": "lhs_span", + "type": "index" + }, + { + "name": "rhs_span", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "lhs_span", + "type": "index" + }, + { + "name": "rhs_span", + "type": "index" + } + ], + "bindings": [ + "lhs", + "rhs", + "output" + ], + "binding_access": [ + "read", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/copy_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_scale_bias_f32", + "family": "loom_libs", + "symbol": "ggml_scale_bias_f32", + "source": "ops/scale_bias_f32.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/scale_bias_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_unary_f32", + "family": "loom_libs", + "symbol": "ggml_unary_f32", + "source": "ops/unary_f32.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/unary_f32.loom" + ], + "library_sources": [ + "motifs/unary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/unary_f32_apply.loom" + ] + }, + { + "name": "ggml_scale_f32", + "family": "loom_libs", + "symbol": "ggml_scale_f32", + "source": "ops/scale_f32.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/scale_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_binary_f32", + "family": "loom_libs", + "symbol": "ggml_binary_f32", + "source": "ops/binary_f32.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "bindings": [ + "lhs", + "rhs", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/binary_f32.loom" + ], + "library_sources": [ + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_binary_swiglu_symmetric_i4_k32", + "family": "loom_libs", + "symbol": "ggml_binary_swiglu_symmetric_i4_k32", + "source": "ops/binary_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "lhs", + "rhs", + "output", + "quantized_values", + "scales", + "sums" + ], + "binding_access": [ + "read", + "read", + "write", + "write", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/binary_f32.loom" + ], + "library_sources": [ + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_binary_bc_f32", + "family": "loom_libs", + "symbol": "ggml_binary_bc_f32", + "source": "ops/binary_bc_f32.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + }, + { + "name": "ne0", + "type": "index" + }, + { + "name": "ne1", + "type": "index" + }, + { + "name": "ne2", + "type": "index" + }, + { + "name": "ne3", + "type": "index" + }, + { + "name": "src0_element_count", + "type": "index" + }, + { + "name": "src1_element_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + }, + { + "name": "ne0", + "type": "index" + }, + { + "name": "ne1", + "type": "index" + }, + { + "name": "ne2", + "type": "index" + }, + { + "name": "ne3", + "type": "index" + }, + { + "name": "src0_element_count", + "type": "index" + }, + { + "name": "src1_element_count", + "type": "index" + } + ], + "bindings": [ + "lhs", + "rhs", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/binary_bc_f32.loom" + ], + "library_sources": [ + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_rmsnorm_f32", + "family": "loom_libs", + "symbol": "ggml_rmsnorm_f32", + "source": "ops/rmsnorm_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom" + ] + }, + { + "name": "ggml_rmsnorm_mul_rope_f32", + "family": "loom_libs", + "symbol": "ggml_rmsnorm_mul_rope_f32", + "source": "ops/rmsnorm_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "input", + "weight", + "positions", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom" + ] + }, + { + "name": "ggml_rmsnorm_binary_f32", + "family": "loom_libs", + "symbol": "ggml_rmsnorm_binary_f32", + "source": "ops/rmsnorm_binary_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "rhs", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_rmsnorm_binary_f32_k16", + "family": "loom_libs", + "symbol": "ggml_rmsnorm_binary_f32_k16", + "source": "ops/rmsnorm_binary_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "rhs", + "output", + "f16_output" + ], + "binding_access": [ + "read", + "read", + "read_write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_rmsnorm_binary_q8_1_x4", + "family": "loom_libs", + "symbol": "ggml_rmsnorm_binary_q8_1_x4", + "source": "ops/rmsnorm_binary_q8_1_x4.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "rhs", + "output", + "q8_output" + ], + "binding_access": [ + "read", + "read", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/rmsnorm_binary_q8_1_x4.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_rmsnorm_binary_symmetric_i4_k32", + "family": "loom_libs", + "symbol": "ggml_rmsnorm_binary_symmetric_i4_k32", + "source": "ops/rmsnorm_binary_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output", + "quantized_values", + "scales", + "sums" + ], + "binding_access": [ + "read", + "read", + "write", + "write", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_add_rmsnorm_binary_symmetric_i4_k32", + "family": "loom_libs", + "symbol": "ggml_add_rmsnorm_binary_symmetric_i4_k32", + "source": "ops/rmsnorm_binary_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "lhs", + "rhs", + "residual_output", + "weight", + "output", + "quantized_values", + "scales", + "sums" + ], + "binding_access": [ + "read", + "read", + "write", + "read", + "write", + "write", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "family": "loom_libs", + "symbol": "ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "source": "ops/rmsnorm_binary_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "raw_gate", + "output", + "quantized_values", + "scales", + "sums" + ], + "binding_access": [ + "read", + "read", + "read", + "write", + "write", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_rmsnorm_gate_f32_f16", + "family": "loom_libs", + "symbol": "ggml_rmsnorm_gate_f32_f16", + "source": "ops/rmsnorm_binary_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "raw_gate", + "output", + "f16_output" + ], + "binding_access": [ + "read", + "read", + "read", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_rmsnorm_gate_f32_q8_1_x4", + "family": "loom_libs", + "symbol": "ggml_rmsnorm_gate_f32_q8_1_x4", + "source": "ops/rmsnorm_binary_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "raw_gate", + "output", + "q8_output" + ], + "binding_access": [ + "read", + "read", + "read", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_mul_mat_quantized_f16_wmma_prefill_conv4_interior", + "family": "loom_libs", + "symbol": "ggml_mul_mat_quantized_f16_wmma_prefill_conv4_interior", + "source": "ops/mul_mat_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "filter", + "output", + "edges" + ], + "binding_access": [ + "read", + "read", + "read", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_mul_mat_quantized_f16_wmma_prefill_conv4", + "family": "loom_libs", + "symbol": "ggml_mul_mat_quantized_f16_wmma_prefill_conv4", + "source": "ops/mul_mat_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "state", + "filter", + "output", + "cache" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_mul_mat_q4_k_f16_wmma_prefill_wave32", + "family": "loom_libs", + "symbol": "ggml_mul_mat_q4_k_f16_wmma_prefill_wave32", + "source": "ops/mul_mat_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_mul_mat_f32_f32_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_f32_f32_wmma", + "source": "ops/mul_mat_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_mul_mat_f32_f32_narrow_split_k4", + "family": "loom_libs", + "symbol": "ggml_mul_mat_f32_f32_narrow_split_k4", + "source": "ops/mul_mat_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "partial", + "output", + "completion_counters" + ], + "binding_access": [ + "read", + "read", + "read_write", + "write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "llm_attention_q_matmul_rope_f32_f32_wmma", + "family": "loom_libs", + "symbol": "llm_attention_q_matmul_rope_f32_f32_wmma", + "source": "ops/llm_attention_q_matmul_rope_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "positions", + "theta", + "freq_factors", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/llm_attention_q_matmul_rope_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/rope_f32.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/rope_f32.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "llm_attention_k_matmul_rope_set_rows_f32_f32_wmma", + "family": "loom_libs", + "symbol": "llm_attention_k_matmul_rope_set_rows_f32_f32_wmma", + "source": "ops/llm_attention_k_matmul_rope_set_rows_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "positions", + "indices", + "theta", + "freq_factors", + "cache" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/llm_attention_k_matmul_rope_set_rows_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/rope_f32.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/rope_f32.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "llm_attention_v_matmul_set_rows_f32_f32_wmma", + "family": "loom_libs", + "symbol": "llm_attention_v_matmul_set_rows_f32_f32_wmma", + "source": "ops/llm_attention_v_matmul_set_rows_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "indices", + "cache" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/llm_attention_v_matmul_set_rows_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/llm_attention_qkv_matmul_postops.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "llm_attention_q_matmul_rope_decode_f32_f32", + "family": "loom_libs", + "symbol": "llm_attention_q_matmul_rope_decode_f32_f32", + "source": "ops/llm_attention_q_matmul_rope_decode_f32_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "positions", + "theta", + "freq_factors", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/llm_attention_q_matmul_rope_decode_f32_f32.loom" + ], + "library_sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/rope_f32.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/rope_f32.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "llm_attention_k_matmul_rope_set_rows_decode_f32_f32", + "family": "loom_libs", + "symbol": "llm_attention_k_matmul_rope_set_rows_decode_f32_f32", + "source": "ops/llm_attention_k_matmul_rope_set_rows_decode_f32_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "positions", + "indices", + "theta", + "freq_factors", + "cache" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/llm_attention_k_matmul_rope_set_rows_decode_f32_f32.loom" + ], + "library_sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/rope_f32.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/rope_f32.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "llm_attention_v_matmul_set_rows_decode_f32_f32", + "family": "loom_libs", + "symbol": "llm_attention_v_matmul_set_rows_decode_f32_f32", + "source": "ops/llm_attention_v_matmul_set_rows_decode_f32_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "indices", + "cache" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/llm_attention_v_matmul_set_rows_decode_f32_f32.loom" + ], + "library_sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_mul_mat_id_f32_f32_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_id_f32_f32_wmma", + "source": "ops/mul_mat_id_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "expert_table", + "partition_table", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_id_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_id_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_id_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_mul_mat_id_f16_f16_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_id_f16_f16_wmma", + "source": "ops/mul_mat_id_f16_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "expert_table", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_id_f16_f16_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_id_f16_f16_wmma_core.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_id_f16_f16_wmma_core.loom" + ] + }, + { + "name": "ggml_mul_mat_id_swiglu_f32_f32_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_id_swiglu_f32_f32_wmma", + "source": "ops/mul_mat_id_swiglu_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "expert_table", + "partition_table", + "gate_weight", + "up_weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_id_swiglu_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_mul_mat_id_swiglu_f16_f16_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_id_swiglu_f16_f16_wmma", + "source": "ops/mul_mat_id_swiglu_f16_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "expert_table", + "partition_table", + "gate_weight", + "up_weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_id_swiglu_f16_f16_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_id_swiglu_f16_f16_wmma_body.loom", + "motifs/mul_mat_id_swiglu_f16_f16_accumulate.loom", + "motifs/mul_mat_id_q4k_f16_wmma_projection.loom", + "motifs/mul_mat_id_q5k_f16_wmma_projection.loom", + "motifs/mul_mat_id_q6k_f16_wmma_projection.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_id_swiglu_f16_f16_wmma_body.loom", + "motifs/mul_mat_id_swiglu_f16_f16_accumulate.loom", + "motifs/mul_mat_id_q4k_f16_wmma_projection.loom", + "motifs/mul_mat_id_q5k_f16_wmma_projection.loom", + "motifs/mul_mat_id_q6k_f16_wmma_projection.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_mul_mat_id_postops_f32_f32_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_id_postops_f32_f32_wmma", + "source": "ops/mul_mat_id_postops_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "expert_table", + "partition_table", + "weight", + "bias", + "residual_input", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_id_postops_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_id_f32_f32_wmma_core.loom", + "motifs/mul_mat_id_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_id_f32_f32_wmma_core.loom", + "motifs/mul_mat_id_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_mul_mat_id_postops_next_rmsnorm_f32_f32_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_id_postops_next_rmsnorm_f32_f32_wmma", + "source": "ops/mul_mat_id_postops_next_rmsnorm_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "expert_table", + "partition_table", + "weight", + "bias", + "residual_input", + "residual_output", + "norm_weight", + "normalized_output", + "completion_counters" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "read_write", + "read", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_id_postops_next_rmsnorm_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_id_f32_f32_wmma_core.loom", + "motifs/mul_mat_id_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_id_f32_f32_wmma_core.loom", + "motifs/mul_mat_id_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_moe_build_expert_table", + "family": "loom_libs", + "symbol": "ggml_moe_build_expert_table", + "source": "ops/moe_routing_tables.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ], + "bindings": [ + "route_ids", + "expert_table" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/moe_routing_tables.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_moe_build_expert_partition_table", + "family": "loom_libs", + "symbol": "ggml_moe_build_expert_partition_table", + "source": "ops/moe_routing_tables.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ], + "bindings": [ + "expert_table", + "partition_table" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/moe_routing_tables.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_add_next_rmsnorm_f32_f32_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_add_next_rmsnorm_f32_f32_wmma", + "source": "ops/mul_mat_add_next_rmsnorm_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "residual_input", + "residual_output", + "norm_weight", + "normalized_output", + "completion_counters" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write", + "read", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_add_next_rmsnorm_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_mul_mat_bias_f32_f32_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_bias_f32_f32_wmma", + "source": "ops/mul_mat_bias_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "bias", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_bias_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_mul_mat_add_f32_f32_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_add_f32_f32_wmma", + "source": "ops/mul_mat_add_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "residual_input", + "residual_output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_add_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_mul_mat_bias_add_f32_f32_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_bias_add_f32_f32_wmma", + "source": "ops/mul_mat_bias_add_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "bias", + "residual_input", + "residual_output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_bias_add_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_mul_mat_bias_add_next_rmsnorm_f32_f32_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_bias_add_next_rmsnorm_f32_f32_wmma", + "source": "ops/mul_mat_bias_add_next_rmsnorm_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "bias", + "residual_input", + "residual_output", + "norm_weight", + "normalized_output", + "completion_counters" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write", + "read", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_bias_add_next_rmsnorm_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/mul_mat_f32_f32_postops.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_mul_mat_dual_q4_f32_decode", + "family": "loom_libs", + "symbol": "ggml_mul_mat_dual_q4_f32_decode", + "source": "ops/mul_mat_dual_q4_f32_decode.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "input", + "first_weight", + "second_weight", + "first_output", + "second_output" + ], + "binding_access": [ + "read", + "read", + "read", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_dual_q4_f32_decode.loom" + ], + "library_sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_mul_mat_f32_f32_decode_wave64", + "family": "loom_libs", + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "source": "ops/mul_mat_f32_f32_decode.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "library_sources": [ + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_mul_mat_add_f32_f32_decode_wave64", + "family": "loom_libs", + "symbol": "ggml_mul_mat_add_f32_f32_decode_wave64", + "source": "ops/mul_mat_add_f32_f32_decode.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "residual_input", + "residual_output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_add_f32_f32_decode.loom" + ], + "library_sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_mul_mat_swiglu_f32_f32_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_swiglu_f32_f32_wmma", + "source": "ops/mul_mat_swiglu_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "gate_weight", + "up_weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_swiglu_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_mul_mat_swiglu_q4_k_f16_wmma_prefill_wave32", + "family": "loom_libs", + "symbol": "ggml_mul_mat_swiglu_q4_k_f16_wmma_prefill_wave32", + "source": "ops/mul_mat_swiglu_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "up_weight", + "output", + "f16_output" + ], + "binding_access": [ + "read", + "read", + "read", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_swiglu_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_mul_mat_swiglu_q4_q8_1_x4_lowtoken_dot", + "family": "loom_libs", + "symbol": "ggml_mul_mat_swiglu_q4_q8_1_x4_lowtoken_dot", + "source": "ops/mul_mat_swiglu_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "gate_weight", + "up_weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_swiglu_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_mul_mat_swiglu_q4_q8_1_x4_lowtoken_dot_q8_output", + "family": "loom_libs", + "symbol": "ggml_mul_mat_swiglu_q4_q8_1_x4_lowtoken_dot_q8_output", + "source": "ops/mul_mat_swiglu_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "gate_weight", + "up_weight", + "output", + "q8_output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_swiglu_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_mul_mat_swiglu_f32_f32_lowtoken_dot", + "family": "loom_libs", + "symbol": "ggml_mul_mat_swiglu_f32_f32_lowtoken_dot", + "source": "ops/mul_mat_swiglu_f32_f32_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "gate_weight", + "up_weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_swiglu_f32_f32_wmma.loom" + ], + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_mul_mat_swiglu_f32_f32_decode_wave64", + "family": "loom_libs", + "symbol": "ggml_mul_mat_swiglu_f32_f32_decode_wave64", + "source": "ops/mul_mat_swiglu_f32_f32_decode.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "gate_weight", + "up_weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_swiglu_f32_f32_decode.loom" + ], + "library_sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/binary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/binary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_quantize_f32_symmetric_i4_k32", + "family": "loom_libs", + "symbol": "ggml_quantize_f32_symmetric_i4_k32", + "source": "ops/mul_mat_swiglu_symmetric_i4_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "input", + "quantized_values", + "scales", + "sums" + ], + "binding_access": [ + "read", + "read_write", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_swiglu_symmetric_i4_wmma_q8_plane", + "family": "loom_libs", + "symbol": "ggml_mul_mat_swiglu_symmetric_i4_wmma_q8_plane", + "source": "ops/mul_mat_swiglu_symmetric_i4_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "gate_weight", + "up_weight", + "input", + "output", + "quantized_values", + "scales", + "sums", + "q8_output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_quantize_f32_symmetric_i4_k64_plane", + "family": "loom_libs", + "symbol": "ggml_quantize_f32_symmetric_i4_k64_plane", + "source": "ops/mul_mat_symmetric_i4_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "input", + "quantized_values", + "scales", + "sums" + ], + "binding_access": [ + "read", + "write", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_symmetric_i4_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_symmetric_i4_wmma", + "source": "ops/mul_mat_symmetric_i4_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "weight", + "input", + "output", + "quantized_values", + "scales", + "sums" + ], + "binding_access": [ + "read", + "read", + "write", + "read", + "read", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "source": "ops/mul_mat_symmetric_i4_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "weight", + "input", + "output", + "quantized_values", + "scales", + "sums" + ], + "binding_access": [ + "read", + "read", + "write", + "read", + "read", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "source": "ops/mul_mat_symmetric_i4_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "first_weight", + "second_weight", + "first_output", + "second_output", + "quantized_values", + "scales", + "sums" + ], + "binding_access": [ + "read", + "read", + "write", + "write", + "read", + "read", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c1", + "family": "loom_libs", + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c1", + "source": "ops/mul_mat_symmetric_i4_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "first_weight", + "second_weight", + "first_output", + "second_output", + "quantized_values", + "scales", + "sums" + ], + "binding_access": [ + "read", + "read", + "write", + "write", + "read", + "read", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c2", + "family": "loom_libs", + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c2", + "source": "ops/mul_mat_symmetric_i4_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "first_weight", + "second_weight", + "first_output", + "second_output", + "quantized_values", + "scales", + "sums" + ], + "binding_access": [ + "read", + "read", + "write", + "write", + "read", + "read", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c3", + "family": "loom_libs", + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c3", + "source": "ops/mul_mat_symmetric_i4_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "first_weight", + "second_weight", + "first_output", + "second_output", + "quantized_values", + "scales", + "sums" + ], + "binding_access": [ + "read", + "read", + "write", + "write", + "read", + "read", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c4", + "family": "loom_libs", + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c4", + "source": "ops/mul_mat_symmetric_i4_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "first_weight", + "second_weight", + "first_output", + "second_output", + "quantized_values", + "scales", + "sums" + ], + "binding_access": [ + "read", + "read", + "write", + "write", + "read", + "read", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c5", + "family": "loom_libs", + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c5", + "source": "ops/mul_mat_symmetric_i4_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "first_weight", + "second_weight", + "first_output", + "second_output", + "quantized_values", + "scales", + "sums" + ], + "binding_access": [ + "read", + "read", + "write", + "write", + "read", + "read", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c1", + "family": "loom_libs", + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c1", + "source": "ops/mul_mat_symmetric_i4_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "weight", + "input", + "output", + "quantized_values", + "scales", + "sums", + "partial", + "completion_counters" + ], + "binding_access": [ + "read", + "read", + "read_write", + "read", + "read", + "read", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c2", + "family": "loom_libs", + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c2", + "source": "ops/mul_mat_symmetric_i4_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "weight", + "input", + "output", + "quantized_values", + "scales", + "sums", + "partial", + "completion_counters" + ], + "binding_access": [ + "read", + "read", + "read_write", + "read", + "read", + "read", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c3", + "family": "loom_libs", + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c3", + "source": "ops/mul_mat_symmetric_i4_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "weight", + "input", + "output", + "quantized_values", + "scales", + "sums", + "partial", + "completion_counters" + ], + "binding_access": [ + "read", + "read", + "read_write", + "read", + "read", + "read", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c4", + "family": "loom_libs", + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c4", + "source": "ops/mul_mat_symmetric_i4_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "weight", + "input", + "output", + "quantized_values", + "scales", + "sums", + "partial", + "completion_counters" + ], + "binding_access": [ + "read", + "read", + "read_write", + "read", + "read", + "read", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c5", + "family": "loom_libs", + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c5", + "source": "ops/mul_mat_symmetric_i4_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "weight", + "input", + "output", + "quantized_values", + "scales", + "sums", + "partial", + "completion_counters" + ], + "binding_access": [ + "read", + "read", + "read_write", + "read", + "read", + "read", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "source": "ops/mul_mat_symmetric_i4_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "weight", + "input", + "output", + "quantized_values", + "scales", + "sums", + "partial", + "completion_counters" + ], + "binding_access": [ + "read", + "read", + "read_write", + "read", + "read", + "read", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_quantize_f32_symmetric_i8_k256", + "family": "loom_libs", + "symbol": "ggml_quantize_f32_symmetric_i8_k256", + "source": "ops/mul_mat_symmetric_i8_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size_arg", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size_arg", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_symmetric_i8_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_q5_k_symmetric_i8_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_q5_k_symmetric_i8_wmma", + "source": "ops/mul_mat_symmetric_i8_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_symmetric_i8_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_q5_k_q8_plane_wmmai8_token256", + "family": "loom_libs", + "symbol": "ggml_mul_mat_q5_k_q8_plane_wmmai8_token256", + "source": "ops/mul_mat_q5_k_q8_plane_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_q5_k_q8_plane_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_q5_k_iq4_xs_q8_1_x4_wmma_token256", + "family": "loom_libs", + "symbol": "ggml_mul_mat_q5_k_iq4_xs_q8_1_x4_wmma_token256", + "source": "ops/mul_mat_q5_k_q8_plane_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_q5_k_q8_plane_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_q4_k_q8_1_x4_swiglu_wmma_token256", + "family": "loom_libs", + "symbol": "ggml_mul_mat_q4_k_q8_1_x4_swiglu_wmma_token256", + "source": "ops/mul_mat_q5_k_q8_plane_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "weight", + "up_weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_q5_k_q8_plane_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_q6_k_f32_wmma_prefill_wave32", + "family": "loom_libs", + "symbol": "ggml_mul_mat_q6_k_f32_wmma_prefill_wave32", + "source": "ops/mul_mat_q6_k_packed_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "library_sources": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_mul_mat_q6_k_f16_wmma_prefill_wave32", + "family": "loom_libs", + "symbol": "ggml_mul_mat_q6_k_f16_wmma_prefill_wave32", + "source": "ops/mul_mat_q6_k_packed_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "library_sources": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "source": "ops/mul_mat_q6_k_packed_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "library_sources": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "family": "loom_libs", + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "source": "ops/mul_mat_q6_k_packed_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "library_sources": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_select_symmetric_i4_k32_groups", + "family": "loom_libs", + "symbol": "ggml_select_symmetric_i4_k32_groups", + "source": "ops/mul_mat_q6_k_packed_f16_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "input", + "selected_groups" + ], + "binding_access": [ + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "library_sources": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_mul_mat_q6_k_symmetric_i2_scan_token1", + "family": "loom_libs", + "symbol": "ggml_mul_mat_q6_k_symmetric_i2_scan_token1", + "source": "ops/mul_mat_q6_k_packed_f16_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "weight", + "output", + "qact", + "scales", + "selected_mask" + ], + "binding_access": [ + "read", + "write", + "read", + "read", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "library_sources": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_top_k8_f32_partitions_register", + "family": "loom_libs", + "symbol": "ggml_top_k8_f32_partitions_register", + "source": "ops/mul_mat_q6_k_packed_f16_wmma.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "bindings": [ + "values", + "partial_values", + "partial_ids" + ], + "binding_access": [ + "read", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "library_sources": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_top_k128_f32_reduce_gather_register", + "family": "loom_libs", + "symbol": "ggml_top_k128_f32_reduce_gather_register", + "source": "ops/mul_mat_q6_k_packed_f16_wmma.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "bindings": [ + "partial_values", + "partial_ids", + "candidate_output", + "value_output" + ], + "binding_access": [ + "read", + "read", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "library_sources": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_fill_negative_f32", + "family": "loom_libs", + "symbol": "ggml_fill_negative_f32", + "source": "ops/mul_mat_q6_k_packed_f16_wmma.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "bindings": [ + "output" + ], + "binding_access": [ + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "library_sources": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_mul_mat_q6_k_packed_selected_refine_token1", + "family": "loom_libs", + "symbol": "ggml_mul_mat_q6_k_packed_selected_refine_token1", + "source": "ops/mul_mat_q6_k_packed_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "candidate_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "candidate_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "candidates", + "exact_output", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "library_sources": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/q6_k_f16.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/mul_mat_quantized_f16_prefill.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "ggml_flash_attention_f32_f16_wmma", + "family": "loom_libs", + "symbol": "ggml_flash_attention_f32_f16_wmma", + "source": "ops/flash_attention_f32_f16_wmma.loom", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ], + "bindings": [ + "query", + "key", + "value", + "mask", + "gate", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_flash_attention_decode_split_f32_f16_wmma_next_q8", + "family": "loom_libs", + "symbol": "ggml_flash_attention_decode_split_f32_f16_wmma_next_q8", + "source": "ops/flash_attention_decode_split_f32_f16_wmma.loom", + "workload_parameters": [ + { + "name": "key_value_token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "key_value_token_count", + "type": "index" + } + ], + "bindings": [ + "query", + "key", + "value", + "mask", + "partial_max", + "partial_sum", + "partial_output", + "completion_counter", + "output", + "next_q8_output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write", + "read_write", + "read_write", + "read_write", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/flash_attention_decode_split_f32_f16_wmma.loom" + ], + "library_sources": [ + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_get_rows_f32", + "family": "loom_libs", + "symbol": "ggml_get_rows_f32", + "source": "ops/get_rows_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ], + "bindings": [ + "token_ids", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "library_sources": [ + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_get_rows_f32_next", + "family": "loom_libs", + "symbol": "ggml_get_rows_f32_next", + "source": "ops/get_rows_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ], + "bindings": [ + "token_ids", + "weight", + "output", + "next_output" + ], + "binding_access": [ + "read", + "read", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "library_sources": [ + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_rope_f32", + "family": "loom_libs", + "symbol": "ggml_rope_f32", + "source": "ops/rope_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "positions", + "input", + "theta", + "freq_factors", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/rope_f32.loom" + ], + "library_sources": [ + "motifs/rope_f32.loom" + ] + }, + "compile_dependencies": [ + "motifs/rope_f32.loom" + ] + }, + { + "name": "ggml_set_rows", + "family": "loom_libs", + "symbol": "ggml_set_rows", + "source": "ops/set_rows.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ], + "bindings": [ + "rows", + "indices", + "cache" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/set_rows.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_set_rows_scatter", + "family": "loom_libs", + "symbol": "ggml_set_rows_scatter", + "source": "ops/set_rows.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + } + ], + "bindings": [ + "rows", + "indices", + "cache" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/set_rows.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_rope_set_rows_f32", + "family": "loom_libs", + "symbol": "ggml_rope_set_rows_f32", + "source": "ops/rope_set_rows_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + } + ], + "bindings": [ + "positions", + "indices", + "input", + "theta", + "freq_factors", + "cache" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/rope_set_rows_f32.loom" + ], + "library_sources": [ + "motifs/rope_f32.loom" + ] + }, + "compile_dependencies": [ + "motifs/rope_f32.loom" + ] + }, + { + "name": "llm_ssm_conv_f32", + "family": "loom_libs", + "symbol": "llm_ssm_conv_f32", + "source": "ops/ssm_conv_generic_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "window", + "filter", + "output" + ], + "binding_access": [ + "read", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/ssm_conv_generic_f32.loom" + ], + "library_sources": [ + "motifs/unary_f32_apply.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/unary_f32_apply.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "llm_ssm_conv_binary_f32", + "family": "loom_libs", + "symbol": "llm_ssm_conv_binary_f32", + "source": "ops/ssm_conv_generic_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "window", + "filter", + "operand", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/ssm_conv_generic_f32.loom" + ], + "library_sources": [ + "motifs/unary_f32_apply.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/unary_f32_apply.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "llm_ssm_conv_snapshot_window_tail_f32", + "family": "loom_libs", + "symbol": "llm_ssm_conv_snapshot_window_tail_f32", + "source": "ops/ssm_conv_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "state", + "x", + "snapshot", + "cache" + ], + "binding_access": [ + "read", + "read", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/ssm_conv_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "llm_ssm_conv_dconv4_silu_prefill_finish_f32", + "family": "loom_libs", + "symbol": "llm_ssm_conv_dconv4_silu_prefill_finish_f32", + "source": "ops/ssm_conv_generic_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "state", + "filter", + "edges", + "output", + "cache" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/ssm_conv_generic_f32.loom" + ], + "library_sources": [ + "motifs/unary_f32_apply.loom", + "motifs/binary_f32_apply.loom" + ] + }, + "compile_dependencies": [ + "motifs/unary_f32_apply.loom", + "motifs/binary_f32_apply.loom" + ] + }, + { + "name": "llm_ssm_conv_dconv4_silu_prefill_512_wg1024", + "family": "loom_libs", + "symbol": "llm_ssm_conv_dconv4_silu_prefill_512_wg1024", + "source": "ops/ssm_conv_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "state_snapshot", + "x", + "filter", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/ssm_conv_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "llm_ssm_conv_dconv4_silu_decode_f32", + "family": "loom_libs", + "symbol": "llm_ssm_conv_dconv4_silu_decode_f32", + "source": "ops/ssm_conv_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "state", + "x", + "filter", + "output", + "cache" + ], + "binding_access": [ + "read", + "read", + "read", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/ssm_conv_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "llm_ssm_conv_dconv4_silu_rollback_f32", + "family": "loom_libs", + "symbol": "llm_ssm_conv_dconv4_silu_rollback_f32", + "source": "ops/ssm_conv_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "state", + "x", + "filter", + "output", + "cache0", + "cache1", + "cache2", + "cache3", + "cache4" + ], + "binding_access": [ + "read", + "read", + "read", + "write", + "write", + "write", + "write", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/ssm_conv_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "llm_gated_delta_net_f32_wmma_head128", + "family": "loom_libs", + "symbol": "llm_gated_delta_net_f32_wmma_head128", + "source": "ops/gated_delta_net_f32_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "q", + "k", + "v", + "g", + "beta", + "state_in", + "dst" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "llm_gated_delta_net_f32_wmma_head128_rmsnorm_gate", + "family": "loom_libs", + "symbol": "llm_gated_delta_net_f32_wmma_head128_rmsnorm_gate", + "source": "ops/gated_delta_net_f32_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "q", + "k", + "v", + "g", + "beta", + "state_in", + "dst", + "rms_weight", + "raw_gate", + "norm_output", + "half_output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "write", + "read", + "read", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "llm_gated_delta_net_f32_wmma_head128_projection_epilogue", + "family": "loom_libs", + "symbol": "llm_gated_delta_net_f32_wmma_head128_projection_epilogue", + "source": "ops/gated_delta_net_f32_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "q", + "k", + "v", + "alpha_raw", + "beta_raw", + "bias", + "a_scale", + "state_in", + "dst" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "llm_gated_delta_net_f32_wmma_head128_inplace", + "family": "loom_libs", + "symbol": "llm_gated_delta_net_f32_wmma_head128_inplace", + "source": "ops/gated_delta_net_f32_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "q", + "k", + "v", + "g", + "beta", + "state_inout", + "dst" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read_write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "llm_gated_delta_net_f32_wmma_head128_inplace_projection_epilogue", + "family": "loom_libs", + "symbol": "llm_gated_delta_net_f32_wmma_head128_inplace_projection_epilogue", + "source": "ops/gated_delta_net_f32_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "q", + "k", + "v", + "alpha_raw", + "beta_raw", + "bias", + "a_scale", + "state_inout", + "dst" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read_write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "llm_gated_delta_net_f32_wmma_head128_snapshot", + "family": "loom_libs", + "symbol": "llm_gated_delta_net_f32_wmma_head128_snapshot", + "source": "ops/gated_delta_net_f32_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "q", + "k", + "v", + "g", + "beta", + "state_in", + "snapshot_cache", + "dst" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "llm_gated_delta_net_f32_wmma_head128_snapshot_projection_epilogue", + "family": "loom_libs", + "symbol": "llm_gated_delta_net_f32_wmma_head128_snapshot_projection_epilogue", + "source": "ops/gated_delta_net_f32_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "q", + "k", + "v", + "alpha_raw", + "beta_raw", + "bias", + "a_scale", + "state_in", + "snapshot_cache", + "dst" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "llm_gated_delta_net_f32_wmma_head128_selected_snapshot_projection_epilogue", + "family": "loom_libs", + "symbol": "llm_gated_delta_net_f32_wmma_head128_selected_snapshot_projection_epilogue", + "source": "ops/gated_delta_net_f32_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "q", + "k", + "v", + "alpha_raw", + "beta_raw", + "bias", + "a_scale", + "state_in", + "snapshot_cache", + "dst", + "state_ids" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "write", + "write", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "llm_gated_delta_net_f32_wmma_head128_selected_snapshot_projection_rms_gate_q8", + "family": "loom_libs", + "symbol": "llm_gated_delta_net_f32_wmma_head128_selected_snapshot_projection_rms_gate_q8", + "source": "ops/gated_delta_net_f32_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "q", + "k", + "v", + "alpha_raw", + "beta_raw", + "bias", + "a_scale", + "state_in", + "snapshot_cache", + "dst", + "state_ids", + "rms_weight", + "raw_gate", + "q8_output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "write", + "write", + "read", + "read", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "llm_gated_delta_net_projection_epilogue_f32", + "family": "loom_libs", + "symbol": "llm_gated_delta_net_projection_epilogue_f32", + "source": "ops/gated_delta_net_f32_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "alpha_raw", + "beta_raw", + "bias", + "a_scale", + "gate_dst", + "beta_dst" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "motifs/rmsnorm_f32.loom", + "motifs/unary_f32_apply.loom", + "motifs/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_grouped_mul_mat_f16_f32", + "family": "loom_libs", + "symbol": "ggml_grouped_mul_mat_f16_f32", + "source": "ops/grouped_mul_mat_f16_f32.loom", + "workload_parameters": [ + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + }, + { + "name": "token_count", + "type": "index" + }, + { + "name": "group_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + }, + { + "name": "token_count", + "type": "index" + }, + { + "name": "group_count", + "type": "index" + } + ], + "bindings": [ + "weight", + "input", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/grouped_mul_mat_f16_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_softmax_rows_f32", + "family": "loom_libs", + "symbol": "ggml_softmax_rows_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "column_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "column_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_sum_rows_f32", + "family": "loom_libs", + "symbol": "ggml_sum_rows_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "column_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "column_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_argsort_rows_f32", + "family": "loom_libs", + "symbol": "ggml_argsort_rows_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "column_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "descending", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "column_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "descending", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_get_rows_small_f32", + "family": "loom_libs", + "symbol": "ggml_get_rows_small_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "width", + "type": "index" + }, + { + "name": "id_count", + "type": "index" + }, + { + "name": "batch_count", + "type": "index" + }, + { + "name": "source_rows", + "type": "index" + }, + { + "name": "id_stride", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "width", + "type": "index" + }, + { + "name": "id_count", + "type": "index" + }, + { + "name": "batch_count", + "type": "index" + }, + { + "name": "source_rows", + "type": "index" + }, + { + "name": "id_stride", + "type": "index" + } + ], + "bindings": [ + "input", + "ids", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_copy_strided_f32", + "family": "loom_libs", + "symbol": "ggml_copy_strided_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "ne0", + "type": "index" + }, + { + "name": "ne1", + "type": "index" + }, + { + "name": "ne2", + "type": "index" + }, + { + "name": "ne3", + "type": "index" + }, + { + "name": "s0", + "type": "index" + }, + { + "name": "s1", + "type": "index" + }, + { + "name": "s2", + "type": "index" + }, + { + "name": "s3", + "type": "index" + }, + { + "name": "source_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "ne0", + "type": "index" + }, + { + "name": "ne1", + "type": "index" + }, + { + "name": "ne2", + "type": "index" + }, + { + "name": "ne3", + "type": "index" + }, + { + "name": "s0", + "type": "index" + }, + { + "name": "s1", + "type": "index" + }, + { + "name": "s2", + "type": "index" + }, + { + "name": "s3", + "type": "index" + }, + { + "name": "source_extent", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_norm_rows_f32", + "family": "loom_libs", + "symbol": "ggml_norm_rows_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "column_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "column_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_binary_strided_f32", + "family": "loom_libs", + "symbol": "ggml_binary_strided_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "ne0", + "type": "index" + }, + { + "name": "ne1", + "type": "index" + }, + { + "name": "ne2", + "type": "index" + }, + { + "name": "ne3", + "type": "index" + }, + { + "name": "a0", + "type": "index" + }, + { + "name": "a1", + "type": "index" + }, + { + "name": "a2", + "type": "index" + }, + { + "name": "a3", + "type": "index" + }, + { + "name": "b0", + "type": "index" + }, + { + "name": "b1", + "type": "index" + }, + { + "name": "b2", + "type": "index" + }, + { + "name": "b3", + "type": "index" + }, + { + "name": "a_extent", + "type": "index" + }, + { + "name": "b_extent", + "type": "index" + }, + { + "name": "op", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "ne0", + "type": "index" + }, + { + "name": "ne1", + "type": "index" + }, + { + "name": "ne2", + "type": "index" + }, + { + "name": "ne3", + "type": "index" + }, + { + "name": "a0", + "type": "index" + }, + { + "name": "a1", + "type": "index" + }, + { + "name": "a2", + "type": "index" + }, + { + "name": "a3", + "type": "index" + }, + { + "name": "b0", + "type": "index" + }, + { + "name": "b1", + "type": "index" + }, + { + "name": "b2", + "type": "index" + }, + { + "name": "b3", + "type": "index" + }, + { + "name": "a_extent", + "type": "index" + }, + { + "name": "b_extent", + "type": "index" + }, + { + "name": "op", + "type": "index" + } + ], + "bindings": [ + "lhs", + "rhs", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_clamp_f32", + "family": "loom_libs", + "symbol": "ggml_clamp_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_clamp_inplace_f32", + "family": "loom_libs", + "symbol": "ggml_clamp_inplace_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "bindings": [ + "data" + ], + "binding_access": [ + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_copy_strided_f32_f16", + "family": "loom_libs", + "symbol": "ggml_copy_strided_f32_f16", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "ne0", + "type": "index" + }, + { + "name": "ne1", + "type": "index" + }, + { + "name": "ne2", + "type": "index" + }, + { + "name": "ne3", + "type": "index" + }, + { + "name": "s0", + "type": "index" + }, + { + "name": "s1", + "type": "index" + }, + { + "name": "s2", + "type": "index" + }, + { + "name": "s3", + "type": "index" + }, + { + "name": "source_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "ne0", + "type": "index" + }, + { + "name": "ne1", + "type": "index" + }, + { + "name": "ne2", + "type": "index" + }, + { + "name": "ne3", + "type": "index" + }, + { + "name": "s0", + "type": "index" + }, + { + "name": "s1", + "type": "index" + }, + { + "name": "s2", + "type": "index" + }, + { + "name": "s3", + "type": "index" + }, + { + "name": "source_extent", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_attention_strided_f32_f16", + "family": "loom_libs", + "symbol": "ggml_attention_strided_f32_f16", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "qk_size", + "type": "index" + }, + { + "name": "v_size", + "type": "index" + }, + { + "name": "q_count", + "type": "index" + }, + { + "name": "kv_count", + "type": "index" + }, + { + "name": "head_count", + "type": "index" + }, + { + "name": "kv_head_count", + "type": "index" + }, + { + "name": "q_s1", + "type": "index" + }, + { + "name": "q_s2", + "type": "index" + }, + { + "name": "k_s1", + "type": "index" + }, + { + "name": "k_s2", + "type": "index" + }, + { + "name": "v_s1", + "type": "index" + }, + { + "name": "v_s2", + "type": "index" + }, + { + "name": "m_s1", + "type": "index" + }, + { + "name": "q_extent", + "type": "index" + }, + { + "name": "k_extent", + "type": "index" + }, + { + "name": "v_extent", + "type": "index" + }, + { + "name": "m_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "qk_size", + "type": "index" + }, + { + "name": "v_size", + "type": "index" + }, + { + "name": "q_count", + "type": "index" + }, + { + "name": "kv_count", + "type": "index" + }, + { + "name": "head_count", + "type": "index" + }, + { + "name": "kv_head_count", + "type": "index" + }, + { + "name": "q_s1", + "type": "index" + }, + { + "name": "q_s2", + "type": "index" + }, + { + "name": "k_s1", + "type": "index" + }, + { + "name": "k_s2", + "type": "index" + }, + { + "name": "v_s1", + "type": "index" + }, + { + "name": "v_s2", + "type": "index" + }, + { + "name": "m_s1", + "type": "index" + }, + { + "name": "q_extent", + "type": "index" + }, + { + "name": "k_extent", + "type": "index" + }, + { + "name": "v_extent", + "type": "index" + }, + { + "name": "m_extent", + "type": "index" + } + ], + "bindings": [ + "query", + "key", + "value", + "mask", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_small_f16_f32", + "family": "loom_libs", + "symbol": "ggml_mul_mat_small_f16_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "k_size", + "type": "index" + }, + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "w_s1", + "type": "index" + }, + { + "name": "x_s1", + "type": "index" + }, + { + "name": "w_extent", + "type": "index" + }, + { + "name": "x_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "k_size", + "type": "index" + }, + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "w_s1", + "type": "index" + }, + { + "name": "x_s1", + "type": "index" + }, + { + "name": "w_extent", + "type": "index" + }, + { + "name": "x_extent", + "type": "index" + } + ], + "bindings": [ + "weight", + "input", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_small_f32_f32", + "family": "loom_libs", + "symbol": "ggml_mul_mat_small_f32_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "k_size", + "type": "index" + }, + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "w_s1", + "type": "index" + }, + { + "name": "x_s1", + "type": "index" + }, + { + "name": "w_extent", + "type": "index" + }, + { + "name": "x_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "k_size", + "type": "index" + }, + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "w_s1", + "type": "index" + }, + { + "name": "x_s1", + "type": "index" + }, + { + "name": "w_extent", + "type": "index" + }, + { + "name": "x_extent", + "type": "index" + } + ], + "bindings": [ + "weight", + "input", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_small_q8_0_f32", + "family": "loom_libs", + "symbol": "ggml_mul_mat_small_q8_0_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "k_size", + "type": "index" + }, + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "w_s1", + "type": "index" + }, + { + "name": "x_s1", + "type": "index" + }, + { + "name": "w_extent", + "type": "index" + }, + { + "name": "x_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "k_size", + "type": "index" + }, + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "w_s1", + "type": "index" + }, + { + "name": "x_s1", + "type": "index" + }, + { + "name": "w_extent", + "type": "index" + }, + { + "name": "x_extent", + "type": "index" + } + ], + "bindings": [ + "weight", + "input", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_rows_q8_0_f32", + "family": "loom_libs", + "symbol": "ggml_mul_mat_rows_q8_0_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "k_size", + "type": "index" + }, + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "w_s1", + "type": "index" + }, + { + "name": "x_s1", + "type": "index" + }, + { + "name": "w_extent", + "type": "index" + }, + { + "name": "x_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "k_size", + "type": "index" + }, + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "w_s1", + "type": "index" + }, + { + "name": "x_s1", + "type": "index" + }, + { + "name": "w_extent", + "type": "index" + }, + { + "name": "x_extent", + "type": "index" + } + ], + "bindings": [ + "weight", + "input", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_attention_rows_f32_f16", + "family": "loom_libs", + "symbol": "ggml_attention_rows_f32_f16", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "qk_size", + "type": "index" + }, + { + "name": "v_size", + "type": "index" + }, + { + "name": "q_count", + "type": "index" + }, + { + "name": "kv_count", + "type": "index" + }, + { + "name": "head_count", + "type": "index" + }, + { + "name": "kv_head_count", + "type": "index" + }, + { + "name": "q_s1", + "type": "index" + }, + { + "name": "q_s2", + "type": "index" + }, + { + "name": "k_s1", + "type": "index" + }, + { + "name": "k_s2", + "type": "index" + }, + { + "name": "v_s1", + "type": "index" + }, + { + "name": "v_s2", + "type": "index" + }, + { + "name": "m_s1", + "type": "index" + }, + { + "name": "q_extent", + "type": "index" + }, + { + "name": "k_extent", + "type": "index" + }, + { + "name": "v_extent", + "type": "index" + }, + { + "name": "m_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "qk_size", + "type": "index" + }, + { + "name": "v_size", + "type": "index" + }, + { + "name": "q_count", + "type": "index" + }, + { + "name": "kv_count", + "type": "index" + }, + { + "name": "head_count", + "type": "index" + }, + { + "name": "kv_head_count", + "type": "index" + }, + { + "name": "q_s1", + "type": "index" + }, + { + "name": "q_s2", + "type": "index" + }, + { + "name": "k_s1", + "type": "index" + }, + { + "name": "k_s2", + "type": "index" + }, + { + "name": "v_s1", + "type": "index" + }, + { + "name": "v_s2", + "type": "index" + }, + { + "name": "m_s1", + "type": "index" + }, + { + "name": "q_extent", + "type": "index" + }, + { + "name": "k_extent", + "type": "index" + }, + { + "name": "v_extent", + "type": "index" + }, + { + "name": "m_extent", + "type": "index" + } + ], + "bindings": [ + "query", + "key", + "value", + "mask", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_rope_rotate_half_f32", + "family": "loom_libs", + "symbol": "ggml_rope_rotate_half_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "ne0", + "type": "index" + }, + { + "name": "ne1", + "type": "index" + }, + { + "name": "ne2", + "type": "index" + }, + { + "name": "cos_s1", + "type": "index" + }, + { + "name": "sin_s1", + "type": "index" + }, + { + "name": "cos_extent", + "type": "index" + }, + { + "name": "sin_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "ne0", + "type": "index" + }, + { + "name": "ne1", + "type": "index" + }, + { + "name": "ne2", + "type": "index" + }, + { + "name": "cos_s1", + "type": "index" + }, + { + "name": "sin_s1", + "type": "index" + }, + { + "name": "cos_extent", + "type": "index" + }, + { + "name": "sin_extent", + "type": "index" + } + ], + "bindings": [ + "input", + "cos", + "sin", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_geglu_strided_f32", + "family": "loom_libs", + "symbol": "ggml_geglu_strided_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "a_s1", + "type": "index" + }, + { + "name": "b_s1", + "type": "index" + }, + { + "name": "a_extent", + "type": "index" + }, + { + "name": "b_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "a_s1", + "type": "index" + }, + { + "name": "b_s1", + "type": "index" + }, + { + "name": "a_extent", + "type": "index" + }, + { + "name": "b_extent", + "type": "index" + } + ], + "bindings": [ + "gate", + "up", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_id_decode_f32_wave64", + "family": "loom_libs", + "symbol": "ggml_mul_mat_id_decode_f32_wave64", + "source": "ops/mul_mat_id_decode_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "slot_count", + "type": "index" + }, + { + "name": "input_rows", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "slot_count", + "type": "index" + }, + { + "name": "input_rows", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "ids", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/mul_mat_id_decode_f32.loom" + ], + "library_sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + "compile_dependencies": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ] + }, + { + "name": "ggml_res_scale_pair_f32", + "family": "loom_libs", + "symbol": "ggml_res_scale_pair_f32", + "source": "ops/res_scale_pair_f32.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + }, + { + "name": "row_size", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + }, + { + "name": "row_size", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + } + ], + "bindings": [ + "a", + "bias", + "scale", + "addend", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/res_scale_pair_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_zaya_cca_conv_decode_f32", + "family": "loom_libs", + "symbol": "ggml_zaya_cca_conv_decode_f32", + "source": "ops/zaya_cca_conv_decode_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "qraw", + "kraw", + "state", + "dw", + "dw_bias", + "weight", + "grp_bias", + "output", + "new_state" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/zaya_cca_conv_decode_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_zaya_cca_qk_norm_decode_f32", + "family": "loom_libs", + "symbol": "ggml_zaya_cca_qk_norm_decode_f32", + "source": "ops/zaya_cca_qk_norm_decode_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "conv", + "qraw", + "kraw", + "k_scale", + "q_out", + "k_out" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/zaya_cca_qk_norm_decode_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_kquant_swiglu_decode_f32", + "family": "loom_libs", + "symbol": "ggml_kquant_swiglu_decode_f32", + "source": "ops/kquant_decode_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "input", + "gate", + "up", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/kquant_decode_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_kquant_mul_mat_decode_f32", + "family": "loom_libs", + "symbol": "ggml_kquant_mul_mat_decode_f32", + "source": "ops/kquant_decode_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "input", + "weight", + "addend", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/kquant_decode_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_kquant_swiglu_decode_tokens_f32", + "family": "loom_libs", + "symbol": "ggml_kquant_swiglu_decode_tokens_f32", + "source": "ops/kquant_decode_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "input", + "gate", + "up", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/kquant_decode_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_kquant_mul_mat_decode_tokens_f32", + "family": "loom_libs", + "symbol": "ggml_kquant_mul_mat_decode_tokens_f32", + "source": "ops/kquant_decode_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "input", + "weight", + "addend", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/kquant_decode_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_softplus_f32", + "family": "loom_libs", + "symbol": "ggml_softplus_f32", + "source": "ops/softplus_f32.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/softplus_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_hadamard_f32", + "family": "loom_libs", + "symbol": "ggml_hadamard_f32", + "source": "ops/hadamard_f32.loom", + "workload_parameters": [ + { + "name": "row_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "row_count", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/hadamard_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_add_id_f32", + "family": "loom_libs", + "symbol": "ggml_add_id_f32", + "source": "ops/add_id_f32.loom", + "workload_parameters": [ + { + "name": "width", + "type": "index" + }, + { + "name": "rows", + "type": "index" + }, + { + "name": "tokens", + "type": "index" + }, + { + "name": "ids_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "width", + "type": "index" + }, + { + "name": "rows", + "type": "index" + }, + { + "name": "tokens", + "type": "index" + }, + { + "name": "ids_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ], + "bindings": [ + "input", + "bias", + "ids", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/add_id_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_swiglu_oai_f32", + "family": "loom_libs", + "symbol": "ggml_swiglu_oai_f32", + "source": "ops/swiglu_oai_f32.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "bindings": [ + "gate", + "up", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/swiglu_oai_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_attention_sink_f32", + "family": "loom_libs", + "symbol": "ggml_attention_sink_f32", + "source": "ops/attention_sink_f32.loom", + "workload_parameters": [ + { + "name": "tokens", + "type": "index" + }, + { + "name": "key_count", + "type": "index" + }, + { + "name": "key_capacity", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "tokens", + "type": "index" + }, + { + "name": "key_count", + "type": "index" + }, + { + "name": "key_capacity", + "type": "index" + } + ], + "bindings": [ + "query", + "key", + "mask", + "sinks", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/attention_sink_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + } + ], + "link_modules": [], + "plan_cases": [] +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/bf16_f16.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/bf16_f16.loom new file mode 100644 index 000000000000..1a38444c8476 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/bf16_f16.loom @@ -0,0 +1,12 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +// Loads four adjacent BF16 weights and truncates them for FP16 matrix staging. +func.def inline @ggml_bf16_f16_vector4(%weight: buffer, %row_byte_base: offset, %input_size: index, %k: index) -> (vector<4xf16>) { + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %row_view = buffer.view %weight[%row_byte_base] : buffer -> view<[%bounded_input_size]xbf16> + %values_bf16 = vector.load %row_view[%k] : view<[%bounded_input_size]xbf16> -> vector<4xbf16> + %values_f32 = vector.extf %values_bf16 : vector<4xbf16> to vector<4xf32> + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/binary_f32_apply.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/binary_f32_apply.loom new file mode 100644 index 000000000000..bdbb77a86432 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/binary_f32_apply.loom @@ -0,0 +1,141 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +template.decl @ggml.binary_f32.apply(%op: index, %lhs: f32, %rhs: f32) -> (f32) + +template.def<@ggml.binary_f32.apply> device priority(1) @ggml_binary_f32_apply(%op: index, %lhs: f32, %rhs: f32) -> (f32) { + %c0_f32 = scalar.constant 0.0 : f32 + %c0_5_f32 = scalar.constant 0.5 : f32 + %c1_f32 = scalar.constant 1.0 : f32 + %c2_f32 = scalar.constant 2.0 : f32 + %cn2_f32 = scalar.constant -2.0 : f32 + %gelu_coef = scalar.constant 0.044715 : f32 + %gelu_quick_coef = scalar.constant -1.702 : f32 + %sqrt_2_over_pi = scalar.constant 0.79788456080286544 : f32 + %sqrt_2_inv = scalar.constant 0.70710678118654745 : f32 + + // ADD (op 0): x + y. + %op_add = index.constant 0 : index + %sum = scalar.addf %lhs, %rhs : f32 + %is_add = index.cmp eq, %op, %op_add : index + %add_selected = scf.select %is_add, %sum, %lhs : f32 + + // SUB (op 1): x - y. + %op_sub = index.constant 1 : index + %difference = scalar.subf %lhs, %rhs : f32 + %is_sub = index.cmp eq, %op, %op_sub : index + %sub_selected = scf.select %is_sub, %difference, %add_selected : f32 + + // MUL (op 2): x * y. + %op_mul = index.constant 2 : index + %product = scalar.mulf %lhs, %rhs : f32 + %is_mul = index.cmp eq, %op, %op_mul : index + %mul_selected = scf.select %is_mul, %product, %sub_selected : f32 + + // DIV (op 3): x / y. + %op_div = index.constant 3 : index + %quotient = scalar.divf %lhs, %rhs : f32 + %is_div = index.cmp eq, %op, %op_div : index + %div_selected = scf.select %is_div, %quotient, %mul_selected : f32 + + // SWIGLU (op 4): silu(x) * y. + %op_swiglu = index.constant 4 : index + %silu_negative = scalar.negf %lhs : f32 + %silu_exp = scalar.expf %silu_negative : f32 + %silu_denominator = scalar.addf %c1_f32, %silu_exp : f32 + %silu_sigmoid = scalar.divf %c1_f32, %silu_denominator : f32 + %silu = scalar.mulf %lhs, %silu_sigmoid : f32 + %swiglu = scalar.mulf %silu, %rhs : f32 + %is_swiglu = index.cmp eq, %op, %op_swiglu : index + %swiglu_selected = scf.select %is_swiglu, %swiglu, %div_selected : f32 + + // GEGLU (op 5): gelu(x) * y. + %op_geglu = index.constant 5 : index + %gelu_x2 = scalar.mulf %lhs, %lhs : f32 + %gelu_poly0 = scalar.mulf %gelu_coef, %gelu_x2 : f32 + %gelu_poly = scalar.addf %c1_f32, %gelu_poly0 : f32 + %gelu_inner0 = scalar.mulf %lhs, %gelu_poly : f32 + %gelu_inner = scalar.mulf %sqrt_2_over_pi, %gelu_inner0 : f32 + %gelu_tanh_scaled = scalar.mulf %cn2_f32, %gelu_inner : f32 + %gelu_tanh_exp = scalar.expf %gelu_tanh_scaled : f32 + %gelu_tanh_denominator = scalar.addf %c1_f32, %gelu_tanh_exp : f32 + %gelu_tanh_ratio = scalar.divf %c2_f32, %gelu_tanh_denominator : f32 + %gelu_tanh = scalar.subf %gelu_tanh_ratio, %c1_f32 : f32 + %gelu_one_plus = scalar.addf %c1_f32, %gelu_tanh : f32 + %gelu_half_x = scalar.mulf %c0_5_f32, %lhs : f32 + %gelu = scalar.mulf %gelu_half_x, %gelu_one_plus : f32 + %geglu = scalar.mulf %gelu, %rhs : f32 + %is_geglu = index.cmp eq, %op, %op_geglu : index + %geglu_selected = scf.select %is_geglu, %geglu, %swiglu_selected : f32 + + // REGLU (op 6): relu(x) * y. + %op_reglu = index.constant 6 : index + %relu_is_positive = scalar.cmpf ogt, %lhs, %c0_f32 : f32 + %relu = scf.select %relu_is_positive, %lhs, %c0_f32 : f32 + %reglu = scalar.mulf %relu, %rhs : f32 + %is_reglu = index.cmp eq, %op, %op_reglu : index + %reglu_selected = scf.select %is_reglu, %reglu, %geglu_selected : f32 + + // GEGLU_ERF (op 7): gelu_erf(x) * y. + %op_geglu_erf = index.constant 7 : index + %gelu_erf_scaled = scalar.mulf %sqrt_2_inv, %lhs : f32 + %gelu_erf_value = scalar.erff %gelu_erf_scaled : f32 + %gelu_erf_one_plus = scalar.addf %c1_f32, %gelu_erf_value : f32 + %gelu_erf_half_x = scalar.mulf %c0_5_f32, %lhs : f32 + %gelu_erf = scalar.mulf %gelu_erf_half_x, %gelu_erf_one_plus : f32 + %geglu_erf = scalar.mulf %gelu_erf, %rhs : f32 + %is_geglu_erf = index.cmp eq, %op, %op_geglu_erf : index + %geglu_erf_selected = scf.select %is_geglu_erf, %geglu_erf, %reglu_selected : f32 + + // GEGLU_QUICK (op 8): gelu_quick(x) * y. + %op_geglu_quick = index.constant 8 : index + %gelu_quick_scaled = scalar.mulf %gelu_quick_coef, %lhs : f32 + %gelu_quick_exp = scalar.expf %gelu_quick_scaled : f32 + %gelu_quick_denominator = scalar.addf %c1_f32, %gelu_quick_exp : f32 + %gelu_quick_sigmoid = scalar.divf %c1_f32, %gelu_quick_denominator : f32 + %gelu_quick = scalar.mulf %lhs, %gelu_quick_sigmoid : f32 + %geglu_quick = scalar.mulf %gelu_quick, %rhs : f32 + %is_geglu_quick = index.cmp eq, %op, %op_geglu_quick : index + %result = scf.select %is_geglu_quick, %geglu_quick, %geglu_erf_selected : f32 + template.return %result : f32 +} + +template.decl @ggml.binary_f32.apply_vector8(%op: index, %lhs: vector<8xf32>, %rhs: vector<8xf32>) -> (vector<8xf32>) + +template.def<@ggml.binary_f32.apply_vector8> device @ggml_binary_f32_apply_vector8(%op: index, %lhs: vector<8xf32>, %rhs: vector<8xf32>) -> (vector<8xf32>) { + %op_swiglu = index.constant 4 : index + %is_swiglu = index.cmp eq, %op, %op_swiglu : index + %result = scf.if %is_swiglu -> (vector<8xf32>) { + %silu = vector.siluf %lhs : vector<8xf32> + %value = vector.mulf %silu, %rhs : vector<8xf32> + scf.yield %value : vector<8xf32> + } else { + %lhs0 = vector.extract %lhs[0] : vector<8xf32> -> f32 + %rhs0 = vector.extract %rhs[0] : vector<8xf32> -> f32 + %value0 = template.apply<@ggml.binary_f32.apply>(%op, %lhs0, %rhs0) : (index, f32, f32) -> (f32) + %lhs1 = vector.extract %lhs[1] : vector<8xf32> -> f32 + %rhs1 = vector.extract %rhs[1] : vector<8xf32> -> f32 + %value1 = template.apply<@ggml.binary_f32.apply>(%op, %lhs1, %rhs1) : (index, f32, f32) -> (f32) + %lhs2 = vector.extract %lhs[2] : vector<8xf32> -> f32 + %rhs2 = vector.extract %rhs[2] : vector<8xf32> -> f32 + %value2 = template.apply<@ggml.binary_f32.apply>(%op, %lhs2, %rhs2) : (index, f32, f32) -> (f32) + %lhs3 = vector.extract %lhs[3] : vector<8xf32> -> f32 + %rhs3 = vector.extract %rhs[3] : vector<8xf32> -> f32 + %value3 = template.apply<@ggml.binary_f32.apply>(%op, %lhs3, %rhs3) : (index, f32, f32) -> (f32) + %lhs4 = vector.extract %lhs[4] : vector<8xf32> -> f32 + %rhs4 = vector.extract %rhs[4] : vector<8xf32> -> f32 + %value4 = template.apply<@ggml.binary_f32.apply>(%op, %lhs4, %rhs4) : (index, f32, f32) -> (f32) + %lhs5 = vector.extract %lhs[5] : vector<8xf32> -> f32 + %rhs5 = vector.extract %rhs[5] : vector<8xf32> -> f32 + %value5 = template.apply<@ggml.binary_f32.apply>(%op, %lhs5, %rhs5) : (index, f32, f32) -> (f32) + %lhs6 = vector.extract %lhs[6] : vector<8xf32> -> f32 + %rhs6 = vector.extract %rhs[6] : vector<8xf32> -> f32 + %value6 = template.apply<@ggml.binary_f32.apply>(%op, %lhs6, %rhs6) : (index, f32, f32) -> (f32) + %lhs7 = vector.extract %lhs[7] : vector<8xf32> -> f32 + %rhs7 = vector.extract %rhs[7] : vector<8xf32> -> f32 + %value7 = template.apply<@ggml.binary_f32.apply>(%op, %lhs7, %rhs7) : (index, f32, f32) -> (f32) + %value = vector.from_elements %value0, %value1, %value2, %value3, %value4, %value5, %value6, %value7 : vector<8xf32> + scf.yield %value : vector<8xf32> + } + template.return %result : vector<8xf32> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/dequant.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/dequant.loom new file mode 100644 index 000000000000..b324789a8a97 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/dequant.loom @@ -0,0 +1,11894 @@ +func.decl @ggml_pq2_0_f32_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf32>) +func.decl @ggml_pq2_0_f16_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf16>) +func.decl @ggml_ptq1_0_f32_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf32>) +func.decl @ggml_ptq1_0_f16_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf16>) + +func.decl @ggml_tq1_0_f16_vector4(%weight: buffer, %row_byte_base: offset, %tq_block: index, %tq_group: index, %packet: index) -> (vector<4xf16>) +func.decl @ggml_tq1_0_f32_vector4(%weight: buffer, %row_byte_base: offset, %tq_block: index, %tq_group: index, %packet: index) -> (vector<4xf32>) +func.decl @ggml_tq2_0_f16_vector4(%weight: buffer, %row_byte_base: offset, %tq_block: index, %tq_group: index, %packet: index) -> (vector<4xf16>) +func.decl @ggml_tq2_0_f32_vector4(%weight: buffer, %row_byte_base: offset, %tq_block: index, %tq_group: index, %packet: index) -> (vector<4xf32>) +func.decl @ggml_mxfp4_f16_vector4(%weight: buffer, %row_byte_base: offset, %mx_block: index, %packet: index) -> (vector<4xf16>) +func.decl @ggml_mxfp4_f32_vector4(%weight: buffer, %row_byte_base: offset, %mx_block: index, %packet: index) -> (vector<4xf32>) +func.def inline @ggml_paired_weight_words(%half_packet: i1, %paired: i1, %load_up: i1, %weight: buffer, %peer: buffer, %offset: offset) -> (vector<4xi32>) { + %zero = index.constant 0 : index + %padding = vector.constant 0 : vector<2xi32> + %weight_words = scf.if %half_packet -> (vector<4xi32>) { + %view = buffer.view %weight[%offset] : buffer -> view<2xi32> + %loaded = vector.load %view[%zero] : view<2xi32> -> vector<2xi32> + %words = vector.concat<0> %loaded, %padding : vector<2xi32>, vector<2xi32> -> vector<4xi32> + scf.yield %words : vector<4xi32> + } else { + %view = buffer.view %weight[%offset] : buffer -> view<4xi32> + %words = vector.load %view[%zero] : view<4xi32> -> vector<4xi32> + scf.yield %words : vector<4xi32> + } + %peer_words = scf.if %paired -> (vector<4xi32>) { + %words = scf.if %half_packet -> (vector<4xi32>) { + %view = buffer.view %peer[%offset] : buffer -> view<2xi32> + %loaded = vector.load %view[%zero] : view<2xi32> -> vector<2xi32> + %packet = vector.concat<0> %loaded, %padding : vector<2xi32>, vector<2xi32> -> vector<4xi32> + scf.yield %packet : vector<4xi32> + } else { + %view = buffer.view %peer[%offset] : buffer -> view<4xi32> + %packet = vector.load %view[%zero] : view<4xi32> -> vector<4xi32> + scf.yield %packet : vector<4xi32> + } + scf.yield %words : vector<4xi32> + } else { + scf.yield %weight_words : vector<4xi32> + } + %words = scf.select %load_up, %peer_words, %weight_words : vector<4xi32> + func.return %words : vector<4xi32> +} + +func.def public inline @ggml_q4k_native_row64_f16_pair(%half_packet: i1, %paired: i1, %weight: buffer, %peer: buffer, %is_up: i1, %row_base: offset, %group0: index, %packet: index) -> (vector<16xf16>, vector<16xf16>) { + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %field_bytes = index.constant 1024 : offset + %false = scalar.constant false : i1 + %header = func.call @ggml_paired_weight_words(%false, %paired, %is_up, %weight, %peer, %row_base) : (i1, i1, i1, buffer, buffer, offset) -> (vector<4xi32>) + %halves = vector.bitcast %header : vector<4xi32> to vector<8xf16> + %d_half = vector.extract %halves[0] : vector<8xf16> -> f16 + %m_half = vector.extract %halves[1] : vector<8xf16> -> f16 + %d = scalar.extf %d_half : f16 to f32 + %m = scalar.extf %m_half : f16 to f32 + %h0 = vector.extract %header[1] : vector<4xi32> -> i32 + %h1 = vector.extract %header[2] : vector<4xi32> -> i32 + %h2 = vector.extract %header[3] : vector<4xi32> -> i32 + %group1 = index.add %group0, %c1 : index + %s0, %m0 = func.call @ggml_q4k_scale_min_from_header(%h0, %h1, %h2, %group0) : (i32, i32, i32, index) -> (i32, i32) + %s1, %m1 = func.call @ggml_q4k_scale_min_from_header(%h0, %h1, %h2, %group1) : (i32, i32, i32, index) -> (i32, i32) + %pair = index.div %group0, %c2 : index + %pair_field = index.mul %pair, %c2 : index + %packet_field = index.div %packet, %c4 : index + %field0 = index.add %pair_field, %packet_field : index + %field = index.add %field0, %c1 : index + %field_add = index.scale %field, %field_bytes : index, offset -> offset + %field_base = index.add %row_base, %field_add : offset + %word_bytes = index.constant 4 : offset + %packet_word = index.rem %packet, %c4 : index + %packet_add = index.scale %packet_word, %word_bytes : index, offset -> offset + %address = index.add %field_base, %packet_add : offset + %codes = func.call @ggml_paired_weight_words(%half_packet, %paired, %is_up, %weight, %peer, %address) : (i1, i1, i1, buffer, buffer, offset) -> (vector<4xi32>) + %mask = vector.constant 252645135 : vector<4xi32> + %shift = vector.constant 4 : vector<4xi32> + %high = vector.shrui %codes, %shift : vector<4xi32> + %nibbles0 = vector.andi %codes, %mask : vector<4xi32> + %bytes0 = vector.bitcast %nibbles0 : vector<4xi32> to vector<16xi8> + %sf0 = scalar.uitofp %s0 : i32 to f32 + %mf0 = scalar.uitofp %m0 : i32 to f32 + %ds0 = scalar.mulf %d, %sf0 : f32 + %ms0 = scalar.mulf %m, %mf0 : f32 + %negative0 = scalar.negf %ms0 : f32 + %dv0 = vector.splat %ds0 : vector<16xf32> + %mv0 = vector.splat %negative0 : vector<16xf32> + %qf0 = vector.uitofp %bytes0 : vector<16xi8> to vector<16xf32> + %values0 = vector.fmaf %qf0, %dv0, %mv0 : vector<16xf32> + %result0 = vector.fptrunc %values0 : vector<16xf32> to vector<16xf16> + %nibbles1 = vector.andi %high, %mask : vector<4xi32> + %bytes1 = vector.bitcast %nibbles1 : vector<4xi32> to vector<16xi8> + %sf1 = scalar.uitofp %s1 : i32 to f32 + %mf1 = scalar.uitofp %m1 : i32 to f32 + %ds1 = scalar.mulf %d, %sf1 : f32 + %ms1 = scalar.mulf %m, %mf1 : f32 + %negative1 = scalar.negf %ms1 : f32 + %dv1 = vector.splat %ds1 : vector<16xf32> + %mv1 = vector.splat %negative1 : vector<16xf32> + %qf1 = vector.uitofp %bytes1 : vector<16xi8> to vector<16xf32> + %values1 = vector.fmaf %qf1, %dv1, %mv1 : vector<16xf32> + %result1 = vector.fptrunc %values1 : vector<16xf32> to vector<16xf16> + func.return %result0, %result1 : vector<16xf16>, vector<16xf16> +} + + +func.def inline @ggml_q4k_scale_min_from_header(%scale0: i32, %scale1: i32, %scale2: i32, %q4_group: index) -> (i32, i32) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c48_i32 = scalar.constant 48 : i32 + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + %is_low_group = index.cmp ult, %bounded_group, %c4 : index + %scale_lane = index.rem %bounded_group, %c4 : index + %scale_shift_index = index.mul %scale_lane, %c8 : index + %scale_shift = index.cast %scale_shift_index : index to i32 + %high_shift = scalar.addi %scale_shift, %c2_i32 : i32 + %minimum_shift = scalar.addi %scale_shift, %c4_i32 : i32 + %selected_scale_source = scf.select %is_low_group, %scale0, %scale2 : i32 + %selected_minimum_source = scf.select %is_low_group, %scale1, %scale2 : i32 + %selected_scale_high_shift = scf.select %is_low_group, %scale_shift, %high_shift : i32 + %selected_minimum_low_shift = scf.select %is_low_group, %scale_shift, %minimum_shift : i32 + %scale_low0 = scalar.shrui %selected_scale_source, %scale_shift : i32 + %scale_low = scalar.andi %scale_low0, %c15_i32 : i32 + %scale_high0 = scalar.shrui %scale0, %selected_scale_high_shift : i32 + %scale_high = scalar.andi %scale_high0, %c48_i32 : i32 + %scale = scalar.ori %scale_low, %scale_high : i32 + %minimum_low0 = scalar.shrui %selected_minimum_source, %selected_minimum_low_shift : i32 + %minimum_low = scalar.andi %minimum_low0, %c15_i32 : i32 + %minimum_high0 = scalar.shrui %scale1, %selected_scale_high_shift : i32 + %minimum_high = scalar.andi %minimum_high0, %c48_i32 : i32 + %minimum = scalar.ori %minimum_low, %minimum_high : i32 + func.return %scale, %minimum : i32, i32 +} + +func.def inline @ggml_dot_u8_s8_vector8_f32(%weight: vector<2xi32>, %activation: vector<2xi32>) -> (f32) { + %zero = vector.constant 0 : vector<2xi32> + %seed = scalar.constant 0 : i32 + %w = vector.bitcast %weight : vector<2xi32> to vector<8xi8> + %a = vector.bitcast %activation : vector<2xi32> to vector<8xi8> + %d = vector.dot4i %w, %a, %zero : vector<8xi8>, vector<8xi8>, vector<2xi32> + %sum = vector.reduce %d, %seed : vector<2xi32>, i32 + %result = scalar.sitofp %sum : i32 to f32 + func.return %result : f32 +} + +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 +func.decl @ggml_q4k_f16_vector4(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group: index, %packet: index) -> (vector<4xf16>) + +func.decl @ggml_q6k_f16_vector4(%weight: buffer, %weight_row_byte_base: offset, %q6_block: index, %q6_group: index, %packet: index) -> (vector<4xf16>) + +func.decl @ggml_q8_0_f16_vector4(%weight: buffer, %row_byte_base: offset, %q8_block: index, %packet: index) -> (vector<4xf16>) + +func.decl @ggml_q8_1_f16_vector4(%weight: buffer, %row_byte_base: offset, %q8_block: index, %packet: index) -> (vector<4xf16>) + +func.decl @ggml_f16_f16_vector4(%weight: buffer, %row_byte_base: offset, %input_size: index, %k: index) -> (vector<4xf16>) + +func.decl @ggml_bf16_f16_vector4(%weight: buffer, %row_byte_base: offset, %input_size: index, %k: index) -> (vector<4xf16>) + +func.decl @ggml_f32_f16_vector4(%weight: buffer, %row_byte_base: offset, %input_size: index, %k: index) -> (vector<4xf16>) + +func.def inline @ggml_q5k_high_bit_f32(%qh: i8, %qh_bit: i32) -> (f32) { + %c0_i32 = scalar.constant 0 : i32 + %c16_i32 = scalar.constant 16 : i32 + %qh_i32 = scalar.extui %qh : i8 to i32 + %masked = scalar.andi %qh_i32, %qh_bit : i32 + %is_set = scalar.cmpi ne, %masked, %c0_i32 : i32 + %high = scf.select %is_set, %c16_i32, %c0_i32 : i32 + %high_f32 = scalar.uitofp %high : i32 to f32 + func.return %high_f32 : f32 +} + +func.def inline @ggml_q5_legacy_high_bit_f32(%qh: i32, %bit: i32) -> (f32) { + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c16_i32 = scalar.constant 16 : i32 + %shifted = scalar.shrui %qh, %bit : i32 + %masked = scalar.andi %shifted, %c1_i32 : i32 + %is_set = scalar.cmpi ne, %masked, %c0_i32 : i32 + %high = scf.select %is_set, %c16_i32, %c0_i32 : i32 + %high_f32 = scalar.uitofp %high : i32 to f32 + func.return %high_f32 : f32 +} + +func.def inline @ggml_q1_0_f32_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c128 = index.constant 128 : index + %c0_i32 = scalar.constant 0 : i32 + %c0_f32 = scalar.constant 0.0 : f32 + %block_bytes = index.constant 18 : offset + %code_offset = index.constant 2 : offset + %block = index.div %k, %c128 : index + %k_in_block = index.rem %k, %c128 : index + %packet = index.div %k_in_block, %c4 : index + %byte_index0 = index.div %k_in_block, %c8 : index + %byte_index = index.assume %byte_index0 [range(%byte_index0, 0, 15)] : index + %bit_base_index = index.rem %k_in_block, %c8 : index + %bit_base = index.cast %bit_base_index : index to i32 + %block_byte_add = index.scale %block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %d_view = buffer.view %weight[%block_byte_base] : buffer -> view<1xf16> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<16xi8> + %d_vector = vector.load %d_view[%c0] : view<1xf16> -> vector<1xf16> + %d_f16 = vector.extract %d_vector[0] : vector<1xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %neg_d = scalar.subf %c0_f32, %d : f32 + %q_i8 = view.load %code_view[%byte_index] : view<16xi8> -> i8 + %q = scalar.extui %q_i8 : i8 to i32 + %bit1 = scalar.addi %bit_base, %c1_i32 : i32 + %bit2 = scalar.addi %bit_base, %c2_i32 : i32 + %bit3 = scalar.addi %bit_base, %c3_i32 : i32 + %q0_shift = scalar.shrui %q, %bit_base : i32 + %q1_shift = scalar.shrui %q, %bit1 : i32 + %q2_shift = scalar.shrui %q, %bit2 : i32 + %q3_shift = scalar.shrui %q, %bit3 : i32 + %q0_mask = scalar.andi %q0_shift, %c1_i32 : i32 + %q1_mask = scalar.andi %q1_shift, %c1_i32 : i32 + %q2_mask = scalar.andi %q2_shift, %c1_i32 : i32 + %q3_mask = scalar.andi %q3_shift, %c1_i32 : i32 + %q0_set = scalar.cmpi ne, %q0_mask, %c0_i32 : i32 + %q1_set = scalar.cmpi ne, %q1_mask, %c0_i32 : i32 + %q2_set = scalar.cmpi ne, %q2_mask, %c0_i32 : i32 + %q3_set = scalar.cmpi ne, %q3_mask, %c0_i32 : i32 + %v0 = scf.select %q0_set, %d, %neg_d : f32 + %v1 = scf.select %q1_set, %d, %neg_d : f32 + %v2 = scf.select %q2_set, %d, %neg_d : f32 + %v3 = scf.select %q3_set, %d, %neg_d : f32 + %result = vector.from_elements %v0, %v1, %v2, %v3 : vector<4xf32> + func.return %result : vector<4xf32> +} + +func.def inline @ggml_q1_0_f16_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_q1_0_f32_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +func.def inline @ggml_q4_0_f32_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %c32 = index.constant 32 : index + %c0_i32 = scalar.constant 0 : i32 + %c4_i32 = scalar.constant 4 : i32 + %block_bytes = index.constant 18 : offset + %code_offset = index.constant 2 : offset + %q4_mask = vector.constant 252645135 : vector<1xi32> + %zero_point = vector.constant 8.0 : vector<4xf32> + %block = index.div %k, %c32 : index + %k_in_block = index.rem %k, %c32 : index + %packet = index.div %k_in_block, %c4 : index + %word_index0 = index.rem %packet, %c4 : index + %word_index = index.assume %word_index0 [range(%word_index0, 0, 3)] : index + %is_high = index.cmp uge, %packet, %c4 : index + %shift_i32 = scf.select %is_high, %c4_i32, %c0_i32 : i32 + %shift = vector.splat %shift_i32 : vector<1xi32> + %block_byte_add = index.scale %block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %d_view = buffer.view %weight[%block_byte_base] : buffer -> view<1xf16> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<4xi32> + %d_vector = vector.load %d_view[%c0] : view<1xf16> -> vector<1xf16> + %d_f16 = vector.extract %d_vector[0] : vector<1xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %d_splat = vector.splat %d : vector<4xf32> + %q_word = vector.load %code_view[%word_index] : view<4xi32> -> vector<1xi32> + %shifted_q = vector.shrui %q_word, %shift : vector<1xi32> + %masked_q = vector.andi %shifted_q, %q4_mask : vector<1xi32> + %q_i8 = vector.bitcast %masked_q : vector<1xi32> to vector<4xi8> + %q_unsigned = vector.uitofp %q_i8 : vector<4xi8> to vector<4xf32> + %q = vector.subf %q_unsigned, %zero_point : vector<4xf32> + %result = vector.mulf %q, %d_splat : vector<4xf32> + func.return %result : vector<4xf32> +} + +func.def inline @ggml_q4_0_f16_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_q4_0_f32_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +func.def inline @ggml_q4_1_f32_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %c32 = index.constant 32 : index + %c0_i32 = scalar.constant 0 : i32 + %c4_i32 = scalar.constant 4 : i32 + %block_bytes = index.constant 20 : offset + %code_offset = index.constant 4 : offset + %q4_mask = vector.constant 252645135 : vector<1xi32> + %block = index.div %k, %c32 : index + %k_in_block = index.rem %k, %c32 : index + %packet = index.div %k_in_block, %c4 : index + %word_index0 = index.rem %packet, %c4 : index + %word_index = index.assume %word_index0 [range(%word_index0, 0, 3)] : index + %is_high = index.cmp uge, %packet, %c4 : index + %shift_i32 = scf.select %is_high, %c4_i32, %c0_i32 : i32 + %shift = vector.splat %shift_i32 : vector<1xi32> + %block_byte_add = index.scale %block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %dm_view = buffer.view %weight[%block_byte_base] : buffer -> view<2xf16> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<4xi32> + %dm = vector.load %dm_view[%c0] : view<2xf16> -> vector<2xf16> + %d_f16 = vector.extract %dm[0] : vector<2xf16> -> f16 + %m_f16 = vector.extract %dm[1] : vector<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %m = scalar.extf %m_f16 : f16 to f32 + %d_splat = vector.splat %d : vector<4xf32> + %m_splat = vector.splat %m : vector<4xf32> + %q_word = vector.load %code_view[%word_index] : view<4xi32> -> vector<1xi32> + %shifted_q = vector.shrui %q_word, %shift : vector<1xi32> + %masked_q = vector.andi %shifted_q, %q4_mask : vector<1xi32> + %q_i8 = vector.bitcast %masked_q : vector<1xi32> to vector<4xi8> + %q = vector.uitofp %q_i8 : vector<4xi8> to vector<4xf32> + %scaled = vector.mulf %q, %d_splat : vector<4xf32> + %result = vector.addf %scaled, %m_splat : vector<4xf32> + func.return %result : vector<4xf32> +} + +func.def inline @ggml_q4_1_f16_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_q4_1_f32_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +func.def inline @ggml_q5_0_f32_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c1_i32 = scalar.constant 1 : i32 + %c2 = index.constant 2 : index + %c2_i32 = scalar.constant 2 : i32 + %c4 = index.constant 4 : index + %c16_i32 = scalar.constant 16 : i32 + %c32 = index.constant 32 : index + %c0_i32 = scalar.constant 0 : i32 + %c4_i32 = scalar.constant 4 : i32 + %block_bytes = index.constant 22 : offset + %qh_offset = index.constant 2 : offset + %code_offset = index.constant 6 : offset + %q4_mask = vector.constant 252645135 : vector<1xi32> + %zero_point = vector.constant 16.0 : vector<4xf32> + %block = index.div %k, %c32 : index + %k_in_block = index.rem %k, %c32 : index + %packet = index.div %k_in_block, %c4 : index + %word_index0 = index.rem %packet, %c4 : index + %word_index = index.assume %word_index0 [range(%word_index0, 0, 3)] : index + %is_high = index.cmp uge, %packet, %c4 : index + %shift_i32 = scf.select %is_high, %c4_i32, %c0_i32 : i32 + %high_half_add = scf.select %is_high, %c16_i32, %c0_i32 : i32 + %shift = vector.splat %shift_i32 : vector<1xi32> + %block_byte_add = index.scale %block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %qh_byte_base = index.add %block_byte_base, %qh_offset : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %d_view = buffer.view %weight[%block_byte_base] : buffer -> view<1xf16> + %qh_view = buffer.view %weight[%qh_byte_base] : buffer -> view<1xi32> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<4xi32> + %d_vector = vector.load %d_view[%c0] : view<1xf16> -> vector<1xf16> + %d_f16 = vector.extract %d_vector[0] : vector<1xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %d_splat = vector.splat %d : vector<4xf32> + %qh_vector = vector.load %qh_view[%c0] : view<1xi32> -> vector<1xi32> + %qh = vector.extract %qh_vector[0] : vector<1xi32> -> i32 + %q_word = vector.load %code_view[%word_index] : view<4xi32> -> vector<1xi32> + %shifted_q = vector.shrui %q_word, %shift : vector<1xi32> + %masked_q = vector.andi %shifted_q, %q4_mask : vector<1xi32> + %q_i8 = vector.bitcast %masked_q : vector<1xi32> to vector<4xi8> + %q_low = vector.uitofp %q_i8 : vector<4xi8> to vector<4xf32> + %packet_base = index.mul %word_index, %c4 : index + %packet_base_i32 = index.cast %packet_base : index to i32 + %bit0 = scalar.addi %packet_base_i32, %high_half_add : i32 + %bit1 = scalar.addi %bit0, %c1_i32 : i32 + %bit2 = scalar.addi %bit0, %c2_i32 : i32 + %bit3 = scalar.addi %bit2, %c1_i32 : i32 + %high0 = func.call @ggml_q5_legacy_high_bit_f32(%qh, %bit0) : (i32, i32) -> (f32) + %high1 = func.call @ggml_q5_legacy_high_bit_f32(%qh, %bit1) : (i32, i32) -> (f32) + %high2 = func.call @ggml_q5_legacy_high_bit_f32(%qh, %bit2) : (i32, i32) -> (f32) + %high3 = func.call @ggml_q5_legacy_high_bit_f32(%qh, %bit3) : (i32, i32) -> (f32) + %high = vector.from_elements %high0, %high1, %high2, %high3 : vector<4xf32> + %q_unsigned = vector.addf %q_low, %high : vector<4xf32> + %q = vector.subf %q_unsigned, %zero_point : vector<4xf32> + %result = vector.mulf %q, %d_splat : vector<4xf32> + func.return %result : vector<4xf32> +} + +func.def inline @ggml_q5_0_f16_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_q5_0_f32_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +func.def inline @ggml_q5_1_f32_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c1_i32 = scalar.constant 1 : i32 + %c2 = index.constant 2 : index + %c2_i32 = scalar.constant 2 : i32 + %c4 = index.constant 4 : index + %c16_i32 = scalar.constant 16 : i32 + %c32 = index.constant 32 : index + %c0_i32 = scalar.constant 0 : i32 + %c4_i32 = scalar.constant 4 : i32 + %block_bytes = index.constant 24 : offset + %qh_offset = index.constant 4 : offset + %code_offset = index.constant 8 : offset + %q4_mask = vector.constant 252645135 : vector<1xi32> + %block = index.div %k, %c32 : index + %k_in_block = index.rem %k, %c32 : index + %packet = index.div %k_in_block, %c4 : index + %word_index0 = index.rem %packet, %c4 : index + %word_index = index.assume %word_index0 [range(%word_index0, 0, 3)] : index + %is_high = index.cmp uge, %packet, %c4 : index + %shift_i32 = scf.select %is_high, %c4_i32, %c0_i32 : i32 + %high_half_add = scf.select %is_high, %c16_i32, %c0_i32 : i32 + %shift = vector.splat %shift_i32 : vector<1xi32> + %block_byte_add = index.scale %block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %qh_byte_base = index.add %block_byte_base, %qh_offset : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %dm_view = buffer.view %weight[%block_byte_base] : buffer -> view<2xf16> + %qh_view = buffer.view %weight[%qh_byte_base] : buffer -> view<1xi32> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<4xi32> + %dm = vector.load %dm_view[%c0] : view<2xf16> -> vector<2xf16> + %d_f16 = vector.extract %dm[0] : vector<2xf16> -> f16 + %m_f16 = vector.extract %dm[1] : vector<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %m = scalar.extf %m_f16 : f16 to f32 + %d_splat = vector.splat %d : vector<4xf32> + %m_splat = vector.splat %m : vector<4xf32> + %qh_vector = vector.load %qh_view[%c0] : view<1xi32> -> vector<1xi32> + %qh = vector.extract %qh_vector[0] : vector<1xi32> -> i32 + %q_word = vector.load %code_view[%word_index] : view<4xi32> -> vector<1xi32> + %shifted_q = vector.shrui %q_word, %shift : vector<1xi32> + %masked_q = vector.andi %shifted_q, %q4_mask : vector<1xi32> + %q_i8 = vector.bitcast %masked_q : vector<1xi32> to vector<4xi8> + %q_low = vector.uitofp %q_i8 : vector<4xi8> to vector<4xf32> + %packet_base = index.mul %word_index, %c4 : index + %packet_base_i32 = index.cast %packet_base : index to i32 + %bit0 = scalar.addi %packet_base_i32, %high_half_add : i32 + %bit1 = scalar.addi %bit0, %c1_i32 : i32 + %bit2 = scalar.addi %bit0, %c2_i32 : i32 + %bit3 = scalar.addi %bit2, %c1_i32 : i32 + %high0 = func.call @ggml_q5_legacy_high_bit_f32(%qh, %bit0) : (i32, i32) -> (f32) + %high1 = func.call @ggml_q5_legacy_high_bit_f32(%qh, %bit1) : (i32, i32) -> (f32) + %high2 = func.call @ggml_q5_legacy_high_bit_f32(%qh, %bit2) : (i32, i32) -> (f32) + %high3 = func.call @ggml_q5_legacy_high_bit_f32(%qh, %bit3) : (i32, i32) -> (f32) + %high = vector.from_elements %high0, %high1, %high2, %high3 : vector<4xf32> + %q = vector.addf %q_low, %high : vector<4xf32> + %scaled = vector.mulf %q, %d_splat : vector<4xf32> + %result = vector.addf %scaled, %m_splat : vector<4xf32> + func.return %result : vector<4xf32> +} + +func.def inline @ggml_q5_1_f16_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_q5_1_f32_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +// Decodes four adjacent Q4_K values to F32 without the FP16 staging round used +// by WMMA paths. +func.def inline @ggml_q4k_f32_vector4(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c0_i32 = scalar.constant 0 : i32 + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %block_bytes = index.constant 144 : offset + %scale_offset = index.constant 4 : offset + %code_offset = index.constant 16 : offset + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15 = vector.constant 15 : vector<1xi32> + %c48 = vector.constant 48 : vector<1xi32> + %q4_mask = vector.constant 252645135 : vector<1xi32> + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q4_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %scale_byte_base = index.add %block_byte_base, %scale_offset : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %dm_view = buffer.view %weight[%block_byte_base] : buffer -> view<2xf16> + %scale_view = buffer.view %weight[%scale_byte_base] : buffer -> view<3xi32> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<32xi32> + %dm = vector.load %dm_view[%c0] : view<2xf16> -> vector<2xf16> + %scales = vector.load %scale_view[%c0] : view<3xi32> -> vector<3xi32> + %d_f16 = vector.extract %dm[0] : vector<2xf16> -> f16 + %dmin_f16 = vector.extract %dm[1] : vector<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %q_page0 = index.div %bounded_group, %c2 : index + %q_page = index.mul %q_page0, %c8 : index + %q_word_index0 = index.add %q_page, %bounded_packet : index + %q_word_index = index.assume %q_word_index0 [range(%q_word_index0, 0, 31)] : index + %is_low = index.cmp ult, %bounded_group, %c4 : index + %scale_lane = index.rem %bounded_group, %c4 : index + %scale_shift_index = index.mul %scale_lane, %c8 : index + %scale_shift_i32 = index.cast %scale_shift_index : index to i32 + %scale_shift = vector.splat %scale_shift_i32 : vector<1xi32> + %scale0_i32 = vector.extract %scales[0] : vector<3xi32> -> i32 + %scale1_i32 = vector.extract %scales[1] : vector<3xi32> -> i32 + %scale2_i32 = vector.extract %scales[2] : vector<3xi32> -> i32 + %scale0 = vector.splat %scale0_i32 : vector<1xi32> + %scale1 = vector.splat %scale1_i32 : vector<1xi32> + %scale2 = vector.splat %scale2_i32 : vector<1xi32> + %high_shift_i32 = scalar.addi %scale_shift_i32, %c2_i32 : i32 + %minimum_shift_i32 = scalar.addi %scale_shift_i32, %c4_i32 : i32 + %selected_scale_source = scf.select %is_low, %scale0, %scale2 : vector<1xi32> + %selected_minimum_source = scf.select %is_low, %scale1, %scale2 : vector<1xi32> + %selected_scale_high_shift_i32 = scf.select %is_low, %scale_shift_i32, %high_shift_i32 : i32 + %selected_minimum_low_shift_i32 = scf.select %is_low, %scale_shift_i32, %minimum_shift_i32 : i32 + %selected_scale_high_shift = vector.splat %selected_scale_high_shift_i32 : vector<1xi32> + %selected_minimum_low_shift = vector.splat %selected_minimum_low_shift_i32 : vector<1xi32> + %scale_low0 = vector.shrui %selected_scale_source, %scale_shift : vector<1xi32> + %scale_low = vector.andi %scale_low0, %c15 : vector<1xi32> + %scale_high0 = vector.shrui %scale0, %selected_scale_high_shift : vector<1xi32> + %scale_high = vector.andi %scale_high0, %c48 : vector<1xi32> + %scale = vector.ori %scale_low, %scale_high : vector<1xi32> + %minimum_low0 = vector.shrui %selected_minimum_source, %selected_minimum_low_shift : vector<1xi32> + %minimum_low = vector.andi %minimum_low0, %c15 : vector<1xi32> + %minimum_high0 = vector.shrui %scale1, %selected_scale_high_shift : vector<1xi32> + %minimum_high = vector.andi %minimum_high0, %c48 : vector<1xi32> + %minimum = vector.ori %minimum_low, %minimum_high : vector<1xi32> + %scale_f32 = vector.uitofp %scale : vector<1xi32> to vector<1xf32> + %minimum_f32 = vector.uitofp %minimum : vector<1xi32> to vector<1xf32> + %d_vector1 = vector.splat %d : vector<1xf32> + %dmin_vector1 = vector.splat %dmin : vector<1xf32> + %d_scale_vector1 = vector.mulf %d_vector1, %scale_f32 : vector<1xf32> + %minimum_scale_vector1 = vector.mulf %dmin_vector1, %minimum_f32 : vector<1xf32> + %d_scale = vector.extract %d_scale_vector1[0] : vector<1xf32> -> f32 + %minimum_scale = vector.extract %minimum_scale_vector1[0] : vector<1xf32> -> f32 + %q_word = vector.load %code_view[%q_word_index] : view<32xi32> -> vector<1xi32> + %q_half = index.rem %bounded_group, %c2 : index + %q_shift_index = index.mul %q_half, %c4 : index + %q_shift_i32 = index.cast %q_shift_index : index to i32 + %q_shift = vector.splat %q_shift_i32 : vector<1xi32> + %shifted_q = vector.shrui %q_word, %q_shift : vector<1xi32> + %masked_q = vector.andi %shifted_q, %q4_mask : vector<1xi32> + %q_i8 = vector.bitcast %masked_q : vector<1xi32> to vector<4xi8> + %q_f32 = vector.uitofp %q_i8 : vector<4xi8> to vector<4xf32> + %negative_minimum_scale = scalar.negf %minimum_scale : f32 + %q0 = vector.extract %q_f32[0] : vector<4xf32> -> f32 + %q1 = vector.extract %q_f32[1] : vector<4xf32> -> f32 + %q2 = vector.extract %q_f32[2] : vector<4xf32> -> f32 + %q3 = vector.extract %q_f32[3] : vector<4xf32> -> f32 + %value0 = scalar.fmaf %q0, %d_scale, %negative_minimum_scale : f32 + %value1 = scalar.fmaf %q1, %d_scale, %negative_minimum_scale : f32 + %value2 = scalar.fmaf %q2, %d_scale, %negative_minimum_scale : f32 + %value3 = scalar.fmaf %q3, %d_scale, %negative_minimum_scale : f32 + %result = vector.from_elements %value0, %value1, %value2, %value3 : vector<4xf32> + func.return %result : vector<4xf32> +} + +// Decode the byte-identical row64 native Q4_K carrier used by low-token +// contractions. The header and paired-nibble fields retain GGUF values. +func.def inline @ggml_q4k_native_row64_f32_vector4(%weight: buffer, %row: index, %input_size: index, %q4_block: index, %q4_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %block_bytes = index.constant 144 : offset + %scale_offset = index.constant 4 : offset + %code_offset = index.constant 16 : offset + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15 = vector.constant 15 : vector<1xi32> + %c48 = vector.constant 48 : vector<1xi32> + %q4_mask = vector.constant 252645135 : vector<1xi32> + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %record_bytes = index.constant 9216 : offset + %row_bytes = index.constant 16 : offset + %payload_offset = index.constant 1024 : offset + %blocks = index.div %input_size, %c256 : index + %row_group = index.div %row, %c64 : index + %row_lane = index.rem %row, %c64 : index + %group_block = index.madd %row_group, %blocks, %q4_block : index + %group_base = index.scale %group_block, %record_bytes : index, offset -> offset + %row_add = index.scale %row_lane, %row_bytes : index, offset -> offset + %block_byte_base = index.add %group_base, %row_add : offset + %scale_byte_base = index.add %block_byte_base, %scale_offset : offset + %code_byte_base = index.add %group_base, %payload_offset : offset + %dm_view = buffer.view %weight[%block_byte_base] : buffer -> view<2xf16> + %scale_view = buffer.view %weight[%scale_byte_base] : buffer -> view<3xi32> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<8x64x4xi32> + %dm = vector.load %dm_view[%c0] : view<2xf16> -> vector<2xf16> + %scales = vector.load %scale_view[%c0] : view<3xi32> -> vector<3xi32> + %d_f16 = vector.extract %dm[0] : vector<2xf16> -> f16 + %dmin_f16 = vector.extract %dm[1] : vector<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %q_page0 = index.div %bounded_group, %c2 : index + %q_page = index.mul %q_page0, %c8 : index + %q_word_index0 = index.add %q_page, %bounded_packet : index + %q_word_index = index.assume %q_word_index0 [range(%q_word_index0, 0, 31)] : index + %is_low = index.cmp ult, %bounded_group, %c4 : index + %scale_lane = index.rem %bounded_group, %c4 : index + %scale_shift_index = index.mul %scale_lane, %c8 : index + %scale_shift_i32 = index.cast %scale_shift_index : index to i32 + %scale_shift = vector.splat %scale_shift_i32 : vector<1xi32> + %scale0_i32 = vector.extract %scales[0] : vector<3xi32> -> i32 + %scale1_i32 = vector.extract %scales[1] : vector<3xi32> -> i32 + %scale2_i32 = vector.extract %scales[2] : vector<3xi32> -> i32 + %scale0 = vector.splat %scale0_i32 : vector<1xi32> + %scale1 = vector.splat %scale1_i32 : vector<1xi32> + %scale2 = vector.splat %scale2_i32 : vector<1xi32> + %high_shift_i32 = scalar.addi %scale_shift_i32, %c2_i32 : i32 + %minimum_shift_i32 = scalar.addi %scale_shift_i32, %c4_i32 : i32 + %selected_scale_source = scf.select %is_low, %scale0, %scale2 : vector<1xi32> + %selected_minimum_source = scf.select %is_low, %scale1, %scale2 : vector<1xi32> + %selected_scale_high_shift_i32 = scf.select %is_low, %scale_shift_i32, %high_shift_i32 : i32 + %selected_minimum_low_shift_i32 = scf.select %is_low, %scale_shift_i32, %minimum_shift_i32 : i32 + %selected_scale_high_shift = vector.splat %selected_scale_high_shift_i32 : vector<1xi32> + %selected_minimum_low_shift = vector.splat %selected_minimum_low_shift_i32 : vector<1xi32> + %scale_low0 = vector.shrui %selected_scale_source, %scale_shift : vector<1xi32> + %scale_low = vector.andi %scale_low0, %c15 : vector<1xi32> + %scale_high0 = vector.shrui %scale0, %selected_scale_high_shift : vector<1xi32> + %scale_high = vector.andi %scale_high0, %c48 : vector<1xi32> + %scale = vector.ori %scale_low, %scale_high : vector<1xi32> + %minimum_low0 = vector.shrui %selected_minimum_source, %selected_minimum_low_shift : vector<1xi32> + %minimum_low = vector.andi %minimum_low0, %c15 : vector<1xi32> + %minimum_high0 = vector.shrui %scale1, %selected_scale_high_shift : vector<1xi32> + %minimum_high = vector.andi %minimum_high0, %c48 : vector<1xi32> + %minimum = vector.ori %minimum_low, %minimum_high : vector<1xi32> + %scale_f32 = vector.uitofp %scale : vector<1xi32> to vector<1xf32> + %minimum_f32 = vector.uitofp %minimum : vector<1xi32> to vector<1xf32> + %d_vector1 = vector.splat %d : vector<1xf32> + %dmin_vector1 = vector.splat %dmin : vector<1xf32> + %d_scale_vector1 = vector.mulf %d_vector1, %scale_f32 : vector<1xf32> + %minimum_scale_vector1 = vector.mulf %dmin_vector1, %minimum_f32 : vector<1xf32> + %d_scale = vector.extract %d_scale_vector1[0] : vector<1xf32> -> f32 + %minimum_scale = vector.extract %minimum_scale_vector1[0] : vector<1xf32> -> f32 + %payload_field = index.div %q_word_index, %c4 : index + %payload_word = index.rem %q_word_index, %c4 : index + %q_word = vector.load %code_view[%payload_field, %row_lane, %payload_word] : view<8x64x4xi32> -> vector<1xi32> + %q_half = index.rem %bounded_group, %c2 : index + %q_shift_index = index.mul %q_half, %c4 : index + %q_shift_i32 = index.cast %q_shift_index : index to i32 + %q_shift = vector.splat %q_shift_i32 : vector<1xi32> + %shifted_q = vector.shrui %q_word, %q_shift : vector<1xi32> + %masked_q = vector.andi %shifted_q, %q4_mask : vector<1xi32> + %q_i8 = vector.bitcast %masked_q : vector<1xi32> to vector<4xi8> + %q_f32 = vector.uitofp %q_i8 : vector<4xi8> to vector<4xf32> + %negative_minimum_scale = scalar.negf %minimum_scale : f32 + %q0 = vector.extract %q_f32[0] : vector<4xf32> -> f32 + %q1 = vector.extract %q_f32[1] : vector<4xf32> -> f32 + %q2 = vector.extract %q_f32[2] : vector<4xf32> -> f32 + %q3 = vector.extract %q_f32[3] : vector<4xf32> -> f32 + %value0 = scalar.fmaf %q0, %d_scale, %negative_minimum_scale : f32 + %value1 = scalar.fmaf %q1, %d_scale, %negative_minimum_scale : f32 + %value2 = scalar.fmaf %q2, %d_scale, %negative_minimum_scale : f32 + %value3 = scalar.fmaf %q3, %d_scale, %negative_minimum_scale : f32 + %result = vector.from_elements %value0, %value1, %value2, %value3 : vector<4xf32> + func.return %result : vector<4xf32> +} + +// Q5_K uses the Q4_K affine header and a 32-byte fifth-bit plane. +func.def inline @ggml_q5k_f32_vector4(%weight: buffer, %row_byte_base: offset, %q5_block: index, %q5_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %block_bytes = index.constant 176 : offset + %scale_offset = index.constant 4 : offset + %qh_offset = index.constant 16 : offset + %code_offset = index.constant 48 : offset + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15 = vector.constant 15 : vector<1xi32> + %c48 = vector.constant 48 : vector<1xi32> + %q4_mask = vector.constant 252645135 : vector<1xi32> + %bounded_group = index.assume %q5_group [range(%q5_group, 0, 7)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q5_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %scale_byte_base = index.add %block_byte_base, %scale_offset : offset + %qh_byte_base = index.add %block_byte_base, %qh_offset : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %dm_view = buffer.view %weight[%block_byte_base] : buffer -> view<2xf16> + %scale_view = buffer.view %weight[%scale_byte_base] : buffer -> view<3xi32> + %qh_view = buffer.view %weight[%qh_byte_base] : buffer -> view<32xi8> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<32xi32> + %dm = vector.load %dm_view[%c0] : view<2xf16> -> vector<2xf16> + %scales = vector.load %scale_view[%c0] : view<3xi32> -> vector<3xi32> + %d_f16 = vector.extract %dm[0] : vector<2xf16> -> f16 + %dmin_f16 = vector.extract %dm[1] : vector<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %q_page0 = index.div %bounded_group, %c2 : index + %q_page = index.mul %q_page0, %c8 : index + %q_word_index0 = index.add %q_page, %bounded_packet : index + %q_word_index = index.assume %q_word_index0 [range(%q_word_index0, 0, 31)] : index + %is_low = index.cmp ult, %bounded_group, %c4 : index + %scale_lane = index.rem %bounded_group, %c4 : index + %scale_shift_index = index.mul %scale_lane, %c8 : index + %scale_shift_i32 = index.cast %scale_shift_index : index to i32 + %scale_shift = vector.splat %scale_shift_i32 : vector<1xi32> + %scale0_i32 = vector.extract %scales[0] : vector<3xi32> -> i32 + %scale1_i32 = vector.extract %scales[1] : vector<3xi32> -> i32 + %scale2_i32 = vector.extract %scales[2] : vector<3xi32> -> i32 + %scale0 = vector.splat %scale0_i32 : vector<1xi32> + %scale1 = vector.splat %scale1_i32 : vector<1xi32> + %scale2 = vector.splat %scale2_i32 : vector<1xi32> + %high_shift_i32 = scalar.addi %scale_shift_i32, %c2_i32 : i32 + %minimum_shift_i32 = scalar.addi %scale_shift_i32, %c4_i32 : i32 + %selected_scale_source = scf.select %is_low, %scale0, %scale2 : vector<1xi32> + %selected_minimum_source = scf.select %is_low, %scale1, %scale2 : vector<1xi32> + %selected_scale_high_shift_i32 = scf.select %is_low, %scale_shift_i32, %high_shift_i32 : i32 + %selected_minimum_low_shift_i32 = scf.select %is_low, %scale_shift_i32, %minimum_shift_i32 : i32 + %selected_scale_high_shift = vector.splat %selected_scale_high_shift_i32 : vector<1xi32> + %selected_minimum_low_shift = vector.splat %selected_minimum_low_shift_i32 : vector<1xi32> + %scale_low0 = vector.shrui %selected_scale_source, %scale_shift : vector<1xi32> + %scale_low = vector.andi %scale_low0, %c15 : vector<1xi32> + %scale_high0 = vector.shrui %scale0, %selected_scale_high_shift : vector<1xi32> + %scale_high = vector.andi %scale_high0, %c48 : vector<1xi32> + %scale = vector.ori %scale_low, %scale_high : vector<1xi32> + %minimum_low0 = vector.shrui %selected_minimum_source, %selected_minimum_low_shift : vector<1xi32> + %minimum_low = vector.andi %minimum_low0, %c15 : vector<1xi32> + %minimum_high0 = vector.shrui %scale1, %selected_scale_high_shift : vector<1xi32> + %minimum_high = vector.andi %minimum_high0, %c48 : vector<1xi32> + %minimum = vector.ori %minimum_low, %minimum_high : vector<1xi32> + %scale_f32 = vector.uitofp %scale : vector<1xi32> to vector<1xf32> + %minimum_f32 = vector.uitofp %minimum : vector<1xi32> to vector<1xf32> + %d_vector1 = vector.splat %d : vector<1xf32> + %dmin_vector1 = vector.splat %dmin : vector<1xf32> + %d_scale_vector1 = vector.mulf %d_vector1, %scale_f32 : vector<1xf32> + %minimum_scale_vector1 = vector.mulf %dmin_vector1, %minimum_f32 : vector<1xf32> + %d_scale = vector.extract %d_scale_vector1[0] : vector<1xf32> -> f32 + %minimum_scale = vector.extract %minimum_scale_vector1[0] : vector<1xf32> -> f32 + %q_word = vector.load %code_view[%q_word_index] : view<32xi32> -> vector<1xi32> + %q_half = index.rem %bounded_group, %c2 : index + %q_shift_index = index.mul %q_half, %c4 : index + %q_shift_i32 = index.cast %q_shift_index : index to i32 + %q_shift = vector.splat %q_shift_i32 : vector<1xi32> + %shifted_q = vector.shrui %q_word, %q_shift : vector<1xi32> + %masked_q = vector.andi %shifted_q, %q4_mask : vector<1xi32> + %q_i8 = vector.bitcast %masked_q : vector<1xi32> to vector<4xi8> + %q_f32 = vector.uitofp %q_i8 : vector<4xi8> to vector<4xf32> + %packet_base = index.mul %bounded_packet, %c4 : index + %qh_bytes = vector.load %qh_view[%packet_base] : view<32xi8> -> vector<4xi8> + %qh_shift_i32 = index.cast %bounded_group : index to i32 + %qh_bit = scalar.shli %c1_i32, %qh_shift_i32 : i32 + %qh0 = vector.extract %qh_bytes[0] : vector<4xi8> -> i8 + %qh1 = vector.extract %qh_bytes[1] : vector<4xi8> -> i8 + %qh2 = vector.extract %qh_bytes[2] : vector<4xi8> -> i8 + %qh3 = vector.extract %qh_bytes[3] : vector<4xi8> -> i8 + %high0 = func.call @ggml_q5k_high_bit_f32(%qh0, %qh_bit) : (i8, i32) -> (f32) + %high1 = func.call @ggml_q5k_high_bit_f32(%qh1, %qh_bit) : (i8, i32) -> (f32) + %high2 = func.call @ggml_q5k_high_bit_f32(%qh2, %qh_bit) : (i8, i32) -> (f32) + %high3 = func.call @ggml_q5k_high_bit_f32(%qh3, %qh_bit) : (i8, i32) -> (f32) + %q0_low = vector.extract %q_f32[0] : vector<4xf32> -> f32 + %q1_low = vector.extract %q_f32[1] : vector<4xf32> -> f32 + %q2_low = vector.extract %q_f32[2] : vector<4xf32> -> f32 + %q3_low = vector.extract %q_f32[3] : vector<4xf32> -> f32 + %q0 = scalar.addf %q0_low, %high0 : f32 + %q1 = scalar.addf %q1_low, %high1 : f32 + %q2 = scalar.addf %q2_low, %high2 : f32 + %q3 = scalar.addf %q3_low, %high3 : f32 + %negative_minimum_scale = scalar.negf %minimum_scale : f32 + %value0 = scalar.fmaf %q0, %d_scale, %negative_minimum_scale : f32 + %value1 = scalar.fmaf %q1, %d_scale, %negative_minimum_scale : f32 + %value2 = scalar.fmaf %q2, %d_scale, %negative_minimum_scale : f32 + %value3 = scalar.fmaf %q3, %d_scale, %negative_minimum_scale : f32 + %result = vector.from_elements %value0, %value1, %value2, %value3 : vector<4xf32> + func.return %result : vector<4xf32> +} + +func.def inline @ggml_q5k_f16_vector4(%weight: buffer, %row_byte_base: offset, %q5_block: index, %q5_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_q5k_f32_vector4(%weight, %row_byte_base, %q5_block, %q5_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +func.def inline @ggml_q3k_value_f32(%q_byte: i8, %hmask_byte: i8, %q_shift: i32, %hmask_bit: i32, %combined_scale: f32) -> (f32) { + %c0_i32 = scalar.constant 0 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %q_i32 = scalar.extui %q_byte : i8 to i32 + %q_shifted = scalar.shrui %q_i32, %q_shift : i32 + %q_low = scalar.andi %q_shifted, %c3_i32 : i32 + %hmask_i32 = scalar.extui %hmask_byte : i8 to i32 + %hmask_masked = scalar.andi %hmask_i32, %hmask_bit : i32 + %has_high_bit = scalar.cmpi ne, %hmask_masked, %c0_i32 : i32 + %subtrahend = scf.select %has_high_bit, %c0_i32, %c4_i32 : i32 + %q = scalar.subi %q_low, %subtrahend : i32 + %q_f32 = scalar.sitofp %q : i32 to f32 + %value = scalar.mulf %q_f32, %combined_scale : f32 + func.return %value : f32 +} + +func.def inline @ggml_q3k_f32_vector4(%weight: buffer, %row_byte_base: offset, %q3_block: index, %q3_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c32 = index.constant 32 : index + %block_bytes = index.constant 110 : offset + %code_offset = index.constant 32 : offset + %scale_offset = index.constant 96 : offset + %d_offset = index.constant 108 : offset + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c32_i32 = scalar.constant 32 : i32 + %bounded_group = index.assume %q3_group [range(%q3_group, 0, 7)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q3_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %scale_byte_base = index.add %block_byte_base, %scale_offset : offset + %d_byte_base = index.add %block_byte_base, %d_offset : offset + %d_view = buffer.view %weight[%d_byte_base] : buffer -> view<1xf16> + %hmask_view = buffer.view %weight[%block_byte_base] : buffer -> view<32xi8> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<64xi8> + %scale_view = buffer.view %weight[%scale_byte_base] : buffer -> view<12xi8> + %scale_subgroup = index.div %bounded_packet, %c4 : index + %scale_group_base = index.mul %bounded_group, %c2 : index + %scale_index0 = index.add %scale_group_base, %scale_subgroup : index + %scale_index = index.assume %scale_index0 [range(%scale_index0, 0, 15)] : index + %uses_high_nibble = index.cmp uge, %scale_index, %c8 : index + %scale_low_index0 = index.rem %scale_index, %c8 : index + %scale_low_index = index.assume %scale_low_index0 [range(%scale_low_index0, 0, 7)] : index + %scale_high_index_rem = index.rem %scale_index, %c4 : index + %scale_high_index0 = index.add %scale_high_index_rem, %c8 : index + %scale_high_index = index.assume %scale_high_index0 [range(%scale_high_index0, 8, 11)] : index + %scale_low_byte = view.load %scale_view[%scale_low_index] : view<12xi8> -> i8 + %scale_high_byte = view.load %scale_view[%scale_high_index] : view<12xi8> -> i8 + %scale_low_i32 = scalar.extui %scale_low_byte : i8 to i32 + %scale_high_i32 = scalar.extui %scale_high_byte : i8 to i32 + %scale_low_shift = scf.select %uses_high_nibble, %c4_i32, %c0_i32 : i32 + %scale_low_shifted = scalar.shrui %scale_low_i32, %scale_low_shift : i32 + %scale_low = scalar.andi %scale_low_shifted, %c15_i32 : i32 + %scale_group_quartile = index.div %scale_index, %c4 : index + %scale_high_shift_index = index.mul %scale_group_quartile, %c2 : index + %scale_high_shift = index.cast %scale_high_shift_index : index to i32 + %scale_high_shifted = scalar.shrui %scale_high_i32, %scale_high_shift : i32 + %scale_high_low = scalar.andi %scale_high_shifted, %c3_i32 : i32 + %scale_high = scalar.shli %scale_high_low, %c4_i32 : i32 + %scale_code = scalar.ori %scale_low, %scale_high : i32 + %centered_scale = scalar.subi %scale_code, %c32_i32 : i32 + %centered_scale_f32 = scalar.sitofp %centered_scale : i32 to f32 + %d_f16 = view.load %d_view[%c0] : view<1xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %combined_scale = scalar.mulf %d, %centered_scale_f32 : f32 + %group_in_half = index.rem %bounded_group, %c4 : index + %half = index.div %bounded_group, %c4 : index + %packet_base = index.mul %bounded_packet, %c4 : index + %code_half_base = index.mul %half, %c32 : index + %code_index0 = index.add %code_half_base, %packet_base : index + %code_index = index.assume %code_index0 [range(%code_index0, 0, 63)] : index + %hmask_index = index.assume %packet_base [range(%packet_base, 0, 28), mul(%packet_base, 4)] : index + %q_shift_index = index.mul %group_in_half, %c2 : index + %q_shift = index.cast %q_shift_index : index to i32 + %hmask_bit_shift = index.cast %bounded_group : index to i32 + %hmask_bit = scalar.shli %c1_i32, %hmask_bit_shift : i32 + %q_bytes = vector.load %code_view[%code_index] : view<64xi8> -> vector<4xi8> + %hmask_bytes = vector.load %hmask_view[%hmask_index] : view<32xi8> -> vector<4xi8> + %q0 = vector.extract %q_bytes[0] : vector<4xi8> -> i8 + %q1 = vector.extract %q_bytes[1] : vector<4xi8> -> i8 + %q2 = vector.extract %q_bytes[2] : vector<4xi8> -> i8 + %q3 = vector.extract %q_bytes[3] : vector<4xi8> -> i8 + %hm0 = vector.extract %hmask_bytes[0] : vector<4xi8> -> i8 + %hm1 = vector.extract %hmask_bytes[1] : vector<4xi8> -> i8 + %hm2 = vector.extract %hmask_bytes[2] : vector<4xi8> -> i8 + %hm3 = vector.extract %hmask_bytes[3] : vector<4xi8> -> i8 + %value0 = func.call @ggml_q3k_value_f32(%q0, %hm0, %q_shift, %hmask_bit, %combined_scale) : (i8, i8, i32, i32, f32) -> (f32) + %value1 = func.call @ggml_q3k_value_f32(%q1, %hm1, %q_shift, %hmask_bit, %combined_scale) : (i8, i8, i32, i32, f32) -> (f32) + %value2 = func.call @ggml_q3k_value_f32(%q2, %hm2, %q_shift, %hmask_bit, %combined_scale) : (i8, i8, i32, i32, f32) -> (f32) + %value3 = func.call @ggml_q3k_value_f32(%q3, %hm3, %q_shift, %hmask_bit, %combined_scale) : (i8, i8, i32, i32, f32) -> (f32) + %result = vector.from_elements %value0, %value1, %value2, %value3 : vector<4xf32> + func.return %result : vector<4xf32> +} + +func.def inline @ggml_q3k_f16_vector4(%weight: buffer, %row_byte_base: offset, %q3_block: index, %q3_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_q3k_f32_vector4(%weight, %row_byte_base, %q3_block, %q3_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +// Q2_K (84 bytes: scales[16], qs[64], d, dmin; ggml-quants.c dequantize_row_q2_K): the same value +// order as Q3_K. Group g, packet p hold values 128 (g / 4) + 32 (g % 4) + 4p + i, codes +// (qs[32 (g / 4) + 4p + i] >> 2 (g % 4)) & 3, scale byte sc = scales[2g + p / 4]; +// value = d * (sc & 15) * q - dmin * (sc >> 4). +func.def inline @ggml_q2k_f32_vector4(%weight: buffer, %row_byte_base: offset, %q2_block: index, %q2_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c32 = index.constant 32 : index + %block_bytes = index.constant 84 : offset + %code_offset = index.constant 16 : offset + %d_offset = index.constant 80 : offset + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15_i32 = scalar.constant 15 : i32 + %bounded_group = index.assume %q2_group [range(%q2_group, 0, 7)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q2_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %d_byte_base = index.add %block_byte_base, %d_offset : offset + %d_view = buffer.view %weight[%d_byte_base] : buffer -> view<2xf16> + %scale_view = buffer.view %weight[%block_byte_base] : buffer -> view<16xi8> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<64xi8> + %scale_subgroup = index.div %bounded_packet, %c4 : index + %scale_group_base = index.mul %bounded_group, %c2 : index + %scale_index0 = index.add %scale_group_base, %scale_subgroup : index + %scale_index = index.assume %scale_index0 [range(%scale_index0, 0, 15)] : index + %scale_byte = view.load %scale_view[%scale_index] : view<16xi8> -> i8 + %scale_i32 = scalar.extui %scale_byte : i8 to i32 + %scale_low = scalar.andi %scale_i32, %c15_i32 : i32 + %min_high = scalar.shrui %scale_i32, %c4_i32 : i32 + %scale_f32 = scalar.uitofp %scale_low : i32 to f32 + %min_f32 = scalar.uitofp %min_high : i32 to f32 + %c1 = index.constant 1 : index + %d_f16 = view.load %d_view[%c0] : view<2xf16> -> f16 + %dmin_f16 = view.load %d_view[%c1] : view<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %d_scale = scalar.mulf %d, %scale_f32 : f32 + %m_scale = scalar.mulf %dmin, %min_f32 : f32 + %neg_m_scale = scalar.negf %m_scale : f32 + %group_in_half = index.rem %bounded_group, %c4 : index + %half = index.div %bounded_group, %c4 : index + %packet_base = index.mul %bounded_packet, %c4 : index + %code_half_base = index.mul %half, %c32 : index + %code_index0 = index.add %code_half_base, %packet_base : index + %code_index = index.assume %code_index0 [range(%code_index0, 0, 60), mul(%code_index0, 4)] : index + %q_shift_index = index.mul %group_in_half, %c2 : index + %q_shift = index.cast %q_shift_index : index to i32 + %q_bytes = vector.load %code_view[%code_index] : view<64xi8> -> vector<4xi8> + %q_word_vector = vector.bitcast %q_bytes : vector<4xi8> to vector<1xi32> + %q_word = vector.extract %q_word_vector[0] : vector<1xi32> -> i32 + %q_shifted = scalar.shrui %q_word, %q_shift : i32 + %mask2 = scalar.constant 50529027 : i32 + %q_masked = scalar.andi %q_shifted, %mask2 : i32 + %q_masked_vector = vector.from_elements %q_masked : vector<1xi32> + %q_codes = vector.bitcast %q_masked_vector : vector<1xi32> to vector<4xi8> + %q_f32 = vector.uitofp %q_codes : vector<4xi8> to vector<4xf32> + %d_scale_v = vector.splat %d_scale : vector<4xf32> + %neg_m_v = vector.splat %neg_m_scale : vector<4xf32> + %result = vector.fmaf %q_f32, %d_scale_v, %neg_m_v : vector<4xf32> + func.return %result : vector<4xf32> +} + +func.def inline @ggml_q2k_f16_vector4(%weight: buffer, %row_byte_base: offset, %q2_block: index, %q2_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_q2k_f32_vector4(%weight, %row_byte_base, %q2_block, %q2_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + + +func.def inline @ggml_iq2s_grid_lookup_i32(%grid_index: i32, %word_index: i32) -> (i32) { + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c5_i32 = scalar.constant 5 : i32 + %c6_i32 = scalar.constant 6 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c9_i32 = scalar.constant 9 : i32 + %c10_i32 = scalar.constant 10 : i32 + %c11_i32 = scalar.constant 11 : i32 + %c12_i32 = scalar.constant 12 : i32 + %c13_i32 = scalar.constant 13 : i32 + %c14_i32 = scalar.constant 14 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c16_i32 = scalar.constant 16 : i32 + %c17_i32 = scalar.constant 17 : i32 + %c18_i32 = scalar.constant 18 : i32 + %c19_i32 = scalar.constant 19 : i32 + %c20_i32 = scalar.constant 20 : i32 + %c21_i32 = scalar.constant 21 : i32 + %c22_i32 = scalar.constant 22 : i32 + %c23_i32 = scalar.constant 23 : i32 + %c24_i32 = scalar.constant 24 : i32 + %c25_i32 = scalar.constant 25 : i32 + %c26_i32 = scalar.constant 26 : i32 + %c27_i32 = scalar.constant 27 : i32 + %c28_i32 = scalar.constant 28 : i32 + %c29_i32 = scalar.constant 29 : i32 + %c30_i32 = scalar.constant 30 : i32 + %c31_i32 = scalar.constant 31 : i32 + %chunk_i32 = scalar.shrui %grid_index, %c5_i32 : i32 + %lane_i32 = scalar.andi %grid_index, %c31_i32 : i32 + %codes = vector.from_elements %lane_i32 : vector<1xi32> + %is_chunk1 = scalar.cmpi eq, %chunk_i32, %c1_i32 : i32 + %is_chunk2 = scalar.cmpi eq, %chunk_i32, %c2_i32 : i32 + %is_chunk3 = scalar.cmpi eq, %chunk_i32, %c3_i32 : i32 + %is_chunk4 = scalar.cmpi eq, %chunk_i32, %c4_i32 : i32 + %is_chunk5 = scalar.cmpi eq, %chunk_i32, %c5_i32 : i32 + %is_chunk6 = scalar.cmpi eq, %chunk_i32, %c6_i32 : i32 + %is_chunk7 = scalar.cmpi eq, %chunk_i32, %c7_i32 : i32 + %is_chunk8 = scalar.cmpi eq, %chunk_i32, %c8_i32 : i32 + %is_chunk9 = scalar.cmpi eq, %chunk_i32, %c9_i32 : i32 + %is_chunk10 = scalar.cmpi eq, %chunk_i32, %c10_i32 : i32 + %is_chunk11 = scalar.cmpi eq, %chunk_i32, %c11_i32 : i32 + %is_chunk12 = scalar.cmpi eq, %chunk_i32, %c12_i32 : i32 + %is_chunk13 = scalar.cmpi eq, %chunk_i32, %c13_i32 : i32 + %is_chunk14 = scalar.cmpi eq, %chunk_i32, %c14_i32 : i32 + %is_chunk15 = scalar.cmpi eq, %chunk_i32, %c15_i32 : i32 + %is_chunk16 = scalar.cmpi eq, %chunk_i32, %c16_i32 : i32 + %is_chunk17 = scalar.cmpi eq, %chunk_i32, %c17_i32 : i32 + %is_chunk18 = scalar.cmpi eq, %chunk_i32, %c18_i32 : i32 + %is_chunk19 = scalar.cmpi eq, %chunk_i32, %c19_i32 : i32 + %is_chunk20 = scalar.cmpi eq, %chunk_i32, %c20_i32 : i32 + %is_chunk21 = scalar.cmpi eq, %chunk_i32, %c21_i32 : i32 + %is_chunk22 = scalar.cmpi eq, %chunk_i32, %c22_i32 : i32 + %is_chunk23 = scalar.cmpi eq, %chunk_i32, %c23_i32 : i32 + %is_chunk24 = scalar.cmpi eq, %chunk_i32, %c24_i32 : i32 + %is_chunk25 = scalar.cmpi eq, %chunk_i32, %c25_i32 : i32 + %is_chunk26 = scalar.cmpi eq, %chunk_i32, %c26_i32 : i32 + %is_chunk27 = scalar.cmpi eq, %chunk_i32, %c27_i32 : i32 + %is_chunk28 = scalar.cmpi eq, %chunk_i32, %c28_i32 : i32 + %is_chunk29 = scalar.cmpi eq, %chunk_i32, %c29_i32 : i32 + %is_chunk30 = scalar.cmpi eq, %chunk_i32, %c30_i32 : i32 + %is_chunk31 = scalar.cmpi eq, %chunk_i32, %c31_i32 : i32 + %iq2s_word0_low_0_0 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_0_1 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_0_2 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_0_3 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_0_4 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_0_5 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_0_6 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_0_7 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_0_8 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_0_9 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_0_10 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_0_11 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_0_12 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_0_13 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_0_14 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_0_15 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_0_16 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_0_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_0_18 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_0_19 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_0_20 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_0_21 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_0_22 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_0_23 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_0_24 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_0_25 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_0_26 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_0_27 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_0_28 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_0_29 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_0_30 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_0_31 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_0 = vector.from_elements %iq2s_word0_low_0_0, %iq2s_word0_low_0_1, %iq2s_word0_low_0_2, %iq2s_word0_low_0_3, %iq2s_word0_low_0_4, %iq2s_word0_low_0_5, %iq2s_word0_low_0_6, %iq2s_word0_low_0_7, %iq2s_word0_low_0_8, %iq2s_word0_low_0_9, %iq2s_word0_low_0_10, %iq2s_word0_low_0_11, %iq2s_word0_low_0_12, %iq2s_word0_low_0_13, %iq2s_word0_low_0_14, %iq2s_word0_low_0_15, %iq2s_word0_low_0_16, %iq2s_word0_low_0_17, %iq2s_word0_low_0_18, %iq2s_word0_low_0_19, %iq2s_word0_low_0_20, %iq2s_word0_low_0_21, %iq2s_word0_low_0_22, %iq2s_word0_low_0_23, %iq2s_word0_low_0_24, %iq2s_word0_low_0_25, %iq2s_word0_low_0_26, %iq2s_word0_low_0_27, %iq2s_word0_low_0_28, %iq2s_word0_low_0_29, %iq2s_word0_low_0_30, %iq2s_word0_low_0_31 : vector<32xf32> + %iq2s_word0_low_1_0 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_1_1 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_1_2 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_1_3 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_1_4 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_1_5 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_1_6 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_1_7 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_1_8 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_1_9 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_1_10 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_1_11 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_1_12 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_1_13 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_1_14 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_1_15 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_1_16 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_1_17 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_1_18 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_1_19 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_1_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_1_21 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_1_22 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_1_23 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_1_24 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_1_25 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_1_26 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_1_27 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_1_28 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_1_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_1_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_1_31 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_1 = vector.from_elements %iq2s_word0_low_1_0, %iq2s_word0_low_1_1, %iq2s_word0_low_1_2, %iq2s_word0_low_1_3, %iq2s_word0_low_1_4, %iq2s_word0_low_1_5, %iq2s_word0_low_1_6, %iq2s_word0_low_1_7, %iq2s_word0_low_1_8, %iq2s_word0_low_1_9, %iq2s_word0_low_1_10, %iq2s_word0_low_1_11, %iq2s_word0_low_1_12, %iq2s_word0_low_1_13, %iq2s_word0_low_1_14, %iq2s_word0_low_1_15, %iq2s_word0_low_1_16, %iq2s_word0_low_1_17, %iq2s_word0_low_1_18, %iq2s_word0_low_1_19, %iq2s_word0_low_1_20, %iq2s_word0_low_1_21, %iq2s_word0_low_1_22, %iq2s_word0_low_1_23, %iq2s_word0_low_1_24, %iq2s_word0_low_1_25, %iq2s_word0_low_1_26, %iq2s_word0_low_1_27, %iq2s_word0_low_1_28, %iq2s_word0_low_1_29, %iq2s_word0_low_1_30, %iq2s_word0_low_1_31 : vector<32xf32> + %iq2s_word0_low_2_0 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_2_1 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_2_2 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_2_3 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_2_4 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_2_5 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_2_6 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_2_7 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_2_8 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_2_9 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_2_10 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_2_11 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_2_12 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_2_13 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_2_14 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_2_15 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_2_16 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_2_17 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_2_18 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_2_19 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_2_20 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_2_21 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_2_22 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_2_23 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_2_24 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_2_25 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_2_26 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_2_27 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_2_28 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_2_29 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_2_30 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_2_31 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_2 = vector.from_elements %iq2s_word0_low_2_0, %iq2s_word0_low_2_1, %iq2s_word0_low_2_2, %iq2s_word0_low_2_3, %iq2s_word0_low_2_4, %iq2s_word0_low_2_5, %iq2s_word0_low_2_6, %iq2s_word0_low_2_7, %iq2s_word0_low_2_8, %iq2s_word0_low_2_9, %iq2s_word0_low_2_10, %iq2s_word0_low_2_11, %iq2s_word0_low_2_12, %iq2s_word0_low_2_13, %iq2s_word0_low_2_14, %iq2s_word0_low_2_15, %iq2s_word0_low_2_16, %iq2s_word0_low_2_17, %iq2s_word0_low_2_18, %iq2s_word0_low_2_19, %iq2s_word0_low_2_20, %iq2s_word0_low_2_21, %iq2s_word0_low_2_22, %iq2s_word0_low_2_23, %iq2s_word0_low_2_24, %iq2s_word0_low_2_25, %iq2s_word0_low_2_26, %iq2s_word0_low_2_27, %iq2s_word0_low_2_28, %iq2s_word0_low_2_29, %iq2s_word0_low_2_30, %iq2s_word0_low_2_31 : vector<32xf32> + %iq2s_word0_low_3_0 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_3_1 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_3_2 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_3_3 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_3_4 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_3_5 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_3_6 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_3_7 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_3_8 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_3_9 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_3_10 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_3_11 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_3_12 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_3_13 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_3_14 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_3_15 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_3_16 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_3_17 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_3_18 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_3_19 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_3_20 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_3_21 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_3_22 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_3_23 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_3_24 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_3_25 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_3_26 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_3_27 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_3_28 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_3_29 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_3_30 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_3_31 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_3 = vector.from_elements %iq2s_word0_low_3_0, %iq2s_word0_low_3_1, %iq2s_word0_low_3_2, %iq2s_word0_low_3_3, %iq2s_word0_low_3_4, %iq2s_word0_low_3_5, %iq2s_word0_low_3_6, %iq2s_word0_low_3_7, %iq2s_word0_low_3_8, %iq2s_word0_low_3_9, %iq2s_word0_low_3_10, %iq2s_word0_low_3_11, %iq2s_word0_low_3_12, %iq2s_word0_low_3_13, %iq2s_word0_low_3_14, %iq2s_word0_low_3_15, %iq2s_word0_low_3_16, %iq2s_word0_low_3_17, %iq2s_word0_low_3_18, %iq2s_word0_low_3_19, %iq2s_word0_low_3_20, %iq2s_word0_low_3_21, %iq2s_word0_low_3_22, %iq2s_word0_low_3_23, %iq2s_word0_low_3_24, %iq2s_word0_low_3_25, %iq2s_word0_low_3_26, %iq2s_word0_low_3_27, %iq2s_word0_low_3_28, %iq2s_word0_low_3_29, %iq2s_word0_low_3_30, %iq2s_word0_low_3_31 : vector<32xf32> + %iq2s_word0_low_4_0 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_4_1 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_4_2 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_4_3 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_4_4 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_4_5 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_4_6 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_4_7 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_4_8 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_4_9 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_4_10 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_4_11 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_4_12 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_4_13 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_4_14 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_4_15 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_4_16 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_4_17 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_4_18 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_4_19 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_4_20 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_4_21 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_4_22 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_4_23 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_4_24 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_4_25 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_4_26 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_4_27 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_4_28 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_4_29 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_4_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_4_31 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_4 = vector.from_elements %iq2s_word0_low_4_0, %iq2s_word0_low_4_1, %iq2s_word0_low_4_2, %iq2s_word0_low_4_3, %iq2s_word0_low_4_4, %iq2s_word0_low_4_5, %iq2s_word0_low_4_6, %iq2s_word0_low_4_7, %iq2s_word0_low_4_8, %iq2s_word0_low_4_9, %iq2s_word0_low_4_10, %iq2s_word0_low_4_11, %iq2s_word0_low_4_12, %iq2s_word0_low_4_13, %iq2s_word0_low_4_14, %iq2s_word0_low_4_15, %iq2s_word0_low_4_16, %iq2s_word0_low_4_17, %iq2s_word0_low_4_18, %iq2s_word0_low_4_19, %iq2s_word0_low_4_20, %iq2s_word0_low_4_21, %iq2s_word0_low_4_22, %iq2s_word0_low_4_23, %iq2s_word0_low_4_24, %iq2s_word0_low_4_25, %iq2s_word0_low_4_26, %iq2s_word0_low_4_27, %iq2s_word0_low_4_28, %iq2s_word0_low_4_29, %iq2s_word0_low_4_30, %iq2s_word0_low_4_31 : vector<32xf32> + %iq2s_word0_low_5_0 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_5_1 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_5_2 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_5_3 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_5_4 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_5_5 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_5_6 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_5_7 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_5_8 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_5_9 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_5_10 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_5_11 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_5_12 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_5_13 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_5_14 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_5_15 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_5_16 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_5_17 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_5_18 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_5_19 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_5_20 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_5_21 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_5_22 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_5_23 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_5_24 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_5_25 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_5_26 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_5_27 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_5_28 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_5_29 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_5_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_5_31 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_5 = vector.from_elements %iq2s_word0_low_5_0, %iq2s_word0_low_5_1, %iq2s_word0_low_5_2, %iq2s_word0_low_5_3, %iq2s_word0_low_5_4, %iq2s_word0_low_5_5, %iq2s_word0_low_5_6, %iq2s_word0_low_5_7, %iq2s_word0_low_5_8, %iq2s_word0_low_5_9, %iq2s_word0_low_5_10, %iq2s_word0_low_5_11, %iq2s_word0_low_5_12, %iq2s_word0_low_5_13, %iq2s_word0_low_5_14, %iq2s_word0_low_5_15, %iq2s_word0_low_5_16, %iq2s_word0_low_5_17, %iq2s_word0_low_5_18, %iq2s_word0_low_5_19, %iq2s_word0_low_5_20, %iq2s_word0_low_5_21, %iq2s_word0_low_5_22, %iq2s_word0_low_5_23, %iq2s_word0_low_5_24, %iq2s_word0_low_5_25, %iq2s_word0_low_5_26, %iq2s_word0_low_5_27, %iq2s_word0_low_5_28, %iq2s_word0_low_5_29, %iq2s_word0_low_5_30, %iq2s_word0_low_5_31 : vector<32xf32> + %iq2s_word0_low_6_0 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_6_1 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_6_2 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_6_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_6_4 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_6_5 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_6_6 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_6_7 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_6_8 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_6_9 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_6_10 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_6_11 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_6_12 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_6_13 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_6_14 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_6_15 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_6_16 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_6_17 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_6_18 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_6_19 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_6_20 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_6_21 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_6_22 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_6_23 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_6_24 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_6_25 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_6_26 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_6_27 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_6_28 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_6_29 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_6_30 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_6_31 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_6 = vector.from_elements %iq2s_word0_low_6_0, %iq2s_word0_low_6_1, %iq2s_word0_low_6_2, %iq2s_word0_low_6_3, %iq2s_word0_low_6_4, %iq2s_word0_low_6_5, %iq2s_word0_low_6_6, %iq2s_word0_low_6_7, %iq2s_word0_low_6_8, %iq2s_word0_low_6_9, %iq2s_word0_low_6_10, %iq2s_word0_low_6_11, %iq2s_word0_low_6_12, %iq2s_word0_low_6_13, %iq2s_word0_low_6_14, %iq2s_word0_low_6_15, %iq2s_word0_low_6_16, %iq2s_word0_low_6_17, %iq2s_word0_low_6_18, %iq2s_word0_low_6_19, %iq2s_word0_low_6_20, %iq2s_word0_low_6_21, %iq2s_word0_low_6_22, %iq2s_word0_low_6_23, %iq2s_word0_low_6_24, %iq2s_word0_low_6_25, %iq2s_word0_low_6_26, %iq2s_word0_low_6_27, %iq2s_word0_low_6_28, %iq2s_word0_low_6_29, %iq2s_word0_low_6_30, %iq2s_word0_low_6_31 : vector<32xf32> + %iq2s_word0_low_7_0 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_7_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_7_2 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_7_3 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_7_4 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_7_5 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_7_6 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_7_7 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_7_8 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_7_9 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_7_10 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_7_11 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_7_12 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_7_13 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_7_14 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_7_15 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_7_16 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_7_17 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_7_18 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_7_19 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_7_20 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_7_21 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_7_22 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_7_23 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_7_24 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_7_25 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_7_26 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_7_27 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_7_28 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_7_29 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_7_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_7_31 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_7 = vector.from_elements %iq2s_word0_low_7_0, %iq2s_word0_low_7_1, %iq2s_word0_low_7_2, %iq2s_word0_low_7_3, %iq2s_word0_low_7_4, %iq2s_word0_low_7_5, %iq2s_word0_low_7_6, %iq2s_word0_low_7_7, %iq2s_word0_low_7_8, %iq2s_word0_low_7_9, %iq2s_word0_low_7_10, %iq2s_word0_low_7_11, %iq2s_word0_low_7_12, %iq2s_word0_low_7_13, %iq2s_word0_low_7_14, %iq2s_word0_low_7_15, %iq2s_word0_low_7_16, %iq2s_word0_low_7_17, %iq2s_word0_low_7_18, %iq2s_word0_low_7_19, %iq2s_word0_low_7_20, %iq2s_word0_low_7_21, %iq2s_word0_low_7_22, %iq2s_word0_low_7_23, %iq2s_word0_low_7_24, %iq2s_word0_low_7_25, %iq2s_word0_low_7_26, %iq2s_word0_low_7_27, %iq2s_word0_low_7_28, %iq2s_word0_low_7_29, %iq2s_word0_low_7_30, %iq2s_word0_low_7_31 : vector<32xf32> + %iq2s_word0_low_8_0 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_8_1 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_8_2 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_8_3 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_8_4 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_8_5 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_8_6 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_8_7 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_8_8 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_8_9 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_8_10 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_8_11 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_8_12 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_8_13 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_8_14 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_8_15 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_8_16 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_8_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_8_18 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_8_19 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_8_20 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_8_21 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_8_22 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_8_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_8_24 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_8_25 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_8_26 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_8_27 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_8_28 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_8_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_8_30 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_8_31 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_8 = vector.from_elements %iq2s_word0_low_8_0, %iq2s_word0_low_8_1, %iq2s_word0_low_8_2, %iq2s_word0_low_8_3, %iq2s_word0_low_8_4, %iq2s_word0_low_8_5, %iq2s_word0_low_8_6, %iq2s_word0_low_8_7, %iq2s_word0_low_8_8, %iq2s_word0_low_8_9, %iq2s_word0_low_8_10, %iq2s_word0_low_8_11, %iq2s_word0_low_8_12, %iq2s_word0_low_8_13, %iq2s_word0_low_8_14, %iq2s_word0_low_8_15, %iq2s_word0_low_8_16, %iq2s_word0_low_8_17, %iq2s_word0_low_8_18, %iq2s_word0_low_8_19, %iq2s_word0_low_8_20, %iq2s_word0_low_8_21, %iq2s_word0_low_8_22, %iq2s_word0_low_8_23, %iq2s_word0_low_8_24, %iq2s_word0_low_8_25, %iq2s_word0_low_8_26, %iq2s_word0_low_8_27, %iq2s_word0_low_8_28, %iq2s_word0_low_8_29, %iq2s_word0_low_8_30, %iq2s_word0_low_8_31 : vector<32xf32> + %iq2s_word0_low_9_0 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_9_1 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_9_2 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_9_3 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_9_4 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_9_5 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_9_6 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_9_7 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_9_8 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_9_9 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_9_10 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_9_11 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_9_12 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_9_13 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_9_14 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_9_15 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_9_16 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_9_17 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_9_18 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_9_19 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_9_20 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_9_21 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_9_22 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_9_23 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_9_24 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_9_25 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_9_26 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_9_27 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_9_28 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_9_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_9_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_9_31 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_9 = vector.from_elements %iq2s_word0_low_9_0, %iq2s_word0_low_9_1, %iq2s_word0_low_9_2, %iq2s_word0_low_9_3, %iq2s_word0_low_9_4, %iq2s_word0_low_9_5, %iq2s_word0_low_9_6, %iq2s_word0_low_9_7, %iq2s_word0_low_9_8, %iq2s_word0_low_9_9, %iq2s_word0_low_9_10, %iq2s_word0_low_9_11, %iq2s_word0_low_9_12, %iq2s_word0_low_9_13, %iq2s_word0_low_9_14, %iq2s_word0_low_9_15, %iq2s_word0_low_9_16, %iq2s_word0_low_9_17, %iq2s_word0_low_9_18, %iq2s_word0_low_9_19, %iq2s_word0_low_9_20, %iq2s_word0_low_9_21, %iq2s_word0_low_9_22, %iq2s_word0_low_9_23, %iq2s_word0_low_9_24, %iq2s_word0_low_9_25, %iq2s_word0_low_9_26, %iq2s_word0_low_9_27, %iq2s_word0_low_9_28, %iq2s_word0_low_9_29, %iq2s_word0_low_9_30, %iq2s_word0_low_9_31 : vector<32xf32> + %iq2s_word0_low_10_0 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_10_1 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_10_2 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_10_3 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_10_4 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_10_5 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_10_6 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_10_7 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_10_8 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_10_9 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_10_10 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_10_11 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_10_12 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_10_13 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_10_14 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_10_15 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_10_16 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_10_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_10_18 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_10_19 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_10_20 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_10_21 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_10_22 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_10_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_10_24 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_10_25 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_10_26 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_10_27 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_10_28 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_10_29 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_10_30 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_10_31 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_10 = vector.from_elements %iq2s_word0_low_10_0, %iq2s_word0_low_10_1, %iq2s_word0_low_10_2, %iq2s_word0_low_10_3, %iq2s_word0_low_10_4, %iq2s_word0_low_10_5, %iq2s_word0_low_10_6, %iq2s_word0_low_10_7, %iq2s_word0_low_10_8, %iq2s_word0_low_10_9, %iq2s_word0_low_10_10, %iq2s_word0_low_10_11, %iq2s_word0_low_10_12, %iq2s_word0_low_10_13, %iq2s_word0_low_10_14, %iq2s_word0_low_10_15, %iq2s_word0_low_10_16, %iq2s_word0_low_10_17, %iq2s_word0_low_10_18, %iq2s_word0_low_10_19, %iq2s_word0_low_10_20, %iq2s_word0_low_10_21, %iq2s_word0_low_10_22, %iq2s_word0_low_10_23, %iq2s_word0_low_10_24, %iq2s_word0_low_10_25, %iq2s_word0_low_10_26, %iq2s_word0_low_10_27, %iq2s_word0_low_10_28, %iq2s_word0_low_10_29, %iq2s_word0_low_10_30, %iq2s_word0_low_10_31 : vector<32xf32> + %iq2s_word0_low_11_0 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_11_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_11_2 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_11_3 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_11_4 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_11_5 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_11_6 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_11_7 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_11_8 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_11_9 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_11_10 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_11_11 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_11_12 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_11_13 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_11_14 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_11_15 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_11_16 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_11_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_11_18 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_11_19 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_11_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_11_21 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_11_22 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_11_23 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_11_24 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_11_25 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_11_26 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_11_27 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_11_28 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_11_29 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_11_30 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_11_31 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_11 = vector.from_elements %iq2s_word0_low_11_0, %iq2s_word0_low_11_1, %iq2s_word0_low_11_2, %iq2s_word0_low_11_3, %iq2s_word0_low_11_4, %iq2s_word0_low_11_5, %iq2s_word0_low_11_6, %iq2s_word0_low_11_7, %iq2s_word0_low_11_8, %iq2s_word0_low_11_9, %iq2s_word0_low_11_10, %iq2s_word0_low_11_11, %iq2s_word0_low_11_12, %iq2s_word0_low_11_13, %iq2s_word0_low_11_14, %iq2s_word0_low_11_15, %iq2s_word0_low_11_16, %iq2s_word0_low_11_17, %iq2s_word0_low_11_18, %iq2s_word0_low_11_19, %iq2s_word0_low_11_20, %iq2s_word0_low_11_21, %iq2s_word0_low_11_22, %iq2s_word0_low_11_23, %iq2s_word0_low_11_24, %iq2s_word0_low_11_25, %iq2s_word0_low_11_26, %iq2s_word0_low_11_27, %iq2s_word0_low_11_28, %iq2s_word0_low_11_29, %iq2s_word0_low_11_30, %iq2s_word0_low_11_31 : vector<32xf32> + %iq2s_word0_low_12_0 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_12_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_12_2 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_12_3 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_12_4 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_12_5 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_12_6 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_12_7 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_12_8 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_12_9 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_12_10 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_12_11 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_12_12 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_12_13 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_12_14 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_12_15 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_12_16 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_12_17 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_12_18 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_12_19 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_12_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_12_21 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_12_22 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_12_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_12_24 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_12_25 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_12_26 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_12_27 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_12_28 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_12_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_12_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_12_31 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_12 = vector.from_elements %iq2s_word0_low_12_0, %iq2s_word0_low_12_1, %iq2s_word0_low_12_2, %iq2s_word0_low_12_3, %iq2s_word0_low_12_4, %iq2s_word0_low_12_5, %iq2s_word0_low_12_6, %iq2s_word0_low_12_7, %iq2s_word0_low_12_8, %iq2s_word0_low_12_9, %iq2s_word0_low_12_10, %iq2s_word0_low_12_11, %iq2s_word0_low_12_12, %iq2s_word0_low_12_13, %iq2s_word0_low_12_14, %iq2s_word0_low_12_15, %iq2s_word0_low_12_16, %iq2s_word0_low_12_17, %iq2s_word0_low_12_18, %iq2s_word0_low_12_19, %iq2s_word0_low_12_20, %iq2s_word0_low_12_21, %iq2s_word0_low_12_22, %iq2s_word0_low_12_23, %iq2s_word0_low_12_24, %iq2s_word0_low_12_25, %iq2s_word0_low_12_26, %iq2s_word0_low_12_27, %iq2s_word0_low_12_28, %iq2s_word0_low_12_29, %iq2s_word0_low_12_30, %iq2s_word0_low_12_31 : vector<32xf32> + %iq2s_word0_low_13_0 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_13_1 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_13_2 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_13_3 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_13_4 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_13_5 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_13_6 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_13_7 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_13_8 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_13_9 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_13_10 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_13_11 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_13_12 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_13_13 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_13_14 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_13_15 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_13_16 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_13_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_13_18 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_13_19 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_13_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_13_21 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_13_22 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_13_23 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_13_24 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_13_25 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_13_26 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_13_27 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_13_28 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_13_29 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_13_30 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_13_31 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_13 = vector.from_elements %iq2s_word0_low_13_0, %iq2s_word0_low_13_1, %iq2s_word0_low_13_2, %iq2s_word0_low_13_3, %iq2s_word0_low_13_4, %iq2s_word0_low_13_5, %iq2s_word0_low_13_6, %iq2s_word0_low_13_7, %iq2s_word0_low_13_8, %iq2s_word0_low_13_9, %iq2s_word0_low_13_10, %iq2s_word0_low_13_11, %iq2s_word0_low_13_12, %iq2s_word0_low_13_13, %iq2s_word0_low_13_14, %iq2s_word0_low_13_15, %iq2s_word0_low_13_16, %iq2s_word0_low_13_17, %iq2s_word0_low_13_18, %iq2s_word0_low_13_19, %iq2s_word0_low_13_20, %iq2s_word0_low_13_21, %iq2s_word0_low_13_22, %iq2s_word0_low_13_23, %iq2s_word0_low_13_24, %iq2s_word0_low_13_25, %iq2s_word0_low_13_26, %iq2s_word0_low_13_27, %iq2s_word0_low_13_28, %iq2s_word0_low_13_29, %iq2s_word0_low_13_30, %iq2s_word0_low_13_31 : vector<32xf32> + %iq2s_word0_low_14_0 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_14_1 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_14_2 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_14_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_14_4 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_14_5 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_14_6 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_14_7 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_14_8 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_14_9 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_14_10 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_14_11 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_14_12 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_14_13 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_14_14 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_14_15 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_14_16 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_14_17 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_14_18 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_14_19 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_14_20 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_14_21 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_14_22 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_14_23 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_14_24 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_14_25 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_14_26 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_14_27 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_14_28 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_14_29 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_14_30 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_14_31 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_14 = vector.from_elements %iq2s_word0_low_14_0, %iq2s_word0_low_14_1, %iq2s_word0_low_14_2, %iq2s_word0_low_14_3, %iq2s_word0_low_14_4, %iq2s_word0_low_14_5, %iq2s_word0_low_14_6, %iq2s_word0_low_14_7, %iq2s_word0_low_14_8, %iq2s_word0_low_14_9, %iq2s_word0_low_14_10, %iq2s_word0_low_14_11, %iq2s_word0_low_14_12, %iq2s_word0_low_14_13, %iq2s_word0_low_14_14, %iq2s_word0_low_14_15, %iq2s_word0_low_14_16, %iq2s_word0_low_14_17, %iq2s_word0_low_14_18, %iq2s_word0_low_14_19, %iq2s_word0_low_14_20, %iq2s_word0_low_14_21, %iq2s_word0_low_14_22, %iq2s_word0_low_14_23, %iq2s_word0_low_14_24, %iq2s_word0_low_14_25, %iq2s_word0_low_14_26, %iq2s_word0_low_14_27, %iq2s_word0_low_14_28, %iq2s_word0_low_14_29, %iq2s_word0_low_14_30, %iq2s_word0_low_14_31 : vector<32xf32> + %iq2s_word0_low_15_0 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_15_1 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_15_2 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_15_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_15_4 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_15_5 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_15_6 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_15_7 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_15_8 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_15_9 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_15_10 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_15_11 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_15_12 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_15_13 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_15_14 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_15_15 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_15_16 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_15_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_15_18 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_15_19 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_15_20 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_15_21 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_15_22 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_15_23 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_15_24 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_15_25 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_15_26 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_15_27 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_15_28 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_15_29 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_15_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_15_31 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_15 = vector.from_elements %iq2s_word0_low_15_0, %iq2s_word0_low_15_1, %iq2s_word0_low_15_2, %iq2s_word0_low_15_3, %iq2s_word0_low_15_4, %iq2s_word0_low_15_5, %iq2s_word0_low_15_6, %iq2s_word0_low_15_7, %iq2s_word0_low_15_8, %iq2s_word0_low_15_9, %iq2s_word0_low_15_10, %iq2s_word0_low_15_11, %iq2s_word0_low_15_12, %iq2s_word0_low_15_13, %iq2s_word0_low_15_14, %iq2s_word0_low_15_15, %iq2s_word0_low_15_16, %iq2s_word0_low_15_17, %iq2s_word0_low_15_18, %iq2s_word0_low_15_19, %iq2s_word0_low_15_20, %iq2s_word0_low_15_21, %iq2s_word0_low_15_22, %iq2s_word0_low_15_23, %iq2s_word0_low_15_24, %iq2s_word0_low_15_25, %iq2s_word0_low_15_26, %iq2s_word0_low_15_27, %iq2s_word0_low_15_28, %iq2s_word0_low_15_29, %iq2s_word0_low_15_30, %iq2s_word0_low_15_31 : vector<32xf32> + %iq2s_word0_low_16_0 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_16_1 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_16_2 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_16_3 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_16_4 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_16_5 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_16_6 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_16_7 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_16_8 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_16_9 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_16_10 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_16_11 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_16_12 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_16_13 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_16_14 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_16_15 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_16_16 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_16_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_16_18 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_16_19 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_16_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_16_21 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_16_22 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_16_23 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_16_24 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_16_25 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_16_26 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_16_27 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_16_28 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_16_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_16_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_16_31 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_16 = vector.from_elements %iq2s_word0_low_16_0, %iq2s_word0_low_16_1, %iq2s_word0_low_16_2, %iq2s_word0_low_16_3, %iq2s_word0_low_16_4, %iq2s_word0_low_16_5, %iq2s_word0_low_16_6, %iq2s_word0_low_16_7, %iq2s_word0_low_16_8, %iq2s_word0_low_16_9, %iq2s_word0_low_16_10, %iq2s_word0_low_16_11, %iq2s_word0_low_16_12, %iq2s_word0_low_16_13, %iq2s_word0_low_16_14, %iq2s_word0_low_16_15, %iq2s_word0_low_16_16, %iq2s_word0_low_16_17, %iq2s_word0_low_16_18, %iq2s_word0_low_16_19, %iq2s_word0_low_16_20, %iq2s_word0_low_16_21, %iq2s_word0_low_16_22, %iq2s_word0_low_16_23, %iq2s_word0_low_16_24, %iq2s_word0_low_16_25, %iq2s_word0_low_16_26, %iq2s_word0_low_16_27, %iq2s_word0_low_16_28, %iq2s_word0_low_16_29, %iq2s_word0_low_16_30, %iq2s_word0_low_16_31 : vector<32xf32> + %iq2s_word0_low_17_0 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_17_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_17_2 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_17_3 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_17_4 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_17_5 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_17_6 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_17_7 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_17_8 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_17_9 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_17_10 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_17_11 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_17_12 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_17_13 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_17_14 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_17_15 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_17_16 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_17_17 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_17_18 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_17_19 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_17_20 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_17_21 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_17_22 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_17_23 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_17_24 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_17_25 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_17_26 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_17_27 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_17_28 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_17_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_17_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_17_31 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_17 = vector.from_elements %iq2s_word0_low_17_0, %iq2s_word0_low_17_1, %iq2s_word0_low_17_2, %iq2s_word0_low_17_3, %iq2s_word0_low_17_4, %iq2s_word0_low_17_5, %iq2s_word0_low_17_6, %iq2s_word0_low_17_7, %iq2s_word0_low_17_8, %iq2s_word0_low_17_9, %iq2s_word0_low_17_10, %iq2s_word0_low_17_11, %iq2s_word0_low_17_12, %iq2s_word0_low_17_13, %iq2s_word0_low_17_14, %iq2s_word0_low_17_15, %iq2s_word0_low_17_16, %iq2s_word0_low_17_17, %iq2s_word0_low_17_18, %iq2s_word0_low_17_19, %iq2s_word0_low_17_20, %iq2s_word0_low_17_21, %iq2s_word0_low_17_22, %iq2s_word0_low_17_23, %iq2s_word0_low_17_24, %iq2s_word0_low_17_25, %iq2s_word0_low_17_26, %iq2s_word0_low_17_27, %iq2s_word0_low_17_28, %iq2s_word0_low_17_29, %iq2s_word0_low_17_30, %iq2s_word0_low_17_31 : vector<32xf32> + %iq2s_word0_low_18_0 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_18_1 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_18_2 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_18_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_18_4 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_18_5 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_18_6 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_18_7 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_18_8 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_18_9 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_18_10 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_18_11 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_18_12 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_18_13 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_18_14 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_18_15 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_18_16 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_18_17 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_18_18 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_18_19 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_18_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_18_21 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_18_22 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_18_23 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_18_24 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_18_25 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_18_26 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_18_27 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_18_28 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_18_29 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_18_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_18_31 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_18 = vector.from_elements %iq2s_word0_low_18_0, %iq2s_word0_low_18_1, %iq2s_word0_low_18_2, %iq2s_word0_low_18_3, %iq2s_word0_low_18_4, %iq2s_word0_low_18_5, %iq2s_word0_low_18_6, %iq2s_word0_low_18_7, %iq2s_word0_low_18_8, %iq2s_word0_low_18_9, %iq2s_word0_low_18_10, %iq2s_word0_low_18_11, %iq2s_word0_low_18_12, %iq2s_word0_low_18_13, %iq2s_word0_low_18_14, %iq2s_word0_low_18_15, %iq2s_word0_low_18_16, %iq2s_word0_low_18_17, %iq2s_word0_low_18_18, %iq2s_word0_low_18_19, %iq2s_word0_low_18_20, %iq2s_word0_low_18_21, %iq2s_word0_low_18_22, %iq2s_word0_low_18_23, %iq2s_word0_low_18_24, %iq2s_word0_low_18_25, %iq2s_word0_low_18_26, %iq2s_word0_low_18_27, %iq2s_word0_low_18_28, %iq2s_word0_low_18_29, %iq2s_word0_low_18_30, %iq2s_word0_low_18_31 : vector<32xf32> + %iq2s_word0_low_19_0 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_19_1 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_19_2 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_19_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_19_4 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_19_5 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_19_6 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_19_7 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_19_8 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_19_9 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_19_10 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_19_11 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_19_12 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_19_13 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_19_14 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_19_15 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_19_16 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_19_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_19_18 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_19_19 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_19_20 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_19_21 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_19_22 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_19_23 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_19_24 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_19_25 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_19_26 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_19_27 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_19_28 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_19_29 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_19_30 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_19_31 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_19 = vector.from_elements %iq2s_word0_low_19_0, %iq2s_word0_low_19_1, %iq2s_word0_low_19_2, %iq2s_word0_low_19_3, %iq2s_word0_low_19_4, %iq2s_word0_low_19_5, %iq2s_word0_low_19_6, %iq2s_word0_low_19_7, %iq2s_word0_low_19_8, %iq2s_word0_low_19_9, %iq2s_word0_low_19_10, %iq2s_word0_low_19_11, %iq2s_word0_low_19_12, %iq2s_word0_low_19_13, %iq2s_word0_low_19_14, %iq2s_word0_low_19_15, %iq2s_word0_low_19_16, %iq2s_word0_low_19_17, %iq2s_word0_low_19_18, %iq2s_word0_low_19_19, %iq2s_word0_low_19_20, %iq2s_word0_low_19_21, %iq2s_word0_low_19_22, %iq2s_word0_low_19_23, %iq2s_word0_low_19_24, %iq2s_word0_low_19_25, %iq2s_word0_low_19_26, %iq2s_word0_low_19_27, %iq2s_word0_low_19_28, %iq2s_word0_low_19_29, %iq2s_word0_low_19_30, %iq2s_word0_low_19_31 : vector<32xf32> + %iq2s_word0_low_20_0 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_20_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_20_2 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_20_3 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_20_4 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_20_5 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_20_6 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_20_7 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_20_8 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_20_9 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_20_10 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_20_11 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_20_12 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_20_13 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_20_14 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_20_15 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_20_16 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_20_17 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_20_18 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_20_19 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_20_20 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_20_21 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_20_22 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_20_23 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_20_24 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_20_25 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_20_26 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_20_27 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_20_28 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_20_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_20_30 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_20_31 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_20 = vector.from_elements %iq2s_word0_low_20_0, %iq2s_word0_low_20_1, %iq2s_word0_low_20_2, %iq2s_word0_low_20_3, %iq2s_word0_low_20_4, %iq2s_word0_low_20_5, %iq2s_word0_low_20_6, %iq2s_word0_low_20_7, %iq2s_word0_low_20_8, %iq2s_word0_low_20_9, %iq2s_word0_low_20_10, %iq2s_word0_low_20_11, %iq2s_word0_low_20_12, %iq2s_word0_low_20_13, %iq2s_word0_low_20_14, %iq2s_word0_low_20_15, %iq2s_word0_low_20_16, %iq2s_word0_low_20_17, %iq2s_word0_low_20_18, %iq2s_word0_low_20_19, %iq2s_word0_low_20_20, %iq2s_word0_low_20_21, %iq2s_word0_low_20_22, %iq2s_word0_low_20_23, %iq2s_word0_low_20_24, %iq2s_word0_low_20_25, %iq2s_word0_low_20_26, %iq2s_word0_low_20_27, %iq2s_word0_low_20_28, %iq2s_word0_low_20_29, %iq2s_word0_low_20_30, %iq2s_word0_low_20_31 : vector<32xf32> + %iq2s_word0_low_21_0 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_21_1 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_21_2 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_21_3 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_21_4 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_21_5 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_21_6 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_21_7 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_21_8 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_21_9 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_21_10 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_21_11 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_21_12 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_21_13 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_21_14 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_21_15 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_21_16 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_21_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_21_18 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_21_19 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_21_20 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_21_21 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_21_22 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_21_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_21_24 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_21_25 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_21_26 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_21_27 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_21_28 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_21_29 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_21_30 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_21_31 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_21 = vector.from_elements %iq2s_word0_low_21_0, %iq2s_word0_low_21_1, %iq2s_word0_low_21_2, %iq2s_word0_low_21_3, %iq2s_word0_low_21_4, %iq2s_word0_low_21_5, %iq2s_word0_low_21_6, %iq2s_word0_low_21_7, %iq2s_word0_low_21_8, %iq2s_word0_low_21_9, %iq2s_word0_low_21_10, %iq2s_word0_low_21_11, %iq2s_word0_low_21_12, %iq2s_word0_low_21_13, %iq2s_word0_low_21_14, %iq2s_word0_low_21_15, %iq2s_word0_low_21_16, %iq2s_word0_low_21_17, %iq2s_word0_low_21_18, %iq2s_word0_low_21_19, %iq2s_word0_low_21_20, %iq2s_word0_low_21_21, %iq2s_word0_low_21_22, %iq2s_word0_low_21_23, %iq2s_word0_low_21_24, %iq2s_word0_low_21_25, %iq2s_word0_low_21_26, %iq2s_word0_low_21_27, %iq2s_word0_low_21_28, %iq2s_word0_low_21_29, %iq2s_word0_low_21_30, %iq2s_word0_low_21_31 : vector<32xf32> + %iq2s_word0_low_22_0 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_22_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_22_2 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_22_3 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_22_4 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_22_5 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_22_6 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_22_7 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_22_8 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_22_9 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_22_10 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_22_11 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_22_12 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_22_13 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_22_14 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_22_15 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_22_16 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_22_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_22_18 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_22_19 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_22_20 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_22_21 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_22_22 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_22_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_22_24 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_22_25 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_22_26 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_22_27 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_22_28 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_22_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_22_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_22_31 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_22 = vector.from_elements %iq2s_word0_low_22_0, %iq2s_word0_low_22_1, %iq2s_word0_low_22_2, %iq2s_word0_low_22_3, %iq2s_word0_low_22_4, %iq2s_word0_low_22_5, %iq2s_word0_low_22_6, %iq2s_word0_low_22_7, %iq2s_word0_low_22_8, %iq2s_word0_low_22_9, %iq2s_word0_low_22_10, %iq2s_word0_low_22_11, %iq2s_word0_low_22_12, %iq2s_word0_low_22_13, %iq2s_word0_low_22_14, %iq2s_word0_low_22_15, %iq2s_word0_low_22_16, %iq2s_word0_low_22_17, %iq2s_word0_low_22_18, %iq2s_word0_low_22_19, %iq2s_word0_low_22_20, %iq2s_word0_low_22_21, %iq2s_word0_low_22_22, %iq2s_word0_low_22_23, %iq2s_word0_low_22_24, %iq2s_word0_low_22_25, %iq2s_word0_low_22_26, %iq2s_word0_low_22_27, %iq2s_word0_low_22_28, %iq2s_word0_low_22_29, %iq2s_word0_low_22_30, %iq2s_word0_low_22_31 : vector<32xf32> + %iq2s_word0_low_23_0 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_23_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_23_2 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_23_3 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_23_4 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_23_5 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_23_6 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_23_7 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_23_8 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_23_9 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_23_10 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_23_11 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_23_12 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_23_13 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_23_14 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_23_15 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_23_16 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_23_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_23_18 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_23_19 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_23_20 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_23_21 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_23_22 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_23_23 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_23_24 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_23_25 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_23_26 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_23_27 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_23_28 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_23_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_23_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_23_31 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_23 = vector.from_elements %iq2s_word0_low_23_0, %iq2s_word0_low_23_1, %iq2s_word0_low_23_2, %iq2s_word0_low_23_3, %iq2s_word0_low_23_4, %iq2s_word0_low_23_5, %iq2s_word0_low_23_6, %iq2s_word0_low_23_7, %iq2s_word0_low_23_8, %iq2s_word0_low_23_9, %iq2s_word0_low_23_10, %iq2s_word0_low_23_11, %iq2s_word0_low_23_12, %iq2s_word0_low_23_13, %iq2s_word0_low_23_14, %iq2s_word0_low_23_15, %iq2s_word0_low_23_16, %iq2s_word0_low_23_17, %iq2s_word0_low_23_18, %iq2s_word0_low_23_19, %iq2s_word0_low_23_20, %iq2s_word0_low_23_21, %iq2s_word0_low_23_22, %iq2s_word0_low_23_23, %iq2s_word0_low_23_24, %iq2s_word0_low_23_25, %iq2s_word0_low_23_26, %iq2s_word0_low_23_27, %iq2s_word0_low_23_28, %iq2s_word0_low_23_29, %iq2s_word0_low_23_30, %iq2s_word0_low_23_31 : vector<32xf32> + %iq2s_word0_low_24_0 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_24_1 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_24_2 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_24_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_24_4 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_24_5 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_24_6 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_24_7 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_24_8 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_24_9 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_24_10 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_24_11 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_24_12 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_24_13 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_24_14 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_24_15 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_24_16 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_24_17 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_24_18 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_24_19 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_24_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_24_21 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_24_22 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_24_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_24_24 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_24_25 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_24_26 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_24_27 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_24_28 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_24_29 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_24_30 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_24_31 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_24 = vector.from_elements %iq2s_word0_low_24_0, %iq2s_word0_low_24_1, %iq2s_word0_low_24_2, %iq2s_word0_low_24_3, %iq2s_word0_low_24_4, %iq2s_word0_low_24_5, %iq2s_word0_low_24_6, %iq2s_word0_low_24_7, %iq2s_word0_low_24_8, %iq2s_word0_low_24_9, %iq2s_word0_low_24_10, %iq2s_word0_low_24_11, %iq2s_word0_low_24_12, %iq2s_word0_low_24_13, %iq2s_word0_low_24_14, %iq2s_word0_low_24_15, %iq2s_word0_low_24_16, %iq2s_word0_low_24_17, %iq2s_word0_low_24_18, %iq2s_word0_low_24_19, %iq2s_word0_low_24_20, %iq2s_word0_low_24_21, %iq2s_word0_low_24_22, %iq2s_word0_low_24_23, %iq2s_word0_low_24_24, %iq2s_word0_low_24_25, %iq2s_word0_low_24_26, %iq2s_word0_low_24_27, %iq2s_word0_low_24_28, %iq2s_word0_low_24_29, %iq2s_word0_low_24_30, %iq2s_word0_low_24_31 : vector<32xf32> + %iq2s_word0_low_25_0 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_25_1 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_25_2 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_25_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_25_4 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_25_5 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_25_6 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_25_7 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_25_8 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_25_9 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_25_10 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_25_11 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_25_12 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_25_13 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_25_14 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_25_15 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_25_16 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_25_17 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_25_18 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_25_19 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_25_20 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_25_21 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_25_22 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_25_23 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_25_24 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_25_25 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_25_26 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_25_27 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_25_28 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_25_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_25_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_25_31 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_25 = vector.from_elements %iq2s_word0_low_25_0, %iq2s_word0_low_25_1, %iq2s_word0_low_25_2, %iq2s_word0_low_25_3, %iq2s_word0_low_25_4, %iq2s_word0_low_25_5, %iq2s_word0_low_25_6, %iq2s_word0_low_25_7, %iq2s_word0_low_25_8, %iq2s_word0_low_25_9, %iq2s_word0_low_25_10, %iq2s_word0_low_25_11, %iq2s_word0_low_25_12, %iq2s_word0_low_25_13, %iq2s_word0_low_25_14, %iq2s_word0_low_25_15, %iq2s_word0_low_25_16, %iq2s_word0_low_25_17, %iq2s_word0_low_25_18, %iq2s_word0_low_25_19, %iq2s_word0_low_25_20, %iq2s_word0_low_25_21, %iq2s_word0_low_25_22, %iq2s_word0_low_25_23, %iq2s_word0_low_25_24, %iq2s_word0_low_25_25, %iq2s_word0_low_25_26, %iq2s_word0_low_25_27, %iq2s_word0_low_25_28, %iq2s_word0_low_25_29, %iq2s_word0_low_25_30, %iq2s_word0_low_25_31 : vector<32xf32> + %iq2s_word0_low_26_0 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_26_1 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_26_2 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_26_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_26_4 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_26_5 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_26_6 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_26_7 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_26_8 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_26_9 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_26_10 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_26_11 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_26_12 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_26_13 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_26_14 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_26_15 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_26_16 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_26_17 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_26_18 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_26_19 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_26_20 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_26_21 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_26_22 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_26_23 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_26_24 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_26_25 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_26_26 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_26_27 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_26_28 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_26_29 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_26_30 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_26_31 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_26 = vector.from_elements %iq2s_word0_low_26_0, %iq2s_word0_low_26_1, %iq2s_word0_low_26_2, %iq2s_word0_low_26_3, %iq2s_word0_low_26_4, %iq2s_word0_low_26_5, %iq2s_word0_low_26_6, %iq2s_word0_low_26_7, %iq2s_word0_low_26_8, %iq2s_word0_low_26_9, %iq2s_word0_low_26_10, %iq2s_word0_low_26_11, %iq2s_word0_low_26_12, %iq2s_word0_low_26_13, %iq2s_word0_low_26_14, %iq2s_word0_low_26_15, %iq2s_word0_low_26_16, %iq2s_word0_low_26_17, %iq2s_word0_low_26_18, %iq2s_word0_low_26_19, %iq2s_word0_low_26_20, %iq2s_word0_low_26_21, %iq2s_word0_low_26_22, %iq2s_word0_low_26_23, %iq2s_word0_low_26_24, %iq2s_word0_low_26_25, %iq2s_word0_low_26_26, %iq2s_word0_low_26_27, %iq2s_word0_low_26_28, %iq2s_word0_low_26_29, %iq2s_word0_low_26_30, %iq2s_word0_low_26_31 : vector<32xf32> + %iq2s_word0_low_27_0 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_27_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_27_2 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_27_3 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_27_4 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_27_5 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_27_6 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_27_7 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_27_8 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_27_9 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_27_10 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_27_11 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_27_12 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_27_13 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_27_14 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_27_15 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_27_16 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_27_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_27_18 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_27_19 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_27_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_27_21 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_27_22 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_27_23 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_27_24 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_27_25 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_27_26 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_27_27 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_27_28 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_27_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_27_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_27_31 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_27 = vector.from_elements %iq2s_word0_low_27_0, %iq2s_word0_low_27_1, %iq2s_word0_low_27_2, %iq2s_word0_low_27_3, %iq2s_word0_low_27_4, %iq2s_word0_low_27_5, %iq2s_word0_low_27_6, %iq2s_word0_low_27_7, %iq2s_word0_low_27_8, %iq2s_word0_low_27_9, %iq2s_word0_low_27_10, %iq2s_word0_low_27_11, %iq2s_word0_low_27_12, %iq2s_word0_low_27_13, %iq2s_word0_low_27_14, %iq2s_word0_low_27_15, %iq2s_word0_low_27_16, %iq2s_word0_low_27_17, %iq2s_word0_low_27_18, %iq2s_word0_low_27_19, %iq2s_word0_low_27_20, %iq2s_word0_low_27_21, %iq2s_word0_low_27_22, %iq2s_word0_low_27_23, %iq2s_word0_low_27_24, %iq2s_word0_low_27_25, %iq2s_word0_low_27_26, %iq2s_word0_low_27_27, %iq2s_word0_low_27_28, %iq2s_word0_low_27_29, %iq2s_word0_low_27_30, %iq2s_word0_low_27_31 : vector<32xf32> + %iq2s_word0_low_28_0 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_28_1 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_28_2 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_28_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_28_4 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_28_5 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_28_6 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_28_7 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_28_8 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_28_9 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_28_10 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_28_11 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_28_12 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_28_13 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_28_14 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_28_15 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_28_16 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_28_17 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_28_18 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_28_19 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_28_20 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_28_21 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_28_22 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_28_23 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_28_24 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_28_25 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_28_26 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_28_27 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_28_28 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_28_29 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_28_30 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_28_31 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_28 = vector.from_elements %iq2s_word0_low_28_0, %iq2s_word0_low_28_1, %iq2s_word0_low_28_2, %iq2s_word0_low_28_3, %iq2s_word0_low_28_4, %iq2s_word0_low_28_5, %iq2s_word0_low_28_6, %iq2s_word0_low_28_7, %iq2s_word0_low_28_8, %iq2s_word0_low_28_9, %iq2s_word0_low_28_10, %iq2s_word0_low_28_11, %iq2s_word0_low_28_12, %iq2s_word0_low_28_13, %iq2s_word0_low_28_14, %iq2s_word0_low_28_15, %iq2s_word0_low_28_16, %iq2s_word0_low_28_17, %iq2s_word0_low_28_18, %iq2s_word0_low_28_19, %iq2s_word0_low_28_20, %iq2s_word0_low_28_21, %iq2s_word0_low_28_22, %iq2s_word0_low_28_23, %iq2s_word0_low_28_24, %iq2s_word0_low_28_25, %iq2s_word0_low_28_26, %iq2s_word0_low_28_27, %iq2s_word0_low_28_28, %iq2s_word0_low_28_29, %iq2s_word0_low_28_30, %iq2s_word0_low_28_31 : vector<32xf32> + %iq2s_word0_low_29_0 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_29_1 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_29_2 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_29_3 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_29_4 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_29_5 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_29_6 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_29_7 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_29_8 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_29_9 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_29_10 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_29_11 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_29_12 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_29_13 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_29_14 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_29_15 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_29_16 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_29_17 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_29_18 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_29_19 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_29_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_29_21 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_29_22 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_29_23 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_29_24 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_29_25 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_29_26 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_29_27 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_29_28 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_29_29 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_29_30 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_29_31 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_29 = vector.from_elements %iq2s_word0_low_29_0, %iq2s_word0_low_29_1, %iq2s_word0_low_29_2, %iq2s_word0_low_29_3, %iq2s_word0_low_29_4, %iq2s_word0_low_29_5, %iq2s_word0_low_29_6, %iq2s_word0_low_29_7, %iq2s_word0_low_29_8, %iq2s_word0_low_29_9, %iq2s_word0_low_29_10, %iq2s_word0_low_29_11, %iq2s_word0_low_29_12, %iq2s_word0_low_29_13, %iq2s_word0_low_29_14, %iq2s_word0_low_29_15, %iq2s_word0_low_29_16, %iq2s_word0_low_29_17, %iq2s_word0_low_29_18, %iq2s_word0_low_29_19, %iq2s_word0_low_29_20, %iq2s_word0_low_29_21, %iq2s_word0_low_29_22, %iq2s_word0_low_29_23, %iq2s_word0_low_29_24, %iq2s_word0_low_29_25, %iq2s_word0_low_29_26, %iq2s_word0_low_29_27, %iq2s_word0_low_29_28, %iq2s_word0_low_29_29, %iq2s_word0_low_29_30, %iq2s_word0_low_29_31 : vector<32xf32> + %iq2s_word0_low_30_0 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_30_1 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_30_2 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_30_3 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_30_4 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_30_5 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_30_6 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_30_7 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_30_8 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_30_9 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_30_10 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_30_11 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_30_12 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_30_13 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_30_14 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_30_15 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_30_16 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_30_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_30_18 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_30_19 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_30_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_30_21 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_30_22 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_30_23 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_30_24 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_30_25 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_30_26 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_30_27 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_30_28 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_30_29 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_30_30 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_30_31 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_30 = vector.from_elements %iq2s_word0_low_30_0, %iq2s_word0_low_30_1, %iq2s_word0_low_30_2, %iq2s_word0_low_30_3, %iq2s_word0_low_30_4, %iq2s_word0_low_30_5, %iq2s_word0_low_30_6, %iq2s_word0_low_30_7, %iq2s_word0_low_30_8, %iq2s_word0_low_30_9, %iq2s_word0_low_30_10, %iq2s_word0_low_30_11, %iq2s_word0_low_30_12, %iq2s_word0_low_30_13, %iq2s_word0_low_30_14, %iq2s_word0_low_30_15, %iq2s_word0_low_30_16, %iq2s_word0_low_30_17, %iq2s_word0_low_30_18, %iq2s_word0_low_30_19, %iq2s_word0_low_30_20, %iq2s_word0_low_30_21, %iq2s_word0_low_30_22, %iq2s_word0_low_30_23, %iq2s_word0_low_30_24, %iq2s_word0_low_30_25, %iq2s_word0_low_30_26, %iq2s_word0_low_30_27, %iq2s_word0_low_30_28, %iq2s_word0_low_30_29, %iq2s_word0_low_30_30, %iq2s_word0_low_30_31 : vector<32xf32> + %iq2s_word0_low_31_0 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_31_1 = scalar.constant 6425.0 : f32 + %iq2s_word0_low_31_2 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_31_3 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_31_4 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_31_5 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_31_6 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_31_7 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_31_8 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_31_9 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_31_10 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_31_11 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_31_12 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_31_13 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_31_14 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_31_15 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_31_16 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_31_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_31_18 = scalar.constant 11033.0 : f32 + %iq2s_word0_low_31_19 = scalar.constant 2073.0 : f32 + %iq2s_word0_low_31_20 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_31_21 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_31_22 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_31_23 = scalar.constant 6408.0 : f32 + %iq2s_word0_low_31_24 = scalar.constant 6443.0 : f32 + %iq2s_word0_low_31_25 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_31_26 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_31_27 = scalar.constant 2056.0 : f32 + %iq2s_word0_low_31_28 = scalar.constant 2091.0 : f32 + %iq2s_word0_low_31_29 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_31_30 = scalar.constant 11016.0 : f32 + %iq2s_word0_low_31_31 = scalar.constant 11051.0 : f32 + %iq2s_word0_low_31 = vector.from_elements %iq2s_word0_low_31_0, %iq2s_word0_low_31_1, %iq2s_word0_low_31_2, %iq2s_word0_low_31_3, %iq2s_word0_low_31_4, %iq2s_word0_low_31_5, %iq2s_word0_low_31_6, %iq2s_word0_low_31_7, %iq2s_word0_low_31_8, %iq2s_word0_low_31_9, %iq2s_word0_low_31_10, %iq2s_word0_low_31_11, %iq2s_word0_low_31_12, %iq2s_word0_low_31_13, %iq2s_word0_low_31_14, %iq2s_word0_low_31_15, %iq2s_word0_low_31_16, %iq2s_word0_low_31_17, %iq2s_word0_low_31_18, %iq2s_word0_low_31_19, %iq2s_word0_low_31_20, %iq2s_word0_low_31_21, %iq2s_word0_low_31_22, %iq2s_word0_low_31_23, %iq2s_word0_low_31_24, %iq2s_word0_low_31_25, %iq2s_word0_low_31_26, %iq2s_word0_low_31_27, %iq2s_word0_low_31_28, %iq2s_word0_low_31_29, %iq2s_word0_low_31_30, %iq2s_word0_low_31_31 : vector<32xf32> + %selected_word0_low1 = scf.select %is_chunk1, %iq2s_word0_low_1, %iq2s_word0_low_0 : vector<32xf32> + %selected_word0_low2 = scf.select %is_chunk2, %iq2s_word0_low_2, %selected_word0_low1 : vector<32xf32> + %selected_word0_low3 = scf.select %is_chunk3, %iq2s_word0_low_3, %selected_word0_low2 : vector<32xf32> + %selected_word0_low4 = scf.select %is_chunk4, %iq2s_word0_low_4, %selected_word0_low3 : vector<32xf32> + %selected_word0_low5 = scf.select %is_chunk5, %iq2s_word0_low_5, %selected_word0_low4 : vector<32xf32> + %selected_word0_low6 = scf.select %is_chunk6, %iq2s_word0_low_6, %selected_word0_low5 : vector<32xf32> + %selected_word0_low7 = scf.select %is_chunk7, %iq2s_word0_low_7, %selected_word0_low6 : vector<32xf32> + %selected_word0_low8 = scf.select %is_chunk8, %iq2s_word0_low_8, %selected_word0_low7 : vector<32xf32> + %selected_word0_low9 = scf.select %is_chunk9, %iq2s_word0_low_9, %selected_word0_low8 : vector<32xf32> + %selected_word0_low10 = scf.select %is_chunk10, %iq2s_word0_low_10, %selected_word0_low9 : vector<32xf32> + %selected_word0_low11 = scf.select %is_chunk11, %iq2s_word0_low_11, %selected_word0_low10 : vector<32xf32> + %selected_word0_low12 = scf.select %is_chunk12, %iq2s_word0_low_12, %selected_word0_low11 : vector<32xf32> + %selected_word0_low13 = scf.select %is_chunk13, %iq2s_word0_low_13, %selected_word0_low12 : vector<32xf32> + %selected_word0_low14 = scf.select %is_chunk14, %iq2s_word0_low_14, %selected_word0_low13 : vector<32xf32> + %selected_word0_low15 = scf.select %is_chunk15, %iq2s_word0_low_15, %selected_word0_low14 : vector<32xf32> + %selected_word0_low16 = scf.select %is_chunk16, %iq2s_word0_low_16, %selected_word0_low15 : vector<32xf32> + %selected_word0_low17 = scf.select %is_chunk17, %iq2s_word0_low_17, %selected_word0_low16 : vector<32xf32> + %selected_word0_low18 = scf.select %is_chunk18, %iq2s_word0_low_18, %selected_word0_low17 : vector<32xf32> + %selected_word0_low19 = scf.select %is_chunk19, %iq2s_word0_low_19, %selected_word0_low18 : vector<32xf32> + %selected_word0_low20 = scf.select %is_chunk20, %iq2s_word0_low_20, %selected_word0_low19 : vector<32xf32> + %selected_word0_low21 = scf.select %is_chunk21, %iq2s_word0_low_21, %selected_word0_low20 : vector<32xf32> + %selected_word0_low22 = scf.select %is_chunk22, %iq2s_word0_low_22, %selected_word0_low21 : vector<32xf32> + %selected_word0_low23 = scf.select %is_chunk23, %iq2s_word0_low_23, %selected_word0_low22 : vector<32xf32> + %selected_word0_low24 = scf.select %is_chunk24, %iq2s_word0_low_24, %selected_word0_low23 : vector<32xf32> + %selected_word0_low25 = scf.select %is_chunk25, %iq2s_word0_low_25, %selected_word0_low24 : vector<32xf32> + %selected_word0_low26 = scf.select %is_chunk26, %iq2s_word0_low_26, %selected_word0_low25 : vector<32xf32> + %selected_word0_low27 = scf.select %is_chunk27, %iq2s_word0_low_27, %selected_word0_low26 : vector<32xf32> + %selected_word0_low28 = scf.select %is_chunk28, %iq2s_word0_low_28, %selected_word0_low27 : vector<32xf32> + %selected_word0_low29 = scf.select %is_chunk29, %iq2s_word0_low_29, %selected_word0_low28 : vector<32xf32> + %selected_word0_low30 = scf.select %is_chunk30, %iq2s_word0_low_30, %selected_word0_low29 : vector<32xf32> + %selected_word0_low31 = scf.select %is_chunk31, %iq2s_word0_low_31, %selected_word0_low30 : vector<32xf32> + %iq2s_word0_high_0_0 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_0_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_0_2 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_0_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_0_4 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_0_5 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_0_6 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_0_7 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_0_8 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_0_9 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_0_10 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_0_11 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_0_12 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_0_13 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_0_14 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_0_15 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_0_16 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_0_17 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_0_18 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_0_19 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_0_20 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_0_21 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_0_22 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_0_23 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_0_24 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_0_25 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_0_26 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_0_27 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_0_28 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_0_29 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_0_30 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_0_31 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_0 = vector.from_elements %iq2s_word0_high_0_0, %iq2s_word0_high_0_1, %iq2s_word0_high_0_2, %iq2s_word0_high_0_3, %iq2s_word0_high_0_4, %iq2s_word0_high_0_5, %iq2s_word0_high_0_6, %iq2s_word0_high_0_7, %iq2s_word0_high_0_8, %iq2s_word0_high_0_9, %iq2s_word0_high_0_10, %iq2s_word0_high_0_11, %iq2s_word0_high_0_12, %iq2s_word0_high_0_13, %iq2s_word0_high_0_14, %iq2s_word0_high_0_15, %iq2s_word0_high_0_16, %iq2s_word0_high_0_17, %iq2s_word0_high_0_18, %iq2s_word0_high_0_19, %iq2s_word0_high_0_20, %iq2s_word0_high_0_21, %iq2s_word0_high_0_22, %iq2s_word0_high_0_23, %iq2s_word0_high_0_24, %iq2s_word0_high_0_25, %iq2s_word0_high_0_26, %iq2s_word0_high_0_27, %iq2s_word0_high_0_28, %iq2s_word0_high_0_29, %iq2s_word0_high_0_30, %iq2s_word0_high_0_31 : vector<32xf32> + %iq2s_word0_high_1_0 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_1_1 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_1_2 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_1_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_1_4 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_1_5 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_1_6 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_1_7 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_1_8 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_1_9 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_1_10 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_1_11 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_1_12 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_1_13 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_1_14 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_1_15 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_1_16 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_1_17 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_1_18 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_1_19 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_1_20 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_1_21 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_1_22 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_1_23 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_1_24 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_1_25 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_1_26 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_1_27 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_1_28 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_1_29 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_1_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_1_31 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_1 = vector.from_elements %iq2s_word0_high_1_0, %iq2s_word0_high_1_1, %iq2s_word0_high_1_2, %iq2s_word0_high_1_3, %iq2s_word0_high_1_4, %iq2s_word0_high_1_5, %iq2s_word0_high_1_6, %iq2s_word0_high_1_7, %iq2s_word0_high_1_8, %iq2s_word0_high_1_9, %iq2s_word0_high_1_10, %iq2s_word0_high_1_11, %iq2s_word0_high_1_12, %iq2s_word0_high_1_13, %iq2s_word0_high_1_14, %iq2s_word0_high_1_15, %iq2s_word0_high_1_16, %iq2s_word0_high_1_17, %iq2s_word0_high_1_18, %iq2s_word0_high_1_19, %iq2s_word0_high_1_20, %iq2s_word0_high_1_21, %iq2s_word0_high_1_22, %iq2s_word0_high_1_23, %iq2s_word0_high_1_24, %iq2s_word0_high_1_25, %iq2s_word0_high_1_26, %iq2s_word0_high_1_27, %iq2s_word0_high_1_28, %iq2s_word0_high_1_29, %iq2s_word0_high_1_30, %iq2s_word0_high_1_31 : vector<32xf32> + %iq2s_word0_high_2_0 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_2_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_2_2 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_2_3 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_2_4 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_2_5 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_2_6 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_2_7 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_2_8 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_2_9 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_2_10 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_2_11 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_2_12 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_2_13 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_2_14 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_2_15 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_2_16 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_2_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_2_18 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_2_19 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_2_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_2_21 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_2_22 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_2_23 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_2_24 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_2_25 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_2_26 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_2_27 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_2_28 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_2_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_2_30 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_2_31 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_2 = vector.from_elements %iq2s_word0_high_2_0, %iq2s_word0_high_2_1, %iq2s_word0_high_2_2, %iq2s_word0_high_2_3, %iq2s_word0_high_2_4, %iq2s_word0_high_2_5, %iq2s_word0_high_2_6, %iq2s_word0_high_2_7, %iq2s_word0_high_2_8, %iq2s_word0_high_2_9, %iq2s_word0_high_2_10, %iq2s_word0_high_2_11, %iq2s_word0_high_2_12, %iq2s_word0_high_2_13, %iq2s_word0_high_2_14, %iq2s_word0_high_2_15, %iq2s_word0_high_2_16, %iq2s_word0_high_2_17, %iq2s_word0_high_2_18, %iq2s_word0_high_2_19, %iq2s_word0_high_2_20, %iq2s_word0_high_2_21, %iq2s_word0_high_2_22, %iq2s_word0_high_2_23, %iq2s_word0_high_2_24, %iq2s_word0_high_2_25, %iq2s_word0_high_2_26, %iq2s_word0_high_2_27, %iq2s_word0_high_2_28, %iq2s_word0_high_2_29, %iq2s_word0_high_2_30, %iq2s_word0_high_2_31 : vector<32xf32> + %iq2s_word0_high_3_0 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_3_1 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_3_2 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_3_3 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_3_4 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_3_5 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_3_6 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_3_7 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_3_8 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_3_9 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_3_10 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_3_11 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_3_12 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_3_13 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_3_14 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_3_15 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_3_16 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_3_17 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_3_18 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_3_19 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_3_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_3_21 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_3_22 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_3_23 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_3_24 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_3_25 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_3_26 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_3_27 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_3_28 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_3_29 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_3_30 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_3_31 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_3 = vector.from_elements %iq2s_word0_high_3_0, %iq2s_word0_high_3_1, %iq2s_word0_high_3_2, %iq2s_word0_high_3_3, %iq2s_word0_high_3_4, %iq2s_word0_high_3_5, %iq2s_word0_high_3_6, %iq2s_word0_high_3_7, %iq2s_word0_high_3_8, %iq2s_word0_high_3_9, %iq2s_word0_high_3_10, %iq2s_word0_high_3_11, %iq2s_word0_high_3_12, %iq2s_word0_high_3_13, %iq2s_word0_high_3_14, %iq2s_word0_high_3_15, %iq2s_word0_high_3_16, %iq2s_word0_high_3_17, %iq2s_word0_high_3_18, %iq2s_word0_high_3_19, %iq2s_word0_high_3_20, %iq2s_word0_high_3_21, %iq2s_word0_high_3_22, %iq2s_word0_high_3_23, %iq2s_word0_high_3_24, %iq2s_word0_high_3_25, %iq2s_word0_high_3_26, %iq2s_word0_high_3_27, %iq2s_word0_high_3_28, %iq2s_word0_high_3_29, %iq2s_word0_high_3_30, %iq2s_word0_high_3_31 : vector<32xf32> + %iq2s_word0_high_4_0 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_4_1 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_4_2 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_4_3 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_4_4 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_4_5 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_4_6 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_4_7 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_4_8 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_4_9 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_4_10 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_4_11 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_4_12 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_4_13 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_4_14 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_4_15 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_4_16 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_4_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_4_18 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_4_19 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_4_20 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_4_21 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_4_22 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_4_23 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_4_24 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_4_25 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_4_26 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_4_27 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_4_28 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_4_29 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_4_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_4_31 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_4 = vector.from_elements %iq2s_word0_high_4_0, %iq2s_word0_high_4_1, %iq2s_word0_high_4_2, %iq2s_word0_high_4_3, %iq2s_word0_high_4_4, %iq2s_word0_high_4_5, %iq2s_word0_high_4_6, %iq2s_word0_high_4_7, %iq2s_word0_high_4_8, %iq2s_word0_high_4_9, %iq2s_word0_high_4_10, %iq2s_word0_high_4_11, %iq2s_word0_high_4_12, %iq2s_word0_high_4_13, %iq2s_word0_high_4_14, %iq2s_word0_high_4_15, %iq2s_word0_high_4_16, %iq2s_word0_high_4_17, %iq2s_word0_high_4_18, %iq2s_word0_high_4_19, %iq2s_word0_high_4_20, %iq2s_word0_high_4_21, %iq2s_word0_high_4_22, %iq2s_word0_high_4_23, %iq2s_word0_high_4_24, %iq2s_word0_high_4_25, %iq2s_word0_high_4_26, %iq2s_word0_high_4_27, %iq2s_word0_high_4_28, %iq2s_word0_high_4_29, %iq2s_word0_high_4_30, %iq2s_word0_high_4_31 : vector<32xf32> + %iq2s_word0_high_5_0 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_5_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_5_2 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_5_3 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_5_4 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_5_5 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_5_6 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_5_7 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_5_8 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_5_9 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_5_10 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_5_11 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_5_12 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_5_13 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_5_14 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_5_15 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_5_16 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_5_17 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_5_18 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_5_19 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_5_20 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_5_21 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_5_22 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_5_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_5_24 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_5_25 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_5_26 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_5_27 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_5_28 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_5_29 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_5_30 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_5_31 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_5 = vector.from_elements %iq2s_word0_high_5_0, %iq2s_word0_high_5_1, %iq2s_word0_high_5_2, %iq2s_word0_high_5_3, %iq2s_word0_high_5_4, %iq2s_word0_high_5_5, %iq2s_word0_high_5_6, %iq2s_word0_high_5_7, %iq2s_word0_high_5_8, %iq2s_word0_high_5_9, %iq2s_word0_high_5_10, %iq2s_word0_high_5_11, %iq2s_word0_high_5_12, %iq2s_word0_high_5_13, %iq2s_word0_high_5_14, %iq2s_word0_high_5_15, %iq2s_word0_high_5_16, %iq2s_word0_high_5_17, %iq2s_word0_high_5_18, %iq2s_word0_high_5_19, %iq2s_word0_high_5_20, %iq2s_word0_high_5_21, %iq2s_word0_high_5_22, %iq2s_word0_high_5_23, %iq2s_word0_high_5_24, %iq2s_word0_high_5_25, %iq2s_word0_high_5_26, %iq2s_word0_high_5_27, %iq2s_word0_high_5_28, %iq2s_word0_high_5_29, %iq2s_word0_high_5_30, %iq2s_word0_high_5_31 : vector<32xf32> + %iq2s_word0_high_6_0 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_6_1 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_6_2 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_6_3 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_6_4 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_6_5 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_6_6 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_6_7 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_6_8 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_6_9 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_6_10 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_6_11 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_6_12 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_6_13 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_6_14 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_6_15 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_6_16 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_6_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_6_18 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_6_19 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_6_20 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_6_21 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_6_22 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_6_23 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_6_24 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_6_25 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_6_26 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_6_27 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_6_28 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_6_29 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_6_30 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_6_31 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_6 = vector.from_elements %iq2s_word0_high_6_0, %iq2s_word0_high_6_1, %iq2s_word0_high_6_2, %iq2s_word0_high_6_3, %iq2s_word0_high_6_4, %iq2s_word0_high_6_5, %iq2s_word0_high_6_6, %iq2s_word0_high_6_7, %iq2s_word0_high_6_8, %iq2s_word0_high_6_9, %iq2s_word0_high_6_10, %iq2s_word0_high_6_11, %iq2s_word0_high_6_12, %iq2s_word0_high_6_13, %iq2s_word0_high_6_14, %iq2s_word0_high_6_15, %iq2s_word0_high_6_16, %iq2s_word0_high_6_17, %iq2s_word0_high_6_18, %iq2s_word0_high_6_19, %iq2s_word0_high_6_20, %iq2s_word0_high_6_21, %iq2s_word0_high_6_22, %iq2s_word0_high_6_23, %iq2s_word0_high_6_24, %iq2s_word0_high_6_25, %iq2s_word0_high_6_26, %iq2s_word0_high_6_27, %iq2s_word0_high_6_28, %iq2s_word0_high_6_29, %iq2s_word0_high_6_30, %iq2s_word0_high_6_31 : vector<32xf32> + %iq2s_word0_high_7_0 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_7_1 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_7_2 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_7_3 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_7_4 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_7_5 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_7_6 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_7_7 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_7_8 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_7_9 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_7_10 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_7_11 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_7_12 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_7_13 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_7_14 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_7_15 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_7_16 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_7_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_7_18 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_7_19 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_7_20 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_7_21 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_7_22 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_7_23 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_7_24 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_7_25 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_7_26 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_7_27 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_7_28 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_7_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_7_30 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_7_31 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_7 = vector.from_elements %iq2s_word0_high_7_0, %iq2s_word0_high_7_1, %iq2s_word0_high_7_2, %iq2s_word0_high_7_3, %iq2s_word0_high_7_4, %iq2s_word0_high_7_5, %iq2s_word0_high_7_6, %iq2s_word0_high_7_7, %iq2s_word0_high_7_8, %iq2s_word0_high_7_9, %iq2s_word0_high_7_10, %iq2s_word0_high_7_11, %iq2s_word0_high_7_12, %iq2s_word0_high_7_13, %iq2s_word0_high_7_14, %iq2s_word0_high_7_15, %iq2s_word0_high_7_16, %iq2s_word0_high_7_17, %iq2s_word0_high_7_18, %iq2s_word0_high_7_19, %iq2s_word0_high_7_20, %iq2s_word0_high_7_21, %iq2s_word0_high_7_22, %iq2s_word0_high_7_23, %iq2s_word0_high_7_24, %iq2s_word0_high_7_25, %iq2s_word0_high_7_26, %iq2s_word0_high_7_27, %iq2s_word0_high_7_28, %iq2s_word0_high_7_29, %iq2s_word0_high_7_30, %iq2s_word0_high_7_31 : vector<32xf32> + %iq2s_word0_high_8_0 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_8_1 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_8_2 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_8_3 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_8_4 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_8_5 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_8_6 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_8_7 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_8_8 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_8_9 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_8_10 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_8_11 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_8_12 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_8_13 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_8_14 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_8_15 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_8_16 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_8_17 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_8_18 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_8_19 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_8_20 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_8_21 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_8_22 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_8_23 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_8_24 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_8_25 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_8_26 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_8_27 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_8_28 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_8_29 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_8_30 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_8_31 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_8 = vector.from_elements %iq2s_word0_high_8_0, %iq2s_word0_high_8_1, %iq2s_word0_high_8_2, %iq2s_word0_high_8_3, %iq2s_word0_high_8_4, %iq2s_word0_high_8_5, %iq2s_word0_high_8_6, %iq2s_word0_high_8_7, %iq2s_word0_high_8_8, %iq2s_word0_high_8_9, %iq2s_word0_high_8_10, %iq2s_word0_high_8_11, %iq2s_word0_high_8_12, %iq2s_word0_high_8_13, %iq2s_word0_high_8_14, %iq2s_word0_high_8_15, %iq2s_word0_high_8_16, %iq2s_word0_high_8_17, %iq2s_word0_high_8_18, %iq2s_word0_high_8_19, %iq2s_word0_high_8_20, %iq2s_word0_high_8_21, %iq2s_word0_high_8_22, %iq2s_word0_high_8_23, %iq2s_word0_high_8_24, %iq2s_word0_high_8_25, %iq2s_word0_high_8_26, %iq2s_word0_high_8_27, %iq2s_word0_high_8_28, %iq2s_word0_high_8_29, %iq2s_word0_high_8_30, %iq2s_word0_high_8_31 : vector<32xf32> + %iq2s_word0_high_9_0 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_9_1 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_9_2 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_9_3 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_9_4 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_9_5 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_9_6 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_9_7 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_9_8 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_9_9 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_9_10 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_9_11 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_9_12 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_9_13 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_9_14 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_9_15 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_9_16 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_9_17 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_9_18 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_9_19 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_9_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_9_21 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_9_22 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_9_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_9_24 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_9_25 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_9_26 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_9_27 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_9_28 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_9_29 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_9_30 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_9_31 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_9 = vector.from_elements %iq2s_word0_high_9_0, %iq2s_word0_high_9_1, %iq2s_word0_high_9_2, %iq2s_word0_high_9_3, %iq2s_word0_high_9_4, %iq2s_word0_high_9_5, %iq2s_word0_high_9_6, %iq2s_word0_high_9_7, %iq2s_word0_high_9_8, %iq2s_word0_high_9_9, %iq2s_word0_high_9_10, %iq2s_word0_high_9_11, %iq2s_word0_high_9_12, %iq2s_word0_high_9_13, %iq2s_word0_high_9_14, %iq2s_word0_high_9_15, %iq2s_word0_high_9_16, %iq2s_word0_high_9_17, %iq2s_word0_high_9_18, %iq2s_word0_high_9_19, %iq2s_word0_high_9_20, %iq2s_word0_high_9_21, %iq2s_word0_high_9_22, %iq2s_word0_high_9_23, %iq2s_word0_high_9_24, %iq2s_word0_high_9_25, %iq2s_word0_high_9_26, %iq2s_word0_high_9_27, %iq2s_word0_high_9_28, %iq2s_word0_high_9_29, %iq2s_word0_high_9_30, %iq2s_word0_high_9_31 : vector<32xf32> + %iq2s_word0_high_10_0 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_10_1 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_10_2 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_10_3 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_10_4 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_10_5 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_10_6 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_10_7 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_10_8 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_10_9 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_10_10 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_10_11 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_10_12 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_10_13 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_10_14 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_10_15 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_10_16 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_10_17 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_10_18 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_10_19 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_10_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_10_21 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_10_22 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_10_23 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_10_24 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_10_25 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_10_26 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_10_27 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_10_28 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_10_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_10_30 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_10_31 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_10 = vector.from_elements %iq2s_word0_high_10_0, %iq2s_word0_high_10_1, %iq2s_word0_high_10_2, %iq2s_word0_high_10_3, %iq2s_word0_high_10_4, %iq2s_word0_high_10_5, %iq2s_word0_high_10_6, %iq2s_word0_high_10_7, %iq2s_word0_high_10_8, %iq2s_word0_high_10_9, %iq2s_word0_high_10_10, %iq2s_word0_high_10_11, %iq2s_word0_high_10_12, %iq2s_word0_high_10_13, %iq2s_word0_high_10_14, %iq2s_word0_high_10_15, %iq2s_word0_high_10_16, %iq2s_word0_high_10_17, %iq2s_word0_high_10_18, %iq2s_word0_high_10_19, %iq2s_word0_high_10_20, %iq2s_word0_high_10_21, %iq2s_word0_high_10_22, %iq2s_word0_high_10_23, %iq2s_word0_high_10_24, %iq2s_word0_high_10_25, %iq2s_word0_high_10_26, %iq2s_word0_high_10_27, %iq2s_word0_high_10_28, %iq2s_word0_high_10_29, %iq2s_word0_high_10_30, %iq2s_word0_high_10_31 : vector<32xf32> + %iq2s_word0_high_11_0 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_11_1 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_11_2 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_11_3 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_11_4 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_11_5 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_11_6 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_11_7 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_11_8 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_11_9 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_11_10 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_11_11 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_11_12 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_11_13 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_11_14 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_11_15 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_11_16 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_11_17 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_11_18 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_11_19 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_11_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_11_21 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_11_22 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_11_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_11_24 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_11_25 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_11_26 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_11_27 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_11_28 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_11_29 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_11_30 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_11_31 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_11 = vector.from_elements %iq2s_word0_high_11_0, %iq2s_word0_high_11_1, %iq2s_word0_high_11_2, %iq2s_word0_high_11_3, %iq2s_word0_high_11_4, %iq2s_word0_high_11_5, %iq2s_word0_high_11_6, %iq2s_word0_high_11_7, %iq2s_word0_high_11_8, %iq2s_word0_high_11_9, %iq2s_word0_high_11_10, %iq2s_word0_high_11_11, %iq2s_word0_high_11_12, %iq2s_word0_high_11_13, %iq2s_word0_high_11_14, %iq2s_word0_high_11_15, %iq2s_word0_high_11_16, %iq2s_word0_high_11_17, %iq2s_word0_high_11_18, %iq2s_word0_high_11_19, %iq2s_word0_high_11_20, %iq2s_word0_high_11_21, %iq2s_word0_high_11_22, %iq2s_word0_high_11_23, %iq2s_word0_high_11_24, %iq2s_word0_high_11_25, %iq2s_word0_high_11_26, %iq2s_word0_high_11_27, %iq2s_word0_high_11_28, %iq2s_word0_high_11_29, %iq2s_word0_high_11_30, %iq2s_word0_high_11_31 : vector<32xf32> + %iq2s_word0_high_12_0 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_12_1 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_12_2 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_12_3 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_12_4 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_12_5 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_12_6 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_12_7 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_12_8 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_12_9 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_12_10 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_12_11 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_12_12 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_12_13 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_12_14 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_12_15 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_12_16 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_12_17 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_12_18 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_12_19 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_12_20 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_12_21 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_12_22 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_12_23 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_12_24 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_12_25 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_12_26 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_12_27 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_12_28 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_12_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_12_30 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_12_31 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_12 = vector.from_elements %iq2s_word0_high_12_0, %iq2s_word0_high_12_1, %iq2s_word0_high_12_2, %iq2s_word0_high_12_3, %iq2s_word0_high_12_4, %iq2s_word0_high_12_5, %iq2s_word0_high_12_6, %iq2s_word0_high_12_7, %iq2s_word0_high_12_8, %iq2s_word0_high_12_9, %iq2s_word0_high_12_10, %iq2s_word0_high_12_11, %iq2s_word0_high_12_12, %iq2s_word0_high_12_13, %iq2s_word0_high_12_14, %iq2s_word0_high_12_15, %iq2s_word0_high_12_16, %iq2s_word0_high_12_17, %iq2s_word0_high_12_18, %iq2s_word0_high_12_19, %iq2s_word0_high_12_20, %iq2s_word0_high_12_21, %iq2s_word0_high_12_22, %iq2s_word0_high_12_23, %iq2s_word0_high_12_24, %iq2s_word0_high_12_25, %iq2s_word0_high_12_26, %iq2s_word0_high_12_27, %iq2s_word0_high_12_28, %iq2s_word0_high_12_29, %iq2s_word0_high_12_30, %iq2s_word0_high_12_31 : vector<32xf32> + %iq2s_word0_high_13_0 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_13_1 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_13_2 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_13_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_13_4 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_13_5 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_13_6 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_13_7 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_13_8 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_13_9 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_13_10 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_13_11 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_13_12 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_13_13 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_13_14 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_13_15 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_13_16 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_13_17 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_13_18 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_13_19 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_13_20 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_13_21 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_13_22 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_13_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_13_24 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_13_25 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_13_26 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_13_27 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_13_28 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_13_29 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_13_30 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_13_31 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_13 = vector.from_elements %iq2s_word0_high_13_0, %iq2s_word0_high_13_1, %iq2s_word0_high_13_2, %iq2s_word0_high_13_3, %iq2s_word0_high_13_4, %iq2s_word0_high_13_5, %iq2s_word0_high_13_6, %iq2s_word0_high_13_7, %iq2s_word0_high_13_8, %iq2s_word0_high_13_9, %iq2s_word0_high_13_10, %iq2s_word0_high_13_11, %iq2s_word0_high_13_12, %iq2s_word0_high_13_13, %iq2s_word0_high_13_14, %iq2s_word0_high_13_15, %iq2s_word0_high_13_16, %iq2s_word0_high_13_17, %iq2s_word0_high_13_18, %iq2s_word0_high_13_19, %iq2s_word0_high_13_20, %iq2s_word0_high_13_21, %iq2s_word0_high_13_22, %iq2s_word0_high_13_23, %iq2s_word0_high_13_24, %iq2s_word0_high_13_25, %iq2s_word0_high_13_26, %iq2s_word0_high_13_27, %iq2s_word0_high_13_28, %iq2s_word0_high_13_29, %iq2s_word0_high_13_30, %iq2s_word0_high_13_31 : vector<32xf32> + %iq2s_word0_high_14_0 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_14_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_14_2 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_14_3 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_14_4 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_14_5 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_14_6 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_14_7 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_14_8 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_14_9 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_14_10 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_14_11 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_14_12 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_14_13 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_14_14 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_14_15 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_14_16 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_14_17 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_14_18 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_14_19 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_14_20 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_14_21 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_14_22 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_14_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_14_24 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_14_25 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_14_26 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_14_27 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_14_28 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_14_29 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_14_30 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_14_31 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_14 = vector.from_elements %iq2s_word0_high_14_0, %iq2s_word0_high_14_1, %iq2s_word0_high_14_2, %iq2s_word0_high_14_3, %iq2s_word0_high_14_4, %iq2s_word0_high_14_5, %iq2s_word0_high_14_6, %iq2s_word0_high_14_7, %iq2s_word0_high_14_8, %iq2s_word0_high_14_9, %iq2s_word0_high_14_10, %iq2s_word0_high_14_11, %iq2s_word0_high_14_12, %iq2s_word0_high_14_13, %iq2s_word0_high_14_14, %iq2s_word0_high_14_15, %iq2s_word0_high_14_16, %iq2s_word0_high_14_17, %iq2s_word0_high_14_18, %iq2s_word0_high_14_19, %iq2s_word0_high_14_20, %iq2s_word0_high_14_21, %iq2s_word0_high_14_22, %iq2s_word0_high_14_23, %iq2s_word0_high_14_24, %iq2s_word0_high_14_25, %iq2s_word0_high_14_26, %iq2s_word0_high_14_27, %iq2s_word0_high_14_28, %iq2s_word0_high_14_29, %iq2s_word0_high_14_30, %iq2s_word0_high_14_31 : vector<32xf32> + %iq2s_word0_high_15_0 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_15_1 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_15_2 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_15_3 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_15_4 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_15_5 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_15_6 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_15_7 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_15_8 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_15_9 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_15_10 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_15_11 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_15_12 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_15_13 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_15_14 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_15_15 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_15_16 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_15_17 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_15_18 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_15_19 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_15_20 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_15_21 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_15_22 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_15_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_15_24 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_15_25 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_15_26 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_15_27 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_15_28 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_15_29 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_15_30 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_15_31 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_15 = vector.from_elements %iq2s_word0_high_15_0, %iq2s_word0_high_15_1, %iq2s_word0_high_15_2, %iq2s_word0_high_15_3, %iq2s_word0_high_15_4, %iq2s_word0_high_15_5, %iq2s_word0_high_15_6, %iq2s_word0_high_15_7, %iq2s_word0_high_15_8, %iq2s_word0_high_15_9, %iq2s_word0_high_15_10, %iq2s_word0_high_15_11, %iq2s_word0_high_15_12, %iq2s_word0_high_15_13, %iq2s_word0_high_15_14, %iq2s_word0_high_15_15, %iq2s_word0_high_15_16, %iq2s_word0_high_15_17, %iq2s_word0_high_15_18, %iq2s_word0_high_15_19, %iq2s_word0_high_15_20, %iq2s_word0_high_15_21, %iq2s_word0_high_15_22, %iq2s_word0_high_15_23, %iq2s_word0_high_15_24, %iq2s_word0_high_15_25, %iq2s_word0_high_15_26, %iq2s_word0_high_15_27, %iq2s_word0_high_15_28, %iq2s_word0_high_15_29, %iq2s_word0_high_15_30, %iq2s_word0_high_15_31 : vector<32xf32> + %iq2s_word0_high_16_0 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_16_1 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_16_2 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_16_3 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_16_4 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_16_5 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_16_6 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_16_7 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_16_8 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_16_9 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_16_10 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_16_11 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_16_12 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_16_13 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_16_14 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_16_15 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_16_16 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_16_17 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_16_18 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_16_19 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_16_20 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_16_21 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_16_22 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_16_23 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_16_24 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_16_25 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_16_26 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_16_27 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_16_28 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_16_29 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_16_30 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_16_31 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_16 = vector.from_elements %iq2s_word0_high_16_0, %iq2s_word0_high_16_1, %iq2s_word0_high_16_2, %iq2s_word0_high_16_3, %iq2s_word0_high_16_4, %iq2s_word0_high_16_5, %iq2s_word0_high_16_6, %iq2s_word0_high_16_7, %iq2s_word0_high_16_8, %iq2s_word0_high_16_9, %iq2s_word0_high_16_10, %iq2s_word0_high_16_11, %iq2s_word0_high_16_12, %iq2s_word0_high_16_13, %iq2s_word0_high_16_14, %iq2s_word0_high_16_15, %iq2s_word0_high_16_16, %iq2s_word0_high_16_17, %iq2s_word0_high_16_18, %iq2s_word0_high_16_19, %iq2s_word0_high_16_20, %iq2s_word0_high_16_21, %iq2s_word0_high_16_22, %iq2s_word0_high_16_23, %iq2s_word0_high_16_24, %iq2s_word0_high_16_25, %iq2s_word0_high_16_26, %iq2s_word0_high_16_27, %iq2s_word0_high_16_28, %iq2s_word0_high_16_29, %iq2s_word0_high_16_30, %iq2s_word0_high_16_31 : vector<32xf32> + %iq2s_word0_high_17_0 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_17_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_17_2 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_17_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_17_4 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_17_5 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_17_6 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_17_7 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_17_8 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_17_9 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_17_10 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_17_11 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_17_12 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_17_13 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_17_14 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_17_15 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_17_16 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_17_17 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_17_18 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_17_19 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_17_20 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_17_21 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_17_22 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_17_23 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_17_24 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_17_25 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_17_26 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_17_27 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_17_28 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_17_29 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_17_30 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_17_31 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_17 = vector.from_elements %iq2s_word0_high_17_0, %iq2s_word0_high_17_1, %iq2s_word0_high_17_2, %iq2s_word0_high_17_3, %iq2s_word0_high_17_4, %iq2s_word0_high_17_5, %iq2s_word0_high_17_6, %iq2s_word0_high_17_7, %iq2s_word0_high_17_8, %iq2s_word0_high_17_9, %iq2s_word0_high_17_10, %iq2s_word0_high_17_11, %iq2s_word0_high_17_12, %iq2s_word0_high_17_13, %iq2s_word0_high_17_14, %iq2s_word0_high_17_15, %iq2s_word0_high_17_16, %iq2s_word0_high_17_17, %iq2s_word0_high_17_18, %iq2s_word0_high_17_19, %iq2s_word0_high_17_20, %iq2s_word0_high_17_21, %iq2s_word0_high_17_22, %iq2s_word0_high_17_23, %iq2s_word0_high_17_24, %iq2s_word0_high_17_25, %iq2s_word0_high_17_26, %iq2s_word0_high_17_27, %iq2s_word0_high_17_28, %iq2s_word0_high_17_29, %iq2s_word0_high_17_30, %iq2s_word0_high_17_31 : vector<32xf32> + %iq2s_word0_high_18_0 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_18_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_18_2 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_18_3 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_18_4 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_18_5 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_18_6 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_18_7 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_18_8 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_18_9 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_18_10 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_18_11 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_18_12 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_18_13 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_18_14 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_18_15 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_18_16 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_18_17 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_18_18 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_18_19 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_18_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_18_21 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_18_22 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_18_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_18_24 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_18_25 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_18_26 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_18_27 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_18_28 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_18_29 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_18_30 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_18_31 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_18 = vector.from_elements %iq2s_word0_high_18_0, %iq2s_word0_high_18_1, %iq2s_word0_high_18_2, %iq2s_word0_high_18_3, %iq2s_word0_high_18_4, %iq2s_word0_high_18_5, %iq2s_word0_high_18_6, %iq2s_word0_high_18_7, %iq2s_word0_high_18_8, %iq2s_word0_high_18_9, %iq2s_word0_high_18_10, %iq2s_word0_high_18_11, %iq2s_word0_high_18_12, %iq2s_word0_high_18_13, %iq2s_word0_high_18_14, %iq2s_word0_high_18_15, %iq2s_word0_high_18_16, %iq2s_word0_high_18_17, %iq2s_word0_high_18_18, %iq2s_word0_high_18_19, %iq2s_word0_high_18_20, %iq2s_word0_high_18_21, %iq2s_word0_high_18_22, %iq2s_word0_high_18_23, %iq2s_word0_high_18_24, %iq2s_word0_high_18_25, %iq2s_word0_high_18_26, %iq2s_word0_high_18_27, %iq2s_word0_high_18_28, %iq2s_word0_high_18_29, %iq2s_word0_high_18_30, %iq2s_word0_high_18_31 : vector<32xf32> + %iq2s_word0_high_19_0 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_19_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_19_2 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_19_3 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_19_4 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_19_5 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_19_6 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_19_7 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_19_8 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_19_9 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_19_10 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_19_11 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_19_12 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_19_13 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_19_14 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_19_15 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_19_16 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_19_17 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_19_18 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_19_19 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_19_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_19_21 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_19_22 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_19_23 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_19_24 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_19_25 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_19_26 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_19_27 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_19_28 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_19_29 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_19_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_19_31 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_19 = vector.from_elements %iq2s_word0_high_19_0, %iq2s_word0_high_19_1, %iq2s_word0_high_19_2, %iq2s_word0_high_19_3, %iq2s_word0_high_19_4, %iq2s_word0_high_19_5, %iq2s_word0_high_19_6, %iq2s_word0_high_19_7, %iq2s_word0_high_19_8, %iq2s_word0_high_19_9, %iq2s_word0_high_19_10, %iq2s_word0_high_19_11, %iq2s_word0_high_19_12, %iq2s_word0_high_19_13, %iq2s_word0_high_19_14, %iq2s_word0_high_19_15, %iq2s_word0_high_19_16, %iq2s_word0_high_19_17, %iq2s_word0_high_19_18, %iq2s_word0_high_19_19, %iq2s_word0_high_19_20, %iq2s_word0_high_19_21, %iq2s_word0_high_19_22, %iq2s_word0_high_19_23, %iq2s_word0_high_19_24, %iq2s_word0_high_19_25, %iq2s_word0_high_19_26, %iq2s_word0_high_19_27, %iq2s_word0_high_19_28, %iq2s_word0_high_19_29, %iq2s_word0_high_19_30, %iq2s_word0_high_19_31 : vector<32xf32> + %iq2s_word0_high_20_0 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_20_1 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_20_2 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_20_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_20_4 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_20_5 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_20_6 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_20_7 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_20_8 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_20_9 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_20_10 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_20_11 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_20_12 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_20_13 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_20_14 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_20_15 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_20_16 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_20_17 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_20_18 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_20_19 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_20_20 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_20_21 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_20_22 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_20_23 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_20_24 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_20_25 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_20_26 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_20_27 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_20_28 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_20_29 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_20_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_20_31 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_20 = vector.from_elements %iq2s_word0_high_20_0, %iq2s_word0_high_20_1, %iq2s_word0_high_20_2, %iq2s_word0_high_20_3, %iq2s_word0_high_20_4, %iq2s_word0_high_20_5, %iq2s_word0_high_20_6, %iq2s_word0_high_20_7, %iq2s_word0_high_20_8, %iq2s_word0_high_20_9, %iq2s_word0_high_20_10, %iq2s_word0_high_20_11, %iq2s_word0_high_20_12, %iq2s_word0_high_20_13, %iq2s_word0_high_20_14, %iq2s_word0_high_20_15, %iq2s_word0_high_20_16, %iq2s_word0_high_20_17, %iq2s_word0_high_20_18, %iq2s_word0_high_20_19, %iq2s_word0_high_20_20, %iq2s_word0_high_20_21, %iq2s_word0_high_20_22, %iq2s_word0_high_20_23, %iq2s_word0_high_20_24, %iq2s_word0_high_20_25, %iq2s_word0_high_20_26, %iq2s_word0_high_20_27, %iq2s_word0_high_20_28, %iq2s_word0_high_20_29, %iq2s_word0_high_20_30, %iq2s_word0_high_20_31 : vector<32xf32> + %iq2s_word0_high_21_0 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_21_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_21_2 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_21_3 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_21_4 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_21_5 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_21_6 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_21_7 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_21_8 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_21_9 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_21_10 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_21_11 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_21_12 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_21_13 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_21_14 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_21_15 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_21_16 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_21_17 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_21_18 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_21_19 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_21_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_21_21 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_21_22 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_21_23 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_21_24 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_21_25 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_21_26 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_21_27 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_21_28 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_21_29 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_21_30 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_21_31 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_21 = vector.from_elements %iq2s_word0_high_21_0, %iq2s_word0_high_21_1, %iq2s_word0_high_21_2, %iq2s_word0_high_21_3, %iq2s_word0_high_21_4, %iq2s_word0_high_21_5, %iq2s_word0_high_21_6, %iq2s_word0_high_21_7, %iq2s_word0_high_21_8, %iq2s_word0_high_21_9, %iq2s_word0_high_21_10, %iq2s_word0_high_21_11, %iq2s_word0_high_21_12, %iq2s_word0_high_21_13, %iq2s_word0_high_21_14, %iq2s_word0_high_21_15, %iq2s_word0_high_21_16, %iq2s_word0_high_21_17, %iq2s_word0_high_21_18, %iq2s_word0_high_21_19, %iq2s_word0_high_21_20, %iq2s_word0_high_21_21, %iq2s_word0_high_21_22, %iq2s_word0_high_21_23, %iq2s_word0_high_21_24, %iq2s_word0_high_21_25, %iq2s_word0_high_21_26, %iq2s_word0_high_21_27, %iq2s_word0_high_21_28, %iq2s_word0_high_21_29, %iq2s_word0_high_21_30, %iq2s_word0_high_21_31 : vector<32xf32> + %iq2s_word0_high_22_0 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_22_1 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_22_2 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_22_3 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_22_4 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_22_5 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_22_6 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_22_7 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_22_8 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_22_9 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_22_10 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_22_11 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_22_12 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_22_13 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_22_14 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_22_15 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_22_16 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_22_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_22_18 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_22_19 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_22_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_22_21 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_22_22 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_22_23 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_22_24 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_22_25 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_22_26 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_22_27 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_22_28 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_22_29 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_22_30 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_22_31 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_22 = vector.from_elements %iq2s_word0_high_22_0, %iq2s_word0_high_22_1, %iq2s_word0_high_22_2, %iq2s_word0_high_22_3, %iq2s_word0_high_22_4, %iq2s_word0_high_22_5, %iq2s_word0_high_22_6, %iq2s_word0_high_22_7, %iq2s_word0_high_22_8, %iq2s_word0_high_22_9, %iq2s_word0_high_22_10, %iq2s_word0_high_22_11, %iq2s_word0_high_22_12, %iq2s_word0_high_22_13, %iq2s_word0_high_22_14, %iq2s_word0_high_22_15, %iq2s_word0_high_22_16, %iq2s_word0_high_22_17, %iq2s_word0_high_22_18, %iq2s_word0_high_22_19, %iq2s_word0_high_22_20, %iq2s_word0_high_22_21, %iq2s_word0_high_22_22, %iq2s_word0_high_22_23, %iq2s_word0_high_22_24, %iq2s_word0_high_22_25, %iq2s_word0_high_22_26, %iq2s_word0_high_22_27, %iq2s_word0_high_22_28, %iq2s_word0_high_22_29, %iq2s_word0_high_22_30, %iq2s_word0_high_22_31 : vector<32xf32> + %iq2s_word0_high_23_0 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_23_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_23_2 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_23_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_23_4 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_23_5 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_23_6 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_23_7 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_23_8 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_23_9 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_23_10 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_23_11 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_23_12 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_23_13 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_23_14 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_23_15 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_23_16 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_23_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_23_18 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_23_19 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_23_20 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_23_21 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_23_22 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_23_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_23_24 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_23_25 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_23_26 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_23_27 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_23_28 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_23_29 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_23_30 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_23_31 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_23 = vector.from_elements %iq2s_word0_high_23_0, %iq2s_word0_high_23_1, %iq2s_word0_high_23_2, %iq2s_word0_high_23_3, %iq2s_word0_high_23_4, %iq2s_word0_high_23_5, %iq2s_word0_high_23_6, %iq2s_word0_high_23_7, %iq2s_word0_high_23_8, %iq2s_word0_high_23_9, %iq2s_word0_high_23_10, %iq2s_word0_high_23_11, %iq2s_word0_high_23_12, %iq2s_word0_high_23_13, %iq2s_word0_high_23_14, %iq2s_word0_high_23_15, %iq2s_word0_high_23_16, %iq2s_word0_high_23_17, %iq2s_word0_high_23_18, %iq2s_word0_high_23_19, %iq2s_word0_high_23_20, %iq2s_word0_high_23_21, %iq2s_word0_high_23_22, %iq2s_word0_high_23_23, %iq2s_word0_high_23_24, %iq2s_word0_high_23_25, %iq2s_word0_high_23_26, %iq2s_word0_high_23_27, %iq2s_word0_high_23_28, %iq2s_word0_high_23_29, %iq2s_word0_high_23_30, %iq2s_word0_high_23_31 : vector<32xf32> + %iq2s_word0_high_24_0 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_24_1 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_24_2 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_24_3 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_24_4 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_24_5 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_24_6 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_24_7 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_24_8 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_24_9 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_24_10 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_24_11 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_24_12 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_24_13 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_24_14 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_24_15 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_24_16 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_24_17 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_24_18 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_24_19 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_24_20 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_24_21 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_24_22 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_24_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_24_24 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_24_25 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_24_26 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_24_27 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_24_28 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_24_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_24_30 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_24_31 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_24 = vector.from_elements %iq2s_word0_high_24_0, %iq2s_word0_high_24_1, %iq2s_word0_high_24_2, %iq2s_word0_high_24_3, %iq2s_word0_high_24_4, %iq2s_word0_high_24_5, %iq2s_word0_high_24_6, %iq2s_word0_high_24_7, %iq2s_word0_high_24_8, %iq2s_word0_high_24_9, %iq2s_word0_high_24_10, %iq2s_word0_high_24_11, %iq2s_word0_high_24_12, %iq2s_word0_high_24_13, %iq2s_word0_high_24_14, %iq2s_word0_high_24_15, %iq2s_word0_high_24_16, %iq2s_word0_high_24_17, %iq2s_word0_high_24_18, %iq2s_word0_high_24_19, %iq2s_word0_high_24_20, %iq2s_word0_high_24_21, %iq2s_word0_high_24_22, %iq2s_word0_high_24_23, %iq2s_word0_high_24_24, %iq2s_word0_high_24_25, %iq2s_word0_high_24_26, %iq2s_word0_high_24_27, %iq2s_word0_high_24_28, %iq2s_word0_high_24_29, %iq2s_word0_high_24_30, %iq2s_word0_high_24_31 : vector<32xf32> + %iq2s_word0_high_25_0 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_25_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_25_2 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_25_3 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_25_4 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_25_5 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_25_6 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_25_7 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_25_8 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_25_9 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_25_10 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_25_11 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_25_12 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_25_13 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_25_14 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_25_15 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_25_16 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_25_17 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_25_18 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_25_19 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_25_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_25_21 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_25_22 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_25_23 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_25_24 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_25_25 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_25_26 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_25_27 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_25_28 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_25_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_25_30 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_25_31 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_25 = vector.from_elements %iq2s_word0_high_25_0, %iq2s_word0_high_25_1, %iq2s_word0_high_25_2, %iq2s_word0_high_25_3, %iq2s_word0_high_25_4, %iq2s_word0_high_25_5, %iq2s_word0_high_25_6, %iq2s_word0_high_25_7, %iq2s_word0_high_25_8, %iq2s_word0_high_25_9, %iq2s_word0_high_25_10, %iq2s_word0_high_25_11, %iq2s_word0_high_25_12, %iq2s_word0_high_25_13, %iq2s_word0_high_25_14, %iq2s_word0_high_25_15, %iq2s_word0_high_25_16, %iq2s_word0_high_25_17, %iq2s_word0_high_25_18, %iq2s_word0_high_25_19, %iq2s_word0_high_25_20, %iq2s_word0_high_25_21, %iq2s_word0_high_25_22, %iq2s_word0_high_25_23, %iq2s_word0_high_25_24, %iq2s_word0_high_25_25, %iq2s_word0_high_25_26, %iq2s_word0_high_25_27, %iq2s_word0_high_25_28, %iq2s_word0_high_25_29, %iq2s_word0_high_25_30, %iq2s_word0_high_25_31 : vector<32xf32> + %iq2s_word0_high_26_0 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_26_1 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_26_2 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_26_3 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_26_4 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_26_5 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_26_6 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_26_7 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_26_8 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_26_9 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_26_10 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_26_11 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_26_12 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_26_13 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_26_14 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_26_15 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_26_16 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_26_17 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_26_18 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_26_19 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_26_20 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_26_21 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_26_22 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_26_23 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_26_24 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_26_25 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_26_26 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_26_27 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_26_28 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_26_29 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_26_30 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_26_31 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_26 = vector.from_elements %iq2s_word0_high_26_0, %iq2s_word0_high_26_1, %iq2s_word0_high_26_2, %iq2s_word0_high_26_3, %iq2s_word0_high_26_4, %iq2s_word0_high_26_5, %iq2s_word0_high_26_6, %iq2s_word0_high_26_7, %iq2s_word0_high_26_8, %iq2s_word0_high_26_9, %iq2s_word0_high_26_10, %iq2s_word0_high_26_11, %iq2s_word0_high_26_12, %iq2s_word0_high_26_13, %iq2s_word0_high_26_14, %iq2s_word0_high_26_15, %iq2s_word0_high_26_16, %iq2s_word0_high_26_17, %iq2s_word0_high_26_18, %iq2s_word0_high_26_19, %iq2s_word0_high_26_20, %iq2s_word0_high_26_21, %iq2s_word0_high_26_22, %iq2s_word0_high_26_23, %iq2s_word0_high_26_24, %iq2s_word0_high_26_25, %iq2s_word0_high_26_26, %iq2s_word0_high_26_27, %iq2s_word0_high_26_28, %iq2s_word0_high_26_29, %iq2s_word0_high_26_30, %iq2s_word0_high_26_31 : vector<32xf32> + %iq2s_word0_high_27_0 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_27_1 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_27_2 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_27_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_27_4 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_27_5 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_27_6 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_27_7 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_27_8 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_27_9 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_27_10 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_27_11 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_27_12 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_27_13 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_27_14 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_27_15 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_27_16 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_27_17 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_27_18 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_27_19 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_27_20 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_27_21 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_27_22 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_27_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_27_24 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_27_25 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_27_26 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_27_27 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_27_28 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_27_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_27_30 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_27_31 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_27 = vector.from_elements %iq2s_word0_high_27_0, %iq2s_word0_high_27_1, %iq2s_word0_high_27_2, %iq2s_word0_high_27_3, %iq2s_word0_high_27_4, %iq2s_word0_high_27_5, %iq2s_word0_high_27_6, %iq2s_word0_high_27_7, %iq2s_word0_high_27_8, %iq2s_word0_high_27_9, %iq2s_word0_high_27_10, %iq2s_word0_high_27_11, %iq2s_word0_high_27_12, %iq2s_word0_high_27_13, %iq2s_word0_high_27_14, %iq2s_word0_high_27_15, %iq2s_word0_high_27_16, %iq2s_word0_high_27_17, %iq2s_word0_high_27_18, %iq2s_word0_high_27_19, %iq2s_word0_high_27_20, %iq2s_word0_high_27_21, %iq2s_word0_high_27_22, %iq2s_word0_high_27_23, %iq2s_word0_high_27_24, %iq2s_word0_high_27_25, %iq2s_word0_high_27_26, %iq2s_word0_high_27_27, %iq2s_word0_high_27_28, %iq2s_word0_high_27_29, %iq2s_word0_high_27_30, %iq2s_word0_high_27_31 : vector<32xf32> + %iq2s_word0_high_28_0 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_28_1 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_28_2 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_28_3 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_28_4 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_28_5 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_28_6 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_28_7 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_28_8 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_28_9 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_28_10 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_28_11 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_28_12 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_28_13 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_28_14 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_28_15 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_28_16 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_28_17 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_28_18 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_28_19 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_28_20 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_28_21 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_28_22 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_28_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_28_24 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_28_25 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_28_26 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_28_27 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_28_28 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_28_29 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_28_30 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_28_31 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_28 = vector.from_elements %iq2s_word0_high_28_0, %iq2s_word0_high_28_1, %iq2s_word0_high_28_2, %iq2s_word0_high_28_3, %iq2s_word0_high_28_4, %iq2s_word0_high_28_5, %iq2s_word0_high_28_6, %iq2s_word0_high_28_7, %iq2s_word0_high_28_8, %iq2s_word0_high_28_9, %iq2s_word0_high_28_10, %iq2s_word0_high_28_11, %iq2s_word0_high_28_12, %iq2s_word0_high_28_13, %iq2s_word0_high_28_14, %iq2s_word0_high_28_15, %iq2s_word0_high_28_16, %iq2s_word0_high_28_17, %iq2s_word0_high_28_18, %iq2s_word0_high_28_19, %iq2s_word0_high_28_20, %iq2s_word0_high_28_21, %iq2s_word0_high_28_22, %iq2s_word0_high_28_23, %iq2s_word0_high_28_24, %iq2s_word0_high_28_25, %iq2s_word0_high_28_26, %iq2s_word0_high_28_27, %iq2s_word0_high_28_28, %iq2s_word0_high_28_29, %iq2s_word0_high_28_30, %iq2s_word0_high_28_31 : vector<32xf32> + %iq2s_word0_high_29_0 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_29_1 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_29_2 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_29_3 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_29_4 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_29_5 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_29_6 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_29_7 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_29_8 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_29_9 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_29_10 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_29_11 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_29_12 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_29_13 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_29_14 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_29_15 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_29_16 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_29_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_29_18 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_29_19 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_29_20 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_29_21 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_29_22 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_29_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_29_24 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_29_25 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_29_26 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_29_27 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_29_28 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_29_29 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_29_30 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_29_31 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_29 = vector.from_elements %iq2s_word0_high_29_0, %iq2s_word0_high_29_1, %iq2s_word0_high_29_2, %iq2s_word0_high_29_3, %iq2s_word0_high_29_4, %iq2s_word0_high_29_5, %iq2s_word0_high_29_6, %iq2s_word0_high_29_7, %iq2s_word0_high_29_8, %iq2s_word0_high_29_9, %iq2s_word0_high_29_10, %iq2s_word0_high_29_11, %iq2s_word0_high_29_12, %iq2s_word0_high_29_13, %iq2s_word0_high_29_14, %iq2s_word0_high_29_15, %iq2s_word0_high_29_16, %iq2s_word0_high_29_17, %iq2s_word0_high_29_18, %iq2s_word0_high_29_19, %iq2s_word0_high_29_20, %iq2s_word0_high_29_21, %iq2s_word0_high_29_22, %iq2s_word0_high_29_23, %iq2s_word0_high_29_24, %iq2s_word0_high_29_25, %iq2s_word0_high_29_26, %iq2s_word0_high_29_27, %iq2s_word0_high_29_28, %iq2s_word0_high_29_29, %iq2s_word0_high_29_30, %iq2s_word0_high_29_31 : vector<32xf32> + %iq2s_word0_high_30_0 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_30_1 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_30_2 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_30_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_30_4 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_30_5 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_30_6 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_30_7 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_30_8 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_30_9 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_30_10 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_30_11 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_30_12 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_30_13 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_30_14 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_30_15 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_30_16 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_30_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_30_18 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_30_19 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_30_20 = scalar.constant 11033.0 : f32 + %iq2s_word0_high_30_21 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_30_22 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_30_23 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_30_24 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_30_25 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_30_26 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_30_27 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_30_28 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_30_29 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_30_30 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_30_31 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_30 = vector.from_elements %iq2s_word0_high_30_0, %iq2s_word0_high_30_1, %iq2s_word0_high_30_2, %iq2s_word0_high_30_3, %iq2s_word0_high_30_4, %iq2s_word0_high_30_5, %iq2s_word0_high_30_6, %iq2s_word0_high_30_7, %iq2s_word0_high_30_8, %iq2s_word0_high_30_9, %iq2s_word0_high_30_10, %iq2s_word0_high_30_11, %iq2s_word0_high_30_12, %iq2s_word0_high_30_13, %iq2s_word0_high_30_14, %iq2s_word0_high_30_15, %iq2s_word0_high_30_16, %iq2s_word0_high_30_17, %iq2s_word0_high_30_18, %iq2s_word0_high_30_19, %iq2s_word0_high_30_20, %iq2s_word0_high_30_21, %iq2s_word0_high_30_22, %iq2s_word0_high_30_23, %iq2s_word0_high_30_24, %iq2s_word0_high_30_25, %iq2s_word0_high_30_26, %iq2s_word0_high_30_27, %iq2s_word0_high_30_28, %iq2s_word0_high_30_29, %iq2s_word0_high_30_30, %iq2s_word0_high_30_31 : vector<32xf32> + %iq2s_word0_high_31_0 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_31_1 = scalar.constant 6443.0 : f32 + %iq2s_word0_high_31_2 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_31_3 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_31_4 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_31_5 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_31_6 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_31_7 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_31_8 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_31_9 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_31_10 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_31_11 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_31_12 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_31_13 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_31_14 = scalar.constant 6408.0 : f32 + %iq2s_word0_high_31_15 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_31_16 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_31_17 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_31_18 = scalar.constant 2073.0 : f32 + %iq2s_word0_high_31_19 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_31_20 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_31_21 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_31_22 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_31_23 = scalar.constant 6425.0 : f32 + %iq2s_word0_high_31_24 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_31_25 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_31_26 = scalar.constant 2056.0 : f32 + %iq2s_word0_high_31_27 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_31_28 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_31_29 = scalar.constant 2091.0 : f32 + %iq2s_word0_high_31_30 = scalar.constant 11016.0 : f32 + %iq2s_word0_high_31_31 = scalar.constant 11051.0 : f32 + %iq2s_word0_high_31 = vector.from_elements %iq2s_word0_high_31_0, %iq2s_word0_high_31_1, %iq2s_word0_high_31_2, %iq2s_word0_high_31_3, %iq2s_word0_high_31_4, %iq2s_word0_high_31_5, %iq2s_word0_high_31_6, %iq2s_word0_high_31_7, %iq2s_word0_high_31_8, %iq2s_word0_high_31_9, %iq2s_word0_high_31_10, %iq2s_word0_high_31_11, %iq2s_word0_high_31_12, %iq2s_word0_high_31_13, %iq2s_word0_high_31_14, %iq2s_word0_high_31_15, %iq2s_word0_high_31_16, %iq2s_word0_high_31_17, %iq2s_word0_high_31_18, %iq2s_word0_high_31_19, %iq2s_word0_high_31_20, %iq2s_word0_high_31_21, %iq2s_word0_high_31_22, %iq2s_word0_high_31_23, %iq2s_word0_high_31_24, %iq2s_word0_high_31_25, %iq2s_word0_high_31_26, %iq2s_word0_high_31_27, %iq2s_word0_high_31_28, %iq2s_word0_high_31_29, %iq2s_word0_high_31_30, %iq2s_word0_high_31_31 : vector<32xf32> + %selected_word0_high1 = scf.select %is_chunk1, %iq2s_word0_high_1, %iq2s_word0_high_0 : vector<32xf32> + %selected_word0_high2 = scf.select %is_chunk2, %iq2s_word0_high_2, %selected_word0_high1 : vector<32xf32> + %selected_word0_high3 = scf.select %is_chunk3, %iq2s_word0_high_3, %selected_word0_high2 : vector<32xf32> + %selected_word0_high4 = scf.select %is_chunk4, %iq2s_word0_high_4, %selected_word0_high3 : vector<32xf32> + %selected_word0_high5 = scf.select %is_chunk5, %iq2s_word0_high_5, %selected_word0_high4 : vector<32xf32> + %selected_word0_high6 = scf.select %is_chunk6, %iq2s_word0_high_6, %selected_word0_high5 : vector<32xf32> + %selected_word0_high7 = scf.select %is_chunk7, %iq2s_word0_high_7, %selected_word0_high6 : vector<32xf32> + %selected_word0_high8 = scf.select %is_chunk8, %iq2s_word0_high_8, %selected_word0_high7 : vector<32xf32> + %selected_word0_high9 = scf.select %is_chunk9, %iq2s_word0_high_9, %selected_word0_high8 : vector<32xf32> + %selected_word0_high10 = scf.select %is_chunk10, %iq2s_word0_high_10, %selected_word0_high9 : vector<32xf32> + %selected_word0_high11 = scf.select %is_chunk11, %iq2s_word0_high_11, %selected_word0_high10 : vector<32xf32> + %selected_word0_high12 = scf.select %is_chunk12, %iq2s_word0_high_12, %selected_word0_high11 : vector<32xf32> + %selected_word0_high13 = scf.select %is_chunk13, %iq2s_word0_high_13, %selected_word0_high12 : vector<32xf32> + %selected_word0_high14 = scf.select %is_chunk14, %iq2s_word0_high_14, %selected_word0_high13 : vector<32xf32> + %selected_word0_high15 = scf.select %is_chunk15, %iq2s_word0_high_15, %selected_word0_high14 : vector<32xf32> + %selected_word0_high16 = scf.select %is_chunk16, %iq2s_word0_high_16, %selected_word0_high15 : vector<32xf32> + %selected_word0_high17 = scf.select %is_chunk17, %iq2s_word0_high_17, %selected_word0_high16 : vector<32xf32> + %selected_word0_high18 = scf.select %is_chunk18, %iq2s_word0_high_18, %selected_word0_high17 : vector<32xf32> + %selected_word0_high19 = scf.select %is_chunk19, %iq2s_word0_high_19, %selected_word0_high18 : vector<32xf32> + %selected_word0_high20 = scf.select %is_chunk20, %iq2s_word0_high_20, %selected_word0_high19 : vector<32xf32> + %selected_word0_high21 = scf.select %is_chunk21, %iq2s_word0_high_21, %selected_word0_high20 : vector<32xf32> + %selected_word0_high22 = scf.select %is_chunk22, %iq2s_word0_high_22, %selected_word0_high21 : vector<32xf32> + %selected_word0_high23 = scf.select %is_chunk23, %iq2s_word0_high_23, %selected_word0_high22 : vector<32xf32> + %selected_word0_high24 = scf.select %is_chunk24, %iq2s_word0_high_24, %selected_word0_high23 : vector<32xf32> + %selected_word0_high25 = scf.select %is_chunk25, %iq2s_word0_high_25, %selected_word0_high24 : vector<32xf32> + %selected_word0_high26 = scf.select %is_chunk26, %iq2s_word0_high_26, %selected_word0_high25 : vector<32xf32> + %selected_word0_high27 = scf.select %is_chunk27, %iq2s_word0_high_27, %selected_word0_high26 : vector<32xf32> + %selected_word0_high28 = scf.select %is_chunk28, %iq2s_word0_high_28, %selected_word0_high27 : vector<32xf32> + %selected_word0_high29 = scf.select %is_chunk29, %iq2s_word0_high_29, %selected_word0_high28 : vector<32xf32> + %selected_word0_high30 = scf.select %is_chunk30, %iq2s_word0_high_30, %selected_word0_high29 : vector<32xf32> + %selected_word0_high31 = scf.select %is_chunk31, %iq2s_word0_high_31, %selected_word0_high30 : vector<32xf32> + %iq2s_word1_low_0_0 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_1 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_2 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_3 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_4 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_5 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_6 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_7 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_8 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_9 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_10 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_11 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_12 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_13 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_14 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_15 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_16 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_17 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_18 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_19 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_20 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_21 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_22 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_23 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_24 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_25 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_26 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_27 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_28 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_29 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_30 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0_31 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_0 = vector.from_elements %iq2s_word1_low_0_0, %iq2s_word1_low_0_1, %iq2s_word1_low_0_2, %iq2s_word1_low_0_3, %iq2s_word1_low_0_4, %iq2s_word1_low_0_5, %iq2s_word1_low_0_6, %iq2s_word1_low_0_7, %iq2s_word1_low_0_8, %iq2s_word1_low_0_9, %iq2s_word1_low_0_10, %iq2s_word1_low_0_11, %iq2s_word1_low_0_12, %iq2s_word1_low_0_13, %iq2s_word1_low_0_14, %iq2s_word1_low_0_15, %iq2s_word1_low_0_16, %iq2s_word1_low_0_17, %iq2s_word1_low_0_18, %iq2s_word1_low_0_19, %iq2s_word1_low_0_20, %iq2s_word1_low_0_21, %iq2s_word1_low_0_22, %iq2s_word1_low_0_23, %iq2s_word1_low_0_24, %iq2s_word1_low_0_25, %iq2s_word1_low_0_26, %iq2s_word1_low_0_27, %iq2s_word1_low_0_28, %iq2s_word1_low_0_29, %iq2s_word1_low_0_30, %iq2s_word1_low_0_31 : vector<32xf32> + %iq2s_word1_low_1_0 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_1_1 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_1_2 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_3 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_4 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_5 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_6 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_7 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_8 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_9 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_10 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_11 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_12 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_13 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_14 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_15 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_16 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_17 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_18 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_19 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_20 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_21 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_22 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_23 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_24 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_25 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_26 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_27 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_28 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_29 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_1_30 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_1_31 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_1 = vector.from_elements %iq2s_word1_low_1_0, %iq2s_word1_low_1_1, %iq2s_word1_low_1_2, %iq2s_word1_low_1_3, %iq2s_word1_low_1_4, %iq2s_word1_low_1_5, %iq2s_word1_low_1_6, %iq2s_word1_low_1_7, %iq2s_word1_low_1_8, %iq2s_word1_low_1_9, %iq2s_word1_low_1_10, %iq2s_word1_low_1_11, %iq2s_word1_low_1_12, %iq2s_word1_low_1_13, %iq2s_word1_low_1_14, %iq2s_word1_low_1_15, %iq2s_word1_low_1_16, %iq2s_word1_low_1_17, %iq2s_word1_low_1_18, %iq2s_word1_low_1_19, %iq2s_word1_low_1_20, %iq2s_word1_low_1_21, %iq2s_word1_low_1_22, %iq2s_word1_low_1_23, %iq2s_word1_low_1_24, %iq2s_word1_low_1_25, %iq2s_word1_low_1_26, %iq2s_word1_low_1_27, %iq2s_word1_low_1_28, %iq2s_word1_low_1_29, %iq2s_word1_low_1_30, %iq2s_word1_low_1_31 : vector<32xf32> + %iq2s_word1_low_2_0 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_2_1 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_2_2 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_2_3 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_2_4 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_2_5 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_2_6 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_2_7 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_2_8 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_2_9 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_2_10 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_2_11 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_2_12 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_2_13 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_2_14 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_2_15 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_2_16 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_2_17 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_2_18 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_2_19 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_2_20 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_2_21 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_2_22 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_2_23 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_2_24 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_2_25 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_2_26 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_2_27 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_2_28 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_2_29 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_2_30 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_2_31 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_2 = vector.from_elements %iq2s_word1_low_2_0, %iq2s_word1_low_2_1, %iq2s_word1_low_2_2, %iq2s_word1_low_2_3, %iq2s_word1_low_2_4, %iq2s_word1_low_2_5, %iq2s_word1_low_2_6, %iq2s_word1_low_2_7, %iq2s_word1_low_2_8, %iq2s_word1_low_2_9, %iq2s_word1_low_2_10, %iq2s_word1_low_2_11, %iq2s_word1_low_2_12, %iq2s_word1_low_2_13, %iq2s_word1_low_2_14, %iq2s_word1_low_2_15, %iq2s_word1_low_2_16, %iq2s_word1_low_2_17, %iq2s_word1_low_2_18, %iq2s_word1_low_2_19, %iq2s_word1_low_2_20, %iq2s_word1_low_2_21, %iq2s_word1_low_2_22, %iq2s_word1_low_2_23, %iq2s_word1_low_2_24, %iq2s_word1_low_2_25, %iq2s_word1_low_2_26, %iq2s_word1_low_2_27, %iq2s_word1_low_2_28, %iq2s_word1_low_2_29, %iq2s_word1_low_2_30, %iq2s_word1_low_2_31 : vector<32xf32> + %iq2s_word1_low_3_0 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_1 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_2 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_3 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_4 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_5 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_6 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_7 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_8 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_9 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_10 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_11 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_12 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_13 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_14 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_15 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_16 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_17 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_3_18 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_3_19 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_3_20 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_3_21 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_3_22 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_3_23 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_3_24 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_3_25 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_3_26 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_3_27 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_3_28 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_3_29 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_3_30 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_3_31 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_3 = vector.from_elements %iq2s_word1_low_3_0, %iq2s_word1_low_3_1, %iq2s_word1_low_3_2, %iq2s_word1_low_3_3, %iq2s_word1_low_3_4, %iq2s_word1_low_3_5, %iq2s_word1_low_3_6, %iq2s_word1_low_3_7, %iq2s_word1_low_3_8, %iq2s_word1_low_3_9, %iq2s_word1_low_3_10, %iq2s_word1_low_3_11, %iq2s_word1_low_3_12, %iq2s_word1_low_3_13, %iq2s_word1_low_3_14, %iq2s_word1_low_3_15, %iq2s_word1_low_3_16, %iq2s_word1_low_3_17, %iq2s_word1_low_3_18, %iq2s_word1_low_3_19, %iq2s_word1_low_3_20, %iq2s_word1_low_3_21, %iq2s_word1_low_3_22, %iq2s_word1_low_3_23, %iq2s_word1_low_3_24, %iq2s_word1_low_3_25, %iq2s_word1_low_3_26, %iq2s_word1_low_3_27, %iq2s_word1_low_3_28, %iq2s_word1_low_3_29, %iq2s_word1_low_3_30, %iq2s_word1_low_3_31 : vector<32xf32> + %iq2s_word1_low_4_0 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_4_1 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_4_2 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_4_3 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_4_4 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_4_5 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_4_6 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_4_7 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_4_8 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_4_9 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_4_10 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_4_11 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_4_12 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_4_13 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_4_14 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_4_15 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_4_16 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_4_17 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_4_18 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_4_19 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_4_20 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_4_21 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_4_22 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_4_23 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_4_24 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_4_25 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_4_26 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_4_27 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_4_28 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_4_29 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_4_30 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_4_31 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_4 = vector.from_elements %iq2s_word1_low_4_0, %iq2s_word1_low_4_1, %iq2s_word1_low_4_2, %iq2s_word1_low_4_3, %iq2s_word1_low_4_4, %iq2s_word1_low_4_5, %iq2s_word1_low_4_6, %iq2s_word1_low_4_7, %iq2s_word1_low_4_8, %iq2s_word1_low_4_9, %iq2s_word1_low_4_10, %iq2s_word1_low_4_11, %iq2s_word1_low_4_12, %iq2s_word1_low_4_13, %iq2s_word1_low_4_14, %iq2s_word1_low_4_15, %iq2s_word1_low_4_16, %iq2s_word1_low_4_17, %iq2s_word1_low_4_18, %iq2s_word1_low_4_19, %iq2s_word1_low_4_20, %iq2s_word1_low_4_21, %iq2s_word1_low_4_22, %iq2s_word1_low_4_23, %iq2s_word1_low_4_24, %iq2s_word1_low_4_25, %iq2s_word1_low_4_26, %iq2s_word1_low_4_27, %iq2s_word1_low_4_28, %iq2s_word1_low_4_29, %iq2s_word1_low_4_30, %iq2s_word1_low_4_31 : vector<32xf32> + %iq2s_word1_low_5_0 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_1 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_2 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_3 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_4 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_5 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_6 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_7 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_8 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_9 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_10 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_11 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_12 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_13 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_14 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_15 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_16 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_17 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_18 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_19 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_20 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_21 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_22 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_5_23 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_5_24 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_5_25 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_5_26 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_5_27 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_5_28 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_5_29 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_5_30 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_5_31 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_5 = vector.from_elements %iq2s_word1_low_5_0, %iq2s_word1_low_5_1, %iq2s_word1_low_5_2, %iq2s_word1_low_5_3, %iq2s_word1_low_5_4, %iq2s_word1_low_5_5, %iq2s_word1_low_5_6, %iq2s_word1_low_5_7, %iq2s_word1_low_5_8, %iq2s_word1_low_5_9, %iq2s_word1_low_5_10, %iq2s_word1_low_5_11, %iq2s_word1_low_5_12, %iq2s_word1_low_5_13, %iq2s_word1_low_5_14, %iq2s_word1_low_5_15, %iq2s_word1_low_5_16, %iq2s_word1_low_5_17, %iq2s_word1_low_5_18, %iq2s_word1_low_5_19, %iq2s_word1_low_5_20, %iq2s_word1_low_5_21, %iq2s_word1_low_5_22, %iq2s_word1_low_5_23, %iq2s_word1_low_5_24, %iq2s_word1_low_5_25, %iq2s_word1_low_5_26, %iq2s_word1_low_5_27, %iq2s_word1_low_5_28, %iq2s_word1_low_5_29, %iq2s_word1_low_5_30, %iq2s_word1_low_5_31 : vector<32xf32> + %iq2s_word1_low_6_0 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_6_1 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_6_2 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_6_3 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_6_4 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_6_5 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_6_6 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_6_7 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_6_8 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_6_9 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_6_10 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_6_11 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_6_12 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_6_13 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_6_14 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6_15 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6_16 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6_17 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6_18 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6_19 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6_20 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6_21 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6_22 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6_23 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6_24 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6_25 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6_26 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6_27 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6_28 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6_29 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6_30 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6_31 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_6 = vector.from_elements %iq2s_word1_low_6_0, %iq2s_word1_low_6_1, %iq2s_word1_low_6_2, %iq2s_word1_low_6_3, %iq2s_word1_low_6_4, %iq2s_word1_low_6_5, %iq2s_word1_low_6_6, %iq2s_word1_low_6_7, %iq2s_word1_low_6_8, %iq2s_word1_low_6_9, %iq2s_word1_low_6_10, %iq2s_word1_low_6_11, %iq2s_word1_low_6_12, %iq2s_word1_low_6_13, %iq2s_word1_low_6_14, %iq2s_word1_low_6_15, %iq2s_word1_low_6_16, %iq2s_word1_low_6_17, %iq2s_word1_low_6_18, %iq2s_word1_low_6_19, %iq2s_word1_low_6_20, %iq2s_word1_low_6_21, %iq2s_word1_low_6_22, %iq2s_word1_low_6_23, %iq2s_word1_low_6_24, %iq2s_word1_low_6_25, %iq2s_word1_low_6_26, %iq2s_word1_low_6_27, %iq2s_word1_low_6_28, %iq2s_word1_low_6_29, %iq2s_word1_low_6_30, %iq2s_word1_low_6_31 : vector<32xf32> + %iq2s_word1_low_7_0 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_7_1 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_7_2 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_7_3 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_7_4 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_7_5 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_7_6 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_7_7 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_7_8 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_7_9 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_7_10 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_7_11 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_7_12 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_7_13 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_14 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_15 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_16 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_17 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_18 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_19 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_20 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_21 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_22 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_23 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_24 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_25 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_26 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_27 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_28 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_29 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_30 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7_31 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_7 = vector.from_elements %iq2s_word1_low_7_0, %iq2s_word1_low_7_1, %iq2s_word1_low_7_2, %iq2s_word1_low_7_3, %iq2s_word1_low_7_4, %iq2s_word1_low_7_5, %iq2s_word1_low_7_6, %iq2s_word1_low_7_7, %iq2s_word1_low_7_8, %iq2s_word1_low_7_9, %iq2s_word1_low_7_10, %iq2s_word1_low_7_11, %iq2s_word1_low_7_12, %iq2s_word1_low_7_13, %iq2s_word1_low_7_14, %iq2s_word1_low_7_15, %iq2s_word1_low_7_16, %iq2s_word1_low_7_17, %iq2s_word1_low_7_18, %iq2s_word1_low_7_19, %iq2s_word1_low_7_20, %iq2s_word1_low_7_21, %iq2s_word1_low_7_22, %iq2s_word1_low_7_23, %iq2s_word1_low_7_24, %iq2s_word1_low_7_25, %iq2s_word1_low_7_26, %iq2s_word1_low_7_27, %iq2s_word1_low_7_28, %iq2s_word1_low_7_29, %iq2s_word1_low_7_30, %iq2s_word1_low_7_31 : vector<32xf32> + %iq2s_word1_low_8_0 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_8_1 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_8_2 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_8_3 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_8_4 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_8_5 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_8_6 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_8_7 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_8_8 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_8_9 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_8_10 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_8_11 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_8_12 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_8_13 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_8_14 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_8_15 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_8_16 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_8_17 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_8_18 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_8_19 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_8_20 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_8_21 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_8_22 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_8_23 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_8_24 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_8_25 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_8_26 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_8_27 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_8_28 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_8_29 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_8_30 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_8_31 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_8 = vector.from_elements %iq2s_word1_low_8_0, %iq2s_word1_low_8_1, %iq2s_word1_low_8_2, %iq2s_word1_low_8_3, %iq2s_word1_low_8_4, %iq2s_word1_low_8_5, %iq2s_word1_low_8_6, %iq2s_word1_low_8_7, %iq2s_word1_low_8_8, %iq2s_word1_low_8_9, %iq2s_word1_low_8_10, %iq2s_word1_low_8_11, %iq2s_word1_low_8_12, %iq2s_word1_low_8_13, %iq2s_word1_low_8_14, %iq2s_word1_low_8_15, %iq2s_word1_low_8_16, %iq2s_word1_low_8_17, %iq2s_word1_low_8_18, %iq2s_word1_low_8_19, %iq2s_word1_low_8_20, %iq2s_word1_low_8_21, %iq2s_word1_low_8_22, %iq2s_word1_low_8_23, %iq2s_word1_low_8_24, %iq2s_word1_low_8_25, %iq2s_word1_low_8_26, %iq2s_word1_low_8_27, %iq2s_word1_low_8_28, %iq2s_word1_low_8_29, %iq2s_word1_low_8_30, %iq2s_word1_low_8_31 : vector<32xf32> + %iq2s_word1_low_9_0 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_1 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_2 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_3 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_4 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_5 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_6 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_7 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_8 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_9 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_10 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_11 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_12 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_13 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_14 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_15 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_16 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_17 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_18 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_19 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_9_20 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_9_21 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_9_22 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_9_23 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_9_24 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_9_25 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_9_26 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_9_27 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_9_28 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_9_29 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_9_30 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_9_31 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_9 = vector.from_elements %iq2s_word1_low_9_0, %iq2s_word1_low_9_1, %iq2s_word1_low_9_2, %iq2s_word1_low_9_3, %iq2s_word1_low_9_4, %iq2s_word1_low_9_5, %iq2s_word1_low_9_6, %iq2s_word1_low_9_7, %iq2s_word1_low_9_8, %iq2s_word1_low_9_9, %iq2s_word1_low_9_10, %iq2s_word1_low_9_11, %iq2s_word1_low_9_12, %iq2s_word1_low_9_13, %iq2s_word1_low_9_14, %iq2s_word1_low_9_15, %iq2s_word1_low_9_16, %iq2s_word1_low_9_17, %iq2s_word1_low_9_18, %iq2s_word1_low_9_19, %iq2s_word1_low_9_20, %iq2s_word1_low_9_21, %iq2s_word1_low_9_22, %iq2s_word1_low_9_23, %iq2s_word1_low_9_24, %iq2s_word1_low_9_25, %iq2s_word1_low_9_26, %iq2s_word1_low_9_27, %iq2s_word1_low_9_28, %iq2s_word1_low_9_29, %iq2s_word1_low_9_30, %iq2s_word1_low_9_31 : vector<32xf32> + %iq2s_word1_low_10_0 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_10_1 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_10_2 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_10_3 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_10_4 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_10_5 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_10_6 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_10_7 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_10_8 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_10_9 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_10_10 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_10_11 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_10_12 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_10_13 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_10_14 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_10_15 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_10_16 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_10_17 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_10_18 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_10_19 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_10_20 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_10_21 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_10_22 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_10_23 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_10_24 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_10_25 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_10_26 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_10_27 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_10_28 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_10_29 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_10_30 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_10_31 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_10 = vector.from_elements %iq2s_word1_low_10_0, %iq2s_word1_low_10_1, %iq2s_word1_low_10_2, %iq2s_word1_low_10_3, %iq2s_word1_low_10_4, %iq2s_word1_low_10_5, %iq2s_word1_low_10_6, %iq2s_word1_low_10_7, %iq2s_word1_low_10_8, %iq2s_word1_low_10_9, %iq2s_word1_low_10_10, %iq2s_word1_low_10_11, %iq2s_word1_low_10_12, %iq2s_word1_low_10_13, %iq2s_word1_low_10_14, %iq2s_word1_low_10_15, %iq2s_word1_low_10_16, %iq2s_word1_low_10_17, %iq2s_word1_low_10_18, %iq2s_word1_low_10_19, %iq2s_word1_low_10_20, %iq2s_word1_low_10_21, %iq2s_word1_low_10_22, %iq2s_word1_low_10_23, %iq2s_word1_low_10_24, %iq2s_word1_low_10_25, %iq2s_word1_low_10_26, %iq2s_word1_low_10_27, %iq2s_word1_low_10_28, %iq2s_word1_low_10_29, %iq2s_word1_low_10_30, %iq2s_word1_low_10_31 : vector<32xf32> + %iq2s_word1_low_11_0 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_11_1 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_11_2 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_11_3 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_11_4 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_11_5 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_11_6 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_11_7 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_11_8 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_11_9 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_11_10 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_11_11 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_11_12 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_11_13 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_11_14 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_11_15 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_11_16 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_11_17 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_11_18 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_11_19 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_11_20 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_11_21 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_11_22 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_11_23 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_11_24 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_11_25 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_11_26 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_11_27 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_11_28 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_11_29 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_11_30 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_11_31 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_11 = vector.from_elements %iq2s_word1_low_11_0, %iq2s_word1_low_11_1, %iq2s_word1_low_11_2, %iq2s_word1_low_11_3, %iq2s_word1_low_11_4, %iq2s_word1_low_11_5, %iq2s_word1_low_11_6, %iq2s_word1_low_11_7, %iq2s_word1_low_11_8, %iq2s_word1_low_11_9, %iq2s_word1_low_11_10, %iq2s_word1_low_11_11, %iq2s_word1_low_11_12, %iq2s_word1_low_11_13, %iq2s_word1_low_11_14, %iq2s_word1_low_11_15, %iq2s_word1_low_11_16, %iq2s_word1_low_11_17, %iq2s_word1_low_11_18, %iq2s_word1_low_11_19, %iq2s_word1_low_11_20, %iq2s_word1_low_11_21, %iq2s_word1_low_11_22, %iq2s_word1_low_11_23, %iq2s_word1_low_11_24, %iq2s_word1_low_11_25, %iq2s_word1_low_11_26, %iq2s_word1_low_11_27, %iq2s_word1_low_11_28, %iq2s_word1_low_11_29, %iq2s_word1_low_11_30, %iq2s_word1_low_11_31 : vector<32xf32> + %iq2s_word1_low_12_0 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_12_1 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_12_2 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_12_3 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_12_4 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_12_5 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_12_6 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_12_7 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_12_8 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_12_9 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_12_10 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_12_11 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_12_12 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_12_13 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_12_14 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_12_15 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_12_16 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_12_17 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_12_18 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_12_19 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_12_20 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_12_21 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_12_22 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_12_23 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_12_24 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_12_25 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_12_26 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_12_27 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_12_28 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_12_29 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_12_30 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_12_31 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_12 = vector.from_elements %iq2s_word1_low_12_0, %iq2s_word1_low_12_1, %iq2s_word1_low_12_2, %iq2s_word1_low_12_3, %iq2s_word1_low_12_4, %iq2s_word1_low_12_5, %iq2s_word1_low_12_6, %iq2s_word1_low_12_7, %iq2s_word1_low_12_8, %iq2s_word1_low_12_9, %iq2s_word1_low_12_10, %iq2s_word1_low_12_11, %iq2s_word1_low_12_12, %iq2s_word1_low_12_13, %iq2s_word1_low_12_14, %iq2s_word1_low_12_15, %iq2s_word1_low_12_16, %iq2s_word1_low_12_17, %iq2s_word1_low_12_18, %iq2s_word1_low_12_19, %iq2s_word1_low_12_20, %iq2s_word1_low_12_21, %iq2s_word1_low_12_22, %iq2s_word1_low_12_23, %iq2s_word1_low_12_24, %iq2s_word1_low_12_25, %iq2s_word1_low_12_26, %iq2s_word1_low_12_27, %iq2s_word1_low_12_28, %iq2s_word1_low_12_29, %iq2s_word1_low_12_30, %iq2s_word1_low_12_31 : vector<32xf32> + %iq2s_word1_low_13_0 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_13_1 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_13_2 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_3 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_4 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_5 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_6 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_7 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_8 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_9 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_10 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_11 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_12 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_13 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_14 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_15 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_16 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_17 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_18 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_19 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_20 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_13_21 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_13_22 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_13_23 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_13_24 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_13_25 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_13_26 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_13_27 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_13_28 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_13_29 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_13_30 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_13_31 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_13 = vector.from_elements %iq2s_word1_low_13_0, %iq2s_word1_low_13_1, %iq2s_word1_low_13_2, %iq2s_word1_low_13_3, %iq2s_word1_low_13_4, %iq2s_word1_low_13_5, %iq2s_word1_low_13_6, %iq2s_word1_low_13_7, %iq2s_word1_low_13_8, %iq2s_word1_low_13_9, %iq2s_word1_low_13_10, %iq2s_word1_low_13_11, %iq2s_word1_low_13_12, %iq2s_word1_low_13_13, %iq2s_word1_low_13_14, %iq2s_word1_low_13_15, %iq2s_word1_low_13_16, %iq2s_word1_low_13_17, %iq2s_word1_low_13_18, %iq2s_word1_low_13_19, %iq2s_word1_low_13_20, %iq2s_word1_low_13_21, %iq2s_word1_low_13_22, %iq2s_word1_low_13_23, %iq2s_word1_low_13_24, %iq2s_word1_low_13_25, %iq2s_word1_low_13_26, %iq2s_word1_low_13_27, %iq2s_word1_low_13_28, %iq2s_word1_low_13_29, %iq2s_word1_low_13_30, %iq2s_word1_low_13_31 : vector<32xf32> + %iq2s_word1_low_14_0 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_14_1 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_14_2 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_14_3 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_14_4 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_14_5 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_14_6 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_14_7 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_14_8 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_14_9 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_14_10 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_14_11 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_14_12 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_14_13 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_14_14 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_14_15 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_14_16 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_14_17 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_14_18 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_14_19 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_14_20 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_14_21 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_14_22 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_14_23 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_14_24 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_14_25 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_14_26 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_14_27 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_14_28 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_14_29 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_14_30 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_14_31 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_14 = vector.from_elements %iq2s_word1_low_14_0, %iq2s_word1_low_14_1, %iq2s_word1_low_14_2, %iq2s_word1_low_14_3, %iq2s_word1_low_14_4, %iq2s_word1_low_14_5, %iq2s_word1_low_14_6, %iq2s_word1_low_14_7, %iq2s_word1_low_14_8, %iq2s_word1_low_14_9, %iq2s_word1_low_14_10, %iq2s_word1_low_14_11, %iq2s_word1_low_14_12, %iq2s_word1_low_14_13, %iq2s_word1_low_14_14, %iq2s_word1_low_14_15, %iq2s_word1_low_14_16, %iq2s_word1_low_14_17, %iq2s_word1_low_14_18, %iq2s_word1_low_14_19, %iq2s_word1_low_14_20, %iq2s_word1_low_14_21, %iq2s_word1_low_14_22, %iq2s_word1_low_14_23, %iq2s_word1_low_14_24, %iq2s_word1_low_14_25, %iq2s_word1_low_14_26, %iq2s_word1_low_14_27, %iq2s_word1_low_14_28, %iq2s_word1_low_14_29, %iq2s_word1_low_14_30, %iq2s_word1_low_14_31 : vector<32xf32> + %iq2s_word1_low_15_0 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_1 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_2 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_3 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_4 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_5 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_6 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_7 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_8 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_9 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_10 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_11 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_12 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_13 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_14 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_15 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_16 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_17 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_18 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_19 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_20 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_21 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_15_22 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_15_23 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_15_24 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_15_25 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_15_26 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_15_27 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_15_28 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_15_29 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_15_30 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_15_31 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_15 = vector.from_elements %iq2s_word1_low_15_0, %iq2s_word1_low_15_1, %iq2s_word1_low_15_2, %iq2s_word1_low_15_3, %iq2s_word1_low_15_4, %iq2s_word1_low_15_5, %iq2s_word1_low_15_6, %iq2s_word1_low_15_7, %iq2s_word1_low_15_8, %iq2s_word1_low_15_9, %iq2s_word1_low_15_10, %iq2s_word1_low_15_11, %iq2s_word1_low_15_12, %iq2s_word1_low_15_13, %iq2s_word1_low_15_14, %iq2s_word1_low_15_15, %iq2s_word1_low_15_16, %iq2s_word1_low_15_17, %iq2s_word1_low_15_18, %iq2s_word1_low_15_19, %iq2s_word1_low_15_20, %iq2s_word1_low_15_21, %iq2s_word1_low_15_22, %iq2s_word1_low_15_23, %iq2s_word1_low_15_24, %iq2s_word1_low_15_25, %iq2s_word1_low_15_26, %iq2s_word1_low_15_27, %iq2s_word1_low_15_28, %iq2s_word1_low_15_29, %iq2s_word1_low_15_30, %iq2s_word1_low_15_31 : vector<32xf32> + %iq2s_word1_low_16_0 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_1 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_2 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_3 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_4 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_5 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_6 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_7 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_8 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_9 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_10 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_11 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_12 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_13 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_14 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_15 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_16 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_17 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_16_18 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_16_19 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_16_20 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_16_21 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_16_22 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_16_23 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_16_24 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_16_25 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_16_26 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_16_27 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_16_28 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_16_29 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_16_30 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_16_31 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_16 = vector.from_elements %iq2s_word1_low_16_0, %iq2s_word1_low_16_1, %iq2s_word1_low_16_2, %iq2s_word1_low_16_3, %iq2s_word1_low_16_4, %iq2s_word1_low_16_5, %iq2s_word1_low_16_6, %iq2s_word1_low_16_7, %iq2s_word1_low_16_8, %iq2s_word1_low_16_9, %iq2s_word1_low_16_10, %iq2s_word1_low_16_11, %iq2s_word1_low_16_12, %iq2s_word1_low_16_13, %iq2s_word1_low_16_14, %iq2s_word1_low_16_15, %iq2s_word1_low_16_16, %iq2s_word1_low_16_17, %iq2s_word1_low_16_18, %iq2s_word1_low_16_19, %iq2s_word1_low_16_20, %iq2s_word1_low_16_21, %iq2s_word1_low_16_22, %iq2s_word1_low_16_23, %iq2s_word1_low_16_24, %iq2s_word1_low_16_25, %iq2s_word1_low_16_26, %iq2s_word1_low_16_27, %iq2s_word1_low_16_28, %iq2s_word1_low_16_29, %iq2s_word1_low_16_30, %iq2s_word1_low_16_31 : vector<32xf32> + %iq2s_word1_low_17_0 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_17_1 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_2 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_3 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_4 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_5 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_6 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_7 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_8 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_9 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_10 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_11 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_12 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_13 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_14 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_15 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_16 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_17 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_18 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_19 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_20 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_21 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_22 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_23 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_24 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_25 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_26 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_27 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_28 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_29 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_30 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_17_31 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_17 = vector.from_elements %iq2s_word1_low_17_0, %iq2s_word1_low_17_1, %iq2s_word1_low_17_2, %iq2s_word1_low_17_3, %iq2s_word1_low_17_4, %iq2s_word1_low_17_5, %iq2s_word1_low_17_6, %iq2s_word1_low_17_7, %iq2s_word1_low_17_8, %iq2s_word1_low_17_9, %iq2s_word1_low_17_10, %iq2s_word1_low_17_11, %iq2s_word1_low_17_12, %iq2s_word1_low_17_13, %iq2s_word1_low_17_14, %iq2s_word1_low_17_15, %iq2s_word1_low_17_16, %iq2s_word1_low_17_17, %iq2s_word1_low_17_18, %iq2s_word1_low_17_19, %iq2s_word1_low_17_20, %iq2s_word1_low_17_21, %iq2s_word1_low_17_22, %iq2s_word1_low_17_23, %iq2s_word1_low_17_24, %iq2s_word1_low_17_25, %iq2s_word1_low_17_26, %iq2s_word1_low_17_27, %iq2s_word1_low_17_28, %iq2s_word1_low_17_29, %iq2s_word1_low_17_30, %iq2s_word1_low_17_31 : vector<32xf32> + %iq2s_word1_low_18_0 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_1 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_2 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_3 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_4 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_5 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_6 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_7 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_8 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_9 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_10 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_11 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_12 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_13 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_14 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_15 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_16 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_17 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_18 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_19 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_18_20 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_18_21 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_18_22 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_18_23 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_18_24 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_18_25 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_18_26 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_18_27 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_18_28 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_18_29 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_18_30 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_18_31 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_18 = vector.from_elements %iq2s_word1_low_18_0, %iq2s_word1_low_18_1, %iq2s_word1_low_18_2, %iq2s_word1_low_18_3, %iq2s_word1_low_18_4, %iq2s_word1_low_18_5, %iq2s_word1_low_18_6, %iq2s_word1_low_18_7, %iq2s_word1_low_18_8, %iq2s_word1_low_18_9, %iq2s_word1_low_18_10, %iq2s_word1_low_18_11, %iq2s_word1_low_18_12, %iq2s_word1_low_18_13, %iq2s_word1_low_18_14, %iq2s_word1_low_18_15, %iq2s_word1_low_18_16, %iq2s_word1_low_18_17, %iq2s_word1_low_18_18, %iq2s_word1_low_18_19, %iq2s_word1_low_18_20, %iq2s_word1_low_18_21, %iq2s_word1_low_18_22, %iq2s_word1_low_18_23, %iq2s_word1_low_18_24, %iq2s_word1_low_18_25, %iq2s_word1_low_18_26, %iq2s_word1_low_18_27, %iq2s_word1_low_18_28, %iq2s_word1_low_18_29, %iq2s_word1_low_18_30, %iq2s_word1_low_18_31 : vector<32xf32> + %iq2s_word1_low_19_0 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_1 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_2 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_3 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_4 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_5 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_6 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_7 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_8 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_9 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_10 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_11 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_12 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_13 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_14 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_15 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_16 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_17 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_19_18 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_19_19 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_19_20 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_19_21 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_19_22 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_19_23 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_19_24 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_19_25 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_19_26 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_19_27 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_19_28 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_19_29 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_19_30 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_19_31 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_19 = vector.from_elements %iq2s_word1_low_19_0, %iq2s_word1_low_19_1, %iq2s_word1_low_19_2, %iq2s_word1_low_19_3, %iq2s_word1_low_19_4, %iq2s_word1_low_19_5, %iq2s_word1_low_19_6, %iq2s_word1_low_19_7, %iq2s_word1_low_19_8, %iq2s_word1_low_19_9, %iq2s_word1_low_19_10, %iq2s_word1_low_19_11, %iq2s_word1_low_19_12, %iq2s_word1_low_19_13, %iq2s_word1_low_19_14, %iq2s_word1_low_19_15, %iq2s_word1_low_19_16, %iq2s_word1_low_19_17, %iq2s_word1_low_19_18, %iq2s_word1_low_19_19, %iq2s_word1_low_19_20, %iq2s_word1_low_19_21, %iq2s_word1_low_19_22, %iq2s_word1_low_19_23, %iq2s_word1_low_19_24, %iq2s_word1_low_19_25, %iq2s_word1_low_19_26, %iq2s_word1_low_19_27, %iq2s_word1_low_19_28, %iq2s_word1_low_19_29, %iq2s_word1_low_19_30, %iq2s_word1_low_19_31 : vector<32xf32> + %iq2s_word1_low_20_0 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_20_1 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_20_2 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_3 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_4 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_5 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_6 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_7 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_8 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_9 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_10 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_11 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_12 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_13 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_14 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_15 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_16 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_17 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_18 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_19 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_20 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_21 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_22 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_23 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_24 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_25 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_26 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_27 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_28 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_29 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_20_30 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_20_31 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_20 = vector.from_elements %iq2s_word1_low_20_0, %iq2s_word1_low_20_1, %iq2s_word1_low_20_2, %iq2s_word1_low_20_3, %iq2s_word1_low_20_4, %iq2s_word1_low_20_5, %iq2s_word1_low_20_6, %iq2s_word1_low_20_7, %iq2s_word1_low_20_8, %iq2s_word1_low_20_9, %iq2s_word1_low_20_10, %iq2s_word1_low_20_11, %iq2s_word1_low_20_12, %iq2s_word1_low_20_13, %iq2s_word1_low_20_14, %iq2s_word1_low_20_15, %iq2s_word1_low_20_16, %iq2s_word1_low_20_17, %iq2s_word1_low_20_18, %iq2s_word1_low_20_19, %iq2s_word1_low_20_20, %iq2s_word1_low_20_21, %iq2s_word1_low_20_22, %iq2s_word1_low_20_23, %iq2s_word1_low_20_24, %iq2s_word1_low_20_25, %iq2s_word1_low_20_26, %iq2s_word1_low_20_27, %iq2s_word1_low_20_28, %iq2s_word1_low_20_29, %iq2s_word1_low_20_30, %iq2s_word1_low_20_31 : vector<32xf32> + %iq2s_word1_low_21_0 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_1 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_2 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_3 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_4 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_5 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_6 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_7 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_8 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_9 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_10 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_11 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_12 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_13 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_14 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_15 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_16 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_17 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_21_18 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_21_19 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_21_20 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_21_21 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_21_22 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_21_23 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_21_24 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_21_25 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_21_26 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_21_27 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_21_28 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_21_29 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_21_30 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_21_31 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_21 = vector.from_elements %iq2s_word1_low_21_0, %iq2s_word1_low_21_1, %iq2s_word1_low_21_2, %iq2s_word1_low_21_3, %iq2s_word1_low_21_4, %iq2s_word1_low_21_5, %iq2s_word1_low_21_6, %iq2s_word1_low_21_7, %iq2s_word1_low_21_8, %iq2s_word1_low_21_9, %iq2s_word1_low_21_10, %iq2s_word1_low_21_11, %iq2s_word1_low_21_12, %iq2s_word1_low_21_13, %iq2s_word1_low_21_14, %iq2s_word1_low_21_15, %iq2s_word1_low_21_16, %iq2s_word1_low_21_17, %iq2s_word1_low_21_18, %iq2s_word1_low_21_19, %iq2s_word1_low_21_20, %iq2s_word1_low_21_21, %iq2s_word1_low_21_22, %iq2s_word1_low_21_23, %iq2s_word1_low_21_24, %iq2s_word1_low_21_25, %iq2s_word1_low_21_26, %iq2s_word1_low_21_27, %iq2s_word1_low_21_28, %iq2s_word1_low_21_29, %iq2s_word1_low_21_30, %iq2s_word1_low_21_31 : vector<32xf32> + %iq2s_word1_low_22_0 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_22_1 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_22_2 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_22_3 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_22_4 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_22_5 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_22_6 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_22_7 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_22_8 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_22_9 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_22_10 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_22_11 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_22_12 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_22_13 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_22_14 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_22_15 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_22_16 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_22_17 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_22_18 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_22_19 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_22_20 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_22_21 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_22_22 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_22_23 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_22_24 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_22_25 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_22_26 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_22_27 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_22_28 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_22_29 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_22_30 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_22_31 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_22 = vector.from_elements %iq2s_word1_low_22_0, %iq2s_word1_low_22_1, %iq2s_word1_low_22_2, %iq2s_word1_low_22_3, %iq2s_word1_low_22_4, %iq2s_word1_low_22_5, %iq2s_word1_low_22_6, %iq2s_word1_low_22_7, %iq2s_word1_low_22_8, %iq2s_word1_low_22_9, %iq2s_word1_low_22_10, %iq2s_word1_low_22_11, %iq2s_word1_low_22_12, %iq2s_word1_low_22_13, %iq2s_word1_low_22_14, %iq2s_word1_low_22_15, %iq2s_word1_low_22_16, %iq2s_word1_low_22_17, %iq2s_word1_low_22_18, %iq2s_word1_low_22_19, %iq2s_word1_low_22_20, %iq2s_word1_low_22_21, %iq2s_word1_low_22_22, %iq2s_word1_low_22_23, %iq2s_word1_low_22_24, %iq2s_word1_low_22_25, %iq2s_word1_low_22_26, %iq2s_word1_low_22_27, %iq2s_word1_low_22_28, %iq2s_word1_low_22_29, %iq2s_word1_low_22_30, %iq2s_word1_low_22_31 : vector<32xf32> + %iq2s_word1_low_23_0 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_23_1 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_23_2 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_23_3 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_23_4 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_23_5 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_23_6 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_23_7 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_23_8 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_23_9 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_23_10 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_23_11 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_23_12 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_23_13 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_23_14 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_23_15 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_23_16 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_23_17 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_23_18 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_23_19 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_23_20 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_23_21 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_23_22 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_23_23 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_23_24 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_23_25 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_23_26 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_23_27 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_23_28 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_23_29 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_23_30 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_23_31 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_23 = vector.from_elements %iq2s_word1_low_23_0, %iq2s_word1_low_23_1, %iq2s_word1_low_23_2, %iq2s_word1_low_23_3, %iq2s_word1_low_23_4, %iq2s_word1_low_23_5, %iq2s_word1_low_23_6, %iq2s_word1_low_23_7, %iq2s_word1_low_23_8, %iq2s_word1_low_23_9, %iq2s_word1_low_23_10, %iq2s_word1_low_23_11, %iq2s_word1_low_23_12, %iq2s_word1_low_23_13, %iq2s_word1_low_23_14, %iq2s_word1_low_23_15, %iq2s_word1_low_23_16, %iq2s_word1_low_23_17, %iq2s_word1_low_23_18, %iq2s_word1_low_23_19, %iq2s_word1_low_23_20, %iq2s_word1_low_23_21, %iq2s_word1_low_23_22, %iq2s_word1_low_23_23, %iq2s_word1_low_23_24, %iq2s_word1_low_23_25, %iq2s_word1_low_23_26, %iq2s_word1_low_23_27, %iq2s_word1_low_23_28, %iq2s_word1_low_23_29, %iq2s_word1_low_23_30, %iq2s_word1_low_23_31 : vector<32xf32> + %iq2s_word1_low_24_0 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_24_1 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_24_2 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_24_3 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_24_4 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_24_5 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_24_6 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_24_7 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_24_8 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_24_9 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_24_10 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_24_11 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_24_12 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_24_13 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_24_14 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_24_15 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_24_16 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_24_17 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_24_18 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_24_19 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_24_20 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_24_21 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_24_22 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_24_23 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_24_24 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_24_25 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_24_26 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_24_27 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_24_28 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_24_29 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_24_30 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_24_31 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_24 = vector.from_elements %iq2s_word1_low_24_0, %iq2s_word1_low_24_1, %iq2s_word1_low_24_2, %iq2s_word1_low_24_3, %iq2s_word1_low_24_4, %iq2s_word1_low_24_5, %iq2s_word1_low_24_6, %iq2s_word1_low_24_7, %iq2s_word1_low_24_8, %iq2s_word1_low_24_9, %iq2s_word1_low_24_10, %iq2s_word1_low_24_11, %iq2s_word1_low_24_12, %iq2s_word1_low_24_13, %iq2s_word1_low_24_14, %iq2s_word1_low_24_15, %iq2s_word1_low_24_16, %iq2s_word1_low_24_17, %iq2s_word1_low_24_18, %iq2s_word1_low_24_19, %iq2s_word1_low_24_20, %iq2s_word1_low_24_21, %iq2s_word1_low_24_22, %iq2s_word1_low_24_23, %iq2s_word1_low_24_24, %iq2s_word1_low_24_25, %iq2s_word1_low_24_26, %iq2s_word1_low_24_27, %iq2s_word1_low_24_28, %iq2s_word1_low_24_29, %iq2s_word1_low_24_30, %iq2s_word1_low_24_31 : vector<32xf32> + %iq2s_word1_low_25_0 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_25_1 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_25_2 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_25_3 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_25_4 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_25_5 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_25_6 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_25_7 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_25_8 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_25_9 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_25_10 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_25_11 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_25_12 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_25_13 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_25_14 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_25_15 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_25_16 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_25_17 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_25_18 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_25_19 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_25_20 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_25_21 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_25_22 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_25_23 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_25_24 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_25_25 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_25_26 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_25_27 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_25_28 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_25_29 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_25_30 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_25_31 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_25 = vector.from_elements %iq2s_word1_low_25_0, %iq2s_word1_low_25_1, %iq2s_word1_low_25_2, %iq2s_word1_low_25_3, %iq2s_word1_low_25_4, %iq2s_word1_low_25_5, %iq2s_word1_low_25_6, %iq2s_word1_low_25_7, %iq2s_word1_low_25_8, %iq2s_word1_low_25_9, %iq2s_word1_low_25_10, %iq2s_word1_low_25_11, %iq2s_word1_low_25_12, %iq2s_word1_low_25_13, %iq2s_word1_low_25_14, %iq2s_word1_low_25_15, %iq2s_word1_low_25_16, %iq2s_word1_low_25_17, %iq2s_word1_low_25_18, %iq2s_word1_low_25_19, %iq2s_word1_low_25_20, %iq2s_word1_low_25_21, %iq2s_word1_low_25_22, %iq2s_word1_low_25_23, %iq2s_word1_low_25_24, %iq2s_word1_low_25_25, %iq2s_word1_low_25_26, %iq2s_word1_low_25_27, %iq2s_word1_low_25_28, %iq2s_word1_low_25_29, %iq2s_word1_low_25_30, %iq2s_word1_low_25_31 : vector<32xf32> + %iq2s_word1_low_26_0 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_26_1 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_26_2 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_26_3 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_26_4 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_26_5 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_26_6 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_26_7 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_8 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_9 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_10 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_11 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_12 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_13 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_14 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_15 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_16 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_17 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_18 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_19 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_20 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_21 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_22 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_23 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_24 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_25 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_26_26 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_26_27 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_26_28 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_26_29 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_26_30 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_26_31 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_26 = vector.from_elements %iq2s_word1_low_26_0, %iq2s_word1_low_26_1, %iq2s_word1_low_26_2, %iq2s_word1_low_26_3, %iq2s_word1_low_26_4, %iq2s_word1_low_26_5, %iq2s_word1_low_26_6, %iq2s_word1_low_26_7, %iq2s_word1_low_26_8, %iq2s_word1_low_26_9, %iq2s_word1_low_26_10, %iq2s_word1_low_26_11, %iq2s_word1_low_26_12, %iq2s_word1_low_26_13, %iq2s_word1_low_26_14, %iq2s_word1_low_26_15, %iq2s_word1_low_26_16, %iq2s_word1_low_26_17, %iq2s_word1_low_26_18, %iq2s_word1_low_26_19, %iq2s_word1_low_26_20, %iq2s_word1_low_26_21, %iq2s_word1_low_26_22, %iq2s_word1_low_26_23, %iq2s_word1_low_26_24, %iq2s_word1_low_26_25, %iq2s_word1_low_26_26, %iq2s_word1_low_26_27, %iq2s_word1_low_26_28, %iq2s_word1_low_26_29, %iq2s_word1_low_26_30, %iq2s_word1_low_26_31 : vector<32xf32> + %iq2s_word1_low_27_0 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_27_1 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_27_2 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_3 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_4 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_5 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_6 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_7 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_8 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_9 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_10 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_11 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_12 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_13 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_14 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_15 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_16 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_17 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_18 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_19 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_20 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_27_21 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_27_22 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_27_23 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_27_24 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_27_25 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_27_26 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_27_27 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_27_28 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_27_29 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_27_30 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_27_31 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_27 = vector.from_elements %iq2s_word1_low_27_0, %iq2s_word1_low_27_1, %iq2s_word1_low_27_2, %iq2s_word1_low_27_3, %iq2s_word1_low_27_4, %iq2s_word1_low_27_5, %iq2s_word1_low_27_6, %iq2s_word1_low_27_7, %iq2s_word1_low_27_8, %iq2s_word1_low_27_9, %iq2s_word1_low_27_10, %iq2s_word1_low_27_11, %iq2s_word1_low_27_12, %iq2s_word1_low_27_13, %iq2s_word1_low_27_14, %iq2s_word1_low_27_15, %iq2s_word1_low_27_16, %iq2s_word1_low_27_17, %iq2s_word1_low_27_18, %iq2s_word1_low_27_19, %iq2s_word1_low_27_20, %iq2s_word1_low_27_21, %iq2s_word1_low_27_22, %iq2s_word1_low_27_23, %iq2s_word1_low_27_24, %iq2s_word1_low_27_25, %iq2s_word1_low_27_26, %iq2s_word1_low_27_27, %iq2s_word1_low_27_28, %iq2s_word1_low_27_29, %iq2s_word1_low_27_30, %iq2s_word1_low_27_31 : vector<32xf32> + %iq2s_word1_low_28_0 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_28_1 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_28_2 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_28_3 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_28_4 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_28_5 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_28_6 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_28_7 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_28_8 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_28_9 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_28_10 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_28_11 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_28_12 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_28_13 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_28_14 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_28_15 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_28_16 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_28_17 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_28_18 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_28_19 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_28_20 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_28_21 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_28_22 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_28_23 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_28_24 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_28_25 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_28_26 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_28_27 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_28_28 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_28_29 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_28_30 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_28_31 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_28 = vector.from_elements %iq2s_word1_low_28_0, %iq2s_word1_low_28_1, %iq2s_word1_low_28_2, %iq2s_word1_low_28_3, %iq2s_word1_low_28_4, %iq2s_word1_low_28_5, %iq2s_word1_low_28_6, %iq2s_word1_low_28_7, %iq2s_word1_low_28_8, %iq2s_word1_low_28_9, %iq2s_word1_low_28_10, %iq2s_word1_low_28_11, %iq2s_word1_low_28_12, %iq2s_word1_low_28_13, %iq2s_word1_low_28_14, %iq2s_word1_low_28_15, %iq2s_word1_low_28_16, %iq2s_word1_low_28_17, %iq2s_word1_low_28_18, %iq2s_word1_low_28_19, %iq2s_word1_low_28_20, %iq2s_word1_low_28_21, %iq2s_word1_low_28_22, %iq2s_word1_low_28_23, %iq2s_word1_low_28_24, %iq2s_word1_low_28_25, %iq2s_word1_low_28_26, %iq2s_word1_low_28_27, %iq2s_word1_low_28_28, %iq2s_word1_low_28_29, %iq2s_word1_low_28_30, %iq2s_word1_low_28_31 : vector<32xf32> + %iq2s_word1_low_29_0 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_29_1 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_29_2 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_29_3 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_29_4 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_29_5 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_29_6 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_29_7 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_29_8 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_29_9 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_29_10 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_29_11 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_29_12 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_29_13 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_29_14 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_29_15 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_29_16 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_29_17 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_29_18 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_29_19 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_29_20 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_29_21 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_29_22 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_29_23 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_29_24 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_29_25 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_29_26 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_29_27 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_29_28 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_29_29 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_29_30 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_29_31 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_29 = vector.from_elements %iq2s_word1_low_29_0, %iq2s_word1_low_29_1, %iq2s_word1_low_29_2, %iq2s_word1_low_29_3, %iq2s_word1_low_29_4, %iq2s_word1_low_29_5, %iq2s_word1_low_29_6, %iq2s_word1_low_29_7, %iq2s_word1_low_29_8, %iq2s_word1_low_29_9, %iq2s_word1_low_29_10, %iq2s_word1_low_29_11, %iq2s_word1_low_29_12, %iq2s_word1_low_29_13, %iq2s_word1_low_29_14, %iq2s_word1_low_29_15, %iq2s_word1_low_29_16, %iq2s_word1_low_29_17, %iq2s_word1_low_29_18, %iq2s_word1_low_29_19, %iq2s_word1_low_29_20, %iq2s_word1_low_29_21, %iq2s_word1_low_29_22, %iq2s_word1_low_29_23, %iq2s_word1_low_29_24, %iq2s_word1_low_29_25, %iq2s_word1_low_29_26, %iq2s_word1_low_29_27, %iq2s_word1_low_29_28, %iq2s_word1_low_29_29, %iq2s_word1_low_29_30, %iq2s_word1_low_29_31 : vector<32xf32> + %iq2s_word1_low_30_0 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_30_1 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_30_2 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_30_3 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_30_4 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_30_5 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_30_6 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_30_7 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_30_8 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_30_9 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_30_10 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_30_11 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_30_12 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_30_13 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_30_14 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_30_15 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_30_16 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_30_17 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_30_18 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_30_19 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_30_20 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_30_21 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_30_22 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_30_23 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_30_24 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_30_25 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_30_26 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_30_27 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_30_28 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_30_29 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_30_30 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_30_31 = scalar.constant 2056.0 : f32 + %iq2s_word1_low_30 = vector.from_elements %iq2s_word1_low_30_0, %iq2s_word1_low_30_1, %iq2s_word1_low_30_2, %iq2s_word1_low_30_3, %iq2s_word1_low_30_4, %iq2s_word1_low_30_5, %iq2s_word1_low_30_6, %iq2s_word1_low_30_7, %iq2s_word1_low_30_8, %iq2s_word1_low_30_9, %iq2s_word1_low_30_10, %iq2s_word1_low_30_11, %iq2s_word1_low_30_12, %iq2s_word1_low_30_13, %iq2s_word1_low_30_14, %iq2s_word1_low_30_15, %iq2s_word1_low_30_16, %iq2s_word1_low_30_17, %iq2s_word1_low_30_18, %iq2s_word1_low_30_19, %iq2s_word1_low_30_20, %iq2s_word1_low_30_21, %iq2s_word1_low_30_22, %iq2s_word1_low_30_23, %iq2s_word1_low_30_24, %iq2s_word1_low_30_25, %iq2s_word1_low_30_26, %iq2s_word1_low_30_27, %iq2s_word1_low_30_28, %iq2s_word1_low_30_29, %iq2s_word1_low_30_30, %iq2s_word1_low_30_31 : vector<32xf32> + %iq2s_word1_low_31_0 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_31_1 = scalar.constant 2073.0 : f32 + %iq2s_word1_low_31_2 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_31_3 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_31_4 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_31_5 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_31_6 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_31_7 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_31_8 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_31_9 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_31_10 = scalar.constant 2091.0 : f32 + %iq2s_word1_low_31_11 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_31_12 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_31_13 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_31_14 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_31_15 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_31_16 = scalar.constant 6408.0 : f32 + %iq2s_word1_low_31_17 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_31_18 = scalar.constant 6425.0 : f32 + %iq2s_word1_low_31_19 = scalar.constant 6443.0 : f32 + %iq2s_word1_low_31_20 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_31_21 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_31_22 = scalar.constant 11016.0 : f32 + %iq2s_word1_low_31_23 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_31_24 = scalar.constant 11033.0 : f32 + %iq2s_word1_low_31_25 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_31_26 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_31_27 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_31_28 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_31_29 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_31_30 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_31_31 = scalar.constant 11051.0 : f32 + %iq2s_word1_low_31 = vector.from_elements %iq2s_word1_low_31_0, %iq2s_word1_low_31_1, %iq2s_word1_low_31_2, %iq2s_word1_low_31_3, %iq2s_word1_low_31_4, %iq2s_word1_low_31_5, %iq2s_word1_low_31_6, %iq2s_word1_low_31_7, %iq2s_word1_low_31_8, %iq2s_word1_low_31_9, %iq2s_word1_low_31_10, %iq2s_word1_low_31_11, %iq2s_word1_low_31_12, %iq2s_word1_low_31_13, %iq2s_word1_low_31_14, %iq2s_word1_low_31_15, %iq2s_word1_low_31_16, %iq2s_word1_low_31_17, %iq2s_word1_low_31_18, %iq2s_word1_low_31_19, %iq2s_word1_low_31_20, %iq2s_word1_low_31_21, %iq2s_word1_low_31_22, %iq2s_word1_low_31_23, %iq2s_word1_low_31_24, %iq2s_word1_low_31_25, %iq2s_word1_low_31_26, %iq2s_word1_low_31_27, %iq2s_word1_low_31_28, %iq2s_word1_low_31_29, %iq2s_word1_low_31_30, %iq2s_word1_low_31_31 : vector<32xf32> + %selected_word1_low1 = scf.select %is_chunk1, %iq2s_word1_low_1, %iq2s_word1_low_0 : vector<32xf32> + %selected_word1_low2 = scf.select %is_chunk2, %iq2s_word1_low_2, %selected_word1_low1 : vector<32xf32> + %selected_word1_low3 = scf.select %is_chunk3, %iq2s_word1_low_3, %selected_word1_low2 : vector<32xf32> + %selected_word1_low4 = scf.select %is_chunk4, %iq2s_word1_low_4, %selected_word1_low3 : vector<32xf32> + %selected_word1_low5 = scf.select %is_chunk5, %iq2s_word1_low_5, %selected_word1_low4 : vector<32xf32> + %selected_word1_low6 = scf.select %is_chunk6, %iq2s_word1_low_6, %selected_word1_low5 : vector<32xf32> + %selected_word1_low7 = scf.select %is_chunk7, %iq2s_word1_low_7, %selected_word1_low6 : vector<32xf32> + %selected_word1_low8 = scf.select %is_chunk8, %iq2s_word1_low_8, %selected_word1_low7 : vector<32xf32> + %selected_word1_low9 = scf.select %is_chunk9, %iq2s_word1_low_9, %selected_word1_low8 : vector<32xf32> + %selected_word1_low10 = scf.select %is_chunk10, %iq2s_word1_low_10, %selected_word1_low9 : vector<32xf32> + %selected_word1_low11 = scf.select %is_chunk11, %iq2s_word1_low_11, %selected_word1_low10 : vector<32xf32> + %selected_word1_low12 = scf.select %is_chunk12, %iq2s_word1_low_12, %selected_word1_low11 : vector<32xf32> + %selected_word1_low13 = scf.select %is_chunk13, %iq2s_word1_low_13, %selected_word1_low12 : vector<32xf32> + %selected_word1_low14 = scf.select %is_chunk14, %iq2s_word1_low_14, %selected_word1_low13 : vector<32xf32> + %selected_word1_low15 = scf.select %is_chunk15, %iq2s_word1_low_15, %selected_word1_low14 : vector<32xf32> + %selected_word1_low16 = scf.select %is_chunk16, %iq2s_word1_low_16, %selected_word1_low15 : vector<32xf32> + %selected_word1_low17 = scf.select %is_chunk17, %iq2s_word1_low_17, %selected_word1_low16 : vector<32xf32> + %selected_word1_low18 = scf.select %is_chunk18, %iq2s_word1_low_18, %selected_word1_low17 : vector<32xf32> + %selected_word1_low19 = scf.select %is_chunk19, %iq2s_word1_low_19, %selected_word1_low18 : vector<32xf32> + %selected_word1_low20 = scf.select %is_chunk20, %iq2s_word1_low_20, %selected_word1_low19 : vector<32xf32> + %selected_word1_low21 = scf.select %is_chunk21, %iq2s_word1_low_21, %selected_word1_low20 : vector<32xf32> + %selected_word1_low22 = scf.select %is_chunk22, %iq2s_word1_low_22, %selected_word1_low21 : vector<32xf32> + %selected_word1_low23 = scf.select %is_chunk23, %iq2s_word1_low_23, %selected_word1_low22 : vector<32xf32> + %selected_word1_low24 = scf.select %is_chunk24, %iq2s_word1_low_24, %selected_word1_low23 : vector<32xf32> + %selected_word1_low25 = scf.select %is_chunk25, %iq2s_word1_low_25, %selected_word1_low24 : vector<32xf32> + %selected_word1_low26 = scf.select %is_chunk26, %iq2s_word1_low_26, %selected_word1_low25 : vector<32xf32> + %selected_word1_low27 = scf.select %is_chunk27, %iq2s_word1_low_27, %selected_word1_low26 : vector<32xf32> + %selected_word1_low28 = scf.select %is_chunk28, %iq2s_word1_low_28, %selected_word1_low27 : vector<32xf32> + %selected_word1_low29 = scf.select %is_chunk29, %iq2s_word1_low_29, %selected_word1_low28 : vector<32xf32> + %selected_word1_low30 = scf.select %is_chunk30, %iq2s_word1_low_30, %selected_word1_low29 : vector<32xf32> + %selected_word1_low31 = scf.select %is_chunk31, %iq2s_word1_low_31, %selected_word1_low30 : vector<32xf32> + %iq2s_word1_high_0_0 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_1 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_2 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_3 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_4 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_5 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_6 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_7 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_8 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_9 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_10 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_11 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_12 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_13 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_14 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_15 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_16 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_17 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_18 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_19 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_20 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_21 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_22 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_23 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_24 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_25 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_26 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_27 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_28 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_29 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_30 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0_31 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_0 = vector.from_elements %iq2s_word1_high_0_0, %iq2s_word1_high_0_1, %iq2s_word1_high_0_2, %iq2s_word1_high_0_3, %iq2s_word1_high_0_4, %iq2s_word1_high_0_5, %iq2s_word1_high_0_6, %iq2s_word1_high_0_7, %iq2s_word1_high_0_8, %iq2s_word1_high_0_9, %iq2s_word1_high_0_10, %iq2s_word1_high_0_11, %iq2s_word1_high_0_12, %iq2s_word1_high_0_13, %iq2s_word1_high_0_14, %iq2s_word1_high_0_15, %iq2s_word1_high_0_16, %iq2s_word1_high_0_17, %iq2s_word1_high_0_18, %iq2s_word1_high_0_19, %iq2s_word1_high_0_20, %iq2s_word1_high_0_21, %iq2s_word1_high_0_22, %iq2s_word1_high_0_23, %iq2s_word1_high_0_24, %iq2s_word1_high_0_25, %iq2s_word1_high_0_26, %iq2s_word1_high_0_27, %iq2s_word1_high_0_28, %iq2s_word1_high_0_29, %iq2s_word1_high_0_30, %iq2s_word1_high_0_31 : vector<32xf32> + %iq2s_word1_high_1_0 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_1 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_2 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_3 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_4 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_5 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_6 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_7 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_8 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_9 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_10 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_11 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_12 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_13 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_14 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_15 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_16 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_17 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_18 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_19 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_20 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_21 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_22 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_23 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_24 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_25 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_26 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_27 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_28 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_29 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_30 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1_31 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_1 = vector.from_elements %iq2s_word1_high_1_0, %iq2s_word1_high_1_1, %iq2s_word1_high_1_2, %iq2s_word1_high_1_3, %iq2s_word1_high_1_4, %iq2s_word1_high_1_5, %iq2s_word1_high_1_6, %iq2s_word1_high_1_7, %iq2s_word1_high_1_8, %iq2s_word1_high_1_9, %iq2s_word1_high_1_10, %iq2s_word1_high_1_11, %iq2s_word1_high_1_12, %iq2s_word1_high_1_13, %iq2s_word1_high_1_14, %iq2s_word1_high_1_15, %iq2s_word1_high_1_16, %iq2s_word1_high_1_17, %iq2s_word1_high_1_18, %iq2s_word1_high_1_19, %iq2s_word1_high_1_20, %iq2s_word1_high_1_21, %iq2s_word1_high_1_22, %iq2s_word1_high_1_23, %iq2s_word1_high_1_24, %iq2s_word1_high_1_25, %iq2s_word1_high_1_26, %iq2s_word1_high_1_27, %iq2s_word1_high_1_28, %iq2s_word1_high_1_29, %iq2s_word1_high_1_30, %iq2s_word1_high_1_31 : vector<32xf32> + %iq2s_word1_high_2_0 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_1 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_2 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_3 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_4 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_5 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_6 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_7 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_8 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_9 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_10 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_11 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_12 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_13 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_14 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_15 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_16 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_17 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_18 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_19 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_20 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_21 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_22 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_23 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_24 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_25 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_26 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_27 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_28 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_29 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_30 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2_31 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_2 = vector.from_elements %iq2s_word1_high_2_0, %iq2s_word1_high_2_1, %iq2s_word1_high_2_2, %iq2s_word1_high_2_3, %iq2s_word1_high_2_4, %iq2s_word1_high_2_5, %iq2s_word1_high_2_6, %iq2s_word1_high_2_7, %iq2s_word1_high_2_8, %iq2s_word1_high_2_9, %iq2s_word1_high_2_10, %iq2s_word1_high_2_11, %iq2s_word1_high_2_12, %iq2s_word1_high_2_13, %iq2s_word1_high_2_14, %iq2s_word1_high_2_15, %iq2s_word1_high_2_16, %iq2s_word1_high_2_17, %iq2s_word1_high_2_18, %iq2s_word1_high_2_19, %iq2s_word1_high_2_20, %iq2s_word1_high_2_21, %iq2s_word1_high_2_22, %iq2s_word1_high_2_23, %iq2s_word1_high_2_24, %iq2s_word1_high_2_25, %iq2s_word1_high_2_26, %iq2s_word1_high_2_27, %iq2s_word1_high_2_28, %iq2s_word1_high_2_29, %iq2s_word1_high_2_30, %iq2s_word1_high_2_31 : vector<32xf32> + %iq2s_word1_high_3_0 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_1 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_2 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_3 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_4 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_5 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_6 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_7 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_8 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_9 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_10 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_11 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_12 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_13 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_14 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_15 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_16 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_17 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_18 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_19 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_20 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_21 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_22 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_23 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_24 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_25 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_26 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_27 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_28 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_29 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_30 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3_31 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_3 = vector.from_elements %iq2s_word1_high_3_0, %iq2s_word1_high_3_1, %iq2s_word1_high_3_2, %iq2s_word1_high_3_3, %iq2s_word1_high_3_4, %iq2s_word1_high_3_5, %iq2s_word1_high_3_6, %iq2s_word1_high_3_7, %iq2s_word1_high_3_8, %iq2s_word1_high_3_9, %iq2s_word1_high_3_10, %iq2s_word1_high_3_11, %iq2s_word1_high_3_12, %iq2s_word1_high_3_13, %iq2s_word1_high_3_14, %iq2s_word1_high_3_15, %iq2s_word1_high_3_16, %iq2s_word1_high_3_17, %iq2s_word1_high_3_18, %iq2s_word1_high_3_19, %iq2s_word1_high_3_20, %iq2s_word1_high_3_21, %iq2s_word1_high_3_22, %iq2s_word1_high_3_23, %iq2s_word1_high_3_24, %iq2s_word1_high_3_25, %iq2s_word1_high_3_26, %iq2s_word1_high_3_27, %iq2s_word1_high_3_28, %iq2s_word1_high_3_29, %iq2s_word1_high_3_30, %iq2s_word1_high_3_31 : vector<32xf32> + %iq2s_word1_high_4_0 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_1 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_2 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_3 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_4 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_5 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_6 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_7 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_8 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_9 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_10 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_11 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_12 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_13 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_14 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_15 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_16 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_17 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_18 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_19 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_20 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_21 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_22 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_23 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_24 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_25 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_26 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_27 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_28 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_29 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_30 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4_31 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_4 = vector.from_elements %iq2s_word1_high_4_0, %iq2s_word1_high_4_1, %iq2s_word1_high_4_2, %iq2s_word1_high_4_3, %iq2s_word1_high_4_4, %iq2s_word1_high_4_5, %iq2s_word1_high_4_6, %iq2s_word1_high_4_7, %iq2s_word1_high_4_8, %iq2s_word1_high_4_9, %iq2s_word1_high_4_10, %iq2s_word1_high_4_11, %iq2s_word1_high_4_12, %iq2s_word1_high_4_13, %iq2s_word1_high_4_14, %iq2s_word1_high_4_15, %iq2s_word1_high_4_16, %iq2s_word1_high_4_17, %iq2s_word1_high_4_18, %iq2s_word1_high_4_19, %iq2s_word1_high_4_20, %iq2s_word1_high_4_21, %iq2s_word1_high_4_22, %iq2s_word1_high_4_23, %iq2s_word1_high_4_24, %iq2s_word1_high_4_25, %iq2s_word1_high_4_26, %iq2s_word1_high_4_27, %iq2s_word1_high_4_28, %iq2s_word1_high_4_29, %iq2s_word1_high_4_30, %iq2s_word1_high_4_31 : vector<32xf32> + %iq2s_word1_high_5_0 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_1 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_2 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_3 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_4 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_5 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_6 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_7 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_8 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_9 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_10 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_11 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_12 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_13 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_14 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_15 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_16 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_17 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_18 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_19 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_20 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_21 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_22 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_23 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_24 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_25 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_26 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_27 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_28 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_29 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_30 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5_31 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_5 = vector.from_elements %iq2s_word1_high_5_0, %iq2s_word1_high_5_1, %iq2s_word1_high_5_2, %iq2s_word1_high_5_3, %iq2s_word1_high_5_4, %iq2s_word1_high_5_5, %iq2s_word1_high_5_6, %iq2s_word1_high_5_7, %iq2s_word1_high_5_8, %iq2s_word1_high_5_9, %iq2s_word1_high_5_10, %iq2s_word1_high_5_11, %iq2s_word1_high_5_12, %iq2s_word1_high_5_13, %iq2s_word1_high_5_14, %iq2s_word1_high_5_15, %iq2s_word1_high_5_16, %iq2s_word1_high_5_17, %iq2s_word1_high_5_18, %iq2s_word1_high_5_19, %iq2s_word1_high_5_20, %iq2s_word1_high_5_21, %iq2s_word1_high_5_22, %iq2s_word1_high_5_23, %iq2s_word1_high_5_24, %iq2s_word1_high_5_25, %iq2s_word1_high_5_26, %iq2s_word1_high_5_27, %iq2s_word1_high_5_28, %iq2s_word1_high_5_29, %iq2s_word1_high_5_30, %iq2s_word1_high_5_31 : vector<32xf32> + %iq2s_word1_high_6_0 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_6_1 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_6_2 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_6_3 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_6_4 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_6_5 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_6_6 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_6_7 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_6_8 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_6_9 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_6_10 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_6_11 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_6_12 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_6_13 = scalar.constant 2056.0 : f32 + %iq2s_word1_high_6_14 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6_15 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6_16 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6_17 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6_18 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6_19 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6_20 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6_21 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6_22 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6_23 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6_24 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6_25 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6_26 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6_27 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6_28 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6_29 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6_30 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6_31 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_6 = vector.from_elements %iq2s_word1_high_6_0, %iq2s_word1_high_6_1, %iq2s_word1_high_6_2, %iq2s_word1_high_6_3, %iq2s_word1_high_6_4, %iq2s_word1_high_6_5, %iq2s_word1_high_6_6, %iq2s_word1_high_6_7, %iq2s_word1_high_6_8, %iq2s_word1_high_6_9, %iq2s_word1_high_6_10, %iq2s_word1_high_6_11, %iq2s_word1_high_6_12, %iq2s_word1_high_6_13, %iq2s_word1_high_6_14, %iq2s_word1_high_6_15, %iq2s_word1_high_6_16, %iq2s_word1_high_6_17, %iq2s_word1_high_6_18, %iq2s_word1_high_6_19, %iq2s_word1_high_6_20, %iq2s_word1_high_6_21, %iq2s_word1_high_6_22, %iq2s_word1_high_6_23, %iq2s_word1_high_6_24, %iq2s_word1_high_6_25, %iq2s_word1_high_6_26, %iq2s_word1_high_6_27, %iq2s_word1_high_6_28, %iq2s_word1_high_6_29, %iq2s_word1_high_6_30, %iq2s_word1_high_6_31 : vector<32xf32> + %iq2s_word1_high_7_0 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_1 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_2 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_3 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_4 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_5 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_6 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_7 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_8 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_9 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_10 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_11 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_12 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_13 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_14 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_15 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_16 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_17 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_18 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_19 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_20 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_21 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_22 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_23 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_24 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_25 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_26 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_27 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_28 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_29 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_30 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7_31 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_7 = vector.from_elements %iq2s_word1_high_7_0, %iq2s_word1_high_7_1, %iq2s_word1_high_7_2, %iq2s_word1_high_7_3, %iq2s_word1_high_7_4, %iq2s_word1_high_7_5, %iq2s_word1_high_7_6, %iq2s_word1_high_7_7, %iq2s_word1_high_7_8, %iq2s_word1_high_7_9, %iq2s_word1_high_7_10, %iq2s_word1_high_7_11, %iq2s_word1_high_7_12, %iq2s_word1_high_7_13, %iq2s_word1_high_7_14, %iq2s_word1_high_7_15, %iq2s_word1_high_7_16, %iq2s_word1_high_7_17, %iq2s_word1_high_7_18, %iq2s_word1_high_7_19, %iq2s_word1_high_7_20, %iq2s_word1_high_7_21, %iq2s_word1_high_7_22, %iq2s_word1_high_7_23, %iq2s_word1_high_7_24, %iq2s_word1_high_7_25, %iq2s_word1_high_7_26, %iq2s_word1_high_7_27, %iq2s_word1_high_7_28, %iq2s_word1_high_7_29, %iq2s_word1_high_7_30, %iq2s_word1_high_7_31 : vector<32xf32> + %iq2s_word1_high_8_0 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_1 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_2 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_3 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_4 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_5 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_6 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_7 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_8 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_9 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_10 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_11 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_12 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_13 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_14 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_15 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_16 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_17 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_18 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_19 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_20 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_21 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_22 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_23 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_24 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_25 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_26 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_27 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_28 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_29 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_30 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8_31 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_8 = vector.from_elements %iq2s_word1_high_8_0, %iq2s_word1_high_8_1, %iq2s_word1_high_8_2, %iq2s_word1_high_8_3, %iq2s_word1_high_8_4, %iq2s_word1_high_8_5, %iq2s_word1_high_8_6, %iq2s_word1_high_8_7, %iq2s_word1_high_8_8, %iq2s_word1_high_8_9, %iq2s_word1_high_8_10, %iq2s_word1_high_8_11, %iq2s_word1_high_8_12, %iq2s_word1_high_8_13, %iq2s_word1_high_8_14, %iq2s_word1_high_8_15, %iq2s_word1_high_8_16, %iq2s_word1_high_8_17, %iq2s_word1_high_8_18, %iq2s_word1_high_8_19, %iq2s_word1_high_8_20, %iq2s_word1_high_8_21, %iq2s_word1_high_8_22, %iq2s_word1_high_8_23, %iq2s_word1_high_8_24, %iq2s_word1_high_8_25, %iq2s_word1_high_8_26, %iq2s_word1_high_8_27, %iq2s_word1_high_8_28, %iq2s_word1_high_8_29, %iq2s_word1_high_8_30, %iq2s_word1_high_8_31 : vector<32xf32> + %iq2s_word1_high_9_0 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_1 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_2 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_3 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_4 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_5 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_6 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_7 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_8 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_9 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_10 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_11 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_12 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_13 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_14 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_15 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_16 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_17 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_18 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_19 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_20 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_21 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_22 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_23 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_24 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_25 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_26 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_27 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_28 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_29 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_30 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9_31 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_9 = vector.from_elements %iq2s_word1_high_9_0, %iq2s_word1_high_9_1, %iq2s_word1_high_9_2, %iq2s_word1_high_9_3, %iq2s_word1_high_9_4, %iq2s_word1_high_9_5, %iq2s_word1_high_9_6, %iq2s_word1_high_9_7, %iq2s_word1_high_9_8, %iq2s_word1_high_9_9, %iq2s_word1_high_9_10, %iq2s_word1_high_9_11, %iq2s_word1_high_9_12, %iq2s_word1_high_9_13, %iq2s_word1_high_9_14, %iq2s_word1_high_9_15, %iq2s_word1_high_9_16, %iq2s_word1_high_9_17, %iq2s_word1_high_9_18, %iq2s_word1_high_9_19, %iq2s_word1_high_9_20, %iq2s_word1_high_9_21, %iq2s_word1_high_9_22, %iq2s_word1_high_9_23, %iq2s_word1_high_9_24, %iq2s_word1_high_9_25, %iq2s_word1_high_9_26, %iq2s_word1_high_9_27, %iq2s_word1_high_9_28, %iq2s_word1_high_9_29, %iq2s_word1_high_9_30, %iq2s_word1_high_9_31 : vector<32xf32> + %iq2s_word1_high_10_0 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_1 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_2 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_3 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_4 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_5 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_6 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_7 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_8 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_9 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_10 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_11 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_12 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_13 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_14 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_15 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_16 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_17 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_18 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_19 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_20 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_21 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_22 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_23 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_24 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_25 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_26 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_27 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_28 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_29 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_30 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10_31 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_10 = vector.from_elements %iq2s_word1_high_10_0, %iq2s_word1_high_10_1, %iq2s_word1_high_10_2, %iq2s_word1_high_10_3, %iq2s_word1_high_10_4, %iq2s_word1_high_10_5, %iq2s_word1_high_10_6, %iq2s_word1_high_10_7, %iq2s_word1_high_10_8, %iq2s_word1_high_10_9, %iq2s_word1_high_10_10, %iq2s_word1_high_10_11, %iq2s_word1_high_10_12, %iq2s_word1_high_10_13, %iq2s_word1_high_10_14, %iq2s_word1_high_10_15, %iq2s_word1_high_10_16, %iq2s_word1_high_10_17, %iq2s_word1_high_10_18, %iq2s_word1_high_10_19, %iq2s_word1_high_10_20, %iq2s_word1_high_10_21, %iq2s_word1_high_10_22, %iq2s_word1_high_10_23, %iq2s_word1_high_10_24, %iq2s_word1_high_10_25, %iq2s_word1_high_10_26, %iq2s_word1_high_10_27, %iq2s_word1_high_10_28, %iq2s_word1_high_10_29, %iq2s_word1_high_10_30, %iq2s_word1_high_10_31 : vector<32xf32> + %iq2s_word1_high_11_0 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_1 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_2 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_3 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_4 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_5 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_6 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_7 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_8 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_9 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_10 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_11 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_12 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_13 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_14 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_15 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_16 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_17 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_18 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_19 = scalar.constant 2073.0 : f32 + %iq2s_word1_high_11_20 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_11_21 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_11_22 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_11_23 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_11_24 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_11_25 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_11_26 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_11_27 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_11_28 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_11_29 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_11_30 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_11_31 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_11 = vector.from_elements %iq2s_word1_high_11_0, %iq2s_word1_high_11_1, %iq2s_word1_high_11_2, %iq2s_word1_high_11_3, %iq2s_word1_high_11_4, %iq2s_word1_high_11_5, %iq2s_word1_high_11_6, %iq2s_word1_high_11_7, %iq2s_word1_high_11_8, %iq2s_word1_high_11_9, %iq2s_word1_high_11_10, %iq2s_word1_high_11_11, %iq2s_word1_high_11_12, %iq2s_word1_high_11_13, %iq2s_word1_high_11_14, %iq2s_word1_high_11_15, %iq2s_word1_high_11_16, %iq2s_word1_high_11_17, %iq2s_word1_high_11_18, %iq2s_word1_high_11_19, %iq2s_word1_high_11_20, %iq2s_word1_high_11_21, %iq2s_word1_high_11_22, %iq2s_word1_high_11_23, %iq2s_word1_high_11_24, %iq2s_word1_high_11_25, %iq2s_word1_high_11_26, %iq2s_word1_high_11_27, %iq2s_word1_high_11_28, %iq2s_word1_high_11_29, %iq2s_word1_high_11_30, %iq2s_word1_high_11_31 : vector<32xf32> + %iq2s_word1_high_12_0 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_1 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_2 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_3 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_4 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_5 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_6 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_7 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_8 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_9 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_10 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_11 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_12 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_13 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_14 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_15 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_16 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_17 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_18 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_19 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_20 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_21 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_22 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_23 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_24 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_25 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_26 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_27 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_28 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_29 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_30 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12_31 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_12 = vector.from_elements %iq2s_word1_high_12_0, %iq2s_word1_high_12_1, %iq2s_word1_high_12_2, %iq2s_word1_high_12_3, %iq2s_word1_high_12_4, %iq2s_word1_high_12_5, %iq2s_word1_high_12_6, %iq2s_word1_high_12_7, %iq2s_word1_high_12_8, %iq2s_word1_high_12_9, %iq2s_word1_high_12_10, %iq2s_word1_high_12_11, %iq2s_word1_high_12_12, %iq2s_word1_high_12_13, %iq2s_word1_high_12_14, %iq2s_word1_high_12_15, %iq2s_word1_high_12_16, %iq2s_word1_high_12_17, %iq2s_word1_high_12_18, %iq2s_word1_high_12_19, %iq2s_word1_high_12_20, %iq2s_word1_high_12_21, %iq2s_word1_high_12_22, %iq2s_word1_high_12_23, %iq2s_word1_high_12_24, %iq2s_word1_high_12_25, %iq2s_word1_high_12_26, %iq2s_word1_high_12_27, %iq2s_word1_high_12_28, %iq2s_word1_high_12_29, %iq2s_word1_high_12_30, %iq2s_word1_high_12_31 : vector<32xf32> + %iq2s_word1_high_13_0 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_1 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_2 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_3 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_4 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_5 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_6 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_7 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_8 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_9 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_10 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_11 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_12 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_13 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_14 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_15 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_16 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_17 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_18 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_19 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_20 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_21 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_22 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_23 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_24 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_25 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_26 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_27 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_28 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_29 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_30 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13_31 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_13 = vector.from_elements %iq2s_word1_high_13_0, %iq2s_word1_high_13_1, %iq2s_word1_high_13_2, %iq2s_word1_high_13_3, %iq2s_word1_high_13_4, %iq2s_word1_high_13_5, %iq2s_word1_high_13_6, %iq2s_word1_high_13_7, %iq2s_word1_high_13_8, %iq2s_word1_high_13_9, %iq2s_word1_high_13_10, %iq2s_word1_high_13_11, %iq2s_word1_high_13_12, %iq2s_word1_high_13_13, %iq2s_word1_high_13_14, %iq2s_word1_high_13_15, %iq2s_word1_high_13_16, %iq2s_word1_high_13_17, %iq2s_word1_high_13_18, %iq2s_word1_high_13_19, %iq2s_word1_high_13_20, %iq2s_word1_high_13_21, %iq2s_word1_high_13_22, %iq2s_word1_high_13_23, %iq2s_word1_high_13_24, %iq2s_word1_high_13_25, %iq2s_word1_high_13_26, %iq2s_word1_high_13_27, %iq2s_word1_high_13_28, %iq2s_word1_high_13_29, %iq2s_word1_high_13_30, %iq2s_word1_high_13_31 : vector<32xf32> + %iq2s_word1_high_14_0 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_1 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_2 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_3 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_4 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_5 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_6 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_7 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_8 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_9 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_10 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_11 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_12 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_13 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_14 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_15 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_16 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_17 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_18 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_19 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_20 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_21 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_22 = scalar.constant 2091.0 : f32 + %iq2s_word1_high_14_23 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_14_24 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_14_25 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_14_26 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_14_27 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_14_28 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_14_29 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_14_30 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_14_31 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_14 = vector.from_elements %iq2s_word1_high_14_0, %iq2s_word1_high_14_1, %iq2s_word1_high_14_2, %iq2s_word1_high_14_3, %iq2s_word1_high_14_4, %iq2s_word1_high_14_5, %iq2s_word1_high_14_6, %iq2s_word1_high_14_7, %iq2s_word1_high_14_8, %iq2s_word1_high_14_9, %iq2s_word1_high_14_10, %iq2s_word1_high_14_11, %iq2s_word1_high_14_12, %iq2s_word1_high_14_13, %iq2s_word1_high_14_14, %iq2s_word1_high_14_15, %iq2s_word1_high_14_16, %iq2s_word1_high_14_17, %iq2s_word1_high_14_18, %iq2s_word1_high_14_19, %iq2s_word1_high_14_20, %iq2s_word1_high_14_21, %iq2s_word1_high_14_22, %iq2s_word1_high_14_23, %iq2s_word1_high_14_24, %iq2s_word1_high_14_25, %iq2s_word1_high_14_26, %iq2s_word1_high_14_27, %iq2s_word1_high_14_28, %iq2s_word1_high_14_29, %iq2s_word1_high_14_30, %iq2s_word1_high_14_31 : vector<32xf32> + %iq2s_word1_high_15_0 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_1 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_2 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_3 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_4 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_5 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_6 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_7 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_8 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_9 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_10 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_11 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_12 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_13 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_14 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_15 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_16 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_17 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_18 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_19 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_20 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_21 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_22 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_23 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_24 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_25 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_26 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_27 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_28 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_29 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_30 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15_31 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_15 = vector.from_elements %iq2s_word1_high_15_0, %iq2s_word1_high_15_1, %iq2s_word1_high_15_2, %iq2s_word1_high_15_3, %iq2s_word1_high_15_4, %iq2s_word1_high_15_5, %iq2s_word1_high_15_6, %iq2s_word1_high_15_7, %iq2s_word1_high_15_8, %iq2s_word1_high_15_9, %iq2s_word1_high_15_10, %iq2s_word1_high_15_11, %iq2s_word1_high_15_12, %iq2s_word1_high_15_13, %iq2s_word1_high_15_14, %iq2s_word1_high_15_15, %iq2s_word1_high_15_16, %iq2s_word1_high_15_17, %iq2s_word1_high_15_18, %iq2s_word1_high_15_19, %iq2s_word1_high_15_20, %iq2s_word1_high_15_21, %iq2s_word1_high_15_22, %iq2s_word1_high_15_23, %iq2s_word1_high_15_24, %iq2s_word1_high_15_25, %iq2s_word1_high_15_26, %iq2s_word1_high_15_27, %iq2s_word1_high_15_28, %iq2s_word1_high_15_29, %iq2s_word1_high_15_30, %iq2s_word1_high_15_31 : vector<32xf32> + %iq2s_word1_high_16_0 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_1 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_2 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_3 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_4 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_5 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_6 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_7 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_8 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_9 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_10 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_11 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_12 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_13 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_14 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_15 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_16 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_17 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_18 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_19 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_20 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_21 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_22 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_23 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_24 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_25 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_26 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_27 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_28 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_29 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_30 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16_31 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_16 = vector.from_elements %iq2s_word1_high_16_0, %iq2s_word1_high_16_1, %iq2s_word1_high_16_2, %iq2s_word1_high_16_3, %iq2s_word1_high_16_4, %iq2s_word1_high_16_5, %iq2s_word1_high_16_6, %iq2s_word1_high_16_7, %iq2s_word1_high_16_8, %iq2s_word1_high_16_9, %iq2s_word1_high_16_10, %iq2s_word1_high_16_11, %iq2s_word1_high_16_12, %iq2s_word1_high_16_13, %iq2s_word1_high_16_14, %iq2s_word1_high_16_15, %iq2s_word1_high_16_16, %iq2s_word1_high_16_17, %iq2s_word1_high_16_18, %iq2s_word1_high_16_19, %iq2s_word1_high_16_20, %iq2s_word1_high_16_21, %iq2s_word1_high_16_22, %iq2s_word1_high_16_23, %iq2s_word1_high_16_24, %iq2s_word1_high_16_25, %iq2s_word1_high_16_26, %iq2s_word1_high_16_27, %iq2s_word1_high_16_28, %iq2s_word1_high_16_29, %iq2s_word1_high_16_30, %iq2s_word1_high_16_31 : vector<32xf32> + %iq2s_word1_high_17_0 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_1 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_2 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_3 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_4 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_5 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_6 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_7 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_8 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_9 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_10 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_11 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_12 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_13 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_14 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_15 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_16 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_17 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_18 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_19 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_20 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_21 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_22 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_23 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_24 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_25 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_26 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_27 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_28 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_29 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_30 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17_31 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_17 = vector.from_elements %iq2s_word1_high_17_0, %iq2s_word1_high_17_1, %iq2s_word1_high_17_2, %iq2s_word1_high_17_3, %iq2s_word1_high_17_4, %iq2s_word1_high_17_5, %iq2s_word1_high_17_6, %iq2s_word1_high_17_7, %iq2s_word1_high_17_8, %iq2s_word1_high_17_9, %iq2s_word1_high_17_10, %iq2s_word1_high_17_11, %iq2s_word1_high_17_12, %iq2s_word1_high_17_13, %iq2s_word1_high_17_14, %iq2s_word1_high_17_15, %iq2s_word1_high_17_16, %iq2s_word1_high_17_17, %iq2s_word1_high_17_18, %iq2s_word1_high_17_19, %iq2s_word1_high_17_20, %iq2s_word1_high_17_21, %iq2s_word1_high_17_22, %iq2s_word1_high_17_23, %iq2s_word1_high_17_24, %iq2s_word1_high_17_25, %iq2s_word1_high_17_26, %iq2s_word1_high_17_27, %iq2s_word1_high_17_28, %iq2s_word1_high_17_29, %iq2s_word1_high_17_30, %iq2s_word1_high_17_31 : vector<32xf32> + %iq2s_word1_high_18_0 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_1 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_2 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_3 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_4 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_5 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_6 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_7 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_8 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_9 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_10 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_11 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_12 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_13 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_14 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_15 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_16 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_17 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_18 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_19 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_20 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_21 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_22 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_23 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_24 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_25 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_26 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_27 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_28 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_29 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_30 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18_31 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_18 = vector.from_elements %iq2s_word1_high_18_0, %iq2s_word1_high_18_1, %iq2s_word1_high_18_2, %iq2s_word1_high_18_3, %iq2s_word1_high_18_4, %iq2s_word1_high_18_5, %iq2s_word1_high_18_6, %iq2s_word1_high_18_7, %iq2s_word1_high_18_8, %iq2s_word1_high_18_9, %iq2s_word1_high_18_10, %iq2s_word1_high_18_11, %iq2s_word1_high_18_12, %iq2s_word1_high_18_13, %iq2s_word1_high_18_14, %iq2s_word1_high_18_15, %iq2s_word1_high_18_16, %iq2s_word1_high_18_17, %iq2s_word1_high_18_18, %iq2s_word1_high_18_19, %iq2s_word1_high_18_20, %iq2s_word1_high_18_21, %iq2s_word1_high_18_22, %iq2s_word1_high_18_23, %iq2s_word1_high_18_24, %iq2s_word1_high_18_25, %iq2s_word1_high_18_26, %iq2s_word1_high_18_27, %iq2s_word1_high_18_28, %iq2s_word1_high_18_29, %iq2s_word1_high_18_30, %iq2s_word1_high_18_31 : vector<32xf32> + %iq2s_word1_high_19_0 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_1 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_2 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_3 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_4 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_5 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_6 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_7 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_8 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_9 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_10 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_11 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_12 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_13 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_14 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_15 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_16 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_17 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_18 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_19 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_20 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_21 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_22 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_23 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_24 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_25 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_26 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_27 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_28 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_29 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_30 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19_31 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_19 = vector.from_elements %iq2s_word1_high_19_0, %iq2s_word1_high_19_1, %iq2s_word1_high_19_2, %iq2s_word1_high_19_3, %iq2s_word1_high_19_4, %iq2s_word1_high_19_5, %iq2s_word1_high_19_6, %iq2s_word1_high_19_7, %iq2s_word1_high_19_8, %iq2s_word1_high_19_9, %iq2s_word1_high_19_10, %iq2s_word1_high_19_11, %iq2s_word1_high_19_12, %iq2s_word1_high_19_13, %iq2s_word1_high_19_14, %iq2s_word1_high_19_15, %iq2s_word1_high_19_16, %iq2s_word1_high_19_17, %iq2s_word1_high_19_18, %iq2s_word1_high_19_19, %iq2s_word1_high_19_20, %iq2s_word1_high_19_21, %iq2s_word1_high_19_22, %iq2s_word1_high_19_23, %iq2s_word1_high_19_24, %iq2s_word1_high_19_25, %iq2s_word1_high_19_26, %iq2s_word1_high_19_27, %iq2s_word1_high_19_28, %iq2s_word1_high_19_29, %iq2s_word1_high_19_30, %iq2s_word1_high_19_31 : vector<32xf32> + %iq2s_word1_high_20_0 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_20_1 = scalar.constant 6408.0 : f32 + %iq2s_word1_high_20_2 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_3 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_4 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_5 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_6 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_7 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_8 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_9 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_10 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_11 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_12 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_13 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_14 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_15 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_16 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_17 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_18 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_19 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_20 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_21 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_22 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_23 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_24 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_25 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_26 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_27 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_28 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_29 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_30 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20_31 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_20 = vector.from_elements %iq2s_word1_high_20_0, %iq2s_word1_high_20_1, %iq2s_word1_high_20_2, %iq2s_word1_high_20_3, %iq2s_word1_high_20_4, %iq2s_word1_high_20_5, %iq2s_word1_high_20_6, %iq2s_word1_high_20_7, %iq2s_word1_high_20_8, %iq2s_word1_high_20_9, %iq2s_word1_high_20_10, %iq2s_word1_high_20_11, %iq2s_word1_high_20_12, %iq2s_word1_high_20_13, %iq2s_word1_high_20_14, %iq2s_word1_high_20_15, %iq2s_word1_high_20_16, %iq2s_word1_high_20_17, %iq2s_word1_high_20_18, %iq2s_word1_high_20_19, %iq2s_word1_high_20_20, %iq2s_word1_high_20_21, %iq2s_word1_high_20_22, %iq2s_word1_high_20_23, %iq2s_word1_high_20_24, %iq2s_word1_high_20_25, %iq2s_word1_high_20_26, %iq2s_word1_high_20_27, %iq2s_word1_high_20_28, %iq2s_word1_high_20_29, %iq2s_word1_high_20_30, %iq2s_word1_high_20_31 : vector<32xf32> + %iq2s_word1_high_21_0 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_1 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_2 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_3 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_4 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_5 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_6 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_7 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_8 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_9 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_10 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_11 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_12 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_13 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_14 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_15 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_16 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_17 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_18 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_19 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_20 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_21 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_22 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_23 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_24 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_25 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_26 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_27 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_28 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_29 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_30 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21_31 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_21 = vector.from_elements %iq2s_word1_high_21_0, %iq2s_word1_high_21_1, %iq2s_word1_high_21_2, %iq2s_word1_high_21_3, %iq2s_word1_high_21_4, %iq2s_word1_high_21_5, %iq2s_word1_high_21_6, %iq2s_word1_high_21_7, %iq2s_word1_high_21_8, %iq2s_word1_high_21_9, %iq2s_word1_high_21_10, %iq2s_word1_high_21_11, %iq2s_word1_high_21_12, %iq2s_word1_high_21_13, %iq2s_word1_high_21_14, %iq2s_word1_high_21_15, %iq2s_word1_high_21_16, %iq2s_word1_high_21_17, %iq2s_word1_high_21_18, %iq2s_word1_high_21_19, %iq2s_word1_high_21_20, %iq2s_word1_high_21_21, %iq2s_word1_high_21_22, %iq2s_word1_high_21_23, %iq2s_word1_high_21_24, %iq2s_word1_high_21_25, %iq2s_word1_high_21_26, %iq2s_word1_high_21_27, %iq2s_word1_high_21_28, %iq2s_word1_high_21_29, %iq2s_word1_high_21_30, %iq2s_word1_high_21_31 : vector<32xf32> + %iq2s_word1_high_22_0 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_1 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_2 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_3 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_4 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_5 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_6 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_7 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_8 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_9 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_10 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_11 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_12 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_13 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_14 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_15 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_16 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_17 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_18 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_19 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_20 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_21 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_22 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_23 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_24 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_25 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_26 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_27 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_28 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_29 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_30 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22_31 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_22 = vector.from_elements %iq2s_word1_high_22_0, %iq2s_word1_high_22_1, %iq2s_word1_high_22_2, %iq2s_word1_high_22_3, %iq2s_word1_high_22_4, %iq2s_word1_high_22_5, %iq2s_word1_high_22_6, %iq2s_word1_high_22_7, %iq2s_word1_high_22_8, %iq2s_word1_high_22_9, %iq2s_word1_high_22_10, %iq2s_word1_high_22_11, %iq2s_word1_high_22_12, %iq2s_word1_high_22_13, %iq2s_word1_high_22_14, %iq2s_word1_high_22_15, %iq2s_word1_high_22_16, %iq2s_word1_high_22_17, %iq2s_word1_high_22_18, %iq2s_word1_high_22_19, %iq2s_word1_high_22_20, %iq2s_word1_high_22_21, %iq2s_word1_high_22_22, %iq2s_word1_high_22_23, %iq2s_word1_high_22_24, %iq2s_word1_high_22_25, %iq2s_word1_high_22_26, %iq2s_word1_high_22_27, %iq2s_word1_high_22_28, %iq2s_word1_high_22_29, %iq2s_word1_high_22_30, %iq2s_word1_high_22_31 : vector<32xf32> + %iq2s_word1_high_23_0 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_1 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_2 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_3 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_4 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_5 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_6 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_7 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_8 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_9 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_10 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_11 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_12 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_13 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_14 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_15 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_16 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_17 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_18 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_19 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_20 = scalar.constant 6425.0 : f32 + %iq2s_word1_high_23_21 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_23_22 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_23_23 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_23_24 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_23_25 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_23_26 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_23_27 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_23_28 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_23_29 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_23_30 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_23_31 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_23 = vector.from_elements %iq2s_word1_high_23_0, %iq2s_word1_high_23_1, %iq2s_word1_high_23_2, %iq2s_word1_high_23_3, %iq2s_word1_high_23_4, %iq2s_word1_high_23_5, %iq2s_word1_high_23_6, %iq2s_word1_high_23_7, %iq2s_word1_high_23_8, %iq2s_word1_high_23_9, %iq2s_word1_high_23_10, %iq2s_word1_high_23_11, %iq2s_word1_high_23_12, %iq2s_word1_high_23_13, %iq2s_word1_high_23_14, %iq2s_word1_high_23_15, %iq2s_word1_high_23_16, %iq2s_word1_high_23_17, %iq2s_word1_high_23_18, %iq2s_word1_high_23_19, %iq2s_word1_high_23_20, %iq2s_word1_high_23_21, %iq2s_word1_high_23_22, %iq2s_word1_high_23_23, %iq2s_word1_high_23_24, %iq2s_word1_high_23_25, %iq2s_word1_high_23_26, %iq2s_word1_high_23_27, %iq2s_word1_high_23_28, %iq2s_word1_high_23_29, %iq2s_word1_high_23_30, %iq2s_word1_high_23_31 : vector<32xf32> + %iq2s_word1_high_24_0 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_1 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_2 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_3 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_4 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_5 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_6 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_7 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_8 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_9 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_10 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_11 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_12 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_13 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_14 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_15 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_16 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_17 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_18 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_19 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_20 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_21 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_22 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_23 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_24 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_25 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_26 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_27 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_28 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_29 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_30 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24_31 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_24 = vector.from_elements %iq2s_word1_high_24_0, %iq2s_word1_high_24_1, %iq2s_word1_high_24_2, %iq2s_word1_high_24_3, %iq2s_word1_high_24_4, %iq2s_word1_high_24_5, %iq2s_word1_high_24_6, %iq2s_word1_high_24_7, %iq2s_word1_high_24_8, %iq2s_word1_high_24_9, %iq2s_word1_high_24_10, %iq2s_word1_high_24_11, %iq2s_word1_high_24_12, %iq2s_word1_high_24_13, %iq2s_word1_high_24_14, %iq2s_word1_high_24_15, %iq2s_word1_high_24_16, %iq2s_word1_high_24_17, %iq2s_word1_high_24_18, %iq2s_word1_high_24_19, %iq2s_word1_high_24_20, %iq2s_word1_high_24_21, %iq2s_word1_high_24_22, %iq2s_word1_high_24_23, %iq2s_word1_high_24_24, %iq2s_word1_high_24_25, %iq2s_word1_high_24_26, %iq2s_word1_high_24_27, %iq2s_word1_high_24_28, %iq2s_word1_high_24_29, %iq2s_word1_high_24_30, %iq2s_word1_high_24_31 : vector<32xf32> + %iq2s_word1_high_25_0 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_1 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_2 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_3 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_4 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_5 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_6 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_7 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_8 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_9 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_10 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_11 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_12 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_13 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_14 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_15 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_16 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_17 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_18 = scalar.constant 6443.0 : f32 + %iq2s_word1_high_25_19 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_25_20 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_25_21 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_25_22 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_25_23 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_25_24 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_25_25 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_25_26 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_25_27 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_25_28 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_25_29 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_25_30 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_25_31 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_25 = vector.from_elements %iq2s_word1_high_25_0, %iq2s_word1_high_25_1, %iq2s_word1_high_25_2, %iq2s_word1_high_25_3, %iq2s_word1_high_25_4, %iq2s_word1_high_25_5, %iq2s_word1_high_25_6, %iq2s_word1_high_25_7, %iq2s_word1_high_25_8, %iq2s_word1_high_25_9, %iq2s_word1_high_25_10, %iq2s_word1_high_25_11, %iq2s_word1_high_25_12, %iq2s_word1_high_25_13, %iq2s_word1_high_25_14, %iq2s_word1_high_25_15, %iq2s_word1_high_25_16, %iq2s_word1_high_25_17, %iq2s_word1_high_25_18, %iq2s_word1_high_25_19, %iq2s_word1_high_25_20, %iq2s_word1_high_25_21, %iq2s_word1_high_25_22, %iq2s_word1_high_25_23, %iq2s_word1_high_25_24, %iq2s_word1_high_25_25, %iq2s_word1_high_25_26, %iq2s_word1_high_25_27, %iq2s_word1_high_25_28, %iq2s_word1_high_25_29, %iq2s_word1_high_25_30, %iq2s_word1_high_25_31 : vector<32xf32> + %iq2s_word1_high_26_0 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_1 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_2 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_3 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_4 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_5 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_6 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_7 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_8 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_9 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_10 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_11 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_12 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_13 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_14 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_15 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_16 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_17 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_18 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_19 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_20 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_21 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_22 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_23 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_24 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_25 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_26 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_27 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_28 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_29 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_30 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26_31 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_26 = vector.from_elements %iq2s_word1_high_26_0, %iq2s_word1_high_26_1, %iq2s_word1_high_26_2, %iq2s_word1_high_26_3, %iq2s_word1_high_26_4, %iq2s_word1_high_26_5, %iq2s_word1_high_26_6, %iq2s_word1_high_26_7, %iq2s_word1_high_26_8, %iq2s_word1_high_26_9, %iq2s_word1_high_26_10, %iq2s_word1_high_26_11, %iq2s_word1_high_26_12, %iq2s_word1_high_26_13, %iq2s_word1_high_26_14, %iq2s_word1_high_26_15, %iq2s_word1_high_26_16, %iq2s_word1_high_26_17, %iq2s_word1_high_26_18, %iq2s_word1_high_26_19, %iq2s_word1_high_26_20, %iq2s_word1_high_26_21, %iq2s_word1_high_26_22, %iq2s_word1_high_26_23, %iq2s_word1_high_26_24, %iq2s_word1_high_26_25, %iq2s_word1_high_26_26, %iq2s_word1_high_26_27, %iq2s_word1_high_26_28, %iq2s_word1_high_26_29, %iq2s_word1_high_26_30, %iq2s_word1_high_26_31 : vector<32xf32> + %iq2s_word1_high_27_0 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_1 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_2 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_3 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_4 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_5 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_6 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_7 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_8 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_9 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_10 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_11 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_12 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_13 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_14 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_15 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_16 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_17 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_18 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_19 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_20 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_21 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_22 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_23 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_24 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_25 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_26 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_27 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_28 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_29 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_30 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27_31 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_27 = vector.from_elements %iq2s_word1_high_27_0, %iq2s_word1_high_27_1, %iq2s_word1_high_27_2, %iq2s_word1_high_27_3, %iq2s_word1_high_27_4, %iq2s_word1_high_27_5, %iq2s_word1_high_27_6, %iq2s_word1_high_27_7, %iq2s_word1_high_27_8, %iq2s_word1_high_27_9, %iq2s_word1_high_27_10, %iq2s_word1_high_27_11, %iq2s_word1_high_27_12, %iq2s_word1_high_27_13, %iq2s_word1_high_27_14, %iq2s_word1_high_27_15, %iq2s_word1_high_27_16, %iq2s_word1_high_27_17, %iq2s_word1_high_27_18, %iq2s_word1_high_27_19, %iq2s_word1_high_27_20, %iq2s_word1_high_27_21, %iq2s_word1_high_27_22, %iq2s_word1_high_27_23, %iq2s_word1_high_27_24, %iq2s_word1_high_27_25, %iq2s_word1_high_27_26, %iq2s_word1_high_27_27, %iq2s_word1_high_27_28, %iq2s_word1_high_27_29, %iq2s_word1_high_27_30, %iq2s_word1_high_27_31 : vector<32xf32> + %iq2s_word1_high_28_0 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_1 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_2 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_3 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_4 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_5 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_6 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_7 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_8 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_9 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_10 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_11 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_12 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_13 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_14 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_15 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_16 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_17 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_18 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_19 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_20 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_21 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_22 = scalar.constant 11016.0 : f32 + %iq2s_word1_high_28_23 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_28_24 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_28_25 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_28_26 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_28_27 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_28_28 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_28_29 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_28_30 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_28_31 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_28 = vector.from_elements %iq2s_word1_high_28_0, %iq2s_word1_high_28_1, %iq2s_word1_high_28_2, %iq2s_word1_high_28_3, %iq2s_word1_high_28_4, %iq2s_word1_high_28_5, %iq2s_word1_high_28_6, %iq2s_word1_high_28_7, %iq2s_word1_high_28_8, %iq2s_word1_high_28_9, %iq2s_word1_high_28_10, %iq2s_word1_high_28_11, %iq2s_word1_high_28_12, %iq2s_word1_high_28_13, %iq2s_word1_high_28_14, %iq2s_word1_high_28_15, %iq2s_word1_high_28_16, %iq2s_word1_high_28_17, %iq2s_word1_high_28_18, %iq2s_word1_high_28_19, %iq2s_word1_high_28_20, %iq2s_word1_high_28_21, %iq2s_word1_high_28_22, %iq2s_word1_high_28_23, %iq2s_word1_high_28_24, %iq2s_word1_high_28_25, %iq2s_word1_high_28_26, %iq2s_word1_high_28_27, %iq2s_word1_high_28_28, %iq2s_word1_high_28_29, %iq2s_word1_high_28_30, %iq2s_word1_high_28_31 : vector<32xf32> + %iq2s_word1_high_29_0 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_1 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_2 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_3 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_4 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_5 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_6 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_7 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_8 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_9 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_10 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_11 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_12 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_13 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_14 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_15 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_16 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_17 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_18 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_19 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_20 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_21 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_22 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_23 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_24 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_25 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_26 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_27 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_28 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_29 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_30 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29_31 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_29 = vector.from_elements %iq2s_word1_high_29_0, %iq2s_word1_high_29_1, %iq2s_word1_high_29_2, %iq2s_word1_high_29_3, %iq2s_word1_high_29_4, %iq2s_word1_high_29_5, %iq2s_word1_high_29_6, %iq2s_word1_high_29_7, %iq2s_word1_high_29_8, %iq2s_word1_high_29_9, %iq2s_word1_high_29_10, %iq2s_word1_high_29_11, %iq2s_word1_high_29_12, %iq2s_word1_high_29_13, %iq2s_word1_high_29_14, %iq2s_word1_high_29_15, %iq2s_word1_high_29_16, %iq2s_word1_high_29_17, %iq2s_word1_high_29_18, %iq2s_word1_high_29_19, %iq2s_word1_high_29_20, %iq2s_word1_high_29_21, %iq2s_word1_high_29_22, %iq2s_word1_high_29_23, %iq2s_word1_high_29_24, %iq2s_word1_high_29_25, %iq2s_word1_high_29_26, %iq2s_word1_high_29_27, %iq2s_word1_high_29_28, %iq2s_word1_high_29_29, %iq2s_word1_high_29_30, %iq2s_word1_high_29_31 : vector<32xf32> + %iq2s_word1_high_30_0 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_1 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_2 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_3 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_4 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_5 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_6 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_7 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_8 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_9 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_10 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_11 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_12 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_13 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_14 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_15 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_16 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_17 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_18 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_19 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_20 = scalar.constant 11033.0 : f32 + %iq2s_word1_high_30_21 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_30_22 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_30_23 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_30_24 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_30_25 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_30_26 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_30_27 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_30_28 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_30_29 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_30_30 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_30_31 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_30 = vector.from_elements %iq2s_word1_high_30_0, %iq2s_word1_high_30_1, %iq2s_word1_high_30_2, %iq2s_word1_high_30_3, %iq2s_word1_high_30_4, %iq2s_word1_high_30_5, %iq2s_word1_high_30_6, %iq2s_word1_high_30_7, %iq2s_word1_high_30_8, %iq2s_word1_high_30_9, %iq2s_word1_high_30_10, %iq2s_word1_high_30_11, %iq2s_word1_high_30_12, %iq2s_word1_high_30_13, %iq2s_word1_high_30_14, %iq2s_word1_high_30_15, %iq2s_word1_high_30_16, %iq2s_word1_high_30_17, %iq2s_word1_high_30_18, %iq2s_word1_high_30_19, %iq2s_word1_high_30_20, %iq2s_word1_high_30_21, %iq2s_word1_high_30_22, %iq2s_word1_high_30_23, %iq2s_word1_high_30_24, %iq2s_word1_high_30_25, %iq2s_word1_high_30_26, %iq2s_word1_high_30_27, %iq2s_word1_high_30_28, %iq2s_word1_high_30_29, %iq2s_word1_high_30_30, %iq2s_word1_high_30_31 : vector<32xf32> + %iq2s_word1_high_31_0 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_1 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_2 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_3 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_4 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_5 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_6 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_7 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_8 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_9 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_10 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_11 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_12 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_13 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_14 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_15 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_16 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_17 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_18 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_19 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_20 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_21 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_22 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_23 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_24 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_25 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_26 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_27 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_28 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_29 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_30 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31_31 = scalar.constant 11051.0 : f32 + %iq2s_word1_high_31 = vector.from_elements %iq2s_word1_high_31_0, %iq2s_word1_high_31_1, %iq2s_word1_high_31_2, %iq2s_word1_high_31_3, %iq2s_word1_high_31_4, %iq2s_word1_high_31_5, %iq2s_word1_high_31_6, %iq2s_word1_high_31_7, %iq2s_word1_high_31_8, %iq2s_word1_high_31_9, %iq2s_word1_high_31_10, %iq2s_word1_high_31_11, %iq2s_word1_high_31_12, %iq2s_word1_high_31_13, %iq2s_word1_high_31_14, %iq2s_word1_high_31_15, %iq2s_word1_high_31_16, %iq2s_word1_high_31_17, %iq2s_word1_high_31_18, %iq2s_word1_high_31_19, %iq2s_word1_high_31_20, %iq2s_word1_high_31_21, %iq2s_word1_high_31_22, %iq2s_word1_high_31_23, %iq2s_word1_high_31_24, %iq2s_word1_high_31_25, %iq2s_word1_high_31_26, %iq2s_word1_high_31_27, %iq2s_word1_high_31_28, %iq2s_word1_high_31_29, %iq2s_word1_high_31_30, %iq2s_word1_high_31_31 : vector<32xf32> + %selected_word1_high1 = scf.select %is_chunk1, %iq2s_word1_high_1, %iq2s_word1_high_0 : vector<32xf32> + %selected_word1_high2 = scf.select %is_chunk2, %iq2s_word1_high_2, %selected_word1_high1 : vector<32xf32> + %selected_word1_high3 = scf.select %is_chunk3, %iq2s_word1_high_3, %selected_word1_high2 : vector<32xf32> + %selected_word1_high4 = scf.select %is_chunk4, %iq2s_word1_high_4, %selected_word1_high3 : vector<32xf32> + %selected_word1_high5 = scf.select %is_chunk5, %iq2s_word1_high_5, %selected_word1_high4 : vector<32xf32> + %selected_word1_high6 = scf.select %is_chunk6, %iq2s_word1_high_6, %selected_word1_high5 : vector<32xf32> + %selected_word1_high7 = scf.select %is_chunk7, %iq2s_word1_high_7, %selected_word1_high6 : vector<32xf32> + %selected_word1_high8 = scf.select %is_chunk8, %iq2s_word1_high_8, %selected_word1_high7 : vector<32xf32> + %selected_word1_high9 = scf.select %is_chunk9, %iq2s_word1_high_9, %selected_word1_high8 : vector<32xf32> + %selected_word1_high10 = scf.select %is_chunk10, %iq2s_word1_high_10, %selected_word1_high9 : vector<32xf32> + %selected_word1_high11 = scf.select %is_chunk11, %iq2s_word1_high_11, %selected_word1_high10 : vector<32xf32> + %selected_word1_high12 = scf.select %is_chunk12, %iq2s_word1_high_12, %selected_word1_high11 : vector<32xf32> + %selected_word1_high13 = scf.select %is_chunk13, %iq2s_word1_high_13, %selected_word1_high12 : vector<32xf32> + %selected_word1_high14 = scf.select %is_chunk14, %iq2s_word1_high_14, %selected_word1_high13 : vector<32xf32> + %selected_word1_high15 = scf.select %is_chunk15, %iq2s_word1_high_15, %selected_word1_high14 : vector<32xf32> + %selected_word1_high16 = scf.select %is_chunk16, %iq2s_word1_high_16, %selected_word1_high15 : vector<32xf32> + %selected_word1_high17 = scf.select %is_chunk17, %iq2s_word1_high_17, %selected_word1_high16 : vector<32xf32> + %selected_word1_high18 = scf.select %is_chunk18, %iq2s_word1_high_18, %selected_word1_high17 : vector<32xf32> + %selected_word1_high19 = scf.select %is_chunk19, %iq2s_word1_high_19, %selected_word1_high18 : vector<32xf32> + %selected_word1_high20 = scf.select %is_chunk20, %iq2s_word1_high_20, %selected_word1_high19 : vector<32xf32> + %selected_word1_high21 = scf.select %is_chunk21, %iq2s_word1_high_21, %selected_word1_high20 : vector<32xf32> + %selected_word1_high22 = scf.select %is_chunk22, %iq2s_word1_high_22, %selected_word1_high21 : vector<32xf32> + %selected_word1_high23 = scf.select %is_chunk23, %iq2s_word1_high_23, %selected_word1_high22 : vector<32xf32> + %selected_word1_high24 = scf.select %is_chunk24, %iq2s_word1_high_24, %selected_word1_high23 : vector<32xf32> + %selected_word1_high25 = scf.select %is_chunk25, %iq2s_word1_high_25, %selected_word1_high24 : vector<32xf32> + %selected_word1_high26 = scf.select %is_chunk26, %iq2s_word1_high_26, %selected_word1_high25 : vector<32xf32> + %selected_word1_high27 = scf.select %is_chunk27, %iq2s_word1_high_27, %selected_word1_high26 : vector<32xf32> + %selected_word1_high28 = scf.select %is_chunk28, %iq2s_word1_high_28, %selected_word1_high27 : vector<32xf32> + %selected_word1_high29 = scf.select %is_chunk29, %iq2s_word1_high_29, %selected_word1_high28 : vector<32xf32> + %selected_word1_high30 = scf.select %is_chunk30, %iq2s_word1_high_30, %selected_word1_high29 : vector<32xf32> + %selected_word1_high31 = scf.select %is_chunk31, %iq2s_word1_high_31, %selected_word1_high30 : vector<32xf32> + %use_word1 = scalar.cmpi eq, %word_index, %c1_i32 : i32 + %selected_low = scf.select %use_word1, %selected_word1_low31, %selected_word0_low31 : vector<32xf32> + %selected_high = scf.select %use_word1, %selected_word1_high31, %selected_word0_high31 : vector<32xf32> + %low_values = vector.table.lookup %selected_low[%codes] : vector<32xf32>, vector<1xi32> -> vector<1xf32> + %high_values = vector.table.lookup %selected_high[%codes] : vector<32xf32>, vector<1xi32> -> vector<1xf32> + %low_f32 = vector.extract %low_values[0] : vector<1xf32> -> f32 + %high_f32 = vector.extract %high_values[0] : vector<1xf32> -> f32 + %low = scalar.fptoui %low_f32 : f32 to i32 + %high = scalar.fptoui %high_f32 : f32 to i32 + %high_shifted = scalar.shli %high, %c16_i32 : i32 + %word = scalar.ori %low, %high_shifted : i32 + func.return %word : i32 +} +func.def inline @ggml_iq2s_value_f32(%grid_word: i32, %shift: i32, %signs: i32, %sign_bit: i32, %scale: f32) -> (f32) { + %c0_i32 = scalar.constant 0 : i32 + %c255_i32 = scalar.constant 255 : i32 + %c1_f32 = scalar.constant 1.0 : f32 + %cneg1_f32 = scalar.constant -1.0 : f32 + %shifted_word = scalar.shrui %grid_word, %shift : i32 + %value_i32 = scalar.andi %shifted_word, %c255_i32 : i32 + %value_f32 = scalar.uitofp %value_i32 : i32 to f32 + %sign_mask = scalar.andi %signs, %sign_bit : i32 + %negative = scalar.cmpi ne, %sign_mask, %c0_i32 : i32 + %sign = scf.select %negative, %cneg1_f32, %c1_f32 : f32 + %signed = scalar.mulf %value_f32, %sign : f32 + %result = scalar.mulf %signed, %scale : f32 + func.return %result : f32 +} + +func.def inline @ggml_iq2s_f32_vector4(%weight: buffer, %row_byte_base: offset, %iq2_block: index, %iq2_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c16_i32 = scalar.constant 16 : i32 + %c24_i32 = scalar.constant 24 : i32 + %c255_i32 = scalar.constant 255 : i32 + %c32_i32 = scalar.constant 32 : i32 + %c64_i32 = scalar.constant 64 : i32 + %c128_i32 = scalar.constant 128 : i32 + %c768_i32 = scalar.constant 768 : i32 + %c05_f32 = scalar.constant 0.5 : f32 + %c025_f32 = scalar.constant 0.25 : f32 + %block_bytes = index.constant 82 : offset + %qs_offset = index.constant 2 : offset + %qh_offset = index.constant 66 : offset + %scales_offset = index.constant 74 : offset + %block_byte_add = index.scale %iq2_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %qs_byte_base = index.add %block_byte_base, %qs_offset : offset + %qh_byte_base = index.add %block_byte_base, %qh_offset : offset + %scales_byte_base = index.add %block_byte_base, %scales_offset : offset + %d_view = buffer.view %weight[%block_byte_base] : buffer -> view<1xf16> + %qs_view = buffer.view %weight[%qs_byte_base] : buffer -> view<64xi8> + %qh_view = buffer.view %weight[%qh_byte_base] : buffer -> view<8xi8> + %scales_view = buffer.view %weight[%scales_byte_base] : buffer -> view<8xi8> + %bounded_group = index.assume %iq2_group [range(%iq2_group, 0, 7)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 63)] : index + %packet_in_group0 = index.rem %bounded_packet, %c8 : index + %l0 = index.div %packet_in_group0, %c2 : index + %l = index.assume %l0 [range(%l0, 0, 3)] : index + %word_index = index.rem %packet_in_group0, %c2 : index + %qs_group_base = index.mul %bounded_group, %c4 : index + %qs_index = index.add %qs_group_base, %l : index + %signs_index = index.add %qs_index, %c32 : index + %scale_half = index.div %l, %c2 : index + %uses_high_scale = index.cmp eq, %scale_half, %c1 : index + %d_f16 = view.load %d_view[%c0] : view<1xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %qs_i8 = view.load %qs_view[%qs_index] : view<64xi8> -> i8 + %qh_i8 = view.load %qh_view[%bounded_group] : view<8xi8> -> i8 + %signs_i8 = view.load %qs_view[%signs_index] : view<64xi8> -> i8 + %scale_i8 = view.load %scales_view[%bounded_group] : view<8xi8> -> i8 + %qs = scalar.extui %qs_i8 : i8 to i32 + %qh = scalar.extui %qh_i8 : i8 to i32 + %signs = scalar.extui %signs_i8 : i8 to i32 + %scale_byte = scalar.extui %scale_i8 : i8 to i32 + %scale_low = scalar.andi %scale_byte, %c15_i32 : i32 + %scale_high_shift = scalar.shrui %scale_byte, %c4_i32 : i32 + %scale_high = scalar.andi %scale_high_shift, %c15_i32 : i32 + %scale_i32 = scf.select %uses_high_scale, %scale_high, %scale_low : i32 + %scale_base = scalar.uitofp %scale_i32 : i32 to f32 + %scale_plus = scalar.addf %scale_base, %c05_f32 : f32 + %scale_mul = scalar.mulf %d, %scale_plus : f32 + %scale = scalar.mulf %scale_mul, %c025_f32 : f32 + %l_i32 = index.cast %l : index to i32 + %word_index_i32 = index.cast %word_index : index to i32 + %two_l = scalar.muli %l_i32, %c2_i32 : i32 + %qh_shift = scalar.subi %c8_i32, %two_l : i32 + %qh_shifted = scalar.shli %qh, %qh_shift : i32 + %qh_masked = scalar.andi %qh_shifted, %c768_i32 : i32 + %grid_index = scalar.ori %qs, %qh_masked : i32 + %grid_word = func.call @ggml_iq2s_grid_lookup_i32(%grid_index, %word_index_i32) : (i32, i32) -> (i32) + %uses_word1 = index.cmp eq, %word_index, %c1 : index + %sign_bit0 = scf.select %uses_word1, %c16_i32, %c1_i32 : i32 + %sign_bit1 = scf.select %uses_word1, %c32_i32, %c2_i32 : i32 + %sign_bit2 = scf.select %uses_word1, %c64_i32, %c4_i32 : i32 + %sign_bit3 = scf.select %uses_word1, %c128_i32, %c8_i32 : i32 + %v0 = func.call @ggml_iq2s_value_f32(%grid_word, %c0_i32, %signs, %sign_bit0, %scale) : (i32, i32, i32, i32, f32) -> (f32) + %v1 = func.call @ggml_iq2s_value_f32(%grid_word, %c8_i32, %signs, %sign_bit1, %scale) : (i32, i32, i32, i32, f32) -> (f32) + %v2 = func.call @ggml_iq2s_value_f32(%grid_word, %c16_i32, %signs, %sign_bit2, %scale) : (i32, i32, i32, i32, f32) -> (f32) + %v3 = func.call @ggml_iq2s_value_f32(%grid_word, %c24_i32, %signs, %sign_bit3, %scale) : (i32, i32, i32, i32, f32) -> (f32) + %result = vector.from_elements %v0, %v1, %v2, %v3 : vector<4xf32> + func.return %result : vector<4xf32> +} + +func.def inline @ggml_iq2s_f16_vector4(%weight: buffer, %row_byte_base: offset, %iq2_block: index, %iq2_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_iq2s_f32_vector4(%weight, %row_byte_base, %iq2_block, %iq2_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} +// iq2xxs_code: 256 grid entries as 16-bit codes (2 bits per value: 0 = 8, 1 = 25, 2 = 43) +func.def inline @ggml_iq2xxs_grid_code_i32(%grid_index: i32) -> (i32) { + %c5_i32 = scalar.constant 5 : i32 + %c31_i32 = scalar.constant 31 : i32 + %chunk_id1 = scalar.constant 1 : i32 + %chunk_id2 = scalar.constant 2 : i32 + %chunk_id3 = scalar.constant 3 : i32 + %chunk_id4 = scalar.constant 4 : i32 + %chunk_id5 = scalar.constant 5 : i32 + %chunk_id6 = scalar.constant 6 : i32 + %chunk_id7 = scalar.constant 7 : i32 + %chunk_i32 = scalar.shrui %grid_index, %c5_i32 : i32 + %lane_i32 = scalar.andi %grid_index, %c31_i32 : i32 + %codes = vector.from_elements %lane_i32 : vector<1xi32> + %iq2xxs_code_0_0 = scalar.constant 0.0 : f32 + %iq2xxs_code_0_1 = scalar.constant 2.0 : f32 + %iq2xxs_code_0_2 = scalar.constant 5.0 : f32 + %iq2xxs_code_0_3 = scalar.constant 8.0 : f32 + %iq2xxs_code_0_4 = scalar.constant 10.0 : f32 + %iq2xxs_code_0_5 = scalar.constant 17.0 : f32 + %iq2xxs_code_0_6 = scalar.constant 20.0 : f32 + %iq2xxs_code_0_7 = scalar.constant 32.0 : f32 + %iq2xxs_code_0_8 = scalar.constant 34.0 : f32 + %iq2xxs_code_0_9 = scalar.constant 40.0 : f32 + %iq2xxs_code_0_10 = scalar.constant 42.0 : f32 + %iq2xxs_code_0_11 = scalar.constant 65.0 : f32 + %iq2xxs_code_0_12 = scalar.constant 68.0 : f32 + %iq2xxs_code_0_13 = scalar.constant 80.0 : f32 + %iq2xxs_code_0_14 = scalar.constant 88.0 : f32 + %iq2xxs_code_0_15 = scalar.constant 97.0 : f32 + %iq2xxs_code_0_16 = scalar.constant 100.0 : f32 + %iq2xxs_code_0_17 = scalar.constant 128.0 : f32 + %iq2xxs_code_0_18 = scalar.constant 130.0 : f32 + %iq2xxs_code_0_19 = scalar.constant 138.0 : f32 + %iq2xxs_code_0_20 = scalar.constant 162.0 : f32 + %iq2xxs_code_0_21 = scalar.constant 257.0 : f32 + %iq2xxs_code_0_22 = scalar.constant 260.0 : f32 + %iq2xxs_code_0_23 = scalar.constant 272.0 : f32 + %iq2xxs_code_0_24 = scalar.constant 277.0 : f32 + %iq2xxs_code_0_25 = scalar.constant 320.0 : f32 + %iq2xxs_code_0_26 = scalar.constant 388.0 : f32 + %iq2xxs_code_0_27 = scalar.constant 408.0 : f32 + %iq2xxs_code_0_28 = scalar.constant 512.0 : f32 + %iq2xxs_code_0_29 = scalar.constant 514.0 : f32 + %iq2xxs_code_0_30 = scalar.constant 546.0 : f32 + %iq2xxs_code_0_31 = scalar.constant 642.0 : f32 + %iq2xxs_code_0 = vector.from_elements %iq2xxs_code_0_0, %iq2xxs_code_0_1, %iq2xxs_code_0_2, %iq2xxs_code_0_3, %iq2xxs_code_0_4, %iq2xxs_code_0_5, %iq2xxs_code_0_6, %iq2xxs_code_0_7, %iq2xxs_code_0_8, %iq2xxs_code_0_9, %iq2xxs_code_0_10, %iq2xxs_code_0_11, %iq2xxs_code_0_12, %iq2xxs_code_0_13, %iq2xxs_code_0_14, %iq2xxs_code_0_15, %iq2xxs_code_0_16, %iq2xxs_code_0_17, %iq2xxs_code_0_18, %iq2xxs_code_0_19, %iq2xxs_code_0_20, %iq2xxs_code_0_21, %iq2xxs_code_0_22, %iq2xxs_code_0_23, %iq2xxs_code_0_24, %iq2xxs_code_0_25, %iq2xxs_code_0_26, %iq2xxs_code_0_27, %iq2xxs_code_0_28, %iq2xxs_code_0_29, %iq2xxs_code_0_30, %iq2xxs_code_0_31 : vector<32xf32> + %iq2xxs_code_1_0 = scalar.constant 1025.0 : f32 + %iq2xxs_code_1_1 = scalar.constant 1028.0 : f32 + %iq2xxs_code_1_2 = scalar.constant 1040.0 : f32 + %iq2xxs_code_1_3 = scalar.constant 1057.0 : f32 + %iq2xxs_code_1_4 = scalar.constant 1060.0 : f32 + %iq2xxs_code_1_5 = scalar.constant 1088.0 : f32 + %iq2xxs_code_1_6 = scalar.constant 1090.0 : f32 + %iq2xxs_code_1_7 = scalar.constant 1096.0 : f32 + %iq2xxs_code_1_8 = scalar.constant 1120.0 : f32 + %iq2xxs_code_1_9 = scalar.constant 1153.0 : f32 + %iq2xxs_code_1_10 = scalar.constant 1156.0 : f32 + %iq2xxs_code_1_11 = scalar.constant 1168.0 : f32 + %iq2xxs_code_1_12 = scalar.constant 1188.0 : f32 + %iq2xxs_code_1_13 = scalar.constant 1280.0 : f32 + %iq2xxs_code_1_14 = scalar.constant 1282.0 : f32 + %iq2xxs_code_1_15 = scalar.constant 1288.0 : f32 + %iq2xxs_code_1_16 = scalar.constant 1312.0 : f32 + %iq2xxs_code_1_17 = scalar.constant 1350.0 : f32 + %iq2xxs_code_1_18 = scalar.constant 1385.0 : f32 + %iq2xxs_code_1_19 = scalar.constant 1408.0 : f32 + %iq2xxs_code_1_20 = scalar.constant 1425.0 : f32 + %iq2xxs_code_1_21 = scalar.constant 1545.0 : f32 + %iq2xxs_code_1_22 = scalar.constant 1552.0 : f32 + %iq2xxs_code_1_23 = scalar.constant 1600.0 : f32 + %iq2xxs_code_1_24 = scalar.constant 1668.0 : f32 + %iq2xxs_code_1_25 = scalar.constant 1700.0 : f32 + %iq2xxs_code_1_26 = scalar.constant 2048.0 : f32 + %iq2xxs_code_1_27 = scalar.constant 2053.0 : f32 + %iq2xxs_code_1_28 = scalar.constant 2056.0 : f32 + %iq2xxs_code_1_29 = scalar.constant 2068.0 : f32 + %iq2xxs_code_1_30 = scalar.constant 2088.0 : f32 + %iq2xxs_code_1_31 = scalar.constant 2113.0 : f32 + %iq2xxs_code_1 = vector.from_elements %iq2xxs_code_1_0, %iq2xxs_code_1_1, %iq2xxs_code_1_2, %iq2xxs_code_1_3, %iq2xxs_code_1_4, %iq2xxs_code_1_5, %iq2xxs_code_1_6, %iq2xxs_code_1_7, %iq2xxs_code_1_8, %iq2xxs_code_1_9, %iq2xxs_code_1_10, %iq2xxs_code_1_11, %iq2xxs_code_1_12, %iq2xxs_code_1_13, %iq2xxs_code_1_14, %iq2xxs_code_1_15, %iq2xxs_code_1_16, %iq2xxs_code_1_17, %iq2xxs_code_1_18, %iq2xxs_code_1_19, %iq2xxs_code_1_20, %iq2xxs_code_1_21, %iq2xxs_code_1_22, %iq2xxs_code_1_23, %iq2xxs_code_1_24, %iq2xxs_code_1_25, %iq2xxs_code_1_26, %iq2xxs_code_1_27, %iq2xxs_code_1_28, %iq2xxs_code_1_29, %iq2xxs_code_1_30, %iq2xxs_code_1_31 : vector<32xf32> + %iq2xxs_code_2_0 = scalar.constant 2116.0 : f32 + %iq2xxs_code_2_1 = scalar.constant 2128.0 : f32 + %iq2xxs_code_2_2 = scalar.constant 2130.0 : f32 + %iq2xxs_code_2_3 = scalar.constant 2184.0 : f32 + %iq2xxs_code_2_4 = scalar.constant 2308.0 : f32 + %iq2xxs_code_2_5 = scalar.constant 2368.0 : f32 + %iq2xxs_code_2_6 = scalar.constant 2562.0 : f32 + %iq2xxs_code_2_7 = scalar.constant 2580.0 : f32 + %iq2xxs_code_2_8 = scalar.constant 4097.0 : f32 + %iq2xxs_code_2_9 = scalar.constant 4100.0 : f32 + %iq2xxs_code_2_10 = scalar.constant 4112.0 : f32 + %iq2xxs_code_2_11 = scalar.constant 4129.0 : f32 + %iq2xxs_code_2_12 = scalar.constant 4160.0 : f32 + %iq2xxs_code_2_13 = scalar.constant 4192.0 : f32 + %iq2xxs_code_2_14 = scalar.constant 4228.0 : f32 + %iq2xxs_code_2_15 = scalar.constant 4240.0 : f32 + %iq2xxs_code_2_16 = scalar.constant 4245.0 : f32 + %iq2xxs_code_2_17 = scalar.constant 4352.0 : f32 + %iq2xxs_code_2_18 = scalar.constant 4360.0 : f32 + %iq2xxs_code_2_19 = scalar.constant 4384.0 : f32 + %iq2xxs_code_2_20 = scalar.constant 4432.0 : f32 + %iq2xxs_code_2_21 = scalar.constant 4442.0 : f32 + %iq2xxs_code_2_22 = scalar.constant 4480.0 : f32 + %iq2xxs_code_2_23 = scalar.constant 4644.0 : f32 + %iq2xxs_code_2_24 = scalar.constant 4677.0 : f32 + %iq2xxs_code_2_25 = scalar.constant 5120.0 : f32 + %iq2xxs_code_2_26 = scalar.constant 5128.0 : f32 + %iq2xxs_code_2_27 = scalar.constant 5152.0 : f32 + %iq2xxs_code_2_28 = scalar.constant 5157.0 : f32 + %iq2xxs_code_2_29 = scalar.constant 5193.0 : f32 + %iq2xxs_code_2_30 = scalar.constant 5248.0 : f32 + %iq2xxs_code_2_31 = scalar.constant 5400.0 : f32 + %iq2xxs_code_2 = vector.from_elements %iq2xxs_code_2_0, %iq2xxs_code_2_1, %iq2xxs_code_2_2, %iq2xxs_code_2_3, %iq2xxs_code_2_4, %iq2xxs_code_2_5, %iq2xxs_code_2_6, %iq2xxs_code_2_7, %iq2xxs_code_2_8, %iq2xxs_code_2_9, %iq2xxs_code_2_10, %iq2xxs_code_2_11, %iq2xxs_code_2_12, %iq2xxs_code_2_13, %iq2xxs_code_2_14, %iq2xxs_code_2_15, %iq2xxs_code_2_16, %iq2xxs_code_2_17, %iq2xxs_code_2_18, %iq2xxs_code_2_19, %iq2xxs_code_2_20, %iq2xxs_code_2_21, %iq2xxs_code_2_22, %iq2xxs_code_2_23, %iq2xxs_code_2_24, %iq2xxs_code_2_25, %iq2xxs_code_2_26, %iq2xxs_code_2_27, %iq2xxs_code_2_28, %iq2xxs_code_2_29, %iq2xxs_code_2_30, %iq2xxs_code_2_31 : vector<32xf32> + %iq2xxs_code_3_0 = scalar.constant 5474.0 : f32 + %iq2xxs_code_3_1 = scalar.constant 5632.0 : f32 + %iq2xxs_code_3_2 = scalar.constant 5654.0 : f32 + %iq2xxs_code_3_3 = scalar.constant 6145.0 : f32 + %iq2xxs_code_3_4 = scalar.constant 6148.0 : f32 + %iq2xxs_code_3_5 = scalar.constant 6160.0 : f32 + %iq2xxs_code_3_6 = scalar.constant 6208.0 : f32 + %iq2xxs_code_3_7 = scalar.constant 6273.0 : f32 + %iq2xxs_code_3_8 = scalar.constant 6400.0 : f32 + %iq2xxs_code_3_9 = scalar.constant 6405.0 : f32 + %iq2xxs_code_3_10 = scalar.constant 6560.0 : f32 + %iq2xxs_code_3_11 = scalar.constant 6737.0 : f32 + %iq2xxs_code_3_12 = scalar.constant 8192.0 : f32 + %iq2xxs_code_3_13 = scalar.constant 8194.0 : f32 + %iq2xxs_code_3_14 = scalar.constant 8202.0 : f32 + %iq2xxs_code_3_15 = scalar.constant 8260.0 : f32 + %iq2xxs_code_3_16 = scalar.constant 8289.0 : f32 + %iq2xxs_code_3_17 = scalar.constant 8320.0 : f32 + %iq2xxs_code_3_18 = scalar.constant 8322.0 : f32 + %iq2xxs_code_3_19 = scalar.constant 8489.0 : f32 + %iq2xxs_code_3_20 = scalar.constant 8520.0 : f32 + %iq2xxs_code_3_21 = scalar.constant 8704.0 : f32 + %iq2xxs_code_3_22 = scalar.constant 8706.0 : f32 + %iq2xxs_code_3_23 = scalar.constant 9217.0 : f32 + %iq2xxs_code_3_24 = scalar.constant 9220.0 : f32 + %iq2xxs_code_3_25 = scalar.constant 9232.0 : f32 + %iq2xxs_code_3_26 = scalar.constant 9280.0 : f32 + %iq2xxs_code_3_27 = scalar.constant 9302.0 : f32 + %iq2xxs_code_3_28 = scalar.constant 9472.0 : f32 + %iq2xxs_code_3_29 = scalar.constant 9537.0 : f32 + %iq2xxs_code_3_30 = scalar.constant 9572.0 : f32 + %iq2xxs_code_3_31 = scalar.constant 9872.0 : f32 + %iq2xxs_code_3 = vector.from_elements %iq2xxs_code_3_0, %iq2xxs_code_3_1, %iq2xxs_code_3_2, %iq2xxs_code_3_3, %iq2xxs_code_3_4, %iq2xxs_code_3_5, %iq2xxs_code_3_6, %iq2xxs_code_3_7, %iq2xxs_code_3_8, %iq2xxs_code_3_9, %iq2xxs_code_3_10, %iq2xxs_code_3_11, %iq2xxs_code_3_12, %iq2xxs_code_3_13, %iq2xxs_code_3_14, %iq2xxs_code_3_15, %iq2xxs_code_3_16, %iq2xxs_code_3_17, %iq2xxs_code_3_18, %iq2xxs_code_3_19, %iq2xxs_code_3_20, %iq2xxs_code_3_21, %iq2xxs_code_3_22, %iq2xxs_code_3_23, %iq2xxs_code_3_24, %iq2xxs_code_3_25, %iq2xxs_code_3_26, %iq2xxs_code_3_27, %iq2xxs_code_3_28, %iq2xxs_code_3_29, %iq2xxs_code_3_30, %iq2xxs_code_3_31 : vector<32xf32> + %iq2xxs_code_4_0 = scalar.constant 10248.0 : f32 + %iq2xxs_code_4_1 = scalar.constant 10272.0 : f32 + %iq2xxs_code_4_2 = scalar.constant 10388.0 : f32 + %iq2xxs_code_4_3 = scalar.constant 10820.0 : f32 + %iq2xxs_code_4_4 = scalar.constant 16385.0 : f32 + %iq2xxs_code_4_5 = scalar.constant 16388.0 : f32 + %iq2xxs_code_4_6 = scalar.constant 16400.0 : f32 + %iq2xxs_code_4_7 = scalar.constant 16408.0 : f32 + %iq2xxs_code_4_8 = scalar.constant 16417.0 : f32 + %iq2xxs_code_4_9 = scalar.constant 16420.0 : f32 + %iq2xxs_code_4_10 = scalar.constant 16448.0 : f32 + %iq2xxs_code_4_11 = scalar.constant 16456.0 : f32 + %iq2xxs_code_4_12 = scalar.constant 16470.0 : f32 + %iq2xxs_code_4_13 = scalar.constant 16480.0 : f32 + %iq2xxs_code_4_14 = scalar.constant 16513.0 : f32 + %iq2xxs_code_4_15 = scalar.constant 16516.0 : f32 + %iq2xxs_code_4_16 = scalar.constant 16528.0 : f32 + %iq2xxs_code_4_17 = scalar.constant 16640.0 : f32 + %iq2xxs_code_4_18 = scalar.constant 16672.0 : f32 + %iq2xxs_code_4_19 = scalar.constant 16737.0 : f32 + %iq2xxs_code_4_20 = scalar.constant 16768.0 : f32 + %iq2xxs_code_4_21 = scalar.constant 16773.0 : f32 + %iq2xxs_code_4_22 = scalar.constant 16897.0 : f32 + %iq2xxs_code_4_23 = scalar.constant 16912.0 : f32 + %iq2xxs_code_4_24 = scalar.constant 16968.0 : f32 + %iq2xxs_code_4_25 = scalar.constant 16982.0 : f32 + %iq2xxs_code_4_26 = scalar.constant 17000.0 : f32 + %iq2xxs_code_4_27 = scalar.constant 17408.0 : f32 + %iq2xxs_code_4_28 = scalar.constant 17416.0 : f32 + %iq2xxs_code_4_29 = scalar.constant 17440.0 : f32 + %iq2xxs_code_4_30 = scalar.constant 17536.0 : f32 + %iq2xxs_code_4_31 = scalar.constant 17561.0 : f32 + %iq2xxs_code_4 = vector.from_elements %iq2xxs_code_4_0, %iq2xxs_code_4_1, %iq2xxs_code_4_2, %iq2xxs_code_4_3, %iq2xxs_code_4_4, %iq2xxs_code_4_5, %iq2xxs_code_4_6, %iq2xxs_code_4_7, %iq2xxs_code_4_8, %iq2xxs_code_4_9, %iq2xxs_code_4_10, %iq2xxs_code_4_11, %iq2xxs_code_4_12, %iq2xxs_code_4_13, %iq2xxs_code_4_14, %iq2xxs_code_4_15, %iq2xxs_code_4_16, %iq2xxs_code_4_17, %iq2xxs_code_4_18, %iq2xxs_code_4_19, %iq2xxs_code_4_20, %iq2xxs_code_4_21, %iq2xxs_code_4_22, %iq2xxs_code_4_23, %iq2xxs_code_4_24, %iq2xxs_code_4_25, %iq2xxs_code_4_26, %iq2xxs_code_4_27, %iq2xxs_code_4_28, %iq2xxs_code_4_29, %iq2xxs_code_4_30, %iq2xxs_code_4_31 : vector<32xf32> + %iq2xxs_code_5_0 = scalar.constant 17682.0 : f32 + %iq2xxs_code_5_1 = scalar.constant 17700.0 : f32 + %iq2xxs_code_5_2 = scalar.constant 17920.0 : f32 + %iq2xxs_code_5_3 = scalar.constant 18433.0 : f32 + %iq2xxs_code_5_4 = scalar.constant 18436.0 : f32 + %iq2xxs_code_5_5 = scalar.constant 18448.0 : f32 + %iq2xxs_code_5_6 = scalar.constant 18496.0 : f32 + %iq2xxs_code_5_7 = scalar.constant 18501.0 : f32 + %iq2xxs_code_5_8 = scalar.constant 18688.0 : f32 + %iq2xxs_code_5_9 = scalar.constant 18776.0 : f32 + %iq2xxs_code_5_10 = scalar.constant 18785.0 : f32 + %iq2xxs_code_5_11 = scalar.constant 18818.0 : f32 + %iq2xxs_code_5_12 = scalar.constant 19013.0 : f32 + %iq2xxs_code_5_13 = scalar.constant 19088.0 : f32 + %iq2xxs_code_5_14 = scalar.constant 20480.0 : f32 + %iq2xxs_code_5_15 = scalar.constant 20488.0 : f32 + %iq2xxs_code_5_16 = scalar.constant 20497.0 : f32 + %iq2xxs_code_5_17 = scalar.constant 20505.0 : f32 + %iq2xxs_code_5_18 = scalar.constant 20512.0 : f32 + %iq2xxs_code_5_19 = scalar.constant 20608.0 : f32 + %iq2xxs_code_5_20 = scalar.constant 20616.0 : f32 + %iq2xxs_code_5_21 = scalar.constant 20740.0 : f32 + %iq2xxs_code_5_22 = scalar.constant 20802.0 : f32 + %iq2xxs_code_5_23 = scalar.constant 20900.0 : f32 + %iq2xxs_code_5_24 = scalar.constant 21137.0 : f32 + %iq2xxs_code_5_25 = scalar.constant 21648.0 : f32 + %iq2xxs_code_5_26 = scalar.constant 21650.0 : f32 + %iq2xxs_code_5_27 = scalar.constant 21770.0 : f32 + %iq2xxs_code_5_28 = scalar.constant 22017.0 : f32 + %iq2xxs_code_5_29 = scalar.constant 22100.0 : f32 + %iq2xxs_code_5_30 = scalar.constant 22528.0 : f32 + %iq2xxs_code_5_31 = scalar.constant 22545.0 : f32 + %iq2xxs_code_5 = vector.from_elements %iq2xxs_code_5_0, %iq2xxs_code_5_1, %iq2xxs_code_5_2, %iq2xxs_code_5_3, %iq2xxs_code_5_4, %iq2xxs_code_5_5, %iq2xxs_code_5_6, %iq2xxs_code_5_7, %iq2xxs_code_5_8, %iq2xxs_code_5_9, %iq2xxs_code_5_10, %iq2xxs_code_5_11, %iq2xxs_code_5_12, %iq2xxs_code_5_13, %iq2xxs_code_5_14, %iq2xxs_code_5_15, %iq2xxs_code_5_16, %iq2xxs_code_5_17, %iq2xxs_code_5_18, %iq2xxs_code_5_19, %iq2xxs_code_5_20, %iq2xxs_code_5_21, %iq2xxs_code_5_22, %iq2xxs_code_5_23, %iq2xxs_code_5_24, %iq2xxs_code_5_25, %iq2xxs_code_5_26, %iq2xxs_code_5_27, %iq2xxs_code_5_28, %iq2xxs_code_5_29, %iq2xxs_code_5_30, %iq2xxs_code_5_31 : vector<32xf32> + %iq2xxs_code_6_0 = scalar.constant 22553.0 : f32 + %iq2xxs_code_6_1 = scalar.constant 22628.0 : f32 + %iq2xxs_code_6_2 = scalar.constant 22848.0 : f32 + %iq2xxs_code_6_3 = scalar.constant 23048.0 : f32 + %iq2xxs_code_6_4 = scalar.constant 24580.0 : f32 + %iq2xxs_code_6_5 = scalar.constant 24592.0 : f32 + %iq2xxs_code_6_6 = scalar.constant 24640.0 : f32 + %iq2xxs_code_6_7 = scalar.constant 24680.0 : f32 + %iq2xxs_code_6_8 = scalar.constant 24832.0 : f32 + %iq2xxs_code_6_9 = scalar.constant 24917.0 : f32 + %iq2xxs_code_6_10 = scalar.constant 25112.0 : f32 + %iq2xxs_code_6_11 = scalar.constant 25184.0 : f32 + %iq2xxs_code_6_12 = scalar.constant 25600.0 : f32 + %iq2xxs_code_6_13 = scalar.constant 25605.0 : f32 + %iq2xxs_code_6_14 = scalar.constant 25872.0 : f32 + %iq2xxs_code_6_15 = scalar.constant 25874.0 : f32 + %iq2xxs_code_6_16 = scalar.constant 25988.0 : f32 + %iq2xxs_code_6_17 = scalar.constant 26690.0 : f32 + %iq2xxs_code_6_18 = scalar.constant 32768.0 : f32 + %iq2xxs_code_6_19 = scalar.constant 32770.0 : f32 + %iq2xxs_code_6_20 = scalar.constant 32778.0 : f32 + %iq2xxs_code_6_21 = scalar.constant 32833.0 : f32 + %iq2xxs_code_6_22 = scalar.constant 32898.0 : f32 + %iq2xxs_code_6_23 = scalar.constant 33028.0 : f32 + %iq2xxs_code_6_24 = scalar.constant 33048.0 : f32 + %iq2xxs_code_6_25 = scalar.constant 33088.0 : f32 + %iq2xxs_code_6_26 = scalar.constant 33297.0 : f32 + %iq2xxs_code_6_27 = scalar.constant 33793.0 : f32 + %iq2xxs_code_6_28 = scalar.constant 33796.0 : f32 + %iq2xxs_code_6_29 = scalar.constant 33808.0 : f32 + %iq2xxs_code_6_30 = scalar.constant 33813.0 : f32 + %iq2xxs_code_6_31 = scalar.constant 33856.0 : f32 + %iq2xxs_code_6 = vector.from_elements %iq2xxs_code_6_0, %iq2xxs_code_6_1, %iq2xxs_code_6_2, %iq2xxs_code_6_3, %iq2xxs_code_6_4, %iq2xxs_code_6_5, %iq2xxs_code_6_6, %iq2xxs_code_6_7, %iq2xxs_code_6_8, %iq2xxs_code_6_9, %iq2xxs_code_6_10, %iq2xxs_code_6_11, %iq2xxs_code_6_12, %iq2xxs_code_6_13, %iq2xxs_code_6_14, %iq2xxs_code_6_15, %iq2xxs_code_6_16, %iq2xxs_code_6_17, %iq2xxs_code_6_18, %iq2xxs_code_6_19, %iq2xxs_code_6_20, %iq2xxs_code_6_21, %iq2xxs_code_6_22, %iq2xxs_code_6_23, %iq2xxs_code_6_24, %iq2xxs_code_6_25, %iq2xxs_code_6_26, %iq2xxs_code_6_27, %iq2xxs_code_6_28, %iq2xxs_code_6_29, %iq2xxs_code_6_30, %iq2xxs_code_6_31 : vector<32xf32> + %iq2xxs_code_7_0 = scalar.constant 33888.0 : f32 + %iq2xxs_code_7_1 = scalar.constant 34048.0 : f32 + %iq2xxs_code_7_2 = scalar.constant 34118.0 : f32 + %iq2xxs_code_7_3 = scalar.constant 34196.0 : f32 + %iq2xxs_code_7_4 = scalar.constant 34313.0 : f32 + %iq2xxs_code_7_5 = scalar.constant 34368.0 : f32 + %iq2xxs_code_7_6 = scalar.constant 34400.0 : f32 + %iq2xxs_code_7_7 = scalar.constant 34818.0 : f32 + %iq2xxs_code_7_8 = scalar.constant 35076.0 : f32 + %iq2xxs_code_7_9 = scalar.constant 35345.0 : f32 + %iq2xxs_code_7_10 = scalar.constant 36868.0 : f32 + %iq2xxs_code_7_11 = scalar.constant 36880.0 : f32 + %iq2xxs_code_7_12 = scalar.constant 36900.0 : f32 + %iq2xxs_code_7_13 = scalar.constant 36928.0 : f32 + %iq2xxs_code_7_14 = scalar.constant 37025.0 : f32 + %iq2xxs_code_7_15 = scalar.constant 37142.0 : f32 + %iq2xxs_code_7_16 = scalar.constant 37248.0 : f32 + %iq2xxs_code_7_17 = scalar.constant 37445.0 : f32 + %iq2xxs_code_7_18 = scalar.constant 37888.0 : f32 + %iq2xxs_code_7_19 = scalar.constant 37922.0 : f32 + %iq2xxs_code_7_20 = scalar.constant 37956.0 : f32 + %iq2xxs_code_7_21 = scalar.constant 38225.0 : f32 + %iq2xxs_code_7_22 = scalar.constant 39041.0 : f32 + %iq2xxs_code_7_23 = scalar.constant 39200.0 : f32 + %iq2xxs_code_7_24 = scalar.constant 40962.0 : f32 + %iq2xxs_code_7_25 = scalar.constant 41040.0 : f32 + %iq2xxs_code_7_26 = scalar.constant 41093.0 : f32 + %iq2xxs_code_7_27 = scalar.constant 41225.0 : f32 + %iq2xxs_code_7_28 = scalar.constant 41472.0 : f32 + %iq2xxs_code_7_29 = scalar.constant 42008.0 : f32 + %iq2xxs_code_7_30 = scalar.constant 43088.0 : f32 + %iq2xxs_code_7_31 = scalar.constant 43268.0 : f32 + %iq2xxs_code_7 = vector.from_elements %iq2xxs_code_7_0, %iq2xxs_code_7_1, %iq2xxs_code_7_2, %iq2xxs_code_7_3, %iq2xxs_code_7_4, %iq2xxs_code_7_5, %iq2xxs_code_7_6, %iq2xxs_code_7_7, %iq2xxs_code_7_8, %iq2xxs_code_7_9, %iq2xxs_code_7_10, %iq2xxs_code_7_11, %iq2xxs_code_7_12, %iq2xxs_code_7_13, %iq2xxs_code_7_14, %iq2xxs_code_7_15, %iq2xxs_code_7_16, %iq2xxs_code_7_17, %iq2xxs_code_7_18, %iq2xxs_code_7_19, %iq2xxs_code_7_20, %iq2xxs_code_7_21, %iq2xxs_code_7_22, %iq2xxs_code_7_23, %iq2xxs_code_7_24, %iq2xxs_code_7_25, %iq2xxs_code_7_26, %iq2xxs_code_7_27, %iq2xxs_code_7_28, %iq2xxs_code_7_29, %iq2xxs_code_7_30, %iq2xxs_code_7_31 : vector<32xf32> + %is_chunk1 = scalar.cmpi eq, %chunk_i32, %chunk_id1 : i32 + %sel1 = scf.select %is_chunk1, %iq2xxs_code_1, %iq2xxs_code_0 : vector<32xf32> + %is_chunk2 = scalar.cmpi eq, %chunk_i32, %chunk_id2 : i32 + %sel2 = scf.select %is_chunk2, %iq2xxs_code_2, %sel1 : vector<32xf32> + %is_chunk3 = scalar.cmpi eq, %chunk_i32, %chunk_id3 : i32 + %sel3 = scf.select %is_chunk3, %iq2xxs_code_3, %sel2 : vector<32xf32> + %is_chunk4 = scalar.cmpi eq, %chunk_i32, %chunk_id4 : i32 + %sel4 = scf.select %is_chunk4, %iq2xxs_code_4, %sel3 : vector<32xf32> + %is_chunk5 = scalar.cmpi eq, %chunk_i32, %chunk_id5 : i32 + %sel5 = scf.select %is_chunk5, %iq2xxs_code_5, %sel4 : vector<32xf32> + %is_chunk6 = scalar.cmpi eq, %chunk_i32, %chunk_id6 : i32 + %sel6 = scf.select %is_chunk6, %iq2xxs_code_6, %sel5 : vector<32xf32> + %is_chunk7 = scalar.cmpi eq, %chunk_i32, %chunk_id7 : i32 + %sel7 = scf.select %is_chunk7, %iq2xxs_code_7, %sel6 : vector<32xf32> + %v = vector.table.lookup %sel7[%codes] : vector<32xf32>, vector<1xi32> -> vector<1xf32> + %v_f32 = vector.extract %v[0] : vector<1xf32> -> f32 + %code = scalar.fptoui %v_f32 : f32 to i32 + func.return %code : i32 +} + +// iq2xs_code: 512 grid entries as 16-bit codes (2 bits per value: 0 = 8, 1 = 25, 2 = 43) +func.def inline @ggml_iq2xs_grid_code_i32(%grid_index: i32) -> (i32) { + %c5_i32 = scalar.constant 5 : i32 + %c31_i32 = scalar.constant 31 : i32 + %chunk_id1 = scalar.constant 1 : i32 + %chunk_id2 = scalar.constant 2 : i32 + %chunk_id3 = scalar.constant 3 : i32 + %chunk_id4 = scalar.constant 4 : i32 + %chunk_id5 = scalar.constant 5 : i32 + %chunk_id6 = scalar.constant 6 : i32 + %chunk_id7 = scalar.constant 7 : i32 + %chunk_id8 = scalar.constant 8 : i32 + %chunk_id9 = scalar.constant 9 : i32 + %chunk_id10 = scalar.constant 10 : i32 + %chunk_id11 = scalar.constant 11 : i32 + %chunk_id12 = scalar.constant 12 : i32 + %chunk_id13 = scalar.constant 13 : i32 + %chunk_id14 = scalar.constant 14 : i32 + %chunk_id15 = scalar.constant 15 : i32 + %chunk_i32 = scalar.shrui %grid_index, %c5_i32 : i32 + %lane_i32 = scalar.andi %grid_index, %c31_i32 : i32 + %codes = vector.from_elements %lane_i32 : vector<1xi32> + %iq2xs_code_0_0 = scalar.constant 0.0 : f32 + %iq2xs_code_0_1 = scalar.constant 2.0 : f32 + %iq2xs_code_0_2 = scalar.constant 5.0 : f32 + %iq2xs_code_0_3 = scalar.constant 8.0 : f32 + %iq2xs_code_0_4 = scalar.constant 10.0 : f32 + %iq2xs_code_0_5 = scalar.constant 17.0 : f32 + %iq2xs_code_0_6 = scalar.constant 20.0 : f32 + %iq2xs_code_0_7 = scalar.constant 22.0 : f32 + %iq2xs_code_0_8 = scalar.constant 25.0 : f32 + %iq2xs_code_0_9 = scalar.constant 32.0 : f32 + %iq2xs_code_0_10 = scalar.constant 34.0 : f32 + %iq2xs_code_0_11 = scalar.constant 37.0 : f32 + %iq2xs_code_0_12 = scalar.constant 40.0 : f32 + %iq2xs_code_0_13 = scalar.constant 65.0 : f32 + %iq2xs_code_0_14 = scalar.constant 68.0 : f32 + %iq2xs_code_0_15 = scalar.constant 70.0 : f32 + %iq2xs_code_0_16 = scalar.constant 73.0 : f32 + %iq2xs_code_0_17 = scalar.constant 80.0 : f32 + %iq2xs_code_0_18 = scalar.constant 82.0 : f32 + %iq2xs_code_0_19 = scalar.constant 85.0 : f32 + %iq2xs_code_0_20 = scalar.constant 88.0 : f32 + %iq2xs_code_0_21 = scalar.constant 97.0 : f32 + %iq2xs_code_0_22 = scalar.constant 100.0 : f32 + %iq2xs_code_0_23 = scalar.constant 128.0 : f32 + %iq2xs_code_0_24 = scalar.constant 130.0 : f32 + %iq2xs_code_0_25 = scalar.constant 133.0 : f32 + %iq2xs_code_0_26 = scalar.constant 136.0 : f32 + %iq2xs_code_0_27 = scalar.constant 145.0 : f32 + %iq2xs_code_0_28 = scalar.constant 148.0 : f32 + %iq2xs_code_0_29 = scalar.constant 153.0 : f32 + %iq2xs_code_0_30 = scalar.constant 160.0 : f32 + %iq2xs_code_0_31 = scalar.constant 257.0 : f32 + %iq2xs_code_0 = vector.from_elements %iq2xs_code_0_0, %iq2xs_code_0_1, %iq2xs_code_0_2, %iq2xs_code_0_3, %iq2xs_code_0_4, %iq2xs_code_0_5, %iq2xs_code_0_6, %iq2xs_code_0_7, %iq2xs_code_0_8, %iq2xs_code_0_9, %iq2xs_code_0_10, %iq2xs_code_0_11, %iq2xs_code_0_12, %iq2xs_code_0_13, %iq2xs_code_0_14, %iq2xs_code_0_15, %iq2xs_code_0_16, %iq2xs_code_0_17, %iq2xs_code_0_18, %iq2xs_code_0_19, %iq2xs_code_0_20, %iq2xs_code_0_21, %iq2xs_code_0_22, %iq2xs_code_0_23, %iq2xs_code_0_24, %iq2xs_code_0_25, %iq2xs_code_0_26, %iq2xs_code_0_27, %iq2xs_code_0_28, %iq2xs_code_0_29, %iq2xs_code_0_30, %iq2xs_code_0_31 : vector<32xf32> + %iq2xs_code_1_0 = scalar.constant 260.0 : f32 + %iq2xs_code_1_1 = scalar.constant 262.0 : f32 + %iq2xs_code_1_2 = scalar.constant 265.0 : f32 + %iq2xs_code_1_3 = scalar.constant 272.0 : f32 + %iq2xs_code_1_4 = scalar.constant 274.0 : f32 + %iq2xs_code_1_5 = scalar.constant 277.0 : f32 + %iq2xs_code_1_6 = scalar.constant 280.0 : f32 + %iq2xs_code_1_7 = scalar.constant 282.0 : f32 + %iq2xs_code_1_8 = scalar.constant 289.0 : f32 + %iq2xs_code_1_9 = scalar.constant 292.0 : f32 + %iq2xs_code_1_10 = scalar.constant 320.0 : f32 + %iq2xs_code_1_11 = scalar.constant 322.0 : f32 + %iq2xs_code_1_12 = scalar.constant 325.0 : f32 + %iq2xs_code_1_13 = scalar.constant 328.0 : f32 + %iq2xs_code_1_14 = scalar.constant 337.0 : f32 + %iq2xs_code_1_15 = scalar.constant 340.0 : f32 + %iq2xs_code_1_16 = scalar.constant 352.0 : f32 + %iq2xs_code_1_17 = scalar.constant 360.0 : f32 + %iq2xs_code_1_18 = scalar.constant 385.0 : f32 + %iq2xs_code_1_19 = scalar.constant 388.0 : f32 + %iq2xs_code_1_20 = scalar.constant 400.0 : f32 + %iq2xs_code_1_21 = scalar.constant 512.0 : f32 + %iq2xs_code_1_22 = scalar.constant 514.0 : f32 + %iq2xs_code_1_23 = scalar.constant 517.0 : f32 + %iq2xs_code_1_24 = scalar.constant 520.0 : f32 + %iq2xs_code_1_25 = scalar.constant 529.0 : f32 + %iq2xs_code_1_26 = scalar.constant 532.0 : f32 + %iq2xs_code_1_27 = scalar.constant 544.0 : f32 + %iq2xs_code_1_28 = scalar.constant 577.0 : f32 + %iq2xs_code_1_29 = scalar.constant 580.0 : f32 + %iq2xs_code_1_30 = scalar.constant 592.0 : f32 + %iq2xs_code_1_31 = scalar.constant 597.0 : f32 + %iq2xs_code_1 = vector.from_elements %iq2xs_code_1_0, %iq2xs_code_1_1, %iq2xs_code_1_2, %iq2xs_code_1_3, %iq2xs_code_1_4, %iq2xs_code_1_5, %iq2xs_code_1_6, %iq2xs_code_1_7, %iq2xs_code_1_8, %iq2xs_code_1_9, %iq2xs_code_1_10, %iq2xs_code_1_11, %iq2xs_code_1_12, %iq2xs_code_1_13, %iq2xs_code_1_14, %iq2xs_code_1_15, %iq2xs_code_1_16, %iq2xs_code_1_17, %iq2xs_code_1_18, %iq2xs_code_1_19, %iq2xs_code_1_20, %iq2xs_code_1_21, %iq2xs_code_1_22, %iq2xs_code_1_23, %iq2xs_code_1_24, %iq2xs_code_1_25, %iq2xs_code_1_26, %iq2xs_code_1_27, %iq2xs_code_1_28, %iq2xs_code_1_29, %iq2xs_code_1_30, %iq2xs_code_1_31 : vector<32xf32> + %iq2xs_code_2_0 = scalar.constant 640.0 : f32 + %iq2xs_code_2_1 = scalar.constant 650.0 : f32 + %iq2xs_code_2_2 = scalar.constant 1025.0 : f32 + %iq2xs_code_2_3 = scalar.constant 1028.0 : f32 + %iq2xs_code_2_4 = scalar.constant 1030.0 : f32 + %iq2xs_code_2_5 = scalar.constant 1033.0 : f32 + %iq2xs_code_2_6 = scalar.constant 1040.0 : f32 + %iq2xs_code_2_7 = scalar.constant 1042.0 : f32 + %iq2xs_code_2_8 = scalar.constant 1045.0 : f32 + %iq2xs_code_2_9 = scalar.constant 1048.0 : f32 + %iq2xs_code_2_10 = scalar.constant 1057.0 : f32 + %iq2xs_code_2_11 = scalar.constant 1060.0 : f32 + %iq2xs_code_2_12 = scalar.constant 1088.0 : f32 + %iq2xs_code_2_13 = scalar.constant 1090.0 : f32 + %iq2xs_code_2_14 = scalar.constant 1093.0 : f32 + %iq2xs_code_2_15 = scalar.constant 1096.0 : f32 + %iq2xs_code_2_16 = scalar.constant 1105.0 : f32 + %iq2xs_code_2_17 = scalar.constant 1108.0 : f32 + %iq2xs_code_2_18 = scalar.constant 1110.0 : f32 + %iq2xs_code_2_19 = scalar.constant 1120.0 : f32 + %iq2xs_code_2_20 = scalar.constant 1153.0 : f32 + %iq2xs_code_2_21 = scalar.constant 1156.0 : f32 + %iq2xs_code_2_22 = scalar.constant 1168.0 : f32 + %iq2xs_code_2_23 = scalar.constant 1280.0 : f32 + %iq2xs_code_2_24 = scalar.constant 1282.0 : f32 + %iq2xs_code_2_25 = scalar.constant 1285.0 : f32 + %iq2xs_code_2_26 = scalar.constant 1288.0 : f32 + %iq2xs_code_2_27 = scalar.constant 1297.0 : f32 + %iq2xs_code_2_28 = scalar.constant 1300.0 : f32 + %iq2xs_code_2_29 = scalar.constant 1312.0 : f32 + %iq2xs_code_2_30 = scalar.constant 1345.0 : f32 + %iq2xs_code_2_31 = scalar.constant 1348.0 : f32 + %iq2xs_code_2 = vector.from_elements %iq2xs_code_2_0, %iq2xs_code_2_1, %iq2xs_code_2_2, %iq2xs_code_2_3, %iq2xs_code_2_4, %iq2xs_code_2_5, %iq2xs_code_2_6, %iq2xs_code_2_7, %iq2xs_code_2_8, %iq2xs_code_2_9, %iq2xs_code_2_10, %iq2xs_code_2_11, %iq2xs_code_2_12, %iq2xs_code_2_13, %iq2xs_code_2_14, %iq2xs_code_2_15, %iq2xs_code_2_16, %iq2xs_code_2_17, %iq2xs_code_2_18, %iq2xs_code_2_19, %iq2xs_code_2_20, %iq2xs_code_2_21, %iq2xs_code_2_22, %iq2xs_code_2_23, %iq2xs_code_2_24, %iq2xs_code_2_25, %iq2xs_code_2_26, %iq2xs_code_2_27, %iq2xs_code_2_28, %iq2xs_code_2_29, %iq2xs_code_2_30, %iq2xs_code_2_31 : vector<32xf32> + %iq2xs_code_3_0 = scalar.constant 1360.0 : f32 + %iq2xs_code_3_1 = scalar.constant 1377.0 : f32 + %iq2xs_code_3_2 = scalar.constant 1408.0 : f32 + %iq2xs_code_3_3 = scalar.constant 1537.0 : f32 + %iq2xs_code_3_4 = scalar.constant 1540.0 : f32 + %iq2xs_code_3_5 = scalar.constant 1552.0 : f32 + %iq2xs_code_3_6 = scalar.constant 1574.0 : f32 + %iq2xs_code_3_7 = scalar.constant 1600.0 : f32 + %iq2xs_code_3_8 = scalar.constant 1602.0 : f32 + %iq2xs_code_3_9 = scalar.constant 1668.0 : f32 + %iq2xs_code_3_10 = scalar.constant 2048.0 : f32 + %iq2xs_code_3_11 = scalar.constant 2050.0 : f32 + %iq2xs_code_3_12 = scalar.constant 2053.0 : f32 + %iq2xs_code_3_13 = scalar.constant 2056.0 : f32 + %iq2xs_code_3_14 = scalar.constant 2058.0 : f32 + %iq2xs_code_3_15 = scalar.constant 2065.0 : f32 + %iq2xs_code_3_16 = scalar.constant 2068.0 : f32 + %iq2xs_code_3_17 = scalar.constant 2080.0 : f32 + %iq2xs_code_3_18 = scalar.constant 2085.0 : f32 + %iq2xs_code_3_19 = scalar.constant 2113.0 : f32 + %iq2xs_code_3_20 = scalar.constant 2116.0 : f32 + %iq2xs_code_3_21 = scalar.constant 2128.0 : f32 + %iq2xs_code_3_22 = scalar.constant 2136.0 : f32 + %iq2xs_code_3_23 = scalar.constant 2176.0 : f32 + %iq2xs_code_3_24 = scalar.constant 2208.0 : f32 + %iq2xs_code_3_25 = scalar.constant 2218.0 : f32 + %iq2xs_code_3_26 = scalar.constant 2305.0 : f32 + %iq2xs_code_3_27 = scalar.constant 2308.0 : f32 + %iq2xs_code_3_28 = scalar.constant 2320.0 : f32 + %iq2xs_code_3_29 = scalar.constant 2368.0 : f32 + %iq2xs_code_3_30 = scalar.constant 2433.0 : f32 + %iq2xs_code_3_31 = scalar.constant 2441.0 : f32 + %iq2xs_code_3 = vector.from_elements %iq2xs_code_3_0, %iq2xs_code_3_1, %iq2xs_code_3_2, %iq2xs_code_3_3, %iq2xs_code_3_4, %iq2xs_code_3_5, %iq2xs_code_3_6, %iq2xs_code_3_7, %iq2xs_code_3_8, %iq2xs_code_3_9, %iq2xs_code_3_10, %iq2xs_code_3_11, %iq2xs_code_3_12, %iq2xs_code_3_13, %iq2xs_code_3_14, %iq2xs_code_3_15, %iq2xs_code_3_16, %iq2xs_code_3_17, %iq2xs_code_3_18, %iq2xs_code_3_19, %iq2xs_code_3_20, %iq2xs_code_3_21, %iq2xs_code_3_22, %iq2xs_code_3_23, %iq2xs_code_3_24, %iq2xs_code_3_25, %iq2xs_code_3_26, %iq2xs_code_3_27, %iq2xs_code_3_28, %iq2xs_code_3_29, %iq2xs_code_3_30, %iq2xs_code_3_31 : vector<32xf32> + %iq2xs_code_4_0 = scalar.constant 2560.0 : f32 + %iq2xs_code_4_1 = scalar.constant 2592.0 : f32 + %iq2xs_code_4_2 = scalar.constant 2600.0 : f32 + %iq2xs_code_4_3 = scalar.constant 2710.0 : f32 + %iq2xs_code_4_4 = scalar.constant 2720.0 : f32 + %iq2xs_code_4_5 = scalar.constant 4097.0 : f32 + %iq2xs_code_4_6 = scalar.constant 4100.0 : f32 + %iq2xs_code_4_7 = scalar.constant 4102.0 : f32 + %iq2xs_code_4_8 = scalar.constant 4105.0 : f32 + %iq2xs_code_4_9 = scalar.constant 4112.0 : f32 + %iq2xs_code_4_10 = scalar.constant 4114.0 : f32 + %iq2xs_code_4_11 = scalar.constant 4117.0 : f32 + %iq2xs_code_4_12 = scalar.constant 4120.0 : f32 + %iq2xs_code_4_13 = scalar.constant 4129.0 : f32 + %iq2xs_code_4_14 = scalar.constant 4132.0 : f32 + %iq2xs_code_4_15 = scalar.constant 4160.0 : f32 + %iq2xs_code_4_16 = scalar.constant 4162.0 : f32 + %iq2xs_code_4_17 = scalar.constant 4165.0 : f32 + %iq2xs_code_4_18 = scalar.constant 4168.0 : f32 + %iq2xs_code_4_19 = scalar.constant 4177.0 : f32 + %iq2xs_code_4_20 = scalar.constant 4180.0 : f32 + %iq2xs_code_4_21 = scalar.constant 4192.0 : f32 + %iq2xs_code_4_22 = scalar.constant 4202.0 : f32 + %iq2xs_code_4_23 = scalar.constant 4225.0 : f32 + %iq2xs_code_4_24 = scalar.constant 4228.0 : f32 + %iq2xs_code_4_25 = scalar.constant 4240.0 : f32 + %iq2xs_code_4_26 = scalar.constant 4352.0 : f32 + %iq2xs_code_4_27 = scalar.constant 4354.0 : f32 + %iq2xs_code_4_28 = scalar.constant 4357.0 : f32 + %iq2xs_code_4_29 = scalar.constant 4360.0 : f32 + %iq2xs_code_4_30 = scalar.constant 4369.0 : f32 + %iq2xs_code_4_31 = scalar.constant 4372.0 : f32 + %iq2xs_code_4 = vector.from_elements %iq2xs_code_4_0, %iq2xs_code_4_1, %iq2xs_code_4_2, %iq2xs_code_4_3, %iq2xs_code_4_4, %iq2xs_code_4_5, %iq2xs_code_4_6, %iq2xs_code_4_7, %iq2xs_code_4_8, %iq2xs_code_4_9, %iq2xs_code_4_10, %iq2xs_code_4_11, %iq2xs_code_4_12, %iq2xs_code_4_13, %iq2xs_code_4_14, %iq2xs_code_4_15, %iq2xs_code_4_16, %iq2xs_code_4_17, %iq2xs_code_4_18, %iq2xs_code_4_19, %iq2xs_code_4_20, %iq2xs_code_4_21, %iq2xs_code_4_22, %iq2xs_code_4_23, %iq2xs_code_4_24, %iq2xs_code_4_25, %iq2xs_code_4_26, %iq2xs_code_4_27, %iq2xs_code_4_28, %iq2xs_code_4_29, %iq2xs_code_4_30, %iq2xs_code_4_31 : vector<32xf32> + %iq2xs_code_5_0 = scalar.constant 4384.0 : f32 + %iq2xs_code_5_1 = scalar.constant 4417.0 : f32 + %iq2xs_code_5_2 = scalar.constant 4420.0 : f32 + %iq2xs_code_5_3 = scalar.constant 4432.0 : f32 + %iq2xs_code_5_4 = scalar.constant 4480.0 : f32 + %iq2xs_code_5_5 = scalar.constant 4500.0 : f32 + %iq2xs_code_5_6 = scalar.constant 4502.0 : f32 + %iq2xs_code_5_7 = scalar.constant 4609.0 : f32 + %iq2xs_code_5_8 = scalar.constant 4612.0 : f32 + %iq2xs_code_5_9 = scalar.constant 4614.0 : f32 + %iq2xs_code_5_10 = scalar.constant 4624.0 : f32 + %iq2xs_code_5_11 = scalar.constant 4672.0 : f32 + %iq2xs_code_5_12 = scalar.constant 4704.0 : f32 + %iq2xs_code_5_13 = scalar.constant 5120.0 : f32 + %iq2xs_code_5_14 = scalar.constant 5122.0 : f32 + %iq2xs_code_5_15 = scalar.constant 5125.0 : f32 + %iq2xs_code_5_16 = scalar.constant 5128.0 : f32 + %iq2xs_code_5_17 = scalar.constant 5137.0 : f32 + %iq2xs_code_5_18 = scalar.constant 5140.0 : f32 + %iq2xs_code_5_19 = scalar.constant 5152.0 : f32 + %iq2xs_code_5_20 = scalar.constant 5185.0 : f32 + %iq2xs_code_5_21 = scalar.constant 5188.0 : f32 + %iq2xs_code_5_22 = scalar.constant 5193.0 : f32 + %iq2xs_code_5_23 = scalar.constant 5200.0 : f32 + %iq2xs_code_5_24 = scalar.constant 5220.0 : f32 + %iq2xs_code_5_25 = scalar.constant 5248.0 : f32 + %iq2xs_code_5_26 = scalar.constant 5377.0 : f32 + %iq2xs_code_5_27 = scalar.constant 5380.0 : f32 + %iq2xs_code_5_28 = scalar.constant 5392.0 : f32 + %iq2xs_code_5_29 = scalar.constant 5440.0 : f32 + %iq2xs_code_5_30 = scalar.constant 5632.0 : f32 + %iq2xs_code_5_31 = scalar.constant 5652.0 : f32 + %iq2xs_code_5 = vector.from_elements %iq2xs_code_5_0, %iq2xs_code_5_1, %iq2xs_code_5_2, %iq2xs_code_5_3, %iq2xs_code_5_4, %iq2xs_code_5_5, %iq2xs_code_5_6, %iq2xs_code_5_7, %iq2xs_code_5_8, %iq2xs_code_5_9, %iq2xs_code_5_10, %iq2xs_code_5_11, %iq2xs_code_5_12, %iq2xs_code_5_13, %iq2xs_code_5_14, %iq2xs_code_5_15, %iq2xs_code_5_16, %iq2xs_code_5_17, %iq2xs_code_5_18, %iq2xs_code_5_19, %iq2xs_code_5_20, %iq2xs_code_5_21, %iq2xs_code_5_22, %iq2xs_code_5_23, %iq2xs_code_5_24, %iq2xs_code_5_25, %iq2xs_code_5_26, %iq2xs_code_5_27, %iq2xs_code_5_28, %iq2xs_code_5_29, %iq2xs_code_5_30, %iq2xs_code_5_31 : vector<32xf32> + %iq2xs_code_6_0 = scalar.constant 5705.0 : f32 + %iq2xs_code_6_1 = scalar.constant 6145.0 : f32 + %iq2xs_code_6_2 = scalar.constant 6148.0 : f32 + %iq2xs_code_6_3 = scalar.constant 6160.0 : f32 + %iq2xs_code_6_4 = scalar.constant 6162.0 : f32 + %iq2xs_code_6_5 = scalar.constant 6208.0 : f32 + %iq2xs_code_6_6 = scalar.constant 6228.0 : f32 + %iq2xs_code_6_7 = scalar.constant 6278.0 : f32 + %iq2xs_code_6_8 = scalar.constant 6400.0 : f32 + %iq2xs_code_6_9 = scalar.constant 6405.0 : f32 + %iq2xs_code_6_10 = scalar.constant 6502.0 : f32 + %iq2xs_code_6_11 = scalar.constant 6737.0 : f32 + %iq2xs_code_6_12 = scalar.constant 6825.0 : f32 + %iq2xs_code_6_13 = scalar.constant 8192.0 : f32 + %iq2xs_code_6_14 = scalar.constant 8194.0 : f32 + %iq2xs_code_6_15 = scalar.constant 8197.0 : f32 + %iq2xs_code_6_16 = scalar.constant 8200.0 : f32 + %iq2xs_code_6_17 = scalar.constant 8202.0 : f32 + %iq2xs_code_6_18 = scalar.constant 8209.0 : f32 + %iq2xs_code_6_19 = scalar.constant 8212.0 : f32 + %iq2xs_code_6_20 = scalar.constant 8224.0 : f32 + %iq2xs_code_6_21 = scalar.constant 8257.0 : f32 + %iq2xs_code_6_22 = scalar.constant 8260.0 : f32 + %iq2xs_code_6_23 = scalar.constant 8272.0 : f32 + %iq2xs_code_6_24 = scalar.constant 8320.0 : f32 + %iq2xs_code_6_25 = scalar.constant 8352.0 : f32 + %iq2xs_code_6_26 = scalar.constant 8449.0 : f32 + %iq2xs_code_6_27 = scalar.constant 8452.0 : f32 + %iq2xs_code_6_28 = scalar.constant 8464.0 : f32 + %iq2xs_code_6_29 = scalar.constant 8512.0 : f32 + %iq2xs_code_6_30 = scalar.constant 8520.0 : f32 + %iq2xs_code_6_31 = scalar.constant 8549.0 : f32 + %iq2xs_code_6 = vector.from_elements %iq2xs_code_6_0, %iq2xs_code_6_1, %iq2xs_code_6_2, %iq2xs_code_6_3, %iq2xs_code_6_4, %iq2xs_code_6_5, %iq2xs_code_6_6, %iq2xs_code_6_7, %iq2xs_code_6_8, %iq2xs_code_6_9, %iq2xs_code_6_10, %iq2xs_code_6_11, %iq2xs_code_6_12, %iq2xs_code_6_13, %iq2xs_code_6_14, %iq2xs_code_6_15, %iq2xs_code_6_16, %iq2xs_code_6_17, %iq2xs_code_6_18, %iq2xs_code_6_19, %iq2xs_code_6_20, %iq2xs_code_6_21, %iq2xs_code_6_22, %iq2xs_code_6_23, %iq2xs_code_6_24, %iq2xs_code_6_25, %iq2xs_code_6_26, %iq2xs_code_6_27, %iq2xs_code_6_28, %iq2xs_code_6_29, %iq2xs_code_6_30, %iq2xs_code_6_31 : vector<32xf32> + %iq2xs_code_7_0 = scalar.constant 8704.0 : f32 + %iq2xs_code_7_1 = scalar.constant 8738.0 : f32 + %iq2xs_code_7_2 = scalar.constant 8832.0 : f32 + %iq2xs_code_7_3 = scalar.constant 8872.0 : f32 + %iq2xs_code_7_4 = scalar.constant 9217.0 : f32 + %iq2xs_code_7_5 = scalar.constant 9220.0 : f32 + %iq2xs_code_7_6 = scalar.constant 9232.0 : f32 + %iq2xs_code_7_7 = scalar.constant 9257.0 : f32 + %iq2xs_code_7_8 = scalar.constant 9280.0 : f32 + %iq2xs_code_7_9 = scalar.constant 9472.0 : f32 + %iq2xs_code_7_10 = scalar.constant 9537.0 : f32 + %iq2xs_code_7_11 = scalar.constant 9554.0 : f32 + %iq2xs_code_7_12 = scalar.constant 9625.0 : f32 + %iq2xs_code_7_13 = scalar.constant 9729.0 : f32 + %iq2xs_code_7_14 = scalar.constant 9754.0 : f32 + %iq2xs_code_7_15 = scalar.constant 9894.0 : f32 + %iq2xs_code_7_16 = scalar.constant 10240.0 : f32 + %iq2xs_code_7_17 = scalar.constant 10248.0 : f32 + %iq2xs_code_7_18 = scalar.constant 10250.0 : f32 + %iq2xs_code_7_19 = scalar.constant 10272.0 : f32 + %iq2xs_code_7_20 = scalar.constant 10325.0 : f32 + %iq2xs_code_7_21 = scalar.constant 10376.0 : f32 + %iq2xs_code_7_22 = scalar.constant 10402.0 : f32 + %iq2xs_code_7_23 = scalar.constant 10600.0 : f32 + %iq2xs_code_7_24 = scalar.constant 10640.0 : f32 + %iq2xs_code_7_25 = scalar.constant 10760.0 : f32 + %iq2xs_code_7_26 = scalar.constant 10784.0 : f32 + %iq2xs_code_7_27 = scalar.constant 10882.0 : f32 + %iq2xs_code_7_28 = scalar.constant 10888.0 : f32 + %iq2xs_code_7_29 = scalar.constant 10890.0 : f32 + %iq2xs_code_7_30 = scalar.constant 16385.0 : f32 + %iq2xs_code_7_31 = scalar.constant 16388.0 : f32 + %iq2xs_code_7 = vector.from_elements %iq2xs_code_7_0, %iq2xs_code_7_1, %iq2xs_code_7_2, %iq2xs_code_7_3, %iq2xs_code_7_4, %iq2xs_code_7_5, %iq2xs_code_7_6, %iq2xs_code_7_7, %iq2xs_code_7_8, %iq2xs_code_7_9, %iq2xs_code_7_10, %iq2xs_code_7_11, %iq2xs_code_7_12, %iq2xs_code_7_13, %iq2xs_code_7_14, %iq2xs_code_7_15, %iq2xs_code_7_16, %iq2xs_code_7_17, %iq2xs_code_7_18, %iq2xs_code_7_19, %iq2xs_code_7_20, %iq2xs_code_7_21, %iq2xs_code_7_22, %iq2xs_code_7_23, %iq2xs_code_7_24, %iq2xs_code_7_25, %iq2xs_code_7_26, %iq2xs_code_7_27, %iq2xs_code_7_28, %iq2xs_code_7_29, %iq2xs_code_7_30, %iq2xs_code_7_31 : vector<32xf32> + %iq2xs_code_8_0 = scalar.constant 16390.0 : f32 + %iq2xs_code_8_1 = scalar.constant 16393.0 : f32 + %iq2xs_code_8_2 = scalar.constant 16400.0 : f32 + %iq2xs_code_8_3 = scalar.constant 16402.0 : f32 + %iq2xs_code_8_4 = scalar.constant 16405.0 : f32 + %iq2xs_code_8_5 = scalar.constant 16408.0 : f32 + %iq2xs_code_8_6 = scalar.constant 16417.0 : f32 + %iq2xs_code_8_7 = scalar.constant 16420.0 : f32 + %iq2xs_code_8_8 = scalar.constant 16448.0 : f32 + %iq2xs_code_8_9 = scalar.constant 16450.0 : f32 + %iq2xs_code_8_10 = scalar.constant 16453.0 : f32 + %iq2xs_code_8_11 = scalar.constant 16456.0 : f32 + %iq2xs_code_8_12 = scalar.constant 16458.0 : f32 + %iq2xs_code_8_13 = scalar.constant 16465.0 : f32 + %iq2xs_code_8_14 = scalar.constant 16468.0 : f32 + %iq2xs_code_8_15 = scalar.constant 16480.0 : f32 + %iq2xs_code_8_16 = scalar.constant 16485.0 : f32 + %iq2xs_code_8_17 = scalar.constant 16513.0 : f32 + %iq2xs_code_8_18 = scalar.constant 16516.0 : f32 + %iq2xs_code_8_19 = scalar.constant 16528.0 : f32 + %iq2xs_code_8_20 = scalar.constant 16640.0 : f32 + %iq2xs_code_8_21 = scalar.constant 16642.0 : f32 + %iq2xs_code_8_22 = scalar.constant 16645.0 : f32 + %iq2xs_code_8_23 = scalar.constant 16648.0 : f32 + %iq2xs_code_8_24 = scalar.constant 16657.0 : f32 + %iq2xs_code_8_25 = scalar.constant 16660.0 : f32 + %iq2xs_code_8_26 = scalar.constant 16672.0 : f32 + %iq2xs_code_8_27 = scalar.constant 16705.0 : f32 + %iq2xs_code_8_28 = scalar.constant 16708.0 : f32 + %iq2xs_code_8_29 = scalar.constant 16720.0 : f32 + %iq2xs_code_8_30 = scalar.constant 16768.0 : f32 + %iq2xs_code_8_31 = scalar.constant 16773.0 : f32 + %iq2xs_code_8 = vector.from_elements %iq2xs_code_8_0, %iq2xs_code_8_1, %iq2xs_code_8_2, %iq2xs_code_8_3, %iq2xs_code_8_4, %iq2xs_code_8_5, %iq2xs_code_8_6, %iq2xs_code_8_7, %iq2xs_code_8_8, %iq2xs_code_8_9, %iq2xs_code_8_10, %iq2xs_code_8_11, %iq2xs_code_8_12, %iq2xs_code_8_13, %iq2xs_code_8_14, %iq2xs_code_8_15, %iq2xs_code_8_16, %iq2xs_code_8_17, %iq2xs_code_8_18, %iq2xs_code_8_19, %iq2xs_code_8_20, %iq2xs_code_8_21, %iq2xs_code_8_22, %iq2xs_code_8_23, %iq2xs_code_8_24, %iq2xs_code_8_25, %iq2xs_code_8_26, %iq2xs_code_8_27, %iq2xs_code_8_28, %iq2xs_code_8_29, %iq2xs_code_8_30, %iq2xs_code_8_31 : vector<32xf32> + %iq2xs_code_9_0 = scalar.constant 16802.0 : f32 + %iq2xs_code_9_1 = scalar.constant 16897.0 : f32 + %iq2xs_code_9_2 = scalar.constant 16900.0 : f32 + %iq2xs_code_9_3 = scalar.constant 16912.0 : f32 + %iq2xs_code_9_4 = scalar.constant 16914.0 : f32 + %iq2xs_code_9_5 = scalar.constant 16937.0 : f32 + %iq2xs_code_9_6 = scalar.constant 16960.0 : f32 + %iq2xs_code_9_7 = scalar.constant 17408.0 : f32 + %iq2xs_code_9_8 = scalar.constant 17410.0 : f32 + %iq2xs_code_9_9 = scalar.constant 17413.0 : f32 + %iq2xs_code_9_10 = scalar.constant 17416.0 : f32 + %iq2xs_code_9_11 = scalar.constant 17425.0 : f32 + %iq2xs_code_9_12 = scalar.constant 17428.0 : f32 + %iq2xs_code_9_13 = scalar.constant 17433.0 : f32 + %iq2xs_code_9_14 = scalar.constant 17440.0 : f32 + %iq2xs_code_9_15 = scalar.constant 17473.0 : f32 + %iq2xs_code_9_16 = scalar.constant 17476.0 : f32 + %iq2xs_code_9_17 = scalar.constant 17488.0 : f32 + %iq2xs_code_9_18 = scalar.constant 17536.0 : f32 + %iq2xs_code_9_19 = scalar.constant 17556.0 : f32 + %iq2xs_code_9_20 = scalar.constant 17665.0 : f32 + %iq2xs_code_9_21 = scalar.constant 17668.0 : f32 + %iq2xs_code_9_22 = scalar.constant 17680.0 : f32 + %iq2xs_code_9_23 = scalar.constant 17700.0 : f32 + %iq2xs_code_9_24 = scalar.constant 17728.0 : f32 + %iq2xs_code_9_25 = scalar.constant 17818.0 : f32 + %iq2xs_code_9_26 = scalar.constant 17920.0 : f32 + %iq2xs_code_9_27 = scalar.constant 17930.0 : f32 + %iq2xs_code_9_28 = scalar.constant 17988.0 : f32 + %iq2xs_code_9_29 = scalar.constant 18000.0 : f32 + %iq2xs_code_9_30 = scalar.constant 18433.0 : f32 + %iq2xs_code_9_31 = scalar.constant 18436.0 : f32 + %iq2xs_code_9 = vector.from_elements %iq2xs_code_9_0, %iq2xs_code_9_1, %iq2xs_code_9_2, %iq2xs_code_9_3, %iq2xs_code_9_4, %iq2xs_code_9_5, %iq2xs_code_9_6, %iq2xs_code_9_7, %iq2xs_code_9_8, %iq2xs_code_9_9, %iq2xs_code_9_10, %iq2xs_code_9_11, %iq2xs_code_9_12, %iq2xs_code_9_13, %iq2xs_code_9_14, %iq2xs_code_9_15, %iq2xs_code_9_16, %iq2xs_code_9_17, %iq2xs_code_9_18, %iq2xs_code_9_19, %iq2xs_code_9_20, %iq2xs_code_9_21, %iq2xs_code_9_22, %iq2xs_code_9_23, %iq2xs_code_9_24, %iq2xs_code_9_25, %iq2xs_code_9_26, %iq2xs_code_9_27, %iq2xs_code_9_28, %iq2xs_code_9_29, %iq2xs_code_9_30, %iq2xs_code_9_31 : vector<32xf32> + %iq2xs_code_10_0 = scalar.constant 18448.0 : f32 + %iq2xs_code_10_1 = scalar.constant 18496.0 : f32 + %iq2xs_code_10_2 = scalar.constant 18501.0 : f32 + %iq2xs_code_10_3 = scalar.constant 18516.0 : f32 + %iq2xs_code_10_4 = scalar.constant 18530.0 : f32 + %iq2xs_code_10_5 = scalar.constant 18688.0 : f32 + %iq2xs_code_10_6 = scalar.constant 18705.0 : f32 + %iq2xs_code_10_7 = scalar.constant 18756.0 : f32 + %iq2xs_code_10_8 = scalar.constant 18768.0 : f32 + %iq2xs_code_10_9 = scalar.constant 18793.0 : f32 + %iq2xs_code_10_10 = scalar.constant 18948.0 : f32 + %iq2xs_code_10_11 = scalar.constant 20480.0 : f32 + %iq2xs_code_10_12 = scalar.constant 20482.0 : f32 + %iq2xs_code_10_13 = scalar.constant 20485.0 : f32 + %iq2xs_code_10_14 = scalar.constant 20488.0 : f32 + %iq2xs_code_10_15 = scalar.constant 20497.0 : f32 + %iq2xs_code_10_16 = scalar.constant 20500.0 : f32 + %iq2xs_code_10_17 = scalar.constant 20512.0 : f32 + %iq2xs_code_10_18 = scalar.constant 20520.0 : f32 + %iq2xs_code_10_19 = scalar.constant 20545.0 : f32 + %iq2xs_code_10_20 = scalar.constant 20548.0 : f32 + %iq2xs_code_10_21 = scalar.constant 20560.0 : f32 + %iq2xs_code_10_22 = scalar.constant 20608.0 : f32 + %iq2xs_code_10_23 = scalar.constant 20737.0 : f32 + %iq2xs_code_10_24 = scalar.constant 20740.0 : f32 + %iq2xs_code_10_25 = scalar.constant 20752.0 : f32 + %iq2xs_code_10_26 = scalar.constant 20757.0 : f32 + %iq2xs_code_10_27 = scalar.constant 20800.0 : f32 + %iq2xs_code_10_28 = scalar.constant 20802.0 : f32 + %iq2xs_code_10_29 = scalar.constant 20992.0 : f32 + %iq2xs_code_10_30 = scalar.constant 21060.0 : f32 + %iq2xs_code_10_31 = scalar.constant 21162.0 : f32 + %iq2xs_code_10 = vector.from_elements %iq2xs_code_10_0, %iq2xs_code_10_1, %iq2xs_code_10_2, %iq2xs_code_10_3, %iq2xs_code_10_4, %iq2xs_code_10_5, %iq2xs_code_10_6, %iq2xs_code_10_7, %iq2xs_code_10_8, %iq2xs_code_10_9, %iq2xs_code_10_10, %iq2xs_code_10_11, %iq2xs_code_10_12, %iq2xs_code_10_13, %iq2xs_code_10_14, %iq2xs_code_10_15, %iq2xs_code_10_16, %iq2xs_code_10_17, %iq2xs_code_10_18, %iq2xs_code_10_19, %iq2xs_code_10_20, %iq2xs_code_10_21, %iq2xs_code_10_22, %iq2xs_code_10_23, %iq2xs_code_10_24, %iq2xs_code_10_25, %iq2xs_code_10_26, %iq2xs_code_10_27, %iq2xs_code_10_28, %iq2xs_code_10_29, %iq2xs_code_10_30, %iq2xs_code_10_31 : vector<32xf32> + %iq2xs_code_11_0 = scalar.constant 21505.0 : f32 + %iq2xs_code_11_1 = scalar.constant 21508.0 : f32 + %iq2xs_code_11_2 = scalar.constant 21520.0 : f32 + %iq2xs_code_11_3 = scalar.constant 21537.0 : f32 + %iq2xs_code_11_4 = scalar.constant 21568.0 : f32 + %iq2xs_code_11_5 = scalar.constant 21600.0 : f32 + %iq2xs_code_11_6 = scalar.constant 21633.0 : f32 + %iq2xs_code_11_7 = scalar.constant 21665.0 : f32 + %iq2xs_code_11_8 = scalar.constant 21760.0 : f32 + %iq2xs_code_11_9 = scalar.constant 21768.0 : f32 + %iq2xs_code_11_10 = scalar.constant 21888.0 : f32 + %iq2xs_code_11_11 = scalar.constant 21896.0 : f32 + %iq2xs_code_11_12 = scalar.constant 22049.0 : f32 + %iq2xs_code_11_13 = scalar.constant 22120.0 : f32 + %iq2xs_code_11_14 = scalar.constant 22177.0 : f32 + %iq2xs_code_11_15 = scalar.constant 22528.0 : f32 + %iq2xs_code_11_16 = scalar.constant 22548.0 : f32 + %iq2xs_code_11_17 = scalar.constant 22593.0 : f32 + %iq2xs_code_11_18 = scalar.constant 22608.0 : f32 + %iq2xs_code_11_19 = scalar.constant 22681.0 : f32 + %iq2xs_code_11_20 = scalar.constant 22810.0 : f32 + %iq2xs_code_11_21 = scalar.constant 22848.0 : f32 + %iq2xs_code_11_22 = scalar.constant 22850.0 : f32 + %iq2xs_code_11_23 = scalar.constant 23173.0 : f32 + %iq2xs_code_11_24 = scalar.constant 24577.0 : f32 + %iq2xs_code_11_25 = scalar.constant 24580.0 : f32 + %iq2xs_code_11_26 = scalar.constant 24592.0 : f32 + %iq2xs_code_11_27 = scalar.constant 24640.0 : f32 + %iq2xs_code_11_28 = scalar.constant 24660.0 : f32 + %iq2xs_code_11_29 = scalar.constant 24674.0 : f32 + %iq2xs_code_11_30 = scalar.constant 24710.0 : f32 + %iq2xs_code_11_31 = scalar.constant 24745.0 : f32 + %iq2xs_code_11 = vector.from_elements %iq2xs_code_11_0, %iq2xs_code_11_1, %iq2xs_code_11_2, %iq2xs_code_11_3, %iq2xs_code_11_4, %iq2xs_code_11_5, %iq2xs_code_11_6, %iq2xs_code_11_7, %iq2xs_code_11_8, %iq2xs_code_11_9, %iq2xs_code_11_10, %iq2xs_code_11_11, %iq2xs_code_11_12, %iq2xs_code_11_13, %iq2xs_code_11_14, %iq2xs_code_11_15, %iq2xs_code_11_16, %iq2xs_code_11_17, %iq2xs_code_11_18, %iq2xs_code_11_19, %iq2xs_code_11_20, %iq2xs_code_11_21, %iq2xs_code_11_22, %iq2xs_code_11_23, %iq2xs_code_11_24, %iq2xs_code_11_25, %iq2xs_code_11_26, %iq2xs_code_11_27, %iq2xs_code_11_28, %iq2xs_code_11_29, %iq2xs_code_11_30, %iq2xs_code_11_31 : vector<32xf32> + %iq2xs_code_12_0 = scalar.constant 24832.0 : f32 + %iq2xs_code_12_1 = scalar.constant 25124.0 : f32 + %iq2xs_code_12_2 = scalar.constant 25162.0 : f32 + %iq2xs_code_12_3 = scalar.constant 25234.0 : f32 + %iq2xs_code_12_4 = scalar.constant 25600.0 : f32 + %iq2xs_code_12_5 = scalar.constant 25622.0 : f32 + %iq2xs_code_12_6 = scalar.constant 25872.0 : f32 + %iq2xs_code_12_7 = scalar.constant 25920.0 : f32 + %iq2xs_code_12_8 = scalar.constant 25925.0 : f32 + %iq2xs_code_12_9 = scalar.constant 26020.0 : f32 + %iq2xs_code_12_10 = scalar.constant 26625.0 : f32 + %iq2xs_code_12_11 = scalar.constant 26730.0 : f32 + %iq2xs_code_12_12 = scalar.constant 26917.0 : f32 + %iq2xs_code_12_13 = scalar.constant 27142.0 : f32 + %iq2xs_code_12_14 = scalar.constant 27220.0 : f32 + %iq2xs_code_12_15 = scalar.constant 27234.0 : f32 + %iq2xs_code_12_16 = scalar.constant 32768.0 : f32 + %iq2xs_code_12_17 = scalar.constant 32770.0 : f32 + %iq2xs_code_12_18 = scalar.constant 32773.0 : f32 + %iq2xs_code_12_19 = scalar.constant 32776.0 : f32 + %iq2xs_code_12_20 = scalar.constant 32785.0 : f32 + %iq2xs_code_12_21 = scalar.constant 32788.0 : f32 + %iq2xs_code_12_22 = scalar.constant 32800.0 : f32 + %iq2xs_code_12_23 = scalar.constant 32810.0 : f32 + %iq2xs_code_12_24 = scalar.constant 32833.0 : f32 + %iq2xs_code_12_25 = scalar.constant 32836.0 : f32 + %iq2xs_code_12_26 = scalar.constant 32848.0 : f32 + %iq2xs_code_12_27 = scalar.constant 32896.0 : f32 + %iq2xs_code_12_28 = scalar.constant 32898.0 : f32 + %iq2xs_code_12_29 = scalar.constant 32936.0 : f32 + %iq2xs_code_12_30 = scalar.constant 32938.0 : f32 + %iq2xs_code_12_31 = scalar.constant 33025.0 : f32 + %iq2xs_code_12 = vector.from_elements %iq2xs_code_12_0, %iq2xs_code_12_1, %iq2xs_code_12_2, %iq2xs_code_12_3, %iq2xs_code_12_4, %iq2xs_code_12_5, %iq2xs_code_12_6, %iq2xs_code_12_7, %iq2xs_code_12_8, %iq2xs_code_12_9, %iq2xs_code_12_10, %iq2xs_code_12_11, %iq2xs_code_12_12, %iq2xs_code_12_13, %iq2xs_code_12_14, %iq2xs_code_12_15, %iq2xs_code_12_16, %iq2xs_code_12_17, %iq2xs_code_12_18, %iq2xs_code_12_19, %iq2xs_code_12_20, %iq2xs_code_12_21, %iq2xs_code_12_22, %iq2xs_code_12_23, %iq2xs_code_12_24, %iq2xs_code_12_25, %iq2xs_code_12_26, %iq2xs_code_12_27, %iq2xs_code_12_28, %iq2xs_code_12_29, %iq2xs_code_12_30, %iq2xs_code_12_31 : vector<32xf32> + %iq2xs_code_13_0 = scalar.constant 33028.0 : f32 + %iq2xs_code_13_1 = scalar.constant 33030.0 : f32 + %iq2xs_code_13_2 = scalar.constant 33040.0 : f32 + %iq2xs_code_13_3 = scalar.constant 33088.0 : f32 + %iq2xs_code_13_4 = scalar.constant 33105.0 : f32 + %iq2xs_code_13_5 = scalar.constant 33113.0 : f32 + %iq2xs_code_13_6 = scalar.constant 33280.0 : f32 + %iq2xs_code_13_7 = scalar.constant 33312.0 : f32 + %iq2xs_code_13_8 = scalar.constant 33408.0 : f32 + %iq2xs_code_13_9 = scalar.constant 33410.0 : f32 + %iq2xs_code_13_10 = scalar.constant 33440.0 : f32 + %iq2xs_code_13_11 = scalar.constant 33448.0 : f32 + %iq2xs_code_13_12 = scalar.constant 33793.0 : f32 + %iq2xs_code_13_13 = scalar.constant 33796.0 : f32 + %iq2xs_code_13_14 = scalar.constant 33808.0 : f32 + %iq2xs_code_13_15 = scalar.constant 33810.0 : f32 + %iq2xs_code_13_16 = scalar.constant 33813.0 : f32 + %iq2xs_code_13_17 = scalar.constant 33856.0 : f32 + %iq2xs_code_13_18 = scalar.constant 33888.0 : f32 + %iq2xs_code_13_19 = scalar.constant 33929.0 : f32 + %iq2xs_code_13_20 = scalar.constant 34048.0 : f32 + %iq2xs_code_13_21 = scalar.constant 34116.0 : f32 + %iq2xs_code_13_22 = scalar.constant 34213.0 : f32 + %iq2xs_code_13_23 = scalar.constant 34328.0 : f32 + %iq2xs_code_13_24 = scalar.constant 34410.0 : f32 + %iq2xs_code_13_25 = scalar.constant 34816.0 : f32 + %iq2xs_code_13_26 = scalar.constant 34824.0 : f32 + %iq2xs_code_13_27 = scalar.constant 34853.0 : f32 + %iq2xs_code_13_28 = scalar.constant 34906.0 : f32 + %iq2xs_code_13_29 = scalar.constant 34944.0 : f32 + %iq2xs_code_13_30 = scalar.constant 34946.0 : f32 + %iq2xs_code_13_31 = scalar.constant 34984.0 : f32 + %iq2xs_code_13 = vector.from_elements %iq2xs_code_13_0, %iq2xs_code_13_1, %iq2xs_code_13_2, %iq2xs_code_13_3, %iq2xs_code_13_4, %iq2xs_code_13_5, %iq2xs_code_13_6, %iq2xs_code_13_7, %iq2xs_code_13_8, %iq2xs_code_13_9, %iq2xs_code_13_10, %iq2xs_code_13_11, %iq2xs_code_13_12, %iq2xs_code_13_13, %iq2xs_code_13_14, %iq2xs_code_13_15, %iq2xs_code_13_16, %iq2xs_code_13_17, %iq2xs_code_13_18, %iq2xs_code_13_19, %iq2xs_code_13_20, %iq2xs_code_13_21, %iq2xs_code_13_22, %iq2xs_code_13_23, %iq2xs_code_13_24, %iq2xs_code_13_25, %iq2xs_code_13_26, %iq2xs_code_13_27, %iq2xs_code_13_28, %iq2xs_code_13_29, %iq2xs_code_13_30, %iq2xs_code_13_31 : vector<32xf32> + %iq2xs_code_14_0 = scalar.constant 35078.0 : f32 + %iq2xs_code_14_1 = scalar.constant 35362.0 : f32 + %iq2xs_code_14_2 = scalar.constant 35456.0 : f32 + %iq2xs_code_14_3 = scalar.constant 35464.0 : f32 + %iq2xs_code_14_4 = scalar.constant 35478.0 : f32 + %iq2xs_code_14_5 = scalar.constant 35496.0 : f32 + %iq2xs_code_14_6 = scalar.constant 36865.0 : f32 + %iq2xs_code_14_7 = scalar.constant 36868.0 : f32 + %iq2xs_code_14_8 = scalar.constant 36880.0 : f32 + %iq2xs_code_14_9 = scalar.constant 36928.0 : f32 + %iq2xs_code_14_10 = scalar.constant 36950.0 : f32 + %iq2xs_code_14_11 = scalar.constant 36996.0 : f32 + %iq2xs_code_14_12 = scalar.constant 37120.0 : f32 + %iq2xs_code_14_13 = scalar.constant 37154.0 : f32 + %iq2xs_code_14_14 = scalar.constant 37220.0 : f32 + %iq2xs_code_14_15 = scalar.constant 37462.0 : f32 + %iq2xs_code_14_16 = scalar.constant 37513.0 : f32 + %iq2xs_code_14_17 = scalar.constant 37888.0 : f32 + %iq2xs_code_14_18 = scalar.constant 37893.0 : f32 + %iq2xs_code_14_19 = scalar.constant 37956.0 : f32 + %iq2xs_code_14_20 = scalar.constant 37968.0 : f32 + %iq2xs_code_14_21 = scalar.constant 37976.0 : f32 + %iq2xs_code_14_22 = scalar.constant 38185.0 : f32 + %iq2xs_code_14_23 = scalar.constant 38288.0 : f32 + %iq2xs_code_14_24 = scalar.constant 38290.0 : f32 + %iq2xs_code_14_25 = scalar.constant 38465.0 : f32 + %iq2xs_code_14_26 = scalar.constant 38993.0 : f32 + %iq2xs_code_14_27 = scalar.constant 39078.0 : f32 + %iq2xs_code_14_28 = scalar.constant 39241.0 : f32 + %iq2xs_code_14_29 = scalar.constant 39445.0 : f32 + %iq2xs_code_14_30 = scalar.constant 39520.0 : f32 + %iq2xs_code_14_31 = scalar.constant 40960.0 : f32 + %iq2xs_code_14 = vector.from_elements %iq2xs_code_14_0, %iq2xs_code_14_1, %iq2xs_code_14_2, %iq2xs_code_14_3, %iq2xs_code_14_4, %iq2xs_code_14_5, %iq2xs_code_14_6, %iq2xs_code_14_7, %iq2xs_code_14_8, %iq2xs_code_14_9, %iq2xs_code_14_10, %iq2xs_code_14_11, %iq2xs_code_14_12, %iq2xs_code_14_13, %iq2xs_code_14_14, %iq2xs_code_14_15, %iq2xs_code_14_16, %iq2xs_code_14_17, %iq2xs_code_14_18, %iq2xs_code_14_19, %iq2xs_code_14_20, %iq2xs_code_14_21, %iq2xs_code_14_22, %iq2xs_code_14_23, %iq2xs_code_14_24, %iq2xs_code_14_25, %iq2xs_code_14_26, %iq2xs_code_14_27, %iq2xs_code_14_28, %iq2xs_code_14_29, %iq2xs_code_14_30, %iq2xs_code_14_31 : vector<32xf32> + %iq2xs_code_15_0 = scalar.constant 40962.0 : f32 + %iq2xs_code_15_1 = scalar.constant 40968.0 : f32 + %iq2xs_code_15_2 = scalar.constant 40970.0 : f32 + %iq2xs_code_15_3 = scalar.constant 40992.0 : f32 + %iq2xs_code_15_4 = scalar.constant 41002.0 : f32 + %iq2xs_code_15_5 = scalar.constant 41120.0 : f32 + %iq2xs_code_15_6 = scalar.constant 41297.0 : f32 + %iq2xs_code_15_7 = scalar.constant 41305.0 : f32 + %iq2xs_code_15_8 = scalar.constant 41382.0 : f32 + %iq2xs_code_15_9 = scalar.constant 41472.0 : f32 + %iq2xs_code_15_10 = scalar.constant 41474.0 : f32 + %iq2xs_code_15_11 = scalar.constant 41480.0 : f32 + %iq2xs_code_15_12 = scalar.constant 41514.0 : f32 + %iq2xs_code_15_13 = scalar.constant 41600.0 : f32 + %iq2xs_code_15_14 = scalar.constant 41632.0 : f32 + %iq2xs_code_15_15 = scalar.constant 42048.0 : f32 + %iq2xs_code_15_16 = scalar.constant 42133.0 : f32 + %iq2xs_code_15_17 = scalar.constant 42597.0 : f32 + %iq2xs_code_15_18 = scalar.constant 42648.0 : f32 + %iq2xs_code_15_19 = scalar.constant 43018.0 : f32 + %iq2xs_code_15_20 = scalar.constant 43040.0 : f32 + %iq2xs_code_15_21 = scalar.constant 43042.0 : f32 + %iq2xs_code_15_22 = scalar.constant 43048.0 : f32 + %iq2xs_code_15_23 = scalar.constant 43168.0 : f32 + %iq2xs_code_15_24 = scalar.constant 43176.0 : f32 + %iq2xs_code_15_25 = scalar.constant 43268.0 : f32 + %iq2xs_code_15_26 = scalar.constant 43396.0 : f32 + %iq2xs_code_15_27 = scalar.constant 43398.0 : f32 + %iq2xs_code_15_28 = scalar.constant 43560.0 : f32 + %iq2xs_code_15_29 = scalar.constant 43562.0 : f32 + %iq2xs_code_15_30 = scalar.constant 43665.0 : f32 + %iq2xs_code_15_31 = scalar.constant 43690.0 : f32 + %iq2xs_code_15 = vector.from_elements %iq2xs_code_15_0, %iq2xs_code_15_1, %iq2xs_code_15_2, %iq2xs_code_15_3, %iq2xs_code_15_4, %iq2xs_code_15_5, %iq2xs_code_15_6, %iq2xs_code_15_7, %iq2xs_code_15_8, %iq2xs_code_15_9, %iq2xs_code_15_10, %iq2xs_code_15_11, %iq2xs_code_15_12, %iq2xs_code_15_13, %iq2xs_code_15_14, %iq2xs_code_15_15, %iq2xs_code_15_16, %iq2xs_code_15_17, %iq2xs_code_15_18, %iq2xs_code_15_19, %iq2xs_code_15_20, %iq2xs_code_15_21, %iq2xs_code_15_22, %iq2xs_code_15_23, %iq2xs_code_15_24, %iq2xs_code_15_25, %iq2xs_code_15_26, %iq2xs_code_15_27, %iq2xs_code_15_28, %iq2xs_code_15_29, %iq2xs_code_15_30, %iq2xs_code_15_31 : vector<32xf32> + %is_chunk1 = scalar.cmpi eq, %chunk_i32, %chunk_id1 : i32 + %sel1 = scf.select %is_chunk1, %iq2xs_code_1, %iq2xs_code_0 : vector<32xf32> + %is_chunk2 = scalar.cmpi eq, %chunk_i32, %chunk_id2 : i32 + %sel2 = scf.select %is_chunk2, %iq2xs_code_2, %sel1 : vector<32xf32> + %is_chunk3 = scalar.cmpi eq, %chunk_i32, %chunk_id3 : i32 + %sel3 = scf.select %is_chunk3, %iq2xs_code_3, %sel2 : vector<32xf32> + %is_chunk4 = scalar.cmpi eq, %chunk_i32, %chunk_id4 : i32 + %sel4 = scf.select %is_chunk4, %iq2xs_code_4, %sel3 : vector<32xf32> + %is_chunk5 = scalar.cmpi eq, %chunk_i32, %chunk_id5 : i32 + %sel5 = scf.select %is_chunk5, %iq2xs_code_5, %sel4 : vector<32xf32> + %is_chunk6 = scalar.cmpi eq, %chunk_i32, %chunk_id6 : i32 + %sel6 = scf.select %is_chunk6, %iq2xs_code_6, %sel5 : vector<32xf32> + %is_chunk7 = scalar.cmpi eq, %chunk_i32, %chunk_id7 : i32 + %sel7 = scf.select %is_chunk7, %iq2xs_code_7, %sel6 : vector<32xf32> + %is_chunk8 = scalar.cmpi eq, %chunk_i32, %chunk_id8 : i32 + %sel8 = scf.select %is_chunk8, %iq2xs_code_8, %sel7 : vector<32xf32> + %is_chunk9 = scalar.cmpi eq, %chunk_i32, %chunk_id9 : i32 + %sel9 = scf.select %is_chunk9, %iq2xs_code_9, %sel8 : vector<32xf32> + %is_chunk10 = scalar.cmpi eq, %chunk_i32, %chunk_id10 : i32 + %sel10 = scf.select %is_chunk10, %iq2xs_code_10, %sel9 : vector<32xf32> + %is_chunk11 = scalar.cmpi eq, %chunk_i32, %chunk_id11 : i32 + %sel11 = scf.select %is_chunk11, %iq2xs_code_11, %sel10 : vector<32xf32> + %is_chunk12 = scalar.cmpi eq, %chunk_i32, %chunk_id12 : i32 + %sel12 = scf.select %is_chunk12, %iq2xs_code_12, %sel11 : vector<32xf32> + %is_chunk13 = scalar.cmpi eq, %chunk_i32, %chunk_id13 : i32 + %sel13 = scf.select %is_chunk13, %iq2xs_code_13, %sel12 : vector<32xf32> + %is_chunk14 = scalar.cmpi eq, %chunk_i32, %chunk_id14 : i32 + %sel14 = scf.select %is_chunk14, %iq2xs_code_14, %sel13 : vector<32xf32> + %is_chunk15 = scalar.cmpi eq, %chunk_i32, %chunk_id15 : i32 + %sel15 = scf.select %is_chunk15, %iq2xs_code_15, %sel14 : vector<32xf32> + %v = vector.table.lookup %sel15[%codes] : vector<32xf32>, vector<1xi32> -> vector<1xf32> + %v_f32 = vector.extract %v[0] : vector<1xf32> -> f32 + %code = scalar.fptoui %v_f32 : f32 to i32 + func.return %code : i32 +} + +// ksigns_iq2xs[i] (ggml-common.h) is i with its parity as bit 7. +func.def inline @ggml_iq2_signs8(%signs7: i32) -> (i32) { + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c7_i32 = scalar.constant 7 : i32 + %p4 = scalar.shrui %signs7, %c4_i32 : i32 + %x4 = scalar.xori %signs7, %p4 : i32 + %p2 = scalar.shrui %x4, %c2_i32 : i32 + %x2 = scalar.xori %x4, %p2 : i32 + %p1 = scalar.shrui %x2, %c1_i32 : i32 + %x1 = scalar.xori %x2, %p1 : i32 + %parity = scalar.andi %x1, %c1_i32 : i32 + %high = scalar.shli %parity, %c7_i32 : i32 + %signs8 = scalar.ori %signs7, %high : i32 + func.return %signs8 : i32 +} + +// value j (0..7) of a grid code: level (8, 25, 43) with sign bit j of %signs8, times %scale +func.def inline @ggml_iq2_code_value_f32(%code: i32, %j: i32, %signs8: i32, %scale: f32) -> (f32) { + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %l8 = scalar.constant 8.0 : f32 + %l25 = scalar.constant 25.0 : f32 + %l43 = scalar.constant 43.0 : f32 + %two_j = scalar.addi %j, %j : i32 + %shifted = scalar.shrui %code, %two_j : i32 + %level = scalar.andi %shifted, %c3_i32 : i32 + %is0 = scalar.cmpi eq, %level, %c0_i32 : i32 + %is1 = scalar.cmpi eq, %level, %c1_i32 : i32 + %v01 = scf.select %is1, %l25, %l43 : f32 + %v = scf.select %is0, %l8, %v01 : f32 + %bit = scalar.shli %c1_i32, %j : i32 + %sign_mask = scalar.andi %signs8, %bit : i32 + %negative = scalar.cmpi ne, %sign_mask, %c0_i32 : i32 + %neg_v = scalar.negf %v : f32 + %signed = scf.select %negative, %neg_v, %v : f32 + %result = scalar.mulf %signed, %scale : f32 + func.return %result : f32 +} + +// four values 4 (p % 2) .. +3 of grid slot p / 2 +func.def inline @ggml_iq2_code_vector4(%code: i32, %half: i32, %signs8: i32, %scale: f32) -> (vector<4xf32>) { + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %j0 = scalar.muli %half, %c4_i32 : i32 + %j1 = scalar.addi %j0, %c1_i32 : i32 + %j2 = scalar.addi %j0, %c2_i32 : i32 + %j3 = scalar.addi %j0, %c3_i32 : i32 + %v0 = func.call @ggml_iq2_code_value_f32(%code, %j0, %signs8, %scale) : (i32, i32, i32, f32) -> (f32) + %v1 = func.call @ggml_iq2_code_value_f32(%code, %j1, %signs8, %scale) : (i32, i32, i32, f32) -> (f32) + %v2 = func.call @ggml_iq2_code_value_f32(%code, %j2, %signs8, %scale) : (i32, i32, i32, f32) -> (f32) + %v3 = func.call @ggml_iq2_code_value_f32(%code, %j3, %signs8, %scale) : (i32, i32, i32, f32) -> (f32) + %result = vector.from_elements %v0, %v1, %v2, %v3 : vector<4xf32> + func.return %result : vector<4xf32> +} + +// IQ2_XXS (66 bytes: d, qs[32] u16; dequantize_row_iq2_xxs). Group g (ib32) is 8 bytes at 2 + 8g: +// four grid indices, then a u32 with four 7-bit sign groups and a 4-bit scale on top: +// value = d * (0.5 + (aux >> 28)) * 0.25 * grid * sign. Packet p covers slot p / 2, values 4 (p % 2) .. +func.def inline @ggml_iq2xxs_f32_vector4(%weight: buffer, %row_byte_base: offset, %iq2_block: index, %iq2_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c7_i32 = scalar.constant 7 : i32 + %c16_i32 = scalar.constant 16 : i32 + %c28_i32 = scalar.constant 28 : i32 + %c127_i32 = scalar.constant 127 : i32 + %c65535_i32 = scalar.constant 65535 : i32 + %c05_f32 = scalar.constant 0.5 : f32 + %c025_f32 = scalar.constant 0.25 : f32 + %block_bytes = index.constant 66 : offset + %block_byte_add = index.scale %iq2_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %hv = buffer.view %weight[%block_byte_base] : buffer -> view<33xf16> + %wv = buffer.view %weight[%block_byte_base] : buffer -> view<33xi16> + %bv = buffer.view %weight[%block_byte_base] : buffer -> view<66xi8> + %g = index.assume %iq2_group [range(%iq2_group, 0, 7)] : index + %p = index.assume %packet [range(%packet, 0, 7)] : index + %slot = index.div %p, %c2 : index + %half = index.rem %p, %c2 : index + %g8 = index.mul %g, %c4 : index + %gw = index.mul %g, %c4 : index + %idx_byte0 = index.add %g8, %g8 : index + %idx_byte1 = index.add %idx_byte0, %c2 : index + %idx_byte = index.add %idx_byte1, %slot : index + %aux_w0_0 = index.add %gw, %c3 : index + %aux_w1_0 = index.add %aux_w0_0, %c1 : index + %d_f16 = view.load %hv[%c0] : view<33xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %gi_i8 = view.load %bv[%idx_byte] : view<66xi8> -> i8 + %gi = scalar.extui %gi_i8 : i8 to i32 + %w0_i16 = view.load %wv[%aux_w0_0] : view<33xi16> -> i16 + %w1_i16 = view.load %wv[%aux_w1_0] : view<33xi16> -> i16 + %w0 = scalar.extui %w0_i16 : i16 to i32 + %w1 = scalar.extui %w1_i16 : i16 to i32 + %w1s = scalar.shli %w1, %c16_i32 : i32 + %aux = scalar.ori %w0, %w1s : i32 + %sc4 = scalar.shrui %aux, %c28_i32 : i32 + %sc_f = scalar.uitofp %sc4 : i32 to f32 + %sc_plus = scalar.addf %sc_f, %c05_f32 : f32 + %ds = scalar.mulf %d, %sc_plus : f32 + %scale = scalar.mulf %ds, %c025_f32 : f32 + %slot_i32 = index.cast %slot : index to i32 + %sshift = scalar.muli %slot_i32, %c7_i32 : i32 + %sgrp0 = scalar.shrui %aux, %sshift : i32 + %signs7 = scalar.andi %sgrp0, %c127_i32 : i32 + %signs8 = func.call @ggml_iq2_signs8(%signs7) : (i32) -> (i32) + %code = func.call @ggml_iq2xxs_grid_code_i32(%gi) : (i32) -> (i32) + %half_i32 = index.cast %half : index to i32 + %result = func.call @ggml_iq2_code_vector4(%code, %half_i32, %signs8, %scale) : (i32, i32, i32, f32) -> (vector<4xf32>) + func.return %result : vector<4xf32> +} + +func.def inline @ggml_iq2xxs_f16_vector4(%weight: buffer, %row_byte_base: offset, %iq2_block: index, %iq2_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_iq2xxs_f32_vector4(%weight, %row_byte_base, %iq2_block, %iq2_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +// IQ2_XS (74 bytes: d, qs[32] u16, scales[8]; dequantize_row_iq2_xs). Slot l of group g is +// q = qs[4g + l]: grid index q & 511, 7 sign bits q >> 9; the scale nibble (l / 2) of scales[g]: +// value = d * (0.5 + nibble) * 0.25 * grid * sign. +func.def inline @ggml_iq2xs_f32_vector4(%weight: buffer, %row_byte_base: offset, %iq2_block: index, %iq2_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c66 = index.constant 66 : index + %c4_i32 = scalar.constant 4 : i32 + %c9_i32 = scalar.constant 9 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c511_i32 = scalar.constant 511 : i32 + %c127_i32 = scalar.constant 127 : i32 + %c05_f32 = scalar.constant 0.5 : f32 + %c025_f32 = scalar.constant 0.25 : f32 + %block_bytes = index.constant 74 : offset + %block_byte_add = index.scale %iq2_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %hv = buffer.view %weight[%block_byte_base] : buffer -> view<37xf16> + %wv = buffer.view %weight[%block_byte_base] : buffer -> view<37xi16> + %bv = buffer.view %weight[%block_byte_base] : buffer -> view<74xi8> + %g = index.assume %iq2_group [range(%iq2_group, 0, 7)] : index + %p = index.assume %packet [range(%packet, 0, 7)] : index + %slot = index.div %p, %c2 : index + %half = index.rem %p, %c2 : index + %g4 = index.mul %g, %c4 : index + %q_at0 = index.add %g4, %slot : index + %q_at = index.add %q_at0, %c1 : index + %sc_at = index.add %c66, %g : index + %d_f16 = view.load %hv[%c0] : view<37xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %q_i16 = view.load %wv[%q_at] : view<37xi16> -> i16 + %q = scalar.extui %q_i16 : i16 to i32 + %gi = scalar.andi %q, %c511_i32 : i32 + %signs7_0 = scalar.shrui %q, %c9_i32 : i32 + %signs7 = scalar.andi %signs7_0, %c127_i32 : i32 + %sc_i8 = view.load %bv[%sc_at] : view<74xi8> -> i8 + %sc = scalar.extui %sc_i8 : i8 to i32 + %nib_sel = index.div %slot, %c2 : index + %nib_sel_i32 = index.cast %nib_sel : index to i32 + %nib_shift = scalar.muli %nib_sel_i32, %c4_i32 : i32 + %nib0 = scalar.shrui %sc, %nib_shift : i32 + %nib = scalar.andi %nib0, %c15_i32 : i32 + %nib_f = scalar.uitofp %nib : i32 to f32 + %nib_plus = scalar.addf %nib_f, %c05_f32 : f32 + %ds = scalar.mulf %d, %nib_plus : f32 + %scale = scalar.mulf %ds, %c025_f32 : f32 + %signs8 = func.call @ggml_iq2_signs8(%signs7) : (i32) -> (i32) + %code = func.call @ggml_iq2xs_grid_code_i32(%gi) : (i32) -> (i32) + %half_i32 = index.cast %half : index to i32 + %result = func.call @ggml_iq2_code_vector4(%code, %half_i32, %signs8, %scale) : (i32, i32, i32, f32) -> (vector<4xf32>) + func.return %result : vector<4xf32> +} + +func.def inline @ggml_iq2xs_f16_vector4(%weight: buffer, %row_byte_base: offset, %iq2_block: index, %iq2_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_iq2xs_f32_vector4(%weight, %row_byte_base, %iq2_block, %iq2_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} +// iq1s_code: 2048 grid entries as 16-bit codes (2 bits per value: value + 1) +func.def inline @ggml_iq1s_grid_code_i32(%grid_index: i32) -> (i32) { + %c5_i32 = scalar.constant 5 : i32 + %c31_i32 = scalar.constant 31 : i32 + %chunk_id1 = scalar.constant 1 : i32 + %chunk_id2 = scalar.constant 2 : i32 + %chunk_id3 = scalar.constant 3 : i32 + %chunk_id4 = scalar.constant 4 : i32 + %chunk_id5 = scalar.constant 5 : i32 + %chunk_id6 = scalar.constant 6 : i32 + %chunk_id7 = scalar.constant 7 : i32 + %chunk_id8 = scalar.constant 8 : i32 + %chunk_id9 = scalar.constant 9 : i32 + %chunk_id10 = scalar.constant 10 : i32 + %chunk_id11 = scalar.constant 11 : i32 + %chunk_id12 = scalar.constant 12 : i32 + %chunk_id13 = scalar.constant 13 : i32 + %chunk_id14 = scalar.constant 14 : i32 + %chunk_id15 = scalar.constant 15 : i32 + %chunk_id16 = scalar.constant 16 : i32 + %chunk_id17 = scalar.constant 17 : i32 + %chunk_id18 = scalar.constant 18 : i32 + %chunk_id19 = scalar.constant 19 : i32 + %chunk_id20 = scalar.constant 20 : i32 + %chunk_id21 = scalar.constant 21 : i32 + %chunk_id22 = scalar.constant 22 : i32 + %chunk_id23 = scalar.constant 23 : i32 + %chunk_id24 = scalar.constant 24 : i32 + %chunk_id25 = scalar.constant 25 : i32 + %chunk_id26 = scalar.constant 26 : i32 + %chunk_id27 = scalar.constant 27 : i32 + %chunk_id28 = scalar.constant 28 : i32 + %chunk_id29 = scalar.constant 29 : i32 + %chunk_id30 = scalar.constant 30 : i32 + %chunk_id31 = scalar.constant 31 : i32 + %chunk_id32 = scalar.constant 32 : i32 + %chunk_id33 = scalar.constant 33 : i32 + %chunk_id34 = scalar.constant 34 : i32 + %chunk_id35 = scalar.constant 35 : i32 + %chunk_id36 = scalar.constant 36 : i32 + %chunk_id37 = scalar.constant 37 : i32 + %chunk_id38 = scalar.constant 38 : i32 + %chunk_id39 = scalar.constant 39 : i32 + %chunk_id40 = scalar.constant 40 : i32 + %chunk_id41 = scalar.constant 41 : i32 + %chunk_id42 = scalar.constant 42 : i32 + %chunk_id43 = scalar.constant 43 : i32 + %chunk_id44 = scalar.constant 44 : i32 + %chunk_id45 = scalar.constant 45 : i32 + %chunk_id46 = scalar.constant 46 : i32 + %chunk_id47 = scalar.constant 47 : i32 + %chunk_id48 = scalar.constant 48 : i32 + %chunk_id49 = scalar.constant 49 : i32 + %chunk_id50 = scalar.constant 50 : i32 + %chunk_id51 = scalar.constant 51 : i32 + %chunk_id52 = scalar.constant 52 : i32 + %chunk_id53 = scalar.constant 53 : i32 + %chunk_id54 = scalar.constant 54 : i32 + %chunk_id55 = scalar.constant 55 : i32 + %chunk_id56 = scalar.constant 56 : i32 + %chunk_id57 = scalar.constant 57 : i32 + %chunk_id58 = scalar.constant 58 : i32 + %chunk_id59 = scalar.constant 59 : i32 + %chunk_id60 = scalar.constant 60 : i32 + %chunk_id61 = scalar.constant 61 : i32 + %chunk_id62 = scalar.constant 62 : i32 + %chunk_id63 = scalar.constant 63 : i32 + %chunk_i32 = scalar.shrui %grid_index, %c5_i32 : i32 + %lane_i32 = scalar.andi %grid_index, %c31_i32 : i32 + %codes = vector.from_elements %lane_i32 : vector<1xi32> + %iq1s_code_0_0 = scalar.constant 0.0 : f32 + %iq1s_code_0_1 = scalar.constant 2.0 : f32 + %iq1s_code_0_2 = scalar.constant 5.0 : f32 + %iq1s_code_0_3 = scalar.constant 8.0 : f32 + %iq1s_code_0_4 = scalar.constant 10.0 : f32 + %iq1s_code_0_5 = scalar.constant 17.0 : f32 + %iq1s_code_0_6 = scalar.constant 21.0 : f32 + %iq1s_code_0_7 = scalar.constant 32.0 : f32 + %iq1s_code_0_8 = scalar.constant 34.0 : f32 + %iq1s_code_0_9 = scalar.constant 40.0 : f32 + %iq1s_code_0_10 = scalar.constant 42.0 : f32 + %iq1s_code_0_11 = scalar.constant 69.0 : f32 + %iq1s_code_0_12 = scalar.constant 81.0 : f32 + %iq1s_code_0_13 = scalar.constant 84.0 : f32 + %iq1s_code_0_14 = scalar.constant 86.0 : f32 + %iq1s_code_0_15 = scalar.constant 101.0 : f32 + %iq1s_code_0_16 = scalar.constant 128.0 : f32 + %iq1s_code_0_17 = scalar.constant 130.0 : f32 + %iq1s_code_0_18 = scalar.constant 136.0 : f32 + %iq1s_code_0_19 = scalar.constant 138.0 : f32 + %iq1s_code_0_20 = scalar.constant 149.0 : f32 + %iq1s_code_0_21 = scalar.constant 160.0 : f32 + %iq1s_code_0_22 = scalar.constant 162.0 : f32 + %iq1s_code_0_23 = scalar.constant 168.0 : f32 + %iq1s_code_0_24 = scalar.constant 170.0 : f32 + %iq1s_code_0_25 = scalar.constant 260.0 : f32 + %iq1s_code_0_26 = scalar.constant 261.0 : f32 + %iq1s_code_0_27 = scalar.constant 273.0 : f32 + %iq1s_code_0_28 = scalar.constant 276.0 : f32 + %iq1s_code_0_29 = scalar.constant 278.0 : f32 + %iq1s_code_0_30 = scalar.constant 281.0 : f32 + %iq1s_code_0_31 = scalar.constant 282.0 : f32 + %iq1s_code_0 = vector.from_elements %iq1s_code_0_0, %iq1s_code_0_1, %iq1s_code_0_2, %iq1s_code_0_3, %iq1s_code_0_4, %iq1s_code_0_5, %iq1s_code_0_6, %iq1s_code_0_7, %iq1s_code_0_8, %iq1s_code_0_9, %iq1s_code_0_10, %iq1s_code_0_11, %iq1s_code_0_12, %iq1s_code_0_13, %iq1s_code_0_14, %iq1s_code_0_15, %iq1s_code_0_16, %iq1s_code_0_17, %iq1s_code_0_18, %iq1s_code_0_19, %iq1s_code_0_20, %iq1s_code_0_21, %iq1s_code_0_22, %iq1s_code_0_23, %iq1s_code_0_24, %iq1s_code_0_25, %iq1s_code_0_26, %iq1s_code_0_27, %iq1s_code_0_28, %iq1s_code_0_29, %iq1s_code_0_30, %iq1s_code_0_31 : vector<32xf32> + %iq1s_code_1_0 = scalar.constant 293.0 : f32 + %iq1s_code_1_1 = scalar.constant 321.0 : f32 + %iq1s_code_1_2 = scalar.constant 326.0 : f32 + %iq1s_code_1_3 = scalar.constant 329.0 : f32 + %iq1s_code_1_4 = scalar.constant 338.0 : f32 + %iq1s_code_1_5 = scalar.constant 341.0 : f32 + %iq1s_code_1_6 = scalar.constant 346.0 : f32 + %iq1s_code_1_7 = scalar.constant 353.0 : f32 + %iq1s_code_1_8 = scalar.constant 356.0 : f32 + %iq1s_code_1_9 = scalar.constant 358.0 : f32 + %iq1s_code_1_10 = scalar.constant 360.0 : f32 + %iq1s_code_1_11 = scalar.constant 389.0 : f32 + %iq1s_code_1_12 = scalar.constant 401.0 : f32 + %iq1s_code_1_13 = scalar.constant 404.0 : f32 + %iq1s_code_1_14 = scalar.constant 406.0 : f32 + %iq1s_code_1_15 = scalar.constant 421.0 : f32 + %iq1s_code_1_16 = scalar.constant 512.0 : f32 + %iq1s_code_1_17 = scalar.constant 514.0 : f32 + %iq1s_code_1_18 = scalar.constant 520.0 : f32 + %iq1s_code_1_19 = scalar.constant 522.0 : f32 + %iq1s_code_1_20 = scalar.constant 533.0 : f32 + %iq1s_code_1_21 = scalar.constant 544.0 : f32 + %iq1s_code_1_22 = scalar.constant 546.0 : f32 + %iq1s_code_1_23 = scalar.constant 552.0 : f32 + %iq1s_code_1_24 = scalar.constant 554.0 : f32 + %iq1s_code_1_25 = scalar.constant 581.0 : f32 + %iq1s_code_1_26 = scalar.constant 593.0 : f32 + %iq1s_code_1_27 = scalar.constant 601.0 : f32 + %iq1s_code_1_28 = scalar.constant 612.0 : f32 + %iq1s_code_1_29 = scalar.constant 617.0 : f32 + %iq1s_code_1_30 = scalar.constant 640.0 : f32 + %iq1s_code_1_31 = scalar.constant 642.0 : f32 + %iq1s_code_1 = vector.from_elements %iq1s_code_1_0, %iq1s_code_1_1, %iq1s_code_1_2, %iq1s_code_1_3, %iq1s_code_1_4, %iq1s_code_1_5, %iq1s_code_1_6, %iq1s_code_1_7, %iq1s_code_1_8, %iq1s_code_1_9, %iq1s_code_1_10, %iq1s_code_1_11, %iq1s_code_1_12, %iq1s_code_1_13, %iq1s_code_1_14, %iq1s_code_1_15, %iq1s_code_1_16, %iq1s_code_1_17, %iq1s_code_1_18, %iq1s_code_1_19, %iq1s_code_1_20, %iq1s_code_1_21, %iq1s_code_1_22, %iq1s_code_1_23, %iq1s_code_1_24, %iq1s_code_1_25, %iq1s_code_1_26, %iq1s_code_1_27, %iq1s_code_1_28, %iq1s_code_1_29, %iq1s_code_1_30, %iq1s_code_1_31 : vector<32xf32> + %iq1s_code_2_0 = scalar.constant 648.0 : f32 + %iq1s_code_2_1 = scalar.constant 650.0 : f32 + %iq1s_code_2_2 = scalar.constant 657.0 : f32 + %iq1s_code_2_3 = scalar.constant 661.0 : f32 + %iq1s_code_2_4 = scalar.constant 665.0 : f32 + %iq1s_code_2_5 = scalar.constant 672.0 : f32 + %iq1s_code_2_6 = scalar.constant 674.0 : f32 + %iq1s_code_2_7 = scalar.constant 680.0 : f32 + %iq1s_code_2_8 = scalar.constant 682.0 : f32 + %iq1s_code_2_9 = scalar.constant 1041.0 : f32 + %iq1s_code_2_10 = scalar.constant 1044.0 : f32 + %iq1s_code_2_11 = scalar.constant 1046.0 : f32 + %iq1s_code_2_12 = scalar.constant 1061.0 : f32 + %iq1s_code_2_13 = scalar.constant 1089.0 : f32 + %iq1s_code_2_14 = scalar.constant 1097.0 : f32 + %iq1s_code_2_15 = scalar.constant 1109.0 : f32 + %iq1s_code_2_16 = scalar.constant 1114.0 : f32 + %iq1s_code_2_17 = scalar.constant 1124.0 : f32 + %iq1s_code_2_18 = scalar.constant 1125.0 : f32 + %iq1s_code_2_19 = scalar.constant 1169.0 : f32 + %iq1s_code_2_20 = scalar.constant 1177.0 : f32 + %iq1s_code_2_21 = scalar.constant 1189.0 : f32 + %iq1s_code_2_22 = scalar.constant 1281.0 : f32 + %iq1s_code_2_23 = scalar.constant 1284.0 : f32 + %iq1s_code_2_24 = scalar.constant 1285.0 : f32 + %iq1s_code_2_25 = scalar.constant 1286.0 : f32 + %iq1s_code_2_26 = scalar.constant 1301.0 : f32 + %iq1s_code_2_27 = scalar.constant 1304.0 : f32 + %iq1s_code_2_28 = scalar.constant 1306.0 : f32 + %iq1s_code_2_29 = scalar.constant 1321.0 : f32 + %iq1s_code_2_30 = scalar.constant 1344.0 : f32 + %iq1s_code_2_31 = scalar.constant 1349.0 : f32 + %iq1s_code_2 = vector.from_elements %iq1s_code_2_0, %iq1s_code_2_1, %iq1s_code_2_2, %iq1s_code_2_3, %iq1s_code_2_4, %iq1s_code_2_5, %iq1s_code_2_6, %iq1s_code_2_7, %iq1s_code_2_8, %iq1s_code_2_9, %iq1s_code_2_10, %iq1s_code_2_11, %iq1s_code_2_12, %iq1s_code_2_13, %iq1s_code_2_14, %iq1s_code_2_15, %iq1s_code_2_16, %iq1s_code_2_17, %iq1s_code_2_18, %iq1s_code_2_19, %iq1s_code_2_20, %iq1s_code_2_21, %iq1s_code_2_22, %iq1s_code_2_23, %iq1s_code_2_24, %iq1s_code_2_25, %iq1s_code_2_26, %iq1s_code_2_27, %iq1s_code_2_28, %iq1s_code_2_29, %iq1s_code_2_30, %iq1s_code_2_31 : vector<32xf32> + %iq1s_code_3_0 = scalar.constant 1354.0 : f32 + %iq1s_code_3_1 = scalar.constant 1360.0 : f32 + %iq1s_code_3_2 = scalar.constant 1361.0 : f32 + %iq1s_code_3_3 = scalar.constant 1364.0 : f32 + %iq1s_code_3_4 = scalar.constant 1365.0 : f32 + %iq1s_code_3_5 = scalar.constant 1366.0 : f32 + %iq1s_code_3_6 = scalar.constant 1369.0 : f32 + %iq1s_code_3_7 = scalar.constant 1376.0 : f32 + %iq1s_code_3_8 = scalar.constant 1378.0 : f32 + %iq1s_code_3_9 = scalar.constant 1381.0 : f32 + %iq1s_code_3_10 = scalar.constant 1384.0 : f32 + %iq1s_code_3_11 = scalar.constant 1386.0 : f32 + %iq1s_code_3_12 = scalar.constant 1409.0 : f32 + %iq1s_code_3_13 = scalar.constant 1425.0 : f32 + %iq1s_code_3_14 = scalar.constant 1429.0 : f32 + %iq1s_code_3_15 = scalar.constant 1432.0 : f32 + %iq1s_code_3_16 = scalar.constant 1434.0 : f32 + %iq1s_code_3_17 = scalar.constant 1441.0 : f32 + %iq1s_code_3_18 = scalar.constant 1444.0 : f32 + %iq1s_code_3_19 = scalar.constant 1445.0 : f32 + %iq1s_code_3_20 = scalar.constant 1446.0 : f32 + %iq1s_code_3_21 = scalar.constant 1449.0 : f32 + %iq1s_code_3_22 = scalar.constant 1556.0 : f32 + %iq1s_code_3_23 = scalar.constant 1561.0 : f32 + %iq1s_code_3_24 = scalar.constant 1601.0 : f32 + %iq1s_code_3_25 = scalar.constant 1604.0 : f32 + %iq1s_code_3_26 = scalar.constant 1616.0 : f32 + %iq1s_code_3_27 = scalar.constant 1618.0 : f32 + %iq1s_code_3_28 = scalar.constant 1621.0 : f32 + %iq1s_code_3_29 = scalar.constant 1624.0 : f32 + %iq1s_code_3_30 = scalar.constant 1632.0 : f32 + %iq1s_code_3_31 = scalar.constant 1633.0 : f32 + %iq1s_code_3 = vector.from_elements %iq1s_code_3_0, %iq1s_code_3_1, %iq1s_code_3_2, %iq1s_code_3_3, %iq1s_code_3_4, %iq1s_code_3_5, %iq1s_code_3_6, %iq1s_code_3_7, %iq1s_code_3_8, %iq1s_code_3_9, %iq1s_code_3_10, %iq1s_code_3_11, %iq1s_code_3_12, %iq1s_code_3_13, %iq1s_code_3_14, %iq1s_code_3_15, %iq1s_code_3_16, %iq1s_code_3_17, %iq1s_code_3_18, %iq1s_code_3_19, %iq1s_code_3_20, %iq1s_code_3_21, %iq1s_code_3_22, %iq1s_code_3_23, %iq1s_code_3_24, %iq1s_code_3_25, %iq1s_code_3_26, %iq1s_code_3_27, %iq1s_code_3_28, %iq1s_code_3_29, %iq1s_code_3_30, %iq1s_code_3_31 : vector<32xf32> + %iq1s_code_4_0 = scalar.constant 1638.0 : f32 + %iq1s_code_4_1 = scalar.constant 1641.0 : f32 + %iq1s_code_4_2 = scalar.constant 1669.0 : f32 + %iq1s_code_4_3 = scalar.constant 1681.0 : f32 + %iq1s_code_4_4 = scalar.constant 1684.0 : f32 + %iq1s_code_4_5 = scalar.constant 1689.0 : f32 + %iq1s_code_4_6 = scalar.constant 2048.0 : f32 + %iq1s_code_4_7 = scalar.constant 2050.0 : f32 + %iq1s_code_4_8 = scalar.constant 2056.0 : f32 + %iq1s_code_4_9 = scalar.constant 2058.0 : f32 + %iq1s_code_4_10 = scalar.constant 2069.0 : f32 + %iq1s_code_4_11 = scalar.constant 2080.0 : f32 + %iq1s_code_4_12 = scalar.constant 2082.0 : f32 + %iq1s_code_4_13 = scalar.constant 2088.0 : f32 + %iq1s_code_4_14 = scalar.constant 2090.0 : f32 + %iq1s_code_4_15 = scalar.constant 2117.0 : f32 + %iq1s_code_4_16 = scalar.constant 2129.0 : f32 + %iq1s_code_4_17 = scalar.constant 2134.0 : f32 + %iq1s_code_4_18 = scalar.constant 2149.0 : f32 + %iq1s_code_4_19 = scalar.constant 2176.0 : f32 + %iq1s_code_4_20 = scalar.constant 2178.0 : f32 + %iq1s_code_4_21 = scalar.constant 2184.0 : f32 + %iq1s_code_4_22 = scalar.constant 2186.0 : f32 + %iq1s_code_4_23 = scalar.constant 2197.0 : f32 + %iq1s_code_4_24 = scalar.constant 2208.0 : f32 + %iq1s_code_4_25 = scalar.constant 2210.0 : f32 + %iq1s_code_4_26 = scalar.constant 2216.0 : f32 + %iq1s_code_4_27 = scalar.constant 2218.0 : f32 + %iq1s_code_4_28 = scalar.constant 2309.0 : f32 + %iq1s_code_4_29 = scalar.constant 2321.0 : f32 + %iq1s_code_4_30 = scalar.constant 2324.0 : f32 + %iq1s_code_4_31 = scalar.constant 2329.0 : f32 + %iq1s_code_4 = vector.from_elements %iq1s_code_4_0, %iq1s_code_4_1, %iq1s_code_4_2, %iq1s_code_4_3, %iq1s_code_4_4, %iq1s_code_4_5, %iq1s_code_4_6, %iq1s_code_4_7, %iq1s_code_4_8, %iq1s_code_4_9, %iq1s_code_4_10, %iq1s_code_4_11, %iq1s_code_4_12, %iq1s_code_4_13, %iq1s_code_4_14, %iq1s_code_4_15, %iq1s_code_4_16, %iq1s_code_4_17, %iq1s_code_4_18, %iq1s_code_4_19, %iq1s_code_4_20, %iq1s_code_4_21, %iq1s_code_4_22, %iq1s_code_4_23, %iq1s_code_4_24, %iq1s_code_4_25, %iq1s_code_4_26, %iq1s_code_4_27, %iq1s_code_4_28, %iq1s_code_4_29, %iq1s_code_4_30, %iq1s_code_4_31 : vector<32xf32> + %iq1s_code_5_0 = scalar.constant 2340.0 : f32 + %iq1s_code_5_1 = scalar.constant 2341.0 : f32 + %iq1s_code_5_2 = scalar.constant 2369.0 : f32 + %iq1s_code_5_3 = scalar.constant 2384.0 : f32 + %iq1s_code_5_4 = scalar.constant 2385.0 : f32 + %iq1s_code_5_5 = scalar.constant 2389.0 : f32 + %iq1s_code_5_6 = scalar.constant 2401.0 : f32 + %iq1s_code_5_7 = scalar.constant 2404.0 : f32 + %iq1s_code_5_8 = scalar.constant 2409.0 : f32 + %iq1s_code_5_9 = scalar.constant 2449.0 : f32 + %iq1s_code_5_10 = scalar.constant 2452.0 : f32 + %iq1s_code_5_11 = scalar.constant 2454.0 : f32 + %iq1s_code_5_12 = scalar.constant 2457.0 : f32 + %iq1s_code_5_13 = scalar.constant 2469.0 : f32 + %iq1s_code_5_14 = scalar.constant 2560.0 : f32 + %iq1s_code_5_15 = scalar.constant 2562.0 : f32 + %iq1s_code_5_16 = scalar.constant 2568.0 : f32 + %iq1s_code_5_17 = scalar.constant 2570.0 : f32 + %iq1s_code_5_18 = scalar.constant 2581.0 : f32 + %iq1s_code_5_19 = scalar.constant 2592.0 : f32 + %iq1s_code_5_20 = scalar.constant 2594.0 : f32 + %iq1s_code_5_21 = scalar.constant 2600.0 : f32 + %iq1s_code_5_22 = scalar.constant 2602.0 : f32 + %iq1s_code_5_23 = scalar.constant 2629.0 : f32 + %iq1s_code_5_24 = scalar.constant 2641.0 : f32 + %iq1s_code_5_25 = scalar.constant 2649.0 : f32 + %iq1s_code_5_26 = scalar.constant 2657.0 : f32 + %iq1s_code_5_27 = scalar.constant 2661.0 : f32 + %iq1s_code_5_28 = scalar.constant 2688.0 : f32 + %iq1s_code_5_29 = scalar.constant 2690.0 : f32 + %iq1s_code_5_30 = scalar.constant 2693.0 : f32 + %iq1s_code_5_31 = scalar.constant 2696.0 : f32 + %iq1s_code_5 = vector.from_elements %iq1s_code_5_0, %iq1s_code_5_1, %iq1s_code_5_2, %iq1s_code_5_3, %iq1s_code_5_4, %iq1s_code_5_5, %iq1s_code_5_6, %iq1s_code_5_7, %iq1s_code_5_8, %iq1s_code_5_9, %iq1s_code_5_10, %iq1s_code_5_11, %iq1s_code_5_12, %iq1s_code_5_13, %iq1s_code_5_14, %iq1s_code_5_15, %iq1s_code_5_16, %iq1s_code_5_17, %iq1s_code_5_18, %iq1s_code_5_19, %iq1s_code_5_20, %iq1s_code_5_21, %iq1s_code_5_22, %iq1s_code_5_23, %iq1s_code_5_24, %iq1s_code_5_25, %iq1s_code_5_26, %iq1s_code_5_27, %iq1s_code_5_28, %iq1s_code_5_29, %iq1s_code_5_30, %iq1s_code_5_31 : vector<32xf32> + %iq1s_code_6_0 = scalar.constant 2698.0 : f32 + %iq1s_code_6_1 = scalar.constant 2709.0 : f32 + %iq1s_code_6_2 = scalar.constant 2720.0 : f32 + %iq1s_code_6_3 = scalar.constant 2722.0 : f32 + %iq1s_code_6_4 = scalar.constant 2728.0 : f32 + %iq1s_code_6_5 = scalar.constant 2730.0 : f32 + %iq1s_code_6_6 = scalar.constant 4112.0 : f32 + %iq1s_code_6_7 = scalar.constant 4113.0 : f32 + %iq1s_code_6_8 = scalar.constant 4116.0 : f32 + %iq1s_code_6_9 = scalar.constant 4121.0 : f32 + %iq1s_code_6_10 = scalar.constant 4132.0 : f32 + %iq1s_code_6_11 = scalar.constant 4133.0 : f32 + %iq1s_code_6_12 = scalar.constant 4161.0 : f32 + %iq1s_code_6_13 = scalar.constant 4164.0 : f32 + %iq1s_code_6_14 = scalar.constant 4176.0 : f32 + %iq1s_code_6_15 = scalar.constant 4181.0 : f32 + %iq1s_code_6_16 = scalar.constant 4184.0 : f32 + %iq1s_code_6_17 = scalar.constant 4193.0 : f32 + %iq1s_code_6_18 = scalar.constant 4196.0 : f32 + %iq1s_code_6_19 = scalar.constant 4197.0 : f32 + %iq1s_code_6_20 = scalar.constant 4201.0 : f32 + %iq1s_code_6_21 = scalar.constant 4241.0 : f32 + %iq1s_code_6_22 = scalar.constant 4244.0 : f32 + %iq1s_code_6_23 = scalar.constant 4246.0 : f32 + %iq1s_code_6_24 = scalar.constant 4257.0 : f32 + %iq1s_code_6_25 = scalar.constant 4261.0 : f32 + %iq1s_code_6_26 = scalar.constant 4353.0 : f32 + %iq1s_code_6_27 = scalar.constant 4356.0 : f32 + %iq1s_code_6_28 = scalar.constant 4358.0 : f32 + %iq1s_code_6_29 = scalar.constant 4361.0 : f32 + %iq1s_code_6_30 = scalar.constant 4368.0 : f32 + %iq1s_code_6_31 = scalar.constant 4370.0 : f32 + %iq1s_code_6 = vector.from_elements %iq1s_code_6_0, %iq1s_code_6_1, %iq1s_code_6_2, %iq1s_code_6_3, %iq1s_code_6_4, %iq1s_code_6_5, %iq1s_code_6_6, %iq1s_code_6_7, %iq1s_code_6_8, %iq1s_code_6_9, %iq1s_code_6_10, %iq1s_code_6_11, %iq1s_code_6_12, %iq1s_code_6_13, %iq1s_code_6_14, %iq1s_code_6_15, %iq1s_code_6_16, %iq1s_code_6_17, %iq1s_code_6_18, %iq1s_code_6_19, %iq1s_code_6_20, %iq1s_code_6_21, %iq1s_code_6_22, %iq1s_code_6_23, %iq1s_code_6_24, %iq1s_code_6_25, %iq1s_code_6_26, %iq1s_code_6_27, %iq1s_code_6_28, %iq1s_code_6_29, %iq1s_code_6_30, %iq1s_code_6_31 : vector<32xf32> + %iq1s_code_7_0 = scalar.constant 4373.0 : f32 + %iq1s_code_7_1 = scalar.constant 4376.0 : f32 + %iq1s_code_7_2 = scalar.constant 4385.0 : f32 + %iq1s_code_7_3 = scalar.constant 4388.0 : f32 + %iq1s_code_7_4 = scalar.constant 4393.0 : f32 + %iq1s_code_7_5 = scalar.constant 4421.0 : f32 + %iq1s_code_7_6 = scalar.constant 4426.0 : f32 + %iq1s_code_7_7 = scalar.constant 4432.0 : f32 + %iq1s_code_7_8 = scalar.constant 4433.0 : f32 + %iq1s_code_7_9 = scalar.constant 4434.0 : f32 + %iq1s_code_7_10 = scalar.constant 4436.0 : f32 + %iq1s_code_7_11 = scalar.constant 4437.0 : f32 + %iq1s_code_7_12 = scalar.constant 4438.0 : f32 + %iq1s_code_7_13 = scalar.constant 4441.0 : f32 + %iq1s_code_7_14 = scalar.constant 4448.0 : f32 + %iq1s_code_7_15 = scalar.constant 4453.0 : f32 + %iq1s_code_7_16 = scalar.constant 4484.0 : f32 + %iq1s_code_7_17 = scalar.constant 4498.0 : f32 + %iq1s_code_7_18 = scalar.constant 4501.0 : f32 + %iq1s_code_7_19 = scalar.constant 4513.0 : f32 + %iq1s_code_7_20 = scalar.constant 4516.0 : f32 + %iq1s_code_7_21 = scalar.constant 4625.0 : f32 + %iq1s_code_7_22 = scalar.constant 4628.0 : f32 + %iq1s_code_7_23 = scalar.constant 4630.0 : f32 + %iq1s_code_7_24 = scalar.constant 4645.0 : f32 + %iq1s_code_7_25 = scalar.constant 4672.0 : f32 + %iq1s_code_7_26 = scalar.constant 4678.0 : f32 + %iq1s_code_7_27 = scalar.constant 4681.0 : f32 + %iq1s_code_7_28 = scalar.constant 4690.0 : f32 + %iq1s_code_7_29 = scalar.constant 4693.0 : f32 + %iq1s_code_7_30 = scalar.constant 4696.0 : f32 + %iq1s_code_7_31 = scalar.constant 4698.0 : f32 + %iq1s_code_7 = vector.from_elements %iq1s_code_7_0, %iq1s_code_7_1, %iq1s_code_7_2, %iq1s_code_7_3, %iq1s_code_7_4, %iq1s_code_7_5, %iq1s_code_7_6, %iq1s_code_7_7, %iq1s_code_7_8, %iq1s_code_7_9, %iq1s_code_7_10, %iq1s_code_7_11, %iq1s_code_7_12, %iq1s_code_7_13, %iq1s_code_7_14, %iq1s_code_7_15, %iq1s_code_7_16, %iq1s_code_7_17, %iq1s_code_7_18, %iq1s_code_7_19, %iq1s_code_7_20, %iq1s_code_7_21, %iq1s_code_7_22, %iq1s_code_7_23, %iq1s_code_7_24, %iq1s_code_7_25, %iq1s_code_7_26, %iq1s_code_7_27, %iq1s_code_7_28, %iq1s_code_7_29, %iq1s_code_7_30, %iq1s_code_7_31 : vector<32xf32> + %iq1s_code_8_0 = scalar.constant 4708.0 : f32 + %iq1s_code_8_1 = scalar.constant 4710.0 : f32 + %iq1s_code_8_2 = scalar.constant 4741.0 : f32 + %iq1s_code_8_3 = scalar.constant 4753.0 : f32 + %iq1s_code_8_4 = scalar.constant 4756.0 : f32 + %iq1s_code_8_5 = scalar.constant 4758.0 : f32 + %iq1s_code_8_6 = scalar.constant 4773.0 : f32 + %iq1s_code_8_7 = scalar.constant 5121.0 : f32 + %iq1s_code_8_8 = scalar.constant 5126.0 : f32 + %iq1s_code_8_9 = scalar.constant 5129.0 : f32 + %iq1s_code_8_10 = scalar.constant 5140.0 : f32 + %iq1s_code_8_11 = scalar.constant 5141.0 : f32 + %iq1s_code_8_12 = scalar.constant 5144.0 : f32 + %iq1s_code_8_13 = scalar.constant 5145.0 : f32 + %iq1s_code_8_14 = scalar.constant 5153.0 : f32 + %iq1s_code_8_15 = scalar.constant 5158.0 : f32 + %iq1s_code_8_16 = scalar.constant 5185.0 : f32 + %iq1s_code_8_17 = scalar.constant 5189.0 : f32 + %iq1s_code_8_18 = scalar.constant 5190.0 : f32 + %iq1s_code_8_19 = scalar.constant 5192.0 : f32 + %iq1s_code_8_20 = scalar.constant 5194.0 : f32 + %iq1s_code_8_21 = scalar.constant 5201.0 : f32 + %iq1s_code_8_22 = scalar.constant 5204.0 : f32 + %iq1s_code_8_23 = scalar.constant 5205.0 : f32 + %iq1s_code_8_24 = scalar.constant 5206.0 : f32 + %iq1s_code_8_25 = scalar.constant 5209.0 : f32 + %iq1s_code_8_26 = scalar.constant 5218.0 : f32 + %iq1s_code_8_27 = scalar.constant 5221.0 : f32 + %iq1s_code_8_28 = scalar.constant 5224.0 : f32 + %iq1s_code_8_29 = scalar.constant 5252.0 : f32 + %iq1s_code_8_30 = scalar.constant 5257.0 : f32 + %iq1s_code_8_31 = scalar.constant 5264.0 : f32 + %iq1s_code_8 = vector.from_elements %iq1s_code_8_0, %iq1s_code_8_1, %iq1s_code_8_2, %iq1s_code_8_3, %iq1s_code_8_4, %iq1s_code_8_5, %iq1s_code_8_6, %iq1s_code_8_7, %iq1s_code_8_8, %iq1s_code_8_9, %iq1s_code_8_10, %iq1s_code_8_11, %iq1s_code_8_12, %iq1s_code_8_13, %iq1s_code_8_14, %iq1s_code_8_15, %iq1s_code_8_16, %iq1s_code_8_17, %iq1s_code_8_18, %iq1s_code_8_19, %iq1s_code_8_20, %iq1s_code_8_21, %iq1s_code_8_22, %iq1s_code_8_23, %iq1s_code_8_24, %iq1s_code_8_25, %iq1s_code_8_26, %iq1s_code_8_27, %iq1s_code_8_28, %iq1s_code_8_29, %iq1s_code_8_30, %iq1s_code_8_31 : vector<32xf32> + %iq1s_code_9_0 = scalar.constant 5268.0 : f32 + %iq1s_code_9_1 = scalar.constant 5269.0 : f32 + %iq1s_code_9_2 = scalar.constant 5272.0 : f32 + %iq1s_code_9_3 = scalar.constant 5273.0 : f32 + %iq1s_code_9_4 = scalar.constant 5274.0 : f32 + %iq1s_code_9_5 = scalar.constant 5281.0 : f32 + %iq1s_code_9_6 = scalar.constant 5284.0 : f32 + %iq1s_code_9_7 = scalar.constant 5285.0 : f32 + %iq1s_code_9_8 = scalar.constant 5289.0 : f32 + %iq1s_code_9_9 = scalar.constant 5378.0 : f32 + %iq1s_code_9_10 = scalar.constant 5381.0 : f32 + %iq1s_code_9_11 = scalar.constant 5386.0 : f32 + %iq1s_code_9_12 = scalar.constant 5393.0 : f32 + %iq1s_code_9_13 = scalar.constant 5396.0 : f32 + %iq1s_code_9_14 = scalar.constant 5397.0 : f32 + %iq1s_code_9_15 = scalar.constant 5398.0 : f32 + %iq1s_code_9_16 = scalar.constant 5401.0 : f32 + %iq1s_code_9_17 = scalar.constant 5408.0 : f32 + %iq1s_code_9_18 = scalar.constant 5410.0 : f32 + %iq1s_code_9_19 = scalar.constant 5413.0 : f32 + %iq1s_code_9_20 = scalar.constant 5416.0 : f32 + %iq1s_code_9_21 = scalar.constant 5418.0 : f32 + %iq1s_code_9_22 = scalar.constant 5441.0 : f32 + %iq1s_code_9_23 = scalar.constant 5444.0 : f32 + %iq1s_code_9_24 = scalar.constant 5445.0 : f32 + %iq1s_code_9_25 = scalar.constant 5446.0 : f32 + %iq1s_code_9_26 = scalar.constant 5457.0 : f32 + %iq1s_code_9_27 = scalar.constant 5458.0 : f32 + %iq1s_code_9_28 = scalar.constant 5460.0 : f32 + %iq1s_code_9_29 = scalar.constant 5461.0 : f32 + %iq1s_code_9_30 = scalar.constant 5462.0 : f32 + %iq1s_code_9_31 = scalar.constant 5465.0 : f32 + %iq1s_code_9 = vector.from_elements %iq1s_code_9_0, %iq1s_code_9_1, %iq1s_code_9_2, %iq1s_code_9_3, %iq1s_code_9_4, %iq1s_code_9_5, %iq1s_code_9_6, %iq1s_code_9_7, %iq1s_code_9_8, %iq1s_code_9_9, %iq1s_code_9_10, %iq1s_code_9_11, %iq1s_code_9_12, %iq1s_code_9_13, %iq1s_code_9_14, %iq1s_code_9_15, %iq1s_code_9_16, %iq1s_code_9_17, %iq1s_code_9_18, %iq1s_code_9_19, %iq1s_code_9_20, %iq1s_code_9_21, %iq1s_code_9_22, %iq1s_code_9_23, %iq1s_code_9_24, %iq1s_code_9_25, %iq1s_code_9_26, %iq1s_code_9_27, %iq1s_code_9_28, %iq1s_code_9_29, %iq1s_code_9_30, %iq1s_code_9_31 : vector<32xf32> + %iq1s_code_10_0 = scalar.constant 5466.0 : f32 + %iq1s_code_10_1 = scalar.constant 5473.0 : f32 + %iq1s_code_10_2 = scalar.constant 5476.0 : f32 + %iq1s_code_10_3 = scalar.constant 5477.0 : f32 + %iq1s_code_10_4 = scalar.constant 5478.0 : f32 + %iq1s_code_10_5 = scalar.constant 5481.0 : f32 + %iq1s_code_10_6 = scalar.constant 5504.0 : f32 + %iq1s_code_10_7 = scalar.constant 5506.0 : f32 + %iq1s_code_10_8 = scalar.constant 5508.0 : f32 + %iq1s_code_10_9 = scalar.constant 5509.0 : f32 + %iq1s_code_10_10 = scalar.constant 5512.0 : f32 + %iq1s_code_10_11 = scalar.constant 5514.0 : f32 + %iq1s_code_10_12 = scalar.constant 5520.0 : f32 + %iq1s_code_10_13 = scalar.constant 5521.0 : f32 + %iq1s_code_10_14 = scalar.constant 5524.0 : f32 + %iq1s_code_10_15 = scalar.constant 5525.0 : f32 + %iq1s_code_10_16 = scalar.constant 5526.0 : f32 + %iq1s_code_10_17 = scalar.constant 5529.0 : f32 + %iq1s_code_10_18 = scalar.constant 5530.0 : f32 + %iq1s_code_10_19 = scalar.constant 5536.0 : f32 + %iq1s_code_10_20 = scalar.constant 5538.0 : f32 + %iq1s_code_10_21 = scalar.constant 5541.0 : f32 + %iq1s_code_10_22 = scalar.constant 5633.0 : f32 + %iq1s_code_10_23 = scalar.constant 5636.0 : f32 + %iq1s_code_10_24 = scalar.constant 5637.0 : f32 + %iq1s_code_10_25 = scalar.constant 5638.0 : f32 + %iq1s_code_10_26 = scalar.constant 5653.0 : f32 + %iq1s_code_10_27 = scalar.constant 5654.0 : f32 + %iq1s_code_10_28 = scalar.constant 5656.0 : f32 + %iq1s_code_10_29 = scalar.constant 5658.0 : f32 + %iq1s_code_10_30 = scalar.constant 5665.0 : f32 + %iq1s_code_10_31 = scalar.constant 5670.0 : f32 + %iq1s_code_10 = vector.from_elements %iq1s_code_10_0, %iq1s_code_10_1, %iq1s_code_10_2, %iq1s_code_10_3, %iq1s_code_10_4, %iq1s_code_10_5, %iq1s_code_10_6, %iq1s_code_10_7, %iq1s_code_10_8, %iq1s_code_10_9, %iq1s_code_10_10, %iq1s_code_10_11, %iq1s_code_10_12, %iq1s_code_10_13, %iq1s_code_10_14, %iq1s_code_10_15, %iq1s_code_10_16, %iq1s_code_10_17, %iq1s_code_10_18, %iq1s_code_10_19, %iq1s_code_10_20, %iq1s_code_10_21, %iq1s_code_10_22, %iq1s_code_10_23, %iq1s_code_10_24, %iq1s_code_10_25, %iq1s_code_10_26, %iq1s_code_10_27, %iq1s_code_10_28, %iq1s_code_10_29, %iq1s_code_10_30, %iq1s_code_10_31 : vector<32xf32> + %iq1s_code_11_0 = scalar.constant 5696.0 : f32 + %iq1s_code_11_1 = scalar.constant 5698.0 : f32 + %iq1s_code_11_2 = scalar.constant 5700.0 : f32 + %iq1s_code_11_3 = scalar.constant 5701.0 : f32 + %iq1s_code_11_4 = scalar.constant 5704.0 : f32 + %iq1s_code_11_5 = scalar.constant 5706.0 : f32 + %iq1s_code_11_6 = scalar.constant 5713.0 : f32 + %iq1s_code_11_7 = scalar.constant 5717.0 : f32 + %iq1s_code_11_8 = scalar.constant 5718.0 : f32 + %iq1s_code_11_9 = scalar.constant 5720.0 : f32 + %iq1s_code_11_10 = scalar.constant 5721.0 : f32 + %iq1s_code_11_11 = scalar.constant 5729.0 : f32 + %iq1s_code_11_12 = scalar.constant 5732.0 : f32 + %iq1s_code_11_13 = scalar.constant 5733.0 : f32 + %iq1s_code_11_14 = scalar.constant 5736.0 : f32 + %iq1s_code_11_15 = scalar.constant 5737.0 : f32 + %iq1s_code_11_16 = scalar.constant 5738.0 : f32 + %iq1s_code_11_17 = scalar.constant 5766.0 : f32 + %iq1s_code_11_18 = scalar.constant 5770.0 : f32 + %iq1s_code_11_19 = scalar.constant 5778.0 : f32 + %iq1s_code_11_20 = scalar.constant 5781.0 : f32 + %iq1s_code_11_21 = scalar.constant 5796.0 : f32 + %iq1s_code_11_22 = scalar.constant 5801.0 : f32 + %iq1s_code_11_23 = scalar.constant 6161.0 : f32 + %iq1s_code_11_24 = scalar.constant 6166.0 : f32 + %iq1s_code_11_25 = scalar.constant 6181.0 : f32 + %iq1s_code_11_26 = scalar.constant 6209.0 : f32 + %iq1s_code_11_27 = scalar.constant 6212.0 : f32 + %iq1s_code_11_28 = scalar.constant 6214.0 : f32 + %iq1s_code_11_29 = scalar.constant 6217.0 : f32 + %iq1s_code_11_30 = scalar.constant 6224.0 : f32 + %iq1s_code_11_31 = scalar.constant 6229.0 : f32 + %iq1s_code_11 = vector.from_elements %iq1s_code_11_0, %iq1s_code_11_1, %iq1s_code_11_2, %iq1s_code_11_3, %iq1s_code_11_4, %iq1s_code_11_5, %iq1s_code_11_6, %iq1s_code_11_7, %iq1s_code_11_8, %iq1s_code_11_9, %iq1s_code_11_10, %iq1s_code_11_11, %iq1s_code_11_12, %iq1s_code_11_13, %iq1s_code_11_14, %iq1s_code_11_15, %iq1s_code_11_16, %iq1s_code_11_17, %iq1s_code_11_18, %iq1s_code_11_19, %iq1s_code_11_20, %iq1s_code_11_21, %iq1s_code_11_22, %iq1s_code_11_23, %iq1s_code_11_24, %iq1s_code_11_25, %iq1s_code_11_26, %iq1s_code_11_27, %iq1s_code_11_28, %iq1s_code_11_29, %iq1s_code_11_30, %iq1s_code_11_31 : vector<32xf32> + %iq1s_code_12_0 = scalar.constant 6232.0 : f32 + %iq1s_code_12_1 = scalar.constant 6234.0 : f32 + %iq1s_code_12_2 = scalar.constant 6240.0 : f32 + %iq1s_code_12_3 = scalar.constant 6241.0 : f32 + %iq1s_code_12_4 = scalar.constant 6244.0 : f32 + %iq1s_code_12_5 = scalar.constant 6246.0 : f32 + %iq1s_code_12_6 = scalar.constant 6249.0 : f32 + %iq1s_code_12_7 = scalar.constant 6277.0 : f32 + %iq1s_code_12_8 = scalar.constant 6289.0 : f32 + %iq1s_code_12_9 = scalar.constant 6292.0 : f32 + %iq1s_code_12_10 = scalar.constant 6309.0 : f32 + %iq1s_code_12_11 = scalar.constant 6416.0 : f32 + %iq1s_code_12_12 = scalar.constant 6418.0 : f32 + %iq1s_code_12_13 = scalar.constant 6421.0 : f32 + %iq1s_code_12_14 = scalar.constant 6426.0 : f32 + %iq1s_code_12_15 = scalar.constant 6433.0 : f32 + %iq1s_code_12_16 = scalar.constant 6437.0 : f32 + %iq1s_code_12_17 = scalar.constant 6466.0 : f32 + %iq1s_code_12_18 = scalar.constant 6468.0 : f32 + %iq1s_code_12_19 = scalar.constant 6469.0 : f32 + %iq1s_code_12_20 = scalar.constant 6472.0 : f32 + %iq1s_code_12_21 = scalar.constant 6481.0 : f32 + %iq1s_code_12_22 = scalar.constant 6484.0 : f32 + %iq1s_code_12_23 = scalar.constant 6485.0 : f32 + %iq1s_code_12_24 = scalar.constant 6486.0 : f32 + %iq1s_code_12_25 = scalar.constant 6489.0 : f32 + %iq1s_code_12_26 = scalar.constant 6490.0 : f32 + %iq1s_code_12_27 = scalar.constant 6496.0 : f32 + %iq1s_code_12_28 = scalar.constant 6501.0 : f32 + %iq1s_code_12_29 = scalar.constant 6506.0 : f32 + %iq1s_code_12_30 = scalar.constant 6537.0 : f32 + %iq1s_code_12_31 = scalar.constant 6545.0 : f32 + %iq1s_code_12 = vector.from_elements %iq1s_code_12_0, %iq1s_code_12_1, %iq1s_code_12_2, %iq1s_code_12_3, %iq1s_code_12_4, %iq1s_code_12_5, %iq1s_code_12_6, %iq1s_code_12_7, %iq1s_code_12_8, %iq1s_code_12_9, %iq1s_code_12_10, %iq1s_code_12_11, %iq1s_code_12_12, %iq1s_code_12_13, %iq1s_code_12_14, %iq1s_code_12_15, %iq1s_code_12_16, %iq1s_code_12_17, %iq1s_code_12_18, %iq1s_code_12_19, %iq1s_code_12_20, %iq1s_code_12_21, %iq1s_code_12_22, %iq1s_code_12_23, %iq1s_code_12_24, %iq1s_code_12_25, %iq1s_code_12_26, %iq1s_code_12_27, %iq1s_code_12_28, %iq1s_code_12_29, %iq1s_code_12_30, %iq1s_code_12_31 : vector<32xf32> + %iq1s_code_13_0 = scalar.constant 6546.0 : f32 + %iq1s_code_13_1 = scalar.constant 6549.0 : f32 + %iq1s_code_13_2 = scalar.constant 6552.0 : f32 + %iq1s_code_13_3 = scalar.constant 6561.0 : f32 + %iq1s_code_13_4 = scalar.constant 6566.0 : f32 + %iq1s_code_13_5 = scalar.constant 6569.0 : f32 + %iq1s_code_13_6 = scalar.constant 6665.0 : f32 + %iq1s_code_13_7 = scalar.constant 6678.0 : f32 + %iq1s_code_13_8 = scalar.constant 6692.0 : f32 + %iq1s_code_13_9 = scalar.constant 6694.0 : f32 + %iq1s_code_13_10 = scalar.constant 6724.0 : f32 + %iq1s_code_13_11 = scalar.constant 6726.0 : f32 + %iq1s_code_13_12 = scalar.constant 6729.0 : f32 + %iq1s_code_13_13 = scalar.constant 6736.0 : f32 + %iq1s_code_13_14 = scalar.constant 6738.0 : f32 + %iq1s_code_13_15 = scalar.constant 6741.0 : f32 + %iq1s_code_13_16 = scalar.constant 6744.0 : f32 + %iq1s_code_13_17 = scalar.constant 6753.0 : f32 + %iq1s_code_13_18 = scalar.constant 6758.0 : f32 + %iq1s_code_13_19 = scalar.constant 6761.0 : f32 + %iq1s_code_13_20 = scalar.constant 6789.0 : f32 + %iq1s_code_13_21 = scalar.constant 6801.0 : f32 + %iq1s_code_13_22 = scalar.constant 6806.0 : f32 + %iq1s_code_13_23 = scalar.constant 6810.0 : f32 + %iq1s_code_13_24 = scalar.constant 8192.0 : f32 + %iq1s_code_13_25 = scalar.constant 8194.0 : f32 + %iq1s_code_13_26 = scalar.constant 8200.0 : f32 + %iq1s_code_13_27 = scalar.constant 8202.0 : f32 + %iq1s_code_13_28 = scalar.constant 8213.0 : f32 + %iq1s_code_13_29 = scalar.constant 8224.0 : f32 + %iq1s_code_13_30 = scalar.constant 8226.0 : f32 + %iq1s_code_13_31 = scalar.constant 8229.0 : f32 + %iq1s_code_13 = vector.from_elements %iq1s_code_13_0, %iq1s_code_13_1, %iq1s_code_13_2, %iq1s_code_13_3, %iq1s_code_13_4, %iq1s_code_13_5, %iq1s_code_13_6, %iq1s_code_13_7, %iq1s_code_13_8, %iq1s_code_13_9, %iq1s_code_13_10, %iq1s_code_13_11, %iq1s_code_13_12, %iq1s_code_13_13, %iq1s_code_13_14, %iq1s_code_13_15, %iq1s_code_13_16, %iq1s_code_13_17, %iq1s_code_13_18, %iq1s_code_13_19, %iq1s_code_13_20, %iq1s_code_13_21, %iq1s_code_13_22, %iq1s_code_13_23, %iq1s_code_13_24, %iq1s_code_13_25, %iq1s_code_13_26, %iq1s_code_13_27, %iq1s_code_13_28, %iq1s_code_13_29, %iq1s_code_13_30, %iq1s_code_13_31 : vector<32xf32> + %iq1s_code_14_0 = scalar.constant 8232.0 : f32 + %iq1s_code_14_1 = scalar.constant 8234.0 : f32 + %iq1s_code_14_2 = scalar.constant 8261.0 : f32 + %iq1s_code_14_3 = scalar.constant 8273.0 : f32 + %iq1s_code_14_4 = scalar.constant 8281.0 : f32 + %iq1s_code_14_5 = scalar.constant 8289.0 : f32 + %iq1s_code_14_6 = scalar.constant 8293.0 : f32 + %iq1s_code_14_7 = scalar.constant 8320.0 : f32 + %iq1s_code_14_8 = scalar.constant 8322.0 : f32 + %iq1s_code_14_9 = scalar.constant 8328.0 : f32 + %iq1s_code_14_10 = scalar.constant 8330.0 : f32 + %iq1s_code_14_11 = scalar.constant 8341.0 : f32 + %iq1s_code_14_12 = scalar.constant 8352.0 : f32 + %iq1s_code_14_13 = scalar.constant 8354.0 : f32 + %iq1s_code_14_14 = scalar.constant 8357.0 : f32 + %iq1s_code_14_15 = scalar.constant 8360.0 : f32 + %iq1s_code_14_16 = scalar.constant 8362.0 : f32 + %iq1s_code_14_17 = scalar.constant 8453.0 : f32 + %iq1s_code_14_18 = scalar.constant 8465.0 : f32 + %iq1s_code_14_19 = scalar.constant 8468.0 : f32 + %iq1s_code_14_20 = scalar.constant 8473.0 : f32 + %iq1s_code_14_21 = scalar.constant 8485.0 : f32 + %iq1s_code_14_22 = scalar.constant 8514.0 : f32 + %iq1s_code_14_23 = scalar.constant 8516.0 : f32 + %iq1s_code_14_24 = scalar.constant 8521.0 : f32 + %iq1s_code_14_25 = scalar.constant 8533.0 : f32 + %iq1s_code_14_26 = scalar.constant 8536.0 : f32 + %iq1s_code_14_27 = scalar.constant 8538.0 : f32 + %iq1s_code_14_28 = scalar.constant 8545.0 : f32 + %iq1s_code_14_29 = scalar.constant 8548.0 : f32 + %iq1s_code_14_30 = scalar.constant 8549.0 : f32 + %iq1s_code_14_31 = scalar.constant 8550.0 : f32 + %iq1s_code_14 = vector.from_elements %iq1s_code_14_0, %iq1s_code_14_1, %iq1s_code_14_2, %iq1s_code_14_3, %iq1s_code_14_4, %iq1s_code_14_5, %iq1s_code_14_6, %iq1s_code_14_7, %iq1s_code_14_8, %iq1s_code_14_9, %iq1s_code_14_10, %iq1s_code_14_11, %iq1s_code_14_12, %iq1s_code_14_13, %iq1s_code_14_14, %iq1s_code_14_15, %iq1s_code_14_16, %iq1s_code_14_17, %iq1s_code_14_18, %iq1s_code_14_19, %iq1s_code_14_20, %iq1s_code_14_21, %iq1s_code_14_22, %iq1s_code_14_23, %iq1s_code_14_24, %iq1s_code_14_25, %iq1s_code_14_26, %iq1s_code_14_27, %iq1s_code_14_28, %iq1s_code_14_29, %iq1s_code_14_30, %iq1s_code_14_31 : vector<32xf32> + %iq1s_code_15_0 = scalar.constant 8581.0 : f32 + %iq1s_code_15_1 = scalar.constant 8592.0 : f32 + %iq1s_code_15_2 = scalar.constant 8598.0 : f32 + %iq1s_code_15_3 = scalar.constant 8601.0 : f32 + %iq1s_code_15_4 = scalar.constant 8613.0 : f32 + %iq1s_code_15_5 = scalar.constant 8705.0 : f32 + %iq1s_code_15_6 = scalar.constant 8712.0 : f32 + %iq1s_code_15_7 = scalar.constant 8714.0 : f32 + %iq1s_code_15_8 = scalar.constant 8721.0 : f32 + %iq1s_code_15_9 = scalar.constant 8725.0 : f32 + %iq1s_code_15_10 = scalar.constant 8736.0 : f32 + %iq1s_code_15_11 = scalar.constant 8738.0 : f32 + %iq1s_code_15_12 = scalar.constant 8744.0 : f32 + %iq1s_code_15_13 = scalar.constant 8746.0 : f32 + %iq1s_code_15_14 = scalar.constant 8773.0 : f32 + %iq1s_code_15_15 = scalar.constant 8785.0 : f32 + %iq1s_code_15_16 = scalar.constant 8790.0 : f32 + %iq1s_code_15_17 = scalar.constant 8793.0 : f32 + %iq1s_code_15_18 = scalar.constant 8805.0 : f32 + %iq1s_code_15_19 = scalar.constant 8833.0 : f32 + %iq1s_code_15_20 = scalar.constant 8840.0 : f32 + %iq1s_code_15_21 = scalar.constant 8842.0 : f32 + %iq1s_code_15_22 = scalar.constant 8849.0 : f32 + %iq1s_code_15_23 = scalar.constant 8853.0 : f32 + %iq1s_code_15_24 = scalar.constant 8864.0 : f32 + %iq1s_code_15_25 = scalar.constant 8866.0 : f32 + %iq1s_code_15_26 = scalar.constant 8872.0 : f32 + %iq1s_code_15_27 = scalar.constant 8874.0 : f32 + %iq1s_code_15_28 = scalar.constant 9221.0 : f32 + %iq1s_code_15_29 = scalar.constant 9236.0 : f32 + %iq1s_code_15_30 = scalar.constant 9238.0 : f32 + %iq1s_code_15_31 = scalar.constant 9241.0 : f32 + %iq1s_code_15 = vector.from_elements %iq1s_code_15_0, %iq1s_code_15_1, %iq1s_code_15_2, %iq1s_code_15_3, %iq1s_code_15_4, %iq1s_code_15_5, %iq1s_code_15_6, %iq1s_code_15_7, %iq1s_code_15_8, %iq1s_code_15_9, %iq1s_code_15_10, %iq1s_code_15_11, %iq1s_code_15_12, %iq1s_code_15_13, %iq1s_code_15_14, %iq1s_code_15_15, %iq1s_code_15_16, %iq1s_code_15_17, %iq1s_code_15_18, %iq1s_code_15_19, %iq1s_code_15_20, %iq1s_code_15_21, %iq1s_code_15_22, %iq1s_code_15_23, %iq1s_code_15_24, %iq1s_code_15_25, %iq1s_code_15_26, %iq1s_code_15_27, %iq1s_code_15_28, %iq1s_code_15_29, %iq1s_code_15_30, %iq1s_code_15_31 : vector<32xf32> + %iq1s_code_16_0 = scalar.constant 9253.0 : f32 + %iq1s_code_16_1 = scalar.constant 9284.0 : f32 + %iq1s_code_16_2 = scalar.constant 9285.0 : f32 + %iq1s_code_16_3 = scalar.constant 9286.0 : f32 + %iq1s_code_16_4 = scalar.constant 9289.0 : f32 + %iq1s_code_16_5 = scalar.constant 9298.0 : f32 + %iq1s_code_16_6 = scalar.constant 9301.0 : f32 + %iq1s_code_16_7 = scalar.constant 9304.0 : f32 + %iq1s_code_16_8 = scalar.constant 9306.0 : f32 + %iq1s_code_16_9 = scalar.constant 9318.0 : f32 + %iq1s_code_16_10 = scalar.constant 9349.0 : f32 + %iq1s_code_16_11 = scalar.constant 9361.0 : f32 + %iq1s_code_16_12 = scalar.constant 9364.0 : f32 + %iq1s_code_16_13 = scalar.constant 9369.0 : f32 + %iq1s_code_16_14 = scalar.constant 9377.0 : f32 + %iq1s_code_16_15 = scalar.constant 9381.0 : f32 + %iq1s_code_16_16 = scalar.constant 9481.0 : f32 + %iq1s_code_16_17 = scalar.constant 9493.0 : f32 + %iq1s_code_16_18 = scalar.constant 9505.0 : f32 + %iq1s_code_16_19 = scalar.constant 9513.0 : f32 + %iq1s_code_16_20 = scalar.constant 9536.0 : f32 + %iq1s_code_16_21 = scalar.constant 9541.0 : f32 + %iq1s_code_16_22 = scalar.constant 9544.0 : f32 + %iq1s_code_16_23 = scalar.constant 9553.0 : f32 + %iq1s_code_16_24 = scalar.constant 9556.0 : f32 + %iq1s_code_16_25 = scalar.constant 9557.0 : f32 + %iq1s_code_16_26 = scalar.constant 9561.0 : f32 + %iq1s_code_16_27 = scalar.constant 9570.0 : f32 + %iq1s_code_16_28 = scalar.constant 9573.0 : f32 + %iq1s_code_16_29 = scalar.constant 9576.0 : f32 + %iq1s_code_16_30 = scalar.constant 9609.0 : f32 + %iq1s_code_16_31 = scalar.constant 9616.0 : f32 + %iq1s_code_16 = vector.from_elements %iq1s_code_16_0, %iq1s_code_16_1, %iq1s_code_16_2, %iq1s_code_16_3, %iq1s_code_16_4, %iq1s_code_16_5, %iq1s_code_16_6, %iq1s_code_16_7, %iq1s_code_16_8, %iq1s_code_16_9, %iq1s_code_16_10, %iq1s_code_16_11, %iq1s_code_16_12, %iq1s_code_16_13, %iq1s_code_16_14, %iq1s_code_16_15, %iq1s_code_16_16, %iq1s_code_16_17, %iq1s_code_16_18, %iq1s_code_16_19, %iq1s_code_16_20, %iq1s_code_16_21, %iq1s_code_16_22, %iq1s_code_16_23, %iq1s_code_16_24, %iq1s_code_16_25, %iq1s_code_16_26, %iq1s_code_16_27, %iq1s_code_16_28, %iq1s_code_16_29, %iq1s_code_16_30, %iq1s_code_16_31 : vector<32xf32> + %iq1s_code_17_0 = scalar.constant 9620.0 : f32 + %iq1s_code_17_1 = scalar.constant 9621.0 : f32 + %iq1s_code_17_2 = scalar.constant 9624.0 : f32 + %iq1s_code_17_3 = scalar.constant 9626.0 : f32 + %iq1s_code_17_4 = scalar.constant 9633.0 : f32 + %iq1s_code_17_5 = scalar.constant 9636.0 : f32 + %iq1s_code_17_6 = scalar.constant 9638.0 : f32 + %iq1s_code_17_7 = scalar.constant 9641.0 : f32 + %iq1s_code_17_8 = scalar.constant 9733.0 : f32 + %iq1s_code_17_9 = scalar.constant 9744.0 : f32 + %iq1s_code_17_10 = scalar.constant 9746.0 : f32 + %iq1s_code_17_11 = scalar.constant 9753.0 : f32 + %iq1s_code_17_12 = scalar.constant 9765.0 : f32 + %iq1s_code_17_13 = scalar.constant 9793.0 : f32 + %iq1s_code_17_14 = scalar.constant 9801.0 : f32 + %iq1s_code_17_15 = scalar.constant 9813.0 : f32 + %iq1s_code_17_16 = scalar.constant 9824.0 : f32 + %iq1s_code_17_17 = scalar.constant 9825.0 : f32 + %iq1s_code_17_18 = scalar.constant 9833.0 : f32 + %iq1s_code_17_19 = scalar.constant 9860.0 : f32 + %iq1s_code_17_20 = scalar.constant 9862.0 : f32 + %iq1s_code_17_21 = scalar.constant 9872.0 : f32 + %iq1s_code_17_22 = scalar.constant 9882.0 : f32 + %iq1s_code_17_23 = scalar.constant 10240.0 : f32 + %iq1s_code_17_24 = scalar.constant 10242.0 : f32 + %iq1s_code_17_25 = scalar.constant 10248.0 : f32 + %iq1s_code_17_26 = scalar.constant 10250.0 : f32 + %iq1s_code_17_27 = scalar.constant 10261.0 : f32 + %iq1s_code_17_28 = scalar.constant 10272.0 : f32 + %iq1s_code_17_29 = scalar.constant 10274.0 : f32 + %iq1s_code_17_30 = scalar.constant 10280.0 : f32 + %iq1s_code_17_31 = scalar.constant 10282.0 : f32 + %iq1s_code_17 = vector.from_elements %iq1s_code_17_0, %iq1s_code_17_1, %iq1s_code_17_2, %iq1s_code_17_3, %iq1s_code_17_4, %iq1s_code_17_5, %iq1s_code_17_6, %iq1s_code_17_7, %iq1s_code_17_8, %iq1s_code_17_9, %iq1s_code_17_10, %iq1s_code_17_11, %iq1s_code_17_12, %iq1s_code_17_13, %iq1s_code_17_14, %iq1s_code_17_15, %iq1s_code_17_16, %iq1s_code_17_17, %iq1s_code_17_18, %iq1s_code_17_19, %iq1s_code_17_20, %iq1s_code_17_21, %iq1s_code_17_22, %iq1s_code_17_23, %iq1s_code_17_24, %iq1s_code_17_25, %iq1s_code_17_26, %iq1s_code_17_27, %iq1s_code_17_28, %iq1s_code_17_29, %iq1s_code_17_30, %iq1s_code_17_31 : vector<32xf32> + %iq1s_code_18_0 = scalar.constant 10309.0 : f32 + %iq1s_code_18_1 = scalar.constant 10321.0 : f32 + %iq1s_code_18_2 = scalar.constant 10324.0 : f32 + %iq1s_code_18_3 = scalar.constant 10341.0 : f32 + %iq1s_code_18_4 = scalar.constant 10368.0 : f32 + %iq1s_code_18_5 = scalar.constant 10370.0 : f32 + %iq1s_code_18_6 = scalar.constant 10376.0 : f32 + %iq1s_code_18_7 = scalar.constant 10378.0 : f32 + %iq1s_code_18_8 = scalar.constant 10400.0 : f32 + %iq1s_code_18_9 = scalar.constant 10402.0 : f32 + %iq1s_code_18_10 = scalar.constant 10408.0 : f32 + %iq1s_code_18_11 = scalar.constant 10410.0 : f32 + %iq1s_code_18_12 = scalar.constant 10505.0 : f32 + %iq1s_code_18_13 = scalar.constant 10513.0 : f32 + %iq1s_code_18_14 = scalar.constant 10516.0 : f32 + %iq1s_code_18_15 = scalar.constant 10521.0 : f32 + %iq1s_code_18_16 = scalar.constant 10533.0 : f32 + %iq1s_code_18_17 = scalar.constant 10566.0 : f32 + %iq1s_code_18_18 = scalar.constant 10569.0 : f32 + %iq1s_code_18_19 = scalar.constant 10578.0 : f32 + %iq1s_code_18_20 = scalar.constant 10581.0 : f32 + %iq1s_code_18_21 = scalar.constant 10593.0 : f32 + %iq1s_code_18_22 = scalar.constant 10596.0 : f32 + %iq1s_code_18_23 = scalar.constant 10598.0 : f32 + %iq1s_code_18_24 = scalar.constant 10601.0 : f32 + %iq1s_code_18_25 = scalar.constant 10629.0 : f32 + %iq1s_code_18_26 = scalar.constant 10640.0 : f32 + %iq1s_code_18_27 = scalar.constant 10646.0 : f32 + %iq1s_code_18_28 = scalar.constant 10649.0 : f32 + %iq1s_code_18_29 = scalar.constant 10660.0 : f32 + %iq1s_code_18_30 = scalar.constant 10661.0 : f32 + %iq1s_code_18_31 = scalar.constant 10752.0 : f32 + %iq1s_code_18 = vector.from_elements %iq1s_code_18_0, %iq1s_code_18_1, %iq1s_code_18_2, %iq1s_code_18_3, %iq1s_code_18_4, %iq1s_code_18_5, %iq1s_code_18_6, %iq1s_code_18_7, %iq1s_code_18_8, %iq1s_code_18_9, %iq1s_code_18_10, %iq1s_code_18_11, %iq1s_code_18_12, %iq1s_code_18_13, %iq1s_code_18_14, %iq1s_code_18_15, %iq1s_code_18_16, %iq1s_code_18_17, %iq1s_code_18_18, %iq1s_code_18_19, %iq1s_code_18_20, %iq1s_code_18_21, %iq1s_code_18_22, %iq1s_code_18_23, %iq1s_code_18_24, %iq1s_code_18_25, %iq1s_code_18_26, %iq1s_code_18_27, %iq1s_code_18_28, %iq1s_code_18_29, %iq1s_code_18_30, %iq1s_code_18_31 : vector<32xf32> + %iq1s_code_19_0 = scalar.constant 10754.0 : f32 + %iq1s_code_19_1 = scalar.constant 10760.0 : f32 + %iq1s_code_19_2 = scalar.constant 10762.0 : f32 + %iq1s_code_19_3 = scalar.constant 10784.0 : f32 + %iq1s_code_19_4 = scalar.constant 10786.0 : f32 + %iq1s_code_19_5 = scalar.constant 10792.0 : f32 + %iq1s_code_19_6 = scalar.constant 10794.0 : f32 + %iq1s_code_19_7 = scalar.constant 10821.0 : f32 + %iq1s_code_19_8 = scalar.constant 10833.0 : f32 + %iq1s_code_19_9 = scalar.constant 10838.0 : f32 + %iq1s_code_19_10 = scalar.constant 10841.0 : f32 + %iq1s_code_19_11 = scalar.constant 10853.0 : f32 + %iq1s_code_19_12 = scalar.constant 10880.0 : f32 + %iq1s_code_19_13 = scalar.constant 10882.0 : f32 + %iq1s_code_19_14 = scalar.constant 10888.0 : f32 + %iq1s_code_19_15 = scalar.constant 10890.0 : f32 + %iq1s_code_19_16 = scalar.constant 10901.0 : f32 + %iq1s_code_19_17 = scalar.constant 10912.0 : f32 + %iq1s_code_19_18 = scalar.constant 10914.0 : f32 + %iq1s_code_19_19 = scalar.constant 10920.0 : f32 + %iq1s_code_19_20 = scalar.constant 10922.0 : f32 + %iq1s_code_19_21 = scalar.constant 16389.0 : f32 + %iq1s_code_19_22 = scalar.constant 16401.0 : f32 + %iq1s_code_19_23 = scalar.constant 16406.0 : f32 + %iq1s_code_19_24 = scalar.constant 16421.0 : f32 + %iq1s_code_19_25 = scalar.constant 16457.0 : f32 + %iq1s_code_19_26 = scalar.constant 16466.0 : f32 + %iq1s_code_19_27 = scalar.constant 16469.0 : f32 + %iq1s_code_19_28 = scalar.constant 16472.0 : f32 + %iq1s_code_19_29 = scalar.constant 16474.0 : f32 + %iq1s_code_19_30 = scalar.constant 16481.0 : f32 + %iq1s_code_19_31 = scalar.constant 16484.0 : f32 + %iq1s_code_19 = vector.from_elements %iq1s_code_19_0, %iq1s_code_19_1, %iq1s_code_19_2, %iq1s_code_19_3, %iq1s_code_19_4, %iq1s_code_19_5, %iq1s_code_19_6, %iq1s_code_19_7, %iq1s_code_19_8, %iq1s_code_19_9, %iq1s_code_19_10, %iq1s_code_19_11, %iq1s_code_19_12, %iq1s_code_19_13, %iq1s_code_19_14, %iq1s_code_19_15, %iq1s_code_19_16, %iq1s_code_19_17, %iq1s_code_19_18, %iq1s_code_19_19, %iq1s_code_19_20, %iq1s_code_19_21, %iq1s_code_19_22, %iq1s_code_19_23, %iq1s_code_19_24, %iq1s_code_19_25, %iq1s_code_19_26, %iq1s_code_19_27, %iq1s_code_19_28, %iq1s_code_19_29, %iq1s_code_19_30, %iq1s_code_19_31 : vector<32xf32> + %iq1s_code_20_0 = scalar.constant 16486.0 : f32 + %iq1s_code_20_1 = scalar.constant 16532.0 : f32 + %iq1s_code_20_2 = scalar.constant 16537.0 : f32 + %iq1s_code_20_3 = scalar.constant 16545.0 : f32 + %iq1s_code_20_4 = scalar.constant 16550.0 : f32 + %iq1s_code_20_5 = scalar.constant 16640.0 : f32 + %iq1s_code_20_6 = scalar.constant 16641.0 : f32 + %iq1s_code_20_7 = scalar.constant 16644.0 : f32 + %iq1s_code_20_8 = scalar.constant 16646.0 : f32 + %iq1s_code_20_9 = scalar.constant 16649.0 : f32 + %iq1s_code_20_10 = scalar.constant 16658.0 : f32 + %iq1s_code_20_11 = scalar.constant 16661.0 : f32 + %iq1s_code_20_12 = scalar.constant 16662.0 : f32 + %iq1s_code_20_13 = scalar.constant 16664.0 : f32 + %iq1s_code_20_14 = scalar.constant 16666.0 : f32 + %iq1s_code_20_15 = scalar.constant 16673.0 : f32 + %iq1s_code_20_16 = scalar.constant 16678.0 : f32 + %iq1s_code_20_17 = scalar.constant 16681.0 : f32 + %iq1s_code_20_18 = scalar.constant 16709.0 : f32 + %iq1s_code_20_19 = scalar.constant 16712.0 : f32 + %iq1s_code_20_20 = scalar.constant 16714.0 : f32 + %iq1s_code_20_21 = scalar.constant 16721.0 : f32 + %iq1s_code_20_22 = scalar.constant 16724.0 : f32 + %iq1s_code_20_23 = scalar.constant 16725.0 : f32 + %iq1s_code_20_24 = scalar.constant 16726.0 : f32 + %iq1s_code_20_25 = scalar.constant 16729.0 : f32 + %iq1s_code_20_26 = scalar.constant 16730.0 : f32 + %iq1s_code_20_27 = scalar.constant 16741.0 : f32 + %iq1s_code_20_28 = scalar.constant 16744.0 : f32 + %iq1s_code_20_29 = scalar.constant 16746.0 : f32 + %iq1s_code_20_30 = scalar.constant 16769.0 : f32 + %iq1s_code_20_31 = scalar.constant 16772.0 : f32 + %iq1s_code_20 = vector.from_elements %iq1s_code_20_0, %iq1s_code_20_1, %iq1s_code_20_2, %iq1s_code_20_3, %iq1s_code_20_4, %iq1s_code_20_5, %iq1s_code_20_6, %iq1s_code_20_7, %iq1s_code_20_8, %iq1s_code_20_9, %iq1s_code_20_10, %iq1s_code_20_11, %iq1s_code_20_12, %iq1s_code_20_13, %iq1s_code_20_14, %iq1s_code_20_15, %iq1s_code_20_16, %iq1s_code_20_17, %iq1s_code_20_18, %iq1s_code_20_19, %iq1s_code_20_20, %iq1s_code_20_21, %iq1s_code_20_22, %iq1s_code_20_23, %iq1s_code_20_24, %iq1s_code_20_25, %iq1s_code_20_26, %iq1s_code_20_27, %iq1s_code_20_28, %iq1s_code_20_29, %iq1s_code_20_30, %iq1s_code_20_31 : vector<32xf32> + %iq1s_code_21_0 = scalar.constant 16774.0 : f32 + %iq1s_code_21_1 = scalar.constant 16784.0 : f32 + %iq1s_code_21_2 = scalar.constant 16786.0 : f32 + %iq1s_code_21_3 = scalar.constant 16789.0 : f32 + %iq1s_code_21_4 = scalar.constant 16800.0 : f32 + %iq1s_code_21_5 = scalar.constant 16801.0 : f32 + %iq1s_code_21_6 = scalar.constant 16802.0 : f32 + %iq1s_code_21_7 = scalar.constant 16901.0 : f32 + %iq1s_code_21_8 = scalar.constant 16913.0 : f32 + %iq1s_code_21_9 = scalar.constant 16916.0 : f32 + %iq1s_code_21_10 = scalar.constant 16918.0 : f32 + %iq1s_code_21_11 = scalar.constant 16933.0 : f32 + %iq1s_code_21_12 = scalar.constant 16961.0 : f32 + %iq1s_code_21_13 = scalar.constant 16978.0 : f32 + %iq1s_code_21_14 = scalar.constant 16981.0 : f32 + %iq1s_code_21_15 = scalar.constant 16986.0 : f32 + %iq1s_code_21_16 = scalar.constant 16996.0 : f32 + %iq1s_code_21_17 = scalar.constant 17001.0 : f32 + %iq1s_code_21_18 = scalar.constant 17033.0 : f32 + %iq1s_code_21_19 = scalar.constant 17044.0 : f32 + %iq1s_code_21_20 = scalar.constant 17061.0 : f32 + %iq1s_code_21_21 = scalar.constant 17409.0 : f32 + %iq1s_code_21_22 = scalar.constant 17429.0 : f32 + %iq1s_code_21_23 = scalar.constant 17433.0 : f32 + %iq1s_code_21_24 = scalar.constant 17449.0 : f32 + %iq1s_code_21_25 = scalar.constant 17477.0 : f32 + %iq1s_code_21_26 = scalar.constant 17480.0 : f32 + %iq1s_code_21_27 = scalar.constant 17482.0 : f32 + %iq1s_code_21_28 = scalar.constant 17489.0 : f32 + %iq1s_code_21_29 = scalar.constant 17492.0 : f32 + %iq1s_code_21_30 = scalar.constant 17493.0 : f32 + %iq1s_code_21_31 = scalar.constant 17494.0 : f32 + %iq1s_code_21 = vector.from_elements %iq1s_code_21_0, %iq1s_code_21_1, %iq1s_code_21_2, %iq1s_code_21_3, %iq1s_code_21_4, %iq1s_code_21_5, %iq1s_code_21_6, %iq1s_code_21_7, %iq1s_code_21_8, %iq1s_code_21_9, %iq1s_code_21_10, %iq1s_code_21_11, %iq1s_code_21_12, %iq1s_code_21_13, %iq1s_code_21_14, %iq1s_code_21_15, %iq1s_code_21_16, %iq1s_code_21_17, %iq1s_code_21_18, %iq1s_code_21_19, %iq1s_code_21_20, %iq1s_code_21_21, %iq1s_code_21_22, %iq1s_code_21_23, %iq1s_code_21_24, %iq1s_code_21_25, %iq1s_code_21_26, %iq1s_code_21_27, %iq1s_code_21_28, %iq1s_code_21_29, %iq1s_code_21_30, %iq1s_code_21_31 : vector<32xf32> + %iq1s_code_22_0 = scalar.constant 17505.0 : f32 + %iq1s_code_22_1 = scalar.constant 17506.0 : f32 + %iq1s_code_22_2 = scalar.constant 17509.0 : f32 + %iq1s_code_22_3 = scalar.constant 17512.0 : f32 + %iq1s_code_22_4 = scalar.constant 17514.0 : f32 + %iq1s_code_22_5 = scalar.constant 17537.0 : f32 + %iq1s_code_22_6 = scalar.constant 17542.0 : f32 + %iq1s_code_22_7 = scalar.constant 17545.0 : f32 + %iq1s_code_22_8 = scalar.constant 17552.0 : f32 + %iq1s_code_22_9 = scalar.constant 17554.0 : f32 + %iq1s_code_22_10 = scalar.constant 17557.0 : f32 + %iq1s_code_22_11 = scalar.constant 17568.0 : f32 + %iq1s_code_22_12 = scalar.constant 17569.0 : f32 + %iq1s_code_22_13 = scalar.constant 17577.0 : f32 + %iq1s_code_22_14 = scalar.constant 17665.0 : f32 + %iq1s_code_22_15 = scalar.constant 17666.0 : f32 + %iq1s_code_22_16 = scalar.constant 17669.0 : f32 + %iq1s_code_22_17 = scalar.constant 17674.0 : f32 + %iq1s_code_22_18 = scalar.constant 17681.0 : f32 + %iq1s_code_22_19 = scalar.constant 17684.0 : f32 + %iq1s_code_22_20 = scalar.constant 17685.0 : f32 + %iq1s_code_22_21 = scalar.constant 17686.0 : f32 + %iq1s_code_22_22 = scalar.constant 17689.0 : f32 + %iq1s_code_22_23 = scalar.constant 17696.0 : f32 + %iq1s_code_22_24 = scalar.constant 17701.0 : f32 + %iq1s_code_22_25 = scalar.constant 17706.0 : f32 + %iq1s_code_22_26 = scalar.constant 17729.0 : f32 + %iq1s_code_22_27 = scalar.constant 17732.0 : f32 + %iq1s_code_22_28 = scalar.constant 17733.0 : f32 + %iq1s_code_22_29 = scalar.constant 17734.0 : f32 + %iq1s_code_22_30 = scalar.constant 17737.0 : f32 + %iq1s_code_22_31 = scalar.constant 17744.0 : f32 + %iq1s_code_22 = vector.from_elements %iq1s_code_22_0, %iq1s_code_22_1, %iq1s_code_22_2, %iq1s_code_22_3, %iq1s_code_22_4, %iq1s_code_22_5, %iq1s_code_22_6, %iq1s_code_22_7, %iq1s_code_22_8, %iq1s_code_22_9, %iq1s_code_22_10, %iq1s_code_22_11, %iq1s_code_22_12, %iq1s_code_22_13, %iq1s_code_22_14, %iq1s_code_22_15, %iq1s_code_22_16, %iq1s_code_22_17, %iq1s_code_22_18, %iq1s_code_22_19, %iq1s_code_22_20, %iq1s_code_22_21, %iq1s_code_22_22, %iq1s_code_22_23, %iq1s_code_22_24, %iq1s_code_22_25, %iq1s_code_22_26, %iq1s_code_22_27, %iq1s_code_22_28, %iq1s_code_22_29, %iq1s_code_22_30, %iq1s_code_22_31 : vector<32xf32> + %iq1s_code_23_0 = scalar.constant 17745.0 : f32 + %iq1s_code_23_1 = scalar.constant 17748.0 : f32 + %iq1s_code_23_2 = scalar.constant 17749.0 : f32 + %iq1s_code_23_3 = scalar.constant 17750.0 : f32 + %iq1s_code_23_4 = scalar.constant 17752.0 : f32 + %iq1s_code_23_5 = scalar.constant 17753.0 : f32 + %iq1s_code_23_6 = scalar.constant 17761.0 : f32 + %iq1s_code_23_7 = scalar.constant 17764.0 : f32 + %iq1s_code_23_8 = scalar.constant 17765.0 : f32 + %iq1s_code_23_9 = scalar.constant 17766.0 : f32 + %iq1s_code_23_10 = scalar.constant 17769.0 : f32 + %iq1s_code_23_11 = scalar.constant 17794.0 : f32 + %iq1s_code_23_12 = scalar.constant 17796.0 : f32 + %iq1s_code_23_13 = scalar.constant 17797.0 : f32 + %iq1s_code_23_14 = scalar.constant 17800.0 : f32 + %iq1s_code_23_15 = scalar.constant 17809.0 : f32 + %iq1s_code_23_16 = scalar.constant 17812.0 : f32 + %iq1s_code_23_17 = scalar.constant 17813.0 : f32 + %iq1s_code_23_18 = scalar.constant 17814.0 : f32 + %iq1s_code_23_19 = scalar.constant 17817.0 : f32 + %iq1s_code_23_20 = scalar.constant 17818.0 : f32 + %iq1s_code_23_21 = scalar.constant 17829.0 : f32 + %iq1s_code_23_22 = scalar.constant 17832.0 : f32 + %iq1s_code_23_23 = scalar.constant 17834.0 : f32 + %iq1s_code_23_24 = scalar.constant 17921.0 : f32 + %iq1s_code_23_25 = scalar.constant 17925.0 : f32 + %iq1s_code_23_26 = scalar.constant 17929.0 : f32 + %iq1s_code_23_27 = scalar.constant 17940.0 : f32 + %iq1s_code_23_28 = scalar.constant 17941.0 : f32 + %iq1s_code_23_29 = scalar.constant 17944.0 : f32 + %iq1s_code_23_30 = scalar.constant 17946.0 : f32 + %iq1s_code_23_31 = scalar.constant 17953.0 : f32 + %iq1s_code_23 = vector.from_elements %iq1s_code_23_0, %iq1s_code_23_1, %iq1s_code_23_2, %iq1s_code_23_3, %iq1s_code_23_4, %iq1s_code_23_5, %iq1s_code_23_6, %iq1s_code_23_7, %iq1s_code_23_8, %iq1s_code_23_9, %iq1s_code_23_10, %iq1s_code_23_11, %iq1s_code_23_12, %iq1s_code_23_13, %iq1s_code_23_14, %iq1s_code_23_15, %iq1s_code_23_16, %iq1s_code_23_17, %iq1s_code_23_18, %iq1s_code_23_19, %iq1s_code_23_20, %iq1s_code_23_21, %iq1s_code_23_22, %iq1s_code_23_23, %iq1s_code_23_24, %iq1s_code_23_25, %iq1s_code_23_26, %iq1s_code_23_27, %iq1s_code_23_28, %iq1s_code_23_29, %iq1s_code_23_30, %iq1s_code_23_31 : vector<32xf32> + %iq1s_code_24_0 = scalar.constant 17956.0 : f32 + %iq1s_code_24_1 = scalar.constant 17961.0 : f32 + %iq1s_code_24_2 = scalar.constant 17984.0 : f32 + %iq1s_code_24_3 = scalar.constant 17986.0 : f32 + %iq1s_code_24_4 = scalar.constant 17989.0 : f32 + %iq1s_code_24_5 = scalar.constant 17992.0 : f32 + %iq1s_code_24_6 = scalar.constant 18000.0 : f32 + %iq1s_code_24_7 = scalar.constant 18001.0 : f32 + %iq1s_code_24_8 = scalar.constant 18002.0 : f32 + %iq1s_code_24_9 = scalar.constant 18005.0 : f32 + %iq1s_code_24_10 = scalar.constant 18006.0 : f32 + %iq1s_code_24_11 = scalar.constant 18009.0 : f32 + %iq1s_code_24_12 = scalar.constant 18018.0 : f32 + %iq1s_code_24_13 = scalar.constant 18021.0 : f32 + %iq1s_code_24_14 = scalar.constant 18024.0 : f32 + %iq1s_code_24_15 = scalar.constant 18049.0 : f32 + %iq1s_code_24_16 = scalar.constant 18053.0 : f32 + %iq1s_code_24_17 = scalar.constant 18058.0 : f32 + %iq1s_code_24_18 = scalar.constant 18068.0 : f32 + %iq1s_code_24_19 = scalar.constant 18069.0 : f32 + %iq1s_code_24_20 = scalar.constant 18081.0 : f32 + %iq1s_code_24_21 = scalar.constant 18084.0 : f32 + %iq1s_code_24_22 = scalar.constant 18086.0 : f32 + %iq1s_code_24_23 = scalar.constant 18437.0 : f32 + %iq1s_code_24_24 = scalar.constant 18449.0 : f32 + %iq1s_code_24_25 = scalar.constant 18453.0 : f32 + %iq1s_code_24_26 = scalar.constant 18458.0 : f32 + %iq1s_code_24_27 = scalar.constant 18469.0 : f32 + %iq1s_code_24_28 = scalar.constant 18498.0 : f32 + %iq1s_code_24_29 = scalar.constant 18505.0 : f32 + %iq1s_code_24_30 = scalar.constant 18512.0 : f32 + %iq1s_code_24_31 = scalar.constant 18517.0 : f32 + %iq1s_code_24 = vector.from_elements %iq1s_code_24_0, %iq1s_code_24_1, %iq1s_code_24_2, %iq1s_code_24_3, %iq1s_code_24_4, %iq1s_code_24_5, %iq1s_code_24_6, %iq1s_code_24_7, %iq1s_code_24_8, %iq1s_code_24_9, %iq1s_code_24_10, %iq1s_code_24_11, %iq1s_code_24_12, %iq1s_code_24_13, %iq1s_code_24_14, %iq1s_code_24_15, %iq1s_code_24_16, %iq1s_code_24_17, %iq1s_code_24_18, %iq1s_code_24_19, %iq1s_code_24_20, %iq1s_code_24_21, %iq1s_code_24_22, %iq1s_code_24_23, %iq1s_code_24_24, %iq1s_code_24_25, %iq1s_code_24_26, %iq1s_code_24_27, %iq1s_code_24_28, %iq1s_code_24_29, %iq1s_code_24_30, %iq1s_code_24_31 : vector<32xf32> + %iq1s_code_25_0 = scalar.constant 18520.0 : f32 + %iq1s_code_25_1 = scalar.constant 18529.0 : f32 + %iq1s_code_25_2 = scalar.constant 18532.0 : f32 + %iq1s_code_25_3 = scalar.constant 18534.0 : f32 + %iq1s_code_25_4 = scalar.constant 18537.0 : f32 + %iq1s_code_25_5 = scalar.constant 18565.0 : f32 + %iq1s_code_25_6 = scalar.constant 18577.0 : f32 + %iq1s_code_25_7 = scalar.constant 18580.0 : f32 + %iq1s_code_25_8 = scalar.constant 18582.0 : f32 + %iq1s_code_25_9 = scalar.constant 18585.0 : f32 + %iq1s_code_25_10 = scalar.constant 18597.0 : f32 + %iq1s_code_25_11 = scalar.constant 18689.0 : f32 + %iq1s_code_25_12 = scalar.constant 18693.0 : f32 + %iq1s_code_25_13 = scalar.constant 18694.0 : f32 + %iq1s_code_25_14 = scalar.constant 18698.0 : f32 + %iq1s_code_25_15 = scalar.constant 18704.0 : f32 + %iq1s_code_25_16 = scalar.constant 18708.0 : f32 + %iq1s_code_25_17 = scalar.constant 18709.0 : f32 + %iq1s_code_25_18 = scalar.constant 18712.0 : f32 + %iq1s_code_25_19 = scalar.constant 18721.0 : f32 + %iq1s_code_25_20 = scalar.constant 18724.0 : f32 + %iq1s_code_25_21 = scalar.constant 18726.0 : f32 + %iq1s_code_25_22 = scalar.constant 18752.0 : f32 + %iq1s_code_25_23 = scalar.constant 18757.0 : f32 + %iq1s_code_25_24 = scalar.constant 18762.0 : f32 + %iq1s_code_25_25 = scalar.constant 18769.0 : f32 + %iq1s_code_25_26 = scalar.constant 18770.0 : f32 + %iq1s_code_25_27 = scalar.constant 18772.0 : f32 + %iq1s_code_25_28 = scalar.constant 18773.0 : f32 + %iq1s_code_25_29 = scalar.constant 18774.0 : f32 + %iq1s_code_25_30 = scalar.constant 18777.0 : f32 + %iq1s_code_25_31 = scalar.constant 18784.0 : f32 + %iq1s_code_25 = vector.from_elements %iq1s_code_25_0, %iq1s_code_25_1, %iq1s_code_25_2, %iq1s_code_25_3, %iq1s_code_25_4, %iq1s_code_25_5, %iq1s_code_25_6, %iq1s_code_25_7, %iq1s_code_25_8, %iq1s_code_25_9, %iq1s_code_25_10, %iq1s_code_25_11, %iq1s_code_25_12, %iq1s_code_25_13, %iq1s_code_25_14, %iq1s_code_25_15, %iq1s_code_25_16, %iq1s_code_25_17, %iq1s_code_25_18, %iq1s_code_25_19, %iq1s_code_25_20, %iq1s_code_25_21, %iq1s_code_25_22, %iq1s_code_25_23, %iq1s_code_25_24, %iq1s_code_25_25, %iq1s_code_25_26, %iq1s_code_25_27, %iq1s_code_25_28, %iq1s_code_25_29, %iq1s_code_25_30, %iq1s_code_25_31 : vector<32xf32> + %iq1s_code_26_0 = scalar.constant 18786.0 : f32 + %iq1s_code_26_1 = scalar.constant 18789.0 : f32 + %iq1s_code_26_2 = scalar.constant 18790.0 : f32 + %iq1s_code_26_3 = scalar.constant 18794.0 : f32 + %iq1s_code_26_4 = scalar.constant 18822.0 : f32 + %iq1s_code_26_5 = scalar.constant 18825.0 : f32 + %iq1s_code_26_6 = scalar.constant 18834.0 : f32 + %iq1s_code_26_7 = scalar.constant 18837.0 : f32 + %iq1s_code_26_8 = scalar.constant 18838.0 : f32 + %iq1s_code_26_9 = scalar.constant 18840.0 : f32 + %iq1s_code_26_10 = scalar.constant 18849.0 : f32 + %iq1s_code_26_11 = scalar.constant 18852.0 : f32 + %iq1s_code_26_12 = scalar.constant 18854.0 : f32 + %iq1s_code_26_13 = scalar.constant 18857.0 : f32 + %iq1s_code_26_14 = scalar.constant 18966.0 : f32 + %iq1s_code_26_15 = scalar.constant 19012.0 : f32 + %iq1s_code_26_16 = scalar.constant 19014.0 : f32 + %iq1s_code_26_17 = scalar.constant 19017.0 : f32 + %iq1s_code_26_18 = scalar.constant 19029.0 : f32 + %iq1s_code_26_19 = scalar.constant 19032.0 : f32 + %iq1s_code_26_20 = scalar.constant 19034.0 : f32 + %iq1s_code_26_21 = scalar.constant 19044.0 : f32 + %iq1s_code_26_22 = scalar.constant 19049.0 : f32 + %iq1s_code_26_23 = scalar.constant 19092.0 : f32 + %iq1s_code_26_24 = scalar.constant 19109.0 : f32 + %iq1s_code_26_25 = scalar.constant 20481.0 : f32 + %iq1s_code_26_26 = scalar.constant 20484.0 : f32 + %iq1s_code_26_27 = scalar.constant 20485.0 : f32 + %iq1s_code_26_28 = scalar.constant 20486.0 : f32 + %iq1s_code_26_29 = scalar.constant 20489.0 : f32 + %iq1s_code_26_30 = scalar.constant 20498.0 : f32 + %iq1s_code_26_31 = scalar.constant 20501.0 : f32 + %iq1s_code_26 = vector.from_elements %iq1s_code_26_0, %iq1s_code_26_1, %iq1s_code_26_2, %iq1s_code_26_3, %iq1s_code_26_4, %iq1s_code_26_5, %iq1s_code_26_6, %iq1s_code_26_7, %iq1s_code_26_8, %iq1s_code_26_9, %iq1s_code_26_10, %iq1s_code_26_11, %iq1s_code_26_12, %iq1s_code_26_13, %iq1s_code_26_14, %iq1s_code_26_15, %iq1s_code_26_16, %iq1s_code_26_17, %iq1s_code_26_18, %iq1s_code_26_19, %iq1s_code_26_20, %iq1s_code_26_21, %iq1s_code_26_22, %iq1s_code_26_23, %iq1s_code_26_24, %iq1s_code_26_25, %iq1s_code_26_26, %iq1s_code_26_27, %iq1s_code_26_28, %iq1s_code_26_29, %iq1s_code_26_30, %iq1s_code_26_31 : vector<32xf32> + %iq1s_code_27_0 = scalar.constant 20506.0 : f32 + %iq1s_code_27_1 = scalar.constant 20513.0 : f32 + %iq1s_code_27_2 = scalar.constant 20516.0 : f32 + %iq1s_code_27_3 = scalar.constant 20521.0 : f32 + %iq1s_code_27_4 = scalar.constant 20544.0 : f32 + %iq1s_code_27_5 = scalar.constant 20549.0 : f32 + %iq1s_code_27_6 = scalar.constant 20552.0 : f32 + %iq1s_code_27_7 = scalar.constant 20561.0 : f32 + %iq1s_code_27_8 = scalar.constant 20564.0 : f32 + %iq1s_code_27_9 = scalar.constant 20565.0 : f32 + %iq1s_code_27_10 = scalar.constant 20566.0 : f32 + %iq1s_code_27_11 = scalar.constant 20569.0 : f32 + %iq1s_code_27_12 = scalar.constant 20581.0 : f32 + %iq1s_code_27_13 = scalar.constant 20584.0 : f32 + %iq1s_code_27_14 = scalar.constant 20614.0 : f32 + %iq1s_code_27_15 = scalar.constant 20617.0 : f32 + %iq1s_code_27_16 = scalar.constant 20629.0 : f32 + %iq1s_code_27_17 = scalar.constant 20632.0 : f32 + %iq1s_code_27_18 = scalar.constant 20640.0 : f32 + %iq1s_code_27_19 = scalar.constant 20641.0 : f32 + %iq1s_code_27_20 = scalar.constant 20646.0 : f32 + %iq1s_code_27_21 = scalar.constant 20649.0 : f32 + %iq1s_code_27_22 = scalar.constant 20741.0 : f32 + %iq1s_code_27_23 = scalar.constant 20744.0 : f32 + %iq1s_code_27_24 = scalar.constant 20745.0 : f32 + %iq1s_code_27_25 = scalar.constant 20746.0 : f32 + %iq1s_code_27_26 = scalar.constant 20753.0 : f32 + %iq1s_code_27_27 = scalar.constant 20756.0 : f32 + %iq1s_code_27_28 = scalar.constant 20757.0 : f32 + %iq1s_code_27_29 = scalar.constant 20758.0 : f32 + %iq1s_code_27_30 = scalar.constant 20760.0 : f32 + %iq1s_code_27_31 = scalar.constant 20761.0 : f32 + %iq1s_code_27 = vector.from_elements %iq1s_code_27_0, %iq1s_code_27_1, %iq1s_code_27_2, %iq1s_code_27_3, %iq1s_code_27_4, %iq1s_code_27_5, %iq1s_code_27_6, %iq1s_code_27_7, %iq1s_code_27_8, %iq1s_code_27_9, %iq1s_code_27_10, %iq1s_code_27_11, %iq1s_code_27_12, %iq1s_code_27_13, %iq1s_code_27_14, %iq1s_code_27_15, %iq1s_code_27_16, %iq1s_code_27_17, %iq1s_code_27_18, %iq1s_code_27_19, %iq1s_code_27_20, %iq1s_code_27_21, %iq1s_code_27_22, %iq1s_code_27_23, %iq1s_code_27_24, %iq1s_code_27_25, %iq1s_code_27_26, %iq1s_code_27_27, %iq1s_code_27_28, %iq1s_code_27_29, %iq1s_code_27_30, %iq1s_code_27_31 : vector<32xf32> + %iq1s_code_28_0 = scalar.constant 20768.0 : f32 + %iq1s_code_28_1 = scalar.constant 20773.0 : f32 + %iq1s_code_28_2 = scalar.constant 20774.0 : f32 + %iq1s_code_28_3 = scalar.constant 20776.0 : f32 + %iq1s_code_28_4 = scalar.constant 20778.0 : f32 + %iq1s_code_28_5 = scalar.constant 20801.0 : f32 + %iq1s_code_28_6 = scalar.constant 20804.0 : f32 + %iq1s_code_28_7 = scalar.constant 20805.0 : f32 + %iq1s_code_28_8 = scalar.constant 20806.0 : f32 + %iq1s_code_28_9 = scalar.constant 20809.0 : f32 + %iq1s_code_28_10 = scalar.constant 20816.0 : f32 + %iq1s_code_28_11 = scalar.constant 20817.0 : f32 + %iq1s_code_28_12 = scalar.constant 20818.0 : f32 + %iq1s_code_28_13 = scalar.constant 20820.0 : f32 + %iq1s_code_28_14 = scalar.constant 20821.0 : f32 + %iq1s_code_28_15 = scalar.constant 20822.0 : f32 + %iq1s_code_28_16 = scalar.constant 20824.0 : f32 + %iq1s_code_28_17 = scalar.constant 20825.0 : f32 + %iq1s_code_28_18 = scalar.constant 20826.0 : f32 + %iq1s_code_28_19 = scalar.constant 20833.0 : f32 + %iq1s_code_28_20 = scalar.constant 20836.0 : f32 + %iq1s_code_28_21 = scalar.constant 20837.0 : f32 + %iq1s_code_28_22 = scalar.constant 20838.0 : f32 + %iq1s_code_28_23 = scalar.constant 20841.0 : f32 + %iq1s_code_28_24 = scalar.constant 20866.0 : f32 + %iq1s_code_28_25 = scalar.constant 20869.0 : f32 + %iq1s_code_28_26 = scalar.constant 20881.0 : f32 + %iq1s_code_28_27 = scalar.constant 20884.0 : f32 + %iq1s_code_28_28 = scalar.constant 20885.0 : f32 + %iq1s_code_28_29 = scalar.constant 20886.0 : f32 + %iq1s_code_28_30 = scalar.constant 20889.0 : f32 + %iq1s_code_28_31 = scalar.constant 20896.0 : f32 + %iq1s_code_28 = vector.from_elements %iq1s_code_28_0, %iq1s_code_28_1, %iq1s_code_28_2, %iq1s_code_28_3, %iq1s_code_28_4, %iq1s_code_28_5, %iq1s_code_28_6, %iq1s_code_28_7, %iq1s_code_28_8, %iq1s_code_28_9, %iq1s_code_28_10, %iq1s_code_28_11, %iq1s_code_28_12, %iq1s_code_28_13, %iq1s_code_28_14, %iq1s_code_28_15, %iq1s_code_28_16, %iq1s_code_28_17, %iq1s_code_28_18, %iq1s_code_28_19, %iq1s_code_28_20, %iq1s_code_28_21, %iq1s_code_28_22, %iq1s_code_28_23, %iq1s_code_28_24, %iq1s_code_28_25, %iq1s_code_28_26, %iq1s_code_28_27, %iq1s_code_28_28, %iq1s_code_28_29, %iq1s_code_28_30, %iq1s_code_28_31 : vector<32xf32> + %iq1s_code_29_0 = scalar.constant 20901.0 : f32 + %iq1s_code_29_1 = scalar.constant 20906.0 : f32 + %iq1s_code_29_2 = scalar.constant 20993.0 : f32 + %iq1s_code_29_3 = scalar.constant 20998.0 : f32 + %iq1s_code_29_4 = scalar.constant 21010.0 : f32 + %iq1s_code_29_5 = scalar.constant 21013.0 : f32 + %iq1s_code_29_6 = scalar.constant 21018.0 : f32 + %iq1s_code_29_7 = scalar.constant 21025.0 : f32 + %iq1s_code_29_8 = scalar.constant 21028.0 : f32 + %iq1s_code_29_9 = scalar.constant 21058.0 : f32 + %iq1s_code_29_10 = scalar.constant 21061.0 : f32 + %iq1s_code_29_11 = scalar.constant 21066.0 : f32 + %iq1s_code_29_12 = scalar.constant 21073.0 : f32 + %iq1s_code_29_13 = scalar.constant 21076.0 : f32 + %iq1s_code_29_14 = scalar.constant 21077.0 : f32 + %iq1s_code_29_15 = scalar.constant 21078.0 : f32 + %iq1s_code_29_16 = scalar.constant 21081.0 : f32 + %iq1s_code_29_17 = scalar.constant 21090.0 : f32 + %iq1s_code_29_18 = scalar.constant 21093.0 : f32 + %iq1s_code_29_19 = scalar.constant 21125.0 : f32 + %iq1s_code_29_20 = scalar.constant 21136.0 : f32 + %iq1s_code_29_21 = scalar.constant 21138.0 : f32 + %iq1s_code_29_22 = scalar.constant 21141.0 : f32 + %iq1s_code_29_23 = scalar.constant 21145.0 : f32 + %iq1s_code_29_24 = scalar.constant 21146.0 : f32 + %iq1s_code_29_25 = scalar.constant 21156.0 : f32 + %iq1s_code_29_26 = scalar.constant 21508.0 : f32 + %iq1s_code_29_27 = scalar.constant 21509.0 : f32 + %iq1s_code_29_28 = scalar.constant 21521.0 : f32 + %iq1s_code_29_29 = scalar.constant 21524.0 : f32 + %iq1s_code_29_30 = scalar.constant 21525.0 : f32 + %iq1s_code_29_31 = scalar.constant 21526.0 : f32 + %iq1s_code_29 = vector.from_elements %iq1s_code_29_0, %iq1s_code_29_1, %iq1s_code_29_2, %iq1s_code_29_3, %iq1s_code_29_4, %iq1s_code_29_5, %iq1s_code_29_6, %iq1s_code_29_7, %iq1s_code_29_8, %iq1s_code_29_9, %iq1s_code_29_10, %iq1s_code_29_11, %iq1s_code_29_12, %iq1s_code_29_13, %iq1s_code_29_14, %iq1s_code_29_15, %iq1s_code_29_16, %iq1s_code_29_17, %iq1s_code_29_18, %iq1s_code_29_19, %iq1s_code_29_20, %iq1s_code_29_21, %iq1s_code_29_22, %iq1s_code_29_23, %iq1s_code_29_24, %iq1s_code_29_25, %iq1s_code_29_26, %iq1s_code_29_27, %iq1s_code_29_28, %iq1s_code_29_29, %iq1s_code_29_30, %iq1s_code_29_31 : vector<32xf32> + %iq1s_code_30_0 = scalar.constant 21528.0 : f32 + %iq1s_code_30_1 = scalar.constant 21529.0 : f32 + %iq1s_code_30_2 = scalar.constant 21537.0 : f32 + %iq1s_code_30_3 = scalar.constant 21541.0 : f32 + %iq1s_code_30_4 = scalar.constant 21544.0 : f32 + %iq1s_code_30_5 = scalar.constant 21546.0 : f32 + %iq1s_code_30_6 = scalar.constant 21569.0 : f32 + %iq1s_code_30_7 = scalar.constant 21572.0 : f32 + %iq1s_code_30_8 = scalar.constant 21573.0 : f32 + %iq1s_code_30_9 = scalar.constant 21574.0 : f32 + %iq1s_code_30_10 = scalar.constant 21577.0 : f32 + %iq1s_code_30_11 = scalar.constant 21578.0 : f32 + %iq1s_code_30_12 = scalar.constant 21584.0 : f32 + %iq1s_code_30_13 = scalar.constant 21585.0 : f32 + %iq1s_code_30_14 = scalar.constant 21588.0 : f32 + %iq1s_code_30_15 = scalar.constant 21589.0 : f32 + %iq1s_code_30_16 = scalar.constant 21590.0 : f32 + %iq1s_code_30_17 = scalar.constant 21592.0 : f32 + %iq1s_code_30_18 = scalar.constant 21593.0 : f32 + %iq1s_code_30_19 = scalar.constant 21594.0 : f32 + %iq1s_code_30_20 = scalar.constant 21601.0 : f32 + %iq1s_code_30_21 = scalar.constant 21602.0 : f32 + %iq1s_code_30_22 = scalar.constant 21604.0 : f32 + %iq1s_code_30_23 = scalar.constant 21605.0 : f32 + %iq1s_code_30_24 = scalar.constant 21606.0 : f32 + %iq1s_code_30_25 = scalar.constant 21609.0 : f32 + %iq1s_code_30_26 = scalar.constant 21632.0 : f32 + %iq1s_code_30_27 = scalar.constant 21640.0 : f32 + %iq1s_code_30_28 = scalar.constant 21642.0 : f32 + %iq1s_code_30_29 = scalar.constant 21649.0 : f32 + %iq1s_code_30_30 = scalar.constant 21652.0 : f32 + %iq1s_code_30_31 = scalar.constant 21653.0 : f32 + %iq1s_code_30 = vector.from_elements %iq1s_code_30_0, %iq1s_code_30_1, %iq1s_code_30_2, %iq1s_code_30_3, %iq1s_code_30_4, %iq1s_code_30_5, %iq1s_code_30_6, %iq1s_code_30_7, %iq1s_code_30_8, %iq1s_code_30_9, %iq1s_code_30_10, %iq1s_code_30_11, %iq1s_code_30_12, %iq1s_code_30_13, %iq1s_code_30_14, %iq1s_code_30_15, %iq1s_code_30_16, %iq1s_code_30_17, %iq1s_code_30_18, %iq1s_code_30_19, %iq1s_code_30_20, %iq1s_code_30_21, %iq1s_code_30_22, %iq1s_code_30_23, %iq1s_code_30_24, %iq1s_code_30_25, %iq1s_code_30_26, %iq1s_code_30_27, %iq1s_code_30_28, %iq1s_code_30_29, %iq1s_code_30_30, %iq1s_code_30_31 : vector<32xf32> + %iq1s_code_31_0 = scalar.constant 21654.0 : f32 + %iq1s_code_31_1 = scalar.constant 21657.0 : f32 + %iq1s_code_31_2 = scalar.constant 21665.0 : f32 + %iq1s_code_31_3 = scalar.constant 21668.0 : f32 + %iq1s_code_31_4 = scalar.constant 21669.0 : f32 + %iq1s_code_31_5 = scalar.constant 21674.0 : f32 + %iq1s_code_31_6 = scalar.constant 21761.0 : f32 + %iq1s_code_31_7 = scalar.constant 21762.0 : f32 + %iq1s_code_31_8 = scalar.constant 21764.0 : f32 + %iq1s_code_31_9 = scalar.constant 21765.0 : f32 + %iq1s_code_31_10 = scalar.constant 21766.0 : f32 + %iq1s_code_31_11 = scalar.constant 21769.0 : f32 + %iq1s_code_31_12 = scalar.constant 21776.0 : f32 + %iq1s_code_31_13 = scalar.constant 21777.0 : f32 + %iq1s_code_31_14 = scalar.constant 21778.0 : f32 + %iq1s_code_31_15 = scalar.constant 21780.0 : f32 + %iq1s_code_31_16 = scalar.constant 21781.0 : f32 + %iq1s_code_31_17 = scalar.constant 21782.0 : f32 + %iq1s_code_31_18 = scalar.constant 21785.0 : f32 + %iq1s_code_31_19 = scalar.constant 21786.0 : f32 + %iq1s_code_31_20 = scalar.constant 21793.0 : f32 + %iq1s_code_31_21 = scalar.constant 21796.0 : f32 + %iq1s_code_31_22 = scalar.constant 21797.0 : f32 + %iq1s_code_31_23 = scalar.constant 21798.0 : f32 + %iq1s_code_31_24 = scalar.constant 21801.0 : f32 + %iq1s_code_31_25 = scalar.constant 21824.0 : f32 + %iq1s_code_31_26 = scalar.constant 21825.0 : f32 + %iq1s_code_31_27 = scalar.constant 21826.0 : f32 + %iq1s_code_31_28 = scalar.constant 21828.0 : f32 + %iq1s_code_31_29 = scalar.constant 21829.0 : f32 + %iq1s_code_31_30 = scalar.constant 21830.0 : f32 + %iq1s_code_31_31 = scalar.constant 21832.0 : f32 + %iq1s_code_31 = vector.from_elements %iq1s_code_31_0, %iq1s_code_31_1, %iq1s_code_31_2, %iq1s_code_31_3, %iq1s_code_31_4, %iq1s_code_31_5, %iq1s_code_31_6, %iq1s_code_31_7, %iq1s_code_31_8, %iq1s_code_31_9, %iq1s_code_31_10, %iq1s_code_31_11, %iq1s_code_31_12, %iq1s_code_31_13, %iq1s_code_31_14, %iq1s_code_31_15, %iq1s_code_31_16, %iq1s_code_31_17, %iq1s_code_31_18, %iq1s_code_31_19, %iq1s_code_31_20, %iq1s_code_31_21, %iq1s_code_31_22, %iq1s_code_31_23, %iq1s_code_31_24, %iq1s_code_31_25, %iq1s_code_31_26, %iq1s_code_31_27, %iq1s_code_31_28, %iq1s_code_31_29, %iq1s_code_31_30, %iq1s_code_31_31 : vector<32xf32> + %iq1s_code_32_0 = scalar.constant 21833.0 : f32 + %iq1s_code_32_1 = scalar.constant 21840.0 : f32 + %iq1s_code_32_2 = scalar.constant 21841.0 : f32 + %iq1s_code_32_3 = scalar.constant 21842.0 : f32 + %iq1s_code_32_4 = scalar.constant 21844.0 : f32 + %iq1s_code_32_5 = scalar.constant 21845.0 : f32 + %iq1s_code_32_6 = scalar.constant 21846.0 : f32 + %iq1s_code_32_7 = scalar.constant 21848.0 : f32 + %iq1s_code_32_8 = scalar.constant 21849.0 : f32 + %iq1s_code_32_9 = scalar.constant 21850.0 : f32 + %iq1s_code_32_10 = scalar.constant 21856.0 : f32 + %iq1s_code_32_11 = scalar.constant 21857.0 : f32 + %iq1s_code_32_12 = scalar.constant 21860.0 : f32 + %iq1s_code_32_13 = scalar.constant 21861.0 : f32 + %iq1s_code_32_14 = scalar.constant 21862.0 : f32 + %iq1s_code_32_15 = scalar.constant 21864.0 : f32 + %iq1s_code_32_16 = scalar.constant 21865.0 : f32 + %iq1s_code_32_17 = scalar.constant 21866.0 : f32 + %iq1s_code_32_18 = scalar.constant 21889.0 : f32 + %iq1s_code_32_19 = scalar.constant 21892.0 : f32 + %iq1s_code_32_20 = scalar.constant 21893.0 : f32 + %iq1s_code_32_21 = scalar.constant 21897.0 : f32 + %iq1s_code_32_22 = scalar.constant 21898.0 : f32 + %iq1s_code_32_23 = scalar.constant 21904.0 : f32 + %iq1s_code_32_24 = scalar.constant 21905.0 : f32 + %iq1s_code_32_25 = scalar.constant 21908.0 : f32 + %iq1s_code_32_26 = scalar.constant 21909.0 : f32 + %iq1s_code_32_27 = scalar.constant 21910.0 : f32 + %iq1s_code_32_28 = scalar.constant 21912.0 : f32 + %iq1s_code_32_29 = scalar.constant 21913.0 : f32 + %iq1s_code_32_30 = scalar.constant 21921.0 : f32 + %iq1s_code_32_31 = scalar.constant 21924.0 : f32 + %iq1s_code_32 = vector.from_elements %iq1s_code_32_0, %iq1s_code_32_1, %iq1s_code_32_2, %iq1s_code_32_3, %iq1s_code_32_4, %iq1s_code_32_5, %iq1s_code_32_6, %iq1s_code_32_7, %iq1s_code_32_8, %iq1s_code_32_9, %iq1s_code_32_10, %iq1s_code_32_11, %iq1s_code_32_12, %iq1s_code_32_13, %iq1s_code_32_14, %iq1s_code_32_15, %iq1s_code_32_16, %iq1s_code_32_17, %iq1s_code_32_18, %iq1s_code_32_19, %iq1s_code_32_20, %iq1s_code_32_21, %iq1s_code_32_22, %iq1s_code_32_23, %iq1s_code_32_24, %iq1s_code_32_25, %iq1s_code_32_26, %iq1s_code_32_27, %iq1s_code_32_28, %iq1s_code_32_29, %iq1s_code_32_30, %iq1s_code_32_31 : vector<32xf32> + %iq1s_code_33_0 = scalar.constant 21925.0 : f32 + %iq1s_code_33_1 = scalar.constant 21926.0 : f32 + %iq1s_code_33_2 = scalar.constant 21929.0 : f32 + %iq1s_code_33_3 = scalar.constant 22016.0 : f32 + %iq1s_code_33_4 = scalar.constant 22017.0 : f32 + %iq1s_code_33_5 = scalar.constant 22018.0 : f32 + %iq1s_code_33_6 = scalar.constant 22020.0 : f32 + %iq1s_code_33_7 = scalar.constant 22022.0 : f32 + %iq1s_code_33_8 = scalar.constant 22024.0 : f32 + %iq1s_code_33_9 = scalar.constant 22025.0 : f32 + %iq1s_code_33_10 = scalar.constant 22033.0 : f32 + %iq1s_code_33_11 = scalar.constant 22036.0 : f32 + %iq1s_code_33_12 = scalar.constant 22037.0 : f32 + %iq1s_code_33_13 = scalar.constant 22040.0 : f32 + %iq1s_code_33_14 = scalar.constant 22041.0 : f32 + %iq1s_code_33_15 = scalar.constant 22048.0 : f32 + %iq1s_code_33_16 = scalar.constant 22049.0 : f32 + %iq1s_code_33_17 = scalar.constant 22050.0 : f32 + %iq1s_code_33_18 = scalar.constant 22052.0 : f32 + %iq1s_code_33_19 = scalar.constant 22053.0 : f32 + %iq1s_code_33_20 = scalar.constant 22054.0 : f32 + %iq1s_code_33_21 = scalar.constant 22056.0 : f32 + %iq1s_code_33_22 = scalar.constant 22057.0 : f32 + %iq1s_code_33_23 = scalar.constant 22081.0 : f32 + %iq1s_code_33_24 = scalar.constant 22085.0 : f32 + %iq1s_code_33_25 = scalar.constant 22086.0 : f32 + %iq1s_code_33_26 = scalar.constant 22088.0 : f32 + %iq1s_code_33_27 = scalar.constant 22089.0 : f32 + %iq1s_code_33_28 = scalar.constant 22090.0 : f32 + %iq1s_code_33_29 = scalar.constant 22096.0 : f32 + %iq1s_code_33_30 = scalar.constant 22097.0 : f32 + %iq1s_code_33_31 = scalar.constant 22098.0 : f32 + %iq1s_code_33 = vector.from_elements %iq1s_code_33_0, %iq1s_code_33_1, %iq1s_code_33_2, %iq1s_code_33_3, %iq1s_code_33_4, %iq1s_code_33_5, %iq1s_code_33_6, %iq1s_code_33_7, %iq1s_code_33_8, %iq1s_code_33_9, %iq1s_code_33_10, %iq1s_code_33_11, %iq1s_code_33_12, %iq1s_code_33_13, %iq1s_code_33_14, %iq1s_code_33_15, %iq1s_code_33_16, %iq1s_code_33_17, %iq1s_code_33_18, %iq1s_code_33_19, %iq1s_code_33_20, %iq1s_code_33_21, %iq1s_code_33_22, %iq1s_code_33_23, %iq1s_code_33_24, %iq1s_code_33_25, %iq1s_code_33_26, %iq1s_code_33_27, %iq1s_code_33_28, %iq1s_code_33_29, %iq1s_code_33_30, %iq1s_code_33_31 : vector<32xf32> + %iq1s_code_34_0 = scalar.constant 22100.0 : f32 + %iq1s_code_34_1 = scalar.constant 22101.0 : f32 + %iq1s_code_34_2 = scalar.constant 22102.0 : f32 + %iq1s_code_34_3 = scalar.constant 22104.0 : f32 + %iq1s_code_34_4 = scalar.constant 22105.0 : f32 + %iq1s_code_34_5 = scalar.constant 22106.0 : f32 + %iq1s_code_34_6 = scalar.constant 22113.0 : f32 + %iq1s_code_34_7 = scalar.constant 22116.0 : f32 + %iq1s_code_34_8 = scalar.constant 22117.0 : f32 + %iq1s_code_34_9 = scalar.constant 22121.0 : f32 + %iq1s_code_34_10 = scalar.constant 22146.0 : f32 + %iq1s_code_34_11 = scalar.constant 22149.0 : f32 + %iq1s_code_34_12 = scalar.constant 22150.0 : f32 + %iq1s_code_34_13 = scalar.constant 22152.0 : f32 + %iq1s_code_34_14 = scalar.constant 22153.0 : f32 + %iq1s_code_34_15 = scalar.constant 22154.0 : f32 + %iq1s_code_34_16 = scalar.constant 22161.0 : f32 + %iq1s_code_34_17 = scalar.constant 22165.0 : f32 + %iq1s_code_34_18 = scalar.constant 22170.0 : f32 + %iq1s_code_34_19 = scalar.constant 22178.0 : f32 + %iq1s_code_34_20 = scalar.constant 22181.0 : f32 + %iq1s_code_34_21 = scalar.constant 22182.0 : f32 + %iq1s_code_34_22 = scalar.constant 22184.0 : f32 + %iq1s_code_34_23 = scalar.constant 22185.0 : f32 + %iq1s_code_34_24 = scalar.constant 22532.0 : f32 + %iq1s_code_34_25 = scalar.constant 22533.0 : f32 + %iq1s_code_34_26 = scalar.constant 22534.0 : f32 + %iq1s_code_34_27 = scalar.constant 22537.0 : f32 + %iq1s_code_34_28 = scalar.constant 22544.0 : f32 + %iq1s_code_34_29 = scalar.constant 22549.0 : f32 + %iq1s_code_34_30 = scalar.constant 22552.0 : f32 + %iq1s_code_34_31 = scalar.constant 22561.0 : f32 + %iq1s_code_34 = vector.from_elements %iq1s_code_34_0, %iq1s_code_34_1, %iq1s_code_34_2, %iq1s_code_34_3, %iq1s_code_34_4, %iq1s_code_34_5, %iq1s_code_34_6, %iq1s_code_34_7, %iq1s_code_34_8, %iq1s_code_34_9, %iq1s_code_34_10, %iq1s_code_34_11, %iq1s_code_34_12, %iq1s_code_34_13, %iq1s_code_34_14, %iq1s_code_34_15, %iq1s_code_34_16, %iq1s_code_34_17, %iq1s_code_34_18, %iq1s_code_34_19, %iq1s_code_34_20, %iq1s_code_34_21, %iq1s_code_34_22, %iq1s_code_34_23, %iq1s_code_34_24, %iq1s_code_34_25, %iq1s_code_34_26, %iq1s_code_34_27, %iq1s_code_34_28, %iq1s_code_34_29, %iq1s_code_34_30, %iq1s_code_34_31 : vector<32xf32> + %iq1s_code_35_0 = scalar.constant 22570.0 : f32 + %iq1s_code_35_1 = scalar.constant 22597.0 : f32 + %iq1s_code_35_2 = scalar.constant 22600.0 : f32 + %iq1s_code_35_3 = scalar.constant 22602.0 : f32 + %iq1s_code_35_4 = scalar.constant 22609.0 : f32 + %iq1s_code_35_5 = scalar.constant 22612.0 : f32 + %iq1s_code_35_6 = scalar.constant 22613.0 : f32 + %iq1s_code_35_7 = scalar.constant 22614.0 : f32 + %iq1s_code_35_8 = scalar.constant 22616.0 : f32 + %iq1s_code_35_9 = scalar.constant 22617.0 : f32 + %iq1s_code_35_10 = scalar.constant 22624.0 : f32 + %iq1s_code_35_11 = scalar.constant 22626.0 : f32 + %iq1s_code_35_12 = scalar.constant 22628.0 : f32 + %iq1s_code_35_13 = scalar.constant 22629.0 : f32 + %iq1s_code_35_14 = scalar.constant 22658.0 : f32 + %iq1s_code_35_15 = scalar.constant 22665.0 : f32 + %iq1s_code_35_16 = scalar.constant 22672.0 : f32 + %iq1s_code_35_17 = scalar.constant 22674.0 : f32 + %iq1s_code_35_18 = scalar.constant 22677.0 : f32 + %iq1s_code_35_19 = scalar.constant 22680.0 : f32 + %iq1s_code_35_20 = scalar.constant 22689.0 : f32 + %iq1s_code_35_21 = scalar.constant 22697.0 : f32 + %iq1s_code_35_22 = scalar.constant 22785.0 : f32 + %iq1s_code_35_23 = scalar.constant 22786.0 : f32 + %iq1s_code_35_24 = scalar.constant 22789.0 : f32 + %iq1s_code_35_25 = scalar.constant 22794.0 : f32 + %iq1s_code_35_26 = scalar.constant 22801.0 : f32 + %iq1s_code_35_27 = scalar.constant 22804.0 : f32 + %iq1s_code_35_28 = scalar.constant 22805.0 : f32 + %iq1s_code_35_29 = scalar.constant 22806.0 : f32 + %iq1s_code_35_30 = scalar.constant 22809.0 : f32 + %iq1s_code_35_31 = scalar.constant 22821.0 : f32 + %iq1s_code_35 = vector.from_elements %iq1s_code_35_0, %iq1s_code_35_1, %iq1s_code_35_2, %iq1s_code_35_3, %iq1s_code_35_4, %iq1s_code_35_5, %iq1s_code_35_6, %iq1s_code_35_7, %iq1s_code_35_8, %iq1s_code_35_9, %iq1s_code_35_10, %iq1s_code_35_11, %iq1s_code_35_12, %iq1s_code_35_13, %iq1s_code_35_14, %iq1s_code_35_15, %iq1s_code_35_16, %iq1s_code_35_17, %iq1s_code_35_18, %iq1s_code_35_19, %iq1s_code_35_20, %iq1s_code_35_21, %iq1s_code_35_22, %iq1s_code_35_23, %iq1s_code_35_24, %iq1s_code_35_25, %iq1s_code_35_26, %iq1s_code_35_27, %iq1s_code_35_28, %iq1s_code_35_29, %iq1s_code_35_30, %iq1s_code_35_31 : vector<32xf32> + %iq1s_code_36_0 = scalar.constant 22849.0 : f32 + %iq1s_code_36_1 = scalar.constant 22852.0 : f32 + %iq1s_code_36_2 = scalar.constant 22853.0 : f32 + %iq1s_code_36_3 = scalar.constant 22854.0 : f32 + %iq1s_code_36_4 = scalar.constant 22857.0 : f32 + %iq1s_code_36_5 = scalar.constant 22864.0 : f32 + %iq1s_code_36_6 = scalar.constant 22865.0 : f32 + %iq1s_code_36_7 = scalar.constant 22866.0 : f32 + %iq1s_code_36_8 = scalar.constant 22868.0 : f32 + %iq1s_code_36_9 = scalar.constant 22869.0 : f32 + %iq1s_code_36_10 = scalar.constant 22870.0 : f32 + %iq1s_code_36_11 = scalar.constant 22872.0 : f32 + %iq1s_code_36_12 = scalar.constant 22873.0 : f32 + %iq1s_code_36_13 = scalar.constant 22874.0 : f32 + %iq1s_code_36_14 = scalar.constant 22881.0 : f32 + %iq1s_code_36_15 = scalar.constant 22884.0 : f32 + %iq1s_code_36_16 = scalar.constant 22885.0 : f32 + %iq1s_code_36_17 = scalar.constant 22886.0 : f32 + %iq1s_code_36_18 = scalar.constant 22889.0 : f32 + %iq1s_code_36_19 = scalar.constant 22913.0 : f32 + %iq1s_code_36_20 = scalar.constant 22917.0 : f32 + %iq1s_code_36_21 = scalar.constant 22921.0 : f32 + %iq1s_code_36_22 = scalar.constant 22929.0 : f32 + %iq1s_code_36_23 = scalar.constant 22932.0 : f32 + %iq1s_code_36_24 = scalar.constant 22933.0 : f32 + %iq1s_code_36_25 = scalar.constant 22934.0 : f32 + %iq1s_code_36_26 = scalar.constant 22936.0 : f32 + %iq1s_code_36_27 = scalar.constant 22937.0 : f32 + %iq1s_code_36_28 = scalar.constant 22949.0 : f32 + %iq1s_code_36_29 = scalar.constant 23044.0 : f32 + %iq1s_code_36_30 = scalar.constant 23048.0 : f32 + %iq1s_code_36_31 = scalar.constant 23061.0 : f32 + %iq1s_code_36 = vector.from_elements %iq1s_code_36_0, %iq1s_code_36_1, %iq1s_code_36_2, %iq1s_code_36_3, %iq1s_code_36_4, %iq1s_code_36_5, %iq1s_code_36_6, %iq1s_code_36_7, %iq1s_code_36_8, %iq1s_code_36_9, %iq1s_code_36_10, %iq1s_code_36_11, %iq1s_code_36_12, %iq1s_code_36_13, %iq1s_code_36_14, %iq1s_code_36_15, %iq1s_code_36_16, %iq1s_code_36_17, %iq1s_code_36_18, %iq1s_code_36_19, %iq1s_code_36_20, %iq1s_code_36_21, %iq1s_code_36_22, %iq1s_code_36_23, %iq1s_code_36_24, %iq1s_code_36_25, %iq1s_code_36_26, %iq1s_code_36_27, %iq1s_code_36_28, %iq1s_code_36_29, %iq1s_code_36_30, %iq1s_code_36_31 : vector<32xf32> + %iq1s_code_37_0 = scalar.constant 23066.0 : f32 + %iq1s_code_37_1 = scalar.constant 23072.0 : f32 + %iq1s_code_37_2 = scalar.constant 23077.0 : f32 + %iq1s_code_37_3 = scalar.constant 23078.0 : f32 + %iq1s_code_37_4 = scalar.constant 23081.0 : f32 + %iq1s_code_37_5 = scalar.constant 23109.0 : f32 + %iq1s_code_37_6 = scalar.constant 23112.0 : f32 + %iq1s_code_37_7 = scalar.constant 23113.0 : f32 + %iq1s_code_37_8 = scalar.constant 23121.0 : f32 + %iq1s_code_37_9 = scalar.constant 23125.0 : f32 + %iq1s_code_37_10 = scalar.constant 23126.0 : f32 + %iq1s_code_37_11 = scalar.constant 23128.0 : f32 + %iq1s_code_37_12 = scalar.constant 23129.0 : f32 + %iq1s_code_37_13 = scalar.constant 23138.0 : f32 + %iq1s_code_37_14 = scalar.constant 23141.0 : f32 + %iq1s_code_37_15 = scalar.constant 23144.0 : f32 + %iq1s_code_37_16 = scalar.constant 23146.0 : f32 + %iq1s_code_37_17 = scalar.constant 23169.0 : f32 + %iq1s_code_37_18 = scalar.constant 23178.0 : f32 + %iq1s_code_37_19 = scalar.constant 23186.0 : f32 + %iq1s_code_37_20 = scalar.constant 23189.0 : f32 + %iq1s_code_37_21 = scalar.constant 23190.0 : f32 + %iq1s_code_37_22 = scalar.constant 23192.0 : f32 + %iq1s_code_37_23 = scalar.constant 23194.0 : f32 + %iq1s_code_37_24 = scalar.constant 23201.0 : f32 + %iq1s_code_37_25 = scalar.constant 24581.0 : f32 + %iq1s_code_37_26 = scalar.constant 24596.0 : f32 + %iq1s_code_37_27 = scalar.constant 24598.0 : f32 + %iq1s_code_37_28 = scalar.constant 24601.0 : f32 + %iq1s_code_37_29 = scalar.constant 24613.0 : f32 + %iq1s_code_37_30 = scalar.constant 24644.0 : f32 + %iq1s_code_37_31 = scalar.constant 24656.0 : f32 + %iq1s_code_37 = vector.from_elements %iq1s_code_37_0, %iq1s_code_37_1, %iq1s_code_37_2, %iq1s_code_37_3, %iq1s_code_37_4, %iq1s_code_37_5, %iq1s_code_37_6, %iq1s_code_37_7, %iq1s_code_37_8, %iq1s_code_37_9, %iq1s_code_37_10, %iq1s_code_37_11, %iq1s_code_37_12, %iq1s_code_37_13, %iq1s_code_37_14, %iq1s_code_37_15, %iq1s_code_37_16, %iq1s_code_37_17, %iq1s_code_37_18, %iq1s_code_37_19, %iq1s_code_37_20, %iq1s_code_37_21, %iq1s_code_37_22, %iq1s_code_37_23, %iq1s_code_37_24, %iq1s_code_37_25, %iq1s_code_37_26, %iq1s_code_37_27, %iq1s_code_37_28, %iq1s_code_37_29, %iq1s_code_37_30, %iq1s_code_37_31 : vector<32xf32> + %iq1s_code_38_0 = scalar.constant 24661.0 : f32 + %iq1s_code_38_1 = scalar.constant 24662.0 : f32 + %iq1s_code_38_2 = scalar.constant 24664.0 : f32 + %iq1s_code_38_3 = scalar.constant 24666.0 : f32 + %iq1s_code_38_4 = scalar.constant 24673.0 : f32 + %iq1s_code_38_5 = scalar.constant 24676.0 : f32 + %iq1s_code_38_6 = scalar.constant 24678.0 : f32 + %iq1s_code_38_7 = scalar.constant 24681.0 : f32 + %iq1s_code_38_8 = scalar.constant 24705.0 : f32 + %iq1s_code_38_9 = scalar.constant 24726.0 : f32 + %iq1s_code_38_10 = scalar.constant 24741.0 : f32 + %iq1s_code_38_11 = scalar.constant 24833.0 : f32 + %iq1s_code_38_12 = scalar.constant 24836.0 : f32 + %iq1s_code_38_13 = scalar.constant 24838.0 : f32 + %iq1s_code_38_14 = scalar.constant 24841.0 : f32 + %iq1s_code_38_15 = scalar.constant 24850.0 : f32 + %iq1s_code_38_16 = scalar.constant 24853.0 : f32 + %iq1s_code_38_17 = scalar.constant 24865.0 : f32 + %iq1s_code_38_18 = scalar.constant 24866.0 : f32 + %iq1s_code_38_19 = scalar.constant 24870.0 : f32 + %iq1s_code_38_20 = scalar.constant 24873.0 : f32 + %iq1s_code_38_21 = scalar.constant 24901.0 : f32 + %iq1s_code_38_22 = scalar.constant 24905.0 : f32 + %iq1s_code_38_23 = scalar.constant 24913.0 : f32 + %iq1s_code_38_24 = scalar.constant 24917.0 : f32 + %iq1s_code_38_25 = scalar.constant 24918.0 : f32 + %iq1s_code_38_26 = scalar.constant 24921.0 : f32 + %iq1s_code_38_27 = scalar.constant 24933.0 : f32 + %iq1s_code_38_28 = scalar.constant 24934.0 : f32 + %iq1s_code_38_29 = scalar.constant 24938.0 : f32 + %iq1s_code_38_30 = scalar.constant 24964.0 : f32 + %iq1s_code_38_31 = scalar.constant 24970.0 : f32 + %iq1s_code_38 = vector.from_elements %iq1s_code_38_0, %iq1s_code_38_1, %iq1s_code_38_2, %iq1s_code_38_3, %iq1s_code_38_4, %iq1s_code_38_5, %iq1s_code_38_6, %iq1s_code_38_7, %iq1s_code_38_8, %iq1s_code_38_9, %iq1s_code_38_10, %iq1s_code_38_11, %iq1s_code_38_12, %iq1s_code_38_13, %iq1s_code_38_14, %iq1s_code_38_15, %iq1s_code_38_16, %iq1s_code_38_17, %iq1s_code_38_18, %iq1s_code_38_19, %iq1s_code_38_20, %iq1s_code_38_21, %iq1s_code_38_22, %iq1s_code_38_23, %iq1s_code_38_24, %iq1s_code_38_25, %iq1s_code_38_26, %iq1s_code_38_27, %iq1s_code_38_28, %iq1s_code_38_29, %iq1s_code_38_30, %iq1s_code_38_31 : vector<32xf32> + %iq1s_code_39_0 = scalar.constant 24978.0 : f32 + %iq1s_code_39_1 = scalar.constant 24981.0 : f32 + %iq1s_code_39_2 = scalar.constant 24993.0 : f32 + %iq1s_code_39_3 = scalar.constant 24998.0 : f32 + %iq1s_code_39_4 = scalar.constant 25001.0 : f32 + %iq1s_code_39_5 = scalar.constant 25105.0 : f32 + %iq1s_code_39_6 = scalar.constant 25110.0 : f32 + %iq1s_code_39_7 = scalar.constant 25113.0 : f32 + %iq1s_code_39_8 = scalar.constant 25152.0 : f32 + %iq1s_code_39_9 = scalar.constant 25153.0 : f32 + %iq1s_code_39_10 = scalar.constant 25158.0 : f32 + %iq1s_code_39_11 = scalar.constant 25173.0 : f32 + %iq1s_code_39_12 = scalar.constant 25174.0 : f32 + %iq1s_code_39_13 = scalar.constant 25176.0 : f32 + %iq1s_code_39_14 = scalar.constant 25184.0 : f32 + %iq1s_code_39_15 = scalar.constant 25221.0 : f32 + %iq1s_code_39_16 = scalar.constant 25233.0 : f32 + %iq1s_code_39_17 = scalar.constant 25238.0 : f32 + %iq1s_code_39_18 = scalar.constant 25253.0 : f32 + %iq1s_code_39_19 = scalar.constant 25617.0 : f32 + %iq1s_code_39_20 = scalar.constant 25618.0 : f32 + %iq1s_code_39_21 = scalar.constant 25621.0 : f32 + %iq1s_code_39_22 = scalar.constant 25622.0 : f32 + %iq1s_code_39_23 = scalar.constant 25626.0 : f32 + %iq1s_code_39_24 = scalar.constant 25633.0 : f32 + %iq1s_code_39_25 = scalar.constant 25638.0 : f32 + %iq1s_code_39_26 = scalar.constant 25641.0 : f32 + %iq1s_code_39_27 = scalar.constant 25664.0 : f32 + %iq1s_code_39_28 = scalar.constant 25666.0 : f32 + %iq1s_code_39_29 = scalar.constant 25669.0 : f32 + %iq1s_code_39_30 = scalar.constant 25672.0 : f32 + %iq1s_code_39_31 = scalar.constant 25674.0 : f32 + %iq1s_code_39 = vector.from_elements %iq1s_code_39_0, %iq1s_code_39_1, %iq1s_code_39_2, %iq1s_code_39_3, %iq1s_code_39_4, %iq1s_code_39_5, %iq1s_code_39_6, %iq1s_code_39_7, %iq1s_code_39_8, %iq1s_code_39_9, %iq1s_code_39_10, %iq1s_code_39_11, %iq1s_code_39_12, %iq1s_code_39_13, %iq1s_code_39_14, %iq1s_code_39_15, %iq1s_code_39_16, %iq1s_code_39_17, %iq1s_code_39_18, %iq1s_code_39_19, %iq1s_code_39_20, %iq1s_code_39_21, %iq1s_code_39_22, %iq1s_code_39_23, %iq1s_code_39_24, %iq1s_code_39_25, %iq1s_code_39_26, %iq1s_code_39_27, %iq1s_code_39_28, %iq1s_code_39_29, %iq1s_code_39_30, %iq1s_code_39_31 : vector<32xf32> + %iq1s_code_40_0 = scalar.constant 25681.0 : f32 + %iq1s_code_40_1 = scalar.constant 25684.0 : f32 + %iq1s_code_40_2 = scalar.constant 25685.0 : f32 + %iq1s_code_40_3 = scalar.constant 25686.0 : f32 + %iq1s_code_40_4 = scalar.constant 25689.0 : f32 + %iq1s_code_40_5 = scalar.constant 25690.0 : f32 + %iq1s_code_40_6 = scalar.constant 25696.0 : f32 + %iq1s_code_40_7 = scalar.constant 25698.0 : f32 + %iq1s_code_40_8 = scalar.constant 25701.0 : f32 + %iq1s_code_40_9 = scalar.constant 25732.0 : f32 + %iq1s_code_40_10 = scalar.constant 25733.0 : f32 + %iq1s_code_40_11 = scalar.constant 25737.0 : f32 + %iq1s_code_40_12 = scalar.constant 25744.0 : f32 + %iq1s_code_40_13 = scalar.constant 25746.0 : f32 + %iq1s_code_40_14 = scalar.constant 25748.0 : f32 + %iq1s_code_40_15 = scalar.constant 25749.0 : f32 + %iq1s_code_40_16 = scalar.constant 25750.0 : f32 + %iq1s_code_40_17 = scalar.constant 25752.0 : f32 + %iq1s_code_40_18 = scalar.constant 25754.0 : f32 + %iq1s_code_40_19 = scalar.constant 25761.0 : f32 + %iq1s_code_40_20 = scalar.constant 25764.0 : f32 + %iq1s_code_40_21 = scalar.constant 25769.0 : f32 + %iq1s_code_40_22 = scalar.constant 25861.0 : f32 + %iq1s_code_40_23 = scalar.constant 25864.0 : f32 + %iq1s_code_40_24 = scalar.constant 25866.0 : f32 + %iq1s_code_40_25 = scalar.constant 25873.0 : f32 + %iq1s_code_40_26 = scalar.constant 25877.0 : f32 + %iq1s_code_40_27 = scalar.constant 25878.0 : f32 + %iq1s_code_40_28 = scalar.constant 25881.0 : f32 + %iq1s_code_40_29 = scalar.constant 25924.0 : f32 + %iq1s_code_40_30 = scalar.constant 25925.0 : f32 + %iq1s_code_40_31 = scalar.constant 25926.0 : f32 + %iq1s_code_40 = vector.from_elements %iq1s_code_40_0, %iq1s_code_40_1, %iq1s_code_40_2, %iq1s_code_40_3, %iq1s_code_40_4, %iq1s_code_40_5, %iq1s_code_40_6, %iq1s_code_40_7, %iq1s_code_40_8, %iq1s_code_40_9, %iq1s_code_40_10, %iq1s_code_40_11, %iq1s_code_40_12, %iq1s_code_40_13, %iq1s_code_40_14, %iq1s_code_40_15, %iq1s_code_40_16, %iq1s_code_40_17, %iq1s_code_40_18, %iq1s_code_40_19, %iq1s_code_40_20, %iq1s_code_40_21, %iq1s_code_40_22, %iq1s_code_40_23, %iq1s_code_40_24, %iq1s_code_40_25, %iq1s_code_40_26, %iq1s_code_40_27, %iq1s_code_40_28, %iq1s_code_40_29, %iq1s_code_40_30, %iq1s_code_40_31 : vector<32xf32> + %iq1s_code_41_0 = scalar.constant 25929.0 : f32 + %iq1s_code_41_1 = scalar.constant 25936.0 : f32 + %iq1s_code_41_2 = scalar.constant 25937.0 : f32 + %iq1s_code_41_3 = scalar.constant 25940.0 : f32 + %iq1s_code_41_4 = scalar.constant 25941.0 : f32 + %iq1s_code_41_5 = scalar.constant 25942.0 : f32 + %iq1s_code_41_6 = scalar.constant 25945.0 : f32 + %iq1s_code_41_7 = scalar.constant 25953.0 : f32 + %iq1s_code_41_8 = scalar.constant 25956.0 : f32 + %iq1s_code_41_9 = scalar.constant 25957.0 : f32 + %iq1s_code_41_10 = scalar.constant 25958.0 : f32 + %iq1s_code_41_11 = scalar.constant 25961.0 : f32 + %iq1s_code_41_12 = scalar.constant 25990.0 : f32 + %iq1s_code_41_13 = scalar.constant 25993.0 : f32 + %iq1s_code_41_14 = scalar.constant 25994.0 : f32 + %iq1s_code_41_15 = scalar.constant 26001.0 : f32 + %iq1s_code_41_16 = scalar.constant 26005.0 : f32 + %iq1s_code_41_17 = scalar.constant 26006.0 : f32 + %iq1s_code_41_18 = scalar.constant 26009.0 : f32 + %iq1s_code_41_19 = scalar.constant 26010.0 : f32 + %iq1s_code_41_20 = scalar.constant 26018.0 : f32 + %iq1s_code_41_21 = scalar.constant 26021.0 : f32 + %iq1s_code_41_22 = scalar.constant 26022.0 : f32 + %iq1s_code_41_23 = scalar.constant 26024.0 : f32 + %iq1s_code_41_24 = scalar.constant 26114.0 : f32 + %iq1s_code_41_25 = scalar.constant 26121.0 : f32 + %iq1s_code_41_26 = scalar.constant 26133.0 : f32 + %iq1s_code_41_27 = scalar.constant 26144.0 : f32 + %iq1s_code_41_28 = scalar.constant 26150.0 : f32 + %iq1s_code_41_29 = scalar.constant 26152.0 : f32 + %iq1s_code_41_30 = scalar.constant 26153.0 : f32 + %iq1s_code_41_31 = scalar.constant 26176.0 : f32 + %iq1s_code_41 = vector.from_elements %iq1s_code_41_0, %iq1s_code_41_1, %iq1s_code_41_2, %iq1s_code_41_3, %iq1s_code_41_4, %iq1s_code_41_5, %iq1s_code_41_6, %iq1s_code_41_7, %iq1s_code_41_8, %iq1s_code_41_9, %iq1s_code_41_10, %iq1s_code_41_11, %iq1s_code_41_12, %iq1s_code_41_13, %iq1s_code_41_14, %iq1s_code_41_15, %iq1s_code_41_16, %iq1s_code_41_17, %iq1s_code_41_18, %iq1s_code_41_19, %iq1s_code_41_20, %iq1s_code_41_21, %iq1s_code_41_22, %iq1s_code_41_23, %iq1s_code_41_24, %iq1s_code_41_25, %iq1s_code_41_26, %iq1s_code_41_27, %iq1s_code_41_28, %iq1s_code_41_29, %iq1s_code_41_30, %iq1s_code_41_31 : vector<32xf32> + %iq1s_code_42_0 = scalar.constant 26181.0 : f32 + %iq1s_code_42_1 = scalar.constant 26184.0 : f32 + %iq1s_code_42_2 = scalar.constant 26186.0 : f32 + %iq1s_code_42_3 = scalar.constant 26193.0 : f32 + %iq1s_code_42_4 = scalar.constant 26196.0 : f32 + %iq1s_code_42_5 = scalar.constant 26197.0 : f32 + %iq1s_code_42_6 = scalar.constant 26198.0 : f32 + %iq1s_code_42_7 = scalar.constant 26200.0 : f32 + %iq1s_code_42_8 = scalar.constant 26202.0 : f32 + %iq1s_code_42_9 = scalar.constant 26208.0 : f32 + %iq1s_code_42_10 = scalar.constant 26213.0 : f32 + %iq1s_code_42_11 = scalar.constant 26216.0 : f32 + %iq1s_code_42_12 = scalar.constant 26240.0 : f32 + %iq1s_code_42_13 = scalar.constant 26242.0 : f32 + %iq1s_code_42_14 = scalar.constant 26245.0 : f32 + %iq1s_code_42_15 = scalar.constant 26250.0 : f32 + %iq1s_code_42_16 = scalar.constant 26260.0 : f32 + %iq1s_code_42_17 = scalar.constant 26262.0 : f32 + %iq1s_code_42_18 = scalar.constant 26264.0 : f32 + %iq1s_code_42_19 = scalar.constant 26265.0 : f32 + %iq1s_code_42_20 = scalar.constant 26272.0 : f32 + %iq1s_code_42_21 = scalar.constant 26276.0 : f32 + %iq1s_code_42_22 = scalar.constant 26278.0 : f32 + %iq1s_code_42_23 = scalar.constant 26282.0 : f32 + %iq1s_code_42_24 = scalar.constant 26646.0 : f32 + %iq1s_code_42_25 = scalar.constant 26649.0 : f32 + %iq1s_code_42_26 = scalar.constant 26661.0 : f32 + %iq1s_code_42_27 = scalar.constant 26689.0 : f32 + %iq1s_code_42_28 = scalar.constant 26706.0 : f32 + %iq1s_code_42_29 = scalar.constant 26709.0 : f32 + %iq1s_code_42_30 = scalar.constant 26714.0 : f32 + %iq1s_code_42_31 = scalar.constant 26721.0 : f32 + %iq1s_code_42 = vector.from_elements %iq1s_code_42_0, %iq1s_code_42_1, %iq1s_code_42_2, %iq1s_code_42_3, %iq1s_code_42_4, %iq1s_code_42_5, %iq1s_code_42_6, %iq1s_code_42_7, %iq1s_code_42_8, %iq1s_code_42_9, %iq1s_code_42_10, %iq1s_code_42_11, %iq1s_code_42_12, %iq1s_code_42_13, %iq1s_code_42_14, %iq1s_code_42_15, %iq1s_code_42_16, %iq1s_code_42_17, %iq1s_code_42_18, %iq1s_code_42_19, %iq1s_code_42_20, %iq1s_code_42_21, %iq1s_code_42_22, %iq1s_code_42_23, %iq1s_code_42_24, %iq1s_code_42_25, %iq1s_code_42_26, %iq1s_code_42_27, %iq1s_code_42_28, %iq1s_code_42_29, %iq1s_code_42_30, %iq1s_code_42_31 : vector<32xf32> + %iq1s_code_43_0 = scalar.constant 26729.0 : f32 + %iq1s_code_43_1 = scalar.constant 26757.0 : f32 + %iq1s_code_43_2 = scalar.constant 26769.0 : f32 + %iq1s_code_43_3 = scalar.constant 26776.0 : f32 + %iq1s_code_43_4 = scalar.constant 26790.0 : f32 + %iq1s_code_43_5 = scalar.constant 26881.0 : f32 + %iq1s_code_43_6 = scalar.constant 26884.0 : f32 + %iq1s_code_43_7 = scalar.constant 26896.0 : f32 + %iq1s_code_43_8 = scalar.constant 26901.0 : f32 + %iq1s_code_43_9 = scalar.constant 26913.0 : f32 + %iq1s_code_43_10 = scalar.constant 26916.0 : f32 + %iq1s_code_43_11 = scalar.constant 26918.0 : f32 + %iq1s_code_43_12 = scalar.constant 26921.0 : f32 + %iq1s_code_43_13 = scalar.constant 26944.0 : f32 + %iq1s_code_43_14 = scalar.constant 26945.0 : f32 + %iq1s_code_43_15 = scalar.constant 26949.0 : f32 + %iq1s_code_43_16 = scalar.constant 26950.0 : f32 + %iq1s_code_43_17 = scalar.constant 26952.0 : f32 + %iq1s_code_43_18 = scalar.constant 26961.0 : f32 + %iq1s_code_43_19 = scalar.constant 26964.0 : f32 + %iq1s_code_43_20 = scalar.constant 26965.0 : f32 + %iq1s_code_43_21 = scalar.constant 26966.0 : f32 + %iq1s_code_43_22 = scalar.constant 26969.0 : f32 + %iq1s_code_43_23 = scalar.constant 26976.0 : f32 + %iq1s_code_43_24 = scalar.constant 26981.0 : f32 + %iq1s_code_43_25 = scalar.constant 26986.0 : f32 + %iq1s_code_43_26 = scalar.constant 27010.0 : f32 + %iq1s_code_43_27 = scalar.constant 27012.0 : f32 + %iq1s_code_43_28 = scalar.constant 27018.0 : f32 + %iq1s_code_43_29 = scalar.constant 27029.0 : f32 + %iq1s_code_43_30 = scalar.constant 27041.0 : f32 + %iq1s_code_43_31 = scalar.constant 27044.0 : f32 + %iq1s_code_43 = vector.from_elements %iq1s_code_43_0, %iq1s_code_43_1, %iq1s_code_43_2, %iq1s_code_43_3, %iq1s_code_43_4, %iq1s_code_43_5, %iq1s_code_43_6, %iq1s_code_43_7, %iq1s_code_43_8, %iq1s_code_43_9, %iq1s_code_43_10, %iq1s_code_43_11, %iq1s_code_43_12, %iq1s_code_43_13, %iq1s_code_43_14, %iq1s_code_43_15, %iq1s_code_43_16, %iq1s_code_43_17, %iq1s_code_43_18, %iq1s_code_43_19, %iq1s_code_43_20, %iq1s_code_43_21, %iq1s_code_43_22, %iq1s_code_43_23, %iq1s_code_43_24, %iq1s_code_43_25, %iq1s_code_43_26, %iq1s_code_43_27, %iq1s_code_43_28, %iq1s_code_43_29, %iq1s_code_43_30, %iq1s_code_43_31 : vector<32xf32> + %iq1s_code_44_0 = scalar.constant 27045.0 : f32 + %iq1s_code_44_1 = scalar.constant 27049.0 : f32 + %iq1s_code_44_2 = scalar.constant 27153.0 : f32 + %iq1s_code_44_3 = scalar.constant 27158.0 : f32 + %iq1s_code_44_4 = scalar.constant 27160.0 : f32 + %iq1s_code_44_5 = scalar.constant 27201.0 : f32 + %iq1s_code_44_6 = scalar.constant 27204.0 : f32 + %iq1s_code_44_7 = scalar.constant 27209.0 : f32 + %iq1s_code_44_8 = scalar.constant 27216.0 : f32 + %iq1s_code_44_9 = scalar.constant 27221.0 : f32 + %iq1s_code_44_10 = scalar.constant 27224.0 : f32 + %iq1s_code_44_11 = scalar.constant 27226.0 : f32 + %iq1s_code_44_12 = scalar.constant 27236.0 : f32 + %iq1s_code_44_13 = scalar.constant 27237.0 : f32 + %iq1s_code_44_14 = scalar.constant 27241.0 : f32 + %iq1s_code_44_15 = scalar.constant 27270.0 : f32 + %iq1s_code_44_16 = scalar.constant 27284.0 : f32 + %iq1s_code_44_17 = scalar.constant 27288.0 : f32 + %iq1s_code_44_18 = scalar.constant 27290.0 : f32 + %iq1s_code_44_19 = scalar.constant 27302.0 : f32 + %iq1s_code_44_20 = scalar.constant 32768.0 : f32 + %iq1s_code_44_21 = scalar.constant 32770.0 : f32 + %iq1s_code_44_22 = scalar.constant 32776.0 : f32 + %iq1s_code_44_23 = scalar.constant 32778.0 : f32 + %iq1s_code_44_24 = scalar.constant 32800.0 : f32 + %iq1s_code_44_25 = scalar.constant 32802.0 : f32 + %iq1s_code_44_26 = scalar.constant 32808.0 : f32 + %iq1s_code_44_27 = scalar.constant 32810.0 : f32 + %iq1s_code_44_28 = scalar.constant 32837.0 : f32 + %iq1s_code_44_29 = scalar.constant 32848.0 : f32 + %iq1s_code_44_30 = scalar.constant 32849.0 : f32 + %iq1s_code_44_31 = scalar.constant 32852.0 : f32 + %iq1s_code_44 = vector.from_elements %iq1s_code_44_0, %iq1s_code_44_1, %iq1s_code_44_2, %iq1s_code_44_3, %iq1s_code_44_4, %iq1s_code_44_5, %iq1s_code_44_6, %iq1s_code_44_7, %iq1s_code_44_8, %iq1s_code_44_9, %iq1s_code_44_10, %iq1s_code_44_11, %iq1s_code_44_12, %iq1s_code_44_13, %iq1s_code_44_14, %iq1s_code_44_15, %iq1s_code_44_16, %iq1s_code_44_17, %iq1s_code_44_18, %iq1s_code_44_19, %iq1s_code_44_20, %iq1s_code_44_21, %iq1s_code_44_22, %iq1s_code_44_23, %iq1s_code_44_24, %iq1s_code_44_25, %iq1s_code_44_26, %iq1s_code_44_27, %iq1s_code_44_28, %iq1s_code_44_29, %iq1s_code_44_30, %iq1s_code_44_31 : vector<32xf32> + %iq1s_code_45_0 = scalar.constant 32854.0 : f32 + %iq1s_code_45_1 = scalar.constant 32857.0 : f32 + %iq1s_code_45_2 = scalar.constant 32869.0 : f32 + %iq1s_code_45_3 = scalar.constant 32896.0 : f32 + %iq1s_code_45_4 = scalar.constant 32898.0 : f32 + %iq1s_code_45_5 = scalar.constant 32904.0 : f32 + %iq1s_code_45_6 = scalar.constant 32906.0 : f32 + %iq1s_code_45_7 = scalar.constant 32917.0 : f32 + %iq1s_code_45_8 = scalar.constant 32928.0 : f32 + %iq1s_code_45_9 = scalar.constant 32930.0 : f32 + %iq1s_code_45_10 = scalar.constant 32936.0 : f32 + %iq1s_code_45_11 = scalar.constant 32938.0 : f32 + %iq1s_code_45_12 = scalar.constant 33029.0 : f32 + %iq1s_code_45_13 = scalar.constant 33041.0 : f32 + %iq1s_code_45_14 = scalar.constant 33044.0 : f32 + %iq1s_code_45_15 = scalar.constant 33046.0 : f32 + %iq1s_code_45_16 = scalar.constant 33049.0 : f32 + %iq1s_code_45_17 = scalar.constant 33061.0 : f32 + %iq1s_code_45_18 = scalar.constant 33089.0 : f32 + %iq1s_code_45_19 = scalar.constant 33092.0 : f32 + %iq1s_code_45_20 = scalar.constant 33097.0 : f32 + %iq1s_code_45_21 = scalar.constant 33104.0 : f32 + %iq1s_code_45_22 = scalar.constant 33106.0 : f32 + %iq1s_code_45_23 = scalar.constant 33109.0 : f32 + %iq1s_code_45_24 = scalar.constant 33110.0 : f32 + %iq1s_code_45_25 = scalar.constant 33112.0 : f32 + %iq1s_code_45_26 = scalar.constant 33113.0 : f32 + %iq1s_code_45_27 = scalar.constant 33124.0 : f32 + %iq1s_code_45_28 = scalar.constant 33126.0 : f32 + %iq1s_code_45_29 = scalar.constant 33129.0 : f32 + %iq1s_code_45_30 = scalar.constant 33157.0 : f32 + %iq1s_code_45_31 = scalar.constant 33161.0 : f32 + %iq1s_code_45 = vector.from_elements %iq1s_code_45_0, %iq1s_code_45_1, %iq1s_code_45_2, %iq1s_code_45_3, %iq1s_code_45_4, %iq1s_code_45_5, %iq1s_code_45_6, %iq1s_code_45_7, %iq1s_code_45_8, %iq1s_code_45_9, %iq1s_code_45_10, %iq1s_code_45_11, %iq1s_code_45_12, %iq1s_code_45_13, %iq1s_code_45_14, %iq1s_code_45_15, %iq1s_code_45_16, %iq1s_code_45_17, %iq1s_code_45_18, %iq1s_code_45_19, %iq1s_code_45_20, %iq1s_code_45_21, %iq1s_code_45_22, %iq1s_code_45_23, %iq1s_code_45_24, %iq1s_code_45_25, %iq1s_code_45_26, %iq1s_code_45_27, %iq1s_code_45_28, %iq1s_code_45_29, %iq1s_code_45_30, %iq1s_code_45_31 : vector<32xf32> + %iq1s_code_46_0 = scalar.constant 33172.0 : f32 + %iq1s_code_46_1 = scalar.constant 33174.0 : f32 + %iq1s_code_46_2 = scalar.constant 33177.0 : f32 + %iq1s_code_46_3 = scalar.constant 33189.0 : f32 + %iq1s_code_46_4 = scalar.constant 33280.0 : f32 + %iq1s_code_46_5 = scalar.constant 33282.0 : f32 + %iq1s_code_46_6 = scalar.constant 33288.0 : f32 + %iq1s_code_46_7 = scalar.constant 33290.0 : f32 + %iq1s_code_46_8 = scalar.constant 33301.0 : f32 + %iq1s_code_46_9 = scalar.constant 33312.0 : f32 + %iq1s_code_46_10 = scalar.constant 33314.0 : f32 + %iq1s_code_46_11 = scalar.constant 33320.0 : f32 + %iq1s_code_46_12 = scalar.constant 33322.0 : f32 + %iq1s_code_46_13 = scalar.constant 33361.0 : f32 + %iq1s_code_46_14 = scalar.constant 33364.0 : f32 + %iq1s_code_46_15 = scalar.constant 33369.0 : f32 + %iq1s_code_46_16 = scalar.constant 33381.0 : f32 + %iq1s_code_46_17 = scalar.constant 33408.0 : f32 + %iq1s_code_46_18 = scalar.constant 33410.0 : f32 + %iq1s_code_46_19 = scalar.constant 33416.0 : f32 + %iq1s_code_46_20 = scalar.constant 33418.0 : f32 + %iq1s_code_46_21 = scalar.constant 33429.0 : f32 + %iq1s_code_46_22 = scalar.constant 33440.0 : f32 + %iq1s_code_46_23 = scalar.constant 33442.0 : f32 + %iq1s_code_46_24 = scalar.constant 33448.0 : f32 + %iq1s_code_46_25 = scalar.constant 33450.0 : f32 + %iq1s_code_46_26 = scalar.constant 33812.0 : f32 + %iq1s_code_46_27 = scalar.constant 33817.0 : f32 + %iq1s_code_46_28 = scalar.constant 33857.0 : f32 + %iq1s_code_46_29 = scalar.constant 33860.0 : f32 + %iq1s_code_46_30 = scalar.constant 33873.0 : f32 + %iq1s_code_46_31 = scalar.constant 33877.0 : f32 + %iq1s_code_46 = vector.from_elements %iq1s_code_46_0, %iq1s_code_46_1, %iq1s_code_46_2, %iq1s_code_46_3, %iq1s_code_46_4, %iq1s_code_46_5, %iq1s_code_46_6, %iq1s_code_46_7, %iq1s_code_46_8, %iq1s_code_46_9, %iq1s_code_46_10, %iq1s_code_46_11, %iq1s_code_46_12, %iq1s_code_46_13, %iq1s_code_46_14, %iq1s_code_46_15, %iq1s_code_46_16, %iq1s_code_46_17, %iq1s_code_46_18, %iq1s_code_46_19, %iq1s_code_46_20, %iq1s_code_46_21, %iq1s_code_46_22, %iq1s_code_46_23, %iq1s_code_46_24, %iq1s_code_46_25, %iq1s_code_46_26, %iq1s_code_46_27, %iq1s_code_46_28, %iq1s_code_46_29, %iq1s_code_46_30, %iq1s_code_46_31 : vector<32xf32> + %iq1s_code_47_0 = scalar.constant 33882.0 : f32 + %iq1s_code_47_1 = scalar.constant 33889.0 : f32 + %iq1s_code_47_2 = scalar.constant 33892.0 : f32 + %iq1s_code_47_3 = scalar.constant 33897.0 : f32 + %iq1s_code_47_4 = scalar.constant 33940.0 : f32 + %iq1s_code_47_5 = scalar.constant 33945.0 : f32 + %iq1s_code_47_6 = scalar.constant 34049.0 : f32 + %iq1s_code_47_7 = scalar.constant 34057.0 : f32 + %iq1s_code_47_8 = scalar.constant 34066.0 : f32 + %iq1s_code_47_9 = scalar.constant 34069.0 : f32 + %iq1s_code_47_10 = scalar.constant 34074.0 : f32 + %iq1s_code_47_11 = scalar.constant 34086.0 : f32 + %iq1s_code_47_12 = scalar.constant 34089.0 : f32 + %iq1s_code_47_13 = scalar.constant 34112.0 : f32 + %iq1s_code_47_14 = scalar.constant 34113.0 : f32 + %iq1s_code_47_15 = scalar.constant 34117.0 : f32 + %iq1s_code_47_16 = scalar.constant 34120.0 : f32 + %iq1s_code_47_17 = scalar.constant 34129.0 : f32 + %iq1s_code_47_18 = scalar.constant 34132.0 : f32 + %iq1s_code_47_19 = scalar.constant 34133.0 : f32 + %iq1s_code_47_20 = scalar.constant 34134.0 : f32 + %iq1s_code_47_21 = scalar.constant 34137.0 : f32 + %iq1s_code_47_22 = scalar.constant 34138.0 : f32 + %iq1s_code_47_23 = scalar.constant 34149.0 : f32 + %iq1s_code_47_24 = scalar.constant 34150.0 : f32 + %iq1s_code_47_25 = scalar.constant 34152.0 : f32 + %iq1s_code_47_26 = scalar.constant 34154.0 : f32 + %iq1s_code_47_27 = scalar.constant 34177.0 : f32 + %iq1s_code_47_28 = scalar.constant 34180.0 : f32 + %iq1s_code_47_29 = scalar.constant 34182.0 : f32 + %iq1s_code_47_30 = scalar.constant 34185.0 : f32 + %iq1s_code_47_31 = scalar.constant 34192.0 : f32 + %iq1s_code_47 = vector.from_elements %iq1s_code_47_0, %iq1s_code_47_1, %iq1s_code_47_2, %iq1s_code_47_3, %iq1s_code_47_4, %iq1s_code_47_5, %iq1s_code_47_6, %iq1s_code_47_7, %iq1s_code_47_8, %iq1s_code_47_9, %iq1s_code_47_10, %iq1s_code_47_11, %iq1s_code_47_12, %iq1s_code_47_13, %iq1s_code_47_14, %iq1s_code_47_15, %iq1s_code_47_16, %iq1s_code_47_17, %iq1s_code_47_18, %iq1s_code_47_19, %iq1s_code_47_20, %iq1s_code_47_21, %iq1s_code_47_22, %iq1s_code_47_23, %iq1s_code_47_24, %iq1s_code_47_25, %iq1s_code_47_26, %iq1s_code_47_27, %iq1s_code_47_28, %iq1s_code_47_29, %iq1s_code_47_30, %iq1s_code_47_31 : vector<32xf32> + %iq1s_code_48_0 = scalar.constant 34194.0 : f32 + %iq1s_code_48_1 = scalar.constant 34197.0 : f32 + %iq1s_code_48_2 = scalar.constant 34200.0 : f32 + %iq1s_code_48_3 = scalar.constant 34214.0 : f32 + %iq1s_code_48_4 = scalar.constant 34321.0 : f32 + %iq1s_code_48_5 = scalar.constant 34326.0 : f32 + %iq1s_code_48_6 = scalar.constant 34329.0 : f32 + %iq1s_code_48_7 = scalar.constant 34341.0 : f32 + %iq1s_code_48_8 = scalar.constant 34369.0 : f32 + %iq1s_code_48_9 = scalar.constant 34372.0 : f32 + %iq1s_code_48_10 = scalar.constant 34377.0 : f32 + %iq1s_code_48_11 = scalar.constant 34378.0 : f32 + %iq1s_code_48_12 = scalar.constant 34384.0 : f32 + %iq1s_code_48_13 = scalar.constant 34389.0 : f32 + %iq1s_code_48_14 = scalar.constant 34393.0 : f32 + %iq1s_code_48_15 = scalar.constant 34394.0 : f32 + %iq1s_code_48_16 = scalar.constant 34401.0 : f32 + %iq1s_code_48_17 = scalar.constant 34406.0 : f32 + %iq1s_code_48_18 = scalar.constant 34410.0 : f32 + %iq1s_code_48_19 = scalar.constant 34437.0 : f32 + %iq1s_code_48_20 = scalar.constant 34449.0 : f32 + %iq1s_code_48_21 = scalar.constant 34458.0 : f32 + %iq1s_code_48_22 = scalar.constant 34468.0 : f32 + %iq1s_code_48_23 = scalar.constant 34816.0 : f32 + %iq1s_code_48_24 = scalar.constant 34818.0 : f32 + %iq1s_code_48_25 = scalar.constant 34824.0 : f32 + %iq1s_code_48_26 = scalar.constant 34826.0 : f32 + %iq1s_code_48_27 = scalar.constant 34837.0 : f32 + %iq1s_code_48_28 = scalar.constant 34848.0 : f32 + %iq1s_code_48_29 = scalar.constant 34850.0 : f32 + %iq1s_code_48_30 = scalar.constant 34856.0 : f32 + %iq1s_code_48_31 = scalar.constant 34858.0 : f32 + %iq1s_code_48 = vector.from_elements %iq1s_code_48_0, %iq1s_code_48_1, %iq1s_code_48_2, %iq1s_code_48_3, %iq1s_code_48_4, %iq1s_code_48_5, %iq1s_code_48_6, %iq1s_code_48_7, %iq1s_code_48_8, %iq1s_code_48_9, %iq1s_code_48_10, %iq1s_code_48_11, %iq1s_code_48_12, %iq1s_code_48_13, %iq1s_code_48_14, %iq1s_code_48_15, %iq1s_code_48_16, %iq1s_code_48_17, %iq1s_code_48_18, %iq1s_code_48_19, %iq1s_code_48_20, %iq1s_code_48_21, %iq1s_code_48_22, %iq1s_code_48_23, %iq1s_code_48_24, %iq1s_code_48_25, %iq1s_code_48_26, %iq1s_code_48_27, %iq1s_code_48_28, %iq1s_code_48_29, %iq1s_code_48_30, %iq1s_code_48_31 : vector<32xf32> + %iq1s_code_49_0 = scalar.constant 34881.0 : f32 + %iq1s_code_49_1 = scalar.constant 34885.0 : f32 + %iq1s_code_49_2 = scalar.constant 34897.0 : f32 + %iq1s_code_49_3 = scalar.constant 34900.0 : f32 + %iq1s_code_49_4 = scalar.constant 34905.0 : f32 + %iq1s_code_49_5 = scalar.constant 34917.0 : f32 + %iq1s_code_49_6 = scalar.constant 34921.0 : f32 + %iq1s_code_49_7 = scalar.constant 34944.0 : f32 + %iq1s_code_49_8 = scalar.constant 34946.0 : f32 + %iq1s_code_49_9 = scalar.constant 34952.0 : f32 + %iq1s_code_49_10 = scalar.constant 34954.0 : f32 + %iq1s_code_49_11 = scalar.constant 34965.0 : f32 + %iq1s_code_49_12 = scalar.constant 34976.0 : f32 + %iq1s_code_49_13 = scalar.constant 34978.0 : f32 + %iq1s_code_49_14 = scalar.constant 34984.0 : f32 + %iq1s_code_49_15 = scalar.constant 34986.0 : f32 + %iq1s_code_49_16 = scalar.constant 35077.0 : f32 + %iq1s_code_49_17 = scalar.constant 35078.0 : f32 + %iq1s_code_49_18 = scalar.constant 35089.0 : f32 + %iq1s_code_49_19 = scalar.constant 35092.0 : f32 + %iq1s_code_49_20 = scalar.constant 35094.0 : f32 + %iq1s_code_49_21 = scalar.constant 35109.0 : f32 + %iq1s_code_49_22 = scalar.constant 35137.0 : f32 + %iq1s_code_49_23 = scalar.constant 35140.0 : f32 + %iq1s_code_49_24 = scalar.constant 35142.0 : f32 + %iq1s_code_49_25 = scalar.constant 35145.0 : f32 + %iq1s_code_49_26 = scalar.constant 35152.0 : f32 + %iq1s_code_49_27 = scalar.constant 35154.0 : f32 + %iq1s_code_49_28 = scalar.constant 35157.0 : f32 + %iq1s_code_49_29 = scalar.constant 35162.0 : f32 + %iq1s_code_49_30 = scalar.constant 35169.0 : f32 + %iq1s_code_49_31 = scalar.constant 35172.0 : f32 + %iq1s_code_49 = vector.from_elements %iq1s_code_49_0, %iq1s_code_49_1, %iq1s_code_49_2, %iq1s_code_49_3, %iq1s_code_49_4, %iq1s_code_49_5, %iq1s_code_49_6, %iq1s_code_49_7, %iq1s_code_49_8, %iq1s_code_49_9, %iq1s_code_49_10, %iq1s_code_49_11, %iq1s_code_49_12, %iq1s_code_49_13, %iq1s_code_49_14, %iq1s_code_49_15, %iq1s_code_49_16, %iq1s_code_49_17, %iq1s_code_49_18, %iq1s_code_49_19, %iq1s_code_49_20, %iq1s_code_49_21, %iq1s_code_49_22, %iq1s_code_49_23, %iq1s_code_49_24, %iq1s_code_49_25, %iq1s_code_49_26, %iq1s_code_49_27, %iq1s_code_49_28, %iq1s_code_49_29, %iq1s_code_49_30, %iq1s_code_49_31 : vector<32xf32> + %iq1s_code_50_0 = scalar.constant 35205.0 : f32 + %iq1s_code_50_1 = scalar.constant 35222.0 : f32 + %iq1s_code_50_2 = scalar.constant 35225.0 : f32 + %iq1s_code_50_3 = scalar.constant 35237.0 : f32 + %iq1s_code_50_4 = scalar.constant 35328.0 : f32 + %iq1s_code_50_5 = scalar.constant 35330.0 : f32 + %iq1s_code_50_6 = scalar.constant 35336.0 : f32 + %iq1s_code_50_7 = scalar.constant 35338.0 : f32 + %iq1s_code_50_8 = scalar.constant 35349.0 : f32 + %iq1s_code_50_9 = scalar.constant 35360.0 : f32 + %iq1s_code_50_10 = scalar.constant 35362.0 : f32 + %iq1s_code_50_11 = scalar.constant 35368.0 : f32 + %iq1s_code_50_12 = scalar.constant 35370.0 : f32 + %iq1s_code_50_13 = scalar.constant 35397.0 : f32 + %iq1s_code_50_14 = scalar.constant 35409.0 : f32 + %iq1s_code_50_15 = scalar.constant 35412.0 : f32 + %iq1s_code_50_16 = scalar.constant 35414.0 : f32 + %iq1s_code_50_17 = scalar.constant 35456.0 : f32 + %iq1s_code_50_18 = scalar.constant 35458.0 : f32 + %iq1s_code_50_19 = scalar.constant 35464.0 : f32 + %iq1s_code_50_20 = scalar.constant 35466.0 : f32 + %iq1s_code_50_21 = scalar.constant 35477.0 : f32 + %iq1s_code_50_22 = scalar.constant 35488.0 : f32 + %iq1s_code_50_23 = scalar.constant 35490.0 : f32 + %iq1s_code_50_24 = scalar.constant 35496.0 : f32 + %iq1s_code_50_25 = scalar.constant 35498.0 : f32 + %iq1s_code_50_26 = scalar.constant 36869.0 : f32 + %iq1s_code_50_27 = scalar.constant 36881.0 : f32 + %iq1s_code_50_28 = scalar.constant 36886.0 : f32 + %iq1s_code_50_29 = scalar.constant 36888.0 : f32 + %iq1s_code_50_30 = scalar.constant 36889.0 : f32 + %iq1s_code_50_31 = scalar.constant 36901.0 : f32 + %iq1s_code_50 = vector.from_elements %iq1s_code_50_0, %iq1s_code_50_1, %iq1s_code_50_2, %iq1s_code_50_3, %iq1s_code_50_4, %iq1s_code_50_5, %iq1s_code_50_6, %iq1s_code_50_7, %iq1s_code_50_8, %iq1s_code_50_9, %iq1s_code_50_10, %iq1s_code_50_11, %iq1s_code_50_12, %iq1s_code_50_13, %iq1s_code_50_14, %iq1s_code_50_15, %iq1s_code_50_16, %iq1s_code_50_17, %iq1s_code_50_18, %iq1s_code_50_19, %iq1s_code_50_20, %iq1s_code_50_21, %iq1s_code_50_22, %iq1s_code_50_23, %iq1s_code_50_24, %iq1s_code_50_25, %iq1s_code_50_26, %iq1s_code_50_27, %iq1s_code_50_28, %iq1s_code_50_29, %iq1s_code_50_30, %iq1s_code_50_31 : vector<32xf32> + %iq1s_code_51_0 = scalar.constant 36929.0 : f32 + %iq1s_code_51_1 = scalar.constant 36934.0 : f32 + %iq1s_code_51_2 = scalar.constant 36937.0 : f32 + %iq1s_code_51_3 = scalar.constant 36949.0 : f32 + %iq1s_code_51_4 = scalar.constant 36952.0 : f32 + %iq1s_code_51_5 = scalar.constant 36954.0 : f32 + %iq1s_code_51_6 = scalar.constant 36969.0 : f32 + %iq1s_code_51_7 = scalar.constant 36970.0 : f32 + %iq1s_code_51_8 = scalar.constant 36997.0 : f32 + %iq1s_code_51_9 = scalar.constant 37009.0 : f32 + %iq1s_code_51_10 = scalar.constant 37012.0 : f32 + %iq1s_code_51_11 = scalar.constant 37014.0 : f32 + %iq1s_code_51_12 = scalar.constant 37017.0 : f32 + %iq1s_code_51_13 = scalar.constant 37029.0 : f32 + %iq1s_code_51_14 = scalar.constant 37121.0 : f32 + %iq1s_code_51_15 = scalar.constant 37124.0 : f32 + %iq1s_code_51_16 = scalar.constant 37126.0 : f32 + %iq1s_code_51_17 = scalar.constant 37129.0 : f32 + %iq1s_code_51_18 = scalar.constant 37136.0 : f32 + %iq1s_code_51_19 = scalar.constant 37141.0 : f32 + %iq1s_code_51_20 = scalar.constant 37144.0 : f32 + %iq1s_code_51_21 = scalar.constant 37146.0 : f32 + %iq1s_code_51_22 = scalar.constant 37153.0 : f32 + %iq1s_code_51_23 = scalar.constant 37156.0 : f32 + %iq1s_code_51_24 = scalar.constant 37158.0 : f32 + %iq1s_code_51_25 = scalar.constant 37161.0 : f32 + %iq1s_code_51_26 = scalar.constant 37184.0 : f32 + %iq1s_code_51_27 = scalar.constant 37189.0 : f32 + %iq1s_code_51_28 = scalar.constant 37200.0 : f32 + %iq1s_code_51_29 = scalar.constant 37201.0 : f32 + %iq1s_code_51_30 = scalar.constant 37204.0 : f32 + %iq1s_code_51_31 = scalar.constant 37205.0 : f32 + %iq1s_code_51 = vector.from_elements %iq1s_code_51_0, %iq1s_code_51_1, %iq1s_code_51_2, %iq1s_code_51_3, %iq1s_code_51_4, %iq1s_code_51_5, %iq1s_code_51_6, %iq1s_code_51_7, %iq1s_code_51_8, %iq1s_code_51_9, %iq1s_code_51_10, %iq1s_code_51_11, %iq1s_code_51_12, %iq1s_code_51_13, %iq1s_code_51_14, %iq1s_code_51_15, %iq1s_code_51_16, %iq1s_code_51_17, %iq1s_code_51_18, %iq1s_code_51_19, %iq1s_code_51_20, %iq1s_code_51_21, %iq1s_code_51_22, %iq1s_code_51_23, %iq1s_code_51_24, %iq1s_code_51_25, %iq1s_code_51_26, %iq1s_code_51_27, %iq1s_code_51_28, %iq1s_code_51_29, %iq1s_code_51_30, %iq1s_code_51_31 : vector<32xf32> + %iq1s_code_52_0 = scalar.constant 37206.0 : f32 + %iq1s_code_52_1 = scalar.constant 37209.0 : f32 + %iq1s_code_52_2 = scalar.constant 37218.0 : f32 + %iq1s_code_52_3 = scalar.constant 37221.0 : f32 + %iq1s_code_52_4 = scalar.constant 37252.0 : f32 + %iq1s_code_52_5 = scalar.constant 37254.0 : f32 + %iq1s_code_52_6 = scalar.constant 37266.0 : f32 + %iq1s_code_52_7 = scalar.constant 37269.0 : f32 + %iq1s_code_52_8 = scalar.constant 37272.0 : f32 + %iq1s_code_52_9 = scalar.constant 37281.0 : f32 + %iq1s_code_52_10 = scalar.constant 37284.0 : f32 + %iq1s_code_52_11 = scalar.constant 37286.0 : f32 + %iq1s_code_52_12 = scalar.constant 37289.0 : f32 + %iq1s_code_52_13 = scalar.constant 37381.0 : f32 + %iq1s_code_52_14 = scalar.constant 37393.0 : f32 + %iq1s_code_52_15 = scalar.constant 37396.0 : f32 + %iq1s_code_52_16 = scalar.constant 37401.0 : f32 + %iq1s_code_52_17 = scalar.constant 37413.0 : f32 + %iq1s_code_52_18 = scalar.constant 37444.0 : f32 + %iq1s_code_52_19 = scalar.constant 37446.0 : f32 + %iq1s_code_52_20 = scalar.constant 37449.0 : f32 + %iq1s_code_52_21 = scalar.constant 37456.0 : f32 + %iq1s_code_52_22 = scalar.constant 37458.0 : f32 + %iq1s_code_52_23 = scalar.constant 37461.0 : f32 + %iq1s_code_52_24 = scalar.constant 37464.0 : f32 + %iq1s_code_52_25 = scalar.constant 37478.0 : f32 + %iq1s_code_52_26 = scalar.constant 37481.0 : f32 + %iq1s_code_52_27 = scalar.constant 37509.0 : f32 + %iq1s_code_52_28 = scalar.constant 37524.0 : f32 + %iq1s_code_52_29 = scalar.constant 37526.0 : f32 + %iq1s_code_52_30 = scalar.constant 37545.0 : f32 + %iq1s_code_52_31 = scalar.constant 37889.0 : f32 + %iq1s_code_52 = vector.from_elements %iq1s_code_52_0, %iq1s_code_52_1, %iq1s_code_52_2, %iq1s_code_52_3, %iq1s_code_52_4, %iq1s_code_52_5, %iq1s_code_52_6, %iq1s_code_52_7, %iq1s_code_52_8, %iq1s_code_52_9, %iq1s_code_52_10, %iq1s_code_52_11, %iq1s_code_52_12, %iq1s_code_52_13, %iq1s_code_52_14, %iq1s_code_52_15, %iq1s_code_52_16, %iq1s_code_52_17, %iq1s_code_52_18, %iq1s_code_52_19, %iq1s_code_52_20, %iq1s_code_52_21, %iq1s_code_52_22, %iq1s_code_52_23, %iq1s_code_52_24, %iq1s_code_52_25, %iq1s_code_52_26, %iq1s_code_52_27, %iq1s_code_52_28, %iq1s_code_52_29, %iq1s_code_52_30, %iq1s_code_52_31 : vector<32xf32> + %iq1s_code_53_0 = scalar.constant 37892.0 : f32 + %iq1s_code_53_1 = scalar.constant 37894.0 : f32 + %iq1s_code_53_2 = scalar.constant 37904.0 : f32 + %iq1s_code_53_3 = scalar.constant 37909.0 : f32 + %iq1s_code_53_4 = scalar.constant 37912.0 : f32 + %iq1s_code_53_5 = scalar.constant 37926.0 : f32 + %iq1s_code_53_6 = scalar.constant 37952.0 : f32 + %iq1s_code_53_7 = scalar.constant 37962.0 : f32 + %iq1s_code_53_8 = scalar.constant 37969.0 : f32 + %iq1s_code_53_9 = scalar.constant 37972.0 : f32 + %iq1s_code_53_10 = scalar.constant 37973.0 : f32 + %iq1s_code_53_11 = scalar.constant 37974.0 : f32 + %iq1s_code_53_12 = scalar.constant 37976.0 : f32 + %iq1s_code_53_13 = scalar.constant 37977.0 : f32 + %iq1s_code_53_14 = scalar.constant 37984.0 : f32 + %iq1s_code_53_15 = scalar.constant 37985.0 : f32 + %iq1s_code_53_16 = scalar.constant 37986.0 : f32 + %iq1s_code_53_17 = scalar.constant 37989.0 : f32 + %iq1s_code_53_18 = scalar.constant 38020.0 : f32 + %iq1s_code_53_19 = scalar.constant 38022.0 : f32 + %iq1s_code_53_20 = scalar.constant 38034.0 : f32 + %iq1s_code_53_21 = scalar.constant 38036.0 : f32 + %iq1s_code_53_22 = scalar.constant 38037.0 : f32 + %iq1s_code_53_23 = scalar.constant 38040.0 : f32 + %iq1s_code_53_24 = scalar.constant 38049.0 : f32 + %iq1s_code_53_25 = scalar.constant 38057.0 : f32 + %iq1s_code_53_26 = scalar.constant 38144.0 : f32 + %iq1s_code_53_27 = scalar.constant 38149.0 : f32 + %iq1s_code_53_28 = scalar.constant 38152.0 : f32 + %iq1s_code_53_29 = scalar.constant 38154.0 : f32 + %iq1s_code_53_30 = scalar.constant 38160.0 : f32 + %iq1s_code_53_31 = scalar.constant 38161.0 : f32 + %iq1s_code_53 = vector.from_elements %iq1s_code_53_0, %iq1s_code_53_1, %iq1s_code_53_2, %iq1s_code_53_3, %iq1s_code_53_4, %iq1s_code_53_5, %iq1s_code_53_6, %iq1s_code_53_7, %iq1s_code_53_8, %iq1s_code_53_9, %iq1s_code_53_10, %iq1s_code_53_11, %iq1s_code_53_12, %iq1s_code_53_13, %iq1s_code_53_14, %iq1s_code_53_15, %iq1s_code_53_16, %iq1s_code_53_17, %iq1s_code_53_18, %iq1s_code_53_19, %iq1s_code_53_20, %iq1s_code_53_21, %iq1s_code_53_22, %iq1s_code_53_23, %iq1s_code_53_24, %iq1s_code_53_25, %iq1s_code_53_26, %iq1s_code_53_27, %iq1s_code_53_28, %iq1s_code_53_29, %iq1s_code_53_30, %iq1s_code_53_31 : vector<32xf32> + %iq1s_code_54_0 = scalar.constant 38164.0 : f32 + %iq1s_code_54_1 = scalar.constant 38165.0 : f32 + %iq1s_code_54_2 = scalar.constant 38166.0 : f32 + %iq1s_code_54_3 = scalar.constant 38169.0 : f32 + %iq1s_code_54_4 = scalar.constant 38177.0 : f32 + %iq1s_code_54_5 = scalar.constant 38181.0 : f32 + %iq1s_code_54_6 = scalar.constant 38185.0 : f32 + %iq1s_code_54_7 = scalar.constant 38186.0 : f32 + %iq1s_code_54_8 = scalar.constant 38209.0 : f32 + %iq1s_code_54_9 = scalar.constant 38212.0 : f32 + %iq1s_code_54_10 = scalar.constant 38213.0 : f32 + %iq1s_code_54_11 = scalar.constant 38214.0 : f32 + %iq1s_code_54_12 = scalar.constant 38217.0 : f32 + %iq1s_code_54_13 = scalar.constant 38224.0 : f32 + %iq1s_code_54_14 = scalar.constant 38225.0 : f32 + %iq1s_code_54_15 = scalar.constant 38226.0 : f32 + %iq1s_code_54_16 = scalar.constant 38228.0 : f32 + %iq1s_code_54_17 = scalar.constant 38229.0 : f32 + %iq1s_code_54_18 = scalar.constant 38230.0 : f32 + %iq1s_code_54_19 = scalar.constant 38232.0 : f32 + %iq1s_code_54_20 = scalar.constant 38233.0 : f32 + %iq1s_code_54_21 = scalar.constant 38234.0 : f32 + %iq1s_code_54_22 = scalar.constant 38241.0 : f32 + %iq1s_code_54_23 = scalar.constant 38244.0 : f32 + %iq1s_code_54_24 = scalar.constant 38245.0 : f32 + %iq1s_code_54_25 = scalar.constant 38246.0 : f32 + %iq1s_code_54_26 = scalar.constant 38249.0 : f32 + %iq1s_code_54_27 = scalar.constant 38273.0 : f32 + %iq1s_code_54_28 = scalar.constant 38277.0 : f32 + %iq1s_code_54_29 = scalar.constant 38280.0 : f32 + %iq1s_code_54_30 = scalar.constant 38289.0 : f32 + %iq1s_code_54_31 = scalar.constant 38290.0 : f32 + %iq1s_code_54 = vector.from_elements %iq1s_code_54_0, %iq1s_code_54_1, %iq1s_code_54_2, %iq1s_code_54_3, %iq1s_code_54_4, %iq1s_code_54_5, %iq1s_code_54_6, %iq1s_code_54_7, %iq1s_code_54_8, %iq1s_code_54_9, %iq1s_code_54_10, %iq1s_code_54_11, %iq1s_code_54_12, %iq1s_code_54_13, %iq1s_code_54_14, %iq1s_code_54_15, %iq1s_code_54_16, %iq1s_code_54_17, %iq1s_code_54_18, %iq1s_code_54_19, %iq1s_code_54_20, %iq1s_code_54_21, %iq1s_code_54_22, %iq1s_code_54_23, %iq1s_code_54_24, %iq1s_code_54_25, %iq1s_code_54_26, %iq1s_code_54_27, %iq1s_code_54_28, %iq1s_code_54_29, %iq1s_code_54_30, %iq1s_code_54_31 : vector<32xf32> + %iq1s_code_55_0 = scalar.constant 38292.0 : f32 + %iq1s_code_55_1 = scalar.constant 38293.0 : f32 + %iq1s_code_55_2 = scalar.constant 38294.0 : f32 + %iq1s_code_55_3 = scalar.constant 38297.0 : f32 + %iq1s_code_55_4 = scalar.constant 38298.0 : f32 + %iq1s_code_55_5 = scalar.constant 38304.0 : f32 + %iq1s_code_55_6 = scalar.constant 38306.0 : f32 + %iq1s_code_55_7 = scalar.constant 38309.0 : f32 + %iq1s_code_55_8 = scalar.constant 38312.0 : f32 + %iq1s_code_55_9 = scalar.constant 38314.0 : f32 + %iq1s_code_55_10 = scalar.constant 38401.0 : f32 + %iq1s_code_55_11 = scalar.constant 38404.0 : f32 + %iq1s_code_55_12 = scalar.constant 38416.0 : f32 + %iq1s_code_55_13 = scalar.constant 38421.0 : f32 + %iq1s_code_55_14 = scalar.constant 38425.0 : f32 + %iq1s_code_55_15 = scalar.constant 38432.0 : f32 + %iq1s_code_55_16 = scalar.constant 38438.0 : f32 + %iq1s_code_55_17 = scalar.constant 38441.0 : f32 + %iq1s_code_55_18 = scalar.constant 38469.0 : f32 + %iq1s_code_55_19 = scalar.constant 38472.0 : f32 + %iq1s_code_55_20 = scalar.constant 38473.0 : f32 + %iq1s_code_55_21 = scalar.constant 38481.0 : f32 + %iq1s_code_55_22 = scalar.constant 38482.0 : f32 + %iq1s_code_55_23 = scalar.constant 38485.0 : f32 + %iq1s_code_55_24 = scalar.constant 38486.0 : f32 + %iq1s_code_55_25 = scalar.constant 38489.0 : f32 + %iq1s_code_55_26 = scalar.constant 38501.0 : f32 + %iq1s_code_55_27 = scalar.constant 38504.0 : f32 + %iq1s_code_55_28 = scalar.constant 38530.0 : f32 + %iq1s_code_55_29 = scalar.constant 38532.0 : f32 + %iq1s_code_55_30 = scalar.constant 38537.0 : f32 + %iq1s_code_55_31 = scalar.constant 38538.0 : f32 + %iq1s_code_55 = vector.from_elements %iq1s_code_55_0, %iq1s_code_55_1, %iq1s_code_55_2, %iq1s_code_55_3, %iq1s_code_55_4, %iq1s_code_55_5, %iq1s_code_55_6, %iq1s_code_55_7, %iq1s_code_55_8, %iq1s_code_55_9, %iq1s_code_55_10, %iq1s_code_55_11, %iq1s_code_55_12, %iq1s_code_55_13, %iq1s_code_55_14, %iq1s_code_55_15, %iq1s_code_55_16, %iq1s_code_55_17, %iq1s_code_55_18, %iq1s_code_55_19, %iq1s_code_55_20, %iq1s_code_55_21, %iq1s_code_55_22, %iq1s_code_55_23, %iq1s_code_55_24, %iq1s_code_55_25, %iq1s_code_55_26, %iq1s_code_55_27, %iq1s_code_55_28, %iq1s_code_55_29, %iq1s_code_55_30, %iq1s_code_55_31 : vector<32xf32> + %iq1s_code_56_0 = scalar.constant 38546.0 : f32 + %iq1s_code_56_1 = scalar.constant 38548.0 : f32 + %iq1s_code_56_2 = scalar.constant 38549.0 : f32 + %iq1s_code_56_3 = scalar.constant 38564.0 : f32 + %iq1s_code_56_4 = scalar.constant 38566.0 : f32 + %iq1s_code_56_5 = scalar.constant 38569.0 : f32 + %iq1s_code_56_6 = scalar.constant 38917.0 : f32 + %iq1s_code_56_7 = scalar.constant 38934.0 : f32 + %iq1s_code_56_8 = scalar.constant 38937.0 : f32 + %iq1s_code_56_9 = scalar.constant 38949.0 : f32 + %iq1s_code_56_10 = scalar.constant 38977.0 : f32 + %iq1s_code_56_11 = scalar.constant 38982.0 : f32 + %iq1s_code_56_12 = scalar.constant 38992.0 : f32 + %iq1s_code_56_13 = scalar.constant 38994.0 : f32 + %iq1s_code_56_14 = scalar.constant 38997.0 : f32 + %iq1s_code_56_15 = scalar.constant 38998.0 : f32 + %iq1s_code_56_16 = scalar.constant 39002.0 : f32 + %iq1s_code_56_17 = scalar.constant 39012.0 : f32 + %iq1s_code_56_18 = scalar.constant 39013.0 : f32 + %iq1s_code_56_19 = scalar.constant 39045.0 : f32 + %iq1s_code_56_20 = scalar.constant 39057.0 : f32 + %iq1s_code_56_21 = scalar.constant 39062.0 : f32 + %iq1s_code_56_22 = scalar.constant 39065.0 : f32 + %iq1s_code_56_23 = scalar.constant 39077.0 : f32 + %iq1s_code_56_24 = scalar.constant 39172.0 : f32 + %iq1s_code_56_25 = scalar.constant 39174.0 : f32 + %iq1s_code_56_26 = scalar.constant 39177.0 : f32 + %iq1s_code_56_27 = scalar.constant 39184.0 : f32 + %iq1s_code_56_28 = scalar.constant 39186.0 : f32 + %iq1s_code_56_29 = scalar.constant 39189.0 : f32 + %iq1s_code_56_30 = scalar.constant 39192.0 : f32 + %iq1s_code_56_31 = scalar.constant 39194.0 : f32 + %iq1s_code_56 = vector.from_elements %iq1s_code_56_0, %iq1s_code_56_1, %iq1s_code_56_2, %iq1s_code_56_3, %iq1s_code_56_4, %iq1s_code_56_5, %iq1s_code_56_6, %iq1s_code_56_7, %iq1s_code_56_8, %iq1s_code_56_9, %iq1s_code_56_10, %iq1s_code_56_11, %iq1s_code_56_12, %iq1s_code_56_13, %iq1s_code_56_14, %iq1s_code_56_15, %iq1s_code_56_16, %iq1s_code_56_17, %iq1s_code_56_18, %iq1s_code_56_19, %iq1s_code_56_20, %iq1s_code_56_21, %iq1s_code_56_22, %iq1s_code_56_23, %iq1s_code_56_24, %iq1s_code_56_25, %iq1s_code_56_26, %iq1s_code_56_27, %iq1s_code_56_28, %iq1s_code_56_29, %iq1s_code_56_30, %iq1s_code_56_31 : vector<32xf32> + %iq1s_code_57_0 = scalar.constant 39200.0 : f32 + %iq1s_code_57_1 = scalar.constant 39201.0 : f32 + %iq1s_code_57_2 = scalar.constant 39204.0 : f32 + %iq1s_code_57_3 = scalar.constant 39206.0 : f32 + %iq1s_code_57_4 = scalar.constant 39232.0 : f32 + %iq1s_code_57_5 = scalar.constant 39234.0 : f32 + %iq1s_code_57_6 = scalar.constant 39237.0 : f32 + %iq1s_code_57_7 = scalar.constant 39240.0 : f32 + %iq1s_code_57_8 = scalar.constant 39242.0 : f32 + %iq1s_code_57_9 = scalar.constant 39249.0 : f32 + %iq1s_code_57_10 = scalar.constant 39252.0 : f32 + %iq1s_code_57_11 = scalar.constant 39253.0 : f32 + %iq1s_code_57_12 = scalar.constant 39254.0 : f32 + %iq1s_code_57_13 = scalar.constant 39257.0 : f32 + %iq1s_code_57_14 = scalar.constant 39266.0 : f32 + %iq1s_code_57_15 = scalar.constant 39269.0 : f32 + %iq1s_code_57_16 = scalar.constant 39270.0 : f32 + %iq1s_code_57_17 = scalar.constant 39274.0 : f32 + %iq1s_code_57_18 = scalar.constant 39297.0 : f32 + %iq1s_code_57_19 = scalar.constant 39300.0 : f32 + %iq1s_code_57_20 = scalar.constant 39312.0 : f32 + %iq1s_code_57_21 = scalar.constant 39314.0 : f32 + %iq1s_code_57_22 = scalar.constant 39317.0 : f32 + %iq1s_code_57_23 = scalar.constant 39322.0 : f32 + %iq1s_code_57_24 = scalar.constant 39329.0 : f32 + %iq1s_code_57_25 = scalar.constant 39334.0 : f32 + %iq1s_code_57_26 = scalar.constant 39429.0 : f32 + %iq1s_code_57_27 = scalar.constant 39445.0 : f32 + %iq1s_code_57_28 = scalar.constant 39461.0 : f32 + %iq1s_code_57_29 = scalar.constant 39492.0 : f32 + %iq1s_code_57_30 = scalar.constant 39494.0 : f32 + %iq1s_code_57_31 = scalar.constant 39497.0 : f32 + %iq1s_code_57 = vector.from_elements %iq1s_code_57_0, %iq1s_code_57_1, %iq1s_code_57_2, %iq1s_code_57_3, %iq1s_code_57_4, %iq1s_code_57_5, %iq1s_code_57_6, %iq1s_code_57_7, %iq1s_code_57_8, %iq1s_code_57_9, %iq1s_code_57_10, %iq1s_code_57_11, %iq1s_code_57_12, %iq1s_code_57_13, %iq1s_code_57_14, %iq1s_code_57_15, %iq1s_code_57_16, %iq1s_code_57_17, %iq1s_code_57_18, %iq1s_code_57_19, %iq1s_code_57_20, %iq1s_code_57_21, %iq1s_code_57_22, %iq1s_code_57_23, %iq1s_code_57_24, %iq1s_code_57_25, %iq1s_code_57_26, %iq1s_code_57_27, %iq1s_code_57_28, %iq1s_code_57_29, %iq1s_code_57_30, %iq1s_code_57_31 : vector<32xf32> + %iq1s_code_58_0 = scalar.constant 39504.0 : f32 + %iq1s_code_58_1 = scalar.constant 39509.0 : f32 + %iq1s_code_58_2 = scalar.constant 39512.0 : f32 + %iq1s_code_58_3 = scalar.constant 39521.0 : f32 + %iq1s_code_58_4 = scalar.constant 39557.0 : f32 + %iq1s_code_58_5 = scalar.constant 39569.0 : f32 + %iq1s_code_58_6 = scalar.constant 39572.0 : f32 + %iq1s_code_58_7 = scalar.constant 39573.0 : f32 + %iq1s_code_58_8 = scalar.constant 39574.0 : f32 + %iq1s_code_58_9 = scalar.constant 40960.0 : f32 + %iq1s_code_58_10 = scalar.constant 40962.0 : f32 + %iq1s_code_58_11 = scalar.constant 40968.0 : f32 + %iq1s_code_58_12 = scalar.constant 40970.0 : f32 + %iq1s_code_58_13 = scalar.constant 40981.0 : f32 + %iq1s_code_58_14 = scalar.constant 40992.0 : f32 + %iq1s_code_58_15 = scalar.constant 40994.0 : f32 + %iq1s_code_58_16 = scalar.constant 41000.0 : f32 + %iq1s_code_58_17 = scalar.constant 41002.0 : f32 + %iq1s_code_58_18 = scalar.constant 41029.0 : f32 + %iq1s_code_58_19 = scalar.constant 41041.0 : f32 + %iq1s_code_58_20 = scalar.constant 41044.0 : f32 + %iq1s_code_58_21 = scalar.constant 41046.0 : f32 + %iq1s_code_58_22 = scalar.constant 41049.0 : f32 + %iq1s_code_58_23 = scalar.constant 41088.0 : f32 + %iq1s_code_58_24 = scalar.constant 41090.0 : f32 + %iq1s_code_58_25 = scalar.constant 41096.0 : f32 + %iq1s_code_58_26 = scalar.constant 41098.0 : f32 + %iq1s_code_58_27 = scalar.constant 41109.0 : f32 + %iq1s_code_58_28 = scalar.constant 41120.0 : f32 + %iq1s_code_58_29 = scalar.constant 41122.0 : f32 + %iq1s_code_58_30 = scalar.constant 41128.0 : f32 + %iq1s_code_58_31 = scalar.constant 41130.0 : f32 + %iq1s_code_58 = vector.from_elements %iq1s_code_58_0, %iq1s_code_58_1, %iq1s_code_58_2, %iq1s_code_58_3, %iq1s_code_58_4, %iq1s_code_58_5, %iq1s_code_58_6, %iq1s_code_58_7, %iq1s_code_58_8, %iq1s_code_58_9, %iq1s_code_58_10, %iq1s_code_58_11, %iq1s_code_58_12, %iq1s_code_58_13, %iq1s_code_58_14, %iq1s_code_58_15, %iq1s_code_58_16, %iq1s_code_58_17, %iq1s_code_58_18, %iq1s_code_58_19, %iq1s_code_58_20, %iq1s_code_58_21, %iq1s_code_58_22, %iq1s_code_58_23, %iq1s_code_58_24, %iq1s_code_58_25, %iq1s_code_58_26, %iq1s_code_58_27, %iq1s_code_58_28, %iq1s_code_58_29, %iq1s_code_58_30, %iq1s_code_58_31 : vector<32xf32> + %iq1s_code_59_0 = scalar.constant 41221.0 : f32 + %iq1s_code_59_1 = scalar.constant 41225.0 : f32 + %iq1s_code_59_2 = scalar.constant 41233.0 : f32 + %iq1s_code_59_3 = scalar.constant 41236.0 : f32 + %iq1s_code_59_4 = scalar.constant 41238.0 : f32 + %iq1s_code_59_5 = scalar.constant 41241.0 : f32 + %iq1s_code_59_6 = scalar.constant 41242.0 : f32 + %iq1s_code_59_7 = scalar.constant 41286.0 : f32 + %iq1s_code_59_8 = scalar.constant 41289.0 : f32 + %iq1s_code_59_9 = scalar.constant 41297.0 : f32 + %iq1s_code_59_10 = scalar.constant 41301.0 : f32 + %iq1s_code_59_11 = scalar.constant 41304.0 : f32 + %iq1s_code_59_12 = scalar.constant 41306.0 : f32 + %iq1s_code_59_13 = scalar.constant 41313.0 : f32 + %iq1s_code_59_14 = scalar.constant 41316.0 : f32 + %iq1s_code_59_15 = scalar.constant 41349.0 : f32 + %iq1s_code_59_16 = scalar.constant 41360.0 : f32 + %iq1s_code_59_17 = scalar.constant 41362.0 : f32 + %iq1s_code_59_18 = scalar.constant 41366.0 : f32 + %iq1s_code_59_19 = scalar.constant 41369.0 : f32 + %iq1s_code_59_20 = scalar.constant 41474.0 : f32 + %iq1s_code_59_21 = scalar.constant 41480.0 : f32 + %iq1s_code_59_22 = scalar.constant 41482.0 : f32 + %iq1s_code_59_23 = scalar.constant 41488.0 : f32 + %iq1s_code_59_24 = scalar.constant 41497.0 : f32 + %iq1s_code_59_25 = scalar.constant 41506.0 : f32 + %iq1s_code_59_26 = scalar.constant 41512.0 : f32 + %iq1s_code_59_27 = scalar.constant 41514.0 : f32 + %iq1s_code_59_28 = scalar.constant 41541.0 : f32 + %iq1s_code_59_29 = scalar.constant 41553.0 : f32 + %iq1s_code_59_30 = scalar.constant 41558.0 : f32 + %iq1s_code_59_31 = scalar.constant 41561.0 : f32 + %iq1s_code_59 = vector.from_elements %iq1s_code_59_0, %iq1s_code_59_1, %iq1s_code_59_2, %iq1s_code_59_3, %iq1s_code_59_4, %iq1s_code_59_5, %iq1s_code_59_6, %iq1s_code_59_7, %iq1s_code_59_8, %iq1s_code_59_9, %iq1s_code_59_10, %iq1s_code_59_11, %iq1s_code_59_12, %iq1s_code_59_13, %iq1s_code_59_14, %iq1s_code_59_15, %iq1s_code_59_16, %iq1s_code_59_17, %iq1s_code_59_18, %iq1s_code_59_19, %iq1s_code_59_20, %iq1s_code_59_21, %iq1s_code_59_22, %iq1s_code_59_23, %iq1s_code_59_24, %iq1s_code_59_25, %iq1s_code_59_26, %iq1s_code_59_27, %iq1s_code_59_28, %iq1s_code_59_29, %iq1s_code_59_30, %iq1s_code_59_31 : vector<32xf32> + %iq1s_code_60_0 = scalar.constant 41573.0 : f32 + %iq1s_code_60_1 = scalar.constant 41600.0 : f32 + %iq1s_code_60_2 = scalar.constant 41602.0 : f32 + %iq1s_code_60_3 = scalar.constant 41608.0 : f32 + %iq1s_code_60_4 = scalar.constant 41610.0 : f32 + %iq1s_code_60_5 = scalar.constant 41621.0 : f32 + %iq1s_code_60_6 = scalar.constant 41632.0 : f32 + %iq1s_code_60_7 = scalar.constant 41634.0 : f32 + %iq1s_code_60_8 = scalar.constant 41640.0 : f32 + %iq1s_code_60_9 = scalar.constant 41642.0 : f32 + %iq1s_code_60_10 = scalar.constant 42009.0 : f32 + %iq1s_code_60_11 = scalar.constant 42021.0 : f32 + %iq1s_code_60_12 = scalar.constant 42049.0 : f32 + %iq1s_code_60_13 = scalar.constant 42052.0 : f32 + %iq1s_code_60_14 = scalar.constant 42064.0 : f32 + %iq1s_code_60_15 = scalar.constant 42068.0 : f32 + %iq1s_code_60_16 = scalar.constant 42069.0 : f32 + %iq1s_code_60_17 = scalar.constant 42072.0 : f32 + %iq1s_code_60_18 = scalar.constant 42074.0 : f32 + %iq1s_code_60_19 = scalar.constant 42081.0 : f32 + %iq1s_code_60_20 = scalar.constant 42085.0 : f32 + %iq1s_code_60_21 = scalar.constant 42086.0 : f32 + %iq1s_code_60_22 = scalar.constant 42088.0 : f32 + %iq1s_code_60_23 = scalar.constant 42089.0 : f32 + %iq1s_code_60_24 = scalar.constant 42117.0 : f32 + %iq1s_code_60_25 = scalar.constant 42246.0 : f32 + %iq1s_code_60_26 = scalar.constant 42249.0 : f32 + %iq1s_code_60_27 = scalar.constant 42256.0 : f32 + %iq1s_code_60_28 = scalar.constant 42258.0 : f32 + %iq1s_code_60_29 = scalar.constant 42261.0 : f32 + %iq1s_code_60_30 = scalar.constant 42264.0 : f32 + %iq1s_code_60_31 = scalar.constant 42278.0 : f32 + %iq1s_code_60 = vector.from_elements %iq1s_code_60_0, %iq1s_code_60_1, %iq1s_code_60_2, %iq1s_code_60_3, %iq1s_code_60_4, %iq1s_code_60_5, %iq1s_code_60_6, %iq1s_code_60_7, %iq1s_code_60_8, %iq1s_code_60_9, %iq1s_code_60_10, %iq1s_code_60_11, %iq1s_code_60_12, %iq1s_code_60_13, %iq1s_code_60_14, %iq1s_code_60_15, %iq1s_code_60_16, %iq1s_code_60_17, %iq1s_code_60_18, %iq1s_code_60_19, %iq1s_code_60_20, %iq1s_code_60_21, %iq1s_code_60_22, %iq1s_code_60_23, %iq1s_code_60_24, %iq1s_code_60_25, %iq1s_code_60_26, %iq1s_code_60_27, %iq1s_code_60_28, %iq1s_code_60_29, %iq1s_code_60_30, %iq1s_code_60_31 : vector<32xf32> + %iq1s_code_61_0 = scalar.constant 42281.0 : f32 + %iq1s_code_61_1 = scalar.constant 42306.0 : f32 + %iq1s_code_61_2 = scalar.constant 42309.0 : f32 + %iq1s_code_61_3 = scalar.constant 42321.0 : f32 + %iq1s_code_61_4 = scalar.constant 42324.0 : f32 + %iq1s_code_61_5 = scalar.constant 42325.0 : f32 + %iq1s_code_61_6 = scalar.constant 42326.0 : f32 + %iq1s_code_61_7 = scalar.constant 42329.0 : f32 + %iq1s_code_61_8 = scalar.constant 42341.0 : f32 + %iq1s_code_61_9 = scalar.constant 42346.0 : f32 + %iq1s_code_61_10 = scalar.constant 42369.0 : f32 + %iq1s_code_61_11 = scalar.constant 42372.0 : f32 + %iq1s_code_61_12 = scalar.constant 42373.0 : f32 + %iq1s_code_61_13 = scalar.constant 42374.0 : f32 + %iq1s_code_61_14 = scalar.constant 42377.0 : f32 + %iq1s_code_61_15 = scalar.constant 42386.0 : f32 + %iq1s_code_61_16 = scalar.constant 42389.0 : f32 + %iq1s_code_61_17 = scalar.constant 42392.0 : f32 + %iq1s_code_61_18 = scalar.constant 42501.0 : f32 + %iq1s_code_61_19 = scalar.constant 42513.0 : f32 + %iq1s_code_61_20 = scalar.constant 42518.0 : f32 + %iq1s_code_61_21 = scalar.constant 42522.0 : f32 + %iq1s_code_61_22 = scalar.constant 42529.0 : f32 + %iq1s_code_61_23 = scalar.constant 42533.0 : f32 + %iq1s_code_61_24 = scalar.constant 42564.0 : f32 + %iq1s_code_61_25 = scalar.constant 42566.0 : f32 + %iq1s_code_61_26 = scalar.constant 42570.0 : f32 + %iq1s_code_61_27 = scalar.constant 42578.0 : f32 + %iq1s_code_61_28 = scalar.constant 42581.0 : f32 + %iq1s_code_61_29 = scalar.constant 42582.0 : f32 + %iq1s_code_61_30 = scalar.constant 42584.0 : f32 + %iq1s_code_61_31 = scalar.constant 42592.0 : f32 + %iq1s_code_61 = vector.from_elements %iq1s_code_61_0, %iq1s_code_61_1, %iq1s_code_61_2, %iq1s_code_61_3, %iq1s_code_61_4, %iq1s_code_61_5, %iq1s_code_61_6, %iq1s_code_61_7, %iq1s_code_61_8, %iq1s_code_61_9, %iq1s_code_61_10, %iq1s_code_61_11, %iq1s_code_61_12, %iq1s_code_61_13, %iq1s_code_61_14, %iq1s_code_61_15, %iq1s_code_61_16, %iq1s_code_61_17, %iq1s_code_61_18, %iq1s_code_61_19, %iq1s_code_61_20, %iq1s_code_61_21, %iq1s_code_61_22, %iq1s_code_61_23, %iq1s_code_61_24, %iq1s_code_61_25, %iq1s_code_61_26, %iq1s_code_61_27, %iq1s_code_61_28, %iq1s_code_61_29, %iq1s_code_61_30, %iq1s_code_61_31 : vector<32xf32> + %iq1s_code_62_0 = scalar.constant 42594.0 : f32 + %iq1s_code_62_1 = scalar.constant 42630.0 : f32 + %iq1s_code_62_2 = scalar.constant 42640.0 : f32 + %iq1s_code_62_3 = scalar.constant 42645.0 : f32 + %iq1s_code_62_4 = scalar.constant 42646.0 : f32 + %iq1s_code_62_5 = scalar.constant 42649.0 : f32 + %iq1s_code_62_6 = scalar.constant 42657.0 : f32 + %iq1s_code_62_7 = scalar.constant 42660.0 : f32 + %iq1s_code_62_8 = scalar.constant 42662.0 : f32 + %iq1s_code_62_9 = scalar.constant 43008.0 : f32 + %iq1s_code_62_10 = scalar.constant 43010.0 : f32 + %iq1s_code_62_11 = scalar.constant 43016.0 : f32 + %iq1s_code_62_12 = scalar.constant 43018.0 : f32 + %iq1s_code_62_13 = scalar.constant 43040.0 : f32 + %iq1s_code_62_14 = scalar.constant 43042.0 : f32 + %iq1s_code_62_15 = scalar.constant 43048.0 : f32 + %iq1s_code_62_16 = scalar.constant 43050.0 : f32 + %iq1s_code_62_17 = scalar.constant 43089.0 : f32 + %iq1s_code_62_18 = scalar.constant 43092.0 : f32 + %iq1s_code_62_19 = scalar.constant 43094.0 : f32 + %iq1s_code_62_20 = scalar.constant 43097.0 : f32 + %iq1s_code_62_21 = scalar.constant 43136.0 : f32 + %iq1s_code_62_22 = scalar.constant 43138.0 : f32 + %iq1s_code_62_23 = scalar.constant 43144.0 : f32 + %iq1s_code_62_24 = scalar.constant 43146.0 : f32 + %iq1s_code_62_25 = scalar.constant 43157.0 : f32 + %iq1s_code_62_26 = scalar.constant 43168.0 : f32 + %iq1s_code_62_27 = scalar.constant 43170.0 : f32 + %iq1s_code_62_28 = scalar.constant 43176.0 : f32 + %iq1s_code_62_29 = scalar.constant 43178.0 : f32 + %iq1s_code_62_30 = scalar.constant 43269.0 : f32 + %iq1s_code_62_31 = scalar.constant 43284.0 : f32 + %iq1s_code_62 = vector.from_elements %iq1s_code_62_0, %iq1s_code_62_1, %iq1s_code_62_2, %iq1s_code_62_3, %iq1s_code_62_4, %iq1s_code_62_5, %iq1s_code_62_6, %iq1s_code_62_7, %iq1s_code_62_8, %iq1s_code_62_9, %iq1s_code_62_10, %iq1s_code_62_11, %iq1s_code_62_12, %iq1s_code_62_13, %iq1s_code_62_14, %iq1s_code_62_15, %iq1s_code_62_16, %iq1s_code_62_17, %iq1s_code_62_18, %iq1s_code_62_19, %iq1s_code_62_20, %iq1s_code_62_21, %iq1s_code_62_22, %iq1s_code_62_23, %iq1s_code_62_24, %iq1s_code_62_25, %iq1s_code_62_26, %iq1s_code_62_27, %iq1s_code_62_28, %iq1s_code_62_29, %iq1s_code_62_30, %iq1s_code_62_31 : vector<32xf32> + %iq1s_code_63_0 = scalar.constant 43289.0 : f32 + %iq1s_code_63_1 = scalar.constant 43297.0 : f32 + %iq1s_code_63_2 = scalar.constant 43301.0 : f32 + %iq1s_code_63_3 = scalar.constant 43329.0 : f32 + %iq1s_code_63_4 = scalar.constant 43344.0 : f32 + %iq1s_code_63_5 = scalar.constant 43349.0 : f32 + %iq1s_code_63_6 = scalar.constant 43354.0 : f32 + %iq1s_code_63_7 = scalar.constant 43361.0 : f32 + %iq1s_code_63_8 = scalar.constant 43366.0 : f32 + %iq1s_code_63_9 = scalar.constant 43369.0 : f32 + %iq1s_code_63_10 = scalar.constant 43408.0 : f32 + %iq1s_code_63_11 = scalar.constant 43414.0 : f32 + %iq1s_code_63_12 = scalar.constant 43520.0 : f32 + %iq1s_code_63_13 = scalar.constant 43522.0 : f32 + %iq1s_code_63_14 = scalar.constant 43528.0 : f32 + %iq1s_code_63_15 = scalar.constant 43530.0 : f32 + %iq1s_code_63_16 = scalar.constant 43552.0 : f32 + %iq1s_code_63_17 = scalar.constant 43554.0 : f32 + %iq1s_code_63_18 = scalar.constant 43560.0 : f32 + %iq1s_code_63_19 = scalar.constant 43562.0 : f32 + %iq1s_code_63_20 = scalar.constant 43601.0 : f32 + %iq1s_code_63_21 = scalar.constant 43604.0 : f32 + %iq1s_code_63_22 = scalar.constant 43606.0 : f32 + %iq1s_code_63_23 = scalar.constant 43648.0 : f32 + %iq1s_code_63_24 = scalar.constant 43650.0 : f32 + %iq1s_code_63_25 = scalar.constant 43656.0 : f32 + %iq1s_code_63_26 = scalar.constant 43658.0 : f32 + %iq1s_code_63_27 = scalar.constant 43669.0 : f32 + %iq1s_code_63_28 = scalar.constant 43680.0 : f32 + %iq1s_code_63_29 = scalar.constant 43682.0 : f32 + %iq1s_code_63_30 = scalar.constant 43688.0 : f32 + %iq1s_code_63_31 = scalar.constant 43690.0 : f32 + %iq1s_code_63 = vector.from_elements %iq1s_code_63_0, %iq1s_code_63_1, %iq1s_code_63_2, %iq1s_code_63_3, %iq1s_code_63_4, %iq1s_code_63_5, %iq1s_code_63_6, %iq1s_code_63_7, %iq1s_code_63_8, %iq1s_code_63_9, %iq1s_code_63_10, %iq1s_code_63_11, %iq1s_code_63_12, %iq1s_code_63_13, %iq1s_code_63_14, %iq1s_code_63_15, %iq1s_code_63_16, %iq1s_code_63_17, %iq1s_code_63_18, %iq1s_code_63_19, %iq1s_code_63_20, %iq1s_code_63_21, %iq1s_code_63_22, %iq1s_code_63_23, %iq1s_code_63_24, %iq1s_code_63_25, %iq1s_code_63_26, %iq1s_code_63_27, %iq1s_code_63_28, %iq1s_code_63_29, %iq1s_code_63_30, %iq1s_code_63_31 : vector<32xf32> + %is_chunk1 = scalar.cmpi eq, %chunk_i32, %chunk_id1 : i32 + %sel1 = scf.select %is_chunk1, %iq1s_code_1, %iq1s_code_0 : vector<32xf32> + %is_chunk2 = scalar.cmpi eq, %chunk_i32, %chunk_id2 : i32 + %sel2 = scf.select %is_chunk2, %iq1s_code_2, %sel1 : vector<32xf32> + %is_chunk3 = scalar.cmpi eq, %chunk_i32, %chunk_id3 : i32 + %sel3 = scf.select %is_chunk3, %iq1s_code_3, %sel2 : vector<32xf32> + %is_chunk4 = scalar.cmpi eq, %chunk_i32, %chunk_id4 : i32 + %sel4 = scf.select %is_chunk4, %iq1s_code_4, %sel3 : vector<32xf32> + %is_chunk5 = scalar.cmpi eq, %chunk_i32, %chunk_id5 : i32 + %sel5 = scf.select %is_chunk5, %iq1s_code_5, %sel4 : vector<32xf32> + %is_chunk6 = scalar.cmpi eq, %chunk_i32, %chunk_id6 : i32 + %sel6 = scf.select %is_chunk6, %iq1s_code_6, %sel5 : vector<32xf32> + %is_chunk7 = scalar.cmpi eq, %chunk_i32, %chunk_id7 : i32 + %sel7 = scf.select %is_chunk7, %iq1s_code_7, %sel6 : vector<32xf32> + %is_chunk8 = scalar.cmpi eq, %chunk_i32, %chunk_id8 : i32 + %sel8 = scf.select %is_chunk8, %iq1s_code_8, %sel7 : vector<32xf32> + %is_chunk9 = scalar.cmpi eq, %chunk_i32, %chunk_id9 : i32 + %sel9 = scf.select %is_chunk9, %iq1s_code_9, %sel8 : vector<32xf32> + %is_chunk10 = scalar.cmpi eq, %chunk_i32, %chunk_id10 : i32 + %sel10 = scf.select %is_chunk10, %iq1s_code_10, %sel9 : vector<32xf32> + %is_chunk11 = scalar.cmpi eq, %chunk_i32, %chunk_id11 : i32 + %sel11 = scf.select %is_chunk11, %iq1s_code_11, %sel10 : vector<32xf32> + %is_chunk12 = scalar.cmpi eq, %chunk_i32, %chunk_id12 : i32 + %sel12 = scf.select %is_chunk12, %iq1s_code_12, %sel11 : vector<32xf32> + %is_chunk13 = scalar.cmpi eq, %chunk_i32, %chunk_id13 : i32 + %sel13 = scf.select %is_chunk13, %iq1s_code_13, %sel12 : vector<32xf32> + %is_chunk14 = scalar.cmpi eq, %chunk_i32, %chunk_id14 : i32 + %sel14 = scf.select %is_chunk14, %iq1s_code_14, %sel13 : vector<32xf32> + %is_chunk15 = scalar.cmpi eq, %chunk_i32, %chunk_id15 : i32 + %sel15 = scf.select %is_chunk15, %iq1s_code_15, %sel14 : vector<32xf32> + %is_chunk16 = scalar.cmpi eq, %chunk_i32, %chunk_id16 : i32 + %sel16 = scf.select %is_chunk16, %iq1s_code_16, %sel15 : vector<32xf32> + %is_chunk17 = scalar.cmpi eq, %chunk_i32, %chunk_id17 : i32 + %sel17 = scf.select %is_chunk17, %iq1s_code_17, %sel16 : vector<32xf32> + %is_chunk18 = scalar.cmpi eq, %chunk_i32, %chunk_id18 : i32 + %sel18 = scf.select %is_chunk18, %iq1s_code_18, %sel17 : vector<32xf32> + %is_chunk19 = scalar.cmpi eq, %chunk_i32, %chunk_id19 : i32 + %sel19 = scf.select %is_chunk19, %iq1s_code_19, %sel18 : vector<32xf32> + %is_chunk20 = scalar.cmpi eq, %chunk_i32, %chunk_id20 : i32 + %sel20 = scf.select %is_chunk20, %iq1s_code_20, %sel19 : vector<32xf32> + %is_chunk21 = scalar.cmpi eq, %chunk_i32, %chunk_id21 : i32 + %sel21 = scf.select %is_chunk21, %iq1s_code_21, %sel20 : vector<32xf32> + %is_chunk22 = scalar.cmpi eq, %chunk_i32, %chunk_id22 : i32 + %sel22 = scf.select %is_chunk22, %iq1s_code_22, %sel21 : vector<32xf32> + %is_chunk23 = scalar.cmpi eq, %chunk_i32, %chunk_id23 : i32 + %sel23 = scf.select %is_chunk23, %iq1s_code_23, %sel22 : vector<32xf32> + %is_chunk24 = scalar.cmpi eq, %chunk_i32, %chunk_id24 : i32 + %sel24 = scf.select %is_chunk24, %iq1s_code_24, %sel23 : vector<32xf32> + %is_chunk25 = scalar.cmpi eq, %chunk_i32, %chunk_id25 : i32 + %sel25 = scf.select %is_chunk25, %iq1s_code_25, %sel24 : vector<32xf32> + %is_chunk26 = scalar.cmpi eq, %chunk_i32, %chunk_id26 : i32 + %sel26 = scf.select %is_chunk26, %iq1s_code_26, %sel25 : vector<32xf32> + %is_chunk27 = scalar.cmpi eq, %chunk_i32, %chunk_id27 : i32 + %sel27 = scf.select %is_chunk27, %iq1s_code_27, %sel26 : vector<32xf32> + %is_chunk28 = scalar.cmpi eq, %chunk_i32, %chunk_id28 : i32 + %sel28 = scf.select %is_chunk28, %iq1s_code_28, %sel27 : vector<32xf32> + %is_chunk29 = scalar.cmpi eq, %chunk_i32, %chunk_id29 : i32 + %sel29 = scf.select %is_chunk29, %iq1s_code_29, %sel28 : vector<32xf32> + %is_chunk30 = scalar.cmpi eq, %chunk_i32, %chunk_id30 : i32 + %sel30 = scf.select %is_chunk30, %iq1s_code_30, %sel29 : vector<32xf32> + %is_chunk31 = scalar.cmpi eq, %chunk_i32, %chunk_id31 : i32 + %sel31 = scf.select %is_chunk31, %iq1s_code_31, %sel30 : vector<32xf32> + %is_chunk32 = scalar.cmpi eq, %chunk_i32, %chunk_id32 : i32 + %sel32 = scf.select %is_chunk32, %iq1s_code_32, %sel31 : vector<32xf32> + %is_chunk33 = scalar.cmpi eq, %chunk_i32, %chunk_id33 : i32 + %sel33 = scf.select %is_chunk33, %iq1s_code_33, %sel32 : vector<32xf32> + %is_chunk34 = scalar.cmpi eq, %chunk_i32, %chunk_id34 : i32 + %sel34 = scf.select %is_chunk34, %iq1s_code_34, %sel33 : vector<32xf32> + %is_chunk35 = scalar.cmpi eq, %chunk_i32, %chunk_id35 : i32 + %sel35 = scf.select %is_chunk35, %iq1s_code_35, %sel34 : vector<32xf32> + %is_chunk36 = scalar.cmpi eq, %chunk_i32, %chunk_id36 : i32 + %sel36 = scf.select %is_chunk36, %iq1s_code_36, %sel35 : vector<32xf32> + %is_chunk37 = scalar.cmpi eq, %chunk_i32, %chunk_id37 : i32 + %sel37 = scf.select %is_chunk37, %iq1s_code_37, %sel36 : vector<32xf32> + %is_chunk38 = scalar.cmpi eq, %chunk_i32, %chunk_id38 : i32 + %sel38 = scf.select %is_chunk38, %iq1s_code_38, %sel37 : vector<32xf32> + %is_chunk39 = scalar.cmpi eq, %chunk_i32, %chunk_id39 : i32 + %sel39 = scf.select %is_chunk39, %iq1s_code_39, %sel38 : vector<32xf32> + %is_chunk40 = scalar.cmpi eq, %chunk_i32, %chunk_id40 : i32 + %sel40 = scf.select %is_chunk40, %iq1s_code_40, %sel39 : vector<32xf32> + %is_chunk41 = scalar.cmpi eq, %chunk_i32, %chunk_id41 : i32 + %sel41 = scf.select %is_chunk41, %iq1s_code_41, %sel40 : vector<32xf32> + %is_chunk42 = scalar.cmpi eq, %chunk_i32, %chunk_id42 : i32 + %sel42 = scf.select %is_chunk42, %iq1s_code_42, %sel41 : vector<32xf32> + %is_chunk43 = scalar.cmpi eq, %chunk_i32, %chunk_id43 : i32 + %sel43 = scf.select %is_chunk43, %iq1s_code_43, %sel42 : vector<32xf32> + %is_chunk44 = scalar.cmpi eq, %chunk_i32, %chunk_id44 : i32 + %sel44 = scf.select %is_chunk44, %iq1s_code_44, %sel43 : vector<32xf32> + %is_chunk45 = scalar.cmpi eq, %chunk_i32, %chunk_id45 : i32 + %sel45 = scf.select %is_chunk45, %iq1s_code_45, %sel44 : vector<32xf32> + %is_chunk46 = scalar.cmpi eq, %chunk_i32, %chunk_id46 : i32 + %sel46 = scf.select %is_chunk46, %iq1s_code_46, %sel45 : vector<32xf32> + %is_chunk47 = scalar.cmpi eq, %chunk_i32, %chunk_id47 : i32 + %sel47 = scf.select %is_chunk47, %iq1s_code_47, %sel46 : vector<32xf32> + %is_chunk48 = scalar.cmpi eq, %chunk_i32, %chunk_id48 : i32 + %sel48 = scf.select %is_chunk48, %iq1s_code_48, %sel47 : vector<32xf32> + %is_chunk49 = scalar.cmpi eq, %chunk_i32, %chunk_id49 : i32 + %sel49 = scf.select %is_chunk49, %iq1s_code_49, %sel48 : vector<32xf32> + %is_chunk50 = scalar.cmpi eq, %chunk_i32, %chunk_id50 : i32 + %sel50 = scf.select %is_chunk50, %iq1s_code_50, %sel49 : vector<32xf32> + %is_chunk51 = scalar.cmpi eq, %chunk_i32, %chunk_id51 : i32 + %sel51 = scf.select %is_chunk51, %iq1s_code_51, %sel50 : vector<32xf32> + %is_chunk52 = scalar.cmpi eq, %chunk_i32, %chunk_id52 : i32 + %sel52 = scf.select %is_chunk52, %iq1s_code_52, %sel51 : vector<32xf32> + %is_chunk53 = scalar.cmpi eq, %chunk_i32, %chunk_id53 : i32 + %sel53 = scf.select %is_chunk53, %iq1s_code_53, %sel52 : vector<32xf32> + %is_chunk54 = scalar.cmpi eq, %chunk_i32, %chunk_id54 : i32 + %sel54 = scf.select %is_chunk54, %iq1s_code_54, %sel53 : vector<32xf32> + %is_chunk55 = scalar.cmpi eq, %chunk_i32, %chunk_id55 : i32 + %sel55 = scf.select %is_chunk55, %iq1s_code_55, %sel54 : vector<32xf32> + %is_chunk56 = scalar.cmpi eq, %chunk_i32, %chunk_id56 : i32 + %sel56 = scf.select %is_chunk56, %iq1s_code_56, %sel55 : vector<32xf32> + %is_chunk57 = scalar.cmpi eq, %chunk_i32, %chunk_id57 : i32 + %sel57 = scf.select %is_chunk57, %iq1s_code_57, %sel56 : vector<32xf32> + %is_chunk58 = scalar.cmpi eq, %chunk_i32, %chunk_id58 : i32 + %sel58 = scf.select %is_chunk58, %iq1s_code_58, %sel57 : vector<32xf32> + %is_chunk59 = scalar.cmpi eq, %chunk_i32, %chunk_id59 : i32 + %sel59 = scf.select %is_chunk59, %iq1s_code_59, %sel58 : vector<32xf32> + %is_chunk60 = scalar.cmpi eq, %chunk_i32, %chunk_id60 : i32 + %sel60 = scf.select %is_chunk60, %iq1s_code_60, %sel59 : vector<32xf32> + %is_chunk61 = scalar.cmpi eq, %chunk_i32, %chunk_id61 : i32 + %sel61 = scf.select %is_chunk61, %iq1s_code_61, %sel60 : vector<32xf32> + %is_chunk62 = scalar.cmpi eq, %chunk_i32, %chunk_id62 : i32 + %sel62 = scf.select %is_chunk62, %iq1s_code_62, %sel61 : vector<32xf32> + %is_chunk63 = scalar.cmpi eq, %chunk_i32, %chunk_id63 : i32 + %sel63 = scf.select %is_chunk63, %iq1s_code_63, %sel62 : vector<32xf32> + %v = vector.table.lookup %sel63[%codes] : vector<32xf32>, vector<1xi32> -> vector<1xf32> + %v_f32 = vector.extract %v[0] : vector<1xf32> -> f32 + %code = scalar.fptoui %v_f32 : f32 to i32 + func.return %code : i32 +} + +// four values 4 (p % 2) .. +3 of an IQ1 grid code: ((code >> 2j) & 3) - 1 + delta, times %scale +func.def inline @ggml_iq1_code_vector4(%code: i32, %half: i32, %delta: f32, %scale: f32) -> (vector<4xf32>) { + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c2v = scalar.constant 2 : i32 + %c4v = scalar.constant 4 : i32 + %c6v = scalar.constant 6 : i32 + %c0v = scalar.constant 0 : i32 + %base = scalar.muli %half, %c8_i32 : i32 + %cv = vector.splat %code : vector<4xi32> + %bv = vector.splat %base : vector<4xi32> + %sh0 = vector.from_elements %c0v, %c2v, %c4v, %c6v : vector<4xi32> + %sh = vector.addi %sh0, %bv : vector<4xi32> + %s = vector.shrui %cv, %sh : vector<4xi32> + %three = vector.splat %c3_i32 : vector<4xi32> + %one = vector.splat %c1_i32 : vector<4xi32> + %c = vector.andi %s, %three : vector<4xi32> + %t = vector.subi %c, %one : vector<4xi32> + %tf = vector.sitofp %t : vector<4xi32> to vector<4xf32> + %dv = vector.splat %delta : vector<4xf32> + %sv = vector.splat %scale : vector<4xf32> + %td = vector.addf %tf, %dv : vector<4xf32> + %result = vector.mulf %td, %sv : vector<4xf32> + func.return %result : vector<4xf32> +} + +// IQ1_S (50 bytes: d, qs[32], qh[8] u16; dequantize_row_iq1_s). Group g: qh[g] holds three 3-bit +// high index parts (slot l at bit 3l), a 3-bit scale at bit 12 and the delta sign at bit 15: +// value = d * (2 s + 1) * (grid + delta), delta = +-0.125. Packet p covers slot p / 2. +func.def inline @ggml_iq1s_f32_vector4(%weight: buffer, %row_byte_base: offset, %iq1_block: index, %iq1_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c17 = index.constant 17 : index + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c12_i32 = scalar.constant 12 : i32 + %c32768_i32 = scalar.constant 32768 : i32 + %c0_i32 = scalar.constant 0 : i32 + %pos_delta = scalar.constant 0.125 : f32 + %neg_delta = scalar.constant -0.125 : f32 + %block_bytes = index.constant 50 : offset + %block_byte_add = index.scale %iq1_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %hv = buffer.view %weight[%block_byte_base] : buffer -> view<25xf16> + %wv = buffer.view %weight[%block_byte_base] : buffer -> view<25xi16> + %bv = buffer.view %weight[%block_byte_base] : buffer -> view<50xi8> + %g = index.assume %iq1_group [range(%iq1_group, 0, 7)] : index + %p = index.assume %packet [range(%packet, 0, 7)] : index + %slot = index.div %p, %c2 : index + %half = index.rem %p, %c2 : index + %g4 = index.mul %g, %c4 : index + %qs_at0 = index.add %g4, %c2 : index + %qs_at = index.add %qs_at0, %slot : index + %qh_at = index.add %c17, %g : index + %d_f16 = view.load %hv[%c0] : view<25xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %qs_i8 = view.load %bv[%qs_at] : view<50xi8> -> i8 + %qs = scalar.extui %qs_i8 : i8 to i32 + %qh_i16 = view.load %wv[%qh_at] : view<25xi16> -> i16 + %qh = scalar.extui %qh_i16 : i16 to i32 + %slot_i32 = index.cast %slot : index to i32 + %hsh = scalar.muli %slot_i32, %c3_i32 : i32 + %hi0 = scalar.shrui %qh, %hsh : i32 + %hi1 = scalar.andi %hi0, %c7_i32 : i32 + %hi = scalar.shli %hi1, %c8_i32 : i32 + %gi = scalar.ori %qs, %hi : i32 + %s0 = scalar.shrui %qh, %c12_i32 : i32 + %s = scalar.andi %s0, %c7_i32 : i32 + %s2 = scalar.addi %s, %s : i32 + %s21 = scalar.addi %s2, %c1_i32 : i32 + %sf = scalar.uitofp %s21 : i32 to f32 + %scale = scalar.mulf %d, %sf : f32 + %neg_bit = scalar.andi %qh, %c32768_i32 : i32 + %neg = scalar.cmpi ne, %neg_bit, %c0_i32 : i32 + %delta = scf.select %neg, %neg_delta, %pos_delta : f32 + %code = func.call @ggml_iq1s_grid_code_i32(%gi) : (i32) -> (i32) + %half_i32 = index.cast %half : index to i32 + %result = func.call @ggml_iq1_code_vector4(%code, %half_i32, %delta, %scale) : (i32, i32, f32, f32) -> (vector<4xf32>) + func.return %result : vector<4xf32> +} + +func.def inline @ggml_iq1s_f16_vector4(%weight: buffer, %row_byte_base: offset, %iq1_block: index, %iq1_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_iq1s_f32_vector4(%weight, %row_byte_base, %iq1_block, %iq1_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +// IQ1_M's fp16 block scale: the top nibbles of its four u16 scale words (bytes 48..55). +func.def inline @ggml_iq1m_block_scale(%weight: buffer, %block_byte_base: offset) -> (f32) { + %wv = buffer.view %weight[%block_byte_base] : buffer -> view<28xi16> + %c24 = index.constant 24 : index + %c25 = index.constant 25 : index + %c26 = index.constant 26 : index + %c27 = index.constant 27 : index + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c12_i32 = scalar.constant 12 : i32 + %c240_i32 = scalar.constant 240 : i32 + %c3840_i32 = scalar.constant 3840 : i32 + %c61440_i32 = scalar.constant 61440 : i32 + %c255_i32 = scalar.constant 255 : i32 + %w0_i16 = view.load %wv[%c24] : view<28xi16> -> i16 + %w1_i16 = view.load %wv[%c25] : view<28xi16> -> i16 + %w2_i16 = view.load %wv[%c26] : view<28xi16> -> i16 + %w3_i16 = view.load %wv[%c27] : view<28xi16> -> i16 + %w0 = scalar.extui %w0_i16 : i16 to i32 + %w1 = scalar.extui %w1_i16 : i16 to i32 + %w2 = scalar.extui %w2_i16 : i16 to i32 + %w3 = scalar.extui %w3_i16 : i16 to i32 + %n0 = scalar.shrui %w0, %c12_i32 : i32 + %n1a = scalar.shrui %w1, %c8_i32 : i32 + %n1 = scalar.andi %n1a, %c240_i32 : i32 + %n2a = scalar.shrui %w2, %c4_i32 : i32 + %n2 = scalar.andi %n2a, %c3840_i32 : i32 + %n3 = scalar.andi %w3, %c61440_i32 : i32 + %u01 = scalar.ori %n0, %n1 : i32 + %u012 = scalar.ori %u01, %n2 : i32 + %u = scalar.ori %u012, %n3 : i32 + %lo = scalar.andi %u, %c255_i32 : i32 + %hi = scalar.shrui %u, %c8_i32 : i32 + %lo8 = scalar.trunci %lo : i32 to i8 + %hi8 = scalar.trunci %hi : i32 to i8 + %bytes = vector.from_elements %lo8, %hi8 : vector<2xi8> + %h = vector.bitcast %bytes : vector<2xi8> to vector<1xf16> + %h0 = vector.extract %h[0] : vector<1xf16> -> f16 + %f = scalar.extf %h0 : f16 to f32 + func.return %f : f32 +} + +// IQ1_M (56 bytes: qs[32], qh[16], scales[8]; dequantize_row_iq1_m). Slot l of group g: qh byte +// 2g + l / 2 holds the high index bits (bits 0..2 for even l, 4..6 for odd) and the delta sign +// (bit 3 / bit 7); the 3-bit scale of slots 0-1 / 2-3 sits at bit 6 (g % 2) / +3 of u16 scale +// word g / 2: value = d * (2 s + 1) * (grid + delta). +func.def inline @ggml_iq1m_f32_vector4(%weight: buffer, %row_byte_base: offset, %iq1_block: index, %iq1_group: index, %packet: index) -> (vector<4xf32>) { + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c24 = index.constant 24 : index + %c32 = index.constant 32 : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c6_i32 = scalar.constant 6 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c1792_i32 = scalar.constant 1792 : i32 + %pos_delta = scalar.constant 0.125 : f32 + %neg_delta = scalar.constant -0.125 : f32 + %block_bytes = index.constant 56 : offset + %block_byte_add = index.scale %iq1_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %wv = buffer.view %weight[%block_byte_base] : buffer -> view<28xi16> + %bv = buffer.view %weight[%block_byte_base] : buffer -> view<56xi8> + %g = index.assume %iq1_group [range(%iq1_group, 0, 7)] : index + %p = index.assume %packet [range(%packet, 0, 7)] : index + %slot = index.div %p, %c2 : index + %half = index.rem %p, %c2 : index + %pair = index.div %slot, %c2 : index + %odd = index.rem %slot, %c2 : index + %g4 = index.mul %g, %c4 : index + %qs_at = index.add %g4, %slot : index + %g2 = index.add %g, %g : index + %qh_at0 = index.add %c32, %g2 : index + %qh_at = index.add %qh_at0, %pair : index + %gh = index.div %g, %c2 : index + %gp = index.rem %g, %c2 : index + %sc_at = index.add %c24, %gh : index + %d = func.call @ggml_iq1m_block_scale(%weight, %block_byte_base) : (buffer, offset) -> (f32) + %qs_i8 = view.load %bv[%qs_at] : view<56xi8> -> i8 + %qs = scalar.extui %qs_i8 : i8 to i32 + %qh_i8 = view.load %bv[%qh_at] : view<56xi8> -> i8 + %qh = scalar.extui %qh_i8 : i8 to i32 + %odd_i32 = index.cast %odd : index to i32 + %odd4 = scalar.muli %odd_i32, %c4_i32 : i32 + %qh_n = scalar.shrui %qh, %odd4 : i32 + %hi0 = scalar.shli %qh_n, %c8_i32 : i32 + %hi = scalar.andi %hi0, %c1792_i32 : i32 + %gi = scalar.ori %qs, %hi : i32 + %neg_bit0 = scalar.shrui %qh_n, %c3_i32 : i32 + %neg_bit = scalar.andi %neg_bit0, %c1_i32 : i32 + %neg = scalar.cmpi ne, %neg_bit, %c0_i32 : i32 + %delta = scf.select %neg, %neg_delta, %pos_delta : f32 + %sc_i16 = view.load %wv[%sc_at] : view<28xi16> -> i16 + %sc = scalar.extui %sc_i16 : i16 to i32 + %gp_i32 = index.cast %gp : index to i32 + %pair_i32 = index.cast %pair : index to i32 + %ssh0 = scalar.muli %gp_i32, %c6_i32 : i32 + %ssh1 = scalar.muli %pair_i32, %c3_i32 : i32 + %ssh = scalar.addi %ssh0, %ssh1 : i32 + %s0 = scalar.shrui %sc, %ssh : i32 + %s = scalar.andi %s0, %c7_i32 : i32 + %s2 = scalar.addi %s, %s : i32 + %s21 = scalar.addi %s2, %c1_i32 : i32 + %sf = scalar.uitofp %s21 : i32 to f32 + %scale = scalar.mulf %d, %sf : f32 + %code = func.call @ggml_iq1s_grid_code_i32(%gi) : (i32) -> (i32) + %half_i32 = index.cast %half : index to i32 + %result = func.call @ggml_iq1_code_vector4(%code, %half_i32, %delta, %scale) : (i32, i32, f32, f32) -> (vector<4xf32>) + func.return %result : vector<4xf32> +} + +func.def inline @ggml_iq1m_f16_vector4(%weight: buffer, %row_byte_base: offset, %iq1_block: index, %iq1_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_iq1m_f32_vector4(%weight, %row_byte_base, %iq1_block, %iq1_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} +// iq3xxs_code: 256 grid entries as 12-bit codes (3 bits per value: level L, value 4 + 8 L, 62 for L = 7) +func.def inline @ggml_iq3xxs_grid_code_i32(%grid_index: i32) -> (i32) { + %c5_i32 = scalar.constant 5 : i32 + %c31_i32 = scalar.constant 31 : i32 + %chunk_id1 = scalar.constant 1 : i32 + %chunk_id2 = scalar.constant 2 : i32 + %chunk_id3 = scalar.constant 3 : i32 + %chunk_id4 = scalar.constant 4 : i32 + %chunk_id5 = scalar.constant 5 : i32 + %chunk_id6 = scalar.constant 6 : i32 + %chunk_id7 = scalar.constant 7 : i32 + %chunk_i32 = scalar.shrui %grid_index, %c5_i32 : i32 + %lane_i32 = scalar.andi %grid_index, %c31_i32 : i32 + %codes = vector.from_elements %lane_i32 : vector<1xi32> + %iq3xxs_code_0_0 = scalar.constant 0.0 : f32 + %iq3xxs_code_0_1 = scalar.constant 2.0 : f32 + %iq3xxs_code_0_2 = scalar.constant 4.0 : f32 + %iq3xxs_code_0_3 = scalar.constant 9.0 : f32 + %iq3xxs_code_0_4 = scalar.constant 11.0 : f32 + %iq3xxs_code_0_5 = scalar.constant 15.0 : f32 + %iq3xxs_code_0_6 = scalar.constant 16.0 : f32 + %iq3xxs_code_0_7 = scalar.constant 18.0 : f32 + %iq3xxs_code_0_8 = scalar.constant 25.0 : f32 + %iq3xxs_code_0_9 = scalar.constant 34.0 : f32 + %iq3xxs_code_0_10 = scalar.constant 59.0 : f32 + %iq3xxs_code_0_11 = scalar.constant 61.0 : f32 + %iq3xxs_code_0_12 = scalar.constant 65.0 : f32 + %iq3xxs_code_0_13 = scalar.constant 67.0 : f32 + %iq3xxs_code_0_14 = scalar.constant 72.0 : f32 + %iq3xxs_code_0_15 = scalar.constant 74.0 : f32 + %iq3xxs_code_0_16 = scalar.constant 81.0 : f32 + %iq3xxs_code_0_17 = scalar.constant 85.0 : f32 + %iq3xxs_code_0_18 = scalar.constant 88.0 : f32 + %iq3xxs_code_0_19 = scalar.constant 90.0 : f32 + %iq3xxs_code_0_20 = scalar.constant 97.0 : f32 + %iq3xxs_code_0_21 = scalar.constant 108.0 : f32 + %iq3xxs_code_0_22 = scalar.constant 120.0 : f32 + %iq3xxs_code_0_23 = scalar.constant 128.0 : f32 + %iq3xxs_code_0_24 = scalar.constant 130.0 : f32 + %iq3xxs_code_0_25 = scalar.constant 132.0 : f32 + %iq3xxs_code_0_26 = scalar.constant 137.0 : f32 + %iq3xxs_code_0_27 = scalar.constant 144.0 : f32 + %iq3xxs_code_0_28 = scalar.constant 146.0 : f32 + %iq3xxs_code_0_29 = scalar.constant 153.0 : f32 + %iq3xxs_code_0_30 = scalar.constant 155.0 : f32 + %iq3xxs_code_0_31 = scalar.constant 159.0 : f32 + %iq3xxs_code_0 = vector.from_elements %iq3xxs_code_0_0, %iq3xxs_code_0_1, %iq3xxs_code_0_2, %iq3xxs_code_0_3, %iq3xxs_code_0_4, %iq3xxs_code_0_5, %iq3xxs_code_0_6, %iq3xxs_code_0_7, %iq3xxs_code_0_8, %iq3xxs_code_0_9, %iq3xxs_code_0_10, %iq3xxs_code_0_11, %iq3xxs_code_0_12, %iq3xxs_code_0_13, %iq3xxs_code_0_14, %iq3xxs_code_0_15, %iq3xxs_code_0_16, %iq3xxs_code_0_17, %iq3xxs_code_0_18, %iq3xxs_code_0_19, %iq3xxs_code_0_20, %iq3xxs_code_0_21, %iq3xxs_code_0_22, %iq3xxs_code_0_23, %iq3xxs_code_0_24, %iq3xxs_code_0_25, %iq3xxs_code_0_26, %iq3xxs_code_0_27, %iq3xxs_code_0_28, %iq3xxs_code_0_29, %iq3xxs_code_0_30, %iq3xxs_code_0_31 : vector<32xf32> + %iq3xxs_code_1_0 = scalar.constant 169.0 : f32 + %iq3xxs_code_1_1 = scalar.constant 175.0 : f32 + %iq3xxs_code_1_2 = scalar.constant 189.0 : f32 + %iq3xxs_code_1_3 = scalar.constant 193.0 : f32 + %iq3xxs_code_1_4 = scalar.constant 199.0 : f32 + %iq3xxs_code_1_5 = scalar.constant 200.0 : f32 + %iq3xxs_code_1_6 = scalar.constant 202.0 : f32 + %iq3xxs_code_1_7 = scalar.constant 213.0 : f32 + %iq3xxs_code_1_8 = scalar.constant 248.0 : f32 + %iq3xxs_code_1_9 = scalar.constant 267.0 : f32 + %iq3xxs_code_1_10 = scalar.constant 287.0 : f32 + %iq3xxs_code_1_11 = scalar.constant 292.0 : f32 + %iq3xxs_code_1_12 = scalar.constant 303.0 : f32 + %iq3xxs_code_1_13 = scalar.constant 315.0 : f32 + %iq3xxs_code_1_14 = scalar.constant 317.0 : f32 + %iq3xxs_code_1_15 = scalar.constant 321.0 : f32 + %iq3xxs_code_1_16 = scalar.constant 327.0 : f32 + %iq3xxs_code_1_17 = scalar.constant 346.0 : f32 + %iq3xxs_code_1_18 = scalar.constant 362.0 : f32 + %iq3xxs_code_1_19 = scalar.constant 413.0 : f32 + %iq3xxs_code_1_20 = scalar.constant 436.0 : f32 + %iq3xxs_code_1_21 = scalar.constant 456.0 : f32 + %iq3xxs_code_1_22 = scalar.constant 460.0 : f32 + %iq3xxs_code_1_23 = scalar.constant 462.0 : f32 + %iq3xxs_code_1_24 = scalar.constant 483.0 : f32 + %iq3xxs_code_1_25 = scalar.constant 497.0 : f32 + %iq3xxs_code_1_26 = scalar.constant 513.0 : f32 + %iq3xxs_code_1_27 = scalar.constant 515.0 : f32 + %iq3xxs_code_1_28 = scalar.constant 520.0 : f32 + %iq3xxs_code_1_29 = scalar.constant 522.0 : f32 + %iq3xxs_code_1_30 = scalar.constant 529.0 : f32 + %iq3xxs_code_1_31 = scalar.constant 531.0 : f32 + %iq3xxs_code_1 = vector.from_elements %iq3xxs_code_1_0, %iq3xxs_code_1_1, %iq3xxs_code_1_2, %iq3xxs_code_1_3, %iq3xxs_code_1_4, %iq3xxs_code_1_5, %iq3xxs_code_1_6, %iq3xxs_code_1_7, %iq3xxs_code_1_8, %iq3xxs_code_1_9, %iq3xxs_code_1_10, %iq3xxs_code_1_11, %iq3xxs_code_1_12, %iq3xxs_code_1_13, %iq3xxs_code_1_14, %iq3xxs_code_1_15, %iq3xxs_code_1_16, %iq3xxs_code_1_17, %iq3xxs_code_1_18, %iq3xxs_code_1_19, %iq3xxs_code_1_20, %iq3xxs_code_1_21, %iq3xxs_code_1_22, %iq3xxs_code_1_23, %iq3xxs_code_1_24, %iq3xxs_code_1_25, %iq3xxs_code_1_26, %iq3xxs_code_1_27, %iq3xxs_code_1_28, %iq3xxs_code_1_29, %iq3xxs_code_1_30, %iq3xxs_code_1_31 : vector<32xf32> + %iq3xxs_code_2_0 = scalar.constant 536.0 : f32 + %iq3xxs_code_2_1 = scalar.constant 538.0 : f32 + %iq3xxs_code_2_2 = scalar.constant 540.0 : f32 + %iq3xxs_code_2_3 = scalar.constant 551.0 : f32 + %iq3xxs_code_2_4 = scalar.constant 552.0 : f32 + %iq3xxs_code_2_5 = scalar.constant 576.0 : f32 + %iq3xxs_code_2_6 = scalar.constant 578.0 : f32 + %iq3xxs_code_2_7 = scalar.constant 585.0 : f32 + %iq3xxs_code_2_8 = scalar.constant 592.0 : f32 + %iq3xxs_code_2_9 = scalar.constant 594.0 : f32 + %iq3xxs_code_2_10 = scalar.constant 641.0 : f32 + %iq3xxs_code_2_11 = scalar.constant 643.0 : f32 + %iq3xxs_code_2_12 = scalar.constant 648.0 : f32 + %iq3xxs_code_2_13 = scalar.constant 650.0 : f32 + %iq3xxs_code_2_14 = scalar.constant 657.0 : f32 + %iq3xxs_code_2_15 = scalar.constant 664.0 : f32 + %iq3xxs_code_2_16 = scalar.constant 698.0 : f32 + %iq3xxs_code_2_17 = scalar.constant 704.0 : f32 + %iq3xxs_code_2_18 = scalar.constant 706.0 : f32 + %iq3xxs_code_2_19 = scalar.constant 720.0 : f32 + %iq3xxs_code_2_20 = scalar.constant 729.0 : f32 + %iq3xxs_code_2_21 = scalar.constant 742.0 : f32 + %iq3xxs_code_2_22 = scalar.constant 758.0 : f32 + %iq3xxs_code_2_23 = scalar.constant 769.0 : f32 + %iq3xxs_code_2_24 = scalar.constant 773.0 : f32 + %iq3xxs_code_2_25 = scalar.constant 808.0 : f32 + %iq3xxs_code_2_26 = scalar.constant 848.0 : f32 + %iq3xxs_code_2_27 = scalar.constant 852.0 : f32 + %iq3xxs_code_2_28 = scalar.constant 870.0 : f32 + %iq3xxs_code_2_29 = scalar.constant 889.0 : f32 + %iq3xxs_code_2_30 = scalar.constant 901.0 : f32 + %iq3xxs_code_2_31 = scalar.constant 978.0 : f32 + %iq3xxs_code_2 = vector.from_elements %iq3xxs_code_2_0, %iq3xxs_code_2_1, %iq3xxs_code_2_2, %iq3xxs_code_2_3, %iq3xxs_code_2_4, %iq3xxs_code_2_5, %iq3xxs_code_2_6, %iq3xxs_code_2_7, %iq3xxs_code_2_8, %iq3xxs_code_2_9, %iq3xxs_code_2_10, %iq3xxs_code_2_11, %iq3xxs_code_2_12, %iq3xxs_code_2_13, %iq3xxs_code_2_14, %iq3xxs_code_2_15, %iq3xxs_code_2_16, %iq3xxs_code_2_17, %iq3xxs_code_2_18, %iq3xxs_code_2_19, %iq3xxs_code_2_20, %iq3xxs_code_2_21, %iq3xxs_code_2_22, %iq3xxs_code_2_23, %iq3xxs_code_2_24, %iq3xxs_code_2_25, %iq3xxs_code_2_26, %iq3xxs_code_2_27, %iq3xxs_code_2_28, %iq3xxs_code_2_29, %iq3xxs_code_2_30, %iq3xxs_code_2_31 : vector<32xf32> + %iq3xxs_code_3_0 = scalar.constant 992.0 : f32 + %iq3xxs_code_3_1 = scalar.constant 1024.0 : f32 + %iq3xxs_code_3_2 = scalar.constant 1026.0 : f32 + %iq3xxs_code_3_3 = scalar.constant 1033.0 : f32 + %iq3xxs_code_3_4 = scalar.constant 1035.0 : f32 + %iq3xxs_code_3_5 = scalar.constant 1040.0 : f32 + %iq3xxs_code_3_6 = scalar.constant 1042.0 : f32 + %iq3xxs_code_3_7 = scalar.constant 1046.0 : f32 + %iq3xxs_code_3_8 = scalar.constant 1049.0 : f32 + %iq3xxs_code_3_9 = scalar.constant 1058.0 : f32 + %iq3xxs_code_3_10 = scalar.constant 1089.0 : f32 + %iq3xxs_code_3_11 = scalar.constant 1091.0 : f32 + %iq3xxs_code_3_12 = scalar.constant 1093.0 : f32 + %iq3xxs_code_3_13 = scalar.constant 1096.0 : f32 + %iq3xxs_code_3_14 = scalar.constant 1098.0 : f32 + %iq3xxs_code_3_15 = scalar.constant 1105.0 : f32 + %iq3xxs_code_3_16 = scalar.constant 1112.0 : f32 + %iq3xxs_code_3_17 = scalar.constant 1139.0 : f32 + %iq3xxs_code_3_18 = scalar.constant 1143.0 : f32 + %iq3xxs_code_3_19 = scalar.constant 1144.0 : f32 + %iq3xxs_code_3_20 = scalar.constant 1152.0 : f32 + %iq3xxs_code_3_21 = scalar.constant 1154.0 : f32 + %iq3xxs_code_3_22 = scalar.constant 1161.0 : f32 + %iq3xxs_code_3_23 = scalar.constant 1167.0 : f32 + %iq3xxs_code_3_24 = scalar.constant 1168.0 : f32 + %iq3xxs_code_3_25 = scalar.constant 1170.0 : f32 + %iq3xxs_code_3_26 = scalar.constant 1183.0 : f32 + %iq3xxs_code_3_27 = scalar.constant 1184.0 : f32 + %iq3xxs_code_3_28 = scalar.constant 1197.0 : f32 + %iq3xxs_code_3_29 = scalar.constant 1217.0 : f32 + %iq3xxs_code_3_30 = scalar.constant 1224.0 : f32 + %iq3xxs_code_3_31 = scalar.constant 1228.0 : f32 + %iq3xxs_code_3 = vector.from_elements %iq3xxs_code_3_0, %iq3xxs_code_3_1, %iq3xxs_code_3_2, %iq3xxs_code_3_3, %iq3xxs_code_3_4, %iq3xxs_code_3_5, %iq3xxs_code_3_6, %iq3xxs_code_3_7, %iq3xxs_code_3_8, %iq3xxs_code_3_9, %iq3xxs_code_3_10, %iq3xxs_code_3_11, %iq3xxs_code_3_12, %iq3xxs_code_3_13, %iq3xxs_code_3_14, %iq3xxs_code_3_15, %iq3xxs_code_3_16, %iq3xxs_code_3_17, %iq3xxs_code_3_18, %iq3xxs_code_3_19, %iq3xxs_code_3_20, %iq3xxs_code_3_21, %iq3xxs_code_3_22, %iq3xxs_code_3_23, %iq3xxs_code_3_24, %iq3xxs_code_3_25, %iq3xxs_code_3_26, %iq3xxs_code_3_27, %iq3xxs_code_3_28, %iq3xxs_code_3_29, %iq3xxs_code_3_30, %iq3xxs_code_3_31 : vector<32xf32> + %iq3xxs_code_4_0 = scalar.constant 1272.0 : f32 + %iq3xxs_code_4_1 = scalar.constant 1276.0 : f32 + %iq3xxs_code_4_2 = scalar.constant 1309.0 : f32 + %iq3xxs_code_4_3 = scalar.constant 1323.0 : f32 + %iq3xxs_code_4_4 = scalar.constant 1347.0 : f32 + %iq3xxs_code_4_5 = scalar.constant 1367.0 : f32 + %iq3xxs_code_4_6 = scalar.constant 1377.0 : f32 + %iq3xxs_code_4_7 = scalar.constant 1404.0 : f32 + %iq3xxs_code_4_8 = scalar.constant 1473.0 : f32 + %iq3xxs_code_4_9 = scalar.constant 1475.0 : f32 + %iq3xxs_code_4_10 = scalar.constant 1486.0 : f32 + %iq3xxs_code_4_11 = scalar.constant 1509.0 : f32 + %iq3xxs_code_4_12 = scalar.constant 1537.0 : f32 + %iq3xxs_code_4_13 = scalar.constant 1544.0 : f32 + %iq3xxs_code_4_14 = scalar.constant 1546.0 : f32 + %iq3xxs_code_4_15 = scalar.constant 1553.0 : f32 + %iq3xxs_code_4_16 = scalar.constant 1555.0 : f32 + %iq3xxs_code_4_17 = scalar.constant 1576.0 : f32 + %iq3xxs_code_4_18 = scalar.constant 1589.0 : f32 + %iq3xxs_code_4_19 = scalar.constant 1594.0 : f32 + %iq3xxs_code_4_20 = scalar.constant 1600.0 : f32 + %iq3xxs_code_4_21 = scalar.constant 1602.0 : f32 + %iq3xxs_code_4_22 = scalar.constant 1616.0 : f32 + %iq3xxs_code_4_23 = scalar.constant 1625.0 : f32 + %iq3xxs_code_4_24 = scalar.constant 1636.0 : f32 + %iq3xxs_code_4_25 = scalar.constant 1638.0 : f32 + %iq3xxs_code_4_26 = scalar.constant 1665.0 : f32 + %iq3xxs_code_4_27 = scalar.constant 1667.0 : f32 + %iq3xxs_code_4_28 = scalar.constant 1672.0 : f32 + %iq3xxs_code_4_29 = scalar.constant 1685.0 : f32 + %iq3xxs_code_4_30 = scalar.constant 1706.0 : f32 + %iq3xxs_code_4_31 = scalar.constant 1722.0 : f32 + %iq3xxs_code_4 = vector.from_elements %iq3xxs_code_4_0, %iq3xxs_code_4_1, %iq3xxs_code_4_2, %iq3xxs_code_4_3, %iq3xxs_code_4_4, %iq3xxs_code_4_5, %iq3xxs_code_4_6, %iq3xxs_code_4_7, %iq3xxs_code_4_8, %iq3xxs_code_4_9, %iq3xxs_code_4_10, %iq3xxs_code_4_11, %iq3xxs_code_4_12, %iq3xxs_code_4_13, %iq3xxs_code_4_14, %iq3xxs_code_4_15, %iq3xxs_code_4_16, %iq3xxs_code_4_17, %iq3xxs_code_4_18, %iq3xxs_code_4_19, %iq3xxs_code_4_20, %iq3xxs_code_4_21, %iq3xxs_code_4_22, %iq3xxs_code_4_23, %iq3xxs_code_4_24, %iq3xxs_code_4_25, %iq3xxs_code_4_26, %iq3xxs_code_4_27, %iq3xxs_code_4_28, %iq3xxs_code_4_29, %iq3xxs_code_4_30, %iq3xxs_code_4_31 : vector<32xf32> + %iq3xxs_code_5_0 = scalar.constant 1737.0 : f32 + %iq3xxs_code_5_1 = scalar.constant 1755.0 : f32 + %iq3xxs_code_5_2 = scalar.constant 1816.0 : f32 + %iq3xxs_code_5_3 = scalar.constant 1831.0 : f32 + %iq3xxs_code_5_4 = scalar.constant 1850.0 : f32 + %iq3xxs_code_5_5 = scalar.constant 1856.0 : f32 + %iq3xxs_code_5_6 = scalar.constant 1862.0 : f32 + %iq3xxs_code_5_7 = scalar.constant 1874.0 : f32 + %iq3xxs_code_5_8 = scalar.constant 1901.0 : f32 + %iq3xxs_code_5_9 = scalar.constant 1932.0 : f32 + %iq3xxs_code_5_10 = scalar.constant 1950.0 : f32 + %iq3xxs_code_5_11 = scalar.constant 1971.0 : f32 + %iq3xxs_code_5_12 = scalar.constant 2011.0 : f32 + %iq3xxs_code_5_13 = scalar.constant 2032.0 : f32 + %iq3xxs_code_5_14 = scalar.constant 2052.0 : f32 + %iq3xxs_code_5_15 = scalar.constant 2063.0 : f32 + %iq3xxs_code_5_16 = scalar.constant 2077.0 : f32 + %iq3xxs_code_5_17 = scalar.constant 2079.0 : f32 + %iq3xxs_code_5_18 = scalar.constant 2091.0 : f32 + %iq3xxs_code_5_19 = scalar.constant 2095.0 : f32 + %iq3xxs_code_5_20 = scalar.constant 2172.0 : f32 + %iq3xxs_code_5_21 = scalar.constant 2192.0 : f32 + %iq3xxs_code_5_22 = scalar.constant 2207.0 : f32 + %iq3xxs_code_5_23 = scalar.constant 2208.0 : f32 + %iq3xxs_code_5_24 = scalar.constant 2224.0 : f32 + %iq3xxs_code_5_25 = scalar.constant 2230.0 : f32 + %iq3xxs_code_5_26 = scalar.constant 2247.0 : f32 + %iq3xxs_code_5_27 = scalar.constant 2277.0 : f32 + %iq3xxs_code_5_28 = scalar.constant 2308.0 : f32 + %iq3xxs_code_5_29 = scalar.constant 2345.0 : f32 + %iq3xxs_code_5_30 = scalar.constant 2356.0 : f32 + %iq3xxs_code_5_31 = scalar.constant 2389.0 : f32 + %iq3xxs_code_5 = vector.from_elements %iq3xxs_code_5_0, %iq3xxs_code_5_1, %iq3xxs_code_5_2, %iq3xxs_code_5_3, %iq3xxs_code_5_4, %iq3xxs_code_5_5, %iq3xxs_code_5_6, %iq3xxs_code_5_7, %iq3xxs_code_5_8, %iq3xxs_code_5_9, %iq3xxs_code_5_10, %iq3xxs_code_5_11, %iq3xxs_code_5_12, %iq3xxs_code_5_13, %iq3xxs_code_5_14, %iq3xxs_code_5_15, %iq3xxs_code_5_16, %iq3xxs_code_5_17, %iq3xxs_code_5_18, %iq3xxs_code_5_19, %iq3xxs_code_5_20, %iq3xxs_code_5_21, %iq3xxs_code_5_22, %iq3xxs_code_5_23, %iq3xxs_code_5_24, %iq3xxs_code_5_25, %iq3xxs_code_5_26, %iq3xxs_code_5_27, %iq3xxs_code_5_28, %iq3xxs_code_5_29, %iq3xxs_code_5_30, %iq3xxs_code_5_31 : vector<32xf32> + %iq3xxs_code_6_0 = scalar.constant 2403.0 : f32 + %iq3xxs_code_6_1 = scalar.constant 2424.0 : f32 + %iq3xxs_code_6_2 = scalar.constant 2501.0 : f32 + %iq3xxs_code_6_3 = scalar.constant 2504.0 : f32 + %iq3xxs_code_6_4 = scalar.constant 2506.0 : f32 + %iq3xxs_code_6_5 = scalar.constant 2520.0 : f32 + %iq3xxs_code_6_6 = scalar.constant 2570.0 : f32 + %iq3xxs_code_6_7 = scalar.constant 2593.0 : f32 + %iq3xxs_code_6_8 = scalar.constant 2616.0 : f32 + %iq3xxs_code_6_9 = scalar.constant 2624.0 : f32 + %iq3xxs_code_6_10 = scalar.constant 2630.0 : f32 + %iq3xxs_code_6_11 = scalar.constant 2646.0 : f32 + %iq3xxs_code_6_12 = scalar.constant 2669.0 : f32 + %iq3xxs_code_6_13 = scalar.constant 2700.0 : f32 + %iq3xxs_code_6_14 = scalar.constant 2714.0 : f32 + %iq3xxs_code_6_15 = scalar.constant 2746.0 : f32 + %iq3xxs_code_6_16 = scalar.constant 2754.0 : f32 + %iq3xxs_code_6_17 = scalar.constant 2795.0 : f32 + %iq3xxs_code_6_18 = scalar.constant 2824.0 : f32 + %iq3xxs_code_6_19 = scalar.constant 2835.0 : f32 + %iq3xxs_code_6_20 = scalar.constant 2839.0 : f32 + %iq3xxs_code_6_21 = scalar.constant 2874.0 : f32 + %iq3xxs_code_6_22 = scalar.constant 2882.0 : f32 + %iq3xxs_code_6_23 = scalar.constant 2905.0 : f32 + %iq3xxs_code_6_24 = scalar.constant 2984.0 : f32 + %iq3xxs_code_6_25 = scalar.constant 3028.0 : f32 + %iq3xxs_code_6_26 = scalar.constant 3042.0 : f32 + %iq3xxs_code_6_27 = scalar.constant 3092.0 : f32 + %iq3xxs_code_6_28 = scalar.constant 3108.0 : f32 + %iq3xxs_code_6_29 = scalar.constant 3110.0 : f32 + %iq3xxs_code_6_30 = scalar.constant 3124.0 : f32 + %iq3xxs_code_6_31 = scalar.constant 3153.0 : f32 + %iq3xxs_code_6 = vector.from_elements %iq3xxs_code_6_0, %iq3xxs_code_6_1, %iq3xxs_code_6_2, %iq3xxs_code_6_3, %iq3xxs_code_6_4, %iq3xxs_code_6_5, %iq3xxs_code_6_6, %iq3xxs_code_6_7, %iq3xxs_code_6_8, %iq3xxs_code_6_9, %iq3xxs_code_6_10, %iq3xxs_code_6_11, %iq3xxs_code_6_12, %iq3xxs_code_6_13, %iq3xxs_code_6_14, %iq3xxs_code_6_15, %iq3xxs_code_6_16, %iq3xxs_code_6_17, %iq3xxs_code_6_18, %iq3xxs_code_6_19, %iq3xxs_code_6_20, %iq3xxs_code_6_21, %iq3xxs_code_6_22, %iq3xxs_code_6_23, %iq3xxs_code_6_24, %iq3xxs_code_6_25, %iq3xxs_code_6_26, %iq3xxs_code_6_27, %iq3xxs_code_6_28, %iq3xxs_code_6_29, %iq3xxs_code_6_30, %iq3xxs_code_6_31 : vector<32xf32> + %iq3xxs_code_7_0 = scalar.constant 3185.0 : f32 + %iq3xxs_code_7_1 = scalar.constant 3215.0 : f32 + %iq3xxs_code_7_2 = scalar.constant 3252.0 : f32 + %iq3xxs_code_7_3 = scalar.constant 3288.0 : f32 + %iq3xxs_code_7_4 = scalar.constant 3294.0 : f32 + %iq3xxs_code_7_5 = scalar.constant 3364.0 : f32 + %iq3xxs_code_7_6 = scalar.constant 3397.0 : f32 + %iq3xxs_code_7_7 = scalar.constant 3434.0 : f32 + %iq3xxs_code_7_8 = scalar.constant 3483.0 : f32 + %iq3xxs_code_7_9 = scalar.constant 3523.0 : f32 + %iq3xxs_code_7_10 = scalar.constant 3537.0 : f32 + %iq3xxs_code_7_11 = scalar.constant 3587.0 : f32 + %iq3xxs_code_7_12 = scalar.constant 3589.0 : f32 + %iq3xxs_code_7_13 = scalar.constant 3591.0 : f32 + %iq3xxs_code_7_14 = scalar.constant 3592.0 : f32 + %iq3xxs_code_7_15 = scalar.constant 3610.0 : f32 + %iq3xxs_code_7_16 = scalar.constant 3626.0 : f32 + %iq3xxs_code_7_17 = scalar.constant 3670.0 : f32 + %iq3xxs_code_7_18 = scalar.constant 3680.0 : f32 + %iq3xxs_code_7_19 = scalar.constant 3722.0 : f32 + %iq3xxs_code_7_20 = scalar.constant 3749.0 : f32 + %iq3xxs_code_7_21 = scalar.constant 3754.0 : f32 + %iq3xxs_code_7_22 = scalar.constant 3776.0 : f32 + %iq3xxs_code_7_23 = scalar.constant 3789.0 : f32 + %iq3xxs_code_7_24 = scalar.constant 3803.0 : f32 + %iq3xxs_code_7_25 = scalar.constant 3824.0 : f32 + %iq3xxs_code_7_26 = scalar.constant 3857.0 : f32 + %iq3xxs_code_7_27 = scalar.constant 3873.0 : f32 + %iq3xxs_code_7_28 = scalar.constant 3904.0 : f32 + %iq3xxs_code_7_29 = scalar.constant 3906.0 : f32 + %iq3xxs_code_7_30 = scalar.constant 3924.0 : f32 + %iq3xxs_code_7_31 = scalar.constant 3992.0 : f32 + %iq3xxs_code_7 = vector.from_elements %iq3xxs_code_7_0, %iq3xxs_code_7_1, %iq3xxs_code_7_2, %iq3xxs_code_7_3, %iq3xxs_code_7_4, %iq3xxs_code_7_5, %iq3xxs_code_7_6, %iq3xxs_code_7_7, %iq3xxs_code_7_8, %iq3xxs_code_7_9, %iq3xxs_code_7_10, %iq3xxs_code_7_11, %iq3xxs_code_7_12, %iq3xxs_code_7_13, %iq3xxs_code_7_14, %iq3xxs_code_7_15, %iq3xxs_code_7_16, %iq3xxs_code_7_17, %iq3xxs_code_7_18, %iq3xxs_code_7_19, %iq3xxs_code_7_20, %iq3xxs_code_7_21, %iq3xxs_code_7_22, %iq3xxs_code_7_23, %iq3xxs_code_7_24, %iq3xxs_code_7_25, %iq3xxs_code_7_26, %iq3xxs_code_7_27, %iq3xxs_code_7_28, %iq3xxs_code_7_29, %iq3xxs_code_7_30, %iq3xxs_code_7_31 : vector<32xf32> + %is_chunk1 = scalar.cmpi eq, %chunk_i32, %chunk_id1 : i32 + %sel1 = scf.select %is_chunk1, %iq3xxs_code_1, %iq3xxs_code_0 : vector<32xf32> + %is_chunk2 = scalar.cmpi eq, %chunk_i32, %chunk_id2 : i32 + %sel2 = scf.select %is_chunk2, %iq3xxs_code_2, %sel1 : vector<32xf32> + %is_chunk3 = scalar.cmpi eq, %chunk_i32, %chunk_id3 : i32 + %sel3 = scf.select %is_chunk3, %iq3xxs_code_3, %sel2 : vector<32xf32> + %is_chunk4 = scalar.cmpi eq, %chunk_i32, %chunk_id4 : i32 + %sel4 = scf.select %is_chunk4, %iq3xxs_code_4, %sel3 : vector<32xf32> + %is_chunk5 = scalar.cmpi eq, %chunk_i32, %chunk_id5 : i32 + %sel5 = scf.select %is_chunk5, %iq3xxs_code_5, %sel4 : vector<32xf32> + %is_chunk6 = scalar.cmpi eq, %chunk_i32, %chunk_id6 : i32 + %sel6 = scf.select %is_chunk6, %iq3xxs_code_6, %sel5 : vector<32xf32> + %is_chunk7 = scalar.cmpi eq, %chunk_i32, %chunk_id7 : i32 + %sel7 = scf.select %is_chunk7, %iq3xxs_code_7, %sel6 : vector<32xf32> + %v = vector.table.lookup %sel7[%codes] : vector<32xf32>, vector<1xi32> -> vector<1xf32> + %v_f32 = vector.extract %v[0] : vector<1xf32> -> f32 + %code = scalar.fptoui %v_f32 : f32 to i32 + func.return %code : i32 +} + +// the four values of a 12-bit IQ3_XXS grid code, with sign bits 4 half .. +3 of %signs8, times %scale +func.def inline @ggml_iq3xxs_code_vector4(%code: i32, %half: i32, %signs8: i32, %scale: f32) -> (vector<4xf32>) { + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c8_i32 = scalar.constant 8 : i32 + %s0 = scalar.constant 0 : i32 + %s3 = scalar.constant 3 : i32 + %s6 = scalar.constant 6 : i32 + %s9 = scalar.constant 9 : i32 + %shift3 = vector.from_elements %s0, %s3, %s6, %s9 : vector<4xi32> + %cv = vector.splat %code : vector<4xi32> + %lvs = vector.shrui %cv, %shift3 : vector<4xi32> + %seven = vector.splat %c7_i32 : vector<4xi32> + %lv = vector.andi %lvs, %seven : vector<4xi32> + %eight = vector.splat %c8_i32 : vector<4xi32> + %four = vector.splat %c4_i32 : vector<4xi32> + %lv8 = vector.muli %lv, %eight : vector<4xi32> + %base = vector.addi %lv8, %four : vector<4xi32> + // level 7 is the only one with all three bits set: bump = 2 (L & L >> 1 & L >> 2 & 1) + %one7 = vector.splat %c1_i32 : vector<4xi32> + %two7 = vector.splat %c2_i32 : vector<4xi32> + %l1 = vector.shrui %lv, %one7 : vector<4xi32> + %l2 = vector.shrui %lv, %two7 : vector<4xi32> + %a01 = vector.andi %lv, %l1 : vector<4xi32> + %a012 = vector.andi %a01, %l2 : vector<4xi32> + %is7 = vector.andi %a012, %one7 : vector<4xi32> + %bump = vector.shli %is7, %one7 : vector<4xi32> + %mag = vector.addi %base, %bump : vector<4xi32> + %sbase = scalar.muli %half, %c4_i32 : i32 + %c1v = scalar.constant 1 : i32 + %c2v = scalar.constant 2 : i32 + %c3v = scalar.constant 3 : i32 + %sh0 = vector.from_elements %s0, %c1v, %c2v, %c3v : vector<4xi32> + %sbv = vector.splat %sbase : vector<4xi32> + %sh = vector.addi %sh0, %sbv : vector<4xi32> + %sv = vector.splat %signs8 : vector<4xi32> + %sb0 = vector.shrui %sv, %sh : vector<4xi32> + %one = vector.splat %c1_i32 : vector<4xi32> + %sb = vector.andi %sb0, %one : vector<4xi32> + %sb2 = vector.shli %sb, %one : vector<4xi32> + %sgn = vector.subi %one, %sb2 : vector<4xi32> + %signed = vector.muli %mag, %sgn : vector<4xi32> + %f = vector.sitofp %signed : vector<4xi32> to vector<4xf32> + %scv = vector.splat %scale : vector<4xf32> + %result = vector.mulf %f, %scv : vector<4xf32> + func.return %result : vector<4xf32> +} + +// IQ3_XXS (98 bytes: d, qs[64] grid indices, 8 x u32 scales and signs; dequantize_row_iq3_xxs). Group g +// has grid indices qs[8 g .. 8 g + 7] (bytes 2 + 8 g ..) and aux = u32 at byte 66 + 4 g: four 7-bit sign +// groups (slot l at bit 7 l) and a 4-bit scale on top: value = d * (0.5 + (aux >> 28)) * 0.5 * grid * sign. +// Packet p covers slot p / 2, grid entry 2 (p / 2) + p % 2 (its four values, sign bits 4 (p % 2) ..). +func.def inline @ggml_iq3xxs_f32_vector4(%weight: buffer, %row_byte_base: offset, %iq3_block: index, %iq3_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c8 = index.constant 8 : index + %c33 = index.constant 33 : index + %c7_i32 = scalar.constant 7 : i32 + %c16_i32 = scalar.constant 16 : i32 + %c28_i32 = scalar.constant 28 : i32 + %c127_i32 = scalar.constant 127 : i32 + %c05_f32 = scalar.constant 0.5 : f32 + %block_bytes = index.constant 98 : offset + %block_byte_add = index.scale %iq3_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %hv = buffer.view %weight[%block_byte_base] : buffer -> view<49xf16> + %wv = buffer.view %weight[%block_byte_base] : buffer -> view<49xi16> + %bv = buffer.view %weight[%block_byte_base] : buffer -> view<98xi8> + %g = index.assume %iq3_group [range(%iq3_group, 0, 7)] : index + %p = index.assume %packet [range(%packet, 0, 7)] : index + %slot = index.div %p, %c2 : index + %half = index.rem %p, %c2 : index + %g8 = index.mul %g, %c8 : index + %idx0 = index.add %g8, %c2 : index + %idx_at = index.add %idx0, %p : index + %g2 = index.add %g, %g : index + %aux_w0 = index.add %c33, %g2 : index + %aux_w1 = index.add %aux_w0, %c1 : index + %d_f16 = view.load %hv[%c0] : view<49xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %gi_i8 = view.load %bv[%idx_at] : view<98xi8> -> i8 + %gi = scalar.extui %gi_i8 : i8 to i32 + %w0_i16 = view.load %wv[%aux_w0] : view<49xi16> -> i16 + %w1_i16 = view.load %wv[%aux_w1] : view<49xi16> -> i16 + %w0 = scalar.extui %w0_i16 : i16 to i32 + %w1 = scalar.extui %w1_i16 : i16 to i32 + %w1s = scalar.shli %w1, %c16_i32 : i32 + %aux = scalar.ori %w0, %w1s : i32 + %sc4 = scalar.shrui %aux, %c28_i32 : i32 + %sc_f = scalar.uitofp %sc4 : i32 to f32 + %sc_plus = scalar.addf %sc_f, %c05_f32 : f32 + %ds = scalar.mulf %d, %sc_plus : f32 + %scale = scalar.mulf %ds, %c05_f32 : f32 + %slot_i32 = index.cast %slot : index to i32 + %sshift = scalar.muli %slot_i32, %c7_i32 : i32 + %sgrp = scalar.shrui %aux, %sshift : i32 + %signs7 = scalar.andi %sgrp, %c127_i32 : i32 + %signs8 = func.call @ggml_iq2_signs8(%signs7) : (i32) -> (i32) + %code = func.call @ggml_iq3xxs_grid_code_i32(%gi) : (i32) -> (i32) + %half_i32 = index.cast %half : index to i32 + %result = func.call @ggml_iq3xxs_code_vector4(%code, %half_i32, %signs8, %scale) : (i32, i32, i32, f32) -> (vector<4xf32>) + func.return %result : vector<4xf32> +} + +func.def inline @ggml_iq3xxs_f16_vector4(%weight: buffer, %row_byte_base: offset, %iq3_block: index, %iq3_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_iq3xxs_f32_vector4(%weight, %row_byte_base, %iq3_block, %iq3_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +func.def inline @ggml_iq3s_grid_f32_halves() -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) { + %grid_low0 = scalar.constant 257.0 : f32 + %grid_high0 = scalar.constant 257.0 : f32 + %grid_low1 = scalar.constant 259.0 : f32 + %grid_high1 = scalar.constant 257.0 : f32 + %grid_low2 = scalar.constant 261.0 : f32 + %grid_high2 = scalar.constant 257.0 : f32 + %grid_low3 = scalar.constant 267.0 : f32 + %grid_high3 = scalar.constant 257.0 : f32 + %grid_low4 = scalar.constant 271.0 : f32 + %grid_high4 = scalar.constant 257.0 : f32 + %grid_low5 = scalar.constant 769.0 : f32 + %grid_high5 = scalar.constant 257.0 : f32 + %grid_low6 = scalar.constant 771.0 : f32 + %grid_high6 = scalar.constant 257.0 : f32 + %grid_low7 = scalar.constant 773.0 : f32 + %grid_high7 = scalar.constant 257.0 : f32 + %grid_low8 = scalar.constant 777.0 : f32 + %grid_high8 = scalar.constant 257.0 : f32 + %grid_low9 = scalar.constant 781.0 : f32 + %grid_high9 = scalar.constant 257.0 : f32 + %grid_low10 = scalar.constant 1281.0 : f32 + %grid_high10 = scalar.constant 257.0 : f32 + %grid_low11 = scalar.constant 1283.0 : f32 + %grid_high11 = scalar.constant 257.0 : f32 + %grid_low12 = scalar.constant 1291.0 : f32 + %grid_high12 = scalar.constant 257.0 : f32 + %grid_low13 = scalar.constant 1799.0 : f32 + %grid_high13 = scalar.constant 257.0 : f32 + %grid_low14 = scalar.constant 2305.0 : f32 + %grid_high14 = scalar.constant 257.0 : f32 + %grid_low15 = scalar.constant 2309.0 : f32 + %grid_high15 = scalar.constant 257.0 : f32 + %grid_low16 = scalar.constant 2315.0 : f32 + %grid_high16 = scalar.constant 257.0 : f32 + %grid_low17 = scalar.constant 2319.0 : f32 + %grid_high17 = scalar.constant 257.0 : f32 + %grid_low18 = scalar.constant 2819.0 : f32 + %grid_high18 = scalar.constant 257.0 : f32 + %grid_low19 = scalar.constant 2823.0 : f32 + %grid_high19 = scalar.constant 257.0 : f32 + %grid_low20 = scalar.constant 3329.0 : f32 + %grid_high20 = scalar.constant 257.0 : f32 + %grid_low21 = scalar.constant 3333.0 : f32 + %grid_high21 = scalar.constant 257.0 : f32 + %grid_low22 = scalar.constant 3843.0 : f32 + %grid_high22 = scalar.constant 257.0 : f32 + %grid_low23 = scalar.constant 3849.0 : f32 + %grid_high23 = scalar.constant 257.0 : f32 + %grid_low24 = scalar.constant 3855.0 : f32 + %grid_high24 = scalar.constant 257.0 : f32 + %grid_low25 = scalar.constant 257.0 : f32 + %grid_high25 = scalar.constant 259.0 : f32 + %grid_low26 = scalar.constant 259.0 : f32 + %grid_high26 = scalar.constant 259.0 : f32 + %grid_low27 = scalar.constant 261.0 : f32 + %grid_high27 = scalar.constant 259.0 : f32 + %grid_low28 = scalar.constant 265.0 : f32 + %grid_high28 = scalar.constant 259.0 : f32 + %grid_low29 = scalar.constant 769.0 : f32 + %grid_high29 = scalar.constant 259.0 : f32 + %grid_low30 = scalar.constant 771.0 : f32 + %grid_high30 = scalar.constant 259.0 : f32 + %grid_low31 = scalar.constant 779.0 : f32 + %grid_high31 = scalar.constant 259.0 : f32 + %grid_low32 = scalar.constant 1281.0 : f32 + %grid_high32 = scalar.constant 259.0 : f32 + %grid_low33 = scalar.constant 1287.0 : f32 + %grid_high33 = scalar.constant 259.0 : f32 + %grid_low34 = scalar.constant 1295.0 : f32 + %grid_high34 = scalar.constant 259.0 : f32 + %grid_low35 = scalar.constant 1795.0 : f32 + %grid_high35 = scalar.constant 259.0 : f32 + %grid_low36 = scalar.constant 1803.0 : f32 + %grid_high36 = scalar.constant 259.0 : f32 + %grid_low37 = scalar.constant 2313.0 : f32 + %grid_high37 = scalar.constant 259.0 : f32 + %grid_low38 = scalar.constant 3331.0 : f32 + %grid_high38 = scalar.constant 259.0 : f32 + %grid_low39 = scalar.constant 3339.0 : f32 + %grid_high39 = scalar.constant 259.0 : f32 + %grid_low40 = scalar.constant 3845.0 : f32 + %grid_high40 = scalar.constant 259.0 : f32 + %grid_low41 = scalar.constant 257.0 : f32 + %grid_high41 = scalar.constant 261.0 : f32 + %grid_low42 = scalar.constant 259.0 : f32 + %grid_high42 = scalar.constant 261.0 : f32 + %grid_low43 = scalar.constant 267.0 : f32 + %grid_high43 = scalar.constant 261.0 : f32 + %grid_low44 = scalar.constant 271.0 : f32 + %grid_high44 = scalar.constant 261.0 : f32 + %grid_low45 = scalar.constant 769.0 : f32 + %grid_high45 = scalar.constant 261.0 : f32 + %grid_low46 = scalar.constant 775.0 : f32 + %grid_high46 = scalar.constant 261.0 : f32 + %grid_low47 = scalar.constant 781.0 : f32 + %grid_high47 = scalar.constant 261.0 : f32 + %grid_low48 = scalar.constant 1283.0 : f32 + %grid_high48 = scalar.constant 261.0 : f32 + %grid_low49 = scalar.constant 1291.0 : f32 + %grid_high49 = scalar.constant 261.0 : f32 + %grid_low50 = scalar.constant 1793.0 : f32 + %grid_high50 = scalar.constant 261.0 : f32 + %grid_low51 = scalar.constant 1801.0 : f32 + %grid_high51 = scalar.constant 261.0 : f32 + %grid_low52 = scalar.constant 2309.0 : f32 + %grid_high52 = scalar.constant 261.0 : f32 + %grid_low53 = scalar.constant 2315.0 : f32 + %grid_high53 = scalar.constant 261.0 : f32 + %grid_low54 = scalar.constant 2319.0 : f32 + %grid_high54 = scalar.constant 261.0 : f32 + %grid_low55 = scalar.constant 2819.0 : f32 + %grid_high55 = scalar.constant 261.0 : f32 + %grid_low56 = scalar.constant 2823.0 : f32 + %grid_high56 = scalar.constant 261.0 : f32 + %grid_low57 = scalar.constant 3841.0 : f32 + %grid_high57 = scalar.constant 261.0 : f32 + %grid_low58 = scalar.constant 3847.0 : f32 + %grid_high58 = scalar.constant 261.0 : f32 + %grid_low59 = scalar.constant 263.0 : f32 + %grid_high59 = scalar.constant 263.0 : f32 + %grid_low60 = scalar.constant 771.0 : f32 + %grid_high60 = scalar.constant 263.0 : f32 + %grid_low61 = scalar.constant 779.0 : f32 + %grid_high61 = scalar.constant 263.0 : f32 + %grid_low62 = scalar.constant 1281.0 : f32 + %grid_high62 = scalar.constant 263.0 : f32 + %grid_low63 = scalar.constant 1285.0 : f32 + %grid_high63 = scalar.constant 263.0 : f32 + %grid_low64 = scalar.constant 1795.0 : f32 + %grid_high64 = scalar.constant 263.0 : f32 + %grid_low65 = scalar.constant 1799.0 : f32 + %grid_high65 = scalar.constant 263.0 : f32 + %grid_low66 = scalar.constant 1805.0 : f32 + %grid_high66 = scalar.constant 263.0 : f32 + %grid_low67 = scalar.constant 2313.0 : f32 + %grid_high67 = scalar.constant 263.0 : f32 + %grid_low68 = scalar.constant 2817.0 : f32 + %grid_high68 = scalar.constant 263.0 : f32 + %grid_low69 = scalar.constant 2821.0 : f32 + %grid_high69 = scalar.constant 263.0 : f32 + %grid_low70 = scalar.constant 3343.0 : f32 + %grid_high70 = scalar.constant 263.0 : f32 + %grid_low71 = scalar.constant 3843.0 : f32 + %grid_high71 = scalar.constant 263.0 : f32 + %grid_low72 = scalar.constant 3851.0 : f32 + %grid_high72 = scalar.constant 263.0 : f32 + %grid_low73 = scalar.constant 257.0 : f32 + %grid_high73 = scalar.constant 265.0 : f32 + %grid_low74 = scalar.constant 775.0 : f32 + %grid_high74 = scalar.constant 265.0 : f32 + %grid_low75 = scalar.constant 783.0 : f32 + %grid_high75 = scalar.constant 265.0 : f32 + %grid_low76 = scalar.constant 1283.0 : f32 + %grid_high76 = scalar.constant 265.0 : f32 + %grid_low77 = scalar.constant 1289.0 : f32 + %grid_high77 = scalar.constant 265.0 : f32 + %grid_low78 = scalar.constant 1797.0 : f32 + %grid_high78 = scalar.constant 265.0 : f32 + %grid_low79 = scalar.constant 2305.0 : f32 + %grid_high79 = scalar.constant 265.0 : f32 + %grid_low80 = scalar.constant 2311.0 : f32 + %grid_high80 = scalar.constant 265.0 : f32 + %grid_low81 = scalar.constant 2819.0 : f32 + %grid_high81 = scalar.constant 265.0 : f32 + %grid_low82 = scalar.constant 3841.0 : f32 + %grid_high82 = scalar.constant 265.0 : f32 + %grid_low83 = scalar.constant 261.0 : f32 + %grid_high83 = scalar.constant 267.0 : f32 + %grid_low84 = scalar.constant 265.0 : f32 + %grid_high84 = scalar.constant 267.0 : f32 + %grid_low85 = scalar.constant 1281.0 : f32 + %grid_high85 = scalar.constant 267.0 : f32 + %grid_low86 = scalar.constant 1285.0 : f32 + %grid_high86 = scalar.constant 267.0 : f32 + %grid_low87 = scalar.constant 1293.0 : f32 + %grid_high87 = scalar.constant 267.0 : f32 + %grid_low88 = scalar.constant 1799.0 : f32 + %grid_high88 = scalar.constant 267.0 : f32 + %grid_low89 = scalar.constant 2307.0 : f32 + %grid_high89 = scalar.constant 267.0 : f32 + %grid_low90 = scalar.constant 2315.0 : f32 + %grid_high90 = scalar.constant 267.0 : f32 + %grid_low91 = scalar.constant 2319.0 : f32 + %grid_high91 = scalar.constant 267.0 : f32 + %grid_low92 = scalar.constant 3341.0 : f32 + %grid_high92 = scalar.constant 267.0 : f32 + %grid_low93 = scalar.constant 3847.0 : f32 + %grid_high93 = scalar.constant 267.0 : f32 + %grid_low94 = scalar.constant 269.0 : f32 + %grid_high94 = scalar.constant 269.0 : f32 + %grid_low95 = scalar.constant 771.0 : f32 + %grid_high95 = scalar.constant 269.0 : f32 + %grid_low96 = scalar.constant 775.0 : f32 + %grid_high96 = scalar.constant 269.0 : f32 + %grid_low97 = scalar.constant 1795.0 : f32 + %grid_high97 = scalar.constant 269.0 : f32 + %grid_low98 = scalar.constant 2821.0 : f32 + %grid_high98 = scalar.constant 269.0 : f32 + %grid_low99 = scalar.constant 3843.0 : f32 + %grid_high99 = scalar.constant 269.0 : f32 + %grid_low100 = scalar.constant 257.0 : f32 + %grid_high100 = scalar.constant 271.0 : f32 + %grid_low101 = scalar.constant 261.0 : f32 + %grid_high101 = scalar.constant 271.0 : f32 + %grid_low102 = scalar.constant 265.0 : f32 + %grid_high102 = scalar.constant 271.0 : f32 + %grid_low103 = scalar.constant 1281.0 : f32 + %grid_high103 = scalar.constant 271.0 : f32 + %grid_low104 = scalar.constant 1285.0 : f32 + %grid_high104 = scalar.constant 271.0 : f32 + %grid_low105 = scalar.constant 1293.0 : f32 + %grid_high105 = scalar.constant 271.0 : f32 + %grid_low106 = scalar.constant 1799.0 : f32 + %grid_high106 = scalar.constant 271.0 : f32 + %grid_low107 = scalar.constant 2817.0 : f32 + %grid_high107 = scalar.constant 271.0 : f32 + %grid_low108 = scalar.constant 2825.0 : f32 + %grid_high108 = scalar.constant 271.0 : f32 + %grid_low109 = scalar.constant 257.0 : f32 + %grid_high109 = scalar.constant 769.0 : f32 + %grid_low110 = scalar.constant 259.0 : f32 + %grid_high110 = scalar.constant 769.0 : f32 + %grid_low111 = scalar.constant 261.0 : f32 + %grid_high111 = scalar.constant 769.0 : f32 + %grid_low112 = scalar.constant 265.0 : f32 + %grid_high112 = scalar.constant 769.0 : f32 + %grid_low113 = scalar.constant 769.0 : f32 + %grid_high113 = scalar.constant 769.0 : f32 + %grid_low114 = scalar.constant 771.0 : f32 + %grid_high114 = scalar.constant 769.0 : f32 + %grid_low115 = scalar.constant 775.0 : f32 + %grid_high115 = scalar.constant 769.0 : f32 + %grid_low116 = scalar.constant 779.0 : f32 + %grid_high116 = scalar.constant 769.0 : f32 + %grid_low117 = scalar.constant 783.0 : f32 + %grid_high117 = scalar.constant 769.0 : f32 + %grid_low118 = scalar.constant 1281.0 : f32 + %grid_high118 = scalar.constant 769.0 : f32 + %grid_low119 = scalar.constant 1285.0 : f32 + %grid_high119 = scalar.constant 769.0 : f32 + %grid_low120 = scalar.constant 1795.0 : f32 + %grid_high120 = scalar.constant 769.0 : f32 + %grid_low121 = scalar.constant 1801.0 : f32 + %grid_high121 = scalar.constant 769.0 : f32 + %grid_low122 = scalar.constant 1805.0 : f32 + %grid_high122 = scalar.constant 769.0 : f32 + %grid_low123 = scalar.constant 2825.0 : f32 + %grid_high123 = scalar.constant 769.0 : f32 + %grid_low124 = scalar.constant 2829.0 : f32 + %grid_high124 = scalar.constant 769.0 : f32 + %grid_low125 = scalar.constant 3331.0 : f32 + %grid_high125 = scalar.constant 769.0 : f32 + %grid_low126 = scalar.constant 3845.0 : f32 + %grid_high126 = scalar.constant 769.0 : f32 + %grid_low127 = scalar.constant 257.0 : f32 + %grid_high127 = scalar.constant 771.0 : f32 + %grid_low128 = scalar.constant 259.0 : f32 + %grid_high128 = scalar.constant 771.0 : f32 + %grid_low129 = scalar.constant 263.0 : f32 + %grid_high129 = scalar.constant 771.0 : f32 + %grid_low130 = scalar.constant 269.0 : f32 + %grid_high130 = scalar.constant 771.0 : f32 + %grid_low131 = scalar.constant 769.0 : f32 + %grid_high131 = scalar.constant 771.0 : f32 + %grid_low132 = scalar.constant 777.0 : f32 + %grid_high132 = scalar.constant 771.0 : f32 + %grid_low133 = scalar.constant 1283.0 : f32 + %grid_high133 = scalar.constant 771.0 : f32 + %grid_low134 = scalar.constant 1793.0 : f32 + %grid_high134 = scalar.constant 771.0 : f32 + %grid_low135 = scalar.constant 1799.0 : f32 + %grid_high135 = scalar.constant 771.0 : f32 + %grid_low136 = scalar.constant 2307.0 : f32 + %grid_high136 = scalar.constant 771.0 : f32 + %grid_low137 = scalar.constant 2817.0 : f32 + %grid_high137 = scalar.constant 771.0 : f32 + %grid_low138 = scalar.constant 2821.0 : f32 + %grid_high138 = scalar.constant 771.0 : f32 + %grid_low139 = scalar.constant 3841.0 : f32 + %grid_high139 = scalar.constant 771.0 : f32 + %grid_low140 = scalar.constant 3853.0 : f32 + %grid_high140 = scalar.constant 771.0 : f32 + %grid_low141 = scalar.constant 257.0 : f32 + %grid_high141 = scalar.constant 773.0 : f32 + %grid_low142 = scalar.constant 773.0 : f32 + %grid_high142 = scalar.constant 773.0 : f32 + %grid_low143 = scalar.constant 779.0 : f32 + %grid_high143 = scalar.constant 773.0 : f32 + %grid_low144 = scalar.constant 783.0 : f32 + %grid_high144 = scalar.constant 773.0 : f32 + %grid_low145 = scalar.constant 1281.0 : f32 + %grid_high145 = scalar.constant 773.0 : f32 + %grid_low146 = scalar.constant 1289.0 : f32 + %grid_high146 = scalar.constant 773.0 : f32 + %grid_low147 = scalar.constant 1797.0 : f32 + %grid_high147 = scalar.constant 773.0 : f32 + %grid_low148 = scalar.constant 2305.0 : f32 + %grid_high148 = scalar.constant 773.0 : f32 + %grid_low149 = scalar.constant 2311.0 : f32 + %grid_high149 = scalar.constant 773.0 : f32 + %grid_low150 = scalar.constant 2827.0 : f32 + %grid_high150 = scalar.constant 773.0 : f32 + %grid_low151 = scalar.constant 3329.0 : f32 + %grid_high151 = scalar.constant 773.0 : f32 + %grid_low152 = scalar.constant 3845.0 : f32 + %grid_high152 = scalar.constant 773.0 : f32 + %grid_low153 = scalar.constant 259.0 : f32 + %grid_high153 = scalar.constant 775.0 : f32 + %grid_low154 = scalar.constant 265.0 : f32 + %grid_high154 = scalar.constant 775.0 : f32 + %grid_low155 = scalar.constant 271.0 : f32 + %grid_high155 = scalar.constant 775.0 : f32 + %grid_low156 = scalar.constant 769.0 : f32 + %grid_high156 = scalar.constant 775.0 : f32 + %grid_low157 = scalar.constant 775.0 : f32 + %grid_high157 = scalar.constant 775.0 : f32 + %grid_low158 = scalar.constant 1283.0 : f32 + %grid_high158 = scalar.constant 775.0 : f32 + %grid_low159 = scalar.constant 1295.0 : f32 + %grid_high159 = scalar.constant 775.0 : f32 + %grid_low160 = scalar.constant 1793.0 : f32 + %grid_high160 = scalar.constant 775.0 : f32 + %grid_low161 = scalar.constant 1801.0 : f32 + %grid_high161 = scalar.constant 775.0 : f32 + %grid_low162 = scalar.constant 2307.0 : f32 + %grid_high162 = scalar.constant 775.0 : f32 + %grid_low163 = scalar.constant 3333.0 : f32 + %grid_high163 = scalar.constant 775.0 : f32 + %grid_low164 = scalar.constant 3841.0 : f32 + %grid_high164 = scalar.constant 775.0 : f32 + %grid_low165 = scalar.constant 263.0 : f32 + %grid_high165 = scalar.constant 777.0 : f32 + %grid_low166 = scalar.constant 267.0 : f32 + %grid_high166 = scalar.constant 777.0 : f32 + %grid_low167 = scalar.constant 773.0 : f32 + %grid_high167 = scalar.constant 777.0 : f32 + %grid_low168 = scalar.constant 777.0 : f32 + %grid_high168 = scalar.constant 777.0 : f32 + %grid_low169 = scalar.constant 1795.0 : f32 + %grid_high169 = scalar.constant 777.0 : f32 + %grid_low170 = scalar.constant 1799.0 : f32 + %grid_high170 = scalar.constant 777.0 : f32 + %grid_low171 = scalar.constant 2309.0 : f32 + %grid_high171 = scalar.constant 777.0 : f32 + %grid_low172 = scalar.constant 2317.0 : f32 + %grid_high172 = scalar.constant 777.0 : f32 + %grid_low173 = scalar.constant 2817.0 : f32 + %grid_high173 = scalar.constant 777.0 : f32 + %grid_low174 = scalar.constant 2825.0 : f32 + %grid_high174 = scalar.constant 777.0 : f32 + %grid_low175 = scalar.constant 259.0 : f32 + %grid_high175 = scalar.constant 779.0 : f32 + %grid_low176 = scalar.constant 769.0 : f32 + %grid_high176 = scalar.constant 779.0 : f32 + %grid_low177 = scalar.constant 775.0 : f32 + %grid_high177 = scalar.constant 779.0 : f32 + %grid_low178 = scalar.constant 1283.0 : f32 + %grid_high178 = scalar.constant 779.0 : f32 + %grid_low179 = scalar.constant 1793.0 : f32 + %grid_high179 = scalar.constant 779.0 : f32 + %grid_low180 = scalar.constant 1797.0 : f32 + %grid_high180 = scalar.constant 779.0 : f32 + %grid_low181 = scalar.constant 2819.0 : f32 + %grid_high181 = scalar.constant 779.0 : f32 + %grid_low182 = scalar.constant 1281.0 : f32 + %grid_high182 = scalar.constant 781.0 : f32 + %grid_low183 = scalar.constant 1289.0 : f32 + %grid_high183 = scalar.constant 781.0 : f32 + %grid_low184 = scalar.constant 1295.0 : f32 + %grid_high184 = scalar.constant 781.0 : f32 + %grid_low185 = scalar.constant 2313.0 : f32 + %grid_high185 = scalar.constant 781.0 : f32 + %grid_low186 = scalar.constant 2317.0 : f32 + %grid_high186 = scalar.constant 781.0 : f32 + %grid_low187 = scalar.constant 259.0 : f32 + %grid_high187 = scalar.constant 783.0 : f32 + %grid_low188 = scalar.constant 263.0 : f32 + %grid_high188 = scalar.constant 783.0 : f32 + %grid_low189 = scalar.constant 769.0 : f32 + %grid_high189 = scalar.constant 783.0 : f32 + %grid_low190 = scalar.constant 773.0 : f32 + %grid_high190 = scalar.constant 783.0 : f32 + %grid_low191 = scalar.constant 1283.0 : f32 + %grid_high191 = scalar.constant 783.0 : f32 + %grid_low192 = scalar.constant 1803.0 : f32 + %grid_high192 = scalar.constant 783.0 : f32 + %grid_low193 = scalar.constant 2307.0 : f32 + %grid_high193 = scalar.constant 783.0 : f32 + %grid_low194 = scalar.constant 3333.0 : f32 + %grid_high194 = scalar.constant 783.0 : f32 + %grid_low195 = scalar.constant 3841.0 : f32 + %grid_high195 = scalar.constant 783.0 : f32 + %grid_low196 = scalar.constant 257.0 : f32 + %grid_high196 = scalar.constant 1281.0 : f32 + %grid_low197 = scalar.constant 259.0 : f32 + %grid_high197 = scalar.constant 1281.0 : f32 + %grid_low198 = scalar.constant 263.0 : f32 + %grid_high198 = scalar.constant 1281.0 : f32 + %grid_low199 = scalar.constant 267.0 : f32 + %grid_high199 = scalar.constant 1281.0 : f32 + %grid_low200 = scalar.constant 271.0 : f32 + %grid_high200 = scalar.constant 1281.0 : f32 + %grid_low201 = scalar.constant 769.0 : f32 + %grid_high201 = scalar.constant 1281.0 : f32 + %grid_low202 = scalar.constant 773.0 : f32 + %grid_high202 = scalar.constant 1281.0 : f32 + %grid_low203 = scalar.constant 777.0 : f32 + %grid_high203 = scalar.constant 1281.0 : f32 + %grid_low204 = scalar.constant 781.0 : f32 + %grid_high204 = scalar.constant 1281.0 : f32 + %grid_low205 = scalar.constant 1283.0 : f32 + %grid_high205 = scalar.constant 1281.0 : f32 + %grid_low206 = scalar.constant 1287.0 : f32 + %grid_high206 = scalar.constant 1281.0 : f32 + %grid_low207 = scalar.constant 1295.0 : f32 + %grid_high207 = scalar.constant 1281.0 : f32 + %grid_low208 = scalar.constant 1793.0 : f32 + %grid_high208 = scalar.constant 1281.0 : f32 + %grid_low209 = scalar.constant 1797.0 : f32 + %grid_high209 = scalar.constant 1281.0 : f32 + %grid_low210 = scalar.constant 2307.0 : f32 + %grid_high210 = scalar.constant 1281.0 : f32 + %grid_low211 = scalar.constant 2311.0 : f32 + %grid_high211 = scalar.constant 1281.0 : f32 + %grid_low212 = scalar.constant 2315.0 : f32 + %grid_high212 = scalar.constant 1281.0 : f32 + %grid_low213 = scalar.constant 2817.0 : f32 + %grid_high213 = scalar.constant 1281.0 : f32 + %grid_low214 = scalar.constant 2821.0 : f32 + %grid_high214 = scalar.constant 1281.0 : f32 + %grid_low215 = scalar.constant 3343.0 : f32 + %grid_high215 = scalar.constant 1281.0 : f32 + %grid_low216 = scalar.constant 3841.0 : f32 + %grid_high216 = scalar.constant 1281.0 : f32 + %grid_low217 = scalar.constant 3847.0 : f32 + %grid_high217 = scalar.constant 1281.0 : f32 + %grid_low218 = scalar.constant 3851.0 : f32 + %grid_high218 = scalar.constant 1281.0 : f32 + %grid_low219 = scalar.constant 257.0 : f32 + %grid_high219 = scalar.constant 1283.0 : f32 + %grid_low220 = scalar.constant 261.0 : f32 + %grid_high220 = scalar.constant 1283.0 : f32 + %grid_low221 = scalar.constant 769.0 : f32 + %grid_high221 = scalar.constant 1283.0 : f32 + %grid_low222 = scalar.constant 775.0 : f32 + %grid_high222 = scalar.constant 1283.0 : f32 + %grid_low223 = scalar.constant 783.0 : f32 + %grid_high223 = scalar.constant 1283.0 : f32 + %grid_low224 = scalar.constant 1285.0 : f32 + %grid_high224 = scalar.constant 1283.0 : f32 + %grid_low225 = scalar.constant 1291.0 : f32 + %grid_high225 = scalar.constant 1283.0 : f32 + %grid_low226 = scalar.constant 1795.0 : f32 + %grid_high226 = scalar.constant 1283.0 : f32 + %grid_low227 = scalar.constant 1801.0 : f32 + %grid_high227 = scalar.constant 1283.0 : f32 + %grid_low228 = scalar.constant 2309.0 : f32 + %grid_high228 = scalar.constant 1283.0 : f32 + %grid_low229 = scalar.constant 2819.0 : f32 + %grid_high229 = scalar.constant 1283.0 : f32 + %grid_low230 = scalar.constant 259.0 : f32 + %grid_high230 = scalar.constant 1285.0 : f32 + %grid_low231 = scalar.constant 265.0 : f32 + %grid_high231 = scalar.constant 1285.0 : f32 + %grid_low232 = scalar.constant 271.0 : f32 + %grid_high232 = scalar.constant 1285.0 : f32 + %grid_low233 = scalar.constant 1283.0 : f32 + %grid_high233 = scalar.constant 1285.0 : f32 + %grid_low234 = scalar.constant 1287.0 : f32 + %grid_high234 = scalar.constant 1285.0 : f32 + %grid_low235 = scalar.constant 1793.0 : f32 + %grid_high235 = scalar.constant 1285.0 : f32 + %grid_low236 = scalar.constant 1807.0 : f32 + %grid_high236 = scalar.constant 1285.0 : f32 + %grid_low237 = scalar.constant 2307.0 : f32 + %grid_high237 = scalar.constant 1285.0 : f32 + %grid_low238 = scalar.constant 2823.0 : f32 + %grid_high238 = scalar.constant 1285.0 : f32 + %grid_low239 = scalar.constant 2831.0 : f32 + %grid_high239 = scalar.constant 1285.0 : f32 + %grid_low240 = scalar.constant 3843.0 : f32 + %grid_high240 = scalar.constant 1285.0 : f32 + %grid_low241 = scalar.constant 3849.0 : f32 + %grid_high241 = scalar.constant 1285.0 : f32 + %grid_low242 = scalar.constant 257.0 : f32 + %grid_high242 = scalar.constant 1287.0 : f32 + %grid_low243 = scalar.constant 261.0 : f32 + %grid_high243 = scalar.constant 1287.0 : f32 + %grid_low244 = scalar.constant 267.0 : f32 + %grid_high244 = scalar.constant 1287.0 : f32 + %grid_low245 = scalar.constant 771.0 : f32 + %grid_high245 = scalar.constant 1287.0 : f32 + %grid_low246 = scalar.constant 1285.0 : f32 + %grid_high246 = scalar.constant 1287.0 : f32 + %grid_low247 = scalar.constant 1289.0 : f32 + %grid_high247 = scalar.constant 1287.0 : f32 + %grid_low248 = scalar.constant 1795.0 : f32 + %grid_high248 = scalar.constant 1287.0 : f32 + %grid_low249 = scalar.constant 1799.0 : f32 + %grid_high249 = scalar.constant 1287.0 : f32 + %grid_low250 = scalar.constant 2309.0 : f32 + %grid_high250 = scalar.constant 1287.0 : f32 + %grid_low251 = scalar.constant 2817.0 : f32 + %grid_high251 = scalar.constant 1287.0 : f32 + %grid_low252 = scalar.constant 3341.0 : f32 + %grid_high252 = scalar.constant 1287.0 : f32 + %grid_low253 = scalar.constant 259.0 : f32 + %grid_high253 = scalar.constant 1289.0 : f32 + %grid_low254 = scalar.constant 271.0 : f32 + %grid_high254 = scalar.constant 1289.0 : f32 + %grid_low255 = scalar.constant 1281.0 : f32 + %grid_high255 = scalar.constant 1289.0 : f32 + %grid_low256 = scalar.constant 1287.0 : f32 + %grid_high256 = scalar.constant 1289.0 : f32 + %grid_low257 = scalar.constant 1797.0 : f32 + %grid_high257 = scalar.constant 1289.0 : f32 + %grid_low258 = scalar.constant 1803.0 : f32 + %grid_high258 = scalar.constant 1289.0 : f32 + %grid_low259 = scalar.constant 2307.0 : f32 + %grid_high259 = scalar.constant 1289.0 : f32 + %grid_low260 = scalar.constant 3845.0 : f32 + %grid_high260 = scalar.constant 1289.0 : f32 + %grid_low261 = scalar.constant 3851.0 : f32 + %grid_high261 = scalar.constant 1289.0 : f32 + %grid_low262 = scalar.constant 265.0 : f32 + %grid_high262 = scalar.constant 1291.0 : f32 + %grid_low263 = scalar.constant 771.0 : f32 + %grid_high263 = scalar.constant 1291.0 : f32 + %grid_low264 = scalar.constant 1285.0 : f32 + %grid_high264 = scalar.constant 1291.0 : f32 + %grid_low265 = scalar.constant 1807.0 : f32 + %grid_high265 = scalar.constant 1291.0 : f32 + %grid_low266 = scalar.constant 2305.0 : f32 + %grid_high266 = scalar.constant 1291.0 : f32 + %grid_low267 = scalar.constant 2823.0 : f32 + %grid_high267 = scalar.constant 1291.0 : f32 + %grid_low268 = scalar.constant 3841.0 : f32 + %grid_high268 = scalar.constant 1291.0 : f32 + %grid_low269 = scalar.constant 257.0 : f32 + %grid_high269 = scalar.constant 1293.0 : f32 + %grid_low270 = scalar.constant 261.0 : f32 + %grid_high270 = scalar.constant 1293.0 : f32 + %grid_low271 = scalar.constant 271.0 : f32 + %grid_high271 = scalar.constant 1293.0 : f32 + %grid_low272 = scalar.constant 1283.0 : f32 + %grid_high272 = scalar.constant 1293.0 : f32 + %grid_low273 = scalar.constant 2827.0 : f32 + %grid_high273 = scalar.constant 1293.0 : f32 + %grid_low274 = scalar.constant 3331.0 : f32 + %grid_high274 = scalar.constant 1293.0 : f32 + %grid_low275 = scalar.constant 267.0 : f32 + %grid_high275 = scalar.constant 1295.0 : f32 + %grid_low276 = scalar.constant 771.0 : f32 + %grid_high276 = scalar.constant 1295.0 : f32 + %grid_low277 = scalar.constant 1293.0 : f32 + %grid_high277 = scalar.constant 1295.0 : f32 + %grid_low278 = scalar.constant 1793.0 : f32 + %grid_high278 = scalar.constant 1295.0 : f32 + %grid_low279 = scalar.constant 2311.0 : f32 + %grid_high279 = scalar.constant 1295.0 : f32 + %grid_low280 = scalar.constant 2817.0 : f32 + %grid_high280 = scalar.constant 1295.0 : f32 + %grid_low281 = scalar.constant 261.0 : f32 + %grid_high281 = scalar.constant 1793.0 : f32 + %grid_low282 = scalar.constant 771.0 : f32 + %grid_high282 = scalar.constant 1793.0 : f32 + %grid_low283 = scalar.constant 775.0 : f32 + %grid_high283 = scalar.constant 1793.0 : f32 + %grid_low284 = scalar.constant 779.0 : f32 + %grid_high284 = scalar.constant 1793.0 : f32 + %grid_low285 = scalar.constant 783.0 : f32 + %grid_high285 = scalar.constant 1793.0 : f32 + %grid_low286 = scalar.constant 1285.0 : f32 + %grid_high286 = scalar.constant 1793.0 : f32 + %grid_low287 = scalar.constant 1795.0 : f32 + %grid_high287 = scalar.constant 1793.0 : f32 + %grid_low288 = scalar.constant 1799.0 : f32 + %grid_high288 = scalar.constant 1793.0 : f32 + %grid_low289 = scalar.constant 1803.0 : f32 + %grid_high289 = scalar.constant 1793.0 : f32 + %grid_low290 = scalar.constant 2309.0 : f32 + %grid_high290 = scalar.constant 1793.0 : f32 + %grid_low291 = scalar.constant 2313.0 : f32 + %grid_high291 = scalar.constant 1793.0 : f32 + %grid_low292 = scalar.constant 2319.0 : f32 + %grid_high292 = scalar.constant 1793.0 : f32 + %grid_low293 = scalar.constant 2819.0 : f32 + %grid_high293 = scalar.constant 1793.0 : f32 + %grid_low294 = scalar.constant 3335.0 : f32 + %grid_high294 = scalar.constant 1793.0 : f32 + %grid_low295 = scalar.constant 3843.0 : f32 + %grid_high295 = scalar.constant 1793.0 : f32 + %grid_low296 = scalar.constant 259.0 : f32 + %grid_high296 = scalar.constant 1795.0 : f32 + %grid_low297 = scalar.constant 263.0 : f32 + %grid_high297 = scalar.constant 1795.0 : f32 + %grid_low298 = scalar.constant 267.0 : f32 + %grid_high298 = scalar.constant 1795.0 : f32 + %grid_low299 = scalar.constant 777.0 : f32 + %grid_high299 = scalar.constant 1795.0 : f32 + %grid_low300 = scalar.constant 1283.0 : f32 + %grid_high300 = scalar.constant 1795.0 : f32 + %grid_low301 = scalar.constant 1287.0 : f32 + %grid_high301 = scalar.constant 1795.0 : f32 + %grid_low302 = scalar.constant 2305.0 : f32 + %grid_high302 = scalar.constant 1795.0 : f32 + %grid_low303 = scalar.constant 3329.0 : f32 + %grid_high303 = scalar.constant 1795.0 : f32 + %grid_low304 = scalar.constant 3845.0 : f32 + %grid_high304 = scalar.constant 1795.0 : f32 + %grid_low305 = scalar.constant 3853.0 : f32 + %grid_high305 = scalar.constant 1795.0 : f32 + %grid_low306 = scalar.constant 257.0 : f32 + %grid_high306 = scalar.constant 1797.0 : f32 + %grid_low307 = scalar.constant 773.0 : f32 + %grid_high307 = scalar.constant 1797.0 : f32 + %grid_low308 = scalar.constant 1281.0 : f32 + %grid_high308 = scalar.constant 1797.0 : f32 + %grid_low309 = scalar.constant 1797.0 : f32 + %grid_high309 = scalar.constant 1797.0 : f32 + %grid_low310 = scalar.constant 1801.0 : f32 + %grid_high310 = scalar.constant 1797.0 : f32 + %grid_low311 = scalar.constant 2817.0 : f32 + %grid_high311 = scalar.constant 1797.0 : f32 + %grid_low312 = scalar.constant 259.0 : f32 + %grid_high312 = scalar.constant 1799.0 : f32 + %grid_low313 = scalar.constant 769.0 : f32 + %grid_high313 = scalar.constant 1799.0 : f32 + %grid_low314 = scalar.constant 777.0 : f32 + %grid_high314 = scalar.constant 1799.0 : f32 + %grid_low315 = scalar.constant 1283.0 : f32 + %grid_high315 = scalar.constant 1799.0 : f32 + %grid_low316 = scalar.constant 1287.0 : f32 + %grid_high316 = scalar.constant 1799.0 : f32 + %grid_low317 = scalar.constant 1295.0 : f32 + %grid_high317 = scalar.constant 1799.0 : f32 + %grid_low318 = scalar.constant 1793.0 : f32 + %grid_high318 = scalar.constant 1799.0 : f32 + %grid_low319 = scalar.constant 2307.0 : f32 + %grid_high319 = scalar.constant 1799.0 : f32 + %grid_low320 = scalar.constant 2311.0 : f32 + %grid_high320 = scalar.constant 1799.0 : f32 + %grid_low321 = scalar.constant 2319.0 : f32 + %grid_high321 = scalar.constant 1799.0 : f32 + %grid_low322 = scalar.constant 2827.0 : f32 + %grid_high322 = scalar.constant 1799.0 : f32 + %grid_low323 = scalar.constant 3847.0 : f32 + %grid_high323 = scalar.constant 1799.0 : f32 + %grid_low324 = scalar.constant 263.0 : f32 + %grid_high324 = scalar.constant 1801.0 : f32 + %grid_low325 = scalar.constant 771.0 : f32 + %grid_high325 = scalar.constant 1801.0 : f32 + %grid_low326 = scalar.constant 781.0 : f32 + %grid_high326 = scalar.constant 1801.0 : f32 + %grid_low327 = scalar.constant 1285.0 : f32 + %grid_high327 = scalar.constant 1801.0 : f32 + %grid_low328 = scalar.constant 1795.0 : f32 + %grid_high328 = scalar.constant 1801.0 : f32 + %grid_low329 = scalar.constant 2821.0 : f32 + %grid_high329 = scalar.constant 1801.0 : f32 + %grid_low330 = scalar.constant 3329.0 : f32 + %grid_high330 = scalar.constant 1801.0 : f32 + %grid_low331 = scalar.constant 3337.0 : f32 + %grid_high331 = scalar.constant 1801.0 : f32 + %grid_low332 = scalar.constant 259.0 : f32 + %grid_high332 = scalar.constant 1803.0 : f32 + %grid_low333 = scalar.constant 769.0 : f32 + %grid_high333 = scalar.constant 1803.0 : f32 + %grid_low334 = scalar.constant 773.0 : f32 + %grid_high334 = scalar.constant 1803.0 : f32 + %grid_low335 = scalar.constant 1291.0 : f32 + %grid_high335 = scalar.constant 1803.0 : f32 + %grid_low336 = scalar.constant 1797.0 : f32 + %grid_high336 = scalar.constant 1803.0 : f32 + %grid_low337 = scalar.constant 2313.0 : f32 + %grid_high337 = scalar.constant 1803.0 : f32 + %grid_low338 = scalar.constant 2829.0 : f32 + %grid_high338 = scalar.constant 1803.0 : f32 + %grid_low339 = scalar.constant 3847.0 : f32 + %grid_high339 = scalar.constant 1803.0 : f32 + %grid_low340 = scalar.constant 781.0 : f32 + %grid_high340 = scalar.constant 1805.0 : f32 + %grid_low341 = scalar.constant 2307.0 : f32 + %grid_high341 = scalar.constant 1805.0 : f32 + %grid_low342 = scalar.constant 259.0 : f32 + %grid_high342 = scalar.constant 1807.0 : f32 + %grid_low343 = scalar.constant 263.0 : f32 + %grid_high343 = scalar.constant 1807.0 : f32 + %grid_low344 = scalar.constant 1281.0 : f32 + %grid_high344 = scalar.constant 1807.0 : f32 + %grid_low345 = scalar.constant 1285.0 : f32 + %grid_high345 = scalar.constant 1807.0 : f32 + %grid_low346 = scalar.constant 1803.0 : f32 + %grid_high346 = scalar.constant 1807.0 : f32 + %grid_low347 = scalar.constant 257.0 : f32 + %grid_high347 = scalar.constant 2305.0 : f32 + %grid_low348 = scalar.constant 265.0 : f32 + %grid_high348 = scalar.constant 2305.0 : f32 + %grid_low349 = scalar.constant 773.0 : f32 + %grid_high349 = scalar.constant 2305.0 : f32 + %grid_low350 = scalar.constant 1281.0 : f32 + %grid_high350 = scalar.constant 2305.0 : f32 + %grid_low351 = scalar.constant 1289.0 : f32 + %grid_high351 = scalar.constant 2305.0 : f32 + %grid_low352 = scalar.constant 1295.0 : f32 + %grid_high352 = scalar.constant 2305.0 : f32 + %grid_low353 = scalar.constant 1797.0 : f32 + %grid_high353 = scalar.constant 2305.0 : f32 + %grid_low354 = scalar.constant 2307.0 : f32 + %grid_high354 = scalar.constant 2305.0 : f32 + %grid_low355 = scalar.constant 2817.0 : f32 + %grid_high355 = scalar.constant 2305.0 : f32 + %grid_low356 = scalar.constant 3841.0 : f32 + %grid_high356 = scalar.constant 2305.0 : f32 + %grid_low357 = scalar.constant 261.0 : f32 + %grid_high357 = scalar.constant 2307.0 : f32 + %grid_low358 = scalar.constant 271.0 : f32 + %grid_high358 = scalar.constant 2307.0 : f32 + %grid_low359 = scalar.constant 771.0 : f32 + %grid_high359 = scalar.constant 2307.0 : f32 + %grid_low360 = scalar.constant 775.0 : f32 + %grid_high360 = scalar.constant 2307.0 : f32 + %grid_low361 = scalar.constant 1285.0 : f32 + %grid_high361 = scalar.constant 2307.0 : f32 + %grid_low362 = scalar.constant 1793.0 : f32 + %grid_high362 = scalar.constant 2307.0 : f32 + %grid_low363 = scalar.constant 1803.0 : f32 + %grid_high363 = scalar.constant 2307.0 : f32 + %grid_low364 = scalar.constant 2311.0 : f32 + %grid_high364 = scalar.constant 2307.0 : f32 + %grid_low365 = scalar.constant 2819.0 : f32 + %grid_high365 = scalar.constant 2307.0 : f32 + %grid_low366 = scalar.constant 2827.0 : f32 + %grid_high366 = scalar.constant 2307.0 : f32 + %grid_low367 = scalar.constant 259.0 : f32 + %grid_high367 = scalar.constant 2309.0 : f32 + %grid_low368 = scalar.constant 263.0 : f32 + %grid_high368 = scalar.constant 2309.0 : f32 + %grid_low369 = scalar.constant 769.0 : f32 + %grid_high369 = scalar.constant 2309.0 : f32 + %grid_low370 = scalar.constant 779.0 : f32 + %grid_high370 = scalar.constant 2309.0 : f32 + %grid_low371 = scalar.constant 1283.0 : f32 + %grid_high371 = scalar.constant 2309.0 : f32 + %grid_low372 = scalar.constant 1799.0 : f32 + %grid_high372 = scalar.constant 2309.0 : f32 + %grid_low373 = scalar.constant 2305.0 : f32 + %grid_high373 = scalar.constant 2309.0 : f32 + %grid_low374 = scalar.constant 2831.0 : f32 + %grid_high374 = scalar.constant 2309.0 : f32 + %grid_low375 = scalar.constant 3333.0 : f32 + %grid_high375 = scalar.constant 2309.0 : f32 + %grid_low376 = scalar.constant 3841.0 : f32 + %grid_high376 = scalar.constant 2309.0 : f32 + %grid_low377 = scalar.constant 265.0 : f32 + %grid_high377 = scalar.constant 2311.0 : f32 + %grid_low378 = scalar.constant 771.0 : f32 + %grid_high378 = scalar.constant 2311.0 : f32 + %grid_low379 = scalar.constant 775.0 : f32 + %grid_high379 = scalar.constant 2311.0 : f32 + %grid_low380 = scalar.constant 1281.0 : f32 + %grid_high380 = scalar.constant 2311.0 : f32 + %grid_low381 = scalar.constant 1285.0 : f32 + %grid_high381 = scalar.constant 2311.0 : f32 + %grid_low382 = scalar.constant 1795.0 : f32 + %grid_high382 = scalar.constant 2311.0 : f32 + %grid_low383 = scalar.constant 1803.0 : f32 + %grid_high383 = scalar.constant 2311.0 : f32 + %grid_low384 = scalar.constant 257.0 : f32 + %grid_high384 = scalar.constant 2313.0 : f32 + %grid_low385 = scalar.constant 261.0 : f32 + %grid_high385 = scalar.constant 2313.0 : f32 + %grid_low386 = scalar.constant 1289.0 : f32 + %grid_high386 = scalar.constant 2313.0 : f32 + %grid_low387 = scalar.constant 1807.0 : f32 + %grid_high387 = scalar.constant 2313.0 : f32 + %grid_low388 = scalar.constant 2305.0 : f32 + %grid_high388 = scalar.constant 2313.0 : f32 + %grid_low389 = scalar.constant 3843.0 : f32 + %grid_high389 = scalar.constant 2313.0 : f32 + %grid_low390 = scalar.constant 267.0 : f32 + %grid_high390 = scalar.constant 2315.0 : f32 + %grid_low391 = scalar.constant 271.0 : f32 + %grid_high391 = scalar.constant 2315.0 : f32 + %grid_low392 = scalar.constant 1283.0 : f32 + %grid_high392 = scalar.constant 2315.0 : f32 + %grid_low393 = scalar.constant 3333.0 : f32 + %grid_high393 = scalar.constant 2315.0 : f32 + %grid_low394 = scalar.constant 775.0 : f32 + %grid_high394 = scalar.constant 2317.0 : f32 + %grid_low395 = scalar.constant 1801.0 : f32 + %grid_high395 = scalar.constant 2317.0 : f32 + %grid_low396 = scalar.constant 3329.0 : f32 + %grid_high396 = scalar.constant 2317.0 : f32 + %grid_low397 = scalar.constant 769.0 : f32 + %grid_high397 = scalar.constant 2319.0 : f32 + %grid_low398 = scalar.constant 779.0 : f32 + %grid_high398 = scalar.constant 2319.0 : f32 + %grid_low399 = scalar.constant 1793.0 : f32 + %grid_high399 = scalar.constant 2319.0 : f32 + %grid_low400 = scalar.constant 2311.0 : f32 + %grid_high400 = scalar.constant 2319.0 : f32 + %grid_low401 = scalar.constant 2819.0 : f32 + %grid_high401 = scalar.constant 2319.0 : f32 + %grid_low402 = scalar.constant 261.0 : f32 + %grid_high402 = scalar.constant 2817.0 : f32 + %grid_low403 = scalar.constant 769.0 : f32 + %grid_high403 = scalar.constant 2817.0 : f32 + %grid_low404 = scalar.constant 777.0 : f32 + %grid_high404 = scalar.constant 2817.0 : f32 + %grid_low405 = scalar.constant 1285.0 : f32 + %grid_high405 = scalar.constant 2817.0 : f32 + %grid_low406 = scalar.constant 2305.0 : f32 + %grid_high406 = scalar.constant 2817.0 : f32 + %grid_low407 = scalar.constant 2313.0 : f32 + %grid_high407 = scalar.constant 2817.0 : f32 + %grid_low408 = scalar.constant 2319.0 : f32 + %grid_high408 = scalar.constant 2817.0 : f32 + %grid_low409 = scalar.constant 2821.0 : f32 + %grid_high409 = scalar.constant 2817.0 : f32 + %grid_low410 = scalar.constant 3341.0 : f32 + %grid_high410 = scalar.constant 2817.0 : f32 + %grid_low411 = scalar.constant 3849.0 : f32 + %grid_high411 = scalar.constant 2817.0 : f32 + %grid_low412 = scalar.constant 259.0 : f32 + %grid_high412 = scalar.constant 2819.0 : f32 + %grid_low413 = scalar.constant 263.0 : f32 + %grid_high413 = scalar.constant 2819.0 : f32 + %grid_low414 = scalar.constant 267.0 : f32 + %grid_high414 = scalar.constant 2819.0 : f32 + %grid_low415 = scalar.constant 773.0 : f32 + %grid_high415 = scalar.constant 2819.0 : f32 + %grid_low416 = scalar.constant 1283.0 : f32 + %grid_high416 = scalar.constant 2819.0 : f32 + %grid_low417 = scalar.constant 1797.0 : f32 + %grid_high417 = scalar.constant 2819.0 : f32 + %grid_low418 = scalar.constant 3845.0 : f32 + %grid_high418 = scalar.constant 2819.0 : f32 + %grid_low419 = scalar.constant 257.0 : f32 + %grid_high419 = scalar.constant 2821.0 : f32 + %grid_low420 = scalar.constant 771.0 : f32 + %grid_high420 = scalar.constant 2821.0 : f32 + %grid_low421 = scalar.constant 1287.0 : f32 + %grid_high421 = scalar.constant 2821.0 : f32 + %grid_low422 = scalar.constant 1793.0 : f32 + %grid_high422 = scalar.constant 2821.0 : f32 + %grid_low423 = scalar.constant 1805.0 : f32 + %grid_high423 = scalar.constant 2821.0 : f32 + %grid_low424 = scalar.constant 2823.0 : f32 + %grid_high424 = scalar.constant 2821.0 : f32 + %grid_low425 = scalar.constant 261.0 : f32 + %grid_high425 = scalar.constant 2823.0 : f32 + %grid_low426 = scalar.constant 271.0 : f32 + %grid_high426 = scalar.constant 2823.0 : f32 + %grid_low427 = scalar.constant 769.0 : f32 + %grid_high427 = scalar.constant 2823.0 : f32 + %grid_low428 = scalar.constant 1295.0 : f32 + %grid_high428 = scalar.constant 2823.0 : f32 + %grid_low429 = scalar.constant 2313.0 : f32 + %grid_high429 = scalar.constant 2823.0 : f32 + %grid_low430 = scalar.constant 2819.0 : f32 + %grid_high430 = scalar.constant 2823.0 : f32 + %grid_low431 = scalar.constant 3339.0 : f32 + %grid_high431 = scalar.constant 2823.0 : f32 + %grid_low432 = scalar.constant 3847.0 : f32 + %grid_high432 = scalar.constant 2823.0 : f32 + %grid_low433 = scalar.constant 259.0 : f32 + %grid_high433 = scalar.constant 2825.0 : f32 + %grid_low434 = scalar.constant 265.0 : f32 + %grid_high434 = scalar.constant 2825.0 : f32 + %grid_low435 = scalar.constant 1281.0 : f32 + %grid_high435 = scalar.constant 2825.0 : f32 + %grid_low436 = scalar.constant 1797.0 : f32 + %grid_high436 = scalar.constant 2825.0 : f32 + %grid_low437 = scalar.constant 2317.0 : f32 + %grid_high437 = scalar.constant 2825.0 : f32 + %grid_low438 = scalar.constant 773.0 : f32 + %grid_high438 = scalar.constant 2827.0 : f32 + %grid_low439 = scalar.constant 1293.0 : f32 + %grid_high439 = scalar.constant 2827.0 : f32 + %grid_low440 = scalar.constant 2819.0 : f32 + %grid_high440 = scalar.constant 2827.0 : f32 + %grid_low441 = scalar.constant 2823.0 : f32 + %grid_high441 = scalar.constant 2827.0 : f32 + %grid_low442 = scalar.constant 2309.0 : f32 + %grid_high442 = scalar.constant 2829.0 : f32 + %grid_low443 = scalar.constant 261.0 : f32 + %grid_high443 = scalar.constant 2831.0 : f32 + %grid_low444 = scalar.constant 265.0 : f32 + %grid_high444 = scalar.constant 2831.0 : f32 + %grid_low445 = scalar.constant 1285.0 : f32 + %grid_high445 = scalar.constant 2831.0 : f32 + %grid_low446 = scalar.constant 771.0 : f32 + %grid_high446 = scalar.constant 3329.0 : f32 + %grid_low447 = scalar.constant 775.0 : f32 + %grid_high447 = scalar.constant 3329.0 : f32 + %grid_low448 = scalar.constant 779.0 : f32 + %grid_high448 = scalar.constant 3329.0 : f32 + %grid_low449 = scalar.constant 1795.0 : f32 + %grid_high449 = scalar.constant 3329.0 : f32 + %grid_low450 = scalar.constant 1799.0 : f32 + %grid_high450 = scalar.constant 3329.0 : f32 + %grid_low451 = scalar.constant 3329.0 : f32 + %grid_high451 = scalar.constant 3329.0 : f32 + %grid_low452 = scalar.constant 257.0 : f32 + %grid_high452 = scalar.constant 3331.0 : f32 + %grid_low453 = scalar.constant 1281.0 : f32 + %grid_high453 = scalar.constant 3331.0 : f32 + %grid_low454 = scalar.constant 1295.0 : f32 + %grid_high454 = scalar.constant 3331.0 : f32 + %grid_low455 = scalar.constant 3337.0 : f32 + %grid_high455 = scalar.constant 3331.0 : f32 + %grid_low456 = scalar.constant 773.0 : f32 + %grid_high456 = scalar.constant 3333.0 : f32 + %grid_low457 = scalar.constant 1801.0 : f32 + %grid_high457 = scalar.constant 3333.0 : f32 + %grid_low458 = scalar.constant 2309.0 : f32 + %grid_high458 = scalar.constant 3333.0 : f32 + %grid_low459 = scalar.constant 2827.0 : f32 + %grid_high459 = scalar.constant 3333.0 : f32 + %grid_low460 = scalar.constant 3333.0 : f32 + %grid_high460 = scalar.constant 3333.0 : f32 + %grid_low461 = scalar.constant 3841.0 : f32 + %grid_high461 = scalar.constant 3333.0 : f32 + %grid_low462 = scalar.constant 257.0 : f32 + %grid_high462 = scalar.constant 3335.0 : f32 + %grid_low463 = scalar.constant 777.0 : f32 + %grid_high463 = scalar.constant 3335.0 : f32 + %grid_low464 = scalar.constant 1283.0 : f32 + %grid_high464 = scalar.constant 3335.0 : f32 + %grid_low465 = scalar.constant 2305.0 : f32 + %grid_high465 = scalar.constant 3335.0 : f32 + %grid_low466 = scalar.constant 1291.0 : f32 + %grid_high466 = scalar.constant 3337.0 : f32 + %grid_low467 = scalar.constant 2311.0 : f32 + %grid_high467 = scalar.constant 3337.0 : f32 + %grid_low468 = scalar.constant 3333.0 : f32 + %grid_high468 = scalar.constant 3337.0 : f32 + %grid_low469 = scalar.constant 257.0 : f32 + %grid_high469 = scalar.constant 3339.0 : f32 + %grid_low470 = scalar.constant 263.0 : f32 + %grid_high470 = scalar.constant 3339.0 : f32 + %grid_low471 = scalar.constant 1801.0 : f32 + %grid_high471 = scalar.constant 3339.0 : f32 + %grid_low472 = scalar.constant 3329.0 : f32 + %grid_high472 = scalar.constant 3339.0 : f32 + %grid_low473 = scalar.constant 267.0 : f32 + %grid_high473 = scalar.constant 3341.0 : f32 + %grid_low474 = scalar.constant 2305.0 : f32 + %grid_high474 = scalar.constant 3341.0 : f32 + %grid_low475 = scalar.constant 771.0 : f32 + %grid_high475 = scalar.constant 3343.0 : f32 + %grid_low476 = scalar.constant 775.0 : f32 + %grid_high476 = scalar.constant 3343.0 : f32 + %grid_low477 = scalar.constant 257.0 : f32 + %grid_high477 = scalar.constant 3841.0 : f32 + %grid_low478 = scalar.constant 265.0 : f32 + %grid_high478 = scalar.constant 3841.0 : f32 + %grid_low479 = scalar.constant 271.0 : f32 + %grid_high479 = scalar.constant 3841.0 : f32 + %grid_low480 = scalar.constant 1281.0 : f32 + %grid_high480 = scalar.constant 3841.0 : f32 + %grid_low481 = scalar.constant 1285.0 : f32 + %grid_high481 = scalar.constant 3841.0 : f32 + %grid_low482 = scalar.constant 1805.0 : f32 + %grid_high482 = scalar.constant 3841.0 : f32 + %grid_low483 = scalar.constant 2305.0 : f32 + %grid_high483 = scalar.constant 3841.0 : f32 + %grid_low484 = scalar.constant 2825.0 : f32 + %grid_high484 = scalar.constant 3841.0 : f32 + %grid_low485 = scalar.constant 3333.0 : f32 + %grid_high485 = scalar.constant 3841.0 : f32 + %grid_low486 = scalar.constant 261.0 : f32 + %grid_high486 = scalar.constant 3843.0 : f32 + %grid_low487 = scalar.constant 771.0 : f32 + %grid_high487 = scalar.constant 3843.0 : f32 + %grid_low488 = scalar.constant 1289.0 : f32 + %grid_high488 = scalar.constant 3843.0 : f32 + %grid_low489 = scalar.constant 2311.0 : f32 + %grid_high489 = scalar.constant 3843.0 : f32 + %grid_low490 = scalar.constant 2315.0 : f32 + %grid_high490 = scalar.constant 3843.0 : f32 + %grid_low491 = scalar.constant 259.0 : f32 + %grid_high491 = scalar.constant 3845.0 : f32 + %grid_low492 = scalar.constant 265.0 : f32 + %grid_high492 = scalar.constant 3845.0 : f32 + %grid_low493 = scalar.constant 769.0 : f32 + %grid_high493 = scalar.constant 3845.0 : f32 + %grid_low494 = scalar.constant 781.0 : f32 + %grid_high494 = scalar.constant 3845.0 : f32 + %grid_low495 = scalar.constant 1283.0 : f32 + %grid_high495 = scalar.constant 3845.0 : f32 + %grid_low496 = scalar.constant 1793.0 : f32 + %grid_high496 = scalar.constant 3845.0 : f32 + %grid_low497 = scalar.constant 2819.0 : f32 + %grid_high497 = scalar.constant 3845.0 : f32 + %grid_low498 = scalar.constant 261.0 : f32 + %grid_high498 = scalar.constant 3847.0 : f32 + %grid_low499 = scalar.constant 1797.0 : f32 + %grid_high499 = scalar.constant 3847.0 : f32 + %grid_low500 = scalar.constant 1803.0 : f32 + %grid_high500 = scalar.constant 3847.0 : f32 + %grid_low501 = scalar.constant 2823.0 : f32 + %grid_high501 = scalar.constant 3847.0 : f32 + %grid_low502 = scalar.constant 259.0 : f32 + %grid_high502 = scalar.constant 3849.0 : f32 + %grid_low503 = scalar.constant 267.0 : f32 + %grid_high503 = scalar.constant 3849.0 : f32 + %grid_low504 = scalar.constant 775.0 : f32 + %grid_high504 = scalar.constant 3849.0 : f32 + %grid_low505 = scalar.constant 1281.0 : f32 + %grid_high505 = scalar.constant 3849.0 : f32 + %grid_low506 = scalar.constant 2817.0 : f32 + %grid_high506 = scalar.constant 3849.0 : f32 + %grid_low507 = scalar.constant 1285.0 : f32 + %grid_high507 = scalar.constant 3851.0 : f32 + %grid_low508 = scalar.constant 2309.0 : f32 + %grid_high508 = scalar.constant 3851.0 : f32 + %grid_low509 = scalar.constant 261.0 : f32 + %grid_high509 = scalar.constant 3853.0 : f32 + %grid_low510 = scalar.constant 1795.0 : f32 + %grid_high510 = scalar.constant 3853.0 : f32 + %grid_low511 = scalar.constant 257.0 : f32 + %grid_high511 = scalar.constant 3855.0 : f32 + %grid_low0_chunk = vector.from_elements %grid_low0, %grid_low1, %grid_low2, %grid_low3, %grid_low4, %grid_low5, %grid_low6, %grid_low7, %grid_low8, %grid_low9, %grid_low10, %grid_low11, %grid_low12, %grid_low13, %grid_low14, %grid_low15, %grid_low16, %grid_low17, %grid_low18, %grid_low19, %grid_low20, %grid_low21, %grid_low22, %grid_low23, %grid_low24, %grid_low25, %grid_low26, %grid_low27, %grid_low28, %grid_low29, %grid_low30, %grid_low31 : vector<32xf32> + %grid_low1_chunk = vector.from_elements %grid_low32, %grid_low33, %grid_low34, %grid_low35, %grid_low36, %grid_low37, %grid_low38, %grid_low39, %grid_low40, %grid_low41, %grid_low42, %grid_low43, %grid_low44, %grid_low45, %grid_low46, %grid_low47, %grid_low48, %grid_low49, %grid_low50, %grid_low51, %grid_low52, %grid_low53, %grid_low54, %grid_low55, %grid_low56, %grid_low57, %grid_low58, %grid_low59, %grid_low60, %grid_low61, %grid_low62, %grid_low63 : vector<32xf32> + %grid_low2_chunk = vector.from_elements %grid_low64, %grid_low65, %grid_low66, %grid_low67, %grid_low68, %grid_low69, %grid_low70, %grid_low71, %grid_low72, %grid_low73, %grid_low74, %grid_low75, %grid_low76, %grid_low77, %grid_low78, %grid_low79, %grid_low80, %grid_low81, %grid_low82, %grid_low83, %grid_low84, %grid_low85, %grid_low86, %grid_low87, %grid_low88, %grid_low89, %grid_low90, %grid_low91, %grid_low92, %grid_low93, %grid_low94, %grid_low95 : vector<32xf32> + %grid_low3_chunk = vector.from_elements %grid_low96, %grid_low97, %grid_low98, %grid_low99, %grid_low100, %grid_low101, %grid_low102, %grid_low103, %grid_low104, %grid_low105, %grid_low106, %grid_low107, %grid_low108, %grid_low109, %grid_low110, %grid_low111, %grid_low112, %grid_low113, %grid_low114, %grid_low115, %grid_low116, %grid_low117, %grid_low118, %grid_low119, %grid_low120, %grid_low121, %grid_low122, %grid_low123, %grid_low124, %grid_low125, %grid_low126, %grid_low127 : vector<32xf32> + %grid_low4_chunk = vector.from_elements %grid_low128, %grid_low129, %grid_low130, %grid_low131, %grid_low132, %grid_low133, %grid_low134, %grid_low135, %grid_low136, %grid_low137, %grid_low138, %grid_low139, %grid_low140, %grid_low141, %grid_low142, %grid_low143, %grid_low144, %grid_low145, %grid_low146, %grid_low147, %grid_low148, %grid_low149, %grid_low150, %grid_low151, %grid_low152, %grid_low153, %grid_low154, %grid_low155, %grid_low156, %grid_low157, %grid_low158, %grid_low159 : vector<32xf32> + %grid_low5_chunk = vector.from_elements %grid_low160, %grid_low161, %grid_low162, %grid_low163, %grid_low164, %grid_low165, %grid_low166, %grid_low167, %grid_low168, %grid_low169, %grid_low170, %grid_low171, %grid_low172, %grid_low173, %grid_low174, %grid_low175, %grid_low176, %grid_low177, %grid_low178, %grid_low179, %grid_low180, %grid_low181, %grid_low182, %grid_low183, %grid_low184, %grid_low185, %grid_low186, %grid_low187, %grid_low188, %grid_low189, %grid_low190, %grid_low191 : vector<32xf32> + %grid_low6_chunk = vector.from_elements %grid_low192, %grid_low193, %grid_low194, %grid_low195, %grid_low196, %grid_low197, %grid_low198, %grid_low199, %grid_low200, %grid_low201, %grid_low202, %grid_low203, %grid_low204, %grid_low205, %grid_low206, %grid_low207, %grid_low208, %grid_low209, %grid_low210, %grid_low211, %grid_low212, %grid_low213, %grid_low214, %grid_low215, %grid_low216, %grid_low217, %grid_low218, %grid_low219, %grid_low220, %grid_low221, %grid_low222, %grid_low223 : vector<32xf32> + %grid_low7_chunk = vector.from_elements %grid_low224, %grid_low225, %grid_low226, %grid_low227, %grid_low228, %grid_low229, %grid_low230, %grid_low231, %grid_low232, %grid_low233, %grid_low234, %grid_low235, %grid_low236, %grid_low237, %grid_low238, %grid_low239, %grid_low240, %grid_low241, %grid_low242, %grid_low243, %grid_low244, %grid_low245, %grid_low246, %grid_low247, %grid_low248, %grid_low249, %grid_low250, %grid_low251, %grid_low252, %grid_low253, %grid_low254, %grid_low255 : vector<32xf32> + %grid_low8_chunk = vector.from_elements %grid_low256, %grid_low257, %grid_low258, %grid_low259, %grid_low260, %grid_low261, %grid_low262, %grid_low263, %grid_low264, %grid_low265, %grid_low266, %grid_low267, %grid_low268, %grid_low269, %grid_low270, %grid_low271, %grid_low272, %grid_low273, %grid_low274, %grid_low275, %grid_low276, %grid_low277, %grid_low278, %grid_low279, %grid_low280, %grid_low281, %grid_low282, %grid_low283, %grid_low284, %grid_low285, %grid_low286, %grid_low287 : vector<32xf32> + %grid_low9_chunk = vector.from_elements %grid_low288, %grid_low289, %grid_low290, %grid_low291, %grid_low292, %grid_low293, %grid_low294, %grid_low295, %grid_low296, %grid_low297, %grid_low298, %grid_low299, %grid_low300, %grid_low301, %grid_low302, %grid_low303, %grid_low304, %grid_low305, %grid_low306, %grid_low307, %grid_low308, %grid_low309, %grid_low310, %grid_low311, %grid_low312, %grid_low313, %grid_low314, %grid_low315, %grid_low316, %grid_low317, %grid_low318, %grid_low319 : vector<32xf32> + %grid_low10_chunk = vector.from_elements %grid_low320, %grid_low321, %grid_low322, %grid_low323, %grid_low324, %grid_low325, %grid_low326, %grid_low327, %grid_low328, %grid_low329, %grid_low330, %grid_low331, %grid_low332, %grid_low333, %grid_low334, %grid_low335, %grid_low336, %grid_low337, %grid_low338, %grid_low339, %grid_low340, %grid_low341, %grid_low342, %grid_low343, %grid_low344, %grid_low345, %grid_low346, %grid_low347, %grid_low348, %grid_low349, %grid_low350, %grid_low351 : vector<32xf32> + %grid_low11_chunk = vector.from_elements %grid_low352, %grid_low353, %grid_low354, %grid_low355, %grid_low356, %grid_low357, %grid_low358, %grid_low359, %grid_low360, %grid_low361, %grid_low362, %grid_low363, %grid_low364, %grid_low365, %grid_low366, %grid_low367, %grid_low368, %grid_low369, %grid_low370, %grid_low371, %grid_low372, %grid_low373, %grid_low374, %grid_low375, %grid_low376, %grid_low377, %grid_low378, %grid_low379, %grid_low380, %grid_low381, %grid_low382, %grid_low383 : vector<32xf32> + %grid_low12_chunk = vector.from_elements %grid_low384, %grid_low385, %grid_low386, %grid_low387, %grid_low388, %grid_low389, %grid_low390, %grid_low391, %grid_low392, %grid_low393, %grid_low394, %grid_low395, %grid_low396, %grid_low397, %grid_low398, %grid_low399, %grid_low400, %grid_low401, %grid_low402, %grid_low403, %grid_low404, %grid_low405, %grid_low406, %grid_low407, %grid_low408, %grid_low409, %grid_low410, %grid_low411, %grid_low412, %grid_low413, %grid_low414, %grid_low415 : vector<32xf32> + %grid_low13_chunk = vector.from_elements %grid_low416, %grid_low417, %grid_low418, %grid_low419, %grid_low420, %grid_low421, %grid_low422, %grid_low423, %grid_low424, %grid_low425, %grid_low426, %grid_low427, %grid_low428, %grid_low429, %grid_low430, %grid_low431, %grid_low432, %grid_low433, %grid_low434, %grid_low435, %grid_low436, %grid_low437, %grid_low438, %grid_low439, %grid_low440, %grid_low441, %grid_low442, %grid_low443, %grid_low444, %grid_low445, %grid_low446, %grid_low447 : vector<32xf32> + %grid_low14_chunk = vector.from_elements %grid_low448, %grid_low449, %grid_low450, %grid_low451, %grid_low452, %grid_low453, %grid_low454, %grid_low455, %grid_low456, %grid_low457, %grid_low458, %grid_low459, %grid_low460, %grid_low461, %grid_low462, %grid_low463, %grid_low464, %grid_low465, %grid_low466, %grid_low467, %grid_low468, %grid_low469, %grid_low470, %grid_low471, %grid_low472, %grid_low473, %grid_low474, %grid_low475, %grid_low476, %grid_low477, %grid_low478, %grid_low479 : vector<32xf32> + %grid_low15_chunk = vector.from_elements %grid_low480, %grid_low481, %grid_low482, %grid_low483, %grid_low484, %grid_low485, %grid_low486, %grid_low487, %grid_low488, %grid_low489, %grid_low490, %grid_low491, %grid_low492, %grid_low493, %grid_low494, %grid_low495, %grid_low496, %grid_low497, %grid_low498, %grid_low499, %grid_low500, %grid_low501, %grid_low502, %grid_low503, %grid_low504, %grid_low505, %grid_low506, %grid_low507, %grid_low508, %grid_low509, %grid_low510, %grid_low511 : vector<32xf32> + %grid_high0_chunk = vector.from_elements %grid_high0, %grid_high1, %grid_high2, %grid_high3, %grid_high4, %grid_high5, %grid_high6, %grid_high7, %grid_high8, %grid_high9, %grid_high10, %grid_high11, %grid_high12, %grid_high13, %grid_high14, %grid_high15, %grid_high16, %grid_high17, %grid_high18, %grid_high19, %grid_high20, %grid_high21, %grid_high22, %grid_high23, %grid_high24, %grid_high25, %grid_high26, %grid_high27, %grid_high28, %grid_high29, %grid_high30, %grid_high31 : vector<32xf32> + %grid_high1_chunk = vector.from_elements %grid_high32, %grid_high33, %grid_high34, %grid_high35, %grid_high36, %grid_high37, %grid_high38, %grid_high39, %grid_high40, %grid_high41, %grid_high42, %grid_high43, %grid_high44, %grid_high45, %grid_high46, %grid_high47, %grid_high48, %grid_high49, %grid_high50, %grid_high51, %grid_high52, %grid_high53, %grid_high54, %grid_high55, %grid_high56, %grid_high57, %grid_high58, %grid_high59, %grid_high60, %grid_high61, %grid_high62, %grid_high63 : vector<32xf32> + %grid_high2_chunk = vector.from_elements %grid_high64, %grid_high65, %grid_high66, %grid_high67, %grid_high68, %grid_high69, %grid_high70, %grid_high71, %grid_high72, %grid_high73, %grid_high74, %grid_high75, %grid_high76, %grid_high77, %grid_high78, %grid_high79, %grid_high80, %grid_high81, %grid_high82, %grid_high83, %grid_high84, %grid_high85, %grid_high86, %grid_high87, %grid_high88, %grid_high89, %grid_high90, %grid_high91, %grid_high92, %grid_high93, %grid_high94, %grid_high95 : vector<32xf32> + %grid_high3_chunk = vector.from_elements %grid_high96, %grid_high97, %grid_high98, %grid_high99, %grid_high100, %grid_high101, %grid_high102, %grid_high103, %grid_high104, %grid_high105, %grid_high106, %grid_high107, %grid_high108, %grid_high109, %grid_high110, %grid_high111, %grid_high112, %grid_high113, %grid_high114, %grid_high115, %grid_high116, %grid_high117, %grid_high118, %grid_high119, %grid_high120, %grid_high121, %grid_high122, %grid_high123, %grid_high124, %grid_high125, %grid_high126, %grid_high127 : vector<32xf32> + %grid_high4_chunk = vector.from_elements %grid_high128, %grid_high129, %grid_high130, %grid_high131, %grid_high132, %grid_high133, %grid_high134, %grid_high135, %grid_high136, %grid_high137, %grid_high138, %grid_high139, %grid_high140, %grid_high141, %grid_high142, %grid_high143, %grid_high144, %grid_high145, %grid_high146, %grid_high147, %grid_high148, %grid_high149, %grid_high150, %grid_high151, %grid_high152, %grid_high153, %grid_high154, %grid_high155, %grid_high156, %grid_high157, %grid_high158, %grid_high159 : vector<32xf32> + %grid_high5_chunk = vector.from_elements %grid_high160, %grid_high161, %grid_high162, %grid_high163, %grid_high164, %grid_high165, %grid_high166, %grid_high167, %grid_high168, %grid_high169, %grid_high170, %grid_high171, %grid_high172, %grid_high173, %grid_high174, %grid_high175, %grid_high176, %grid_high177, %grid_high178, %grid_high179, %grid_high180, %grid_high181, %grid_high182, %grid_high183, %grid_high184, %grid_high185, %grid_high186, %grid_high187, %grid_high188, %grid_high189, %grid_high190, %grid_high191 : vector<32xf32> + %grid_high6_chunk = vector.from_elements %grid_high192, %grid_high193, %grid_high194, %grid_high195, %grid_high196, %grid_high197, %grid_high198, %grid_high199, %grid_high200, %grid_high201, %grid_high202, %grid_high203, %grid_high204, %grid_high205, %grid_high206, %grid_high207, %grid_high208, %grid_high209, %grid_high210, %grid_high211, %grid_high212, %grid_high213, %grid_high214, %grid_high215, %grid_high216, %grid_high217, %grid_high218, %grid_high219, %grid_high220, %grid_high221, %grid_high222, %grid_high223 : vector<32xf32> + %grid_high7_chunk = vector.from_elements %grid_high224, %grid_high225, %grid_high226, %grid_high227, %grid_high228, %grid_high229, %grid_high230, %grid_high231, %grid_high232, %grid_high233, %grid_high234, %grid_high235, %grid_high236, %grid_high237, %grid_high238, %grid_high239, %grid_high240, %grid_high241, %grid_high242, %grid_high243, %grid_high244, %grid_high245, %grid_high246, %grid_high247, %grid_high248, %grid_high249, %grid_high250, %grid_high251, %grid_high252, %grid_high253, %grid_high254, %grid_high255 : vector<32xf32> + %grid_high8_chunk = vector.from_elements %grid_high256, %grid_high257, %grid_high258, %grid_high259, %grid_high260, %grid_high261, %grid_high262, %grid_high263, %grid_high264, %grid_high265, %grid_high266, %grid_high267, %grid_high268, %grid_high269, %grid_high270, %grid_high271, %grid_high272, %grid_high273, %grid_high274, %grid_high275, %grid_high276, %grid_high277, %grid_high278, %grid_high279, %grid_high280, %grid_high281, %grid_high282, %grid_high283, %grid_high284, %grid_high285, %grid_high286, %grid_high287 : vector<32xf32> + %grid_high9_chunk = vector.from_elements %grid_high288, %grid_high289, %grid_high290, %grid_high291, %grid_high292, %grid_high293, %grid_high294, %grid_high295, %grid_high296, %grid_high297, %grid_high298, %grid_high299, %grid_high300, %grid_high301, %grid_high302, %grid_high303, %grid_high304, %grid_high305, %grid_high306, %grid_high307, %grid_high308, %grid_high309, %grid_high310, %grid_high311, %grid_high312, %grid_high313, %grid_high314, %grid_high315, %grid_high316, %grid_high317, %grid_high318, %grid_high319 : vector<32xf32> + %grid_high10_chunk = vector.from_elements %grid_high320, %grid_high321, %grid_high322, %grid_high323, %grid_high324, %grid_high325, %grid_high326, %grid_high327, %grid_high328, %grid_high329, %grid_high330, %grid_high331, %grid_high332, %grid_high333, %grid_high334, %grid_high335, %grid_high336, %grid_high337, %grid_high338, %grid_high339, %grid_high340, %grid_high341, %grid_high342, %grid_high343, %grid_high344, %grid_high345, %grid_high346, %grid_high347, %grid_high348, %grid_high349, %grid_high350, %grid_high351 : vector<32xf32> + %grid_high11_chunk = vector.from_elements %grid_high352, %grid_high353, %grid_high354, %grid_high355, %grid_high356, %grid_high357, %grid_high358, %grid_high359, %grid_high360, %grid_high361, %grid_high362, %grid_high363, %grid_high364, %grid_high365, %grid_high366, %grid_high367, %grid_high368, %grid_high369, %grid_high370, %grid_high371, %grid_high372, %grid_high373, %grid_high374, %grid_high375, %grid_high376, %grid_high377, %grid_high378, %grid_high379, %grid_high380, %grid_high381, %grid_high382, %grid_high383 : vector<32xf32> + %grid_high12_chunk = vector.from_elements %grid_high384, %grid_high385, %grid_high386, %grid_high387, %grid_high388, %grid_high389, %grid_high390, %grid_high391, %grid_high392, %grid_high393, %grid_high394, %grid_high395, %grid_high396, %grid_high397, %grid_high398, %grid_high399, %grid_high400, %grid_high401, %grid_high402, %grid_high403, %grid_high404, %grid_high405, %grid_high406, %grid_high407, %grid_high408, %grid_high409, %grid_high410, %grid_high411, %grid_high412, %grid_high413, %grid_high414, %grid_high415 : vector<32xf32> + %grid_high13_chunk = vector.from_elements %grid_high416, %grid_high417, %grid_high418, %grid_high419, %grid_high420, %grid_high421, %grid_high422, %grid_high423, %grid_high424, %grid_high425, %grid_high426, %grid_high427, %grid_high428, %grid_high429, %grid_high430, %grid_high431, %grid_high432, %grid_high433, %grid_high434, %grid_high435, %grid_high436, %grid_high437, %grid_high438, %grid_high439, %grid_high440, %grid_high441, %grid_high442, %grid_high443, %grid_high444, %grid_high445, %grid_high446, %grid_high447 : vector<32xf32> + %grid_high14_chunk = vector.from_elements %grid_high448, %grid_high449, %grid_high450, %grid_high451, %grid_high452, %grid_high453, %grid_high454, %grid_high455, %grid_high456, %grid_high457, %grid_high458, %grid_high459, %grid_high460, %grid_high461, %grid_high462, %grid_high463, %grid_high464, %grid_high465, %grid_high466, %grid_high467, %grid_high468, %grid_high469, %grid_high470, %grid_high471, %grid_high472, %grid_high473, %grid_high474, %grid_high475, %grid_high476, %grid_high477, %grid_high478, %grid_high479 : vector<32xf32> + %grid_high15_chunk = vector.from_elements %grid_high480, %grid_high481, %grid_high482, %grid_high483, %grid_high484, %grid_high485, %grid_high486, %grid_high487, %grid_high488, %grid_high489, %grid_high490, %grid_high491, %grid_high492, %grid_high493, %grid_high494, %grid_high495, %grid_high496, %grid_high497, %grid_high498, %grid_high499, %grid_high500, %grid_high501, %grid_high502, %grid_high503, %grid_high504, %grid_high505, %grid_high506, %grid_high507, %grid_high508, %grid_high509, %grid_high510, %grid_high511 : vector<32xf32> + func.return %grid_low0_chunk, %grid_low1_chunk, %grid_low2_chunk, %grid_low3_chunk, %grid_low4_chunk, %grid_low5_chunk, %grid_low6_chunk, %grid_low7_chunk, %grid_low8_chunk, %grid_low9_chunk, %grid_low10_chunk, %grid_low11_chunk, %grid_low12_chunk, %grid_low13_chunk, %grid_low14_chunk, %grid_low15_chunk, %grid_high0_chunk, %grid_high1_chunk, %grid_high2_chunk, %grid_high3_chunk, %grid_high4_chunk, %grid_high5_chunk, %grid_high6_chunk, %grid_high7_chunk, %grid_high8_chunk, %grid_high9_chunk, %grid_high10_chunk, %grid_high11_chunk, %grid_high12_chunk, %grid_high13_chunk, %grid_high14_chunk, %grid_high15_chunk : vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32> +} + +func.def inline @ggml_iq3s_grid_lookup_i32(%iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %grid_index: i32) -> (i32) { + %c5_i32 = scalar.constant 5 : i32 + %c31_i32 = scalar.constant 31 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c6_i32 = scalar.constant 6 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c9_i32 = scalar.constant 9 : i32 + %c10_i32 = scalar.constant 10 : i32 + %c11_i32 = scalar.constant 11 : i32 + %c12_i32 = scalar.constant 12 : i32 + %c13_i32 = scalar.constant 13 : i32 + %c14_i32 = scalar.constant 14 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c16_i32 = scalar.constant 16 : i32 + %chunk_i32 = scalar.shrui %grid_index, %c5_i32 : i32 + %lane_i32 = scalar.andi %grid_index, %c31_i32 : i32 + %codes = vector.from_elements %lane_i32 : vector<1xi32> + %is_chunk1 = scalar.cmpi eq, %chunk_i32, %c1_i32 : i32 + %is_chunk2 = scalar.cmpi eq, %chunk_i32, %c2_i32 : i32 + %is_chunk3 = scalar.cmpi eq, %chunk_i32, %c3_i32 : i32 + %is_chunk4 = scalar.cmpi eq, %chunk_i32, %c4_i32 : i32 + %is_chunk5 = scalar.cmpi eq, %chunk_i32, %c5_i32 : i32 + %is_chunk6 = scalar.cmpi eq, %chunk_i32, %c6_i32 : i32 + %is_chunk7 = scalar.cmpi eq, %chunk_i32, %c7_i32 : i32 + %is_chunk8 = scalar.cmpi eq, %chunk_i32, %c8_i32 : i32 + %is_chunk9 = scalar.cmpi eq, %chunk_i32, %c9_i32 : i32 + %is_chunk10 = scalar.cmpi eq, %chunk_i32, %c10_i32 : i32 + %is_chunk11 = scalar.cmpi eq, %chunk_i32, %c11_i32 : i32 + %is_chunk12 = scalar.cmpi eq, %chunk_i32, %c12_i32 : i32 + %is_chunk13 = scalar.cmpi eq, %chunk_i32, %c13_i32 : i32 + %is_chunk14 = scalar.cmpi eq, %chunk_i32, %c14_i32 : i32 + %is_chunk15 = scalar.cmpi eq, %chunk_i32, %c15_i32 : i32 + %selected_low1 = scf.select %is_chunk1, %iq3s_grid_low1, %iq3s_grid_low0 : vector<32xf32> + %selected_low2 = scf.select %is_chunk2, %iq3s_grid_low2, %selected_low1 : vector<32xf32> + %selected_low3 = scf.select %is_chunk3, %iq3s_grid_low3, %selected_low2 : vector<32xf32> + %selected_low4 = scf.select %is_chunk4, %iq3s_grid_low4, %selected_low3 : vector<32xf32> + %selected_low5 = scf.select %is_chunk5, %iq3s_grid_low5, %selected_low4 : vector<32xf32> + %selected_low6 = scf.select %is_chunk6, %iq3s_grid_low6, %selected_low5 : vector<32xf32> + %selected_low7 = scf.select %is_chunk7, %iq3s_grid_low7, %selected_low6 : vector<32xf32> + %selected_low8 = scf.select %is_chunk8, %iq3s_grid_low8, %selected_low7 : vector<32xf32> + %selected_low9 = scf.select %is_chunk9, %iq3s_grid_low9, %selected_low8 : vector<32xf32> + %selected_low10 = scf.select %is_chunk10, %iq3s_grid_low10, %selected_low9 : vector<32xf32> + %selected_low11 = scf.select %is_chunk11, %iq3s_grid_low11, %selected_low10 : vector<32xf32> + %selected_low12 = scf.select %is_chunk12, %iq3s_grid_low12, %selected_low11 : vector<32xf32> + %selected_low13 = scf.select %is_chunk13, %iq3s_grid_low13, %selected_low12 : vector<32xf32> + %selected_low14 = scf.select %is_chunk14, %iq3s_grid_low14, %selected_low13 : vector<32xf32> + %selected_low = scf.select %is_chunk15, %iq3s_grid_low15, %selected_low14 : vector<32xf32> + %selected_high1 = scf.select %is_chunk1, %iq3s_grid_high1, %iq3s_grid_high0 : vector<32xf32> + %selected_high2 = scf.select %is_chunk2, %iq3s_grid_high2, %selected_high1 : vector<32xf32> + %selected_high3 = scf.select %is_chunk3, %iq3s_grid_high3, %selected_high2 : vector<32xf32> + %selected_high4 = scf.select %is_chunk4, %iq3s_grid_high4, %selected_high3 : vector<32xf32> + %selected_high5 = scf.select %is_chunk5, %iq3s_grid_high5, %selected_high4 : vector<32xf32> + %selected_high6 = scf.select %is_chunk6, %iq3s_grid_high6, %selected_high5 : vector<32xf32> + %selected_high7 = scf.select %is_chunk7, %iq3s_grid_high7, %selected_high6 : vector<32xf32> + %selected_high8 = scf.select %is_chunk8, %iq3s_grid_high8, %selected_high7 : vector<32xf32> + %selected_high9 = scf.select %is_chunk9, %iq3s_grid_high9, %selected_high8 : vector<32xf32> + %selected_high10 = scf.select %is_chunk10, %iq3s_grid_high10, %selected_high9 : vector<32xf32> + %selected_high11 = scf.select %is_chunk11, %iq3s_grid_high11, %selected_high10 : vector<32xf32> + %selected_high12 = scf.select %is_chunk12, %iq3s_grid_high12, %selected_high11 : vector<32xf32> + %selected_high13 = scf.select %is_chunk13, %iq3s_grid_high13, %selected_high12 : vector<32xf32> + %selected_high14 = scf.select %is_chunk14, %iq3s_grid_high14, %selected_high13 : vector<32xf32> + %selected_high = scf.select %is_chunk15, %iq3s_grid_high15, %selected_high14 : vector<32xf32> + %low_values = vector.table.lookup %selected_low[%codes] : vector<32xf32>, vector<1xi32> -> vector<1xf32> + %high_values = vector.table.lookup %selected_high[%codes] : vector<32xf32>, vector<1xi32> -> vector<1xf32> + %low_f32 = vector.extract %low_values[0] : vector<1xf32> -> f32 + %high_f32 = vector.extract %high_values[0] : vector<1xf32> -> f32 + %low = scalar.fptoui %low_f32 : f32 to i32 + %high = scalar.fptoui %high_f32 : f32 to i32 + %high_shifted = scalar.shli %high, %c16_i32 : i32 + %word = scalar.ori %low, %high_shifted : i32 + func.return %word : i32 +} + +func.def inline @ggml_iq3s_signed_value_f32(%grid_word: i32, %shift: i32, %shifted_signs: i32, %sign_bit: i32) -> (f32) { + %c0_i32 = scalar.constant 0 : i32 + %c255_i32 = scalar.constant 255 : i32 + %shifted_word = scalar.shrui %grid_word, %shift : i32 + %value_i32 = scalar.andi %shifted_word, %c255_i32 : i32 + %negative_value = scalar.subi %c0_i32, %value_i32 : i32 + %masked_sign = scalar.andi %shifted_signs, %sign_bit : i32 + %is_negative = scalar.cmpi ne, %masked_sign, %c0_i32 : i32 + %signed_value = scf.select %is_negative, %negative_value, %value_i32 : i32 + %value = scalar.sitofp %signed_value : i32 to f32 + func.return %value : f32 +} + +func.def inline @ggml_iq3s_f32_vector4(%iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight: buffer, %row_byte_base: offset, %iq3_block: index, %iq3_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c0_i32 = scalar.constant 0 : i32 + %c2 = index.constant 2 : index + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c4 = index.constant 4 : index + %c4_i32 = scalar.constant 4 : i32 + %c8 = index.constant 8 : index + %c8_i32 = scalar.constant 8 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c16 = index.constant 16 : index + %c16_i32 = scalar.constant 16 : i32 + %c24_i32 = scalar.constant 24 : i32 + %c256_i32 = scalar.constant 256 : i32 + %block_bytes = index.constant 110 : offset + %qs_offset = index.constant 2 : offset + %qh_offset = index.constant 66 : offset + %signs_offset = index.constant 74 : offset + %scales_offset = index.constant 106 : offset + %bounded_group = index.assume %iq3_group [range(%iq3_group, 0, 7)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %iq3_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %qs_byte_base = index.add %block_byte_base, %qs_offset : offset + %qh_byte_base = index.add %block_byte_base, %qh_offset : offset + %signs_byte_base = index.add %block_byte_base, %signs_offset : offset + %scales_byte_base = index.add %block_byte_base, %scales_offset : offset + %d_view = buffer.view %weight[%block_byte_base] : buffer -> view<1xf16> + %qs_view = buffer.view %weight[%qs_byte_base] : buffer -> view<64xi8> + %qh_view = buffer.view %weight[%qh_byte_base] : buffer -> view<8xi8> + %signs_view = buffer.view %weight[%signs_byte_base] : buffer -> view<32xi8> + %scales_view = buffer.view %weight[%scales_byte_base] : buffer -> view<4xi8> + %iqs_group_base = index.mul %bounded_group, %c8 : index + %iqs0 = index.add %iqs_group_base, %bounded_packet : index + %iqs = index.assume %iqs0 [range(%iqs0, 0, 63)] : index + %signs_index0 = index.div %iqs, %c2 : index + %signs_index = index.assume %signs_index0 [range(%signs_index0, 0, 31)] : index + %scale_index0 = index.div %iqs, %c16 : index + %scale_index = index.assume %scale_index0 [range(%scale_index0, 0, 3)] : index + %iqs_mod8_index = index.rem %iqs, %c8 : index + %high_shift_index = index.sub %c8, %iqs_mod8_index : index + %high_shift = index.cast %high_shift_index : index to i32 + %scale_shift_index = index.rem %bounded_group, %c2 : index + %scale_shift_base = index.cast %scale_shift_index : index to i32 + %scale_shift = scalar.muli %scale_shift_base, %c4_i32 : i32 + %sign_shift_index = index.rem %bounded_packet, %c2 : index + %sign_shift_base = index.cast %sign_shift_index : index to i32 + %sign_shift = scalar.muli %sign_shift_base, %c4_i32 : i32 + %q_i8 = view.load %qs_view[%iqs] : view<64xi8> -> i8 + %qh_i8 = view.load %qh_view[%bounded_group] : view<8xi8> -> i8 + %signs_i8 = view.load %signs_view[%signs_index] : view<32xi8> -> i8 + %scale_i8 = view.load %scales_view[%scale_index] : view<4xi8> -> i8 + %q_i32 = scalar.extui %q_i8 : i8 to i32 + %qh_i32 = scalar.extui %qh_i8 : i8 to i32 + %signs_i32 = scalar.extui %signs_i8 : i8 to i32 + %scale_i32 = scalar.extui %scale_i8 : i8 to i32 + %qh_shifted = scalar.shli %qh_i32, %high_shift : i32 + %qh_high = scalar.andi %qh_shifted, %c256_i32 : i32 + %grid_index = scalar.ori %q_i32, %qh_high : i32 + %grid_word = func.call @ggml_iq3s_grid_lookup_i32(%iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %grid_index) : (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, i32) -> (i32) + %scale_shifted = scalar.shrui %scale_i32, %scale_shift : i32 + %local_scale = scalar.andi %scale_shifted, %c15_i32 : i32 + %local_scale2 = scalar.muli %local_scale, %c2_i32 : i32 + %scale_code = scalar.addi %local_scale2, %c1_i32 : i32 + %scale_f32 = scalar.uitofp %scale_code : i32 to f32 + %d_f16 = view.load %d_view[%c0] : view<1xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %db = scalar.mulf %d, %scale_f32 : f32 + %shifted_signs = scalar.shrui %signs_i32, %sign_shift : i32 + %v0_base = func.call @ggml_iq3s_signed_value_f32(%grid_word, %c0_i32, %shifted_signs, %c1_i32) : (i32, i32, i32, i32) -> (f32) + %v1_base = func.call @ggml_iq3s_signed_value_f32(%grid_word, %c8_i32, %shifted_signs, %c2_i32) : (i32, i32, i32, i32) -> (f32) + %v2_base = func.call @ggml_iq3s_signed_value_f32(%grid_word, %c16_i32, %shifted_signs, %c4_i32) : (i32, i32, i32, i32) -> (f32) + %v3_base = func.call @ggml_iq3s_signed_value_f32(%grid_word, %c24_i32, %shifted_signs, %c8_i32) : (i32, i32, i32, i32) -> (f32) + %v0 = scalar.mulf %db, %v0_base : f32 + %v1 = scalar.mulf %db, %v1_base : f32 + %v2 = scalar.mulf %db, %v2_base : f32 + %v3 = scalar.mulf %db, %v3_base : f32 + %result = vector.from_elements %v0, %v1, %v2, %v3 : vector<4xf32> + func.return %result : vector<4xf32> +} + +func.def inline @ggml_iq3s_f16_vector4(%iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight: buffer, %row_byte_base: offset, %iq3_block: index, %iq3_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_iq3s_f32_vector4(%iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight, %row_byte_base, %iq3_block, %iq3_group, %packet) : (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +func.def inline @ggml_iq4nl_table_i8() -> (vector<16xi8>) { + %v0 = scalar.constant -127 : i8 + %v1 = scalar.constant -104 : i8 + %v2 = scalar.constant -83 : i8 + %v3 = scalar.constant -65 : i8 + %v4 = scalar.constant -49 : i8 + %v5 = scalar.constant -35 : i8 + %v6 = scalar.constant -22 : i8 + %v7 = scalar.constant -10 : i8 + %v8 = scalar.constant 1 : i8 + %v9 = scalar.constant 13 : i8 + %v10 = scalar.constant 25 : i8 + %v11 = scalar.constant 38 : i8 + %v12 = scalar.constant 53 : i8 + %v13 = scalar.constant 69 : i8 + %v14 = scalar.constant 89 : i8 + %v15 = scalar.constant 113 : i8 + %value_table = vector.from_elements %v0, %v1, %v2, %v3, %v4, %v5, %v6, %v7, %v8, %v9, %v10, %v11, %v12, %v13, %v14, %v15 : vector<16xi8> + func.return %value_table : vector<16xi8> +} + +func.def inline @ggml_iq4nl_lookup_i8_vector4(%value_table: vector<16xi8>, %codes: vector<4xi8>) -> (vector<4xi8>) { + %values = vector.table.lookup %value_table[%codes] : vector<16xi8>, vector<4xi8> -> vector<4xi8> + func.return %values : vector<4xi8> +} + +func.def inline @ggml_iq4nl_table_codes4(%q_bytes: vector<4xi8>, %uses_high: i1) -> (vector<4xi8>) { + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c12_i32 = scalar.constant 12 : i32 + %c16_i32 = scalar.constant 16 : i32 + %mask0_i32 = scalar.constant 15 : i32 + %mask1_i32 = scalar.constant 240 : i32 + %mask2_i32 = scalar.constant 3840 : i32 + %mask3_i32 = scalar.constant 61440 : i32 + %q_word_vector = vector.bitcast %q_bytes : vector<4xi8> to vector<1xi32> + %word = vector.extract %q_word_vector[0] : vector<1xi32> -> i32 + %word_shr4 = scalar.shrui %word, %c4_i32 : i32 + %word_shr8 = scalar.shrui %word, %c8_i32 : i32 + %word_shr12 = scalar.shrui %word, %c12_i32 : i32 + %word_shr16 = scalar.shrui %word, %c16_i32 : i32 + %low0 = scalar.andi %word, %mask0_i32 : i32 + %low1 = scalar.andi %word_shr4, %mask1_i32 : i32 + %low2 = scalar.andi %word_shr8, %mask2_i32 : i32 + %low3 = scalar.andi %word_shr12, %mask3_i32 : i32 + %low01 = scalar.ori %low0, %low1 : i32 + %low23 = scalar.ori %low2, %low3 : i32 + %low = scalar.ori %low01, %low23 : i32 + %high0 = scalar.andi %word_shr4, %mask0_i32 : i32 + %high1 = scalar.andi %word_shr8, %mask1_i32 : i32 + %high2 = scalar.andi %word_shr12, %mask2_i32 : i32 + %high3 = scalar.andi %word_shr16, %mask3_i32 : i32 + %high01 = scalar.ori %high0, %high1 : i32 + %high23 = scalar.ori %high2, %high3 : i32 + %high = scalar.ori %high01, %high23 : i32 + %selected = scf.select %uses_high, %high, %low : i32 + %target_selected_shr8 = scalar.shrui %selected, %c8_i32 : i32 + %packed0 = scalar.trunci %selected : i32 to i8 + %packed1 = scalar.trunci %target_selected_shr8 : i32 to i8 + %packed = vector.from_elements %packed0, %packed1 : vector<2xi8> + %codes = vector.bitunpacku<4> %packed : vector<2xi8> -> vector<4xi8> + func.return %codes : vector<4xi8> +} + +func.def inline @ggml_iq4nl_f32_vector4(%iq4nl_table: vector<16xi8>, %weight: buffer, %row_byte_base: offset, %iq4_block: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %block_bytes = index.constant 18 : offset + %code_offset = index.constant 2 : offset + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %iq4_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %d_view = buffer.view %weight[%block_byte_base] : buffer -> view<1xf16> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<16xi8> + %packet_page = index.rem %bounded_packet, %c4 : index + %code_index0 = index.mul %packet_page, %c4 : index + %code_index = index.assume %code_index0 [range(%code_index0, 0, 12), mul(%code_index0, 4)] : index + %uses_high = index.cmp uge, %bounded_packet, %c4 : index + %q_bytes = vector.load %code_view[%code_index] : view<16xi8> -> vector<4xi8> + %codes = func.call @ggml_iq4nl_table_codes4(%q_bytes, %uses_high) : (vector<4xi8>, i1) -> (vector<4xi8>) + %lookup_i8 = func.call @ggml_iq4nl_lookup_i8_vector4(%iq4nl_table, %codes) : (vector<16xi8>, vector<4xi8>) -> (vector<4xi8>) + %lookup = vector.sitofp %lookup_i8 : vector<4xi8> to vector<4xf32> + %d_f16 = view.load %d_view[%c0] : view<1xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %d_vector = vector.splat %d : vector<4xf32> + %result = vector.mulf %d_vector, %lookup : vector<4xf32> + func.return %result : vector<4xf32> +} + +func.def inline @ggml_iq4nl_f16_vector4(%iq4nl_table: vector<16xi8>, %weight: buffer, %row_byte_base: offset, %iq4_block: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_iq4nl_f32_vector4(%iq4nl_table, %weight, %row_byte_base, %iq4_block, %packet) : (vector<16xi8>, buffer, offset, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +// Packs either nibble from four adjacent GGUF IQ4_XS bytes for the gfx11 table lookup. +// The nibble is the table index as it is: vector.table.lookup indexes the table in order on gfx1151 +// (an XOR 12 remap for a reversing V_PERM made every IQ4_NL / IQ4_XS weight wrong there). +func.def inline @ggml_iq4xs_table_codes4(%code_word: vector<1xi32>, %uses_high: i1) -> (vector<4xi8>) { + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c12_i32 = scalar.constant 12 : i32 + %c16_i32 = scalar.constant 16 : i32 + %mask0_i32 = scalar.constant 15 : i32 + %mask1_i32 = scalar.constant 240 : i32 + %mask2_i32 = scalar.constant 3840 : i32 + %mask3_i32 = scalar.constant 61440 : i32 + %word = vector.extract %code_word[0] : vector<1xi32> -> i32 + %word_shr4 = scalar.shrui %word, %c4_i32 : i32 + %word_shr8 = scalar.shrui %word, %c8_i32 : i32 + %word_shr12 = scalar.shrui %word, %c12_i32 : i32 + %word_shr16 = scalar.shrui %word, %c16_i32 : i32 + %low0 = scalar.andi %word, %mask0_i32 : i32 + %low1 = scalar.andi %word_shr4, %mask1_i32 : i32 + %low2 = scalar.andi %word_shr8, %mask2_i32 : i32 + %low3 = scalar.andi %word_shr12, %mask3_i32 : i32 + %low01 = scalar.ori %low0, %low1 : i32 + %low23 = scalar.ori %low2, %low3 : i32 + %low = scalar.ori %low01, %low23 : i32 + %high0 = scalar.andi %word_shr4, %mask0_i32 : i32 + %high1 = scalar.andi %word_shr8, %mask1_i32 : i32 + %high2 = scalar.andi %word_shr12, %mask2_i32 : i32 + %high3 = scalar.andi %word_shr16, %mask3_i32 : i32 + %high01 = scalar.ori %high0, %high1 : i32 + %high23 = scalar.ori %high2, %high3 : i32 + %high = scalar.ori %high01, %high23 : i32 + %selected = scf.select %uses_high, %high, %low : i32 + %target_selected_shr8 = scalar.shrui %selected, %c8_i32 : i32 + %packed0 = scalar.trunci %selected : i32 to i8 + %packed1 = scalar.trunci %target_selected_shr8 : i32 to i8 + %packed = vector.from_elements %packed0, %packed1 : vector<2xi8> + %codes = vector.bitunpacku<4> %packed : vector<2xi8> -> vector<4xi8> + func.return %codes : vector<4xi8> +} + +func.def inline @ggml_iq4xs_f32_vector4(%iq4nl_table: vector<16xi8>, %weight: buffer, %row_byte_base: offset, %iq4_block: index, %iq4_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %block_bytes = index.constant 136 : offset + %scales_l_offset = index.constant 4 : offset + %code_offset = index.constant 8 : offset + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c16_i32 = scalar.constant 16 : i32 + %c32_i32 = scalar.constant 32 : i32 + %c65535_i32 = scalar.constant 65535 : i32 + %bounded_group = index.assume %iq4_group [range(%iq4_group, 0, 7)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %iq4_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %scales_l_byte_base = index.add %block_byte_base, %scales_l_offset : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %header_view = buffer.view %weight[%block_byte_base] : buffer -> view<1xi32> + %scales_l_view = buffer.view %weight[%scales_l_byte_base] : buffer -> view<1xi32> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<32xi32> + %header = view.load %header_view[%c0] : view<1xi32> -> i32 + %header_vector = vector.from_elements %header : vector<1xi32> + %header_halves = vector.bitcast %header_vector : vector<1xi32> to vector<2xf16> + %d_f16 = vector.extract %header_halves[0] : vector<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %scales_l = view.load %scales_l_view[%c0] : view<1xi32> -> i32 + %scales_h_shifted = scalar.shrui %header, %c16_i32 : i32 + %scales_h = scalar.andi %scales_h_shifted, %c65535_i32 : i32 + %scale_l_shift_index = index.mul %bounded_group, %c4 : index + %scale_l_shift = index.cast %scale_l_shift_index : index to i32 + %scale_l_shifted = scalar.shrui %scales_l, %scale_l_shift : i32 + %scale_l = scalar.andi %scale_l_shifted, %c15_i32 : i32 + %scale_h_shift_index = index.mul %bounded_group, %c2 : index + %scale_h_shift = index.cast %scale_h_shift_index : index to i32 + %scale_h_shifted = scalar.shrui %scales_h, %scale_h_shift : i32 + %scale_h = scalar.andi %scale_h_shifted, %c3_i32 : i32 + %scale_h_high = scalar.shli %scale_h, %c4_i32 : i32 + %local_scale = scalar.ori %scale_l, %scale_h_high : i32 + %centered_scale = scalar.subi %local_scale, %c32_i32 : i32 + %centered_scale_f32 = scalar.sitofp %centered_scale : i32 to f32 + %block_scale = scalar.mulf %d, %centered_scale_f32 : f32 + %packet_page = index.rem %bounded_packet, %c4 : index + %q_word_base = index.mul %bounded_group, %c4 : index + %q_word_index0 = index.add %q_word_base, %packet_page : index + %q_word_index = index.assume %q_word_index0 [range(%q_word_index0, 0, 31)] : index + %q_word = vector.load %code_view[%q_word_index] : view<32xi32> -> vector<1xi32> + %uses_high = index.cmp uge, %bounded_packet, %c4 : index + %codes = func.call @ggml_iq4xs_table_codes4(%q_word, %uses_high) : (vector<1xi32>, i1) -> (vector<4xi8>) + %lookup_i8 = func.call @ggml_iq4nl_lookup_i8_vector4(%iq4nl_table, %codes) : (vector<16xi8>, vector<4xi8>) -> (vector<4xi8>) + %lookup0_i8 = vector.extract %lookup_i8[0] : vector<4xi8> -> i8 + %lookup1_i8 = vector.extract %lookup_i8[1] : vector<4xi8> -> i8 + %lookup2_i8 = vector.extract %lookup_i8[2] : vector<4xi8> -> i8 + %lookup3_i8 = vector.extract %lookup_i8[3] : vector<4xi8> -> i8 + %lookup0_i32 = scalar.extsi %lookup0_i8 : i8 to i32 + %lookup1_i32 = scalar.extsi %lookup1_i8 : i8 to i32 + %lookup2_i32 = scalar.extsi %lookup2_i8 : i8 to i32 + %lookup3_i32 = scalar.extsi %lookup3_i8 : i8 to i32 + %lookup0 = scalar.sitofp %lookup0_i32 : i32 to f32 + %lookup1 = scalar.sitofp %lookup1_i32 : i32 to f32 + %lookup2 = scalar.sitofp %lookup2_i32 : i32 to f32 + %lookup3 = scalar.sitofp %lookup3_i32 : i32 to f32 + %value0 = scalar.mulf %block_scale, %lookup0 : f32 + %value1 = scalar.mulf %block_scale, %lookup1 : f32 + %value2 = scalar.mulf %block_scale, %lookup2 : f32 + %value3 = scalar.mulf %block_scale, %lookup3 : f32 + %result = vector.from_elements %value0, %value1, %value2, %value3 : vector<4xf32> + func.return %result : vector<4xf32> +} + +func.def inline @ggml_iq4xs_f16_vector4(%iq4nl_table: vector<16xi8>, %weight: buffer, %row_byte_base: offset, %iq4_block: index, %iq4_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_iq4xs_f32_vector4(%iq4nl_table, %weight, %row_byte_base, %iq4_block, %iq4_group, %packet) : (vector<16xi8>, buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +func.def inline @ggml_q6k_f32_vector4(%weight: buffer, %weight_row_byte_base: offset, %q6_block: index, %q6_group: index, %packet: index) -> (vector<4xf32>) { + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %block_bytes = index.constant 210 : offset + %qh_byte_add = index.constant 128 : offset + %scale_byte_add = index.constant 192 : offset + %d_byte_add = index.constant 208 : offset + %c4_i32v = vector.constant 4 : vector<1xi32> + %nibble_mask = vector.constant 252645135 : vector<1xi32> + %high_mask = vector.constant 50529027 : vector<1xi32> + %c32_f32v = vector.constant 32.0 : vector<4xf32> + %bounded_group = index.assume %q6_group [range(%q6_group, 0, 7)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q6_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %weight_row_byte_base, %block_byte_add : offset + %qh_byte_base = index.add %block_byte_base, %qh_byte_add : offset + %scale_byte_base = index.add %block_byte_base, %scale_byte_add : offset + %d_byte_base = index.add %block_byte_base, %d_byte_add : offset + %ql_view = buffer.view %weight[%block_byte_base] : buffer -> view<32xi32> + %qh_view = buffer.view %weight[%qh_byte_base] : buffer -> view<16xi32> + %scale_view = buffer.view %weight[%scale_byte_base] : buffer -> view<16xi8> + %d_view = buffer.view %weight[%d_byte_base] : buffer -> view<1xf16> + %group_in_half = index.rem %bounded_group, %c4 : index + %half = index.div %bounded_group, %c4 : index + %ql_side = index.rem %group_in_half, %c2 : index + %ql_half_word_base = index.mul %half, %c16 : index + %ql_side_word_add = index.mul %ql_side, %c8 : index + %ql_word_base = index.add %ql_half_word_base, %ql_side_word_add : index + %ql_word_index = index.add %ql_word_base, %bounded_packet : index + %qh_half_word_base = index.mul %half, %c8 : index + %qh_word_index = index.add %qh_half_word_base, %bounded_packet : index + %nibble = index.div %group_in_half, %c2 : index + %nibble_shift_index = index.mul %nibble, %c4 : index + %nibble_shift_i32 = index.cast %nibble_shift_index : index to i32 + %nibble_shift = vector.splat %nibble_shift_i32 : vector<1xi32> + %qh_shift_index = index.mul %group_in_half, %c2 : index + %qh_shift_i32 = index.cast %qh_shift_index : index to i32 + %qh_shift = vector.splat %qh_shift_i32 : vector<1xi32> + %scale_packet_half = index.div %bounded_packet, %c4 : index + %scale_group_base = index.mul %bounded_group, %c2 : index + %scale_index = index.add %scale_group_base, %scale_packet_half : index + %ql_word = vector.load %ql_view[%ql_word_index] : view<32xi32> -> vector<1xi32> + %qh_word = vector.load %qh_view[%qh_word_index] : view<16xi32> -> vector<1xi32> + %ql_shifted = vector.shrui %ql_word, %nibble_shift : vector<1xi32> + %ql = vector.andi %ql_shifted, %nibble_mask : vector<1xi32> + %qh_shifted = vector.shrui %qh_word, %qh_shift : vector<1xi32> + %qh_low = vector.andi %qh_shifted, %high_mask : vector<1xi32> + %qh = vector.shli %qh_low, %c4_i32v : vector<1xi32> + %code = vector.ori %ql, %qh : vector<1xi32> + %code_i8 = vector.bitcast %code : vector<1xi32> to vector<4xi8> + %code_f32 = vector.uitofp %code_i8 : vector<4xi8> to vector<4xf32> + %centered = vector.subf %code_f32, %c32_f32v : vector<4xf32> + %scale_i8 = view.load %scale_view[%scale_index] : view<16xi8> -> i8 + %d_f16 = view.load %d_view[0] : view<1xf16> -> f16 + %scale = scalar.sitofp %scale_i8 : i8 to f32 + %d = scalar.extf %d_f16 : f16 to f32 + %combined_scale = scalar.mulf %scale, %d : f32 + %combined_scale_vector = vector.splat %combined_scale : vector<4xf32> + %values = vector.mulf %centered, %combined_scale_vector : vector<4xf32> + func.return %values : vector<4xf32> +} + +func.def inline @ggml_q8_0_f32_vector4(%weight: buffer, %row_byte_base: offset, %q8_block: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %block_bytes = index.constant 34 : offset + %code_offset = index.constant 2 : offset + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q8_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %d_view = buffer.view %weight[%block_byte_base] : buffer -> view<1xf16> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<32xi8> + %packet_base = index.mul %bounded_packet, %c4 : index + %codes = vector.load %code_view[%packet_base] : view<32xi8> -> vector<4xi8> + %codes_f32 = vector.sitofp %codes : vector<4xi8> to vector<4xf32> + %d_f16 = view.load %d_view[%c0] : view<1xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %d_vector = vector.splat %d : vector<4xf32> + %values = vector.mulf %codes_f32, %d_vector : vector<4xf32> + func.return %values : vector<4xf32> +} + +func.def inline @ggml_q8_1_f32_vector4(%weight: buffer, %row_byte_base: offset, %q8_block: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %block_bytes = index.constant 36 : offset + %code_offset = index.constant 4 : offset + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q8_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %ds_view = buffer.view %weight[%block_byte_base] : buffer -> view<2xf16> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<32xi8> + %packet_base = index.mul %bounded_packet, %c4 : index + %codes = vector.load %code_view[%packet_base] : view<32xi8> -> vector<4xi8> + %codes_f32 = vector.sitofp %codes : vector<4xi8> to vector<4xf32> + %d_f16 = view.load %ds_view[%c0] : view<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %d_vector = vector.splat %d : vector<4xf32> + %values = vector.mulf %codes_f32, %d_vector : vector<4xf32> + func.return %values : vector<4xf32> +} + +func.def inline @ggml_f16_f32_vector4(%weight: buffer, %row_byte_base: offset, %input_size: index, %k: index) -> (vector<4xf32>) { + %bounded_input_size = index.assume %input_size [range(%input_size, 4, 1073741824), mul(%input_size, 4)] : index + %row_view = buffer.view %weight[%row_byte_base] : buffer -> view<[%bounded_input_size]xf16> + %values_f16 = vector.load %row_view[%k] : view<[%bounded_input_size]xf16> -> vector<4xf16> + %values = vector.extf %values_f16 : vector<4xf16> to vector<4xf32> + func.return %values : vector<4xf32> +} + +func.def inline @ggml_bf16_f32_vector4(%weight: buffer, %row_byte_base: offset, %input_size: index, %k: index) -> (vector<4xf32>) { + %bounded_input_size = index.assume %input_size [range(%input_size, 4, 1073741824), mul(%input_size, 4)] : index + %row_view = buffer.view %weight[%row_byte_base] : buffer -> view<[%bounded_input_size]xbf16> + %values_bf16 = vector.load %row_view[%k] : view<[%bounded_input_size]xbf16> -> vector<4xbf16> + %values = vector.extf %values_bf16 : vector<4xbf16> to vector<4xf32> + func.return %values : vector<4xf32> +} + +func.def inline @ggml_f32_f32_vector4(%weight: buffer, %row_byte_base: offset, %input_size: index, %k: index) -> (vector<4xf32>) { + %bounded_input_size = index.assume %input_size [range(%input_size, 4, 1073741824), mul(%input_size, 4)] : index + %row_view = buffer.view %weight[%row_byte_base] : buffer -> view<[%bounded_input_size]xf32> + %values = vector.load %row_view[%k] : view<[%bounded_input_size]xf32> -> vector<4xf32> + func.return %values : vector<4xf32> +} + +func.def inline @ggml_dequant_weight_tile_bytes(%weight_format: index) -> (offset) { + %q1_0_format = index.constant 10 : index + %c3 = index.constant 11 : index + %c4 = index.constant 4 : index + %c5 = index.constant 5 : index + %c6 = index.constant 6 : index + %c20 = index.constant 20 : index + %c21 = index.constant 21 : index + %c22 = index.constant 22 : index + %c23 = index.constant 23 : index + %c16 = index.constant 16 : index + %bf16_format = index.constant 30 : index + %c32 = index.constant 32 : index + %q4_0_format = index.constant 40 : index + %q4_1_format = index.constant 41 : index + %q5_0_format = index.constant 50 : index + %q5_1_format = index.constant 51 : index + %q8_0_format = index.constant 80 : index + %q8_1_format = index.constant 81 : index + %zero_bytes = index.constant 0 : offset + %q1_0_tile_bytes = index.constant 36 : offset + %q2_tile_bytes = index.constant 84 : offset + %iq2_xs_tile_bytes = index.constant 74 : offset + %iq2_xs_format = index.constant 25 : index + %iq2_xxs_tile_bytes = index.constant 66 : offset + %iq2_xxs_format = index.constant 24 : index + %iq3_xxs_tile_bytes = index.constant 98 : offset + %iq3_xxs_format = index.constant 28 : index + %iq1_s_tile_bytes = index.constant 50 : offset + %iq1_s_format = index.constant 26 : index + %iq1_m_tile_bytes = index.constant 56 : offset + %iq1_m_format = index.constant 27 : index + %q3_tile_bytes = index.constant 110 : offset + %q4_tile_bytes = index.constant 144 : offset + %q5_tile_bytes = index.constant 176 : offset + %q6_tile_bytes = index.constant 210 : offset + %q4_0_tile_bytes = index.constant 144 : offset + %q4_1_tile_bytes = index.constant 160 : offset + %q5_0_tile_bytes = index.constant 176 : offset + %q5_1_tile_bytes = index.constant 192 : offset + %iq2_s_tile_bytes = index.constant 82 : offset + %iq4_nl_tile_bytes = index.constant 144 : offset + %iq3_s_tile_bytes = index.constant 110 : offset + %iq4_xs_tile_bytes = index.constant 136 : offset + %q8_0_tile_bytes = index.constant 272 : offset + %q8_1_tile_bytes = index.constant 288 : offset + %f16_tile_bytes = index.constant 512 : offset + %bf16_tile_bytes = index.constant 512 : offset + %f32_tile_bytes = index.constant 1024 : offset + %is_q1_0 = index.cmp eq, %weight_format, %q1_0_format : index + %q2_format = index.constant 12 : index + %is_q2 = index.cmp eq, %weight_format, %q2_format : index + %is_q3 = index.cmp eq, %weight_format, %c3 : index + %is_q4 = index.cmp eq, %weight_format, %c4 : index + %is_q5 = index.cmp eq, %weight_format, %c5 : index + %is_q6 = index.cmp eq, %weight_format, %c6 : index + %is_q4_0 = index.cmp eq, %weight_format, %q4_0_format : index + %is_q4_1 = index.cmp eq, %weight_format, %q4_1_format : index + %is_q5_0 = index.cmp eq, %weight_format, %q5_0_format : index + %is_q5_1 = index.cmp eq, %weight_format, %q5_1_format : index + %is_iq2_s = index.cmp eq, %weight_format, %c22 : index + %is_iq4_nl = index.cmp eq, %weight_format, %c20 : index + %is_iq3_s = index.cmp eq, %weight_format, %c21 : index + %is_iq4_xs = index.cmp eq, %weight_format, %c23 : index + %is_q8_0 = index.cmp eq, %weight_format, %q8_0_format : index + %is_q8_1 = index.cmp eq, %weight_format, %q8_1_format : index + %is_f16 = index.cmp eq, %weight_format, %c16 : index + %is_bf16 = index.cmp eq, %weight_format, %bf16_format : index + %is_f32 = index.cmp eq, %weight_format, %c32 : index + %selected_q1_0 = scf.select %is_q1_0, %q1_0_tile_bytes, %zero_bytes : offset + %selected_q2 = scf.select %is_q2, %q2_tile_bytes, %selected_q1_0 : offset + %is_iq2_xxs = index.cmp eq, %weight_format, %iq2_xxs_format : index + %is_iq2_xs = index.cmp eq, %weight_format, %iq2_xs_format : index + %selected_iq2_xxs = scf.select %is_iq2_xxs, %iq2_xxs_tile_bytes, %selected_q2 : offset + %selected_iq2_xs = scf.select %is_iq2_xs, %iq2_xs_tile_bytes, %selected_iq2_xxs : offset + %is_iq3_xxs = index.cmp eq, %weight_format, %iq3_xxs_format : index + %selected_iq3_xxs = scf.select %is_iq3_xxs, %iq3_xxs_tile_bytes, %selected_iq2_xs : offset + %is_iq1_s = index.cmp eq, %weight_format, %iq1_s_format : index + %is_iq1_m = index.cmp eq, %weight_format, %iq1_m_format : index + %selected_iq1_s = scf.select %is_iq1_s, %iq1_s_tile_bytes, %selected_iq3_xxs : offset + %selected_iq1_m = scf.select %is_iq1_m, %iq1_m_tile_bytes, %selected_iq1_s : offset + %selected_q3 = scf.select %is_q3, %q3_tile_bytes, %selected_iq1_m : offset + %selected_q4 = scf.select %is_q4, %q4_tile_bytes, %selected_q3 : offset + %selected_q5 = scf.select %is_q5, %q5_tile_bytes, %selected_q4 : offset + %selected_q6 = scf.select %is_q6, %q6_tile_bytes, %selected_q5 : offset + %selected_q4_0 = scf.select %is_q4_0, %q4_0_tile_bytes, %selected_q6 : offset + %selected_q4_1 = scf.select %is_q4_1, %q4_1_tile_bytes, %selected_q4_0 : offset + %selected_q5_0 = scf.select %is_q5_0, %q5_0_tile_bytes, %selected_q4_1 : offset + %selected_q5_1 = scf.select %is_q5_1, %q5_1_tile_bytes, %selected_q5_0 : offset + %selected_iq2_s = scf.select %is_iq2_s, %iq2_s_tile_bytes, %selected_q5_1 : offset + %selected_iq4_nl = scf.select %is_iq4_nl, %iq4_nl_tile_bytes, %selected_iq2_s : offset + %mxfp4_format_t = index.constant 39 : index + %mxfp4_tile_bytes = index.constant 136 : offset + %is_mxfp4_t = index.cmp eq, %weight_format, %mxfp4_format_t : index + %selected_mxfp4_t = scf.select %is_mxfp4_t, %mxfp4_tile_bytes, %selected_iq4_nl : offset + %tq1_format_t = index.constant 34 : index + %tq2_format_t = index.constant 35 : index + %tq1_tile_bytes = index.constant 54 : offset + %tq2_tile_bytes = index.constant 66 : offset + %is_tq1_t = index.cmp eq, %weight_format, %tq1_format_t : index + %is_tq2_t = index.cmp eq, %weight_format, %tq2_format_t : index + %selected_tq1_t = scf.select %is_tq1_t, %tq1_tile_bytes, %selected_mxfp4_t : offset + %selected_tq2_t = scf.select %is_tq2_t, %tq2_tile_bytes, %selected_tq1_t : offset + // PrismML group-128 formats: two 34-byte PQ2_0 blocks or two 28-byte PTQ1_0 blocks per 256 values. Without them + // the tile size was 0 and the low-token SwiGLU (row bytes = K/256 x tile bytes) read row 0 for every row. + %pq2_format_t = index.constant 72 : index + %ptq1_format_t = index.constant 73 : index + %pq2_tile_bytes = index.constant 68 : offset + %ptq1_tile_bytes = index.constant 56 : offset + %is_pq2_t = index.cmp eq, %weight_format, %pq2_format_t : index + %is_ptq1_t = index.cmp eq, %weight_format, %ptq1_format_t : index + %selected_pq2_t = scf.select %is_pq2_t, %pq2_tile_bytes, %selected_tq2_t : offset + %selected_ptq1_t = scf.select %is_ptq1_t, %ptq1_tile_bytes, %selected_pq2_t : offset + %selected_iq3_s = scf.select %is_iq3_s, %iq3_s_tile_bytes, %selected_ptq1_t : offset + %selected_iq4_xs = scf.select %is_iq4_xs, %iq4_xs_tile_bytes, %selected_iq3_s : offset + %selected_q8_0 = scf.select %is_q8_0, %q8_0_tile_bytes, %selected_iq4_xs : offset + %selected_q8_1 = scf.select %is_q8_1, %q8_1_tile_bytes, %selected_q8_0 : offset + %selected_f16 = scf.select %is_f16, %f16_tile_bytes, %selected_q8_1 : offset + %selected_bf16 = scf.select %is_bf16, %bf16_tile_bytes, %selected_f16 : offset + %selected = scf.select %is_f32, %f32_tile_bytes, %selected_bf16 : offset + func.return %selected : offset +} + +func.def inline @ggml_dequant_weight_row_bytes(%weight_format: index, %hidden_size: index) -> (offset) { + %q1_0_format = index.constant 10 : index + %c3 = index.constant 11 : index + %c4 = index.constant 4 : index + %c5 = index.constant 5 : index + %c6 = index.constant 6 : index + %c20 = index.constant 20 : index + %c21 = index.constant 21 : index + %c22 = index.constant 22 : index + %c23 = index.constant 23 : index + %c16 = index.constant 16 : index + %bf16_format = index.constant 30 : index + %c32_format = index.constant 32 : index + %q4_0_format = index.constant 40 : index + %q4_1_format = index.constant 41 : index + %q5_0_format = index.constant 50 : index + %q5_1_format = index.constant 51 : index + %q8_0_format = index.constant 80 : index + %q8_1_format = index.constant 81 : index + %c32 = index.constant 32 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %zero_bytes = index.constant 0 : offset + %q1_0_block_bytes = index.constant 18 : offset + %q4_0_block_bytes = index.constant 18 : offset + %q4_1_block_bytes = index.constant 20 : offset + %q5_0_block_bytes = index.constant 22 : offset + %q5_1_block_bytes = index.constant 24 : offset + %iq4_nl_block_bytes = index.constant 18 : offset + %q8_0_block_bytes = index.constant 34 : offset + %q8_1_block_bytes = index.constant 36 : offset + %f16_element_bytes = index.constant 2 : offset + %f32_element_bytes = index.constant 4 : offset + %uses_k_tile0 = index.cmp eq, %weight_format, %c3 : index + %uses_k_tile1 = index.cmp eq, %weight_format, %c4 : index + %uses_k_tile2 = index.cmp eq, %weight_format, %c5 : index + %uses_k_tile3 = index.cmp eq, %weight_format, %c6 : index + %uses_k_tile4 = index.cmp eq, %weight_format, %c20 : index + %uses_k_tile5 = index.cmp eq, %weight_format, %c21 : index + %uses_k_tile6 = index.cmp eq, %weight_format, %c23 : index + %uses_k_tile7 = index.cmp eq, %weight_format, %c22 : index + %uses_k_tile01 = scalar.ori %uses_k_tile0, %uses_k_tile1 : i1 + %uses_k_tile23 = scalar.ori %uses_k_tile2, %uses_k_tile3 : i1 + %uses_k_tile45 = scalar.ori %uses_k_tile4, %uses_k_tile5 : i1 + %uses_k_tile67 = scalar.ori %uses_k_tile6, %uses_k_tile7 : i1 + %uses_k_tile0123 = scalar.ori %uses_k_tile01, %uses_k_tile23 : i1 + %uses_k_tile4567 = scalar.ori %uses_k_tile45, %uses_k_tile67 : i1 + %q2_format = index.constant 12 : index + %uses_k_tile8 = index.cmp eq, %weight_format, %q2_format : index + %uses_k_tile07 = scalar.ori %uses_k_tile0123, %uses_k_tile4567 : i1 + %iq2_xxs_format = index.constant 24 : index + %iq2_xs_format = index.constant 25 : index + %uses_k_tile9 = index.cmp eq, %weight_format, %iq2_xxs_format : index + %uses_k_tile10 = index.cmp eq, %weight_format, %iq2_xs_format : index + %uses_k_tile9_10 = scalar.ori %uses_k_tile9, %uses_k_tile10 : i1 + %uses_k_tile08 = scalar.ori %uses_k_tile07, %uses_k_tile8 : i1 + %iq3_xxs_format = index.constant 28 : index + %uses_k_tile13 = index.cmp eq, %weight_format, %iq3_xxs_format : index + %uses_k_tile0810 = scalar.ori %uses_k_tile08, %uses_k_tile9_10 : i1 + %iq1_s_format = index.constant 26 : index + %iq1_m_format = index.constant 27 : index + %uses_k_tile11 = index.cmp eq, %weight_format, %iq1_s_format : index + %uses_k_tile12 = index.cmp eq, %weight_format, %iq1_m_format : index + %uses_k_tile11_12 = scalar.ori %uses_k_tile11, %uses_k_tile12 : i1 + %uses_k_tile0813 = scalar.ori %uses_k_tile0810, %uses_k_tile13 : i1 + %uses_k_tile = scalar.ori %uses_k_tile0813, %uses_k_tile11_12 : i1 + %tile_bytes = func.call @ggml_dequant_weight_tile_bytes(%weight_format) : (index) -> (offset) + %tile_count = index.div %hidden_size, %c256 : index + %tile_row_bytes = index.scale %tile_count, %tile_bytes : index, offset -> offset + %is_q1_0 = index.cmp eq, %weight_format, %q1_0_format : index + %is_q4_0 = index.cmp eq, %weight_format, %q4_0_format : index + %is_q4_1 = index.cmp eq, %weight_format, %q4_1_format : index + %is_q5_0 = index.cmp eq, %weight_format, %q5_0_format : index + %is_q5_1 = index.cmp eq, %weight_format, %q5_1_format : index + %is_iq4_nl = index.cmp eq, %weight_format, %c20 : index + %is_q8_0 = index.cmp eq, %weight_format, %q8_0_format : index + %is_q8_1 = index.cmp eq, %weight_format, %q8_1_format : index + %is_f16 = index.cmp eq, %weight_format, %c16 : index + %is_bf16 = index.cmp eq, %weight_format, %bf16_format : index + %is_f32 = index.cmp eq, %weight_format, %c32_format : index + %q1_0_block_count = index.div %hidden_size, %c128 : index + %q1_0_row_bytes = index.scale %q1_0_block_count, %q1_0_block_bytes : index, offset -> offset + %legacy_block_count = index.div %hidden_size, %c32 : index + %q4_0_row_bytes = index.scale %legacy_block_count, %q4_0_block_bytes : index, offset -> offset + %q4_1_row_bytes = index.scale %legacy_block_count, %q4_1_block_bytes : index, offset -> offset + %q5_0_row_bytes = index.scale %legacy_block_count, %q5_0_block_bytes : index, offset -> offset + %q5_1_row_bytes = index.scale %legacy_block_count, %q5_1_block_bytes : index, offset -> offset + %iq4_nl_row_bytes = index.scale %legacy_block_count, %iq4_nl_block_bytes : index, offset -> offset + %q8_0_row_bytes = index.scale %legacy_block_count, %q8_0_block_bytes : index, offset -> offset + %q8_1_row_bytes = index.scale %legacy_block_count, %q8_1_block_bytes : index, offset -> offset + %f16_row_bytes = index.scale %hidden_size, %f16_element_bytes : index, offset -> offset + %f32_row_bytes = index.scale %hidden_size, %f32_element_bytes : index, offset -> offset + %selected_tile = scf.select %uses_k_tile, %tile_row_bytes, %zero_bytes : offset + %selected_q1_0 = scf.select %is_q1_0, %q1_0_row_bytes, %selected_tile : offset + %selected_q4_0 = scf.select %is_q4_0, %q4_0_row_bytes, %selected_q1_0 : offset + %selected_q4_1 = scf.select %is_q4_1, %q4_1_row_bytes, %selected_q4_0 : offset + %selected_q5_0 = scf.select %is_q5_0, %q5_0_row_bytes, %selected_q4_1 : offset + %selected_q5_1 = scf.select %is_q5_1, %q5_1_row_bytes, %selected_q5_0 : offset + %selected_iq4_nl = scf.select %is_iq4_nl, %iq4_nl_row_bytes, %selected_q5_1 : offset + %mxfp4_format_r = index.constant 39 : index + %mxfp4_block_bytes = index.constant 17 : offset + %is_mxfp4_r = index.cmp eq, %weight_format, %mxfp4_format_r : index + %mxfp4_row_bytes = index.scale %legacy_block_count, %mxfp4_block_bytes : index, offset -> offset + %selected_mxfp4_r = scf.select %is_mxfp4_r, %mxfp4_row_bytes, %selected_iq4_nl : offset + %tq1_format_r = index.constant 34 : index + %tq2_format_r = index.constant 35 : index + %is_tq1_r = index.cmp eq, %weight_format, %tq1_format_r : index + %is_tq2_r = index.cmp eq, %weight_format, %tq2_format_r : index + %is_tq_r = scalar.ori %is_tq1_r, %is_tq2_r : i1 + %selected_tq_r = scf.select %is_tq_r, %tile_row_bytes, %selected_mxfp4_r : offset + %selected_q8_0 = scf.select %is_q8_0, %q8_0_row_bytes, %selected_tq_r : offset + %selected_q8_1 = scf.select %is_q8_1, %q8_1_row_bytes, %selected_q8_0 : offset + %selected_f16 = scf.select %is_f16, %f16_row_bytes, %selected_q8_1 : offset + %selected_bf16 = scf.select %is_bf16, %f16_row_bytes, %selected_f16 : offset + %selected0 = scf.select %is_f32, %f32_row_bytes, %selected_bf16 : offset + %prism_pq2_format = index.constant 72 : index + %prism_ptq1_format = index.constant 73 : index + %is_prism_pq2 = index.cmp eq, %weight_format, %prism_pq2_format : index + %is_prism_ptq1 = index.cmp eq, %weight_format, %prism_ptq1_format : index + %prism_pq2_block_bytes = index.constant 34 : offset + %prism_ptq1_block_bytes = index.constant 28 : offset + %prism_pq2_row_bytes = index.scale %q1_0_block_count, %prism_pq2_block_bytes : index, offset -> offset + %prism_ptq1_row_bytes = index.scale %q1_0_block_count, %prism_ptq1_block_bytes : index, offset -> offset + %selected1 = scf.select %is_prism_pq2, %prism_pq2_row_bytes, %selected0 : offset + %selected = scf.select %is_prism_ptq1, %prism_ptq1_row_bytes, %selected1 : offset + func.return %selected : offset +} + +// Centralizes row-format decoding for FP16 WMMA staging. Each format path is +// independently guarded so inactive layouts do not issue memory loads. +// +// Format ids: +// 10: Q1_0 +// 12/11/4/5/6: Q2_K/Q3_K/Q4_K/Q5_K/Q6_K +// 40/41/50/51: Q4_0/Q4_1/Q5_0/Q5_1 +// 26/27/24/25/22/28/20/21/23: IQ1_S/IQ1_M/IQ2_XXS/IQ2_XS/IQ2_S/IQ3_XXS/IQ4_NL/IQ3_S/IQ4_XS +// 39/34/35: MXFP4/TQ1_0/TQ2_0 (motifs/dequant_1bit.loom) +// 80/81: Q8_0/Q8_1 +// 16/30/32: F16/BF16/F32 +func.def inline @ggml_dequant_f16_vector4(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight_format: index, %weight: buffer, %row_byte_base: offset, %input_size: index, %quant_block: index, %quant_group: index, %packet: index, %k: index) -> (vector<4xf16>) { + %q1_0_format = index.constant 10 : index + %c3 = index.constant 11 : index + %c4 = index.constant 4 : index + %c5 = index.constant 5 : index + %c6 = index.constant 6 : index + %c8 = index.constant 8 : index + %c20 = index.constant 20 : index + %c21 = index.constant 21 : index + %c22 = index.constant 22 : index + %c23 = index.constant 23 : index + %c16 = index.constant 16 : index + %bf16_format = index.constant 30 : index + %c32 = index.constant 32 : index + %q4_0_format = index.constant 40 : index + %q4_1_format = index.constant 41 : index + %q5_0_format = index.constant 50 : index + %q5_1_format = index.constant 51 : index + %q8_0_format = index.constant 80 : index + %q8_1_format = index.constant 81 : index + %zero = vector.constant 0.0 : vector<4xf16> + // Decode the compact integer format id into one predicate per layout. + %is_q1_0 = index.cmp eq, %weight_format, %q1_0_format : index + %q2_format = index.constant 12 : index + %iq2_xxs_format = index.constant 24 : index + %iq2_xs_format = index.constant 25 : index + %is_iq2_xxs = index.cmp eq, %weight_format, %iq2_xxs_format : index + %is_iq2_xs = index.cmp eq, %weight_format, %iq2_xs_format : index + %iq3_xxs_format = index.constant 28 : index + %is_iq3_xxs = index.cmp eq, %weight_format, %iq3_xxs_format : index + %iq1_s_format = index.constant 26 : index + %iq1_m_format = index.constant 27 : index + %is_iq1_s = index.cmp eq, %weight_format, %iq1_s_format : index + %is_iq1_m = index.cmp eq, %weight_format, %iq1_m_format : index + %is_q2 = index.cmp eq, %weight_format, %q2_format : index + %is_q3 = index.cmp eq, %weight_format, %c3 : index + %is_q4 = index.cmp eq, %weight_format, %c4 : index + %is_q5 = index.cmp eq, %weight_format, %c5 : index + %is_q6 = index.cmp eq, %weight_format, %c6 : index + %is_q4_0 = index.cmp eq, %weight_format, %q4_0_format : index + %is_q4_1 = index.cmp eq, %weight_format, %q4_1_format : index + %is_q5_0 = index.cmp eq, %weight_format, %q5_0_format : index + %is_q5_1 = index.cmp eq, %weight_format, %q5_1_format : index + %is_iq2_s = index.cmp eq, %weight_format, %c22 : index + %is_iq4_nl = index.cmp eq, %weight_format, %c20 : index + %is_iq3_s = index.cmp eq, %weight_format, %c21 : index + %is_iq4_xs = index.cmp eq, %weight_format, %c23 : index + %is_q8_0 = index.cmp eq, %weight_format, %q8_0_format : index + %is_q8_1 = index.cmp eq, %weight_format, %q8_1_format : index + %is_f16 = index.cmp eq, %weight_format, %c16 : index + %is_bf16 = index.cmp eq, %weight_format, %bf16_format : index + %is_f32 = index.cmp eq, %weight_format, %c32 : index + // Shared block index used by 32-value formats that pack eight groups per + // 256-value tile. + %q8_block_base = index.mul %quant_block, %c8 : index + %q8_block = index.add %q8_block_base, %quant_group : index + // Legacy 1-bit format. + %q1_0_values = scf.if %is_q1_0 -> (vector<4xf16>) { + %values = func.call @ggml_q1_0_f16_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + // K-quant family. + %iq2_xxs_values = scf.if %is_iq2_xxs -> (vector<4xf16>) { + %values = func.call @ggml_iq2xxs_f16_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %iq2_xs_values = scf.if %is_iq2_xs -> (vector<4xf16>) { + %values = func.call @ggml_iq2xs_f16_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %iq3_xxs_values = scf.if %is_iq3_xxs -> (vector<4xf16>) { + %values = func.call @ggml_iq3xxs_f16_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %iq1_s_values = scf.if %is_iq1_s -> (vector<4xf16>) { + %values = func.call @ggml_iq1s_f16_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %iq1_m_values = scf.if %is_iq1_m -> (vector<4xf16>) { + %values = func.call @ggml_iq1m_f16_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %q2_values = scf.if %is_q2 -> (vector<4xf16>) { + %values = func.call @ggml_q2k_f16_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %q3_values = scf.if %is_q3 -> (vector<4xf16>) { + %values = func.call @ggml_q3k_f16_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %q4_values = scf.if %is_q4 -> (vector<4xf16>) { + %values = func.call @ggml_q4k_f16_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %q6_values = scf.if %is_q6 -> (vector<4xf16>) { + %values = func.call @ggml_q6k_f16_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %q5_values = scf.if %is_q5 -> (vector<4xf16>) { + %values = func.call @ggml_q5k_f16_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + // Legacy 32-value block family. + %q4_0_values = scf.if %is_q4_0 -> (vector<4xf16>) { + %values = func.call @ggml_q4_0_f16_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %q4_1_values = scf.if %is_q4_1 -> (vector<4xf16>) { + %values = func.call @ggml_q4_1_f16_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %q5_0_values = scf.if %is_q5_0 -> (vector<4xf16>) { + %values = func.call @ggml_q5_0_f16_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %q5_1_values = scf.if %is_q5_1 -> (vector<4xf16>) { + %values = func.call @ggml_q5_1_f16_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + // IQ family. + %iq2_s_values = scf.if %is_iq2_s -> (vector<4xf16>) { + %values = func.call @ggml_iq2s_f16_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %mxfp4_format_v = index.constant 39 : index + %is_mxfp4_v = index.cmp eq, %weight_format, %mxfp4_format_v : index + %tq1_0_format_v = index.constant 34 : index + %is_tq1_0_v = index.cmp eq, %weight_format, %tq1_0_format_v : index + %tq1_0_values = scf.if %is_tq1_0_v -> (vector<4xf16>) { + %values = func.call @ggml_tq1_0_f16_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %tq2_0_format_v = index.constant 35 : index + %is_tq2_0_v = index.cmp eq, %weight_format, %tq2_0_format_v : index + %tq2_0_values = scf.if %is_tq2_0_v -> (vector<4xf16>) { + %values = func.call @ggml_tq2_0_f16_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %mxfp4_values = scf.if %is_mxfp4_v -> (vector<4xf16>) { + %values = func.call @ggml_mxfp4_f16_vector4(%weight, %row_byte_base, %q8_block, %packet) : (buffer, offset, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %iq4_nl_values = scf.if %is_iq4_nl -> (vector<4xf16>) { + %values = func.call @ggml_iq4nl_f16_vector4(%iq4nl_table, %weight, %row_byte_base, %q8_block, %packet) : (vector<16xi8>, buffer, offset, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %iq3_s_values = scf.if %is_iq3_s -> (vector<4xf16>) { + %values = func.call @ggml_iq3s_f16_vector4(%iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight, %row_byte_base, %quant_block, %quant_group, %packet) : (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %iq4_xs_values = scf.if %is_iq4_xs -> (vector<4xf16>) { + %values = func.call @ggml_iq4xs_f16_vector4(%iq4nl_table, %weight, %row_byte_base, %quant_block, %quant_group, %packet) : (vector<16xi8>, buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + // Q8 row formats. + %q8_0_values = scf.if %is_q8_0 -> (vector<4xf16>) { + %values = func.call @ggml_q8_0_f16_vector4(%weight, %row_byte_base, %q8_block, %packet) : (buffer, offset, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %q8_1_values = scf.if %is_q8_1 -> (vector<4xf16>) { + %values = func.call @ggml_q8_1_f16_vector4(%weight, %row_byte_base, %q8_block, %packet) : (buffer, offset, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + // Dense element formats. + %f16_values = scf.if %is_f16 -> (vector<4xf16>) { + %values = func.call @ggml_f16_f16_vector4(%weight, %row_byte_base, %input_size, %k) : (buffer, offset, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %bf16_values = scf.if %is_bf16 -> (vector<4xf16>) { + %values = func.call @ggml_bf16_f16_vector4(%weight, %row_byte_base, %input_size, %k) : (buffer, offset, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %f32_values = scf.if %is_f32 -> (vector<4xf16>) { + %values = func.call @ggml_f32_f16_vector4(%weight, %row_byte_base, %input_size, %k) : (buffer, offset, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + // Select the active decoded vector. Exactly one predicate should be true. + %selected_q1_0 = scf.select %is_q1_0, %q1_0_values, %zero : vector<4xf16> + %selected_q2 = scf.select %is_q2, %q2_values, %selected_q1_0 : vector<4xf16> + %selected_iq2_xxs = scf.select %is_iq2_xxs, %iq2_xxs_values, %selected_q2 : vector<4xf16> + %selected_iq2_xs = scf.select %is_iq2_xs, %iq2_xs_values, %selected_iq2_xxs : vector<4xf16> + %selected_iq3_xxs = scf.select %is_iq3_xxs, %iq3_xxs_values, %selected_iq2_xs : vector<4xf16> + %selected_iq1_s = scf.select %is_iq1_s, %iq1_s_values, %selected_iq3_xxs : vector<4xf16> + %selected_iq1_m = scf.select %is_iq1_m, %iq1_m_values, %selected_iq1_s : vector<4xf16> + %selected_q3 = scf.select %is_q3, %q3_values, %selected_iq1_m : vector<4xf16> + %selected_q4 = scf.select %is_q4, %q4_values, %selected_q3 : vector<4xf16> + %selected_q5 = scf.select %is_q5, %q5_values, %selected_q4 : vector<4xf16> + %selected_q6 = scf.select %is_q6, %q6_values, %selected_q5 : vector<4xf16> + %selected_q4_0 = scf.select %is_q4_0, %q4_0_values, %selected_q6 : vector<4xf16> + %selected_q4_1 = scf.select %is_q4_1, %q4_1_values, %selected_q4_0 : vector<4xf16> + %selected_q5_0 = scf.select %is_q5_0, %q5_0_values, %selected_q4_1 : vector<4xf16> + %selected_q5_1 = scf.select %is_q5_1, %q5_1_values, %selected_q5_0 : vector<4xf16> + %selected_iq2_s = scf.select %is_iq2_s, %iq2_s_values, %selected_q5_1 : vector<4xf16> + %selected_iq4_nl = scf.select %is_iq4_nl, %iq4_nl_values, %selected_iq2_s : vector<4xf16> + %selected_mxfp4_v = scf.select %is_mxfp4_v, %mxfp4_values, %selected_iq4_nl : vector<4xf16> + %selected_tq1_0_v = scf.select %is_tq1_0_v, %tq1_0_values, %selected_mxfp4_v : vector<4xf16> + %selected_tq2_0_v = scf.select %is_tq2_0_v, %tq2_0_values, %selected_tq1_0_v : vector<4xf16> + %selected_iq3_s = scf.select %is_iq3_s, %iq3_s_values, %selected_tq2_0_v : vector<4xf16> + %selected_iq4_xs = scf.select %is_iq4_xs, %iq4_xs_values, %selected_iq3_s : vector<4xf16> + %selected_q8_0 = scf.select %is_q8_0, %q8_0_values, %selected_iq4_xs : vector<4xf16> + %selected_q8_1 = scf.select %is_q8_1, %q8_1_values, %selected_q8_0 : vector<4xf16> + %selected_f16 = scf.select %is_f16, %f16_values, %selected_q8_1 : vector<4xf16> + %selected_bf16 = scf.select %is_bf16, %bf16_values, %selected_f16 : vector<4xf16> + %selected0 = scf.select %is_f32, %f32_values, %selected_bf16 : vector<4xf16> + %prism_pq2_format = index.constant 72 : index + %prism_ptq1_format = index.constant 73 : index + %is_prism_pq2 = index.cmp eq, %weight_format, %prism_pq2_format : index + %is_prism_ptq1 = index.cmp eq, %weight_format, %prism_ptq1_format : index + %prism_pq2_values = scf.if %is_prism_pq2 -> (vector<4xf16>) { + %values = func.call @ggml_pq2_0_f16_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %prism_ptq1_values = scf.if %is_prism_ptq1 -> (vector<4xf16>) { + %values = func.call @ggml_ptq1_0_f16_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } + %selected1 = scf.select %is_prism_pq2, %prism_pq2_values, %selected0 : vector<4xf16> + %selected = scf.select %is_prism_ptq1, %prism_ptq1_values, %selected1 : vector<4xf16> + func.return %selected : vector<4xf16> +} + +// Centralizes row-format decoding for F32 row-gather outputs. +func.def inline @ggml_dequant_f32_vector4(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight_format: index, %weight: buffer, %row_byte_base: offset, %input_size: index, %quant_block: index, %quant_group: index, %packet: index, %k: index) -> (vector<4xf32>) { + %q1_0_format = index.constant 10 : index + %c3 = index.constant 11 : index + %c4 = index.constant 4 : index + %c5 = index.constant 5 : index + %c6 = index.constant 6 : index + %c8 = index.constant 8 : index + %c20 = index.constant 20 : index + %c21 = index.constant 21 : index + %c22 = index.constant 22 : index + %c23 = index.constant 23 : index + %c16 = index.constant 16 : index + %bf16_format = index.constant 30 : index + %c32 = index.constant 32 : index + %q4_0_format = index.constant 40 : index + %q4_1_format = index.constant 41 : index + %q5_0_format = index.constant 50 : index + %q5_1_format = index.constant 51 : index + %q8_0_format = index.constant 80 : index + %q8_1_format = index.constant 81 : index + %zero = vector.constant 0.0 : vector<4xf32> + %is_q1_0 = index.cmp eq, %weight_format, %q1_0_format : index + %q2_format = index.constant 12 : index + %iq2_xxs_format = index.constant 24 : index + %iq2_xs_format = index.constant 25 : index + %is_iq2_xxs = index.cmp eq, %weight_format, %iq2_xxs_format : index + %is_iq2_xs = index.cmp eq, %weight_format, %iq2_xs_format : index + %iq3_xxs_format = index.constant 28 : index + %is_iq3_xxs = index.cmp eq, %weight_format, %iq3_xxs_format : index + %iq1_s_format = index.constant 26 : index + %iq1_m_format = index.constant 27 : index + %is_iq1_s = index.cmp eq, %weight_format, %iq1_s_format : index + %is_iq1_m = index.cmp eq, %weight_format, %iq1_m_format : index + %is_q2 = index.cmp eq, %weight_format, %q2_format : index + %is_q3 = index.cmp eq, %weight_format, %c3 : index + %is_q4 = index.cmp eq, %weight_format, %c4 : index + %is_q5 = index.cmp eq, %weight_format, %c5 : index + %is_q6 = index.cmp eq, %weight_format, %c6 : index + %is_q4_0 = index.cmp eq, %weight_format, %q4_0_format : index + %is_q4_1 = index.cmp eq, %weight_format, %q4_1_format : index + %is_q5_0 = index.cmp eq, %weight_format, %q5_0_format : index + %is_q5_1 = index.cmp eq, %weight_format, %q5_1_format : index + %is_iq2_s = index.cmp eq, %weight_format, %c22 : index + %is_iq4_nl = index.cmp eq, %weight_format, %c20 : index + %is_iq3_s = index.cmp eq, %weight_format, %c21 : index + %is_iq4_xs = index.cmp eq, %weight_format, %c23 : index + %is_q8_0 = index.cmp eq, %weight_format, %q8_0_format : index + %is_q8_1 = index.cmp eq, %weight_format, %q8_1_format : index + %is_f16 = index.cmp eq, %weight_format, %c16 : index + %is_bf16 = index.cmp eq, %weight_format, %bf16_format : index + %is_f32 = index.cmp eq, %weight_format, %c32 : index + %q8_block_base = index.mul %quant_block, %c8 : index + %q8_block = index.add %q8_block_base, %quant_group : index + %q1_0_values = scf.if %is_q1_0 -> (vector<4xf32>) { + %values = func.call @ggml_q1_0_f32_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %iq2_xxs_values = scf.if %is_iq2_xxs -> (vector<4xf32>) { + %values = func.call @ggml_iq2xxs_f32_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %iq2_xs_values = scf.if %is_iq2_xs -> (vector<4xf32>) { + %values = func.call @ggml_iq2xs_f32_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %iq3_xxs_values = scf.if %is_iq3_xxs -> (vector<4xf32>) { + %values = func.call @ggml_iq3xxs_f32_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %iq1_s_values = scf.if %is_iq1_s -> (vector<4xf32>) { + %values = func.call @ggml_iq1s_f32_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %iq1_m_values = scf.if %is_iq1_m -> (vector<4xf32>) { + %values = func.call @ggml_iq1m_f32_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %q2_values = scf.if %is_q2 -> (vector<4xf32>) { + %values = func.call @ggml_q2k_f32_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %q3_values = scf.if %is_q3 -> (vector<4xf32>) { + %values = func.call @ggml_q3k_f32_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %q4_values = scf.if %is_q4 -> (vector<4xf32>) { + %values = func.call @ggml_q4k_f32_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %q6_values = scf.if %is_q6 -> (vector<4xf32>) { + %values = func.call @ggml_q6k_f32_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %q5_values = scf.if %is_q5 -> (vector<4xf32>) { + %values = func.call @ggml_q5k_f32_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %q4_0_values = scf.if %is_q4_0 -> (vector<4xf32>) { + %values = func.call @ggml_q4_0_f32_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %q4_1_values = scf.if %is_q4_1 -> (vector<4xf32>) { + %values = func.call @ggml_q4_1_f32_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %q5_0_values = scf.if %is_q5_0 -> (vector<4xf32>) { + %values = func.call @ggml_q5_0_f32_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %q5_1_values = scf.if %is_q5_1 -> (vector<4xf32>) { + %values = func.call @ggml_q5_1_f32_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %iq2_s_values = scf.if %is_iq2_s -> (vector<4xf32>) { + %values = func.call @ggml_iq2s_f32_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %mxfp4_format_v = index.constant 39 : index + %is_mxfp4_v = index.cmp eq, %weight_format, %mxfp4_format_v : index + %tq1_0_format_v = index.constant 34 : index + %is_tq1_0_v = index.cmp eq, %weight_format, %tq1_0_format_v : index + %tq1_0_values = scf.if %is_tq1_0_v -> (vector<4xf32>) { + %values = func.call @ggml_tq1_0_f32_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %tq2_0_format_v = index.constant 35 : index + %is_tq2_0_v = index.cmp eq, %weight_format, %tq2_0_format_v : index + %tq2_0_values = scf.if %is_tq2_0_v -> (vector<4xf32>) { + %values = func.call @ggml_tq2_0_f32_vector4(%weight, %row_byte_base, %quant_block, %quant_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %mxfp4_values = scf.if %is_mxfp4_v -> (vector<4xf32>) { + %values = func.call @ggml_mxfp4_f32_vector4(%weight, %row_byte_base, %q8_block, %packet) : (buffer, offset, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %iq4_nl_values = scf.if %is_iq4_nl -> (vector<4xf32>) { + %values = func.call @ggml_iq4nl_f32_vector4(%iq4nl_table, %weight, %row_byte_base, %q8_block, %packet) : (vector<16xi8>, buffer, offset, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %iq3_s_values = scf.if %is_iq3_s -> (vector<4xf32>) { + %values = func.call @ggml_iq3s_f32_vector4(%iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight, %row_byte_base, %quant_block, %quant_group, %packet) : (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, buffer, offset, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %iq4_xs_values = scf.if %is_iq4_xs -> (vector<4xf32>) { + %values = func.call @ggml_iq4xs_f32_vector4(%iq4nl_table, %weight, %row_byte_base, %quant_block, %quant_group, %packet) : (vector<16xi8>, buffer, offset, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %q8_0_values = scf.if %is_q8_0 -> (vector<4xf32>) { + %values = func.call @ggml_q8_0_f32_vector4(%weight, %row_byte_base, %q8_block, %packet) : (buffer, offset, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %q8_1_values = scf.if %is_q8_1 -> (vector<4xf32>) { + %values = func.call @ggml_q8_1_f32_vector4(%weight, %row_byte_base, %q8_block, %packet) : (buffer, offset, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %f16_values = scf.if %is_f16 -> (vector<4xf32>) { + %values = func.call @ggml_f16_f32_vector4(%weight, %row_byte_base, %input_size, %k) : (buffer, offset, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %bf16_values = scf.if %is_bf16 -> (vector<4xf32>) { + %values = func.call @ggml_bf16_f32_vector4(%weight, %row_byte_base, %input_size, %k) : (buffer, offset, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %f32_values = scf.if %is_f32 -> (vector<4xf32>) { + %values = func.call @ggml_f32_f32_vector4(%weight, %row_byte_base, %input_size, %k) : (buffer, offset, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %selected_q1_0 = scf.select %is_q1_0, %q1_0_values, %zero : vector<4xf32> + %selected_q2 = scf.select %is_q2, %q2_values, %selected_q1_0 : vector<4xf32> + %selected_iq2_xxs = scf.select %is_iq2_xxs, %iq2_xxs_values, %selected_q2 : vector<4xf32> + %selected_iq2_xs = scf.select %is_iq2_xs, %iq2_xs_values, %selected_iq2_xxs : vector<4xf32> + %selected_iq3_xxs = scf.select %is_iq3_xxs, %iq3_xxs_values, %selected_iq2_xs : vector<4xf32> + %selected_iq1_s = scf.select %is_iq1_s, %iq1_s_values, %selected_iq3_xxs : vector<4xf32> + %selected_iq1_m = scf.select %is_iq1_m, %iq1_m_values, %selected_iq1_s : vector<4xf32> + %selected_q3 = scf.select %is_q3, %q3_values, %selected_iq1_m : vector<4xf32> + %selected_q4 = scf.select %is_q4, %q4_values, %selected_q3 : vector<4xf32> + %selected_q5 = scf.select %is_q5, %q5_values, %selected_q4 : vector<4xf32> + %selected_q6 = scf.select %is_q6, %q6_values, %selected_q5 : vector<4xf32> + %selected_q4_0 = scf.select %is_q4_0, %q4_0_values, %selected_q6 : vector<4xf32> + %selected_q4_1 = scf.select %is_q4_1, %q4_1_values, %selected_q4_0 : vector<4xf32> + %selected_q5_0 = scf.select %is_q5_0, %q5_0_values, %selected_q4_1 : vector<4xf32> + %selected_q5_1 = scf.select %is_q5_1, %q5_1_values, %selected_q5_0 : vector<4xf32> + %selected_iq2_s = scf.select %is_iq2_s, %iq2_s_values, %selected_q5_1 : vector<4xf32> + %selected_iq4_nl = scf.select %is_iq4_nl, %iq4_nl_values, %selected_iq2_s : vector<4xf32> + %selected_mxfp4_v = scf.select %is_mxfp4_v, %mxfp4_values, %selected_iq4_nl : vector<4xf32> + %selected_tq1_0_v = scf.select %is_tq1_0_v, %tq1_0_values, %selected_mxfp4_v : vector<4xf32> + %selected_tq2_0_v = scf.select %is_tq2_0_v, %tq2_0_values, %selected_tq1_0_v : vector<4xf32> + %selected_iq3_s = scf.select %is_iq3_s, %iq3_s_values, %selected_tq2_0_v : vector<4xf32> + %selected_iq4_xs = scf.select %is_iq4_xs, %iq4_xs_values, %selected_iq3_s : vector<4xf32> + %selected_q8_0 = scf.select %is_q8_0, %q8_0_values, %selected_iq4_xs : vector<4xf32> + %selected_q8_1 = scf.select %is_q8_1, %q8_1_values, %selected_q8_0 : vector<4xf32> + %selected_f16 = scf.select %is_f16, %f16_values, %selected_q8_1 : vector<4xf32> + %selected_bf16 = scf.select %is_bf16, %bf16_values, %selected_f16 : vector<4xf32> + %selected0 = scf.select %is_f32, %f32_values, %selected_bf16 : vector<4xf32> + %prism_pq2_format = index.constant 72 : index + %prism_ptq1_format = index.constant 73 : index + %is_prism_pq2 = index.cmp eq, %weight_format, %prism_pq2_format : index + %is_prism_ptq1 = index.cmp eq, %weight_format, %prism_ptq1_format : index + %prism_pq2_values = scf.if %is_prism_pq2 -> (vector<4xf32>) { + %values = func.call @ggml_pq2_0_f32_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %prism_ptq1_values = scf.if %is_prism_ptq1 -> (vector<4xf32>) { + %values = func.call @ggml_ptq1_0_f32_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %selected1 = scf.select %is_prism_pq2, %prism_pq2_values, %selected0 : vector<4xf32> + %selected = scf.select %is_prism_ptq1, %prism_ptq1_values, %selected1 : vector<4xf32> + func.return %selected : vector<4xf32> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/dequant_1bit.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/dequant_1bit.loom new file mode 100644 index 000000000000..5b8744329a53 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/dequant_1bit.loom @@ -0,0 +1,300 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// Weight decoders for the shared dequantizer (motifs/dequant.loom calls these from its format switch) that are not +// part of the upstream corpus. +// +// MXFP4 (format 39, gpt-oss): 17 bytes per 32 values = e (E8M0 shared exponent), qs[16]; value j is the low nibble +// of qs[j] and value j + 16 the high nibble, as in IQ4_NL. A value is kvalues_mxfp4[code] * 2^(e - 127) / 2 +// (ggml GGML_E8M0_TO_FP32_HALF), with the E2M1 table {0, 1, 2, 3, 4, 6, 8, 12, 0, -1, -2, -3, -4, -6, -8, -12}. +// %packet (0..7) picks four consecutive values: 4 (packet % 4)..+4, from the high nibbles when packet >= 4. +// +// TQ1_0 (format 34) and TQ2_0 (format 35): ggml ternary types, 256 values per block; see their decoders below. + +// Nibble unpacking: four codes from four bytes, low nibbles or (uses_high) high nibbles. +func.def inline @ggml_1bit_nibble_codes4(%q_bytes: vector<4xi8>, %uses_high: i1) -> (vector<4xi8>) { + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c12_i32 = scalar.constant 12 : i32 + %c16_i32 = scalar.constant 16 : i32 + %mask0_i32 = scalar.constant 15 : i32 + %mask1_i32 = scalar.constant 240 : i32 + %mask2_i32 = scalar.constant 3840 : i32 + %mask3_i32 = scalar.constant 61440 : i32 + %q_word_vector = vector.bitcast %q_bytes : vector<4xi8> to vector<1xi32> + %word = vector.extract %q_word_vector[0] : vector<1xi32> -> i32 + %word_shr4 = scalar.shrui %word, %c4_i32 : i32 + %word_shr8 = scalar.shrui %word, %c8_i32 : i32 + %word_shr12 = scalar.shrui %word, %c12_i32 : i32 + %word_shr16 = scalar.shrui %word, %c16_i32 : i32 + %low0 = scalar.andi %word, %mask0_i32 : i32 + %low1 = scalar.andi %word_shr4, %mask1_i32 : i32 + %low2 = scalar.andi %word_shr8, %mask2_i32 : i32 + %low3 = scalar.andi %word_shr12, %mask3_i32 : i32 + %low01 = scalar.ori %low0, %low1 : i32 + %low23 = scalar.ori %low2, %low3 : i32 + %low = scalar.ori %low01, %low23 : i32 + %high0 = scalar.andi %word_shr4, %mask0_i32 : i32 + %high1 = scalar.andi %word_shr8, %mask1_i32 : i32 + %high2 = scalar.andi %word_shr12, %mask2_i32 : i32 + %high3 = scalar.andi %word_shr16, %mask3_i32 : i32 + %high01 = scalar.ori %high0, %high1 : i32 + %high23 = scalar.ori %high2, %high3 : i32 + %high = scalar.ori %high01, %high23 : i32 + %selected = scf.select %uses_high, %high, %low : i32 + %target_selected_shr8 = scalar.shrui %selected, %c8_i32 : i32 + %packed0 = scalar.trunci %selected : i32 to i8 + %packed1 = scalar.trunci %target_selected_shr8 : i32 to i8 + %packed = vector.from_elements %packed0, %packed1 : vector<2xi8> + %codes = vector.bitunpacku<4> %packed : vector<2xi8> -> vector<4xi8> + func.return %codes : vector<4xi8> +} + +func.def inline @ggml_mxfp4_table_i8() -> (vector<16xi8>) { + %v0 = scalar.constant 0 : i8 + %v1 = scalar.constant 1 : i8 + %v2 = scalar.constant 2 : i8 + %v3 = scalar.constant 3 : i8 + %v4 = scalar.constant 4 : i8 + %v5 = scalar.constant 6 : i8 + %v6 = scalar.constant 8 : i8 + %v7 = scalar.constant 12 : i8 + %v8 = scalar.constant 0 : i8 + %v9 = scalar.constant -1 : i8 + %v10 = scalar.constant -2 : i8 + %v11 = scalar.constant -3 : i8 + %v12 = scalar.constant -4 : i8 + %v13 = scalar.constant -6 : i8 + %v14 = scalar.constant -8 : i8 + %v15 = scalar.constant -12 : i8 + %table = vector.from_elements %v0, %v1, %v2, %v3, %v4, %v5, %v6, %v7, %v8, %v9, %v10, %v11, %v12, %v13, %v14, %v15 : vector<16xi8> + func.return %table : vector<16xi8> +} + +// 2^(e - 127) / 2, exact, built from its f32 bits as ggml's GGML_E8M0_TO_FP32_HALF: (e - 1) << 23 for e >= 2, +// the subnormals 0x00200000 << e below that. +func.def inline @ggml_mxfp4_half_scale(%e: i32) -> (f32) { + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c23_i32 = scalar.constant 23 : i32 + %sub_unit = scalar.constant 2097152 : i32 + %is_small = scalar.cmpi ult, %e, %c2_i32 : i32 + %e_m1 = scalar.subi %e, %c1_i32 : i32 + %normal_bits = scalar.shli %e_m1, %c23_i32 : i32 + %sub_bits = scalar.shli %sub_unit, %e : i32 + %bits = scf.select %is_small, %sub_bits, %normal_bits : i32 + %bits_v = vector.from_elements %bits : vector<1xi32> + %scale_v = vector.bitcast %bits_v : vector<1xi32> to vector<1xf32> + %scale = vector.extract %scale_v[0] : vector<1xf32> -> f32 + func.return %scale : f32 +} + +func.def inline @ggml_mxfp4_f32_vector4(%weight: buffer, %row_byte_base: offset, %mx_block: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %block_bytes = index.constant 17 : offset + %code_offset = index.constant 1 : offset + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %mx_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %e_view = buffer.view %weight[%block_byte_base] : buffer -> view<1xi8> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<16xi8> + %packet_page = index.rem %bounded_packet, %c4 : index + %code_index0 = index.mul %packet_page, %c4 : index + %code_index = index.assume %code_index0 [range(%code_index0, 0, 12), mul(%code_index0, 4)] : index + %uses_high = index.cmp uge, %bounded_packet, %c4 : index + %q_bytes = vector.load %code_view[%code_index] : view<16xi8> -> vector<4xi8> + %codes = func.call @ggml_1bit_nibble_codes4(%q_bytes, %uses_high) : (vector<4xi8>, i1) -> (vector<4xi8>) + %table = func.call @ggml_mxfp4_table_i8() : () -> (vector<16xi8>) + %lookup_i8 = vector.table.lookup %table[%codes] : vector<16xi8>, vector<4xi8> -> vector<4xi8> + %lookup = vector.sitofp %lookup_i8 : vector<4xi8> to vector<4xf32> + %e_i8 = view.load %e_view[%c0] : view<1xi8> -> i8 + %e = scalar.extui %e_i8 : i8 to i32 + %scale = func.call @ggml_mxfp4_half_scale(%e) : (i32) -> (f32) + %scale_v = vector.splat %scale : vector<4xf32> + %result = vector.mulf %scale_v, %lookup : vector<4xf32> + func.return %result : vector<4xf32> +} + +func.def inline @ggml_mxfp4_f16_vector4(%weight: buffer, %row_byte_base: offset, %mx_block: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_mxfp4_f32_vector4(%weight, %row_byte_base, %mx_block, %packet) : (buffer, offset, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + + +// TQ2_0 (format 35): 66 bytes per 256 values = qs[64], d (f16). Value v = 128 a + 32 l + m is +// ((qs[32 a + m] >> 2 l) & 3) - 1, times d (ggml dequantize_row_tq2_0). %tq_group (0..7) = v / 32 and %packet (0..7) +// = (v % 32) / 4: four consecutive values are four consecutive bytes, one shift. +func.def inline @ggml_tq2_0_f32_vector4(%weight: buffer, %row_byte_base: offset, %tq_block: index, %tq_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %c32 = index.constant 32 : index + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c1_i32 = scalar.constant 1 : i32 + %block_bytes = index.constant 66 : offset + %d_offset = index.constant 64 : offset + %g = index.assume %tq_group [range(%tq_group, 0, 7)] : index + %p = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %tq_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %d_base = index.add %block_byte_base, %d_offset : offset + %d_view = buffer.view %weight[%d_base] : buffer -> view<1xf16> + %qv = buffer.view %weight[%block_byte_base] : buffer -> view<64xi8> + %a = index.div %g, %c4 : index + %l = index.rem %g, %c4 : index + %a32 = index.mul %a, %c32 : index + %p4 = index.mul %p, %c4 : index + %byte0 = index.add %a32, %p4 : index + %byte = index.assume %byte0 [range(%byte0, 0, 60), mul(%byte0, 4)] : index + %bytes = vector.load %qv[%byte] : view<64xi8> -> vector<4xi8> + %word_v = vector.bitcast %bytes : vector<4xi8> to vector<1xi32> + %word = vector.extract %word_v[0] : vector<1xi32> -> i32 + %bsh0 = scalar.constant 0 : i32 + %bsh8 = scalar.constant 8 : i32 + %bsh16 = scalar.constant 16 : i32 + %bsh24 = scalar.constant 24 : i32 + %byte_shift = vector.from_elements %bsh0, %bsh8, %bsh16, %bsh24 : vector<4xi32> + %byte_mask = scalar.constant 255 : i32 + %byte_mask_v = vector.splat %byte_mask : vector<4xi32> + %word_s = vector.splat %word : vector<4xi32> + %word_sh = vector.shrui %word_s, %byte_shift : vector<4xi32> + %b = vector.andi %word_sh, %byte_mask_v : vector<4xi32> + %l_i32 = index.cast %l : index to i32 + %shift = scalar.muli %l_i32, %c2_i32 : i32 + %shift_v = vector.splat %shift : vector<4xi32> + %three = vector.splat %c3_i32 : vector<4xi32> + %one = vector.splat %c1_i32 : vector<4xi32> + %bs = vector.shrui %b, %shift_v : vector<4xi32> + %q = vector.andi %bs, %three : vector<4xi32> + %t = vector.subi %q, %one : vector<4xi32> + %tf = vector.sitofp %t : vector<4xi32> to vector<4xf32> + %d_f16 = view.load %d_view[%c0] : view<1xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %dv = vector.splat %d : vector<4xf32> + %r = vector.mulf %tf, %dv : vector<4xf32> + func.return %r : vector<4xf32> +} + +// TQ1_0 (format 34): 54 bytes per 256 values = qs[48], qh[4], d (f16), trits packed 5 per byte (4 in qh) in base 3: +// the n-th trit of byte b is ((b * 3^n mod 256) * 3) >> 8, minus 1 (ggml dequantize_row_tq1_0). Value order: +// 0..159 = trit n of qs[0..31] (n = v / 32), 160..239 = trit n of qs[32..47] (n = (v - 160) / 16), 240..255 = trit n +// of qh[0..3] (n = (v - 240) / 4). Every packet of four values is four consecutive bytes and one power. +func.def inline @ggml_tq1_0_f32_vector4(%weight: buffer, %row_byte_base: offset, %tq_block: index, %tq_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %c5 = index.constant 5 : index + %c7 = index.constant 7 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c48 = index.constant 48 : index + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c255_i32 = scalar.constant 255 : i32 + %block_bytes = index.constant 54 : offset + %d_offset = index.constant 52 : offset + %g = index.assume %tq_group [range(%tq_group, 0, 7)] : index + %p = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %tq_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %d_base = index.add %block_byte_base, %d_offset : offset + %d_view = buffer.view %weight[%d_base] : buffer -> view<1xf16> + %qv = buffer.view %weight[%block_byte_base] : buffer -> view<52xi8> + %p4 = index.mul %p, %c4 : index + // groups 0..4: qs[4p..], power g + %first = index.cmp ult, %g, %c5 : index + // groups 5, 6: k = 32 (g - 5) + 4p in 0..60, power k / 16, byte 32 + k % 16 + %g_mid = scf.select %first, %c5, %g : index + %g5 = index.sub %g_mid, %c5 : index + %g5_32 = index.mul %g5, %c32 : index + %k = index.add %g5_32, %p4 : index + %mid_n = index.div %k, %c16 : index + %k16 = index.rem %k, %c16 : index + %mid_byte = index.add %c32, %k16 : index + %middle = index.cmp ult, %g, %c7 : index + // group 7: packets 0..3 = qs[32 + 4p..], power 4; packets 4..7 = qh[0..3], power p - 4 + %low_half = index.cmp ult, %p, %c4 : index + %last_byte_lo = index.add %c32, %p4 : index + %p_hi = scf.select %low_half, %c4, %p : index + %last_n_hi = index.sub %p_hi, %c4 : index + %last_byte = scf.select %low_half, %last_byte_lo, %c48 : index + %last_n = scf.select %low_half, %c4, %last_n_hi : index + %rest_byte = scf.select %middle, %mid_byte, %last_byte : index + %rest_n = scf.select %middle, %mid_n, %last_n : index + %byte0 = scf.select %first, %p4, %rest_byte : index + %n = scf.select %first, %g, %rest_n : index + %byte = index.assume %byte0 [range(%byte0, 0, 48)] : index + %bytes = vector.load %qv[%byte] : view<52xi8> -> vector<4xi8> + %word_v = vector.bitcast %bytes : vector<4xi8> to vector<1xi32> + %word = vector.extract %word_v[0] : vector<1xi32> -> i32 + %bsh0 = scalar.constant 0 : i32 + %bsh8 = scalar.constant 8 : i32 + %bsh16 = scalar.constant 16 : i32 + %bsh24 = scalar.constant 24 : i32 + %byte_shift = vector.from_elements %bsh0, %bsh8, %bsh16, %bsh24 : vector<4xi32> + %byte_mask = scalar.constant 255 : i32 + %byte_mask_v = vector.splat %byte_mask : vector<4xi32> + %word_s = vector.splat %word : vector<4xi32> + %word_sh = vector.shrui %word_s, %byte_shift : vector<4xi32> + %b = vector.andi %word_sh, %byte_mask_v : vector<4xi32> + // 3^n for n in 0..4 + %n_i32 = index.cast %n : index to i32 + %pw1 = scalar.constant 1 : i32 + %pw3 = scalar.constant 3 : i32 + %pw9 = scalar.constant 9 : i32 + %pw27 = scalar.constant 27 : i32 + %pw81 = scalar.constant 81 : i32 + %is1 = scalar.cmpi eq, %n_i32, %c1_i32 : i32 + %c2_n = scalar.constant 2 : i32 + %is2 = scalar.cmpi eq, %n_i32, %c2_n : i32 + %is3 = scalar.cmpi eq, %n_i32, %c3_i32 : i32 + %c4_n = scalar.constant 4 : i32 + %is4 = scalar.cmpi eq, %n_i32, %c4_n : i32 + %s1 = scf.select %is1, %pw3, %pw1 : i32 + %s2 = scf.select %is2, %pw9, %s1 : i32 + %s3 = scf.select %is3, %pw27, %s2 : i32 + %pow = scf.select %is4, %pw81, %s3 : i32 + %pow_v = vector.splat %pow : vector<4xi32> + %mask = vector.splat %c255_i32 : vector<4xi32> + %three = vector.splat %c3_i32 : vector<4xi32> + %eight = vector.splat %c8_i32 : vector<4xi32> + %one = vector.splat %c1_i32 : vector<4xi32> + %bp = vector.muli %b, %pow_v : vector<4xi32> + %q = vector.andi %bp, %mask : vector<4xi32> + %q3 = vector.muli %q, %three : vector<4xi32> + %xi = vector.shrui %q3, %eight : vector<4xi32> + %t = vector.subi %xi, %one : vector<4xi32> + %tf = vector.sitofp %t : vector<4xi32> to vector<4xf32> + %d_f16 = view.load %d_view[%c0] : view<1xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %dv = vector.splat %d : vector<4xf32> + %r = vector.mulf %tf, %dv : vector<4xf32> + func.return %r : vector<4xf32> +} + +func.def inline @ggml_tq1_0_f16_vector4(%weight: buffer, %row_byte_base: offset, %tq_block: index, %tq_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_tq1_0_f32_vector4(%weight, %row_byte_base, %tq_block, %tq_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +func.def inline @ggml_tq2_0_f16_vector4(%weight: buffer, %row_byte_base: offset, %tq_block: index, %tq_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_tq2_0_f32_vector4(%weight, %row_byte_base, %tq_block, %tq_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/dequant_prism.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/dequant_prism.loom new file mode 100644 index 000000000000..c31111a2969c --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/dequant_prism.loom @@ -0,0 +1,198 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// Decoders for PrismML's group-128 weight types (ggml-prism.h) in the shared dequantizer: motifs/dequant.loom +// calls these from its format switch (prompt WMMA staging, the generic decode GEMV and GET_ROWS). Each returns +// the four values k..k+3 (k a multiple of 4) of a row starting at %row_byte_base, like the Q1_0 decoder. +// +// PQ2_0 (format 72): 34 bytes per 128 values = d (fp16), qs[32]; value j is code (qs[j / 4] >> 2 (j % 4)) & 3, +// w = (code - 1) * d. +// PTQ1_0 (format 73): 28 bytes per 128 values = qs[24], qh[2], d (fp16 at byte 26); trits t = -1, 0, +1, w = t * d. +// A byte b holds trit n as ((uint8_t)(b * 3^n) * 3) >> 8 (minus 1). Values 0..79: byte k % 16, trit k / 16; +// 80..119: byte 16 + (k - 80) % 8, trit (k - 80) / 8; 120..127: byte 24 + (k - 120) % 2, trit (k - 120) / 2. +// Blocks are 28 bytes, so with 4-byte aligned rows (28 * K / 128) each group of four values is one 32-bit +// word: qs words 0..3 for the first range, 4..5 for the second; the last eight values take qh from word 6 +// (bytes qh0, qh1, qh0, qh1 with trits n, n, n + 1, n + 1). + +func.def inline @ggml_pq2_0_f32_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %c128 = index.constant 128 : index + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c6_i32 = scalar.constant 6 : i32 + %block_bytes = index.constant 34 : offset + %code_offset = index.constant 2 : offset + %block = index.div %k, %c128 : index + %k_in_block = index.rem %k, %c128 : index + %byte_index0 = index.div %k_in_block, %c4 : index + %byte_index = index.assume %byte_index0 [range(%byte_index0, 0, 31)] : index + %block_byte_add = index.scale %block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %d_view = buffer.view %weight[%block_byte_base] : buffer -> view<1xf16> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<32xi8> + %d_vector = vector.load %d_view[%c0] : view<1xf16> -> vector<1xf16> + %d_f16 = vector.extract %d_vector[0] : vector<1xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %q_i8 = view.load %code_view[%byte_index] : view<32xi8> -> i8 + %q = scalar.extui %q_i8 : i8 to i32 + %q1_shift = scalar.shrui %q, %c2_i32 : i32 + %q2_shift = scalar.shrui %q, %c4_i32 : i32 + %q3_shift = scalar.shrui %q, %c6_i32 : i32 + %code0 = scalar.andi %q, %c3_i32 : i32 + %code1 = scalar.andi %q1_shift, %c3_i32 : i32 + %code2 = scalar.andi %q2_shift, %c3_i32 : i32 + %code3 = scalar.andi %q3_shift, %c3_i32 : i32 + %t0 = scalar.subi %code0, %c1_i32 : i32 + %t1 = scalar.subi %code1, %c1_i32 : i32 + %t2 = scalar.subi %code2, %c1_i32 : i32 + %t3 = scalar.subi %code3, %c1_i32 : i32 + %t0_f32 = scalar.sitofp %t0 : i32 to f32 + %t1_f32 = scalar.sitofp %t1 : i32 to f32 + %t2_f32 = scalar.sitofp %t2 : i32 to f32 + %t3_f32 = scalar.sitofp %t3 : i32 to f32 + %v0 = scalar.mulf %t0_f32, %d : f32 + %v1 = scalar.mulf %t1_f32, %d : f32 + %v2 = scalar.mulf %t2_f32, %d : f32 + %v3 = scalar.mulf %t3_f32, %d : f32 + %result = vector.from_elements %v0, %v1, %v2, %v3 : vector<4xf32> + func.return %result : vector<4xf32> +} + +func.def inline @ggml_pq2_0_f16_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_pq2_0_f32_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +// 3^n for n = 0..4. +func.def inline @ggml_ptq1_0_pow3(%n: index) -> (i32) { + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %p1 = scalar.constant 1 : i32 + %p3 = scalar.constant 3 : i32 + %p9 = scalar.constant 9 : i32 + %p27 = scalar.constant 27 : i32 + %p81 = scalar.constant 81 : i32 + %is1 = index.cmp eq, %n, %c1 : index + %is2 = index.cmp eq, %n, %c2 : index + %is3 = index.cmp eq, %n, %c3 : index + %is4 = index.cmp eq, %n, %c4 : index + %s1 = scf.select %is1, %p3, %p1 : i32 + %s2 = scf.select %is2, %p9, %s1 : i32 + %s3 = scf.select %is3, %p27, %s2 : i32 + %s4 = scf.select %is4, %p81, %s3 : i32 + func.return %s4 : i32 +} + +// Trit (-1, 0, +1) n of byte b: ((uint8_t)(b * 3^n) * 3) >> 8, minus 1, as f32. +func.def inline @ggml_ptq1_0_trit_f32(%byte: i32, %pow: i32) -> (f32) { + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c255_i32 = scalar.constant 255 : i32 + %m = scalar.muli %byte, %pow : i32 + %q = scalar.andi %m, %c255_i32 : i32 + %q3 = scalar.muli %q, %c3_i32 : i32 + %u = scalar.shrui %q3, %c8_i32 : i32 + %t = scalar.subi %u, %c1_i32 : i32 + %t_f32 = scalar.sitofp %t : i32 to f32 + func.return %t_f32 : f32 +} + +func.def inline @ggml_ptq1_0_f32_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c6 = index.constant 6 : index + %c8 = index.constant 8 : index + %c13 = index.constant 13 : index + %c16 = index.constant 16 : index + %c80 = index.constant 80 : index + %c120 = index.constant 120 : index + %c128 = index.constant 128 : index + %c0_i32 = scalar.constant 0 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c16_i32 = scalar.constant 16 : i32 + %c24_i32 = scalar.constant 24 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c255_i32 = scalar.constant 255 : i32 + %block_bytes = index.constant 28 : offset + %block = index.div %k, %c128 : index + %kb = index.rem %k, %c128 : index + %block_byte_add = index.scale %block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %word_view = buffer.view %weight[%block_byte_base] : buffer -> view<7xi32> + %half_view = buffer.view %weight[%block_byte_base] : buffer -> view<14xf16> + %d_f16 = view.load %half_view[%c13] : view<14xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + // Range of the four values: qs[0..15] (a), qs[16..23] (b) or qh (c). + %in_a = index.cmp ult, %kb, %c80 : index + %in_c = index.cmp uge, %kb, %c120 : index + %kb_b0 = scf.select %in_a, %c80, %kb : index + %kb_b = index.sub %kb_b0, %c80 : index + %kb_c0 = scf.select %in_c, %kb, %c120 : index + %kb_c = index.sub %kb_c0, %c120 : index + %n_a = index.div %kb, %c16 : index + %m_a = index.rem %kb, %c16 : index + %w_a = index.div %m_a, %c4 : index + %n_b = index.div %kb_b, %c8 : index + %m_b = index.rem %kb_b, %c8 : index + %w_b0 = index.div %m_b, %c4 : index + %w_b = index.add %w_b0, %c4 : index + %n_c = index.div %kb_c, %c2 : index + %n_ab = scf.select %in_a, %n_a, %n_b : index + %n = scf.select %in_c, %n_c, %n_ab : index + %w_ab = scf.select %in_a, %w_a, %w_b : index + %w0 = scf.select %in_c, %c6, %w_ab : index + %w = index.assume %w0 [range(%w0, 0, 6)] : index + %word = view.load %word_view[%w] : view<7xi32> -> i32 + // Bytes 0..3 of the word, or qh0, qh1, qh0, qh1 in range c. + %shift2 = scf.select %in_c, %c0_i32, %c16_i32 : i32 + %shift3 = scf.select %in_c, %c8_i32, %c24_i32 : i32 + %b1_shift = scalar.shrui %word, %c8_i32 : i32 + %b2_shift = scalar.shrui %word, %shift2 : i32 + %b3_shift = scalar.shrui %word, %shift3 : i32 + %b0 = scalar.andi %word, %c255_i32 : i32 + %b1 = scalar.andi %b1_shift, %c255_i32 : i32 + %b2 = scalar.andi %b2_shift, %c255_i32 : i32 + %b3 = scalar.andi %b3_shift, %c255_i32 : i32 + // Trit n for all four, or n, n, n + 1, n + 1 in range c. + %p_lo = func.call @ggml_ptq1_0_pow3(%n) : (index) -> (i32) + %p_lo3 = scalar.muli %p_lo, %c3_i32 : i32 + %p_hi = scf.select %in_c, %p_lo3, %p_lo : i32 + %t0 = func.call @ggml_ptq1_0_trit_f32(%b0, %p_lo) : (i32, i32) -> (f32) + %t1 = func.call @ggml_ptq1_0_trit_f32(%b1, %p_lo) : (i32, i32) -> (f32) + %t2 = func.call @ggml_ptq1_0_trit_f32(%b2, %p_hi) : (i32, i32) -> (f32) + %t3 = func.call @ggml_ptq1_0_trit_f32(%b3, %p_hi) : (i32, i32) -> (f32) + %v0 = scalar.mulf %t0, %d : f32 + %v1 = scalar.mulf %t1, %d : f32 + %v2 = scalar.mulf %t2, %d : f32 + %v3 = scalar.mulf %t3, %d : f32 + %result = vector.from_elements %v0, %v1, %v2, %v3 : vector<4xf32> + func.return %result : vector<4xf32> +} + +func.def inline @ggml_ptq1_0_f16_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_ptq1_0_f32_vector4(%weight, %row_byte_base, %k) : (buffer, offset, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/f16_f16.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/f16_f16.loom new file mode 100644 index 000000000000..e64b41ef147a --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/f16_f16.loom @@ -0,0 +1,10 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +// Loads four adjacent F16 weights for FP16 matrix staging. +func.def inline @ggml_f16_f16_vector4(%weight: buffer, %row_byte_base: offset, %input_size: index, %k: index) -> (vector<4xf16>) { + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %row_view = buffer.view %weight[%row_byte_base] : buffer -> view<[%bounded_input_size]xf16> + %values = vector.load %row_view[%k] : view<[%bounded_input_size]xf16> -> vector<4xf16> + func.return %values : vector<4xf16> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/f32_f16.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/f32_f16.loom new file mode 100644 index 000000000000..1de466245d5f --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/f32_f16.loom @@ -0,0 +1,11 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +// Loads four adjacent F32 weights and truncates them for FP16 matrix staging. +func.def inline @ggml_f32_f16_vector4(%weight: buffer, %row_byte_base: offset, %input_size: index, %k: index) -> (vector<4xf16>) { + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %row_view = buffer.view %weight[%row_byte_base] : buffer -> view<[%bounded_input_size]xf32> + %values_f32 = vector.load %row_view[%k] : view<[%bounded_input_size]xf32> -> vector<4xf32> + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/llm_attention_qkv_matmul_postops.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/llm_attention_qkv_matmul_postops.loom new file mode 100644 index 000000000000..3d1868a39790 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/llm_attention_qkv_matmul_postops.loom @@ -0,0 +1,101 @@ +template.decl @llm.attention_qkv.normal_rope_vector4(%position: f32, %theta: vector<2xf32>, %freq_factors: vector<2xf32>, %values: vector<4xf32>) -> (vector<4xf32>) + +template.decl @llm.attention_qkv.rope_vector4(%token: index, %channel: index, %token_count0: index, %head_size0: index, %head_count0: index, %positions: buffer, %theta: buffer, %freq_factors: buffer, %values: vector<4xf32>) -> (vector<4xf32>) + +template.decl @llm.attention_qkv.store_cache_vector4(%output_format: index, %token0: index, %channel: index, %token_count0: index, %cache_row_count0: index, %output_size0: index, %indices: buffer, %values: vector<4xf32>, %cache: buffer) + +template.decl @llm.attention_qkv.store_query_vector4(%token0: index, %channel: index, %token_count0: index, %head_size0: index, %head_count0: index, %values: vector<4xf32>, %output: buffer) + +func.decl @ggml_rope_f32_pair_packet(%position: f32, %theta: vector<2xf32>, %freq_factors: vector<2xf32>, %x_values: vector<2xf32>, %y_values: vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) + +template.def<@llm.attention_qkv.normal_rope_vector4> device @llm_attention_qkv_normal_rope_vector4(%position: f32, %theta: vector<2xf32>, %freq_factors: vector<2xf32>, %values: vector<4xf32>) -> (vector<4xf32>) { + %x0 = vector.extract %values[0] : vector<4xf32> -> f32 + %y0 = vector.extract %values[1] : vector<4xf32> -> f32 + %x1 = vector.extract %values[2] : vector<4xf32> -> f32 + %y1 = vector.extract %values[3] : vector<4xf32> -> f32 + %x_values = vector.from_elements %x0, %x1 : vector<2xf32> + %y_values = vector.from_elements %y0, %y1 : vector<2xf32> + %rotated_x, %rotated_y = func.call @ggml_rope_f32_pair_packet(%position, %theta, %freq_factors, %x_values, %y_values) : (f32, vector<2xf32>, vector<2xf32>, vector<2xf32>, vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) + %rotated_x0 = vector.extract %rotated_x[0] : vector<2xf32> -> f32 + %rotated_x1 = vector.extract %rotated_x[1] : vector<2xf32> -> f32 + %rotated_y0 = vector.extract %rotated_y[0] : vector<2xf32> -> f32 + %rotated_y1 = vector.extract %rotated_y[1] : vector<2xf32> -> f32 + %rotated = vector.from_elements %rotated_x0, %rotated_y0, %rotated_x1, %rotated_y1 : vector<4xf32> + template.return %rotated : vector<4xf32> +} + +template.def<@llm.attention_qkv.rope_vector4> device @llm_attention_qkv_rope_vector4(%token: index, %channel: index, %token_count0: index, %head_size0: index, %head_count0: index, %positions: buffer, %theta: buffer, %freq_factors: buffer, %values: vector<4xf32>) -> (vector<4xf32>) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %head_size = index.assume %head_size0 [range(%head_size0, 4, 1024), mul(%head_size0, 4)] : index + %head_count = index.assume %head_count0 [range(%head_count0, 1, 64)] : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c0_offset = index.constant 0 : offset + %half_head_size = index.div %head_size, %c2 : index + %head_channel = index.rem %channel, %head_size : index + %channel_packet = index.div %head_channel, %c4 : index + %theta_channel = index.mul %channel_packet, %c2 : index + %positions_noalias, %theta_noalias, %freq_factors_noalias = buffer.assume.noalias %positions, %theta, %freq_factors : buffer, buffer, buffer + %positions_view = buffer.view %positions_noalias[%c0_offset] : buffer -> view<[%token_count]xi32> + %theta_view = buffer.view %theta_noalias[%c0_offset] : buffer -> view<[%half_head_size]xf32> + %freq_factors_view = buffer.view %freq_factors_noalias[%c0_offset] : buffer -> view<[%half_head_size]xf32> + %position_i32 = view.load %positions_view[%token] : view<[%token_count]xi32> -> i32 + %position = scalar.sitofp %position_i32 : i32 to f32 + %theta_packet = vector.load %theta_view[%theta_channel] : view<[%half_head_size]xf32> -> vector<2xf32> + %freq_factors_packet = vector.load %freq_factors_view[%theta_channel] : view<[%half_head_size]xf32> -> vector<2xf32> + %rotated = template.apply<@llm.attention_qkv.normal_rope_vector4>(%position, %theta_packet, %freq_factors_packet, %values) : (f32, vector<2xf32>, vector<2xf32>, vector<4xf32>) -> (vector<4xf32>) + template.return %rotated : vector<4xf32> +} + +template.def<@llm.attention_qkv.store_query_vector4> device @llm_attention_qkv_store_query_vector4(%token0: index, %channel: index, %token_count0: index, %head_size0: index, %head_count0: index, %values: vector<4xf32>, %output: buffer) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %head_size = index.assume %head_size0 [range(%head_size0, 4, 1024), mul(%head_size0, 4)] : index + %head_count = index.assume %head_count0 [range(%head_count0, 1, 64)] : index + %c0_offset = index.constant 0 : offset + %head = index.div %channel, %head_size : index + %head_channel = index.rem %channel, %head_size : index + %token, %launch_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %output_noalias = buffer.assume.noalias %output : buffer + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%token_count]x[%head_count]x[%head_size]xf32> + vector.store %values, %output_view[%token, %head, %head_channel] : vector<4xf32>, view<[%token_count]x[%head_count]x[%head_size]xf32> + template.return +} + +template.def<@llm.attention_qkv.store_cache_vector4> device @llm_attention_qkv_store_cache_vector4(%output_format: index, %token0: index, %channel: index, %token_count0: index, %cache_row_count0: index, %output_size0: index, %indices: buffer, %values: vector<4xf32>, %cache: buffer) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %cache_row_count = index.assume %cache_row_count0 [range(%cache_row_count0, 1, 1048576)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 4, 32768), mul(%output_size0, 4)] : index + %c0 = index.constant 0 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c0_i64 = scalar.constant 0 : i64 + %c1048575_i64 = scalar.constant 1048575 : i64 + %c0_offset = index.constant 0 : offset + %is_f16 = index.cmp eq, %output_format, %c16 : index + %is_f32 = index.cmp eq, %output_format, %c32 : index + %token, %launch_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %indices_noalias, %cache_noalias = buffer.assume.noalias %indices, %cache : buffer, buffer + %indices_view = buffer.view %indices_noalias[%c0_offset] : buffer -> view<[%token_count]xi64> + %index_raw = view.load %indices_view[%token] : view<[%token_count]xi64> -> i64 + %index_nonnegative = scalar.cmpi sge, %index_raw, %c0_i64 : i64 + %index_in_cast_range = scalar.cmpi sle, %index_raw, %c1048575_i64 : i64 + %valid_index = scalar.andi %index_nonnegative, %index_in_cast_range : i1 + %safe_index0_i64 = scf.select %valid_index, %index_raw, %c0_i64 : i64 + %safe_index_i64 = scalar.assume %safe_index0_i64 [range(%safe_index0_i64, 0, 1048575)] : i64 + %cache_row0 = index.cast %safe_index_i64 : i64 to index + %valid_row = index.cmp ult, %cache_row0, %cache_row_count : index + %cache_row = scf.select %valid_row, %cache_row0, %c0 : index + %publish = scalar.andi %valid_index, %valid_row : i1 + %publish_f16 = scalar.andi %publish, %is_f16 : i1 + %publish_f32 = scalar.andi %publish, %is_f32 : i1 + scf.if %publish_f16 { + %cache_view = buffer.view %cache_noalias[%c0_offset] : buffer -> view<[%cache_row_count]x[%output_size]xf16> + %truncated = vector.fptrunc %values : vector<4xf32> to vector<4xf16> + vector.store %truncated, %cache_view[%cache_row, %channel] : vector<4xf16>, view<[%cache_row_count]x[%output_size]xf16> + } + scf.if %publish_f32 { + %cache_view = buffer.view %cache_noalias[%c0_offset] : buffer -> view<[%cache_row_count]x[%output_size]xf32> + vector.store %values, %cache_view[%cache_row, %channel] : vector<4xf32>, view<[%cache_row_count]x[%output_size]xf32> + } + template.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_f32_f32_postops.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_f32_f32_postops.loom new file mode 100644 index 000000000000..ade604180adb --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_f32_f32_postops.loom @@ -0,0 +1,149 @@ +template.decl @ggml.mul_mat_f32_f32_wmma.add_bias_vector4(%output_size0: index, %channel: index, %values: vector<4xf32>, %bias: buffer) -> (vector<4xf32>) + +template.decl @ggml.mul_mat_f32_f32_wmma.add_residual_vector4(%token_count0: index, %output_size0: index, %token0: index, %channel: index, %values: vector<4xf32>, %residual_input: buffer) -> (vector<4xf32>) + +template.decl @ggml.mul_mat_f32_f32_wmma.finish_rmsnorm_weight(%token_count0: index, %output_size0: index, %output_tile_count0: index, %token_tile_count0: index, %epsilon: f32, %token_tile: index, %source: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.store_vector4(%token_count0: index, %output_size0: index, %token0: index, %channel: index, %values: vector<4xf32>, %output: buffer) + +template.def<@ggml.mul_mat_f32_f32_wmma.add_bias_vector4> device @ggml_mul_mat_f32_f32_wmma_add_bias_vector4(%output_size0: index, %channel: index, %values: vector<4xf32>, %bias: buffer) -> (vector<4xf32>) { + %output_size = index.assume %output_size0 [range(%output_size0, 1, 262144)] : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %bias_noalias = buffer.assume.noalias %bias : buffer + %bias_view = buffer.view %bias_noalias[%c0_offset] : buffer -> view<[%output_size]xf32> + %mask = vector.mask.range [%channel to %output_size step %c1] : index -> vector<4xi1> + %bias_values = vector.load.mask %bias_view[%channel], %mask, %c0_f32x4 : view<[%output_size]xf32>, vector<4xi1>, vector<4xf32> + %sum = vector.addf %values, %bias_values : vector<4xf32> + template.return %sum : vector<4xf32> +} + +template.def<@ggml.mul_mat_f32_f32_wmma.add_residual_vector4> device @ggml_mul_mat_f32_f32_wmma_add_residual_vector4(%token_count0: index, %output_size0: index, %token0: index, %channel: index, %values: vector<4xf32>, %residual_input: buffer) -> (vector<4xf32>) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 1, 262144)] : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %residual_input_noalias = buffer.assume.noalias %residual_input : buffer + %residual_input_view = buffer.view %residual_input_noalias[%c0_offset] : buffer -> view<[%token_count]x[%output_size]xf32> + %token, %launch_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %mask = vector.mask.range [%channel to %output_size step %c1] : index -> vector<4xi1> + %residual = vector.load.mask %residual_input_view[%token, %channel], %mask, %c0_f32x4 : view<[%token_count]x[%output_size]xf32>, vector<4xi1>, vector<4xf32> + %sum = vector.addf %values, %residual : vector<4xf32> + template.return %sum : vector<4xf32> +} + +template.def<@ggml.mul_mat_f32_f32_wmma.store_vector4> device @ggml_mul_mat_f32_f32_wmma_store_vector4(%token_count0: index, %output_size0: index, %token0: index, %channel: index, %values: vector<4xf32>, %output: buffer) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 1, 262144)] : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %output_noalias = buffer.assume.noalias %output : buffer + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%token_count]x[%output_size]xf32> + %token, %launch_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %mask = vector.mask.range [%channel to %output_size step %c1] : index -> vector<4xi1> + vector.store.mask %values, %output_view[%token, %channel], %mask : vector<4xf32>, view<[%token_count]x[%output_size]xf32>, vector<4xi1> + template.return +} + +template.def<@ggml.mul_mat_f32_f32_wmma.finish_rmsnorm_weight> device @ggml_mul_mat_f32_f32_wmma_finish_rmsnorm_weight(%token_count0: index, %output_size0: index, %output_tile_count0: index, %token_tile_count0: index, %epsilon: f32, %token_tile: index, %source: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 128, 32768), mul(%output_size0, 128)] : index + %output_tile_count = index.assume %output_tile_count0 [range(%output_tile_count0, 1, 4096)] : index + %token_tile_count = index.assume %token_tile_count0 [range(%token_tile_count0, 1, 64)] : index + %workitem = kernel.workitem.id : index + %workgroup_size = kernel.workgroup.size : index + %subgroup = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %subgroup_count = index.div %workgroup_size, %c64 : index + %c0_offset = index.constant 0 : offset + %counter_scratch_bytes = index.constant 4 : offset + %rms_scratch_bytes = index.constant 512 : offset + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c0_f32 = scalar.constant 0.0 : f32 + %hidden_size_i32 = index.cast %output_size : index to i32 + %hidden_size_f32 = scalar.sitofp %hidden_size_i32 : i32 to f32 + %output_tile_count_i32 = index.cast %output_tile_count : index to i32 + %last_channel_ordinal_i32 = scalar.subi %output_tile_count_i32, %c1_i32 : i32 + %negative_output_tile_count_i32 = scalar.subi %c0_i32, %output_tile_count_i32 : i32 + %source_noalias, %norm_weight_noalias, %normalized_output_noalias, %completion_counters_noalias = buffer.assume.noalias %source, %norm_weight, %normalized_output, %completion_counters : buffer, buffer, buffer, buffer + %completion_counters_aligned = buffer.assume.alignment %completion_counters_noalias {minimum_alignment = 16} : buffer + %source_view = buffer.view %source_noalias[%c0_offset] : buffer -> view<[%token_count]x[%output_size]xf32> + %norm_weight_view = buffer.view %norm_weight_noalias[%c0_offset] : buffer -> view<[%output_size]xf32> + %normalized_output_view = buffer.view %normalized_output_noalias[%c0_offset] : buffer -> view<[%token_count]x[%output_size]xf32> + %completion_counter_view = buffer.view %completion_counters_aligned[%c0_offset] : buffer -> view<[%token_tile_count]xi32> + %counter_scratch = buffer.alloca align(4) %counter_scratch_bytes : buffer + %rms_scratch = buffer.alloca align(16) %rms_scratch_bytes : buffer + %counter_scratch_view = buffer.view %counter_scratch[%c0_offset] : buffer -> view<1xi32> + %rms_scratch_view = buffer.view %rms_scratch[%c0_offset] : buffer -> view<128xf32> + %token_tile_base = index.mul %token_tile, %c32 : index + kernel.barrier scope(workgroup) ordering(release) + %workitem_is_zero = index.cmp eq, %workitem, %c0 : index + scf.if %workitem_is_zero { + %old_counter = view.atomic.rmw %c1_i32, %completion_counter_view[%token_tile] {ordering = acq_rel, scope = device} : i32, view<[%token_tile_count]xi32> -> i32 + view.store %old_counter, %counter_scratch_view[%c0] : i32, view<1xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %old_counter = view.load %counter_scratch_view[%c0] : view<1xi32> -> i32 + %is_last_channel_tile = scalar.cmpi eq, %old_counter, %last_channel_ordinal_i32 : i32 + scf.if %is_last_channel_tile { + kernel.barrier scope(workgroup) ordering(acquire) + scf.for %token_offset = [%c0 to %c32 step %c1] { + %token0 = index.add %token_tile_base, %token_offset : index + %valid_token = index.cmp ult, %token0, %token_count : index + scf.if %valid_token { + %token, %launch_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %thread_sum = scf.for %channel = [%workitem to %output_size step %workgroup_size](%running_sum = %c0_f32 : f32) -> (f32) { + %value = view.load %source_view[%token, %channel] : view<[%token_count]x[%output_size]xf32> -> f32 + %square = scalar.mulf %value, %value : f32 + %next_sum = scalar.addf %running_sum, %square : f32 + scf.yield %next_sum : f32 + } + %subgroup_sum = kernel.subgroup.reduce %thread_sum : f32 + %is_subgroup_leader = index.cmp eq, %lane, %c0 : index + scf.if %is_subgroup_leader { + view.store %subgroup_sum, %rms_scratch_view[%subgroup] : f32, view<128xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %is_reduction_subgroup = index.cmp eq, %subgroup, %c0 : index + %is_reduction_lane = index.cmp ult, %lane, %subgroup_count : index + %loads_subgroup_sum = scalar.andi %is_reduction_subgroup, %is_reduction_lane : i1 + %subgroup_partial = scf.if %loads_subgroup_sum -> (f32) { + %value = view.load %rms_scratch_view[%lane] : view<128xf32> -> f32 + scf.yield %value : f32 + } else { + scf.yield %c0_f32 : f32 + } + %row_sum = kernel.subgroup.reduce %subgroup_partial : f32 + %writes_scale = scalar.andi %is_reduction_subgroup, %is_subgroup_leader : i1 + scf.if %writes_scale { + %mean = scalar.divf %row_sum, %hidden_size_f32 : f32 + %biased_mean = scalar.addf %mean, %epsilon : f32 + %scale = scalar.rsqrtf %biased_mean : f32 + view.store %scale, %rms_scratch_view[%c0] : f32, view<128xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %scale = view.load %rms_scratch_view[%c0] : view<128xf32> -> f32 + scf.for %channel = [%workitem to %output_size step %workgroup_size] { + %value = view.load %source_view[%token, %channel] : view<[%token_count]x[%output_size]xf32> -> f32 + %weight_value = view.load %norm_weight_view[%channel] : view<[%output_size]xf32> -> f32 + %scaled = scalar.mulf %value, %scale : f32 + %normalized = scalar.mulf %scaled, %weight_value : f32 + view.store %normalized, %normalized_output_view[%token, %channel] : f32, view<[%token_count]x[%output_size]xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + } + } + kernel.barrier scope(workgroup) ordering(release) + scf.if %workitem_is_zero { + view.atomic.reduce %negative_output_tile_count_i32, %completion_counter_view[%token_tile] {ordering = release, scope = device} : i32, view<[%token_tile_count]xi32> + } + } + template.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_f32_f32_wmma_core.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_f32_f32_wmma_core.loom new file mode 100644 index 000000000000..ea4fd32dc9df --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_f32_f32_wmma_core.loom @@ -0,0 +1,1607 @@ +func.decl @ggml_q4k_scale_min_from_header(%scale0: i32, %scale1: i32, %scale2: i32, %q4_group: index) -> (i32, i32) + +func.decl @ggml_dot_u8_s8_vector8_f32(%weight: vector<2xi32>, %activation: vector<2xi32>) -> (f32) + +config.def @ggml.mul_mat.activation_format = 0 : index + +func.decl @ggml_q4k_native_row64_f32_vector4(%weight: buffer, %row: index, %input_size: index, %q4_block: index, %q4_group: index, %packet: index) -> (vector<4xf32>) + +func.decl @ggml_dequant_f32_vector4(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight_format: index, %weight: buffer, %row_byte_base: offset, %input_size: index, %quant_block: index, %quant_group: index, %packet: index, %k: index) -> (vector<4xf32>) + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +// Shared contiguous 2D F32 activation WMMA matmul pipeline. +// +// Op wrappers provide publication and tile-finalization hooks by defining: +// ggml.mul_mat_f32_f32_wmma.publish_vector4 +// ggml.mul_mat_f32_f32_wmma.finish_tile +template.decl @ggml.mul_mat_f32_f32_wmma.core(%weight_format: index, %token_count: index, %input_size0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %input: buffer, %weight: buffer, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.finish_tile(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: f32, %arg7: index, %arg8: index, %arg9: buffer, %arg10: buffer, %arg11: buffer, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.launch(%token_capacity: index, %input_size: index, %output_size: index, %weight_format: index) -> (index, index, index, index) + +template.decl @ggml.mul_mat_f32_f32_wmma.publish_vector4(%arg0: i1, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: index, %arg7: vector<4xf32>, %arg8: buffer, %arg9: buffer, %arg10: buffer, %arg11: buffer, %arg12: buffer, %arg13: buffer, %arg14: buffer) + +func.decl @ggml_dequant_weight_tile_bytes(%weight_format: index) -> (offset) + +func.decl @ggml_dequant_weight_row_bytes(%weight_format: index, %hidden_size: index) -> (offset) + +func.decl @ggml_iq4nl_table_i8() -> (vector<16xi8>) + +func.decl @ggml_iq3s_grid_f32_halves() -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + +func.decl @ggml_dequant_f16_vector4(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight_format: index, %weight: buffer, %row_byte_base: offset, %input_size: index, %quant_block: index, %quant_group: index, %packet: index, %k: index) -> (vector<4xf16>) + +func.def inline @ggml_mul_mat_f32_f32_wmma_tiled(%weight_format: index, %token_count: index, %input_size0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %input: buffer, %weight: buffer, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_input_size = index.assume %input_size0 [range(%input_size0, 256, 32768), mul(%input_size0, 32)] : index + %bounded_output_size = index.assume %output_size0 [range(%output_size0, 1, 262144)] : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 1)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c40 = index.constant 40 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %weight_stage_bytes = index.constant 5120 : offset + %activation_stage_bytes = index.constant 2560 : offset + %wave_result_stage_bytes = index.constant 1024 : offset + %result_stage_bytes = index.constant 2048 : offset + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %c0_f16x4 = vector.constant 0.0 : vector<4xf16> + %zero_accumulator = vector.constant 0.0 : vector<4xf32> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %padded_output_size = index.add %bounded_output_size, %c63 : index + %output_tile_count = index.div %padded_output_size, %c64 : index + %padded_token_count = index.add %bounded_token_count, %c31 : index + %token_tile_count = index.div %padded_token_count, %c32 : index + %padded_input_size = index.add %bounded_input_size, %c255 : index + %quant_block_count = index.div %padded_input_size, %c256 : index + %weight_row_bytes = func.call @ggml_dequant_weight_row_bytes(%weight_format, %bounded_input_size) : (index, index) -> (offset) + %split_count = kernel.workgroup.count : index + %partition = kernel.workgroup.id : index + %blocks_per_partition = index.div %quant_block_count, %split_count : index + %quant_block_begin = index.mul %partition, %blocks_per_partition : index + %quant_block_end = index.add %quant_block_begin, %blocks_per_partition : index + %input_noalias, %weight_noalias = buffer.assume.noalias %input, %weight : buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_input_size]xf32> + %weight_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %activation_stage = buffer.alloca align(16) %activation_stage_bytes : buffer + %result_stage = buffer.alloca align(16) %result_stage_bytes : buffer + %weight_stage_view = buffer.view %weight_stage[%c0_offset] : buffer -> view<64x40xf16> + %activation_stage_physical_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x40xf16> + %activation_fragment_layout = encoding.layout.strided [1, %c40] : encoding + %activation_fragment_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x32xf16, %activation_fragment_layout> + %wave_result_stage_offset = index.scale %subgroup, %wave_result_stage_bytes : index, offset -> offset + %result_fragment_layout = encoding.layout.strided [1, %c16] : encoding + %result_fragment_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf32, %result_fragment_layout> + %result_physical_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf32> + %channel_tile_base = index.mul %channel_tile, %c64 : index + %token_tile_base = index.mul %token_tile, %c32 : index + %load_packet = index.rem %workitem, %c8 : index + %load_k = index.mul %load_packet, %c4 : index + %load_row0 = index.div %workitem, %c8 : index + %load_row = index.assume %load_row0 [range(%load_row0, 0, 15)] : index + %subgroup_channel_add = index.mul %subgroup, %c32 : index + %subgroup_channel1 = index.add %subgroup_channel_add, %c16 : index + %init00 = vector.fragment %zero_accumulator shape [%m, %n] : vector<4xf32> + %init01 = vector.fragment %zero_accumulator shape [%m, %n] : vector<4xf32> + %init10 = vector.fragment %zero_accumulator shape [%m, %n] : vector<4xf32> + %init11 = vector.fragment %zero_accumulator shape [%m, %n] : vector<4xf32> + %iq4nl_table = func.call @ggml_iq4nl_table_i8() : () -> (vector<16xi8>) + %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15 = func.call @ggml_iq3s_grid_f32_halves() : () -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + %result00, %result01, %result10, %result11 = scf.for %quant_block = [%quant_block_begin to %quant_block_end step %c1](%block_acc00 = %init00 : vector<4xf32>, %block_acc01 = %init01 : vector<4xf32>, %block_acc10 = %init10 : vector<4xf32>, %block_acc11 = %init11 : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + %block_result00, %block_result01, %block_result10, %block_result11 = scf.for %quant_group = [%c0 to %c8 step %c1](%acc00 = %block_acc00 : vector<4xf32>, %acc01 = %block_acc01 : vector<4xf32>, %acc10 = %block_acc10 : vector<4xf32>, %acc11 = %block_acc11 : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + %block_k_base = index.mul %quant_block, %c256 : index + %group_k_add = index.mul %quant_group, %c32 : index + %k_origin = index.add %block_k_base, %group_k_add : index + scf.for %row_offset = [%c0 to %c64 step %c16] unroll { + %local_row0 = index.add %load_row, %row_offset : index + %local_row = index.assume %local_row0 [range(%local_row0, 0, 63)] : index + %channel = index.add %channel_tile_base, %local_row : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %weight_k = index.add %k_origin, %load_k : index + %valid_k = index.cmp ult, %weight_k, %bounded_input_size : index + %valid_weight = scalar.andi %valid_channel, %valid_k : i1 + %weight_values = scf.if %valid_weight -> (vector<4xf16>) { + %row_byte_base = index.scale %channel, %weight_row_bytes : index, offset -> offset + %decoded = func.call @ggml_dequant_f16_vector4(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight_format, %weight_noalias, %row_byte_base, %bounded_input_size, %quant_block, %quant_group, %load_packet, %weight_k) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, buffer, offset, index, index, index, index, index) -> (vector<4xf16>) + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %weight_values, %weight_stage_view[%local_row, %load_k] : vector<4xf16>, view<64x40xf16> + %is_activation_row = index.cmp ult, %local_row, %c32 : index + scf.if %is_activation_row { + %activation_row = index.assume %local_row [range(%local_row, 0, 31)] : index + %token = index.add %token_tile_base, %activation_row : index + %valid_token = index.cmp ult, %token, %bounded_token_count : index + %input_k = index.add %k_origin, %load_k : index + %valid_input = index.cmp ult, %input_k, %bounded_input_size : index + %valid_activation = scalar.andi %valid_token, %valid_input : i1 + %activation_values = scf.if %valid_activation -> (vector<4xf16>) { + %bounded_token, %input_token_count = index.assume %token, %bounded_token_count [lt(%token, %bounded_token_count)] : index, index + %mask = vector.mask.range [%input_k to %bounded_input_size step %c1] : index -> vector<4xi1> + %loaded = vector.load.mask %input_view[%bounded_token, %input_k], %mask, %c0_f32x4 : view<[%bounded_token_count]x[%bounded_input_size]xf32>, vector<4xi1>, vector<4xf32> + %converted = vector.fptrunc %loaded : vector<4xf32> to vector<4xf16> + scf.yield %converted : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %activation_values, %activation_stage_physical_view[%activation_row, %load_k] : vector<4xf16>, view<32x40xf16> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %next00, %next01, %next10, %next11 = scf.for %k_half = [%c0 to %c32 step %c16](%half_acc00 = %acc00 : vector<4xf32>, %half_acc01 = %acc01 : vector<4xf32>, %half_acc10 = %acc10 : vector<4xf32>, %half_acc11 = %acc11 : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) unroll { + %lhs0 = vector.fragment.load %weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %lhs1 = vector.fragment.load %weight_stage_view[%subgroup_channel1, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %rhs0 = vector.fragment.load %activation_fragment_view[%k_half, %c0] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %rhs1 = vector.fragment.load %activation_fragment_view[%k_half, %c16] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %half_next00 = vector.mma %lhs0, %rhs0, %half_acc00 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %half_next01 = vector.mma %lhs0, %rhs1, %half_acc01 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %half_next10 = vector.mma %lhs1, %rhs0, %half_acc10 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %half_next11 = vector.mma %lhs1, %rhs1, %half_acc11 : vector<16xf16>, vector<16xf16>, vector<4xf32> + scf.yield %half_next00, %half_next01, %half_next10, %half_next11 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %next00, %next01, %next10, %next11 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + scf.yield %block_result00, %block_result01, %block_result10, %block_result11 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + %publish_token0 = index.div %lane, %c4 : index + %publish_token = index.assume %publish_token0 [range(%publish_token0, 0, 15)] : index + %publish_packet0 = index.rem %lane, %c4 : index + %publish_packet = index.assume %publish_packet0 [range(%publish_packet0, 0, 3)] : index + %publish_channel_add = index.mul %publish_packet, %c4 : index + %token0 = index.add %token_tile_base, %publish_token : index + %publish_token1 = index.add %publish_token, %c16 : index + %token1 = index.add %token_tile_base, %publish_token1 : index + %subgroup_channel_base = index.add %channel_tile_base, %subgroup_channel_add : index + %channel0 = index.add %subgroup_channel_base, %publish_channel_add : index + %channel1_base = index.add %subgroup_channel_base, %c16 : index + %channel1 = index.add %channel1_base, %publish_channel_add : index + %valid_token0 = index.cmp ult, %token0, %bounded_token_count : index + %valid_token1 = index.cmp ult, %token1, %bounded_token_count : index + %valid_channel0 = index.cmp ult, %channel0, %bounded_output_size : index + %valid_channel1 = index.cmp ult, %channel1, %bounded_output_size : index + %writes00 = scalar.andi %valid_token0, %valid_channel0 : i1 + %writes01 = scalar.andi %valid_token1, %valid_channel0 : i1 + %writes10 = scalar.andi %valid_token0, %valid_channel1 : i1 + %writes11 = scalar.andi %valid_token1, %valid_channel1 : i1 + vector.fragment.store %result00, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes00 { + %values = vector.load %result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + template.apply<@ggml.mul_mat_f32_f32_wmma.publish_vector4>(%writes00, %token0, %channel0, %bounded_token_count, %bounded_output_size, %output_accumulation, %output_unary_op, %values, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (i1, index, index, index, index, index, index, vector<4xf32>, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result01, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes01 { + %values = vector.load %result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + template.apply<@ggml.mul_mat_f32_f32_wmma.publish_vector4>(%writes01, %token1, %channel0, %bounded_token_count, %bounded_output_size, %output_accumulation, %output_unary_op, %values, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (i1, index, index, index, index, index, index, vector<4xf32>, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result10, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes10 { + %values = vector.load %result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + template.apply<@ggml.mul_mat_f32_f32_wmma.publish_vector4>(%writes10, %token0, %channel1, %bounded_token_count, %bounded_output_size, %output_accumulation, %output_unary_op, %values, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (i1, index, index, index, index, index, index, vector<4xf32>, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result11, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes11 { + %values = vector.load %result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + template.apply<@ggml.mul_mat_f32_f32_wmma.publish_vector4>(%writes11, %token1, %channel1, %bounded_token_count, %bounded_output_size, %output_accumulation, %output_unary_op, %values, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (i1, index, index, index, index, index, index, vector<4xf32>, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + } + template.apply<@ggml.mul_mat_f32_f32_wmma.finish_tile>(%bounded_token_count, %bounded_output_size, %output_tile_count, %token_tile_count, %output_accumulation, %output_unary_op, %epsilon, %channel_tile, %token_tile, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + func.return +} + +func.def pure inline @ggml_q8_lowtoken_geometry(%token_capacity: index, %input_size: index, %output_size: index, %weight_format: index) -> (index, index) { + %c0 = index.constant 0 : index + %c5 = index.constant 5 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %q4_format = index.constant 44 : index + %q6_format = index.constant 46 : index + %q8_format = index.constant 9 : index + %activation_format = config.get @ggml.mul_mat.activation_format : index + %native_q4 = index.cmp eq, %weight_format, %q4_format : index + %native_q6 = index.cmp eq, %weight_format, %q6_format : index + %q8_input = index.cmp eq, %activation_format, %q8_format : index + %q4_min_blocks = index.constant 327680 : index + %blocks = index.div %input_size, %c256 : index + %weight_blocks = index.mul %blocks, %output_size : index + %large_payload = index.cmp uge, %weight_blocks, %q4_min_blocks : index + %q4_can_stage = scalar.andi %native_q4, %large_payload : i1 + %native_staged = scalar.ori %q4_can_stage, %native_q6 : i1 + %native_q8 = scalar.andi %native_staged, %q8_input : i1 + %q4_min_k = index.constant 16384 : index + %q6_min_k = index.constant 8192 : index + %min_k = scf.select %native_q4, %q4_min_k, %q6_min_k : index + %min_n = index.constant 4096 : index + %five_tokens = index.cmp eq, %token_capacity, %c5 : index + %long_k = index.cmp uge, %input_size, %min_k : index + %wide_n = index.cmp uge, %output_size, %min_n : index + %contracts = index.cmp ule, %output_size, %input_size : index + %large = scalar.andi %long_k, %wide_n : i1 + %long_reduction = scalar.andi %large, %contracts : i1 + %staged_shape = scalar.andi %five_tokens, %long_reduction : i1 + %staged = scalar.andi %native_q8, %staged_shape : i1 + %rows = scf.select %native_q4, %c64, %c32 : index + %result = scf.select %staged, %rows, %c0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %block_bytes = index.constant 144 : index + %half_cache_bytes = index.constant 16777216 : index + %cache_bytes = index.constant 33554432 : index + %payload_bytes = index.mul %weight_blocks, %block_bytes : index + %large_half = index.cmp uge, %payload_bytes, %half_cache_bytes : index + %fits_cache = index.cmp ule, %payload_bytes, %cache_bytes : index + %cache_range = scalar.andi %large_half, %fits_cache : i1 + %remainder = index.rem %blocks, %c2 : index + %even_blocks = index.cmp eq, %remainder, %c0 : index + %reduces_rows = index.cmp ult, %output_size, %input_size : index + %native_q4_q8 = scalar.andi %native_q4, %q8_input : i1 + %plain = index.cmp eq, %result, %c0 : index + %q4_plain = scalar.andi %native_q4_q8, %plain : i1 + %five_wide = scalar.andi %five_tokens, %wide_n : i1 + %reduction = scalar.andi %reduces_rows, %even_blocks : i1 + %partition_shape = scalar.andi %five_wide, %reduction : i1 + %partition_payload = scalar.andi %partition_shape, %cache_range : i1 + %partitioned = scalar.andi %q4_plain, %partition_payload : i1 + %k_partitions = scf.select %partitioned, %c2, %c1 : index + func.return %result, %k_partitions : index, index +} + +func.def pure inline @ggml_mul_mat_uses_lowtoken_dot(%token_capacity: index, %input_size: index) -> (i1) { + %c0 = index.constant 0 : index + %c5 = index.constant 5 : index + %c256 = index.constant 256 : index + %few_tokens = index.cmp ule, %token_capacity, %c5 : index + %tail = index.rem %input_size, %c256 : index + %whole_blocks = index.cmp eq, %tail, %c0 : index + %use_dot = scalar.andi %few_tokens, %whole_blocks : i1 + func.return %use_dot : i1 +} + +template.def<@ggml.mul_mat_f32_f32_wmma.launch> pure @ggml_mul_mat_f32_f32_wmma_launch(%token_capacity: index, %input_size: index, %output_size: index, %weight_format: index) -> (index, index, index, index) { + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %padded_token_count = index.add %token_capacity, %c31 : index + %token_tiles = index.div %padded_token_count, %c32 : index + %c5 = index.constant 5 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %lowtoken = func.call pure @ggml_mul_mat_uses_lowtoken_dot(%token_capacity, %input_size) : (index, index) -> (i1) + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c256 = index.constant 256 : index + %small_output = index.cmp ult, %output_size, %c64 : index + %narrow = scalar.andi %lowtoken, %small_output : i1 + %c0 = index.constant 0 : index + %c16 = index.constant 16 : index + %staged_rows, %k_partitions = func.call pure @ggml_q8_lowtoken_geometry(%token_capacity, %input_size, %output_size, %weight_format) : (index, index, index, index) -> (index, index) + %staged = index.cmp ugt, %staged_rows, %c0 : index + %staged_rounding = index.sub %staged_rows, %c1 : index + %staged_threads = index.mul %staged_rows, %c16 : index + %cohort_rows = index.div %c8, %k_partitions : index + %cohort_rounding = index.sub %cohort_rows, %c1 : index + %base_row_tile = scf.select %narrow, %c4, %cohort_rows : index + %c1024 = index.constant 1024 : index + %c4096 = index.constant 4096 : index + %activation_format = config.get @ggml.mul_mat.activation_format : index + %f32_input = index.cmp eq, %activation_format, %c0 : index + %batched_tokens = index.cmp ugt, %token_capacity, %c1 : index + %long_reduction = index.cmp uge, %input_size, %c4096 : index + %partition_tail = index.rem %input_size, %c1024 : index + %full_partitions = index.cmp eq, %partition_tail, %c0 : index + %batch_f32 = scalar.andi %batched_tokens, %f32_input : i1 + %partitionable = scalar.andi %long_reduction, %full_partitions : i1 + %partition_rows = scalar.andi %batch_f32, %partitionable : i1 + %narrow_threads = scf.select %partition_rows, %c1024, %c256 : index + %base_workgroup_size = scf.select %narrow, %narrow_threads, %c128 : index + %base_rounding = scf.select %narrow, %c3, %cohort_rounding : index + %row_rounding = scf.select %staged, %staged_rounding, %base_rounding : index + %row_tile = scf.select %staged, %staged_rows, %base_row_tile : index + %workgroup_size = scf.select %staged, %staged_threads, %base_workgroup_size : index + %rounded_n = index.add %output_size, %row_rounding : index + %direct_tiles = index.div %rounded_n, %row_tile : index + %launch_x = scf.select %lowtoken, %direct_tiles, %output_tiles : index + %launch_y = scf.select %lowtoken, %c1, %token_tiles : index + template.return %launch_x, %launch_y, %c1, %workgroup_size : index, index, index, index +} + +func.def inline @ggml_lowtoken_combine_k_partitions(%value: f32, %partitions: index) -> (f32) { + %c1 = index.constant 1 : index + %partitioned = index.cmp ugt, %partitions, %c1 : index + %result = scf.if %partitioned -> (f32) { + %offset = scalar.constant 16 : i32 + %width = scalar.constant 32 : i32 + %peer, %ok = kernel.subgroup.shuffle %value, %offset, %width : f32, i32, i32 + %sum = scalar.addf %value, %peer : f32 + scf.yield %sum : f32 + } else { + scf.yield %value : f32 + } + func.return %result : f32 +} + +func.def inline @ggml_lowtoken_reduce_cohort_f32(%value: f32) -> (f32) { + %width = scalar.constant 32 : i32 + %off1 = scalar.constant 1 : i32 + %peer1, %ok1 = kernel.subgroup.shuffle %value, %off1, %width : f32, i32, i32 + %sum1 = scalar.addf %value, %peer1 : f32 + %off2 = scalar.constant 2 : i32 + %peer2, %ok2 = kernel.subgroup.shuffle %sum1, %off2, %width : f32, i32, i32 + %sum2 = scalar.addf %sum1, %peer2 : f32 + %off4 = scalar.constant 4 : i32 + %peer4, %ok4 = kernel.subgroup.shuffle %sum2, %off4, %width : f32, i32, i32 + %sum4 = scalar.addf %sum2, %peer4 : f32 + %off8 = scalar.constant 8 : i32 + %peer8, %ok8 = kernel.subgroup.shuffle %sum4, %off8, %width : f32, i32, i32 + %sum8 = scalar.addf %sum4, %peer8 : f32 + func.return %sum8 : f32 +} + +// Narrow low-token projections assign a full wave to each row. +// Both cohort sizes use the existing wrapper publication hooks. +func.def inline @ggml_lowtoken_reduce_f32(%values: vector<4xf32>, %full_wave: i1) -> (f32) { + %zero = scalar.constant 0.0 : f32 + %width = scalar.constant 32 : i32 + %xor1 = scalar.constant 1 : i32 + %xor2 = scalar.constant 2 : i32 + %xor4 = scalar.constant 4 : i32 + %xor8 = scalar.constant 8 : i32 + %sum = vector.reduce %values, %zero : vector<4xf32>, f32 + %peer1, %ok1 = kernel.subgroup.shuffle %sum, %xor1, %width : f32, i32, i32 + %sum1 = scalar.addf %sum, %peer1 : f32 + %peer2, %ok2 = kernel.subgroup.shuffle %sum1, %xor2, %width : f32, i32, i32 + %sum2 = scalar.addf %sum1, %peer2 : f32 + %peer4, %ok4 = kernel.subgroup.shuffle %sum2, %xor4, %width : f32, i32, i32 + %sum4 = scalar.addf %sum2, %peer4 : f32 + %peer8, %ok8 = kernel.subgroup.shuffle %sum4, %xor8, %width : f32, i32, i32 + %sum8 = scalar.addf %sum4, %peer8 : f32 + %result = scf.if %full_wave -> (f32) { + %xor16 = scalar.constant 16 : i32 + %peer16, %ok16 = kernel.subgroup.shuffle %sum8, %xor16, %width : f32, i32, i32 + %sum16 = scalar.addf %sum8, %peer16 : f32 + %lane32 = scalar.constant 32 : i32 + %peer32 = kernel.subgroup.broadcast %sum16 from %lane32 : f32, i32 + %sum64 = scalar.addf %sum16, %peer32 : f32 + scf.yield %sum64 : f32 + } else { + scf.yield %sum8 : f32 + } + func.return %result : f32 +} + +func.def inline @ggml_mul_mat_f32_f32_lowtoken_dot(%weight_format: index, %token_count: index, %input_size0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %input: buffer, %weight: buffer, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %iq4nl_table = func.call @ggml_iq4nl_table_i8() : () -> (vector<16xi8>) + %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15 = func.call @ggml_iq3s_grid_f32_halves() : () -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %base = index.constant 0 : offset + %zero = vector.constant 0.0 : vector<4xf32> + %tokens = index.assume %token_count [range(%token_count, 1, 5)] : index + %k = index.assume %input_size0 [range(%input_size0, 256, 32768), mul(%input_size0, 256)] : index + %n = index.assume %output_size0 [range(%output_size0, 1, 262144)] : index + %workgroup_size = kernel.workgroup.size : index + %one_wave_per_row = index.cmp eq, %workgroup_size, %c256 : index + %c1024 = index.constant 1024 : index + %partition_group = index.cmp eq, %workgroup_size, %c1024 : index + %narrow_output = index.cmp ult, %n, %c64 : index + %partitioned = scalar.andi %partition_group, %narrow_output : i1 + %full_wave = scalar.ori %one_wave_per_row, %partitioned : i1 + %k_partitions = scf.select %partitioned, %c4, %c1 : index + %cohort_width = scf.select %full_wave, %c64, %c16 : index + %tile_rows = scf.select %full_wave, %c4, %c8 : index + %wave_rows = scf.select %full_wave, %c1, %c4 : index + %step_width = scf.select %full_wave, %c256, %c64 : index + %steps_per_block = scf.select %full_wave, %c1, %c4 : index + %subgroup = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %lane16 = index.rem %lane, %cohort_width : index + %cohort = index.div %lane, %cohort_width : index + %tile_base = index.mul %channel_tile, %tile_rows : index + %row_subgroup = index.div %subgroup, %k_partitions : index + %k_partition = index.rem %subgroup, %k_partitions : index + %wave_add = index.mul %row_subgroup, %wave_rows : index + %wave_base = index.add %tile_base, %wave_add : index + %row = index.add %wave_base, %cohort : index + %valid_row = index.cmp ult, %row, %n : index + %packet = index.rem %lane16, %c8 : index + %lane_group = index.div %lane16, %c8 : index + %lane_k = index.mul %lane16, %c4 : index + %blocks = index.div %k, %c256 : index + %steps = index.div %k, %step_width : index + %partition_steps = index.div %steps, %k_partitions : index + %step_begin = index.mul %k_partition, %partition_steps : index + %step_end = index.add %step_begin, %partition_steps : index + %weight_bytes = func.call @ggml_dequant_weight_tile_bytes(%weight_format) : (index) -> (offset) + %row_bytes = index.scale %blocks, %weight_bytes : index, offset -> offset + %row_base = index.scale %row, %row_bytes : index, offset -> offset + %a_na, %w_na = buffer.assume.noalias %input, %weight : buffer, buffer + %av = buffer.view %a_na[%base] : buffer -> view<[%tokens]x[%k]xf32> + %total0, %total1, %total2, %total3, %total4 = scf.for %step = [%step_begin to %step_end step %c1](%acc0 = %zero : vector<4xf32>, %acc1 = %zero : vector<4xf32>, %acc2 = %zero : vector<4xf32>, %acc3 = %zero : vector<4xf32>, %acc4 = %zero : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + %block = index.div %step, %steps_per_block : index + %quarter = index.rem %step, %steps_per_block : index + %quarter_group = index.mul %quarter, %c2 : index + %group = index.add %quarter_group, %lane_group : index + %step_k = index.mul %step, %step_width : index + %kk = index.add %step_k, %lane_k : index + %ww = scf.if %valid_row -> (vector<4xf32>) { + %format44 = index.constant 44 : index + %packed = index.cmp eq, %weight_format, %format44 : index + %wh = scf.if %packed -> (vector<4xf32>) { + %v = func.call @ggml_q4k_native_row64_f32_vector4(%w_na, %row, %k, %block, %group, %packet) : (buffer, index, index, index, index, index) -> (vector<4xf32>) + scf.yield %v : vector<4xf32> + } else { + %v = func.call @ggml_dequant_f32_vector4(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight_format, %w_na, %row_base, %k, %block, %group, %packet, %kk) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, buffer, offset, index, index, index, index, index) -> (vector<4xf32>) + scf.yield %v : vector<4xf32> + } + scf.yield %wh : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %t0 = index.constant 0 : index + %valid0 = index.cmp ult, %t0, %tokens : index + %next0 = scf.if %valid0 -> (vector<4xf32>) { + %a = vector.load %av[%t0, %kk] : view<[%tokens]x[%k]xf32> -> vector<4xf32> + %r = vector.fmaf %ww, %a, %acc0 : vector<4xf32> + scf.yield %r : vector<4xf32> + } else { + scf.yield %acc0 : vector<4xf32> + } + %t1 = index.constant 1 : index + %valid1 = index.cmp ult, %t1, %tokens : index + %next1 = scf.if %valid1 -> (vector<4xf32>) { + %a = vector.load %av[%t1, %kk] : view<[%tokens]x[%k]xf32> -> vector<4xf32> + %r = vector.fmaf %ww, %a, %acc1 : vector<4xf32> + scf.yield %r : vector<4xf32> + } else { + scf.yield %acc1 : vector<4xf32> + } + %t2 = index.constant 2 : index + %valid2 = index.cmp ult, %t2, %tokens : index + %next2 = scf.if %valid2 -> (vector<4xf32>) { + %a = vector.load %av[%t2, %kk] : view<[%tokens]x[%k]xf32> -> vector<4xf32> + %r = vector.fmaf %ww, %a, %acc2 : vector<4xf32> + scf.yield %r : vector<4xf32> + } else { + scf.yield %acc2 : vector<4xf32> + } + %t3 = index.constant 3 : index + %valid3 = index.cmp ult, %t3, %tokens : index + %next3 = scf.if %valid3 -> (vector<4xf32>) { + %a = vector.load %av[%t3, %kk] : view<[%tokens]x[%k]xf32> -> vector<4xf32> + %r = vector.fmaf %ww, %a, %acc3 : vector<4xf32> + scf.yield %r : vector<4xf32> + } else { + scf.yield %acc3 : vector<4xf32> + } + %t4 = index.constant 4 : index + %valid4 = index.cmp ult, %t4, %tokens : index + %next4 = scf.if %valid4 -> (vector<4xf32>) { + %a = vector.load %av[%t4, %kk] : view<[%tokens]x[%k]xf32> -> vector<4xf32> + %r = vector.fmaf %ww, %a, %acc4 : vector<4xf32> + scf.yield %r : vector<4xf32> + } else { + scf.yield %acc4 : vector<4xf32> + } + scf.yield %next0, %next1, %next2, %next3, %next4 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + %leader = index.cmp eq, %lane16, %c0 : index + %c5 = index.constant 5 : index + %stage_width = scf.select %partitioned, %c16, %c8 : index + %stage_elements = index.mul %stage_width, %c5 : index + %element_bytes = index.constant 4 : offset + %stage_bytes = index.scale %stage_elements, %element_bytes : index, offset -> offset + %stage = buffer.alloca align(16) %stage_bytes : buffer + %sv = buffer.view %stage[%base] : buffer -> view<5x[%stage_width]xf32> + %local_row0 = index.add %wave_add, %cohort : index + %partial_offset = index.mul %k_partition, %c4 : index + %local_row = index.add %local_row0, %partial_offset : index + %p0 = index.constant 0 : index + %publishes0 = index.cmp ult, %p0, %tokens : index + scf.if %publishes0 { + %sum8 = func.call @ggml_lowtoken_reduce_f32(%total0, %full_wave) : (vector<4xf32>, i1) -> (f32) + scf.if %leader { + view.store %sum8, %sv[%p0, %local_row] : f32, view<5x[%stage_width]xf32> + } + } + %p1 = index.constant 1 : index + %publishes1 = index.cmp ult, %p1, %tokens : index + scf.if %publishes1 { + %sum8 = func.call @ggml_lowtoken_reduce_f32(%total1, %full_wave) : (vector<4xf32>, i1) -> (f32) + scf.if %leader { + view.store %sum8, %sv[%p1, %local_row] : f32, view<5x[%stage_width]xf32> + } + } + %p2 = index.constant 2 : index + %publishes2 = index.cmp ult, %p2, %tokens : index + scf.if %publishes2 { + %sum8 = func.call @ggml_lowtoken_reduce_f32(%total2, %full_wave) : (vector<4xf32>, i1) -> (f32) + scf.if %leader { + view.store %sum8, %sv[%p2, %local_row] : f32, view<5x[%stage_width]xf32> + } + } + %p3 = index.constant 3 : index + %publishes3 = index.cmp ult, %p3, %tokens : index + scf.if %publishes3 { + %sum8 = func.call @ggml_lowtoken_reduce_f32(%total3, %full_wave) : (vector<4xf32>, i1) -> (f32) + scf.if %leader { + view.store %sum8, %sv[%p3, %local_row] : f32, view<5x[%stage_width]xf32> + } + } + %p4 = index.constant 4 : index + %publishes4 = index.cmp ult, %p4, %tokens : index + scf.if %publishes4 { + %sum8 = func.call @ggml_lowtoken_reduce_f32(%total4, %full_wave) : (vector<4xf32>, i1) -> (f32) + scf.if %leader { + view.store %sum8, %sv[%p4, %local_row] : f32, view<5x[%stage_width]xf32> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %workitem = kernel.workitem.id : index + %publish_packets = index.div %tile_rows, %c4 : index + %publish_token = index.div %workitem, %publish_packets : index + %publish_packet = index.rem %workitem, %publish_packets : index + %publish_add = index.mul %publish_packet, %c4 : index + %publish_channel = index.add %tile_base, %publish_add : index + %valid_token = index.cmp ult, %publish_token, %tokens : index + %valid_channel = index.cmp ult, %publish_channel, %n : index + %publishes = scalar.andi %valid_token, %valid_channel : i1 + scf.if %publishes { + %safe_token, %safe_count = index.assume %publish_token, %tokens [lt(%publish_token, %tokens)] : index, index + %stage_token = index.assume %safe_token [range(%safe_token, 0, 4)] : index + %stage_add = index.assume %publish_add [range(%publish_add, 0, 4), mul(%publish_add, 4)] : index + %stage_row_offset = index.mul %stage_token, %stage_width : index + %stage_origin0 = index.add %stage_row_offset, %stage_add : index + %stage_origin = scf.if %partitioned -> (index) { + %origin = index.assume %stage_origin0 [range(%stage_origin0, 0, 64), mul(%stage_origin0, 4)] : index + scf.yield %origin : index + } else { + %origin = index.assume %stage_origin0 [range(%stage_origin0, 0, 36), mul(%stage_origin0, 4)] : index + scf.yield %origin : index + } + %flat_stage = buffer.view %stage[%base] : buffer -> view<[%stage_elements]xf32> + %values0 = vector.load %flat_stage[%stage_origin] : view<[%stage_elements]xf32> -> vector<4xf32> + %values = scf.if %partitioned -> (vector<4xf32>) { + %stage_origin1 = index.add %stage_origin, %c4 : index + %stage_origin2 = index.add %stage_origin1, %c4 : index + %stage_origin3 = index.add %stage_origin2, %c4 : index + %values1 = vector.load %flat_stage[%stage_origin1] : view<[%stage_elements]xf32> -> vector<4xf32> + %values2 = vector.load %flat_stage[%stage_origin2] : view<[%stage_elements]xf32> -> vector<4xf32> + %values3 = vector.load %flat_stage[%stage_origin3] : view<[%stage_elements]xf32> -> vector<4xf32> + %sum01 = vector.addf %values0, %values1 : vector<4xf32> + %sum012 = vector.addf %sum01, %values2 : vector<4xf32> + %sum = vector.addf %sum012, %values3 : vector<4xf32> + scf.yield %sum : vector<4xf32> + } else { + scf.yield %values0 : vector<4xf32> + } + template.apply<@ggml.mul_mat_f32_f32_wmma.publish_vector4>(%publishes, %safe_token, %publish_channel, %safe_count, %n, %output_accumulation, %output_unary_op, %values, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (i1, index, index, index, index, index, index, vector<4xf32>, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + } + %row_rounding = index.sub %tile_rows, %c1 : index + %rounded_n = index.add %n, %row_rounding : index + %output_tiles = index.div %rounded_n, %tile_rows : index + template.apply<@ggml.mul_mat_f32_f32_wmma.finish_tile>(%tokens, %n, %output_tiles, %c1, %output_accumulation, %output_unary_op, %epsilon, %channel_tile, %c0, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + func.return +} + +// Copy the original Q8_1_x4 records for up to 2048 reduction elements. +// Call uniformly at a slab boundary; only the final slab may be partial. +func.def inline @ggml_q8_1_x4_stage_lowtoken_slab(%input: buffer, %stage: buffer, %capacity: index, %slab_begin: index, %blocks: index, %global_row_bytes: offset) { + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c144 = index.constant 144 : index + %c576 = index.constant 576 : index + %c18 = index.constant 18 : index + %base = index.constant 0 : offset + %packet_bytes = index.constant 16 : offset + %block_bytes = index.constant 288 : offset + %workgroup_size = kernel.workgroup.size : index + %thread = kernel.workitem.id : index + %stage_packets = index.mul %capacity, %c144 : index + %stage_elements = index.mul %capacity, %c576 : index + %stage_view = buffer.view %stage[%base] : buffer -> view<[%stage_elements]xi32> + %slab_limit = index.add %slab_begin, %c8 : index + %full_slab = index.cmp ule, %slab_limit, %blocks : index + %slab_end = scf.select %full_slab, %slab_limit, %blocks : index + %active_blocks0 = index.sub %slab_end, %slab_begin : index + %active_blocks = index.assume %active_blocks0 [range(%active_blocks0, 1, 8)] : index + %active_packets = index.mul %active_blocks, %c18 : index + %slab_byte_add = index.scale %slab_begin, %block_bytes : index, offset -> offset + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.for %copy_packet = [%thread to %stage_packets step %workgroup_size] { + %copy_token = index.div %copy_packet, %c144 : index + %copy_inner = index.rem %copy_packet, %c144 : index + %copy_valid = index.cmp ult, %copy_inner, %active_packets : index + scf.if %copy_valid { + %copy_row = index.scale %copy_token, %global_row_bytes : index, offset -> offset + %copy_bytes = index.scale %copy_inner, %packet_bytes : index, offset -> offset + %copy_base0 = index.add %copy_row, %slab_byte_add : offset + %copy_base = index.add %copy_base0, %copy_bytes : offset + %source = buffer.view %input[%copy_base] : buffer -> view<4xi32> + %values = vector.load %source[%c0] : view<4xi32> -> vector<4xi32> + %copy_origin = index.mul %copy_packet, %c4 : index + vector.store %values, %stage_view[%copy_origin] : vector<4xi32>, view<[%stage_elements]xi32> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + func.return +} + +func.def inline @ggml_mul_mat_q8_lowtoken_publish(%tokens: index, %n: index, %row_tile: index, %row_base: index, %local_row: index, %block_lane: index, %channel_tile: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %total_g0: f32, %total_g1: f32, %total_g2: f32, %total_g3: f32, %total_g4: f32, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %base = index.constant 0 : offset + %zero = scalar.constant 0.0 : f32 + %leader = index.cmp eq, %block_lane, %c0 : index + %width = scalar.constant 32 : i32 + %stage_rows = index.constant 5 : index + %stage_elements = index.mul %stage_rows, %row_tile : index + %element_bytes = index.constant 4 : offset + %stage_bytes = index.scale %stage_elements, %element_bytes : index, offset -> offset + %stage = buffer.alloca align(16) %stage_bytes : buffer + %sv = buffer.view %stage[%base] : buffer -> view<5x[%row_tile]xf32> + %p0 = index.constant 0 : index + %publishes0 = index.cmp ult, %p0, %tokens : index + scf.if %publishes0 { + %sum = scalar.addf %total_g0, %zero : f32 + %sum8 = func.call @ggml_lowtoken_reduce_cohort_f32(%sum) : (f32) -> (f32) + scf.if %leader { + view.store %sum8, %sv[%p0, %local_row] : f32, view<5x[%row_tile]xf32> + } + } + %p1 = index.constant 1 : index + %publishes1 = index.cmp ult, %p1, %tokens : index + scf.if %publishes1 { + %sum = scalar.addf %total_g1, %zero : f32 + %sum8 = func.call @ggml_lowtoken_reduce_cohort_f32(%sum) : (f32) -> (f32) + scf.if %leader { + view.store %sum8, %sv[%p1, %local_row] : f32, view<5x[%row_tile]xf32> + } + } + %p2 = index.constant 2 : index + %publishes2 = index.cmp ult, %p2, %tokens : index + scf.if %publishes2 { + %sum = scalar.addf %total_g2, %zero : f32 + %sum8 = func.call @ggml_lowtoken_reduce_cohort_f32(%sum) : (f32) -> (f32) + scf.if %leader { + view.store %sum8, %sv[%p2, %local_row] : f32, view<5x[%row_tile]xf32> + } + } + %p3 = index.constant 3 : index + %publishes3 = index.cmp ult, %p3, %tokens : index + scf.if %publishes3 { + %sum = scalar.addf %total_g3, %zero : f32 + %sum8 = func.call @ggml_lowtoken_reduce_cohort_f32(%sum) : (f32) -> (f32) + scf.if %leader { + view.store %sum8, %sv[%p3, %local_row] : f32, view<5x[%row_tile]xf32> + } + } + %p4 = index.constant 4 : index + %publishes4 = index.cmp ult, %p4, %tokens : index + scf.if %publishes4 { + %sum = scalar.addf %total_g4, %zero : f32 + %sum8 = func.call @ggml_lowtoken_reduce_cohort_f32(%sum) : (f32) -> (f32) + scf.if %leader { + view.store %sum8, %sv[%p4, %local_row] : f32, view<5x[%row_tile]xf32> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %workitem = kernel.workitem.id : index + %publish_packets = index.div %row_tile, %c4 : index + %publish_token = index.div %workitem, %publish_packets : index + %publish_packet = index.rem %workitem, %publish_packets : index + %publish_add = index.mul %publish_packet, %c4 : index + %publish_channel = index.add %row_base, %publish_add : index + %valid_token = index.cmp ult, %publish_token, %tokens : index + %valid_channel = index.cmp ult, %publish_channel, %n : index + %publishes = scalar.andi %valid_token, %valid_channel : i1 + scf.if %publishes { + %safe_token, %safe_count = index.assume %publish_token, %tokens [lt(%publish_token, %tokens)] : index, index + %stage_token = index.assume %safe_token [range(%safe_token, 0, 4)] : index + %row_limit = index.sub %row_tile, %c4 : index + %stage_add = index.assume %publish_add [range(%publish_add, 0, %row_limit), mul(%publish_add, 4)] : index + %stage_row_offset = index.mul %stage_token, %row_tile : index + %stage_origin0 = index.add %stage_row_offset, %stage_add : index + %stage_limit = index.sub %stage_elements, %c4 : index + %stage_origin = index.assume %stage_origin0 [range(%stage_origin0, 0, %stage_limit), mul(%stage_origin0, 4)] : index + %flat_stage = buffer.view %stage[%base] : buffer -> view<[%stage_elements]xf32> + %values = vector.load %flat_stage[%stage_origin] : view<[%stage_elements]xf32> -> vector<4xf32> + template.apply<@ggml.mul_mat_f32_f32_wmma.publish_vector4>(%publishes, %safe_token, %publish_channel, %safe_count, %n, %output_accumulation, %output_unary_op, %values, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (i1, index, index, index, index, index, index, vector<4xf32>, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + } + %rounding = index.sub %row_tile, %c1 : index + %rounded_n = index.add %n, %rounding : index + %output_tiles = index.div %rounded_n, %row_tile : index + template.apply<@ggml.mul_mat_f32_f32_wmma.finish_tile>(%tokens, %n, %output_tiles, %c1, %output_accumulation, %output_unary_op, %epsilon, %channel_tile, %c0, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + func.return +} + +func.def inline @ggml_q4_prefetch_native_words(%weight: buffer, %row_group_block: index, %block: index, %row_lane: index, %payload_field: index, %payload_word: index) -> (vector<4xi32>, vector<2xi32>) { + %zero = index.constant 0 : index + %record_bytes = index.constant 9216 : offset + %row_field_bytes = index.constant 16 : offset + %payload_add = index.constant 1024 : offset + %group_block = index.add %row_group_block, %block : index + %group_base = index.scale %group_block, %record_bytes : index, offset -> offset + %row_offset = index.scale %row_lane, %row_field_bytes : index, offset -> offset + %wb = index.add %group_base, %row_offset : offset + %wc = index.add %group_base, %payload_add : offset + %header_view = buffer.view %weight[%wb] : buffer -> view<4xi32> + %code_view = buffer.view %weight[%wc] : buffer -> view<8x64x4xi32> + %headers = vector.load %header_view[%zero] : view<4xi32> -> vector<4xi32> + %raw = vector.load %code_view[%payload_field, %row_lane, %payload_word] : view<8x64x4xi32> -> vector<2xi32> + func.return %headers, %raw : vector<4xi32>, vector<2xi32> +} + +// Native Q4_K bytes and Q8_1_x4 activations, with one 16-lane cohort per row. +// Retain the wrapper publication hooks so bias, residual and RMSNorm fusions +// share this contraction without a separate dispatch or intermediate output. +func.def pure inline @ggml_q4_q8_lowtoken_token_pairs(%capacity: index, %tokens: index, %k: index, %n: index) -> (i1) { + %c4 = index.constant 4 : index + %c4096 = index.constant 4096 : index + %capacity4 = index.cmp eq, %capacity, %c4 : index + %tokens4 = index.cmp eq, %tokens, %c4 : index + %four_tokens = scalar.andi %capacity4, %tokens4 : i1 + %long_input = index.cmp uge, %k, %c4096 : index + %wide_output = index.cmp uge, %n, %c4096 : index + %contracts = index.cmp ule, %n, %k : index + %wide = scalar.andi %long_input, %wide_output : i1 + %contracting = scalar.andi %wide, %contracts : i1 + %eligible = scalar.andi %four_tokens, %contracting : i1 + func.return %eligible : i1 +} + +func.def inline @ggml_q4_q8_lowtoken_adjacent_scale_min(%scale0: i32, %scale1: i32, %scale2: i32, %pair0: index) -> (i32, i32, i32, i32) { + %pair = index.assume %pair0 [range(%pair0, 0, 3)] : index + %two = index.constant 2 : index + %sixteen = index.constant 16 : index + %low = index.cmp ult, %pair, %two : index + %pair_lane = index.rem %pair, %two : index + %shift_index = index.mul %pair_lane, %sixteen : index + %shift = index.cast %shift_index : index to i32 + %two_i32 = scalar.constant 2 : i32 + %four_i32 = scalar.constant 4 : i32 + %eight_i32 = scalar.constant 8 : i32 + %low_mask = scalar.constant 3855 : i32 + %high_mask = scalar.constant 12336 : i32 + %byte_mask = scalar.constant 255 : i32 + %high_shift = scalar.addi %shift, %two_i32 : i32 + %min_shift = scalar.addi %shift, %four_i32 : i32 + %s_source = scf.select %low, %scale0, %scale2 : i32 + %m_source = scf.select %low, %scale1, %scale2 : i32 + %s_high_shift = scf.select %low, %shift, %high_shift : i32 + %m_low_shift = scf.select %low, %shift, %min_shift : i32 + %s_lo_raw = scalar.shrui %s_source, %shift : i32 + %s_lo = scalar.andi %s_lo_raw, %low_mask : i32 + %s_hi_raw = scalar.shrui %scale0, %s_high_shift : i32 + %s_hi = scalar.andi %s_hi_raw, %high_mask : i32 + %s = scalar.ori %s_lo, %s_hi : i32 + %m_lo_raw = scalar.shrui %m_source, %m_low_shift : i32 + %m_lo = scalar.andi %m_lo_raw, %low_mask : i32 + %m_hi_raw = scalar.shrui %scale1, %s_high_shift : i32 + %m_hi = scalar.andi %m_hi_raw, %high_mask : i32 + %m = scalar.ori %m_lo, %m_hi : i32 + %s0 = scalar.andi %s, %byte_mask : i32 + %m0 = scalar.andi %m, %byte_mask : i32 + %s1 = scalar.shrui %s, %eight_i32 : i32 + %m1 = scalar.shrui %m, %eight_i32 : i32 + func.return %s0, %m0, %s1, %m1 : i32, i32, i32, i32 +} + +func.def inline @ggml_q4_q8_lowtoken_dot8(%paired_tokens: i1, %weight: vector<2xi32>, %activation: vector<2xi32>) -> (f32) { + %value = scf.if %paired_tokens -> (f32) { + %zero = vector.constant 0 : vector<1xi32> + %wi0 = vector.extract %weight[0] : vector<2xi32> -> i32 + %ai0 = vector.extract %activation[0] : vector<2xi32> -> i32 + %ww0 = vector.splat %wi0 : vector<1xi32> + %aw0 = vector.splat %ai0 : vector<1xi32> + %w0 = vector.bitcast %ww0 : vector<1xi32> to vector<4xi8> + %a0 = vector.bitcast %aw0 : vector<1xi32> to vector<4xi8> + %d0 = vector.dot4i %w0, %a0, %zero : vector<4xi8>, vector<4xi8>, vector<1xi32> + %wi1 = vector.extract %weight[1] : vector<2xi32> -> i32 + %ai1 = vector.extract %activation[1] : vector<2xi32> -> i32 + %ww1 = vector.splat %wi1 : vector<1xi32> + %aw1 = vector.splat %ai1 : vector<1xi32> + %w1 = vector.bitcast %ww1 : vector<1xi32> to vector<4xi8> + %a1 = vector.bitcast %aw1 : vector<1xi32> to vector<4xi8> + %d1 = vector.dot4i %w1, %a1, %d0 : vector<4xi8>, vector<4xi8>, vector<1xi32> + %sum = vector.extract %d1[0] : vector<1xi32> -> i32 + %result = scalar.sitofp %sum : i32 to f32 + scf.yield %result : f32 + } else { + %original = func.call @ggml_dot_u8_s8_vector8_f32(%weight, %activation) : (vector<2xi32>, vector<2xi32>) -> (f32) + scf.yield %original : f32 + } + func.return %value : f32 +} + +func.def inline @ggml_mul_mat_q4_q8_1_x4_lowtoken_dot(%weight_format: index, %token_count: index, %input_size0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %input: buffer, %weight: buffer, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + + %capacity = config.get @ggml.workload.token_capacity : index + %staged_rows, %k_partitions = func.call pure @ggml_q8_lowtoken_geometry(%capacity, %input_size0, %output_size0, %weight_format) : (index, index, index, index) -> (index, index) + %no_staged_rows = index.constant 0 : index + %staged = index.cmp ugt, %staged_rows, %no_staged_rows : index + %tokens = index.assume %token_count [range(%token_count, 1, 5), le(%token_count, %capacity)] : index + %k = index.assume %input_size0 [range(%input_size0, 256, 32768), mul(%input_size0, 256)] : index + %n = index.assume %output_size0 [range(%output_size0, 64, 262144), mul(%output_size0, 64)] : index + %paired_tokens = func.call pure @ggml_q4_q8_lowtoken_token_pairs(%capacity, %tokens, %k, %n) : (index, index, index, index) -> (i1) + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %base = index.constant 0 : offset + %block_bytes = index.constant 144 : offset + %code_offset = index.constant 16 : offset + %zero = scalar.constant 0.0 : f32 + %lane0 = kernel.subgroup.lane.id : index + %lane = index.assume %lane0 [range(%lane0, 0, 63)] : index + %wave = kernel.subgroup.id : index + %wg = index.add %channel_tile, %c0 : index + %c64_rows = index.constant 64 : index + %direct_rows = index.div %c8, %k_partitions : index + %row_tile = scf.select %staged, %c64_rows, %direct_rows : index + %row_base = index.mul %wg, %row_tile : index + %wave_rows = index.div %c4, %k_partitions : index + %wave_row = index.mul %wave, %wave_rows : index + %cohort_width = index.mul %c16, %k_partitions : index + %row_half = index.div %lane, %cohort_width : index + %cohort_lane = index.rem %lane, %cohort_width : index + %local_row = index.add %wave_row, %row_half : index + %row = index.add %row_base, %local_row : index + %valid_row = index.cmp ult, %row, %n : index + %cohort = index.div %lane, %c16 : index + %block_lane = index.rem %lane, %c16 : index + %pair = index.div %block_lane, %c4 : index + %packet = index.rem %block_lane, %c4 : index + %group0 = index.mul %pair, %c2 : index + %group1 = index.add %group0, %c1 : index + %word_page = index.mul %pair, %c8 : index + %word_add = index.mul %packet, %c2 : index + %word_index = index.add %word_page, %word_add : index + %meta_lane = index.cmp eq, %packet, %c0 : index + %blocks = index.div %k, %c256 : index + %partition_blocks = index.div %blocks, %k_partitions : index + %k_team = index.div %cohort_lane, %c16 : index + %k_begin = index.mul %k_team, %partition_blocks : index + %weight_row_bytes = index.scale %blocks, %block_bytes : index, offset -> offset + %weight_row_base = index.scale %row, %weight_row_bytes : index, offset -> offset + %q8_groups = index.div %k, %c128 : index + %global_q8_row_bytes = index.scale %q8_groups, %block_bytes : index, offset -> offset + %slab_row_bytes = index.constant 2304 : offset + %q8_row_bytes = scf.select %staged, %slab_row_bytes, %global_q8_row_bytes : offset + %q8_half = index.div %pair, %c2 : index + %q8_pair = index.rem %pair, %c2 : index + %q8_group0 = index.mul %q8_pair, %c2 : index + %q8_group1 = index.add %q8_group0, %c1 : index + %q8_words_base = index.mul %q8_group0, %c8 : index + %q8_word0 = index.add %q8_words_base, %word_add : index + %q8_word1 = index.add %q8_word0, %c8 : index + %q8_meta0 = index.mul %q8_group0, %c2 : index + %q8_meta1 = index.mul %q8_group1, %c2 : index + %q8_sum0 = index.add %q8_meta0, %c1 : index + %q8_sum1 = index.add %q8_meta1, %c1 : index + %c64 = index.constant 64 : index + %record_bytes = index.constant 9216 : offset + %row_field_bytes = index.constant 16 : offset + %payload_add = index.constant 1024 : offset + %direct_row_group = index.div %row, %c64 : index + %row_group = scf.select %staged, %channel_tile, %direct_row_group : index + %row_lane = index.rem %row, %c64 : index + %row_group_block = index.mul %row_group, %blocks : index + %row_offset = index.scale %row_lane, %row_field_bytes : index, offset -> offset + %payload_field = index.div %word_index, %c4 : index + %payload_word = index.rem %word_index, %c4 : index + %input_na, %gate_na = buffer.assume.noalias %input, %weight : buffer, buffer + %slab_blocks = index.constant 8 : index + %full_stage_bytes = index.constant 11520 : offset + %q8_stage_bytes = scf.select %staged, %full_stage_bytes, %base : offset + %q8_stage = buffer.alloca align(16) %q8_stage_bytes : buffer + %q8_input = scf.select %staged, %q8_stage, %input_na : buffer + %zero_header = vector.constant 0 : vector<4xi32> + %zero_raw = vector.constant 0 : vector<2xi32> + %prime_header, %prime_raw = scf.if %staged -> (vector<4xi32>, vector<2xi32>) { + %header, %raw = func.call @ggml_q4_prefetch_native_words(%gate_na, %row_group_block, %c0, %row_lane, %payload_field, %payload_word) : (buffer, index, index, index, index, index) -> (vector<4xi32>, vector<2xi32>) + scf.yield %header, %raw : vector<4xi32>, vector<2xi32> + } else { + scf.yield %zero_header, %zero_raw : vector<4xi32>, vector<2xi32> + } + %total_g0, %total_g1, %total_g2, %total_g3, %total_g4, %last_header, %last_raw = scf.for %block_base = [%c0 to %partition_blocks step %c1](%acc_g0 = %zero : f32, %acc_g1 = %zero : f32, %acc_g2 = %zero : f32, %acc_g3 = %zero : f32, %acc_g4 = %zero : f32, %carry_header = %prime_header : vector<4xi32>, %carry_raw = %prime_raw : vector<2xi32>) -> (f32, f32, f32, f32, f32, vector<4xi32>, vector<2xi32>) { + %block_in_slab = index.rem %block_base, %slab_blocks : index + %slab_boundary = index.cmp eq, %block_in_slab, %c0 : index + %new_slab = scalar.andi %staged, %slab_boundary : i1 + scf.if %new_slab { + func.call @ggml_q8_1_x4_stage_lowtoken_slab(%input_na, %q8_stage, %capacity, %block_base, %blocks, %global_q8_row_bytes) : (buffer, buffer, index, index, index, offset) + } + %block = index.add %block_base, %k_begin : index + %valid_block = index.cmp ult, %block, %blocks : index + %safe_block = scf.select %valid_block, %block, %c0 : index + %group_block = index.add %row_group_block, %safe_block : index + %group_base = index.scale %group_block, %record_bytes : index, offset -> offset + %wb = index.add %group_base, %row_offset : offset + %wc = index.add %group_base, %payload_add : offset + %g_header_view = buffer.view %gate_na[%wb] : buffer -> view<4xi32> + %g_code_view = buffer.view %gate_na[%wc] : buffer -> view<8x64x4xi32> + %g_headers = scf.if %staged -> (vector<4xi32>) { + scf.yield %carry_header : vector<4xi32> + } else { + %header = vector.load %g_header_view[%c0] : view<4xi32> -> vector<4xi32> + scf.yield %header : vector<4xi32> + } + %g_dm_i32 = vector.extract %g_headers[0] : vector<4xi32> -> i32 + %g_dm_word = vector.splat %g_dm_i32 : vector<1xi32> + %g_dm = vector.bitcast %g_dm_word : vector<1xi32> to vector<2xf16> + %g_d0h = vector.extract %g_dm[0] : vector<2xf16> -> f16 + %g_d1h = vector.extract %g_dm[1] : vector<2xf16> -> f16 + %g_d0 = scalar.extf %g_d0h : f16 to f32 + %g_d1 = scalar.extf %g_d1h : f16 to f32 + %g_s0 = vector.extract %g_headers[1] : vector<4xi32> -> i32 + %g_s1 = vector.extract %g_headers[2] : vector<4xi32> -> i32 + %g_s2 = vector.extract %g_headers[3] : vector<4xi32> -> i32 + %pair_scale0, %pair_min0, %pair_scale1, %pair_min1 = scf.if %paired_tokens -> (i32, i32, i32, i32) { + %s0, %m0, %s1, %m1 = func.call @ggml_q4_q8_lowtoken_adjacent_scale_min(%g_s0, %g_s1, %g_s2, %pair) : (i32, i32, i32, index) -> (i32, i32, i32, i32) + scf.yield %s0, %m0, %s1, %m1 : i32, i32, i32, i32 + } else { + %zero_i32 = scalar.constant 0 : i32 + scf.yield %zero_i32, %zero_i32, %zero_i32, %zero_i32 : i32, i32, i32, i32 + } + %g_scale0i, %g_min0i = scf.if %paired_tokens -> (i32, i32) { + scf.yield %pair_scale0, %pair_min0 : i32, i32 + } else { + %scale, %minimum = func.call @ggml_q4k_scale_min_from_header(%g_s0, %g_s1, %g_s2, %group0) : (i32, i32, i32, index) -> (i32, i32) + scf.yield %scale, %minimum : i32, i32 + } + %g_scale0f = scalar.uitofp %g_scale0i : i32 to f32 + %g_min0f = scalar.uitofp %g_min0i : i32 to f32 + %g_ws0 = scalar.mulf %g_d0, %g_scale0f : f32 + %g_wm0 = scalar.mulf %g_d1, %g_min0f : f32 + %g_scale1i, %g_min1i = scf.if %paired_tokens -> (i32, i32) { + scf.yield %pair_scale1, %pair_min1 : i32, i32 + } else { + %scale, %minimum = func.call @ggml_q4k_scale_min_from_header(%g_s0, %g_s1, %g_s2, %group1) : (i32, i32, i32, index) -> (i32, i32) + scf.yield %scale, %minimum : i32, i32 + } + %g_scale1f = scalar.uitofp %g_scale1i : i32 to f32 + %g_min1f = scalar.uitofp %g_min1i : i32 to f32 + %g_ws1 = scalar.mulf %g_d0, %g_scale1f : f32 + %g_wm1 = scalar.mulf %g_d1, %g_min1f : f32 + %g_raw = scf.if %staged -> (vector<2xi32>) { + scf.yield %carry_raw : vector<2xi32> + } else { + %raw = vector.load %g_code_view[%payload_field, %row_lane, %payload_word] : view<8x64x4xi32> -> vector<2xi32> + scf.yield %raw : vector<2xi32> + } + %g_mask = vector.constant 252645135 : vector<2xi32> + %g_shift = vector.constant 4 : vector<2xi32> + %g_lo = vector.andi %g_raw, %g_mask : vector<2xi32> + %g_high = vector.shrui %g_raw, %g_shift : vector<2xi32> + %g_hi = vector.andi %g_high, %g_mask : vector<2xi32> + %next_block = index.add %block, %c1 : index + %next_exists = index.cmp ult, %next_block, %blocks : index + %has_next = scalar.andi %staged, %next_exists : i1 + %next_header, %next_raw = scf.if %has_next -> (vector<4xi32>, vector<2xi32>) { + %header, %raw = func.call @ggml_q4_prefetch_native_words(%gate_na, %row_group_block, %next_block, %row_lane, %payload_field, %payload_word) : (buffer, index, index, index, index, index) -> (vector<4xi32>, vector<2xi32>) + scf.yield %header, %raw : vector<4xi32>, vector<2xi32> + } else { + scf.yield %carry_header, %carry_raw : vector<4xi32>, vector<2xi32> + } + %input_block = scf.select %staged, %block_in_slab, %safe_block : index + %q8_block0 = index.mul %input_block, %c2 : index + %q8_block = index.add %q8_block0, %q8_half : index + %q8_block_add = index.scale %q8_block, %block_bytes : index, offset -> offset + %t0 = index.constant 0 : index + %valid0 = index.cmp ult, %t0, %tokens : index + %next_g0 = scf.if %valid0 -> (f32) { + %tb = index.scale %t0, %q8_row_bytes : index, offset -> offset + %qb = index.add %tb, %q8_block_add : offset + %qc = index.add %qb, %code_offset : offset + %q_view = buffer.view %q8_input[%qc] : buffer -> view<32xi32> + %ds_view = buffer.view %q8_input[%qb] : buffer -> view<8xf16> + %a0 = vector.load %q_view[%q8_word0] : view<32xi32> -> vector<2xi32> + %a1 = vector.load %q_view[%q8_word1] : view<32xi32> -> vector<2xi32> + %ad0h = view.load %ds_view[%q8_meta0] : view<8xf16> -> f16 + %ad1h = view.load %ds_view[%q8_meta1] : view<8xf16> -> f16 + %ad0 = scalar.extf %ad0h : f16 to f32 + %ad1 = scalar.extf %ad1h : f16 to f32 + %as0h = view.load %ds_view[%q8_sum0] : view<8xf16> -> f16 + %as1h = view.load %ds_view[%q8_sum1] : view<8xf16> -> f16 + %as0_raw = scalar.extf %as0h : f16 to f32 + %as1_raw = scalar.extf %as1h : f16 to f32 + %as0 = scf.select %meta_lane, %as0_raw, %zero : f32 + %as1 = scf.select %meta_lane, %as1_raw, %zero : f32 + %g_dot0 = func.call @ggml_q4_q8_lowtoken_dot8(%paired_tokens, %g_lo, %a0) : (i1, vector<2xi32>, vector<2xi32>) -> (f32) + %g_dot1 = func.call @ggml_q4_q8_lowtoken_dot8(%paired_tokens, %g_hi, %a1) : (i1, vector<2xi32>, vector<2xi32>) -> (f32) + %g_s0 = scalar.mulf %g_ws0, %ad0 : f32 + %g_s1 = scalar.mulf %g_ws1, %ad1 : f32 + %g_c0 = scalar.mulf %g_dot0, %g_s0 : f32 + %g_contribution = scalar.fmaf %g_dot1, %g_s1, %g_c0 : f32 + %g_cor0 = scalar.mulf %g_wm0, %as0 : f32 + %g_cor = scalar.fmaf %g_wm1, %as1, %g_cor0 : f32 + %g_corrected = scalar.subf %g_contribution, %g_cor : f32 + %g_valid = scf.select %valid_block, %g_corrected, %zero : f32 + %g_updated = scalar.addf %acc_g0, %g_valid : f32 + scf.yield %g_updated : f32 + } else { + scf.yield %acc_g0 : f32 + } + %t1 = index.constant 1 : index + %valid1 = index.cmp ult, %t1, %tokens : index + %next_g1 = scf.if %valid1 -> (f32) { + %tb = index.scale %t1, %q8_row_bytes : index, offset -> offset + %qb = index.add %tb, %q8_block_add : offset + %qc = index.add %qb, %code_offset : offset + %q_view = buffer.view %q8_input[%qc] : buffer -> view<32xi32> + %ds_view = buffer.view %q8_input[%qb] : buffer -> view<8xf16> + %a0 = vector.load %q_view[%q8_word0] : view<32xi32> -> vector<2xi32> + %a1 = vector.load %q_view[%q8_word1] : view<32xi32> -> vector<2xi32> + %ds = vector.load %ds_view[%q8_meta0] : view<8xf16> -> vector<4xf16> + %ad0h = vector.extract %ds[0] : vector<4xf16> -> f16 + %ad1h = vector.extract %ds[2] : vector<4xf16> -> f16 + %ad0 = scalar.extf %ad0h : f16 to f32 + %ad1 = scalar.extf %ad1h : f16 to f32 + %as0h = vector.extract %ds[1] : vector<4xf16> -> f16 + %as1h = vector.extract %ds[3] : vector<4xf16> -> f16 + %as0_raw = scalar.extf %as0h : f16 to f32 + %as1_raw = scalar.extf %as1h : f16 to f32 + %as0 = scf.select %meta_lane, %as0_raw, %zero : f32 + %as1 = scf.select %meta_lane, %as1_raw, %zero : f32 + %g_dot0 = func.call @ggml_q4_q8_lowtoken_dot8(%paired_tokens, %g_lo, %a0) : (i1, vector<2xi32>, vector<2xi32>) -> (f32) + %g_dot1 = func.call @ggml_q4_q8_lowtoken_dot8(%paired_tokens, %g_hi, %a1) : (i1, vector<2xi32>, vector<2xi32>) -> (f32) + %g_s0 = scalar.mulf %g_ws0, %ad0 : f32 + %g_s1 = scalar.mulf %g_ws1, %ad1 : f32 + %g_c0 = scalar.mulf %g_dot0, %g_s0 : f32 + %g_contribution = scalar.fmaf %g_dot1, %g_s1, %g_c0 : f32 + %g_cor0 = scalar.mulf %g_wm0, %as0 : f32 + %g_cor = scalar.fmaf %g_wm1, %as1, %g_cor0 : f32 + %g_corrected = scalar.subf %g_contribution, %g_cor : f32 + %g_valid = scf.select %valid_block, %g_corrected, %zero : f32 + %g_updated = scalar.addf %acc_g1, %g_valid : f32 + scf.yield %g_updated : f32 + } else { + scf.yield %acc_g1 : f32 + } + scf.if %paired_tokens { + scf.schedule.fence + } + %t2 = index.constant 2 : index + %valid2 = index.cmp ult, %t2, %tokens : index + %next_g2 = scf.if %valid2 -> (f32) { + %tb = index.scale %t2, %q8_row_bytes : index, offset -> offset + %qb = index.add %tb, %q8_block_add : offset + %qc = index.add %qb, %code_offset : offset + %q_view = buffer.view %q8_input[%qc] : buffer -> view<32xi32> + %ds_view = buffer.view %q8_input[%qb] : buffer -> view<8xf16> + %a0 = vector.load %q_view[%q8_word0] : view<32xi32> -> vector<2xi32> + %a1 = vector.load %q_view[%q8_word1] : view<32xi32> -> vector<2xi32> + %ds = vector.load %ds_view[%q8_meta0] : view<8xf16> -> vector<4xf16> + %ad0h = vector.extract %ds[0] : vector<4xf16> -> f16 + %ad1h = vector.extract %ds[2] : vector<4xf16> -> f16 + %ad0 = scalar.extf %ad0h : f16 to f32 + %ad1 = scalar.extf %ad1h : f16 to f32 + %as0h = vector.extract %ds[1] : vector<4xf16> -> f16 + %as1h = vector.extract %ds[3] : vector<4xf16> -> f16 + %as0_raw = scalar.extf %as0h : f16 to f32 + %as1_raw = scalar.extf %as1h : f16 to f32 + %as0 = scf.select %meta_lane, %as0_raw, %zero : f32 + %as1 = scf.select %meta_lane, %as1_raw, %zero : f32 + %g_dot0 = func.call @ggml_q4_q8_lowtoken_dot8(%paired_tokens, %g_lo, %a0) : (i1, vector<2xi32>, vector<2xi32>) -> (f32) + %g_dot1 = func.call @ggml_q4_q8_lowtoken_dot8(%paired_tokens, %g_hi, %a1) : (i1, vector<2xi32>, vector<2xi32>) -> (f32) + %g_s0 = scalar.mulf %g_ws0, %ad0 : f32 + %g_s1 = scalar.mulf %g_ws1, %ad1 : f32 + %g_c0 = scalar.mulf %g_dot0, %g_s0 : f32 + %g_contribution = scalar.fmaf %g_dot1, %g_s1, %g_c0 : f32 + %g_cor0 = scalar.mulf %g_wm0, %as0 : f32 + %g_cor = scalar.fmaf %g_wm1, %as1, %g_cor0 : f32 + %g_corrected = scalar.subf %g_contribution, %g_cor : f32 + %g_valid = scf.select %valid_block, %g_corrected, %zero : f32 + %g_updated = scalar.addf %acc_g2, %g_valid : f32 + scf.yield %g_updated : f32 + } else { + scf.yield %acc_g2 : f32 + } + %t3 = index.constant 3 : index + %valid3 = index.cmp ult, %t3, %tokens : index + %next_g3 = scf.if %valid3 -> (f32) { + %tb = index.scale %t3, %q8_row_bytes : index, offset -> offset + %qb = index.add %tb, %q8_block_add : offset + %qc = index.add %qb, %code_offset : offset + %q_view = buffer.view %q8_input[%qc] : buffer -> view<32xi32> + %ds_view = buffer.view %q8_input[%qb] : buffer -> view<8xf16> + %a0 = vector.load %q_view[%q8_word0] : view<32xi32> -> vector<2xi32> + %a1 = vector.load %q_view[%q8_word1] : view<32xi32> -> vector<2xi32> + %ds = vector.load %ds_view[%q8_meta0] : view<8xf16> -> vector<4xf16> + %ad0h = vector.extract %ds[0] : vector<4xf16> -> f16 + %ad1h = vector.extract %ds[2] : vector<4xf16> -> f16 + %ad0 = scalar.extf %ad0h : f16 to f32 + %ad1 = scalar.extf %ad1h : f16 to f32 + %as0h = vector.extract %ds[1] : vector<4xf16> -> f16 + %as1h = vector.extract %ds[3] : vector<4xf16> -> f16 + %as0_raw = scalar.extf %as0h : f16 to f32 + %as1_raw = scalar.extf %as1h : f16 to f32 + %as0 = scf.select %meta_lane, %as0_raw, %zero : f32 + %as1 = scf.select %meta_lane, %as1_raw, %zero : f32 + %g_dot0 = func.call @ggml_q4_q8_lowtoken_dot8(%paired_tokens, %g_lo, %a0) : (i1, vector<2xi32>, vector<2xi32>) -> (f32) + %g_dot1 = func.call @ggml_q4_q8_lowtoken_dot8(%paired_tokens, %g_hi, %a1) : (i1, vector<2xi32>, vector<2xi32>) -> (f32) + %g_s0 = scalar.mulf %g_ws0, %ad0 : f32 + %g_s1 = scalar.mulf %g_ws1, %ad1 : f32 + %g_c0 = scalar.mulf %g_dot0, %g_s0 : f32 + %g_contribution = scalar.fmaf %g_dot1, %g_s1, %g_c0 : f32 + %g_cor0 = scalar.mulf %g_wm0, %as0 : f32 + %g_cor = scalar.fmaf %g_wm1, %as1, %g_cor0 : f32 + %g_corrected = scalar.subf %g_contribution, %g_cor : f32 + %g_valid = scf.select %valid_block, %g_corrected, %zero : f32 + %g_updated = scalar.addf %acc_g3, %g_valid : f32 + scf.yield %g_updated : f32 + } else { + scf.yield %acc_g3 : f32 + } + %t4 = index.constant 4 : index + %valid4 = index.cmp ult, %t4, %tokens : index + %next_g4 = scf.if %valid4 -> (f32) { + %tb = index.scale %t4, %q8_row_bytes : index, offset -> offset + %qb = index.add %tb, %q8_block_add : offset + %qc = index.add %qb, %code_offset : offset + %q_view = buffer.view %q8_input[%qc] : buffer -> view<32xi32> + %ds_view = buffer.view %q8_input[%qb] : buffer -> view<8xf16> + %a0 = vector.load %q_view[%q8_word0] : view<32xi32> -> vector<2xi32> + %a1 = vector.load %q_view[%q8_word1] : view<32xi32> -> vector<2xi32> + %ds = vector.load %ds_view[%q8_meta0] : view<8xf16> -> vector<4xf16> + %ad0h = vector.extract %ds[0] : vector<4xf16> -> f16 + %ad1h = vector.extract %ds[2] : vector<4xf16> -> f16 + %ad0 = scalar.extf %ad0h : f16 to f32 + %ad1 = scalar.extf %ad1h : f16 to f32 + %as0h = vector.extract %ds[1] : vector<4xf16> -> f16 + %as1h = vector.extract %ds[3] : vector<4xf16> -> f16 + %as0_raw = scalar.extf %as0h : f16 to f32 + %as1_raw = scalar.extf %as1h : f16 to f32 + %as0 = scf.select %meta_lane, %as0_raw, %zero : f32 + %as1 = scf.select %meta_lane, %as1_raw, %zero : f32 + %g_dot0 = func.call @ggml_q4_q8_lowtoken_dot8(%paired_tokens, %g_lo, %a0) : (i1, vector<2xi32>, vector<2xi32>) -> (f32) + %g_dot1 = func.call @ggml_q4_q8_lowtoken_dot8(%paired_tokens, %g_hi, %a1) : (i1, vector<2xi32>, vector<2xi32>) -> (f32) + %g_s0 = scalar.mulf %g_ws0, %ad0 : f32 + %g_s1 = scalar.mulf %g_ws1, %ad1 : f32 + %g_c0 = scalar.mulf %g_dot0, %g_s0 : f32 + %g_contribution = scalar.fmaf %g_dot1, %g_s1, %g_c0 : f32 + %g_cor0 = scalar.mulf %g_wm0, %as0 : f32 + %g_cor = scalar.fmaf %g_wm1, %as1, %g_cor0 : f32 + %g_corrected = scalar.subf %g_contribution, %g_cor : f32 + %g_valid = scf.select %valid_block, %g_corrected, %zero : f32 + %g_updated = scalar.addf %acc_g4, %g_valid : f32 + scf.yield %g_updated : f32 + } else { + scf.yield %acc_g4 : f32 + } + scf.yield %next_g0, %next_g1, %next_g2, %next_g3, %next_g4, %next_header, %next_raw : f32, f32, f32, f32, f32, vector<4xi32>, vector<2xi32> + } + %combined0 = func.call @ggml_lowtoken_combine_k_partitions(%total_g0, %k_partitions) : (f32, index) -> (f32) + %combined1 = func.call @ggml_lowtoken_combine_k_partitions(%total_g1, %k_partitions) : (f32, index) -> (f32) + %combined2 = func.call @ggml_lowtoken_combine_k_partitions(%total_g2, %k_partitions) : (f32, index) -> (f32) + %combined3 = func.call @ggml_lowtoken_combine_k_partitions(%total_g3, %k_partitions) : (f32, index) -> (f32) + %combined4 = func.call @ggml_lowtoken_combine_k_partitions(%total_g4, %k_partitions) : (f32, index) -> (f32) + func.call @ggml_mul_mat_q8_lowtoken_publish(%tokens, %n, %row_tile, %row_base, %local_row, %cohort_lane, %channel_tile, %output_accumulation, %output_unary_op, %epsilon, %combined0, %combined1, %combined2, %combined3, %combined4, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (index, index, index, index, index, index, index, index, index, f32, f32, f32, f32, f32, f32, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + func.return +} + +// Native Q6_K carrier with signed I8 dot products and F32 scaling. +func.def inline @ggml_q6_prefetch_native_words(%weight_na: buffer, %row_group_block: index, %block: index, %row_lane: index, %half16: index, %group: index, %half128: index, %ql_side: index, %packet: index) -> (f16, i8, vector<4xi32>, vector<4xi32>) { + %c2 = index.constant 2 : index + %record_bytes = index.constant 13440 : offset + %scale_add = index.constant 128 : offset + %raw_add = index.constant 1152 : offset + %group_block = index.add %row_group_block, %block : index + %block_base = index.scale %group_block, %record_bytes : index, offset -> offset + %scale_base = index.add %block_base, %scale_add : offset + %raw_base = index.add %block_base, %raw_add : offset + %d_view = buffer.view %weight_na[%block_base] : buffer -> view<64xf16> + %s_view = buffer.view %weight_na[%scale_base] : buffer -> view<64x2x8xi8> + %raw_view = buffer.view %weight_na[%raw_base] : buffer -> view<2x3x64x8xi32> + %dh = view.load %d_view[%row_lane] : view<64xf16> -> f16 + %si = view.load %s_view[%row_lane, %half16, %group] : view<64x2x8xi8> -> i8 + %ql = vector.load %raw_view[%half128, %ql_side, %row_lane, %packet] : view<2x3x64x8xi32> -> vector<4xi32> + %qh = vector.load %raw_view[%half128, %c2, %row_lane, %packet] : view<2x3x64x8xi32> -> vector<4xi32> + func.return %dh, %si, %ql, %qh : f16, i8, vector<4xi32>, vector<4xi32> +} + +func.def inline @ggml_q6_q8_lowtoken_activation_values(%pipelined: i1, %token: index, %q8_row_bytes: offset, %q8_block_add: offset, %capacity: index, %current_slot: index, %q8_block0: index, %q8_input: buffer, %q8_word0: index, %q8_meta: index) -> (vector<16xi8>, f32) { + %code_offset = index.constant 16 : offset + %a, %ad = scf.if %pipelined -> (vector<16xi8>, f32) { + %base = index.constant 0 : offset + %c4 = index.constant 4 : index + %q8_block = index.assume %q8_block0 [range(%q8_block0, 0, 7)] : index + %word = index.add %q8_word0, %c4 : index + %q_view = buffer.view %q8_input[%base] : buffer -> view<2x[%capacity]x8x36xi32> + %ds_view = buffer.view %q8_input[%base] : buffer -> view<2x[%capacity]x8x72xf16> + %words = vector.load %q_view[%current_slot, %token, %q8_block, %word] : view<2x[%capacity]x8x36xi32> -> vector<4xi32> + %a = vector.bitcast %words : vector<4xi32> to vector<16xi8> + %scale = view.load %ds_view[%current_slot, %token, %q8_block, %q8_meta] : view<2x[%capacity]x8x72xf16> -> f16 + %ad = scalar.extf %scale : f16 to f32 + scf.yield %a, %ad : vector<16xi8>, f32 + } else { + %tb = index.scale %token, %q8_row_bytes : index, offset -> offset + %qb = index.add %tb, %q8_block_add : offset + %qc = index.add %qb, %code_offset : offset + %q_view = buffer.view %q8_input[%qc] : buffer -> view<32xi32> + %ds_view = buffer.view %q8_input[%qb] : buffer -> view<8xf16> + %words = vector.load %q_view[%q8_word0] : view<32xi32> -> vector<4xi32> + %a = vector.bitcast %words : vector<4xi32> to vector<16xi8> + %scale = view.load %ds_view[%q8_meta] : view<8xf16> -> f16 + %ad = scalar.extf %scale : f16 to f32 + scf.yield %a, %ad : vector<16xi8>, f32 + } + func.return %a, %ad : vector<16xi8>, f32 +} + +func.def inline @ggml_q6_q8_lowtoken_dot16(%pipelined: i1, %w: vector<16xi8>, %a: vector<16xi8>) -> (i32) { + %sum = scf.if %pipelined -> (i32) { + %di0 = vector.constant 0 : vector<1xi32> + %dw0 = vector.slice %w[0] : vector<16xi8> -> vector<4xi8> + %da0 = vector.slice %a[0] : vector<16xi8> -> vector<4xi8> + %di1 = vector.dot4i %dw0, %da0, %di0 : vector<4xi8>, vector<4xi8>, vector<1xi32> + %dw1 = vector.slice %w[4] : vector<16xi8> -> vector<4xi8> + %da1 = vector.slice %a[4] : vector<16xi8> -> vector<4xi8> + %di2 = vector.dot4i %dw1, %da1, %di1 : vector<4xi8>, vector<4xi8>, vector<1xi32> + %dw2 = vector.slice %w[8] : vector<16xi8> -> vector<4xi8> + %da2 = vector.slice %a[8] : vector<16xi8> -> vector<4xi8> + %di3 = vector.dot4i %dw2, %da2, %di2 : vector<4xi8>, vector<4xi8>, vector<1xi32> + %dw3 = vector.slice %w[12] : vector<16xi8> -> vector<4xi8> + %da3 = vector.slice %a[12] : vector<16xi8> -> vector<4xi8> + %di4 = vector.dot4i %dw3, %da3, %di3 : vector<4xi8>, vector<4xi8>, vector<1xi32> + %chain_sum = vector.extract %di4[0] : vector<1xi32> -> i32 + scf.yield %chain_sum : i32 + } else { + %init = vector.constant 0 : vector<4xi32> + %dot = vector.dot4i %w, %a, %init : vector<16xi8>, vector<16xi8>, vector<4xi32> + %seed = scalar.constant 0 : i32 + %direct_sum = vector.reduce %dot, %seed : vector<4xi32>, i32 + scf.yield %direct_sum : i32 + } + func.return %sum : i32 +} + +func.def inline @ggml_q8_1_x4_prefetch_lowtoken_half_slab(%input: buffer, %capacity0: index, %slot_begin: index, %blocks: index, %global_row_bytes: offset) -> (vector<4xi32>) { + %capacity = index.assume %capacity0 [range(%capacity0, 1, 5)] : index + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %c18 = index.constant 18 : index + %c72 = index.constant 72 : index + %packet_bytes = index.constant 16 : offset + %block_bytes = index.constant 288 : offset + %thread = kernel.workitem.id : index + %total_packets = index.mul %capacity, %c72 : index + %copy_packet = index.rem %thread, %total_packets : index + %copy_token = index.div %copy_packet, %c72 : index + %copy_inner = index.rem %copy_packet, %c72 : index + %remaining = index.sub %blocks, %slot_begin : index + %full = index.cmp uge, %remaining, %c4 : index + %active0 = scf.select %full, %c4, %remaining : index + %active = index.assume %active0 [range(%active0, 1, 4)] : index + %active_packets = index.mul %active, %c18 : index + %copy_valid = index.cmp ult, %copy_inner, %active_packets : index + %safe_inner = scf.select %copy_valid, %copy_inner, %c0 : index + %copy_row = index.scale %copy_token, %global_row_bytes : index, offset -> offset + %slot_bytes = index.scale %slot_begin, %block_bytes : index, offset -> offset + %copy_bytes = index.scale %safe_inner, %packet_bytes : index, offset -> offset + %copy_base0 = index.add %copy_row, %slot_bytes : offset + %copy_base = index.add %copy_base0, %copy_bytes : offset + %source = buffer.view %input[%copy_base] : buffer -> view<4xi32> + %values = vector.load %source[%c0] : view<4xi32> -> vector<4xi32> + func.return %values : vector<4xi32> +} +func.def inline @ggml_q8_1_x4_publish_lowtoken_half_slab(%stage: buffer, %capacity0: index, %slot0: index, %values: vector<4xi32>) { + %capacity = index.assume %capacity0 [range(%capacity0, 1, 5)] : index + %slot = index.assume %slot0 [range(%slot0, 0, 1)] : index + %c4 = index.constant 4 : index + %c72 = index.constant 72 : index + %base = index.constant 0 : offset + %thread = kernel.workitem.id : index + %total_packets = index.mul %capacity, %c72 : index + %thread_valid = index.cmp ult, %thread, %total_packets : index + scf.if %thread_valid { + %copy_token0 = index.div %thread, %c72 : index + %copy_token = index.assume %copy_token0 [range(%copy_token0, 0, 4), lt(%copy_token0, %capacity)] : index + %copy_inner = index.rem %thread, %c72 : index + %origin = index.mul %copy_inner, %c4 : index + %destination = buffer.view %stage[%base] : buffer -> view<2x[%capacity]x288xi32> + vector.store %values, %destination[%slot, %copy_token, %origin] : vector<4xi32>, view<2x[%capacity]x288xi32> + } + func.return +} + + +func.def inline @ggml_q6_q8_lowtoken_accumulate(%pipelined: i1, %tokens: index, %q8_row_bytes: offset, %q8_block_add: offset, %capacity: index, %current_slot: index, %q8_block0: index, %q8_input: buffer, %q8_word0: index, %q8_meta: index, %w: vector<16xi8>, %ws: f32, %acc_g0: f32, %acc_g1: f32, %acc_g2: f32, %acc_g3: f32, %acc_g4: f32) -> (f32, f32, f32, f32, f32) { + %code_offset = index.constant 16 : offset + %t0 = index.constant 0 : index + %valid0 = index.cmp ult, %t0, %tokens : index + %computed_g0 = scf.if %valid0 -> (f32) { + %a, %ad = func.call @ggml_q6_q8_lowtoken_activation_values(%pipelined, %t0, %q8_row_bytes, %q8_block_add, %capacity, %current_slot, %q8_block0, %q8_input, %q8_word0, %q8_meta) : (i1, index, offset, offset, index, index, index, buffer, index, index) -> (vector<16xi8>, f32) + %sum = func.call @ggml_q6_q8_lowtoken_dot16(%pipelined, %w, %a) : (i1, vector<16xi8>, vector<16xi8>) -> (i32) + %sumf = scalar.sitofp %sum : i32 to f32 + %scale = scalar.mulf %ws, %ad : f32 + %updated = scalar.fmaf %sumf, %scale, %acc_g0 : f32 + scf.yield %updated : f32 + } else { + scf.yield %acc_g0 : f32 + } + %t1 = index.constant 1 : index + %valid1 = index.cmp ult, %t1, %tokens : index + %computed_g1 = scf.if %valid1 -> (f32) { + %a, %ad = func.call @ggml_q6_q8_lowtoken_activation_values(%pipelined, %t1, %q8_row_bytes, %q8_block_add, %capacity, %current_slot, %q8_block0, %q8_input, %q8_word0, %q8_meta) : (i1, index, offset, offset, index, index, index, buffer, index, index) -> (vector<16xi8>, f32) + %sum = func.call @ggml_q6_q8_lowtoken_dot16(%pipelined, %w, %a) : (i1, vector<16xi8>, vector<16xi8>) -> (i32) + %sumf = scalar.sitofp %sum : i32 to f32 + %scale = scalar.mulf %ws, %ad : f32 + %updated = scalar.fmaf %sumf, %scale, %acc_g1 : f32 + scf.yield %updated : f32 + } else { + scf.yield %acc_g1 : f32 + } + %t2 = index.constant 2 : index + %valid2 = index.cmp ult, %t2, %tokens : index + %computed_g2 = scf.if %valid2 -> (f32) { + %a, %ad = func.call @ggml_q6_q8_lowtoken_activation_values(%pipelined, %t2, %q8_row_bytes, %q8_block_add, %capacity, %current_slot, %q8_block0, %q8_input, %q8_word0, %q8_meta) : (i1, index, offset, offset, index, index, index, buffer, index, index) -> (vector<16xi8>, f32) + %sum = func.call @ggml_q6_q8_lowtoken_dot16(%pipelined, %w, %a) : (i1, vector<16xi8>, vector<16xi8>) -> (i32) + %sumf = scalar.sitofp %sum : i32 to f32 + %scale = scalar.mulf %ws, %ad : f32 + %updated = scalar.fmaf %sumf, %scale, %acc_g2 : f32 + scf.yield %updated : f32 + } else { + scf.yield %acc_g2 : f32 + } + %t3 = index.constant 3 : index + %valid3 = index.cmp ult, %t3, %tokens : index + %computed_g3 = scf.if %valid3 -> (f32) { + %a, %ad = func.call @ggml_q6_q8_lowtoken_activation_values(%pipelined, %t3, %q8_row_bytes, %q8_block_add, %capacity, %current_slot, %q8_block0, %q8_input, %q8_word0, %q8_meta) : (i1, index, offset, offset, index, index, index, buffer, index, index) -> (vector<16xi8>, f32) + %sum = func.call @ggml_q6_q8_lowtoken_dot16(%pipelined, %w, %a) : (i1, vector<16xi8>, vector<16xi8>) -> (i32) + %sumf = scalar.sitofp %sum : i32 to f32 + %scale = scalar.mulf %ws, %ad : f32 + %updated = scalar.fmaf %sumf, %scale, %acc_g3 : f32 + scf.yield %updated : f32 + } else { + scf.yield %acc_g3 : f32 + } + %t4 = index.constant 4 : index + %valid4 = index.cmp ult, %t4, %tokens : index + %computed_g4 = scf.if %valid4 -> (f32) { + %a, %ad = func.call @ggml_q6_q8_lowtoken_activation_values(%pipelined, %t4, %q8_row_bytes, %q8_block_add, %capacity, %current_slot, %q8_block0, %q8_input, %q8_word0, %q8_meta) : (i1, index, offset, offset, index, index, index, buffer, index, index) -> (vector<16xi8>, f32) + %sum = func.call @ggml_q6_q8_lowtoken_dot16(%pipelined, %w, %a) : (i1, vector<16xi8>, vector<16xi8>) -> (i32) + %sumf = scalar.sitofp %sum : i32 to f32 + %scale = scalar.mulf %ws, %ad : f32 + %updated = scalar.fmaf %sumf, %scale, %acc_g4 : f32 + scf.yield %updated : f32 + } else { + scf.yield %acc_g4 : f32 + } + func.return %computed_g0, %computed_g1, %computed_g2, %computed_g3, %computed_g4 : f32, f32, f32, f32, f32 +} + +func.def inline @ggml_mul_mat_q6_q8_1_x4_lowtoken_dot(%weight_format: index, %token_count: index, %input_size0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %input: buffer, %weight: buffer, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %capacity = config.get @ggml.workload.token_capacity : index + %staged_rows, %k_partitions = func.call pure @ggml_q8_lowtoken_geometry(%capacity, %input_size0, %output_size0, %weight_format) : (index, index, index, index) -> (index, index) + %no_staged_rows = index.constant 0 : index + %staged = index.cmp ugt, %staged_rows, %no_staged_rows : index + %tokens = index.assume %token_count [range(%token_count, 1, 5), le(%token_count, %capacity)] : index + %k = index.assume %input_size0 [range(%input_size0, 256, 32768), mul(%input_size0, 256)] : index + %n = index.assume %output_size0 [range(%output_size0, 64, 262144), mul(%output_size0, 64)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %base = index.constant 0 : offset + %zero = scalar.constant 0.0 : f32 + %lane0 = kernel.subgroup.lane.id : index + %lane = index.assume %lane0 [range(%lane0, 0, 63)] : index + %wave = kernel.subgroup.id : index + %row_tile = scf.select %staged, %c32, %c8 : index + %row_base = index.mul %channel_tile, %row_tile : index + %wave_row = index.mul %wave, %c4 : index + %cohort = index.div %lane, %c16 : index + %local_row = index.add %wave_row, %cohort : index + %row = index.add %row_base, %local_row : index + %block_lane = index.rem %lane, %c16 : index + %group = index.div %block_lane, %c2 : index + %half16 = index.rem %block_lane, %c2 : index + %half128 = index.div %group, %c4 : index + %group_in_half = index.rem %group, %c4 : index + %ql_side = index.rem %group_in_half, %c2 : index + %nibble = index.div %group_in_half, %c2 : index + %nibble_shift0 = index.mul %nibble, %c4 : index + %nibble_shift1 = index.cast %nibble_shift0 : index to i32 + %nibble_shift = vector.splat %nibble_shift1 : vector<4xi32> + %high_shift0 = index.mul %group_in_half, %c2 : index + %high_shift1 = index.cast %high_shift0 : index to i32 + %high_shift = vector.splat %high_shift1 : vector<4xi32> + %packet = index.mul %half16, %c4 : index + %q8_group_words = index.mul %group_in_half, %c8 : index + %q8_word0 = index.add %q8_group_words, %packet : index + %q8_meta = index.mul %group_in_half, %c2 : index + %blocks = index.div %k, %c256 : index + %q8_groups = index.div %k, %c128 : index + %q8_block_bytes = index.constant 144 : offset + %global_q8_row_bytes = index.scale %q8_groups, %q8_block_bytes : index, offset -> offset + %native_block_bytes = index.constant 210 : index + %cache_bytes = index.constant 33554432 : index + %weight_blocks = index.mul %blocks, %n : index + %weight_bytes = index.mul %weight_blocks, %native_block_bytes : index + %exceeds_cache = index.cmp uge, %weight_bytes, %cache_bytes : index + %pipelined = scalar.andi %staged, %exceeds_cache : i1 + %full_slab_row_bytes = index.constant 2304 : offset + %half_slab_row_bytes = index.constant 1152 : offset + %slab_row_bytes = scf.select %pipelined, %half_slab_row_bytes, %full_slab_row_bytes : offset + %q8_row_bytes = scf.select %staged, %slab_row_bytes, %global_q8_row_bytes : offset + %code_offset = index.constant 16 : offset + %record_bytes = index.constant 13440 : offset + %scale_add = index.constant 128 : offset + %raw_add = index.constant 1152 : offset + %staged_row_group = index.div %channel_tile, %c2 : index + %direct_row_group = index.div %row, %c64 : index + %row_group = scf.select %staged, %staged_row_group, %direct_row_group : index + %row_lane = index.rem %row, %c64 : index + %row_group_block = index.mul %row_group, %blocks : index + %input_na, %weight_na = buffer.assume.noalias %input, %weight : buffer, buffer + %slab_blocks = scf.select %pipelined, %c4, %c8 : index + %full_stage_bytes = index.constant 11520 : offset + %q8_stage_bytes = scf.select %staged, %full_stage_bytes, %base : offset + %q8_stage = buffer.alloca align(16) %q8_stage_bytes : buffer + %q8_input = scf.select %staged, %q8_stage, %input_na : buffer + %zero_half = scalar.constant 0.0 : f16 + %zero_code = scalar.constant 0 : i8 + %zero_words = vector.constant 0 : vector<4xi32> + %prime_d, %prime_s, %prime_ql, %prime_qh = scf.if %staged -> (f16, i8, vector<4xi32>, vector<4xi32>) { + %pd, %ps, %pql, %pqh = func.call @ggml_q6_prefetch_native_words(%weight_na, %row_group_block, %c0, %row_lane, %half16, %group, %half128, %ql_side, %packet) : (buffer, index, index, index, index, index, index, index, index) -> (f16, i8, vector<4xi32>, vector<4xi32>) + scf.yield %pd, %ps, %pql, %pqh : f16, i8, vector<4xi32>, vector<4xi32> + } else { + scf.yield %zero_half, %zero_code, %zero_words, %zero_words : f16, i8, vector<4xi32>, vector<4xi32> + } + scf.if %pipelined { + %prime_q8 = func.call @ggml_q8_1_x4_prefetch_lowtoken_half_slab(%input_na, %capacity, %c0, %blocks, %global_q8_row_bytes) : (buffer, index, index, index, offset) -> (vector<4xi32>) + func.call @ggml_q8_1_x4_publish_lowtoken_half_slab(%q8_stage, %capacity, %c0, %prime_q8) : (buffer, index, index, vector<4xi32>) + kernel.barrier scope(workgroup) ordering(acq_rel) + } + %total_g0, %total_g1, %total_g2, %total_g3, %total_g4, %last_d, %last_s, %last_ql, %last_qh = scf.for %block = [%c0 to %blocks step %c1](%acc_g0 = %zero : f32, %acc_g1 = %zero : f32, %acc_g2 = %zero : f32, %acc_g3 = %zero : f32, %acc_g4 = %zero : f32, %carry_d = %prime_d : f16, %carry_s = %prime_s : i8, %carry_ql = %prime_ql : vector<4xi32>, %carry_qh = %prime_qh : vector<4xi32>) -> (f32, f32, f32, f32, f32, f16, i8, vector<4xi32>, vector<4xi32>) { + %block_in_slab = index.rem %block, %slab_blocks : index + %slab_boundary = index.cmp eq, %block_in_slab, %c0 : index + %direct_stage = scalar.andi %staged, %slab_boundary : i1 + %true = scalar.constant true : i1 + %not_pipelined = scalar.xori %pipelined, %true : i1 + %stage_current = scalar.andi %direct_stage, %not_pipelined : i1 + scf.if %stage_current { + func.call @ggml_q8_1_x4_stage_lowtoken_slab(%input_na, %q8_stage, %capacity, %block, %blocks, %global_q8_row_bytes) : (buffer, buffer, index, index, index, offset) + } + %slab_number = index.div %block, %slab_blocks : index + %current_slot = index.rem %slab_number, %c2 : index + %following_slab = index.add %slab_number, %c1 : index + %following_begin = index.mul %following_slab, %slab_blocks : index + %following_exists = index.cmp ult, %following_begin, %blocks : index + %next_slot = index.rem %following_slab, %c2 : index + %staged_future = scalar.andi %pipelined, %following_exists : i1 + %current_d, %current_s, %current_ql, %current_qh = scf.if %staged -> (f16, i8, vector<4xi32>, vector<4xi32>) { + scf.yield %carry_d, %carry_s, %carry_ql, %carry_qh : f16, i8, vector<4xi32>, vector<4xi32> + } else { + %wd, %ws8, %wql, %wqh = func.call @ggml_q6_prefetch_native_words(%weight_na, %row_group_block, %block, %row_lane, %half16, %group, %half128, %ql_side, %packet) : (buffer, index, index, index, index, index, index, index, index) -> (f16, i8, vector<4xi32>, vector<4xi32>) + scf.yield %wd, %ws8, %wql, %wqh : f16, i8, vector<4xi32>, vector<4xi32> + } + %d = scalar.extf %current_d : f16 to f32 + %sf = scalar.sitofp %current_s : i8 to f32 + %ws = scalar.mulf %d, %sf : f32 + %nibble_mask = vector.constant 252645135 : vector<4xi32> + %high_mask = vector.constant 50529027 : vector<4xi32> + %four = vector.constant 4 : vector<4xi32> + %ql_shifted = vector.shrui %current_ql, %nibble_shift : vector<4xi32> + %low = vector.andi %ql_shifted, %nibble_mask : vector<4xi32> + %qh_shifted = vector.shrui %current_qh, %high_shift : vector<4xi32> + %high0 = vector.andi %qh_shifted, %high_mask : vector<4xi32> + %high = vector.shli %high0, %four : vector<4xi32> + %code = vector.ori %low, %high : vector<4xi32> + // Guard bits prevent subtraction from borrowing across packed bytes. + %guard_mask = vector.constant -2139062144 : vector<4xi32> + %center = vector.constant 538976288 : vector<4xi32> + %guarded = vector.ori %code, %guard_mask : vector<4xi32> + %centered = vector.subi %guarded, %center : vector<4xi32> + %signed_code = vector.xori %centered, %guard_mask : vector<4xi32> + %w = vector.bitcast %signed_code : vector<4xi32> to vector<16xi8> + %next_block = index.add %block, %c1 : index + %next_exists = index.cmp ult, %next_block, %blocks : index + %has_next = scalar.andi %staged, %next_exists : i1 + %next_d, %next_s, %next_ql, %next_qh = scf.if %has_next -> (f16, i8, vector<4xi32>, vector<4xi32>) { + %wd, %ws8, %wql, %wqh = func.call @ggml_q6_prefetch_native_words(%weight_na, %row_group_block, %next_block, %row_lane, %half16, %group, %half128, %ql_side, %packet) : (buffer, index, index, index, index, index, index, index, index) -> (f16, i8, vector<4xi32>, vector<4xi32>) + scf.yield %wd, %ws8, %wql, %wqh : f16, i8, vector<4xi32>, vector<4xi32> + } else { + scf.yield %carry_d, %carry_s, %carry_ql, %carry_qh : f16, i8, vector<4xi32>, vector<4xi32> + } + %input_block = scf.select %staged, %block_in_slab, %block : index + %q8_block0 = index.mul %input_block, %c2 : index + %q8_block = index.add %q8_block0, %half128 : index + %q8_local_block_add = index.scale %q8_block, %q8_block_bytes : index, offset -> offset + %slot_row = index.mul %current_slot, %capacity : index + %slot_byte_add0 = index.scale %slot_row, %slab_row_bytes : index, offset -> offset + %slot_byte_add = scf.select %pipelined, %slot_byte_add0, %base : offset + %q8_block_add = index.add %slot_byte_add, %q8_local_block_add : offset + %slot_complete = index.cmp eq, %block_in_slab, %c3 : index + %publish_next = scalar.andi %staged_future, %slot_complete : i1 + %next_g0, %next_g1, %next_g2, %next_g3, %next_g4 = scf.if %publish_next -> (f32, f32, f32, f32, f32) { + %next_q8 = func.call @ggml_q8_1_x4_prefetch_lowtoken_half_slab(%input_na, %capacity, %following_begin, %blocks, %global_q8_row_bytes) : (buffer, index, index, index, offset) -> (vector<4xi32>) + %v0, %v1, %v2, %v3, %v4 = func.call @ggml_q6_q8_lowtoken_accumulate(%pipelined, %tokens, %q8_row_bytes, %q8_block_add, %capacity, %current_slot, %q8_block, %q8_input, %q8_word0, %q8_meta, %w, %ws, %acc_g0, %acc_g1, %acc_g2, %acc_g3, %acc_g4) : (i1, index, offset, offset, index, index, index, buffer, index, index, vector<16xi8>, f32, f32, f32, f32, f32, f32) -> (f32, f32, f32, f32, f32) + func.call @ggml_q8_1_x4_publish_lowtoken_half_slab(%q8_stage, %capacity, %next_slot, %next_q8) : (buffer, index, index, vector<4xi32>) + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %v0, %v1, %v2, %v3, %v4 : f32, f32, f32, f32, f32 + } else { + %v0, %v1, %v2, %v3, %v4 = func.call @ggml_q6_q8_lowtoken_accumulate(%pipelined, %tokens, %q8_row_bytes, %q8_block_add, %capacity, %current_slot, %q8_block, %q8_input, %q8_word0, %q8_meta, %w, %ws, %acc_g0, %acc_g1, %acc_g2, %acc_g3, %acc_g4) : (i1, index, offset, offset, index, index, index, buffer, index, index, vector<16xi8>, f32, f32, f32, f32, f32, f32) -> (f32, f32, f32, f32, f32) + scf.yield %v0, %v1, %v2, %v3, %v4 : f32, f32, f32, f32, f32 + } + scf.yield %next_g0, %next_g1, %next_g2, %next_g3, %next_g4, %next_d, %next_s, %next_ql, %next_qh : f32, f32, f32, f32, f32, f16, i8, vector<4xi32>, vector<4xi32> + } + func.call @ggml_mul_mat_q8_lowtoken_publish(%tokens, %n, %row_tile, %row_base, %local_row, %block_lane, %channel_tile, %output_accumulation, %output_unary_op, %epsilon, %total_g0, %total_g1, %total_g2, %total_g3, %total_g4, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (index, index, index, index, index, index, index, index, index, f32, f32, f32, f32, f32, f32, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + func.return +} + +template.def<@ggml.mul_mat_f32_f32_wmma.core> device @ggml_mul_mat_f32_f32_wmma_core(%weight_format: index, %token_count: index, %input_size0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %input: buffer, %weight: buffer, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %token_capacity = config.get @ggml.workload.token_capacity : index + %lowtoken = func.call pure @ggml_mul_mat_uses_lowtoken_dot(%token_capacity, %input_size0) : (index, index) -> (i1) + %activation_format = config.get @ggml.mul_mat.activation_format : index + %q8_format = index.constant 9 : index + %uses_q8 = index.cmp eq, %activation_format, %q8_format : index + scf.if %lowtoken { + scf.if %uses_q8 { + %q6_format = index.constant 46 : index + %uses_q6 = index.cmp eq, %weight_format, %q6_format : index + scf.if %uses_q6 { + func.call @ggml_mul_mat_q6_q8_1_x4_lowtoken_dot(%weight_format, %token_count, %input_size0, %output_size0, %output_accumulation, %output_unary_op, %epsilon, %channel_tile, %token_tile, %input, %weight, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + } else { + func.call @ggml_mul_mat_q4_q8_1_x4_lowtoken_dot(%weight_format, %token_count, %input_size0, %output_size0, %output_accumulation, %output_unary_op, %epsilon, %channel_tile, %token_tile, %input, %weight, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + } + } else { + func.call @ggml_mul_mat_f32_f32_lowtoken_dot(%weight_format, %token_count, %input_size0, %output_size0, %output_accumulation, %output_unary_op, %epsilon, %channel_tile, %token_tile, %input, %weight, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + } + } else { + func.call @ggml_mul_mat_f32_f32_wmma_tiled(%weight_format, %token_count, %input_size0, %output_size0, %output_accumulation, %output_unary_op, %epsilon, %channel_tile, %token_tile, %input, %weight, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + } + template.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_f16_f16_wmma_core.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_f16_f16_wmma_core.loom new file mode 100644 index 000000000000..bd6902a2d703 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_f16_f16_wmma_core.loom @@ -0,0 +1,585 @@ +// Expert-grouped Q4_K and Q6_K routed projections for gfx11 prefill. +// +// Each two-wave workgroup decodes 64 output channels for 32 routed rows of one +// expert. Raw GGUF weights and FP16 routed SwiGLU rows are staged as FP16 WMMA +// operands. Results remain FP16 and are scattered into compact +// [token, route, hidden] order without collisions. A following reduction +// widens each route, applies its normalized weight, and accumulates the +// residual in FP32. +// +// Top-k routing selects an expert at most once per token. The expert table +// therefore contains at most token_count assignments per expert, while each +// assignment ordinal still identifies the compact [token, route] activation +// and route-weight row. +template.decl @ggml.mul_mat_id_f16_f16_wmma.core(%weight_format: index, %token_count: index, %input_size: index, %route_count: index, %expert_count: index, %output_size: index, %channel_tile: index, %route_tile: index, %expert: index, %input: buffer, %expert_table: buffer, %weight: buffer, %output: buffer) + +template.decl @ggml.mul_mat_id_f16_f16_wmma.finish_tile(%token_count: index, %route_count: index, %output_size: index, %expert_count: index, %channel_tile: index, %expert: index, %route_tile_base: index, %expert_route_count: index, %expert_table: buffer, %output: buffer) + +template.decl @ggml.mul_mat_id_f16_f16_wmma.publish_vector4(%publish_word: i1, %assignment: index, %channel: index, %token_count: index, %route_count: index, %output_size: index, %values: vector<4xf16>, %output: buffer) + +func.def inline @ggml_q4k_scale_from_header(%scale0: i32, %scale1: i32, %scale2: i32, %q4_group: index) -> (i32, i32) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c48_i32 = scalar.constant 48 : i32 + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + %is_low_group = index.cmp ult, %bounded_group, %c4 : index + %scale_lane = index.rem %bounded_group, %c4 : index + %scale_shift_index = index.mul %scale_lane, %c8 : index + %scale_shift = index.cast %scale_shift_index : index to i32 + %high_shift = scalar.addi %scale_shift, %c2_i32 : i32 + %minimum_shift = scalar.addi %scale_shift, %c4_i32 : i32 + %selected_scale_source = scf.select %is_low_group, %scale0, %scale2 : i32 + %selected_minimum_source = scf.select %is_low_group, %scale1, %scale2 : i32 + %selected_scale_high_shift = scf.select %is_low_group, %scale_shift, %high_shift : i32 + %selected_minimum_low_shift = scf.select %is_low_group, %scale_shift, %minimum_shift : i32 + %scale_low0 = scalar.shrui %selected_scale_source, %scale_shift : i32 + %scale_low = scalar.andi %scale_low0, %c15_i32 : i32 + %scale_high0 = scalar.shrui %scale0, %selected_scale_high_shift : i32 + %scale_high = scalar.andi %scale_high0, %c48_i32 : i32 + %scale = scalar.ori %scale_low, %scale_high : i32 + %minimum_low0 = scalar.shrui %selected_minimum_source, %selected_minimum_low_shift : i32 + %minimum_low = scalar.andi %minimum_low0, %c15_i32 : i32 + %minimum_high0 = scalar.shrui %scale1, %selected_scale_high_shift : i32 + %minimum_high = scalar.andi %minimum_high0, %c48_i32 : i32 + %minimum = scalar.ori %minimum_low, %minimum_high : i32 + func.return %scale, %minimum : i32, i32 +} + +// Acquires the packed code word shared by one adjacent Q4_K group pair. +func.def inline @ggml_q4k_split_wmma_load_code(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group_pair: index, %packet: index) -> (vector<1xi32>) { + %c8 = index.constant 8 : index + %block_bytes = index.constant 144 : offset + %code_offset = index.constant 16 : offset + %bounded_group_pair = index.assume %q4_group_pair [range(%q4_group_pair, 0, 3)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q4_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<32xi32> + %q_page = index.mul %bounded_group_pair, %c8 : index + %q_word_index0 = index.add %q_page, %bounded_packet : index + %q_word_index = index.assume %q_word_index0 [range(%q_word_index0, 0, 31)] : index + %q_word = vector.load %code_view[%q_word_index] : view<32xi32> -> vector<1xi32> + func.return %q_word : vector<1xi32> +} + +// Decodes the four adjacent Q4_K values owned by one load packet from an +// already-loaded block header and packed code word. Matrix schedules choose +// the lifetime of both immutable packets. +func.def inline @ggml_q4k_split_wmma_vector4_from_header_code(%q4_group: index, %header_words: vector<4xi32>, %q_word: vector<1xi32>) -> (vector<4xf16>) { + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %q4_mask = vector.constant 252645135 : vector<1xi32> + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + %header_halves = vector.bitcast %header_words : vector<4xi32> to vector<8xf16> + %d_f16 = vector.extract %header_halves[0] : vector<8xf16> -> f16 + %dmin_f16 = vector.extract %header_halves[1] : vector<8xf16> -> f16 + %scale0 = vector.extract %header_words[1] : vector<4xi32> -> i32 + %scale1 = vector.extract %header_words[2] : vector<4xi32> -> i32 + %scale2 = vector.extract %header_words[3] : vector<4xi32> -> i32 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %scale, %minimum = func.call @ggml_q4k_scale_from_header(%scale0, %scale1, %scale2, %bounded_group) : (i32, i32, i32, index) -> (i32, i32) + %scale_f32 = scalar.uitofp %scale : i32 to f32 + %minimum_f32 = scalar.uitofp %minimum : i32 to f32 + %d_scale = scalar.mulf %d, %scale_f32 : f32 + %minimum_scale = scalar.mulf %dmin, %minimum_f32 : f32 + %q_half = index.rem %bounded_group, %c2 : index + %q_shift_index = index.mul %q_half, %c4 : index + %q_shift_i32 = index.cast %q_shift_index : index to i32 + %q_shift = vector.splat %q_shift_i32 : vector<1xi32> + %shifted_q = vector.shrui %q_word, %q_shift : vector<1xi32> + %masked_q = vector.andi %shifted_q, %q4_mask : vector<1xi32> + %q_i8 = vector.bitcast %masked_q : vector<1xi32> to vector<4xi8> + %q_f32 = vector.uitofp %q_i8 : vector<4xi8> to vector<4xf32> + // Form adjacent FP16 lanes from fused FP32 affine expressions. AMDGPU maps + // this natural shape to packed mixlo/mixhi instructions where available. + %negative_minimum_scale = scalar.negf %minimum_scale : f32 + %q0 = vector.extract %q_f32[0] : vector<4xf32> -> f32 + %q1 = vector.extract %q_f32[1] : vector<4xf32> -> f32 + %q2 = vector.extract %q_f32[2] : vector<4xf32> -> f32 + %q3 = vector.extract %q_f32[3] : vector<4xf32> -> f32 + %value0 = scalar.fmaf %q0, %d_scale, %negative_minimum_scale : f32 + %value1 = scalar.fmaf %q1, %d_scale, %negative_minimum_scale : f32 + %value2 = scalar.fmaf %q2, %d_scale, %negative_minimum_scale : f32 + %value3 = scalar.fmaf %q3, %d_scale, %negative_minimum_scale : f32 + %half0 = scalar.fptrunc %value0 : f32 to f16 + %half1 = scalar.fptrunc %value1 : f32 to f16 + %half2 = scalar.fptrunc %value2 : f32 to f16 + %half3 = scalar.fptrunc %value3 : f32 to f16 + %result = vector.from_elements %half0, %half1, %half2, %half3 : vector<4xf16> + func.return %result : vector<4xf16> +} + +// Decodes one group when its caller has retained only the Q4_K block header. +func.def inline @ggml_q4k_split_wmma_vector4_from_header(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group: index, %packet: index, %header_words: vector<4xi32>) -> (vector<4xf16>) { + %c2 = index.constant 2 : index + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + %q4_group_pair = index.div %bounded_group, %c2 : index + %q_word = func.call @ggml_q4k_split_wmma_load_code(%weight, %row_byte_base, %q4_block, %q4_group_pair, %packet) : (buffer, offset, index, index, index) -> (vector<1xi32>) + %values = func.call @ggml_q4k_split_wmma_vector4_from_header_code(%bounded_group, %header_words, %q_word) : (index, vector<4xi32>, vector<1xi32>) -> (vector<4xf16>) + func.return %values : vector<4xf16> +} + +// Acquires one naturally aligned Q4_K block header as a single 16-byte packet. +func.def inline @ggml_q4k_split_wmma_load_header(%weight: buffer, %row_byte_base: offset, %q4_block: index) -> (vector<4xi32>) { + %c0 = index.constant 0 : index + %block_bytes = index.constant 144 : offset + %block_byte_add = index.scale %q4_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %header_view = buffer.view %weight[%block_byte_base] : buffer -> view<4xi32> + %header_words = vector.load %header_view[%c0] : view<4xi32> -> vector<4xi32> + func.return %header_words : vector<4xi32> +} + +// Acquires one block header before decoding the selected four-value group +// packet. Matrix schedules that span several groups call the two operations +// separately so the header lifetime matches their complete block loop. +func.def inline @ggml_q4k_split_wmma_vector4(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group: index, %packet: index) -> (vector<4xf16>) { + %header_words = func.call @ggml_q4k_split_wmma_load_header(%weight, %row_byte_base, %q4_block) : (buffer, offset, index) -> (vector<4xi32>) + %values = func.call @ggml_q4k_split_wmma_vector4_from_header(%weight, %row_byte_base, %q4_block, %q4_group, %packet, %header_words) : (buffer, offset, index, index, index, vector<4xi32>) -> (vector<4xf16>) + func.return %values : vector<4xf16> +} + +// Acquires the packed words shared by one four-group half of a Q6_K block. +// Each QL word supplies two groups and the QH word supplies all four. +func.def inline @ggml_q6k_split_wmma_load_half_codes(%weight: buffer, %weight_row_byte_base: offset, %q6_block: index, %q6_half: index, %packet: index) -> (vector<1xi32>, vector<1xi32>, vector<1xi32>) { + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %block_bytes = index.constant 210 : offset + %qh_byte_add = index.constant 128 : offset + %bounded_half = index.assume %q6_half [range(%q6_half, 0, 1)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q6_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %weight_row_byte_base, %block_byte_add : offset + %qh_byte_base = index.add %block_byte_base, %qh_byte_add : offset + %ql_view = buffer.view %weight[%block_byte_base] : buffer -> view<32xi32> + %qh_view = buffer.view %weight[%qh_byte_base] : buffer -> view<16xi32> + %ql_half_word_base = index.mul %bounded_half, %c16 : index + %ql0_word_index = index.add %ql_half_word_base, %bounded_packet : index + %ql1_word_base = index.add %ql_half_word_base, %c8 : index + %ql1_word_index = index.add %ql1_word_base, %bounded_packet : index + %qh_half_word_base = index.mul %bounded_half, %c8 : index + %qh_word_index = index.add %qh_half_word_base, %bounded_packet : index + %ql0_word = vector.load %ql_view[%ql0_word_index] : view<32xi32> -> vector<1xi32> + %ql1_word = vector.load %ql_view[%ql1_word_index] : view<32xi32> -> vector<1xi32> + %qh_word = vector.load %qh_view[%qh_word_index] : view<16xi32> -> vector<1xi32> + func.return %ql0_word, %ql1_word, %qh_word : vector<1xi32>, vector<1xi32>, vector<1xi32> +} + +// Decodes four adjacent values after the surrounding schedule has selected +// the scale and retained the packed code words at their natural lifetimes. +func.def inline @ggml_q6k_split_wmma_vector4_from_scale_codes(%q6_group: index, %scale_i8: i8, %d_f16: f16, %ql_word: vector<1xi32>, %qh_word: vector<1xi32>) -> (vector<4xf16>) { + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c4_i32v = vector.constant 4 : vector<1xi32> + %nibble_mask = vector.constant 252645135 : vector<1xi32> + %high_mask = vector.constant 50529027 : vector<1xi32> + %c32_f32v = vector.constant 32.0 : vector<4xf32> + %bounded_group = index.assume %q6_group [range(%q6_group, 0, 7)] : index + %group_in_half = index.rem %bounded_group, %c4 : index + %nibble = index.div %group_in_half, %c2 : index + %nibble_shift_index = index.mul %nibble, %c4 : index + %nibble_shift_i32 = index.cast %nibble_shift_index : index to i32 + %nibble_shift = vector.splat %nibble_shift_i32 : vector<1xi32> + %qh_shift_index = index.mul %group_in_half, %c2 : index + %qh_shift_i32 = index.cast %qh_shift_index : index to i32 + %qh_shift = vector.splat %qh_shift_i32 : vector<1xi32> + %ql_shifted = vector.shrui %ql_word, %nibble_shift : vector<1xi32> + %ql = vector.andi %ql_shifted, %nibble_mask : vector<1xi32> + %qh_shifted = vector.shrui %qh_word, %qh_shift : vector<1xi32> + %qh_low = vector.andi %qh_shifted, %high_mask : vector<1xi32> + %qh = vector.shli %qh_low, %c4_i32v : vector<1xi32> + %code = vector.ori %ql, %qh : vector<1xi32> + %code_i8 = vector.bitcast %code : vector<1xi32> to vector<4xi8> + %code_f32 = vector.uitofp %code_i8 : vector<4xi8> to vector<4xf32> + %centered = vector.subf %code_f32, %c32_f32v : vector<4xf32> + %scale = scalar.sitofp %scale_i8 : i8 to f32 + %d = scalar.extf %d_f16 : f16 to f32 + %combined_scale = scalar.mulf %scale, %d : f32 + %combined_scale_vector = vector.splat %combined_scale : vector<4xf32> + %values_f32 = vector.mulf %centered, %combined_scale_vector : vector<4xf32> + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +// Loads only the group scale and block multiplier while reusing packed words +// retained by an enclosing four-group half-block schedule. +func.def inline @ggml_q6k_split_wmma_vector4_from_half_codes(%weight: buffer, %weight_row_byte_base: offset, %q6_block: index, %q6_group: index, %packet: index, %ql0_word: vector<1xi32>, %ql1_word: vector<1xi32>, %qh_word: vector<1xi32>) -> (vector<4xf16>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %block_bytes = index.constant 210 : offset + %scale_byte_add = index.constant 192 : offset + %d_byte_add = index.constant 208 : offset + %bounded_group = index.assume %q6_group [range(%q6_group, 0, 7)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q6_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %weight_row_byte_base, %block_byte_add : offset + %scale_byte_base = index.add %block_byte_base, %scale_byte_add : offset + %d_byte_base = index.add %block_byte_base, %d_byte_add : offset + %scale_view = buffer.view %weight[%scale_byte_base] : buffer -> view<16xi8> + %d_view = buffer.view %weight[%d_byte_base] : buffer -> view<1xf16> + %group_in_half = index.rem %bounded_group, %c4 : index + %ql_side = index.rem %group_in_half, %c2 : index + %uses_ql1 = index.cmp eq, %ql_side, %c1 : index + %ql_word = scf.select %uses_ql1, %ql1_word, %ql0_word : vector<1xi32> + %scale_packet_half = index.div %bounded_packet, %c4 : index + %scale_group_base = index.mul %bounded_group, %c2 : index + %scale_index = index.add %scale_group_base, %scale_packet_half : index + %scale_i8 = view.load %scale_view[%scale_index] : view<16xi8> -> i8 + %d_f16 = view.load %d_view[%c0] : view<1xf16> -> f16 + %values = func.call @ggml_q6k_split_wmma_vector4_from_scale_codes(%bounded_group, %scale_i8, %d_f16, %ql_word, %qh_word) : (index, i8, f16, vector<1xi32>, vector<1xi32>) -> (vector<4xf16>) + func.return %values : vector<4xf16> +} + +// Shared raw-quantized matrix schedule. Entry points pass a literal weight +// format so linking and JIT specialization erase the inactive packed decoder. +template.def<@ggml.mul_mat_id_f16_f16_wmma.core> device @ggml_mul_mat_id_f16_f16_wmma_core(%weight_format: index, %token_count: index, %input_size: index, %route_count: index, %expert_count: index, %output_size: index, %channel_tile: index, %route_tile: index, %expert: index, %input: buffer, %expert_table: buffer, %weight: buffer, %output: buffer) { + %bounded_route_count = index.assume %route_count [range(%route_count, 1, 8)] : index + %bounded_expert_count = index.assume %expert_count [range(%expert_count, 1, 128)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 4096)] : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 1)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c6 = index.constant 6 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c40 = index.constant 40 : index + %c48 = index.constant 48 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %c4_bytes = index.constant 4 : offset + %q4_block_bytes = index.constant 144 : offset + %q6_block_bytes = index.constant 210 : offset + %weight_stage_bytes = index.constant 5120 : offset + %activation_stage_bytes = index.constant 2560 : offset + %route_stage_bytes = index.constant 128 : offset + %wave_result_stage_bytes = index.constant 512 : offset + %result_stage_bytes = index.constant 1024 : offset + %c0_i32 = scalar.constant 0 : i32 + %cn1_i32 = scalar.constant -1 : i32 + %c0_i32x1 = vector.constant 0 : vector<1xi32> + %c0_i32x4 = vector.constant 0 : vector<4xi32> + %c0_f16x4 = vector.constant 0.0 : vector<4xf16> + %zero_accumulator = vector.constant 0.0 : vector<8xf16> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %assignment_count = index.mul %token_count, %bounded_route_count : index + %assignment_table_byte_base = index.scale %bounded_expert_count, %c4_bytes : index, offset -> offset + %is_q4 = index.cmp eq, %weight_format, %c4 : index + %is_q6 = index.cmp eq, %weight_format, %c6 : index + %quant_block_bytes = scf.select %is_q6, %q6_block_bytes, %q4_block_bytes : offset + %quant_block_count = index.div %input_size, %c256 : index + %weight_row_bytes = index.scale %quant_block_count, %quant_block_bytes : index, offset -> offset + %weight_expert_bytes = index.scale %bounded_output_size, %weight_row_bytes : index, offset -> offset + %input_noalias, %expert_table_noalias, %weight_noalias = buffer.assume.noalias %input, %expert_table, %weight : buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%assignment_count]x[%input_size]xf16> + %count_view = buffer.view %expert_table_noalias[%c0_offset] : buffer -> view<[%bounded_expert_count]xi32> + %assignment_view = buffer.view %expert_table_noalias[%assignment_table_byte_base] : buffer -> view<[%bounded_expert_count]x[%token_count]xi32> + %weight_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %activation_stage = buffer.alloca align(16) %activation_stage_bytes : buffer + %route_stage = buffer.alloca align(16) %route_stage_bytes : buffer + %result_stage = buffer.alloca align(16) %result_stage_bytes : buffer + %weight_stage_view = buffer.view %weight_stage[%c0_offset] : buffer -> view<64x40xf16> + %activation_stage_physical_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x40xf16> + %activation_fragment_layout = encoding.layout.strided [1, %c40] : encoding + %activation_fragment_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x32xf16, %activation_fragment_layout> + %route_stage_view = buffer.view %route_stage[%c0_offset] : buffer -> view<32xi32> + %wave_result_stage_offset = index.scale %subgroup, %wave_result_stage_bytes : index, offset -> offset + %result_fragment_layout = encoding.layout.strided [1, %c16] : encoding + %result_fragment_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf16, %result_fragment_layout> + %result_physical_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf16> + %channel_tile_base = index.mul %channel_tile, %c64 : index + %initial_route_tile_base = index.mul %route_tile, %c32 : index + // The launch workload fixes the interleaved route partitions for this exact + // command-program specialization. + %padded_token_count = index.add %token_count, %c63 : index + %route_partition_count = index.div %padded_token_count, %c64 : index + %route_partition_step = index.mul %route_partition_count, %c32 : index + %bounded_expert, %table_expert_count = index.assume %expert, %bounded_expert_count [lt(%expert, %bounded_expert_count)] : index, index + %is_workitem_zero = index.cmp eq, %workitem, %c0 : index + %lane_expert_route_count = scf.if %is_workitem_zero -> (i32) { + %loaded = view.load %count_view[%bounded_expert] : view<[%bounded_expert_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %c0_i32 : i32 + } + %expert_route_count_reduced = kernel.workgroup.reduce %lane_expert_route_count : i32 + %expert_route_count_i32 = kernel.subgroup.broadcast.first %expert_route_count_reduced : i32 + %expert_route_count0 = index.cast %expert_route_count_i32 : i32 to index + %expert_route_count = index.assume %expert_route_count0 [range(%expert_route_count0, 0, 2048)] : index + scf.for %route_tile_base = [%initial_route_tile_base to %expert_route_count step %route_partition_step] { + // Snapshot this expert's compact assignment map once, then reuse it for + // every K group and output-channel tile. + %loads_route = index.cmp ult, %workitem, %c32 : index + scf.if %loads_route { + %local_route = index.assume %workitem [range(%workitem, 0, 31)] : index + %assignment_ordinal = index.add %route_tile_base, %local_route : index + %valid_row = index.cmp ult, %assignment_ordinal, %expert_route_count : index + %assignment_i32 = scf.if %valid_row -> (i32) { + %bounded_assignment_ordinal, %table_token_count = index.assume %assignment_ordinal, %token_count [lt(%assignment_ordinal, %token_count)] : index, index + %loaded = view.load %assignment_view[%bounded_expert, %bounded_assignment_ordinal] : view<[%bounded_expert_count]x[%token_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %cn1_i32 : i32 + } + view.store %assignment_i32, %route_stage_view[%local_route] : i32, view<32xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %load_packet = index.rem %workitem, %c8 : index + %load_k = index.mul %load_packet, %c4 : index + %load_row0 = index.div %workitem, %c8 : index + %load_row = index.assume %load_row0 [range(%load_row0, 0, 15)] : index + %expert_byte_base = index.scale %bounded_expert, %weight_expert_bytes : index, offset -> offset + %subgroup_channel_add = index.mul %subgroup, %c32 : index + %subgroup_channel1 = index.add %subgroup_channel_add, %c16 : index + %init00 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf16> + %init01 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf16> + %init10 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf16> + %init11 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf16> + %result00, %result01, %result10, %result11 = scf.for %quant_block = [%c0 to %quant_block_count step %c1](%block_acc00 = %init00 : vector<8xf16>, %block_acc01 = %init01 : vector<8xf16>, %block_acc10 = %init10 : vector<8xf16>, %block_acc11 = %init11 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16>) { + // Four lanes cover the 64 output rows owned by one load packet. Snapshot + // each Q4_K block header once and retain it across all eight quant groups. + // The Q6 specialization erases this loop with %is_q4=false. + %header0, %header1, %header2, %header3 = scf.for %header_row_offset = [%c0 to %c64 step %c16](%prior_header0 = %c0_i32x4 : vector<4xi32>, %prior_header1 = %c0_i32x4 : vector<4xi32>, %prior_header2 = %c0_i32x4 : vector<4xi32>, %prior_header3 = %c0_i32x4 : vector<4xi32>) -> (vector<4xi32>, vector<4xi32>, vector<4xi32>, vector<4xi32>) unroll { + %header_local_row0 = index.add %load_row, %header_row_offset : index + %header_local_row = index.assume %header_local_row0 [range(%header_local_row0, 0, 63)] : index + %header_channel = index.add %channel_tile_base, %header_local_row : index + %valid_header_channel = index.cmp ult, %header_channel, %bounded_output_size : index + %loads_q4_header = scalar.andi %is_q4, %valid_header_channel : i1 + %loaded_header = scf.if %loads_q4_header -> (vector<4xi32>) { + %header_channel_byte_add = index.scale %header_channel, %weight_row_bytes : index, offset -> offset + %header_row_byte_base = index.add %expert_byte_base, %header_channel_byte_add : offset + %header_words = func.call @ggml_q4k_split_wmma_load_header(%weight_noalias, %header_row_byte_base, %quant_block) : (buffer, offset, index) -> (vector<4xi32>) + scf.yield %header_words : vector<4xi32> + } else { + scf.yield %c0_i32x4 : vector<4xi32> + } + %updates_header0 = index.cmp eq, %header_row_offset, %c0 : index + %updates_header1 = index.cmp eq, %header_row_offset, %c16 : index + %updates_header2 = index.cmp eq, %header_row_offset, %c32 : index + %updates_header3 = index.cmp eq, %header_row_offset, %c48 : index + %next_header0 = scf.select %updates_header0, %loaded_header, %prior_header0 : vector<4xi32> + %next_header1 = scf.select %updates_header1, %loaded_header, %prior_header1 : vector<4xi32> + %next_header2 = scf.select %updates_header2, %loaded_header, %prior_header2 : vector<4xi32> + %next_header3 = scf.select %updates_header3, %loaded_header, %prior_header3 : vector<4xi32> + scf.yield %next_header0, %next_header1, %next_header2, %next_header3 : vector<4xi32>, vector<4xi32>, vector<4xi32>, vector<4xi32> + } + %quant_group_outer_count = scf.select %is_q4, %c4, %c2 : index + %groups_per_outer = scf.select %is_q4, %c2, %c4 : index + %block_result00, %block_result01, %block_result10, %block_result11 = scf.for %quant_group_outer = [%c0 to %quant_group_outer_count step %c1](%acc00 = %block_acc00 : vector<8xf16>, %acc01 = %block_acc01 : vector<8xf16>, %acc10 = %block_acc10 : vector<8xf16>, %acc11 = %block_acc11 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16>) { + // Q4_K retains one packed code word across its adjacent low/high group + // pair. The Q6 specialization erases this acquisition path. + %q_word0, %q_word1, %q_word2, %q_word3 = scf.for %code_row_offset = [%c0 to %c64 step %c16](%prior_q_word0 = %c0_i32x1 : vector<1xi32>, %prior_q_word1 = %c0_i32x1 : vector<1xi32>, %prior_q_word2 = %c0_i32x1 : vector<1xi32>, %prior_q_word3 = %c0_i32x1 : vector<1xi32>) -> (vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>) unroll { + %code_local_row0 = index.add %load_row, %code_row_offset : index + %code_local_row = index.assume %code_local_row0 [range(%code_local_row0, 0, 63)] : index + %code_channel = index.add %channel_tile_base, %code_local_row : index + %valid_code_channel = index.cmp ult, %code_channel, %bounded_output_size : index + %loads_q4_code = scalar.andi %is_q4, %valid_code_channel : i1 + %loaded_q_word = scf.if %loads_q4_code -> (vector<1xi32>) { + %q4_group_pair = index.assume %quant_group_outer [range(%quant_group_outer, 0, 3)] : index + %code_channel_byte_add = index.scale %code_channel, %weight_row_bytes : index, offset -> offset + %code_row_byte_base = index.add %expert_byte_base, %code_channel_byte_add : offset + %q_word = func.call @ggml_q4k_split_wmma_load_code(%weight_noalias, %code_row_byte_base, %quant_block, %q4_group_pair, %load_packet) : (buffer, offset, index, index, index) -> (vector<1xi32>) + scf.yield %q_word : vector<1xi32> + } else { + scf.yield %c0_i32x1 : vector<1xi32> + } + %updates_q_word0 = index.cmp eq, %code_row_offset, %c0 : index + %updates_q_word1 = index.cmp eq, %code_row_offset, %c16 : index + %updates_q_word2 = index.cmp eq, %code_row_offset, %c32 : index + %updates_q_word3 = index.cmp eq, %code_row_offset, %c48 : index + %next_q_word0 = scf.select %updates_q_word0, %loaded_q_word, %prior_q_word0 : vector<1xi32> + %next_q_word1 = scf.select %updates_q_word1, %loaded_q_word, %prior_q_word1 : vector<1xi32> + %next_q_word2 = scf.select %updates_q_word2, %loaded_q_word, %prior_q_word2 : vector<1xi32> + %next_q_word3 = scf.select %updates_q_word3, %loaded_q_word, %prior_q_word3 : vector<1xi32> + scf.yield %next_q_word0, %next_q_word1, %next_q_word2, %next_q_word3 : vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32> + } + // Four Q6_K groups in one half-block share a QH word and consume two + // nibbles from each of two QL words. Retain the three packed words for + // all four groups. The Q4 specialization erases this acquisition path. + %q6_ql00, %q6_ql10, %q6_qh0, %q6_ql01, %q6_ql11, %q6_qh1, %q6_ql02, %q6_ql12, %q6_qh2, %q6_ql03, %q6_ql13, %q6_qh3 = scf.for %q6_code_row_offset = [%c0 to %c64 step %c16](%prior_q6_ql00 = %c0_i32x1 : vector<1xi32>, %prior_q6_ql10 = %c0_i32x1 : vector<1xi32>, %prior_q6_qh0 = %c0_i32x1 : vector<1xi32>, %prior_q6_ql01 = %c0_i32x1 : vector<1xi32>, %prior_q6_ql11 = %c0_i32x1 : vector<1xi32>, %prior_q6_qh1 = %c0_i32x1 : vector<1xi32>, %prior_q6_ql02 = %c0_i32x1 : vector<1xi32>, %prior_q6_ql12 = %c0_i32x1 : vector<1xi32>, %prior_q6_qh2 = %c0_i32x1 : vector<1xi32>, %prior_q6_ql03 = %c0_i32x1 : vector<1xi32>, %prior_q6_ql13 = %c0_i32x1 : vector<1xi32>, %prior_q6_qh3 = %c0_i32x1 : vector<1xi32>) -> (vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>) unroll { + %q6_code_local_row0 = index.add %load_row, %q6_code_row_offset : index + %q6_code_local_row = index.assume %q6_code_local_row0 [range(%q6_code_local_row0, 0, 63)] : index + %q6_code_channel = index.add %channel_tile_base, %q6_code_local_row : index + %valid_q6_code_channel = index.cmp ult, %q6_code_channel, %bounded_output_size : index + %loads_q6_code = scalar.andi %is_q6, %valid_q6_code_channel : i1 + %loaded_q6_ql0, %loaded_q6_ql1, %loaded_q6_qh = scf.if %loads_q6_code -> (vector<1xi32>, vector<1xi32>, vector<1xi32>) { + %q6_half = index.assume %quant_group_outer [range(%quant_group_outer, 0, 1)] : index + %q6_code_channel_byte_add = index.scale %q6_code_channel, %weight_row_bytes : index, offset -> offset + %q6_code_row_byte_base = index.add %expert_byte_base, %q6_code_channel_byte_add : offset + %ql0_word, %ql1_word, %qh_word = func.call @ggml_q6k_split_wmma_load_half_codes(%weight_noalias, %q6_code_row_byte_base, %quant_block, %q6_half, %load_packet) : (buffer, offset, index, index, index) -> (vector<1xi32>, vector<1xi32>, vector<1xi32>) + scf.yield %ql0_word, %ql1_word, %qh_word : vector<1xi32>, vector<1xi32>, vector<1xi32> + } else { + scf.yield %c0_i32x1, %c0_i32x1, %c0_i32x1 : vector<1xi32>, vector<1xi32>, vector<1xi32> + } + %updates_q6_code0 = index.cmp eq, %q6_code_row_offset, %c0 : index + %updates_q6_code1 = index.cmp eq, %q6_code_row_offset, %c16 : index + %updates_q6_code2 = index.cmp eq, %q6_code_row_offset, %c32 : index + %updates_q6_code3 = index.cmp eq, %q6_code_row_offset, %c48 : index + %next_q6_ql00 = scf.select %updates_q6_code0, %loaded_q6_ql0, %prior_q6_ql00 : vector<1xi32> + %next_q6_ql10 = scf.select %updates_q6_code0, %loaded_q6_ql1, %prior_q6_ql10 : vector<1xi32> + %next_q6_qh0 = scf.select %updates_q6_code0, %loaded_q6_qh, %prior_q6_qh0 : vector<1xi32> + %next_q6_ql01 = scf.select %updates_q6_code1, %loaded_q6_ql0, %prior_q6_ql01 : vector<1xi32> + %next_q6_ql11 = scf.select %updates_q6_code1, %loaded_q6_ql1, %prior_q6_ql11 : vector<1xi32> + %next_q6_qh1 = scf.select %updates_q6_code1, %loaded_q6_qh, %prior_q6_qh1 : vector<1xi32> + %next_q6_ql02 = scf.select %updates_q6_code2, %loaded_q6_ql0, %prior_q6_ql02 : vector<1xi32> + %next_q6_ql12 = scf.select %updates_q6_code2, %loaded_q6_ql1, %prior_q6_ql12 : vector<1xi32> + %next_q6_qh2 = scf.select %updates_q6_code2, %loaded_q6_qh, %prior_q6_qh2 : vector<1xi32> + %next_q6_ql03 = scf.select %updates_q6_code3, %loaded_q6_ql0, %prior_q6_ql03 : vector<1xi32> + %next_q6_ql13 = scf.select %updates_q6_code3, %loaded_q6_ql1, %prior_q6_ql13 : vector<1xi32> + %next_q6_qh3 = scf.select %updates_q6_code3, %loaded_q6_qh, %prior_q6_qh3 : vector<1xi32> + scf.yield %next_q6_ql00, %next_q6_ql10, %next_q6_qh0, %next_q6_ql01, %next_q6_ql11, %next_q6_qh1, %next_q6_ql02, %next_q6_ql12, %next_q6_qh2, %next_q6_ql03, %next_q6_ql13, %next_q6_qh3 : vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32> + } + %outer_result00, %outer_result01, %outer_result10, %outer_result11 = scf.for %group_within_outer = [%c0 to %groups_per_outer step %c1](%group_acc00 = %acc00 : vector<8xf16>, %group_acc01 = %acc01 : vector<8xf16>, %group_acc10 = %acc10 : vector<8xf16>, %group_acc11 = %acc11 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16>) { + %quant_group_base = index.mul %quant_group_outer, %groups_per_outer : index + %quant_group0 = index.add %quant_group_base, %group_within_outer : index + %quant_group = index.assume %quant_group0 [range(%quant_group0, 0, 7)] : index + %block_k_base = index.mul %quant_block, %c256 : index + %group_k_add = index.mul %quant_group, %c32 : index + %k_origin = index.add %block_k_base, %group_k_add : index + scf.for %row_offset = [%c0 to %c64 step %c16] unroll { + %local_row0 = index.add %load_row, %row_offset : index + %local_row = index.assume %local_row0 [range(%local_row0, 0, 63)] : index + %selects_row0 = index.cmp eq, %row_offset, %c0 : index + %selects_row2 = index.cmp eq, %row_offset, %c32 : index + %selects_low_row_pair = index.cmp ult, %row_offset, %c32 : index + %selected_header01 = scf.select %selects_row0, %header0, %header1 : vector<4xi32> + %selected_header23 = scf.select %selects_row2, %header2, %header3 : vector<4xi32> + %selected_header = scf.select %selects_low_row_pair, %selected_header01, %selected_header23 : vector<4xi32> + %selected_q_word01 = scf.select %selects_row0, %q_word0, %q_word1 : vector<1xi32> + %selected_q_word23 = scf.select %selects_row2, %q_word2, %q_word3 : vector<1xi32> + %selected_q_word = scf.select %selects_low_row_pair, %selected_q_word01, %selected_q_word23 : vector<1xi32> + %selected_q6_ql0_01 = scf.select %selects_row0, %q6_ql00, %q6_ql01 : vector<1xi32> + %selected_q6_ql0_23 = scf.select %selects_row2, %q6_ql02, %q6_ql03 : vector<1xi32> + %selected_q6_ql0 = scf.select %selects_low_row_pair, %selected_q6_ql0_01, %selected_q6_ql0_23 : vector<1xi32> + %selected_q6_ql1_01 = scf.select %selects_row0, %q6_ql10, %q6_ql11 : vector<1xi32> + %selected_q6_ql1_23 = scf.select %selects_row2, %q6_ql12, %q6_ql13 : vector<1xi32> + %selected_q6_ql1 = scf.select %selects_low_row_pair, %selected_q6_ql1_01, %selected_q6_ql1_23 : vector<1xi32> + %selected_q6_qh01 = scf.select %selects_row0, %q6_qh0, %q6_qh1 : vector<1xi32> + %selected_q6_qh23 = scf.select %selects_row2, %q6_qh2, %q6_qh3 : vector<1xi32> + %selected_q6_qh = scf.select %selects_low_row_pair, %selected_q6_qh01, %selected_q6_qh23 : vector<1xi32> + %channel = index.add %channel_tile_base, %local_row : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %weight_values = scf.if %valid_channel -> (vector<4xf16>) { + %decoded = scf.if %is_q6 -> (vector<4xf16>) { + %channel_byte_add = index.scale %channel, %weight_row_bytes : index, offset -> offset + %row_byte_base = index.add %expert_byte_base, %channel_byte_add : offset + %q6_values = func.call @ggml_q6k_split_wmma_vector4_from_half_codes(%weight_noalias, %row_byte_base, %quant_block, %quant_group, %load_packet, %selected_q6_ql0, %selected_q6_ql1, %selected_q6_qh) : (buffer, offset, index, index, index, vector<1xi32>, vector<1xi32>, vector<1xi32>) -> (vector<4xf16>) + scf.yield %q6_values : vector<4xf16> + } else { + %q4_values = func.call @ggml_q4k_split_wmma_vector4_from_header_code(%quant_group, %selected_header, %selected_q_word) : (index, vector<4xi32>, vector<1xi32>) -> (vector<4xf16>) + scf.yield %q4_values : vector<4xf16> + } + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + %is_activation_row = index.cmp ult, %local_row, %c32 : index + vector.store %weight_values, %weight_stage_view[%local_row, %load_k] : vector<4xf16>, view<64x40xf16> + scf.if %is_activation_row { + %activation_row = index.assume %local_row [range(%local_row, 0, 31)] : index + %assignment_i32 = view.load %route_stage_view[%activation_row] : view<32xi32> -> i32 + %valid_assignment = scalar.cmpi sge, %assignment_i32, %c0_i32 : i32 + %activation_values = scf.if %valid_assignment -> (vector<4xf16>) { + %assignment0 = index.cast %assignment_i32 : i32 to index + %assignment = index.assume %assignment0 [range(%assignment0, 0, 16383)] : index + %bounded_assignment, %bounded_assignment_count = index.assume %assignment, %assignment_count [lt(%assignment, %assignment_count)] : index, index + %input_k = index.add %k_origin, %load_k : index + %loaded = vector.load %input_view[%bounded_assignment, %input_k] : view<[%assignment_count]x[%input_size]xf16> -> vector<4xf16> + scf.yield %loaded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %activation_values, %activation_stage_physical_view[%activation_row, %load_k] : vector<4xf16>, view<32x40xf16> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %next00, %next01, %next10, %next11 = scf.for %k_half = [%c0 to %c32 step %c16](%half_acc00 = %group_acc00 : vector<8xf16>, %half_acc01 = %group_acc01 : vector<8xf16>, %half_acc10 = %group_acc10 : vector<8xf16>, %half_acc11 = %group_acc11 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16>) unroll { + %lhs0 = vector.fragment.load %weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %lhs1 = vector.fragment.load %weight_stage_view[%subgroup_channel1, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %rhs0 = vector.fragment.load %activation_fragment_view[%k_half, %c0] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %rhs1 = vector.fragment.load %activation_fragment_view[%k_half, %c16] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %half_next00 = vector.mma %lhs0, %rhs0, %half_acc00 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %half_next01 = vector.mma %lhs0, %rhs1, %half_acc01 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %half_next10 = vector.mma %lhs1, %rhs0, %half_acc10 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %half_next11 = vector.mma %lhs1, %rhs1, %half_acc11 : vector<16xf16>, vector<16xf16>, vector<8xf16> + scf.yield %half_next00, %half_next01, %half_next10, %half_next11 : vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %next00, %next01, %next10, %next11 : vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16> + } + scf.yield %outer_result00, %outer_result01, %outer_result10, %outer_result11 : vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16> + } + scf.yield %block_result00, %block_result01, %block_result10, %block_result11 : vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16> + } + // WMMA produces [channel][route] fragments. Transpose through wave-private + // LDS so each lane publishes four adjacent channels for one assignment. + %publish_route0 = index.div %lane, %c4 : index + %publish_route = index.assume %publish_route0 [range(%publish_route0, 0, 15)] : index + %publish_packet0 = index.rem %lane, %c4 : index + %publish_packet = index.assume %publish_packet0 [range(%publish_packet0, 0, 3)] : index + %publish_channel_add = index.mul %publish_packet, %c4 : index + %local_route1 = index.add %c16, %publish_route : index + %assignment0_i32 = view.load %route_stage_view[%publish_route] : view<32xi32> -> i32 + %assignment1_i32 = view.load %route_stage_view[%local_route1] : view<32xi32> -> i32 + %assignment0_nonnegative = scalar.cmpi sge, %assignment0_i32, %c0_i32 : i32 + %assignment1_nonnegative = scalar.cmpi sge, %assignment1_i32, %c0_i32 : i32 + %safe_assignment0_i32 = scf.select %assignment0_nonnegative, %assignment0_i32, %c0_i32 : i32 + %safe_assignment1_i32 = scf.select %assignment1_nonnegative, %assignment1_i32, %c0_i32 : i32 + %safe_assignment0_0 = index.cast %safe_assignment0_i32 : i32 to index + %safe_assignment1_0 = index.cast %safe_assignment1_i32 : i32 to index + %safe_assignment0 = index.assume %safe_assignment0_0 [range(%safe_assignment0_0, 0, 16383)] : index + %safe_assignment1 = index.assume %safe_assignment1_0 [range(%safe_assignment1_0, 0, 16383)] : index + %bounded_assignment0, %bounded_assignment_count0 = index.assume %safe_assignment0, %assignment_count [lt(%safe_assignment0, %assignment_count)] : index, index + %bounded_assignment1, %bounded_assignment_count1 = index.assume %safe_assignment1, %assignment_count [lt(%safe_assignment1, %assignment_count)] : index, index + %subgroup_channel_base = index.add %channel_tile_base, %subgroup_channel_add : index + %channel0 = index.add %subgroup_channel_base, %publish_channel_add : index + %channel1_base = index.add %subgroup_channel_base, %c16 : index + %channel1 = index.add %channel1_base, %publish_channel_add : index + %valid_channel0 = index.cmp ult, %channel0, %bounded_output_size : index + %valid_channel1 = index.cmp ult, %channel1, %bounded_output_size : index + %writes00 = scalar.andi %assignment0_nonnegative, %valid_channel0 : i1 + %writes01 = scalar.andi %assignment1_nonnegative, %valid_channel0 : i1 + %writes10 = scalar.andi %assignment0_nonnegative, %valid_channel1 : i1 + %writes11 = scalar.andi %assignment1_nonnegative, %valid_channel1 : i1 + vector.fragment.store %result00, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<8xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + %values00 = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf16> -> vector<4xf16> + template.apply<@ggml.mul_mat_id_f16_f16_wmma.publish_vector4>(%writes00, %bounded_assignment0, %channel0, %token_count, %bounded_route_count, %bounded_output_size, %values00, %output) : (i1, index, index, index, index, index, vector<4xf16>, buffer) + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result01, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<8xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + %values01 = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf16> -> vector<4xf16> + template.apply<@ggml.mul_mat_id_f16_f16_wmma.publish_vector4>(%writes01, %bounded_assignment1, %channel0, %token_count, %bounded_route_count, %bounded_output_size, %values01, %output) : (i1, index, index, index, index, index, vector<4xf16>, buffer) + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result10, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<8xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + %values10 = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf16> -> vector<4xf16> + template.apply<@ggml.mul_mat_id_f16_f16_wmma.publish_vector4>(%writes10, %bounded_assignment0, %channel1, %token_count, %bounded_route_count, %bounded_output_size, %values10, %output) : (i1, index, index, index, index, index, vector<4xf16>, buffer) + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result11, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<8xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + %values11 = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf16> -> vector<4xf16> + template.apply<@ggml.mul_mat_id_f16_f16_wmma.publish_vector4>(%writes11, %bounded_assignment1, %channel1, %token_count, %bounded_route_count, %bounded_output_size, %values11, %output) : (i1, index, index, index, index, index, vector<4xf16>, buffer) + template.apply<@ggml.mul_mat_id_f16_f16_wmma.finish_tile>(%token_count, %bounded_route_count, %bounded_output_size, %bounded_expert_count, %channel_tile, %bounded_expert, %route_tile_base, %expert_route_count, %expert_table, %output) : (index, index, index, index, index, index, index, index, buffer, buffer) + // Route, operand, and result stages are reused by the next concentrated + // routing partition. All waves must finish publication before reuse. + kernel.barrier scope(workgroup) ordering(acq_rel) + } + template.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_f32_f32_postops.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_f32_f32_postops.loom new file mode 100644 index 000000000000..d3497b0a1895 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_f32_f32_postops.loom @@ -0,0 +1,167 @@ +template.decl @ggml.mul_mat_id_f32_f32_wmma.add_bias_vector4(%output_size0: index, %channel: index, %values: vector<4xf32>, %bias: buffer) -> (vector<4xf32>) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.add_residual_vector4(%token_count0: index, %route_count0: index, %output_size0: index, %assignment0: index, %channel: index, %values: vector<4xf32>, %residual_input: buffer) -> (vector<4xf32>) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.finish_routed_rmsnorm_weight(%token_count0: index, %route_count0: index, %output_size0: index, %expert_count0: index, %output_tile_count0: index, %maximum_partition_count0: index, %epsilon: f32, %descriptor_ordinal0: index, %expert0: index, %route_tile_base0: index, %partition_row_count0: index, %expert_table: buffer, %source: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.store_vector4(%token_count0: index, %route_count0: index, %output_size0: index, %assignment0: index, %channel: index, %values: vector<4xf32>, %output: buffer) + +template.def<@ggml.mul_mat_id_f32_f32_wmma.add_bias_vector4> device @ggml_mul_mat_id_f32_f32_wmma_add_bias_vector4(%output_size0: index, %channel: index, %values: vector<4xf32>, %bias: buffer) -> (vector<4xf32>) { + %output_size = index.assume %output_size0 [range(%output_size0, 1, 262144)] : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %bias_noalias = buffer.assume.noalias %bias : buffer + %bias_view = buffer.view %bias_noalias[%c0_offset] : buffer -> view<[%output_size]xf32> + %mask = vector.mask.range [%channel to %output_size step %c1] : index -> vector<4xi1> + %bias_values = vector.load.mask %bias_view[%channel], %mask, %c0_f32x4 : view<[%output_size]xf32>, vector<4xi1>, vector<4xf32> + %sum = vector.addf %values, %bias_values : vector<4xf32> + template.return %sum : vector<4xf32> +} + +template.def<@ggml.mul_mat_id_f32_f32_wmma.add_residual_vector4> device @ggml_mul_mat_id_f32_f32_wmma_add_residual_vector4(%token_count0: index, %route_count0: index, %output_size0: index, %assignment0: index, %channel: index, %values: vector<4xf32>, %residual_input: buffer) -> (vector<4xf32>) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %route_count = index.assume %route_count0 [range(%route_count0, 1, 32)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 1, 262144)] : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %row_count = index.mul %token_count, %route_count : index + %assignment, %assignment_count = index.assume %assignment0, %row_count [lt(%assignment0, %row_count)] : index, index + %residual_input_noalias = buffer.assume.noalias %residual_input : buffer + %residual_input_view = buffer.view %residual_input_noalias[%c0_offset] : buffer -> view<[%row_count]x[%output_size]xf32> + %mask = vector.mask.range [%channel to %output_size step %c1] : index -> vector<4xi1> + %residual = vector.load.mask %residual_input_view[%assignment, %channel], %mask, %c0_f32x4 : view<[%row_count]x[%output_size]xf32>, vector<4xi1>, vector<4xf32> + %sum = vector.addf %values, %residual : vector<4xf32> + template.return %sum : vector<4xf32> +} + +template.def<@ggml.mul_mat_id_f32_f32_wmma.store_vector4> device @ggml_mul_mat_id_f32_f32_wmma_store_vector4(%token_count0: index, %route_count0: index, %output_size0: index, %assignment0: index, %channel: index, %values: vector<4xf32>, %output: buffer) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %route_count = index.assume %route_count0 [range(%route_count0, 1, 32)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 1, 262144)] : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %row_count = index.mul %token_count, %route_count : index + %assignment, %assignment_count = index.assume %assignment0, %row_count [lt(%assignment0, %row_count)] : index, index + %output_noalias = buffer.assume.noalias %output : buffer + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%row_count]x[%output_size]xf32> + %mask = vector.mask.range [%channel to %output_size step %c1] : index -> vector<4xi1> + vector.store.mask %values, %output_view[%assignment, %channel], %mask : vector<4xf32>, view<[%row_count]x[%output_size]xf32>, vector<4xi1> + template.return +} + +template.def<@ggml.mul_mat_id_f32_f32_wmma.finish_routed_rmsnorm_weight> device @ggml_mul_mat_id_f32_f32_wmma_finish_routed_rmsnorm_weight(%token_count0: index, %route_count0: index, %output_size0: index, %expert_count0: index, %output_tile_count0: index, %maximum_partition_count0: index, %epsilon: f32, %descriptor_ordinal0: index, %expert0: index, %route_tile_base0: index, %partition_row_count0: index, %expert_table: buffer, %source: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %route_count = index.assume %route_count0 [range(%route_count0, 1, 32)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 128, 32768), mul(%output_size0, 128)] : index + %expert_count = index.assume %expert_count0 [range(%expert_count0, 1, 512)] : index + %output_tile_count = index.assume %output_tile_count0 [range(%output_tile_count0, 1, 512)] : index + %maximum_partition_count = index.assume %maximum_partition_count0 [range(%maximum_partition_count0, 1, 2560)] : index + %descriptor_ordinal, %table_partition_count = index.assume %descriptor_ordinal0, %maximum_partition_count [lt(%descriptor_ordinal0, %maximum_partition_count)] : index, index + %expert, %table_expert_count = index.assume %expert0, %expert_count [lt(%expert0, %expert_count)] : index, index + %route_tile_base, %table_token_count0 = index.assume %route_tile_base0, %token_count [lt(%route_tile_base0, %token_count)] : index, index + %partition_row_count = index.assume %partition_row_count0 [range(%partition_row_count0, 1, 32)] : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 1)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c32 = index.constant 32 : index + %c128 = index.constant 128 : index + %c0_offset = index.constant 0 : offset + %c4_bytes = index.constant 4 : offset + %counter_scratch_bytes = index.constant 4 : offset + %rms_scratch_bytes = index.constant 512 : offset + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c0_f32 = scalar.constant 0.0 : f32 + %assignment_table_byte_base = index.scale %expert_count, %c4_bytes : index, offset -> offset + %row_count = index.mul %token_count, %route_count : index + %hidden_size_i32 = index.cast %output_size : index to i32 + %hidden_size_f32 = scalar.sitofp %hidden_size_i32 : i32 to f32 + %output_tile_count_i32 = index.cast %output_tile_count : index to i32 + %last_channel_ordinal_i32 = scalar.subi %output_tile_count_i32, %c1_i32 : i32 + %negative_output_tile_count_i32 = scalar.subi %c0_i32, %output_tile_count_i32 : i32 + %expert_table_noalias, %source_noalias, %norm_weight_noalias, %normalized_output_noalias, %completion_counters_noalias = buffer.assume.noalias %expert_table, %source, %norm_weight, %normalized_output, %completion_counters : buffer, buffer, buffer, buffer, buffer + %completion_counters_aligned = buffer.assume.alignment %completion_counters_noalias {minimum_alignment = 16} : buffer + %assignment_view = buffer.view %expert_table_noalias[%assignment_table_byte_base] : buffer -> view<[%expert_count]x[%token_count]xi32> + %source_view = buffer.view %source_noalias[%c0_offset] : buffer -> view<[%row_count]x[%output_size]xf32> + %norm_weight_view = buffer.view %norm_weight_noalias[%c0_offset] : buffer -> view<[%output_size]xf32> + %normalized_output_view = buffer.view %normalized_output_noalias[%c0_offset] : buffer -> view<[%row_count]x[%output_size]xf32> + %completion_counter_view = buffer.view %completion_counters_aligned[%c0_offset] : buffer -> view<[%maximum_partition_count]xi32> + %counter_scratch = buffer.alloca align(4) %counter_scratch_bytes : buffer + %rms_scratch = buffer.alloca align(16) %rms_scratch_bytes : buffer + %counter_scratch_view = buffer.view %counter_scratch[%c0_offset] : buffer -> view<1xi32> + %rms_scratch_view = buffer.view %rms_scratch[%c0_offset] : buffer -> view<128xf32> + kernel.barrier scope(workgroup) ordering(release) + %workitem_is_zero = index.cmp eq, %workitem, %c0 : index + scf.if %workitem_is_zero { + %old_counter = view.atomic.rmw %c1_i32, %completion_counter_view[%descriptor_ordinal] {ordering = acq_rel, scope = device} : i32, view<[%maximum_partition_count]xi32> -> i32 + view.store %old_counter, %counter_scratch_view[%c0] : i32, view<1xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %old_counter = view.load %counter_scratch_view[%c0] : view<1xi32> -> i32 + %is_last_channel_tile = scalar.cmpi eq, %old_counter, %last_channel_ordinal_i32 : i32 + scf.if %is_last_channel_tile { + kernel.barrier scope(workgroup) ordering(acquire) + scf.for %row_offset = [%c0 to %c32 step %c1] { + %valid_row = index.cmp ult, %row_offset, %partition_row_count : index + scf.if %valid_row { + %assignment_ordinal0 = index.add %route_tile_base, %row_offset : index + %assignment_ordinal, %table_token_count = index.assume %assignment_ordinal0, %token_count [lt(%assignment_ordinal0, %token_count)] : index, index + %assignment_i32 = view.load %assignment_view[%expert, %assignment_ordinal] : view<[%expert_count]x[%token_count]xi32> -> i32 + %assignment0 = index.cast %assignment_i32 : i32 to index + %assignment = index.assume %assignment0 [range(%assignment0, 0, 65535)] : index + %bounded_assignment, %assignment_count = index.assume %assignment, %row_count [lt(%assignment, %row_count)] : index, index + %thread_sum = scf.for %channel = [%workitem to %output_size step %c128](%running_sum = %c0_f32 : f32) -> (f32) { + %value = view.load %source_view[%bounded_assignment, %channel] : view<[%row_count]x[%output_size]xf32> -> f32 + %square = scalar.mulf %value, %value : f32 + %next_sum = scalar.addf %running_sum, %square : f32 + scf.yield %next_sum : f32 + } + %subgroup_sum = kernel.subgroup.reduce %thread_sum : f32 + %is_subgroup_leader = index.cmp eq, %lane, %c0 : index + scf.if %is_subgroup_leader { + view.store %subgroup_sum, %rms_scratch_view[%subgroup] : f32, view<128xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %is_reduction_subgroup = index.cmp eq, %subgroup, %c0 : index + %is_reduction_lane = index.cmp ult, %lane, %c2 : index + %loads_subgroup_sum = scalar.andi %is_reduction_subgroup, %is_reduction_lane : i1 + %subgroup_partial = scf.if %loads_subgroup_sum -> (f32) { + %value = view.load %rms_scratch_view[%lane] : view<128xf32> -> f32 + scf.yield %value : f32 + } else { + scf.yield %c0_f32 : f32 + } + %row_sum = kernel.subgroup.reduce %subgroup_partial : f32 + %writes_scale = scalar.andi %is_reduction_subgroup, %is_subgroup_leader : i1 + scf.if %writes_scale { + %mean = scalar.divf %row_sum, %hidden_size_f32 : f32 + %biased_mean = scalar.addf %mean, %epsilon : f32 + %scale = scalar.rsqrtf %biased_mean : f32 + view.store %scale, %rms_scratch_view[%c0] : f32, view<128xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %scale = view.load %rms_scratch_view[%c0] : view<128xf32> -> f32 + scf.for %channel = [%workitem to %output_size step %c128] { + %value = view.load %source_view[%bounded_assignment, %channel] : view<[%row_count]x[%output_size]xf32> -> f32 + %weight_value = view.load %norm_weight_view[%channel] : view<[%output_size]xf32> -> f32 + %scaled = scalar.mulf %value, %scale : f32 + %normalized = scalar.mulf %scaled, %weight_value : f32 + view.store %normalized, %normalized_output_view[%bounded_assignment, %channel] : f32, view<[%row_count]x[%output_size]xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + } + } + kernel.barrier scope(workgroup) ordering(release) + scf.if %workitem_is_zero { + view.atomic.reduce %negative_output_tile_count_i32, %completion_counter_view[%descriptor_ordinal] {ordering = release, scope = device} : i32, view<[%maximum_partition_count]xi32> + } + } + template.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_f32_f32_wmma_core.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_f32_f32_wmma_core.loom new file mode 100644 index 000000000000..4be48c72cfed --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_f32_f32_wmma_core.loom @@ -0,0 +1,309 @@ +// Shared routed expert WMMA matmul pipeline for GGML MUL_MAT_ID. +// +// Op wrappers provide publication and tile-finalization hooks by defining: +// ggml.mul_mat_id_f32_f32_wmma.publish_vector4 +// ggml.mul_mat_id_f32_f32_wmma.finish_tile +template.decl @ggml.mul_mat_id_f32_f32_wmma.core(%weight_format: index, %token_count: index, %input_size: index, %output_size: index, %route_count: index, %input_route_count: index, %expert_count: index, %epsilon: f32, %channel_tile: index, %partition_ordinal: index, %input: buffer, %expert_table: buffer, %partition_table: buffer, %weight: buffer, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.finish_tile(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: f32, %arg7: index, %arg8: index, %arg9: index, %arg10: index, %arg11: index, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer, %arg16: buffer, %arg17: buffer, %arg18: buffer) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.publish_vector4(%arg0: i1, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: index, %arg7: index, %arg8: vector<4xf32>, %arg9: buffer, %arg10: buffer, %arg11: buffer, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer) + +func.decl @ggml_dequant_weight_tile_bytes(%weight_format: index) -> (offset) + +func.decl @ggml_dequant_weight_row_bytes(%weight_format: index, %hidden_size: index) -> (offset) + +func.decl @ggml_iq4nl_table_i8() -> (vector<16xi8>) + +func.decl @ggml_iq3s_grid_f32_halves() -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + +func.decl @ggml_dequant_f16_vector4(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight_format: index, %weight: buffer, %row_byte_base: offset, %input_size: index, %quant_block: index, %quant_group: index, %packet: index, %k: index) -> (vector<4xf16>) + +// Common descriptor packing uses a 9-bit expert ordinal, a 6-bit 32-row +// partition ordinal, and a 5-bit row count minus one. +func.def inline @ggml_moe_unpack_expert_partition_descriptor(%descriptor: i32) -> (index, index, index) { + %c1_i32 = scalar.constant 1 : i32 + %c5_i32 = scalar.constant 5 : i32 + %c9_i32 = scalar.constant 9 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c31_i32 = scalar.constant 31 : i32 + %c63_i32 = scalar.constant 63 : i32 + %c511_i32 = scalar.constant 511 : i32 + %expert_i32 = scalar.andi %descriptor, %c511_i32 : i32 + %partition_shifted_i32 = scalar.shrui %descriptor, %c9_i32 : i32 + %partition_i32 = scalar.andi %partition_shifted_i32, %c63_i32 : i32 + %route_tile_base_i32 = scalar.shli %partition_i32, %c5_i32 : i32 + %row_count_shifted_i32 = scalar.shrui %descriptor, %c15_i32 : i32 + %row_count_minus_one_i32 = scalar.andi %row_count_shifted_i32, %c31_i32 : i32 + %partition_row_count_i32 = scalar.addi %row_count_minus_one_i32, %c1_i32 : i32 + %expert0 = index.cast %expert_i32 : i32 to index + %expert = index.assume %expert0 [range(%expert0, 0, 511)] : index + %route_tile_base0 = index.cast %route_tile_base_i32 : i32 to index + %route_tile_base = index.assume %route_tile_base0 [range(%route_tile_base0, 0, 2016)] : index + %partition_row_count0 = index.cast %partition_row_count_i32 : i32 to index + %partition_row_count = index.assume %partition_row_count0 [range(%partition_row_count0, 1, 32)] : index + func.return %expert, %route_tile_base, %partition_row_count : index, index, index +} + +template.def<@ggml.mul_mat_id_f32_f32_wmma.core> device @ggml_mul_mat_id_f32_f32_wmma_core(%weight_format: index, %token_count: index, %input_size: index, %output_size: index, %route_count: index, %input_route_count: index, %expert_count: index, %epsilon: f32, %channel_tile: index, %partition_ordinal: index, %input: buffer, %expert_table: buffer, %partition_table: buffer, %weight: buffer, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_route_count = index.assume %route_count [range(%route_count, 1, 32)] : index + %bounded_input_route_count = index.assume %input_route_count [range(%input_route_count, 1, 32), le(%input_route_count, %bounded_route_count)] : index + %bounded_expert_count = index.assume %expert_count [range(%expert_count, 1, 512)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 262144)] : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 1)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c40 = index.constant 40 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %c4_bytes = index.constant 4 : offset + %weight_stage_bytes = index.constant 5120 : offset + %activation_stage_bytes = index.constant 2560 : offset + %route_stage_bytes = index.constant 128 : offset + %wave_result_stage_bytes = index.constant 1024 : offset + %result_stage_bytes = index.constant 2048 : offset + %c0_i32 = scalar.constant 0 : i32 + %cn1_i32 = scalar.constant -1 : i32 + %c0_f16x4 = vector.constant 0.0 : vector<4xf16> + %zero_accumulator = vector.constant 0.0 : vector<4xf32> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %assignment_count = index.mul %bounded_token_count, %bounded_route_count : index + %padded_output_size = index.add %bounded_output_size, %c63 : index + %output_tile_count = index.div %padded_output_size, %c64 : index + %rounded_assignment_count = index.add %assignment_count, %c31 : index + %assignment_partition_count = index.div %rounded_assignment_count, %c32 : index + %maximum_partition_count = index.add %assignment_partition_count, %bounded_expert_count : index + %has_more_assignment_partitions = index.cmp ugt, %assignment_partition_count, %bounded_expert_count : index + %launch_partition_count = scf.select %has_more_assignment_partitions, %assignment_partition_count, %bounded_expert_count : index + %assignment_table_byte_base = index.scale %bounded_expert_count, %c4_bytes : index, offset -> offset + %c255 = index.constant 255 : index + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %padded_input_size = index.add %input_size, %c255 : index + %quant_block_count = index.div %padded_input_size, %c256 : index + %weight_row_bytes = func.call @ggml_dequant_weight_row_bytes(%weight_format, %input_size) : (index, index) -> (offset) + %weight_expert_bytes = index.scale %bounded_output_size, %weight_row_bytes : index, offset -> offset + %input_noalias, %expert_table_noalias, %partition_table_noalias, %weight_noalias = buffer.assume.noalias %input, %expert_table, %partition_table, %weight : buffer, buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_input_route_count]x[%input_size]xf32> + %assignment_view = buffer.view %expert_table_noalias[%assignment_table_byte_base] : buffer -> view<[%bounded_expert_count]x[%bounded_token_count]xi32> + %partition_count_view = buffer.view %partition_table_noalias[%c0_offset] : buffer -> view<1xi32> + %partition_descriptor_view = buffer.view %partition_table_noalias[%c4_bytes] : buffer -> view<[%maximum_partition_count]xi32> + %weight_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %activation_stage = buffer.alloca align(16) %activation_stage_bytes : buffer + %route_stage = buffer.alloca align(16) %route_stage_bytes : buffer + %result_stage = buffer.alloca align(16) %result_stage_bytes : buffer + %weight_stage_view = buffer.view %weight_stage[%c0_offset] : buffer -> view<64x40xf16> + %activation_stage_physical_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x40xf16> + %activation_fragment_layout = encoding.layout.strided [1, %c40] : encoding + %activation_fragment_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x32xf16, %activation_fragment_layout> + %route_stage_view = buffer.view %route_stage[%c0_offset] : buffer -> view<32xi32> + %wave_result_stage_offset = index.scale %subgroup, %wave_result_stage_bytes : index, offset -> offset + %result_fragment_layout = encoding.layout.strided [1, %c16] : encoding + %result_fragment_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf32, %result_fragment_layout> + %result_physical_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf32> + %channel_tile_base = index.mul %channel_tile, %c64 : index + %is_workitem_zero = index.cmp eq, %workitem, %c0 : index + %lane_partition_count_i32 = scf.if %is_workitem_zero -> (i32) { + %loaded = view.load %partition_count_view[%c0] : view<1xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %c0_i32 : i32 + } + %partition_count_reduced = kernel.workgroup.reduce %lane_partition_count_i32 : i32 + %partition_count_i32 = kernel.subgroup.broadcast.first %partition_count_reduced : i32 + %partition_count0 = index.cast %partition_count_i32 : i32 to index + %partition_count, %partition_capacity = index.assume %partition_count0, %maximum_partition_count [le(%partition_count0, %maximum_partition_count)] : index, index + scf.for %active_partition = [%partition_ordinal to %partition_count step %launch_partition_count] { + %descriptor_ordinal, %descriptor_count = index.assume %active_partition, %partition_count [lt(%active_partition, %partition_count)] : index, index + %table_descriptor_ordinal, %table_descriptor_capacity = index.assume %descriptor_ordinal, %maximum_partition_count [lt(%descriptor_ordinal, %maximum_partition_count)] : index, index + %lane_descriptor_i32 = scf.if %is_workitem_zero -> (i32) { + %loaded = view.load %partition_descriptor_view[%table_descriptor_ordinal] : view<[%maximum_partition_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %c0_i32 : i32 + } + %descriptor_reduced_i32 = kernel.workgroup.reduce %lane_descriptor_i32 : i32 + %descriptor_i32 = kernel.subgroup.broadcast.first %descriptor_reduced_i32 : i32 + %expert, %route_tile_base, %partition_row_count = func.call @ggml_moe_unpack_expert_partition_descriptor(%descriptor_i32) : (i32) -> (index, index, index) + %bounded_expert, %table_expert_count = index.assume %expert, %bounded_expert_count [lt(%expert, %bounded_expert_count)] : index, index + %loads_route = index.cmp ult, %workitem, %c32 : index + scf.if %loads_route { + %local_route = index.assume %workitem [range(%workitem, 0, 31)] : index + %assignment_ordinal = index.add %route_tile_base, %local_route : index + %valid_row = index.cmp ult, %local_route, %partition_row_count : index + %assignment_i32 = scf.if %valid_row -> (i32) { + %bounded_assignment_ordinal, %table_token_count = index.assume %assignment_ordinal, %bounded_token_count [lt(%assignment_ordinal, %bounded_token_count)] : index, index + %loaded = view.load %assignment_view[%bounded_expert, %bounded_assignment_ordinal] : view<[%bounded_expert_count]x[%bounded_token_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %cn1_i32 : i32 + } + view.store %assignment_i32, %route_stage_view[%local_route] : i32, view<32xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %load_packet = index.rem %workitem, %c8 : index + %load_k = index.mul %load_packet, %c4 : index + %load_row0 = index.div %workitem, %c8 : index + %load_row = index.assume %load_row0 [range(%load_row0, 0, 15)] : index + %expert_byte_base = index.scale %bounded_expert, %weight_expert_bytes : index, offset -> offset + %subgroup_channel_add = index.mul %subgroup, %c32 : index + %subgroup_channel1 = index.add %subgroup_channel_add, %c16 : index + %init00 = vector.fragment %zero_accumulator shape [%m, %n] : vector<4xf32> + %init01 = vector.fragment %zero_accumulator shape [%m, %n] : vector<4xf32> + %init10 = vector.fragment %zero_accumulator shape [%m, %n] : vector<4xf32> + %init11 = vector.fragment %zero_accumulator shape [%m, %n] : vector<4xf32> + %iq4nl_table = func.call @ggml_iq4nl_table_i8() : () -> (vector<16xi8>) + %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15 = func.call @ggml_iq3s_grid_f32_halves() : () -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + %result00, %result01, %result10, %result11 = scf.for %quant_block = [%c0 to %quant_block_count step %c1](%block_acc00 = %init00 : vector<4xf32>, %block_acc01 = %init01 : vector<4xf32>, %block_acc10 = %init10 : vector<4xf32>, %block_acc11 = %init11 : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + %block_result00, %block_result01, %block_result10, %block_result11 = scf.for %quant_group = [%c0 to %c8 step %c1](%acc00 = %block_acc00 : vector<4xf32>, %acc01 = %block_acc01 : vector<4xf32>, %acc10 = %block_acc10 : vector<4xf32>, %acc11 = %block_acc11 : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + %block_k_base = index.mul %quant_block, %c256 : index + %group_k_add = index.mul %quant_group, %c32 : index + %k_origin = index.add %block_k_base, %group_k_add : index + scf.for %row_offset = [%c0 to %c64 step %c16] unroll { + %local_row0 = index.add %load_row, %row_offset : index + %local_row = index.assume %local_row0 [range(%local_row0, 0, 63)] : index + %channel = index.add %channel_tile_base, %local_row : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %weight_k_tail = index.add %k_origin, %load_k : index + %valid_k = index.cmp ult, %weight_k_tail, %input_size : index + %valid_weight = scalar.andi %valid_channel, %valid_k : i1 + %weight_values = scf.if %valid_weight -> (vector<4xf16>) { + %channel_byte_add = index.scale %channel, %weight_row_bytes : index, offset -> offset + %row_byte_base = index.add %expert_byte_base, %channel_byte_add : offset + %weight_k = index.add %k_origin, %load_k : index + %decoded = func.call @ggml_dequant_f16_vector4(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight_format, %weight_noalias, %row_byte_base, %input_size, %quant_block, %quant_group, %load_packet, %weight_k) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, buffer, offset, index, index, index, index, index) -> (vector<4xf16>) + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + %is_activation_row = index.cmp ult, %local_row, %c32 : index + vector.store %weight_values, %weight_stage_view[%local_row, %load_k] : vector<4xf16>, view<64x40xf16> + scf.if %is_activation_row { + %activation_row = index.assume %local_row [range(%local_row, 0, 31)] : index + %assignment_i32 = view.load %route_stage_view[%activation_row] : view<32xi32> -> i32 + %valid_assignment = scalar.cmpi sge, %assignment_i32, %c0_i32 : i32 + %activation_k = index.add %k_origin, %load_k : index + %valid_input = index.cmp ult, %activation_k, %input_size : index + %valid_activation = scalar.andi %valid_assignment, %valid_input : i1 + %activation_values = scf.if %valid_activation -> (vector<4xf16>) { + %assignment0 = index.cast %assignment_i32 : i32 to index + %assignment = index.assume %assignment0 [range(%assignment0, 0, 65535)] : index + %bounded_assignment, %bounded_assignment_count = index.assume %assignment, %assignment_count [lt(%assignment, %assignment_count)] : index, index + %token0 = index.div %bounded_assignment, %bounded_route_count : index + %route0 = index.rem %bounded_assignment, %bounded_route_count : index + %token, %input_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %input_route0 = index.rem %route0, %bounded_input_route_count : index + %input_route = index.assume %input_route0 [range(%input_route0, 0, 31), lt(%input_route0, %bounded_input_route_count)] : index + %input_k = index.add %k_origin, %load_k : index + %mask = vector.mask.range [%input_k to %input_size step %c1] : index -> vector<4xi1> + %loaded = vector.load.mask %input_view[%token, %input_route, %input_k], %mask, %c0_f32x4 : view<[%bounded_token_count]x[%bounded_input_route_count]x[%input_size]xf32>, vector<4xi1>, vector<4xf32> + %converted = vector.fptrunc %loaded : vector<4xf32> to vector<4xf16> + scf.yield %converted : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %activation_values, %activation_stage_physical_view[%activation_row, %load_k] : vector<4xf16>, view<32x40xf16> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %next00, %next01, %next10, %next11 = scf.for %k_half = [%c0 to %c32 step %c16](%half_acc00 = %acc00 : vector<4xf32>, %half_acc01 = %acc01 : vector<4xf32>, %half_acc10 = %acc10 : vector<4xf32>, %half_acc11 = %acc11 : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) unroll { + %lhs0 = vector.fragment.load %weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %lhs1 = vector.fragment.load %weight_stage_view[%subgroup_channel1, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %rhs0 = vector.fragment.load %activation_fragment_view[%k_half, %c0] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %rhs1 = vector.fragment.load %activation_fragment_view[%k_half, %c16] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %half_next00 = vector.mma %lhs0, %rhs0, %half_acc00 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %half_next01 = vector.mma %lhs0, %rhs1, %half_acc01 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %half_next10 = vector.mma %lhs1, %rhs0, %half_acc10 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %half_next11 = vector.mma %lhs1, %rhs1, %half_acc11 : vector<16xf16>, vector<16xf16>, vector<4xf32> + scf.yield %half_next00, %half_next01, %half_next10, %half_next11 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %next00, %next01, %next10, %next11 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + scf.yield %block_result00, %block_result01, %block_result10, %block_result11 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + %publish_route0 = index.div %lane, %c4 : index + %publish_route = index.assume %publish_route0 [range(%publish_route0, 0, 15)] : index + %publish_packet0 = index.rem %lane, %c4 : index + %publish_packet = index.assume %publish_packet0 [range(%publish_packet0, 0, 3)] : index + %publish_channel_add = index.mul %publish_packet, %c4 : index + %local_route1 = index.add %c16, %publish_route : index + %assignment0_i32 = view.load %route_stage_view[%publish_route] : view<32xi32> -> i32 + %assignment1_i32 = view.load %route_stage_view[%local_route1] : view<32xi32> -> i32 + %assignment0_nonnegative = scalar.cmpi sge, %assignment0_i32, %c0_i32 : i32 + %assignment1_nonnegative = scalar.cmpi sge, %assignment1_i32, %c0_i32 : i32 + %safe_assignment0_i32 = scf.if %assignment0_nonnegative -> (i32) { + scf.yield %assignment0_i32 : i32 + } else { + scf.yield %c0_i32 : i32 + } + %safe_assignment1_i32 = scf.if %assignment1_nonnegative -> (i32) { + scf.yield %assignment1_i32 : i32 + } else { + scf.yield %c0_i32 : i32 + } + %safe_assignment0_0 = index.cast %safe_assignment0_i32 : i32 to index + %safe_assignment1_0 = index.cast %safe_assignment1_i32 : i32 to index + %safe_assignment0 = index.assume %safe_assignment0_0 [range(%safe_assignment0_0, 0, 65535)] : index + %safe_assignment1 = index.assume %safe_assignment1_0 [range(%safe_assignment1_0, 0, 65535)] : index + %bounded_assignment0, %bounded_assignment_count0 = index.assume %safe_assignment0, %assignment_count [lt(%safe_assignment0, %assignment_count)] : index, index + %bounded_assignment1, %bounded_assignment_count1 = index.assume %safe_assignment1, %assignment_count [lt(%safe_assignment1, %assignment_count)] : index, index + %assignment_token0 = index.div %bounded_assignment0, %bounded_route_count : index + %assignment_route0 = index.rem %bounded_assignment0, %bounded_route_count : index + %assignment_token1 = index.div %bounded_assignment1, %bounded_route_count : index + %assignment_route1 = index.rem %bounded_assignment1, %bounded_route_count : index + %subgroup_channel_base = index.add %channel_tile_base, %subgroup_channel_add : index + %channel0 = index.add %subgroup_channel_base, %publish_channel_add : index + %channel1_base = index.add %subgroup_channel_base, %c16 : index + %channel1 = index.add %channel1_base, %publish_channel_add : index + %valid_channel0 = index.cmp ult, %channel0, %bounded_output_size : index + %valid_channel1 = index.cmp ult, %channel1, %bounded_output_size : index + %writes00 = scalar.andi %assignment0_nonnegative, %valid_channel0 : i1 + %writes01 = scalar.andi %assignment1_nonnegative, %valid_channel0 : i1 + %writes10 = scalar.andi %assignment0_nonnegative, %valid_channel1 : i1 + %writes11 = scalar.andi %assignment1_nonnegative, %valid_channel1 : i1 + vector.fragment.store %result00, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes00 { + %values = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + template.apply<@ggml.mul_mat_id_f32_f32_wmma.publish_vector4>(%writes00, %bounded_assignment0, %assignment_token0, %assignment_route0, %channel0, %bounded_token_count, %bounded_route_count, %bounded_output_size, %values, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (i1, index, index, index, index, index, index, index, vector<4xf32>, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result01, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes01 { + %values = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + template.apply<@ggml.mul_mat_id_f32_f32_wmma.publish_vector4>(%writes01, %bounded_assignment1, %assignment_token1, %assignment_route1, %channel0, %bounded_token_count, %bounded_route_count, %bounded_output_size, %values, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (i1, index, index, index, index, index, index, index, vector<4xf32>, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result10, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes10 { + %values = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + template.apply<@ggml.mul_mat_id_f32_f32_wmma.publish_vector4>(%writes10, %bounded_assignment0, %assignment_token0, %assignment_route0, %channel1, %bounded_token_count, %bounded_route_count, %bounded_output_size, %values, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (i1, index, index, index, index, index, index, index, vector<4xf32>, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result11, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes11 { + %values = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + template.apply<@ggml.mul_mat_id_f32_f32_wmma.publish_vector4>(%writes11, %bounded_assignment1, %assignment_token1, %assignment_route1, %channel1, %bounded_token_count, %bounded_route_count, %bounded_output_size, %values, %bias, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (i1, index, index, index, index, index, index, index, vector<4xf32>, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + } + kernel.barrier scope(workgroup) ordering(acq_rel) + template.apply<@ggml.mul_mat_id_f32_f32_wmma.finish_tile>(%bounded_token_count, %bounded_route_count, %bounded_output_size, %bounded_expert_count, %output_tile_count, %maximum_partition_count, %epsilon, %channel_tile, %descriptor_ordinal, %bounded_expert, %route_tile_base, %partition_row_count, %expert_table, %output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (index, index, index, index, index, index, f32, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.barrier scope(workgroup) ordering(acq_rel) + } + template.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_q4k_f16_wmma_projection.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_q4k_f16_wmma_projection.loom new file mode 100644 index 000000000000..c744d20f85a0 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_q4k_f16_wmma_projection.loom @@ -0,0 +1,200 @@ +// Shared routed Q4_K single-projection WMMA step. +// +// Callers own the outer routing and fusion schedule. This motif stages one +// projected Q4_K weight tile, optionally stages the shared routed activation +// tile, performs the 16x16 WMMA update, and returns the updated accumulator. + +func.def inline @ggml_q4k_scale_from_header(%scale0: i32, %scale1: i32, %scale2: i32, %q4_group: index) -> (i32, i32) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c48_i32 = scalar.constant 48 : i32 + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + %is_low_group = index.cmp ult, %bounded_group, %c4 : index + %scale_lane = index.rem %bounded_group, %c4 : index + %scale_shift_index = index.mul %scale_lane, %c8 : index + %scale_shift = index.cast %scale_shift_index : index to i32 + %high_shift = scalar.addi %scale_shift, %c2_i32 : i32 + %minimum_shift = scalar.addi %scale_shift, %c4_i32 : i32 + %selected_scale_source = scf.select %is_low_group, %scale0, %scale2 : i32 + %selected_minimum_source = scf.select %is_low_group, %scale1, %scale2 : i32 + %selected_scale_high_shift = scf.select %is_low_group, %scale_shift, %high_shift : i32 + %selected_minimum_low_shift = scf.select %is_low_group, %scale_shift, %minimum_shift : i32 + %scale_low0 = scalar.shrui %selected_scale_source, %scale_shift : i32 + %scale_low = scalar.andi %scale_low0, %c15_i32 : i32 + %scale_high0 = scalar.shrui %scale0, %selected_scale_high_shift : i32 + %scale_high = scalar.andi %scale_high0, %c48_i32 : i32 + %scale = scalar.ori %scale_low, %scale_high : i32 + %minimum_low0 = scalar.shrui %selected_minimum_source, %selected_minimum_low_shift : i32 + %minimum_low = scalar.andi %minimum_low0, %c15_i32 : i32 + %minimum_high0 = scalar.shrui %scale1, %selected_scale_high_shift : i32 + %minimum_high = scalar.andi %minimum_high0, %c48_i32 : i32 + %minimum = scalar.ori %minimum_low, %minimum_high : i32 + func.return %scale, %minimum : i32, i32 +} + +// Acquires the packed code word shared by one adjacent Q4_K group pair. +func.def inline @ggml_q4k_split_wmma_load_code(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group_pair: index, %packet: index) -> (vector<1xi32>) { + %c8 = index.constant 8 : index + %block_bytes = index.constant 144 : offset + %code_offset = index.constant 16 : offset + %bounded_group_pair = index.assume %q4_group_pair [range(%q4_group_pair, 0, 3)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q4_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<32xi32> + %q_page = index.mul %bounded_group_pair, %c8 : index + %q_word_index0 = index.add %q_page, %bounded_packet : index + %q_word_index = index.assume %q_word_index0 [range(%q_word_index0, 0, 31)] : index + %q_word = vector.load %code_view[%q_word_index] : view<32xi32> -> vector<1xi32> + func.return %q_word : vector<1xi32> +} + +// Decodes the four adjacent Q4_K values owned by one load packet from an +// already-loaded block header and packed code word. Matrix schedules choose +// the lifetime of both immutable packets. +func.def inline @ggml_q4k_split_wmma_vector4_from_header_code(%q4_group: index, %header_words: vector<4xi32>, %q_word: vector<1xi32>) -> (vector<4xf16>) { + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %q4_mask = vector.constant 252645135 : vector<1xi32> + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + %header_halves = vector.bitcast %header_words : vector<4xi32> to vector<8xf16> + %d_f16 = vector.extract %header_halves[0] : vector<8xf16> -> f16 + %dmin_f16 = vector.extract %header_halves[1] : vector<8xf16> -> f16 + %scale0 = vector.extract %header_words[1] : vector<4xi32> -> i32 + %scale1 = vector.extract %header_words[2] : vector<4xi32> -> i32 + %scale2 = vector.extract %header_words[3] : vector<4xi32> -> i32 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %scale, %minimum = func.call @ggml_q4k_scale_from_header(%scale0, %scale1, %scale2, %bounded_group) : (i32, i32, i32, index) -> (i32, i32) + %scale_f32 = scalar.uitofp %scale : i32 to f32 + %minimum_f32 = scalar.uitofp %minimum : i32 to f32 + %d_scale = scalar.mulf %d, %scale_f32 : f32 + %minimum_scale = scalar.mulf %dmin, %minimum_f32 : f32 + %q_half = index.rem %bounded_group, %c2 : index + %q_shift_index = index.mul %q_half, %c4 : index + %q_shift_i32 = index.cast %q_shift_index : index to i32 + %q_shift = vector.splat %q_shift_i32 : vector<1xi32> + %shifted_q = vector.shrui %q_word, %q_shift : vector<1xi32> + %masked_q = vector.andi %shifted_q, %q4_mask : vector<1xi32> + %q_i8 = vector.bitcast %masked_q : vector<1xi32> to vector<4xi8> + %q_f32 = vector.uitofp %q_i8 : vector<4xi8> to vector<4xf32> + // Form adjacent FP16 lanes from fused FP32 affine expressions. AMDGPU maps + // this natural shape to packed mixlo/mixhi instructions where available. + %negative_minimum_scale = scalar.negf %minimum_scale : f32 + %q0 = vector.extract %q_f32[0] : vector<4xf32> -> f32 + %q1 = vector.extract %q_f32[1] : vector<4xf32> -> f32 + %q2 = vector.extract %q_f32[2] : vector<4xf32> -> f32 + %q3 = vector.extract %q_f32[3] : vector<4xf32> -> f32 + %value0 = scalar.fmaf %q0, %d_scale, %negative_minimum_scale : f32 + %value1 = scalar.fmaf %q1, %d_scale, %negative_minimum_scale : f32 + %value2 = scalar.fmaf %q2, %d_scale, %negative_minimum_scale : f32 + %value3 = scalar.fmaf %q3, %d_scale, %negative_minimum_scale : f32 + %half0 = scalar.fptrunc %value0 : f32 to f16 + %half1 = scalar.fptrunc %value1 : f32 to f16 + %half2 = scalar.fptrunc %value2 : f32 to f16 + %half3 = scalar.fptrunc %value3 : f32 to f16 + %result = vector.from_elements %half0, %half1, %half2, %half3 : vector<4xf16> + func.return %result : vector<4xf16> +} + +// Decodes one group when its caller has retained only the Q4_K block header. +func.def inline @ggml_q4k_split_wmma_vector4_from_header(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group: index, %packet: index, %header_words: vector<4xi32>) -> (vector<4xf16>) { + %c2 = index.constant 2 : index + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + %q4_group_pair = index.div %bounded_group, %c2 : index + %q_word = func.call @ggml_q4k_split_wmma_load_code(%weight, %row_byte_base, %q4_block, %q4_group_pair, %packet) : (buffer, offset, index, index, index) -> (vector<1xi32>) + %values = func.call @ggml_q4k_split_wmma_vector4_from_header_code(%bounded_group, %header_words, %q_word) : (index, vector<4xi32>, vector<1xi32>) -> (vector<4xf16>) + func.return %values : vector<4xf16> +} + +// Acquires one naturally aligned Q4_K block header as a single 16-byte packet. +func.def inline @ggml_q4k_split_wmma_load_header(%weight: buffer, %row_byte_base: offset, %q4_block: index) -> (vector<4xi32>) { + %c0 = index.constant 0 : index + %block_bytes = index.constant 144 : offset + %block_byte_add = index.scale %q4_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %header_view = buffer.view %weight[%block_byte_base] : buffer -> view<4xi32> + %header_words = vector.load %header_view[%c0] : view<4xi32> -> vector<4xi32> + func.return %header_words : vector<4xi32> +} + +// Acquires one block header before decoding the selected four-value group +// packet. Matrix schedules that span several groups call the two operations +// separately so the header lifetime matches their complete block loop. +func.def inline @ggml_q4k_split_wmma_vector4(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group: index, %packet: index) -> (vector<4xf16>) { + %header_words = func.call @ggml_q4k_split_wmma_load_header(%weight, %row_byte_base, %q4_block) : (buffer, offset, index) -> (vector<4xi32>) + %values = func.call @ggml_q4k_split_wmma_vector4_from_header(%weight, %row_byte_base, %q4_block, %q4_group, %packet, %header_words) : (buffer, offset, index, index, index, vector<4xi32>) -> (vector<4xf16>) + func.return %values : vector<4xf16> +} + +template.decl @ggml.mul_mat_id_q4k_f16_wmma.project_vector16(%stage_activation: i1, %q4_group: index, %k_origin: index, %load_row: index, %load_k: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %header0: vector<4xi32>, %header1: vector<4xi32>, %q_word0: vector<1xi32>, %q_word1: vector<1xi32>, %accumulator: vector<16xf16>) -> (vector<16xf16>) + +template.def<@ggml.mul_mat_id_q4k_f16_wmma.project_vector16> device @ggml_mul_mat_id_q4k_f16_wmma_project_vector16(%stage_activation: i1, %q4_group: index, %k_origin: index, %load_row: index, %load_k: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %header0: vector<4xi32>, %header1: vector<4xi32>, %q_word0: vector<1xi32>, %q_word1: vector<1xi32>, %accumulator: vector<16xf16>) -> (vector<16xf16>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c40 = index.constant 40 : index + %c64 = index.constant 64 : index + %c0_offset = index.constant 0 : offset + %c0_i32 = scalar.constant 0 : i32 + %c0_f16x4 = vector.constant 0.0 : vector<4xf16> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %input_noalias, %route_stage_noalias, %weight_stage_noalias, %activation_stage_noalias = buffer.assume.noalias %input, %route_stage, %weight_stage, %activation_stage : buffer, buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%token_count]x[%input_size]xf32> + %route_stage_view = buffer.view %route_stage_noalias[%c0_offset] : buffer -> view<32xi32> + %weight_stage_view = buffer.view %weight_stage_noalias[%c0_offset] : buffer -> view<64x40xf16> + %activation_stage_physical_view = buffer.view %activation_stage_noalias[%c0_offset] : buffer -> view<32x40xf16> + %activation_fragment_layout = encoding.layout.strided [1, %c40] : encoding + %activation_fragment_view = buffer.view %activation_stage_noalias[%c0_offset] : buffer -> view<32x32xf16, %activation_fragment_layout> + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + scf.for %row_offset = [%c0 to %c64 step %c32] unroll { + %local_row0 = index.add %load_row, %row_offset : index + %local_row = index.assume %local_row0 [range(%local_row0, 0, 63)] : index + %selects_first_row = index.cmp eq, %row_offset, %c0 : index + %selected_header = scf.select %selects_first_row, %header0, %header1 : vector<4xi32> + %selected_q_word = scf.select %selects_first_row, %q_word0, %q_word1 : vector<1xi32> + %channel = index.add %channel_tile_base, %local_row : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %weight_values = scf.if %valid_channel -> (vector<4xf16>) { + %decoded = func.call @ggml_q4k_split_wmma_vector4_from_header_code(%bounded_group, %selected_header, %selected_q_word) : (index, vector<4xi32>, vector<1xi32>) -> (vector<4xf16>) + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %weight_values, %weight_stage_view[%local_row, %load_k] : vector<4xf16>, view<64x40xf16> + } + scf.if %stage_activation { + %assignment_i32 = view.load %route_stage_view[%load_row] : view<32xi32> -> i32 + %valid_assignment = scalar.cmpi sge, %assignment_i32, %c0_i32 : i32 + %activation_values = scf.if %valid_assignment -> (vector<4xf16>) { + %assignment0 = index.cast %assignment_i32 : i32 to index + %assignment = index.assume %assignment0 [range(%assignment0, 0, 16383)] : index + %bounded_assignment, %bounded_assignment_count = index.assume %assignment, %assignment_count [lt(%assignment, %assignment_count)] : index, index + %token0 = index.div %bounded_assignment, %route_count : index + %token, %input_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %input_k = index.add %k_origin, %load_k : index + %loaded = vector.load %input_view[%token, %input_k] : view<[%token_count]x[%input_size]xf32> -> vector<4xf32> + %converted = vector.fptrunc %loaded : vector<4xf32> to vector<4xf16> + scf.yield %converted : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %activation_values, %activation_stage_physical_view[%load_row, %load_k] : vector<4xf16>, view<32x40xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %next_accumulator = scf.for %k_half = [%c0 to %c32 step %c16](%half_accumulator = %accumulator : vector<16xf16>) -> (vector<16xf16>) unroll { + %lhs = vector.fragment.load %weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %rhs = vector.fragment.load %activation_fragment_view[%k_half, %subgroup_route_add] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %next = vector.mma %lhs, %rhs, %half_accumulator : vector<16xf16>, vector<16xf16>, vector<16xf16> + scf.yield %next : vector<16xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + template.return %next_accumulator : vector<16xf16> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_q5k_f16_wmma_projection.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_q5k_f16_wmma_projection.loom new file mode 100644 index 000000000000..34714bf1633a --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_q5k_f16_wmma_projection.loom @@ -0,0 +1,221 @@ +// Shared routed Q5_K single-projection WMMA step. +// +// Callers own the outer routing and fusion schedule. This motif stages one +// projected Q5_K weight tile, optionally stages the shared routed activation +// tile, performs the 16x16 WMMA update, and returns the updated accumulator. + +func.def inline @ggml_q5k_split_wmma_high_bit_f32(%qh: i8, %qh_bit: i32) -> (f32) { + %c0_i32 = scalar.constant 0 : i32 + %c16_i32 = scalar.constant 16 : i32 + %qh_i32 = scalar.extui %qh : i8 to i32 + %masked = scalar.andi %qh_i32, %qh_bit : i32 + %is_set = scalar.cmpi ne, %masked, %c0_i32 : i32 + %high = scf.select %is_set, %c16_i32, %c0_i32 : i32 + %high_f32 = scalar.uitofp %high : i32 to f32 + func.return %high_f32 : f32 +} + +func.def inline @ggml_q5k_scale_from_header(%scale0: i32, %scale1: i32, %scale2: i32, %q5_group: index) -> (i32, i32) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c48_i32 = scalar.constant 48 : i32 + %bounded_group = index.assume %q5_group [range(%q5_group, 0, 7)] : index + %is_low_group = index.cmp ult, %bounded_group, %c4 : index + %scale_lane = index.rem %bounded_group, %c4 : index + %scale_shift_index = index.mul %scale_lane, %c8 : index + %scale_shift = index.cast %scale_shift_index : index to i32 + %high_shift = scalar.addi %scale_shift, %c2_i32 : i32 + %minimum_shift = scalar.addi %scale_shift, %c4_i32 : i32 + %selected_scale_source = scf.select %is_low_group, %scale0, %scale2 : i32 + %selected_minimum_source = scf.select %is_low_group, %scale1, %scale2 : i32 + %selected_scale_high_shift = scf.select %is_low_group, %scale_shift, %high_shift : i32 + %selected_minimum_low_shift = scf.select %is_low_group, %scale_shift, %minimum_shift : i32 + %scale_low0 = scalar.shrui %selected_scale_source, %scale_shift : i32 + %scale_low = scalar.andi %scale_low0, %c15_i32 : i32 + %scale_high0 = scalar.shrui %scale0, %selected_scale_high_shift : i32 + %scale_high = scalar.andi %scale_high0, %c48_i32 : i32 + %scale = scalar.ori %scale_low, %scale_high : i32 + %minimum_low0 = scalar.shrui %selected_minimum_source, %selected_minimum_low_shift : i32 + %minimum_low = scalar.andi %minimum_low0, %c15_i32 : i32 + %minimum_high0 = scalar.shrui %scale1, %selected_scale_high_shift : i32 + %minimum_high = scalar.andi %minimum_high0, %c48_i32 : i32 + %minimum = scalar.ori %minimum_low, %minimum_high : i32 + func.return %scale, %minimum : i32, i32 +} + +// Acquires the packed low-bit code word shared by one adjacent Q5_K group pair. +func.def inline @ggml_q5k_split_wmma_load_code(%weight: buffer, %row_byte_base: offset, %q5_block: index, %q5_group_pair: index, %packet: index) -> (vector<1xi32>) { + %c8 = index.constant 8 : index + %block_bytes = index.constant 176 : offset + %code_offset = index.constant 48 : offset + %bounded_group_pair = index.assume %q5_group_pair [range(%q5_group_pair, 0, 3)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q5_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<32xi32> + %q_page = index.mul %bounded_group_pair, %c8 : index + %q_word_index0 = index.add %q_page, %bounded_packet : index + %q_word_index = index.assume %q_word_index0 [range(%q_word_index0, 0, 31)] : index + %q_word = vector.load %code_view[%q_word_index] : view<32xi32> -> vector<1xi32> + func.return %q_word : vector<1xi32> +} + +// Acquires the four Q5_K high-bit bytes owned by one packet. The same bytes are +// reused by all eight groups in the block. +func.def inline @ggml_q5k_split_wmma_load_high_bits(%weight: buffer, %row_byte_base: offset, %q5_block: index, %packet: index) -> (vector<4xi8>) { + %c4 = index.constant 4 : index + %block_bytes = index.constant 176 : offset + %qh_offset = index.constant 16 : offset + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q5_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %qh_byte_base = index.add %block_byte_base, %qh_offset : offset + %qh_view = buffer.view %weight[%qh_byte_base] : buffer -> view<32xi8> + %packet_base = index.mul %bounded_packet, %c4 : index + %qh_bytes = vector.load %qh_view[%packet_base] : view<32xi8> -> vector<4xi8> + func.return %qh_bytes : vector<4xi8> +} + +// Decodes the four adjacent Q5_K values owned by one load packet from retained +// block metadata and packed code. +func.def inline @ggml_q5k_split_wmma_vector4_from_header_code(%q5_group: index, %header_words: vector<4xi32>, %q_word: vector<1xi32>, %qh_bytes: vector<4xi8>) -> (vector<4xf16>) { + %c1_i32 = scalar.constant 1 : i32 + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %q4_mask = vector.constant 252645135 : vector<1xi32> + %bounded_group = index.assume %q5_group [range(%q5_group, 0, 7)] : index + %header_halves = vector.bitcast %header_words : vector<4xi32> to vector<8xf16> + %d_f16 = vector.extract %header_halves[0] : vector<8xf16> -> f16 + %dmin_f16 = vector.extract %header_halves[1] : vector<8xf16> -> f16 + %scale0 = vector.extract %header_words[1] : vector<4xi32> -> i32 + %scale1 = vector.extract %header_words[2] : vector<4xi32> -> i32 + %scale2 = vector.extract %header_words[3] : vector<4xi32> -> i32 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %scale, %minimum = func.call @ggml_q5k_scale_from_header(%scale0, %scale1, %scale2, %bounded_group) : (i32, i32, i32, index) -> (i32, i32) + %scale_f32 = scalar.uitofp %scale : i32 to f32 + %minimum_f32 = scalar.uitofp %minimum : i32 to f32 + %d_scale = scalar.mulf %d, %scale_f32 : f32 + %minimum_scale = scalar.mulf %dmin, %minimum_f32 : f32 + %q_half = index.rem %bounded_group, %c2 : index + %q_shift_index = index.mul %q_half, %c4 : index + %q_shift_i32 = index.cast %q_shift_index : index to i32 + %q_shift = vector.splat %q_shift_i32 : vector<1xi32> + %shifted_q = vector.shrui %q_word, %q_shift : vector<1xi32> + %masked_q = vector.andi %shifted_q, %q4_mask : vector<1xi32> + %q_i8 = vector.bitcast %masked_q : vector<1xi32> to vector<4xi8> + %q_f32 = vector.uitofp %q_i8 : vector<4xi8> to vector<4xf32> + %qh_shift_i32 = index.cast %bounded_group : index to i32 + %qh_bit = scalar.shli %c1_i32, %qh_shift_i32 : i32 + %qh0 = vector.extract %qh_bytes[0] : vector<4xi8> -> i8 + %qh1 = vector.extract %qh_bytes[1] : vector<4xi8> -> i8 + %qh2 = vector.extract %qh_bytes[2] : vector<4xi8> -> i8 + %qh3 = vector.extract %qh_bytes[3] : vector<4xi8> -> i8 + %high0 = func.call @ggml_q5k_split_wmma_high_bit_f32(%qh0, %qh_bit) : (i8, i32) -> (f32) + %high1 = func.call @ggml_q5k_split_wmma_high_bit_f32(%qh1, %qh_bit) : (i8, i32) -> (f32) + %high2 = func.call @ggml_q5k_split_wmma_high_bit_f32(%qh2, %qh_bit) : (i8, i32) -> (f32) + %high3 = func.call @ggml_q5k_split_wmma_high_bit_f32(%qh3, %qh_bit) : (i8, i32) -> (f32) + %q0_low = vector.extract %q_f32[0] : vector<4xf32> -> f32 + %q1_low = vector.extract %q_f32[1] : vector<4xf32> -> f32 + %q2_low = vector.extract %q_f32[2] : vector<4xf32> -> f32 + %q3_low = vector.extract %q_f32[3] : vector<4xf32> -> f32 + %q0 = scalar.addf %q0_low, %high0 : f32 + %q1 = scalar.addf %q1_low, %high1 : f32 + %q2 = scalar.addf %q2_low, %high2 : f32 + %q3 = scalar.addf %q3_low, %high3 : f32 + %negative_minimum_scale = scalar.negf %minimum_scale : f32 + %value0 = scalar.fmaf %q0, %d_scale, %negative_minimum_scale : f32 + %value1 = scalar.fmaf %q1, %d_scale, %negative_minimum_scale : f32 + %value2 = scalar.fmaf %q2, %d_scale, %negative_minimum_scale : f32 + %value3 = scalar.fmaf %q3, %d_scale, %negative_minimum_scale : f32 + %half0 = scalar.fptrunc %value0 : f32 to f16 + %half1 = scalar.fptrunc %value1 : f32 to f16 + %half2 = scalar.fptrunc %value2 : f32 to f16 + %half3 = scalar.fptrunc %value3 : f32 to f16 + %result = vector.from_elements %half0, %half1, %half2, %half3 : vector<4xf16> + func.return %result : vector<4xf16> +} + +// Acquires one naturally aligned Q5_K block header as a single 16-byte packet. +func.def inline @ggml_q5k_split_wmma_load_header(%weight: buffer, %row_byte_base: offset, %q5_block: index) -> (vector<4xi32>) { + %c0 = index.constant 0 : index + %block_bytes = index.constant 176 : offset + %block_byte_add = index.scale %q5_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %header_view = buffer.view %weight[%block_byte_base] : buffer -> view<4xi32> + %header_words = vector.load %header_view[%c0] : view<4xi32> -> vector<4xi32> + func.return %header_words : vector<4xi32> +} + +template.decl @ggml.mul_mat_id_q5k_f16_wmma.project_vector16(%stage_activation: i1, %q5_group: index, %k_origin: index, %load_row: index, %load_k: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %header0: vector<4xi32>, %header1: vector<4xi32>, %q_word0: vector<1xi32>, %q_word1: vector<1xi32>, %qh_bytes0: vector<4xi8>, %qh_bytes1: vector<4xi8>, %accumulator: vector<16xf16>) -> (vector<16xf16>) + +template.def<@ggml.mul_mat_id_q5k_f16_wmma.project_vector16> device @ggml_mul_mat_id_q5k_f16_wmma_project_vector16(%stage_activation: i1, %q5_group: index, %k_origin: index, %load_row: index, %load_k: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %header0: vector<4xi32>, %header1: vector<4xi32>, %q_word0: vector<1xi32>, %q_word1: vector<1xi32>, %qh_bytes0: vector<4xi8>, %qh_bytes1: vector<4xi8>, %accumulator: vector<16xf16>) -> (vector<16xf16>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c40 = index.constant 40 : index + %c64 = index.constant 64 : index + %c0_offset = index.constant 0 : offset + %c0_i32 = scalar.constant 0 : i32 + %c0_f16x4 = vector.constant 0.0 : vector<4xf16> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %input_noalias, %route_stage_noalias, %weight_stage_noalias, %activation_stage_noalias = buffer.assume.noalias %input, %route_stage, %weight_stage, %activation_stage : buffer, buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%token_count]x[%input_size]xf32> + %route_stage_view = buffer.view %route_stage_noalias[%c0_offset] : buffer -> view<32xi32> + %weight_stage_view = buffer.view %weight_stage_noalias[%c0_offset] : buffer -> view<64x40xf16> + %activation_stage_physical_view = buffer.view %activation_stage_noalias[%c0_offset] : buffer -> view<32x40xf16> + %activation_fragment_layout = encoding.layout.strided [1, %c40] : encoding + %activation_fragment_view = buffer.view %activation_stage_noalias[%c0_offset] : buffer -> view<32x32xf16, %activation_fragment_layout> + %bounded_group = index.assume %q5_group [range(%q5_group, 0, 7)] : index + scf.for %row_offset = [%c0 to %c64 step %c32] unroll { + %local_row0 = index.add %load_row, %row_offset : index + %local_row = index.assume %local_row0 [range(%local_row0, 0, 63)] : index + %selects_first_row = index.cmp eq, %row_offset, %c0 : index + %selected_header = scf.select %selects_first_row, %header0, %header1 : vector<4xi32> + %selected_q_word = scf.select %selects_first_row, %q_word0, %q_word1 : vector<1xi32> + %selected_qh_bytes = scf.select %selects_first_row, %qh_bytes0, %qh_bytes1 : vector<4xi8> + %channel = index.add %channel_tile_base, %local_row : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %weight_values = scf.if %valid_channel -> (vector<4xf16>) { + %decoded = func.call @ggml_q5k_split_wmma_vector4_from_header_code(%bounded_group, %selected_header, %selected_q_word, %selected_qh_bytes) : (index, vector<4xi32>, vector<1xi32>, vector<4xi8>) -> (vector<4xf16>) + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %weight_values, %weight_stage_view[%local_row, %load_k] : vector<4xf16>, view<64x40xf16> + } + scf.if %stage_activation { + %assignment_i32 = view.load %route_stage_view[%load_row] : view<32xi32> -> i32 + %valid_assignment = scalar.cmpi sge, %assignment_i32, %c0_i32 : i32 + %activation_values = scf.if %valid_assignment -> (vector<4xf16>) { + %assignment0 = index.cast %assignment_i32 : i32 to index + %assignment = index.assume %assignment0 [range(%assignment0, 0, 16383)] : index + %bounded_assignment, %bounded_assignment_count = index.assume %assignment, %assignment_count [lt(%assignment, %assignment_count)] : index, index + %token0 = index.div %bounded_assignment, %route_count : index + %token, %input_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %input_k = index.add %k_origin, %load_k : index + %loaded = vector.load %input_view[%token, %input_k] : view<[%token_count]x[%input_size]xf32> -> vector<4xf32> + %converted = vector.fptrunc %loaded : vector<4xf32> to vector<4xf16> + scf.yield %converted : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %activation_values, %activation_stage_physical_view[%load_row, %load_k] : vector<4xf16>, view<32x40xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %next_accumulator = scf.for %k_half = [%c0 to %c32 step %c16](%half_accumulator = %accumulator : vector<16xf16>) -> (vector<16xf16>) unroll { + %lhs = vector.fragment.load %weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %rhs = vector.fragment.load %activation_fragment_view[%k_half, %subgroup_route_add] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %next = vector.mma %lhs, %rhs, %half_accumulator : vector<16xf16>, vector<16xf16>, vector<16xf16> + scf.yield %next : vector<16xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + template.return %next_accumulator : vector<16xf16> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_q6k_f16_wmma_projection.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_q6k_f16_wmma_projection.loom new file mode 100644 index 000000000000..284577982a6d --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_q6k_f16_wmma_projection.loom @@ -0,0 +1,169 @@ +// Shared routed Q6_K single-projection WMMA step. +// +// Callers own the outer routing and fusion schedule. This motif stages one +// projected Q6_K weight tile, optionally stages the shared routed activation +// tile, performs the 16x16 WMMA update, and returns the updated accumulator. + +// Acquires the packed words shared by one four-group half of a Q6_K block. +// Each QL word supplies two groups and the QH word supplies all four. +func.def inline @ggml_q6k_split_wmma_load_half_codes(%weight: buffer, %weight_row_byte_base: offset, %q6_block: index, %q6_half: index, %packet: index) -> (vector<1xi32>, vector<1xi32>, vector<1xi32>) { + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %block_bytes = index.constant 210 : offset + %qh_byte_add = index.constant 128 : offset + %bounded_half = index.assume %q6_half [range(%q6_half, 0, 1)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q6_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %weight_row_byte_base, %block_byte_add : offset + %qh_byte_base = index.add %block_byte_base, %qh_byte_add : offset + %ql_view = buffer.view %weight[%block_byte_base] : buffer -> view<32xi32> + %qh_view = buffer.view %weight[%qh_byte_base] : buffer -> view<16xi32> + %ql_half_word_base = index.mul %bounded_half, %c16 : index + %ql0_word_index = index.add %ql_half_word_base, %bounded_packet : index + %ql1_word_base = index.add %ql_half_word_base, %c8 : index + %ql1_word_index = index.add %ql1_word_base, %bounded_packet : index + %qh_half_word_base = index.mul %bounded_half, %c8 : index + %qh_word_index = index.add %qh_half_word_base, %bounded_packet : index + %ql0_word = vector.load %ql_view[%ql0_word_index] : view<32xi32> -> vector<1xi32> + %ql1_word = vector.load %ql_view[%ql1_word_index] : view<32xi32> -> vector<1xi32> + %qh_word = vector.load %qh_view[%qh_word_index] : view<16xi32> -> vector<1xi32> + func.return %ql0_word, %ql1_word, %qh_word : vector<1xi32>, vector<1xi32>, vector<1xi32> +} + +// Decodes four adjacent values after the surrounding schedule has selected +// the scale and retained the packed code words at their natural lifetimes. +func.def inline @ggml_q6k_split_wmma_vector4_from_scale_codes(%q6_group: index, %scale_i8: i8, %d_f16: f16, %ql_word: vector<1xi32>, %qh_word: vector<1xi32>) -> (vector<4xf16>) { + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c4_i32v = vector.constant 4 : vector<1xi32> + %nibble_mask = vector.constant 252645135 : vector<1xi32> + %high_mask = vector.constant 50529027 : vector<1xi32> + %c32_f32v = vector.constant 32.0 : vector<4xf32> + %bounded_group = index.assume %q6_group [range(%q6_group, 0, 7)] : index + %group_in_half = index.rem %bounded_group, %c4 : index + %nibble = index.div %group_in_half, %c2 : index + %nibble_shift_index = index.mul %nibble, %c4 : index + %nibble_shift_i32 = index.cast %nibble_shift_index : index to i32 + %nibble_shift = vector.splat %nibble_shift_i32 : vector<1xi32> + %qh_shift_index = index.mul %group_in_half, %c2 : index + %qh_shift_i32 = index.cast %qh_shift_index : index to i32 + %qh_shift = vector.splat %qh_shift_i32 : vector<1xi32> + %ql_shifted = vector.shrui %ql_word, %nibble_shift : vector<1xi32> + %ql = vector.andi %ql_shifted, %nibble_mask : vector<1xi32> + %qh_shifted = vector.shrui %qh_word, %qh_shift : vector<1xi32> + %qh_low = vector.andi %qh_shifted, %high_mask : vector<1xi32> + %qh = vector.shli %qh_low, %c4_i32v : vector<1xi32> + %code = vector.ori %ql, %qh : vector<1xi32> + %code_i8 = vector.bitcast %code : vector<1xi32> to vector<4xi8> + %code_f32 = vector.uitofp %code_i8 : vector<4xi8> to vector<4xf32> + %centered = vector.subf %code_f32, %c32_f32v : vector<4xf32> + %scale = scalar.sitofp %scale_i8 : i8 to f32 + %d = scalar.extf %d_f16 : f16 to f32 + %combined_scale = scalar.mulf %scale, %d : f32 + %combined_scale_vector = vector.splat %combined_scale : vector<4xf32> + %values_f32 = vector.mulf %centered, %combined_scale_vector : vector<4xf32> + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +// Loads only the group scale and block multiplier while reusing packed words +// retained by an enclosing four-group half-block schedule. +func.def inline @ggml_q6k_split_wmma_vector4_from_half_codes(%weight: buffer, %weight_row_byte_base: offset, %q6_block: index, %q6_group: index, %packet: index, %ql0_word: vector<1xi32>, %ql1_word: vector<1xi32>, %qh_word: vector<1xi32>) -> (vector<4xf16>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %block_bytes = index.constant 210 : offset + %scale_byte_add = index.constant 192 : offset + %d_byte_add = index.constant 208 : offset + %bounded_group = index.assume %q6_group [range(%q6_group, 0, 7)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q6_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %weight_row_byte_base, %block_byte_add : offset + %scale_byte_base = index.add %block_byte_base, %scale_byte_add : offset + %d_byte_base = index.add %block_byte_base, %d_byte_add : offset + %scale_view = buffer.view %weight[%scale_byte_base] : buffer -> view<16xi8> + %d_view = buffer.view %weight[%d_byte_base] : buffer -> view<1xf16> + %group_in_half = index.rem %bounded_group, %c4 : index + %ql_side = index.rem %group_in_half, %c2 : index + %uses_ql1 = index.cmp eq, %ql_side, %c1 : index + %ql_word = scf.select %uses_ql1, %ql1_word, %ql0_word : vector<1xi32> + %scale_packet_half = index.div %bounded_packet, %c4 : index + %scale_group_base = index.mul %bounded_group, %c2 : index + %scale_index = index.add %scale_group_base, %scale_packet_half : index + %scale_i8 = view.load %scale_view[%scale_index] : view<16xi8> -> i8 + %d_f16 = view.load %d_view[%c0] : view<1xf16> -> f16 + %values = func.call @ggml_q6k_split_wmma_vector4_from_scale_codes(%bounded_group, %scale_i8, %d_f16, %ql_word, %qh_word) : (index, i8, f16, vector<1xi32>, vector<1xi32>) -> (vector<4xf16>) + func.return %values : vector<4xf16> +} + +template.decl @ggml.mul_mat_id_q6k_f16_wmma.project_vector16(%stage_activation: i1, %q6_group: index, %q6_block: index, %k_origin: index, %load_row: index, %load_k: index, %load_packet: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %weight: buffer, %expert_byte_base: offset, %weight_row_bytes: offset, %ql0_word0: vector<1xi32>, %ql1_word0: vector<1xi32>, %qh_word0: vector<1xi32>, %ql0_word1: vector<1xi32>, %ql1_word1: vector<1xi32>, %qh_word1: vector<1xi32>, %accumulator: vector<16xf16>) -> (vector<16xf16>) + +template.def<@ggml.mul_mat_id_q6k_f16_wmma.project_vector16> device @ggml_mul_mat_id_q6k_f16_wmma_project_vector16(%stage_activation: i1, %q6_group: index, %q6_block: index, %k_origin: index, %load_row: index, %load_k: index, %load_packet: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %weight: buffer, %expert_byte_base: offset, %weight_row_bytes: offset, %ql0_word0: vector<1xi32>, %ql1_word0: vector<1xi32>, %qh_word0: vector<1xi32>, %ql0_word1: vector<1xi32>, %ql1_word1: vector<1xi32>, %qh_word1: vector<1xi32>, %accumulator: vector<16xf16>) -> (vector<16xf16>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c40 = index.constant 40 : index + %c64 = index.constant 64 : index + %c0_offset = index.constant 0 : offset + %c0_i32 = scalar.constant 0 : i32 + %c0_f16x4 = vector.constant 0.0 : vector<4xf16> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %input_noalias, %route_stage_noalias, %weight_stage_noalias, %activation_stage_noalias, %weight_noalias = buffer.assume.noalias %input, %route_stage, %weight_stage, %activation_stage, %weight : buffer, buffer, buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%token_count]x[%input_size]xf32> + %route_stage_view = buffer.view %route_stage_noalias[%c0_offset] : buffer -> view<32xi32> + %weight_stage_view = buffer.view %weight_stage_noalias[%c0_offset] : buffer -> view<64x40xf16> + %activation_stage_physical_view = buffer.view %activation_stage_noalias[%c0_offset] : buffer -> view<32x40xf16> + %activation_fragment_layout = encoding.layout.strided [1, %c40] : encoding + %activation_fragment_view = buffer.view %activation_stage_noalias[%c0_offset] : buffer -> view<32x32xf16, %activation_fragment_layout> + %bounded_group = index.assume %q6_group [range(%q6_group, 0, 7)] : index + scf.for %row_offset = [%c0 to %c64 step %c32] unroll { + %local_row0 = index.add %load_row, %row_offset : index + %local_row = index.assume %local_row0 [range(%local_row0, 0, 63)] : index + %selects_first_row = index.cmp eq, %row_offset, %c0 : index + %selected_ql0_word = scf.select %selects_first_row, %ql0_word0, %ql0_word1 : vector<1xi32> + %selected_ql1_word = scf.select %selects_first_row, %ql1_word0, %ql1_word1 : vector<1xi32> + %selected_qh_word = scf.select %selects_first_row, %qh_word0, %qh_word1 : vector<1xi32> + %channel = index.add %channel_tile_base, %local_row : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %weight_values = scf.if %valid_channel -> (vector<4xf16>) { + %channel_byte_add = index.scale %channel, %weight_row_bytes : index, offset -> offset + %row_byte_base = index.add %expert_byte_base, %channel_byte_add : offset + %decoded = func.call @ggml_q6k_split_wmma_vector4_from_half_codes(%weight_noalias, %row_byte_base, %q6_block, %bounded_group, %load_packet, %selected_ql0_word, %selected_ql1_word, %selected_qh_word) : (buffer, offset, index, index, index, vector<1xi32>, vector<1xi32>, vector<1xi32>) -> (vector<4xf16>) + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %weight_values, %weight_stage_view[%local_row, %load_k] : vector<4xf16>, view<64x40xf16> + } + scf.if %stage_activation { + %assignment_i32 = view.load %route_stage_view[%load_row] : view<32xi32> -> i32 + %valid_assignment = scalar.cmpi sge, %assignment_i32, %c0_i32 : i32 + %activation_values = scf.if %valid_assignment -> (vector<4xf16>) { + %assignment0 = index.cast %assignment_i32 : i32 to index + %assignment = index.assume %assignment0 [range(%assignment0, 0, 16383)] : index + %bounded_assignment, %bounded_assignment_count = index.assume %assignment, %assignment_count [lt(%assignment, %assignment_count)] : index, index + %token0 = index.div %bounded_assignment, %route_count : index + %token, %input_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %input_k = index.add %k_origin, %load_k : index + %loaded = vector.load %input_view[%token, %input_k] : view<[%token_count]x[%input_size]xf32> -> vector<4xf32> + %converted = vector.fptrunc %loaded : vector<4xf32> to vector<4xf16> + scf.yield %converted : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %activation_values, %activation_stage_physical_view[%load_row, %load_k] : vector<4xf16>, view<32x40xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %next_accumulator = scf.for %k_half = [%c0 to %c32 step %c16](%half_accumulator = %accumulator : vector<16xf16>) -> (vector<16xf16>) unroll { + %lhs = vector.fragment.load %weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %rhs = vector.fragment.load %activation_fragment_view[%k_half, %subgroup_route_add] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %next = vector.mma %lhs, %rhs, %half_accumulator : vector<16xf16>, vector<16xf16>, vector<16xf16> + scf.yield %next : vector<16xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + template.return %next_accumulator : vector<16xf16> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_swiglu_f16_f16_accumulate.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_swiglu_f16_f16_accumulate.loom new file mode 100644 index 000000000000..874bd565ae5a --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_swiglu_f16_f16_accumulate.loom @@ -0,0 +1,395 @@ +// Compile-time selected routed SwiGLU accumulator. + +template.decl @ggml.mul_mat_id_swiglu_f16_f16.accumulate_block(%quant_block: index, %load_row: index, %load_k: index, %load_packet: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %gate_weight: buffer, %up_weight: buffer, %expert_byte_base: offset, %gate_weight_row_bytes: offset, %up_weight_row_bytes: offset, %gate_weight_format: index, %up_weight_format: index, %gate_accumulator: vector<16xf16>, %up_accumulator: vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) + +template.decl @ggml.mul_mat_id_swiglu_f16_f16.accumulate_q4k_block(%quant_block: index, %load_row: index, %load_k: index, %load_packet: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %gate_weight: buffer, %up_weight: buffer, %expert_byte_base: offset, %weight_row_bytes: offset, %gate_accumulator: vector<16xf16>, %up_accumulator: vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) + +template.decl @ggml.mul_mat_id_swiglu_f16_f16.accumulate_q5k_block(%quant_block: index, %load_row: index, %load_k: index, %load_packet: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %gate_weight: buffer, %up_weight: buffer, %expert_byte_base: offset, %weight_row_bytes: offset, %gate_accumulator: vector<16xf16>, %up_accumulator: vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) + +template.decl @ggml.mul_mat_id_swiglu_f16_f16.accumulate_q6k_block(%quant_block: index, %load_row: index, %load_k: index, %load_packet: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %gate_weight: buffer, %up_weight: buffer, %expert_byte_base: offset, %weight_row_bytes: offset, %gate_accumulator: vector<16xf16>, %up_accumulator: vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) + +template.decl @ggml.mul_mat_id_swiglu_f16_f16.accumulate_generic_block(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %quant_block: index, %load_row: index, %load_k: index, %load_packet: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %gate_weight: buffer, %up_weight: buffer, %expert_byte_base: offset, %gate_weight_row_bytes: offset, %up_weight_row_bytes: offset, %gate_weight_format: index, %up_weight_format: index, %gate_accumulator: vector<16xf16>, %up_accumulator: vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) + +func.decl @ggml_iq4nl_table_i8() -> (vector<16xi8>) + +func.decl @ggml_iq3s_grid_f32_halves() -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + +template.def<@ggml.mul_mat_id_swiglu_f16_f16.accumulate_block> device @ggml_mul_mat_id_swiglu_f16_f16_accumulate_block(%quant_block: index, %load_row: index, %load_k: index, %load_packet: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %gate_weight: buffer, %up_weight: buffer, %expert_byte_base: offset, %gate_weight_row_bytes: offset, %up_weight_row_bytes: offset, %gate_weight_format: index, %up_weight_format: index, %gate_accumulator: vector<16xf16>, %up_accumulator: vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) { + %c4 = index.constant 4 : index + %c5 = index.constant 5 : index + %c6 = index.constant 6 : index + %same_weight_format = index.cmp eq, %gate_weight_format, %up_weight_format : index + %gate_is_q4k = index.cmp eq, %gate_weight_format, %c4 : index + %gate_is_q5k = index.cmp eq, %gate_weight_format, %c5 : index + %gate_is_q6k = index.cmp eq, %gate_weight_format, %c6 : index + %is_q4k = scalar.andi %same_weight_format, %gate_is_q4k : i1 + %is_q5k = scalar.andi %same_weight_format, %gate_is_q5k : i1 + %is_q6k = scalar.andi %same_weight_format, %gate_is_q6k : i1 + %q4_gate, %q4_up = scf.if %is_q4k -> (vector<16xf16>, vector<16xf16>) { + %q4_gate, %q4_up = template.apply<@ggml.mul_mat_id_swiglu_f16_f16.accumulate_q4k_block>(%quant_block, %load_row, %load_k, %load_packet, %channel_tile_base, %bounded_output_size, %assignment_count, %token_count, %route_count, %input_size, %subgroup_channel_add, %subgroup_route_add, %input, %route_stage, %weight_stage, %activation_stage, %gate_weight, %up_weight, %expert_byte_base, %gate_weight_row_bytes, %gate_accumulator, %up_accumulator) : (index, index, index, index, index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, offset, offset, vector<16xf16>, vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) + scf.yield %q4_gate, %q4_up : vector<16xf16>, vector<16xf16> + } else { + scf.yield %gate_accumulator, %up_accumulator : vector<16xf16>, vector<16xf16> + } + %q5_gate, %q5_up = scf.if %is_q5k -> (vector<16xf16>, vector<16xf16>) { + %q5_gate, %q5_up = template.apply<@ggml.mul_mat_id_swiglu_f16_f16.accumulate_q5k_block>(%quant_block, %load_row, %load_k, %load_packet, %channel_tile_base, %bounded_output_size, %assignment_count, %token_count, %route_count, %input_size, %subgroup_channel_add, %subgroup_route_add, %input, %route_stage, %weight_stage, %activation_stage, %gate_weight, %up_weight, %expert_byte_base, %gate_weight_row_bytes, %gate_accumulator, %up_accumulator) : (index, index, index, index, index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, offset, offset, vector<16xf16>, vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) + scf.yield %q5_gate, %q5_up : vector<16xf16>, vector<16xf16> + } else { + scf.yield %gate_accumulator, %up_accumulator : vector<16xf16>, vector<16xf16> + } + %q6_gate, %q6_up = scf.if %is_q6k -> (vector<16xf16>, vector<16xf16>) { + %q6_gate, %q6_up = template.apply<@ggml.mul_mat_id_swiglu_f16_f16.accumulate_q6k_block>(%quant_block, %load_row, %load_k, %load_packet, %channel_tile_base, %bounded_output_size, %assignment_count, %token_count, %route_count, %input_size, %subgroup_channel_add, %subgroup_route_add, %input, %route_stage, %weight_stage, %activation_stage, %gate_weight, %up_weight, %expert_byte_base, %gate_weight_row_bytes, %gate_accumulator, %up_accumulator) : (index, index, index, index, index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, offset, offset, vector<16xf16>, vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) + scf.yield %q6_gate, %q6_up : vector<16xf16>, vector<16xf16> + } else { + scf.yield %gate_accumulator, %up_accumulator : vector<16xf16>, vector<16xf16> + } + %is_q4_or_q5 = scalar.ori %is_q4k, %is_q5k : i1 + %uses_fast_path = scalar.ori %is_q4_or_q5, %is_q6k : i1 + %generic_gate, %generic_up = scf.if %uses_fast_path -> (vector<16xf16>, vector<16xf16>) { + scf.yield %gate_accumulator, %up_accumulator : vector<16xf16>, vector<16xf16> + } else { + %iq4nl_table = func.call @ggml_iq4nl_table_i8() : () -> (vector<16xi8>) + %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15 = func.call @ggml_iq3s_grid_f32_halves() : () -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + %generic_gate, %generic_up = template.apply<@ggml.mul_mat_id_swiglu_f16_f16.accumulate_generic_block>(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %quant_block, %load_row, %load_k, %load_packet, %channel_tile_base, %bounded_output_size, %assignment_count, %token_count, %route_count, %input_size, %subgroup_channel_add, %subgroup_route_add, %input, %route_stage, %weight_stage, %activation_stage, %gate_weight, %up_weight, %expert_byte_base, %gate_weight_row_bytes, %up_weight_row_bytes, %gate_weight_format, %up_weight_format, %gate_accumulator, %up_accumulator) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, index, index, index, index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, offset, offset, offset, index, index, vector<16xf16>, vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) + scf.yield %generic_gate, %generic_up : vector<16xf16>, vector<16xf16> + } + %q4_or_generic_gate = scf.select %is_q4k, %q4_gate, %generic_gate : vector<16xf16> + %q4_or_generic_up = scf.select %is_q4k, %q4_up, %generic_up : vector<16xf16> + %q5_or_prior_gate = scf.select %is_q5k, %q5_gate, %q4_or_generic_gate : vector<16xf16> + %q5_or_prior_up = scf.select %is_q5k, %q5_up, %q4_or_generic_up : vector<16xf16> + %next_gate = scf.select %is_q6k, %q6_gate, %q5_or_prior_gate : vector<16xf16> + %next_up = scf.select %is_q6k, %q6_up, %q5_or_prior_up : vector<16xf16> + template.return %next_gate, %next_up : vector<16xf16>, vector<16xf16> +} +// Generic routed gate/up SwiGLU accumulation for FP16 WMMA. +// +// This path covers formats whose values can be staged directly through the +// shared dequant helper. Q4_K and Q6_K keep specialized retained-packet paths. + +template.decl @ggml.mul_mat_id_swiglu_f16_f16.project_generic_vector16(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %stage_activation: i1, %weight_format: index, %quant_block: index, %quant_group: index, %k_origin: index, %load_row: index, %load_k: index, %load_packet: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %weight: buffer, %expert_byte_base: offset, %weight_row_bytes: offset, %accumulator: vector<16xf16>) -> (vector<16xf16>) + +func.decl @ggml_dequant_f16_vector4(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight_format: index, %weight: buffer, %row_byte_base: offset, %input_size: index, %quant_block: index, %quant_group: index, %packet: index, %k: index) -> (vector<4xf16>) + +template.def<@ggml.mul_mat_id_swiglu_f16_f16.project_generic_vector16> device @ggml_mul_mat_id_swiglu_f16_f16_project_generic_vector16(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %stage_activation: i1, %weight_format: index, %quant_block: index, %quant_group: index, %k_origin: index, %load_row: index, %load_k: index, %load_packet: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %weight: buffer, %expert_byte_base: offset, %weight_row_bytes: offset, %accumulator: vector<16xf16>) -> (vector<16xf16>) { + %c0 = index.constant 0 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c40 = index.constant 40 : index + %c64 = index.constant 64 : index + %c0_offset = index.constant 0 : offset + %c0_i32 = scalar.constant 0 : i32 + %c0_f16x4 = vector.constant 0.0 : vector<4xf16> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %input_noalias, %route_stage_noalias, %weight_stage_noalias, %activation_stage_noalias, %weight_noalias = buffer.assume.noalias %input, %route_stage, %weight_stage, %activation_stage, %weight : buffer, buffer, buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%token_count]x[%input_size]xf32> + %route_stage_view = buffer.view %route_stage_noalias[%c0_offset] : buffer -> view<32xi32> + %weight_stage_view = buffer.view %weight_stage_noalias[%c0_offset] : buffer -> view<64x40xf16> + %activation_stage_physical_view = buffer.view %activation_stage_noalias[%c0_offset] : buffer -> view<32x40xf16> + %activation_fragment_layout = encoding.layout.strided [1, %c40] : encoding + %activation_fragment_view = buffer.view %activation_stage_noalias[%c0_offset] : buffer -> view<32x32xf16, %activation_fragment_layout> + %bounded_group = index.assume %quant_group [range(%quant_group, 0, 7)] : index + scf.for %row_offset = [%c0 to %c64 step %c32] unroll { + %local_row0 = index.add %load_row, %row_offset : index + %local_row = index.assume %local_row0 [range(%local_row0, 0, 63)] : index + %channel = index.add %channel_tile_base, %local_row : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %weight_values = scf.if %valid_channel -> (vector<4xf16>) { + %channel_byte_add = index.scale %channel, %weight_row_bytes : index, offset -> offset + %row_byte_base = index.add %expert_byte_base, %channel_byte_add : offset + %weight_k = index.add %k_origin, %load_k : index + %decoded = func.call @ggml_dequant_f16_vector4(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight_format, %weight_noalias, %row_byte_base, %input_size, %quant_block, %bounded_group, %load_packet, %weight_k) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, buffer, offset, index, index, index, index, index) -> (vector<4xf16>) + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %weight_values, %weight_stage_view[%local_row, %load_k] : vector<4xf16>, view<64x40xf16> + } + scf.if %stage_activation { + %assignment_i32 = view.load %route_stage_view[%load_row] : view<32xi32> -> i32 + %valid_assignment = scalar.cmpi sge, %assignment_i32, %c0_i32 : i32 + %activation_values = scf.if %valid_assignment -> (vector<4xf16>) { + %assignment0 = index.cast %assignment_i32 : i32 to index + %assignment = index.assume %assignment0 [range(%assignment0, 0, 16383)] : index + %bounded_assignment, %bounded_assignment_count = index.assume %assignment, %assignment_count [lt(%assignment, %assignment_count)] : index, index + %token0 = index.div %bounded_assignment, %route_count : index + %token, %input_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %input_k = index.add %k_origin, %load_k : index + %loaded = vector.load %input_view[%token, %input_k] : view<[%token_count]x[%input_size]xf32> -> vector<4xf32> + %converted = vector.fptrunc %loaded : vector<4xf32> to vector<4xf16> + scf.yield %converted : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %activation_values, %activation_stage_physical_view[%load_row, %load_k] : vector<4xf16>, view<32x40xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %next_accumulator = scf.for %k_half = [%c0 to %c32 step %c16](%half_accumulator = %accumulator : vector<16xf16>) -> (vector<16xf16>) unroll { + %lhs = vector.fragment.load %weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %rhs = vector.fragment.load %activation_fragment_view[%k_half, %subgroup_route_add] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %next = vector.mma %lhs, %rhs, %half_accumulator : vector<16xf16>, vector<16xf16>, vector<16xf16> + scf.yield %next : vector<16xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + template.return %next_accumulator : vector<16xf16> +} + +template.def<@ggml.mul_mat_id_swiglu_f16_f16.accumulate_generic_block> device @ggml_mul_mat_id_swiglu_f16_f16_accumulate_generic_block(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %quant_block: index, %load_row: index, %load_k: index, %load_packet: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %gate_weight: buffer, %up_weight: buffer, %expert_byte_base: offset, %gate_weight_row_bytes: offset, %up_weight_row_bytes: offset, %gate_weight_format: index, %up_weight_format: index, %gate_accumulator: vector<16xf16>, %up_accumulator: vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c8 = index.constant 8 : index + %c32 = index.constant 32 : index + %c256 = index.constant 256 : index + %true = scalar.constant true : i1 + %false = scalar.constant false : i1 + %next_block_gate, %next_block_up = scf.for %quant_group = [%c0 to %c8 step %c1](%group_gate = %gate_accumulator : vector<16xf16>, %group_up = %up_accumulator : vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) { + %bounded_group = index.assume %quant_group [range(%quant_group, 0, 7)] : index + %block_k_base = index.mul %quant_block, %c256 : index + %group_k_add = index.mul %bounded_group, %c32 : index + %k_origin = index.add %block_k_base, %group_k_add : index + %gate_next = template.apply<@ggml.mul_mat_id_swiglu_f16_f16.project_generic_vector16>(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %true, %gate_weight_format, %quant_block, %bounded_group, %k_origin, %load_row, %load_k, %load_packet, %channel_tile_base, %bounded_output_size, %assignment_count, %token_count, %route_count, %input_size, %subgroup_channel_add, %subgroup_route_add, %input, %route_stage, %weight_stage, %activation_stage, %gate_weight, %expert_byte_base, %gate_weight_row_bytes, %group_gate) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, i1, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer, offset, offset, vector<16xf16>) -> (vector<16xf16>) + %up_next = template.apply<@ggml.mul_mat_id_swiglu_f16_f16.project_generic_vector16>(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %false, %up_weight_format, %quant_block, %bounded_group, %k_origin, %load_row, %load_k, %load_packet, %channel_tile_base, %bounded_output_size, %assignment_count, %token_count, %route_count, %input_size, %subgroup_channel_add, %subgroup_route_add, %input, %route_stage, %weight_stage, %activation_stage, %up_weight, %expert_byte_base, %up_weight_row_bytes, %group_up) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, i1, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer, offset, offset, vector<16xf16>) -> (vector<16xf16>) + scf.yield %gate_next, %up_next : vector<16xf16>, vector<16xf16> + } + template.return %next_block_gate, %next_block_up : vector<16xf16>, vector<16xf16> +} +// Q4_K routed gate/up SwiGLU accumulation. +// +// The shared body owns routing and publish. This file only binds Q4_K block +// accumulation to that common schedule. + +template.decl @ggml.mul_mat_id_q4k_f16_wmma.project_vector16(%arg0: i1, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: index, %arg7: index, %arg8: index, %arg9: index, %arg10: index, %arg11: index, %arg12: index, %arg13: buffer, %arg14: buffer, %arg15: buffer, %arg16: buffer, %arg17: vector<4xi32>, %arg18: vector<4xi32>, %arg19: vector<1xi32>, %arg20: vector<1xi32>, %arg21: vector<16xf16>) -> (vector<16xf16>) + +func.decl @ggml_q4k_split_wmma_load_header(%weight: buffer, %row_byte_base: offset, %q4_block: index) -> (vector<4xi32>) + +func.decl @ggml_q4k_split_wmma_load_code(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group_pair: index, %packet: index) -> (vector<1xi32>) + +template.def<@ggml.mul_mat_id_swiglu_f16_f16.accumulate_q4k_block> device @ggml_mul_mat_id_swiglu_f16_f16_accumulate_q4k_block(%q4_block: index, %load_row: index, %load_k: index, %load_packet: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %gate_weight: buffer, %up_weight: buffer, %expert_byte_base: offset, %weight_row_bytes: offset, %gate_accumulator: vector<16xf16>, %up_accumulator: vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %c0_i32x1 = vector.constant 0 : vector<1xi32> + %c0_i32x4 = vector.constant 0 : vector<4xi32> + %true = scalar.constant true : i1 + %false = scalar.constant false : i1 + %gate_header0, %gate_header1, %up_header0, %up_header1 = scf.for %header_row_offset = [%c0 to %c64 step %c32](%prior_gate_header0 = %c0_i32x4 : vector<4xi32>, %prior_gate_header1 = %c0_i32x4 : vector<4xi32>, %prior_up_header0 = %c0_i32x4 : vector<4xi32>, %prior_up_header1 = %c0_i32x4 : vector<4xi32>) -> (vector<4xi32>, vector<4xi32>, vector<4xi32>, vector<4xi32>) unroll { + %header_local_row0 = index.add %load_row, %header_row_offset : index + %header_local_row = index.assume %header_local_row0 [range(%header_local_row0, 0, 63)] : index + %header_channel = index.add %channel_tile_base, %header_local_row : index + %valid_header_channel = index.cmp ult, %header_channel, %bounded_output_size : index + %loaded_gate_header, %loaded_up_header = scf.if %valid_header_channel -> (vector<4xi32>, vector<4xi32>) { + %header_channel_byte_add = index.scale %header_channel, %weight_row_bytes : index, offset -> offset + %header_row_byte_base = index.add %expert_byte_base, %header_channel_byte_add : offset + %gate_header = func.call @ggml_q4k_split_wmma_load_header(%gate_weight, %header_row_byte_base, %q4_block) : (buffer, offset, index) -> (vector<4xi32>) + %up_header = func.call @ggml_q4k_split_wmma_load_header(%up_weight, %header_row_byte_base, %q4_block) : (buffer, offset, index) -> (vector<4xi32>) + scf.yield %gate_header, %up_header : vector<4xi32>, vector<4xi32> + } else { + scf.yield %c0_i32x4, %c0_i32x4 : vector<4xi32>, vector<4xi32> + } + %updates_header0 = index.cmp eq, %header_row_offset, %c0 : index + %next_gate_header0 = scf.select %updates_header0, %loaded_gate_header, %prior_gate_header0 : vector<4xi32> + %next_gate_header1 = scf.select %updates_header0, %prior_gate_header1, %loaded_gate_header : vector<4xi32> + %next_up_header0 = scf.select %updates_header0, %loaded_up_header, %prior_up_header0 : vector<4xi32> + %next_up_header1 = scf.select %updates_header0, %prior_up_header1, %loaded_up_header : vector<4xi32> + scf.yield %next_gate_header0, %next_gate_header1, %next_up_header0, %next_up_header1 : vector<4xi32>, vector<4xi32>, vector<4xi32>, vector<4xi32> + } + %next_block_gate, %next_block_up = scf.for %q4_group_pair = [%c0 to %c4 step %c1](%pair_gate_acc = %gate_accumulator : vector<16xf16>, %pair_up_acc = %up_accumulator : vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) { + %gate_q_word0, %gate_q_word1, %up_q_word0, %up_q_word1 = scf.for %code_row_offset = [%c0 to %c64 step %c32](%prior_gate_q_word0 = %c0_i32x1 : vector<1xi32>, %prior_gate_q_word1 = %c0_i32x1 : vector<1xi32>, %prior_up_q_word0 = %c0_i32x1 : vector<1xi32>, %prior_up_q_word1 = %c0_i32x1 : vector<1xi32>) -> (vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>) unroll { + %code_local_row0 = index.add %load_row, %code_row_offset : index + %code_local_row = index.assume %code_local_row0 [range(%code_local_row0, 0, 63)] : index + %code_channel = index.add %channel_tile_base, %code_local_row : index + %valid_code_channel = index.cmp ult, %code_channel, %bounded_output_size : index + %loaded_gate_q_word, %loaded_up_q_word = scf.if %valid_code_channel -> (vector<1xi32>, vector<1xi32>) { + %bounded_q4_group_pair = index.assume %q4_group_pair [range(%q4_group_pair, 0, 3)] : index + %code_channel_byte_add = index.scale %code_channel, %weight_row_bytes : index, offset -> offset + %code_row_byte_base = index.add %expert_byte_base, %code_channel_byte_add : offset + %gate_q_word = func.call @ggml_q4k_split_wmma_load_code(%gate_weight, %code_row_byte_base, %q4_block, %bounded_q4_group_pair, %load_packet) : (buffer, offset, index, index, index) -> (vector<1xi32>) + %up_q_word = func.call @ggml_q4k_split_wmma_load_code(%up_weight, %code_row_byte_base, %q4_block, %bounded_q4_group_pair, %load_packet) : (buffer, offset, index, index, index) -> (vector<1xi32>) + scf.yield %gate_q_word, %up_q_word : vector<1xi32>, vector<1xi32> + } else { + scf.yield %c0_i32x1, %c0_i32x1 : vector<1xi32>, vector<1xi32> + } + %updates_q_word0 = index.cmp eq, %code_row_offset, %c0 : index + %next_gate_q_word0 = scf.select %updates_q_word0, %loaded_gate_q_word, %prior_gate_q_word0 : vector<1xi32> + %next_gate_q_word1 = scf.select %updates_q_word0, %prior_gate_q_word1, %loaded_gate_q_word : vector<1xi32> + %next_up_q_word0 = scf.select %updates_q_word0, %loaded_up_q_word, %prior_up_q_word0 : vector<1xi32> + %next_up_q_word1 = scf.select %updates_q_word0, %prior_up_q_word1, %loaded_up_q_word : vector<1xi32> + scf.yield %next_gate_q_word0, %next_gate_q_word1, %next_up_q_word0, %next_up_q_word1 : vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32> + } + %next_pair_gate, %next_pair_up = scf.for %group_within_pair = [%c0 to %c2 step %c1](%group_gate = %pair_gate_acc : vector<16xf16>, %group_up = %pair_up_acc : vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) { + %q4_group_base = index.mul %q4_group_pair, %c2 : index + %q4_group0 = index.add %q4_group_base, %group_within_pair : index + %q4_group = index.assume %q4_group0 [range(%q4_group0, 0, 7)] : index + %block_k_base = index.mul %q4_block, %c256 : index + %group_k_add = index.mul %q4_group, %c32 : index + %k_origin = index.add %block_k_base, %group_k_add : index + %gate_next = template.apply<@ggml.mul_mat_id_q4k_f16_wmma.project_vector16>(%true, %q4_group, %k_origin, %load_row, %load_k, %channel_tile_base, %bounded_output_size, %assignment_count, %token_count, %route_count, %input_size, %subgroup_channel_add, %subgroup_route_add, %input, %route_stage, %weight_stage, %activation_stage, %gate_header0, %gate_header1, %gate_q_word0, %gate_q_word1, %group_gate) : (i1, index, index, index, index, index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, vector<4xi32>, vector<4xi32>, vector<1xi32>, vector<1xi32>, vector<16xf16>) -> (vector<16xf16>) + %up_next = template.apply<@ggml.mul_mat_id_q4k_f16_wmma.project_vector16>(%false, %q4_group, %k_origin, %load_row, %load_k, %channel_tile_base, %bounded_output_size, %assignment_count, %token_count, %route_count, %input_size, %subgroup_channel_add, %subgroup_route_add, %input, %route_stage, %weight_stage, %activation_stage, %up_header0, %up_header1, %up_q_word0, %up_q_word1, %group_up) : (i1, index, index, index, index, index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, vector<4xi32>, vector<4xi32>, vector<1xi32>, vector<1xi32>, vector<16xf16>) -> (vector<16xf16>) + scf.yield %gate_next, %up_next : vector<16xf16>, vector<16xf16> + } + scf.yield %next_pair_gate, %next_pair_up : vector<16xf16>, vector<16xf16> + } + template.return %next_block_gate, %next_block_up : vector<16xf16>, vector<16xf16> +} +// Q5_K routed gate/up SwiGLU accumulation. +// +// Q5_K reuses the Q4_K group-pair low-bit schedule and additionally retains the +// per-packet high-bit bytes across all groups in the block. + +template.decl @ggml.mul_mat_id_q5k_f16_wmma.project_vector16(%arg0: i1, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: index, %arg7: index, %arg8: index, %arg9: index, %arg10: index, %arg11: index, %arg12: index, %arg13: buffer, %arg14: buffer, %arg15: buffer, %arg16: buffer, %arg17: vector<4xi32>, %arg18: vector<4xi32>, %arg19: vector<1xi32>, %arg20: vector<1xi32>, %arg21: vector<4xi8>, %arg22: vector<4xi8>, %arg23: vector<16xf16>) -> (vector<16xf16>) + +func.decl @ggml_q5k_split_wmma_load_header(%weight: buffer, %row_byte_base: offset, %q5_block: index) -> (vector<4xi32>) + +func.decl @ggml_q5k_split_wmma_load_high_bits(%weight: buffer, %row_byte_base: offset, %q5_block: index, %packet: index) -> (vector<4xi8>) + +func.decl @ggml_q5k_split_wmma_load_code(%weight: buffer, %row_byte_base: offset, %q5_block: index, %q5_group_pair: index, %packet: index) -> (vector<1xi32>) + +template.def<@ggml.mul_mat_id_swiglu_f16_f16.accumulate_q5k_block> device @ggml_mul_mat_id_swiglu_f16_f16_accumulate_q5k_block(%q5_block: index, %load_row: index, %load_k: index, %load_packet: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %gate_weight: buffer, %up_weight: buffer, %expert_byte_base: offset, %weight_row_bytes: offset, %gate_accumulator: vector<16xf16>, %up_accumulator: vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %c0_i32x1 = vector.constant 0 : vector<1xi32> + %c0_i32x4 = vector.constant 0 : vector<4xi32> + %c0_i8x4 = vector.constant 0 : vector<4xi8> + %true = scalar.constant true : i1 + %false = scalar.constant false : i1 + %gate_header0, %gate_header1, %up_header0, %up_header1, %gate_qh_bytes0, %gate_qh_bytes1, %up_qh_bytes0, %up_qh_bytes1 = scf.for %header_row_offset = [%c0 to %c64 step %c32](%prior_gate_header0 = %c0_i32x4 : vector<4xi32>, %prior_gate_header1 = %c0_i32x4 : vector<4xi32>, %prior_up_header0 = %c0_i32x4 : vector<4xi32>, %prior_up_header1 = %c0_i32x4 : vector<4xi32>, %prior_gate_qh_bytes0 = %c0_i8x4 : vector<4xi8>, %prior_gate_qh_bytes1 = %c0_i8x4 : vector<4xi8>, %prior_up_qh_bytes0 = %c0_i8x4 : vector<4xi8>, %prior_up_qh_bytes1 = %c0_i8x4 : vector<4xi8>) -> (vector<4xi32>, vector<4xi32>, vector<4xi32>, vector<4xi32>, vector<4xi8>, vector<4xi8>, vector<4xi8>, vector<4xi8>) unroll { + %header_local_row0 = index.add %load_row, %header_row_offset : index + %header_local_row = index.assume %header_local_row0 [range(%header_local_row0, 0, 63)] : index + %header_channel = index.add %channel_tile_base, %header_local_row : index + %valid_header_channel = index.cmp ult, %header_channel, %bounded_output_size : index + %loaded_gate_header, %loaded_up_header, %loaded_gate_qh_bytes, %loaded_up_qh_bytes = scf.if %valid_header_channel -> (vector<4xi32>, vector<4xi32>, vector<4xi8>, vector<4xi8>) { + %header_channel_byte_add = index.scale %header_channel, %weight_row_bytes : index, offset -> offset + %header_row_byte_base = index.add %expert_byte_base, %header_channel_byte_add : offset + %gate_header = func.call @ggml_q5k_split_wmma_load_header(%gate_weight, %header_row_byte_base, %q5_block) : (buffer, offset, index) -> (vector<4xi32>) + %up_header = func.call @ggml_q5k_split_wmma_load_header(%up_weight, %header_row_byte_base, %q5_block) : (buffer, offset, index) -> (vector<4xi32>) + %gate_qh_bytes = func.call @ggml_q5k_split_wmma_load_high_bits(%gate_weight, %header_row_byte_base, %q5_block, %load_packet) : (buffer, offset, index, index) -> (vector<4xi8>) + %up_qh_bytes = func.call @ggml_q5k_split_wmma_load_high_bits(%up_weight, %header_row_byte_base, %q5_block, %load_packet) : (buffer, offset, index, index) -> (vector<4xi8>) + scf.yield %gate_header, %up_header, %gate_qh_bytes, %up_qh_bytes : vector<4xi32>, vector<4xi32>, vector<4xi8>, vector<4xi8> + } else { + scf.yield %c0_i32x4, %c0_i32x4, %c0_i8x4, %c0_i8x4 : vector<4xi32>, vector<4xi32>, vector<4xi8>, vector<4xi8> + } + %updates_header0 = index.cmp eq, %header_row_offset, %c0 : index + %next_gate_header0 = scf.select %updates_header0, %loaded_gate_header, %prior_gate_header0 : vector<4xi32> + %next_gate_header1 = scf.select %updates_header0, %prior_gate_header1, %loaded_gate_header : vector<4xi32> + %next_up_header0 = scf.select %updates_header0, %loaded_up_header, %prior_up_header0 : vector<4xi32> + %next_up_header1 = scf.select %updates_header0, %prior_up_header1, %loaded_up_header : vector<4xi32> + %next_gate_qh_bytes0 = scf.select %updates_header0, %loaded_gate_qh_bytes, %prior_gate_qh_bytes0 : vector<4xi8> + %next_gate_qh_bytes1 = scf.select %updates_header0, %prior_gate_qh_bytes1, %loaded_gate_qh_bytes : vector<4xi8> + %next_up_qh_bytes0 = scf.select %updates_header0, %loaded_up_qh_bytes, %prior_up_qh_bytes0 : vector<4xi8> + %next_up_qh_bytes1 = scf.select %updates_header0, %prior_up_qh_bytes1, %loaded_up_qh_bytes : vector<4xi8> + scf.yield %next_gate_header0, %next_gate_header1, %next_up_header0, %next_up_header1, %next_gate_qh_bytes0, %next_gate_qh_bytes1, %next_up_qh_bytes0, %next_up_qh_bytes1 : vector<4xi32>, vector<4xi32>, vector<4xi32>, vector<4xi32>, vector<4xi8>, vector<4xi8>, vector<4xi8>, vector<4xi8> + } + %next_block_gate, %next_block_up = scf.for %q5_group_pair = [%c0 to %c4 step %c1](%pair_gate_acc = %gate_accumulator : vector<16xf16>, %pair_up_acc = %up_accumulator : vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) { + %gate_q_word0, %gate_q_word1, %up_q_word0, %up_q_word1 = scf.for %code_row_offset = [%c0 to %c64 step %c32](%prior_gate_q_word0 = %c0_i32x1 : vector<1xi32>, %prior_gate_q_word1 = %c0_i32x1 : vector<1xi32>, %prior_up_q_word0 = %c0_i32x1 : vector<1xi32>, %prior_up_q_word1 = %c0_i32x1 : vector<1xi32>) -> (vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>) unroll { + %code_local_row0 = index.add %load_row, %code_row_offset : index + %code_local_row = index.assume %code_local_row0 [range(%code_local_row0, 0, 63)] : index + %code_channel = index.add %channel_tile_base, %code_local_row : index + %valid_code_channel = index.cmp ult, %code_channel, %bounded_output_size : index + %loaded_gate_q_word, %loaded_up_q_word = scf.if %valid_code_channel -> (vector<1xi32>, vector<1xi32>) { + %bounded_q5_group_pair = index.assume %q5_group_pair [range(%q5_group_pair, 0, 3)] : index + %code_channel_byte_add = index.scale %code_channel, %weight_row_bytes : index, offset -> offset + %code_row_byte_base = index.add %expert_byte_base, %code_channel_byte_add : offset + %gate_q_word = func.call @ggml_q5k_split_wmma_load_code(%gate_weight, %code_row_byte_base, %q5_block, %bounded_q5_group_pair, %load_packet) : (buffer, offset, index, index, index) -> (vector<1xi32>) + %up_q_word = func.call @ggml_q5k_split_wmma_load_code(%up_weight, %code_row_byte_base, %q5_block, %bounded_q5_group_pair, %load_packet) : (buffer, offset, index, index, index) -> (vector<1xi32>) + scf.yield %gate_q_word, %up_q_word : vector<1xi32>, vector<1xi32> + } else { + scf.yield %c0_i32x1, %c0_i32x1 : vector<1xi32>, vector<1xi32> + } + %updates_q_word0 = index.cmp eq, %code_row_offset, %c0 : index + %next_gate_q_word0 = scf.select %updates_q_word0, %loaded_gate_q_word, %prior_gate_q_word0 : vector<1xi32> + %next_gate_q_word1 = scf.select %updates_q_word0, %prior_gate_q_word1, %loaded_gate_q_word : vector<1xi32> + %next_up_q_word0 = scf.select %updates_q_word0, %loaded_up_q_word, %prior_up_q_word0 : vector<1xi32> + %next_up_q_word1 = scf.select %updates_q_word0, %prior_up_q_word1, %loaded_up_q_word : vector<1xi32> + scf.yield %next_gate_q_word0, %next_gate_q_word1, %next_up_q_word0, %next_up_q_word1 : vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32> + } + %next_pair_gate, %next_pair_up = scf.for %group_within_pair = [%c0 to %c2 step %c1](%group_gate = %pair_gate_acc : vector<16xf16>, %group_up = %pair_up_acc : vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) { + %q5_group_base = index.mul %q5_group_pair, %c2 : index + %q5_group0 = index.add %q5_group_base, %group_within_pair : index + %q5_group = index.assume %q5_group0 [range(%q5_group0, 0, 7)] : index + %block_k_base = index.mul %q5_block, %c256 : index + %group_k_add = index.mul %q5_group, %c32 : index + %k_origin = index.add %block_k_base, %group_k_add : index + %gate_next = template.apply<@ggml.mul_mat_id_q5k_f16_wmma.project_vector16>(%true, %q5_group, %k_origin, %load_row, %load_k, %channel_tile_base, %bounded_output_size, %assignment_count, %token_count, %route_count, %input_size, %subgroup_channel_add, %subgroup_route_add, %input, %route_stage, %weight_stage, %activation_stage, %gate_header0, %gate_header1, %gate_q_word0, %gate_q_word1, %gate_qh_bytes0, %gate_qh_bytes1, %group_gate) : (i1, index, index, index, index, index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, vector<4xi32>, vector<4xi32>, vector<1xi32>, vector<1xi32>, vector<4xi8>, vector<4xi8>, vector<16xf16>) -> (vector<16xf16>) + %up_next = template.apply<@ggml.mul_mat_id_q5k_f16_wmma.project_vector16>(%false, %q5_group, %k_origin, %load_row, %load_k, %channel_tile_base, %bounded_output_size, %assignment_count, %token_count, %route_count, %input_size, %subgroup_channel_add, %subgroup_route_add, %input, %route_stage, %weight_stage, %activation_stage, %up_header0, %up_header1, %up_q_word0, %up_q_word1, %up_qh_bytes0, %up_qh_bytes1, %group_up) : (i1, index, index, index, index, index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, vector<4xi32>, vector<4xi32>, vector<1xi32>, vector<1xi32>, vector<4xi8>, vector<4xi8>, vector<16xf16>) -> (vector<16xf16>) + scf.yield %gate_next, %up_next : vector<16xf16>, vector<16xf16> + } + scf.yield %next_pair_gate, %next_pair_up : vector<16xf16>, vector<16xf16> + } + template.return %next_block_gate, %next_block_up : vector<16xf16>, vector<16xf16> +} +// Q6_K routed gate/up SwiGLU accumulation. +// +// The shared body owns routing and publish. This file only binds Q6_K block +// accumulation to that common schedule. + +template.decl @ggml.mul_mat_id_q6k_f16_wmma.project_vector16(%arg0: i1, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: index, %arg7: index, %arg8: index, %arg9: index, %arg10: index, %arg11: index, %arg12: index, %arg13: index, %arg14: index, %arg15: buffer, %arg16: buffer, %arg17: buffer, %arg18: buffer, %arg19: buffer, %arg20: offset, %arg21: offset, %arg22: vector<1xi32>, %arg23: vector<1xi32>, %arg24: vector<1xi32>, %arg25: vector<1xi32>, %arg26: vector<1xi32>, %arg27: vector<1xi32>, %arg28: vector<16xf16>) -> (vector<16xf16>) + +func.decl @ggml_q6k_split_wmma_load_half_codes(%weight: buffer, %weight_row_byte_base: offset, %q6_block: index, %q6_half: index, %packet: index) -> (vector<1xi32>, vector<1xi32>, vector<1xi32>) + +template.def<@ggml.mul_mat_id_swiglu_f16_f16.accumulate_q6k_block> device @ggml_mul_mat_id_swiglu_f16_f16_accumulate_q6k_block(%q6_block: index, %load_row: index, %load_k: index, %load_packet: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %gate_weight: buffer, %up_weight: buffer, %expert_byte_base: offset, %weight_row_bytes: offset, %gate_accumulator: vector<16xf16>, %up_accumulator: vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %c0_i32x1 = vector.constant 0 : vector<1xi32> + %true = scalar.constant true : i1 + %false = scalar.constant false : i1 + %next_block_gate, %next_block_up = scf.for %q6_half = [%c0 to %c2 step %c1](%half_gate_acc = %gate_accumulator : vector<16xf16>, %half_up_acc = %up_accumulator : vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) { + %gate_ql0_word0, %gate_ql1_word0, %gate_qh_word0, %gate_ql0_word1, %gate_ql1_word1, %gate_qh_word1, %up_ql0_word0, %up_ql1_word0, %up_qh_word0, %up_ql0_word1, %up_ql1_word1, %up_qh_word1 = scf.for %code_row_offset = [%c0 to %c64 step %c32](%prior_gate_ql0_word0 = %c0_i32x1 : vector<1xi32>, %prior_gate_ql1_word0 = %c0_i32x1 : vector<1xi32>, %prior_gate_qh_word0 = %c0_i32x1 : vector<1xi32>, %prior_gate_ql0_word1 = %c0_i32x1 : vector<1xi32>, %prior_gate_ql1_word1 = %c0_i32x1 : vector<1xi32>, %prior_gate_qh_word1 = %c0_i32x1 : vector<1xi32>, %prior_up_ql0_word0 = %c0_i32x1 : vector<1xi32>, %prior_up_ql1_word0 = %c0_i32x1 : vector<1xi32>, %prior_up_qh_word0 = %c0_i32x1 : vector<1xi32>, %prior_up_ql0_word1 = %c0_i32x1 : vector<1xi32>, %prior_up_ql1_word1 = %c0_i32x1 : vector<1xi32>, %prior_up_qh_word1 = %c0_i32x1 : vector<1xi32>) -> (vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>) unroll { + %code_local_row0 = index.add %load_row, %code_row_offset : index + %code_local_row = index.assume %code_local_row0 [range(%code_local_row0, 0, 63)] : index + %code_channel = index.add %channel_tile_base, %code_local_row : index + %valid_code_channel = index.cmp ult, %code_channel, %bounded_output_size : index + %loaded_gate_ql0, %loaded_gate_ql1, %loaded_gate_qh, %loaded_up_ql0, %loaded_up_ql1, %loaded_up_qh = scf.if %valid_code_channel -> (vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>) { + %bounded_q6_half = index.assume %q6_half [range(%q6_half, 0, 1)] : index + %code_channel_byte_add = index.scale %code_channel, %weight_row_bytes : index, offset -> offset + %code_row_byte_base = index.add %expert_byte_base, %code_channel_byte_add : offset + %gate_ql0, %gate_ql1, %gate_qh = func.call @ggml_q6k_split_wmma_load_half_codes(%gate_weight, %code_row_byte_base, %q6_block, %bounded_q6_half, %load_packet) : (buffer, offset, index, index, index) -> (vector<1xi32>, vector<1xi32>, vector<1xi32>) + %up_ql0, %up_ql1, %up_qh = func.call @ggml_q6k_split_wmma_load_half_codes(%up_weight, %code_row_byte_base, %q6_block, %bounded_q6_half, %load_packet) : (buffer, offset, index, index, index) -> (vector<1xi32>, vector<1xi32>, vector<1xi32>) + scf.yield %gate_ql0, %gate_ql1, %gate_qh, %up_ql0, %up_ql1, %up_qh : vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32> + } else { + scf.yield %c0_i32x1, %c0_i32x1, %c0_i32x1, %c0_i32x1, %c0_i32x1, %c0_i32x1 : vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32> + } + %updates_word0 = index.cmp eq, %code_row_offset, %c0 : index + %next_gate_ql0_word0 = scf.select %updates_word0, %loaded_gate_ql0, %prior_gate_ql0_word0 : vector<1xi32> + %next_gate_ql1_word0 = scf.select %updates_word0, %loaded_gate_ql1, %prior_gate_ql1_word0 : vector<1xi32> + %next_gate_qh_word0 = scf.select %updates_word0, %loaded_gate_qh, %prior_gate_qh_word0 : vector<1xi32> + %next_gate_ql0_word1 = scf.select %updates_word0, %prior_gate_ql0_word1, %loaded_gate_ql0 : vector<1xi32> + %next_gate_ql1_word1 = scf.select %updates_word0, %prior_gate_ql1_word1, %loaded_gate_ql1 : vector<1xi32> + %next_gate_qh_word1 = scf.select %updates_word0, %prior_gate_qh_word1, %loaded_gate_qh : vector<1xi32> + %next_up_ql0_word0 = scf.select %updates_word0, %loaded_up_ql0, %prior_up_ql0_word0 : vector<1xi32> + %next_up_ql1_word0 = scf.select %updates_word0, %loaded_up_ql1, %prior_up_ql1_word0 : vector<1xi32> + %next_up_qh_word0 = scf.select %updates_word0, %loaded_up_qh, %prior_up_qh_word0 : vector<1xi32> + %next_up_ql0_word1 = scf.select %updates_word0, %prior_up_ql0_word1, %loaded_up_ql0 : vector<1xi32> + %next_up_ql1_word1 = scf.select %updates_word0, %prior_up_ql1_word1, %loaded_up_ql1 : vector<1xi32> + %next_up_qh_word1 = scf.select %updates_word0, %prior_up_qh_word1, %loaded_up_qh : vector<1xi32> + scf.yield %next_gate_ql0_word0, %next_gate_ql1_word0, %next_gate_qh_word0, %next_gate_ql0_word1, %next_gate_ql1_word1, %next_gate_qh_word1, %next_up_ql0_word0, %next_up_ql1_word0, %next_up_qh_word0, %next_up_ql0_word1, %next_up_ql1_word1, %next_up_qh_word1 : vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32> + } + %next_half_gate, %next_half_up = scf.for %group_within_half = [%c0 to %c4 step %c1](%group_gate = %half_gate_acc : vector<16xf16>, %group_up = %half_up_acc : vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) { + %q6_group_base = index.mul %q6_half, %c4 : index + %q6_group0 = index.add %q6_group_base, %group_within_half : index + %q6_group = index.assume %q6_group0 [range(%q6_group0, 0, 7)] : index + %block_k_base = index.mul %q6_block, %c256 : index + %group_k_add = index.mul %q6_group, %c32 : index + %k_origin = index.add %block_k_base, %group_k_add : index + %gate_next = template.apply<@ggml.mul_mat_id_q6k_f16_wmma.project_vector16>(%true, %q6_group, %q6_block, %k_origin, %load_row, %load_k, %load_packet, %channel_tile_base, %bounded_output_size, %assignment_count, %token_count, %route_count, %input_size, %subgroup_channel_add, %subgroup_route_add, %input, %route_stage, %weight_stage, %activation_stage, %gate_weight, %expert_byte_base, %weight_row_bytes, %gate_ql0_word0, %gate_ql1_word0, %gate_qh_word0, %gate_ql0_word1, %gate_ql1_word1, %gate_qh_word1, %group_gate) : (i1, index, index, index, index, index, index, index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer, offset, offset, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<16xf16>) -> (vector<16xf16>) + %up_next = template.apply<@ggml.mul_mat_id_q6k_f16_wmma.project_vector16>(%false, %q6_group, %q6_block, %k_origin, %load_row, %load_k, %load_packet, %channel_tile_base, %bounded_output_size, %assignment_count, %token_count, %route_count, %input_size, %subgroup_channel_add, %subgroup_route_add, %input, %route_stage, %weight_stage, %activation_stage, %up_weight, %expert_byte_base, %weight_row_bytes, %up_ql0_word0, %up_ql1_word0, %up_qh_word0, %up_ql0_word1, %up_ql1_word1, %up_qh_word1, %group_up) : (i1, index, index, index, index, index, index, index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer, offset, offset, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<16xf16>) -> (vector<16xf16>) + scf.yield %gate_next, %up_next : vector<16xf16>, vector<16xf16> + } + scf.yield %next_half_gate, %next_half_up : vector<16xf16>, vector<16xf16> + } + template.return %next_block_gate, %next_block_up : vector<16xf16>, vector<16xf16> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_swiglu_f16_f16_wmma_body.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_swiglu_f16_f16_wmma_body.loom new file mode 100644 index 000000000000..21169146fca9 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_id_swiglu_f16_f16_wmma_body.loom @@ -0,0 +1,193 @@ +// Shared routed gate/up SwiGLU schedule for FP16 WMMA fast paths. +// +// The body owns routing, staging lifetime, accumulator flow, SwiGLU, and f16 +// publish. A common accumulator motif selects the dtype-specific inner loop. + +template.decl @ggml.mul_mat_id_swiglu_f16_f16.accumulate_block(%quant_block: index, %load_row: index, %load_k: index, %load_packet: index, %channel_tile_base: index, %bounded_output_size: index, %assignment_count: index, %token_count: index, %route_count: index, %input_size: index, %subgroup_channel_add: index, %subgroup_route_add: index, %input: buffer, %route_stage: buffer, %weight_stage: buffer, %activation_stage: buffer, %gate_weight: buffer, %up_weight: buffer, %expert_byte_base: offset, %gate_weight_row_bytes: offset, %up_weight_row_bytes: offset, %gate_weight_format: index, %up_weight_format: index, %gate_accumulator: vector<16xf16>, %up_accumulator: vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) + +template.decl @ggml.mul_mat_id_swiglu_f16_f16.body(%gate_weight_format: index, %up_weight_format: index, %token_count: index, %input_size: index, %route_count: index, %expert_count: index, %output_size: index, %descriptor_expert_mask: i32, %descriptor_partition_shift: i32, %descriptor_row_count_shift: i32, %channel_tile: index, %partition_ordinal: index, %input: buffer, %expert_table: buffer, %partition_table: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) + +template.decl @ggml.mul_mat_id_swiglu_f16_f16.apply_vector8(%gate: vector<8xf32>, %up: vector<8xf32>) -> (vector<8xf32>) + +func.decl @ggml_dequant_weight_tile_bytes(%weight_format: index) -> (offset) + +func.def inline @ggml_mul_mat_id_swiglu_f16_f16_unpack_expert_partition_descriptor(%descriptor: i32, %expert_mask: i32, %partition_shift: i32, %row_count_shift: i32) -> (index, index, index) { + %c1_i32 = scalar.constant 1 : i32 + %c5_i32 = scalar.constant 5 : i32 + %c31_i32 = scalar.constant 31 : i32 + %c63_i32 = scalar.constant 63 : i32 + %expert_i32 = scalar.andi %descriptor, %expert_mask : i32 + %partition_shifted_i32 = scalar.shrui %descriptor, %partition_shift : i32 + %partition_i32 = scalar.andi %partition_shifted_i32, %c63_i32 : i32 + %route_tile_base_i32 = scalar.shli %partition_i32, %c5_i32 : i32 + %row_count_shifted_i32 = scalar.shrui %descriptor, %row_count_shift : i32 + %row_count_minus_one_i32 = scalar.andi %row_count_shifted_i32, %c31_i32 : i32 + %partition_row_count_i32 = scalar.addi %row_count_minus_one_i32, %c1_i32 : i32 + %expert0 = index.cast %expert_i32 : i32 to index + %expert = index.assume %expert0 [range(%expert0, 0, 127)] : index + %route_tile_base0 = index.cast %route_tile_base_i32 : i32 to index + %route_tile_base = index.assume %route_tile_base0 [range(%route_tile_base0, 0, 2016)] : index + %partition_row_count0 = index.cast %partition_row_count_i32 : i32 to index + %partition_row_count = index.assume %partition_row_count0 [range(%partition_row_count0, 1, 32)] : index + func.return %expert, %route_tile_base, %partition_row_count : index, index, index +} + +template.def<@ggml.mul_mat_id_swiglu_f16_f16.apply_vector8> device @ggml_mul_mat_id_swiglu_f16_f16_apply_swiglu_vector8(%gate: vector<8xf32>, %up: vector<8xf32>) -> (vector<8xf32>) { + %activated = vector.siluf %gate : vector<8xf32> + %result = vector.mulf %activated, %up : vector<8xf32> + template.return %result : vector<8xf32> +} + +template.def<@ggml.mul_mat_id_swiglu_f16_f16.body> device @ggml_mul_mat_id_swiglu_f16_f16_wmma_body(%gate_weight_format: index, %up_weight_format: index, %token_count: index, %input_size: index, %route_count: index, %expert_count: index, %output_size: index, %descriptor_expert_mask: i32, %descriptor_partition_shift: i32, %descriptor_row_count_shift: i32, %channel_tile: index, %partition_ordinal: index, %input: buffer, %expert_table: buffer, %partition_table: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) { + %bounded_route_count = index.assume %route_count [range(%route_count, 1, 8)] : index + %bounded_expert_count = index.assume %expert_count [range(%expert_count, 1, 128)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 4096)] : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %c4_bytes = index.constant 4 : offset + %weight_stage_bytes = index.constant 5120 : offset + %activation_stage_bytes = index.constant 2560 : offset + %route_stage_bytes = index.constant 128 : offset + %wave_result_stage_bytes = index.constant 512 : offset + %result_stage_bytes = index.constant 4096 : offset + %c0_i32 = scalar.constant 0 : i32 + %cn1_i32 = scalar.constant -1 : i32 + %zero_accumulator = vector.constant 0.0 : vector<16xf16> + %m = index.constant 16 : index + %n = index.constant 16 : index + %assignment_count = index.mul %token_count, %bounded_route_count : index + %rounded_assignment_count = index.add %assignment_count, %c31 : index + %assignment_partition_count = index.div %rounded_assignment_count, %c32 : index + %maximum_partition_count = index.add %assignment_partition_count, %bounded_expert_count : index + %has_more_assignment_partitions = index.cmp ugt, %assignment_partition_count, %bounded_expert_count : index + %launch_partition_count = scf.select %has_more_assignment_partitions, %assignment_partition_count, %bounded_expert_count : index + %assignment_table_byte_base = index.scale %bounded_expert_count, %c4_bytes : index, offset -> offset + %gate_weight_block_bytes = func.call @ggml_dequant_weight_tile_bytes(%gate_weight_format) : (index) -> (offset) + %up_weight_block_bytes = func.call @ggml_dequant_weight_tile_bytes(%up_weight_format) : (index) -> (offset) + %quant_block_count = index.div %input_size, %c256 : index + %gate_weight_row_bytes = index.scale %quant_block_count, %gate_weight_block_bytes : index, offset -> offset + %up_weight_row_bytes = index.scale %quant_block_count, %up_weight_block_bytes : index, offset -> offset + %gate_weight_expert_bytes = index.scale %bounded_output_size, %gate_weight_row_bytes : index, offset -> offset + %output_row_count = index.mul %token_count, %bounded_route_count : index + %input_noalias, %expert_table_noalias, %partition_table_noalias, %gate_weight_noalias, %up_weight_noalias, %output_noalias = buffer.assume.noalias %input, %expert_table, %partition_table, %gate_weight, %up_weight, %output : buffer, buffer, buffer, buffer, buffer, buffer + %assignment_view = buffer.view %expert_table_noalias[%assignment_table_byte_base] : buffer -> view<[%bounded_expert_count]x[%token_count]xi32> + %partition_count_view = buffer.view %partition_table_noalias[%c0_offset] : buffer -> view<1xi32> + %partition_descriptor_view = buffer.view %partition_table_noalias[%c4_bytes] : buffer -> view<[%maximum_partition_count]xi32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%output_row_count]x[%bounded_output_size]xf16> + %weight_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %activation_stage = buffer.alloca align(16) %activation_stage_bytes : buffer + %route_stage = buffer.alloca align(16) %route_stage_bytes : buffer + %result_stage = buffer.alloca align(16) %result_stage_bytes : buffer + %route_stage_view = buffer.view %route_stage[%c0_offset] : buffer -> view<32xi32> + %wave_result_stage_offset = index.scale %subgroup, %wave_result_stage_bytes : index, offset -> offset + %result_fragment_layout = encoding.layout.strided [1, %c16] : encoding + %result_fragment_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf16, %result_fragment_layout> + %result_physical_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf16> + %channel_tile_base = index.mul %channel_tile, %c64 : index + %is_workitem_zero = index.cmp eq, %workitem, %c0 : index + %lane_partition_count_i32 = scf.if %is_workitem_zero -> (i32) { + %loaded = view.load %partition_count_view[%c0] : view<1xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %c0_i32 : i32 + } + %partition_count_reduced = kernel.workgroup.reduce %lane_partition_count_i32 : i32 + %partition_count_i32 = kernel.subgroup.broadcast.first %partition_count_reduced : i32 + %partition_count0 = index.cast %partition_count_i32 : i32 to index + %partition_count, %partition_capacity = index.assume %partition_count0, %maximum_partition_count [lt(%partition_count0, %maximum_partition_count)] : index, index + scf.for %active_partition = [%partition_ordinal to %partition_count step %launch_partition_count] { + %descriptor_ordinal, %descriptor_count = index.assume %active_partition, %partition_count [lt(%active_partition, %partition_count)] : index, index + %table_descriptor_ordinal, %table_descriptor_capacity = index.assume %descriptor_ordinal, %maximum_partition_count [lt(%descriptor_ordinal, %maximum_partition_count)] : index, index + %lane_descriptor_i32 = scf.if %is_workitem_zero -> (i32) { + %loaded = view.load %partition_descriptor_view[%table_descriptor_ordinal] : view<[%maximum_partition_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %c0_i32 : i32 + } + %descriptor_reduced_i32 = kernel.workgroup.reduce %lane_descriptor_i32 : i32 + %descriptor_i32 = kernel.subgroup.broadcast.first %descriptor_reduced_i32 : i32 + %expert, %route_tile_base, %partition_row_count = func.call @ggml_mul_mat_id_swiglu_f16_f16_unpack_expert_partition_descriptor(%descriptor_i32, %descriptor_expert_mask, %descriptor_partition_shift, %descriptor_row_count_shift) : (i32, i32, i32, i32) -> (index, index, index) + %bounded_expert, %table_expert_count = index.assume %expert, %bounded_expert_count [lt(%expert, %bounded_expert_count)] : index, index + %loads_route = index.cmp ult, %workitem, %c32 : index + scf.if %loads_route { + %local_route = index.assume %workitem [range(%workitem, 0, 31)] : index + %assignment_ordinal = index.add %route_tile_base, %local_route : index + %valid_row = index.cmp ult, %local_route, %partition_row_count : index + %assignment_i32 = scf.if %valid_row -> (i32) { + %bounded_assignment_ordinal, %table_token_count = index.assume %assignment_ordinal, %token_count [lt(%assignment_ordinal, %token_count)] : index, index + %loaded = view.load %assignment_view[%bounded_expert, %bounded_assignment_ordinal] : view<[%bounded_expert_count]x[%token_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %cn1_i32 : i32 + } + view.store %assignment_i32, %route_stage_view[%local_route] : i32, view<32xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %load_packet = index.rem %workitem, %c8 : index + %load_k = index.mul %load_packet, %c4 : index + %load_row0 = index.div %workitem, %c8 : index + %load_row = index.assume %load_row0 [range(%load_row0, 0, 31)] : index + %expert_byte_base = index.scale %bounded_expert, %gate_weight_expert_bytes : index, offset -> offset + %channel_subgroup = index.rem %subgroup, %c4 : index + %route_subgroup = index.div %subgroup, %c4 : index + %subgroup_channel_add = index.mul %channel_subgroup, %c16 : index + %subgroup_route_add = index.mul %route_subgroup, %c16 : index + %init_gate = vector.fragment %zero_accumulator shape [%m, %n] : vector<16xf16> + %init_up = vector.fragment %zero_accumulator shape [%m, %n] : vector<16xf16> + %gate_result, %up_result = scf.for %quant_block = [%c0 to %quant_block_count step %c1](%block_gate = %init_gate : vector<16xf16>, %block_up = %init_up : vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) { + %next_gate, %next_up = template.apply<@ggml.mul_mat_id_swiglu_f16_f16.accumulate_block>(%quant_block, %load_row, %load_k, %load_packet, %channel_tile_base, %bounded_output_size, %assignment_count, %token_count, %bounded_route_count, %input_size, %subgroup_channel_add, %subgroup_route_add, %input_noalias, %route_stage, %weight_stage, %activation_stage, %gate_weight_noalias, %up_weight_noalias, %expert_byte_base, %gate_weight_row_bytes, %up_weight_row_bytes, %gate_weight_format, %up_weight_format, %block_gate, %block_up) : (index, index, index, index, index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, offset, offset, offset, index, index, vector<16xf16>, vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) + scf.yield %next_gate, %next_up : vector<16xf16>, vector<16xf16> + } + %publish_route0 = index.div %lane, %c2 : index + %publish_route = index.assume %publish_route0 [range(%publish_route0, 0, 15)] : index + %publish_packet0 = index.rem %lane, %c2 : index + %publish_packet = index.assume %publish_packet0 [range(%publish_packet0, 0, 1)] : index + %publish_channel_add = index.mul %publish_packet, %c8 : index + %local_route = index.add %subgroup_route_add, %publish_route : index + %assignment_i32 = view.load %route_stage_view[%local_route] : view<32xi32> -> i32 + %assignment_nonnegative = scalar.cmpi sge, %assignment_i32, %c0_i32 : i32 + %safe_assignment_i32 = scf.if %assignment_nonnegative -> (i32) { + scf.yield %assignment_i32 : i32 + } else { + scf.yield %c0_i32 : i32 + } + %safe_assignment0 = index.cast %safe_assignment_i32 : i32 to index + %safe_assignment = index.assume %safe_assignment0 [range(%safe_assignment0, 0, 16383)] : index + %bounded_assignment, %bounded_assignment_count = index.assume %safe_assignment, %assignment_count [lt(%safe_assignment, %assignment_count)] : index, index + %subgroup_channel_base = index.add %channel_tile_base, %subgroup_channel_add : index + %channel = index.add %subgroup_channel_base, %publish_channel_add : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %writes = scalar.andi %assignment_nonnegative, %valid_channel : i1 + vector.fragment.store %gate_result, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<16xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + %gate_values = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf16> -> vector<8xf16> + %gate_wide = vector.extf %gate_values : vector<8xf16> to vector<8xf32> + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %up_result, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<16xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes { + %up_values = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf16> -> vector<8xf16> + %up_wide = vector.extf %up_values : vector<8xf16> to vector<8xf32> + %wide_values = template.apply<@ggml.mul_mat_id_swiglu_f16_f16.apply_vector8>(%gate_wide, %up_wide) : (vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + %values = vector.fptrunc %wide_values : vector<8xf32> to vector<8xf16> + %mask = vector.mask.range [%channel to %bounded_output_size step %c1] : index -> vector<8xi1> + vector.store.mask %values, %output_view[%bounded_assignment, %channel], %mask : vector<8xf16>, view<[%output_row_count]x[%bounded_output_size]xf16>, vector<8xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + } + template.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_quantized_f16_prefill.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_quantized_f16_prefill.loom new file mode 100644 index 000000000000..07f2ef3fe8aa --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/mul_mat_quantized_f16_prefill.loom @@ -0,0 +1,780 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +template.decl @ggml.binary_f32.apply_vector8(%op: index, %lhs: vector<8xf32>, %rhs: vector<8xf32>) -> (vector<8xf32>) + +func.decl @ggml_q4k_native_row64_f16_pair(%half_packet: i1, %paired: i1, %weight: buffer, %peer: buffer, %is_up: i1, %row_base: offset, %group0: index, %packet: index) -> (vector<16xf16>, vector<16xf16>) +func.decl @ggml_q6k_f16_pair(%half_packet: i1, %weight: buffer, %weight_row_byte_base: offset, %q6_block: index, %q6_group: index, %packet: index) -> (vector<16xf16>, vector<16xf16>) + +func.def inline @ggml_f16_prefill_load_pair(%input: buffer, %token_count: index, %input_size: index, %token_origin: index, %k_origin: index) -> (vector<8xf16>, vector<8xf16>) { + %c0_offset = index.constant 0 : offset + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c1024 = index.constant 1024 : index + %workgroup_size = kernel.workgroup.size : index + %larger_tile = index.cmp eq, %workgroup_size, %c1024 : index + %row_packets = scf.select %larger_tile, %c4, %c2 : index + %packet_width = scf.select %larger_tile, %c8, %c16 : index + %workitem = kernel.workitem.id : index + %row0 = index.div %workitem, %row_packets : index + %row = index.assume %row0 [range(%row0, 0, 255)] : index + %packet = index.rem %workitem, %row_packets : index + %packet_k = index.mul %packet, %packet_width : index + %token = index.add %token_origin, %row : index + %input_k0 = index.add %k_origin, %packet_k : index + %input_k1 = index.add %input_k0, %c8 : index + %view = buffer.view %input[%c0_offset] : buffer -> view<[%token_count]x[%input_size]xf16> + %valid = index.cmp ult, %token, %token_count : index + %zero = vector.constant 0.0 : vector<8xf16> + %lo, %hi = scf.if %valid -> (vector<8xf16>, vector<8xf16>) { + %bounded_token, %bounded_count = index.assume %token, %token_count [lt(%token, %token_count)] : index, index + %v0 = vector.load %view[%bounded_token, %input_k0] : view<[%token_count]x[%input_size]xf16> -> vector<8xf16> + %v1 = scf.if %larger_tile -> (vector<8xf16>) { + scf.yield %zero : vector<8xf16> + } else { + %loaded = vector.load %view[%bounded_token, %input_k1] : view<[%token_count]x[%input_size]xf16> -> vector<8xf16> + scf.yield %loaded : vector<8xf16> + } + scf.yield %v0, %v1 : vector<8xf16>, vector<8xf16> + } else { + scf.yield %zero, %zero : vector<8xf16>, vector<8xf16> + } + func.return %lo, %hi : vector<8xf16>, vector<8xf16> +} + + +func.def inline @ggml_f16_prefill_stage_mma(%paired: i1, %packed_input: i1, %bounded_token_count: index, %bounded_input_size: index, %token_tile_origin: index, %quant_block: index, %weight_group_base: index, %quant_group: index, %input_noalias: buffer, %weight_stage: buffer, %activation_stage: buffer, %stage_offset: offset, %pipeline_weights: i1, %acc00: vector<8xf32>, %acc01: vector<8xf32>, %acc02: vector<8xf32>, %acc03: vector<8xf32>, %acc10: vector<8xf32>, %acc11: vector<8xf32>, %acc12: vector<8xf32>, %acc13: vector<8xf32>, %current_a0: vector<8xf16>, %current_a1: vector<8xf16>) -> (vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf16>, vector<8xf16>) { + %layout_true = scalar.constant true : i1 + %sync_weights = scalar.xori %pipeline_weights, %layout_true : i1 + %staged_input = scalar.xori %packed_input, %layout_true : i1 + %workgroup_size = kernel.workgroup.size : index + %large_workgroup = index.constant 512 : index + %wide_tile = index.cmp uge, %workgroup_size, %large_workgroup : index + %can_prefetch_a = index.cmp eq, %workgroup_size, %large_workgroup : index + %prefetch_a = scalar.andi %can_prefetch_a, %staged_input : i1 + %largest_workgroup = index.constant 1024 : index + %wg1024 = index.cmp eq, %workgroup_size, %largest_workgroup : index + %larger_tile = scalar.andi %wg1024, %staged_input : i1 + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c15 = index.constant 15 : index + %c16 = index.constant 16 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c40 = index.constant 40 : index + %c64 = index.constant 64 : index + %c72 = index.constant 72 : index + %c127 = index.constant 127 : index + %c128 = index.constant 128 : index + %c136 = index.constant 136 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %c512 = index.constant 512 : index + %c0_offset = index.constant 0 : offset + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %true = scalar.constant true : i1 + %not_wide_tile = scalar.xori %wide_tile, %true : i1 + %zero_f16 = vector.constant 0.0 : vector<16xf16> + %weight_stride = scf.select %wide_tile, %c72, %c136 : index + %ordinary_weight_rows = scf.select %wide_tile, %c128, %c64 : index + %staged_weight_rows = scf.select %larger_tile, %c256, %ordinary_weight_rows : index + %packed_weight_rows = scf.select %wg1024, %c128, %c64 : index + %weight_rows = scf.select %packed_input, %packed_weight_rows, %staged_weight_rows : index + %weight_window = scf.select %wide_tile, %c64, %c128 : index + %input_f16_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_input_size]xf16> + %activation_stage_physical_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<256x40xf16> + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %ordinary_subgroup_limit = scf.select %wide_tile, %c15, %c7 : index + %subgroup_limit = scf.select %wg1024, %c31, %ordinary_subgroup_limit : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, %subgroup_limit)] : index + %token_subgroups = scf.select %packed_input, %c16, %c8 : index + %token_subgroup = index.rem %subgroup, %token_subgroups : index + %subgroup_token_add = index.mul %token_subgroup, %c32 : index + %channel_subgroup = index.div %subgroup, %token_subgroups : index + %subgroup_channel_stride = scf.select %paired, %c32, %c64 : index + %subgroup_channel_add = index.mul %channel_subgroup, %subgroup_channel_stride : index + %half_weight_rows = index.div %weight_rows, %c2 : index + %peer_channel_delta = scf.select %paired, %half_weight_rows, %c32 : index + %rhs_column0 = index.add %subgroup_channel_add, %c0 : index + %rhs_column1 = index.add %rhs_column0, %c16 : index + %rhs_column2 = index.add %rhs_column0, %peer_channel_delta : index + %rhs_column3 = index.add %rhs_column2, %c16 : index + %activation_row_packets = scf.select %larger_tile, %c4, %c2 : index + %activation_packet_width = scf.select %larger_tile, %c8, %c16 : index + %activation_load_row0 = index.div %workitem, %activation_row_packets : index + %activation_row_limit = scf.select %wide_tile, %c255, %c127 : index + %activation_load_row = index.assume %activation_load_row0 [range(%activation_load_row0, 0, %activation_row_limit)] : index + %activation_packet = index.rem %workitem, %activation_row_packets : index + %activation_load_k = index.mul %activation_packet, %activation_packet_width : index + %second_token_tile_base = index.add %token_tile_origin, %c128 : index + %synchronous_a = scalar.xori %prefetch_a, %true : i1 + %k_origin = scf.if %synchronous_a -> (index) { + %block_k_base = index.mul %quant_block, %c256 : index + %absolute_group = index.add %weight_group_base, %quant_group : index + %group_k_add = index.mul %absolute_group, %c32 : index + %ordinary_k = index.add %block_k_base, %group_k_add : index + scf.yield %ordinary_k : index + } else { + scf.yield %c0 : index + } + scf.if %staged_input { + %token0 = index.add %token_tile_origin, %activation_load_row : index + %valid_token0 = index.cmp ult, %token0, %bounded_token_count : index + scf.if %larger_tile { + %loaded, %unused = func.call @ggml_f16_prefill_load_pair(%input_noalias, %bounded_token_count, %bounded_input_size, %token_tile_origin, %k_origin) : (buffer, index, index, index, index) -> (vector<8xf16>, vector<8xf16>) + vector.store %loaded, %activation_stage_physical_view[%activation_load_row, %activation_load_k] : vector<8xf16>, view<256x40xf16> + scf.yield + } else { + %activation_values0 = scf.if %prefetch_a -> (vector<16xf16>) { + %values = vector.concat<0> %current_a0, %current_a1 : vector<8xf16>, vector<8xf16> -> vector<16xf16> + scf.yield %values : vector<16xf16> + } else { + %loaded_values = scf.if %valid_token0 -> (vector<16xf16>) { + %bounded_token, %input_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %input_k = index.add %k_origin, %activation_load_k : index + %loaded = vector.load %input_f16_view[%bounded_token, %input_k] : view<[%bounded_token_count]x[%bounded_input_size]xf16> -> vector<16xf16> + scf.yield %loaded : vector<16xf16> + } else { + scf.yield %zero_f16 : vector<16xf16> + } + scf.yield %loaded_values : vector<16xf16> + } + vector.store %activation_values0, %activation_stage_physical_view[%activation_load_row, %activation_load_k] : vector<16xf16>, view<256x40xf16> + scf.yield + } + scf.if %not_wide_tile { + %token1 = index.add %second_token_tile_base, %activation_load_row : index + %valid_token1 = index.cmp ult, %token1, %bounded_token_count : index + %activation_values1 = scf.if %valid_token1 -> (vector<16xf16>) { + %bounded_token, %input_token_count = index.assume %token1, %bounded_token_count [lt(%token1, %bounded_token_count)] : index, index + %input_k = index.add %k_origin, %activation_load_k : index + %loaded = vector.load %input_f16_view[%bounded_token, %input_k] : view<[%bounded_token_count]x[%bounded_input_size]xf16> -> vector<16xf16> + scf.yield %loaded : vector<16xf16> + } else { + scf.yield %zero_f16 : vector<16xf16> + } + %activation_second_row = index.add %activation_load_row, %c128 : index + vector.store %activation_values1, %activation_stage_physical_view[%activation_second_row, %activation_load_k] : vector<16xf16>, view<256x40xf16> + scf.yield + } + kernel.barrier scope(workgroup) ordering(acq_rel) + } else { + %first_group = index.cmp eq, %quant_group, %c0 : index + %publish_stage = scalar.andi %first_group, %sync_weights : i1 + scf.if %publish_stage { + kernel.barrier scope(workgroup) ordering(acq_rel) + } + } + %next_a0, %next_a1 = scf.if %prefetch_a -> (vector<8xf16>, vector<8xf16>) { + %block_k_base = index.mul %quant_block, %c256 : index + %absolute_group = index.add %weight_group_base, %quant_group : index + %group_k_add = index.mul %absolute_group, %c32 : index + %wide_k = index.add %block_k_base, %group_k_add : index + %next_k = index.add %wide_k, %c32 : index + %has_next = index.cmp ult, %next_k, %bounded_input_size : index + %next_input_k = scf.select %has_next, %next_k, %wide_k : index + %a0, %a1 = func.call @ggml_f16_prefill_load_pair(%input_noalias, %bounded_token_count, %bounded_input_size, %token_tile_origin, %next_input_k) : (buffer, index, index, index, index) -> (vector<8xf16>, vector<8xf16>) + scf.yield %a0, %a1 : vector<8xf16>, vector<8xf16> + } else { + %zero = vector.constant 0.0 : vector<8xf16> + scf.yield %zero, %zero : vector<8xf16>, vector<8xf16> + } + %serial_rhs = scalar.ori %larger_tile, %packed_input : i1 + %next00, %next01, %next02, %next03, %next10, %next11, %next12, %next13 = scf.for %k_half = [%c0 to %c32 step %c16](%half_acc00 = %acc00 : vector<8xf32>, %half_acc01 = %acc01 : vector<8xf32>, %half_acc02 = %acc02 : vector<8xf32>, %half_acc03 = %acc03 : vector<8xf32>, %half_acc10 = %acc10 : vector<8xf32>, %half_acc11 = %acc11 : vector<8xf32>, %half_acc12 = %acc12 : vector<8xf32>, %half_acc13 = %acc13 : vector<8xf32>) -> (vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) unroll { + %lhs_m0 = index.add %subgroup_token_add, %c0 : index + %lhs_m1 = index.add %subgroup_token_add, %c16 : index + %lhs0 = scf.if %packed_input -> (vector<16xf16>) { + %token = index.add %token_tile_origin, %lhs_m0 : index + %global_k = index.add %k_origin, %k_half : index + %loaded = func.call @ggml_packed_f16_lhs_fragment_load(%input_noalias, %bounded_token_count, %bounded_input_size, %token, %global_k) : (buffer, index, index, index, index) -> (vector<16xf16>) + scf.yield %loaded : vector<16xf16> + } else { + %loaded = vector.fragment.load %activation_stage_physical_view[%lhs_m0, %k_half] shape [%m, %k] : view<256x40xf16> -> vector<16xf16> + scf.yield %loaded : vector<16xf16> + } + %lhs1 = scf.if %packed_input -> (vector<16xf16>) { + %token = index.add %token_tile_origin, %lhs_m1 : index + %global_k = index.add %k_origin, %k_half : index + %loaded = func.call @ggml_packed_f16_lhs_fragment_load(%input_noalias, %bounded_token_count, %bounded_input_size, %token, %global_k) : (buffer, index, index, index, index) -> (vector<16xf16>) + scf.yield %loaded : vector<16xf16> + } else { + %loaded = vector.fragment.load %activation_stage_physical_view[%lhs_m1, %k_half] shape [%m, %k] : view<256x40xf16> -> vector<16xf16> + scf.yield %loaded : vector<16xf16> + } + %weight_fragment_layout = encoding.layout.strided [1, %weight_stride] : encoding + %weight_fragment_view = buffer.view %weight_stage[%stage_offset] : buffer -> view<[%weight_window]x[%weight_rows]xf16, %weight_fragment_layout> + %weight_k_group = index.mul %quant_group, %c32 : index + %weight_half = index.add %weight_k_group, %k_half : index + %rhs0 = vector.fragment.load %weight_fragment_view[%weight_half, %rhs_column0] shape [%k, %n] : view<[%weight_window]x[%weight_rows]xf16, %weight_fragment_layout> -> vector<16xf16> + %early_rhs1 = scf.if %serial_rhs -> (vector<16xf16>) { + scf.yield %zero_f16 : vector<16xf16> + } else { + %loaded = vector.fragment.load %weight_fragment_view[%weight_half, %rhs_column1] shape [%k, %n] : view<[%weight_window]x[%weight_rows]xf16, %weight_fragment_layout> -> vector<16xf16> + scf.yield %loaded : vector<16xf16> + } + %half_next00 = vector.mma %lhs0, %rhs0, %half_acc00 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %half_next10 = vector.mma %lhs1, %rhs0, %half_acc10 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %rhs1 = scf.if %serial_rhs -> (vector<16xf16>) { + scf.schedule.fence + %loaded = vector.fragment.load %weight_fragment_view[%weight_half, %rhs_column1] shape [%k, %n] : view<[%weight_window]x[%weight_rows]xf16, %weight_fragment_layout> -> vector<16xf16> + scf.yield %loaded : vector<16xf16> + } else { + scf.yield %early_rhs1 : vector<16xf16> + } + %half_next01 = vector.mma %lhs0, %rhs1, %half_acc01 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %half_next11 = vector.mma %lhs1, %rhs1, %half_acc11 : vector<16xf16>, vector<16xf16>, vector<8xf32> + scf.schedule.fence + %rhs2 = vector.fragment.load %weight_fragment_view[%weight_half, %rhs_column2] shape [%k, %n] : view<[%weight_window]x[%weight_rows]xf16, %weight_fragment_layout> -> vector<16xf16> + %early_rhs3 = scf.if %serial_rhs -> (vector<16xf16>) { + scf.yield %zero_f16 : vector<16xf16> + } else { + %loaded = vector.fragment.load %weight_fragment_view[%weight_half, %rhs_column3] shape [%k, %n] : view<[%weight_window]x[%weight_rows]xf16, %weight_fragment_layout> -> vector<16xf16> + scf.yield %loaded : vector<16xf16> + } + %half_next02 = vector.mma %lhs0, %rhs2, %half_acc02 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %half_next12 = vector.mma %lhs1, %rhs2, %half_acc12 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %rhs3 = scf.if %serial_rhs -> (vector<16xf16>) { + scf.schedule.fence + %loaded = vector.fragment.load %weight_fragment_view[%weight_half, %rhs_column3] shape [%k, %n] : view<[%weight_window]x[%weight_rows]xf16, %weight_fragment_layout> -> vector<16xf16> + scf.yield %loaded : vector<16xf16> + } else { + scf.yield %early_rhs3 : vector<16xf16> + } + %half_next03 = vector.mma %lhs0, %rhs3, %half_acc03 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %half_next13 = vector.mma %lhs1, %rhs3, %half_acc13 : vector<16xf16>, vector<16xf16>, vector<8xf32> + scf.schedule.fence + scf.yield %half_next00, %half_next01, %half_next02, %half_next03, %half_next10, %half_next11, %half_next12, %half_next13 : vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32> + } + %last_group = index.cmp eq, %quant_group, %c1 : index + %release_stage = scalar.ori %staged_input, %last_group : i1 + %release_weights = scalar.andi %release_stage, %sync_weights : i1 + scf.if %release_weights { + kernel.barrier scope(workgroup) ordering(acq_rel) + } + func.return %next00, %next01, %next02, %next03, %next10, %next11, %next12, %next13, %next_a0, %next_a1 : vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf16>, vector<8xf16> +} + + +func.def inline @ggml_f16_prefill_publish_weights(%weight_format: index, %paired: i1, %packed_input: i1, %short_weight_packets: i1, %active: i1, %bounded_input_size: index, %bounded_output_size: index, %channel_tile: index, %weight_epoch: index, %weight_noalias: buffer, %up_weight: buffer, %weight_stage: buffer, %stage_offset: offset) { + %q6_format = index.constant 6 : index + %q6_weights = index.cmp eq, %weight_format, %q6_format : index + %layout_true = scalar.constant true : i1 + %staged_input = scalar.xori %packed_input, %layout_true : i1 + %workgroup_size = kernel.workgroup.size : index + %large_workgroup = index.constant 512 : index + %wide_tile = index.cmp uge, %workgroup_size, %large_workgroup : index + %can_prefetch_a = index.cmp eq, %workgroup_size, %large_workgroup : index + %prefetch_a = scalar.andi %can_prefetch_a, %staged_input : i1 + %largest_workgroup = index.constant 1024 : index + %wg1024 = index.cmp eq, %workgroup_size, %largest_workgroup : index + %larger_tile = scalar.andi %wg1024, %staged_input : i1 + %true = scalar.constant true : i1 + %workitem = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c32 = index.constant 32 : index + %c72 = index.constant 72 : index + %c136 = index.constant 136 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %weight_stride = scf.select %wide_tile, %c72, %c136 : index + %ordinary_weight_rows = scf.select %wide_tile, %c128, %c64 : index + %staged_weight_rows = scf.select %larger_tile, %c256, %ordinary_weight_rows : index + %packed_weight_rows = scf.select %wg1024, %c128, %c64 : index + %weight_rows = scf.select %packed_input, %packed_weight_rows, %staged_weight_rows : index + %zero_weight_values = vector.constant 0.0 : vector<16xf16> + %quant_block_count = index.div %bounded_input_size, %c256 : index + %epochs_per_block = scf.select %wide_tile, %c4, %c2 : index + %weight_stage_view = buffer.view %weight_stage[%stage_offset] : buffer -> view<[%weight_rows]x[%weight_stride]xf16> + %half_weight_rows = index.div %weight_rows, %c2 : index + %tile_channels = scf.select %paired, %half_weight_rows, %weight_rows : index + %channel_tile_base = index.mul %channel_tile, %tile_channels : index + %half_weight_producers = scalar.ori %prefetch_a, %packed_input : i1 + %quant_block = index.div %weight_epoch, %epochs_per_block : index + %weight_epoch_half = index.rem %weight_epoch, %epochs_per_block : index + %groups_per_epoch = scf.select %wide_tile, %c2, %c4 : index + %weight_group_base = index.mul %weight_epoch_half, %groups_per_epoch : index + %packed_threads_per_row = scf.select %short_weight_packets, %c4, %c2 : index + %packed_producer_count = index.mul %weight_rows, %packed_threads_per_row : index + %producer_count = scf.select %packed_input, %packed_producer_count, %c256 : index + %lower_workgroup_half = index.cmp ult, %workitem, %producer_count : index + %selected_weight_producer = scf.select %half_weight_producers, %lower_workgroup_half, %true : i1 + %weight_producer = scalar.andi %selected_weight_producer, %active : i1 + scf.if %weight_producer { + %load_row_packets = scf.select %short_weight_packets, %c4, %c2 : index + %load_packet_words = scf.select %short_weight_packets, %c2, %c4 : index + %load_half = index.rem %workitem, %load_row_packets : index + %load_packet = index.mul %load_half, %load_packet_words : index + %load_k = index.mul %load_packet, %c4 : index + %load_row0 = index.div %workitem, %load_row_packets : index + %load_row = index.assume %load_row0 [range(%load_row0, 0, 255)] : index + %staged_threads_per_row = scf.select %half_weight_producers, %c2, %c4 : index + %weight_threads_per_row = scf.select %packed_input, %packed_threads_per_row, %staged_threads_per_row : index + %weight_local_row0 = index.div %workitem, %weight_threads_per_row : index + %weight_row_limit = index.sub %weight_rows, %c1 : index + %weight_local_row = index.assume %weight_local_row0 [range(%weight_local_row0, 0, %weight_row_limit)] : index + %weight_half_row = index.rem %weight_local_row, %half_weight_rows : index + %weight_is_up = index.cmp uge, %weight_local_row, %half_weight_rows : index + %weight_pair = index.rem %load_row, %c2 : index + %weight_local_group0 = index.mul %weight_pair, %c2 : index + %weight_local_group = scf.select %wide_tile, %c0, %weight_local_group0 : index + %weight_group = index.add %weight_group_base, %weight_local_group : index + %weight_k_base = index.mul %weight_local_group, %c32 : index + %weight_k0 = index.add %weight_k_base, %load_k : index + %weight_k1 = index.add %weight_k0, %c32 : index + %channel_in_tile = scf.select %paired, %weight_half_row, %weight_local_row : index + %channel = index.add %channel_tile_base, %channel_in_tile : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %weight_values0, %weight_values1 = scf.if %valid_channel -> (vector<16xf16>, vector<16xf16>) { + %decoded0, %decoded1 = scf.if %q6_weights -> (vector<16xf16>, vector<16xf16>) { + %native_block_bytes = index.constant 210 : offset + %native_row_bytes = index.scale %quant_block_count, %native_block_bytes : index, offset -> offset + %row_byte_base = index.scale %channel, %native_row_bytes : index, offset -> offset + %decoded0, %decoded1 = func.call @ggml_q6k_f16_pair(%short_weight_packets, %weight_noalias, %row_byte_base, %quant_block, %weight_group, %load_packet) : (i1, buffer, offset, index, index, index) -> (vector<16xf16>, vector<16xf16>) + scf.yield %decoded0, %decoded1 : vector<16xf16>, vector<16xf16> + } else { + %packed_record_bytes = index.constant 16 : offset + %packed_block_bytes = index.constant 9216 : offset + %packed_row_group = index.div %channel, %c64 : index + %packed_row_lane = index.rem %channel, %c64 : index + %packed_tile_block0 = index.mul %packed_row_group, %quant_block_count : index + %packed_tile_block = index.add %packed_tile_block0, %quant_block : index + %packed_tile_base = index.scale %packed_tile_block, %packed_block_bytes : index, offset -> offset + %packed_row_add = index.scale %packed_row_lane, %packed_record_bytes : index, offset -> offset + %row_byte_base = index.add %packed_tile_base, %packed_row_add : offset + %decoded0, %decoded1 = func.call @ggml_q4k_native_row64_f16_pair(%short_weight_packets, %paired, %weight_noalias, %up_weight, %weight_is_up, %row_byte_base, %weight_group, %load_packet) : (i1, i1, buffer, buffer, i1, offset, index, index) -> (vector<16xf16>, vector<16xf16>) + scf.yield %decoded0, %decoded1 : vector<16xf16>, vector<16xf16> + } + scf.yield %decoded0, %decoded1 : vector<16xf16>, vector<16xf16> + } else { + scf.yield %zero_weight_values, %zero_weight_values : vector<16xf16>, vector<16xf16> + } + scf.if %short_weight_packets { + %half0 = vector.slice %weight_values0[0] : vector<16xf16> -> vector<8xf16> + %half1 = vector.slice %weight_values1[0] : vector<16xf16> -> vector<8xf16> + vector.store %half0, %weight_stage_view[%weight_local_row, %weight_k0] : vector<8xf16>, view<[%weight_rows]x[%weight_stride]xf16> + vector.store %half1, %weight_stage_view[%weight_local_row, %weight_k1] : vector<8xf16>, view<[%weight_rows]x[%weight_stride]xf16> + scf.yield + } else { + vector.store %weight_values0, %weight_stage_view[%weight_local_row, %weight_k0] : vector<16xf16>, view<[%weight_rows]x[%weight_stride]xf16> + vector.store %weight_values1, %weight_stage_view[%weight_local_row, %weight_k1] : vector<16xf16>, view<[%weight_rows]x[%weight_stride]xf16> + scf.yield + } + scf.yield + } + func.return +} + +func.def inline @ggml_prefill_conv4_strip(%deferred_state: i1, %width: index, %channel_base: index, %stage: buffer, %state: buffer, %filter: buffer, %output: buffer, %cache: buffer, %first: vector<8xf32>, %second: vector<8xf32>) { + %zero = index.constant 0 : index + %one = index.constant 1 : index + %two = index.constant 2 : index + %three = index.constant 3 : index + %six = index.constant 6 : index + %zero_values = vector.constant 0.0 : vector<2xf32> + %cache_rows = scf.select %deferred_state, %six, %three : index + %eight = index.constant 8 : index + %sixteen = index.constant 16 : index + %thirtytwo = index.constant 32 : index + %sixtyfour = index.constant 64 : index + %five09 = index.constant 509 : index + %five12 = index.constant 512 : index + %base = index.constant 0 : offset + %filter_bytes = index.constant 16 : offset + %tid0 = kernel.workitem.id : index + %tid = index.assume %tid0 [range(%tid0, 0, 511)] : index + %wave0 = kernel.subgroup.id : index + %wave = index.assume %wave0 [range(%wave0, 0, 15)] : index + %local_row = index.mul %wave, %thirtytwo : index + %second_row = index.add %local_row, %sixteen : index + %column_lane = index.rem %tid, %eight : index + %local_column = index.mul %column_lane, %two : index + %channel0 = index.add %channel_base, %local_column : index + %channel_last = index.sub %width, %two : index + %channel = index.assume %channel0 [range(%channel0, 0, %channel_last), mul(%channel0, 2)] : index + %channel1 = index.add %channel, %one : index + %thread_row = index.div %tid, %eight : index + %tile = buffer.view %stage[%base] : buffer -> view<512x16xf32> + %output_view = buffer.view %output[%base] : buffer -> view<512x[%width]xf32> + %state_view = buffer.view %state[%base] : buffer -> view<[%width]x3xf32> + %cache_view = buffer.view %cache[%base] : buffer -> view<[%width]x[%cache_rows]xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + vector.fragment.store %first, %tile[%local_row, %zero] shape [%sixteen, %sixteen] : vector<8xf32>, view<512x16xf32> + vector.fragment.store %second, %tile[%second_row, %zero] shape [%sixteen, %sixteen] : vector<8xf32>, view<512x16xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + %filter_offset = index.scale %channel, %filter_bytes : index, offset -> offset + %filter_view = buffer.view %filter[%filter_offset] : buffer -> view<2x4xf32> + %weights0 = vector.load %filter_view[%zero, %zero] : view<2x4xf32> -> vector<2x4xf32> + %weights = vector.transpose<[1, 0]> %weights0 : vector<2x4xf32> -> vector<4x2xf32> + %w0 = vector.extract %weights[0] : vector<4x2xf32> -> vector<2xf32> + %w1 = vector.extract %weights[1] : vector<4x2xf32> -> vector<2xf32> + %w2 = vector.extract %weights[2] : vector<4x2xf32> -> vector<2xf32> + %w3 = vector.extract %weights[3] : vector<4x2xf32> -> vector<2xf32> + scf.for %row_block = [%zero to %five12 step %sixtyfour] unroll { + %row = index.add %thread_row, %row_block : index + %history0 = index.cmp ult, %row, %three : index + %x0 = scf.if %history0 -> (vector<2xf32>) { + %history_values = scf.if %deferred_state -> (vector<2xf32>) { + scf.yield %zero_values : vector<2xf32> + } else { + %history_row0_0 = index.add %row, %zero : index + %history_row0 = index.assume %history_row0_0 [range(%history_row0_0, 0, 2)] : index + %h0_0 = view.load %state_view[%channel, %history_row0] : view<[%width]x3xf32> -> f32 + %h0_1 = view.load %state_view[%channel1, %history_row0] : view<[%width]x3xf32> -> f32 + %hv0_0 = vector.splat %h0_0 : vector<1xf32> + %hv0_1 = vector.splat %h0_1 : vector<1xf32> + %hv0 = vector.concat<0> %hv0_0, %hv0_1 : vector<1xf32>, vector<1xf32> -> vector<2xf32> + scf.yield %hv0 : vector<2xf32> + } + scf.yield %history_values : vector<2xf32> + } else { + %r0_0 = index.sub %row, %three : index + %r0 = index.assume %r0_0 [range(%r0_0, 0, 510)] : index + %local0 = vector.load %tile[%r0, %local_column] : view<512x16xf32> -> vector<2xf32> + scf.yield %local0 : vector<2xf32> + } + %history1 = index.cmp ult, %row, %two : index + %x1 = scf.if %history1 -> (vector<2xf32>) { + %history_values = scf.if %deferred_state -> (vector<2xf32>) { + scf.yield %zero_values : vector<2xf32> + } else { + %history_row1_0 = index.add %row, %one : index + %history_row1 = index.assume %history_row1_0 [range(%history_row1_0, 0, 2)] : index + %h1_0 = view.load %state_view[%channel, %history_row1] : view<[%width]x3xf32> -> f32 + %h1_1 = view.load %state_view[%channel1, %history_row1] : view<[%width]x3xf32> -> f32 + %hv1_0 = vector.splat %h1_0 : vector<1xf32> + %hv1_1 = vector.splat %h1_1 : vector<1xf32> + %hv1 = vector.concat<0> %hv1_0, %hv1_1 : vector<1xf32>, vector<1xf32> -> vector<2xf32> + scf.yield %hv1 : vector<2xf32> + } + scf.yield %history_values : vector<2xf32> + } else { + %r1_0 = index.sub %row, %two : index + %r1 = index.assume %r1_0 [range(%r1_0, 0, 510)] : index + %local1 = vector.load %tile[%r1, %local_column] : view<512x16xf32> -> vector<2xf32> + scf.yield %local1 : vector<2xf32> + } + %history2 = index.cmp ult, %row, %one : index + %x2 = scf.if %history2 -> (vector<2xf32>) { + %history_values = scf.if %deferred_state -> (vector<2xf32>) { + scf.yield %zero_values : vector<2xf32> + } else { + %history_row2_0 = index.add %row, %two : index + %history_row2 = index.assume %history_row2_0 [range(%history_row2_0, 0, 2)] : index + %h2_0 = view.load %state_view[%channel, %history_row2] : view<[%width]x3xf32> -> f32 + %h2_1 = view.load %state_view[%channel1, %history_row2] : view<[%width]x3xf32> -> f32 + %hv2_0 = vector.splat %h2_0 : vector<1xf32> + %hv2_1 = vector.splat %h2_1 : vector<1xf32> + %hv2 = vector.concat<0> %hv2_0, %hv2_1 : vector<1xf32>, vector<1xf32> -> vector<2xf32> + scf.yield %hv2 : vector<2xf32> + } + scf.yield %history_values : vector<2xf32> + } else { + %r2_0 = index.sub %row, %one : index + %r2 = index.assume %r2_0 [range(%r2_0, 0, 510)] : index + %local2 = vector.load %tile[%r2, %local_column] : view<512x16xf32> -> vector<2xf32> + scf.yield %local2 : vector<2xf32> + } + %x3 = vector.load %tile[%row, %local_column] : view<512x16xf32> -> vector<2xf32> + %a0 = vector.mulf %x0, %w0 : vector<2xf32> + %a1 = vector.fmaf %x1, %w1, %a0 : vector<2xf32> + %a2 = vector.fmaf %x2, %w2, %a1 : vector<2xf32> + %a3 = vector.fmaf %x3, %w3, %a2 : vector<2xf32> + %activated = vector.siluf %a3 : vector<2xf32> + vector.store %activated, %output_view[%row, %channel] : vector<2xf32>, view<512x[%width]xf32> + scf.if %deferred_state { + %first_worker = index.cmp ult, %row, %three : index + scf.if %first_worker { + %edge_row = index.assume %row [range(%row, 0, 2)] : index + %edge0 = vector.extract %x3[0] : vector<2xf32> -> f32 + %edge1 = vector.extract %x3[1] : vector<2xf32> -> f32 + view.store %edge0, %cache_view[%channel, %edge_row] : f32, view<[%width]x[%cache_rows]xf32> + view.store %edge1, %cache_view[%channel1, %edge_row] : f32, view<[%width]x[%cache_rows]xf32> + } + } + %cache_worker = index.cmp uge, %row, %five09 : index + scf.if %cache_worker { + %cache_start = scf.select %deferred_state, %three, %zero : index + %tail_row = index.sub %row, %five09 : index + %cache_row0 = index.add %tail_row, %cache_start : index + %cache_row, %cache_size = index.assume %cache_row0, %cache_rows [lt(%cache_row0, %cache_rows)] : index, index + %cache_v0 = vector.extract %x3[0] : vector<2xf32> -> f32 + %cache_v1 = vector.extract %x3[1] : vector<2xf32> -> f32 + view.store %cache_v0, %cache_view[%channel, %cache_row] : f32, view<[%width]x[%cache_rows]xf32> + view.store %cache_v1, %cache_view[%channel1, %cache_row] : f32, view<[%width]x[%cache_rows]xf32> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + func.return +} + +func.def inline @ggml_mul_mat_quantized_f16_prefill_wave32_body(%weight_format: index, %binary_op: index, %paired: i1, %packed_input: i1, %packed_output: i1, %token_count: index, %input_size0: index, %output_size0: index, %channel_tile: index, %token_tile: index, %input: buffer, %weight: buffer, %up_weight: buffer, %output: buffer, %f16_output: buffer, %conv_epilogue: i1, %conv_state: buffer, %conv_filter: buffer, %conv_output: buffer, %conv_cache: buffer, %deferred_state: i1) { + %q6_format = index.constant 6 : index + %q6_weights = index.cmp eq, %weight_format, %q6_format : index + %layout_true = scalar.constant true : i1 + %unpaired = scalar.xori %paired, %layout_true : i1 + %packed_single = scalar.andi %unpaired, %packed_input : i1 + %staged_input = scalar.xori %packed_input, %layout_true : i1 + %workgroup_size = kernel.workgroup.size : index + %large_workgroup = index.constant 512 : index + %wide_tile = index.cmp uge, %workgroup_size, %large_workgroup : index + %can_prefetch_a = index.cmp eq, %workgroup_size, %large_workgroup : index + %prefetch_a = scalar.andi %can_prefetch_a, %staged_input : i1 + %largest_workgroup = index.constant 1024 : index + %wg1024 = index.cmp eq, %workgroup_size, %largest_workgroup : index + %larger_tile = scalar.andi %wg1024, %staged_input : i1 + %pipeline_weights = scalar.andi %packed_input, %can_prefetch_a : i1 + %default_short_packets = scalar.ori %larger_tile, %packed_single : i1 + %short_weight_packets = scalar.ori %default_short_packets, %pipeline_weights : i1 + %true = scalar.constant true : i1 + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_input_size = index.assume %input_size0 [range(%input_size0, 256, 32768), mul(%input_size0, 256)] : index + %bounded_output_size = index.assume %output_size0 [range(%output_size0, 1, 262144)] : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup_small = index.constant 7 : index + %subgroup_large = index.constant 15 : index + %subgroup_larger = index.constant 31 : index + %ordinary_subgroup_limit = scf.select %wide_tile, %subgroup_large, %subgroup_small : index + %subgroup_limit = scf.select %wg1024, %subgroup_larger, %ordinary_subgroup_limit : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, %subgroup_limit)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c72 = index.constant 72 : index + %c136 = index.constant 136 : index + %c48 = index.constant 48 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %c512 = index.constant 512 : index + %c0_offset = index.constant 0 : offset + %weight_stride = scf.select %wide_tile, %c72, %c136 : index + %ordinary_weight_rows = scf.select %wide_tile, %c128, %c64 : index + %staged_weight_rows = scf.select %larger_tile, %c256, %ordinary_weight_rows : index + %packed_weight_rows = scf.select %wg1024, %c128, %c64 : index + %weight_rows = scf.select %packed_input, %packed_weight_rows, %staged_weight_rows : index + %weight_elements = index.mul %weight_rows, %weight_stride : index + %element_bytes = index.constant 2 : offset + %weight_stage_bytes = index.scale %weight_elements, %element_bytes : index, offset -> offset + %activation_stage_bytes = index.constant 20480 : offset + %zero_weight_values = vector.constant 0.0 : vector<16xf16> + %zero_accumulator = vector.constant 0.0 : vector<8xf32> + %m = index.constant 16 : index + %n = index.constant 16 : index + %quant_block_count = index.div %bounded_input_size, %c256 : index + %epochs_per_block = scf.select %wide_tile, %c4, %c2 : index + %weight_epoch_count = index.mul %quant_block_count, %epochs_per_block : index + %input_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input, %weight, %output : buffer, buffer, buffer + %slot_count = scf.select %pipeline_weights, %c2, %c1 : index + %weight_ring_bytes = index.scale %slot_count, %weight_stage_bytes : index, offset -> offset + %conv_stage_bytes = index.constant 32768 : offset + %stage_bytes = scf.select %conv_epilogue, %conv_stage_bytes, %weight_ring_bytes : offset + %weight_stage = buffer.alloca align(16) %stage_bytes : buffer + %activation_stage = buffer.alloca align(16) %activation_stage_bytes : buffer + %weight_stage_view = buffer.view %weight_stage[%c0_offset] : buffer -> view<[%weight_rows]x[%weight_stride]xf16> + %half_weight_rows = index.div %weight_rows, %c2 : index + %tile_channels = scf.select %paired, %half_weight_rows, %weight_rows : index + %channel_tile_base = index.mul %channel_tile, %tile_channels : index + %token_tile_size = scf.select %packed_input, %c512, %c256 : index + %token_tile_origin = index.mul %token_tile, %token_tile_size : index + %half_weight_producers = scalar.ori %prefetch_a, %packed_input : i1 + %init00 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + %init01 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + %init02 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + %init03 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + %init10 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + %init11 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + %init12 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + %init13 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + %zero_prefetch = vector.constant 0.0 : vector<8xf16> + %initial_a0, %initial_a1 = scf.if %prefetch_a -> (vector<8xf16>, vector<8xf16>) { + %a0, %a1 = func.call @ggml_f16_prefill_load_pair(%input_noalias, %bounded_token_count, %bounded_input_size, %token_tile_origin, %c0) : (buffer, index, index, index, index) -> (vector<8xf16>, vector<8xf16>) + scf.yield %a0, %a1 : vector<8xf16>, vector<8xf16> + } else { + scf.yield %zero_prefetch, %zero_prefetch : vector<8xf16>, vector<8xf16> + } + scf.if %pipeline_weights { + func.call @ggml_f16_prefill_publish_weights(%weight_format, %paired, %packed_input, %short_weight_packets, %true, %bounded_input_size, %bounded_output_size, %channel_tile, %c0, %weight_noalias, %up_weight, %weight_stage, %c0_offset) : (index, i1, i1, i1, i1, index, index, index, index, buffer, buffer, buffer, offset) + kernel.barrier scope(workgroup) ordering(acq_rel) + } + %result00, %result01, %result02, %result03, %result10, %result11, %result12, %result13, %final_a0, %final_a1 = scf.for %weight_epoch = [%c0 to %weight_epoch_count step %c1](%block_acc00 = %init00 : vector<8xf32>, %block_acc01 = %init01 : vector<8xf32>, %block_acc02 = %init02 : vector<8xf32>, %block_acc03 = %init03 : vector<8xf32>, %block_acc10 = %init10 : vector<8xf32>, %block_acc11 = %init11 : vector<8xf32>, %block_acc12 = %init12 : vector<8xf32>, %block_acc13 = %init13 : vector<8xf32>, %block_a0 = %initial_a0 : vector<8xf16>, %block_a1 = %initial_a1 : vector<8xf16>) -> (vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf16>, vector<8xf16>) { + %quant_block = index.div %weight_epoch, %epochs_per_block : index + %weight_epoch_half = index.rem %weight_epoch, %epochs_per_block : index + %groups_per_epoch = scf.select %wide_tile, %c2, %c4 : index + %weight_group_base = index.mul %weight_epoch_half, %groups_per_epoch : index + %slot = index.rem %weight_epoch, %c2 : index + %slot_offset = index.scale %slot, %weight_stage_bytes : index, offset -> offset + %stage_offset = scf.select %pipeline_weights, %slot_offset, %c0_offset : offset + %next_epoch = index.add %weight_epoch, %c1 : index + %load_epoch = scf.select %pipeline_weights, %next_epoch, %weight_epoch : index + %has_load = index.cmp ult, %load_epoch, %weight_epoch_count : index + %load_slot = index.rem %load_epoch, %c2 : index + %load_slot_offset = index.scale %load_slot, %weight_stage_bytes : index, offset -> offset + %load_offset = scf.select %pipeline_weights, %load_slot_offset, %c0_offset : offset + func.call @ggml_f16_prefill_publish_weights(%weight_format, %paired, %packed_input, %short_weight_packets, %has_load, %bounded_input_size, %bounded_output_size, %channel_tile, %load_epoch, %weight_noalias, %up_weight, %weight_stage, %load_offset) : (index, i1, i1, i1, i1, index, index, index, index, buffer, buffer, buffer, offset) + %not_wide_tile = scalar.xori %wide_tile, %true : i1 + // First activation publication also publishes the wide weight tile. + scf.if %not_wide_tile { + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield + } + %block_result00, %block_result01, %block_result02, %block_result03, %block_result10, %block_result11, %block_result12, %block_result13, %block_next_a0, %block_next_a1 = scf.if %wide_tile -> (vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf16>, vector<8xf16>) { + %first00, %first01, %first02, %first03, %first10, %first11, %first12, %first13, %first_a0, %first_a1 = func.call @ggml_f16_prefill_stage_mma(%paired, %packed_input, %bounded_token_count, %bounded_input_size, %token_tile_origin, %quant_block, %weight_group_base, %c0, %input_noalias, %weight_stage, %activation_stage, %stage_offset, %pipeline_weights, %block_acc00, %block_acc01, %block_acc02, %block_acc03, %block_acc10, %block_acc11, %block_acc12, %block_acc13, %block_a0, %block_a1) : (i1, i1, index, index, index, index, index, index, buffer, buffer, buffer, offset, i1, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf16>, vector<8xf16>) -> (vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf16>, vector<8xf16>) + %second00, %second01, %second02, %second03, %second10, %second11, %second12, %second13, %second_a0, %second_a1 = func.call @ggml_f16_prefill_stage_mma(%paired, %packed_input, %bounded_token_count, %bounded_input_size, %token_tile_origin, %quant_block, %weight_group_base, %c1, %input_noalias, %weight_stage, %activation_stage, %stage_offset, %pipeline_weights, %first00, %first01, %first02, %first03, %first10, %first11, %first12, %first13, %first_a0, %first_a1) : (i1, i1, index, index, index, index, index, index, buffer, buffer, buffer, offset, i1, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf16>, vector<8xf16>) -> (vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf16>, vector<8xf16>) + scf.yield %second00, %second01, %second02, %second03, %second10, %second11, %second12, %second13, %second_a0, %second_a1 : vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf16>, vector<8xf16> + } else { + %ordinary00, %ordinary01, %ordinary02, %ordinary03, %ordinary10, %ordinary11, %ordinary12, %ordinary13 = scf.for %quant_group = [%c0 to %c4 step %c1](%acc00 = %block_acc00 : vector<8xf32>, %acc01 = %block_acc01 : vector<8xf32>, %acc02 = %block_acc02 : vector<8xf32>, %acc03 = %block_acc03 : vector<8xf32>, %acc10 = %block_acc10 : vector<8xf32>, %acc11 = %block_acc11 : vector<8xf32>, %acc12 = %block_acc12 : vector<8xf32>, %acc13 = %block_acc13 : vector<8xf32>) -> (vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) { + %next00, %next01, %next02, %next03, %next10, %next11, %next12, %next13, %unused_a0, %unused_a1 = func.call @ggml_f16_prefill_stage_mma(%paired, %packed_input, %bounded_token_count, %bounded_input_size, %token_tile_origin, %quant_block, %weight_group_base, %quant_group, %input_noalias, %weight_stage, %activation_stage, %stage_offset, %pipeline_weights, %acc00, %acc01, %acc02, %acc03, %acc10, %acc11, %acc12, %acc13, %zero_prefetch, %zero_prefetch) : (i1, i1, index, index, index, index, index, index, buffer, buffer, buffer, offset, i1, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf16>, vector<8xf16>) -> (vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf16>, vector<8xf16>) + scf.yield %next00, %next01, %next02, %next03, %next10, %next11, %next12, %next13 : vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32> + } + scf.yield %ordinary00, %ordinary01, %ordinary02, %ordinary03, %ordinary10, %ordinary11, %ordinary12, %ordinary13, %zero_prefetch, %zero_prefetch : vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf16>, vector<8xf16> + } + scf.if %pipeline_weights { + kernel.barrier scope(workgroup) ordering(acq_rel) + } + scf.yield %block_result00, %block_result01, %block_result02, %block_result03, %block_result10, %block_result11, %block_result12, %block_result13, %block_next_a0, %block_next_a1 : vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf16>, vector<8xf16> + } + %token_subgroups = scf.select %packed_input, %c16, %c8 : index + %token_subgroup = index.rem %subgroup, %token_subgroups : index + %subgroup_token_add = index.mul %token_subgroup, %c32 : index + %token_tile_base = index.add %token_tile_origin, %subgroup_token_add : index + %channel_subgroup = index.div %subgroup, %token_subgroups : index + %subgroup_channel_stride = scf.select %paired, %c32, %c64 : index + %subgroup_channel_add = index.mul %channel_subgroup, %subgroup_channel_stride : index + %channel_output_base = index.add %channel_tile_base, %subgroup_channel_add : index + scf.if %conv_epilogue { + %conv_channel_0 = index.add %channel_output_base, %c0 : index + func.call @ggml_prefill_conv4_strip(%deferred_state, %bounded_output_size, %conv_channel_0, %weight_stage, %conv_state, %conv_filter, %conv_output, %conv_cache, %result00, %result10) : (i1, index, index, buffer, buffer, buffer, buffer, buffer, vector<8xf32>, vector<8xf32>) + %conv_channel_1 = index.add %channel_output_base, %c16 : index + func.call @ggml_prefill_conv4_strip(%deferred_state, %bounded_output_size, %conv_channel_1, %weight_stage, %conv_state, %conv_filter, %conv_output, %conv_cache, %result01, %result11) : (i1, index, index, buffer, buffer, buffer, buffer, buffer, vector<8xf32>, vector<8xf32>) + %conv_channel_2 = index.add %channel_output_base, %c32 : index + func.call @ggml_prefill_conv4_strip(%deferred_state, %bounded_output_size, %conv_channel_2, %weight_stage, %conv_state, %conv_filter, %conv_output, %conv_cache, %result02, %result12) : (i1, index, index, buffer, buffer, buffer, buffer, buffer, vector<8xf32>, vector<8xf32>) + %conv_channel_3 = index.add %channel_output_base, %c48 : index + func.call @ggml_prefill_conv4_strip(%deferred_state, %bounded_output_size, %conv_channel_3, %weight_stage, %conv_state, %conv_filter, %conv_output, %conv_cache, %result03, %result13) : (i1, index, index, buffer, buffer, buffer, buffer, buffer, vector<8xf32>, vector<8xf32>) + } else { + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_output_size]xf32> + %f16_columns = scf.select %packed_output, %c16, %bounded_output_size : index + %f16_group_elements = index.mul %channel_output_base, %bounded_token_count : index + %f16_group_bytes = index.scale %f16_group_elements, %element_bytes : index, offset -> offset + %f16_next_elements = index.mul %c16, %bounded_token_count : index + %f16_next_bytes = index.scale %f16_next_elements, %element_bytes : index, offset -> offset + %f16_next_group_bytes = index.add %f16_group_bytes, %f16_next_bytes : offset + %f16_base0 = scf.select %packed_output, %f16_group_bytes, %c0_offset : offset + %f16_base1 = scf.select %packed_output, %f16_next_group_bytes, %c0_offset : offset + %f16_output_view0 = buffer.view %f16_output[%f16_base0] : buffer -> view<[%bounded_token_count]x[%f16_columns]xf16> + %f16_output_view1 = buffer.view %f16_output[%f16_base1] : buffer -> view<[%bounded_token_count]x[%f16_columns]xf16> + scf.if %paired { + %out_m00 = index.add %token_tile_base, %c0 : index + %out_n00 = index.add %channel_output_base, %c0 : index + %fused00 = template.apply<@ggml.binary_f32.apply_vector8>(%binary_op, %result00, %result02) : (index, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + vector.fragment.store %fused00, %output_view[%out_m00, %out_n00] shape [%m, %n] : vector<8xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32> + %narrow00 = vector.fptrunc %fused00 : vector<8xf32> to vector<8xf16> + %f16_n00 = scf.select %packed_output, %c0, %out_n00 : index + vector.fragment.store %narrow00, %f16_output_view0[%out_m00, %f16_n00] shape [%m, %n] : vector<8xf16>, view<[%bounded_token_count]x[%f16_columns]xf16> + scf.schedule.fence + %out_m01 = index.add %token_tile_base, %c0 : index + %out_n01 = index.add %channel_output_base, %c16 : index + %fused01 = template.apply<@ggml.binary_f32.apply_vector8>(%binary_op, %result01, %result03) : (index, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + vector.fragment.store %fused01, %output_view[%out_m01, %out_n01] shape [%m, %n] : vector<8xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32> + %narrow01 = vector.fptrunc %fused01 : vector<8xf32> to vector<8xf16> + %f16_n01 = scf.select %packed_output, %c0, %out_n01 : index + vector.fragment.store %narrow01, %f16_output_view1[%out_m01, %f16_n01] shape [%m, %n] : vector<8xf16>, view<[%bounded_token_count]x[%f16_columns]xf16> + scf.schedule.fence + %out_m10 = index.add %token_tile_base, %c16 : index + %out_n10 = index.add %channel_output_base, %c0 : index + %fused10 = template.apply<@ggml.binary_f32.apply_vector8>(%binary_op, %result10, %result12) : (index, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + vector.fragment.store %fused10, %output_view[%out_m10, %out_n10] shape [%m, %n] : vector<8xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32> + %narrow10 = vector.fptrunc %fused10 : vector<8xf32> to vector<8xf16> + %f16_n10 = scf.select %packed_output, %c0, %out_n10 : index + vector.fragment.store %narrow10, %f16_output_view0[%out_m10, %f16_n10] shape [%m, %n] : vector<8xf16>, view<[%bounded_token_count]x[%f16_columns]xf16> + scf.schedule.fence + %out_m11 = index.add %token_tile_base, %c16 : index + %out_n11 = index.add %channel_output_base, %c16 : index + %fused11 = template.apply<@ggml.binary_f32.apply_vector8>(%binary_op, %result11, %result13) : (index, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + vector.fragment.store %fused11, %output_view[%out_m11, %out_n11] shape [%m, %n] : vector<8xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32> + %narrow11 = vector.fptrunc %fused11 : vector<8xf32> to vector<8xf16> + %f16_n11 = scf.select %packed_output, %c0, %out_n11 : index + vector.fragment.store %narrow11, %f16_output_view1[%out_m11, %f16_n11] shape [%m, %n] : vector<8xf16>, view<[%bounded_token_count]x[%f16_columns]xf16> + scf.schedule.fence + } else { + %out_m00 = index.add %token_tile_base, %c0 : index + %out_n00 = index.add %channel_output_base, %c0 : index + vector.fragment.store %result00, %output_view[%out_m00, %out_n00] shape [%m, %n] : vector<8xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32> + %out_m01 = index.add %token_tile_base, %c0 : index + %out_n01 = index.add %channel_output_base, %c16 : index + vector.fragment.store %result01, %output_view[%out_m01, %out_n01] shape [%m, %n] : vector<8xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32> + %out_m02 = index.add %token_tile_base, %c0 : index + %out_n02 = index.add %channel_output_base, %c32 : index + vector.fragment.store %result02, %output_view[%out_m02, %out_n02] shape [%m, %n] : vector<8xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32> + %out_m03 = index.add %token_tile_base, %c0 : index + %out_n03 = index.add %channel_output_base, %c48 : index + vector.fragment.store %result03, %output_view[%out_m03, %out_n03] shape [%m, %n] : vector<8xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32> + %out_m10 = index.add %token_tile_base, %c16 : index + %out_n10 = index.add %channel_output_base, %c0 : index + vector.fragment.store %result10, %output_view[%out_m10, %out_n10] shape [%m, %n] : vector<8xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32> + %out_m11 = index.add %token_tile_base, %c16 : index + %out_n11 = index.add %channel_output_base, %c16 : index + vector.fragment.store %result11, %output_view[%out_m11, %out_n11] shape [%m, %n] : vector<8xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32> + %out_m12 = index.add %token_tile_base, %c16 : index + %out_n12 = index.add %channel_output_base, %c32 : index + vector.fragment.store %result12, %output_view[%out_m12, %out_n12] shape [%m, %n] : vector<8xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32> + %out_m13 = index.add %token_tile_base, %c16 : index + %out_n13 = index.add %channel_output_base, %c48 : index + vector.fragment.store %result13, %output_view[%out_m13, %out_n13] shape [%m, %n] : vector<8xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32> + } + } + func.return +} + +func.def public inline @ggml_mul_mat_quantized_f16_prefill_wave32(%weight_format: index, %binary_op: index, %paired: i1, %packed_input: i1, %packed_output: i1, %token_count: index, %input_size0: index, %output_size0: index, %channel_tile: index, %token_tile: index, %input: buffer, %weight: buffer, %up_weight: buffer, %output: buffer, %f16_output: buffer) { + %no_conv = scalar.constant false : i1 + func.call @ggml_mul_mat_quantized_f16_prefill_wave32_body(%weight_format, %binary_op, %paired, %packed_input, %packed_output, %token_count, %input_size0, %output_size0, %channel_tile, %token_tile, %input, %weight, %up_weight, %output, %f16_output, %no_conv, %output, %output, %output, %output, %no_conv) : (index, index, i1, i1, i1, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer, i1, buffer, buffer, buffer, buffer, i1) + func.return +} + +func.def public inline @ggml_mul_mat_quantized_f16_prefill_conv4(%weight_format: index, %token_count: index, %input_size: index, %output_size: index, %channel_tile: index, %input: buffer, %weight: buffer, %state: buffer, %filter: buffer, %output: buffer, %cache: buffer) { + %zero = index.constant 0 : index + %false = scalar.constant false : i1 + %true = scalar.constant true : i1 + func.call @ggml_mul_mat_quantized_f16_prefill_wave32_body(%weight_format, %zero, %false, %true, %false, %token_count, %input_size, %output_size, %channel_tile, %zero, %input, %weight, %weight, %output, %output, %true, %state, %filter, %output, %cache, %false) : (index, index, i1, i1, i1, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer, i1, buffer, buffer, buffer, buffer, i1) + func.return +} + +func.def public inline @ggml_mul_mat_quantized_f16_prefill_conv4_interior(%weight_format: index, %token_count: index, %input_size: index, %output_size: index, %channel_tile: index, %input: buffer, %weight: buffer, %filter: buffer, %output: buffer, %edges: buffer) { + %zero = index.constant 0 : index + %false = scalar.constant false : i1 + %true = scalar.constant true : i1 + func.call @ggml_mul_mat_quantized_f16_prefill_wave32_body(%weight_format, %zero, %false, %true, %false, %token_count, %input_size, %output_size, %channel_tile, %zero, %input, %weight, %weight, %output, %output, %true, %edges, %filter, %output, %edges, %true) : (index, index, i1, i1, i1, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer, i1, buffer, buffer, buffer, buffer, i1) + func.return +} + +// Exact K16-major F16 activation tiles used by larger token-reuse GEMMs. +func.def inline @ggml_packed_f16_lhs_fragment_load(%input: buffer, %tokens: index, %input_size: index, %token: index, %k_origin: index) -> (vector<16xf16>) { + %zero = index.constant 0 : offset + %c0 = index.constant 0 : index + %c16 = index.constant 16 : index + %k_tiles = index.div %input_size, %c16 : index + %k_tile = index.div %k_origin, %c16 : index + %packed_row = index.madd %k_tile, %tokens, %token : index + %packed_rows = index.mul %tokens, %k_tiles : index + %packed_view = buffer.view %input[%zero] : buffer -> view<[%packed_rows]x16xf16> + %values = vector.fragment.load %packed_view[%packed_row, %c0] shape [%c16, %c16] : view<[%packed_rows]x16xf16> -> vector<16xf16> + func.return %values : vector<16xf16> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/publish_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/publish_f32.loom new file mode 100644 index 000000000000..006a83106659 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/publish_f32.loom @@ -0,0 +1,81 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +template.decl @ggml.publish_f32.f16_vector4(%publish_word: i1, %token: index, %channel: index, %token_count: index, %hidden_size: index, %values: vector<4xf32>, %output: buffer) + +template.decl @ggml.publish_f32.f32_vector4(%publish_word: i1, %token: index, %channel: index, %token_count: index, %hidden_size: index, %values: vector<4xf32>, %output: buffer) + +template.decl @ggml.publish_f32.next_vector4(%format: index, %publish_word: i1, %token: index, %channel: index, %token_count: index, %hidden_size: index, %values: vector<4xf32>, %scratch_values: view<256xf32>, %scratch_d: view<32xf32>, %output: buffer) + +template.decl @ggml.publish_f32.q8_1_x4_vector4(%publish_word: i1, %token_output_byte_base: offset, %channel: index, %values: vector<4xf32>, %scratch_values: view<256xf32>, %scratch_d: view<32xf32>, %output: buffer) + +template.decl @ggml.quantize_q8_1_x4.publish_vector4(%arg0: i1, %arg1: offset, %arg2: index, %arg3: vector<4xf32>, %arg4: view<256xf32>, %arg5: view<32xf32>, %arg6: buffer) + +template.def<@ggml.publish_f32.f32_vector4> device @ggml_publish_f32_f32_vector4(%publish_word: i1, %token: index, %channel: index, %token_count: index, %hidden_size: index, %values: vector<4xf32>, %output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_hidden_size = index.assume %hidden_size [range(%hidden_size, 128, 1073741824), mul(%hidden_size, 128)] : index + %c0_offset = index.constant 0 : offset + scf.if %publish_word { + %output_view = buffer.view %output[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_hidden_size]xf32> + vector.store %values, %output_view[%token, %channel] : vector<4xf32>, view<[%bounded_token_count]x[%bounded_hidden_size]xf32> + } + template.return +} + +template.def<@ggml.publish_f32.f16_vector4> device @ggml_publish_f32_f16_vector4(%publish_word: i1, %token: index, %channel: index, %token_count: index, %hidden_size: index, %values: vector<4xf32>, %output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_hidden_size = index.assume %hidden_size [range(%hidden_size, 128, 32768), mul(%hidden_size, 128)] : index + %c0_offset = index.constant 0 : offset + scf.if %publish_word { + %output_view = buffer.view %output[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_hidden_size]xf16> + %truncated = vector.fptrunc %values : vector<4xf32> to vector<4xf16> + vector.store %truncated, %output_view[%token, %channel] : vector<4xf16>, view<[%bounded_token_count]x[%bounded_hidden_size]xf16> + } + template.return +} + +template.def<@ggml.publish_f32.q8_1_x4_vector4> device @ggml_publish_f32_q8_1_x4_vector4(%publish_word: i1, %token_output_byte_base: offset, %channel: index, %values: vector<4xf32>, %scratch_values: view<256xf32>, %scratch_d: view<32xf32>, %output: buffer) { + template.apply<@ggml.quantize_q8_1_x4.publish_vector4>(%publish_word, %token_output_byte_base, %channel, %values, %scratch_values, %scratch_d, %output) : (i1, offset, index, vector<4xf32>, view<256xf32>, view<32xf32>, buffer) + template.return +} + +func.def inline @ggml_publish_f32_row_bytes(%format: index, %hidden_size: index) -> (offset) { + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c81 = index.constant 81 : index + %c128 = index.constant 128 : index + %zero_bytes = index.constant 0 : offset + %f16_element_bytes = index.constant 2 : offset + %f32_element_bytes = index.constant 4 : offset + %q8_group_bytes = index.constant 144 : offset + %is_f16 = index.cmp eq, %format, %c16 : index + %is_f32 = index.cmp eq, %format, %c32 : index + %is_q8_1 = index.cmp eq, %format, %c81 : index + %f16_row_bytes = index.scale %hidden_size, %f16_element_bytes : index, offset -> offset + %f32_row_bytes = index.scale %hidden_size, %f32_element_bytes : index, offset -> offset + %q8_group_count = index.div %hidden_size, %c128 : index + %q8_row_bytes = index.scale %q8_group_count, %q8_group_bytes : index, offset -> offset + %selected_f16 = scf.select %is_f16, %f16_row_bytes, %zero_bytes : offset + %selected_f32 = scf.select %is_f32, %f32_row_bytes, %selected_f16 : offset + %selected = scf.select %is_q8_1, %q8_row_bytes, %selected_f32 : offset + func.return %selected : offset +} + +template.def<@ggml.publish_f32.next_vector4> device @ggml_publish_f32_next_vector4(%format: index, %publish_word: i1, %token: index, %channel: index, %token_count: index, %hidden_size: index, %values: vector<4xf32>, %scratch_values: view<256xf32>, %scratch_d: view<32xf32>, %output: buffer) { + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c81 = index.constant 81 : index + %is_f16 = index.cmp eq, %format, %c16 : index + %is_f32 = index.cmp eq, %format, %c32 : index + %is_q8_1 = index.cmp eq, %format, %c81 : index + %publish_f16 = scalar.andi %publish_word, %is_f16 : i1 + %publish_f32 = scalar.andi %publish_word, %is_f32 : i1 + template.apply<@ggml.publish_f32.f16_vector4>(%publish_f16, %token, %channel, %token_count, %hidden_size, %values, %output) : (i1, index, index, index, index, vector<4xf32>, buffer) + template.apply<@ggml.publish_f32.f32_vector4>(%publish_f32, %token, %channel, %token_count, %hidden_size, %values, %output) : (i1, index, index, index, index, vector<4xf32>, buffer) + scf.if %is_q8_1 { + %row_bytes = func.call @ggml_publish_f32_row_bytes(%format, %hidden_size) : (index, index) -> (offset) + %token_output_byte_base = index.scale %token, %row_bytes : index, offset -> offset + template.apply<@ggml.publish_f32.q8_1_x4_vector4>(%publish_word, %token_output_byte_base, %channel, %values, %scratch_values, %scratch_d, %output) : (i1, offset, index, vector<4xf32>, view<256xf32>, view<32xf32>, buffer) + } + template.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/q4_k_f16.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/q4_k_f16.loom new file mode 100644 index 000000000000..73e457665b28 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/q4_k_f16.loom @@ -0,0 +1,100 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +// Decodes four adjacent values from one 32-value Q4_K group for FP16 matrix +// staging. The group and packet coordinates match the contiguous K dimension +// consumed by WMMA tiles. +func.def inline @ggml_q4k_f16_vector4(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group: index, %packet: index) -> (vector<4xf16>) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %block_bytes = index.constant 144 : offset + %scale_offset = index.constant 4 : offset + %code_offset = index.constant 16 : offset + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15 = vector.constant 15 : vector<1xi32> + %c48 = vector.constant 48 : vector<1xi32> + %q4_mask = vector.constant 252645135 : vector<1xi32> + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q4_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %scale_byte_base = index.add %block_byte_base, %scale_offset : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %dm_view = buffer.view %weight[%block_byte_base] : buffer -> view<2xf16> + %scale_view = buffer.view %weight[%scale_byte_base] : buffer -> view<3xi32> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<32xi32> + %dm = vector.load %dm_view[%c0] : view<2xf16> -> vector<2xf16> + %scales = vector.load %scale_view[%c0] : view<3xi32> -> vector<3xi32> + %d_f16 = vector.extract %dm[0] : vector<2xf16> -> f16 + %dmin_f16 = vector.extract %dm[1] : vector<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %q_page0 = index.div %bounded_group, %c2 : index + %q_page = index.mul %q_page0, %c8 : index + %q_word_index0 = index.add %q_page, %bounded_packet : index + %q_word_index = index.assume %q_word_index0 [range(%q_word_index0, 0, 31)] : index + %is_low = index.cmp ult, %bounded_group, %c4 : index + %scale_lane = index.rem %bounded_group, %c4 : index + %scale_shift_index = index.mul %scale_lane, %c8 : index + %scale_shift_i32 = index.cast %scale_shift_index : index to i32 + %scale_shift = vector.splat %scale_shift_i32 : vector<1xi32> + %scale0_i32 = vector.extract %scales[0] : vector<3xi32> -> i32 + %scale1_i32 = vector.extract %scales[1] : vector<3xi32> -> i32 + %scale2_i32 = vector.extract %scales[2] : vector<3xi32> -> i32 + %scale0 = vector.splat %scale0_i32 : vector<1xi32> + %scale1 = vector.splat %scale1_i32 : vector<1xi32> + %scale2 = vector.splat %scale2_i32 : vector<1xi32> + %high_shift_i32 = scalar.addi %scale_shift_i32, %c2_i32 : i32 + %minimum_shift_i32 = scalar.addi %scale_shift_i32, %c4_i32 : i32 + %selected_scale_source = scf.select %is_low, %scale0, %scale2 : vector<1xi32> + %selected_minimum_source = scf.select %is_low, %scale1, %scale2 : vector<1xi32> + %selected_scale_high_shift_i32 = scf.select %is_low, %scale_shift_i32, %high_shift_i32 : i32 + %selected_minimum_low_shift_i32 = scf.select %is_low, %scale_shift_i32, %minimum_shift_i32 : i32 + %selected_scale_high_shift = vector.splat %selected_scale_high_shift_i32 : vector<1xi32> + %selected_minimum_low_shift = vector.splat %selected_minimum_low_shift_i32 : vector<1xi32> + %scale_low0 = vector.shrui %selected_scale_source, %scale_shift : vector<1xi32> + %scale_low = vector.andi %scale_low0, %c15 : vector<1xi32> + %scale_high0 = vector.shrui %scale0, %selected_scale_high_shift : vector<1xi32> + %scale_high = vector.andi %scale_high0, %c48 : vector<1xi32> + %scale = vector.ori %scale_low, %scale_high : vector<1xi32> + %minimum_low0 = vector.shrui %selected_minimum_source, %selected_minimum_low_shift : vector<1xi32> + %minimum_low = vector.andi %minimum_low0, %c15 : vector<1xi32> + %minimum_high0 = vector.shrui %scale1, %selected_scale_high_shift : vector<1xi32> + %minimum_high = vector.andi %minimum_high0, %c48 : vector<1xi32> + %minimum = vector.ori %minimum_low, %minimum_high : vector<1xi32> + %scale_f32 = vector.uitofp %scale : vector<1xi32> to vector<1xf32> + %minimum_f32 = vector.uitofp %minimum : vector<1xi32> to vector<1xf32> + %d_vector1 = vector.splat %d : vector<1xf32> + %dmin_vector1 = vector.splat %dmin : vector<1xf32> + %d_scale_vector1 = vector.mulf %d_vector1, %scale_f32 : vector<1xf32> + %minimum_scale_vector1 = vector.mulf %dmin_vector1, %minimum_f32 : vector<1xf32> + %d_scale = vector.extract %d_scale_vector1[0] : vector<1xf32> -> f32 + %minimum_scale = vector.extract %minimum_scale_vector1[0] : vector<1xf32> -> f32 + %q_word = vector.load %code_view[%q_word_index] : view<32xi32> -> vector<1xi32> + %q_half = index.rem %bounded_group, %c2 : index + %q_shift_index = index.mul %q_half, %c4 : index + %q_shift_i32 = index.cast %q_shift_index : index to i32 + %q_shift = vector.splat %q_shift_i32 : vector<1xi32> + %shifted_q = vector.shrui %q_word, %q_shift : vector<1xi32> + %masked_q = vector.andi %shifted_q, %q4_mask : vector<1xi32> + %q_i8 = vector.bitcast %masked_q : vector<1xi32> to vector<4xi8> + %q_f32 = vector.uitofp %q_i8 : vector<4xi8> to vector<4xf32> + %negative_minimum_scale = scalar.negf %minimum_scale : f32 + %q0 = vector.extract %q_f32[0] : vector<4xf32> -> f32 + %q1 = vector.extract %q_f32[1] : vector<4xf32> -> f32 + %q2 = vector.extract %q_f32[2] : vector<4xf32> -> f32 + %q3 = vector.extract %q_f32[3] : vector<4xf32> -> f32 + %value0 = scalar.fmaf %q0, %d_scale, %negative_minimum_scale : f32 + %value1 = scalar.fmaf %q1, %d_scale, %negative_minimum_scale : f32 + %value2 = scalar.fmaf %q2, %d_scale, %negative_minimum_scale : f32 + %value3 = scalar.fmaf %q3, %d_scale, %negative_minimum_scale : f32 + %half0 = scalar.fptrunc %value0 : f32 to f16 + %half1 = scalar.fptrunc %value1 : f32 to f16 + %half2 = scalar.fptrunc %value2 : f32 to f16 + %half3 = scalar.fptrunc %value3 : f32 to f16 + %result = vector.from_elements %half0, %half1, %half2, %half3 : vector<4xf16> + func.return %result : vector<4xf16> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/q6_k_f16.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/q6_k_f16.loom new file mode 100644 index 000000000000..b461ffd1c153 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/q6_k_f16.loom @@ -0,0 +1,182 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +// Decodes four adjacent values from one 32-value Q6_K group for FP16 matrix +// staging. The group and packet coordinates match the contiguous K dimension +// consumed by WMMA tiles. +func.def inline @ggml_q6k_f16_vector4(%weight: buffer, %weight_row_byte_base: offset, %q6_block: index, %q6_group: index, %packet: index) -> (vector<4xf16>) { + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %block_bytes = index.constant 210 : offset + %qh_byte_add = index.constant 128 : offset + %scale_byte_add = index.constant 192 : offset + %d_byte_add = index.constant 208 : offset + %c4_i32v = vector.constant 4 : vector<1xi32> + %nibble_mask = vector.constant 252645135 : vector<1xi32> + %high_mask = vector.constant 50529027 : vector<1xi32> + %c32_f32v = vector.constant 32.0 : vector<4xf32> + %bounded_group = index.assume %q6_group [range(%q6_group, 0, 7)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q6_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %weight_row_byte_base, %block_byte_add : offset + %qh_byte_base = index.add %block_byte_base, %qh_byte_add : offset + %scale_byte_base = index.add %block_byte_base, %scale_byte_add : offset + %d_byte_base = index.add %block_byte_base, %d_byte_add : offset + %ql_view = buffer.view %weight[%block_byte_base] : buffer -> view<32xi32> + %qh_view = buffer.view %weight[%qh_byte_base] : buffer -> view<16xi32> + %scale_view = buffer.view %weight[%scale_byte_base] : buffer -> view<16xi8> + %d_view = buffer.view %weight[%d_byte_base] : buffer -> view<1xf16> + %group_in_half = index.rem %bounded_group, %c4 : index + %half = index.div %bounded_group, %c4 : index + %ql_side = index.rem %group_in_half, %c2 : index + %ql_half_word_base = index.mul %half, %c16 : index + %ql_side_word_add = index.mul %ql_side, %c8 : index + %ql_word_base = index.add %ql_half_word_base, %ql_side_word_add : index + %ql_word_index = index.add %ql_word_base, %bounded_packet : index + %qh_half_word_base = index.mul %half, %c8 : index + %qh_word_index = index.add %qh_half_word_base, %bounded_packet : index + %nibble = index.div %group_in_half, %c2 : index + %nibble_shift_index = index.mul %nibble, %c4 : index + %nibble_shift_i32 = index.cast %nibble_shift_index : index to i32 + %nibble_shift = vector.splat %nibble_shift_i32 : vector<1xi32> + %qh_shift_index = index.mul %group_in_half, %c2 : index + %qh_shift_i32 = index.cast %qh_shift_index : index to i32 + %qh_shift = vector.splat %qh_shift_i32 : vector<1xi32> + %scale_packet_half = index.div %bounded_packet, %c4 : index + %scale_group_base = index.mul %bounded_group, %c2 : index + %scale_index = index.add %scale_group_base, %scale_packet_half : index + %ql_word = vector.load %ql_view[%ql_word_index] : view<32xi32> -> vector<1xi32> + %qh_word = vector.load %qh_view[%qh_word_index] : view<16xi32> -> vector<1xi32> + %ql_shifted = vector.shrui %ql_word, %nibble_shift : vector<1xi32> + %ql = vector.andi %ql_shifted, %nibble_mask : vector<1xi32> + %qh_shifted = vector.shrui %qh_word, %qh_shift : vector<1xi32> + %qh_low = vector.andi %qh_shifted, %high_mask : vector<1xi32> + %qh = vector.shli %qh_low, %c4_i32v : vector<1xi32> + %code = vector.ori %ql, %qh : vector<1xi32> + %code_i8 = vector.bitcast %code : vector<1xi32> to vector<4xi8> + %code_f32 = vector.uitofp %code_i8 : vector<4xi8> to vector<4xf32> + %centered = vector.subf %code_f32, %c32_f32v : vector<4xf32> + %scale_i8 = view.load %scale_view[%scale_index] : view<16xi8> -> i8 + %d_f16 = view.load %d_view[0] : view<1xf16> -> f16 + %scale = scalar.sitofp %scale_i8 : i8 to f32 + %d = scalar.extf %d_f16 : f16 to f32 + %combined_scale = scalar.mulf %scale, %d : f32 + %combined_scale_vector = vector.splat %combined_scale : vector<4xf32> + %values_f32 = vector.mulf %centered, %combined_scale_vector : vector<4xf32> + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +// Decode adjacent K32 groups using their shared high bits and scales. +func.def public inline @ggml_q6k_f16_pair(%half_packet: i1, %weight: buffer, %weight_row_byte_base: offset, %q6_block: index, %q6_group: index, %packet: index) -> (vector<16xf16>, vector<16xf16>) { + %padding = vector.constant 0 : vector<2xi32> + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %block_bytes = index.constant 210 : offset + %qh_byte_add = index.constant 128 : offset + %scale_byte_add = index.constant 192 : offset + %d_byte_add = index.constant 208 : offset + %c4_i32v = vector.constant 4 : vector<4xi32> + %nibble_mask = vector.constant 252645135 : vector<4xi32> + %high_mask = vector.constant 50529027 : vector<4xi32> + %c32_f32v = vector.constant 32.0 : vector<16xf32> + %bounded_group = index.assume %q6_group [range(%q6_group, 0, 6), mul(%q6_group, 2)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 6), mul(%packet, 2)] : index + %block_byte_add = index.scale %q6_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %weight_row_byte_base, %block_byte_add : offset + %qh_byte_base = index.add %block_byte_base, %qh_byte_add : offset + %scale_byte_base = index.add %block_byte_base, %scale_byte_add : offset + %d_byte_base = index.add %block_byte_base, %d_byte_add : offset + %ql_view = buffer.view %weight[%block_byte_base] : buffer -> view<32xi32> + %qh_view = buffer.view %weight[%qh_byte_base] : buffer -> view<16xi32> + %scale_view = buffer.view %weight[%scale_byte_base] : buffer -> view<4xi32> + %d_view = buffer.view %weight[%d_byte_base] : buffer -> view<1xf16> + %group_in_half = index.rem %bounded_group, %c4 : index + %half = index.div %bounded_group, %c4 : index + %ql_half_word_base = index.mul %half, %c16 : index + %ql_word_index0 = index.add %ql_half_word_base, %bounded_packet : index + %ql_word_index1 = index.add %ql_word_index0, %c8 : index + %qh_half_word_base = index.mul %half, %c8 : index + %qh_word_index = index.add %qh_half_word_base, %bounded_packet : index + %nibble = index.div %group_in_half, %c2 : index + %nibble_shift_index = index.mul %nibble, %c4 : index + %nibble_shift_i32 = index.cast %nibble_shift_index : index to i32 + %nibble_shift = vector.splat %nibble_shift_i32 : vector<4xi32> + %qh_shift_index0 = index.mul %group_in_half, %c2 : index + %qh_shift_index1 = index.add %qh_shift_index0, %c2 : index + %qh_shift_i32_0 = index.cast %qh_shift_index0 : index to i32 + %qh_shift_i32_1 = index.cast %qh_shift_index1 : index to i32 + %qh_shift0 = vector.splat %qh_shift_i32_0 : vector<4xi32> + %qh_shift1 = vector.splat %qh_shift_i32_1 : vector<4xi32> + %scale_packet_half = index.div %bounded_packet, %c4 : index + %scale_word_index = index.div %bounded_group, %c2 : index + %scale_shift_index0 = index.mul %scale_packet_half, %c8 : index + %scale_shift_index1 = index.add %scale_shift_index0, %c16 : index + %scale_shift0 = index.cast %scale_shift_index0 : index to i32 + %scale_shift1 = index.cast %scale_shift_index1 : index to i32 + %ql_word0 = scf.if %half_packet -> (vector<4xi32>) { + %words = vector.load %ql_view[%ql_word_index0] : view<32xi32> -> vector<2xi32> + %padded = vector.concat<0> %words, %padding : vector<2xi32>, vector<2xi32> -> vector<4xi32> + scf.yield %padded : vector<4xi32> + } else { + %words = vector.load %ql_view[%ql_word_index0] : view<32xi32> -> vector<4xi32> + scf.yield %words : vector<4xi32> + } + %ql_word1 = scf.if %half_packet -> (vector<4xi32>) { + %words = vector.load %ql_view[%ql_word_index1] : view<32xi32> -> vector<2xi32> + %padded = vector.concat<0> %words, %padding : vector<2xi32>, vector<2xi32> -> vector<4xi32> + scf.yield %padded : vector<4xi32> + } else { + %words = vector.load %ql_view[%ql_word_index1] : view<32xi32> -> vector<4xi32> + scf.yield %words : vector<4xi32> + } + %qh_word = scf.if %half_packet -> (vector<4xi32>) { + %words = vector.load %qh_view[%qh_word_index] : view<16xi32> -> vector<2xi32> + %padded = vector.concat<0> %words, %padding : vector<2xi32>, vector<2xi32> -> vector<4xi32> + scf.yield %padded : vector<4xi32> + } else { + %words = vector.load %qh_view[%qh_word_index] : view<16xi32> -> vector<4xi32> + scf.yield %words : vector<4xi32> + } + %scale_word = view.load %scale_view[%scale_word_index] : view<4xi32> -> i32 + %d_f16 = view.load %d_view[%c0] : view<1xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %ql_shifted0 = vector.shrui %ql_word0, %nibble_shift : vector<4xi32> + %ql0 = vector.andi %ql_shifted0, %nibble_mask : vector<4xi32> + %qh_shifted0 = vector.shrui %qh_word, %qh_shift0 : vector<4xi32> + %qh_low0 = vector.andi %qh_shifted0, %high_mask : vector<4xi32> + %qh0 = vector.shli %qh_low0, %c4_i32v : vector<4xi32> + %code0 = vector.ori %ql0, %qh0 : vector<4xi32> + %code_i8_0 = vector.bitcast %code0 : vector<4xi32> to vector<16xi8> + %code_f32_0 = vector.uitofp %code_i8_0 : vector<16xi8> to vector<16xf32> + %centered0 = vector.subf %code_f32_0, %c32_f32v : vector<16xf32> + %scale_shifted0 = scalar.shrui %scale_word, %scale_shift0 : i32 + %scale_i8_0 = scalar.trunci %scale_shifted0 : i32 to i8 + %scale0 = scalar.sitofp %scale_i8_0 : i8 to f32 + %combined_scale0 = scalar.mulf %scale0, %d : f32 + %combined_scale_vector0 = vector.splat %combined_scale0 : vector<16xf32> + %values_f32_0 = vector.mulf %centered0, %combined_scale_vector0 : vector<16xf32> + %values0 = vector.fptrunc %values_f32_0 : vector<16xf32> to vector<16xf16> + %ql_shifted1 = vector.shrui %ql_word1, %nibble_shift : vector<4xi32> + %ql1 = vector.andi %ql_shifted1, %nibble_mask : vector<4xi32> + %qh_shifted1 = vector.shrui %qh_word, %qh_shift1 : vector<4xi32> + %qh_low1 = vector.andi %qh_shifted1, %high_mask : vector<4xi32> + %qh1 = vector.shli %qh_low1, %c4_i32v : vector<4xi32> + %code1 = vector.ori %ql1, %qh1 : vector<4xi32> + %code_i8_1 = vector.bitcast %code1 : vector<4xi32> to vector<16xi8> + %code_f32_1 = vector.uitofp %code_i8_1 : vector<16xi8> to vector<16xf32> + %centered1 = vector.subf %code_f32_1, %c32_f32v : vector<16xf32> + %scale_shifted1 = scalar.shrui %scale_word, %scale_shift1 : i32 + %scale_i8_1 = scalar.trunci %scale_shifted1 : i32 to i8 + %scale1 = scalar.sitofp %scale_i8_1 : i8 to f32 + %combined_scale1 = scalar.mulf %scale1, %d : f32 + %combined_scale_vector1 = vector.splat %combined_scale1 : vector<16xf32> + %values_f32_1 = vector.mulf %centered1, %combined_scale_vector1 : vector<16xf32> + %values1 = vector.fptrunc %values_f32_1 : vector<16xf32> to vector<16xf16> + func.return %values0, %values1 : vector<16xf16>, vector<16xf16> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/q8_0_f16.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/q8_0_f16.loom new file mode 100644 index 000000000000..26f7760ac4b9 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/q8_0_f16.loom @@ -0,0 +1,26 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +// Decodes four adjacent values from one 32-value Q8_0 block for FP16 matrix +// staging. +func.def inline @ggml_q8_0_f16_vector4(%weight: buffer, %row_byte_base: offset, %q8_block: index, %packet: index) -> (vector<4xf16>) { + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %block_bytes = index.constant 34 : offset + %code_offset = index.constant 2 : offset + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q8_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %d_view = buffer.view %weight[%block_byte_base] : buffer -> view<1xf16> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<32xi8> + %packet_base = index.mul %bounded_packet, %c4 : index + %codes = vector.load %code_view[%packet_base] : view<32xi8> -> vector<4xi8> + %codes_f32 = vector.sitofp %codes : vector<4xi8> to vector<4xf32> + %d_f16 = view.load %d_view[%c0] : view<1xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %d_vector = vector.splat %d : vector<4xf32> + %values_f32 = vector.mulf %codes_f32, %d_vector : vector<4xf32> + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/q8_1_f16.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/q8_1_f16.loom new file mode 100644 index 000000000000..1ccd1fdf3b36 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/q8_1_f16.loom @@ -0,0 +1,26 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +// Decodes four adjacent values from one 32-value Q8_1 block for FP16 matrix +// staging. The block sum field is not part of the row values. +func.def inline @ggml_q8_1_f16_vector4(%weight: buffer, %row_byte_base: offset, %q8_block: index, %packet: index) -> (vector<4xf16>) { + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %block_bytes = index.constant 36 : offset + %code_offset = index.constant 4 : offset + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q8_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %ds_view = buffer.view %weight[%block_byte_base] : buffer -> view<2xf16> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<32xi8> + %packet_base = index.mul %bounded_packet, %c4 : index + %codes = vector.load %code_view[%packet_base] : view<32xi8> -> vector<4xi8> + %codes_f32 = vector.sitofp %codes : vector<4xi8> to vector<4xf32> + %d_f16 = view.load %ds_view[%c0] : view<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %d_vector = vector.splat %d : vector<4xf32> + %values_f32 = vector.mulf %codes_f32, %d_vector : vector<4xf32> + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/quantize_q8_1_x4.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/quantize_q8_1_x4.loom new file mode 100644 index 000000000000..2866f5671e86 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/quantize_q8_1_x4.loom @@ -0,0 +1,120 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +// Publishes four F32 values into GGML's block_q8_1_x4 layout. The caller must +// execute complete eight-lane cohorts with four consecutive channels per lane. +template.decl @ggml.quantize_q8_1_x4.publish_vector4.body(%strict_scaling: i1, %publish_word: i1, %token_output_byte_base: offset, %channel: index, %values: vector<4xf32>, %scratch_values: view<256xf32>, %scratch_d: view<32xf32>, %output: buffer) + +template.decl @ggml.quantize_q8_1_x4.publish_vector4(%publish_word: i1, %token_output_byte_base: offset, %channel: index, %values: vector<4xf32>, %scratch_values: view<256xf32>, %scratch_d: view<32xf32>, %output: buffer) + +template.def<@ggml.quantize_q8_1_x4.publish_vector4> device @ggml_quantize_q8_1_x4_publish_vector4(%publish_word: i1, %token_output_byte_base: offset, %channel: index, %values: vector<4xf32>, %scratch_values: view<256xf32>, %scratch_d: view<32xf32>, %output: buffer) { + %strict_scaling = scalar.constant false : i1 + template.apply<@ggml.quantize_q8_1_x4.publish_vector4.body>(%strict_scaling, %publish_word, %token_output_byte_base, %channel, %values, %scratch_values, %scratch_d, %output) : (i1, i1, offset, index, vector<4xf32>, view<256xf32>, view<32xf32>, buffer) + template.return +} + +template.decl @ggml.quantize_q8_1_x4.publish_vector4_strict(%publish_word: i1, %token_output_byte_base: offset, %channel: index, %values: vector<4xf32>, %scratch_values: view<256xf32>, %scratch_d: view<32xf32>, %output: buffer) + +template.def<@ggml.quantize_q8_1_x4.publish_vector4_strict> device @ggml_quantize_q8_1_x4_publish_vector4_strict(%publish_word: i1, %token_output_byte_base: offset, %channel: index, %values: vector<4xf32>, %scratch_values: view<256xf32>, %scratch_d: view<32xf32>, %output: buffer) { + %strict_scaling = scalar.constant true : i1 + template.apply<@ggml.quantize_q8_1_x4.publish_vector4.body>(%strict_scaling, %publish_word, %token_output_byte_base, %channel, %values, %scratch_values, %scratch_d, %output) : (i1, i1, offset, index, vector<4xf32>, view<256xf32>, view<32xf32>, buffer) + template.return +} + +template.def<@ggml.quantize_q8_1_x4.publish_vector4.body> device @ggml_quantize_q8_1_x4_publish_vector4_body(%strict_scaling: i1, %publish_word: i1, %token_output_byte_base: offset, %channel: index, %values: vector<4xf32>, %scratch_values: view<256xf32>, %scratch_d: view<32xf32>, %output: buffer) { + %workitem = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c32 = index.constant 32 : index + %group_bytes = index.constant 144 : offset + %payload_byte_add = index.constant 16 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %c1_f32 = scalar.constant 1.0 : f32 + %c127 = scalar.constant 127.0 : f32 + %word_in_block = index.rem %workitem, %c8 : index + %absolute_values = vector.absf %values : vector<4xf32> + %thread_max = vector.reduce %absolute_values, %c0_f32 : vector<4xf32>, f32 + %shuffle_width = scalar.constant 8 : i32 + %xor1 = scalar.constant 1 : i32 + %max_peer1, %max_valid1 = kernel.subgroup.shuffle %thread_max, %xor1, %shuffle_width : f32, i32, i32 + %max1 = scalar.maxnumf %thread_max, %max_peer1 : f32 + %xor2 = scalar.constant 2 : i32 + %max_peer2, %max_valid2 = kernel.subgroup.shuffle %max1, %xor2, %shuffle_width : f32, i32, i32 + %max2 = scalar.maxnumf %max1, %max_peer2 : f32 + %xor4 = scalar.constant 4 : i32 + %max_peer4, %max_valid4 = kernel.subgroup.shuffle %max2, %xor4, %shuffle_width : f32, i32, i32 + %max4 = scalar.maxnumf %max2, %max_peer4 : f32 + %d0 = scf.if %strict_scaling -> (f32) { + %precise = scalar.divf %max4, %c127 : f32 + scf.yield %precise : f32 + } else { + %fast = scalar.divf %max4, %c127 : f32 + scf.yield %fast : f32 + } + %d = scf.select %publish_word, %d0, %c0_f32 : f32 + %is_block_leader = index.cmp eq, %word_in_block, %c0 : index + %writes_block_d = scalar.andi %publish_word, %is_block_leader : i1 + %d_nonzero = scalar.cmpf one, %d, %c0_f32 : f32 + %d_inverse = scf.if %d_nonzero -> (f32) { + %inverse = scf.if %strict_scaling -> (f32) { + %precise = scalar.divf %c1_f32, %d : f32 + scf.yield %precise : f32 + } else { + %fast = scalar.divf %c1_f32, %d : f32 + scf.yield %fast : f32 + } + scf.yield %inverse : f32 + } else { + scf.yield %c0_f32 : f32 + } + %d_inverse_vector = vector.splat %d_inverse : vector<4xf32> + %scaled_values = vector.mulf %values, %d_inverse_vector : vector<4xf32> + %rounded_values = vector.roundf %scaled_values : vector<4xf32> + %quantized_values = vector.fptosi %rounded_values : vector<4xf32> to vector<4xi8> + %packed_word = vector.bitcast %quantized_values : vector<4xi8> to vector<1xi32> + scf.if %publish_word { + %q8_block = index.div %channel, %c32 : index + %physical_group = index.div %q8_block, %c4 : index + %block_in_group = index.rem %q8_block, %c4 : index + %group_byte_add = index.scale %physical_group, %group_bytes : index, offset -> offset + %group_byte_offset = index.add %token_output_byte_base, %group_byte_add : offset + %payload_byte_offset = index.add %group_byte_offset, %payload_byte_add : offset + %group_qs = buffer.view %output[%payload_byte_offset] : buffer -> view<32xi32> + %block_word_base = index.mul %block_in_group, %c8 : index + %packed_word_index0 = index.add %block_word_base, %word_in_block : index + %packed_word_index = index.assume %packed_word_index0 [range(%packed_word_index0, 0, 31)] : index + vector.store %packed_word, %group_qs[%packed_word_index] : vector<1xi32>, view<32xi32> + } + %thread_quantized_sum = vector.reduce %rounded_values, %c0_f32 : vector<4xf32>, f32 + %sum_peer1, %sum_valid1 = kernel.subgroup.shuffle %thread_quantized_sum, %xor1, %shuffle_width : f32, i32, i32 + %sum1 = scalar.addf %thread_quantized_sum, %sum_peer1 : f32 + %sum_peer2, %sum_valid2 = kernel.subgroup.shuffle %sum1, %xor2, %shuffle_width : f32, i32, i32 + %sum2 = scalar.addf %sum1, %sum_peer2 : f32 + %sum_peer4, %sum_valid4 = kernel.subgroup.shuffle %sum2, %xor4, %shuffle_width : f32, i32, i32 + %quantized_sum = scalar.addf %sum2, %sum_peer4 : f32 + scf.if %writes_block_d { + %s = scf.if %strict_scaling -> (f32) { + %precise = scalar.mulf %quantized_sum, %d : f32 + scf.yield %precise : f32 + } else { + %fast = scalar.mulf %quantized_sum, %d : f32 + scf.yield %fast : f32 + } + %q8_block = index.div %channel, %c32 : index + %physical_group = index.div %q8_block, %c4 : index + %block_in_group = index.rem %q8_block, %c4 : index + %group_byte_add = index.scale %physical_group, %group_bytes : index, offset -> offset + %group_byte_offset = index.add %token_output_byte_base, %group_byte_add : offset + %group_ds = buffer.view %output[%group_byte_offset] : buffer -> view<8xf16> + %d_f16 = scalar.fptrunc %d : f32 to f16 + %s_f16 = scalar.fptrunc %s : f32 to f16 + %ds_index = index.mul %block_in_group, %c2 : index + %s_index = index.add %ds_index, %c1 : index + view.store %d_f16, %group_ds[%ds_index] : f16, view<8xf16> + view.store %s_f16, %group_ds[%s_index] : f16, view<8xf16> + } + template.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/rmsnorm_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/rmsnorm_f32.loom new file mode 100644 index 000000000000..2ceae48b03ec --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/rmsnorm_f32.loom @@ -0,0 +1,132 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +// Computes the reciprocal RMS scale for one F32 row. The caller owns launch +// geometry; this motif assumes a 256-workitem wave32 workgroup. +template.decl @ggml.rmsnorm_f32.apply(%value: f32, %scale: f32) -> (f32) + +template.decl @ggml.rmsnorm_f32.row_scale(%token_count0: index, %token0: index, %hidden_size0: index, %input_stride0: index, %epsilon: f32, %input: buffer) -> (f32, index, index, i1) + +template.def<@ggml.rmsnorm_f32.row_scale> device @ggml_rmsnorm_f32_row_scale(%token_count0: index, %token0: index, %hidden_size0: index, %input_stride0: index, %epsilon: f32, %input: buffer) -> (f32, index, index, i1) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 1048576)] : index + %hidden_size1 = index.assume %hidden_size0 [range(%hidden_size0, 64, 32768), mul(%hidden_size0, 64)] : index + %input_stride1 = index.assume %input_stride0 [range(%input_stride0, 64, 1048576)] : index + %hidden_size, %input_stride = index.assume %hidden_size1, %input_stride1 [le(%hidden_size1, %input_stride1)] : index, index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c8 = index.constant 8 : index + %c256 = index.constant 256 : index + %scratch_bytes = index.constant 1024 : offset + %c0_offset = index.constant 0 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %valid_token = index.cmp ult, %token0, %token_count : index + %safe_token0 = scf.select %valid_token, %token0, %c0 : index + %token = index.assume %safe_token0 [lt(%safe_token0, %token_count)] : index + %hidden_size_i32 = index.cast %hidden_size : index to i32 + %hidden_size_f32 = scalar.sitofp %hidden_size_i32 : i32 to f32 + %input_view = buffer.view %input[%c0_offset] : buffer -> view<[%token_count]x[%input_stride]xf32> + %c1024 = index.constant 1024 : index + %tail = index.rem %hidden_size, %c1024 : index + %unroll_limit = index.sub %hidden_size, %tail : index + %prefix_sum = scf.for %channel = [%workitem to %unroll_limit step %c1024](%running_sum = %c0_f32 : f32) -> (f32) { + %offset0 = index.constant 0 : index + %channel0 = index.add %channel, %offset0 : index + %value0 = view.load %input_view[%token, %channel0] : view<[%token_count]x[%input_stride]xf32> -> f32 + %offset1 = index.constant 256 : index + %channel1 = index.add %channel, %offset1 : index + %value1 = view.load %input_view[%token, %channel1] : view<[%token_count]x[%input_stride]xf32> -> f32 + %offset2 = index.constant 512 : index + %channel2 = index.add %channel, %offset2 : index + %value2 = view.load %input_view[%token, %channel2] : view<[%token_count]x[%input_stride]xf32> -> f32 + %offset3 = index.constant 768 : index + %channel3 = index.add %channel, %offset3 : index + %value3 = view.load %input_view[%token, %channel3] : view<[%token_count]x[%input_stride]xf32> -> f32 + %square0 = scalar.mulf %value0, %value0 : f32 + %sum0 = scalar.addf %running_sum, %square0 : f32 + %square1 = scalar.mulf %value1, %value1 : f32 + %sum1 = scalar.addf %sum0, %square1 : f32 + %square2 = scalar.mulf %value2, %value2 : f32 + %sum2 = scalar.addf %sum1, %square2 : f32 + %square3 = scalar.mulf %value3, %value3 : f32 + %sum3 = scalar.addf %sum2, %square3 : f32 + scf.yield %sum3 : f32 + } + %tail_start = index.add %unroll_limit, %workitem : index + %thread_sum = scf.for %channel = [%tail_start to %hidden_size step %c256](%running_sum = %prefix_sum : f32) -> (f32) { + %value = view.load %input_view[%token, %channel] : view<[%token_count]x[%input_stride]xf32> -> f32 + %square = scalar.mulf %value, %value : f32 + %next_sum = scalar.addf %running_sum, %square : f32 + scf.yield %next_sum : f32 + } + %subgroup_sum = kernel.subgroup.reduce %thread_sum : f32 + %scratch = buffer.alloca align(16) %scratch_bytes : buffer + %scratch_view = buffer.view %scratch[%c0_offset] : buffer -> view<256xf32> + %is_subgroup_leader = index.cmp eq, %lane, %c0 : index + scf.if %is_subgroup_leader { + view.store %subgroup_sum, %scratch_view[%subgroup] : f32, view<256xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %is_reduction_subgroup = index.cmp eq, %subgroup, %c0 : index + %is_reduction_lane = index.cmp ult, %lane, %c8 : index + %loads_subgroup_sum = scalar.andi %is_reduction_subgroup, %is_reduction_lane : i1 + %subgroup_partial = scf.if %loads_subgroup_sum -> (f32) { + %value = view.load %scratch_view[%lane] : view<256xf32> -> f32 + scf.yield %value : f32 + } else { + scf.yield %c0_f32 : f32 + } + %row_sum = kernel.subgroup.reduce %subgroup_partial : f32 + %writes_scale = scalar.andi %is_reduction_subgroup, %is_subgroup_leader : i1 + scf.if %writes_scale { + %mean = scalar.divf %row_sum, %hidden_size_f32 : f32 + %biased_mean = scalar.addf %mean, %epsilon : f32 + %scale = scalar.rsqrtf %biased_mean : f32 + view.store %scale, %scratch_view[%c0] : f32, view<256xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %scale = view.load %scratch_view[%c0] : view<256xf32> -> f32 + template.return %scale, %token, %token_count, %valid_token : f32, index, index, i1 +} + +// One subgroup owns a short row; callers supply a valid row index and pack +// independent rows into a workgroup. No workgroup scratch or barriers. +template.decl @ggml.rmsnorm_f32.subgroup_row_scale(%token_count: index, %token: index, %hidden_size: index, %epsilon: f32, %input: buffer) -> (f32) + +template.def<@ggml.rmsnorm_f32.subgroup_row_scale> device @ggml_rmsnorm_f32_subgroup_row_scale(%token_count0: index, %token0: index, %hidden0: index, %epsilon: f32, %input_n: buffer) -> (f32) { + %bounded_token_count = index.assume %token_count0 [range(%token_count0, 1, 1048576)] : index + %row = index.assume %token0 [lt(%token0, %bounded_token_count)] : index + %hidden = index.assume %hidden0 [range(%hidden0, 128, 1024), mul(%hidden0, 128)] : index + %lane = kernel.subgroup.lane.id : index + %four = index.constant 4 : index + %c128 = index.constant 128 : index + %zero_offset = index.constant 0 : offset + %channel0 = index.mul %lane, %four : index + %zero = scalar.constant 0.0 : f32 + %zero4 = vector.constant 0.0 : vector<4xf32> + %iv = buffer.view %input_n[%zero_offset] : buffer -> view<[%bounded_token_count]x[%hidden]xf32> + %last_channel = index.sub %hidden, %four : index + %squares = scf.for %channel = [%channel0 to %hidden step %c128](%sum = %zero4 : vector<4xf32>) -> (vector<4xf32>) { + %bounded_channel = index.assume %channel [range(%channel, 0, 1020), mul(%channel, 4), le(%channel, %last_channel)] : index + %x = vector.load %iv[%row, %bounded_channel] : view<[%bounded_token_count]x[%hidden]xf32> -> vector<4xf32> + %xx = vector.mulf %x, %x : vector<4xf32> + %next = vector.addf %sum, %xx : vector<4xf32> + scf.yield %next : vector<4xf32> + } + %lane_sum = vector.reduce %squares, %zero : vector<4xf32>, f32 + %row_sum = kernel.subgroup.reduce %lane_sum : f32 + %hidden_i = index.cast %hidden : index to i32 + %hidden_f = scalar.sitofp %hidden_i : i32 to f32 + %mean = scalar.divf %row_sum, %hidden_f : f32 + %biased = scalar.addf %mean, %epsilon : f32 + %scale = scalar.rsqrtf %biased : f32 + template.return %scale : f32 +} + +template.def<@ggml.rmsnorm_f32.apply> device @ggml_rmsnorm_f32_apply(%value: f32, %scale: f32) -> (f32) { + %normalized = scalar.mulf %value, %scale : f32 + template.return %normalized : f32 +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/rope_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/rope_f32.loom new file mode 100644 index 000000000000..cbbe90dee2c7 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/rope_f32.loom @@ -0,0 +1,20 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +func.def inline @ggml_rope_f32_pair_packet(%position: f32, %theta: vector<2xf32>, %freq_factors: vector<2xf32>, %x_values: vector<2xf32>, %y_values: vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) { + %position_vector = vector.splat %position : vector<2xf32> + %inverse_two_pi = scalar.constant 0.15915494309189535 : f32 + %inverse_two_pi_vector = vector.splat %inverse_two_pi : vector<2xf32> + %scaled_theta = vector.divf %theta, %freq_factors : vector<2xf32> + %angles = vector.mulf %position_vector, %scaled_theta : vector<2xf32> + %turns = vector.mulf %angles, %inverse_two_pi_vector : vector<2xf32> + %cosines = vector.costurnsf %turns : vector<2xf32> + %sines = vector.sinturnsf %turns : vector<2xf32> + %x_cosines = vector.mulf %x_values, %cosines : vector<2xf32> + %y_sines = vector.mulf %y_values, %sines : vector<2xf32> + %x_sines = vector.mulf %x_values, %sines : vector<2xf32> + %y_cosines = vector.mulf %y_values, %cosines : vector<2xf32> + %rotated_x = vector.subf %x_cosines, %y_sines : vector<2xf32> + %rotated_y = vector.addf %x_sines, %y_cosines : vector<2xf32> + func.return %rotated_x, %rotated_y : vector<2xf32>, vector<2xf32> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/unary_f32_apply.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/unary_f32_apply.loom new file mode 100644 index 000000000000..84b1d30a768d --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/unary_f32_apply.loom @@ -0,0 +1,380 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +template.decl @ggml.unary_f32.apply(%op: index, %value: f32) -> (f32) + +template.decl @ggml.unary_f32.apply_vector4(%op: index, %values: vector<4xf32>) -> (vector<4xf32>) + +template.def<@ggml.unary_f32.apply> device priority(1) @ggml_unary_f32_apply(%op: index, %value: f32) -> (f32) { + %c0_f32 = scalar.constant 0.0 : f32 + %c0_5_f32 = scalar.constant 0.5 : f32 + %c1_f32 = scalar.constant 1.0 : f32 + %c2_f32 = scalar.constant 2.0 : f32 + %c3_f32 = scalar.constant 3.0 : f32 + %c6_f32 = scalar.constant 6.0 : f32 + %cn1_f32 = scalar.constant -1.0 : f32 + %cn2_f32 = scalar.constant -2.0 : f32 + %gelu_coef = scalar.constant 0.044715 : f32 + %gelu_quick_coef = scalar.constant -1.702 : f32 + %sqrt_2_over_pi = scalar.constant 0.79788456080286544 : f32 + %sqrt_2_inv = scalar.constant 0.70710678118654745 : f32 + + %op_neg = index.constant 0 : index + %negated = scalar.negf %value : f32 + %is_neg = index.cmp eq, %op, %op_neg : index + %neg_selected = scf.select %is_neg, %negated, %value : f32 + + %op_abs = index.constant 1 : index + %abs_is_negative = scalar.cmpf olt, %value, %c0_f32 : f32 + %absolute = scf.select %abs_is_negative, %negated, %value : f32 + %is_abs = index.cmp eq, %op, %op_abs : index + %abs_selected = scf.select %is_abs, %absolute, %neg_selected : f32 + + %op_relu = index.constant 2 : index + %relu_is_positive = scalar.cmpf ogt, %value, %c0_f32 : f32 + %relu = scf.select %relu_is_positive, %value, %c0_f32 : f32 + %is_relu = index.cmp eq, %op, %op_relu : index + %relu_selected = scf.select %is_relu, %relu, %abs_selected : f32 + + %op_step = index.constant 3 : index + %step_is_positive = scalar.cmpf ogt, %value, %c0_f32 : f32 + %step = scf.select %step_is_positive, %c1_f32, %c0_f32 : f32 + %is_step = index.cmp eq, %op, %op_step : index + %step_selected = scf.select %is_step, %step, %relu_selected : f32 + + %op_sqr = index.constant 4 : index + %squared = scalar.mulf %value, %value : f32 + %is_sqr = index.cmp eq, %op, %op_sqr : index + %sqr_selected = scf.select %is_sqr, %squared, %step_selected : f32 + + %op_sgn = index.constant 5 : index + %sgn_is_negative = scalar.cmpf olt, %value, %c0_f32 : f32 + %sgn_is_positive = scalar.cmpf ogt, %value, %c0_f32 : f32 + %negative_sgn = scf.select %sgn_is_negative, %cn1_f32, %c0_f32 : f32 + %sgn = scf.select %sgn_is_positive, %c1_f32, %negative_sgn : f32 + %is_sgn = index.cmp eq, %op, %op_sgn : index + %sgn_selected = scf.select %is_sgn, %sgn, %sqr_selected : f32 + + %op_floor = index.constant 6 : index + %floored = scalar.floorf %value : f32 + %is_floor = index.cmp eq, %op, %op_floor : index + %floor_selected = scf.select %is_floor, %floored, %sgn_selected : f32 + + %op_ceil = index.constant 7 : index + %ceiled = scalar.ceilf %value : f32 + %is_ceil = index.cmp eq, %op, %op_ceil : index + %ceil_selected = scf.select %is_ceil, %ceiled, %floor_selected : f32 + + %op_round = index.constant 8 : index + %rounded = scalar.roundf %value : f32 + %is_round = index.cmp eq, %op, %op_round : index + %round_selected = scf.select %is_round, %rounded, %ceil_selected : f32 + + %op_trunc = index.constant 9 : index + %truncated = scalar.truncf %value : f32 + %is_trunc = index.cmp eq, %op, %op_trunc : index + %trunc_selected = scf.select %is_trunc, %truncated, %round_selected : f32 + + %op_tanh = index.constant 10 : index + %tanh_scaled = scalar.mulf %cn2_f32, %value : f32 + %tanh_exp = scalar.expf %tanh_scaled : f32 + %tanh_denominator = scalar.addf %c1_f32, %tanh_exp : f32 + %tanh_ratio = scalar.divf %c2_f32, %tanh_denominator : f32 + %tanh = scalar.subf %tanh_ratio, %c1_f32 : f32 + %is_tanh = index.cmp eq, %op, %op_tanh : index + %tanh_selected = scf.select %is_tanh, %tanh, %trunc_selected : f32 + + %op_elu = index.constant 11 : index + %elu_is_positive = scalar.cmpf ogt, %value, %c0_f32 : f32 + %elu_exp = scalar.expf %value : f32 + %elu_expm1 = scalar.subf %elu_exp, %c1_f32 : f32 + %elu = scf.select %elu_is_positive, %value, %elu_expm1 : f32 + %is_elu = index.cmp eq, %op, %op_elu : index + %elu_selected = scf.select %is_elu, %elu, %tanh_selected : f32 + + %op_sigmoid = index.constant 12 : index + %sigmoid_negative = scalar.negf %value : f32 + %sigmoid_exp = scalar.expf %sigmoid_negative : f32 + %sigmoid_denominator = scalar.addf %c1_f32, %sigmoid_exp : f32 + %sigmoid = scalar.divf %c1_f32, %sigmoid_denominator : f32 + %is_sigmoid = index.cmp eq, %op, %op_sigmoid : index + %sigmoid_selected = scf.select %is_sigmoid, %sigmoid, %elu_selected : f32 + + %op_gelu = index.constant 13 : index + %gelu_x2 = scalar.mulf %value, %value : f32 + %gelu_poly0 = scalar.mulf %gelu_coef, %gelu_x2 : f32 + %gelu_poly = scalar.addf %c1_f32, %gelu_poly0 : f32 + %gelu_inner0 = scalar.mulf %value, %gelu_poly : f32 + %gelu_inner = scalar.mulf %sqrt_2_over_pi, %gelu_inner0 : f32 + %gelu_tanh_scaled = scalar.mulf %cn2_f32, %gelu_inner : f32 + %gelu_tanh_exp = scalar.expf %gelu_tanh_scaled : f32 + %gelu_tanh_denominator = scalar.addf %c1_f32, %gelu_tanh_exp : f32 + %gelu_tanh_ratio = scalar.divf %c2_f32, %gelu_tanh_denominator : f32 + %gelu_tanh = scalar.subf %gelu_tanh_ratio, %c1_f32 : f32 + %gelu_one_plus = scalar.addf %c1_f32, %gelu_tanh : f32 + %gelu_half_x = scalar.mulf %c0_5_f32, %value : f32 + %gelu = scalar.mulf %gelu_half_x, %gelu_one_plus : f32 + %is_gelu = index.cmp eq, %op, %op_gelu : index + %gelu_selected = scf.select %is_gelu, %gelu, %sigmoid_selected : f32 + + %op_gelu_quick = index.constant 14 : index + %gelu_quick_scaled = scalar.mulf %gelu_quick_coef, %value : f32 + %gelu_quick_exp = scalar.expf %gelu_quick_scaled : f32 + %gelu_quick_denominator = scalar.addf %c1_f32, %gelu_quick_exp : f32 + %gelu_quick_sigmoid = scalar.divf %c1_f32, %gelu_quick_denominator : f32 + %gelu_quick = scalar.mulf %value, %gelu_quick_sigmoid : f32 + %is_gelu_quick = index.cmp eq, %op, %op_gelu_quick : index + %gelu_quick_selected = scf.select %is_gelu_quick, %gelu_quick, %gelu_selected : f32 + + %op_silu = index.constant 15 : index + %silu_negative = scalar.negf %value : f32 + %silu_exp = scalar.expf %silu_negative : f32 + %silu_denominator = scalar.addf %c1_f32, %silu_exp : f32 + %silu_sigmoid = scalar.divf %c1_f32, %silu_denominator : f32 + %silu = scalar.mulf %value, %silu_sigmoid : f32 + %is_silu = index.cmp eq, %op, %op_silu : index + %silu_selected = scf.select %is_silu, %silu, %gelu_quick_selected : f32 + + %op_hardswish = index.constant 16 : index + %hardswish_bias = scalar.addf %value, %c3_f32 : f32 + %hardswish_scaled = scalar.divf %hardswish_bias, %c6_f32 : f32 + %hardswish_above_zero = scalar.cmpf ogt, %hardswish_scaled, %c0_f32 : f32 + %hardswish_low = scf.select %hardswish_above_zero, %hardswish_scaled, %c0_f32 : f32 + %hardswish_above_one = scalar.cmpf ogt, %hardswish_low, %c1_f32 : f32 + %hardswish_gate = scf.select %hardswish_above_one, %c1_f32, %hardswish_low : f32 + %hardswish = scalar.mulf %value, %hardswish_gate : f32 + %is_hardswish = index.cmp eq, %op, %op_hardswish : index + %hardswish_selected = scf.select %is_hardswish, %hardswish, %silu_selected : f32 + + %op_hardsigmoid = index.constant 17 : index + %hardsigmoid_bias = scalar.addf %value, %c3_f32 : f32 + %hardsigmoid_scaled = scalar.divf %hardsigmoid_bias, %c6_f32 : f32 + %hardsigmoid_above_zero = scalar.cmpf ogt, %hardsigmoid_scaled, %c0_f32 : f32 + %hardsigmoid_low = scf.select %hardsigmoid_above_zero, %hardsigmoid_scaled, %c0_f32 : f32 + %hardsigmoid_above_one = scalar.cmpf ogt, %hardsigmoid_low, %c1_f32 : f32 + %hardsigmoid = scf.select %hardsigmoid_above_one, %c1_f32, %hardsigmoid_low : f32 + %is_hardsigmoid = index.cmp eq, %op, %op_hardsigmoid : index + %hardsigmoid_selected = scf.select %is_hardsigmoid, %hardsigmoid, %hardswish_selected : f32 + + %op_exp = index.constant 18 : index + %exp_value = scalar.expf %value : f32 + %is_exp = index.cmp eq, %op, %op_exp : index + %exp_selected = scf.select %is_exp, %exp_value, %hardsigmoid_selected : f32 + + %op_expm1 = index.constant 19 : index + %expm1_exp = scalar.expf %value : f32 + %expm1 = scalar.subf %expm1_exp, %c1_f32 : f32 + %is_expm1 = index.cmp eq, %op, %op_expm1 : index + %expm1_selected = scf.select %is_expm1, %expm1, %exp_selected : f32 + + %op_gelu_erf = index.constant 20 : index + %gelu_erf_scaled = scalar.mulf %sqrt_2_inv, %value : f32 + %gelu_erf_value = scalar.erff %gelu_erf_scaled : f32 + %gelu_erf_one_plus = scalar.addf %c1_f32, %gelu_erf_value : f32 + %gelu_erf_half_x = scalar.mulf %c0_5_f32, %value : f32 + %gelu_erf = scalar.mulf %gelu_erf_half_x, %gelu_erf_one_plus : f32 + %is_gelu_erf = index.cmp eq, %op, %op_gelu_erf : index + %gelu_erf_selected = scf.select %is_gelu_erf, %gelu_erf, %expm1_selected : f32 + + %op_sqrt = index.constant 21 : index + %sqrt = scalar.sqrtf %value : f32 + %is_sqrt = index.cmp eq, %op, %op_sqrt : index + %sqrt_selected = scf.select %is_sqrt, %sqrt, %gelu_erf_selected : f32 + + %op_log = index.constant 22 : index + %log = scalar.logf %value : f32 + %is_log = index.cmp eq, %op, %op_log : index + %log_selected = scf.select %is_log, %log, %sqrt_selected : f32 + + %op_identity = index.constant 23 : index + %is_identity = index.cmp eq, %op, %op_identity : index + %result = scf.select %is_identity, %value, %log_selected : f32 + template.return %result : f32 +} + +template.def<@ggml.unary_f32.apply_vector4> device priority(1) @ggml_unary_f32_apply_vector4(%op: index, %values: vector<4xf32>) -> (vector<4xf32>) { + %c0_f32 = vector.constant 0.0 : vector<4xf32> + %c0_5_f32 = vector.constant 0.5 : vector<4xf32> + %c1_f32 = vector.constant 1.0 : vector<4xf32> + %c2_f32 = vector.constant 2.0 : vector<4xf32> + %c3_f32 = vector.constant 3.0 : vector<4xf32> + %c6_f32 = vector.constant 6.0 : vector<4xf32> + %cn1_f32 = vector.constant -1.0 : vector<4xf32> + %cn2_f32 = vector.constant -2.0 : vector<4xf32> + %gelu_coef = vector.constant 0.044715 : vector<4xf32> + %gelu_quick_coef = vector.constant -1.702 : vector<4xf32> + %sqrt_2_over_pi = vector.constant 0.79788456080286544 : vector<4xf32> + %sqrt_2_inv = vector.constant 0.70710678118654745 : vector<4xf32> + + %op_neg = index.constant 0 : index + %negated = vector.negf %values : vector<4xf32> + %is_neg = index.cmp eq, %op, %op_neg : index + %neg_selected = scf.select %is_neg, %negated, %values : vector<4xf32> + + %op_abs = index.constant 1 : index + %abs_is_negative = vector.cmpf olt, %values, %c0_f32 : vector<4xf32> -> vector<4xi1> + %absolute = vector.select %abs_is_negative, %negated, %values : vector<4xf32> + %is_abs = index.cmp eq, %op, %op_abs : index + %abs_selected = scf.select %is_abs, %absolute, %neg_selected : vector<4xf32> + + %op_relu = index.constant 2 : index + %relu_is_positive = vector.cmpf ogt, %values, %c0_f32 : vector<4xf32> -> vector<4xi1> + %relu = vector.select %relu_is_positive, %values, %c0_f32 : vector<4xf32> + %is_relu = index.cmp eq, %op, %op_relu : index + %relu_selected = scf.select %is_relu, %relu, %abs_selected : vector<4xf32> + + %op_step = index.constant 3 : index + %step_is_positive = vector.cmpf ogt, %values, %c0_f32 : vector<4xf32> -> vector<4xi1> + %step = vector.select %step_is_positive, %c1_f32, %c0_f32 : vector<4xf32> + %is_step = index.cmp eq, %op, %op_step : index + %step_selected = scf.select %is_step, %step, %relu_selected : vector<4xf32> + + %op_sqr = index.constant 4 : index + %squared = vector.mulf %values, %values : vector<4xf32> + %is_sqr = index.cmp eq, %op, %op_sqr : index + %sqr_selected = scf.select %is_sqr, %squared, %step_selected : vector<4xf32> + + %op_sgn = index.constant 5 : index + %sgn_is_negative = vector.cmpf olt, %values, %c0_f32 : vector<4xf32> -> vector<4xi1> + %sgn_is_positive = vector.cmpf ogt, %values, %c0_f32 : vector<4xf32> -> vector<4xi1> + %negative_sgn = vector.select %sgn_is_negative, %cn1_f32, %c0_f32 : vector<4xf32> + %sgn = vector.select %sgn_is_positive, %c1_f32, %negative_sgn : vector<4xf32> + %is_sgn = index.cmp eq, %op, %op_sgn : index + %sgn_selected = scf.select %is_sgn, %sgn, %sqr_selected : vector<4xf32> + + %op_floor = index.constant 6 : index + %floored = vector.floorf %values : vector<4xf32> + %is_floor = index.cmp eq, %op, %op_floor : index + %floor_selected = scf.select %is_floor, %floored, %sgn_selected : vector<4xf32> + + %op_ceil = index.constant 7 : index + %ceiled = vector.ceilf %values : vector<4xf32> + %is_ceil = index.cmp eq, %op, %op_ceil : index + %ceil_selected = scf.select %is_ceil, %ceiled, %floor_selected : vector<4xf32> + + %op_round = index.constant 8 : index + %rounded = vector.roundf %values : vector<4xf32> + %is_round = index.cmp eq, %op, %op_round : index + %round_selected = scf.select %is_round, %rounded, %ceil_selected : vector<4xf32> + + %op_trunc = index.constant 9 : index + %truncated = vector.truncf %values : vector<4xf32> + %is_trunc = index.cmp eq, %op, %op_trunc : index + %trunc_selected = scf.select %is_trunc, %truncated, %round_selected : vector<4xf32> + + %op_tanh = index.constant 10 : index + %tanh_scaled = vector.mulf %cn2_f32, %values : vector<4xf32> + %tanh_exp = vector.expf %tanh_scaled : vector<4xf32> + %tanh_denominator = vector.addf %c1_f32, %tanh_exp : vector<4xf32> + %tanh_ratio = vector.divf %c2_f32, %tanh_denominator : vector<4xf32> + %tanh = vector.subf %tanh_ratio, %c1_f32 : vector<4xf32> + %is_tanh = index.cmp eq, %op, %op_tanh : index + %tanh_selected = scf.select %is_tanh, %tanh, %trunc_selected : vector<4xf32> + + %op_elu = index.constant 11 : index + %elu_is_positive = vector.cmpf ogt, %values, %c0_f32 : vector<4xf32> -> vector<4xi1> + %elu_exp = vector.expf %values : vector<4xf32> + %elu_expm1 = vector.subf %elu_exp, %c1_f32 : vector<4xf32> + %elu = vector.select %elu_is_positive, %values, %elu_expm1 : vector<4xf32> + %is_elu = index.cmp eq, %op, %op_elu : index + %elu_selected = scf.select %is_elu, %elu, %tanh_selected : vector<4xf32> + + %op_sigmoid = index.constant 12 : index + %sigmoid_negative = vector.negf %values : vector<4xf32> + %sigmoid_exp = vector.expf %sigmoid_negative : vector<4xf32> + %sigmoid_denominator = vector.addf %c1_f32, %sigmoid_exp : vector<4xf32> + %sigmoid = vector.divf %c1_f32, %sigmoid_denominator : vector<4xf32> + %is_sigmoid = index.cmp eq, %op, %op_sigmoid : index + %sigmoid_selected = scf.select %is_sigmoid, %sigmoid, %elu_selected : vector<4xf32> + + %op_gelu = index.constant 13 : index + %gelu_x2 = vector.mulf %values, %values : vector<4xf32> + %gelu_poly0 = vector.mulf %gelu_coef, %gelu_x2 : vector<4xf32> + %gelu_poly = vector.addf %c1_f32, %gelu_poly0 : vector<4xf32> + %gelu_inner0 = vector.mulf %values, %gelu_poly : vector<4xf32> + %gelu_inner = vector.mulf %sqrt_2_over_pi, %gelu_inner0 : vector<4xf32> + %gelu_tanh_scaled = vector.mulf %cn2_f32, %gelu_inner : vector<4xf32> + %gelu_tanh_exp = vector.expf %gelu_tanh_scaled : vector<4xf32> + %gelu_tanh_denominator = vector.addf %c1_f32, %gelu_tanh_exp : vector<4xf32> + %gelu_tanh_ratio = vector.divf %c2_f32, %gelu_tanh_denominator : vector<4xf32> + %gelu_tanh = vector.subf %gelu_tanh_ratio, %c1_f32 : vector<4xf32> + %gelu_one_plus = vector.addf %c1_f32, %gelu_tanh : vector<4xf32> + %gelu_half_x = vector.mulf %c0_5_f32, %values : vector<4xf32> + %gelu = vector.mulf %gelu_half_x, %gelu_one_plus : vector<4xf32> + %is_gelu = index.cmp eq, %op, %op_gelu : index + %gelu_selected = scf.select %is_gelu, %gelu, %sigmoid_selected : vector<4xf32> + + %op_gelu_quick = index.constant 14 : index + %gelu_quick_scaled = vector.mulf %gelu_quick_coef, %values : vector<4xf32> + %gelu_quick_exp = vector.expf %gelu_quick_scaled : vector<4xf32> + %gelu_quick_denominator = vector.addf %c1_f32, %gelu_quick_exp : vector<4xf32> + %gelu_quick_sigmoid = vector.divf %c1_f32, %gelu_quick_denominator : vector<4xf32> + %gelu_quick = vector.mulf %values, %gelu_quick_sigmoid : vector<4xf32> + %is_gelu_quick = index.cmp eq, %op, %op_gelu_quick : index + %gelu_quick_selected = scf.select %is_gelu_quick, %gelu_quick, %gelu_selected : vector<4xf32> + + %op_silu = index.constant 15 : index + %silu_negative = vector.negf %values : vector<4xf32> + %silu_exp = vector.expf %silu_negative : vector<4xf32> + %silu_denominator = vector.addf %c1_f32, %silu_exp : vector<4xf32> + %silu_sigmoid = vector.divf %c1_f32, %silu_denominator : vector<4xf32> + %silu = vector.mulf %values, %silu_sigmoid : vector<4xf32> + %is_silu = index.cmp eq, %op, %op_silu : index + %silu_selected = scf.select %is_silu, %silu, %gelu_quick_selected : vector<4xf32> + + %op_hardswish = index.constant 16 : index + %hardswish_bias = vector.addf %values, %c3_f32 : vector<4xf32> + %hardswish_scaled = vector.divf %hardswish_bias, %c6_f32 : vector<4xf32> + %hardswish_above_zero = vector.cmpf ogt, %hardswish_scaled, %c0_f32 : vector<4xf32> -> vector<4xi1> + %hardswish_low = vector.select %hardswish_above_zero, %hardswish_scaled, %c0_f32 : vector<4xf32> + %hardswish_above_one = vector.cmpf ogt, %hardswish_low, %c1_f32 : vector<4xf32> -> vector<4xi1> + %hardswish_gate = vector.select %hardswish_above_one, %c1_f32, %hardswish_low : vector<4xf32> + %hardswish = vector.mulf %values, %hardswish_gate : vector<4xf32> + %is_hardswish = index.cmp eq, %op, %op_hardswish : index + %hardswish_selected = scf.select %is_hardswish, %hardswish, %silu_selected : vector<4xf32> + + %op_hardsigmoid = index.constant 17 : index + %hardsigmoid_bias = vector.addf %values, %c3_f32 : vector<4xf32> + %hardsigmoid_scaled = vector.divf %hardsigmoid_bias, %c6_f32 : vector<4xf32> + %hardsigmoid_above_zero = vector.cmpf ogt, %hardsigmoid_scaled, %c0_f32 : vector<4xf32> -> vector<4xi1> + %hardsigmoid_low = vector.select %hardsigmoid_above_zero, %hardsigmoid_scaled, %c0_f32 : vector<4xf32> + %hardsigmoid_above_one = vector.cmpf ogt, %hardsigmoid_low, %c1_f32 : vector<4xf32> -> vector<4xi1> + %hardsigmoid = vector.select %hardsigmoid_above_one, %c1_f32, %hardsigmoid_low : vector<4xf32> + %is_hardsigmoid = index.cmp eq, %op, %op_hardsigmoid : index + %hardsigmoid_selected = scf.select %is_hardsigmoid, %hardsigmoid, %hardswish_selected : vector<4xf32> + + %op_exp = index.constant 18 : index + %exp_value = vector.expf %values : vector<4xf32> + %is_exp = index.cmp eq, %op, %op_exp : index + %exp_selected = scf.select %is_exp, %exp_value, %hardsigmoid_selected : vector<4xf32> + + %op_expm1 = index.constant 19 : index + %expm1_exp = vector.expf %values : vector<4xf32> + %expm1 = vector.subf %expm1_exp, %c1_f32 : vector<4xf32> + %is_expm1 = index.cmp eq, %op, %op_expm1 : index + %expm1_selected = scf.select %is_expm1, %expm1, %exp_selected : vector<4xf32> + + %op_gelu_erf = index.constant 20 : index + %gelu_erf_scaled = vector.mulf %sqrt_2_inv, %values : vector<4xf32> + %gelu_erf_value = vector.erff %gelu_erf_scaled : vector<4xf32> + %gelu_erf_one_plus = vector.addf %c1_f32, %gelu_erf_value : vector<4xf32> + %gelu_erf_half_x = vector.mulf %c0_5_f32, %values : vector<4xf32> + %gelu_erf = vector.mulf %gelu_erf_half_x, %gelu_erf_one_plus : vector<4xf32> + %is_gelu_erf = index.cmp eq, %op, %op_gelu_erf : index + %gelu_erf_selected = scf.select %is_gelu_erf, %gelu_erf, %expm1_selected : vector<4xf32> + + %op_sqrt = index.constant 21 : index + %sqrt = vector.sqrtf %values : vector<4xf32> + %is_sqrt = index.cmp eq, %op, %op_sqrt : index + %sqrt_selected = scf.select %is_sqrt, %sqrt, %gelu_erf_selected : vector<4xf32> + + %op_log = index.constant 22 : index + %log = vector.logf %values : vector<4xf32> + %is_log = index.cmp eq, %op, %op_log : index + %log_selected = scf.select %is_log, %log, %sqrt_selected : vector<4xf32> + + %op_identity = index.constant 23 : index + %is_identity = index.cmp eq, %op, %op_identity : index + %result = scf.select %is_identity, %values, %log_selected : vector<4xf32> + template.return %result : vector<4xf32> +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/add_id_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/add_id_f32.loom new file mode 100644 index 000000000000..00ff391096c9 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/add_id_f32.loom @@ -0,0 +1,107 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// GGML_OP_ADD_ID on F32 (gpt-oss's per-expert biases after MUL_MAT_ID): +// output[t][r][i] = input[t][r][i] + bias[ids[t][r]][i] +// input / output are contiguous [tokens][rows][width], bias is [expert_count][width], ids is [tokens][ids_stride] i32 +// with the first `rows` entries of each token used (llama.cpp passes a view of the argsort, so ids_stride is the +// expert count). One workgroup per (row, token); an id outside 0..expert_count-1 reads expert 0. + +amdgpu.target @ggml_add_id_gfx11_wave32 {subgroup_size = 32} + +kernel.def target(@ggml_add_id_gfx11_wave32) export("ggml_add_id_f32") @ggml_add_id_f32(%width: index, %rows: index, %tokens: index, %ids_stride: index, %expert_count: index) { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%rows, %tokens, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%width: index, %rows: index, %tokens: index, %ids_stride: index, %expert_count: index, %input: buffer, %bias: buffer, %ids: buffer, %output: buffer) where [range(%width, 1, 1048576), range(%rows, 1, 4096), range(%tokens, 1, 65536), range(%ids_stride, 1, 4096), range(%expert_count, 1, 4096)] { + %c0 = index.constant 0 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %row0 = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %row, %launch_rows = index.assume %row0, %rows [lt(%row0, %rows)] : index, index + %token, %launch_tokens = index.assume %token0, %tokens [lt(%token0, %tokens)] : index, index + %ids_row, %launch_ids_stride = index.assume %row, %ids_stride [lt(%row, %ids_stride)] : index, index + %input_noalias, %bias_noalias, %ids_noalias, %output_noalias = buffer.assume.noalias %input, %bias, %ids, %output : buffer, buffer, buffer, buffer + %ids_view = buffer.view %ids_noalias[%c0_offset] : buffer -> view<[%launch_tokens]x[%launch_ids_stride]xi32> + %expert_i32 = view.load %ids_view[%token, %ids_row] : view<[%launch_tokens]x[%launch_ids_stride]xi32> -> i32 + %expert_raw = index.cast %expert_i32 : i32 to index + %expert_valid = index.cmp ult, %expert_raw, %expert_count : index + %expert0 = scf.select %expert_valid, %expert_raw, %c0 : index + %expert, %launch_experts = index.assume %expert0, %expert_count [lt(%expert0, %expert_count)] : index, index + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%launch_tokens]x[%launch_rows]x[%width]xf32> + %bias_view = buffer.view %bias_noalias[%c0_offset] : buffer -> view<[%launch_experts]x[%width]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_tokens]x[%launch_rows]x[%width]xf32> + scf.for %i = [%workitem to %width step %c256] { + %a = view.load %input_view[%token, %row, %i] : view<[%launch_tokens]x[%launch_rows]x[%width]xf32> -> f32 + %b = view.load %bias_view[%expert, %i] : view<[%launch_experts]x[%width]xf32> -> f32 + %sum = scalar.addf %a, %b : f32 + view.store %sum, %output_view[%token, %row, %i] : f32, view<[%launch_tokens]x[%launch_rows]x[%width]xf32> + } + kernel.return +} + +// Reference for the cases: one workitem per output value, from flat buffers. +kernel.def target(@ggml_add_id_gfx11_wave32) export("ggml_add_id_reference_f32") @ggml_add_id_reference_f32(%width: index, %rows: index, %tokens: index, %ids_stride: index, %expert_count: index) { + %c1 = index.constant 1 : index + kernel.launch.config workgroups(%width, %rows, %tokens) workgroup_size(%c1, %c1, %c1) : index +} launch(%width: index, %rows: index, %tokens: index, %ids_stride: index, %expert_count: index, %input: buffer, %bias: buffer, %ids: buffer, %output: buffer) where [range(%width, 1, 1048576), range(%rows, 1, 4096), range(%tokens, 1, 65536), range(%ids_stride, 1, 4096), range(%expert_count, 1, 4096)] { + %c0 = index.constant 0 : index + %c0_offset = index.constant 0 : offset + %i0 = kernel.workgroup.id : index + %row0 = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %i, %launch_width = index.assume %i0, %width [lt(%i0, %width)] : index, index + %row, %launch_rows = index.assume %row0, %rows [lt(%row0, %rows)] : index, index + %token, %launch_tokens = index.assume %token0, %tokens [lt(%token0, %tokens)] : index, index + %ids_row, %launch_ids_stride = index.assume %row, %ids_stride [lt(%row, %ids_stride)] : index, index + %ids_view = buffer.view %ids[%c0_offset] : buffer -> view<[%launch_tokens]x[%launch_ids_stride]xi32> + %expert_i32 = view.load %ids_view[%token, %ids_row] : view<[%launch_tokens]x[%launch_ids_stride]xi32> -> i32 + %expert_raw = index.cast %expert_i32 : i32 to index + %expert_valid = index.cmp ult, %expert_raw, %expert_count : index + %expert0 = scf.select %expert_valid, %expert_raw, %c0 : index + %expert, %launch_experts = index.assume %expert0, %expert_count [lt(%expert0, %expert_count)] : index, index + %input_view = buffer.view %input[%c0_offset] : buffer -> view<[%launch_tokens]x[%launch_rows]x[%launch_width]xf32> + %bias_view = buffer.view %bias[%c0_offset] : buffer -> view<[%launch_experts]x[%launch_width]xf32> + %output_view = buffer.view %output[%c0_offset] : buffer -> view<[%launch_tokens]x[%launch_rows]x[%launch_width]xf32> + %a = view.load %input_view[%token, %row, %i] : view<[%launch_tokens]x[%launch_rows]x[%launch_width]xf32> -> f32 + %b = view.load %bias_view[%expert, %i] : view<[%launch_experts]x[%launch_width]xf32> -> f32 + %sum = scalar.addf %a, %b : f32 + view.store %sum, %output_view[%token, %row, %i] : f32, view<[%launch_tokens]x[%launch_rows]x[%launch_width]xf32> + kernel.return +} + +// gpt-oss shape in small: width 300 (not a multiple of 256: the strided loop's last pass is partial), 4 of 8 +// experts per token through an ids view of stride 8, 3 tokens. +check.case public @ggml_add_id_f32_case { + %input_seed = check.param.seed base(7300000000000066001) count(1) : i64 + %bias_seed = check.param.seed base(7300000000000066002) count(1) : i64 + %ids_seed = check.param.seed base(7300000000000066003) count(1) : i64 + %width = check.literal value(300) : index + %rows = check.literal value(4) : index + %tokens = check.literal value(3) : index + %ids_stride = check.literal value(8) : index + %experts = check.literal value(8) : index + %input = check.generate.random.uniform seed(%input_seed) range(-4.0 to 4.0) : tensor<3600xf32> + %bias = check.generate.random.uniform seed(%bias_seed) range(-4.0 to 4.0) : tensor<2400xf32> + %ids = check.generate.random.uniform seed(%ids_seed) range(0 to 7) : tensor<24xi32> + %output = check.generate.fill value(-7.0) : tensor<3600xf32> + %expected = check.generate.fill value(7.0) : tensor<3600xf32> + kernel.launch @ggml_add_id_reference_f32[%width, %rows, %tokens, %ids_stride, %experts](%width, %rows, %tokens, %ids_stride, %experts, %input, %bias, %ids, %expected) : [index, index, index, index, index](index, index, index, index, index, tensor<3600xf32>, tensor<2400xf32>, tensor<24xi32>, tensor<3600xf32>) + kernel.launch @ggml_add_id_f32[%width, %rows, %tokens, %ids_stride, %experts](%width, %rows, %tokens, %ids_stride, %experts, %input, %bias, %ids, %output) : [index, index, index, index, index](index, index, index, index, index, tensor<3600xf32>, tensor<2400xf32>, tensor<24xi32>, tensor<3600xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<3600xf32> + check.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/attention_sink_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/attention_sink_f32.loom new file mode 100644 index 000000000000..8c477cef1610 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/attention_sink_f32.loom @@ -0,0 +1,329 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Attention sinks (gpt-oss) applied after FLASH_ATTN_EXT, in place on its output. A sink adds one logit per head to +// the softmax denominator and nothing to the numerator, so with M the row maximum and S the row sum of +// exp(scale q.k + mask - M): +// output_with_sink = output_without_sink * S / (S + exp(sink - M)) +// This kernel recomputes M and S for each (query token, query head) row with FlashAttention's conventions (scores +// scale * q.k + mask, mask entries below -1e30, i.e. -inf, skipped) and rescales the row. A row with every key +// masked has S = 0 and is written as zeros. +// Layouts (as the HRX FlashAttention matchers require): query [tokens][query_heads][qk_head_size] f32, key +// [key_capacity][key_value_heads][qk_head_size] f16, mask [tokens][key_count] f16, sinks [query_heads] f32, output +// [tokens][query_heads][value_head_size] f32. Query head h reads key-value head h / (query_heads / key_value_heads). + +amdgpu.target @ggml_attention_sink_gfx11_wave32 {subgroup_size = 32} + +config.decl @ggml.attention_sink.query_head_count : %value: index where [range(%value, 1, 256)] + +config.decl @ggml.attention_sink.key_value_head_count : %value: index where [range(%value, 1, 256)] + +config.decl @ggml.attention_sink.qk_head_size : %value: index where [range(%value, 16, 576), mul(%value, 16)] + +config.decl @ggml.attention_sink.value_head_size : %value: index where [range(%value, 16, 576), mul(%value, 16)] + +config.decl @ggml.attention_sink.scale : f32 + +kernel.def target(@ggml_attention_sink_gfx11_wave32) export("ggml_attention_sink_f32") @ggml_attention_sink_f32(%tokens: index, %key_count: index, %key_capacity: index) { + %query_head_count = config.get @ggml.attention_sink.query_head_count : index + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%query_head_count, %tokens, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%tokens: index, %key_count: index, %key_capacity: index, %query: buffer, %key: buffer, %mask: buffer, %sinks: buffer, %output: buffer) where [range(%tokens, 1, 65536), range(%key_count, 1, 1048576), range(%key_capacity, 1, 1048576)] { + %query_head_count0 = config.get @ggml.attention_sink.query_head_count : index + %key_value_head_count0 = config.get @ggml.attention_sink.key_value_head_count : index + %qk_head_size0 = config.get @ggml.attention_sink.qk_head_size : index + %value_head_size0 = config.get @ggml.attention_sink.value_head_size : index + %scale = config.get @ggml.attention_sink.scale : f32 + %query_head_count = index.assume %query_head_count0 [range(%query_head_count0, 1, 256)] : index + %key_value_head_count = index.assume %key_value_head_count0 [range(%key_value_head_count0, 1, 256)] : index + %qk_head_size = index.assume %qk_head_size0 [range(%qk_head_size0, 16, 576), mul(%qk_head_size0, 16)] : index + %value_head_size = index.assume %value_head_size0 [range(%value_head_size0, 16, 576), mul(%value_head_size0, 16)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %stage_bytes = index.constant 64 : offset + %zero = scalar.constant 0.0 : f32 + %one = scalar.constant 1.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %head0 = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %subgroup = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %head, %launch_heads = index.assume %head0, %query_head_count [lt(%head0, %query_head_count)] : index, index + %token, %launch_tokens = index.assume %token0, %tokens [lt(%token0, %tokens)] : index, index + %heads_per_key_value_head0 = index.div %query_head_count, %key_value_head_count : index + %heads_per_key_value_head = index.assume %heads_per_key_value_head0 [range(%heads_per_key_value_head0, 1, 256)] : index + %key_value_head0 = index.div %head, %heads_per_key_value_head : index + %key_value_head, %launch_key_value_heads = index.assume %key_value_head0, %key_value_head_count [lt(%key_value_head0, %key_value_head_count)] : index, index + %bounded_key_count, %launch_key_capacity = index.assume %key_count, %key_capacity [le(%key_count, %key_capacity)] : index, index + %query_view = buffer.view %query[%c0_offset] : buffer -> view<[%launch_tokens]x[%launch_heads]x[%qk_head_size]xf32> + %key_view = buffer.view %key[%c0_offset] : buffer -> view<[%launch_key_capacity]x[%launch_key_value_heads]x[%qk_head_size]xf16> + %mask_view = buffer.view %mask[%c0_offset] : buffer -> view<[%launch_tokens]x[%bounded_key_count]xf16> + %sink_view = buffer.view %sinks[%c0_offset] : buffer -> view<[%launch_heads]xf32> + %output_view = buffer.view %output[%c0_offset] : buffer -> view<[%launch_tokens]x[%launch_heads]x[%value_head_size]xf32> + %stage = buffer.alloca align(16) %stage_bytes : buffer + %stage_view = buffer.view %stage[%c0_offset] : buffer -> view<16xf32> + // Per workitem: online maximum and sum over keys workitem, workitem + 256, ... + %local_max, %local_sum = scf.for %key_index = [%workitem to %bounded_key_count step %c256](%running_max = %negative_large : f32, %running_sum = %zero : f32) -> (f32, f32) { + %mask_f16 = view.load %mask_view[%token, %key_index] : view<[%launch_tokens]x[%bounded_key_count]xf16> -> f16 + %mask_value = scalar.extf %mask_f16 : f16 to f32 + %active = scalar.cmpf ogt, %mask_value, %negative_large : f32 + %next_max, %next_sum = scf.if %active -> (f32, f32) { + %key_row, %row_bound = index.assume %key_index, %key_capacity [lt(%key_index, %key_capacity)] : index, index + %dot = scf.for %channel = [%c0 to %qk_head_size step %c4](%acc = %zero : f32) -> (f32) { + %q = vector.load %query_view[%token, %head, %channel] : view<[%launch_tokens]x[%launch_heads]x[%qk_head_size]xf32> -> vector<4xf32> + %k16 = vector.load %key_view[%key_row, %key_value_head, %channel] : view<[%launch_key_capacity]x[%launch_key_value_heads]x[%qk_head_size]xf16> -> vector<4xf16> + %k = vector.extf %k16 : vector<4xf16> to vector<4xf32> + %qk = vector.mulf %q, %k : vector<4xf32> + %sum4 = vector.reduce %qk, %acc : vector<4xf32>, f32 + scf.yield %sum4 : f32 + } + %scaled = scalar.mulf %dot, %scale : f32 + %score = scalar.addf %scaled, %mask_value : f32 + %grows = scalar.cmpf ogt, %score, %running_max : f32 + %new_max = scf.select %grows, %score, %running_max : f32 + %old_shift = scalar.subf %running_max, %new_max : f32 + %old_scale = scalar.expf %old_shift : f32 + %score_shift = scalar.subf %score, %new_max : f32 + %score_exp = scalar.expf %score_shift : f32 + %kept = scalar.mulf %running_sum, %old_scale : f32 + %new_sum = scalar.addf %kept, %score_exp : f32 + scf.yield %new_max, %new_sum : f32, f32 + } else { + scf.yield %running_max, %running_sum : f32, f32 + } + scf.yield %next_max, %next_sum : f32, f32 + } + // Row maximum: subgroup reduce, then across the eight subgroups through workgroup memory. + %subgroup_max = kernel.subgroup.reduce %local_max : f32 + %lane_is_zero = index.cmp eq, %lane, %c0 : index + %subgroup_slot0 = index.assume %subgroup [range(%subgroup, 0, 7)] : index + scf.if %lane_is_zero { + view.store %subgroup_max, %stage_view[%subgroup_slot0] : f32, view<16xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %row_max = scf.for %slot = [%c0 to %c8 step %c1](%acc_max = %negative_large : f32) -> (f32) { + %slot_max = view.load %stage_view[%slot] : view<16xf32> -> f32 + %m = scalar.maxnumf %acc_max, %slot_max : f32 + scf.yield %m : f32 + } + // Row sum relative to the row maximum. + %local_shift = scalar.subf %local_max, %row_max : f32 + %local_rescale = scalar.expf %local_shift : f32 + %local_rescaled = scalar.mulf %local_sum, %local_rescale : f32 + %subgroup_sum = kernel.subgroup.reduce %local_rescaled : f32 + %subgroup_slot1 = index.add %subgroup_slot0, %c8 : index + scf.if %lane_is_zero { + view.store %subgroup_sum, %stage_view[%subgroup_slot1] : f32, view<16xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %row_sum_total = scf.for %slot0 = [%c0 to %c8 step %c1](%acc_s = %zero : f32) -> (f32) { + %slot = index.add %slot0, %c8 : index + %slot_sum = view.load %stage_view[%slot] : view<16xf32> -> f32 + %s = scalar.addf %acc_s, %slot_sum : f32 + scf.yield %s : f32 + } + %sink = view.load %sink_view[%head] : view<[%launch_heads]xf32> -> f32 + %sink_shift = scalar.subf %sink, %row_max : f32 + %sink_exp = scalar.expf %sink_shift : f32 + %denominator = scalar.addf %row_sum_total, %sink_exp : f32 + %factor0 = scalar.divf %row_sum_total, %denominator : f32 + %any_key = scalar.cmpf ogt, %row_sum_total, %zero : f32 + %factor = scf.select %any_key, %factor0, %zero : f32 + scf.for %channel = [%workitem to %value_head_size step %c256] { + %value = view.load %output_view[%token, %head, %channel] : view<[%launch_tokens]x[%launch_heads]x[%value_head_size]xf32> -> f32 + %scaled_value = scalar.mulf %value, %factor : f32 + %result = scf.select %any_key, %scaled_value, %zero : f32 + view.store %result, %output_view[%token, %head, %channel] : f32, view<[%launch_tokens]x[%launch_heads]x[%value_head_size]xf32> + } + kernel.return +} + +// Reference for the cases: full attention for one (query head, token) row per workitem, with (use_sink = 1) or +// without (use_sink = 0) the sink in the denominator. value is [key_capacity][key_value_heads][value_head_size] f16. +kernel.def target(@ggml_attention_sink_gfx11_wave32) export("ggml_attention_sink_reference_f32") @ggml_attention_sink_reference_f32(%tokens: index, %key_count: index, %key_capacity: index, %use_sink: index) { + %query_head_count = config.get @ggml.attention_sink.query_head_count : index + %c1 = index.constant 1 : index + kernel.launch.config workgroups(%query_head_count, %tokens, %c1) workgroup_size(%c1, %c1, %c1) : index +} launch(%tokens: index, %key_count: index, %key_capacity: index, %use_sink: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %sinks: buffer, %output: buffer) where [range(%tokens, 1, 65536), range(%key_count, 1, 1048576), range(%key_capacity, 1, 1048576), range(%use_sink, 0, 1)] { + %query_head_count0 = config.get @ggml.attention_sink.query_head_count : index + %key_value_head_count0 = config.get @ggml.attention_sink.key_value_head_count : index + %qk_head_size0 = config.get @ggml.attention_sink.qk_head_size : index + %value_head_size0 = config.get @ggml.attention_sink.value_head_size : index + %scale = config.get @ggml.attention_sink.scale : f32 + %query_head_count = index.assume %query_head_count0 [range(%query_head_count0, 1, 256)] : index + %key_value_head_count = index.assume %key_value_head_count0 [range(%key_value_head_count0, 1, 256)] : index + %qk_head_size = index.assume %qk_head_size0 [range(%qk_head_size0, 16, 576), mul(%qk_head_size0, 16)] : index + %value_head_size = index.assume %value_head_size0 [range(%value_head_size0, 16, 576), mul(%value_head_size0, 16)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %zero = scalar.constant 0.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %head0 = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %head, %launch_heads = index.assume %head0, %query_head_count [lt(%head0, %query_head_count)] : index, index + %token, %launch_tokens = index.assume %token0, %tokens [lt(%token0, %tokens)] : index, index + %ratio0 = index.div %query_head_count, %key_value_head_count : index + %ratio = index.assume %ratio0 [range(%ratio0, 1, 256)] : index + %kv_head0 = index.div %head, %ratio : index + %kv_head, %launch_kv_heads = index.assume %kv_head0, %key_value_head_count [lt(%kv_head0, %key_value_head_count)] : index, index + %count, %capacity = index.assume %key_count, %key_capacity [le(%key_count, %key_capacity)] : index, index + %qv = buffer.view %query[%c0_offset] : buffer -> view<[%launch_tokens]x[%launch_heads]x[%qk_head_size]xf32> + %kv = buffer.view %key[%c0_offset] : buffer -> view<[%capacity]x[%launch_kv_heads]x[%qk_head_size]xf16> + %vv = buffer.view %value[%c0_offset] : buffer -> view<[%capacity]x[%launch_kv_heads]x[%value_head_size]xf16> + %mv = buffer.view %mask[%c0_offset] : buffer -> view<[%launch_tokens]x[%count]xf16> + %sv = buffer.view %sinks[%c0_offset] : buffer -> view<[%launch_heads]xf32> + %ov = buffer.view %output[%c0_offset] : buffer -> view<[%launch_tokens]x[%launch_heads]x[%value_head_size]xf32> + %row_max = scf.for %j = [%c0 to %count step %c1](%m = %negative_large : f32) -> (f32) { + %jk, %jcap = index.assume %j, %key_capacity [lt(%j, %key_capacity)] : index, index + %mk16 = view.load %mv[%token, %j] : view<[%launch_tokens]x[%count]xf16> -> f16 + %mk = scalar.extf %mk16 : f16 to f32 + %dot = scf.for %c = [%c0 to %qk_head_size step %c1](%a = %zero : f32) -> (f32) { + %q = view.load %qv[%token, %head, %c] : view<[%launch_tokens]x[%launch_heads]x[%qk_head_size]xf32> -> f32 + %k16 = view.load %kv[%jk, %kv_head, %c] : view<[%capacity]x[%launch_kv_heads]x[%qk_head_size]xf16> -> f16 + %k = scalar.extf %k16 : f16 to f32 + %p = scalar.mulf %q, %k : f32 + %a1 = scalar.addf %a, %p : f32 + scf.yield %a1 : f32 + } + %sd = scalar.mulf %dot, %scale : f32 + %s = scalar.addf %sd, %mk : f32 + %live = scalar.cmpf ogt, %mk, %negative_large : f32 + %bigger = scalar.cmpf ogt, %s, %m : f32 + %take = scalar.andi %live, %bigger : i1 + %m1 = scf.select %take, %s, %m : f32 + scf.yield %m1 : f32 + } + %sink_raw = view.load %sv[%head] : view<[%launch_heads]xf32> -> f32 + %has_sink = index.cmp eq, %use_sink, %c1 : index + scf.for %c = [%c0 to %value_head_size step %c1] { + %num, %den = scf.for %j = [%c0 to %count step %c1](%n = %zero : f32, %d = %zero : f32) -> (f32, f32) { + %jk, %jcap = index.assume %j, %key_capacity [lt(%j, %key_capacity)] : index, index + %mk16 = view.load %mv[%token, %j] : view<[%launch_tokens]x[%count]xf16> -> f16 + %mk = scalar.extf %mk16 : f16 to f32 + %live = scalar.cmpf ogt, %mk, %negative_large : f32 + %dot = scf.for %i = [%c0 to %qk_head_size step %c1](%a = %zero : f32) -> (f32) { + %q = view.load %qv[%token, %head, %i] : view<[%launch_tokens]x[%launch_heads]x[%qk_head_size]xf32> -> f32 + %k16 = view.load %kv[%jk, %kv_head, %i] : view<[%capacity]x[%launch_kv_heads]x[%qk_head_size]xf16> -> f16 + %k = scalar.extf %k16 : f16 to f32 + %p = scalar.mulf %q, %k : f32 + %a1 = scalar.addf %a, %p : f32 + scf.yield %a1 : f32 + } + %sd = scalar.mulf %dot, %scale : f32 + %s = scalar.addf %sd, %mk : f32 + %sh = scalar.subf %s, %row_max : f32 + %e0 = scalar.expf %sh : f32 + %e = scf.select %live, %e0, %zero : f32 + %v16 = view.load %vv[%jk, %kv_head, %c] : view<[%capacity]x[%launch_kv_heads]x[%value_head_size]xf16> -> f16 + %v = scalar.extf %v16 : f16 to f32 + %ev = scalar.mulf %e, %v : f32 + %n1 = scalar.addf %n, %ev : f32 + %d1 = scalar.addf %d, %e : f32 + scf.yield %n1, %d1 : f32, f32 + } + %sink_shift = scalar.subf %sink_raw, %row_max : f32 + %sink_e = scalar.expf %sink_shift : f32 + %sink_term = scf.select %has_sink, %sink_e, %zero : f32 + %den_total = scalar.addf %den, %sink_term : f32 + %ratio_out = scalar.divf %num, %den_total : f32 + %any = scalar.cmpf ogt, %den, %zero : f32 + %out = scf.select %any, %ratio_out, %zero : f32 + view.store %out, %ov[%token, %head, %c] : f32, view<[%launch_tokens]x[%launch_heads]x[%value_head_size]xf32> + } + kernel.return +} + +// Cases: 8 query heads over 2 key-value heads (gpt-oss's 64:8 grouping in small), head sizes 64, scale 0.125. +// Run with --config=ggml.attention_sink.query_head_count=8 --config=ggml.attention_sink.key_value_head_count=2 +// --config=ggml.attention_sink.qk_head_size=64 --config=ggml.attention_sink.value_head_size=64 --config=ggml.attention_sink.scale=0.125 + +check.case public @ggml_attention_sink_prefill_case { + %q_seed = check.param.seed base(7300000000000068001) count(1) : i64 + %k_seed = check.param.seed base(7300000000000068002) count(1) : i64 + %v_seed = check.param.seed base(7300000000000068003) count(1) : i64 + %m_seed = check.param.seed base(7300000000000068004) count(1) : i64 + %s_seed = check.param.seed base(7300000000000068005) count(1) : i64 + %tokens = check.literal value(3) : index + %keys = check.literal value(300) : index + %no_sink = check.literal value(0) : index + %sink = check.literal value(1) : index + %query = check.generate.random.uniform seed(%q_seed) range(-1.0 to 1.0) : tensor<1536xf32> + %key = check.generate.random.uniform seed(%k_seed) range(-1.0 to 1.0) : tensor<38400xf16> + %value = check.generate.random.uniform seed(%v_seed) range(-1.0 to 1.0) : tensor<38400xf16> + %mask = check.generate.random.uniform seed(%m_seed) range(-3.0 to 0.0) : tensor<900xf16> + %sinks = check.generate.random.uniform seed(%s_seed) range(-2.0 to 2.0) : tensor<8xf32> + %actual = check.generate.fill value(-7.0) : tensor<1536xf32> + %expected = check.generate.fill value(7.0) : tensor<1536xf32> + kernel.launch @ggml_attention_sink_reference_f32[%tokens, %keys, %keys, %no_sink](%tokens, %keys, %keys, %no_sink, %query, %key, %value, %mask, %sinks, %actual) : [index, index, index, index](index, index, index, index, tensor<1536xf32>, tensor<38400xf16>, tensor<38400xf16>, tensor<900xf16>, tensor<8xf32>, tensor<1536xf32>) + kernel.launch @ggml_attention_sink_f32[%tokens, %keys, %keys](%tokens, %keys, %keys, %query, %key, %mask, %sinks, %actual) : [index, index, index](index, index, index, tensor<1536xf32>, tensor<38400xf16>, tensor<900xf16>, tensor<8xf32>, tensor<1536xf32>) + kernel.launch @ggml_attention_sink_reference_f32[%tokens, %keys, %keys, %sink](%tokens, %keys, %keys, %sink, %query, %key, %value, %mask, %sinks, %expected) : [index, index, index, index](index, index, index, index, tensor<1536xf32>, tensor<38400xf16>, tensor<38400xf16>, tensor<900xf16>, tensor<8xf32>, tensor<1536xf32>) + check.expect.close actual(%actual) expected(%expected) atol(1.0e-5) rtol(1.0e-4) nan(same) : tensor<1536xf32> + check.return +} + +check.case public @ggml_attention_sink_decode_case { + %q_seed = check.param.seed base(7300000000000068011) count(1) : i64 + %k_seed = check.param.seed base(7300000000000068012) count(1) : i64 + %v_seed = check.param.seed base(7300000000000068013) count(1) : i64 + %m_seed = check.param.seed base(7300000000000068014) count(1) : i64 + %s_seed = check.param.seed base(7300000000000068015) count(1) : i64 + %tokens = check.literal value(1) : index + %keys = check.literal value(1000) : index + %no_sink = check.literal value(0) : index + %sink = check.literal value(1) : index + %query = check.generate.random.uniform seed(%q_seed) range(-1.0 to 1.0) : tensor<512xf32> + %key = check.generate.random.uniform seed(%k_seed) range(-1.0 to 1.0) : tensor<128000xf16> + %value = check.generate.random.uniform seed(%v_seed) range(-1.0 to 1.0) : tensor<128000xf16> + %mask = check.generate.random.uniform seed(%m_seed) range(-3.0 to 0.0) : tensor<1000xf16> + %sinks = check.generate.random.uniform seed(%s_seed) range(-2.0 to 2.0) : tensor<8xf32> + %actual = check.generate.fill value(-7.0) : tensor<512xf32> + %expected = check.generate.fill value(7.0) : tensor<512xf32> + kernel.launch @ggml_attention_sink_reference_f32[%tokens, %keys, %keys, %no_sink](%tokens, %keys, %keys, %no_sink, %query, %key, %value, %mask, %sinks, %actual) : [index, index, index, index](index, index, index, index, tensor<512xf32>, tensor<128000xf16>, tensor<128000xf16>, tensor<1000xf16>, tensor<8xf32>, tensor<512xf32>) + kernel.launch @ggml_attention_sink_f32[%tokens, %keys, %keys](%tokens, %keys, %keys, %query, %key, %mask, %sinks, %actual) : [index, index, index](index, index, index, tensor<512xf32>, tensor<128000xf16>, tensor<1000xf16>, tensor<8xf32>, tensor<512xf32>) + kernel.launch @ggml_attention_sink_reference_f32[%tokens, %keys, %keys, %sink](%tokens, %keys, %keys, %sink, %query, %key, %value, %mask, %sinks, %expected) : [index, index, index, index](index, index, index, index, tensor<512xf32>, tensor<128000xf16>, tensor<128000xf16>, tensor<1000xf16>, tensor<8xf32>, tensor<512xf32>) + check.expect.close actual(%actual) expected(%expected) atol(1.0e-5) rtol(1.0e-4) nan(same) : tensor<512xf32> + check.return +} + +check.case public @ggml_attention_sink_masked_case { + %q_seed = check.param.seed base(7300000000000068021) count(1) : i64 + %k_seed = check.param.seed base(7300000000000068022) count(1) : i64 + %v_seed = check.param.seed base(7300000000000068023) count(1) : i64 + %m_seed = check.param.seed base(7300000000000068024) count(1) : i64 + %s_seed = check.param.seed base(7300000000000068025) count(1) : i64 + %tokens = check.literal value(2) : index + %keys = check.literal value(40) : index + %no_sink = check.literal value(0) : index + %sink = check.literal value(1) : index + %query = check.generate.random.uniform seed(%q_seed) range(-1.0 to 1.0) : tensor<1024xf32> + %key = check.generate.random.uniform seed(%k_seed) range(-1.0 to 1.0) : tensor<5120xf16> + %value = check.generate.random.uniform seed(%v_seed) range(-1.0 to 1.0) : tensor<5120xf16> + %mask = check.generate.fill value(-1.0e30) : tensor<80xf16> + %sinks = check.generate.random.uniform seed(%s_seed) range(-2.0 to 2.0) : tensor<8xf32> + %actual = check.generate.fill value(-7.0) : tensor<1024xf32> + %expected = check.generate.fill value(7.0) : tensor<1024xf32> + kernel.launch @ggml_attention_sink_reference_f32[%tokens, %keys, %keys, %no_sink](%tokens, %keys, %keys, %no_sink, %query, %key, %value, %mask, %sinks, %actual) : [index, index, index, index](index, index, index, index, tensor<1024xf32>, tensor<5120xf16>, tensor<5120xf16>, tensor<80xf16>, tensor<8xf32>, tensor<1024xf32>) + kernel.launch @ggml_attention_sink_f32[%tokens, %keys, %keys](%tokens, %keys, %keys, %query, %key, %mask, %sinks, %actual) : [index, index, index](index, index, index, tensor<1024xf32>, tensor<5120xf16>, tensor<80xf16>, tensor<8xf32>, tensor<1024xf32>) + kernel.launch @ggml_attention_sink_reference_f32[%tokens, %keys, %keys, %sink](%tokens, %keys, %keys, %sink, %query, %key, %value, %mask, %sinks, %expected) : [index, index, index, index](index, index, index, index, tensor<1024xf32>, tensor<5120xf16>, tensor<5120xf16>, tensor<80xf16>, tensor<8xf32>, tensor<1024xf32>) + check.expect.close actual(%actual) expected(%expected) atol(1.0e-5) rtol(1.0e-4) nan(same) : tensor<1024xf32> + check.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/binary_bc_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/binary_bc_f32.loom new file mode 100644 index 000000000000..c536833537a7 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/binary_bc_f32.loom @@ -0,0 +1,128 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +template.decl @ggml.binary_f32.apply(%arg0: index, %arg1: f32, %arg2: f32) -> (f32) + +amdgpu.target @ggml_binary_bc_f32_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.binary_bc_f32.op : %value: index where [range(%value, 0, 8)] + +config.decl @ggml.binary_bc_f32.src0_broadcast_dim0 : %value: index where [range(%value, 0, 1)] + +config.decl @ggml.binary_bc_f32.src0_broadcast_dim1 : %value: index where [range(%value, 0, 1)] + +config.decl @ggml.binary_bc_f32.src0_broadcast_dim2 : %value: index where [range(%value, 0, 1)] + +config.decl @ggml.binary_bc_f32.src0_broadcast_dim3 : %value: index where [range(%value, 0, 1)] + +config.decl @ggml.binary_bc_f32.src1_broadcast_dim0 : %value: index where [range(%value, 0, 1)] + +config.decl @ggml.binary_bc_f32.src1_broadcast_dim1 : %value: index where [range(%value, 0, 1)] + +config.decl @ggml.binary_bc_f32.src1_broadcast_dim2 : %value: index where [range(%value, 0, 1)] + +config.decl @ggml.binary_bc_f32.src1_broadcast_dim3 : %value: index where [range(%value, 0, 1)] + +kernel.def target(@ggml_binary_bc_f32_gfx11_wave64) export("ggml_binary_bc_f32") @ggml_binary_bc_f32(%element_count: index, %ne0: index, %ne1: index, %ne2: index, %ne3: index, %src0_element_count: index, %src1_element_count: index) { + %one = index.constant 1 : index + %twofiftysix = index.constant 256 : index + %rounding = index.constant 255 : index + %rounded = index.add %element_count, %rounding : index + %workgroup_count = index.div %rounded, %twofiftysix : index + kernel.launch.config workgroups(%workgroup_count, %one, %one) workgroup_size(%twofiftysix, %one, %one) : index +} launch(%element_count: index, %ne0: index, %ne1: index, %ne2: index, %ne3: index, %src0_element_count: index, %src1_element_count: index, %lhs: buffer, %rhs: buffer, %output: buffer) { + %count = index.assume %element_count [range(%element_count, 1, 134217728)] : index + %n0 = index.assume %ne0 [range(%ne0, 1, 134217728)] : index + %n1 = index.assume %ne1 [range(%ne1, 1, 134217728)] : index + %n2 = index.assume %ne2 [range(%ne2, 1, 134217728)] : index + %src0_count = index.assume %src0_element_count [range(%src0_element_count, 1, 134217728)] : index + %src1_count = index.assume %src1_element_count [range(%src1_element_count, 1, 134217728)] : index + %op = config.get @ggml.binary_bc_f32.op : index + %src0_broadcast_dim0_flag = config.get @ggml.binary_bc_f32.src0_broadcast_dim0 : index + %src0_broadcast_dim1_flag = config.get @ggml.binary_bc_f32.src0_broadcast_dim1 : index + %src0_broadcast_dim2_flag = config.get @ggml.binary_bc_f32.src0_broadcast_dim2 : index + %src0_broadcast_dim3_flag = config.get @ggml.binary_bc_f32.src0_broadcast_dim3 : index + %src1_broadcast_dim0_flag = config.get @ggml.binary_bc_f32.src1_broadcast_dim0 : index + %src1_broadcast_dim1_flag = config.get @ggml.binary_bc_f32.src1_broadcast_dim1 : index + %src1_broadcast_dim2_flag = config.get @ggml.binary_bc_f32.src1_broadcast_dim2 : index + %src1_broadcast_dim3_flag = config.get @ggml.binary_bc_f32.src1_broadcast_dim3 : index + %workgroup = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %twofiftysix = index.constant 256 : index + %one = index.constant 1 : index + %zero = index.constant 0 : index + %base0 = index.mul %workgroup, %twofiftysix : index + %linear0 = index.add %base0, %workitem : index + %linear = index.assume %linear0 [range(%linear0, 0, 134217983)] : index + %in_bounds = index.cmp ult, %linear, %count : index + %zero_offset = index.constant 0 : offset + %lhs_view = buffer.view %lhs[%zero_offset] : buffer -> view<[%src0_count]xf32> + %rhs_view = buffer.view %rhs[%zero_offset] : buffer -> view<[%src1_count]xf32> + %output_view = buffer.view %output[%zero_offset] : buffer -> view<[%count]xf32> + scf.if %in_bounds { + %coord0 = index.rem %linear, %n0 : index + %linear_div_n0 = index.div %linear, %n0 : index + %coord1 = index.rem %linear_div_n0, %n1 : index + %linear_div_n01 = index.div %linear_div_n0, %n1 : index + %coord2 = index.rem %linear_div_n01, %n2 : index + %coord3 = index.div %linear_div_n01, %n2 : index + + %src0_broadcast_dims01 = index.add %src0_broadcast_dim0_flag, %src0_broadcast_dim1_flag : index + %src0_broadcast_dims23 = index.add %src0_broadcast_dim2_flag, %src0_broadcast_dim3_flag : index + %src0_broadcast_dims = index.add %src0_broadcast_dims01, %src0_broadcast_dims23 : index + %src0_no_broadcast = index.cmp eq, %src0_broadcast_dims, %zero : index + %src0_broadcast_dim0 = index.cmp eq, %src0_broadcast_dim0_flag, %one : index + %src0_broadcast_dim1 = index.cmp eq, %src0_broadcast_dim1_flag, %one : index + %src0_broadcast_dim2 = index.cmp eq, %src0_broadcast_dim2_flag, %one : index + %src0_broadcast_dim3 = index.cmp eq, %src0_broadcast_dim3_flag, %one : index + %src0_ne0 = scf.select %src0_broadcast_dim0, %one, %n0 : index + %src0_ne1 = scf.select %src0_broadcast_dim1, %one, %n1 : index + %src0_ne2 = scf.select %src0_broadcast_dim2, %one, %n2 : index + %src0_coord0 = scf.select %src0_broadcast_dim0, %zero, %coord0 : index + %src0_coord1 = scf.select %src0_broadcast_dim1, %zero, %coord1 : index + %src0_coord2 = scf.select %src0_broadcast_dim2, %zero, %coord2 : index + %src0_coord3 = scf.select %src0_broadcast_dim3, %zero, %coord3 : index + %src0_stride1 = index.add %src0_ne0, %zero : index + %src0_stride2 = index.mul %src0_stride1, %src0_ne1 : index + %src0_stride3 = index.mul %src0_stride2, %src0_ne2 : index + %src0_term1 = index.mul %src0_coord1, %src0_stride1 : index + %src0_term2 = index.mul %src0_coord2, %src0_stride2 : index + %src0_term3 = index.mul %src0_coord3, %src0_stride3 : index + %src0_linear01 = index.add %src0_coord0, %src0_term1 : index + %src0_linear012 = index.add %src0_linear01, %src0_term2 : index + %src0_broadcast_linear = index.add %src0_linear012, %src0_term3 : index + %src0_linear = scf.select %src0_no_broadcast, %linear, %src0_broadcast_linear : index + + %src1_broadcast_dims01 = index.add %src1_broadcast_dim0_flag, %src1_broadcast_dim1_flag : index + %src1_broadcast_dims23 = index.add %src1_broadcast_dim2_flag, %src1_broadcast_dim3_flag : index + %src1_broadcast_dims = index.add %src1_broadcast_dims01, %src1_broadcast_dims23 : index + %src1_no_broadcast = index.cmp eq, %src1_broadcast_dims, %zero : index + %src1_broadcast_dim0 = index.cmp eq, %src1_broadcast_dim0_flag, %one : index + %src1_broadcast_dim1 = index.cmp eq, %src1_broadcast_dim1_flag, %one : index + %src1_broadcast_dim2 = index.cmp eq, %src1_broadcast_dim2_flag, %one : index + %src1_broadcast_dim3 = index.cmp eq, %src1_broadcast_dim3_flag, %one : index + %src1_ne0 = scf.select %src1_broadcast_dim0, %one, %n0 : index + %src1_ne1 = scf.select %src1_broadcast_dim1, %one, %n1 : index + %src1_ne2 = scf.select %src1_broadcast_dim2, %one, %n2 : index + %src1_coord0 = scf.select %src1_broadcast_dim0, %zero, %coord0 : index + %src1_coord1 = scf.select %src1_broadcast_dim1, %zero, %coord1 : index + %src1_coord2 = scf.select %src1_broadcast_dim2, %zero, %coord2 : index + %src1_coord3 = scf.select %src1_broadcast_dim3, %zero, %coord3 : index + %src1_stride1 = index.add %src1_ne0, %zero : index + %src1_stride2 = index.mul %src1_stride1, %src1_ne1 : index + %src1_stride3 = index.mul %src1_stride2, %src1_ne2 : index + %src1_term1 = index.mul %src1_coord1, %src1_stride1 : index + %src1_term2 = index.mul %src1_coord2, %src1_stride2 : index + %src1_term3 = index.mul %src1_coord3, %src1_stride3 : index + %src1_linear01 = index.add %src1_coord0, %src1_term1 : index + %src1_linear012 = index.add %src1_linear01, %src1_term2 : index + %src1_broadcast_linear = index.add %src1_linear012, %src1_term3 : index + %src1_linear = scf.select %src1_no_broadcast, %linear, %src1_broadcast_linear : index + + %lhs_value = view.load %lhs_view[%src0_linear] : view<[%src0_count]xf32> -> f32 + %rhs_value = view.load %rhs_view[%src1_linear] : view<[%src1_count]xf32> -> f32 + %result = template.apply<@ggml.binary_f32.apply>(%op, %lhs_value, %rhs_value) : (index, f32, f32) -> (f32) + view.store %result, %output_view[%linear] : f32, view<[%count]xf32> + } + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/binary_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/binary_f32.loom new file mode 100644 index 000000000000..1064d5cc6f2f --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/binary_f32.loom @@ -0,0 +1,240 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +template.decl @ggml.binary_f32.apply(%arg0: index, %arg1: f32, %arg2: f32) -> (f32) + +amdgpu.target @ggml_binary_f32_gfx11_wave64 {subgroup_size = 64} + +amdgpu.target @ggml_binary_swiglu_i4_gfx11_wave32 {subgroup_size = 32} + +config.decl @ggml.binary_f32.op : %value: index where [range(%value, 0, 8)] + +config.decl @ggml.binary_f32.ne0 : %value: index where [range(%value, 1, 134217728)] + +config.decl @ggml.binary_f32.ne1 : %value: index where [range(%value, 1, 134217728)] + +config.decl @ggml.binary_f32.ne2 : %value: index where [range(%value, 1, 134217728)] + +config.decl @ggml.binary_f32.src0_stride1 : %value: index where [range(%value, 1, 134217728)] + +config.decl @ggml.binary_f32.src0_stride2 : %value: index where [range(%value, 1, 134217728)] + +config.decl @ggml.binary_f32.src0_stride3 : %value: index where [range(%value, 1, 134217728)] + +config.decl @ggml.binary_f32.src1_stride1 : %value: index where [range(%value, 1, 134217728)] + +config.decl @ggml.binary_f32.src1_stride2 : %value: index where [range(%value, 1, 134217728)] + +config.decl @ggml.binary_f32.src1_stride3 : %value: index where [range(%value, 1, 134217728)] + +config.decl @ggml.binary_f32.src0_span : %value: index where [range(%value, 1, 134217728)] + +config.decl @ggml.binary_f32.src1_span : %value: index where [range(%value, 1, 134217728)] + +config.decl @ggml.binary_swiglu_symmetric_i4.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 64)] + +config.decl @ggml.binary_swiglu_symmetric_i4.token_count : %value: index where [range(%value, 1, 16)] + +kernel.def target(@ggml_binary_f32_gfx11_wave64) export("ggml_binary_f32") @ggml_binary_f32(%element_count: index) { + %one = index.constant 1 : index + %twofiftysix = index.constant 256 : index + %rounding = index.constant 255 : index + %rounded = index.add %element_count, %rounding : index + %workgroup_count = index.div %rounded, %twofiftysix : index + kernel.launch.config workgroups(%workgroup_count, %one, %one) workgroup_size(%twofiftysix, %one, %one) : index +} launch(%element_count: index, %lhs: buffer, %rhs: buffer, %output: buffer) { + %count = index.assume %element_count [range(%element_count, 1, 134217728)] : index + %n0 = config.get @ggml.binary_f32.ne0 : index + %n1 = config.get @ggml.binary_f32.ne1 : index + %n2 = config.get @ggml.binary_f32.ne2 : index + %s0_1 = config.get @ggml.binary_f32.src0_stride1 : index + %s0_2 = config.get @ggml.binary_f32.src0_stride2 : index + %s0_3 = config.get @ggml.binary_f32.src0_stride3 : index + %s1_1 = config.get @ggml.binary_f32.src1_stride1 : index + %s1_2 = config.get @ggml.binary_f32.src1_stride2 : index + %s1_3 = config.get @ggml.binary_f32.src1_stride3 : index + %src0_count = config.get @ggml.binary_f32.src0_span : index + %src1_count = config.get @ggml.binary_f32.src1_span : index + %op = config.get @ggml.binary_f32.op : index + %workgroup = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %twofiftysix = index.constant 256 : index + %zero = index.constant 0 : index + %base0 = index.mul %workgroup, %twofiftysix : index + %linear0 = index.add %base0, %workitem : index + %linear = index.assume %linear0 [range(%linear0, 0, 134217983)] : index + %in_bounds = index.cmp ult, %linear, %count : index + %zero_offset = index.constant 0 : offset + %lhs_view = buffer.view %lhs[%zero_offset] : buffer -> view<[%src0_count]xf32> + %rhs_view = buffer.view %rhs[%zero_offset] : buffer -> view<[%src1_count]xf32> + %output_view = buffer.view %output[%zero_offset] : buffer -> view<[%count]xf32> + scf.if %in_bounds { + %coord0 = index.rem %linear, %n0 : index + %linear_div_n0 = index.div %linear, %n0 : index + %coord1 = index.rem %linear_div_n0, %n1 : index + %linear_div_n01 = index.div %linear_div_n0, %n1 : index + %coord2 = index.rem %linear_div_n01, %n2 : index + %coord3 = index.div %linear_div_n01, %n2 : index + %src0_term1 = index.mul %coord1, %s0_1 : index + %src0_term2 = index.mul %coord2, %s0_2 : index + %src0_term3 = index.mul %coord3, %s0_3 : index + %src0_linear01 = index.add %coord0, %src0_term1 : index + %src0_linear012 = index.add %src0_linear01, %src0_term2 : index + %src0_linear = index.add %src0_linear012, %src0_term3 : index + %src1_term1 = index.mul %coord1, %s1_1 : index + %src1_term2 = index.mul %coord2, %s1_2 : index + %src1_term3 = index.mul %coord3, %s1_3 : index + %src1_linear01 = index.add %coord0, %src1_term1 : index + %src1_linear012 = index.add %src1_linear01, %src1_term2 : index + %src1_linear = index.add %src1_linear012, %src1_term3 : index + %src0_in_bounds = index.cmp ult, %src0_linear, %src0_count : index + scf.if %src0_in_bounds { + %src1_in_bounds = index.cmp ult, %src1_linear, %src1_count : index + scf.if %src1_in_bounds { + %lhs_value = view.load %lhs_view[%src0_linear] : view<[%src0_count]xf32> -> f32 + %rhs_value = view.load %rhs_view[%src1_linear] : view<[%src1_count]xf32> -> f32 + %result = template.apply<@ggml.binary_f32.apply>(%op, %lhs_value, %rhs_value) : (index, f32, f32) -> (f32) + view.store %result, %output_view[%linear] : f32, view<[%count]xf32> + } + } + } + kernel.return +} + +kernel.def target(@ggml_binary_swiglu_i4_gfx11_wave32) export("ggml_binary_swiglu_symmetric_i4_k32") @ggml_binary_swiglu_symmetric_i4_k32() { + %c1 = index.constant 1 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %input_size = config.get @ggml.binary_swiglu_symmetric_i4.input_size : index + %token_count = config.get @ggml.binary_swiglu_symmetric_i4.token_count : index + %element_count = index.mul %input_size, %token_count : index + %group_count = index.div %element_count, %c64 : index + %rounded_group_count = index.add %group_count, %c7 : index + %workgroup_count = index.div %rounded_group_count, %c8 : index + kernel.launch.config workgroups(%workgroup_count, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%lhs: buffer, %rhs: buffer, %output: buffer, %i4_qs: buffer, %i4_ds: buffer, %i4_sums: buffer) { + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %c0_f32 = scalar.constant 0.0 : f32 + %amax_epsilon = scalar.constant 1.0000000000000001e-30 : f32 + %one_seventh = scalar.constant 0.14285714285714285 : f32 + %seven = scalar.constant 7.0 : f32 + %xor1 = scalar.constant 1 : i32 + %xor2 = scalar.constant 2 : i32 + %xor4 = scalar.constant 4 : i32 + %xor8 = scalar.constant 8 : i32 + %shuffle_width = scalar.constant 32 : i32 + %shift4 = scalar.constant 4 : i32 + %shift8 = scalar.constant 8 : i32 + %mask15 = vector.constant 15 : vector<2xi32> + + %input_size0 = config.get @ggml.binary_swiglu_symmetric_i4.input_size : index + %token_count0 = config.get @ggml.binary_swiglu_symmetric_i4.token_count : index + %input_size = index.assume %input_size0 [range(%input_size0, 256, 32768), mul(%input_size0, 64)] : index + %token_count = index.assume %token_count0 [range(%token_count0, 1, 16)] : index + %element_count0 = index.mul %input_size, %token_count : index + %element_count = index.assume %element_count0 [range(%element_count0, 256, 524288), mul(%element_count0, 64)] : index + %groups32 = index.div %element_count, %c32 : index + %groups64 = index.div %element_count, %c64 : index + %qs_halfwords = index.div %element_count, %c4 : index + + %workgroup = kernel.workgroup.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %lane0 = kernel.subgroup.lane.id : index + %lane = index.assume %lane0 [range(%lane0, 0, 31)] : index + %group_base = index.mul %workgroup, %c8 : index + %group64 = index.add %group_base, %subgroup : index + %valid_group = index.cmp ult, %group64, %groups64 : index + + %lhs_global = buffer.assume.memory_space %lhs : buffer + %rhs_global = buffer.assume.memory_space %rhs : buffer + %output_global = buffer.assume.memory_space %output : buffer + %i4_qs_global = buffer.assume.memory_space %i4_qs : buffer + %i4_ds_global = buffer.assume.memory_space %i4_ds : buffer + %i4_sums_global = buffer.assume.memory_space %i4_sums : buffer + %lhs_na, %rhs_na, %output_na, %i4_qs_na, %i4_ds_na, %i4_sums_na = buffer.assume.noalias %lhs_global, %rhs_global, %output_global, %i4_qs_global, %i4_ds_global, %i4_sums_global : buffer, buffer, buffer, buffer, buffer, buffer + %lhs_view = buffer.view %lhs_na[%base] : buffer -> view<[%element_count]xf32> + %rhs_view = buffer.view %rhs_na[%base] : buffer -> view<[%element_count]xf32> + %output_view = buffer.view %output_na[%base] : buffer -> view<[%element_count]xf32> + %i4_qs_view = buffer.view %i4_qs_na[%base] : buffer -> view<[%qs_halfwords]xi16> + %i4_ds_view = buffer.view %i4_ds_na[%base] : buffer -> view<[%groups32]xf32> + %i4_sums_view = buffer.view %i4_sums_na[%base] : buffer -> view<[%groups32]xi32> + + scf.if %valid_group { + %group_element_base = index.mul %group64, %c64 : index + %lane_element_offset = index.mul %lane, %c2 : index + %linear0 = index.add %group_element_base, %lane_element_offset : index + %linear = index.assume %linear0 [range(%linear0, 0, 524286), mul(%linear0, 2)] : index + %lhs_values = vector.load %lhs_view[%linear] : view<[%element_count]xf32> -> vector<2xf32> + %rhs_values = vector.load %rhs_view[%linear] : view<[%element_count]xf32> -> vector<2xf32> + %activated = vector.siluf %lhs_values : vector<2xf32> + %result = vector.mulf %activated, %rhs_values : vector<2xf32> + vector.store %result, %output_view[%linear] : vector<2xf32>, view<[%element_count]xf32> + + %absolute_values = vector.absf %result : vector<2xf32> + %lane_max = vector.reduce %absolute_values, %c0_f32 : vector<2xf32>, f32 + %group_max = kernel.subgroup.reduce %lane_max : f32 + %amax = scalar.maxnumf %group_max, %amax_epsilon : f32 + %scale = scalar.mulf %amax, %one_seventh : f32 + %is_scale_leader = index.cmp eq, %lane, %c0 : index + %group32_base = index.mul %group64, %c2 : index + %group32_high = index.add %group32_base, %c1 : index + scf.if %is_scale_leader { + view.store %scale, %i4_ds_view[%group32_base] : f32, view<[%groups32]xf32> + view.store %scale, %i4_ds_view[%group32_high] : f32, view<[%groups32]xf32> + } + + %rscale = scalar.divf %seven, %amax : f32 + %rscale_vector = vector.splat %rscale : vector<2xf32> + %scaled = vector.mulf %result, %rscale_vector : vector<2xf32> + %rounded = vector.roundf %scaled : vector<2xf32> + %quantized = vector.fptosi %rounded : vector<2xf32> to vector<2xi32> + %nibbles = vector.andi %quantized, %mask15 : vector<2xi32> + %q0 = vector.extract %nibbles[0] : vector<2xi32> -> i32 + %q1 = vector.extract %nibbles[1] : vector<2xi32> -> i32 + %q1_shifted = scalar.shli %q1, %shift4 : i32 + %packed_pair = scalar.ori %q0, %q1_shifted : i32 + %peer_pair, %peer_valid = kernel.subgroup.shuffle %packed_pair, %xor1, %shuffle_width : i32, i32, i32 + %peer_shifted = scalar.shli %peer_pair, %shift8 : i32 + %packed_word = scalar.ori %packed_pair, %peer_shifted : i32 + %packed_i16 = scalar.trunci %packed_word : i32 to i16 + %lane_pair = index.div %lane, %c2 : index + %word_base = index.mul %group64, %c16 : index + %word_index = index.add %word_base, %lane_pair : index + %lane_in_pair = index.rem %lane, %c2 : index + %is_word_leader = index.cmp eq, %lane_in_pair, %c0 : index + scf.if %is_word_leader { + view.store %packed_i16, %i4_qs_view[%word_index] : i16, view<[%qs_halfwords]xi16> + } + + %lane_sum = vector.reduce %rounded, %c0_f32 : vector<2xf32>, f32 + %sum_x1_peer, %sum_x1_valid = kernel.subgroup.shuffle %lane_sum, %xor1, %shuffle_width : f32, i32, i32 + %sum_x1 = scalar.addf %lane_sum, %sum_x1_peer : f32 + %sum_x2_peer, %sum_x2_valid = kernel.subgroup.shuffle %sum_x1, %xor2, %shuffle_width : f32, i32, i32 + %sum_x2 = scalar.addf %sum_x1, %sum_x2_peer : f32 + %sum_x4_peer, %sum_x4_valid = kernel.subgroup.shuffle %sum_x2, %xor4, %shuffle_width : f32, i32, i32 + %sum_x4 = scalar.addf %sum_x2, %sum_x4_peer : f32 + %sum_x8_peer, %sum_x8_valid = kernel.subgroup.shuffle %sum_x4, %xor8, %shuffle_width : f32, i32, i32 + %group_sum = scalar.addf %sum_x4, %sum_x8_peer : f32 + %group_half = index.div %lane, %c16 : index + %lane_in_half = index.rem %lane, %c16 : index + %is_sum_leader = index.cmp eq, %lane_in_half, %c0 : index + scf.if %is_sum_leader { + %sum_index = index.add %group32_base, %group_half : index + %sum_value = scalar.fptosi %group_sum : f32 to i32 + view.store %sum_value, %i4_sums_view[%sum_index] : i32, view<[%groups32]xi32> + } + } + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/copy_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/copy_f32.loom new file mode 100644 index 000000000000..28947949bd38 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/copy_f32.loom @@ -0,0 +1,471 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +amdgpu.target @ggml_copy_f32_gfx11_wave64 {subgroup_size = 64} + +kernel.def target(@ggml_copy_f32_gfx11_wave64) export("ggml_copy_f32") @ggml_copy_f32(%element_count: index) { + %one = index.constant 1 : index + %twofiftysix = index.constant 256 : index + %rounding = index.constant 255 : index + %rounded = index.add %element_count, %rounding : index + %workgroup_count = index.div %rounded, %twofiftysix : index + kernel.launch.config workgroups(%workgroup_count, %one, %one) workgroup_size(%twofiftysix, %one, %one) : index +} launch(%element_count: index, %source: buffer, %output: buffer) { + %count = index.assume %element_count [range(%element_count, 1, 1073741824)] : index + %workgroup = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %twofiftysix = index.constant 256 : index + %base = index.mul %workgroup, %twofiftysix : index + %linear0 = index.add %base, %workitem : index + %linear = index.assume %linear0 [range(%linear0, 0, 1073742079)] : index + %in_bounds = index.cmp ult, %linear, %count : index + %zero_offset = index.constant 0 : offset + %source_noalias, %output_noalias = buffer.assume.noalias %source, %output : buffer, buffer + %source_view = buffer.view %source_noalias[%zero_offset] : buffer -> view<[%count]xf32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%count]xf32> + scf.if %in_bounds { + %value = view.load %source_view[%linear] : view<[%count]xf32> -> f32 + view.store %value, %output_view[%linear] : f32, view<[%count]xf32> + } + kernel.return +} + + +// Exact K16-major F16 layout for native matrix fragments. The ordinary F16 +// alternate remains unchanged; this copy is a private dispatch transient. +config.decl @ggml.copy_f16_k16_major.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] +config.decl @ggml.copy_f16_k16_major.token_count : %value: index where [range(%value, 512, 2048), mul(%value, 512)] +config.decl @ggml.copy_f16_k16_major.input_is_f16 : %value: index where [range(%value, 0, 1)] +amdgpu.target @ggml_copy_f16_k16_major_target {subgroup_size = 32} + +kernel.def target(@ggml_copy_f16_k16_major_target) export("ggml_copy_f16_k16_major") @ggml_copy_f16_k16_major(%token_count: index) { + %one = index.constant 1 : index + %threads = index.constant 256 : index + %tile = index.constant 2048 : index + %rounding = index.constant 2047 : index + %width = config.get @ggml.copy_f16_k16_major.input_size : index + %elements = index.mul %token_count, %width : index + %rounded = index.add %elements, %rounding : index + %groups = index.div %rounded, %tile : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%threads, %one, %one) : index +} launch(%token_count: index, %input: buffer, %output: buffer) { + %width = config.get @ggml.copy_f16_k16_major.input_size : index + %token_count_static = config.get @ggml.copy_f16_k16_major.token_count : index + %tokens = index.assume %token_count_static [range(%token_count_static, 16, 2048), mul(%token_count_static, 16)] : index + %count = index.mul %tokens, %width : index + %c0 = index.constant 0 : offset + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c256 = index.constant 256 : index + %thread = kernel.workitem.id : index + %group = kernel.workgroup.id : index + %packet = index.madd %group, %c256, %thread : index + %linear_raw = index.mul %packet, %c8 : index + %linear = index.assume %linear_raw [range(%linear_raw, 0, 67110904), mul(%linear_raw, 8)] : index + %end = index.add %linear, %c8 : index + %valid = index.cmp ule, %end, %count : index + scf.if %valid { + %tile = index.div %linear, %c256 : index + %local = index.rem %linear, %c256 : index + %m_tiles_raw = index.div %tokens, %c16 : index + %m_tiles = index.assume %m_tiles_raw [range(%m_tiles_raw, 1, 128)] : index + %m_tile = index.rem %tile, %m_tiles : index + %k_tile = index.div %tile, %m_tiles : index + %row_in_tile = index.div %local, %c16 : index + %k_in_tile_raw = index.rem %local, %c16 : index + %k_in_tile = index.assume %k_in_tile_raw [range(%k_in_tile_raw, 0, 8), mul(%k_in_tile_raw, 8)] : index + %row = index.madd %m_tile, %c16, %row_in_tile : index + %k = index.madd %k_tile, %c16, %k_in_tile : index + %bounded_row, %bounded_tokens = index.assume %row, %tokens [lt(%row, %tokens)] : index, index + %input_noalias, %output_noalias = buffer.assume.noalias %input, %output : buffer, buffer + %dest = buffer.view %output_noalias[%c0] : buffer -> view<[%count]xf16> + %maximum_k = index.sub %width, %c8 : index + %bounded_k = index.assume %k [range(%k, 0, %maximum_k), mul(%k, 8)] : index + %input_is_f16 = config.get @ggml.copy_f16_k16_major.input_is_f16 : index + %one = index.constant 1 : index + %use_f16 = index.cmp eq, %input_is_f16, %one : index + %values = scf.if %use_f16 -> (vector<8xf16>) { + %source = buffer.view %input_noalias[%c0] : buffer -> view<[%tokens]x[%width]xf16> + %loaded = vector.load %source[%bounded_row, %bounded_k] : view<[%tokens]x[%width]xf16> -> vector<8xf16> + scf.yield %loaded : vector<8xf16> + } else { + %source = buffer.view %input_noalias[%c0] : buffer -> view<[%tokens]x[%width]xf32> + %loaded = vector.load %source[%bounded_row, %bounded_k] : view<[%tokens]x[%width]xf32> -> vector<8xf32> + %converted = vector.fptrunc %loaded : vector<8xf32> to vector<8xf16> + scf.yield %converted : vector<8xf16> + } + vector.store %values, %dest[%linear] : vector<8xf16>, view<[%count]xf16> + } + kernel.return +} + +config.decl @ggml.copy_transpose_f16.row_count : %value: index where [range(%value, 32, 32768), mul(%value, 32)] +config.decl @ggml.copy_transpose_f16.column_count : %value: index where [range(%value, 32, 32768), mul(%value, 32)] +kernel.def target(@ggml_copy_f32_gfx11_wave64) export("ggml_copy_transpose_f16") @ggml_copy_transpose_f16() { + %rows = config.get @ggml.copy_transpose_f16.row_count : index + %columns = config.get @ggml.copy_transpose_f16.column_count : index + %tile_size = index.constant 32 : index + %one = index.constant 1 : index + %x = index.div %columns, %tile_size : index + %gy = index.div %rows, %tile_size : index + %y = index.constant 256 : index + kernel.launch.config workgroups(%x, %gy, %one) workgroup_size(%y, %one, %one) : index +} launch(%input: buffer, %output: buffer) { + %rows = config.get @ggml.copy_transpose_f16.row_count : index + %columns = config.get @ggml.copy_transpose_f16.column_count : index + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c32 = index.constant 32 : index + %zero = index.constant 0 : offset + %bytes = index.constant 2560 : offset + %wi = kernel.workitem.id : index + %gx = kernel.workgroup.id : index + %gy = kernel.workgroup.id : index + %row_origin = index.mul %gy, %c32 : index + %col_origin = index.mul %gx, %c32 : index + %row = index.div %wi, %c8 : index + %col_word = index.rem %wi, %c8 : index + %col = index.mul %col_word, %c4 : index + %source_row = index.add %row_origin, %row : index + %source_col = index.add %col_origin, %col : index + %src, %dst = buffer.assume.noalias %input, %output : buffer, buffer + %input_view = buffer.view %src[%zero] : buffer -> view<[%rows]x[%columns]xf16> + %output_view = buffer.view %dst[%zero] : buffer -> view<[%columns]x[%rows]xf16> + %shared = buffer.alloca align(16) %bytes : buffer + %shared_view = buffer.view %shared[%zero] : buffer -> view<32x40xf16> + %transposed = encoding.layout.strided [1, 40] : encoding + %shared_transpose_view = buffer.view %shared[%zero] : buffer -> view<32x32xf16, %transposed> + %loaded = vector.load %input_view[%source_row, %source_col] : view<[%rows]x[%columns]xf16> -> vector<4xf16> + vector.store %loaded, %shared_view[%row, %col] : vector<4xf16>, view<32x40xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + %add0 = index.constant 0 : index + %read_col0 = index.add %col, %add0 : index + %s0 = view.load %shared_transpose_view[%row, %read_col0] : view<32x32xf16, %transposed> -> f16 + %add1 = index.constant 1 : index + %read_col1 = index.add %col, %add1 : index + %s1 = view.load %shared_transpose_view[%row, %read_col1] : view<32x32xf16, %transposed> -> f16 + %add2 = index.constant 2 : index + %read_col2 = index.add %col, %add2 : index + %s2 = view.load %shared_transpose_view[%row, %read_col2] : view<32x32xf16, %transposed> -> f16 + %add3 = index.constant 3 : index + %read_col3 = index.add %col, %add3 : index + %s3 = view.load %shared_transpose_view[%row, %read_col3] : view<32x32xf16, %transposed> -> f16 + %swapped = vector.from_elements %s0, %s1, %s2, %s3 : vector<4xf16> + %target_row = index.add %col_origin, %row : index + %target_col = index.add %row_origin, %col : index + vector.store %swapped, %output_view[%target_row, %target_col] : vector<4xf16>, view<[%columns]x[%rows]xf16> + kernel.return +} + +check.case public @ggml_copy_f32_small_case { + %four = check.literal value(4) : index + %source = check.generate.iota offset(-2.0) step(1.0) : tensor<4xf32> + %output = check.generate.fill value(0.0) : tensor<4xf32> + %expected = check.generate.iota offset(-2.0) step(1.0) : tensor<4xf32> + kernel.launch @ggml_copy_f32[%four](%four, %source, %output) : [index](index, tensor<4xf32>, tensor<4xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<4xf32> + check.return +} + +check.benchmark<@ggml_copy_f32_small_case> @ggml_copy_f32_small + +// Aligned tensor conversions used by existing F16-consuming matmul entries. +kernel.def target(@ggml_copy_f32_gfx11_wave64) export("ggml_copy_f32_f16") @ggml_copy_f32_f16(%element_count: index) { + %one = index.constant 1 : index + %threads = index.constant 256 : index + %tile = index.constant 1024 : index + %rounding = index.constant 1023 : index + %rounded = index.add %element_count, %rounding : index + %groups = index.div %rounded, %tile : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%threads, %one, %one) : index +} launch(%element_count: index, %source: buffer, %output: buffer) { + %count = index.assume %element_count [range(%element_count, 4, 1073741824), mul(%element_count, 4)] : index + %group = kernel.workgroup.id : index + %thread = kernel.workitem.id : index + %threads = index.constant 256 : index + %four = index.constant 4 : index + %lane = index.madd %group, %threads, %thread : index + %linear0 = index.mul %lane, %four : index + %linear = index.assume %linear0 [range(%linear0, 0, 1073742844), mul(%linear0, 4)] : index + %end = index.add %linear, %four : index + %valid = index.cmp ule, %end, %count : index + %zero = index.constant 0 : offset + %source_noalias, %output_noalias = buffer.assume.noalias %source, %output : buffer, buffer + %source_view = buffer.view %source_noalias[%zero] : buffer -> view<[%count]xf32> + %output_view = buffer.view %output_noalias[%zero] : buffer -> view<[%count]xf16> + scf.if %valid { + %values = vector.load %source_view[%linear] : view<[%count]xf32> -> vector<4xf32> + %converted = vector.fptrunc %values : vector<4xf32> to vector<4xf16> + vector.store %converted, %output_view[%linear] : vector<4xf16>, view<[%count]xf16> + } + kernel.return +} + +check.case public @ggml_copy_f32_f16_small_case { + %count = check.literal value(12) : index + %source = check.generate.iota offset(-2.0) step(0.5) : tensor<12xf32> + %output = check.generate.fill value(0.0) : tensor<12xf16> + %expected = check.generate.iota offset(-2.0) step(0.5) : tensor<12xf16> + kernel.launch @ggml_copy_f32_f16[%count](%count, %source, %output) : [index](index, tensor<12xf32>, tensor<12xf16>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<12xf16> + check.return +} + +config.decl @ggml.copy_strided_source_f32.ne0 : %value: index where [range(%value, 1, 1073741824)] + +config.decl @ggml.copy_strided_source_f32.ne1 : %value: index where [range(%value, 1, 1073741824)] + +config.decl @ggml.copy_strided_source_f32.ne2 : %value: index where [range(%value, 1, 1073741824)] + +config.decl @ggml.copy_strided_source_f32.stride0 : %value: index where [range(%value, 1, 1073741824)] + +config.decl @ggml.copy_strided_source_f32.stride1 : %value: index where [range(%value, 1, 1073741824)] + +config.decl @ggml.copy_strided_source_f32.stride2 : %value: index where [range(%value, 1, 1073741824)] + +config.decl @ggml.copy_strided_source_f32.stride3 : %value: index where [range(%value, 1, 1073741824)] + +kernel.def target(@ggml_copy_f32_gfx11_wave64) export("ggml_copy_strided_source_f32") @ggml_copy_strided_source_f32(%element_count: index, %source_span: index) { + %one = index.constant 1 : index + %twofiftysix = index.constant 256 : index + %rounding = index.constant 255 : index + %rounded = index.add %element_count, %rounding : index + %workgroup_count = index.div %rounded, %twofiftysix : index + kernel.launch.config workgroups(%workgroup_count, %one, %one) workgroup_size(%twofiftysix, %one, %one) : index +} launch(%element_count: index, %source_span: index, %source: buffer, %output: buffer) { + %count = index.assume %element_count [range(%element_count, 1, 1073741824)] : index + %n0 = config.get @ggml.copy_strided_source_f32.ne0 : index + %n1 = config.get @ggml.copy_strided_source_f32.ne1 : index + %n2 = config.get @ggml.copy_strided_source_f32.ne2 : index + %span = index.assume %source_span [range(%source_span, 1, 1073741824)] : index + %s0 = config.get @ggml.copy_strided_source_f32.stride0 : index + %s1 = config.get @ggml.copy_strided_source_f32.stride1 : index + %s2 = config.get @ggml.copy_strided_source_f32.stride2 : index + %s3 = config.get @ggml.copy_strided_source_f32.stride3 : index + %workgroup = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %twofiftysix = index.constant 256 : index + %base = index.mul %workgroup, %twofiftysix : index + %linear0 = index.add %base, %workitem : index + %linear = index.assume %linear0 [range(%linear0, 0, 1073742079)] : index + %in_bounds = index.cmp ult, %linear, %count : index + %zero_offset = index.constant 0 : offset + %source_noalias, %output_noalias = buffer.assume.noalias %source, %output : buffer, buffer + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%count]xf32> + scf.if %in_bounds { + %coord0 = index.rem %linear, %n0 : index + %remaining0 = index.div %linear, %n0 : index + %coord1 = index.rem %remaining0, %n1 : index + %remaining1 = index.div %remaining0, %n1 : index + %coord2 = index.rem %remaining1, %n2 : index + %coord3 = index.div %remaining1, %n2 : index + %offset0 = index.mul %coord0, %s0 : index + %offset1 = index.mul %coord1, %s1 : index + %offset2 = index.mul %coord2, %s2 : index + %offset3 = index.mul %coord3, %s3 : index + %offset01 = index.add %offset0, %offset1 : index + %offset23 = index.add %offset2, %offset3 : index + %source_index0 = index.add %offset01, %offset23 : index + %source_index, %view_span = index.assume %source_index0, %span [lt(%source_index0, %span)] : index, index + %source_view = buffer.view %source_noalias[%zero_offset] : buffer -> view<[%view_span]xf32> + %value = view.load %source_view[%source_index] : view<[%view_span]xf32> -> f32 + view.store %value, %output_view[%linear] : f32, view<[%count]xf32> + } + kernel.return +} + +check.case public @ggml_copy_strided_source_f32_small_case { + %six = check.literal value(6) : index + %eight = check.literal value(8) : index + %source = check.generate.iota offset(0.0) step(1.0) period(5) : tensor<10xf32> + %output = check.generate.fill value(-1.0) : tensor<6xf32> + %expected = check.generate.iota offset(0.0) step(1.0) period(3) : tensor<6xf32> + kernel.launch @ggml_copy_strided_source_f32[%six, %eight](%six, %eight, %source, %output) : [index, index](index, index, tensor<10xf32>, tensor<6xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<6xf32> + check.return +} + +check.benchmark<@ggml_copy_strided_source_f32_small_case> @ggml_copy_strided_source_f32_small + +config.decl @ggml.concat_dim0_f32.lhs_width : %value: index where [range(%value, 1, 65536)] + +config.decl @ggml.concat_dim0_f32.rhs_width : %value: index where [range(%value, 1, 65536)] + +config.decl @ggml.concat_dim0_f32.row_count : %value: index where [range(%value, 1, 1048576)] + +kernel.def target(@ggml_copy_f32_gfx11_wave64) export("ggml_concat_dim0_f32") @ggml_concat_dim0_f32() { + %one = index.constant 1 : index + %workgroup_size = index.constant 256 : index + %rounding = index.constant 255 : index + %lhs_width = config.get @ggml.concat_dim0_f32.lhs_width : index + %rhs_width = config.get @ggml.concat_dim0_f32.rhs_width : index + %row_count = config.get @ggml.concat_dim0_f32.row_count : index + %output_width = index.add %lhs_width, %rhs_width : index + %element_count = index.mul %output_width, %row_count : index + %rounded = index.add %element_count, %rounding : index + %workgroup_count = index.div %rounded, %workgroup_size : index + kernel.launch.config workgroups(%workgroup_count, %one, %one) workgroup_size(%workgroup_size, %one, %one) : index +} launch(%lhs: buffer, %rhs: buffer, %output: buffer) { + %base = index.constant 0 : offset + %workgroup_size = index.constant 256 : index + %lhs_width0 = config.get @ggml.concat_dim0_f32.lhs_width : index + %rhs_width0 = config.get @ggml.concat_dim0_f32.rhs_width : index + %row_count0 = config.get @ggml.concat_dim0_f32.row_count : index + %lhs_width = index.assume %lhs_width0 [range(%lhs_width0, 1, 65536)] : index + %rhs_width = index.assume %rhs_width0 [range(%rhs_width0, 1, 65536)] : index + %row_count = index.assume %row_count0 [range(%row_count0, 1, 1048576)] : index + %output_width0 = index.add %lhs_width, %rhs_width : index + %output_width = index.assume %output_width0 [range(%output_width0, 2, 131072)] : index + %element_count0 = index.mul %output_width, %row_count : index + %element_count = index.assume %element_count0 [range(%element_count0, 1, 1073741824)] : index + %workgroup0 = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %workgroup = index.assume %workgroup0 [range(%workgroup0, 0, 4194303)] : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 255)] : index + %linear_base = index.mul %workgroup, %workgroup_size : index + %linear0 = index.add %linear_base, %workitem : index + %linear = index.assume %linear0 [range(%linear0, 0, 1073742079)] : index + %in_bounds = index.cmp ult, %linear, %element_count : index + + %lhs_noalias, %rhs_noalias, %output_noalias = buffer.assume.noalias %lhs, %rhs, %output : buffer, buffer, buffer + %lhs_view = buffer.view %lhs_noalias[%base] : buffer -> view<1073741824xf32> + %rhs_view = buffer.view %rhs_noalias[%base] : buffer -> view<1073741824xf32> + %output_view = buffer.view %output_noalias[%base] : buffer -> view<1073741824xf32> + + scf.if %in_bounds { + %row = index.div %linear, %output_width : index + %column = index.rem %linear, %output_width : index + %from_lhs = index.cmp ult, %column, %lhs_width : index + %value = scf.if %from_lhs -> (f32) { + %lhs_row = index.mul %row, %lhs_width : index + %lhs_index0 = index.add %lhs_row, %column : index + %lhs_index = index.assume %lhs_index0 [range(%lhs_index0, 0, 1073741823)] : index + %loaded = view.load %lhs_view[%lhs_index] : view<1073741824xf32> -> f32 + scf.yield %loaded : f32 + } else { + %rhs_column = index.sub %column, %lhs_width : index + %rhs_row = index.mul %row, %rhs_width : index + %rhs_index0 = index.add %rhs_row, %rhs_column : index + %rhs_index = index.assume %rhs_index0 [range(%rhs_index0, 0, 1073741823)] : index + %loaded = view.load %rhs_view[%rhs_index] : view<1073741824xf32> -> f32 + scf.yield %loaded : f32 + } + view.store %value, %output_view[%linear] : f32, view<1073741824xf32> + } + kernel.return +} + +config.decl @ggml.concat_dim0_strided_source_f32.lhs_width : %value: index where [range(%value, 1, 65536)] + +config.decl @ggml.concat_dim0_strided_source_f32.rhs_width : %value: index where [range(%value, 1, 65536)] + +config.decl @ggml.concat_dim0_strided_source_f32.row_count : %value: index where [range(%value, 1, 1048576)] + +config.decl @ggml.concat_dim0_strided_source_f32.row_ne1 : %value: index where [range(%value, 1, 1048576)] + +config.decl @ggml.concat_dim0_strided_source_f32.row_ne2 : %value: index where [range(%value, 1, 1048576)] + +config.decl @ggml.concat_dim0_strided_source_f32.lhs_stride0 : %value: index where [range(%value, 1, 1073741824)] + +config.decl @ggml.concat_dim0_strided_source_f32.lhs_stride1 : %value: index where [range(%value, 1, 1073741824)] + +config.decl @ggml.concat_dim0_strided_source_f32.lhs_stride2 : %value: index where [range(%value, 1, 1073741824)] + +config.decl @ggml.concat_dim0_strided_source_f32.lhs_stride3 : %value: index where [range(%value, 1, 1073741824)] + +config.decl @ggml.concat_dim0_strided_source_f32.rhs_stride0 : %value: index where [range(%value, 1, 1073741824)] + +config.decl @ggml.concat_dim0_strided_source_f32.rhs_stride1 : %value: index where [range(%value, 1, 1073741824)] + +config.decl @ggml.concat_dim0_strided_source_f32.rhs_stride2 : %value: index where [range(%value, 1, 1073741824)] + +config.decl @ggml.concat_dim0_strided_source_f32.rhs_stride3 : %value: index where [range(%value, 1, 1073741824)] + +kernel.def target(@ggml_copy_f32_gfx11_wave64) export("ggml_concat_dim0_strided_source_f32") @ggml_concat_dim0_strided_source_f32(%lhs_span: index, %rhs_span: index) { + %one = index.constant 1 : index + %workgroup_size = index.constant 256 : index + %rounding = index.constant 255 : index + %lhs_width = config.get @ggml.concat_dim0_strided_source_f32.lhs_width : index + %rhs_width = config.get @ggml.concat_dim0_strided_source_f32.rhs_width : index + %row_count = config.get @ggml.concat_dim0_strided_source_f32.row_count : index + %output_width = index.add %lhs_width, %rhs_width : index + %element_count = index.mul %output_width, %row_count : index + %rounded = index.add %element_count, %rounding : index + %workgroup_count = index.div %rounded, %workgroup_size : index + kernel.launch.config workgroups(%workgroup_count, %one, %one) workgroup_size(%workgroup_size, %one, %one) : index +} launch(%lhs_span: index, %rhs_span: index, %lhs: buffer, %rhs: buffer, %output: buffer) { + %base = index.constant 0 : offset + %workgroup_size = index.constant 256 : index + %lhs_width0 = config.get @ggml.concat_dim0_strided_source_f32.lhs_width : index + %rhs_width0 = config.get @ggml.concat_dim0_strided_source_f32.rhs_width : index + %row_count0 = config.get @ggml.concat_dim0_strided_source_f32.row_count : index + %row_ne1 = config.get @ggml.concat_dim0_strided_source_f32.row_ne1 : index + %row_ne2 = config.get @ggml.concat_dim0_strided_source_f32.row_ne2 : index + %lhs_s0 = config.get @ggml.concat_dim0_strided_source_f32.lhs_stride0 : index + %lhs_s1 = config.get @ggml.concat_dim0_strided_source_f32.lhs_stride1 : index + %lhs_s2 = config.get @ggml.concat_dim0_strided_source_f32.lhs_stride2 : index + %lhs_s3 = config.get @ggml.concat_dim0_strided_source_f32.lhs_stride3 : index + %rhs_s0 = config.get @ggml.concat_dim0_strided_source_f32.rhs_stride0 : index + %rhs_s1 = config.get @ggml.concat_dim0_strided_source_f32.rhs_stride1 : index + %rhs_s2 = config.get @ggml.concat_dim0_strided_source_f32.rhs_stride2 : index + %rhs_s3 = config.get @ggml.concat_dim0_strided_source_f32.rhs_stride3 : index + %lhs_width = index.assume %lhs_width0 [range(%lhs_width0, 1, 65536)] : index + %rhs_width = index.assume %rhs_width0 [range(%rhs_width0, 1, 65536)] : index + %row_count = index.assume %row_count0 [range(%row_count0, 1, 1048576)] : index + %output_width0 = index.add %lhs_width, %rhs_width : index + %output_width = index.assume %output_width0 [range(%output_width0, 2, 131072)] : index + %element_count0 = index.mul %output_width, %row_count : index + %element_count = index.assume %element_count0 [range(%element_count0, 1, 1073741824)] : index + %lhs_span_checked = index.assume %lhs_span [range(%lhs_span, 1, 1073741824)] : index + %rhs_span_checked = index.assume %rhs_span [range(%rhs_span, 1, 1073741824)] : index + %workgroup0 = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %workgroup = index.assume %workgroup0 [range(%workgroup0, 0, 4194303)] : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 255)] : index + %linear_base = index.mul %workgroup, %workgroup_size : index + %linear0 = index.add %linear_base, %workitem : index + %linear = index.assume %linear0 [range(%linear0, 0, 1073742079)] : index + %in_bounds = index.cmp ult, %linear, %element_count : index + + %lhs_noalias, %rhs_noalias, %output_noalias = buffer.assume.noalias %lhs, %rhs, %output : buffer, buffer, buffer + %output_view = buffer.view %output_noalias[%base] : buffer -> view<[%element_count]xf32> + + scf.if %in_bounds { + %row = index.div %linear, %output_width : index + %column = index.rem %linear, %output_width : index + %coord1 = index.rem %row, %row_ne1 : index + %row_after_ne1 = index.div %row, %row_ne1 : index + %coord2 = index.rem %row_after_ne1, %row_ne2 : index + %coord3 = index.div %row_after_ne1, %row_ne2 : index + %from_lhs = index.cmp ult, %column, %lhs_width : index + %value = scf.if %from_lhs -> (f32) { + %offset0 = index.mul %column, %lhs_s0 : index + %offset1 = index.mul %coord1, %lhs_s1 : index + %offset2 = index.mul %coord2, %lhs_s2 : index + %offset3 = index.mul %coord3, %lhs_s3 : index + %offset01 = index.add %offset0, %offset1 : index + %offset23 = index.add %offset2, %offset3 : index + %lhs_index0 = index.add %offset01, %offset23 : index + %lhs_index, %lhs_view_span = index.assume %lhs_index0, %lhs_span_checked [lt(%lhs_index0, %lhs_span_checked)] : index, index + %lhs_view = buffer.view %lhs_noalias[%base] : buffer -> view<[%lhs_view_span]xf32> + %loaded = view.load %lhs_view[%lhs_index] : view<[%lhs_view_span]xf32> -> f32 + scf.yield %loaded : f32 + } else { + %rhs_column = index.sub %column, %lhs_width : index + %offset0 = index.mul %rhs_column, %rhs_s0 : index + %offset1 = index.mul %coord1, %rhs_s1 : index + %offset2 = index.mul %coord2, %rhs_s2 : index + %offset3 = index.mul %coord3, %rhs_s3 : index + %offset01 = index.add %offset0, %offset1 : index + %offset23 = index.add %offset2, %offset3 : index + %rhs_index0 = index.add %offset01, %offset23 : index + %rhs_index, %rhs_view_span = index.assume %rhs_index0, %rhs_span_checked [lt(%rhs_index0, %rhs_span_checked)] : index, index + %rhs_view = buffer.view %rhs_noalias[%base] : buffer -> view<[%rhs_view_span]xf32> + %loaded = view.load %rhs_view[%rhs_index] : view<[%rhs_view_span]xf32> -> f32 + scf.yield %loaded : f32 + } + view.store %value, %output_view[%linear] : f32, view<[%element_count]xf32> + } + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/flash_attention_decode_split_f32_f16_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/flash_attention_decode_split_f32_f16_wmma.loom new file mode 100644 index 000000000000..e3026cef3f2f --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/flash_attention_decode_split_f32_f16_wmma.loom @@ -0,0 +1,1532 @@ +// Qwen3 MoE grouped-query decode FlashAttention. +// +// Each workgroup processes one 64-token KV block for all GQA query heads that +// share a KV head. The workgroups publish online-softmax state, then the last +// arrival folds every block and resets the per-KV-head completion counter for +// the next invocation. Packing GQA heads removes redundant K/V traffic while +// split-K preserves enough parallelism for decode without another dispatch. +template.decl @ggml.flash_attention.decode_split.pack_completed_q8(%key_value_head: index, %output: buffer, %q8_output: buffer) + +template.decl @ggml.flash_attention.decode_split.produce_partials(%key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer) + +template.decl @ggml.flash_attention.decode_split.produce_partials.active(%key_value_token_count: index, %partial_block_capacity0: index, %launched_block_count0: index, %query: buffer, %key: buffer, %value: buffer, %lane_mask: f32, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer) + +template.decl @ggml.flash_attention.decode_split.reduce_completed.cooperative(%partial_block_capacity0: index, %active_block_count0: index, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %output: buffer) + +template.decl @ggml.flash_attention.decode_split.reduce_completed.direct(%partial_block_capacity0: index, %active_block_count0: index, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %output: buffer) +template.decl @ggml.flash_attention.decode_split.reduce_completed.multipass(%partial_block_capacity0: index, %active_block_count0: index, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %output: buffer) + +template.decl @ggml.flash_attention.decode_split.reduce_fused(%key_value_token_capacity: index, %partial_block_capacity0: index, %producer_block_count0: index, %publish_q8: i1, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %completion_counter: buffer, %output: buffer, %q8_output: buffer) + +template.decl @ggml.quantize_q8_1_x4.publish_vector4(%publish_word: i1, %token_output_byte_base: offset, %channel: index, %values: vector<4xf32>, %scratch_values: view<256xf32>, %scratch_d: view<32xf32>, %output: buffer) + +amdgpu.target @ggml_flash_attention_decode_split_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.flash_attention.query_head_count : %value: index where [range(%value, 1, 64)] + +config.decl @ggml.flash_attention.key_value_head_count : %value: index where [range(%value, 1, 64)] + +config.decl @ggml.flash_attention.qk_head_size : %value: index where [range(%value, 16, 576), mul(%value, 16)] + +config.decl @ggml.flash_attention.value_head_size : %value: index where [range(%value, 64, 512), mul(%value, 64)] + +config.decl @ggml.flash_attention.attention_scale : f32 + +// Maximum K/V storage capacity available to the compiled kernel. +config.decl @ggml.flash_attention.decode.key_value_token_capacity : %value: index where [range(%value, 64, 262144)] + +// Computes one active online-softmax partial for a 64-row KV block. Both the +// fused short-context export and the two-dispatch long-context export reach +// this body through the block-classifying producer below, keeping their +// different binding contracts honest without duplicating the attention math. +// Partial storage retains its capacity-specialized block-axis stride while the +// launched block count bounds issue-time work. Keeping those values distinct +// preserves constant address arithmetic across changing visible prefixes. +template.def<@ggml.flash_attention.decode_split.produce_partials.active> device @ggml_flash_attention_decode_split_produce_active_partials_body_f32_f16_wmma(%key_value_token_count: index, %partial_block_capacity0: index, %launched_block_count0: index, %query: buffer, %key: buffer, %value: buffer, %lane_mask: f32, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer) { + %bounded_key_value_token_count = index.assume %key_value_token_count [range(%key_value_token_count, 1, 262144)] : index + %partial_block_capacity, %launched_block_count = index.assume %partial_block_capacity0, %launched_block_count0 [range(%partial_block_capacity0, 1, 4096), range(%launched_block_count0, 1, 4096), le(%launched_block_count0, %partial_block_capacity0)] : index, index + %query_head_count = config.get @ggml.flash_attention.query_head_count : index + %key_value_head_count = config.get @ggml.flash_attention.key_value_head_count : index + %qk_head_size0 = config.get @ggml.flash_attention.qk_head_size : index + %qk_head_size = index.assume %qk_head_size0 [range(%qk_head_size0, 16, 576), mul(%qk_head_size0, 16)] : index + %value_head_size0 = config.get @ggml.flash_attention.value_head_size : index + %value_head_size = index.assume %value_head_size0 [range(%value_head_size0, 64, 512), mul(%value_head_size0, 64)] : index + %workgroup_x0 = kernel.workgroup.id : index + %workgroup_y0 = kernel.workgroup.id : index + %workgroup_x_in_launch = index.assume %workgroup_x0 [range(%workgroup_x0, 0, 511)] : index + %workgroup_y = index.assume %workgroup_y0 [range(%workgroup_y0, 0, 63)] : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 3)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c127 = index.constant 127 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %f16_bytes = index.constant 2 : offset + %score_stage_bytes = index.constant 6144 : offset + %probability_stage_bytes = index.constant 3072 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %c0_i32 = scalar.constant 0 : i32 + %masked_value_stage_bytes = index.constant 1024 : offset + %c1_i32 = scalar.constant 1 : i32 + %attention_scale = config.get @ggml.flash_attention.attention_scale : f32 + %c0_f16 = scalar.constant 0.0 : f16 + %c0_f16x8 = vector.constant 0.0 : vector<8xf16> + %c0_f16x16 = vector.constant 0.0 : vector<16xf16> + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %workgroup_x, %launch_partial_block_capacity, %launch_launched_block_count = index.assume %workgroup_x_in_launch, %partial_block_capacity, %launched_block_count [lt(%workgroup_x_in_launch, %launched_block_count)] : index, index, index + %tail_key_value_token_count = index.rem %bounded_key_value_token_count, %c64 : index + %has_no_tail = index.cmp eq, %tail_key_value_token_count, %c0 : index + %active_padded_key_value_token_count = index.add %bounded_key_value_token_count, %c63 : index + %active_key_value_block_count = index.div %active_padded_key_value_token_count, %c64 : index + %last_block_ordinal = index.sub %active_key_value_block_count, %c1 : index + %block_ordinal = index.add %workgroup_x, %c0 : index + %is_not_last_block = index.cmp ne, %block_ordinal, %last_block_ordinal : index + %is_full_block = scalar.ori %has_no_tail, %is_not_last_block : i1 + %key_value_head, %launch_key_value_head_count = index.assume %workgroup_y, %key_value_head_count [lt(%workgroup_y, %key_value_head_count)] : index, index + %query_heads_per_key_value_head0 = index.div %query_head_count, %key_value_head_count : index + %query_heads_per_key_value_head = index.assume %query_heads_per_key_value_head0 [range(%query_heads_per_key_value_head0, 1, 16)] : index + %query_head_base = index.mul %key_value_head, %query_heads_per_key_value_head : index + %key_width = index.mul %key_value_head_count, %qk_head_size : index + %value_width = index.mul %key_value_head_count, %value_head_size : index + %key_head_base = index.mul %key_value_head, %qk_head_size : index + %value_head_base = index.mul %key_value_head, %value_head_size : index + %padded_value_head_size = index.add %value_head_size, %c127 : index + %output_tile_count = index.div %padded_value_head_size, %c128 : index + %output_stage_size = index.mul %output_tile_count, %c128 : index + %query_stage_stride = index.add %qk_head_size, %c8 : index + %query_stage_element_count = index.mul %c16, %query_stage_stride : index + %query_stage_bytes = index.scale %query_stage_element_count, %f16_bytes : index, offset -> offset + %query_element_count = index.mul %c16, %qk_head_size : index + %query_load_iteration_count = index.div %query_element_count, %c256 : index + %product_stage_element_count = index.mul %c16, %output_stage_size : index + %product_stage_bytes = index.scale %product_stage_element_count, %f16_bytes : index, offset -> offset + %key_origin = index.mul %block_ordinal, %c64 : index + %subgroup_score_column = index.mul %subgroup, %c16 : index + %subgroup_query_row = index.mul %subgroup, %c4 : index + %query_row0 = index.add %subgroup_query_row, %c0 : index + %query_row1 = index.add %subgroup_query_row, %c1 : index + %query_row2 = index.add %subgroup_query_row, %c2 : index + %query_row3 = index.add %subgroup_query_row, %c3 : index + %query_head0 = index.add %query_head_base, %query_row0 : index + %query_head1 = index.add %query_head_base, %query_row1 : index + %query_head2 = index.add %query_head_base, %query_row2 : index + %query_head3 = index.add %query_head_base, %query_row3 : index + %query_head_valid0 = index.cmp ult, %query_head0, %query_head_count : index + %query_head_valid1 = index.cmp ult, %query_head1, %query_head_count : index + %query_head_valid2 = index.cmp ult, %query_head2, %query_head_count : index + %query_head_valid3 = index.cmp ult, %query_head3, %query_head_count : index + %query_valid = vector.from_elements %query_head_valid0, %query_head_valid1, %query_head_valid2, %query_head_valid3 : vector<4xi1> + %lane_has_output = index.cmp ult, %lane, %c32 : index + %lane_is_zero = index.cmp eq, %lane, %c0 : index + %query_transposed_layout = encoding.layout.strided [1, %query_stage_stride] : encoding + %probability_transposed_layout = encoding.layout.strided [1, 24] : encoding + %query_noalias, %key_noalias, %value_noalias, %partial_max_noalias, %partial_sum_noalias, %partial_output_noalias = buffer.assume.noalias %query, %key, %value, %partial_max, %partial_sum, %partial_output : buffer, buffer, buffer, buffer, buffer, buffer + %query_aligned = buffer.assume.alignment %query_noalias {minimum_alignment = 16} : buffer + %key_aligned = buffer.assume.alignment %key_noalias {minimum_alignment = 16} : buffer + %value_aligned = buffer.assume.alignment %value_noalias {minimum_alignment = 16} : buffer + %partial_max_aligned = buffer.assume.alignment %partial_max_noalias {minimum_alignment = 16} : buffer + %partial_sum_aligned = buffer.assume.alignment %partial_sum_noalias {minimum_alignment = 16} : buffer + %partial_output_aligned = buffer.assume.alignment %partial_output_noalias {minimum_alignment = 16} : buffer + %query_view = buffer.view %query_aligned[%c0_offset] : buffer -> view<[%query_head_count]x[%qk_head_size]xf32> + %key_view = buffer.view %key_aligned[%c0_offset] : buffer -> view<[%bounded_key_value_token_count]x[%key_width]xf16> + %value_view = buffer.view %value_aligned[%c0_offset] : buffer -> view<[%bounded_key_value_token_count]x[%value_width]xf16> + %partial_max_view = buffer.view %partial_max_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + %partial_sum_view = buffer.view %partial_sum_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + %partial_output_view = buffer.view %partial_output_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%launch_partial_block_capacity]x16x[%value_head_size]xf16> + %query_stage = buffer.alloca align(16) %query_stage_bytes : buffer + %score_stage = buffer.alloca align(16) %score_stage_bytes : buffer + %probability_stage = buffer.alloca align(16) %probability_stage_bytes : buffer + %product_stage = buffer.alloca align(16) %product_stage_bytes : buffer + %tail_value_stage = buffer.alloca align(16) %product_stage_bytes : buffer + %visibility_stage_bytes = index.constant 256 : offset + %visibility_stage = buffer.alloca align(16) %visibility_stage_bytes : buffer + %visibility_view = buffer.view %visibility_stage[%c0_offset] : buffer -> view<64xf32> + // Per-subgroup 16x32 V panels for key tiles that mix visible and masked + // keys; the score stage is free once probabilities are published. + %masked_value_stage_offset = index.scale %subgroup, %masked_value_stage_bytes : index, offset -> offset + %masked_value_stage_view = buffer.view %score_stage[%masked_value_stage_offset] : buffer -> view<16x32xf16> + %query_stage_view = buffer.view %query_stage[%c0_offset] : buffer -> view<16x[%query_stage_stride]xf16> + %query_transposed_view = buffer.view %query_stage[%c0_offset] : buffer -> view<[%qk_head_size]x16xf16, %query_transposed_layout> + %score_stage_view = buffer.view %score_stage[%c0_offset] : buffer -> view<64x24xf32> + %probability_stage_view = buffer.view %probability_stage[%c0_offset] : buffer -> view<16x64xf16, %probability_transposed_layout> + %product_stage_view = buffer.view %product_stage[%c0_offset] : buffer -> view<16x[%output_stage_size]xf16> + %tail_key_stage_view = buffer.view %product_stage[%c0_offset] : buffer -> view<64x16xf16> + %tail_value_stage_view = buffer.view %tail_value_stage[%c0_offset] : buffer -> view<16x[%output_stage_size]xf16> + scf.for %load_iteration = [%c0 to %query_load_iteration_count step %c1] unroll { + %linear = index.madd %load_iteration, %c256, %workitem : index + %local_query_row = index.div %linear, %qk_head_size : index + %query_channel = index.rem %linear, %qk_head_size : index + %local_query_head = index.add %query_head_base, %local_query_row : index + %local_query_head_valid = index.cmp ult, %local_query_head, %query_head_count : index + %query_value = scf.if %local_query_head_valid -> (f16) { + %loaded = view.load %query_view[%local_query_head, %query_channel] : view<[%query_head_count]x[%qk_head_size]xf32> -> f32 + %scaled = scalar.mulf %loaded, %attention_scale : f32 + %truncated = scalar.fptrunc %scaled : f32 to f16 + scf.yield %truncated : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %query_value, %query_stage_view[%local_query_row, %query_channel] : f16, view<16x[%query_stage_stride]xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %score_init_values = vector.constant 0.0 : vector<4xf32> + %score_init = vector.fragment %score_init_values shape [%m, %n] : vector<4xf32> + %score_fragment = scf.if %is_full_block -> (vector<4xf32>) { + %full_score_fragment = scf.for %head_tile = [%c0 to %qk_head_size step %c16](%score_accumulator = %score_init : vector<4xf32>) -> (vector<4xf32>) unroll { + // Bound this view by the proven tile end so the full vector footprint is + // visible without treating physical tail padding as logical storage. + %score_key_origin0 = index.add %key_origin, %subgroup_score_column : index + %score_key_end0 = index.add %score_key_origin0, %c16 : index + %score_key_origin, %score_key_end = index.assume %score_key_origin0, %score_key_end0 [le(%score_key_end0, %bounded_key_value_token_count)] : index, index + %full_key_view = buffer.view %key_aligned[%c0_offset] : buffer -> view<[%score_key_end]x[%key_width]xf16> + %key_channel = index.add %key_head_base, %head_tile : index + %key_fragment = vector.fragment.load %full_key_view[%score_key_origin, %key_channel] shape [%m, %k] : view<[%score_key_end]x[%key_width]xf16> -> vector<16xf16> + %query_fragment = vector.fragment.load %query_transposed_view[%head_tile, %c0] shape [%k, %n] : view<[%qk_head_size]x16xf16, %query_transposed_layout> -> vector<16xf16> + %next_score_accumulator = vector.mma %key_fragment, %query_fragment, %score_accumulator : vector<16xf16>, vector<16xf16>, vector<4xf32> + scf.yield %next_score_accumulator : vector<4xf32> + } + scf.yield %full_score_fragment : vector<4xf32> + } else { + // Stage one 64x16 K panel at a time. Every physical load is guarded, so + // callers need no initialized padding beyond the logical KV length. + %tail_score_fragment = scf.for %head_tile = [%c0 to %qk_head_size step %c16](%score_accumulator = %score_init : vector<4xf32>) -> (vector<4xf32>) { + scf.for %load_iteration = [%c0 to %c4 step %c1] unroll { + %linear = index.madd %load_iteration, %c256, %workitem : index + %key_row = index.div %linear, %c16 : index + %head_channel = index.rem %linear, %c16 : index + %key_token = index.add %key_origin, %key_row : index + %key_valid = index.cmp ult, %key_token, %bounded_key_value_token_count : index + %key_value = scf.if %key_valid -> (f16) { + %head_channel0 = index.add %head_tile, %head_channel : index + %key_channel = index.add %key_head_base, %head_channel0 : index + %loaded = view.load %key_view[%key_token, %key_channel] : view<[%bounded_key_value_token_count]x[%key_width]xf16> -> f16 + scf.yield %loaded : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %key_value, %tail_key_stage_view[%key_row, %head_channel] : f16, view<64x16xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %key_fragment = vector.fragment.load %tail_key_stage_view[%subgroup_score_column, %c0] shape [%m, %k] : view<64x16xf16> -> vector<16xf16> + %query_fragment = vector.fragment.load %query_transposed_view[%head_tile, %c0] shape [%k, %n] : view<[%qk_head_size]x16xf16, %query_transposed_layout> -> vector<16xf16> + %next_score_accumulator = vector.mma %key_fragment, %query_fragment, %score_accumulator : vector<16xf16>, vector<16xf16>, vector<4xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %next_score_accumulator : vector<4xf32> + } + scf.yield %tail_score_fragment : vector<4xf32> + } + vector.fragment.store %score_fragment, %score_stage_view[%subgroup_score_column, %c0] shape [%m, %n] : vector<4xf32>, view<64x24xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + %key_token = index.add %key_origin, %lane : index + %raw_score0 = view.load %score_stage_view[%lane, %query_row0] : view<64x24xf32> -> f32 + %raw_score1 = view.load %score_stage_view[%lane, %query_row1] : view<64x24xf32> -> f32 + %raw_score2 = view.load %score_stage_view[%lane, %query_row2] : view<64x24xf32> -> f32 + %raw_score3 = view.load %score_stage_view[%lane, %query_row3] : view<64x24xf32> -> f32 + %key_valid = index.cmp ult, %key_token, %bounded_key_value_token_count : index + %score_valid0 = scalar.andi %query_head_valid0, %key_valid : i1 + %score_valid1 = scalar.andi %query_head_valid1, %key_valid : i1 + %score_valid2 = scalar.andi %query_head_valid2, %key_valid : i1 + %score_valid3 = scalar.andi %query_head_valid3, %key_valid : i1 + %masked_score0 = scf.if %score_valid0 -> (f32) { + %score = scalar.addf %raw_score0, %lane_mask : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score1 = scf.if %score_valid1 -> (f32) { + %score = scalar.addf %raw_score1, %lane_mask : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score2 = scf.if %score_valid2 -> (f32) { + %score = scalar.addf %raw_score2, %lane_mask : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score3 = scf.if %score_valid3 -> (f32) { + %score = scalar.addf %raw_score3, %lane_mask : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_scores = vector.from_elements %masked_score0, %masked_score1, %masked_score2, %masked_score3 : vector<4xf32> + %block_max = kernel.subgroup.reduce %masked_scores : vector<4xf32> + %score_delta = vector.subf %masked_scores, %block_max : vector<4xf32> + %raw_probability = vector.expf %score_delta : vector<4xf32> + %probability = vector.select %query_valid, %raw_probability, %c0_f32x4 : vector<4xf32> + %block_sum = kernel.subgroup.reduce %probability : vector<4xf32> + %probability_f16 = vector.fptrunc %probability : vector<4xf32> to vector<4xf16> + %probability0 = vector.extract %probability_f16[0] : vector<4xf16> -> f16 + %probability1 = vector.extract %probability_f16[1] : vector<4xf16> -> f16 + %probability2 = vector.extract %probability_f16[2] : vector<4xf16> -> f16 + %probability3 = vector.extract %probability_f16[3] : vector<4xf16> -> f16 + view.store %probability0, %probability_stage_view[%query_row0, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability1, %probability_stage_view[%query_row1, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability2, %probability_stage_view[%query_row2, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability3, %probability_stage_view[%query_row3, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + // Masked keys must reach P*V as V = +0, not just P = 0. The F16 WMMA + // result depended on the sign of V at keys whose probability is zero + // (measured: flipping the sign of one masked V row changed one output + // element by one F16 ulp). Rows past the sequence end hold whatever an + // earlier request left there, so identical requests gave different logits. + // Key tiles with every key visible read V directly, fully masked tiles + // use a zero V fragment, and mixed tiles stage V with masked rows zeroed. + %is_first_subgroup = index.cmp eq, %subgroup, %c0 : index + scf.if %is_first_subgroup { + view.store %lane_mask, %visibility_view[%lane] : f32, view<64xf32> + } + %lane_key_visible = scalar.cmpf ogt, %lane_mask, %negative_large : f32 + %lane_key_tile = index.div %lane, %c16 : index + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.for %output_tile = [%c0 to %output_tile_count step %c1] unroll { + %output_tile_channel = index.mul %output_tile, %c128 : index + %subgroup_output_base = index.mul %subgroup, %c32 : index + %subgroup_output_channel0 = index.add %output_tile_channel, %subgroup_output_base : index + %subgroup_output_channel = index.assume %subgroup_output_channel0 [range(%subgroup_output_channel0, 0, 480)] : index + %subgroup_output_channel1_0 = index.add %subgroup_output_channel, %c16 : index + %subgroup_output_channel1 = index.assume %subgroup_output_channel1_0 [range(%subgroup_output_channel1_0, 16, 496)] : index + %subgroup_output_channel_valid = index.cmp ult, %subgroup_output_channel, %value_head_size : index + %subgroup_output_channel1_valid = index.cmp ult, %subgroup_output_channel1, %value_head_size : index + %lane_output_base = index.mul %lane, %c4 : index + %lane_output_channel0 = index.add %output_tile_channel, %lane_output_base : index + %lane_output_channel = index.assume %lane_output_channel0 [range(%lane_output_channel0, 0, 636)] : index + %lane_output_channel_valid = index.cmp ult, %lane_output_channel, %value_head_size : index + %lane_publishes_output = scalar.andi %lane_has_output, %lane_output_channel_valid : i1 + %safe_lane_output_channel0 = scf.select %lane_output_channel_valid, %lane_output_channel, %c0 : index + %safe_lane_output_channel_end0 = index.add %safe_lane_output_channel0, %c4 : index + %safe_lane_output_channel, %safe_lane_output_channel_end = index.assume %safe_lane_output_channel0, %safe_lane_output_channel_end0 [le(%safe_lane_output_channel_end0, %value_head_size), le(%safe_lane_output_channel_end0, %output_stage_size)] : index, index + %product_init0 = vector.fragment %c0_f16x8 shape [%m, %n] : vector<8xf16> + %product_init1 = vector.fragment %c0_f16x8 shape [%m, %n] : vector<8xf16> + %safe_subgroup_output_channel0 = scf.select %subgroup_output_channel_valid, %subgroup_output_channel, %c0 : index + %safe_subgroup_output_channel1_0 = scf.select %subgroup_output_channel1_valid, %subgroup_output_channel1, %c0 : index + %safe_subgroup_output_channel_end0 = index.add %safe_subgroup_output_channel0, %c16 : index + %safe_subgroup_output_channel1_end0 = index.add %safe_subgroup_output_channel1_0, %c16 : index + %safe_subgroup_output_channel, %safe_subgroup_output_channel_end = index.assume %safe_subgroup_output_channel0, %safe_subgroup_output_channel_end0 [le(%safe_subgroup_output_channel_end0, %value_head_size), le(%safe_subgroup_output_channel_end0, %output_stage_size)] : index, index + %safe_subgroup_output_channel1, %safe_subgroup_output_channel1_end = index.assume %safe_subgroup_output_channel1_0, %safe_subgroup_output_channel1_end0 [le(%safe_subgroup_output_channel1_end0, %value_head_size), le(%safe_subgroup_output_channel1_end0, %output_stage_size)] : index, index + %value_channel0 = index.add %value_head_base, %safe_subgroup_output_channel : index + %value_channel1 = index.add %value_head_base, %safe_subgroup_output_channel1 : index + %value_channel0_end0 = index.add %value_head_base, %safe_subgroup_output_channel_end : index + %value_channel1_end0 = index.add %value_head_base, %safe_subgroup_output_channel1_end : index + %safe_value_channel0, %safe_value_channel0_end = index.assume %value_channel0, %value_channel0_end0 [le(%value_channel0_end0, %value_width)] : index, index + %safe_value_channel1, %safe_value_channel1_end = index.assume %value_channel1, %value_channel1_end0 [le(%value_channel1_end0, %value_width)] : index, index + %last_value_fragment_channel = index.sub %value_width, %c16 : index + %clamped_value_channel0 = index.min %safe_value_channel0, %last_value_fragment_channel : index + %clamped_value_channel1 = index.min %safe_value_channel1, %last_value_fragment_channel : index + %last_stage_fragment_channel = index.sub %output_stage_size, %c16 : index + %clamped_stage_channel0 = index.min %safe_subgroup_output_channel, %last_stage_fragment_channel : index + %clamped_stage_channel1 = index.min %safe_subgroup_output_channel1, %last_stage_fragment_channel : index + %product_fragment0, %product_fragment1 = scf.if %is_full_block -> (vector<8xf16>, vector<8xf16>) { + %full_product_fragment0, %full_product_fragment1 = scf.for %key_tile = [%c0 to %c64 step %c16](%product_accumulator0 = %product_init0 : vector<8xf16>, %product_accumulator1 = %product_init1 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>) unroll { + %value_token0 = index.add %key_origin, %key_tile : index + %value_token_end0 = index.add %value_token0, %c16 : index + %value_token, %value_token_end = index.assume %value_token0, %value_token_end0 [le(%value_token_end0, %bounded_key_value_token_count)] : index, index + %full_value_view = buffer.view %value_aligned[%c0_offset] : buffer -> view<[%value_token_end]x[%value_width]xf16> + %probability_fragment = vector.fragment.load %probability_stage_view[%c0, %key_tile] shape [%m, %k] : view<16x64xf16, %probability_transposed_layout> -> vector<16xf16> + %key_tile_ordinal = index.div %key_tile, %c16 : index + %lane_in_key_tile = index.cmp eq, %lane_key_tile, %key_tile_ordinal : index + %lane_outside_key_tile = index.cmp ne, %lane_key_tile, %key_tile_ordinal : index + %lane_tile_visible = scalar.ori %lane_outside_key_tile, %lane_key_visible : i1 + %lane_tile_has_visible = scalar.andi %lane_in_key_tile, %lane_key_visible : i1 + // engine#314 candidate 2: always take the staging path. Constant predicates keep the + // existing if/else structure (so the scf yields stay in the right regions) while + // removing the subgroup collectives the AMDGPU scheduler cannot order. + %key_tile_all_visible = scalar.constant false : i1 + %key_tile_any_visible = scalar.constant true : i1 + %next_product_accumulator0, %next_product_accumulator1 = scf.if %key_tile_all_visible -> (vector<8xf16>, vector<8xf16>) { + %value_fragment0 = vector.fragment.load %full_value_view[%value_token, %clamped_value_channel0] shape [%k, %n] : view<[%value_token_end]x[%value_width]xf16> -> vector<16xf16> + %value_fragment1 = vector.fragment.load %full_value_view[%value_token, %clamped_value_channel1] shape [%k, %n] : view<[%value_token_end]x[%value_width]xf16> -> vector<16xf16> + %direct_accumulator0 = vector.mma %probability_fragment, %value_fragment0, %product_accumulator0 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %direct_accumulator1 = vector.mma %probability_fragment, %value_fragment1, %product_accumulator1 : vector<16xf16>, vector<16xf16>, vector<8xf16> + scf.yield %direct_accumulator0, %direct_accumulator1 : vector<8xf16>, vector<8xf16> + } else { + %masked_accumulator0, %masked_accumulator1 = scf.if %key_tile_any_visible -> (vector<8xf16>, vector<8xf16>) { + scf.for %stage_iteration = [%c0 to %c8 step %c1] unroll { + %stage_linear = index.madd %stage_iteration, %c64, %lane : index + %stage_key_row = index.div %stage_linear, %c32 : index + %stage_column = index.rem %stage_linear, %c32 : index + %stage_second_fragment = index.cmp uge, %stage_column, %c16 : index + %stage_fragment_column = index.rem %stage_column, %c16 : index + %stage_channel_base = scf.select %stage_second_fragment, %clamped_value_channel1, %clamped_value_channel0 : index + %stage_channel0 = index.add %stage_channel_base, %stage_fragment_column : index + %stage_channel = index.assume %stage_channel0 [lt(%stage_channel0, %value_width)] : index + %stage_block_key0 = index.add %key_tile, %stage_key_row : index + %stage_block_key = index.min %stage_block_key0, %c63 : index + %stage_key_mask = view.load %visibility_view[%stage_block_key] : view<64xf32> -> f32 + %stage_key_visible = scalar.cmpf ogt, %stage_key_mask, %negative_large : f32 + %stage_value = scf.if %stage_key_visible -> (f16) { + %stage_token0 = index.add %value_token, %stage_key_row : index + %stage_token = index.assume %stage_token0 [lt(%stage_token0, %value_token_end)] : index + %loaded = view.load %full_value_view[%stage_token, %stage_channel] : view<[%value_token_end]x[%value_width]xf16> -> f16 + scf.yield %loaded : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %stage_value, %masked_value_stage_view[%stage_key_row, %stage_column] : f16, view<16x32xf16> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + %staged_value_fragment0 = vector.fragment.load %masked_value_stage_view[%c0, %c0] shape [%k, %n] : view<16x32xf16> -> vector<16xf16> + %staged_value_fragment1 = vector.fragment.load %masked_value_stage_view[%c0, %c16] shape [%k, %n] : view<16x32xf16> -> vector<16xf16> + %staged_accumulator0 = vector.mma %probability_fragment, %staged_value_fragment0, %product_accumulator0 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %staged_accumulator1 = vector.mma %probability_fragment, %staged_value_fragment1, %product_accumulator1 : vector<16xf16>, vector<16xf16>, vector<8xf16> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.yield %staged_accumulator0, %staged_accumulator1 : vector<8xf16>, vector<8xf16> + } else { + %zero_value_fragment = vector.fragment %c0_f16x16 shape [%k, %n] : vector<16xf16> + %zero_accumulator0 = vector.mma %probability_fragment, %zero_value_fragment, %product_accumulator0 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %zero_accumulator1 = vector.mma %probability_fragment, %zero_value_fragment, %product_accumulator1 : vector<16xf16>, vector<16xf16>, vector<8xf16> + scf.yield %zero_accumulator0, %zero_accumulator1 : vector<8xf16>, vector<8xf16> + } + scf.yield %masked_accumulator0, %masked_accumulator1 : vector<8xf16>, vector<8xf16> + } + scf.yield %next_product_accumulator0, %next_product_accumulator1 : vector<8xf16>, vector<8xf16> + } + scf.yield %full_product_fragment0, %full_product_fragment1 : vector<8xf16>, vector<8xf16> + } else { + %tail_product_fragment0, %tail_product_fragment1 = scf.for %key_tile = [%c0 to %c64 step %c16](%product_accumulator0 = %product_init0 : vector<8xf16>, %product_accumulator1 = %product_init1 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>) unroll { + scf.for %load_iteration = [%c0 to %c8 step %c1] unroll { + %linear = index.madd %load_iteration, %c256, %workitem : index + %key_row = index.div %linear, %c128 : index + %tile_value_channel = index.rem %linear, %c128 : index + %value_channel = index.add %output_tile_channel, %tile_value_channel : index + %value_token0 = index.add %key_origin, %key_tile : index + %value_token = index.add %value_token0, %key_row : index + %value_token_valid = index.cmp ult, %value_token, %bounded_key_value_token_count : index + %value_channel_valid = index.cmp ult, %value_channel, %value_head_size : index + %value_valid0 = scalar.andi %value_token_valid, %value_channel_valid : i1 + %block_value_key0 = index.add %key_tile, %key_row : index + %block_value_key = index.min %block_value_key0, %c63 : index + %value_key_mask = view.load %visibility_view[%block_value_key] : view<64xf32> -> f32 + %value_key_visible = scalar.cmpf ogt, %value_key_mask, %negative_large : f32 + %value_valid = scalar.andi %value_valid0, %value_key_visible : i1 + %value_element = scf.if %value_valid -> (f16) { + %global_value_channel = index.add %value_head_base, %value_channel : index + %loaded = view.load %value_view[%value_token, %global_value_channel] : view<[%bounded_key_value_token_count]x[%value_width]xf16> -> f16 + scf.yield %loaded : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %value_element, %tail_value_stage_view[%key_row, %value_channel] : f16, view<16x[%output_stage_size]xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %probability_fragment = vector.fragment.load %probability_stage_view[%c0, %key_tile] shape [%m, %k] : view<16x64xf16, %probability_transposed_layout> -> vector<16xf16> + // Bound each fragment without changing the physical LDS row stride. + %value_fragment0 = vector.fragment.load %tail_value_stage_view[%c0, %clamped_stage_channel0] shape [%k, %n] : view<16x[%output_stage_size]xf16> -> vector<16xf16> + %value_fragment1 = vector.fragment.load %tail_value_stage_view[%c0, %clamped_stage_channel1] shape [%k, %n] : view<16x[%output_stage_size]xf16> -> vector<16xf16> + %next_product_accumulator0 = vector.mma %probability_fragment, %value_fragment0, %product_accumulator0 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %next_product_accumulator1 = vector.mma %probability_fragment, %value_fragment1, %product_accumulator1 : vector<16xf16>, vector<16xf16>, vector<8xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %next_product_accumulator0, %next_product_accumulator1 : vector<8xf16>, vector<8xf16> + } + scf.yield %tail_product_fragment0, %tail_product_fragment1 : vector<8xf16>, vector<8xf16> + } + vector.fragment.store %product_fragment0, %product_stage_view[%c0, %subgroup_output_channel] shape [%m, %n] : vector<8xf16>, view<16x[%output_stage_size]xf16> + vector.fragment.store %product_fragment1, %product_stage_view[%c0, %subgroup_output_channel1] shape [%m, %n] : vector<8xf16>, view<16x[%output_stage_size]xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.if %lane_publishes_output { + %block_output0 = vector.load %product_stage_view[%query_row0, %safe_lane_output_channel] : view<16x[%output_stage_size]xf16> -> vector<4xf16> + %block_output1 = vector.load %product_stage_view[%query_row1, %safe_lane_output_channel] : view<16x[%output_stage_size]xf16> -> vector<4xf16> + %block_output2 = vector.load %product_stage_view[%query_row2, %safe_lane_output_channel] : view<16x[%output_stage_size]xf16> -> vector<4xf16> + %block_output3 = vector.load %product_stage_view[%query_row3, %safe_lane_output_channel] : view<16x[%output_stage_size]xf16> -> vector<4xf16> + scf.if %query_head_valid0 { + vector.store %block_output0, %partial_output_view[%key_value_head, %block_ordinal, %query_row0, %safe_lane_output_channel] : vector<4xf16>, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16x[%value_head_size]xf16> + } + scf.if %query_head_valid1 { + vector.store %block_output1, %partial_output_view[%key_value_head, %block_ordinal, %query_row1, %safe_lane_output_channel] : vector<4xf16>, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16x[%value_head_size]xf16> + } + scf.if %query_head_valid2 { + vector.store %block_output2, %partial_output_view[%key_value_head, %block_ordinal, %query_row2, %safe_lane_output_channel] : vector<4xf16>, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16x[%value_head_size]xf16> + } + scf.if %query_head_valid3 { + vector.store %block_output3, %partial_output_view[%key_value_head, %block_ordinal, %query_row3, %safe_lane_output_channel] : vector<4xf16>, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16x[%value_head_size]xf16> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + } + scf.if %lane_is_zero { + scf.if %query_head_valid0 { + %maximum = vector.extract %block_max[0] : vector<4xf32> -> f32 + %sum = vector.extract %block_sum[0] : vector<4xf32> -> f32 + view.store %maximum, %partial_max_view[%key_value_head, %block_ordinal, %query_row0] : f32, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + view.store %sum, %partial_sum_view[%key_value_head, %block_ordinal, %query_row0] : f32, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + } + scf.if %query_head_valid1 { + %maximum = vector.extract %block_max[1] : vector<4xf32> -> f32 + %sum = vector.extract %block_sum[1] : vector<4xf32> -> f32 + view.store %maximum, %partial_max_view[%key_value_head, %block_ordinal, %query_row1] : f32, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + view.store %sum, %partial_sum_view[%key_value_head, %block_ordinal, %query_row1] : f32, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + } + scf.if %query_head_valid2 { + %maximum = vector.extract %block_max[2] : vector<4xf32> -> f32 + %sum = vector.extract %block_sum[2] : vector<4xf32> -> f32 + view.store %maximum, %partial_max_view[%key_value_head, %block_ordinal, %query_row2] : f32, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + view.store %sum, %partial_sum_view[%key_value_head, %block_ordinal, %query_row2] : f32, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + } + scf.if %query_head_valid3 { + %maximum = vector.extract %block_max[3] : vector<4xf32> -> f32 + %sum = vector.extract %block_sum[3] : vector<4xf32> -> f32 + view.store %maximum, %partial_max_view[%key_value_head, %block_ordinal, %query_row3] : f32, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + view.store %sum, %partial_sum_view[%key_value_head, %block_ordinal, %query_row3] : f32, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + } + } + template.return +} + +// Fully masked splits publish the online-softmax identity without running QK +// or P*V. Valid additive masks contain finite F16 values or negative infinity, +// and every finite F16 value is greater than -1e30. Comparing after extension +// therefore distinguishes active rows from masked rows exactly. +template.def<@ggml.flash_attention.decode_split.produce_partials> device @ggml_flash_attention_decode_split_produce_partials_body_f32_f16_wmma(%key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer) { + %bounded_key_value_token_count = index.assume %key_value_token_count [range(%key_value_token_count, 1, 262144)] : index + %query_head_count = config.get @ggml.flash_attention.query_head_count : index + %key_value_head_count = config.get @ggml.flash_attention.key_value_head_count : index + %value_head_size0 = config.get @ggml.flash_attention.value_head_size : index + %value_head_size = index.assume %value_head_size0 [range(%value_head_size0, 64, 512), mul(%value_head_size0, 64)] : index + %workgroup_x0 = kernel.workgroup.id : index + %workgroup_y0 = kernel.workgroup.id : index + %workgroup_x_in_launch = index.assume %workgroup_x0 [range(%workgroup_x0, 0, 511)] : index + %workgroup_y = index.assume %workgroup_y0 [range(%workgroup_y0, 0, 63)] : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 3)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c127 = index.constant 127 : index + %c128 = index.constant 128 : index + %c0_offset = index.constant 0 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %c0_f16x4 = vector.constant 0.0 : vector<4xf16> + %key_value_token_capacity = config.get @ggml.flash_attention.decode.key_value_token_capacity : index + %c63 = index.constant 63 : index + %padded_key_value_token_capacity = index.add %key_value_token_capacity, %c63 : index + %key_value_block_count0 = index.div %padded_key_value_token_capacity, %c64 : index + %key_value_block_count = index.assume %key_value_block_count0 [range(%key_value_block_count0, 1, 4096)] : index + %workgroup_x, %launch_key_value_block_count = index.assume %workgroup_x_in_launch, %key_value_block_count [lt(%workgroup_x_in_launch, %key_value_block_count)] : index, index + %block_ordinal = index.add %workgroup_x, %c0 : index + %key_origin = index.mul %block_ordinal, %c64 : index + %key_token = index.add %key_origin, %lane : index + %key_valid = index.cmp ult, %key_token, %bounded_key_value_token_count : index + %mask_aligned = buffer.assume.alignment %mask {minimum_alignment = 16} : buffer + %mask_view = buffer.view %mask_aligned[%c0_offset] : buffer -> view<[%bounded_key_value_token_count]xf16> + %lane_mask = scf.if %key_valid -> (f32) { + %mask_f16 = view.load %mask_view[%key_token] : view<[%bounded_key_value_token_count]xf16> -> f16 + %mask_f32 = scalar.extf %mask_f16 : f16 to f32 + scf.yield %mask_f32 : f32 + } else { + scf.yield %negative_large : f32 + } + %block_mask_maximum = kernel.workgroup.reduce %lane_mask : f32 + %block_has_attention = scalar.cmpf ogt, %block_mask_maximum, %negative_large : f32 + scf.if %block_has_attention { + template.apply<@ggml.flash_attention.decode_split.produce_partials.active>(%bounded_key_value_token_count, %launch_key_value_block_count, %launch_key_value_block_count, %query, %key, %value, %lane_mask, %partial_max, %partial_sum, %partial_output) : (index, index, index, buffer, buffer, buffer, f32, buffer, buffer, buffer) + } else { + %query_heads_per_key_value_head0 = index.div %query_head_count, %key_value_head_count : index + %query_heads_per_key_value_head = index.assume %query_heads_per_key_value_head0 [range(%query_heads_per_key_value_head0, 1, 16)] : index + %key_value_head = index.add %workgroup_y, %c0 : index + %query_head_base = index.mul %key_value_head, %query_heads_per_key_value_head : index + %subgroup_query_row = index.mul %subgroup, %c4 : index + %query_row0 = index.add %subgroup_query_row, %c0 : index + %query_row1 = index.add %subgroup_query_row, %c1 : index + %query_row2 = index.add %subgroup_query_row, %c2 : index + %query_row3 = index.add %subgroup_query_row, %c3 : index + %query_head0 = index.add %query_head_base, %query_row0 : index + %query_head1 = index.add %query_head_base, %query_row1 : index + %query_head2 = index.add %query_head_base, %query_row2 : index + %query_head3 = index.add %query_head_base, %query_row3 : index + %query_head_valid0 = index.cmp ult, %query_head0, %query_head_count : index + %query_head_valid1 = index.cmp ult, %query_head1, %query_head_count : index + %query_head_valid2 = index.cmp ult, %query_head2, %query_head_count : index + %query_head_valid3 = index.cmp ult, %query_head3, %query_head_count : index + %padded_value_head_size = index.add %value_head_size, %c127 : index + %output_tile_count = index.div %padded_value_head_size, %c128 : index + %lane_output_base = index.mul %lane, %c4 : index + %lane_has_output = index.cmp ult, %lane, %c32 : index + %lane_is_zero = index.cmp eq, %lane, %c0 : index + %partial_max_noalias, %partial_sum_noalias, %partial_output_noalias = buffer.assume.noalias %partial_max, %partial_sum, %partial_output : buffer, buffer, buffer + %partial_max_aligned = buffer.assume.alignment %partial_max_noalias {minimum_alignment = 16} : buffer + %partial_sum_aligned = buffer.assume.alignment %partial_sum_noalias {minimum_alignment = 16} : buffer + %partial_output_aligned = buffer.assume.alignment %partial_output_noalias {minimum_alignment = 16} : buffer + %partial_max_view = buffer.view %partial_max_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + %partial_sum_view = buffer.view %partial_sum_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + %partial_output_view = buffer.view %partial_output_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%key_value_block_count]x16x[%value_head_size]xf16> + scf.for %output_tile = [%c0 to %output_tile_count step %c1] unroll { + %output_tile_channel = index.mul %output_tile, %c128 : index + %lane_output_channel = index.add %output_tile_channel, %lane_output_base : index + %lane_output_channel_valid = index.cmp ult, %lane_output_channel, %value_head_size : index + %lane_publishes_output = scalar.andi %lane_has_output, %lane_output_channel_valid : i1 + scf.if %lane_publishes_output { + scf.if %query_head_valid0 { + vector.store %c0_f16x4, %partial_output_view[%key_value_head, %block_ordinal, %query_row0, %lane_output_channel] : vector<4xf16>, view<[%key_value_head_count]x[%key_value_block_count]x16x[%value_head_size]xf16> + } + scf.if %query_head_valid1 { + vector.store %c0_f16x4, %partial_output_view[%key_value_head, %block_ordinal, %query_row1, %lane_output_channel] : vector<4xf16>, view<[%key_value_head_count]x[%key_value_block_count]x16x[%value_head_size]xf16> + } + scf.if %query_head_valid2 { + vector.store %c0_f16x4, %partial_output_view[%key_value_head, %block_ordinal, %query_row2, %lane_output_channel] : vector<4xf16>, view<[%key_value_head_count]x[%key_value_block_count]x16x[%value_head_size]xf16> + } + scf.if %query_head_valid3 { + vector.store %c0_f16x4, %partial_output_view[%key_value_head, %block_ordinal, %query_row3, %lane_output_channel] : vector<4xf16>, view<[%key_value_head_count]x[%key_value_block_count]x16x[%value_head_size]xf16> + } + } + } + scf.if %lane_is_zero { + scf.if %query_head_valid0 { + view.store %negative_large, %partial_max_view[%key_value_head, %block_ordinal, %query_row0] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + view.store %c0_f32, %partial_sum_view[%key_value_head, %block_ordinal, %query_row0] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + } + scf.if %query_head_valid1 { + view.store %negative_large, %partial_max_view[%key_value_head, %block_ordinal, %query_row1] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + view.store %c0_f32, %partial_sum_view[%key_value_head, %block_ordinal, %query_row1] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + } + scf.if %query_head_valid2 { + view.store %negative_large, %partial_max_view[%key_value_head, %block_ordinal, %query_row2] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + view.store %c0_f32, %partial_sum_view[%key_value_head, %block_ordinal, %query_row2] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + } + scf.if %query_head_valid3 { + view.store %negative_large, %partial_max_view[%key_value_head, %block_ordinal, %query_row3] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + view.store %c0_f32, %partial_sum_view[%key_value_head, %block_ordinal, %query_row3] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + } + } + } + template.return +} + +// Up to four split-K blocks are cheapest to fold directly. Each workitem owns +// one output element, so the reducer needs no LDS or subgroup synchronization. +template.def<@ggml.flash_attention.decode_split.reduce_completed.direct> device @ggml_flash_attention_decode_split_reduce_completed_direct_f32(%partial_block_capacity0: index, %active_block_count0: index, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %output: buffer) where [range(%partial_block_capacity0, 1, 4)] { + %partial_block_capacity, %active_block_count = index.assume %partial_block_capacity0, %active_block_count0 [range(%partial_block_capacity0, 1, 4), range(%active_block_count0, 1, 4), le(%active_block_count0, %partial_block_capacity0)] : index, index + %query_head_count = config.get @ggml.flash_attention.query_head_count : index + %key_value_head_count = config.get @ggml.flash_attention.key_value_head_count : index + %value_head_size0 = config.get @ggml.flash_attention.value_head_size : index + %value_head_size = index.assume %value_head_size0 [range(%value_head_size0, 64, 512), mul(%value_head_size0, 64)] : index + %workgroup_y0 = kernel.workgroup.id : index + %key_value_head = index.assume %workgroup_y0 [range(%workgroup_y0, 0, 63)] : index + %workitem = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %query_heads_per_key_value_head0 = index.div %query_head_count, %key_value_head_count : index + %query_heads_per_key_value_head = index.assume %query_heads_per_key_value_head0 [range(%query_heads_per_key_value_head0, 1, 16)] : index + %query_head_base = index.mul %key_value_head, %query_heads_per_key_value_head : index + %partial_max_aligned = buffer.assume.alignment %partial_max {minimum_alignment = 16} : buffer + %partial_sum_aligned = buffer.assume.alignment %partial_sum {minimum_alignment = 16} : buffer + %partial_output_aligned = buffer.assume.alignment %partial_output {minimum_alignment = 16} : buffer + %output_aligned = buffer.assume.alignment %output {minimum_alignment = 16} : buffer + %partial_max_view = buffer.view %partial_max_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> + %partial_sum_view = buffer.view %partial_sum_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> + %partial_output_view = buffer.view %partial_output_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%partial_block_capacity]x16x[%value_head_size]xf16> + %output_view = buffer.view %output_aligned[%c0_offset] : buffer -> view<[%query_head_count]x[%value_head_size]xf32> + %partial_element_count = index.mul %query_heads_per_key_value_head, %value_head_size : index + scf.for %linear = [%workitem to %partial_element_count step %c256] { + %query_row = index.div %linear, %value_head_size : index + %output_channel = index.rem %linear, %value_head_size : index + %query_head = index.add %query_head_base, %query_row : index + %query_head_valid = index.cmp ult, %query_head, %query_head_count : index + scf.if %query_head_valid { + %maximum = scf.for %block = [%c0 to %active_block_count step %c1](%running_maximum = %negative_large : f32) -> (f32) unroll schedule(interleaved) { + %block_maximum = view.load %partial_max_view[%key_value_head, %block, %query_row] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + %next_maximum = scalar.maxnumf %running_maximum, %block_maximum : f32 + scf.yield %next_maximum : f32 + } + %sum, %unnormalized_output = scf.for %block = [%c0 to %active_block_count step %c1](%running_sum = %c0_f32 : f32, %running_output = %c0_f32 : f32) -> (f32, f32) unroll schedule(interleaved) { + %block_maximum = view.load %partial_max_view[%key_value_head, %block, %query_row] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + %delta = scalar.subf %block_maximum, %maximum : f32 + %scale = scalar.expf %delta : f32 + %block_sum = view.load %partial_sum_view[%key_value_head, %block, %query_row] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + %scaled_sum = scalar.mulf %block_sum, %scale : f32 + %next_sum = scalar.addf %running_sum, %scaled_sum : f32 + %block_output_f16 = view.load %partial_output_view[%key_value_head, %block, %query_row, %output_channel] : view<[%key_value_head_count]x[%partial_block_capacity]x16x[%value_head_size]xf16> -> f16 + %block_output = scalar.extf %block_output_f16 : f16 to f32 + %scaled_output = scalar.mulf %block_output, %scale : f32 + %next_output = scalar.addf %running_output, %scaled_output : f32 + scf.yield %next_sum, %next_output : f32, f32 + } + %normalized_output = scalar.divf %unnormalized_output, %sum : f32 + view.store %normalized_output, %output_view[%query_head, %output_channel] : f32, view<[%query_head_count]x[%value_head_size]xf32> + } + } + template.return +} + +// Longer bounded contexts amortize a cooperative reducer. Each wave folds two +// query rows, stages the per-block scales in LDS, and writes two output channels +// per lane. +template.def<@ggml.flash_attention.decode_split.reduce_completed.cooperative> device @ggml_flash_attention_decode_split_reduce_completed_cooperative_f32(%partial_block_capacity0: index, %active_block_count0: index, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %output: buffer) { + %partial_block_capacity, %active_block_count = index.assume %partial_block_capacity0, %active_block_count0 [range(%partial_block_capacity0, 1, 32), range(%active_block_count0, 1, 32), le(%active_block_count0, %partial_block_capacity0)] : index, index + %query_head_count = config.get @ggml.flash_attention.query_head_count : index + %key_value_head_count = config.get @ggml.flash_attention.key_value_head_count : index + %value_head_size0 = config.get @ggml.flash_attention.value_head_size : index + %value_head_size = index.assume %value_head_size0 [range(%value_head_size0, 64, 512), mul(%value_head_size0, 64)] : index + %workgroup_y0 = kernel.workgroup.id : index + %key_value_head = index.assume %workgroup_y0 [range(%workgroup_y0, 0, 63)] : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 3)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c127 = index.constant 127 : index + %c128 = index.constant 128 : index + %c0_offset = index.constant 0 : offset + %scale_stage_bytes = index.constant 1024 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %c0_f16x2 = vector.constant 0.0 : vector<2xf16> + %c0_f32x2 = vector.constant 0.0 : vector<2xf32> + %query_heads_per_key_value_head0 = index.div %query_head_count, %key_value_head_count : index + %query_heads_per_key_value_head = index.assume %query_heads_per_key_value_head0 [range(%query_heads_per_key_value_head0, 1, 16)] : index + %query_head_base = index.mul %key_value_head, %query_heads_per_key_value_head : index + %padded_value_head_size = index.add %value_head_size, %c127 : index + %output_tile_count = index.div %padded_value_head_size, %c128 : index + %lane_output_base = index.mul %lane, %c2 : index + %lane_has_block = index.cmp ult, %lane, %active_block_count : index + %partial_max_aligned = buffer.assume.alignment %partial_max {minimum_alignment = 16} : buffer + %partial_sum_aligned = buffer.assume.alignment %partial_sum {minimum_alignment = 16} : buffer + %partial_output_aligned = buffer.assume.alignment %partial_output {minimum_alignment = 16} : buffer + %output_aligned = buffer.assume.alignment %output {minimum_alignment = 16} : buffer + %partial_max_view = buffer.view %partial_max_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> + %partial_sum_view = buffer.view %partial_sum_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> + %partial_output_view = buffer.view %partial_output_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%partial_block_capacity]x16x[%value_head_size]xf16> + %output_view = buffer.view %output_aligned[%c0_offset] : buffer -> view<[%query_head_count]x[%value_head_size]xf32> + %scale_stage = buffer.alloca align(16) %scale_stage_bytes : buffer + %scale_stage_view = buffer.view %scale_stage[%c0_offset] : buffer -> view<4x32x2xf32> + %padded_reduction_row_count = index.add %query_heads_per_key_value_head, %c7 : index + %reduction_phase_count = index.div %padded_reduction_row_count, %c8 : index + scf.for %phase = [%c0 to %reduction_phase_count step %c1] unroll { + %phase_row_base = index.mul %phase, %c8 : index + %subgroup_row_base = index.mul %subgroup, %c2 : index + %query_row0 = index.add %phase_row_base, %subgroup_row_base : index + %query_row1 = index.add %query_row0, %c1 : index + %query_head0 = index.add %query_head_base, %query_row0 : index + %query_head1 = index.add %query_head0, %c1 : index + // A KV head owns only its own query rows; rows past that belong to the next + // KV head's workgroup, and reducing them here overwrites its output. + %query_row_owned0 = index.cmp ult, %query_row0, %query_heads_per_key_value_head : index + %query_row_owned1 = index.cmp ult, %query_row1, %query_heads_per_key_value_head : index + %query_head_in_range0 = index.cmp ult, %query_head0, %query_head_count : index + %query_head_in_range1 = index.cmp ult, %query_head1, %query_head_count : index + %query_head_valid0 = scalar.andi %query_row_owned0, %query_head_in_range0 : i1 + %query_head_valid1 = scalar.andi %query_row_owned1, %query_head_in_range1 : i1 + %safe_query_row0 = scf.select %query_head_valid0, %query_row0, %c0 : index + %safe_query_row1 = scf.select %query_head_valid1, %query_row1, %c0 : index + %reducer_active0 = scalar.andi %lane_has_block, %query_head_valid0 : i1 + %reducer_active1 = scalar.andi %lane_has_block, %query_head_valid1 : i1 + %local_maximum0 = scf.if %reducer_active0 -> (f32) { + %block_maximum = view.load %partial_max_view[%key_value_head, %lane, %query_row0] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + scf.yield %block_maximum : f32 + } else { + scf.yield %negative_large : f32 + } + %local_maximum1 = scf.if %reducer_active1 -> (f32) { + %block_maximum = view.load %partial_max_view[%key_value_head, %lane, %query_row1] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + scf.yield %block_maximum : f32 + } else { + scf.yield %negative_large : f32 + } + %local_maximums = vector.from_elements %local_maximum0, %local_maximum1 : vector<2xf32> + %maximums = kernel.subgroup.reduce %local_maximums : vector<2xf32> + %maximum0 = vector.extract %maximums[0] : vector<2xf32> -> f32 + %maximum1 = vector.extract %maximums[1] : vector<2xf32> -> f32 + %local_sum0, %scale0 = scf.if %reducer_active0 -> (f32, f32) { + %block_maximum = view.load %partial_max_view[%key_value_head, %lane, %query_row0] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + %delta = scalar.subf %block_maximum, %maximum0 : f32 + %scale = scalar.expf %delta : f32 + %block_sum = view.load %partial_sum_view[%key_value_head, %lane, %query_row0] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + %scaled_sum = scalar.mulf %block_sum, %scale : f32 + scf.yield %scaled_sum, %scale : f32, f32 + } else { + scf.yield %c0_f32, %c0_f32 : f32, f32 + } + %local_sum1, %scale1 = scf.if %reducer_active1 -> (f32, f32) { + %block_maximum = view.load %partial_max_view[%key_value_head, %lane, %query_row1] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + %delta = scalar.subf %block_maximum, %maximum1 : f32 + %scale = scalar.expf %delta : f32 + %block_sum = view.load %partial_sum_view[%key_value_head, %lane, %query_row1] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + %scaled_sum = scalar.mulf %block_sum, %scale : f32 + scf.yield %scaled_sum, %scale : f32, f32 + } else { + scf.yield %c0_f32, %c0_f32 : f32, f32 + } + scf.if %lane_has_block { + %bounded_block_lane = index.assume %lane [range(%lane, 0, 31)] : index + %scales = vector.from_elements %scale0, %scale1 : vector<2xf32> + vector.store %scales, %scale_stage_view[%subgroup, %bounded_block_lane, %c0] : vector<2xf32>, view<4x32x2xf32> + } + %local_sums = vector.from_elements %local_sum0, %local_sum1 : vector<2xf32> + %sums = kernel.subgroup.reduce %local_sums : vector<2xf32> + %sum0 = vector.extract %sums[0] : vector<2xf32> -> f32 + %sum1 = vector.extract %sums[1] : vector<2xf32> -> f32 + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.for %output_tile = [%c0 to %output_tile_count step %c1] unroll { + %output_tile_channel = index.mul %output_tile, %c128 : index + %lane_output_channel = index.add %output_tile_channel, %lane_output_base : index + %lane_output_channel_valid = index.cmp ult, %lane_output_channel, %value_head_size : index + %safe_lane_output_channel = scf.select %lane_output_channel_valid, %lane_output_channel, %c0 : index + %unnormalized_output0, %unnormalized_output1 = scf.for %block = [%c0 to %active_block_count step %c1](%running_output0 = %c0_f32x2 : vector<2xf32>, %running_output1 = %c0_f32x2 : vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) unroll(%c4) schedule(interleaved) { + %scales = vector.load %scale_stage_view[%subgroup, %block, %c0] : view<4x32x2xf32> -> vector<2xf32> + %scale0 = vector.extract %scales[0] : vector<2xf32> -> f32 + %scale1 = vector.extract %scales[1] : vector<2xf32> -> f32 + %block_output0_f16 = vector.load %partial_output_view[%key_value_head, %block, %safe_query_row0, %safe_lane_output_channel] : view<[%key_value_head_count]x[%partial_block_capacity]x16x[%value_head_size]xf16> -> vector<2xf16> + %block_output0 = vector.extf %block_output0_f16 : vector<2xf16> to vector<2xf32> + %scale_vector0 = vector.splat %scale0 : vector<2xf32> + %scaled_output0 = vector.mulf %block_output0, %scale_vector0 : vector<2xf32> + %next_output0 = vector.addf %running_output0, %scaled_output0 : vector<2xf32> + %block_output1_f16 = vector.load %partial_output_view[%key_value_head, %block, %safe_query_row1, %safe_lane_output_channel] : view<[%key_value_head_count]x[%partial_block_capacity]x16x[%value_head_size]xf16> -> vector<2xf16> + %block_output1 = vector.extf %block_output1_f16 : vector<2xf16> to vector<2xf32> + %scale_vector1 = vector.splat %scale1 : vector<2xf32> + %scaled_output1 = vector.mulf %block_output1, %scale_vector1 : vector<2xf32> + %next_output1 = vector.addf %running_output1, %scaled_output1 : vector<2xf32> + scf.yield %next_output0, %next_output1 : vector<2xf32>, vector<2xf32> + } + %sum_vector0 = vector.splat %sum0 : vector<2xf32> + %sum_vector1 = vector.splat %sum1 : vector<2xf32> + %normalized_output0 = vector.divf %unnormalized_output0, %sum_vector0 : vector<2xf32> + %normalized_output1 = vector.divf %unnormalized_output1, %sum_vector1 : vector<2xf32> + %publish_output0 = scalar.andi %query_head_valid0, %lane_output_channel_valid : i1 + %publish_output1 = scalar.andi %query_head_valid1, %lane_output_channel_valid : i1 + scf.if %publish_output0 { + vector.store %normalized_output0, %output_view[%query_head0, %lane_output_channel] : vector<2xf32>, view<[%query_head_count]x[%value_head_size]xf32> + } + scf.if %publish_output1 { + vector.store %normalized_output1, %output_view[%query_head1, %lane_output_channel] : vector<2xf32>, view<[%query_head_count]x[%value_head_size]xf32> + } + } + } + template.return +} + +// Packs the query heads owned by one completed KV head into GGML's Q8_1 x4 +// layout. Each 32-workitem cohort owns one contiguous 128-element query row; +// the production 8:1 GQA ratio therefore fills all 256 workitems. Smaller or +// larger valid ratios use the same phase loop without making barriers +// conditional. +template.def<@ggml.flash_attention.decode_split.pack_completed_q8> device @ggml_flash_attention_decode_pack_completed_key_value_head_q8_1_x4(%key_value_head: index, %output: buffer, %q8_output: buffer) { + %query_head_count = config.get @ggml.flash_attention.query_head_count : index + %key_value_head_count = config.get @ggml.flash_attention.key_value_head_count : index + %value_head_size0 = config.get @ggml.flash_attention.value_head_size : index + %value_head_size = index.assume %value_head_size0 [range(%value_head_size0, 64, 512), mul(%value_head_size0, 64)] : index + %query_heads_per_key_value_head0 = index.div %query_head_count, %key_value_head_count : index + %query_heads_per_key_value_head = index.assume %query_heads_per_key_value_head0 [range(%query_heads_per_key_value_head0, 1, 16)] : index + %workitem0 = kernel.workitem.id : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 255)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c32 = index.constant 32 : index + %c127 = index.constant 127 : index + %c128 = index.constant 128 : index + %scratch_d_byte_add = index.constant 1024 : offset + %scratch_bytes = index.constant 1152 : offset + %c0_offset = index.constant 0 : offset + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %output_noalias, %q8_output_noalias = buffer.assume.noalias %output, %q8_output : buffer, buffer + %output_aligned = buffer.assume.alignment %output_noalias {minimum_alignment = 16} : buffer + %q8_output_aligned = buffer.assume.alignment %q8_output_noalias {minimum_alignment = 16} : buffer + %output_view = buffer.view %output_aligned[%c0_offset] : buffer -> view<[%query_head_count]x[%value_head_size]xf32> + %scratch = buffer.alloca align(16) %scratch_bytes : buffer + %scratch_values = buffer.view %scratch[%c0_offset] : buffer -> view<256xf32> + %scratch_d = buffer.view %scratch[%scratch_d_byte_add] : buffer -> view<32xf32> + %padded_value_head_size = index.add %value_head_size, %c127 : index + %head_tile_count = index.div %padded_value_head_size, %c128 : index + %group_in_phase = index.div %workitem, %c32 : index + %lane_in_group0 = index.rem %workitem, %c32 : index + %lane_in_group = index.assume %lane_in_group0 [range(%lane_in_group0, 0, 31)] : index + %block_in_group0 = index.div %lane_in_group, %c8 : index + %block_in_group = index.assume %block_in_group0 [range(%block_in_group0, 0, 3)] : index + %word_in_block0 = index.rem %lane_in_group, %c8 : index + %word_in_block = index.assume %word_in_block0 [range(%word_in_block0, 0, 7)] : index + %padded_phase_count = index.add %query_heads_per_key_value_head, %c7 : index + %phase_count = index.div %padded_phase_count, %c8 : index + scf.for %phase = [%c0 to %phase_count step %c1] unroll { + %phase_group_base = index.mul %phase, %c8 : index + %group_in_key_value_head = index.add %phase_group_base, %group_in_phase : index + %valid_group = index.cmp ult, %group_in_key_value_head, %query_heads_per_key_value_head : index + %key_value_query_head_base = index.mul %key_value_head, %query_heads_per_key_value_head : index + %query_head0 = index.add %key_value_query_head_base, %group_in_key_value_head : index + %safe_query_head = scf.select %valid_group, %query_head0, %c0 : index + scf.for %head_tile = [%c0 to %head_tile_count step %c1] unroll { + %head_tile_channel = index.mul %head_tile, %c128 : index + %block_element_add = index.mul %block_in_group, %c32 : index + %word_element_add = index.mul %word_in_block, %c4 : index + %input_channel_add = index.add %block_element_add, %word_element_add : index + %input_channel0 = index.add %head_tile_channel, %input_channel_add : index + %input_channel = index.assume %input_channel0 [range(%input_channel0, 0, 508)] : index + %valid_channel = index.cmp ult, %input_channel, %value_head_size : index + %publish_word = scalar.andi %valid_group, %valid_channel : i1 + %input_values = scf.if %publish_word -> (vector<4xf32>) { + %values = vector.load %output_view[%safe_query_head, %input_channel] : view<[%query_head_count]x[%value_head_size]xf32> -> vector<4xf32> + scf.yield %values : vector<4xf32> + } else { + scf.yield %c0_f32x4 : vector<4xf32> + } + %flattened_head_channel = index.mul %safe_query_head, %value_head_size : index + %flattened_channel = index.add %flattened_head_channel, %input_channel : index + template.apply<@ggml.quantize_q8_1_x4.publish_vector4>(%publish_word, %c0_offset, %flattened_channel, %input_values, %scratch_values, %scratch_d, %q8_output_aligned) : (i1, offset, index, vector<4xf32>, view<256xf32>, view<32xf32>, buffer) + } + } + template.return +} + +// Completes bounded decode contexts inside the last arriving producer +// workgroup. Four KV-head workgroups reduce their own query heads while all +// other producers retire, erasing a second dispatch and its execution barrier. +// Capacity selects the algorithm and partial layout; producer count owns the +// issue-time completion threshold and active reduction prefix. +// Multi-pass reducer for > 32 producer blocks (capacity > 2048). Same +// reduction as @..reduce_completed.cooperative, but the block dimension is +// lane-strided instead of one-block-per-lane, so any producer block count +// works; the per-block normalization scale is recomputed from the intact +// partial_max in the output pass, avoiding a global-memory round-trip. +template.def<@ggml.flash_attention.decode_split.reduce_completed.multipass> device @ggml_flash_attention_decode_split_reduce_completed_multipass_f32(%partial_block_capacity0: index, %active_block_count0: index, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %output: buffer) where [range(%partial_block_capacity0, 33, 4096), range(%active_block_count0, 33, 4096)] { + %partial_block_capacity, %active_block_count = index.assume %partial_block_capacity0, %active_block_count0 [range(%partial_block_capacity0, 33, 4096), range(%active_block_count0, 33, 4096), le(%active_block_count0, %partial_block_capacity0)] : index, index + %query_head_count = config.get @ggml.flash_attention.query_head_count : index + %key_value_head_count = config.get @ggml.flash_attention.key_value_head_count : index + %value_head_size0 = config.get @ggml.flash_attention.value_head_size : index + %value_head_size = index.assume %value_head_size0 [range(%value_head_size0, 64, 512), mul(%value_head_size0, 64)] : index + %workgroup_y0 = kernel.workgroup.id : index + %key_value_head = index.assume %workgroup_y0 [range(%workgroup_y0, 0, 63)] : index + %workitem = kernel.workitem.id : index + %subgroup = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c64 = index.constant 64 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %reduction_stage_bytes = index.constant 8 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %query_heads_per_key_value_head0 = index.div %query_head_count, %key_value_head_count : index + %query_heads_per_key_value_head = index.assume %query_heads_per_key_value_head0 [range(%query_heads_per_key_value_head0, 1, 16)] : index + %query_head_base = index.mul %key_value_head, %query_heads_per_key_value_head : index + %quads_per_row = index.div %value_head_size, %c4 : index + %total_quad_count = index.mul %query_heads_per_key_value_head, %quads_per_row : index + %padded_quad_count = index.add %total_quad_count, %c255 : index + %output_tile_count = index.div %padded_quad_count, %c256 : index + %is_first_subgroup = index.cmp eq, %subgroup, %c0 : index + %workitem_is_zero = index.cmp eq, %workitem, %c0 : index + %partial_max_noalias, %partial_sum_noalias, %partial_output_noalias, %output_noalias = buffer.assume.noalias %partial_max, %partial_sum, %partial_output, %output : buffer, buffer, buffer, buffer + %partial_max_aligned = buffer.assume.alignment %partial_max_noalias {minimum_alignment = 16} : buffer + %partial_sum_aligned = buffer.assume.alignment %partial_sum_noalias {minimum_alignment = 16} : buffer + %partial_output_aligned = buffer.assume.alignment %partial_output_noalias {minimum_alignment = 16} : buffer + %output_aligned = buffer.assume.alignment %output_noalias {minimum_alignment = 16} : buffer + %partial_max_view = buffer.view %partial_max_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> + %partial_sum_view = buffer.view %partial_sum_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> + %partial_output_view = buffer.view %partial_output_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%partial_block_capacity]x16x[%value_head_size]xf16> + %output_view = buffer.view %output_aligned[%c0_offset] : buffer -> view<[%query_head_count]x[%value_head_size]xf32> + %reduction_stage = buffer.alloca align(8) %reduction_stage_bytes : buffer + %reduction_stage_view = buffer.view %reduction_stage[%c0_offset] : buffer -> view<2xf32> + // Per-row, per-block normalisation scales and per-row sums, computed once by + // the lane-strided max/sum passes and reused by the vectorised output pass, + // so expf runs once per (row, block) instead of once per output element. The + // per-element accumulation order over blocks is unchanged. + %f32_bytes = index.constant 4 : offset + %scale_stage_elements = index.mul %query_heads_per_key_value_head, %active_block_count : index + %scale_stage_bytes = index.scale %scale_stage_elements, %f32_bytes : index, offset -> offset + %scale_stage = buffer.alloca align(16) %scale_stage_bytes : buffer + %scale_stage_view = buffer.view %scale_stage[%c0_offset] : buffer -> view<[%query_heads_per_key_value_head]x[%active_block_count]xf32> + %sum_stage_bytes = index.scale %query_heads_per_key_value_head, %f32_bytes : index, offset -> offset + %sum_stage = buffer.alloca align(16) %sum_stage_bytes : buffer + %sum_stage_view = buffer.view %sum_stage[%c0_offset] : buffer -> view<[%query_heads_per_key_value_head]xf32> + scf.for %query_row = [%c0 to %query_heads_per_key_value_head step %c1] { + %query_head = index.add %query_head_base, %query_row : index + %lane_maximum = scf.if %is_first_subgroup -> (f32) { + %maximum = scf.for %block = [%lane to %active_block_count step %c64](%running_maximum = %negative_large : f32) -> (f32) { + %block_maximum = view.load %partial_max_view[%key_value_head, %block, %query_row] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + %next_maximum = scalar.maxnumf %running_maximum, %block_maximum : f32 + scf.yield %next_maximum : f32 + } + scf.yield %maximum : f32 + } else { + scf.yield %negative_large : f32 + } + %subgroup_maximum = kernel.subgroup.reduce %lane_maximum : f32 + scf.if %workitem_is_zero { + view.store %subgroup_maximum, %reduction_stage_view[%c0] : f32, view<2xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %maximum = view.load %reduction_stage_view[%c0] : view<2xf32> -> f32 + %lane_sum = scf.if %is_first_subgroup -> (f32) { + %sum = scf.for %block = [%lane to %active_block_count step %c64](%running_sum = %c0_f32 : f32) -> (f32) { + %block_maximum = view.load %partial_max_view[%key_value_head, %block, %query_row] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + %delta = scalar.subf %block_maximum, %maximum : f32 + %scale = scalar.expf %delta : f32 + view.store %scale, %scale_stage_view[%query_row, %block] : f32, view<[%query_heads_per_key_value_head]x[%active_block_count]xf32> + %block_sum = view.load %partial_sum_view[%key_value_head, %block, %query_row] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + %scaled_sum = scalar.mulf %block_sum, %scale : f32 + %next_sum = scalar.addf %running_sum, %scaled_sum : f32 + scf.yield %next_sum : f32 + } + scf.yield %sum : f32 + } else { + scf.yield %c0_f32 : f32 + } + %subgroup_sum = kernel.subgroup.reduce %lane_sum : f32 + scf.if %workitem_is_zero { + view.store %subgroup_sum, %reduction_stage_view[%c1] : f32, view<2xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %sum = view.load %reduction_stage_view[%c1] : view<2xf32> -> f32 + view.store %sum, %sum_stage_view[%query_row] : f32, view<[%query_heads_per_key_value_head]xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + } + // Vectorised output pass: one 4-channel quad per workitem, spanning every + // query row so all four subgroups stay busy. Each channel accumulates its + // blocks in the same order as the scalar pass, so the values are identical. + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + scf.for %output_tile = [%c0 to %output_tile_count step %c1] { + %output_tile_base = index.mul %output_tile, %c256 : index + %quad = index.add %output_tile_base, %workitem : index + %quad_valid = index.cmp ult, %quad, %total_quad_count : index + scf.if %quad_valid { + %output_row0 = index.div %quad, %quads_per_row : index + %output_row = index.assume %output_row0 [range(%output_row0, 0, 16)] : index + %output_row_base = index.mul %output_row, %quads_per_row : index + %quad_in_row = index.sub %quad, %output_row_base : index + %output_channel0 = index.mul %quad_in_row, %c4 : index + %output_channel = index.assume %output_channel0 [range(%output_channel0, 0, 512)] : index + %query_head = index.add %query_head_base, %output_row : index + %row_sum = view.load %sum_stage_view[%output_row] : view<[%query_heads_per_key_value_head]xf32> -> f32 + %row_sum_vector = vector.splat %row_sum : vector<4xf32> + %unnormalized_output = scf.for %block = [%c0 to %active_block_count step %c1](%running_output = %c0_f32x4 : vector<4xf32>) -> (vector<4xf32>) unroll(%c4) schedule(interleaved) { + %scale = view.load %scale_stage_view[%output_row, %block] : view<[%query_heads_per_key_value_head]x[%active_block_count]xf32> -> f32 + %scale_vector = vector.splat %scale : vector<4xf32> + %block_output_f16 = vector.load %partial_output_view[%key_value_head, %block, %output_row, %output_channel] : view<[%key_value_head_count]x[%partial_block_capacity]x16x[%value_head_size]xf16> -> vector<4xf16> + %block_output = vector.extf %block_output_f16 : vector<4xf16> to vector<4xf32> + %scaled_output = vector.mulf %scale_vector, %block_output : vector<4xf32> + %next_output = vector.addf %running_output, %scaled_output : vector<4xf32> + scf.yield %next_output : vector<4xf32> + } + %normalized_output = vector.divf %unnormalized_output, %row_sum_vector : vector<4xf32> + vector.store %normalized_output, %output_view[%query_head, %output_channel] : vector<4xf32>, view<[%query_head_count]x[%value_head_size]xf32> + } + } + template.return +} + +template.def<@ggml.flash_attention.decode_split.reduce_fused> device priority(20) @ggml_flash_attention_decode_split_reduce_fused_direct_f32(%key_value_token_capacity: index, %partial_block_capacity0: index, %producer_block_count0: index, %publish_q8: i1, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %completion_counter: buffer, %output: buffer, %q8_output: buffer) where [range(%key_value_token_capacity, 64, 256)] { + %partial_block_capacity, %producer_block_count = index.assume %partial_block_capacity0, %producer_block_count0 [range(%partial_block_capacity0, 1, 4), range(%producer_block_count0, 1, 4), le(%producer_block_count0, %partial_block_capacity0)] : index, index + %key_value_head_count = config.get @ggml.flash_attention.key_value_head_count : index + %workgroup_y0 = kernel.workgroup.id : index + %key_value_head = index.assume %workgroup_y0 [range(%workgroup_y0, 0, 63)] : index + %workitem = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %counter_scratch_bytes = index.constant 4 : offset + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %partial_max_noalias, %partial_sum_noalias, %partial_output_noalias, %completion_counter_noalias, %output_noalias = buffer.assume.noalias %partial_max, %partial_sum, %partial_output, %completion_counter, %output : buffer, buffer, buffer, buffer, buffer + %completion_counter_aligned = buffer.assume.alignment %completion_counter_noalias {minimum_alignment = 16} : buffer + %completion_counter_view = buffer.view %completion_counter_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]xi32> + %counter_scratch = buffer.alloca align(4) %counter_scratch_bytes : buffer + %counter_scratch_view = buffer.view %counter_scratch[%c0_offset] : buffer -> view<1xi32> + %workitem_is_zero = index.cmp eq, %workitem, %c0 : index + // Publish every producer's partial stores before the leader advances one + // workgroup arrival. The last arrival then acquires every partial. + kernel.barrier scope(workgroup) ordering(release) + scf.if %workitem_is_zero { + %old_counter = view.atomic.rmw %c1_i32, %completion_counter_view[%key_value_head] {ordering = acq_rel, scope = device} : i32, view<[%key_value_head_count]xi32> -> i32 + view.store %old_counter, %counter_scratch_view[%c0] : i32, view<1xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %old_counter = view.load %counter_scratch_view[%c0] : view<1xi32> -> i32 + %key_value_block_count_i32 = index.cast %producer_block_count : index to i32 + %last_block_ordinal_i32 = scalar.subi %key_value_block_count_i32, %c1_i32 : i32 + %negative_key_value_block_count_i32 = scalar.subi %c0_i32, %key_value_block_count_i32 : i32 + %is_last_partition = scalar.cmpi eq, %old_counter, %last_block_ordinal_i32 : i32 + scf.if %is_last_partition { + kernel.barrier scope(workgroup) ordering(acquire) + template.apply<@ggml.flash_attention.decode_split.reduce_completed.direct>(%partial_block_capacity, %producer_block_count, %partial_max_noalias, %partial_sum_noalias, %partial_output_noalias, %output_noalias) : (index, index, buffer, buffer, buffer, buffer) + scf.if %publish_q8 { + kernel.barrier scope(workgroup) ordering(acq_rel) + kernel.barrier scope(workgroup) ordering(acq_rel) + template.apply<@ggml.flash_attention.decode_split.pack_completed_q8>(%key_value_head, %output_noalias, %q8_output) : (index, buffer, buffer) + } + // Do not expose the reset until every final F32 and Q8 store completes. + kernel.barrier scope(workgroup) ordering(release) + scf.if %workitem_is_zero { + view.store %c0_i32, %completion_counter_view[%key_value_head] : i32, view<[%key_value_head_count]xi32> + } + } + template.return +} + +// Contexts without a proven short bound use the cooperative completion path. +template.def<@ggml.flash_attention.decode_split.reduce_fused> device priority(10) @ggml_flash_attention_decode_split_reduce_fused_cooperative_f32(%key_value_token_capacity: index, %partial_block_capacity0: index, %producer_block_count0: index, %publish_q8: i1, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %completion_counter: buffer, %output: buffer, %q8_output: buffer) where [range(%key_value_token_capacity, 257, 2048)] { + %partial_block_capacity, %producer_block_count = index.assume %partial_block_capacity0, %producer_block_count0 [range(%partial_block_capacity0, 1, 32), range(%producer_block_count0, 1, 32), le(%producer_block_count0, %partial_block_capacity0)] : index, index + %key_value_head_count = config.get @ggml.flash_attention.key_value_head_count : index + %workgroup_y0 = kernel.workgroup.id : index + %key_value_head = index.assume %workgroup_y0 [range(%workgroup_y0, 0, 63)] : index + %workitem = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %counter_scratch_bytes = index.constant 4 : offset + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %partial_max_noalias, %partial_sum_noalias, %partial_output_noalias, %completion_counter_noalias, %output_noalias = buffer.assume.noalias %partial_max, %partial_sum, %partial_output, %completion_counter, %output : buffer, buffer, buffer, buffer, buffer + %completion_counter_aligned = buffer.assume.alignment %completion_counter_noalias {minimum_alignment = 16} : buffer + %completion_counter_view = buffer.view %completion_counter_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]xi32> + %counter_scratch = buffer.alloca align(4) %counter_scratch_bytes : buffer + %counter_scratch_view = buffer.view %counter_scratch[%c0_offset] : buffer -> view<1xi32> + %workitem_is_zero = index.cmp eq, %workitem, %c0 : index + // Publish every producer's partial stores before the leader advances one + // workgroup arrival. The last arrival then acquires every partial. + kernel.barrier scope(workgroup) ordering(release) + scf.if %workitem_is_zero { + %old_counter = view.atomic.rmw %c1_i32, %completion_counter_view[%key_value_head] {ordering = acq_rel, scope = device} : i32, view<[%key_value_head_count]xi32> -> i32 + view.store %old_counter, %counter_scratch_view[%c0] : i32, view<1xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %old_counter = view.load %counter_scratch_view[%c0] : view<1xi32> -> i32 + %key_value_block_count_i32 = index.cast %producer_block_count : index to i32 + %last_block_ordinal_i32 = scalar.subi %key_value_block_count_i32, %c1_i32 : i32 + %negative_key_value_block_count_i32 = scalar.subi %c0_i32, %key_value_block_count_i32 : i32 + %is_last_partition = scalar.cmpi eq, %old_counter, %last_block_ordinal_i32 : i32 + scf.if %is_last_partition { + kernel.barrier scope(workgroup) ordering(acquire) + template.apply<@ggml.flash_attention.decode_split.reduce_completed.cooperative>(%partial_block_capacity, %producer_block_count, %partial_max_noalias, %partial_sum_noalias, %partial_output_noalias, %output_noalias) : (index, index, buffer, buffer, buffer, buffer) + scf.if %publish_q8 { + kernel.barrier scope(workgroup) ordering(acq_rel) + kernel.barrier scope(workgroup) ordering(acq_rel) + template.apply<@ggml.flash_attention.decode_split.pack_completed_q8>(%key_value_head, %output_noalias, %q8_output) : (index, buffer, buffer) + } + // Do not expose the reset until every final F32 and Q8 store completes. + kernel.barrier scope(workgroup) ordering(release) + scf.if %workitem_is_zero { + view.store %c0_i32, %completion_counter_view[%key_value_head] : i32, view<[%key_value_head_count]xi32> + } + } + template.return +} + +// Long-context fused reducer: same self-synchronising completion-counter +// protocol as @..reduce_fused.cooperative, but the reduce itself is the +// multi-pass lane-strided reducer above, so any producer block count works. +template.def<@ggml.flash_attention.decode_split.reduce_fused> device priority(5) @ggml_flash_attention_decode_split_reduce_fused_multipass_f32(%key_value_token_capacity: index, %partial_block_capacity0: index, %producer_block_count0: index, %publish_q8: i1, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %completion_counter: buffer, %output: buffer, %q8_output: buffer) where [range(%key_value_token_capacity, 2049, 262144)] { + %partial_block_capacity, %producer_block_count = index.assume %partial_block_capacity0, %producer_block_count0 [range(%partial_block_capacity0, 33, 4096), range(%producer_block_count0, 33, 4096), le(%producer_block_count0, %partial_block_capacity0)] : index, index + %key_value_head_count = config.get @ggml.flash_attention.key_value_head_count : index + %workgroup_y0 = kernel.workgroup.id : index + %key_value_head = index.assume %workgroup_y0 [range(%workgroup_y0, 0, 63)] : index + %workitem = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %counter_scratch_bytes = index.constant 4 : offset + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %partial_max_noalias, %partial_sum_noalias, %partial_output_noalias, %completion_counter_noalias, %output_noalias = buffer.assume.noalias %partial_max, %partial_sum, %partial_output, %completion_counter, %output : buffer, buffer, buffer, buffer, buffer + %completion_counter_aligned = buffer.assume.alignment %completion_counter_noalias {minimum_alignment = 16} : buffer + %completion_counter_view = buffer.view %completion_counter_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]xi32> + %counter_scratch = buffer.alloca align(4) %counter_scratch_bytes : buffer + %counter_scratch_view = buffer.view %counter_scratch[%c0_offset] : buffer -> view<1xi32> + %workitem_is_zero = index.cmp eq, %workitem, %c0 : index + kernel.barrier scope(workgroup) ordering(release) + scf.if %workitem_is_zero { + %old_counter = view.atomic.rmw %c1_i32, %completion_counter_view[%key_value_head] {ordering = acq_rel, scope = device} : i32, view<[%key_value_head_count]xi32> -> i32 + view.store %old_counter, %counter_scratch_view[%c0] : i32, view<1xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %old_counter = view.load %counter_scratch_view[%c0] : view<1xi32> -> i32 + %key_value_block_count_i32 = index.cast %producer_block_count : index to i32 + %last_block_ordinal_i32 = scalar.subi %key_value_block_count_i32, %c1_i32 : i32 + %negative_key_value_block_count_i32 = scalar.subi %c0_i32, %key_value_block_count_i32 : i32 + %is_last_partition = scalar.cmpi eq, %old_counter, %last_block_ordinal_i32 : i32 + scf.if %is_last_partition { + kernel.barrier scope(workgroup) ordering(acquire) + template.apply<@ggml.flash_attention.decode_split.reduce_completed.multipass>(%partial_block_capacity, %producer_block_count, %partial_max_noalias, %partial_sum_noalias, %partial_output_noalias, %output_noalias) : (index, index, buffer, buffer, buffer, buffer) + scf.if %publish_q8 { + kernel.barrier scope(workgroup) ordering(acq_rel) + kernel.barrier scope(workgroup) ordering(acq_rel) + template.apply<@ggml.flash_attention.decode_split.pack_completed_q8>(%key_value_head, %output_noalias, %q8_output) : (index, buffer, buffer) + } + kernel.barrier scope(workgroup) ordering(release) + scf.if %workitem_is_zero { + view.store %c0_i32, %completion_counter_view[%key_value_head] : i32, view<[%key_value_head_count]xi32> + } + } + template.return +} + +// Short-context export: produce and reduce in one dispatch. +kernel.def target(@ggml_flash_attention_decode_split_gfx11_wave64) @ggml_flash_attention_decode_split_f32_f16_wmma(%key_value_token_count: index) { + %key_value_head_count = config.get @ggml.flash_attention.key_value_head_count : index + %key_value_token_capacity = config.get @ggml.flash_attention.decode.key_value_token_capacity : index + %c1 = index.constant 1 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %padded_key_value_token_capacity = index.add %key_value_token_capacity, %c63 : index + %key_value_block_count = index.div %padded_key_value_token_capacity, %c64 : index + kernel.launch.config workgroups(%key_value_block_count, %key_value_head_count, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %completion_counter: buffer, %output: buffer) { + %key_value_token_capacity = config.get @ggml.flash_attention.decode.key_value_token_capacity : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %key_value_token_count_in_range = index.assume %key_value_token_count [range(%key_value_token_count, 1, 262144)] : index + %bounded_key_value_token_count, %launch_key_value_token_capacity = index.assume %key_value_token_count_in_range, %key_value_token_capacity [le(%key_value_token_count_in_range, %key_value_token_capacity)] : index, index + %padded_key_value_token_capacity = index.add %launch_key_value_token_capacity, %c63 : index + %producer_block_count = index.div %padded_key_value_token_capacity, %c64 : index + %publish_q8 = scalar.constant false : i1 + template.apply<@ggml.flash_attention.decode_split.produce_partials>(%bounded_key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output) : (index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + template.apply<@ggml.flash_attention.decode_split.reduce_fused>(%launch_key_value_token_capacity, %producer_block_count, %producer_block_count, %publish_q8, %partial_max, %partial_sum, %partial_output, %completion_counter, %output, %partial_output) : (index, index, index, i1, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +// Short-context export that publishes the F32 attention result and the Q8_1 +// representation consumed by the following output projection. +kernel.def target(@ggml_flash_attention_decode_split_gfx11_wave64) @ggml_flash_attention_decode_split_f32_f16_wmma_next_q8(%key_value_token_count: index) { + %key_value_head_count = config.get @ggml.flash_attention.key_value_head_count : index + %key_value_token_capacity = config.get @ggml.flash_attention.decode.key_value_token_capacity : index + %c1 = index.constant 1 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %padded_key_value_token_capacity = index.add %key_value_token_capacity, %c63 : index + %key_value_block_count = index.div %padded_key_value_token_capacity, %c64 : index + kernel.launch.config workgroups(%key_value_block_count, %key_value_head_count, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %completion_counter: buffer, %output: buffer, %next_q8_output: buffer) { + %key_value_token_capacity = config.get @ggml.flash_attention.decode.key_value_token_capacity : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %key_value_token_count_in_range = index.assume %key_value_token_count [range(%key_value_token_count, 1, 262144)] : index + %bounded_key_value_token_count, %launch_key_value_token_capacity = index.assume %key_value_token_count_in_range, %key_value_token_capacity [le(%key_value_token_count_in_range, %key_value_token_capacity)] : index, index + %padded_key_value_token_capacity = index.add %launch_key_value_token_capacity, %c63 : index + %producer_block_count = index.div %padded_key_value_token_capacity, %c64 : index + %publish_q8 = scalar.constant true : i1 + template.apply<@ggml.flash_attention.decode_split.produce_partials>(%bounded_key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output) : (index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + template.apply<@ggml.flash_attention.decode_split.reduce_fused>(%launch_key_value_token_capacity, %producer_block_count, %producer_block_count, %publish_q8, %partial_max, %partial_sum, %partial_output, %completion_counter, %output, %next_q8_output) : (index, index, index, i1, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +// Long-context producer export: publish partials for a following parallel +// reducer without carrying short-context synchronization bindings. +kernel.def target(@ggml_flash_attention_decode_split_gfx11_wave64) @ggml_flash_attention_decode_split_produce_partials_f32_f16_wmma(%key_value_token_count: index) { + %key_value_head_count = config.get @ggml.flash_attention.key_value_head_count : index + %key_value_token_capacity = config.get @ggml.flash_attention.decode.key_value_token_capacity : index + %c1 = index.constant 1 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %padded_key_value_token_capacity = index.add %key_value_token_capacity, %c63 : index + %key_value_block_count = index.div %padded_key_value_token_capacity, %c64 : index + kernel.launch.config workgroups(%key_value_block_count, %key_value_head_count, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer) { + %key_value_token_capacity = config.get @ggml.flash_attention.decode.key_value_token_capacity : index + %key_value_token_count_in_range = index.assume %key_value_token_count [range(%key_value_token_count, 1, 262144)] : index + %bounded_key_value_token_count = index.assume %key_value_token_count_in_range [le(%key_value_token_count_in_range, %key_value_token_capacity)] : index + template.apply<@ggml.flash_attention.decode_split.produce_partials>(%bounded_key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output) : (index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +// Long contexts have enough split-K partials that assigning the reduction to +// one last-arriving producer serializes useful work. This reducer launches one +// two-wave workgroup per query head after an execution barrier from the +// producer. The first wave computes normalization once and both waves consume +// it, avoiding the duplicate work of independent 64-channel output slices. +kernel.def target(@ggml_flash_attention_decode_split_gfx11_wave64) @ggml_flash_attention_decode_split_reduce_f32(%key_value_token_count: index) { + %query_head_count = config.get @ggml.flash_attention.query_head_count : index + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%c1, %query_head_count, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%key_value_token_count: index, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %output: buffer) { + %bounded_key_value_token_count = index.assume %key_value_token_count [range(%key_value_token_count, 1, 262144)] : index + %query_head_count = config.get @ggml.flash_attention.query_head_count : index + %key_value_head_count = config.get @ggml.flash_attention.key_value_head_count : index + %value_head_size0 = config.get @ggml.flash_attention.value_head_size : index + %value_head_size = index.assume %value_head_size0 [range(%value_head_size0, 64, 512), mul(%value_head_size0, 64)] : index + %query_head0 = kernel.workgroup.id : index + %query_head = index.assume %query_head0 [range(%query_head0, 0, 63)] : index + %workitem = kernel.workitem.id : index + %subgroup = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c64 = index.constant 64 : index + %c0_offset = index.constant 0 : offset + %reduction_stage_bytes = index.constant 8 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %key_value_token_capacity = config.get @ggml.flash_attention.decode.key_value_token_capacity : index + %c63 = index.constant 63 : index + %padded_key_value_token_capacity = index.add %key_value_token_capacity, %c63 : index + %key_value_block_count0 = index.div %padded_key_value_token_capacity, %c64 : index + %key_value_block_count = index.assume %key_value_block_count0 [range(%key_value_block_count0, 1, 4096)] : index + %query_heads_per_key_value_head0 = index.div %query_head_count, %key_value_head_count : index + %query_heads_per_key_value_head = index.assume %query_heads_per_key_value_head0 [range(%query_heads_per_key_value_head0, 1, 16)] : index + %key_value_head = index.div %query_head, %query_heads_per_key_value_head : index + %query_row = index.rem %query_head, %query_heads_per_key_value_head : index + %is_first_subgroup = index.cmp eq, %subgroup, %c0 : index + %workitem_is_zero = index.cmp eq, %workitem, %c0 : index + %output_channel = index.add %workitem, %c0 : index + %output_channel_valid = index.cmp ult, %output_channel, %value_head_size : index + %partial_max_noalias, %partial_sum_noalias, %partial_output_noalias, %output_noalias = buffer.assume.noalias %partial_max, %partial_sum, %partial_output, %output : buffer, buffer, buffer, buffer + %partial_max_aligned = buffer.assume.alignment %partial_max_noalias {minimum_alignment = 16} : buffer + %partial_sum_aligned = buffer.assume.alignment %partial_sum_noalias {minimum_alignment = 16} : buffer + %partial_output_aligned = buffer.assume.alignment %partial_output_noalias {minimum_alignment = 16} : buffer + %output_aligned = buffer.assume.alignment %output_noalias {minimum_alignment = 16} : buffer + %partial_max_view = buffer.view %partial_max_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + %partial_sum_view = buffer.view %partial_sum_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + %partial_output_view = buffer.view %partial_output_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%key_value_block_count]x16x[%value_head_size]xf16> + %output_view = buffer.view %output_aligned[%c0_offset] : buffer -> view<[%query_head_count]x[%value_head_size]xf32> + %reduction_stage = buffer.alloca align(8) %reduction_stage_bytes : buffer + %reduction_stage_view = buffer.view %reduction_stage[%c0_offset] : buffer -> view<2xf32> + %lane_maximum = scf.if %is_first_subgroup -> (f32) { + %maximum = scf.for %block = [%lane to %key_value_block_count step %c64](%running_maximum = %negative_large : f32) -> (f32) { + %block_maximum = view.load %partial_max_view[%key_value_head, %block, %query_row] : view<[%key_value_head_count]x[%key_value_block_count]x16xf32> -> f32 + %next_maximum = scalar.maxnumf %running_maximum, %block_maximum : f32 + scf.yield %next_maximum : f32 + } + scf.yield %maximum : f32 + } else { + scf.yield %negative_large : f32 + } + %subgroup_maximum = kernel.subgroup.reduce %lane_maximum : f32 + scf.if %workitem_is_zero { + view.store %subgroup_maximum, %reduction_stage_view[%c0] : f32, view<2xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %maximum = view.load %reduction_stage_view[%c0] : view<2xf32> -> f32 + %lane_sum = scf.if %is_first_subgroup -> (f32) { + %sum = scf.for %block = [%lane to %key_value_block_count step %c64](%running_sum = %c0_f32 : f32) -> (f32) { + %block_maximum = view.load %partial_max_view[%key_value_head, %block, %query_row] : view<[%key_value_head_count]x[%key_value_block_count]x16xf32> -> f32 + %delta = scalar.subf %block_maximum, %maximum : f32 + %scale = scalar.expf %delta : f32 + view.store %scale, %partial_max_view[%key_value_head, %block, %query_row] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + %block_sum = view.load %partial_sum_view[%key_value_head, %block, %query_row] : view<[%key_value_head_count]x[%key_value_block_count]x16xf32> -> f32 + %scaled_sum = scalar.mulf %block_sum, %scale : f32 + %next_sum = scalar.addf %running_sum, %scaled_sum : f32 + scf.yield %next_sum : f32 + } + scf.yield %sum : f32 + } else { + scf.yield %c0_f32 : f32 + } + %subgroup_sum = kernel.subgroup.reduce %lane_sum : f32 + scf.if %workitem_is_zero { + view.store %subgroup_sum, %reduction_stage_view[%c1] : f32, view<2xf32> + } + // Wave 0 rewrote partial_max in global memory; every wave reads it below. + kernel.barrier scope(workgroup) ordering(acq_rel) + kernel.barrier scope(workgroup) ordering(acq_rel) + %sum = view.load %reduction_stage_view[%c1] : view<2xf32> -> f32 + scf.if %output_channel_valid { + %unnormalized_output = scf.for %block = [%c0 to %key_value_block_count step %c1](%running_output = %c0_f32 : f32) -> (f32) unroll(%c4) schedule(interleaved) { + %scale = view.load %partial_max_view[%key_value_head, %block, %query_row] : view<[%key_value_head_count]x[%key_value_block_count]x16xf32> -> f32 + %block_output_f16 = view.load %partial_output_view[%key_value_head, %block, %query_row, %output_channel] : view<[%key_value_head_count]x[%key_value_block_count]x16x[%value_head_size]xf16> -> f16 + %block_output = scalar.extf %block_output_f16 : f16 to f32 + %scaled_output = scalar.mulf %block_output, %scale : f32 + %next_output = scalar.addf %running_output, %scaled_output : f32 + scf.yield %next_output : f32 + } + %normalized_output = scalar.divf %unnormalized_output, %sum : f32 + view.store %normalized_output, %output_view[%query_head, %output_channel] : f32, view<[%query_head_count]x[%value_head_size]xf32> + } + kernel.return +} + +check.case public @ggml_flash_attention_decode_split_f32_f16_wmma_case { + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(1.0) : tensor<32x128xf32> + %key = check.generate.fill value(1.0) : tensor<256x4x128xf16> + %value = check.generate.fill value(2.0) : tensor<256x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<256xf16> + %partial_max = check.generate.fill value(-1.0) : tensor<4x4x16xf32> + %partial_sum = check.generate.fill value(-1.0) : tensor<4x4x16xf32> + %partial_output = check.generate.fill value(-1.0) : tensor<4x4x16x128xf16> + %completion_counter = check.generate.fill value(0) : tensor<4xi32> + %output = check.generate.fill value(-1.0) : tensor<32x128xf32> + %expected = check.generate.fill value(2.0) : tensor<32x128xf32> + kernel.launch @ggml_flash_attention_decode_split_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output, %completion_counter, %output) : [index](index, tensor<32x128xf32>, tensor<256x4x128xf16>, tensor<256x4x128xf16>, tensor<256xf16>, tensor<4x4x16xf32>, tensor<4x4x16xf32>, tensor<4x4x16x128xf16>, tensor<4xi32>, tensor<32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<32x128xf32> + check.return +} + +// Only KV row zero participates. The finite F32 iota step overflows to F16 +// negative infinity at every later row, leaving blocks one through three +// entirely masked. Each empty split must publish the online-softmax identity +// instead of evaluating -inf - -inf and contaminating the final reduction. +check.case public @ggml_flash_attention_decode_split_f32_f16_wmma_masked_blocks_case { + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(1.0) : tensor<32x128xf32> + %key = check.generate.fill value(1.0) : tensor<256x4x128xf16> + %value = check.generate.fill value(2.0) : tensor<256x4x128xf16> + %mask = check.generate.iota offset(0.0) step(-1e+30) : tensor<256xf16> + %partial_max = check.generate.fill value(-1.0) : tensor<4x4x16xf32> + %partial_sum = check.generate.fill value(-1.0) : tensor<4x4x16xf32> + %partial_output = check.generate.fill value(-1.0) : tensor<4x4x16x128xf16> + %completion_counter = check.generate.fill value(0) : tensor<4xi32> + %output = check.generate.fill value(-1.0) : tensor<32x128xf32> + %expected = check.generate.fill value(2.0) : tensor<32x128xf32> + kernel.launch @ggml_flash_attention_decode_split_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output, %completion_counter, %output) : [index](index, tensor<32x128xf32>, tensor<256x4x128xf16>, tensor<256x4x128xf16>, tensor<256xf16>, tensor<4x4x16xf32>, tensor<4x4x16xf32>, tensor<4x4x16x128xf16>, tensor<4xi32>, tensor<32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<32x128xf32> + check.return +} + +// The production Decode-513 shape combines eight full blocks with a one-row +// tail and selects the cooperative fused reducer. Constant nonzero V keeps the +// expected result exact while all four KV heads, 32 GQA heads, completion +// counters, partial tensors, and tail guards participate. A second invocation +// reuses the partial and counter storage with a different V tensor, making the +// completion-counter reset observable. +check.case public @ggml_flash_attention_decode_split_f32_f16_wmma_decode_513_case { + %key_value_token_count = check.literal value(513) : index + %query = check.generate.fill value(1.0) : tensor<32x128xf32> + %key = check.generate.fill value(1.0) : tensor<513x4x128xf16> + %value0 = check.generate.fill value(2.0) : tensor<513x4x128xf16> + %value1 = check.generate.fill value(3.0) : tensor<513x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<513xf16> + %partial_max = check.generate.fill value(-1.0) : tensor<4x9x16xf32> + %partial_sum = check.generate.fill value(-1.0) : tensor<4x9x16xf32> + %partial_output = check.generate.fill value(-1.0) : tensor<4x9x16x128xf16> + %completion_counter = check.generate.fill value(0) : tensor<4xi32> + %output0 = check.generate.fill value(-1.0) : tensor<32x128xf32> + %output1 = check.generate.fill value(-1.0) : tensor<32x128xf32> + %expected0 = check.generate.fill value(2.0) : tensor<32x128xf32> + %expected1 = check.generate.fill value(3.0) : tensor<32x128xf32> + kernel.launch @ggml_flash_attention_decode_split_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value0, %mask, %partial_max, %partial_sum, %partial_output, %completion_counter, %output0) : [index](index, tensor<32x128xf32>, tensor<513x4x128xf16>, tensor<513x4x128xf16>, tensor<513xf16>, tensor<4x9x16xf32>, tensor<4x9x16xf32>, tensor<4x9x16x128xf16>, tensor<4xi32>, tensor<32x128xf32>) + kernel.launch @ggml_flash_attention_decode_split_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value1, %mask, %partial_max, %partial_sum, %partial_output, %completion_counter, %output1) : [index](index, tensor<32x128xf32>, tensor<513x4x128xf16>, tensor<513x4x128xf16>, tensor<513xf16>, tensor<4x9x16xf32>, tensor<4x9x16xf32>, tensor<4x9x16x128xf16>, tensor<4xi32>, tensor<32x128xf32>) + check.expect.close actual(%output0) expected(%expected0) atol(0.001) rtol(0.001) nan(same) : tensor<32x128xf32> + check.expect.close actual(%output1) expected(%expected1) atol(0.001) rtol(0.001) nan(same) : tensor<32x128xf32> + check.return +} + +// Sixty-five KV rows force a second split containing one valid row. The mask +// selects that final row, whose iota values begin at 1024, so an omitted or +// uninitialized tail cannot accidentally satisfy the check. +check.case public @ggml_flash_attention_decode_split_f32_f16_wmma_tail_case { + %key_value_token_count = check.literal value(65) : index + %query = check.generate.fill value(1.0) : tensor<1x128xf32> + %key = check.generate.fill value(1.0) : tensor<65x1x128xf16> + %value = check.generate.iota offset(0.0) step(0.125) : tensor<65x1x128xf16> + %mask = check.generate.iota offset(-64000.0) step(1000.0) : tensor<65xf16> + %partial_max = check.generate.fill value(-1.0) : tensor<1x2x16xf32> + %partial_sum = check.generate.fill value(-1.0) : tensor<1x2x16xf32> + %partial_output = check.generate.fill value(-1.0) : tensor<1x2x16x128xf16> + %completion_counter = check.generate.fill value(0) : tensor<1xi32> + %output = check.generate.fill value(-1.0) : tensor<1x128xf32> + %expected = check.generate.iota offset(1024.0) step(0.125) : tensor<1x128xf32> + kernel.launch @ggml_flash_attention_decode_split_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output, %completion_counter, %output) : [index](index, tensor<1x128xf32>, tensor<65x1x128xf16>, tensor<65x1x128xf16>, tensor<65xf16>, tensor<1x2x16xf32>, tensor<1x2x16xf32>, tensor<1x2x16x128xf16>, tensor<1xi32>, tensor<1x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<1x128xf32> + check.return +} + +// Thirty-two split-K blocks exercise the separate parallel reducer and the +// execution barrier between its producer and consumer dispatches. +check.case public @ggml_flash_attention_decode_split_f32_f16_wmma_long_case { + %key_value_token_count = check.literal value(2048) : index + %query = check.generate.fill value(0.0) : tensor<32x128xf32> + %key = check.generate.fill value(0.0) : tensor<2048x4x128xf16> + %value = check.generate.fill value(0.0) : tensor<2048x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<2048xf16> + %partial_max = check.generate.fill value(-1.0) : tensor<4x32x16xf32> + %partial_sum = check.generate.fill value(-1.0) : tensor<4x32x16xf32> + %partial_output = check.generate.fill value(-1.0) : tensor<4x32x16x128xf16> + %output = check.generate.fill value(1.0) : tensor<32x128xf32> + %expected = check.generate.fill value(0.0) : tensor<32x128xf32> + kernel.launch @ggml_flash_attention_decode_split_produce_partials_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output) : [index](index, tensor<32x128xf32>, tensor<2048x4x128xf16>, tensor<2048x4x128xf16>, tensor<2048xf16>, tensor<4x32x16xf32>, tensor<4x32x16xf32>, tensor<4x32x16x128xf16>) + kernel.launch @ggml_flash_attention_decode_split_reduce_f32[%key_value_token_count](%key_value_token_count, %partial_max, %partial_sum, %partial_output, %output) : [index](index, tensor<4x32x16xf32>, tensor<4x32x16xf32>, tensor<4x32x16x128xf16>, tensor<32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<32x128xf32> + check.return +} + +check.case public @ggml_flash_attention_decode_split_f32_f16_wmma_benchmark_case { + %key_value_token_count = check.param.choice values([64, 65, 128, 256, 512, 513, 768, 1024, 1280, 2048]) name("key_value_token_count") : index + %query = check.generate.fill value(0.0) : tensor<32x128xf32> + %key = check.generate.fill value(0.0) : tensor<[%key_value_token_count]x4x128xf16> + %value = check.generate.fill value(0.0) : tensor<[%key_value_token_count]x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<[%key_value_token_count]xf16> + // Reserve the bounded scratch capacity once; each specialization addresses + // only ceildiv(key_value_token_count, 64) blocks. + %partial_max = check.generate.fill value(-1.0) : tensor<4x512x16xf32> + %partial_sum = check.generate.fill value(-1.0) : tensor<4x512x16xf32> + %partial_output = check.generate.fill value(-1.0) : tensor<4x512x16x128xf16> + %completion_counter = check.generate.fill value(0) : tensor<4xi32> + %output = check.generate.fill value(1.0) : tensor<32x128xf32> + %expected = check.generate.fill value(0.0) : tensor<32x128xf32> + kernel.launch @ggml_flash_attention_decode_split_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output, %completion_counter, %output) : [index](index, tensor<32x128xf32>, tensor<[%key_value_token_count]x4x128xf16>, tensor<[%key_value_token_count]x4x128xf16>, tensor<[%key_value_token_count]xf16>, tensor<4x512x16xf32>, tensor<4x512x16xf32>, tensor<4x512x16x128xf16>, tensor<4xi32>, tensor<32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<32x128xf32> + check.return +} + +// Long-context execution is one reusable producer/reducer command buffer. The +// harness records an explicit dispatch execution barrier between these calls +// and profiles their complete device-side span as one semantic operation. +check.case public @ggml_flash_attention_decode_split_f32_f16_wmma_long_benchmark_case { + %key_value_token_count = check.param.choice values([2048, 32768]) name("key_value_token_count") : index + %query = check.generate.fill value(0.0) : tensor<32x128xf32> + %key = check.generate.fill value(0.0) : tensor<[%key_value_token_count]x4x128xf16> + %value = check.generate.fill value(0.0) : tensor<[%key_value_token_count]x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<[%key_value_token_count]xf16> + %partial_max = check.generate.fill value(-1.0) : tensor<4x512x16xf32> + %partial_sum = check.generate.fill value(-1.0) : tensor<4x512x16xf32> + %partial_output = check.generate.fill value(-1.0) : tensor<4x512x16x128xf16> + %output = check.generate.fill value(1.0) : tensor<32x128xf32> + %expected = check.generate.fill value(0.0) : tensor<32x128xf32> + kernel.launch @ggml_flash_attention_decode_split_produce_partials_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output) : [index](index, tensor<32x128xf32>, tensor<[%key_value_token_count]x4x128xf16>, tensor<[%key_value_token_count]x4x128xf16>, tensor<[%key_value_token_count]xf16>, tensor<4x512x16xf32>, tensor<4x512x16xf32>, tensor<4x512x16x128xf16>) + kernel.launch @ggml_flash_attention_decode_split_reduce_f32[%key_value_token_count](%key_value_token_count, %partial_max, %partial_sum, %partial_output, %output) : [index](index, tensor<4x512x16xf32>, tensor<4x512x16xf32>, tensor<4x512x16x128xf16>, tensor<32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<32x128xf32> + check.return +} + +check.benchmark<@ggml_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @ggml_flash_attention_decode_split_f32_f16_wmma_decode_64 {key_value_token_count = 64} + +check.benchmark<@ggml_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @ggml_flash_attention_decode_split_f32_f16_wmma_decode_65 {key_value_token_count = 65} + +check.benchmark<@ggml_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @ggml_flash_attention_decode_split_f32_f16_wmma_decode_128 {key_value_token_count = 128} + +check.benchmark<@ggml_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @ggml_flash_attention_decode_split_f32_f16_wmma_decode_256 {key_value_token_count = 256} + +check.benchmark<@ggml_flash_attention_decode_split_f32_f16_wmma_masked_blocks_case> @ggml_flash_attention_decode_split_f32_f16_wmma_decode_256_masked_blocks + +check.benchmark<@ggml_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @ggml_flash_attention_decode_split_f32_f16_wmma_decode_512 {key_value_token_count = 512} + +check.benchmark<@ggml_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @ggml_flash_attention_decode_split_f32_f16_wmma_decode_513 {key_value_token_count = 513} + +check.benchmark<@ggml_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @ggml_flash_attention_decode_split_f32_f16_wmma_decode_768 {key_value_token_count = 768} + +check.benchmark<@ggml_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @ggml_flash_attention_decode_split_f32_f16_wmma_decode_1024 {key_value_token_count = 1024} + +check.benchmark<@ggml_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @ggml_flash_attention_decode_split_f32_f16_wmma_decode_1280 {key_value_token_count = 1280} + +check.benchmark<@ggml_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @ggml_flash_attention_decode_split_f32_f16_wmma_decode_2048_fused {key_value_token_count = 2048} + +check.benchmark<@ggml_flash_attention_decode_split_f32_f16_wmma_long_benchmark_case> @ggml_flash_attention_decode_split_f32_f16_wmma_decode_2048 {key_value_token_count = 2048} + +check.benchmark<@ggml_flash_attention_decode_split_f32_f16_wmma_long_benchmark_case> @ggml_flash_attention_decode_split_f32_f16_wmma_decode_32768 {key_value_token_count = 32768} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/flash_attention_f32_f16_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/flash_attention_f32_f16_wmma.loom new file mode 100644 index 000000000000..d225ccac32f6 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/flash_attention_f32_f16_wmma.loom @@ -0,0 +1,1125 @@ +// ggml grouped-query FlashAttention for the F32-query/F16-cache path. +// +// One four-wave workgroup computes 16 query rows for one query head and one +// 256-channel output group against 64 KV rows at a time. The ownership changes +// mirror the cooperative-matrix schedule used by llama.cpp's Vulkan CM1 kernel: +// +// 1. All workitems stage a scaled 16xQK-head-size F16 query tile. +// 2. Each wave computes one 16x16 QK score slice. +// 3. Scores cross LDS so each wave can normalize four complete query rows. +// 4. F16 probabilities cross LDS for four P*V WMMA steps. +// 5. Each active lane retains one four-channel F16 packet for each of its +// four query rows across subsequent 64-row KV blocks. +// +// K and V remain in the row-major llama.cpp cache layout +// [KV token][KV head][head_size]. K uses the QK head size; V uses the output +// value head size. Their aligned F16 fragments load directly from global +// memory; there is no expanded or repacked persistent allocation. +// QK and the online-softmax statistics remain F32, while the P*V accumulation +// and carried output match the Vulkan oracle's F16 policy. +amdgpu.target @ggml_flash_attention_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.flash_attention.query_head_count : %value: index where [range(%value, 1, 64)] + +config.decl @ggml.flash_attention.key_value_head_count : %value: index where [range(%value, 1, 64)] + +config.decl @ggml.flash_attention.qk_head_size : %value: index where [range(%value, 16, 36864), mul(%value, 16)] + +config.decl @ggml.flash_attention.value_head_size : %value: index where [range(%value, 64, 512), mul(%value, 64)] +config.decl @ggml.flash_attention.value_stride : %value: index where [range(%value, 16, 36864), mul(%value, 16)] + +config.decl @ggml.flash_attention.attention_scale : f32 + +config.def @ggml.flash_attention.value_layout = 0 : index + +config.decl @ggml.flash_attention.apply_gate : %value: index where [range(%value, 0, 1)] + +config.decl @ggml.flash_attention.gate_stride_head : %value: index where [range(%value, 1, 16777216)] + +config.decl @ggml.flash_attention.gate_stride_token : %value: index where [range(%value, 1, 16777216)] + +// Bounds the context copied by the test-only row-extraction kernel below. +config.decl @ggml.flash_attention.test.context_capacity : %value: index where [range(%value, 1, 32768)] + +template.decl @ggml.flash_attention.prefill.stage_query(%query_token_count: index, %query_origin: index, %query_head: index, %query_load_iteration_count: index, %query_stage_stride: index, %qk_head_size: index, %attention_scale: f32, %query: buffer, %query_stage: buffer) + +template.decl @ggml.flash_attention.prefill.publish_output(%query_token_count: index, %query_head_count: index, %value_head_size: index, %lane_has_output: i1, %query_valid0: i1, %query_valid1: i1, %query_valid2: i1, %query_valid3: i1, %query_token0: index, %query_token1: index, %query_token2: index, %query_token3: index, %query_head: index, %lane_output_channel: index, %final_sum: vector<4xf32>, %final_output0: vector<4xf16>, %final_output1: vector<4xf16>, %final_output2: vector<4xf16>, %final_output3: vector<4xf16>, %gate: buffer, %output: buffer) + +template.decl @ggml.flash_attention.prefill.body(%query_token_count: index, %key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %gate: buffer, %output: buffer) + +template.def<@ggml.flash_attention.prefill.stage_query> device @ggml_flash_attention_prefill_stage_query(%query_token_count: index, %query_origin: index, %query_head: index, %query_load_iteration_count: index, %query_stage_stride: index, %qk_head_size: index, %attention_scale: f32, %query: buffer, %query_stage: buffer) { + %query_head_count = config.get @ggml.flash_attention.query_head_count : index + %workitem = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %c0_f16 = scalar.constant 0.0 : f16 + %query_view = buffer.view %query[%c0_offset] : buffer -> view<[%query_token_count]x[%query_head_count]x[%qk_head_size]xf32> + %query_stage_view = buffer.view %query_stage[%c0_offset] : buffer -> view<16x[%query_stage_stride]xf16> + // Scale and truncate Q exactly once. The Vulkan reference does this before + // entering its KV loop, making QK a native F16 WMMA while retaining F32 + // accumulation. + scf.for %load_iteration = [%c0 to %query_load_iteration_count step %c1] unroll { + %linear = index.madd %load_iteration, %c256, %workitem : index + %local_query_row = index.div %linear, %qk_head_size : index + %query_channel = index.rem %linear, %qk_head_size : index + %query_token = index.add %query_origin, %local_query_row : index + %query_valid_load = index.cmp ult, %query_token, %query_token_count : index + %query_value = scf.if %query_valid_load -> (f16) { + %loaded = view.load %query_view[%query_token, %query_head, %query_channel] : view<[%query_token_count]x[%query_head_count]x[%qk_head_size]xf32> -> f32 + %scaled = scalar.mulf %loaded, %attention_scale : f32 + %truncated = scalar.fptrunc %scaled : f32 to f16 + scf.yield %truncated : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %query_value, %query_stage_view[%local_query_row, %query_channel] : f16, view<16x[%query_stage_stride]xf16> + } + template.return +} + +template.def<@ggml.flash_attention.prefill.publish_output> device @ggml_flash_attention_prefill_publish_output(%query_token_count: index, %query_head_count: index, %value_head_size: index, %lane_has_output: i1, %query_valid0: i1, %query_valid1: i1, %query_valid2: i1, %query_valid3: i1, %query_token0: index, %query_token1: index, %query_token2: index, %query_token3: index, %query_head: index, %lane_output_channel: index, %final_sum: vector<4xf32>, %final_output0: vector<4xf16>, %final_output1: vector<4xf16>, %final_output2: vector<4xf16>, %final_output3: vector<4xf16>, %gate: buffer, %output: buffer) { + %c0_offset = index.constant 0 : offset + %c1 = index.constant 1 : index + %c1_f32 = scalar.constant 1.0 : f32 + %c1_f32x4 = vector.constant 1.0 : vector<4xf32> + %apply_gate_index = config.get @ggml.flash_attention.apply_gate : index + %gate_stride_head = config.get @ggml.flash_attention.gate_stride_head : index + %gate_stride_token = config.get @ggml.flash_attention.gate_stride_token : index + %apply_gate = index.cmp eq, %apply_gate_index, %c1 : index + %gate_aligned = buffer.assume.alignment %gate {minimum_alignment = 16} : buffer + %gate_view = buffer.view %gate_aligned[%c0_offset] : buffer -> view<268435456xf32> + %output_view = buffer.view %output[%c0_offset] : buffer -> view<[%query_token_count]x[%query_head_count]x[%value_head_size]xf32> + // Normalize and publish the lane-owned packets. Query-tail rows never + // participate in the output store. + scf.if %lane_has_output { + %sum0_scalar = vector.extract %final_sum[0] : vector<4xf32> -> f32 + %sum1_scalar = vector.extract %final_sum[1] : vector<4xf32> -> f32 + %sum2_scalar = vector.extract %final_sum[2] : vector<4xf32> -> f32 + %sum3_scalar = vector.extract %final_sum[3] : vector<4xf32> -> f32 + %inverse_sum0_f32 = scalar.divf %c1_f32, %sum0_scalar : f32 + %inverse_sum1_f32 = scalar.divf %c1_f32, %sum1_scalar : f32 + %inverse_sum2_f32 = scalar.divf %c1_f32, %sum2_scalar : f32 + %inverse_sum3_f32 = scalar.divf %c1_f32, %sum3_scalar : f32 + %inverse_sum0_f16 = scalar.fptrunc %inverse_sum0_f32 : f32 to f16 + %inverse_sum1_f16 = scalar.fptrunc %inverse_sum1_f32 : f32 to f16 + %inverse_sum2_f16 = scalar.fptrunc %inverse_sum2_f32 : f32 to f16 + %inverse_sum3_f16 = scalar.fptrunc %inverse_sum3_f32 : f32 to f16 + %inverse_sum0 = vector.splat %inverse_sum0_f16 : vector<4xf16> + %inverse_sum1 = vector.splat %inverse_sum1_f16 : vector<4xf16> + %inverse_sum2 = vector.splat %inverse_sum2_f16 : vector<4xf16> + %inverse_sum3 = vector.splat %inverse_sum3_f16 : vector<4xf16> + %normalized0_f16 = vector.mulf %final_output0, %inverse_sum0 : vector<4xf16> + %normalized1_f16 = vector.mulf %final_output1, %inverse_sum1 : vector<4xf16> + %normalized2_f16 = vector.mulf %final_output2, %inverse_sum2 : vector<4xf16> + %normalized3_f16 = vector.mulf %final_output3, %inverse_sum3 : vector<4xf16> + %normalized0 = vector.extf %normalized0_f16 : vector<4xf16> to vector<4xf32> + %normalized1 = vector.extf %normalized1_f16 : vector<4xf16> to vector<4xf32> + %normalized2 = vector.extf %normalized2_f16 : vector<4xf16> to vector<4xf32> + %normalized3 = vector.extf %normalized3_f16 : vector<4xf16> to vector<4xf32> + scf.if %query_valid0 { + %published0 = scf.if %apply_gate -> (vector<4xf32>) { + %gate_token_offset0 = index.mul %query_token0, %gate_stride_token : index + %gate_head_offset0 = index.mul %query_head, %gate_stride_head : index + %gate_head_base0 = index.add %gate_token_offset0, %gate_head_offset0 : index + %gate_index0_0 = index.add %gate_head_base0, %lane_output_channel : index + %gate_index0 = index.assume %gate_index0_0 [range(%gate_index0_0, 0, 268435452)] : index + %raw_gate0 = vector.load %gate_view[%gate_index0] : view<268435456xf32> -> vector<4xf32> + %negative_gate0 = vector.negf %raw_gate0 : vector<4xf32> + %gate_exp0 = vector.expf %negative_gate0 : vector<4xf32> + %gate_denominator0 = vector.addf %c1_f32x4, %gate_exp0 : vector<4xf32> + %sigmoid_gate0 = vector.divf %c1_f32x4, %gate_denominator0 : vector<4xf32> + %gated_output0 = vector.mulf %normalized0, %sigmoid_gate0 : vector<4xf32> + scf.yield %gated_output0 : vector<4xf32> + } else { + scf.yield %normalized0 : vector<4xf32> + } + vector.store %published0, %output_view[%query_token0, %query_head, %lane_output_channel] : vector<4xf32>, view<[%query_token_count]x[%query_head_count]x[%value_head_size]xf32> + } + scf.if %query_valid1 { + %published1 = scf.if %apply_gate -> (vector<4xf32>) { + %gate_token_offset1 = index.mul %query_token1, %gate_stride_token : index + %gate_head_offset1 = index.mul %query_head, %gate_stride_head : index + %gate_head_base1 = index.add %gate_token_offset1, %gate_head_offset1 : index + %gate_index1_0 = index.add %gate_head_base1, %lane_output_channel : index + %gate_index1 = index.assume %gate_index1_0 [range(%gate_index1_0, 0, 268435452)] : index + %raw_gate1 = vector.load %gate_view[%gate_index1] : view<268435456xf32> -> vector<4xf32> + %negative_gate1 = vector.negf %raw_gate1 : vector<4xf32> + %gate_exp1 = vector.expf %negative_gate1 : vector<4xf32> + %gate_denominator1 = vector.addf %c1_f32x4, %gate_exp1 : vector<4xf32> + %sigmoid_gate1 = vector.divf %c1_f32x4, %gate_denominator1 : vector<4xf32> + %gated_output1 = vector.mulf %normalized1, %sigmoid_gate1 : vector<4xf32> + scf.yield %gated_output1 : vector<4xf32> + } else { + scf.yield %normalized1 : vector<4xf32> + } + vector.store %published1, %output_view[%query_token1, %query_head, %lane_output_channel] : vector<4xf32>, view<[%query_token_count]x[%query_head_count]x[%value_head_size]xf32> + } + scf.if %query_valid2 { + %published2 = scf.if %apply_gate -> (vector<4xf32>) { + %gate_token_offset2 = index.mul %query_token2, %gate_stride_token : index + %gate_head_offset2 = index.mul %query_head, %gate_stride_head : index + %gate_head_base2 = index.add %gate_token_offset2, %gate_head_offset2 : index + %gate_index2_0 = index.add %gate_head_base2, %lane_output_channel : index + %gate_index2 = index.assume %gate_index2_0 [range(%gate_index2_0, 0, 268435452)] : index + %raw_gate2 = vector.load %gate_view[%gate_index2] : view<268435456xf32> -> vector<4xf32> + %negative_gate2 = vector.negf %raw_gate2 : vector<4xf32> + %gate_exp2 = vector.expf %negative_gate2 : vector<4xf32> + %gate_denominator2 = vector.addf %c1_f32x4, %gate_exp2 : vector<4xf32> + %sigmoid_gate2 = vector.divf %c1_f32x4, %gate_denominator2 : vector<4xf32> + %gated_output2 = vector.mulf %normalized2, %sigmoid_gate2 : vector<4xf32> + scf.yield %gated_output2 : vector<4xf32> + } else { + scf.yield %normalized2 : vector<4xf32> + } + vector.store %published2, %output_view[%query_token2, %query_head, %lane_output_channel] : vector<4xf32>, view<[%query_token_count]x[%query_head_count]x[%value_head_size]xf32> + } + scf.if %query_valid3 { + %published3 = scf.if %apply_gate -> (vector<4xf32>) { + %gate_token_offset3 = index.mul %query_token3, %gate_stride_token : index + %gate_head_offset3 = index.mul %query_head, %gate_stride_head : index + %gate_head_base3 = index.add %gate_token_offset3, %gate_head_offset3 : index + %gate_index3_0 = index.add %gate_head_base3, %lane_output_channel : index + %gate_index3 = index.assume %gate_index3_0 [range(%gate_index3_0, 0, 268435452)] : index + %raw_gate3 = vector.load %gate_view[%gate_index3] : view<268435456xf32> -> vector<4xf32> + %negative_gate3 = vector.negf %raw_gate3 : vector<4xf32> + %gate_exp3 = vector.expf %negative_gate3 : vector<4xf32> + %gate_denominator3 = vector.addf %c1_f32x4, %gate_exp3 : vector<4xf32> + %sigmoid_gate3 = vector.divf %c1_f32x4, %gate_denominator3 : vector<4xf32> + %gated_output3 = vector.mulf %normalized3, %sigmoid_gate3 : vector<4xf32> + scf.yield %gated_output3 : vector<4xf32> + } else { + scf.yield %normalized3 : vector<4xf32> + } + vector.store %published3, %output_view[%query_token3, %query_head, %lane_output_channel] : vector<4xf32>, view<[%query_token_count]x[%query_head_count]x[%value_head_size]xf32> + } + } + template.return +} + +template.def<@ggml.flash_attention.prefill.body> device @ggml_flash_attention_prefill_body(%query_token_count: index, %key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %gate: buffer, %output: buffer) { + %bounded_key_value_token_count = index.assume %key_value_token_count [range(%key_value_token_count, 1, 262144)] : index + %query_head_count = config.get @ggml.flash_attention.query_head_count : index + %key_value_head_count = config.get @ggml.flash_attention.key_value_head_count : index + %qk_head_size0 = config.get @ggml.flash_attention.qk_head_size : index + %qk_head_size = index.assume %qk_head_size0 [range(%qk_head_size0, 16, 576), mul(%qk_head_size0, 16)] : index + %value_head_size0 = config.get @ggml.flash_attention.value_head_size : index + %value_head_size = index.assume %value_head_size0 [range(%value_head_size0, 64, 512), mul(%value_head_size0, 64)] : index + %query_tile = kernel.workgroup.id : index + %query_head0 = kernel.workgroup.id : index + %output_group0 = kernel.workgroup.id : index + %query_head, %launch_query_head_count = index.assume %query_head0, %query_head_count [lt(%query_head0, %query_head_count)] : index, index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 3)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c15 = index.constant 15 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %f16_bytes = index.constant 2 : offset + %score_stage_bytes = index.constant 5120 : offset + %probability_stage_bytes = index.constant 2304 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %attention_scale = config.get @ggml.flash_attention.attention_scale : f32 + %c0_f16 = scalar.constant 0.0 : f16 + %output_zero0 = vector.constant 0.0 : vector<4xf16> + %output_zero1 = vector.constant 0.0 : vector<4xf16> + %output_zero2 = vector.constant 0.0 : vector<4xf16> + %output_zero3 = vector.constant 0.0 : vector<4xf16> + %c0_f16x8 = vector.constant 0.0 : vector<8xf16> + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %negative_f32x4 = vector.constant -1e+30 : vector<4xf32> + %positive_f32x4 = vector.constant 1e+30 : vector<4xf32> + %positive_large = scalar.constant 1e+30 : f32 + %c63 = index.constant 63 : index + %staged_value_offset = index.constant 2048 : offset + %key_visibility_bytes = index.constant 1024 : offset + %hidden_summary_bytes = index.constant 16 : offset + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %query_heads_per_key_value_head = index.div %query_head_count, %key_value_head_count : index + %key_value_head = index.div %query_head, %query_heads_per_key_value_head : index + %key_width = index.mul %key_value_head_count, %qk_head_size : index + %value_width = index.mul %key_value_head_count, %value_head_size : index + %key_head_base = index.mul %key_value_head, %qk_head_size : index + %value_head_base = index.mul %key_value_head, %value_head_size : index + %output_tile_count = index.div %value_head_size, %c64 : index + %padded_output_tile_count = index.add %output_tile_count, %c3 : index + %output_group_count = index.div %padded_output_tile_count, %c4 : index + %output_group, %launch_output_group_count = index.assume %output_group0, %output_group_count [lt(%output_group0, %output_group_count)] : index, index + %output_group_tile_base = index.mul %output_group, %c4 : index + %remaining_output_tile_count = index.sub %output_tile_count, %output_group_tile_base : index + %group_output_tile_count = index.min %remaining_output_tile_count, %c4 : index + %group_output_channel_count = index.mul %group_output_tile_count, %c64 : index + %query_stage_stride = index.add %qk_head_size, %c8 : index + %query_stage_element_count = index.mul %c16, %query_stage_stride : index + %query_stage_bytes = index.scale %query_stage_element_count, %f16_bytes : index, offset -> offset + %query_element_count = index.mul %c16, %qk_head_size : index + %query_load_iteration_count = index.div %query_element_count, %c256 : index + %tail_stage_head_size_is_value = index.cmp ult, %qk_head_size, %value_head_size : index + %tail_stage_head_size = scf.select %tail_stage_head_size_is_value, %value_head_size, %qk_head_size : index + %tail_key_element_count = index.mul %c32, %qk_head_size : index + %tail_key_load_iteration_count = index.div %tail_key_element_count, %c256 : index + %tail_value_stage_channel_count = index.min %value_head_size, %c256 : index + %tail_value_element_count = index.mul %c32, %tail_value_stage_channel_count : index + %tail_value_load_iteration_count = index.div %tail_value_element_count, %c256 : index + %tail_key_value_stage_element_count = index.mul %c32, %tail_stage_head_size : index + %tail_key_value_stage_capacity = index.scale %tail_key_value_stage_element_count, %f16_bytes : index, offset -> offset + %query_origin = index.mul %query_tile, %c16 : index + %full_key_value_block_count = index.div %bounded_key_value_token_count, %c64 : index + %full_key_value_token_count0 = index.mul %full_key_value_block_count, %c64 : index + %full_key_value_token_count = index.assume %full_key_value_token_count0 [range(%full_key_value_token_count0, 0, 262144), mul(%full_key_value_token_count0, 64)] : index + %tail_key_value_token_count = index.sub %bounded_key_value_token_count, %full_key_value_token_count : index + %has_key_value_tail = index.cmp ne, %tail_key_value_token_count, %c0 : index + %tail_key_value_stage_bytes = scf.select %has_key_value_tail, %tail_key_value_stage_capacity, %c0_offset : offset + %subgroup_score_column = index.mul %subgroup, %c16 : index + %subgroup_query_row = index.mul %subgroup, %c4 : index + %query_row0 = index.add %subgroup_query_row, %c0 : index + %query_row1 = index.add %subgroup_query_row, %c1 : index + %query_row2 = index.add %subgroup_query_row, %c2 : index + %query_row3 = index.add %subgroup_query_row, %c3 : index + %query_token0 = index.add %query_origin, %query_row0 : index + %query_token1 = index.add %query_origin, %query_row1 : index + %query_token2 = index.add %query_origin, %query_row2 : index + %query_token3 = index.add %query_origin, %query_row3 : index + %query_valid0 = index.cmp ult, %query_token0, %query_token_count : index + %query_valid1 = index.cmp ult, %query_token1, %query_token_count : index + %query_valid2 = index.cmp ult, %query_token2, %query_token_count : index + %query_valid3 = index.cmp ult, %query_token3, %query_token_count : index + %query_valid = vector.from_elements %query_valid0, %query_valid1, %query_valid2, %query_valid3 : vector<4xi1> + %subgroup_product_channel = index.mul %subgroup, %c16 : index + %lane_output_tile = index.div %lane, %c16 : index + %lane_product_channel0 = index.rem %lane, %c16 : index + %lane_product_channel = index.mul %lane_product_channel0, %c4 : index + %output_group_channel_base = index.mul %output_group, %c256 : index + %lane_output_channel0 = index.mul %lane, %c4 : index + %lane_output_channel = index.add %output_group_channel_base, %lane_output_channel0 : index + %lane_output_end = index.add %lane_output_channel, %c4 : index + %lane_has_output = index.cmp ule, %lane_output_end, %value_head_size : index + // Q and score rows retain padding for their ownership exchanges. Store + // probabilities query-major with eight spare K columns so native LHS + // fragments use aligned wide reads instead of scalar transposed loads. + %query_transposed_layout = encoding.layout.strided [1, %query_stage_stride] : encoding + %probability_layout = encoding.layout.strided [72, 1] : encoding + %query_noalias, %key_noalias, %value_noalias, %mask_noalias, %output_noalias = buffer.assume.noalias %query, %key, %value, %mask, %output : buffer, buffer, buffer, buffer, buffer + %query_aligned = buffer.assume.alignment %query_noalias {minimum_alignment = 16} : buffer + %key_aligned = buffer.assume.alignment %key_noalias {minimum_alignment = 16} : buffer + %value_aligned = buffer.assume.alignment %value_noalias {minimum_alignment = 16} : buffer + %mask_aligned = buffer.assume.alignment %mask_noalias {minimum_alignment = 16} : buffer + %output_aligned = buffer.assume.alignment %output_noalias {minimum_alignment = 16} : buffer + %key_view = buffer.view %key_aligned[%c0_offset] : buffer -> view<[%bounded_key_value_token_count]x[%key_width]xf16> + %value_layout_tag = config.get @ggml.flash_attention.value_layout : index + %transposed_value = index.cmp eq, %value_layout_tag, %c1 : index + %value_stride0 = config.get @ggml.flash_attention.value_stride : index + %value_stride = index.assume %value_stride0 [range(%value_stride0, 16, 36864), mul(%value_stride0, 16)] : index + %value_row_stride = scf.select %transposed_value, %c1, %value_stride : index + %value_column_stride = scf.select %transposed_value, %bounded_key_value_token_count, %c1 : index + %value_layout = encoding.layout.strided [%value_row_stride, %value_column_stride] : encoding + %value_view = buffer.view %value_aligned[%c0_offset] : buffer -> view<[%bounded_key_value_token_count]x[%value_width]xf16, %value_layout> + %mask_view = buffer.view %mask_aligned[%c0_offset] : buffer -> view<[%query_token_count]x[%bounded_key_value_token_count]xf16> + %query_stage = buffer.alloca align(16) %query_stage_bytes : buffer + %score_stage = buffer.alloca align(16) %score_stage_bytes : buffer + %probability_stage = buffer.alloca align(16) %probability_stage_bytes : buffer + %tail_key_value_stage = buffer.alloca align(16) %tail_key_value_stage_bytes : buffer + // Per key, the largest probability any query row of this tile gives it + // ([key][subgroup]), and the per-subgroup largest negated mask value. + %key_visibility_stage = buffer.alloca align(16) %key_visibility_bytes : buffer + %key_visibility_view = buffer.view %key_visibility_stage[%c0_offset] : buffer -> view<64x4xf32> + %hidden_summary_stage = buffer.alloca align(16) %hidden_summary_bytes : buffer + %hidden_summary_view = buffer.view %hidden_summary_stage[%c0_offset] : buffer -> view<4xf32> + // V rows staged for blocks with masked keys; shares the score stage with + // the 2 KiB product exchange, after score reads have retired. + %staged_value_view = buffer.view %score_stage[%staged_value_offset] : buffer -> view<16x64xf16> + %query_transposed_view = buffer.view %query_stage[%c0_offset] : buffer -> view<[%qk_head_size]x16xf16, %query_transposed_layout> + %score_stage_view = buffer.view %score_stage[%c0_offset] : buffer -> view<64x20xf32> + %mask_summary_view = buffer.view %score_stage[%c0_offset] : buffer -> view<4xf32> + %probability_stage_view = buffer.view %probability_stage[%c0_offset] : buffer -> view<16x64xf16, %probability_layout> + // Score reads retire at probability publication. Product reads retire + // before the next mask/score phase, so the two exchanges share storage. + %product_stage_view = buffer.view %score_stage[%c0_offset] : buffer -> view<16x64xf16> + %tail_key_value_stage_view = buffer.view %tail_key_value_stage[%c0_offset] : buffer -> view<32x[%tail_stage_head_size]xf16> + template.apply<@ggml.flash_attention.prefill.stage_query>(%query_token_count, %query_origin, %query_head, %query_load_iteration_count, %query_stage_stride, %qk_head_size, %attention_scale, %query_aligned, %query_stage) : (index, index, index, index, index, index, f32, buffer, buffer) + kernel.barrier scope(workgroup) ordering(acq_rel) + %full_max, %full_sum, %full_output0, %full_output1, %full_output2, %full_output3 = scf.for %key_origin = [%c0 to %full_key_value_token_count step %c64](%current_max = %negative_f32x4 : vector<4xf32>, %current_sum = %c0_f32x4 : vector<4xf32>, %current_output0 = %output_zero0 : vector<4xf16>, %current_output1 = %output_zero1 : vector<4xf16>, %current_output2 = %output_zero2 : vector<4xf16>, %current_output3 = %output_zero3 : vector<4xf16>) -> (vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + // Cache the complete 16x64 mask tile before QK. Causal masks leave future + // KV blocks entirely at negative infinity; reducing the cached tile lets + // the workgroup skip QK, softmax, and P*V for those blocks. This mirrors + // the Vulkan oracle without requiring its auxiliary compact mask buffer. + %key_token0 = index.add %key_origin, %lane : index + %key_token = index.assume %key_token0 [lt(%key_token0, %bounded_key_value_token_count)] : index + %mask0 = scf.if %query_valid0 -> (f16) { + %value = view.load %mask_view[%query_token0, %key_token] : view<[%query_token_count]x[%bounded_key_value_token_count]xf16> -> f16 + scf.yield %value : f16 + } else { + scf.yield %c0_f16 : f16 + } + %mask1 = scf.if %query_valid1 -> (f16) { + %value = view.load %mask_view[%query_token1, %key_token] : view<[%query_token_count]x[%bounded_key_value_token_count]xf16> -> f16 + scf.yield %value : f16 + } else { + scf.yield %c0_f16 : f16 + } + %mask2 = scf.if %query_valid2 -> (f16) { + %value = view.load %mask_view[%query_token2, %key_token] : view<[%query_token_count]x[%bounded_key_value_token_count]xf16> -> f16 + scf.yield %value : f16 + } else { + scf.yield %c0_f16 : f16 + } + %mask3 = scf.if %query_valid3 -> (f16) { + %value = view.load %mask_view[%query_token3, %key_token] : view<[%query_token_count]x[%bounded_key_value_token_count]xf16> -> f16 + scf.yield %value : f16 + } else { + scf.yield %c0_f16 : f16 + } + %mask_f16 = vector.from_elements %mask0, %mask1, %mask2, %mask3 : vector<4xf16> + %mask_summary_f32 = vector.extf %mask_f16 : vector<4xf16> to vector<4xf32> + %effective_mask = vector.select %query_valid, %mask_summary_f32, %negative_f32x4 : vector<4xf32> + %lane_mask_maximum = vector.reduce %effective_mask, %negative_large : vector<4xf32>, f32 + %subgroup_mask_maximum = kernel.subgroup.reduce %lane_mask_maximum : f32 + %valid_row_mask = vector.select %query_valid, %mask_summary_f32, %positive_f32x4 : vector<4xf32> + %negated_valid_row_mask = vector.negf %valid_row_mask : vector<4xf32> + %lane_hidden_maximum = vector.reduce %negated_valid_row_mask, %negative_large : vector<4xf32>, f32 + %subgroup_hidden_maximum = kernel.subgroup.reduce %lane_hidden_maximum : f32 + %is_subgroup_leader = index.cmp eq, %lane, %c0 : index + scf.if %is_subgroup_leader { + view.store %subgroup_mask_maximum, %mask_summary_view[%subgroup] : f32, view<4xf32> + view.store %subgroup_hidden_maximum, %hidden_summary_view[%subgroup] : f32, view<4xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %subgroup_mask_maxima = vector.load %mask_summary_view[%c0] : view<4xf32> -> vector<4xf32> + %workgroup_mask_maximum = vector.reduce %subgroup_mask_maxima, %negative_large : vector<4xf32>, f32 + %block_has_attention = scalar.cmpf ogt, %workgroup_mask_maximum, %negative_large : f32 + %subgroup_hidden_maxima = vector.load %hidden_summary_view[%c0] : view<4xf32> -> vector<4xf32> + %workgroup_hidden_maximum = vector.reduce %subgroup_hidden_maxima, %negative_large : vector<4xf32>, f32 + %block_fully_visible = scalar.cmpf olt, %workgroup_hidden_maximum, %positive_large : f32 + %next_block_max, %next_block_sum, %next_block_output0, %next_block_output1, %next_block_output2, %next_block_output3 = scf.if %block_has_attention -> (vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + // Four independent wave-level WMMAs produce a 16x64 score tile. + %score_key_origin0 = index.add %key_origin, %subgroup_score_column : index + %last_full_key_tile_start = index.sub %bounded_key_value_token_count, %c15 : index + %score_key_origin = index.assume %score_key_origin0 [lt(%score_key_origin0, %last_full_key_tile_start)] : index + %score_init_values = vector.constant 0.0 : vector<4xf32> + %score_init = vector.fragment %score_init_values shape [%m, %n] : vector<4xf32> + %score_fragment = scf.for %head_tile = [%c0 to %qk_head_size step %c16](%score_accumulator = %score_init : vector<4xf32>) -> (vector<4xf32>) unroll schedule(recurrence) { + %key_channel = index.add %key_head_base, %head_tile : index + %key_fragment = vector.fragment.load %key_view[%score_key_origin, %key_channel] shape [%m, %k] : view<[%bounded_key_value_token_count]x[%key_width]xf16> -> vector<16xf16> + %query_fragment = vector.fragment.load %query_transposed_view[%head_tile, %c0] shape [%k, %n] : view<[%qk_head_size]x16xf16, %query_transposed_layout> -> vector<16xf16> + %next_score_accumulator = vector.mma %key_fragment, %query_fragment, %score_accumulator : vector<16xf16>, vector<16xf16>, vector<4xf32> + scf.yield %next_score_accumulator : vector<4xf32> + } + vector.fragment.store %score_fragment, %score_stage_view[%subgroup_score_column, %c0] shape [%m, %n] : vector<4xf32>, view<64x20xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + // LDS transposes ownership from one 16-column score slice per wave to + // four complete query rows per wave. Every lane contributes one key + // column to each of those rows. + %raw_score0 = view.load %score_stage_view[%lane, %query_row0] : view<64x20xf32> -> f32 + %raw_score1 = view.load %score_stage_view[%lane, %query_row1] : view<64x20xf32> -> f32 + %raw_score2 = view.load %score_stage_view[%lane, %query_row2] : view<64x20xf32> -> f32 + %raw_score3 = view.load %score_stage_view[%lane, %query_row3] : view<64x20xf32> -> f32 + %mask_f32 = vector.extf %mask_f16 : vector<4xf16> to vector<4xf32> + %raw_scores = vector.from_elements %raw_score0, %raw_score1, %raw_score2, %raw_score3 : vector<4xf32> + %masked_scores0 = vector.addf %raw_scores, %mask_f32 : vector<4xf32> + %masked_scores = vector.select %query_valid, %masked_scores0, %negative_f32x4 : vector<4xf32> + %block_max = kernel.subgroup.reduce %masked_scores : vector<4xf32> + %next_max = vector.maxnumf %current_max, %block_max : vector<4xf32> + %score_delta = vector.subf %masked_scores, %next_max : vector<4xf32> + %raw_probability = vector.expf %score_delta : vector<4xf32> + %probability = vector.select %query_valid, %raw_probability, %c0_f32x4 : vector<4xf32> + %block_sum = kernel.subgroup.reduce %probability : vector<4xf32> + %old_delta = vector.subf %current_max, %next_max : vector<4xf32> + %old_scale = vector.expf %old_delta : vector<4xf32> + %scaled_current_sum = vector.mulf %current_sum, %old_scale : vector<4xf32> + %next_sum = vector.addf %scaled_current_sum, %block_sum : vector<4xf32> + %probability_f16 = vector.fptrunc %probability : vector<4xf32> to vector<4xf16> + %probability0 = vector.extract %probability_f16[0] : vector<4xf16> -> f16 + %probability1 = vector.extract %probability_f16[1] : vector<4xf16> -> f16 + %probability2 = vector.extract %probability_f16[2] : vector<4xf16> -> f16 + %probability3 = vector.extract %probability_f16[3] : vector<4xf16> -> f16 + view.store %probability0, %probability_stage_view[%query_row0, %lane] : f16, view<16x64xf16, %probability_layout> + view.store %probability1, %probability_stage_view[%query_row1, %lane] : f16, view<16x64xf16, %probability_layout> + view.store %probability2, %probability_stage_view[%query_row2, %lane] : f16, view<16x64xf16, %probability_layout> + view.store %probability3, %probability_stage_view[%query_row3, %lane] : f16, view<16x64xf16, %probability_layout> + %lane_probability_maximum = vector.reduce %probability, %c0_f32 : vector<4xf32>, f32 + view.store %lane_probability_maximum, %key_visibility_view[%lane, %subgroup] : f32, view<64x4xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + // Compute the 64-channel tiles in this output group sequentially. This + // keeps one P*V accumulator live per wave and reuses a 2 KiB exchange + // tile instead of retaining multiple output groups. + %old_scale_f16 = vector.fptrunc %old_scale : vector<4xf32> to vector<4xf16> + %old_scale0_scalar = vector.extract %old_scale_f16[0] : vector<4xf16> -> f16 + %old_scale1_scalar = vector.extract %old_scale_f16[1] : vector<4xf16> -> f16 + %old_scale2_scalar = vector.extract %old_scale_f16[2] : vector<4xf16> -> f16 + %old_scale3_scalar = vector.extract %old_scale_f16[3] : vector<4xf16> -> f16 + %old_scale0 = vector.splat %old_scale0_scalar : vector<4xf16> + %old_scale1 = vector.splat %old_scale1_scalar : vector<4xf16> + %old_scale2 = vector.splat %old_scale2_scalar : vector<4xf16> + %old_scale3 = vector.splat %old_scale3_scalar : vector<4xf16> + %scaled_current_output0 = vector.mulf %current_output0, %old_scale0 : vector<4xf16> + %scaled_current_output1 = vector.mulf %current_output1, %old_scale1 : vector<4xf16> + %scaled_current_output2 = vector.mulf %current_output2, %old_scale2 : vector<4xf16> + %scaled_current_output3 = vector.mulf %current_output3, %old_scale3 : vector<4xf16> + %next_output0, %next_output1, %next_output2, %next_output3 = scf.for %output_tile = [%c0 to %c4 step %c1](%tile_output0 = %scaled_current_output0 : vector<4xf16>, %tile_output1 = %scaled_current_output1 : vector<4xf16>, %tile_output2 = %scaled_current_output2 : vector<4xf16>, %tile_output3 = %scaled_current_output3 : vector<4xf16>) -> (vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) unroll { + %output_tile_valid = index.cmp ult, %output_tile, %group_output_tile_count : index + %updated_output0, %updated_output1, %updated_output2, %updated_output3 = scf.if %output_tile_valid -> (vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + %global_output_tile0 = index.add %output_group_tile_base, %output_tile : index + %global_output_tile = index.assume %global_output_tile0 [lt(%global_output_tile0, %output_tile_count)] : index + %output_tile_channel = index.mul %global_output_tile, %c64 : index + %value_channel0 = index.add %value_head_base, %output_tile_channel : index + %value_channel = index.add %value_channel0, %subgroup_product_channel : index + %product_init = vector.fragment %c0_f16x8 shape [%m, %n] : vector<8xf16> + // Masked keys must reach P*V as V = +0, not just P = 0. The F16 WMMA + // result depended on the sign of V at keys whose probability is zero + // (measured: flipping the sign of one masked V row changed one output + // element by one F16 ulp). Rows past the sequence end hold whatever an + // earlier request left there, so identical requests gave different + // logits. Blocks where some valid row masks some key stage V through + // LDS and zero the keys no row of this tile gives a probability. + %product_fragment = scf.if %block_fully_visible -> (vector<8xf16>) { + %direct_product_fragment = scf.for %key_tile = [%c0 to %c64 step %c16](%product_accumulator = %product_init : vector<8xf16>) -> (vector<8xf16>) unroll schedule(recurrence) { + %value_token0 = index.add %key_origin, %key_tile : index + %value_token = index.assume %value_token0 [lt(%value_token0, %last_full_key_tile_start)] : index + %probability_fragment = vector.fragment.load %probability_stage_view[%c0, %key_tile] shape [%m, %k] : view<16x64xf16, %probability_layout> -> vector<16xf16> + %value_fragment = vector.fragment.load %value_view[%value_token, %value_channel] shape [%k, %n] : view<[%bounded_key_value_token_count]x[%value_width]xf16, %value_layout> -> vector<16xf16> + %next_product_accumulator = vector.mma %probability_fragment, %value_fragment, %product_accumulator : vector<16xf16>, vector<16xf16>, vector<8xf16> + scf.yield %next_product_accumulator : vector<8xf16> + } + scf.yield %direct_product_fragment : vector<8xf16> + } else { + %staged_product_fragment = scf.for %key_tile = [%c0 to %c64 step %c16](%product_accumulator = %product_init : vector<8xf16>) -> (vector<8xf16>) { + scf.for %load_iteration = [%c0 to %c4 step %c1] unroll { + %linear = index.madd %load_iteration, %c256, %workitem : index + %staged_key_row = index.div %linear, %c64 : index + %staged_channel = index.rem %linear, %c64 : index + %staged_block_key0 = index.add %key_tile, %staged_key_row : index + %staged_block_key = index.min %staged_block_key0, %c63 : index + %staged_key_probabilities = vector.load %key_visibility_view[%staged_block_key, %c0] : view<64x4xf32> -> vector<4xf32> + %staged_key_probability = vector.reduce %staged_key_probabilities, %c0_f32 : vector<4xf32>, f32 + %staged_key_visible = scalar.cmpf ogt, %staged_key_probability, %c0_f32 : f32 + %staged_value = scf.if %staged_key_visible -> (f16) { + %staged_token0 = index.add %key_origin, %staged_block_key : index + %staged_token = index.assume %staged_token0 [lt(%staged_token0, %bounded_key_value_token_count)] : index + %staged_value_channel0 = index.add %value_channel0, %staged_channel : index + %staged_value_channel = index.assume %staged_value_channel0 [lt(%staged_value_channel0, %value_width)] : index + %loaded = view.load %value_view[%staged_token, %staged_value_channel] : view<[%bounded_key_value_token_count]x[%value_width]xf16, %value_layout> -> f16 + scf.yield %loaded : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %staged_value, %staged_value_view[%staged_key_row, %staged_channel] : f16, view<16x64xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %probability_fragment = vector.fragment.load %probability_stage_view[%c0, %key_tile] shape [%m, %k] : view<16x64xf16, %probability_layout> -> vector<16xf16> + %value_fragment = vector.fragment.load %staged_value_view[%c0, %subgroup_product_channel] shape [%k, %n] : view<16x64xf16> -> vector<16xf16> + %next_product_accumulator = vector.mma %probability_fragment, %value_fragment, %product_accumulator : vector<16xf16>, vector<16xf16>, vector<8xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %next_product_accumulator : vector<8xf16> + } + scf.yield %staged_product_fragment : vector<8xf16> + } + vector.fragment.store %product_fragment, %product_stage_view[%c0, %subgroup_product_channel] shape [%m, %n] : vector<8xf16>, view<16x64xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + %owns_output_tile = index.cmp eq, %lane_output_tile, %output_tile : index + %tile_updated_output0, %tile_updated_output1, %tile_updated_output2, %tile_updated_output3 = scf.if %owns_output_tile -> (vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + %block_output0 = vector.load %product_stage_view[%query_row0, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output1 = vector.load %product_stage_view[%query_row1, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output2 = vector.load %product_stage_view[%query_row2, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output3 = vector.load %product_stage_view[%query_row3, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %updated_tile_output0 = vector.addf %tile_output0, %block_output0 : vector<4xf16> + %updated_tile_output1 = vector.addf %tile_output1, %block_output1 : vector<4xf16> + %updated_tile_output2 = vector.addf %tile_output2, %block_output2 : vector<4xf16> + %updated_tile_output3 = vector.addf %tile_output3, %block_output3 : vector<4xf16> + scf.yield %updated_tile_output0, %updated_tile_output1, %updated_tile_output2, %updated_tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } else { + scf.yield %tile_output0, %tile_output1, %tile_output2, %tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + // Complete every read before the next output tile overwrites LDS. + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %tile_updated_output0, %tile_updated_output1, %tile_updated_output2, %tile_updated_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } else { + scf.yield %tile_output0, %tile_output1, %tile_output2, %tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + scf.yield %updated_output0, %updated_output1, %updated_output2, %updated_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + scf.yield %next_max, %next_sum, %next_output0, %next_output1, %next_output2, %next_output3 : vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } else { + scf.yield %current_max, %current_sum, %current_output0, %current_output1, %current_output2, %current_output3 : vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + scf.yield %next_block_max, %next_block_sum, %next_block_output0, %next_block_output1, %next_block_output2, %next_block_output3 : vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + // A masked 32-row tile handles the final 1-63 KV rows with the same WMMA + // fragments as the aligned path. K and V are staged separately into product + // scratch, so every padded element is initialized and no physical padding is + // required of the caller. Each tile rounds probabilities to F16 before P*V. + %tail_score_wave = index.cmp ult, %subgroup, %c2 : index + %final_max, %final_sum, %final_output0, %final_output1, %final_output2, %final_output3 = scf.for %tail_key_origin = [%full_key_value_token_count to %bounded_key_value_token_count step %c32](%current_max = %full_max : vector<4xf32>, %current_sum = %full_sum : vector<4xf32>, %current_output0 = %full_output0 : vector<4xf16>, %current_output1 = %full_output1 : vector<4xf16>, %current_output2 = %full_output2 : vector<4xf16>, %current_output3 = %full_output3 : vector<4xf16>) -> (vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + %tail_remaining = index.sub %bounded_key_value_token_count, %tail_key_origin : index + %tail_key_count = index.min %tail_remaining, %c32 : index + // Cooperatively stage one K tile, explicitly zeroing the padded rows. + scf.for %load_iteration = [%c0 to %tail_key_load_iteration_count step %c1] unroll { + %linear = index.madd %load_iteration, %c256, %workitem : index + %tail_key_row = index.div %linear, %qk_head_size : index + %tail_key_channel = index.rem %linear, %qk_head_size : index + %tail_key_valid = index.cmp ult, %tail_key_row, %tail_key_count : index + %tail_key_value = scf.if %tail_key_valid -> (f16) { + %tail_key_token0 = index.add %tail_key_origin, %tail_key_row : index + %tail_key_token = index.assume %tail_key_token0 [lt(%tail_key_token0, %bounded_key_value_token_count)] : index + %global_key_channel = index.add %key_head_base, %tail_key_channel : index + %loaded = view.load %key_view[%tail_key_token, %global_key_channel] : view<[%bounded_key_value_token_count]x[%key_width]xf16> -> f16 + scf.yield %loaded : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %tail_key_value, %tail_key_value_stage_view[%tail_key_row, %tail_key_channel] : f16, view<32x[%tail_stage_head_size]xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + // Waves zero and one compute the two 16-column QK fragments in this tile. + scf.if %tail_score_wave { + %tail_score_subgroup = index.assume %subgroup [range(%subgroup, 0, 1)] : index + %tail_score_column = index.mul %tail_score_subgroup, %c16 : index + %tail_score_init_values = vector.constant 0.0 : vector<4xf32> + %tail_score_init = vector.fragment %tail_score_init_values shape [%m, %n] : vector<4xf32> + %tail_score_fragment = scf.for %head_tile = [%c0 to %qk_head_size step %c16](%score_accumulator = %tail_score_init : vector<4xf32>) -> (vector<4xf32>) unroll { + %key_fragment = vector.fragment.load %tail_key_value_stage_view[%tail_score_column, %head_tile] shape [%m, %k] : view<32x[%tail_stage_head_size]xf16> -> vector<16xf16> + %query_fragment = vector.fragment.load %query_transposed_view[%head_tile, %c0] shape [%k, %n] : view<[%qk_head_size]x16xf16, %query_transposed_layout> -> vector<16xf16> + %next_score_accumulator = vector.mma %key_fragment, %query_fragment, %score_accumulator : vector<16xf16>, vector<16xf16>, vector<4xf32> + scf.yield %next_score_accumulator : vector<4xf32> + } + vector.fragment.store %tail_score_fragment, %score_stage_view[%tail_score_column, %c0] shape [%m, %n] : vector<4xf32>, view<64x20xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + // LDS changes ownership from the QK wave to four query-row waves. Lanes + // beyond the logical tail never read the score or mask buffers. + %tail_lane_valid = index.cmp ult, %lane, %tail_key_count : index + %tail_valid0 = scalar.andi %tail_lane_valid, %query_valid0 : i1 + %tail_valid1 = scalar.andi %tail_lane_valid, %query_valid1 : i1 + %tail_valid2 = scalar.andi %tail_lane_valid, %query_valid2 : i1 + %tail_valid3 = scalar.andi %tail_lane_valid, %query_valid3 : i1 + %tail_valid = vector.from_elements %tail_valid0, %tail_valid1, %tail_valid2, %tail_valid3 : vector<4xi1> + %tail_key_token0 = index.add %tail_key_origin, %lane : index + %masked_score0 = scf.if %tail_valid0 -> (f32) { + %tail_key_token = index.assume %tail_key_token0 [lt(%tail_key_token0, %bounded_key_value_token_count)] : index + %raw_score = view.load %score_stage_view[%lane, %query_row0] : view<64x20xf32> -> f32 + %mask_f16 = view.load %mask_view[%query_token0, %tail_key_token] : view<[%query_token_count]x[%bounded_key_value_token_count]xf16> -> f16 + %mask_f32 = scalar.extf %mask_f16 : f16 to f32 + %score = scalar.addf %raw_score, %mask_f32 : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score1 = scf.if %tail_valid1 -> (f32) { + %tail_key_token = index.assume %tail_key_token0 [lt(%tail_key_token0, %bounded_key_value_token_count)] : index + %raw_score = view.load %score_stage_view[%lane, %query_row1] : view<64x20xf32> -> f32 + %mask_f16 = view.load %mask_view[%query_token1, %tail_key_token] : view<[%query_token_count]x[%bounded_key_value_token_count]xf16> -> f16 + %mask_f32 = scalar.extf %mask_f16 : f16 to f32 + %score = scalar.addf %raw_score, %mask_f32 : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score2 = scf.if %tail_valid2 -> (f32) { + %tail_key_token = index.assume %tail_key_token0 [lt(%tail_key_token0, %bounded_key_value_token_count)] : index + %raw_score = view.load %score_stage_view[%lane, %query_row2] : view<64x20xf32> -> f32 + %mask_f16 = view.load %mask_view[%query_token2, %tail_key_token] : view<[%query_token_count]x[%bounded_key_value_token_count]xf16> -> f16 + %mask_f32 = scalar.extf %mask_f16 : f16 to f32 + %score = scalar.addf %raw_score, %mask_f32 : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score3 = scf.if %tail_valid3 -> (f32) { + %tail_key_token = index.assume %tail_key_token0 [lt(%tail_key_token0, %bounded_key_value_token_count)] : index + %raw_score = view.load %score_stage_view[%lane, %query_row3] : view<64x20xf32> -> f32 + %mask_f16 = view.load %mask_view[%query_token3, %tail_key_token] : view<[%query_token_count]x[%bounded_key_value_token_count]xf16> -> f16 + %mask_f32 = scalar.extf %mask_f16 : f16 to f32 + %score = scalar.addf %raw_score, %mask_f32 : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_scores = vector.from_elements %masked_score0, %masked_score1, %masked_score2, %masked_score3 : vector<4xf32> + %block_max = kernel.subgroup.reduce %masked_scores : vector<4xf32> + %next_max = vector.maxnumf %current_max, %block_max : vector<4xf32> + %score_delta = vector.subf %masked_scores, %next_max : vector<4xf32> + %raw_probability = vector.expf %score_delta : vector<4xf32> + %probability = vector.select %tail_valid, %raw_probability, %c0_f32x4 : vector<4xf32> + %block_sum = kernel.subgroup.reduce %probability : vector<4xf32> + %old_delta = vector.subf %current_max, %next_max : vector<4xf32> + %old_scale = vector.expf %old_delta : vector<4xf32> + %scaled_current_sum = vector.mulf %current_sum, %old_scale : vector<4xf32> + %next_sum = vector.addf %scaled_current_sum, %block_sum : vector<4xf32> + %probability_f16 = vector.fptrunc %probability : vector<4xf32> to vector<4xf16> + %probability0 = vector.extract %probability_f16[0] : vector<4xf16> -> f16 + %probability1 = vector.extract %probability_f16[1] : vector<4xf16> -> f16 + %probability2 = vector.extract %probability_f16[2] : vector<4xf16> -> f16 + %probability3 = vector.extract %probability_f16[3] : vector<4xf16> -> f16 + view.store %probability0, %probability_stage_view[%query_row0, %lane] : f16, view<16x64xf16, %probability_layout> + view.store %probability1, %probability_stage_view[%query_row1, %lane] : f16, view<16x64xf16, %probability_layout> + view.store %probability2, %probability_stage_view[%query_row2, %lane] : f16, view<16x64xf16, %probability_layout> + view.store %probability3, %probability_stage_view[%query_row3, %lane] : f16, view<16x64xf16, %probability_layout> + %lane_probability_maximum = vector.reduce %probability, %c0_f32 : vector<4xf32>, f32 + view.store %lane_probability_maximum, %key_visibility_view[%lane, %subgroup] : f32, view<64x4xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + // Reuse the tail K scratch for V after probability publication. Each + // workgroup stages only its output group, starting at local column zero. + // Stage the fixed maximum output group and zero-pad a partial group. + // This keeps the unrolled trip count and local coordinates static. + scf.for %load_iteration = [%c0 to %tail_value_load_iteration_count step %c1] unroll { + %linear = index.madd %load_iteration, %c256, %workitem : index + %tail_value_row = index.div %linear, %tail_value_stage_channel_count : index + %tail_value_channel = index.rem %linear, %tail_value_stage_channel_count : index + %tail_value_row_valid = index.cmp ult, %tail_value_row, %tail_key_count : index + %tail_value_channel_valid = index.cmp ult, %tail_value_channel, %group_output_channel_count : index + %tail_value_in_range = scalar.andi %tail_value_row_valid, %tail_value_channel_valid : i1 + %tail_value_block_key = index.min %tail_value_row, %c63 : index + %tail_key_probabilities = vector.load %key_visibility_view[%tail_value_block_key, %c0] : view<64x4xf32> -> vector<4xf32> + %tail_key_probability = vector.reduce %tail_key_probabilities, %c0_f32 : vector<4xf32>, f32 + %tail_key_visible = scalar.cmpf ogt, %tail_key_probability, %c0_f32 : f32 + %tail_value_valid = scalar.andi %tail_value_in_range, %tail_key_visible : i1 + %output_value_channel = index.add %output_group_channel_base, %tail_value_channel : index + %tail_value = scf.if %tail_value_valid -> (f16) { + %tail_value_token0 = index.add %tail_key_origin, %tail_value_row : index + %tail_value_token = index.assume %tail_value_token0 [lt(%tail_value_token0, %bounded_key_value_token_count)] : index + %global_value_channel0 = index.add %value_head_base, %output_value_channel : index + %global_value_channel = index.assume %global_value_channel0 [lt(%global_value_channel0, %value_width)] : index + %loaded = view.load %value_view[%tail_value_token, %global_value_channel] : view<[%bounded_key_value_token_count]x[%value_width]xf16, %value_layout> -> f16 + scf.yield %loaded : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %tail_value, %tail_key_value_stage_view[%tail_value_row, %tail_value_channel] : f16, view<32x[%tail_stage_head_size]xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + // Keep tail V staging separate from the 2 KiB product exchange, then use + // the same grouped 64-channel phases as the aligned path. + %old_scale_f16 = vector.fptrunc %old_scale : vector<4xf32> to vector<4xf16> + %old_scale0_scalar = vector.extract %old_scale_f16[0] : vector<4xf16> -> f16 + %old_scale1_scalar = vector.extract %old_scale_f16[1] : vector<4xf16> -> f16 + %old_scale2_scalar = vector.extract %old_scale_f16[2] : vector<4xf16> -> f16 + %old_scale3_scalar = vector.extract %old_scale_f16[3] : vector<4xf16> -> f16 + %old_scale0 = vector.splat %old_scale0_scalar : vector<4xf16> + %old_scale1 = vector.splat %old_scale1_scalar : vector<4xf16> + %old_scale2 = vector.splat %old_scale2_scalar : vector<4xf16> + %old_scale3 = vector.splat %old_scale3_scalar : vector<4xf16> + %scaled_current_output0 = vector.mulf %current_output0, %old_scale0 : vector<4xf16> + %scaled_current_output1 = vector.mulf %current_output1, %old_scale1 : vector<4xf16> + %scaled_current_output2 = vector.mulf %current_output2, %old_scale2 : vector<4xf16> + %scaled_current_output3 = vector.mulf %current_output3, %old_scale3 : vector<4xf16> + %next_output0, %next_output1, %next_output2, %next_output3 = scf.for %output_tile = [%c0 to %c4 step %c1](%tile_output0 = %scaled_current_output0 : vector<4xf16>, %tile_output1 = %scaled_current_output1 : vector<4xf16>, %tile_output2 = %scaled_current_output2 : vector<4xf16>, %tile_output3 = %scaled_current_output3 : vector<4xf16>) -> (vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) unroll { + %output_tile_valid = index.cmp ult, %output_tile, %group_output_tile_count : index + %updated_output0, %updated_output1, %updated_output2, %updated_output3 = scf.if %output_tile_valid -> (vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + %output_tile_channel = index.mul %output_tile, %c64 : index + %value_channel = index.add %output_tile_channel, %subgroup_product_channel : index + %tail_product_init = vector.fragment %c0_f16x8 shape [%m, %n] : vector<8xf16> + %tail_product_fragment = scf.for %key_tile = [%c0 to %c32 step %c16](%product_accumulator = %tail_product_init : vector<8xf16>) -> (vector<8xf16>) unroll { + %probability_fragment = vector.fragment.load %probability_stage_view[%c0, %key_tile] shape [%m, %k] : view<16x64xf16, %probability_layout> -> vector<16xf16> + %value_fragment = vector.fragment.load %tail_key_value_stage_view[%key_tile, %value_channel] shape [%k, %n] : view<32x[%tail_stage_head_size]xf16> -> vector<16xf16> + %next_product_accumulator = vector.mma %probability_fragment, %value_fragment, %product_accumulator : vector<16xf16>, vector<16xf16>, vector<8xf16> + scf.yield %next_product_accumulator : vector<8xf16> + } + vector.fragment.store %tail_product_fragment, %product_stage_view[%c0, %subgroup_product_channel] shape [%m, %n] : vector<8xf16>, view<16x64xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + %owns_output_tile = index.cmp eq, %lane_output_tile, %output_tile : index + %tile_updated_output0, %tile_updated_output1, %tile_updated_output2, %tile_updated_output3 = scf.if %owns_output_tile -> (vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + %block_output0 = vector.load %product_stage_view[%query_row0, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output1 = vector.load %product_stage_view[%query_row1, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output2 = vector.load %product_stage_view[%query_row2, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output3 = vector.load %product_stage_view[%query_row3, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %updated_tile_output0 = vector.addf %tile_output0, %block_output0 : vector<4xf16> + %updated_tile_output1 = vector.addf %tile_output1, %block_output1 : vector<4xf16> + %updated_tile_output2 = vector.addf %tile_output2, %block_output2 : vector<4xf16> + %updated_tile_output3 = vector.addf %tile_output3, %block_output3 : vector<4xf16> + scf.yield %updated_tile_output0, %updated_tile_output1, %updated_tile_output2, %updated_tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } else { + scf.yield %tile_output0, %tile_output1, %tile_output2, %tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %tile_updated_output0, %tile_updated_output1, %tile_updated_output2, %tile_updated_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } else { + scf.yield %tile_output0, %tile_output1, %tile_output2, %tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + scf.yield %updated_output0, %updated_output1, %updated_output2, %updated_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + scf.yield %next_max, %next_sum, %next_output0, %next_output1, %next_output2, %next_output3 : vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + template.apply<@ggml.flash_attention.prefill.publish_output>(%query_token_count, %query_head_count, %value_head_size, %lane_has_output, %query_valid0, %query_valid1, %query_valid2, %query_valid3, %query_token0, %query_token1, %query_token2, %query_token3, %query_head, %lane_output_channel, %final_sum, %final_output0, %final_output1, %final_output2, %final_output3, %gate, %output_aligned) : (index, index, index, i1, i1, i1, i1, i1, index, index, index, index, index, index, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>, buffer, buffer) + template.return +} + +kernel.def target(@ggml_flash_attention_gfx11_wave64) @ggml_flash_attention_f32_f16_wmma(%query_token_count: index, %key_value_token_count: index) { + %query_head_count = config.get @ggml.flash_attention.query_head_count : index + %value_head_size0 = config.get @ggml.flash_attention.value_head_size : index + %value_head_size = index.assume %value_head_size0 [range(%value_head_size0, 64, 512), mul(%value_head_size0, 64)] : index + %c1 = index.constant 1 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c15 = index.constant 15 : index + %c16 = index.constant 16 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %padded_query_token_count = index.add %query_token_count, %c15 : index + %query_tile_count = index.div %padded_query_token_count, %c16 : index + %output_tile_count = index.div %value_head_size, %c64 : index + %padded_output_tile_count = index.add %output_tile_count, %c3 : index + %output_group_count = index.div %padded_output_tile_count, %c4 : index + kernel.launch.config workgroups(%query_tile_count, %query_head_count, %output_group_count) workgroup_size(%c256, %c1, %c1) : index +} launch(%query_token_count: index, %key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %gate: buffer, %output: buffer) where [range(%query_token_count, 1, 2048)] { + template.apply<@ggml.flash_attention.prefill.body>(%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %gate, %output) : (index, index, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +// Test-only causal-mask construction for the production 14-token witness. +// Finite F16 minima retain exact zero probabilities without requiring infinity +// support from synthetic tensor generators. Selective linking drops this +// helper from production roots. +kernel.def target(@ggml_flash_attention_gfx11_wave64) @ggml_flash_attention_test_make_causal_mask() { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%c1, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%mask: buffer) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c14 = index.constant 14 : index + %c196 = index.constant 196 : index + %c0_offset = index.constant 0 : offset + %c0_f16 = scalar.constant 0.0 : f16 + %negative_f16 = scalar.constant -65504.0 : f16 + %workitem = kernel.workitem.id : index + %in_bounds = index.cmp ult, %workitem, %c196 : index + %mask_noalias = buffer.assume.noalias %mask : buffer + %mask_view = buffer.view %mask_noalias[%c0_offset] : buffer -> view<14x14xf16> + scf.if %in_bounds { + %row = index.div %workitem, %c14 : index + %column = index.rem %workitem, %c14 : index + %row_limit = index.add %row, %c1 : index + %is_visible = index.cmp ult, %column, %row_limit : index + %value = scf.select %is_visible, %c0_f16, %negative_f16 : f16 + view.store %value, %mask_view[%row, %column] : f16, view<14x14xf16> + } + kernel.return +} + +// Test-only row extraction for comparing one multirow attention dispatch with +// independent one-row dispatches over identical data. Selective linking drops +// this helper from production roots. +kernel.def target(@ggml_flash_attention_gfx11_wave64) @ggml_flash_attention_test_extract_row(%query_token_count: index, %context_count: index, %source_row: index) { + %query_head_count = config.get @ggml.flash_attention.query_head_count : index + %qk_head_size0 = config.get @ggml.flash_attention.qk_head_size : index + %qk_head_size = index.assume %qk_head_size0 [range(%qk_head_size0, 16, 576), mul(%qk_head_size0, 16)] : index + %context_capacity = config.get @ggml.flash_attention.test.context_capacity : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %c1 = index.constant 1 : index + %query_element_count = index.mul %query_head_count, %qk_head_size : index + %element_count = index.add %query_element_count, %context_capacity : index + %padded_element_count = index.add %element_count, %c255 : index + %workgroup_count = index.div %padded_element_count, %c256 : index + kernel.launch.config workgroups(%workgroup_count, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%query_token_count: index, %context_count: index, %source_row: index, %source_query: buffer, %source_mask: buffer, %target_query: buffer, %target_mask: buffer) { + %c0 = index.constant 0 : index + %c0_offset = index.constant 0 : offset + %c256 = index.constant 256 : index + %query_head_count = config.get @ggml.flash_attention.query_head_count : index + %qk_head_size0 = config.get @ggml.flash_attention.qk_head_size : index + %qk_head_size = index.assume %qk_head_size0 [range(%qk_head_size0, 16, 576), mul(%qk_head_size0, 16)] : index + %context_capacity = config.get @ggml.flash_attention.test.context_capacity : index + %bounded_query_token_count = index.assume %query_token_count [range(%query_token_count, 1, 2048)] : index + %bounded_context_count = index.assume %context_count [range(%context_count, 1, 32768), le(%context_count, %context_capacity)] : index + %bounded_source_row, %source_query_token_count = index.assume %source_row, %bounded_query_token_count [range(%source_row, 0, 2047), lt(%source_row, %bounded_query_token_count)] : index, index + %workgroup = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %linear = index.madd %workgroup, %c256, %workitem : index + %query_element_count = index.mul %query_head_count, %qk_head_size : index + %element_count = index.add %query_element_count, %bounded_context_count : index + %in_bounds = index.cmp ult, %linear, %element_count : index + %source_query_noalias, %source_mask_noalias, %target_query_noalias, %target_mask_noalias = buffer.assume.noalias %source_query, %source_mask, %target_query, %target_mask : buffer, buffer, buffer, buffer + %source_query_view = buffer.view %source_query_noalias[%c0_offset] : buffer -> view<[%source_query_token_count]x[%query_head_count]x[%qk_head_size]xf32> + %source_mask_view = buffer.view %source_mask_noalias[%c0_offset] : buffer -> view<[%source_query_token_count]x[%bounded_context_count]xf16> + %target_query_view = buffer.view %target_query_noalias[%c0_offset] : buffer -> view<1x[%query_head_count]x[%qk_head_size]xf32> + %target_mask_view = buffer.view %target_mask_noalias[%c0_offset] : buffer -> view<1x[%bounded_context_count]xf16> + scf.if %in_bounds { + %is_query_element = index.cmp ult, %linear, %query_element_count : index + scf.if %is_query_element { + %query_head = index.div %linear, %qk_head_size : index + %query_channel = index.rem %linear, %qk_head_size : index + %value = view.load %source_query_view[%bounded_source_row, %query_head, %query_channel] : view<[%source_query_token_count]x[%query_head_count]x[%qk_head_size]xf32> -> f32 + view.store %value, %target_query_view[%c0, %query_head, %query_channel] : f32, view<1x[%query_head_count]x[%qk_head_size]xf32> + } else { + %mask_column = index.sub %linear, %query_element_count : index + %value = view.load %source_mask_view[%bounded_source_row, %mask_column] : view<[%source_query_token_count]x[%bounded_context_count]xf16> -> f16 + view.store %value, %target_mask_view[%c0, %mask_column] : f16, view<1x[%bounded_context_count]xf16> + } + } + kernel.return +} + +// The mask selects the first KV row exactly. QK, F16 probability conversion, +// GQA addressing, and the P*V path all execute, while the expected result is +// the first V row and remains auditable as an iota. +check.case public @ggml_flash_attention_f32_f16_wmma_selected_row_case { + %query_token_count = check.literal value(1) : index + %key_value_token_count = check.literal value(64) : index + %query = check.generate.fill value(1.0) : tensor<1x1x128xf32> + %key = check.generate.fill value(1.0) : tensor<64x1x128xf16> + %value = check.generate.iota offset(0.0) step(0.125) : tensor<64x1x128xf16> + %mask = check.generate.iota offset(0.0) step(-10000.0) : tensor<1x64xf16> + %output = check.generate.fill value(-1.0) : tensor<1x1x128xf32> + %expected = check.generate.iota offset(0.0) step(0.125) : tensor<1x1x128xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %query, %output) : [index, index](index, index, tensor<1x1x128xf32>, tensor<64x1x128xf16>, tensor<64x1x128xf16>, tensor<1x64xf16>, tensor<1x1x128xf32>, tensor<1x1x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<1x1x128xf32> + check.return +} + +// Four query rows select four distinct KV rows. The 33-element mask period +// maps row-major index row*32+column to zero exactly when row == column for +// rows zero through three. The expected output is therefore the first four +// rows of V, preserving an auditable iota while distinguishing every component +// in the four-row per-wave ownership path. +check.case public @ggml_flash_attention_f32_f16_wmma_multirow_selected_case { + %query_token_count = check.literal value(4) : index + %key_value_token_count = check.literal value(32) : index + %query = check.generate.fill value(1.0) : tensor<4x1x128xf32> + %key = check.generate.fill value(1.0) : tensor<32x1x128xf16> + %value = check.generate.iota offset(0.0) step(0.125) : tensor<32x1x128xf16> + %mask = check.generate.iota offset(0.0) step(-10000.0) period(33) : tensor<4x32xf16> + %output = check.generate.fill value(-1.0) : tensor<4x1x128xf32> + %expected = check.generate.iota offset(0.0) step(0.125) : tensor<4x1x128xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %query, %output) : [index, index](index, index, tensor<4x1x128xf32>, tensor<32x1x128xf16>, tensor<32x1x128xf16>, tensor<4x32xf16>, tensor<4x1x128xf32>, tensor<4x1x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<4x1x128xf32> + check.return +} + +// A unit mask step gives every row the same geometric softmax distribution, +// shifted to begin at the row-matched KV entry. Values advance by 1/4096, so +// each KV row adds exactly 1/32 and the infinite-series weighted row offset is +// 1/(32*(e-1)). Terms that wrap at period 33 are below F16 significance. This +// exercises four independent online-softmax and P*V states while retaining a +// compact closed-form expected iota. +check.case public @ggml_flash_attention_f32_f16_wmma_multirow_online_case { + %query_token_count = check.literal value(4) : index + %key_value_token_count = check.literal value(32) : index + %query = check.generate.fill value(1.0) : tensor<4x1x128xf32> + %key = check.generate.fill value(1.0) : tensor<32x1x128xf16> + %value = check.generate.iota offset(0.0) step(0.000244140625) : tensor<32x1x128xf16> + %mask = check.generate.iota offset(0.0) step(-1.0) period(33) : tensor<4x32xf16> + %output = check.generate.fill value(-1.0) : tensor<4x1x128xf32> + %expected = check.generate.iota offset(0.01818677224124075) step(0.000244140625) : tensor<4x1x128xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %query, %output) : [index, index](index, index, tensor<4x1x128xf32>, tensor<32x1x128xf16>, tensor<32x1x128xf16>, tensor<4x32xf16>, tensor<4x1x128xf32>, tensor<4x1x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<4x1x128xf32> + check.return +} + +// Compare the first four rows of one fourteen-row launch with four one-row +// launches over identical, nonuniform data and the production causal mask. +// Fourteen rows retain the exact dynamic dense-row stride, ownership, and +// query-validity pressure, while the one-row path is the proven containment. +// Row extraction only reshapes bindings and performs no attention arithmetic. +check.case public @ggml_flash_attention_f32_f16_wmma_multirow_differential_case { + %query_token_count = check.literal value(14) : index + %c1 = check.literal value(1) : index + %key_value_token_count = check.literal value(14) : index + %row0 = check.literal value(0) : index + %row1 = check.literal value(1) : index + %row2 = check.literal value(2) : index + %row3 = check.literal value(3) : index + %query_seed = check.param.seed base(5858425849414763858) count(1) : i64 + %key_seed = check.param.seed base(5858425849313057073) count(1) : i64 + %value_seed = check.param.seed base(5858425849497340977) count(1) : i64 + %query = check.generate.random.uniform seed(%query_seed) range(-1.0 to 1.0) : tensor<14x32x128xf32> + %key = check.generate.random.uniform seed(%key_seed) range(-1.0 to 1.0) : tensor<14x4x128xf16> + %value = check.generate.random.uniform seed(%value_seed) range(-1.0 to 1.0) : tensor<14x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<14x14xf16> + %actual = check.generate.fill value(-1.0) : tensor<14x32x128xf32> + %row_query = check.generate.fill value(0.0) : tensor<1x32x128xf32> + %row_mask = check.generate.fill value(0.0) : tensor<1x14xf16> + %expected0 = check.generate.fill value(0.0) : tensor<1x32x128xf32> + %expected1 = check.generate.fill value(0.0) : tensor<1x32x128xf32> + %expected2 = check.generate.fill value(0.0) : tensor<1x32x128xf32> + %expected3 = check.generate.fill value(0.0) : tensor<1x32x128xf32> + %actual0 = check.generate.fill value(-1.0) : tensor<1x32x128xf32> + %actual1 = check.generate.fill value(-1.0) : tensor<1x32x128xf32> + %actual2 = check.generate.fill value(-1.0) : tensor<1x32x128xf32> + %actual3 = check.generate.fill value(-1.0) : tensor<1x32x128xf32> + kernel.launch @ggml_flash_attention_test_make_causal_mask(%mask) : (tensor<14x14xf16>) + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %query, %actual) : [index, index](index, index, tensor<14x32x128xf32>, tensor<14x4x128xf16>, tensor<14x4x128xf16>, tensor<14x14xf16>, tensor<14x32x128xf32>, tensor<14x32x128xf32>) + kernel.launch @ggml_flash_attention_test_extract_row[%query_token_count, %key_value_token_count, %row0](%query_token_count, %key_value_token_count, %row0, %query, %mask, %row_query, %row_mask) : [index, index, index](index, index, index, tensor<14x32x128xf32>, tensor<14x14xf16>, tensor<1x32x128xf32>, tensor<1x14xf16>) + kernel.launch @ggml_flash_attention_f32_f16_wmma[%c1, %key_value_token_count](%c1, %key_value_token_count, %row_query, %key, %value, %row_mask, %row_query, %expected0) : [index, index](index, index, tensor<1x32x128xf32>, tensor<14x4x128xf16>, tensor<14x4x128xf16>, tensor<1x14xf16>, tensor<1x32x128xf32>, tensor<1x32x128xf32>) + kernel.launch @ggml_flash_attention_test_extract_row[%query_token_count, %key_value_token_count, %row0](%query_token_count, %key_value_token_count, %row0, %actual, %mask, %actual0, %row_mask) : [index, index, index](index, index, index, tensor<14x32x128xf32>, tensor<14x14xf16>, tensor<1x32x128xf32>, tensor<1x14xf16>) + kernel.launch @ggml_flash_attention_test_extract_row[%query_token_count, %key_value_token_count, %row1](%query_token_count, %key_value_token_count, %row1, %query, %mask, %row_query, %row_mask) : [index, index, index](index, index, index, tensor<14x32x128xf32>, tensor<14x14xf16>, tensor<1x32x128xf32>, tensor<1x14xf16>) + kernel.launch @ggml_flash_attention_f32_f16_wmma[%c1, %key_value_token_count](%c1, %key_value_token_count, %row_query, %key, %value, %row_mask, %row_query, %expected1) : [index, index](index, index, tensor<1x32x128xf32>, tensor<14x4x128xf16>, tensor<14x4x128xf16>, tensor<1x14xf16>, tensor<1x32x128xf32>, tensor<1x32x128xf32>) + kernel.launch @ggml_flash_attention_test_extract_row[%query_token_count, %key_value_token_count, %row1](%query_token_count, %key_value_token_count, %row1, %actual, %mask, %actual1, %row_mask) : [index, index, index](index, index, index, tensor<14x32x128xf32>, tensor<14x14xf16>, tensor<1x32x128xf32>, tensor<1x14xf16>) + kernel.launch @ggml_flash_attention_test_extract_row[%query_token_count, %key_value_token_count, %row2](%query_token_count, %key_value_token_count, %row2, %query, %mask, %row_query, %row_mask) : [index, index, index](index, index, index, tensor<14x32x128xf32>, tensor<14x14xf16>, tensor<1x32x128xf32>, tensor<1x14xf16>) + kernel.launch @ggml_flash_attention_f32_f16_wmma[%c1, %key_value_token_count](%c1, %key_value_token_count, %row_query, %key, %value, %row_mask, %row_query, %expected2) : [index, index](index, index, tensor<1x32x128xf32>, tensor<14x4x128xf16>, tensor<14x4x128xf16>, tensor<1x14xf16>, tensor<1x32x128xf32>, tensor<1x32x128xf32>) + kernel.launch @ggml_flash_attention_test_extract_row[%query_token_count, %key_value_token_count, %row2](%query_token_count, %key_value_token_count, %row2, %actual, %mask, %actual2, %row_mask) : [index, index, index](index, index, index, tensor<14x32x128xf32>, tensor<14x14xf16>, tensor<1x32x128xf32>, tensor<1x14xf16>) + kernel.launch @ggml_flash_attention_test_extract_row[%query_token_count, %key_value_token_count, %row3](%query_token_count, %key_value_token_count, %row3, %query, %mask, %row_query, %row_mask) : [index, index, index](index, index, index, tensor<14x32x128xf32>, tensor<14x14xf16>, tensor<1x32x128xf32>, tensor<1x14xf16>) + kernel.launch @ggml_flash_attention_f32_f16_wmma[%c1, %key_value_token_count](%c1, %key_value_token_count, %row_query, %key, %value, %row_mask, %row_query, %expected3) : [index, index](index, index, tensor<1x32x128xf32>, tensor<14x4x128xf16>, tensor<14x4x128xf16>, tensor<1x14xf16>, tensor<1x32x128xf32>, tensor<1x32x128xf32>) + kernel.launch @ggml_flash_attention_test_extract_row[%query_token_count, %key_value_token_count, %row3](%query_token_count, %key_value_token_count, %row3, %actual, %mask, %actual3, %row_mask) : [index, index, index](index, index, index, tensor<14x32x128xf32>, tensor<14x14xf16>, tensor<1x32x128xf32>, tensor<1x14xf16>) + check.expect.close actual(%actual0) expected(%expected0) atol(0.001) rtol(0.001) nan(same) : tensor<1x32x128xf32> + check.expect.close actual(%actual1) expected(%expected1) atol(0.001) rtol(0.001) nan(same) : tensor<1x32x128xf32> + check.expect.close actual(%actual2) expected(%expected2) atol(0.001) rtol(0.001) nan(same) : tensor<1x32x128xf32> + check.expect.close actual(%actual3) expected(%expected3) atol(0.001) rtol(0.001) nan(same) : tensor<1x32x128xf32> + check.return +} + +// Seventeen rows force a partial second query tile. Equal scores and constant +// values make every valid output exactly two while still exercising online +// normalization and the query-tail guards. +check.case public @ggml_flash_attention_f32_f16_wmma_query_tail_case { + %query_token_count = check.literal value(17) : index + %key_value_token_count = check.literal value(64) : index + %query = check.generate.fill value(1.0) : tensor<17x1x128xf32> + %key = check.generate.fill value(1.0) : tensor<64x1x128xf16> + %value = check.generate.fill value(2.0) : tensor<64x1x128xf16> + %mask = check.generate.fill value(0.0) : tensor<17x64xf16> + %output = check.generate.fill value(-1.0) : tensor<17x1x128xf32> + %expected = check.generate.fill value(2.0) : tensor<17x1x128xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %query, %output) : [index, index](index, index, tensor<17x1x128xf32>, tensor<64x1x128xf16>, tensor<64x1x128xf16>, tensor<17x64xf16>, tensor<17x1x128xf32>, tensor<17x1x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<17x1x128xf32> + check.return +} + +// Sixty-five KV rows force one masked cleanup tile after a full WMMA block. +// The mask selects only that final row, whose iota values begin at 1024, so +// omitting the cleanup path cannot accidentally satisfy the check. +check.case public @ggml_flash_attention_f32_f16_wmma_key_value_tail_case { + %query_token_count = check.literal value(1) : index + %key_value_token_count = check.literal value(65) : index + %query = check.generate.fill value(1.0) : tensor<1x1x128xf32> + %key = check.generate.fill value(1.0) : tensor<65x1x128xf16> + %value = check.generate.iota offset(0.0) step(0.125) : tensor<65x1x128xf16> + %mask = check.generate.iota offset(-64000.0) step(1000.0) : tensor<1x65xf16> + %output = check.generate.fill value(-1.0) : tensor<1x1x128xf32> + %expected = check.generate.iota offset(1024.0) step(0.125) : tensor<1x1x128xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %query, %output) : [index, index](index, index, tensor<1x1x128xf32>, tensor<65x1x128xf16>, tensor<65x1x128xf16>, tensor<1x65xf16>, tensor<1x1x128xf32>, tensor<1x1x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<1x1x128xf32> + check.return +} + +// The first 64-row block contains the selected value while every mask in the +// second block rounds to negative infinity. This preserves a finite online +// softmax state while exercising the workgroup-wide block-pruning path. +check.case public @ggml_flash_attention_f32_f16_wmma_pruned_block_case { + %query_token_count = check.literal value(1) : index + %key_value_token_count = check.literal value(128) : index + %query = check.generate.fill value(1.0) : tensor<1x1x128xf32> + %key = check.generate.fill value(1.0) : tensor<128x1x128xf16> + %value = check.generate.iota offset(0.0) step(0.125) : tensor<128x1x128xf16> + %mask = check.generate.iota offset(0.0) step(-2000.0) : tensor<1x128xf16> + %output = check.generate.fill value(-1.0) : tensor<1x1x128xf32> + %expected = check.generate.iota offset(0.0) step(0.125) : tensor<1x1x128xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %query, %output) : [index, index](index, index, tensor<1x1x128xf32>, tensor<128x1x128xf16>, tensor<128x1x128xf16>, tensor<1x128xf16>, tensor<1x1x128xf32>, tensor<1x1x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<1x1x128xf32> + check.return +} + +check.case public @ggml_flash_attention_f32_f16_wmma_benchmark_case { + %query_token_count = check.param.choice values([1, 32, 64, 128, 192, 255, 256, 257, 384, 511, 512, 513, 768, 1023, 1024, 1025, 1280, 1536, 1792, 2048]) name("query_token_count") : index + %key_value_token_count = check.param.choice values([64, 128, 192, 255, 256, 257, 384, 511, 512, 513, 768, 1023, 1024, 1025, 1280, 1536, 1792, 2048, 32768]) name("key_value_token_count") : index + %query = check.generate.fill value(0.0) : tensor<[%query_token_count]x32x128xf32> + %key = check.generate.fill value(0.0) : tensor<[%key_value_token_count]x4x128xf16> + %value = check.generate.fill value(0.0) : tensor<[%key_value_token_count]x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<[%query_token_count]x[%key_value_token_count]xf16> + %output = check.generate.fill value(1.0) : tensor<[%query_token_count]x32x128xf32> + %expected = check.generate.fill value(0.0) : tensor<[%query_token_count]x32x128xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %query, %output) : [index, index](index, index, tensor<[%query_token_count]x32x128xf32>, tensor<[%key_value_token_count]x4x128xf16>, tensor<[%key_value_token_count]x4x128xf16>, tensor<[%query_token_count]x[%key_value_token_count]xf16>, tensor<[%query_token_count]x32x128xf32>, tensor<[%query_token_count]x32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%query_token_count]x32x128xf32> + check.return +} + +// Four finite 64-row blocks followed by four negative-infinity blocks model +// the aggregate computed/pruned work ratio of 512-token causal prefill while +// keeping every workgroup's path identical for a stable microbenchmark. +check.case public @ggml_flash_attention_f32_f16_wmma_pruned_half_benchmark_case { + %query_token_count = check.literal value(512) : index + %key_value_token_count = check.literal value(512) : index + %query = check.generate.fill value(0.0) : tensor<512x32x128xf32> + %key = check.generate.fill value(0.0) : tensor<512x4x128xf16> + %value = check.generate.fill value(0.0) : tensor<512x4x128xf16> + %mask = check.generate.iota offset(0.0) step(-300.0) period(512) : tensor<512x512xf16> + %output = check.generate.fill value(1.0) : tensor<512x32x128xf32> + %expected = check.generate.fill value(0.0) : tensor<512x32x128xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %query, %output) : [index, index](index, index, tensor<512x32x128xf32>, tensor<512x4x128xf16>, tensor<512x4x128xf16>, tensor<512x512xf16>, tensor<512x32x128xf32>, tensor<512x32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<512x32x128xf32> + check.return +} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_selected_row_case> @ggml_flash_attention_f32_f16_wmma_selected_row + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_query_tail_case> @ggml_flash_attention_f32_f16_wmma_query_tail + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_key_value_tail_case> @ggml_flash_attention_f32_f16_wmma_key_value_tail + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_pruned_block_case> @ggml_flash_attention_f32_f16_wmma_pruned_block + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_pruned_half_benchmark_case> @ggml_flash_attention_f32_f16_wmma_prefill_512_pruned_half + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_decode_256 {key_value_token_count = 256, query_token_count = 1} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_decode_2048 {key_value_token_count = 2048, query_token_count = 1} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_decode_32768 {key_value_token_count = 32768, query_token_count = 1} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_prefill_32 {key_value_token_count = 256, query_token_count = 32} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_prefill_128 {key_value_token_count = 256, query_token_count = 128} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_prefill_512 {key_value_token_count = 512, query_token_count = 512} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_prefill_512_context_1024 {key_value_token_count = 1024, query_token_count = 512} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_prefill_512_context_1536 {key_value_token_count = 1536, query_token_count = 512} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_prefill_512_context_2048 {key_value_token_count = 2048, query_token_count = 512} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_prefill_1024 {key_value_token_count = 1024, query_token_count = 1024} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_prefill_2048 {key_value_token_count = 2048, query_token_count = 2048} + +// These aligned and boundary-adjacent self-attention shapes expose launch or +// tail cliffs that a powers-of-two-only benchmark would hide. They are +// measurement witnesses for one shape-specialized kernel, not routing buckets. +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_64 {key_value_token_count = 64, query_token_count = 64} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_128 {key_value_token_count = 128, query_token_count = 128} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_192 {key_value_token_count = 192, query_token_count = 192} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_255 {key_value_token_count = 255, query_token_count = 255} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_256 {key_value_token_count = 256, query_token_count = 256} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_257 {key_value_token_count = 257, query_token_count = 257} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_384 {key_value_token_count = 384, query_token_count = 384} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_511 {key_value_token_count = 511, query_token_count = 511} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_512 {key_value_token_count = 512, query_token_count = 512} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_513 {key_value_token_count = 513, query_token_count = 513} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_768 {key_value_token_count = 768, query_token_count = 768} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_1023 {key_value_token_count = 1023, query_token_count = 1023} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_1024 {key_value_token_count = 1024, query_token_count = 1024} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_1025 {key_value_token_count = 1025, query_token_count = 1025} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_1280 {key_value_token_count = 1280, query_token_count = 1280} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_1536 {key_value_token_count = 1536, query_token_count = 1536} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_1792 {key_value_token_count = 1792, query_token_count = 1792} + +check.benchmark<@ggml_flash_attention_f32_f16_wmma_benchmark_case> @ggml_flash_attention_f32_f16_wmma_self_2048 {key_value_token_count = 2048, query_token_count = 2048} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/gated_delta_net_f32_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/gated_delta_net_f32_wmma.loom new file mode 100644 index 000000000000..6356d0ed93eb --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/gated_delta_net_f32_wmma.loom @@ -0,0 +1,4227 @@ +template.decl @ggml.quantize_q8_1_x4.publish_vector4_strict(%publish_word: i1, %token_output_byte_base: offset, %channel: index, %values: vector<4xf32>, %scratch_values: view<256xf32>, %scratch_d: view<32xf32>, %output: buffer) +template.decl @ggml.rmsnorm_f32.subgroup_row_scale(%token_count: index, %token: index, %hidden_size: index, %epsilon: f32, %input: buffer) -> (f32) +template.decl @ggml.unary_f32.apply_vector4(%op: index, %values: vector<4xf32>) -> (vector<4xf32>) +config.decl @ggml.rmsnorm_gate_f32.rms_epsilon : f32 +config.decl @ggml.rmsnorm_gate_f32.gate_op : %value: index where [range(%value, 0, 23)] + +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +template.decl @llm.gated_delta_net.snapshot_prefix.body(%state_factor: f32, %cache_base: index, %state_offset: index, %wcol: index, %row0: index, %rowh: index, %wave: index, %lane: index, %lds: buffer, %snapshot_cache: buffer, %sa0: vector<4xf32>, %sa1: vector<4xf32>, %sa2: vector<4xf32>, %sa3: vector<4xf32>, %sa4: vector<4xf32>, %sa5: vector<4xf32>, %sa6: vector<4xf32>, %sa7: vector<4xf32>, %sa8: vector<4xf32>, %sa9: vector<4xf32>, %sa10: vector<4xf32>, %sa11: vector<4xf32>, %sa12: vector<4xf32>, %sa13: vector<4xf32>, %sa14: vector<4xf32>, %sa15: vector<4xf32>, %fz8: vector<8xf32>) + +template.decl @llm.gated_delta_net.snapshot_fragment_prefix.body(%state_factor: f32, %cache_base: index, %state_offset: index, %wcol: index, %row0: index, %rowh: index, %wave: index, %lane: index, %lds: buffer, %snapshot_cache: buffer, %lo0_state: vector<8xf32>, %lo1_state: vector<8xf32>, %lo2_state: vector<8xf32>, %lo3_state: vector<8xf32>, %hi0_state: vector<8xf32>, %hi1_state: vector<8xf32>, %hi2_state: vector<8xf32>, %hi3_state: vector<8xf32>, %fz8: vector<8xf32>) + +template.decl @llm.gated_delta_net.f32_wmma_head128.body(%state_inplace: i1, %publish_snapshots: i1, %projection_epilogue: i1, %snapshot_stride: index, %scale: f32, %q: buffer, %k: buffer, %v: buffer, %g: buffer, %beta: buffer, %bias: buffer, %a_scale: buffer, %state_in: buffer, %dst: buffer, %snapshot_cache: buffer, %rmsnorm_gate: i1, %rms_epsilon: f32, %rms_gate_op: index, %rms_weight: buffer, %raw_gate: buffer, %norm_output: buffer, %half_output: buffer, %select_state: i1, %state_row_count: index, %state_ids: buffer, %publish_q8: i1) + +// Chunked Gated DeltaNet matrix-core scan with one eight-wave workgroup per head and sequence. +// Each lane retains four rows of every state column in registers and stages the shard to F16 LDS once per 16-token chunk. +// The 63,296-byte LDS footprint permits one full-head workgroup per CU. +config.decl @llm.gated_delta_net.head_width : %value: index where [range(%value, 1, 4096)] + +config.decl @llm.gated_delta_net.head_count : %value: index where [range(%value, 1, 4096)] + +config.decl @llm.gated_delta_net.token_count : %value: index where [range(%value, 1, 1048576)] + +config.decl @llm.gated_delta_net.state_row_count : %value: index where [range(%value, 1, 262208)] + +config.decl @llm.gated_delta_net.sequence_count : %value: index where [range(%value, 1, 4096)] + +config.decl @llm.gated_delta_net.qk_stride1 : %value: index where [range(%value, 0, 1073741823)] + +config.decl @llm.gated_delta_net.qk_stride2 : %value: index where [range(%value, 0, 1073741823), mul(%value, 8)] + +config.decl @llm.gated_delta_net.qk_stride3 : %value: index where [range(%value, 0, 1073741823)] + +config.decl @llm.gated_delta_net.value_stride1 : %value: index where [range(%value, 0, 1073741823)] + +config.decl @llm.gated_delta_net.value_stride2 : %value: index where [range(%value, 0, 1073741823)] + +config.decl @llm.gated_delta_net.value_stride3 : %value: index where [range(%value, 0, 1073741823)] + +config.decl @llm.gated_delta_net.scalar_stride1 : %value: index where [range(%value, 0, 1073741823)] + +config.decl @llm.gated_delta_net.scalar_stride2 : %value: index where [range(%value, 0, 1073741823)] + +config.decl @llm.gated_delta_net.scalar_stride3 : %value: index where [range(%value, 0, 1073741823)] + +config.decl @llm.gated_delta_net.query_head_count : %value: index where [range(%value, 1, 4096)] + +config.decl @llm.gated_delta_net.query_sequence_ratio : %value: index where [range(%value, 1, 1024)] + +config.decl @llm.gated_delta_net.snapshot_stride : %value: index where [range(%value, 1, 1073741823)] + +config.decl @llm.gated_delta_net.l2_epsilon : f32 + +config.decl @llm.gated_delta_net.workgroup_size : %value: index where [range(%value, 32, 1024), mul(%value, 32)] + +kernel.def export("llm_gated_delta_net_f32_wmma_head128") @llm_gated_delta_net_f32_wmma_head128() { + %unit = index.constant 1 : index + %wg = index.constant 256 : index + %blk = index.constant 128 : index + %blk_m = index.constant 127 : index + %s_v = config.get @llm.gated_delta_net.head_width : index + %n_heads = config.get @llm.gated_delta_net.head_count : index + %n_seqs = config.get @llm.gated_delta_net.sequence_count : index + %cr = index.add %s_v, %blk_m : index + %col_blocks = index.div %cr, %blk : index + kernel.launch.config workgroups(%n_heads, %n_seqs, %col_blocks) workgroup_size(%wg, %unit, %unit) : index +} launch(%q: buffer, %k: buffer, %v: buffer, %g: buffer, %beta: buffer, %state_in: buffer, %dst: buffer) { + %state_inplace = scalar.constant false : i1 + %publish_snapshots = scalar.constant false : i1 + %projection_epilogue = scalar.constant false : i1 + %snapshot_stride = index.constant 1 : index + %scale = scalar.constant 0.0883883461356163 : f32 + %rmsnorm_gate = scalar.constant false : i1 + %select_state = scalar.constant false : i1 + %state_row_count = index.constant 1 : index + %publish_q8 = scalar.constant false : i1 + template.apply<@llm.gated_delta_net.f32_wmma_head128.body>(%state_inplace, %publish_snapshots, %projection_epilogue, %snapshot_stride, %scale, %q, %k, %v, %g, %beta, %g, %beta, %state_in, %dst, %state_in, %rmsnorm_gate, %scale, %snapshot_stride, %dst, %dst, %dst, %dst, %select_state, %state_row_count, %dst, %publish_q8) : (i1, i1, i1, index, f32, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, i1, f32, index, buffer, buffer, buffer, buffer, i1, index, buffer, i1) + kernel.return +} + +// Low-row scan with the alpha/beta projection epilogue folded into the +// existing one-thread-per-token scalar load point. +kernel.def export("llm_gated_delta_net_f32_wmma_head128_projection_epilogue") @llm_gated_delta_net_f32_wmma_head128_projection_epilogue() { + %unit = index.constant 1 : index + %wg = index.constant 256 : index + %blk = index.constant 128 : index + %blk_m = index.constant 127 : index + %s_v = config.get @llm.gated_delta_net.head_width : index + %n_heads = config.get @llm.gated_delta_net.head_count : index + %n_seqs = config.get @llm.gated_delta_net.sequence_count : index + %cr = index.add %s_v, %blk_m : index + %col_blocks = index.div %cr, %blk : index + kernel.launch.config workgroups(%n_heads, %n_seqs, %col_blocks) workgroup_size(%wg, %unit, %unit) : index +} launch(%q: buffer, %k: buffer, %v: buffer, %alpha_raw: buffer, %beta_raw: buffer, %bias: buffer, %a_scale: buffer, %state_in: buffer, %dst: buffer) { + %state_inplace = scalar.constant false : i1 + %publish_snapshots = scalar.constant false : i1 + %projection_epilogue = scalar.constant true : i1 + %snapshot_stride = index.constant 1 : index + %scale = scalar.constant 0.0883883461356163 : f32 + %rmsnorm_gate = scalar.constant false : i1 + %select_state = scalar.constant false : i1 + %state_row_count = index.constant 1 : index + %publish_q8 = scalar.constant false : i1 + template.apply<@llm.gated_delta_net.f32_wmma_head128.body>(%state_inplace, %publish_snapshots, %projection_epilogue, %snapshot_stride, %scale, %q, %k, %v, %alpha_raw, %beta_raw, %bias, %a_scale, %state_in, %dst, %state_in, %rmsnorm_gate, %scale, %snapshot_stride, %dst, %dst, %dst, %dst, %select_state, %state_row_count, %dst, %publish_q8) : (i1, i1, i1, index, f32, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, i1, f32, index, buffer, buffer, buffer, buffer, i1, index, buffer, i1) + kernel.return +} + +// Decode specialization keeps recurrent state in the graph-owned cache row. +kernel.def export("llm_gated_delta_net_f32_wmma_head128_inplace") @llm_gated_delta_net_f32_wmma_head128_inplace() { + %unit = index.constant 1 : index + %wg = index.constant 256 : index + %blk = index.constant 128 : index + %blk_m = index.constant 127 : index + %s_v = config.get @llm.gated_delta_net.head_width : index + %n_heads = config.get @llm.gated_delta_net.head_count : index + %n_seqs = config.get @llm.gated_delta_net.sequence_count : index + %cr = index.add %s_v, %blk_m : index + %col_blocks = index.div %cr, %blk : index + kernel.launch.config workgroups(%n_heads, %n_seqs, %col_blocks) workgroup_size(%wg, %unit, %unit) : index +} launch(%q: buffer, %k: buffer, %v: buffer, %g: buffer, %beta: buffer, %state_inout: buffer, %dst: buffer) { + %state_inplace = scalar.constant true : i1 + %publish_snapshots = scalar.constant false : i1 + %projection_epilogue = scalar.constant false : i1 + %snapshot_stride = index.constant 1 : index + %scale = scalar.constant 0.0883883461356163 : f32 + %rmsnorm_gate = scalar.constant false : i1 + %select_state = scalar.constant false : i1 + %state_row_count = index.constant 1 : index + %publish_q8 = scalar.constant false : i1 + template.apply<@llm.gated_delta_net.f32_wmma_head128.body>(%state_inplace, %publish_snapshots, %projection_epilogue, %snapshot_stride, %scale, %q, %k, %v, %g, %beta, %g, %beta, %state_inout, %dst, %state_inout, %rmsnorm_gate, %scale, %snapshot_stride, %dst, %dst, %dst, %dst, %select_state, %state_row_count, %dst, %publish_q8) : (i1, i1, i1, index, f32, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, i1, f32, index, buffer, buffer, buffer, buffer, i1, index, buffer, i1) + kernel.return +} + +// Low-row in-place scan with the alpha/beta projection epilogue folded into +// the existing one-thread-per-token scalar load point. +kernel.def export("llm_gated_delta_net_f32_wmma_head128_inplace_projection_epilogue") @llm_gated_delta_net_f32_wmma_head128_inplace_projection_epilogue() { + %unit = index.constant 1 : index + %wg = index.constant 256 : index + %blk = index.constant 128 : index + %blk_m = index.constant 127 : index + %s_v = config.get @llm.gated_delta_net.head_width : index + %n_heads = config.get @llm.gated_delta_net.head_count : index + %n_seqs = config.get @llm.gated_delta_net.sequence_count : index + %cr = index.add %s_v, %blk_m : index + %col_blocks = index.div %cr, %blk : index + kernel.launch.config workgroups(%n_heads, %n_seqs, %col_blocks) workgroup_size(%wg, %unit, %unit) : index +} launch(%q: buffer, %k: buffer, %v: buffer, %alpha_raw: buffer, %beta_raw: buffer, %bias: buffer, %a_scale: buffer, %state_inout: buffer, %dst: buffer) { + %state_inplace = scalar.constant true : i1 + %publish_snapshots = scalar.constant false : i1 + %projection_epilogue = scalar.constant true : i1 + %snapshot_stride = index.constant 1 : index + %scale = scalar.constant 0.0883883461356163 : f32 + %rmsnorm_gate = scalar.constant false : i1 + %select_state = scalar.constant false : i1 + %state_row_count = index.constant 1 : index + %publish_q8 = scalar.constant false : i1 + template.apply<@llm.gated_delta_net.f32_wmma_head128.body>(%state_inplace, %publish_snapshots, %projection_epilogue, %snapshot_stride, %scale, %q, %k, %v, %alpha_raw, %beta_raw, %bias, %a_scale, %state_inout, %dst, %state_inout, %rmsnorm_gate, %scale, %snapshot_stride, %dst, %dst, %dst, %dst, %select_state, %state_row_count, %dst, %publish_q8) : (i1, i1, i1, index, f32, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, i1, f32, index, buffer, buffer, buffer, buffer, i1, index, buffer, i1) + kernel.return +} + +// Low-row MTP specialization publishes each recurrence prefix and final state to the reverse-ordered rollback cache. +kernel.def export("llm_gated_delta_net_f32_wmma_head128_snapshot") @llm_gated_delta_net_f32_wmma_head128_snapshot() { + %unit = index.constant 1 : index + %wg = index.constant 256 : index + %blk = index.constant 128 : index + %blk_m = index.constant 127 : index + %s_v = config.get @llm.gated_delta_net.head_width : index + %n_heads = config.get @llm.gated_delta_net.head_count : index + %n_seqs = config.get @llm.gated_delta_net.sequence_count : index + %cr = index.add %s_v, %blk_m : index + %col_blocks = index.div %cr, %blk : index + kernel.launch.config workgroups(%n_heads, %n_seqs, %col_blocks) workgroup_size(%wg, %unit, %unit) : index +} launch(%q: buffer, %k: buffer, %v: buffer, %g: buffer, %beta: buffer, %state_in: buffer, %snapshot_cache: buffer, %dst: buffer) { + %state_inplace = scalar.constant false : i1 + %publish_snapshots = scalar.constant true : i1 + %projection_epilogue = scalar.constant false : i1 + %scale = scalar.constant 0.0883883461356163 : f32 + %snapshot_stride = config.get @llm.gated_delta_net.snapshot_stride : index + %rmsnorm_gate = scalar.constant false : i1 + %select_state = scalar.constant false : i1 + %state_row_count = index.constant 1 : index + %publish_q8 = scalar.constant false : i1 + template.apply<@llm.gated_delta_net.f32_wmma_head128.body>(%state_inplace, %publish_snapshots, %projection_epilogue, %snapshot_stride, %scale, %q, %k, %v, %g, %beta, %g, %beta, %state_in, %dst, %snapshot_cache, %rmsnorm_gate, %scale, %snapshot_stride, %dst, %dst, %dst, %dst, %select_state, %state_row_count, %dst, %publish_q8) : (i1, i1, i1, index, f32, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, i1, f32, index, buffer, buffer, buffer, buffer, i1, index, buffer, i1) + kernel.return +} + +// Low-row snapshot scan with the alpha/beta projection epilogue folded into +// the existing one-thread-per-token scalar load point. +kernel.def export("llm_gated_delta_net_f32_wmma_head128_snapshot_projection_epilogue") @llm_gated_delta_net_f32_wmma_head128_snapshot_projection_epilogue() { + %unit = index.constant 1 : index + %wg = index.constant 256 : index + %blk = index.constant 128 : index + %blk_m = index.constant 127 : index + %s_v = config.get @llm.gated_delta_net.head_width : index + %n_heads = config.get @llm.gated_delta_net.head_count : index + %n_seqs = config.get @llm.gated_delta_net.sequence_count : index + %cr = index.add %s_v, %blk_m : index + %col_blocks = index.div %cr, %blk : index + kernel.launch.config workgroups(%n_heads, %n_seqs, %col_blocks) workgroup_size(%wg, %unit, %unit) : index +} launch(%q: buffer, %k: buffer, %v: buffer, %alpha_raw: buffer, %beta_raw: buffer, %bias: buffer, %a_scale: buffer, %state_in: buffer, %snapshot_cache: buffer, %dst: buffer) { + %state_inplace = scalar.constant false : i1 + %publish_snapshots = scalar.constant true : i1 + %projection_epilogue = scalar.constant true : i1 + %scale = scalar.constant 0.0883883461356163 : f32 + %snapshot_stride = config.get @llm.gated_delta_net.snapshot_stride : index + %rmsnorm_gate = scalar.constant false : i1 + %select_state = scalar.constant false : i1 + %state_row_count = index.constant 1 : index + %publish_q8 = scalar.constant false : i1 + template.apply<@llm.gated_delta_net.f32_wmma_head128.body>(%state_inplace, %publish_snapshots, %projection_epilogue, %snapshot_stride, %scale, %q, %k, %v, %alpha_raw, %beta_raw, %bias, %a_scale, %state_in, %dst, %snapshot_cache, %rmsnorm_gate, %scale, %snapshot_stride, %dst, %dst, %dst, %dst, %select_state, %state_row_count, %dst, %publish_q8) : (i1, i1, i1, index, f32, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, i1, f32, index, buffer, buffer, buffer, buffer, i1, index, buffer, i1) + kernel.return +} + +kernel.def export("llm_gated_delta_net_f32_wmma_head128_selected_snapshot_projection_epilogue") @llm_gated_delta_net_f32_wmma_head128_selected_snapshot_projection_epilogue() { + %unit = index.constant 1 : index + %wg = index.constant 256 : index + %blk = index.constant 128 : index + %blk_m = index.constant 127 : index + %s_v = config.get @llm.gated_delta_net.head_width : index + %n_heads = config.get @llm.gated_delta_net.head_count : index + %n_seqs = config.get @llm.gated_delta_net.sequence_count : index + %cr = index.add %s_v, %blk_m : index + %col_blocks = index.div %cr, %blk : index + kernel.launch.config workgroups(%n_heads, %n_seqs, %col_blocks) workgroup_size(%wg, %unit, %unit) : index +} launch(%q: buffer, %k: buffer, %v: buffer, %alpha_raw: buffer, %beta_raw: buffer, %bias: buffer, %a_scale: buffer, %state_in: buffer, %snapshot_cache: buffer, %dst: buffer, %state_ids: buffer) { + %state_inplace = scalar.constant false : i1 + %publish_snapshots = scalar.constant true : i1 + %projection_epilogue = scalar.constant true : i1 + %scale = scalar.constant 0.0883883461356163 : f32 + %snapshot_stride = config.get @llm.gated_delta_net.snapshot_stride : index + %rmsnorm_gate = scalar.constant false : i1 + %select_state = scalar.constant true : i1 + %state_row_count = config.get @llm.gated_delta_net.state_row_count : index + %publish_q8 = scalar.constant false : i1 + template.apply<@llm.gated_delta_net.f32_wmma_head128.body>(%state_inplace, %publish_snapshots, %projection_epilogue, %snapshot_stride, %scale, %q, %k, %v, %alpha_raw, %beta_raw, %bias, %a_scale, %state_in, %dst, %snapshot_cache, %rmsnorm_gate, %scale, %snapshot_stride, %dst, %dst, %dst, %dst, %select_state, %state_row_count, %state_ids, %publish_q8) : (i1, i1, i1, index, f32, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, i1, f32, index, buffer, buffer, buffer, buffer, i1, index, buffer, i1) + kernel.return +} + +kernel.def export("llm_gated_delta_net_f32_wmma_head128_selected_snapshot_projection_rms_gate_q8") @llm_gated_delta_net_f32_wmma_head128_selected_snapshot_projection_rms_gate_q8() { + %unit = index.constant 1 : index + %wg = index.constant 256 : index + %blk = index.constant 128 : index + %blk_m = index.constant 127 : index + %s_v = config.get @llm.gated_delta_net.head_width : index + %n_heads = config.get @llm.gated_delta_net.head_count : index + %n_seqs = config.get @llm.gated_delta_net.sequence_count : index + %cr = index.add %s_v, %blk_m : index + %col_blocks = index.div %cr, %blk : index + kernel.launch.config workgroups(%n_heads, %n_seqs, %col_blocks) workgroup_size(%wg, %unit, %unit) : index +} launch(%q: buffer, %k: buffer, %v: buffer, %alpha_raw: buffer, %beta_raw: buffer, %bias: buffer, %a_scale: buffer, %state_in: buffer, %snapshot_cache: buffer, %dst: buffer, %state_ids: buffer, %rms_weight: buffer, %raw_gate: buffer, %q8_output: buffer) { + %state_inplace = scalar.constant false : i1 + %publish_snapshots = scalar.constant true : i1 + %projection_epilogue = scalar.constant true : i1 + %scale = scalar.constant 0.0883883461356163 : f32 + %snapshot_stride = config.get @llm.gated_delta_net.snapshot_stride : index + %rms_epsilon = config.get @ggml.rmsnorm_gate_f32.rms_epsilon : f32 + %rms_gate_op = config.get @ggml.rmsnorm_gate_f32.gate_op : index + %rmsnorm_gate = scalar.constant true : i1 + %select_state = scalar.constant true : i1 + %state_row_count = config.get @llm.gated_delta_net.state_row_count : index + %publish_q8 = scalar.constant true : i1 + template.apply<@llm.gated_delta_net.f32_wmma_head128.body>(%state_inplace, %publish_snapshots, %projection_epilogue, %snapshot_stride, %scale, %q, %k, %v, %alpha_raw, %beta_raw, %bias, %a_scale, %state_in, %dst, %snapshot_cache, %rmsnorm_gate, %rms_epsilon, %rms_gate_op, %rms_weight, %raw_gate, %dst, %q8_output, %select_state, %state_row_count, %state_ids, %publish_q8) : (i1, i1, i1, index, f32, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, i1, f32, index, buffer, buffer, buffer, buffer, i1, index, buffer, i1) + kernel.return +} + +template.def<@llm.gated_delta_net.f32_wmma_head128.body> device @llm_gated_delta_net_f32_wmma_head128_body(%state_inplace: i1, %publish_snapshots: i1, %projection_epilogue: i1, %snapshot_stride: index, %scale: f32, %q: buffer, %k: buffer, %v: buffer, %g: buffer, %beta: buffer, %bias: buffer, %a_scale: buffer, %state_in: buffer, %dst: buffer, %snapshot_cache: buffer, %rmsnorm_gate: i1, %rms_epsilon: f32, %rms_gate_op: index, %rms_weight: buffer, %raw_gate: buffer, %norm_output: buffer, %half_output: buffer, %select_state: i1, %state_row_count: index, %state_ids: buffer, %publish_q8: i1) { + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %k0 = index.constant 0 : index + %k1 = index.constant 1 : index + %k2 = index.constant 2 : index + %k3 = index.constant 3 : index + %k4 = index.constant 4 : index + %k5 = index.constant 5 : index + %k6 = index.constant 6 : index + %k7 = index.constant 7 : index + %k8 = index.constant 8 : index + %k9 = index.constant 9 : index + %k10 = index.constant 10 : index + %k11 = index.constant 11 : index + %k12 = index.constant 12 : index + %k13 = index.constant 13 : index + %k14 = index.constant 14 : index + %k15 = index.constant 15 : index + %k16 = index.constant 16 : index + %k17 = index.constant 17 : index + %k20 = index.constant 20 : index + %k24 = index.constant 24 : index + %k28 = index.constant 28 : index + %k32 = index.constant 32 : index + %k36 = index.constant 36 : index + %k40 = index.constant 40 : index + %k44 = index.constant 44 : index + %k48 = index.constant 48 : index + %k52 = index.constant 52 : index + %k56 = index.constant 56 : index + %k60 = index.constant 60 : index + %k61 = index.constant 61 : index + %k64 = index.constant 64 : index + %k65 = index.constant 65 : index + %k66 = index.constant 66 : index + %k67 = index.constant 67 : index + %k68 = index.constant 68 : index + %k69 = index.constant 69 : index + %k70 = index.constant 70 : index + %k71 = index.constant 71 : index + %k72 = index.constant 72 : index + %k73 = index.constant 73 : index + %k74 = index.constant 74 : index + %k75 = index.constant 75 : index + %k76 = index.constant 76 : index + %k77 = index.constant 77 : index + %k78 = index.constant 78 : index + %k79 = index.constant 79 : index + %k80 = index.constant 80 : index + %k88 = index.constant 88 : index + %k96 = index.constant 96 : index + %k104 = index.constant 104 : index + %k112 = index.constant 112 : index + %k120 = index.constant 120 : index + %k128 = index.constant 128 : index + %k136 = index.constant 136 : index + %k144 = index.constant 144 : index + %k152 = index.constant 152 : index + %k160 = index.constant 160 : index + %k168 = index.constant 168 : index + %k176 = index.constant 176 : index + %k184 = index.constant 184 : index + %k192 = index.constant 192 : index + %k200 = index.constant 200 : index + %k208 = index.constant 208 : index + %k216 = index.constant 216 : index + %k224 = index.constant 224 : index + %k232 = index.constant 232 : index + %k240 = index.constant 240 : index + %k248 = index.constant 248 : index + %k256 = index.constant 256 : index + %k272 = index.constant 272 : index + %k384 = index.constant 384 : index + %k408 = index.constant 408 : index + %k512 = index.constant 512 : index + %k544 = index.constant 544 : index + %k640 = index.constant 640 : index + %k680 = index.constant 680 : index + %k768 = index.constant 768 : index + %k816 = index.constant 816 : index + %k896 = index.constant 896 : index + %k952 = index.constant 952 : index + %k1024 = index.constant 1024 : index + %k1088 = index.constant 1088 : index + %k1152 = index.constant 1152 : index + %k1224 = index.constant 1224 : index + %k1280 = index.constant 1280 : index + %k1360 = index.constant 1360 : index + %k1408 = index.constant 1408 : index + %k1496 = index.constant 1496 : index + %k1536 = index.constant 1536 : index + %k1632 = index.constant 1632 : index + %k1664 = index.constant 1664 : index + %k1768 = index.constant 1768 : index + %k1792 = index.constant 1792 : index + %k1904 = index.constant 1904 : index + %k1920 = index.constant 1920 : index + %k2040 = index.constant 2040 : index + %k4352 = index.constant 4352 : index + %k5440 = index.constant 5440 : index + %k6528 = index.constant 6528 : index + %k11072 = index.constant 11072 : index + + %s_v0 = config.get @llm.gated_delta_net.head_width : index + %n_heads0 = config.get @llm.gated_delta_net.head_count : index + %n_tokens0 = config.get @llm.gated_delta_net.token_count : index + %n_seqs = config.get @llm.gated_delta_net.sequence_count : index + %sq1 = config.get @llm.gated_delta_net.qk_stride1 : index + %sq2 = config.get @llm.gated_delta_net.qk_stride2 : index + %sq3 = config.get @llm.gated_delta_net.qk_stride3 : index + %sv1 = config.get @llm.gated_delta_net.value_stride1 : index + %sv2 = config.get @llm.gated_delta_net.value_stride2 : index + %sv3 = config.get @llm.gated_delta_net.value_stride3 : index + %sb1 = config.get @llm.gated_delta_net.scalar_stride1 : index + %sb2 = config.get @llm.gated_delta_net.scalar_stride2 : index + %sb3 = config.get @llm.gated_delta_net.scalar_stride3 : index + %neqk1 = config.get @llm.gated_delta_net.query_head_count : index + %rq3 = config.get @llm.gated_delta_net.query_sequence_ratio : index + %s_v = index.assume %s_v0 [range(%s_v0, 128, 128)] : index + %n_heads = index.assume %n_heads0 [range(%n_heads0, 1, 4096)] : index + %n_tokens = index.assume %n_tokens0 [range(%n_tokens0, 1, 1048576)] : index + + %recurrence_true = scalar.constant true : i1 + %recurrence_wide = index.cmp ugt, %n_tokens, %k5 : index + %recurrence_other = scalar.xori %publish_q8, %recurrence_true : i1 + %all_recurrence_rows = scalar.ori %recurrence_wide, %recurrence_other : i1 + + %h_idx0 = kernel.workgroup.id : index + %h_idx = index.assume %h_idx0 [range(%h_idx0, 0, 4095)] : index + %seq0 = kernel.workgroup.id : index + %seq = index.assume %seq0 [range(%seq0, 0, 4095)] : index + %colblk0 = kernel.workgroup.id : index + %colblk = index.assume %colblk0 [range(%colblk0, 0, 1023)] : index + %tid0 = kernel.workitem.id : index + %tid = index.assume %tid0 [range(%tid0, 0, 255)] : index + %lane0 = kernel.subgroup.lane.id : index + %lane = index.assume %lane0 [range(%lane0, 0, 31)] : index + %wave0 = kernel.subgroup.id : index + %wave = index.assume %wave0 [range(%wave0, 0, 7)] : index + %is_w0 = index.cmp eq, %wave, %c0 : index + %is_w1 = index.cmp eq, %wave, %c1 : index + %lane_live = index.cmp ult, %lane, %k16 : index + %lane_lo = index.cmp ult, %lane, %k16 : index + + %bb = index.mul %colblk, %k128 : index + %wo = index.mul %wave, %k16 : index + %waveoff0 = index.assume %wo [range(%wo, 0, 112)] : index + %wcol0 = index.add %bb, %wo : index + %wcol = index.assume %wcol0 [range(%wcol0, 0, 127)] : index + + %iq1 = index.rem %h_idx, %neqk1 : index + %iq3 = index.div %seq, %rq3 : index + + %q_g = buffer.assume.memory_space %q : buffer + %k_g = buffer.assume.memory_space %k : buffer + %v_g = buffer.assume.memory_space %v : buffer + %g_g = buffer.assume.memory_space %g : buffer + %b_g = buffer.assume.memory_space %beta : buffer + %bias_g = buffer.assume.memory_space %bias : buffer + %a_scale_g = buffer.assume.memory_space %a_scale : buffer + %si_g = buffer.assume.memory_space %state_in : buffer + %d_g = buffer.assume.memory_space %dst : buffer + %snapshot_cache_g = buffer.assume.memory_space %snapshot_cache : buffer + %q_view = buffer.view %q_g[%base] : buffer -> view<1073741824xf32> + %k_view = buffer.view %k_g[%base] : buffer -> view<1073741824xf32> + %v_view = buffer.view %v_g[%base] : buffer -> view<1073741824xf32> + %g_view = buffer.view %g_g[%base] : buffer -> view<1073741824xf32> + %b_view = buffer.view %b_g[%base] : buffer -> view<1073741824xf32> + %bias_view = buffer.view %bias_g[%base] : buffer -> view<4096xf32> + %a_scale_view = buffer.view %a_scale_g[%base] : buffer -> view<4096xf32> + %si_view = buffer.view %si_g[%base] : buffer -> view<1073741824xf32> + %dst_view = buffer.view %d_g[%base] : buffer -> view<1073741824xf32> + %snapshot_cache_view = buffer.view %snapshot_cache_g[%base] : buffer -> view<1073741824xf32> + + %ae0 = index.mul %s_v, %n_heads : index + %ae1 = index.mul %ae0, %n_tokens : index + %attn_elems = index.mul %ae1, %n_seqs : index + %sv_sq = index.mul %s_v, %s_v : index + %so0 = index.mul %seq, %n_heads : index + %so1 = index.add %so0, %h_idx : index + %state_offset = index.mul %so1, %sv_sq : index + %selected_state_offset, %selected_valid = scf.if %select_state -> (index, i1) { + %selected_zero_i32 = scalar.constant 0 : i32 + %selected_limit_i32 = index.cast %state_row_count : index to i32 + %selected_ids_view = buffer.view %state_ids[%base] : buffer -> view<[%n_seqs]xi32> + %selected_raw = view.load %selected_ids_view[%seq] : view<[%n_seqs]xi32> -> i32 + %selected_nonnegative = scalar.cmpi sge, %selected_raw, %selected_zero_i32 : i32 + %selected_in_range = scalar.cmpi slt, %selected_raw, %selected_limit_i32 : i32 + %selected_ok = scalar.andi %selected_nonnegative, %selected_in_range : i1 + %selected_safe_i32 = scf.select %selected_ok, %selected_raw, %selected_zero_i32 : i32 + %selected_id0 = index.cast %selected_safe_i32 : i32 to index + %selected_id = index.assume %selected_id0 [range(%selected_id0, 0, 262207), lt(%selected_id0, %state_row_count)] : index + %selected_head0 = index.mul %selected_id, %n_heads : index + %selected_head = index.add %selected_head0, %h_idx : index + %selected_offset = index.mul %selected_head, %sv_sq : index + scf.yield %selected_offset, %selected_ok : index, i1 + } else { + %selected_always_valid = scalar.constant true : i1 + scf.yield %state_offset, %selected_always_valid : index, i1 + } + %selected_zero_state = vector.constant 0.0 : vector<4xf32> + %ab0 = index.mul %seq, %n_tokens : index + %ab1 = index.mul %ab0, %n_heads : index + %ab2 = index.add %ab1, %h_idx : index + %attn_base = index.mul %ab2, %s_v : index + %attn_stride = index.mul %s_v, %n_heads : index + %qk_s = index.mul %iq3, %sq3 : index + %qk_h = index.mul %iq1, %sq1 : index + %qk_base = index.add %qk_s, %qk_h : index + %v_s = index.mul %seq, %sv3 : index + %v_h = index.mul %h_idx, %sv1 : index + %v_head_base = index.add %v_s, %v_h : index + %gb_s = index.mul %seq, %sb3 : index + %gb_h = index.mul %h_idx, %sb1 : index + %gb_r = index.add %gb_s, %gb_h : index + %gb_base = index.assume %gb_r [range(%gb_r, 0, 1073741823)] : index + + %lds_bytes = index.constant 63296 : offset + %lds = buffer.alloca align(16) %lds_bytes : buffer + %km_o = index.constant 0 : offset + %qm_o = index.constant 4352 : offset + %ka_o = index.constant 4352 : offset + %gr_o = index.constant 8704 : offset + %qkt_o = index.constant 9728 : offset + %lay_tokrow = encoding.layout.strided [136, 1] : encoding + %lay_rowtok = encoding.layout.strided [1, 136] : encoding + %lay_tile = encoding.layout.strided [17, 1] : encoding + %lay_gram = encoding.layout.strided [16, 1] : encoding + %lay_colrow = encoding.layout.strided [136, 1] : encoding + %lay_dcol = encoding.layout.strided [24, 1] : encoding + %lay_katok = encoding.layout.strided [1, 24] : encoding + %lay_upd = encoding.layout.strided [68, 1] : encoding + %km_flat = buffer.view %lds[%km_o] : buffer -> view<2176xf16> + %qm_flat = buffer.view %lds[%qm_o] : buffer -> view<2176xf16> + %ka_flat = buffer.view %lds[%ka_o] : buffer -> view<3072xf16> + %gr_flat = buffer.view %lds[%gr_o] : buffer -> view<256xf32> + %qkt_flat = buffer.view %lds[%qkt_o] : buffer -> view<256xf32> + %km_lhs = buffer.view %lds[%km_o] : buffer -> view<16x128xf16, %lay_tokrow> + %km_rhs = buffer.view %lds[%km_o] : buffer -> view<128x16xf16, %lay_rowtok> + %qm_rhs = buffer.view %lds[%qm_o] : buffer -> view<128x16xf16, %lay_rowtok> + %ka_rhs = buffer.view %lds[%ka_o] : buffer -> view<16x128xf16, %lay_katok> + %gr_res = buffer.view %lds[%gr_o] : buffer -> view<16x16xf32, %lay_gram> + %qkt_res = buffer.view %lds[%qkt_o] : buffer -> view<16x16xf32, %lay_gram> + + %wb0 = index.mul %wave, %k6528 : index + %wb1 = index.add %wb0, %k11072 : index + %wb = index.cast %wb1 : index to offset + %ab1o = index.add %wb1, %k4352 : index + %abo = index.cast %ab1o : index to offset + %bb1o = index.add %wb1, %k5440 : index + %bbo = index.cast %bb1o : index to offset + %dd1o = index.add %wb1, %k4352 : index + %ddo = index.cast %dd1o : index to offset + %sh_flat = buffer.view %lds[%wb] : buffer -> view<2176xf16> + %sh_lhs = buffer.view %lds[%wb] : buffer -> view<16x128xf16, %lay_colrow> + %upd_flat = buffer.view %lds[%wb] : buffer -> view<1088xf32> + %upd_tile = buffer.view %lds[%wb] : buffer -> view<16x68xf32, %lay_upd> + %a_flat = buffer.view %lds[%abo] : buffer -> view<272xf32> + %a_res = buffer.view %lds[%abo] : buffer -> view<16x16xf32, %lay_tile> + %b_flat = buffer.view %lds[%bbo] : buffer -> view<272xf32> + %b_res = buffer.view %lds[%bbo] : buffer -> view<16x16xf32, %lay_tile> + %d_flat = buffer.view %lds[%ddo] : buffer -> view<384xf16> + %d_lhs = buffer.view %lds[%ddo] : buffer -> view<16x16xf16, %lay_dcol> + // V and attention output share LDS slots; each recurrence reads V before overwriting its dead slot. + // The full-head plane overlays dead K/Q storage after the matrix products. + %vst_flat = buffer.view %lds[%km_o] : buffer -> view<2048xf32> + %aost_flat = buffer.view %lds[%km_o] : buffer -> view<2048xf32> + %vs_tok0 = index.div %tid, %k128 : index + %vs_tok = index.assume %vs_tok0 [range(%vs_tok0, 0, 1)] : index + %vs_col0 = index.rem %tid, %k128 : index + %vs_col = index.assume %vs_col0 [range(%vs_col0, 0, 127)] : index + %vs_gcol0 = index.add %bb, %vs_col : index + %vs_gcol = index.assume %vs_gcol0 [range(%vs_gcol0, 0, 127)] : index + %gs_o = index.constant 10752 : offset + %gs_flat = buffer.view %lds[%gs_o] : buffer -> view<80xf32> + + // Low-row projection epilogue. Distribute the valid token chains across + // waves and issue them before the state-plane loads below so the scheduler + // can cover their transcendental latency with independent global reads. + scf.if %projection_epilogue { + %projection_zero = scalar.constant 0.0 : f32 + %projection_one = scalar.constant 1.0 : f32 + %projection_worker = index.cmp ult, %lane, %k2 : index + %projection_lane_offset = index.mul %lane, %k8 : index + %projection_slot_r = index.add %projection_lane_offset, %wave : index + %projection_slot = index.assume %projection_slot_r [range(%projection_slot_r, 0, 15)] : index + %projection_token_ok = index.cmp ult, %projection_slot, %n_tokens : index + scf.if %projection_worker { + %gate_value, %beta_value = scf.if %projection_token_ok -> (f32, f32) { + %projection_token = index.assume %projection_slot [range(%projection_slot, 0, 1048575)] : index + %gbi_m = index.mul %projection_token, %sb2 : index + %gbi_a = index.add %gb_base, %gbi_m : index + %gbi = index.assume %gbi_a [range(%gbi_a, 0, 1073741823)] : index + %graw = view.load %g_view[%gbi] : view<1073741824xf32> -> f32 + %braw = view.load %b_view[%gbi] : view<1073741824xf32> -> f32 + %bias_value = view.load %bias_view[%h_idx] : view<4096xf32> -> f32 + %a_value = view.load %a_scale_view[%h_idx] : view<4096xf32> -> f32 + %biased = scalar.addf %graw, %bias_value : f32 + %abs = scalar.absf %biased : f32 + %neg_abs = scalar.negf %abs : f32 + %exp = scalar.expf %neg_abs : f32 + %one_plus = scalar.addf %projection_one, %exp : f32 + %log = scalar.logf %one_plus : f32 + %positive = scalar.maxnumf %biased, %projection_zero : f32 + %softplus = scalar.addf %positive, %log : f32 + %gate_result = scalar.mulf %softplus, %a_value : f32 + %beta_result = scalar.logisticf %braw : f32 + scf.yield %gate_result, %beta_result : f32, f32 + } else { + scf.yield %projection_zero, %projection_zero : f32, f32 + } + %gsi = index.add %projection_slot, %k64 : index + %gsib = index.assume %gsi [range(%gsi, 64, 79)] : index + view.store %gate_value, %gs_flat[%gsib] : f32, view<80xf32> + %bsi_m = index.mul %projection_slot, %k4 : index + %bsib = index.assume %bsi_m [range(%bsi_m, 0, 60)] : index + view.store %beta_value, %gs_flat[%bsib] : f32, view<80xf32> + } + } + %projection_true = scalar.constant true : i1 + %load_processed_scalars = scalar.xori %projection_epilogue, %projection_true : i1 + + // Staging map: 16 consecutive lanes cover one token's 128 rows, 8 rows each, + // so a pass reads whole 512 B token rows and the wave's accesses coalesce. + %tid_lt_c = index.cmp ult, %tid, %k16 : index + %stokl_r = index.div %tid, %k16 : index + %stokl = index.assume %stokl_r [range(%stokl_r, 0, 15)] : index + %srlane = index.rem %tid, %k16 : index + %srow0_m = index.mul %srlane, %k8 : index + %srow0 = index.assume %srow0_m [range(%srow0_m, 0, 120)] : index + + // Lane L owns rows 4L..4L+3 of every column: one f32x4 per column, and the + // update's two 64-row halves split exactly on the lane-16 boundary. + %row0_m = index.mul %lane, %k4 : index + %row0 = index.assume %row0_m [range(%row0_m, 0, 124)] : index + %lc_s = scf.select %lane_live, %lane, %c0 : index + %lc = index.assume %lc_s [range(%lc_s, 0, 15)] : index + %rec_col0 = index.add %waveoff0, %lc : index + %rec_col = index.assume %rec_col0 [range(%rec_col0, 0, 127)] : index + %gcol_r = index.add %wcol, %lc : index + %gcol = index.assume %gcol_r [range(%gcol_r, 0, 127)] : index + %tile_row_m = index.mul %lc, %k17 : index + %tile_row = index.assume %tile_row_m [range(%tile_row_m, 0, 255)] : index + %d_row_m = index.mul %lc, %k24 : index + %d_row = index.assume %d_row_m [range(%d_row_m, 0, 360)] : index + %ic0_r = index.add %wcol, %k0 : index + %ic0 = index.assume %ic0_r [range(%ic0_r, 0, 127)] : index + %icb0_m = index.mul %ic0, %s_v : index + %icb0_a = index.add %state_offset, %icb0_m : index + %icb0_r = index.add %icb0_a, %row0 : index + %icb0 = index.assume %icb0_r [range(%icb0_r, 0, 1073741820)] : index + %selected_icb0_local = index.sub %icb0, %state_offset : index + %selected_icb0 = index.add %selected_icb0_local, %selected_state_offset : index + %selected_s0 = vector.load %si_view[%selected_icb0] : view<1073741824xf32> -> vector<4xf32> + %s0_init = scf.select %selected_valid, %selected_s0, %selected_zero_state : vector<4xf32> + %icb1_r = index.add %icb0, %k128 : index + %icb1 = index.assume %icb1_r [range(%icb1_r, 0, 1073741820)] : index + %selected_icb1_local = index.sub %icb1, %state_offset : index + %selected_icb1 = index.add %selected_icb1_local, %selected_state_offset : index + %selected_s1 = vector.load %si_view[%selected_icb1] : view<1073741824xf32> -> vector<4xf32> + %s1_init = scf.select %selected_valid, %selected_s1, %selected_zero_state : vector<4xf32> + %icb2_r = index.add %icb1, %k128 : index + %icb2 = index.assume %icb2_r [range(%icb2_r, 0, 1073741820)] : index + %selected_icb2_local = index.sub %icb2, %state_offset : index + %selected_icb2 = index.add %selected_icb2_local, %selected_state_offset : index + %selected_s2 = vector.load %si_view[%selected_icb2] : view<1073741824xf32> -> vector<4xf32> + %s2_init = scf.select %selected_valid, %selected_s2, %selected_zero_state : vector<4xf32> + %icb3_r = index.add %icb2, %k128 : index + %icb3 = index.assume %icb3_r [range(%icb3_r, 0, 1073741820)] : index + %selected_icb3_local = index.sub %icb3, %state_offset : index + %selected_icb3 = index.add %selected_icb3_local, %selected_state_offset : index + %selected_s3 = vector.load %si_view[%selected_icb3] : view<1073741824xf32> -> vector<4xf32> + %s3_init = scf.select %selected_valid, %selected_s3, %selected_zero_state : vector<4xf32> + %icb4_r = index.add %icb3, %k128 : index + %icb4 = index.assume %icb4_r [range(%icb4_r, 0, 1073741820)] : index + %selected_icb4_local = index.sub %icb4, %state_offset : index + %selected_icb4 = index.add %selected_icb4_local, %selected_state_offset : index + %selected_s4 = vector.load %si_view[%selected_icb4] : view<1073741824xf32> -> vector<4xf32> + %s4_init = scf.select %selected_valid, %selected_s4, %selected_zero_state : vector<4xf32> + %icb5_r = index.add %icb4, %k128 : index + %icb5 = index.assume %icb5_r [range(%icb5_r, 0, 1073741820)] : index + %selected_icb5_local = index.sub %icb5, %state_offset : index + %selected_icb5 = index.add %selected_icb5_local, %selected_state_offset : index + %selected_s5 = vector.load %si_view[%selected_icb5] : view<1073741824xf32> -> vector<4xf32> + %s5_init = scf.select %selected_valid, %selected_s5, %selected_zero_state : vector<4xf32> + %icb6_r = index.add %icb5, %k128 : index + %icb6 = index.assume %icb6_r [range(%icb6_r, 0, 1073741820)] : index + %selected_icb6_local = index.sub %icb6, %state_offset : index + %selected_icb6 = index.add %selected_icb6_local, %selected_state_offset : index + %selected_s6 = vector.load %si_view[%selected_icb6] : view<1073741824xf32> -> vector<4xf32> + %s6_init = scf.select %selected_valid, %selected_s6, %selected_zero_state : vector<4xf32> + %icb7_r = index.add %icb6, %k128 : index + %icb7 = index.assume %icb7_r [range(%icb7_r, 0, 1073741820)] : index + %selected_icb7_local = index.sub %icb7, %state_offset : index + %selected_icb7 = index.add %selected_icb7_local, %selected_state_offset : index + %selected_s7 = vector.load %si_view[%selected_icb7] : view<1073741824xf32> -> vector<4xf32> + %s7_init = scf.select %selected_valid, %selected_s7, %selected_zero_state : vector<4xf32> + %icb8_r = index.add %icb7, %k128 : index + %icb8 = index.assume %icb8_r [range(%icb8_r, 0, 1073741820)] : index + %selected_icb8_local = index.sub %icb8, %state_offset : index + %selected_icb8 = index.add %selected_icb8_local, %selected_state_offset : index + %selected_s8 = vector.load %si_view[%selected_icb8] : view<1073741824xf32> -> vector<4xf32> + %s8_init = scf.select %selected_valid, %selected_s8, %selected_zero_state : vector<4xf32> + %icb9_r = index.add %icb8, %k128 : index + %icb9 = index.assume %icb9_r [range(%icb9_r, 0, 1073741820)] : index + %selected_icb9_local = index.sub %icb9, %state_offset : index + %selected_icb9 = index.add %selected_icb9_local, %selected_state_offset : index + %selected_s9 = vector.load %si_view[%selected_icb9] : view<1073741824xf32> -> vector<4xf32> + %s9_init = scf.select %selected_valid, %selected_s9, %selected_zero_state : vector<4xf32> + %icb10_r = index.add %icb9, %k128 : index + %icb10 = index.assume %icb10_r [range(%icb10_r, 0, 1073741820)] : index + %selected_icb10_local = index.sub %icb10, %state_offset : index + %selected_icb10 = index.add %selected_icb10_local, %selected_state_offset : index + %selected_s10 = vector.load %si_view[%selected_icb10] : view<1073741824xf32> -> vector<4xf32> + %s10_init = scf.select %selected_valid, %selected_s10, %selected_zero_state : vector<4xf32> + %icb11_r = index.add %icb10, %k128 : index + %icb11 = index.assume %icb11_r [range(%icb11_r, 0, 1073741820)] : index + %selected_icb11_local = index.sub %icb11, %state_offset : index + %selected_icb11 = index.add %selected_icb11_local, %selected_state_offset : index + %selected_s11 = vector.load %si_view[%selected_icb11] : view<1073741824xf32> -> vector<4xf32> + %s11_init = scf.select %selected_valid, %selected_s11, %selected_zero_state : vector<4xf32> + %icb12_r = index.add %icb11, %k128 : index + %icb12 = index.assume %icb12_r [range(%icb12_r, 0, 1073741820)] : index + %selected_icb12_local = index.sub %icb12, %state_offset : index + %selected_icb12 = index.add %selected_icb12_local, %selected_state_offset : index + %selected_s12 = vector.load %si_view[%selected_icb12] : view<1073741824xf32> -> vector<4xf32> + %s12_init = scf.select %selected_valid, %selected_s12, %selected_zero_state : vector<4xf32> + %icb13_r = index.add %icb12, %k128 : index + %icb13 = index.assume %icb13_r [range(%icb13_r, 0, 1073741820)] : index + %selected_icb13_local = index.sub %icb13, %state_offset : index + %selected_icb13 = index.add %selected_icb13_local, %selected_state_offset : index + %selected_s13 = vector.load %si_view[%selected_icb13] : view<1073741824xf32> -> vector<4xf32> + %s13_init = scf.select %selected_valid, %selected_s13, %selected_zero_state : vector<4xf32> + %icb14_r = index.add %icb13, %k128 : index + %icb14 = index.assume %icb14_r [range(%icb14_r, 0, 1073741820)] : index + %selected_icb14_local = index.sub %icb14, %state_offset : index + %selected_icb14 = index.add %selected_icb14_local, %selected_state_offset : index + %selected_s14 = vector.load %si_view[%selected_icb14] : view<1073741824xf32> -> vector<4xf32> + %s14_init = scf.select %selected_valid, %selected_s14, %selected_zero_state : vector<4xf32> + %icb15_r = index.add %icb14, %k128 : index + %icb15 = index.assume %icb15_r [range(%icb15_r, 0, 1073741820)] : index + %selected_icb15_local = index.sub %icb15, %state_offset : index + %selected_icb15 = index.add %selected_icb15_local, %selected_state_offset : index + %selected_s15 = vector.load %si_view[%selected_icb15] : view<1073741824xf32> -> vector<4xf32> + %s15_init = scf.select %selected_valid, %selected_s15, %selected_zero_state : vector<4xf32> + %ch_r = index.add %n_tokens, %k15 : index + %n_chunks = index.div %ch_r, %k16 : index + %f1 = scalar.constant 1.0 : f32 + %f0 = scalar.constant 0.0 : f32 + %norm_eps = config.get @llm.gated_delta_net.l2_epsilon : f32 + %norm_eps_sq = scalar.mulf %norm_eps, %norm_eps : f32 + %fz8 = vector.constant 0.0 : vector<8xf32> + %zv8 = vector.constant 0.0 : vector<8xf32> + %fzh8 = vector.constant 0.0 : vector<8xf16> + %five_snapshot_tokens = index.cmp eq, %n_tokens, %k5 : index + %reuse_snapshot_fragments = scalar.andi %publish_snapshots, %five_snapshot_tokens : i1 + %sf0, %sf1, %sf2, %sf3, %sf4, %sf5, %sf6, %sf7, %sf8, %sf9, %sf10, %sf11, %sf12, %sf13, %sf14, %sf15 = scf.for %ch = [%c0 to %n_chunks step %c1](%sa0 = %s0_init : vector<4xf32>, %sa1 = %s1_init : vector<4xf32>, %sa2 = %s2_init : vector<4xf32>, %sa3 = %s3_init : vector<4xf32>, %sa4 = %s4_init : vector<4xf32>, %sa5 = %s5_init : vector<4xf32>, %sa6 = %s6_init : vector<4xf32>, %sa7 = %s7_init : vector<4xf32>, %sa8 = %s8_init : vector<4xf32>, %sa9 = %s9_init : vector<4xf32>, %sa10 = %s10_init : vector<4xf32>, %sa11 = %s11_init : vector<4xf32>, %sa12 = %s12_init : vector<4xf32>, %sa13 = %s13_init : vector<4xf32>, %sa14 = %s14_init : vector<4xf32>, %sa15 = %s15_init : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + %t0 = index.mul %ch, %k16 : index + // Compute gate prefix and suffix decays once in log space to avoid product underflow. + scf.if %load_processed_scalars { + // Masked tokens use g=1 and beta=0; threads 0..15 each own one token. + %gtok_r = index.add %t0, %tid : index + %gtok_ok = index.cmp ult, %gtok_r, %n_tokens : index + %gtok_s = scf.select %gtok_ok, %gtok_r, %c0 : index + %gtok = index.assume %gtok_s [range(%gtok_s, 0, 1048575)] : index + scf.if %tid_lt_c { + %gbi_m = index.mul %gtok, %sb2 : index + %gbi_a = index.add %gb_base, %gbi_m : index + %gbi = index.assume %gbi_a [range(%gbi_a, 0, 1073741823)] : index + %graw = view.load %g_view[%gbi] : view<1073741824xf32> -> f32 + %braw = view.load %b_view[%gbi] : view<1073741824xf32> -> f32 + %gvv = scf.select %gtok_ok, %graw, %f0 : f32 + %bvv = scf.select %gtok_ok, %braw, %f0 : f32 + %gsi = index.add %tid, %k64 : index + %gsib = index.assume %gsi [range(%gsi, 0, 79)] : index + view.store %gvv, %gs_flat[%gsib] : f32, view<80xf32> + %bsi_m = index.mul %tid, %k4 : index + %bsib = index.assume %bsi_m [range(%bsi_m, 0, 79)] : index + view.store %bvv, %gs_flat[%bsib] : f32, view<80xf32> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + // c_t, the suffix decay, and log(c_t). Every lane reads all 16 + // log-gates and forms prefix and suffix sums without a cross-lane scan. + scf.if %tid_lt_c { + %pg0 = view.load %gs_flat[%k64] : view<80xf32> -> f32 + %pg1 = view.load %gs_flat[%k65] : view<80xf32> -> f32 + %pg2 = view.load %gs_flat[%k66] : view<80xf32> -> f32 + %pg3 = view.load %gs_flat[%k67] : view<80xf32> -> f32 + %pg4 = view.load %gs_flat[%k68] : view<80xf32> -> f32 + %pg5 = view.load %gs_flat[%k69] : view<80xf32> -> f32 + %pg6 = view.load %gs_flat[%k70] : view<80xf32> -> f32 + %pg7 = view.load %gs_flat[%k71] : view<80xf32> -> f32 + %pg8 = view.load %gs_flat[%k72] : view<80xf32> -> f32 + %pg9 = view.load %gs_flat[%k73] : view<80xf32> -> f32 + %pg10 = view.load %gs_flat[%k74] : view<80xf32> -> f32 + %pg11 = view.load %gs_flat[%k75] : view<80xf32> -> f32 + %pg12 = view.load %gs_flat[%k76] : view<80xf32> -> f32 + %pg13 = view.load %gs_flat[%k77] : view<80xf32> -> f32 + %pg14 = view.load %gs_flat[%k78] : view<80xf32> -> f32 + %pg15 = view.load %gs_flat[%k79] : view<80xf32> -> f32 + %ple0 = index.cmp ule, %k0, %tid : index + %pm0 = scf.select %ple0, %pg0, %f0 : f32 + %sm0 = scf.select %ple0, %f0, %pg0 : f32 + %ple1 = index.cmp ule, %k1, %tid : index + %pm1 = scf.select %ple1, %pg1, %f0 : f32 + %sm1 = scf.select %ple1, %f0, %pg1 : f32 + %pp1 = scalar.addf %pm0, %pm1 : f32 + %ss1 = scalar.addf %sm0, %sm1 : f32 + %ple2 = index.cmp ule, %k2, %tid : index + %pm2 = scf.select %ple2, %pg2, %f0 : f32 + %sm2 = scf.select %ple2, %f0, %pg2 : f32 + %pp2 = scalar.addf %pp1, %pm2 : f32 + %ss2 = scalar.addf %ss1, %sm2 : f32 + %ple3 = index.cmp ule, %k3, %tid : index + %pm3 = scf.select %ple3, %pg3, %f0 : f32 + %sm3 = scf.select %ple3, %f0, %pg3 : f32 + %pp3 = scalar.addf %pp2, %pm3 : f32 + %ss3 = scalar.addf %ss2, %sm3 : f32 + %ple4 = index.cmp ule, %k4, %tid : index + %pm4 = scf.select %ple4, %pg4, %f0 : f32 + %sm4 = scf.select %ple4, %f0, %pg4 : f32 + %pp4 = scalar.addf %pp3, %pm4 : f32 + %ss4 = scalar.addf %ss3, %sm4 : f32 + %ple5 = index.cmp ule, %k5, %tid : index + %pm5 = scf.select %ple5, %pg5, %f0 : f32 + %sm5 = scf.select %ple5, %f0, %pg5 : f32 + %pp5 = scalar.addf %pp4, %pm5 : f32 + %ss5 = scalar.addf %ss4, %sm5 : f32 + %ple6 = index.cmp ule, %k6, %tid : index + %pm6 = scf.select %ple6, %pg6, %f0 : f32 + %sm6 = scf.select %ple6, %f0, %pg6 : f32 + %pp6 = scalar.addf %pp5, %pm6 : f32 + %ss6 = scalar.addf %ss5, %sm6 : f32 + %ple7 = index.cmp ule, %k7, %tid : index + %pm7 = scf.select %ple7, %pg7, %f0 : f32 + %sm7 = scf.select %ple7, %f0, %pg7 : f32 + %pp7 = scalar.addf %pp6, %pm7 : f32 + %ss7 = scalar.addf %ss6, %sm7 : f32 + %ple8 = index.cmp ule, %k8, %tid : index + %pm8 = scf.select %ple8, %pg8, %f0 : f32 + %sm8 = scf.select %ple8, %f0, %pg8 : f32 + %pp8 = scalar.addf %pp7, %pm8 : f32 + %ss8 = scalar.addf %ss7, %sm8 : f32 + %ple9 = index.cmp ule, %k9, %tid : index + %pm9 = scf.select %ple9, %pg9, %f0 : f32 + %sm9 = scf.select %ple9, %f0, %pg9 : f32 + %pp9 = scalar.addf %pp8, %pm9 : f32 + %ss9 = scalar.addf %ss8, %sm9 : f32 + %ple10 = index.cmp ule, %k10, %tid : index + %pm10 = scf.select %ple10, %pg10, %f0 : f32 + %sm10 = scf.select %ple10, %f0, %pg10 : f32 + %pp10 = scalar.addf %pp9, %pm10 : f32 + %ss10 = scalar.addf %ss9, %sm10 : f32 + %ple11 = index.cmp ule, %k11, %tid : index + %pm11 = scf.select %ple11, %pg11, %f0 : f32 + %sm11 = scf.select %ple11, %f0, %pg11 : f32 + %pp11 = scalar.addf %pp10, %pm11 : f32 + %ss11 = scalar.addf %ss10, %sm11 : f32 + %ple12 = index.cmp ule, %k12, %tid : index + %pm12 = scf.select %ple12, %pg12, %f0 : f32 + %sm12 = scf.select %ple12, %f0, %pg12 : f32 + %pp12 = scalar.addf %pp11, %pm12 : f32 + %ss12 = scalar.addf %ss11, %sm12 : f32 + %ple13 = index.cmp ule, %k13, %tid : index + %pm13 = scf.select %ple13, %pg13, %f0 : f32 + %sm13 = scf.select %ple13, %f0, %pg13 : f32 + %pp13 = scalar.addf %pp12, %pm13 : f32 + %ss13 = scalar.addf %ss12, %sm13 : f32 + %ple14 = index.cmp ule, %k14, %tid : index + %pm14 = scf.select %ple14, %pg14, %f0 : f32 + %sm14 = scf.select %ple14, %f0, %pg14 : f32 + %pp14 = scalar.addf %pp13, %pm14 : f32 + %ss14 = scalar.addf %ss13, %sm14 : f32 + %ple15 = index.cmp ule, %k15, %tid : index + %pm15 = scf.select %ple15, %pg15, %f0 : f32 + %sm15 = scf.select %ple15, %f0, %pg15 : f32 + %pp15 = scalar.addf %pp14, %pm15 : f32 + %ss15 = scalar.addf %ss14, %sm15 : f32 + %cpv = scalar.expf %pp15 : f32 + %pev = scalar.expf %ss15 : f32 + %quad = index.mul %tid, %k4 : index + %cpi = index.add %quad, %k1 : index + %cpib = index.assume %cpi [range(%cpi, 0, 79)] : index + view.store %cpv, %gs_flat[%cpib] : f32, view<80xf32> + %evi = index.add %quad, %k2 : index + %evib = index.assume %evi [range(%evi, 0, 79)] : index + view.store %pev, %gs_flat[%evib] : f32, view<80xf32> + %ivi = index.add %quad, %k3 : index + %ivib = index.assume %ivi [range(%ivi, 0, 79)] : index + view.store %pp15, %gs_flat[%ivib] : f32, view<80xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + + // Normalize K/Q into F16 LDS while four-lane groups reproduce the exact wave32 reduction DAG. + %stk0_r = index.add %t0, %k0 : index + %stk0_s = index.add %stk0_r, %stokl : index + %stk0_ok = index.cmp ult, %stk0_s, %n_tokens : index + %stk0_c = scf.select %stk0_ok, %stk0_s, %c0 : index + %stk0 = index.assume %stk0_c [range(%stk0_c, 0, 1048575)] : index + %sgi0_m = index.mul %stk0, %sq2 : index + %sgi0_a = index.add %qk_base, %sgi0_m : index + %sgi0_b = index.add %sgi0_a, %srow0 : index + %sgi0 = index.assume %sgi0_b [range(%sgi0_b, 0, 1073741816)] : index + %kvr0 = vector.load %k_view[%sgi0] : view<1073741824xf32> -> vector<8xf32> + %kvs0 = scf.select %stk0_ok, %kvr0, %zv8 : vector<8xf32> + %n0_k_v0 = vector.extract %kvs0[0] : vector<8xf32> -> f32 + %n0_k_sq0 = scalar.mulf %n0_k_v0, %n0_k_v0 : f32 + %n0_k_v1 = vector.extract %kvs0[1] : vector<8xf32> -> f32 + %n0_k_sq1 = scalar.mulf %n0_k_v1, %n0_k_v1 : f32 + %n0_k_v2 = vector.extract %kvs0[2] : vector<8xf32> -> f32 + %n0_k_sq2 = scalar.mulf %n0_k_v2, %n0_k_v2 : f32 + %n0_k_v3 = vector.extract %kvs0[3] : vector<8xf32> -> f32 + %n0_k_sq3 = scalar.mulf %n0_k_v3, %n0_k_v3 : f32 + %n0_k_v4 = vector.extract %kvs0[4] : vector<8xf32> -> f32 + %n0_k_sq4 = scalar.mulf %n0_k_v4, %n0_k_v4 : f32 + %n0_k_v5 = vector.extract %kvs0[5] : vector<8xf32> -> f32 + %n0_k_sq5 = scalar.mulf %n0_k_v5, %n0_k_v5 : f32 + %n0_k_v6 = vector.extract %kvs0[6] : vector<8xf32> -> f32 + %n0_k_sq6 = scalar.mulf %n0_k_v6, %n0_k_v6 : f32 + %n0_k_v7 = vector.extract %kvs0[7] : vector<8xf32> -> f32 + %n0_k_sq7 = scalar.mulf %n0_k_v7, %n0_k_v7 : f32 + %n0_k_p01 = scalar.addf %n0_k_sq0, %n0_k_sq1 : f32 + %n0_k_p23 = scalar.addf %n0_k_sq2, %n0_k_sq3 : f32 + %n0_k_p45 = scalar.addf %n0_k_sq4, %n0_k_sq5 : f32 + %n0_k_p67 = scalar.addf %n0_k_sq6, %n0_k_sq7 : f32 + %n0_k_q03 = scalar.addf %n0_k_p01, %n0_k_p23 : f32 + %n0_k_q47 = scalar.addf %n0_k_p45, %n0_k_p67 : f32 + %n0_k_local8 = scalar.addf %n0_k_q03, %n0_k_q47 : f32 + %n0_k_c0_ge = index.cmp uge, %lane, %k0 : index + %n0_k_c0_lt = index.cmp ult, %lane, %k4 : index + %n0_k_c0_ge_v = scf.select %n0_k_c0_ge, %n0_k_local8, %f0 : f32 + %n0_k_c0_in = scf.select %n0_k_c0_lt, %n0_k_c0_ge_v, %f0 : f32 + %n0_k_w0 = kernel.subgroup.reduce %n0_k_c0_in : f32 + %n0_k_c2_ge = index.cmp uge, %lane, %k8 : index + %n0_k_c2_lt = index.cmp ult, %lane, %k12 : index + %n0_k_c2_ge_v = scf.select %n0_k_c2_ge, %n0_k_local8, %f0 : f32 + %n0_k_c2_in = scf.select %n0_k_c2_lt, %n0_k_c2_ge_v, %f0 : f32 + %n0_k_w2 = kernel.subgroup.reduce %n0_k_c2_in : f32 + %n0_k_c1_ge = index.cmp uge, %lane, %k4 : index + %n0_k_c1_lt = index.cmp ult, %lane, %k8 : index + %n0_k_c1_ge_v = scf.select %n0_k_c1_ge, %n0_k_local8, %f0 : f32 + %n0_k_c1_in = scf.select %n0_k_c1_lt, %n0_k_c1_ge_v, %f0 : f32 + %n0_k_w1 = kernel.subgroup.reduce %n0_k_c1_in : f32 + %n0_k_c3_ge = index.cmp uge, %lane, %k12 : index + %n0_k_c3_lt = index.cmp ult, %lane, %k16 : index + %n0_k_c3_ge_v = scf.select %n0_k_c3_ge, %n0_k_local8, %f0 : f32 + %n0_k_c3_in = scf.select %n0_k_c3_lt, %n0_k_c3_ge_v, %f0 : f32 + %n0_k_w3 = kernel.subgroup.reduce %n0_k_c3_in : f32 + %n0_k_c4_ge = index.cmp uge, %lane, %k16 : index + %n0_k_c4_lt = index.cmp ult, %lane, %k20 : index + %n0_k_c4_ge_v = scf.select %n0_k_c4_ge, %n0_k_local8, %f0 : f32 + %n0_k_c4_in = scf.select %n0_k_c4_lt, %n0_k_c4_ge_v, %f0 : f32 + %n0_k_w4 = kernel.subgroup.reduce %n0_k_c4_in : f32 + %n0_k_c6_ge = index.cmp uge, %lane, %k24 : index + %n0_k_c6_lt = index.cmp ult, %lane, %k28 : index + %n0_k_c6_ge_v = scf.select %n0_k_c6_ge, %n0_k_local8, %f0 : f32 + %n0_k_c6_in = scf.select %n0_k_c6_lt, %n0_k_c6_ge_v, %f0 : f32 + %n0_k_w6 = kernel.subgroup.reduce %n0_k_c6_in : f32 + %n0_k_c5_ge = index.cmp uge, %lane, %k20 : index + %n0_k_c5_lt = index.cmp ult, %lane, %k24 : index + %n0_k_c5_ge_v = scf.select %n0_k_c5_ge, %n0_k_local8, %f0 : f32 + %n0_k_c5_in = scf.select %n0_k_c5_lt, %n0_k_c5_ge_v, %f0 : f32 + %n0_k_w5 = kernel.subgroup.reduce %n0_k_c5_in : f32 + %n0_k_c7_ge = index.cmp uge, %lane, %k28 : index + %n0_k_c7_lt = index.cmp ult, %lane, %k32 : index + %n0_k_c7_ge_v = scf.select %n0_k_c7_ge, %n0_k_local8, %f0 : f32 + %n0_k_c7_in = scf.select %n0_k_c7_lt, %n0_k_c7_ge_v, %f0 : f32 + %n0_k_w7 = kernel.subgroup.reduce %n0_k_c7_in : f32 + %n0_k_lo02 = scalar.addf %n0_k_w0, %n0_k_w2 : f32 + %n0_k_lo13 = scalar.addf %n0_k_w1, %n0_k_w3 : f32 + %n0_k_lo = scalar.addf %n0_k_lo02, %n0_k_lo13 : f32 + %n0_k_hi46 = scalar.addf %n0_k_w4, %n0_k_w6 : f32 + %n0_k_hi57 = scalar.addf %n0_k_w5, %n0_k_w7 : f32 + %n0_k_hi = scalar.addf %n0_k_hi46, %n0_k_hi57 : f32 + %n0_k_sum = scf.select %lane_lo, %n0_k_lo, %n0_k_hi : f32 + %n0_k_clamped = scalar.maxnumf %n0_k_sum, %norm_eps_sq : f32 + %n0_k_scale = scalar.rsqrtf %n0_k_clamped : f32 + %n0_k_scale_v = vector.splat %n0_k_scale : vector<8xf32> + %n0_kn = vector.mulf %n0_k_scale_v, %kvs0 : vector<8xf32> + %kh0 = vector.fptrunc %n0_kn : vector<8xf32> to vector<8xf16> + %sll0_m = index.mul %stokl, %k136 : index + %sll0_a = index.add %sll0_m, %srow0 : index + %sll0_b = index.add %sll0_a, %k0 : index + %sll0 = index.assume %sll0_b [range(%sll0_b, 0, 2168)] : index + vector.store %kh0, %km_flat[%sll0] : vector<8xf16>, view<2176xf16> + // Introduce Q only after K's payload and partials are dead. + %qvr0 = vector.load %q_view[%sgi0] : view<1073741824xf32> -> vector<8xf32> + %qvs0 = scf.select %stk0_ok, %qvr0, %zv8 : vector<8xf32> + %n0_q_v0 = vector.extract %qvs0[0] : vector<8xf32> -> f32 + %n0_q_sq0 = scalar.mulf %n0_q_v0, %n0_q_v0 : f32 + %n0_q_v1 = vector.extract %qvs0[1] : vector<8xf32> -> f32 + %n0_q_sq1 = scalar.mulf %n0_q_v1, %n0_q_v1 : f32 + %n0_q_v2 = vector.extract %qvs0[2] : vector<8xf32> -> f32 + %n0_q_sq2 = scalar.mulf %n0_q_v2, %n0_q_v2 : f32 + %n0_q_v3 = vector.extract %qvs0[3] : vector<8xf32> -> f32 + %n0_q_sq3 = scalar.mulf %n0_q_v3, %n0_q_v3 : f32 + %n0_q_v4 = vector.extract %qvs0[4] : vector<8xf32> -> f32 + %n0_q_sq4 = scalar.mulf %n0_q_v4, %n0_q_v4 : f32 + %n0_q_v5 = vector.extract %qvs0[5] : vector<8xf32> -> f32 + %n0_q_sq5 = scalar.mulf %n0_q_v5, %n0_q_v5 : f32 + %n0_q_v6 = vector.extract %qvs0[6] : vector<8xf32> -> f32 + %n0_q_sq6 = scalar.mulf %n0_q_v6, %n0_q_v6 : f32 + %n0_q_v7 = vector.extract %qvs0[7] : vector<8xf32> -> f32 + %n0_q_sq7 = scalar.mulf %n0_q_v7, %n0_q_v7 : f32 + %n0_q_p01 = scalar.addf %n0_q_sq0, %n0_q_sq1 : f32 + %n0_q_p23 = scalar.addf %n0_q_sq2, %n0_q_sq3 : f32 + %n0_q_p45 = scalar.addf %n0_q_sq4, %n0_q_sq5 : f32 + %n0_q_p67 = scalar.addf %n0_q_sq6, %n0_q_sq7 : f32 + %n0_q_q03 = scalar.addf %n0_q_p01, %n0_q_p23 : f32 + %n0_q_q47 = scalar.addf %n0_q_p45, %n0_q_p67 : f32 + %n0_q_local8 = scalar.addf %n0_q_q03, %n0_q_q47 : f32 + %n0_q_c0_ge = index.cmp uge, %lane, %k0 : index + %n0_q_c0_lt = index.cmp ult, %lane, %k4 : index + %n0_q_c0_ge_v = scf.select %n0_q_c0_ge, %n0_q_local8, %f0 : f32 + %n0_q_c0_in = scf.select %n0_q_c0_lt, %n0_q_c0_ge_v, %f0 : f32 + %n0_q_w0 = kernel.subgroup.reduce %n0_q_c0_in : f32 + %n0_q_c2_ge = index.cmp uge, %lane, %k8 : index + %n0_q_c2_lt = index.cmp ult, %lane, %k12 : index + %n0_q_c2_ge_v = scf.select %n0_q_c2_ge, %n0_q_local8, %f0 : f32 + %n0_q_c2_in = scf.select %n0_q_c2_lt, %n0_q_c2_ge_v, %f0 : f32 + %n0_q_w2 = kernel.subgroup.reduce %n0_q_c2_in : f32 + %n0_q_c1_ge = index.cmp uge, %lane, %k4 : index + %n0_q_c1_lt = index.cmp ult, %lane, %k8 : index + %n0_q_c1_ge_v = scf.select %n0_q_c1_ge, %n0_q_local8, %f0 : f32 + %n0_q_c1_in = scf.select %n0_q_c1_lt, %n0_q_c1_ge_v, %f0 : f32 + %n0_q_w1 = kernel.subgroup.reduce %n0_q_c1_in : f32 + %n0_q_c3_ge = index.cmp uge, %lane, %k12 : index + %n0_q_c3_lt = index.cmp ult, %lane, %k16 : index + %n0_q_c3_ge_v = scf.select %n0_q_c3_ge, %n0_q_local8, %f0 : f32 + %n0_q_c3_in = scf.select %n0_q_c3_lt, %n0_q_c3_ge_v, %f0 : f32 + %n0_q_w3 = kernel.subgroup.reduce %n0_q_c3_in : f32 + %n0_q_c4_ge = index.cmp uge, %lane, %k16 : index + %n0_q_c4_lt = index.cmp ult, %lane, %k20 : index + %n0_q_c4_ge_v = scf.select %n0_q_c4_ge, %n0_q_local8, %f0 : f32 + %n0_q_c4_in = scf.select %n0_q_c4_lt, %n0_q_c4_ge_v, %f0 : f32 + %n0_q_w4 = kernel.subgroup.reduce %n0_q_c4_in : f32 + %n0_q_c6_ge = index.cmp uge, %lane, %k24 : index + %n0_q_c6_lt = index.cmp ult, %lane, %k28 : index + %n0_q_c6_ge_v = scf.select %n0_q_c6_ge, %n0_q_local8, %f0 : f32 + %n0_q_c6_in = scf.select %n0_q_c6_lt, %n0_q_c6_ge_v, %f0 : f32 + %n0_q_w6 = kernel.subgroup.reduce %n0_q_c6_in : f32 + %n0_q_c5_ge = index.cmp uge, %lane, %k20 : index + %n0_q_c5_lt = index.cmp ult, %lane, %k24 : index + %n0_q_c5_ge_v = scf.select %n0_q_c5_ge, %n0_q_local8, %f0 : f32 + %n0_q_c5_in = scf.select %n0_q_c5_lt, %n0_q_c5_ge_v, %f0 : f32 + %n0_q_w5 = kernel.subgroup.reduce %n0_q_c5_in : f32 + %n0_q_c7_ge = index.cmp uge, %lane, %k28 : index + %n0_q_c7_lt = index.cmp ult, %lane, %k32 : index + %n0_q_c7_ge_v = scf.select %n0_q_c7_ge, %n0_q_local8, %f0 : f32 + %n0_q_c7_in = scf.select %n0_q_c7_lt, %n0_q_c7_ge_v, %f0 : f32 + %n0_q_w7 = kernel.subgroup.reduce %n0_q_c7_in : f32 + %n0_q_lo02 = scalar.addf %n0_q_w0, %n0_q_w2 : f32 + %n0_q_lo13 = scalar.addf %n0_q_w1, %n0_q_w3 : f32 + %n0_q_lo = scalar.addf %n0_q_lo02, %n0_q_lo13 : f32 + %n0_q_hi46 = scalar.addf %n0_q_w4, %n0_q_w6 : f32 + %n0_q_hi57 = scalar.addf %n0_q_w5, %n0_q_w7 : f32 + %n0_q_hi = scalar.addf %n0_q_hi46, %n0_q_hi57 : f32 + %n0_q_sum = scf.select %lane_lo, %n0_q_lo, %n0_q_hi : f32 + %n0_q_clamped = scalar.maxnumf %n0_q_sum, %norm_eps_sq : f32 + %n0_q_scale = scalar.rsqrtf %n0_q_clamped : f32 + %n0_q_scale_v = vector.splat %n0_q_scale : vector<8xf32> + %n0_qn = vector.mulf %n0_q_scale_v, %qvs0 : vector<8xf32> + %qh0 = vector.fptrunc %n0_qn : vector<8xf32> to vector<8xf16> + vector.store %qh0, %qm_flat[%sll0] : vector<8xf16>, view<2176xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + + // K' = K^T for the state update. Retain each row's two f16x8 halves + // in registers until the attention output has vacated Q_main. + %kt_live = index.cmp ult, %tid, %k128 : index + %ktv0, %ktv1 = scf.if %kt_live -> (vector<8xf16>, vector<8xf16>) { + %ktrow = index.assume %tid [range(%tid, 0, 127)] : index + %kts0_a = index.add %ktrow, %k0 : index + %kts0 = index.assume %kts0_a [range(%kts0_a, 0, 2175)] : index + %kt0 = view.load %km_flat[%kts0] : view<2176xf16> -> f16 + %kts1_a = index.add %ktrow, %k136 : index + %kts1 = index.assume %kts1_a [range(%kts1_a, 0, 2175)] : index + %kt1 = view.load %km_flat[%kts1] : view<2176xf16> -> f16 + %kts2_a = index.add %ktrow, %k272 : index + %kts2 = index.assume %kts2_a [range(%kts2_a, 0, 2175)] : index + %kt2 = view.load %km_flat[%kts2] : view<2176xf16> -> f16 + %kts3_a = index.add %ktrow, %k408 : index + %kts3 = index.assume %kts3_a [range(%kts3_a, 0, 2175)] : index + %kt3 = view.load %km_flat[%kts3] : view<2176xf16> -> f16 + %kts4_a = index.add %ktrow, %k544 : index + %kts4 = index.assume %kts4_a [range(%kts4_a, 0, 2175)] : index + %kt4 = view.load %km_flat[%kts4] : view<2176xf16> -> f16 + %kts5_a = index.add %ktrow, %k680 : index + %kts5 = index.assume %kts5_a [range(%kts5_a, 0, 2175)] : index + %kt5 = view.load %km_flat[%kts5] : view<2176xf16> -> f16 + %kts6_a = index.add %ktrow, %k816 : index + %kts6 = index.assume %kts6_a [range(%kts6_a, 0, 2175)] : index + %kt6 = view.load %km_flat[%kts6] : view<2176xf16> -> f16 + %kts7_a = index.add %ktrow, %k952 : index + %kts7 = index.assume %kts7_a [range(%kts7_a, 0, 2175)] : index + %kt7 = view.load %km_flat[%kts7] : view<2176xf16> -> f16 + %kts8_a = index.add %ktrow, %k1088 : index + %kts8 = index.assume %kts8_a [range(%kts8_a, 0, 2175)] : index + %kt8 = view.load %km_flat[%kts8] : view<2176xf16> -> f16 + %kts9_a = index.add %ktrow, %k1224 : index + %kts9 = index.assume %kts9_a [range(%kts9_a, 0, 2175)] : index + %kt9 = view.load %km_flat[%kts9] : view<2176xf16> -> f16 + %kts10_a = index.add %ktrow, %k1360 : index + %kts10 = index.assume %kts10_a [range(%kts10_a, 0, 2175)] : index + %kt10 = view.load %km_flat[%kts10] : view<2176xf16> -> f16 + %kts11_a = index.add %ktrow, %k1496 : index + %kts11 = index.assume %kts11_a [range(%kts11_a, 0, 2175)] : index + %kt11 = view.load %km_flat[%kts11] : view<2176xf16> -> f16 + %kts12_a = index.add %ktrow, %k1632 : index + %kts12 = index.assume %kts12_a [range(%kts12_a, 0, 2175)] : index + %kt12 = view.load %km_flat[%kts12] : view<2176xf16> -> f16 + %kts13_a = index.add %ktrow, %k1768 : index + %kts13 = index.assume %kts13_a [range(%kts13_a, 0, 2175)] : index + %kt13 = view.load %km_flat[%kts13] : view<2176xf16> -> f16 + %kts14_a = index.add %ktrow, %k1904 : index + %kts14 = index.assume %kts14_a [range(%kts14_a, 0, 2175)] : index + %kt14 = view.load %km_flat[%kts14] : view<2176xf16> -> f16 + %kts15_a = index.add %ktrow, %k2040 : index + %kts15 = index.assume %kts15_a [range(%kts15_a, 0, 2175)] : index + %kt15 = view.load %km_flat[%kts15] : view<2176xf16> -> f16 + %ktlo = vector.from_elements %kt0, %kt1, %kt2, %kt3, %kt4, %kt5, %kt6, %kt7 : vector<8xf16> + %kthi = vector.from_elements %kt8, %kt9, %kt10, %kt11, %kt12, %kt13, %kt14, %kt15 : vector<8xf16> + scf.yield %ktlo, %kthi : vector<8xf16>, vector<8xf16> + } else { + scf.yield %fzh8, %fzh8 : vector<8xf16>, vector<8xf16> + } + + // State shard -> LDS as f16: the lhs of both A and B. + %shx0 = vector.fptrunc %sa0 : vector<4xf32> to vector<4xf16> + %shi0_m = index.mul %k0, %k136 : index + %shi0_a = index.add %shi0_m, %row0 : index + %shi0 = index.assume %shi0_a [range(%shi0_a, 0, 2172)] : index + vector.store %shx0, %sh_flat[%shi0] : vector<4xf16>, view<2176xf16> + %shx1 = vector.fptrunc %sa1 : vector<4xf32> to vector<4xf16> + %shi1_m = index.mul %k1, %k136 : index + %shi1_a = index.add %shi1_m, %row0 : index + %shi1 = index.assume %shi1_a [range(%shi1_a, 0, 2172)] : index + vector.store %shx1, %sh_flat[%shi1] : vector<4xf16>, view<2176xf16> + %shx2 = vector.fptrunc %sa2 : vector<4xf32> to vector<4xf16> + %shi2_m = index.mul %k2, %k136 : index + %shi2_a = index.add %shi2_m, %row0 : index + %shi2 = index.assume %shi2_a [range(%shi2_a, 0, 2172)] : index + vector.store %shx2, %sh_flat[%shi2] : vector<4xf16>, view<2176xf16> + %shx3 = vector.fptrunc %sa3 : vector<4xf32> to vector<4xf16> + %shi3_m = index.mul %k3, %k136 : index + %shi3_a = index.add %shi3_m, %row0 : index + %shi3 = index.assume %shi3_a [range(%shi3_a, 0, 2172)] : index + vector.store %shx3, %sh_flat[%shi3] : vector<4xf16>, view<2176xf16> + %shx4 = vector.fptrunc %sa4 : vector<4xf32> to vector<4xf16> + %shi4_m = index.mul %k4, %k136 : index + %shi4_a = index.add %shi4_m, %row0 : index + %shi4 = index.assume %shi4_a [range(%shi4_a, 0, 2172)] : index + vector.store %shx4, %sh_flat[%shi4] : vector<4xf16>, view<2176xf16> + %shx5 = vector.fptrunc %sa5 : vector<4xf32> to vector<4xf16> + %shi5_m = index.mul %k5, %k136 : index + %shi5_a = index.add %shi5_m, %row0 : index + %shi5 = index.assume %shi5_a [range(%shi5_a, 0, 2172)] : index + vector.store %shx5, %sh_flat[%shi5] : vector<4xf16>, view<2176xf16> + %shx6 = vector.fptrunc %sa6 : vector<4xf32> to vector<4xf16> + %shi6_m = index.mul %k6, %k136 : index + %shi6_a = index.add %shi6_m, %row0 : index + %shi6 = index.assume %shi6_a [range(%shi6_a, 0, 2172)] : index + vector.store %shx6, %sh_flat[%shi6] : vector<4xf16>, view<2176xf16> + %shx7 = vector.fptrunc %sa7 : vector<4xf32> to vector<4xf16> + %shi7_m = index.mul %k7, %k136 : index + %shi7_a = index.add %shi7_m, %row0 : index + %shi7 = index.assume %shi7_a [range(%shi7_a, 0, 2172)] : index + vector.store %shx7, %sh_flat[%shi7] : vector<4xf16>, view<2176xf16> + %shx8 = vector.fptrunc %sa8 : vector<4xf32> to vector<4xf16> + %shi8_m = index.mul %k8, %k136 : index + %shi8_a = index.add %shi8_m, %row0 : index + %shi8 = index.assume %shi8_a [range(%shi8_a, 0, 2172)] : index + vector.store %shx8, %sh_flat[%shi8] : vector<4xf16>, view<2176xf16> + %shx9 = vector.fptrunc %sa9 : vector<4xf32> to vector<4xf16> + %shi9_m = index.mul %k9, %k136 : index + %shi9_a = index.add %shi9_m, %row0 : index + %shi9 = index.assume %shi9_a [range(%shi9_a, 0, 2172)] : index + vector.store %shx9, %sh_flat[%shi9] : vector<4xf16>, view<2176xf16> + %shx10 = vector.fptrunc %sa10 : vector<4xf32> to vector<4xf16> + %shi10_m = index.mul %k10, %k136 : index + %shi10_a = index.add %shi10_m, %row0 : index + %shi10 = index.assume %shi10_a [range(%shi10_a, 0, 2172)] : index + vector.store %shx10, %sh_flat[%shi10] : vector<4xf16>, view<2176xf16> + %shx11 = vector.fptrunc %sa11 : vector<4xf32> to vector<4xf16> + %shi11_m = index.mul %k11, %k136 : index + %shi11_a = index.add %shi11_m, %row0 : index + %shi11 = index.assume %shi11_a [range(%shi11_a, 0, 2172)] : index + vector.store %shx11, %sh_flat[%shi11] : vector<4xf16>, view<2176xf16> + %shx12 = vector.fptrunc %sa12 : vector<4xf32> to vector<4xf16> + %shi12_m = index.mul %k12, %k136 : index + %shi12_a = index.add %shi12_m, %row0 : index + %shi12 = index.assume %shi12_a [range(%shi12_a, 0, 2172)] : index + vector.store %shx12, %sh_flat[%shi12] : vector<4xf16>, view<2176xf16> + %shx13 = vector.fptrunc %sa13 : vector<4xf32> to vector<4xf16> + %shi13_m = index.mul %k13, %k136 : index + %shi13_a = index.add %shi13_m, %row0 : index + %shi13 = index.assume %shi13_a [range(%shi13_a, 0, 2172)] : index + vector.store %shx13, %sh_flat[%shi13] : vector<4xf16>, view<2176xf16> + %shx14 = vector.fptrunc %sa14 : vector<4xf32> to vector<4xf16> + %shi14_m = index.mul %k14, %k136 : index + %shi14_a = index.add %shi14_m, %row0 : index + %shi14 = index.assume %shi14_a [range(%shi14_a, 0, 2172)] : index + vector.store %shx14, %sh_flat[%shi14] : vector<4xf16>, view<2176xf16> + %shx15 = vector.fptrunc %sa15 : vector<4xf32> to vector<4xf16> + %shi15_m = index.mul %k15, %k136 : index + %shi15_a = index.add %shi15_m, %row0 : index + %shi15 = index.assume %shi15_a [range(%shi15_a, 0, 2172)] : index + vector.store %shx15, %sh_flat[%shi15] : vector<4xf16>, view<2176xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + + // ---- matmuls ---------------------------------------------------------- + %ma_a0 = vector.fragment %fz8 shape [%k16, %k16] : vector<8xf32> + %mb_a0 = vector.fragment %fz8 shape [%k16, %k16] : vector<8xf32> + %ma_l0 = vector.fragment.load %sh_lhs[%k0, %k0] shape [%k16, %k16] : view<16x128xf16, %lay_colrow> -> vector<16xf16> + %ma_r0 = vector.fragment.load %km_rhs[%k0, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %ma_a1 = vector.mma %ma_l0, %ma_r0, %ma_a0 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mb_r0 = vector.fragment.load %qm_rhs[%k0, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mb_a1 = vector.mma %ma_l0, %mb_r0, %mb_a0 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %ma_l1 = vector.fragment.load %sh_lhs[%k0, %k16] shape [%k16, %k16] : view<16x128xf16, %lay_colrow> -> vector<16xf16> + %ma_r1 = vector.fragment.load %km_rhs[%k16, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %ma_a2 = vector.mma %ma_l1, %ma_r1, %ma_a1 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mb_r1 = vector.fragment.load %qm_rhs[%k16, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mb_a2 = vector.mma %ma_l1, %mb_r1, %mb_a1 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %ma_l2 = vector.fragment.load %sh_lhs[%k0, %k32] shape [%k16, %k16] : view<16x128xf16, %lay_colrow> -> vector<16xf16> + %ma_r2 = vector.fragment.load %km_rhs[%k32, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %ma_a3 = vector.mma %ma_l2, %ma_r2, %ma_a2 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mb_r2 = vector.fragment.load %qm_rhs[%k32, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mb_a3 = vector.mma %ma_l2, %mb_r2, %mb_a2 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %ma_l3 = vector.fragment.load %sh_lhs[%k0, %k48] shape [%k16, %k16] : view<16x128xf16, %lay_colrow> -> vector<16xf16> + %ma_r3 = vector.fragment.load %km_rhs[%k48, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %ma_a4 = vector.mma %ma_l3, %ma_r3, %ma_a3 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mb_r3 = vector.fragment.load %qm_rhs[%k48, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mb_a4 = vector.mma %ma_l3, %mb_r3, %mb_a3 : vector<16xf16>, vector<16xf16>, vector<8xf32> + kernel.barrier scope(subgroup) ordering(acq_rel) + %ma_l4 = vector.fragment.load %sh_lhs[%k0, %k64] shape [%k16, %k16] : view<16x128xf16, %lay_colrow> -> vector<16xf16> + %ma_r4 = vector.fragment.load %km_rhs[%k64, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %ma_a5 = vector.mma %ma_l4, %ma_r4, %ma_a4 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mb_r4 = vector.fragment.load %qm_rhs[%k64, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mb_a5 = vector.mma %ma_l4, %mb_r4, %mb_a4 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %ma_l5 = vector.fragment.load %sh_lhs[%k0, %k80] shape [%k16, %k16] : view<16x128xf16, %lay_colrow> -> vector<16xf16> + %ma_r5 = vector.fragment.load %km_rhs[%k80, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %ma_a6 = vector.mma %ma_l5, %ma_r5, %ma_a5 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mb_r5 = vector.fragment.load %qm_rhs[%k80, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mb_a6 = vector.mma %ma_l5, %mb_r5, %mb_a5 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %ma_l6 = vector.fragment.load %sh_lhs[%k0, %k96] shape [%k16, %k16] : view<16x128xf16, %lay_colrow> -> vector<16xf16> + %ma_r6 = vector.fragment.load %km_rhs[%k96, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %ma_a7 = vector.mma %ma_l6, %ma_r6, %ma_a6 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mb_r6 = vector.fragment.load %qm_rhs[%k96, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mb_a7 = vector.mma %ma_l6, %mb_r6, %mb_a6 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %ma_l7 = vector.fragment.load %sh_lhs[%k0, %k112] shape [%k16, %k16] : view<16x128xf16, %lay_colrow> -> vector<16xf16> + %ma_r7 = vector.fragment.load %km_rhs[%k112, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %ma_a8 = vector.mma %ma_l7, %ma_r7, %ma_a7 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mb_r7 = vector.fragment.load %qm_rhs[%k112, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mb_a8 = vector.mma %ma_l7, %mb_r7, %mb_a7 : vector<16xf16>, vector<16xf16>, vector<8xf32> + vector.fragment.store %ma_a8, %a_res[%k0, %k0] shape [%k16, %k16] : vector<8xf32>, view<16x16xf32, %lay_tile> + vector.fragment.store %mb_a8, %b_res[%k0, %k0] shape [%k16, %k16] : vector<8xf32>, view<16x16xf32, %lay_tile> + // G and QK are per head, so the two waves take one each. + scf.if %is_w0 { + %mg_a0 = vector.fragment %fz8 shape [%k16, %k16] : vector<8xf32> + %mg_l0 = vector.fragment.load %km_lhs[%k0, %k0] shape [%k16, %k16] : view<16x128xf16, %lay_tokrow> -> vector<16xf16> + %mg_r0 = vector.fragment.load %km_rhs[%k0, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mg_a1 = vector.mma %mg_l0, %mg_r0, %mg_a0 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mg_l1 = vector.fragment.load %km_lhs[%k0, %k16] shape [%k16, %k16] : view<16x128xf16, %lay_tokrow> -> vector<16xf16> + %mg_r1 = vector.fragment.load %km_rhs[%k16, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mg_a2 = vector.mma %mg_l1, %mg_r1, %mg_a1 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mg_l2 = vector.fragment.load %km_lhs[%k0, %k32] shape [%k16, %k16] : view<16x128xf16, %lay_tokrow> -> vector<16xf16> + %mg_r2 = vector.fragment.load %km_rhs[%k32, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mg_a3 = vector.mma %mg_l2, %mg_r2, %mg_a2 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mg_l3 = vector.fragment.load %km_lhs[%k0, %k48] shape [%k16, %k16] : view<16x128xf16, %lay_tokrow> -> vector<16xf16> + %mg_r3 = vector.fragment.load %km_rhs[%k48, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mg_a4 = vector.mma %mg_l3, %mg_r3, %mg_a3 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mg_l4 = vector.fragment.load %km_lhs[%k0, %k64] shape [%k16, %k16] : view<16x128xf16, %lay_tokrow> -> vector<16xf16> + %mg_r4 = vector.fragment.load %km_rhs[%k64, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mg_a5 = vector.mma %mg_l4, %mg_r4, %mg_a4 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mg_l5 = vector.fragment.load %km_lhs[%k0, %k80] shape [%k16, %k16] : view<16x128xf16, %lay_tokrow> -> vector<16xf16> + %mg_r5 = vector.fragment.load %km_rhs[%k80, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mg_a6 = vector.mma %mg_l5, %mg_r5, %mg_a5 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mg_l6 = vector.fragment.load %km_lhs[%k0, %k96] shape [%k16, %k16] : view<16x128xf16, %lay_tokrow> -> vector<16xf16> + %mg_r6 = vector.fragment.load %km_rhs[%k96, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mg_a7 = vector.mma %mg_l6, %mg_r6, %mg_a6 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mg_l7 = vector.fragment.load %km_lhs[%k0, %k112] shape [%k16, %k16] : view<16x128xf16, %lay_tokrow> -> vector<16xf16> + %mg_r7 = vector.fragment.load %km_rhs[%k112, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mg_a8 = vector.mma %mg_l7, %mg_r7, %mg_a7 : vector<16xf16>, vector<16xf16>, vector<8xf32> + vector.fragment.store %mg_a8, %gr_res[%k0, %k0] shape [%k16, %k16] : vector<8xf32>, view<16x16xf32, %lay_gram> + } + scf.if %is_w1 { + %mq_a0 = vector.fragment %fz8 shape [%k16, %k16] : vector<8xf32> + %mq_l0 = vector.fragment.load %km_lhs[%k0, %k0] shape [%k16, %k16] : view<16x128xf16, %lay_tokrow> -> vector<16xf16> + %mq_r0 = vector.fragment.load %qm_rhs[%k0, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mq_a1 = vector.mma %mq_l0, %mq_r0, %mq_a0 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mq_l1 = vector.fragment.load %km_lhs[%k0, %k16] shape [%k16, %k16] : view<16x128xf16, %lay_tokrow> -> vector<16xf16> + %mq_r1 = vector.fragment.load %qm_rhs[%k16, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mq_a2 = vector.mma %mq_l1, %mq_r1, %mq_a1 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mq_l2 = vector.fragment.load %km_lhs[%k0, %k32] shape [%k16, %k16] : view<16x128xf16, %lay_tokrow> -> vector<16xf16> + %mq_r2 = vector.fragment.load %qm_rhs[%k32, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mq_a3 = vector.mma %mq_l2, %mq_r2, %mq_a2 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mq_l3 = vector.fragment.load %km_lhs[%k0, %k48] shape [%k16, %k16] : view<16x128xf16, %lay_tokrow> -> vector<16xf16> + %mq_r3 = vector.fragment.load %qm_rhs[%k48, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mq_a4 = vector.mma %mq_l3, %mq_r3, %mq_a3 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mq_l4 = vector.fragment.load %km_lhs[%k0, %k64] shape [%k16, %k16] : view<16x128xf16, %lay_tokrow> -> vector<16xf16> + %mq_r4 = vector.fragment.load %qm_rhs[%k64, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mq_a5 = vector.mma %mq_l4, %mq_r4, %mq_a4 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mq_l5 = vector.fragment.load %km_lhs[%k0, %k80] shape [%k16, %k16] : view<16x128xf16, %lay_tokrow> -> vector<16xf16> + %mq_r5 = vector.fragment.load %qm_rhs[%k80, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mq_a6 = vector.mma %mq_l5, %mq_r5, %mq_a5 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mq_l6 = vector.fragment.load %km_lhs[%k0, %k96] shape [%k16, %k16] : view<16x128xf16, %lay_tokrow> -> vector<16xf16> + %mq_r6 = vector.fragment.load %qm_rhs[%k96, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mq_a7 = vector.mma %mq_l6, %mq_r6, %mq_a6 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %mq_l7 = vector.fragment.load %km_lhs[%k0, %k112] shape [%k16, %k16] : view<16x128xf16, %lay_tokrow> -> vector<16xf16> + %mq_r7 = vector.fragment.load %qm_rhs[%k112, %k0] shape [%k16, %k16] : view<128x16xf16, %lay_rowtok> -> vector<16xf16> + %mq_a8 = vector.mma %mq_l7, %mq_r7, %mq_a7 : vector<16xf16>, vector<16xf16>, vector<8xf32> + vector.fragment.store %mq_a8, %qkt_res[%k0, %k0] shape [%k16, %k16] : vector<8xf32>, view<16x16xf32, %lay_gram> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %df0_a = index.add %tid, %k0 : index + %df0 = index.assume %df0_a [range(%df0_a, 0, 255)] : index + %ds0_r = index.div %df0, %k16 : index + %ds0 = index.assume %ds0_r [range(%ds0_r, 0, 15)] : index + %dt0_r = index.rem %df0, %k16 : index + %dt0 = index.assume %dt0_r [range(%dt0_r, 0, 15)] : index + %du0 = index.cmp ult, %ds0, %dt0 : index + scf.if %du0 { + %dsi0_m = index.mul %ds0, %k4 : index + %dsi0_a = index.add %dsi0_m, %k3 : index + %dsi0 = index.assume %dsi0_a [range(%dsi0_a, 3, 63)] : index + %dti0_m = index.mul %dt0, %k4 : index + %dti0_a = index.add %dti0_m, %k3 : index + %dti0 = index.assume %dti0_a [range(%dti0_a, 3, 63)] : index + %dsl0 = view.load %gs_flat[%dsi0] : view<80xf32> -> f32 + %dtl0 = view.load %gs_flat[%dti0] : view<80xf32> -> f32 + %ddl0 = scalar.subf %dtl0, %dsl0 : f32 + %ddc0 = scalar.expf %ddl0 : f32 + %dgv0 = view.load %gr_flat[%df0] : view<256xf32> -> f32 + %dqv0 = view.load %qkt_flat[%df0] : view<256xf32> -> f32 + %dgs0 = scalar.mulf %dgv0, %ddc0 : f32 + %dqs0 = scalar.mulf %dqv0, %ddc0 : f32 + view.store %dgs0, %gr_flat[%df0] : f32, view<256xf32> + view.store %dqs0, %qkt_flat[%df0] : f32, view<256xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %vlt0_r = index.add %t0, %k0 : index + %vlt0_s = index.add %vlt0_r, %vs_tok : index + %vlt0_ok = index.cmp ult, %vlt0_s, %n_tokens : index + %vlt0_c = scf.select %vlt0_ok, %vlt0_s, %c0 : index + %vlt0 = index.assume %vlt0_c [range(%vlt0_c, 0, 1048575)] : index + %vli0_m = index.mul %vlt0, %sv2 : index + %vli0_a = index.add %v_head_base, %vli0_m : index + %vli0_c = index.add %vli0_a, %vs_gcol : index + %vli0 = index.assume %vli0_c [range(%vli0_c, 0, 1073741823)] : index + %vlv0 = view.load %v_view[%vli0] : view<1073741824xf32> -> f32 + %vls0_m = index.mul %vs_tok, %k128 : index + %vls0_a = index.add %vls0_m, %vs_col : index + %vls0_b = index.add %vls0_a, %k0 : index + %vls0 = index.assume %vls0_b [range(%vls0_b, 0, 2047)] : index + view.store %vlv0, %vst_flat[%vls0] : f32, view<2048xf32> + %vlt1_r = index.add %t0, %k2 : index + %vlt1_s = index.add %vlt1_r, %vs_tok : index + %vlt1_ok = index.cmp ult, %vlt1_s, %n_tokens : index + %vlt1_c = scf.select %vlt1_ok, %vlt1_s, %c0 : index + %vlt1 = index.assume %vlt1_c [range(%vlt1_c, 0, 1048575)] : index + %vli1_m = index.mul %vlt1, %sv2 : index + %vli1_a = index.add %v_head_base, %vli1_m : index + %vli1_c = index.add %vli1_a, %vs_gcol : index + %vli1 = index.assume %vli1_c [range(%vli1_c, 0, 1073741823)] : index + %vlv1 = view.load %v_view[%vli1] : view<1073741824xf32> -> f32 + %vls1_m = index.mul %vs_tok, %k128 : index + %vls1_a = index.add %vls1_m, %vs_col : index + %vls1_b = index.add %vls1_a, %k256 : index + %vls1 = index.assume %vls1_b [range(%vls1_b, 0, 2047)] : index + view.store %vlv1, %vst_flat[%vls1] : f32, view<2048xf32> + %vlt2_r = index.add %t0, %k4 : index + %vlt2_s = index.add %vlt2_r, %vs_tok : index + %vlt2_ok = index.cmp ult, %vlt2_s, %n_tokens : index + %vlt2_c = scf.select %vlt2_ok, %vlt2_s, %c0 : index + %vlt2 = index.assume %vlt2_c [range(%vlt2_c, 0, 1048575)] : index + %vli2_m = index.mul %vlt2, %sv2 : index + %vli2_a = index.add %v_head_base, %vli2_m : index + %vli2_c = index.add %vli2_a, %vs_gcol : index + %vli2 = index.assume %vli2_c [range(%vli2_c, 0, 1073741823)] : index + %vlv2 = view.load %v_view[%vli2] : view<1073741824xf32> -> f32 + %vls2_m = index.mul %vs_tok, %k128 : index + %vls2_a = index.add %vls2_m, %vs_col : index + %vls2_b = index.add %vls2_a, %k512 : index + %vls2 = index.assume %vls2_b [range(%vls2_b, 0, 2047)] : index + view.store %vlv2, %vst_flat[%vls2] : f32, view<2048xf32> + %vlt3_r = index.add %t0, %k6 : index + %vlt3_s = index.add %vlt3_r, %vs_tok : index + %vlt3_ok = index.cmp ult, %vlt3_s, %n_tokens : index + %vlt3_c = scf.select %vlt3_ok, %vlt3_s, %c0 : index + %vlt3 = index.assume %vlt3_c [range(%vlt3_c, 0, 1048575)] : index + %vli3_m = index.mul %vlt3, %sv2 : index + %vli3_a = index.add %v_head_base, %vli3_m : index + %vli3_c = index.add %vli3_a, %vs_gcol : index + %vli3 = index.assume %vli3_c [range(%vli3_c, 0, 1073741823)] : index + %vlv3 = view.load %v_view[%vli3] : view<1073741824xf32> -> f32 + %vls3_m = index.mul %vs_tok, %k128 : index + %vls3_a = index.add %vls3_m, %vs_col : index + %vls3_b = index.add %vls3_a, %k768 : index + %vls3 = index.assume %vls3_b [range(%vls3_b, 0, 2047)] : index + view.store %vlv3, %vst_flat[%vls3] : f32, view<2048xf32> + %vlt4_r = index.add %t0, %k8 : index + %vlt4_s = index.add %vlt4_r, %vs_tok : index + %vlt4_ok = index.cmp ult, %vlt4_s, %n_tokens : index + %vlt4_c = scf.select %vlt4_ok, %vlt4_s, %c0 : index + %vlt4 = index.assume %vlt4_c [range(%vlt4_c, 0, 1048575)] : index + %vli4_m = index.mul %vlt4, %sv2 : index + %vli4_a = index.add %v_head_base, %vli4_m : index + %vli4_c = index.add %vli4_a, %vs_gcol : index + %vli4 = index.assume %vli4_c [range(%vli4_c, 0, 1073741823)] : index + %vlv4 = view.load %v_view[%vli4] : view<1073741824xf32> -> f32 + %vls4_m = index.mul %vs_tok, %k128 : index + %vls4_a = index.add %vls4_m, %vs_col : index + %vls4_b = index.add %vls4_a, %k1024 : index + %vls4 = index.assume %vls4_b [range(%vls4_b, 0, 2047)] : index + view.store %vlv4, %vst_flat[%vls4] : f32, view<2048xf32> + %vlt5_r = index.add %t0, %k10 : index + %vlt5_s = index.add %vlt5_r, %vs_tok : index + %vlt5_ok = index.cmp ult, %vlt5_s, %n_tokens : index + %vlt5_c = scf.select %vlt5_ok, %vlt5_s, %c0 : index + %vlt5 = index.assume %vlt5_c [range(%vlt5_c, 0, 1048575)] : index + %vli5_m = index.mul %vlt5, %sv2 : index + %vli5_a = index.add %v_head_base, %vli5_m : index + %vli5_c = index.add %vli5_a, %vs_gcol : index + %vli5 = index.assume %vli5_c [range(%vli5_c, 0, 1073741823)] : index + %vlv5 = view.load %v_view[%vli5] : view<1073741824xf32> -> f32 + %vls5_m = index.mul %vs_tok, %k128 : index + %vls5_a = index.add %vls5_m, %vs_col : index + %vls5_b = index.add %vls5_a, %k1280 : index + %vls5 = index.assume %vls5_b [range(%vls5_b, 0, 2047)] : index + view.store %vlv5, %vst_flat[%vls5] : f32, view<2048xf32> + %vlt6_r = index.add %t0, %k12 : index + %vlt6_s = index.add %vlt6_r, %vs_tok : index + %vlt6_ok = index.cmp ult, %vlt6_s, %n_tokens : index + %vlt6_c = scf.select %vlt6_ok, %vlt6_s, %c0 : index + %vlt6 = index.assume %vlt6_c [range(%vlt6_c, 0, 1048575)] : index + %vli6_m = index.mul %vlt6, %sv2 : index + %vli6_a = index.add %v_head_base, %vli6_m : index + %vli6_c = index.add %vli6_a, %vs_gcol : index + %vli6 = index.assume %vli6_c [range(%vli6_c, 0, 1073741823)] : index + %vlv6 = view.load %v_view[%vli6] : view<1073741824xf32> -> f32 + %vls6_m = index.mul %vs_tok, %k128 : index + %vls6_a = index.add %vls6_m, %vs_col : index + %vls6_b = index.add %vls6_a, %k1536 : index + %vls6 = index.assume %vls6_b [range(%vls6_b, 0, 2047)] : index + view.store %vlv6, %vst_flat[%vls6] : f32, view<2048xf32> + %vlt7_r = index.add %t0, %k14 : index + %vlt7_s = index.add %vlt7_r, %vs_tok : index + %vlt7_ok = index.cmp ult, %vlt7_s, %n_tokens : index + %vlt7_c = scf.select %vlt7_ok, %vlt7_s, %c0 : index + %vlt7 = index.assume %vlt7_c [range(%vlt7_c, 0, 1048575)] : index + %vli7_m = index.mul %vlt7, %sv2 : index + %vli7_a = index.add %v_head_base, %vli7_m : index + %vli7_c = index.add %vli7_a, %vs_gcol : index + %vli7 = index.assume %vli7_c [range(%vli7_c, 0, 1073741823)] : index + %vlv7 = view.load %v_view[%vli7] : view<1073741824xf32> -> f32 + %vls7_m = index.mul %vs_tok, %k128 : index + %vls7_a = index.add %vls7_m, %vs_col : index + %vls7_b = index.add %vls7_a, %k1792 : index + %vls7 = index.assume %vls7_b [range(%vls7_b, 0, 2047)] : index + view.store %vlv7, %vst_flat[%vls7] : f32, view<2048xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + + // Lane c owns column c; recurrence operands are registers or LDS broadcasts with no cross-lane exchange. + %gq0 = vector.load %gs_flat[%k0] : view<80xf32> -> vector<4xf32> + %bv0 = vector.extract %gq0[0] : vector<4xf32> -> f32 + %cp0 = vector.extract %gq0[1] : vector<4xf32> -> f32 + %lc0 = vector.extract %gq0[3] : vector<4xf32> -> f32 + %ev0 = vector.extract %gq0[2] : vector<4xf32> -> f32 + %tk0_r = index.add %t0, %k0 : index + %tk0_ok = index.cmp ult, %tk0_r, %n_tokens : index + %tk0_s = scf.select %tk0_ok, %tk0_r, %c0 : index + %tk0 = index.assume %tk0_s [range(%tk0_s, 0, 1048575)] : index + %av0_i = index.add %tile_row, %k0 : index + %av0_b = index.assume %av0_i [range(%av0_i, 0, 271)] : index + %av0 = view.load %a_flat[%av0_b] : view<272xf32> -> f32 + %bvv0 = view.load %b_flat[%av0_b] : view<272xf32> -> f32 + %vrd0_a = index.add %rec_col, %k0 : index + %vrd0 = index.assume %vrd0_a [range(%vrd0_a, 0, 2047)] : index + %vval0 = view.load %vst_flat[%vrd0] : view<2048xf32> -> f32 + %adec0 = scalar.mulf %av0, %cp0 : f32 + %vsub0 = scalar.subf %vval0, %adec0 : f32 + // Keep the original padded multiply, including its signed zero. + %active_index0 = index.add %t0, %k0 : index + %row_in_range0 = index.cmp ult, %active_index0, %n_tokens : index + %active0 = scalar.ori %row_in_range0, %all_recurrence_rows : i1 + %dl0 = scf.if %active0 -> (f32) { + %value = scalar.mulf %vsub0, %bv0 : f32 + scf.yield %value : f32 + } else { + %padded = scalar.mulf %vval0, %bv0 : f32 + scf.yield %padded : f32 + } + %dh0 = scalar.mulf %dl0, %ev0 : f32 + %grw0_0_a = vector.load %gr_flat[%k0] : view<256xf32> -> vector<4xf32> + %grw0_0_bi = index.constant 4 : index + %grw0_0_b = vector.load %gr_flat[%grw0_0_bi] : view<256xf32> -> vector<4xf32> + %grw0_1_a = vector.load %gr_flat[%k8] : view<256xf32> -> vector<4xf32> + %grw0_1_bi = index.constant 12 : index + %grw0_1_b = vector.load %gr_flat[%grw0_1_bi] : view<256xf32> -> vector<4xf32> + %qkw0_0_a = vector.load %qkt_flat[%k0] : view<256xf32> -> vector<4xf32> + %qkw0_0_bi = index.constant 4 : index + %qkw0_0_b = vector.load %qkt_flat[%qkw0_0_bi] : view<256xf32> -> vector<4xf32> + %qkw0_1_a = vector.load %qkt_flat[%k8] : view<256xf32> -> vector<4xf32> + %qkw0_1_bi = index.constant 12 : index + %qkw0_1_b = vector.load %qkt_flat[%qkw0_1_bi] : view<256xf32> -> vector<4xf32> + %qkd0_0 = vector.extract %qkw0_0_a[0] : vector<4xf32> -> f32 + %qsum0 = scalar.mulf %qkd0_0, %dl0 : f32 + %ao0_r = scalar.fmaf %bvv0, %cp0, %qsum0 : f32 + %ao0 = scf.if %active0 -> (f32) { + %value = scalar.mulf %ao0_r, %scale : f32 + scf.yield %value : f32 + } else { + %zero = scalar.constant 0.0 : f32 + scf.yield %zero : f32 + } + %aoi0_a = index.add %rec_col, %k0 : index + %aoi0 = index.assume %aoi0_a [range(%aoi0_a, 0, 2047)] : index + %ddh0 = scalar.fptrunc %dh0 : f32 to f16 + %publish_active0 = scalar.andi %lane_live, %active0 : i1 + scf.if %publish_active0 { + view.store %ao0, %aost_flat[%aoi0] : f32, view<2048xf32> + } + %gp0_1 = vector.extract %grw0_0_a[1] : vector<4xf32> -> f32 + %ag0_1 = scalar.mulf %gp0_1, %dl0 : f32 + %qp0_1 = vector.extract %qkw0_0_a[1] : vector<4xf32> -> f32 + %aq0_1 = scalar.mulf %qp0_1, %dl0 : f32 + %gp0_2 = vector.extract %grw0_0_a[2] : vector<4xf32> -> f32 + %ag0_2 = scalar.mulf %gp0_2, %dl0 : f32 + %qp0_2 = vector.extract %qkw0_0_a[2] : vector<4xf32> -> f32 + %aq0_2 = scalar.mulf %qp0_2, %dl0 : f32 + %gp0_3 = vector.extract %grw0_0_a[3] : vector<4xf32> -> f32 + %ag0_3 = scalar.mulf %gp0_3, %dl0 : f32 + %qp0_3 = vector.extract %qkw0_0_a[3] : vector<4xf32> -> f32 + %aq0_3 = scalar.mulf %qp0_3, %dl0 : f32 + %gp0_4 = vector.extract %grw0_0_b[0] : vector<4xf32> -> f32 + %ag0_4 = scalar.mulf %gp0_4, %dl0 : f32 + %qp0_4 = vector.extract %qkw0_0_b[0] : vector<4xf32> -> f32 + %aq0_4 = scalar.mulf %qp0_4, %dl0 : f32 + %gp0_5 = vector.extract %grw0_0_b[1] : vector<4xf32> -> f32 + %ag0_5 = scalar.mulf %gp0_5, %dl0 : f32 + %qp0_5 = vector.extract %qkw0_0_b[1] : vector<4xf32> -> f32 + %aq0_5 = scalar.mulf %qp0_5, %dl0 : f32 + %gp0_6 = vector.extract %grw0_0_b[2] : vector<4xf32> -> f32 + %ag0_6 = scalar.mulf %gp0_6, %dl0 : f32 + %qp0_6 = vector.extract %qkw0_0_b[2] : vector<4xf32> -> f32 + %aq0_6 = scalar.mulf %qp0_6, %dl0 : f32 + %gp0_7 = vector.extract %grw0_0_b[3] : vector<4xf32> -> f32 + %ag0_7 = scalar.mulf %gp0_7, %dl0 : f32 + %qp0_7 = vector.extract %qkw0_0_b[3] : vector<4xf32> -> f32 + %aq0_7 = scalar.mulf %qp0_7, %dl0 : f32 + %gp0_8 = vector.extract %grw0_1_a[0] : vector<4xf32> -> f32 + %ag0_8 = scalar.mulf %gp0_8, %dl0 : f32 + %qp0_8 = vector.extract %qkw0_1_a[0] : vector<4xf32> -> f32 + %aq0_8 = scalar.mulf %qp0_8, %dl0 : f32 + %gp0_9 = vector.extract %grw0_1_a[1] : vector<4xf32> -> f32 + %ag0_9 = scalar.mulf %gp0_9, %dl0 : f32 + %qp0_9 = vector.extract %qkw0_1_a[1] : vector<4xf32> -> f32 + %aq0_9 = scalar.mulf %qp0_9, %dl0 : f32 + %gp0_10 = vector.extract %grw0_1_a[2] : vector<4xf32> -> f32 + %ag0_10 = scalar.mulf %gp0_10, %dl0 : f32 + %qp0_10 = vector.extract %qkw0_1_a[2] : vector<4xf32> -> f32 + %aq0_10 = scalar.mulf %qp0_10, %dl0 : f32 + %gp0_11 = vector.extract %grw0_1_a[3] : vector<4xf32> -> f32 + %ag0_11 = scalar.mulf %gp0_11, %dl0 : f32 + %qp0_11 = vector.extract %qkw0_1_a[3] : vector<4xf32> -> f32 + %aq0_11 = scalar.mulf %qp0_11, %dl0 : f32 + %gp0_12 = vector.extract %grw0_1_b[0] : vector<4xf32> -> f32 + %ag0_12 = scalar.mulf %gp0_12, %dl0 : f32 + %qp0_12 = vector.extract %qkw0_1_b[0] : vector<4xf32> -> f32 + %aq0_12 = scalar.mulf %qp0_12, %dl0 : f32 + %gp0_13 = vector.extract %grw0_1_b[1] : vector<4xf32> -> f32 + %ag0_13 = scalar.mulf %gp0_13, %dl0 : f32 + %qp0_13 = vector.extract %qkw0_1_b[1] : vector<4xf32> -> f32 + %aq0_13 = scalar.mulf %qp0_13, %dl0 : f32 + %gp0_14 = vector.extract %grw0_1_b[2] : vector<4xf32> -> f32 + %ag0_14 = scalar.mulf %gp0_14, %dl0 : f32 + %qp0_14 = vector.extract %qkw0_1_b[2] : vector<4xf32> -> f32 + %aq0_14 = scalar.mulf %qp0_14, %dl0 : f32 + %gp0_15 = vector.extract %grw0_1_b[3] : vector<4xf32> -> f32 + %ag0_15 = scalar.mulf %gp0_15, %dl0 : f32 + %qp0_15 = vector.extract %qkw0_1_b[3] : vector<4xf32> -> f32 + %aq0_15 = scalar.mulf %qp0_15, %dl0 : f32 + %gq1 = vector.load %gs_flat[%k4] : view<80xf32> -> vector<4xf32> + %bv1 = vector.extract %gq1[0] : vector<4xf32> -> f32 + %cp1 = vector.extract %gq1[1] : vector<4xf32> -> f32 + %lc1 = vector.extract %gq1[3] : vector<4xf32> -> f32 + %ev1 = vector.extract %gq1[2] : vector<4xf32> -> f32 + %tk1_r = index.add %t0, %k1 : index + %tk1_ok = index.cmp ult, %tk1_r, %n_tokens : index + %tk1_s = scf.select %tk1_ok, %tk1_r, %c0 : index + %tk1 = index.assume %tk1_s [range(%tk1_s, 0, 1048575)] : index + %av1_i = index.add %tile_row, %k1 : index + %av1_b = index.assume %av1_i [range(%av1_i, 0, 271)] : index + %av1 = view.load %a_flat[%av1_b] : view<272xf32> -> f32 + %bvv1 = view.load %b_flat[%av1_b] : view<272xf32> -> f32 + %vrd1_a = index.add %rec_col, %k128 : index + %vrd1 = index.assume %vrd1_a [range(%vrd1_a, 0, 2047)] : index + %vval1 = view.load %vst_flat[%vrd1] : view<2048xf32> -> f32 + %adec1 = scalar.fmaf %av1, %cp1, %ag0_1 : f32 + %vsub1 = scalar.subf %vval1, %adec1 : f32 + %active_index1 = index.add %t0, %k1 : index + %row_in_range1 = index.cmp ult, %active_index1, %n_tokens : index + %active1 = scalar.ori %row_in_range1, %all_recurrence_rows : i1 + %dl1 = scf.if %active1 -> (f32) { + %value = scalar.mulf %vsub1, %bv1 : f32 + scf.yield %value : f32 + } else { + %padded = scalar.mulf %vval1, %bv1 : f32 + scf.yield %padded : f32 + } + %dh1 = scalar.mulf %dl1, %ev1 : f32 + %grw1_0_a = vector.load %gr_flat[%k16] : view<256xf32> -> vector<4xf32> + %grw1_0_bi = index.constant 20 : index + %grw1_0_b = vector.load %gr_flat[%grw1_0_bi] : view<256xf32> -> vector<4xf32> + %grw1_1_a = vector.load %gr_flat[%k24] : view<256xf32> -> vector<4xf32> + %grw1_1_bi = index.constant 28 : index + %grw1_1_b = vector.load %gr_flat[%grw1_1_bi] : view<256xf32> -> vector<4xf32> + %qkw1_0_a = vector.load %qkt_flat[%k16] : view<256xf32> -> vector<4xf32> + %qkw1_0_bi = index.constant 20 : index + %qkw1_0_b = vector.load %qkt_flat[%qkw1_0_bi] : view<256xf32> -> vector<4xf32> + %qkw1_1_a = vector.load %qkt_flat[%k24] : view<256xf32> -> vector<4xf32> + %qkw1_1_bi = index.constant 28 : index + %qkw1_1_b = vector.load %qkt_flat[%qkw1_1_bi] : view<256xf32> -> vector<4xf32> + %qkd1_1 = vector.extract %qkw1_0_a[1] : vector<4xf32> -> f32 + %qsum1 = scalar.fmaf %qkd1_1, %dl1, %aq0_1 : f32 + %ao1_r = scalar.fmaf %bvv1, %cp1, %qsum1 : f32 + %ao1 = scf.if %active1 -> (f32) { + %value = scalar.mulf %ao1_r, %scale : f32 + scf.yield %value : f32 + } else { + %zero = scalar.constant 0.0 : f32 + scf.yield %zero : f32 + } + %aoi1_a = index.add %rec_col, %k128 : index + %aoi1 = index.assume %aoi1_a [range(%aoi1_a, 0, 2047)] : index + %ddh1 = scalar.fptrunc %dh1 : f32 to f16 + %publish_active1 = scalar.andi %lane_live, %active1 : i1 + scf.if %publish_active1 { + view.store %ao1, %aost_flat[%aoi1] : f32, view<2048xf32> + } + %gp1_2 = vector.extract %grw1_0_a[2] : vector<4xf32> -> f32 + %ag1_2 = scalar.fmaf %gp1_2, %dl1, %ag0_2 : f32 + %qp1_2 = vector.extract %qkw1_0_a[2] : vector<4xf32> -> f32 + %aq1_2 = scalar.fmaf %qp1_2, %dl1, %aq0_2 : f32 + %gp1_3 = vector.extract %grw1_0_a[3] : vector<4xf32> -> f32 + %ag1_3 = scalar.fmaf %gp1_3, %dl1, %ag0_3 : f32 + %qp1_3 = vector.extract %qkw1_0_a[3] : vector<4xf32> -> f32 + %aq1_3 = scalar.fmaf %qp1_3, %dl1, %aq0_3 : f32 + %gp1_4 = vector.extract %grw1_0_b[0] : vector<4xf32> -> f32 + %ag1_4 = scalar.fmaf %gp1_4, %dl1, %ag0_4 : f32 + %qp1_4 = vector.extract %qkw1_0_b[0] : vector<4xf32> -> f32 + %aq1_4 = scalar.fmaf %qp1_4, %dl1, %aq0_4 : f32 + %gp1_5 = vector.extract %grw1_0_b[1] : vector<4xf32> -> f32 + %ag1_5 = scalar.fmaf %gp1_5, %dl1, %ag0_5 : f32 + %qp1_5 = vector.extract %qkw1_0_b[1] : vector<4xf32> -> f32 + %aq1_5 = scalar.fmaf %qp1_5, %dl1, %aq0_5 : f32 + %gp1_6 = vector.extract %grw1_0_b[2] : vector<4xf32> -> f32 + %ag1_6 = scalar.fmaf %gp1_6, %dl1, %ag0_6 : f32 + %qp1_6 = vector.extract %qkw1_0_b[2] : vector<4xf32> -> f32 + %aq1_6 = scalar.fmaf %qp1_6, %dl1, %aq0_6 : f32 + %gp1_7 = vector.extract %grw1_0_b[3] : vector<4xf32> -> f32 + %ag1_7 = scalar.fmaf %gp1_7, %dl1, %ag0_7 : f32 + %qp1_7 = vector.extract %qkw1_0_b[3] : vector<4xf32> -> f32 + %aq1_7 = scalar.fmaf %qp1_7, %dl1, %aq0_7 : f32 + %gp1_8 = vector.extract %grw1_1_a[0] : vector<4xf32> -> f32 + %ag1_8 = scalar.fmaf %gp1_8, %dl1, %ag0_8 : f32 + %qp1_8 = vector.extract %qkw1_1_a[0] : vector<4xf32> -> f32 + %aq1_8 = scalar.fmaf %qp1_8, %dl1, %aq0_8 : f32 + %gp1_9 = vector.extract %grw1_1_a[1] : vector<4xf32> -> f32 + %ag1_9 = scalar.fmaf %gp1_9, %dl1, %ag0_9 : f32 + %qp1_9 = vector.extract %qkw1_1_a[1] : vector<4xf32> -> f32 + %aq1_9 = scalar.fmaf %qp1_9, %dl1, %aq0_9 : f32 + %gp1_10 = vector.extract %grw1_1_a[2] : vector<4xf32> -> f32 + %ag1_10 = scalar.fmaf %gp1_10, %dl1, %ag0_10 : f32 + %qp1_10 = vector.extract %qkw1_1_a[2] : vector<4xf32> -> f32 + %aq1_10 = scalar.fmaf %qp1_10, %dl1, %aq0_10 : f32 + %gp1_11 = vector.extract %grw1_1_a[3] : vector<4xf32> -> f32 + %ag1_11 = scalar.fmaf %gp1_11, %dl1, %ag0_11 : f32 + %qp1_11 = vector.extract %qkw1_1_a[3] : vector<4xf32> -> f32 + %aq1_11 = scalar.fmaf %qp1_11, %dl1, %aq0_11 : f32 + %gp1_12 = vector.extract %grw1_1_b[0] : vector<4xf32> -> f32 + %ag1_12 = scalar.fmaf %gp1_12, %dl1, %ag0_12 : f32 + %qp1_12 = vector.extract %qkw1_1_b[0] : vector<4xf32> -> f32 + %aq1_12 = scalar.fmaf %qp1_12, %dl1, %aq0_12 : f32 + %gp1_13 = vector.extract %grw1_1_b[1] : vector<4xf32> -> f32 + %ag1_13 = scalar.fmaf %gp1_13, %dl1, %ag0_13 : f32 + %qp1_13 = vector.extract %qkw1_1_b[1] : vector<4xf32> -> f32 + %aq1_13 = scalar.fmaf %qp1_13, %dl1, %aq0_13 : f32 + %gp1_14 = vector.extract %grw1_1_b[2] : vector<4xf32> -> f32 + %ag1_14 = scalar.fmaf %gp1_14, %dl1, %ag0_14 : f32 + %qp1_14 = vector.extract %qkw1_1_b[2] : vector<4xf32> -> f32 + %aq1_14 = scalar.fmaf %qp1_14, %dl1, %aq0_14 : f32 + %gp1_15 = vector.extract %grw1_1_b[3] : vector<4xf32> -> f32 + %ag1_15 = scalar.fmaf %gp1_15, %dl1, %ag0_15 : f32 + %qp1_15 = vector.extract %qkw1_1_b[3] : vector<4xf32> -> f32 + %aq1_15 = scalar.fmaf %qp1_15, %dl1, %aq0_15 : f32 + %gq2 = vector.load %gs_flat[%k8] : view<80xf32> -> vector<4xf32> + %bv2 = vector.extract %gq2[0] : vector<4xf32> -> f32 + %cp2 = vector.extract %gq2[1] : vector<4xf32> -> f32 + %lc2 = vector.extract %gq2[3] : vector<4xf32> -> f32 + %ev2 = vector.extract %gq2[2] : vector<4xf32> -> f32 + %tk2_r = index.add %t0, %k2 : index + %tk2_ok = index.cmp ult, %tk2_r, %n_tokens : index + %tk2_s = scf.select %tk2_ok, %tk2_r, %c0 : index + %tk2 = index.assume %tk2_s [range(%tk2_s, 0, 1048575)] : index + %av2_i = index.add %tile_row, %k2 : index + %av2_b = index.assume %av2_i [range(%av2_i, 0, 271)] : index + %av2 = view.load %a_flat[%av2_b] : view<272xf32> -> f32 + %bvv2 = view.load %b_flat[%av2_b] : view<272xf32> -> f32 + %vrd2_a = index.add %rec_col, %k256 : index + %vrd2 = index.assume %vrd2_a [range(%vrd2_a, 0, 2047)] : index + %vval2 = view.load %vst_flat[%vrd2] : view<2048xf32> -> f32 + %adec2 = scalar.fmaf %av2, %cp2, %ag1_2 : f32 + %vsub2 = scalar.subf %vval2, %adec2 : f32 + %active_index2 = index.add %t0, %k2 : index + %row_in_range2 = index.cmp ult, %active_index2, %n_tokens : index + %active2 = scalar.ori %row_in_range2, %all_recurrence_rows : i1 + %dl2 = scf.if %active2 -> (f32) { + %value = scalar.mulf %vsub2, %bv2 : f32 + scf.yield %value : f32 + } else { + %padded = scalar.mulf %vval2, %bv2 : f32 + scf.yield %padded : f32 + } + %dh2 = scalar.mulf %dl2, %ev2 : f32 + %grw2_0_a = vector.load %gr_flat[%k32] : view<256xf32> -> vector<4xf32> + %grw2_0_bi = index.constant 36 : index + %grw2_0_b = vector.load %gr_flat[%grw2_0_bi] : view<256xf32> -> vector<4xf32> + %grw2_1_a = vector.load %gr_flat[%k40] : view<256xf32> -> vector<4xf32> + %grw2_1_bi = index.constant 44 : index + %grw2_1_b = vector.load %gr_flat[%grw2_1_bi] : view<256xf32> -> vector<4xf32> + %qkw2_0_a = vector.load %qkt_flat[%k32] : view<256xf32> -> vector<4xf32> + %qkw2_0_bi = index.constant 36 : index + %qkw2_0_b = vector.load %qkt_flat[%qkw2_0_bi] : view<256xf32> -> vector<4xf32> + %qkw2_1_a = vector.load %qkt_flat[%k40] : view<256xf32> -> vector<4xf32> + %qkw2_1_bi = index.constant 44 : index + %qkw2_1_b = vector.load %qkt_flat[%qkw2_1_bi] : view<256xf32> -> vector<4xf32> + %qkd2_2 = vector.extract %qkw2_0_a[2] : vector<4xf32> -> f32 + %qsum2 = scalar.fmaf %qkd2_2, %dl2, %aq1_2 : f32 + %ao2_r = scalar.fmaf %bvv2, %cp2, %qsum2 : f32 + %ao2 = scf.if %active2 -> (f32) { + %value = scalar.mulf %ao2_r, %scale : f32 + scf.yield %value : f32 + } else { + %zero = scalar.constant 0.0 : f32 + scf.yield %zero : f32 + } + %aoi2_a = index.add %rec_col, %k256 : index + %aoi2 = index.assume %aoi2_a [range(%aoi2_a, 0, 2047)] : index + %ddh2 = scalar.fptrunc %dh2 : f32 to f16 + %publish_active2 = scalar.andi %lane_live, %active2 : i1 + scf.if %publish_active2 { + view.store %ao2, %aost_flat[%aoi2] : f32, view<2048xf32> + } + %gp2_3 = vector.extract %grw2_0_a[3] : vector<4xf32> -> f32 + %ag2_3 = scalar.fmaf %gp2_3, %dl2, %ag1_3 : f32 + %qp2_3 = vector.extract %qkw2_0_a[3] : vector<4xf32> -> f32 + %aq2_3 = scalar.fmaf %qp2_3, %dl2, %aq1_3 : f32 + %gp2_4 = vector.extract %grw2_0_b[0] : vector<4xf32> -> f32 + %ag2_4 = scalar.fmaf %gp2_4, %dl2, %ag1_4 : f32 + %qp2_4 = vector.extract %qkw2_0_b[0] : vector<4xf32> -> f32 + %aq2_4 = scalar.fmaf %qp2_4, %dl2, %aq1_4 : f32 + %gp2_5 = vector.extract %grw2_0_b[1] : vector<4xf32> -> f32 + %ag2_5 = scalar.fmaf %gp2_5, %dl2, %ag1_5 : f32 + %qp2_5 = vector.extract %qkw2_0_b[1] : vector<4xf32> -> f32 + %aq2_5 = scalar.fmaf %qp2_5, %dl2, %aq1_5 : f32 + %gp2_6 = vector.extract %grw2_0_b[2] : vector<4xf32> -> f32 + %ag2_6 = scalar.fmaf %gp2_6, %dl2, %ag1_6 : f32 + %qp2_6 = vector.extract %qkw2_0_b[2] : vector<4xf32> -> f32 + %aq2_6 = scalar.fmaf %qp2_6, %dl2, %aq1_6 : f32 + %gp2_7 = vector.extract %grw2_0_b[3] : vector<4xf32> -> f32 + %ag2_7 = scalar.fmaf %gp2_7, %dl2, %ag1_7 : f32 + %qp2_7 = vector.extract %qkw2_0_b[3] : vector<4xf32> -> f32 + %aq2_7 = scalar.fmaf %qp2_7, %dl2, %aq1_7 : f32 + %gp2_8 = vector.extract %grw2_1_a[0] : vector<4xf32> -> f32 + %ag2_8 = scalar.fmaf %gp2_8, %dl2, %ag1_8 : f32 + %qp2_8 = vector.extract %qkw2_1_a[0] : vector<4xf32> -> f32 + %aq2_8 = scalar.fmaf %qp2_8, %dl2, %aq1_8 : f32 + %gp2_9 = vector.extract %grw2_1_a[1] : vector<4xf32> -> f32 + %ag2_9 = scalar.fmaf %gp2_9, %dl2, %ag1_9 : f32 + %qp2_9 = vector.extract %qkw2_1_a[1] : vector<4xf32> -> f32 + %aq2_9 = scalar.fmaf %qp2_9, %dl2, %aq1_9 : f32 + %gp2_10 = vector.extract %grw2_1_a[2] : vector<4xf32> -> f32 + %ag2_10 = scalar.fmaf %gp2_10, %dl2, %ag1_10 : f32 + %qp2_10 = vector.extract %qkw2_1_a[2] : vector<4xf32> -> f32 + %aq2_10 = scalar.fmaf %qp2_10, %dl2, %aq1_10 : f32 + %gp2_11 = vector.extract %grw2_1_a[3] : vector<4xf32> -> f32 + %ag2_11 = scalar.fmaf %gp2_11, %dl2, %ag1_11 : f32 + %qp2_11 = vector.extract %qkw2_1_a[3] : vector<4xf32> -> f32 + %aq2_11 = scalar.fmaf %qp2_11, %dl2, %aq1_11 : f32 + %gp2_12 = vector.extract %grw2_1_b[0] : vector<4xf32> -> f32 + %ag2_12 = scalar.fmaf %gp2_12, %dl2, %ag1_12 : f32 + %qp2_12 = vector.extract %qkw2_1_b[0] : vector<4xf32> -> f32 + %aq2_12 = scalar.fmaf %qp2_12, %dl2, %aq1_12 : f32 + %gp2_13 = vector.extract %grw2_1_b[1] : vector<4xf32> -> f32 + %ag2_13 = scalar.fmaf %gp2_13, %dl2, %ag1_13 : f32 + %qp2_13 = vector.extract %qkw2_1_b[1] : vector<4xf32> -> f32 + %aq2_13 = scalar.fmaf %qp2_13, %dl2, %aq1_13 : f32 + %gp2_14 = vector.extract %grw2_1_b[2] : vector<4xf32> -> f32 + %ag2_14 = scalar.fmaf %gp2_14, %dl2, %ag1_14 : f32 + %qp2_14 = vector.extract %qkw2_1_b[2] : vector<4xf32> -> f32 + %aq2_14 = scalar.fmaf %qp2_14, %dl2, %aq1_14 : f32 + %gp2_15 = vector.extract %grw2_1_b[3] : vector<4xf32> -> f32 + %ag2_15 = scalar.fmaf %gp2_15, %dl2, %ag1_15 : f32 + %qp2_15 = vector.extract %qkw2_1_b[3] : vector<4xf32> -> f32 + %aq2_15 = scalar.fmaf %qp2_15, %dl2, %aq1_15 : f32 + %gq3 = vector.load %gs_flat[%k12] : view<80xf32> -> vector<4xf32> + %bv3 = vector.extract %gq3[0] : vector<4xf32> -> f32 + %cp3 = vector.extract %gq3[1] : vector<4xf32> -> f32 + %lc3 = vector.extract %gq3[3] : vector<4xf32> -> f32 + %ev3 = vector.extract %gq3[2] : vector<4xf32> -> f32 + %tk3_r = index.add %t0, %k3 : index + %tk3_ok = index.cmp ult, %tk3_r, %n_tokens : index + %tk3_s = scf.select %tk3_ok, %tk3_r, %c0 : index + %tk3 = index.assume %tk3_s [range(%tk3_s, 0, 1048575)] : index + %av3_i = index.add %tile_row, %k3 : index + %av3_b = index.assume %av3_i [range(%av3_i, 0, 271)] : index + %av3 = view.load %a_flat[%av3_b] : view<272xf32> -> f32 + %bvv3 = view.load %b_flat[%av3_b] : view<272xf32> -> f32 + %vrd3_a = index.add %rec_col, %k384 : index + %vrd3 = index.assume %vrd3_a [range(%vrd3_a, 0, 2047)] : index + %vval3 = view.load %vst_flat[%vrd3] : view<2048xf32> -> f32 + %adec3 = scalar.fmaf %av3, %cp3, %ag2_3 : f32 + %vsub3 = scalar.subf %vval3, %adec3 : f32 + %active_index3 = index.add %t0, %k3 : index + %row_in_range3 = index.cmp ult, %active_index3, %n_tokens : index + %active3 = scalar.ori %row_in_range3, %all_recurrence_rows : i1 + %dl3 = scf.if %active3 -> (f32) { + %value = scalar.mulf %vsub3, %bv3 : f32 + scf.yield %value : f32 + } else { + %padded = scalar.mulf %vval3, %bv3 : f32 + scf.yield %padded : f32 + } + %dh3 = scalar.mulf %dl3, %ev3 : f32 + %grw3_0_a = vector.load %gr_flat[%k48] : view<256xf32> -> vector<4xf32> + %grw3_0_bi = index.constant 52 : index + %grw3_0_b = vector.load %gr_flat[%grw3_0_bi] : view<256xf32> -> vector<4xf32> + %grw3_1_a = vector.load %gr_flat[%k56] : view<256xf32> -> vector<4xf32> + %grw3_1_bi = index.constant 60 : index + %grw3_1_b = vector.load %gr_flat[%grw3_1_bi] : view<256xf32> -> vector<4xf32> + %qkw3_0_a = vector.load %qkt_flat[%k48] : view<256xf32> -> vector<4xf32> + %qkw3_0_bi = index.constant 52 : index + %qkw3_0_b = vector.load %qkt_flat[%qkw3_0_bi] : view<256xf32> -> vector<4xf32> + %qkw3_1_a = vector.load %qkt_flat[%k56] : view<256xf32> -> vector<4xf32> + %qkw3_1_bi = index.constant 60 : index + %qkw3_1_b = vector.load %qkt_flat[%qkw3_1_bi] : view<256xf32> -> vector<4xf32> + %qkd3_3 = vector.extract %qkw3_0_a[3] : vector<4xf32> -> f32 + %qsum3 = scalar.fmaf %qkd3_3, %dl3, %aq2_3 : f32 + %ao3_r = scalar.fmaf %bvv3, %cp3, %qsum3 : f32 + %ao3 = scf.if %active3 -> (f32) { + %value = scalar.mulf %ao3_r, %scale : f32 + scf.yield %value : f32 + } else { + %zero = scalar.constant 0.0 : f32 + scf.yield %zero : f32 + } + %aoi3_a = index.add %rec_col, %k384 : index + %aoi3 = index.assume %aoi3_a [range(%aoi3_a, 0, 2047)] : index + %ddh3 = scalar.fptrunc %dh3 : f32 to f16 + %publish_active3 = scalar.andi %lane_live, %active3 : i1 + scf.if %publish_active3 { + view.store %ao3, %aost_flat[%aoi3] : f32, view<2048xf32> + } + %gp3_4 = vector.extract %grw3_0_b[0] : vector<4xf32> -> f32 + %ag3_4 = scalar.fmaf %gp3_4, %dl3, %ag2_4 : f32 + %qp3_4 = vector.extract %qkw3_0_b[0] : vector<4xf32> -> f32 + %aq3_4 = scalar.fmaf %qp3_4, %dl3, %aq2_4 : f32 + %gp3_5 = vector.extract %grw3_0_b[1] : vector<4xf32> -> f32 + %ag3_5 = scalar.fmaf %gp3_5, %dl3, %ag2_5 : f32 + %qp3_5 = vector.extract %qkw3_0_b[1] : vector<4xf32> -> f32 + %aq3_5 = scalar.fmaf %qp3_5, %dl3, %aq2_5 : f32 + %gp3_6 = vector.extract %grw3_0_b[2] : vector<4xf32> -> f32 + %ag3_6 = scalar.fmaf %gp3_6, %dl3, %ag2_6 : f32 + %qp3_6 = vector.extract %qkw3_0_b[2] : vector<4xf32> -> f32 + %aq3_6 = scalar.fmaf %qp3_6, %dl3, %aq2_6 : f32 + %gp3_7 = vector.extract %grw3_0_b[3] : vector<4xf32> -> f32 + %ag3_7 = scalar.fmaf %gp3_7, %dl3, %ag2_7 : f32 + %qp3_7 = vector.extract %qkw3_0_b[3] : vector<4xf32> -> f32 + %aq3_7 = scalar.fmaf %qp3_7, %dl3, %aq2_7 : f32 + %gp3_8 = vector.extract %grw3_1_a[0] : vector<4xf32> -> f32 + %ag3_8 = scalar.fmaf %gp3_8, %dl3, %ag2_8 : f32 + %qp3_8 = vector.extract %qkw3_1_a[0] : vector<4xf32> -> f32 + %aq3_8 = scalar.fmaf %qp3_8, %dl3, %aq2_8 : f32 + %gp3_9 = vector.extract %grw3_1_a[1] : vector<4xf32> -> f32 + %ag3_9 = scalar.fmaf %gp3_9, %dl3, %ag2_9 : f32 + %qp3_9 = vector.extract %qkw3_1_a[1] : vector<4xf32> -> f32 + %aq3_9 = scalar.fmaf %qp3_9, %dl3, %aq2_9 : f32 + %gp3_10 = vector.extract %grw3_1_a[2] : vector<4xf32> -> f32 + %ag3_10 = scalar.fmaf %gp3_10, %dl3, %ag2_10 : f32 + %qp3_10 = vector.extract %qkw3_1_a[2] : vector<4xf32> -> f32 + %aq3_10 = scalar.fmaf %qp3_10, %dl3, %aq2_10 : f32 + %gp3_11 = vector.extract %grw3_1_a[3] : vector<4xf32> -> f32 + %ag3_11 = scalar.fmaf %gp3_11, %dl3, %ag2_11 : f32 + %qp3_11 = vector.extract %qkw3_1_a[3] : vector<4xf32> -> f32 + %aq3_11 = scalar.fmaf %qp3_11, %dl3, %aq2_11 : f32 + %gp3_12 = vector.extract %grw3_1_b[0] : vector<4xf32> -> f32 + %ag3_12 = scalar.fmaf %gp3_12, %dl3, %ag2_12 : f32 + %qp3_12 = vector.extract %qkw3_1_b[0] : vector<4xf32> -> f32 + %aq3_12 = scalar.fmaf %qp3_12, %dl3, %aq2_12 : f32 + %gp3_13 = vector.extract %grw3_1_b[1] : vector<4xf32> -> f32 + %ag3_13 = scalar.fmaf %gp3_13, %dl3, %ag2_13 : f32 + %qp3_13 = vector.extract %qkw3_1_b[1] : vector<4xf32> -> f32 + %aq3_13 = scalar.fmaf %qp3_13, %dl3, %aq2_13 : f32 + %gp3_14 = vector.extract %grw3_1_b[2] : vector<4xf32> -> f32 + %ag3_14 = scalar.fmaf %gp3_14, %dl3, %ag2_14 : f32 + %qp3_14 = vector.extract %qkw3_1_b[2] : vector<4xf32> -> f32 + %aq3_14 = scalar.fmaf %qp3_14, %dl3, %aq2_14 : f32 + %gp3_15 = vector.extract %grw3_1_b[3] : vector<4xf32> -> f32 + %ag3_15 = scalar.fmaf %gp3_15, %dl3, %ag2_15 : f32 + %qp3_15 = vector.extract %qkw3_1_b[3] : vector<4xf32> -> f32 + %aq3_15 = scalar.fmaf %qp3_15, %dl3, %aq2_15 : f32 + %gq4 = vector.load %gs_flat[%k16] : view<80xf32> -> vector<4xf32> + %bv4 = vector.extract %gq4[0] : vector<4xf32> -> f32 + %cp4 = vector.extract %gq4[1] : vector<4xf32> -> f32 + %ev4 = vector.extract %gq4[2] : vector<4xf32> -> f32 + %tk4_r = index.add %t0, %k4 : index + %tk4_ok = index.cmp ult, %tk4_r, %n_tokens : index + %tk4_s = scf.select %tk4_ok, %tk4_r, %c0 : index + %tk4 = index.assume %tk4_s [range(%tk4_s, 0, 1048575)] : index + %av4_i = index.add %tile_row, %k4 : index + %av4_b = index.assume %av4_i [range(%av4_i, 0, 271)] : index + %av4 = view.load %a_flat[%av4_b] : view<272xf32> -> f32 + %bvv4 = view.load %b_flat[%av4_b] : view<272xf32> -> f32 + %vrd4_a = index.add %rec_col, %k512 : index + %vrd4 = index.assume %vrd4_a [range(%vrd4_a, 0, 2047)] : index + %vval4 = view.load %vst_flat[%vrd4] : view<2048xf32> -> f32 + %adec4 = scalar.fmaf %av4, %cp4, %ag3_4 : f32 + %vsub4 = scalar.subf %vval4, %adec4 : f32 + %active_index4 = index.add %t0, %k4 : index + %row_in_range4 = index.cmp ult, %active_index4, %n_tokens : index + %active4 = scalar.ori %row_in_range4, %all_recurrence_rows : i1 + %dl4 = scf.if %active4 -> (f32) { + %value = scalar.mulf %vsub4, %bv4 : f32 + scf.yield %value : f32 + } else { + %padded = scalar.mulf %vval4, %bv4 : f32 + scf.yield %padded : f32 + } + %dh4 = scalar.mulf %dl4, %ev4 : f32 + %grw4_0_a = vector.load %gr_flat[%k64] : view<256xf32> -> vector<4xf32> + %grw4_0_bi = index.constant 68 : index + %grw4_0_b = vector.load %gr_flat[%grw4_0_bi] : view<256xf32> -> vector<4xf32> + %grw4_1_a = vector.load %gr_flat[%k72] : view<256xf32> -> vector<4xf32> + %grw4_1_bi = index.constant 76 : index + %grw4_1_b = vector.load %gr_flat[%grw4_1_bi] : view<256xf32> -> vector<4xf32> + %qkw4_0_a = vector.load %qkt_flat[%k64] : view<256xf32> -> vector<4xf32> + %qkw4_0_bi = index.constant 68 : index + %qkw4_0_b = vector.load %qkt_flat[%qkw4_0_bi] : view<256xf32> -> vector<4xf32> + %qkw4_1_a = vector.load %qkt_flat[%k72] : view<256xf32> -> vector<4xf32> + %qkw4_1_bi = index.constant 76 : index + %qkw4_1_b = vector.load %qkt_flat[%qkw4_1_bi] : view<256xf32> -> vector<4xf32> + %qkd4_4 = vector.extract %qkw4_0_b[0] : vector<4xf32> -> f32 + %qsum4 = scalar.fmaf %qkd4_4, %dl4, %aq3_4 : f32 + %ao4_r = scalar.fmaf %bvv4, %cp4, %qsum4 : f32 + %ao4 = scf.if %active4 -> (f32) { + %value = scalar.mulf %ao4_r, %scale : f32 + scf.yield %value : f32 + } else { + %zero = scalar.constant 0.0 : f32 + scf.yield %zero : f32 + } + %aoi4_a = index.add %rec_col, %k512 : index + %aoi4 = index.assume %aoi4_a [range(%aoi4_a, 0, 2047)] : index + %ddh4 = scalar.fptrunc %dh4 : f32 to f16 + %publish_active4 = scalar.andi %lane_live, %active4 : i1 + scf.if %publish_active4 { + view.store %ao4, %aost_flat[%aoi4] : f32, view<2048xf32> + } + %gp4_5 = vector.extract %grw4_0_b[1] : vector<4xf32> -> f32 + %ag4_5 = scalar.fmaf %gp4_5, %dl4, %ag3_5 : f32 + %qp4_5 = vector.extract %qkw4_0_b[1] : vector<4xf32> -> f32 + %aq4_5 = scalar.fmaf %qp4_5, %dl4, %aq3_5 : f32 + %gp4_6 = vector.extract %grw4_0_b[2] : vector<4xf32> -> f32 + %ag4_6 = scalar.fmaf %gp4_6, %dl4, %ag3_6 : f32 + %qp4_6 = vector.extract %qkw4_0_b[2] : vector<4xf32> -> f32 + %aq4_6 = scalar.fmaf %qp4_6, %dl4, %aq3_6 : f32 + %gp4_7 = vector.extract %grw4_0_b[3] : vector<4xf32> -> f32 + %ag4_7 = scalar.fmaf %gp4_7, %dl4, %ag3_7 : f32 + %qp4_7 = vector.extract %qkw4_0_b[3] : vector<4xf32> -> f32 + %aq4_7 = scalar.fmaf %qp4_7, %dl4, %aq3_7 : f32 + %gp4_8 = vector.extract %grw4_1_a[0] : vector<4xf32> -> f32 + %ag4_8 = scalar.fmaf %gp4_8, %dl4, %ag3_8 : f32 + %qp4_8 = vector.extract %qkw4_1_a[0] : vector<4xf32> -> f32 + %aq4_8 = scalar.fmaf %qp4_8, %dl4, %aq3_8 : f32 + %gp4_9 = vector.extract %grw4_1_a[1] : vector<4xf32> -> f32 + %ag4_9 = scalar.fmaf %gp4_9, %dl4, %ag3_9 : f32 + %qp4_9 = vector.extract %qkw4_1_a[1] : vector<4xf32> -> f32 + %aq4_9 = scalar.fmaf %qp4_9, %dl4, %aq3_9 : f32 + %gp4_10 = vector.extract %grw4_1_a[2] : vector<4xf32> -> f32 + %ag4_10 = scalar.fmaf %gp4_10, %dl4, %ag3_10 : f32 + %qp4_10 = vector.extract %qkw4_1_a[2] : vector<4xf32> -> f32 + %aq4_10 = scalar.fmaf %qp4_10, %dl4, %aq3_10 : f32 + %gp4_11 = vector.extract %grw4_1_a[3] : vector<4xf32> -> f32 + %ag4_11 = scalar.fmaf %gp4_11, %dl4, %ag3_11 : f32 + %qp4_11 = vector.extract %qkw4_1_a[3] : vector<4xf32> -> f32 + %aq4_11 = scalar.fmaf %qp4_11, %dl4, %aq3_11 : f32 + %gp4_12 = vector.extract %grw4_1_b[0] : vector<4xf32> -> f32 + %ag4_12 = scalar.fmaf %gp4_12, %dl4, %ag3_12 : f32 + %qp4_12 = vector.extract %qkw4_1_b[0] : vector<4xf32> -> f32 + %aq4_12 = scalar.fmaf %qp4_12, %dl4, %aq3_12 : f32 + %gp4_13 = vector.extract %grw4_1_b[1] : vector<4xf32> -> f32 + %ag4_13 = scalar.fmaf %gp4_13, %dl4, %ag3_13 : f32 + %qp4_13 = vector.extract %qkw4_1_b[1] : vector<4xf32> -> f32 + %aq4_13 = scalar.fmaf %qp4_13, %dl4, %aq3_13 : f32 + %gp4_14 = vector.extract %grw4_1_b[2] : vector<4xf32> -> f32 + %ag4_14 = scalar.fmaf %gp4_14, %dl4, %ag3_14 : f32 + %qp4_14 = vector.extract %qkw4_1_b[2] : vector<4xf32> -> f32 + %aq4_14 = scalar.fmaf %qp4_14, %dl4, %aq3_14 : f32 + %gp4_15 = vector.extract %grw4_1_b[3] : vector<4xf32> -> f32 + %ag4_15 = scalar.fmaf %gp4_15, %dl4, %ag3_15 : f32 + %qp4_15 = vector.extract %qkw4_1_b[3] : vector<4xf32> -> f32 + %aq4_15 = scalar.fmaf %qp4_15, %dl4, %aq3_15 : f32 + %gq5 = vector.load %gs_flat[%k20] : view<80xf32> -> vector<4xf32> + %bv5 = vector.extract %gq5[0] : vector<4xf32> -> f32 + %cp5 = vector.extract %gq5[1] : vector<4xf32> -> f32 + %ev5 = vector.extract %gq5[2] : vector<4xf32> -> f32 + %tk5_r = index.add %t0, %k5 : index + %tk5_ok = index.cmp ult, %tk5_r, %n_tokens : index + %tk5_s = scf.select %tk5_ok, %tk5_r, %c0 : index + %tk5 = index.assume %tk5_s [range(%tk5_s, 0, 1048575)] : index + %av5_i = index.add %tile_row, %k5 : index + %av5_b = index.assume %av5_i [range(%av5_i, 0, 271)] : index + %av5 = view.load %a_flat[%av5_b] : view<272xf32> -> f32 + %bvv5 = view.load %b_flat[%av5_b] : view<272xf32> -> f32 + %vrd5_a = index.add %rec_col, %k640 : index + %vrd5 = index.assume %vrd5_a [range(%vrd5_a, 0, 2047)] : index + %vval5 = view.load %vst_flat[%vrd5] : view<2048xf32> -> f32 + %adec5 = scalar.fmaf %av5, %cp5, %ag4_5 : f32 + %vsub5 = scalar.subf %vval5, %adec5 : f32 + %active_index5 = index.add %t0, %k5 : index + %row_in_range5 = index.cmp ult, %active_index5, %n_tokens : index + %active5 = scalar.ori %row_in_range5, %all_recurrence_rows : i1 + %dl5 = scf.if %active5 -> (f32) { + %value = scalar.mulf %vsub5, %bv5 : f32 + scf.yield %value : f32 + } else { + %padded = scalar.mulf %vval5, %bv5 : f32 + scf.yield %padded : f32 + } + %dh5 = scalar.mulf %dl5, %ev5 : f32 + %grw5_0_a = vector.load %gr_flat[%k80] : view<256xf32> -> vector<4xf32> + %grw5_0_bi = index.constant 84 : index + %grw5_0_b = vector.load %gr_flat[%grw5_0_bi] : view<256xf32> -> vector<4xf32> + %grw5_1_a = vector.load %gr_flat[%k88] : view<256xf32> -> vector<4xf32> + %grw5_1_bi = index.constant 92 : index + %grw5_1_b = vector.load %gr_flat[%grw5_1_bi] : view<256xf32> -> vector<4xf32> + %qkw5_0_a = vector.load %qkt_flat[%k80] : view<256xf32> -> vector<4xf32> + %qkw5_0_bi = index.constant 84 : index + %qkw5_0_b = vector.load %qkt_flat[%qkw5_0_bi] : view<256xf32> -> vector<4xf32> + %qkw5_1_a = vector.load %qkt_flat[%k88] : view<256xf32> -> vector<4xf32> + %qkw5_1_bi = index.constant 92 : index + %qkw5_1_b = vector.load %qkt_flat[%qkw5_1_bi] : view<256xf32> -> vector<4xf32> + %qkd5_5 = vector.extract %qkw5_0_b[1] : vector<4xf32> -> f32 + %qsum5 = scalar.fmaf %qkd5_5, %dl5, %aq4_5 : f32 + %ao5_r = scalar.fmaf %bvv5, %cp5, %qsum5 : f32 + %ao5 = scf.if %active5 -> (f32) { + %value = scalar.mulf %ao5_r, %scale : f32 + scf.yield %value : f32 + } else { + %zero = scalar.constant 0.0 : f32 + scf.yield %zero : f32 + } + %aoi5_a = index.add %rec_col, %k640 : index + %aoi5 = index.assume %aoi5_a [range(%aoi5_a, 0, 2047)] : index + %ddh5 = scalar.fptrunc %dh5 : f32 to f16 + %publish_active5 = scalar.andi %lane_live, %active5 : i1 + scf.if %publish_active5 { + view.store %ao5, %aost_flat[%aoi5] : f32, view<2048xf32> + } + %gp5_6 = vector.extract %grw5_0_b[2] : vector<4xf32> -> f32 + %ag5_6 = scalar.fmaf %gp5_6, %dl5, %ag4_6 : f32 + %qp5_6 = vector.extract %qkw5_0_b[2] : vector<4xf32> -> f32 + %aq5_6 = scalar.fmaf %qp5_6, %dl5, %aq4_6 : f32 + %gp5_7 = vector.extract %grw5_0_b[3] : vector<4xf32> -> f32 + %ag5_7 = scalar.fmaf %gp5_7, %dl5, %ag4_7 : f32 + %qp5_7 = vector.extract %qkw5_0_b[3] : vector<4xf32> -> f32 + %aq5_7 = scalar.fmaf %qp5_7, %dl5, %aq4_7 : f32 + %gp5_8 = vector.extract %grw5_1_a[0] : vector<4xf32> -> f32 + %ag5_8 = scalar.fmaf %gp5_8, %dl5, %ag4_8 : f32 + %qp5_8 = vector.extract %qkw5_1_a[0] : vector<4xf32> -> f32 + %aq5_8 = scalar.fmaf %qp5_8, %dl5, %aq4_8 : f32 + %gp5_9 = vector.extract %grw5_1_a[1] : vector<4xf32> -> f32 + %ag5_9 = scalar.fmaf %gp5_9, %dl5, %ag4_9 : f32 + %qp5_9 = vector.extract %qkw5_1_a[1] : vector<4xf32> -> f32 + %aq5_9 = scalar.fmaf %qp5_9, %dl5, %aq4_9 : f32 + %gp5_10 = vector.extract %grw5_1_a[2] : vector<4xf32> -> f32 + %ag5_10 = scalar.fmaf %gp5_10, %dl5, %ag4_10 : f32 + %qp5_10 = vector.extract %qkw5_1_a[2] : vector<4xf32> -> f32 + %aq5_10 = scalar.fmaf %qp5_10, %dl5, %aq4_10 : f32 + %gp5_11 = vector.extract %grw5_1_a[3] : vector<4xf32> -> f32 + %ag5_11 = scalar.fmaf %gp5_11, %dl5, %ag4_11 : f32 + %qp5_11 = vector.extract %qkw5_1_a[3] : vector<4xf32> -> f32 + %aq5_11 = scalar.fmaf %qp5_11, %dl5, %aq4_11 : f32 + %gp5_12 = vector.extract %grw5_1_b[0] : vector<4xf32> -> f32 + %ag5_12 = scalar.fmaf %gp5_12, %dl5, %ag4_12 : f32 + %qp5_12 = vector.extract %qkw5_1_b[0] : vector<4xf32> -> f32 + %aq5_12 = scalar.fmaf %qp5_12, %dl5, %aq4_12 : f32 + %gp5_13 = vector.extract %grw5_1_b[1] : vector<4xf32> -> f32 + %ag5_13 = scalar.fmaf %gp5_13, %dl5, %ag4_13 : f32 + %qp5_13 = vector.extract %qkw5_1_b[1] : vector<4xf32> -> f32 + %aq5_13 = scalar.fmaf %qp5_13, %dl5, %aq4_13 : f32 + %gp5_14 = vector.extract %grw5_1_b[2] : vector<4xf32> -> f32 + %ag5_14 = scalar.fmaf %gp5_14, %dl5, %ag4_14 : f32 + %qp5_14 = vector.extract %qkw5_1_b[2] : vector<4xf32> -> f32 + %aq5_14 = scalar.fmaf %qp5_14, %dl5, %aq4_14 : f32 + %gp5_15 = vector.extract %grw5_1_b[3] : vector<4xf32> -> f32 + %ag5_15 = scalar.fmaf %gp5_15, %dl5, %ag4_15 : f32 + %qp5_15 = vector.extract %qkw5_1_b[3] : vector<4xf32> -> f32 + %aq5_15 = scalar.fmaf %qp5_15, %dl5, %aq4_15 : f32 + %gq6 = vector.load %gs_flat[%k24] : view<80xf32> -> vector<4xf32> + %bv6 = vector.extract %gq6[0] : vector<4xf32> -> f32 + %cp6 = vector.extract %gq6[1] : vector<4xf32> -> f32 + %ev6 = vector.extract %gq6[2] : vector<4xf32> -> f32 + %tk6_r = index.add %t0, %k6 : index + %tk6_ok = index.cmp ult, %tk6_r, %n_tokens : index + %tk6_s = scf.select %tk6_ok, %tk6_r, %c0 : index + %tk6 = index.assume %tk6_s [range(%tk6_s, 0, 1048575)] : index + %av6_i = index.add %tile_row, %k6 : index + %av6_b = index.assume %av6_i [range(%av6_i, 0, 271)] : index + %av6 = view.load %a_flat[%av6_b] : view<272xf32> -> f32 + %bvv6 = view.load %b_flat[%av6_b] : view<272xf32> -> f32 + %vrd6_a = index.add %rec_col, %k768 : index + %vrd6 = index.assume %vrd6_a [range(%vrd6_a, 0, 2047)] : index + %vval6 = view.load %vst_flat[%vrd6] : view<2048xf32> -> f32 + %adec6 = scalar.fmaf %av6, %cp6, %ag5_6 : f32 + %vsub6 = scalar.subf %vval6, %adec6 : f32 + %active_index6 = index.add %t0, %k6 : index + %row_in_range6 = index.cmp ult, %active_index6, %n_tokens : index + %active6 = scalar.ori %row_in_range6, %all_recurrence_rows : i1 + %dl6 = scf.if %active6 -> (f32) { + %value = scalar.mulf %vsub6, %bv6 : f32 + scf.yield %value : f32 + } else { + %padded = scalar.mulf %vval6, %bv6 : f32 + scf.yield %padded : f32 + } + %dh6 = scalar.mulf %dl6, %ev6 : f32 + %grw6_0_a = vector.load %gr_flat[%k96] : view<256xf32> -> vector<4xf32> + %grw6_0_bi = index.constant 100 : index + %grw6_0_b = vector.load %gr_flat[%grw6_0_bi] : view<256xf32> -> vector<4xf32> + %grw6_1_a = vector.load %gr_flat[%k104] : view<256xf32> -> vector<4xf32> + %grw6_1_bi = index.constant 108 : index + %grw6_1_b = vector.load %gr_flat[%grw6_1_bi] : view<256xf32> -> vector<4xf32> + %qkw6_0_a = vector.load %qkt_flat[%k96] : view<256xf32> -> vector<4xf32> + %qkw6_0_bi = index.constant 100 : index + %qkw6_0_b = vector.load %qkt_flat[%qkw6_0_bi] : view<256xf32> -> vector<4xf32> + %qkw6_1_a = vector.load %qkt_flat[%k104] : view<256xf32> -> vector<4xf32> + %qkw6_1_bi = index.constant 108 : index + %qkw6_1_b = vector.load %qkt_flat[%qkw6_1_bi] : view<256xf32> -> vector<4xf32> + %qkd6_6 = vector.extract %qkw6_0_b[2] : vector<4xf32> -> f32 + %qsum6 = scalar.fmaf %qkd6_6, %dl6, %aq5_6 : f32 + %ao6_r = scalar.fmaf %bvv6, %cp6, %qsum6 : f32 + %ao6 = scf.if %active6 -> (f32) { + %value = scalar.mulf %ao6_r, %scale : f32 + scf.yield %value : f32 + } else { + %zero = scalar.constant 0.0 : f32 + scf.yield %zero : f32 + } + %aoi6_a = index.add %rec_col, %k768 : index + %aoi6 = index.assume %aoi6_a [range(%aoi6_a, 0, 2047)] : index + %ddh6 = scalar.fptrunc %dh6 : f32 to f16 + %publish_active6 = scalar.andi %lane_live, %active6 : i1 + scf.if %publish_active6 { + view.store %ao6, %aost_flat[%aoi6] : f32, view<2048xf32> + } + %gp6_7 = vector.extract %grw6_0_b[3] : vector<4xf32> -> f32 + %ag6_7 = scalar.fmaf %gp6_7, %dl6, %ag5_7 : f32 + %qp6_7 = vector.extract %qkw6_0_b[3] : vector<4xf32> -> f32 + %aq6_7 = scalar.fmaf %qp6_7, %dl6, %aq5_7 : f32 + %gp6_8 = vector.extract %grw6_1_a[0] : vector<4xf32> -> f32 + %ag6_8 = scalar.fmaf %gp6_8, %dl6, %ag5_8 : f32 + %qp6_8 = vector.extract %qkw6_1_a[0] : vector<4xf32> -> f32 + %aq6_8 = scalar.fmaf %qp6_8, %dl6, %aq5_8 : f32 + %gp6_9 = vector.extract %grw6_1_a[1] : vector<4xf32> -> f32 + %ag6_9 = scalar.fmaf %gp6_9, %dl6, %ag5_9 : f32 + %qp6_9 = vector.extract %qkw6_1_a[1] : vector<4xf32> -> f32 + %aq6_9 = scalar.fmaf %qp6_9, %dl6, %aq5_9 : f32 + %gp6_10 = vector.extract %grw6_1_a[2] : vector<4xf32> -> f32 + %ag6_10 = scalar.fmaf %gp6_10, %dl6, %ag5_10 : f32 + %qp6_10 = vector.extract %qkw6_1_a[2] : vector<4xf32> -> f32 + %aq6_10 = scalar.fmaf %qp6_10, %dl6, %aq5_10 : f32 + %gp6_11 = vector.extract %grw6_1_a[3] : vector<4xf32> -> f32 + %ag6_11 = scalar.fmaf %gp6_11, %dl6, %ag5_11 : f32 + %qp6_11 = vector.extract %qkw6_1_a[3] : vector<4xf32> -> f32 + %aq6_11 = scalar.fmaf %qp6_11, %dl6, %aq5_11 : f32 + %gp6_12 = vector.extract %grw6_1_b[0] : vector<4xf32> -> f32 + %ag6_12 = scalar.fmaf %gp6_12, %dl6, %ag5_12 : f32 + %qp6_12 = vector.extract %qkw6_1_b[0] : vector<4xf32> -> f32 + %aq6_12 = scalar.fmaf %qp6_12, %dl6, %aq5_12 : f32 + %gp6_13 = vector.extract %grw6_1_b[1] : vector<4xf32> -> f32 + %ag6_13 = scalar.fmaf %gp6_13, %dl6, %ag5_13 : f32 + %qp6_13 = vector.extract %qkw6_1_b[1] : vector<4xf32> -> f32 + %aq6_13 = scalar.fmaf %qp6_13, %dl6, %aq5_13 : f32 + %gp6_14 = vector.extract %grw6_1_b[2] : vector<4xf32> -> f32 + %ag6_14 = scalar.fmaf %gp6_14, %dl6, %ag5_14 : f32 + %qp6_14 = vector.extract %qkw6_1_b[2] : vector<4xf32> -> f32 + %aq6_14 = scalar.fmaf %qp6_14, %dl6, %aq5_14 : f32 + %gp6_15 = vector.extract %grw6_1_b[3] : vector<4xf32> -> f32 + %ag6_15 = scalar.fmaf %gp6_15, %dl6, %ag5_15 : f32 + %qp6_15 = vector.extract %qkw6_1_b[3] : vector<4xf32> -> f32 + %aq6_15 = scalar.fmaf %qp6_15, %dl6, %aq5_15 : f32 + %gq7 = vector.load %gs_flat[%k28] : view<80xf32> -> vector<4xf32> + %bv7 = vector.extract %gq7[0] : vector<4xf32> -> f32 + %cp7 = vector.extract %gq7[1] : vector<4xf32> -> f32 + %ev7 = vector.extract %gq7[2] : vector<4xf32> -> f32 + %tk7_r = index.add %t0, %k7 : index + %tk7_ok = index.cmp ult, %tk7_r, %n_tokens : index + %tk7_s = scf.select %tk7_ok, %tk7_r, %c0 : index + %tk7 = index.assume %tk7_s [range(%tk7_s, 0, 1048575)] : index + %av7_i = index.add %tile_row, %k7 : index + %av7_b = index.assume %av7_i [range(%av7_i, 0, 271)] : index + %av7 = view.load %a_flat[%av7_b] : view<272xf32> -> f32 + %bvv7 = view.load %b_flat[%av7_b] : view<272xf32> -> f32 + %vrd7_a = index.add %rec_col, %k896 : index + %vrd7 = index.assume %vrd7_a [range(%vrd7_a, 0, 2047)] : index + %vval7 = view.load %vst_flat[%vrd7] : view<2048xf32> -> f32 + %adec7 = scalar.fmaf %av7, %cp7, %ag6_7 : f32 + %vsub7 = scalar.subf %vval7, %adec7 : f32 + %active_index7 = index.add %t0, %k7 : index + %row_in_range7 = index.cmp ult, %active_index7, %n_tokens : index + %active7 = scalar.ori %row_in_range7, %all_recurrence_rows : i1 + %dl7 = scf.if %active7 -> (f32) { + %value = scalar.mulf %vsub7, %bv7 : f32 + scf.yield %value : f32 + } else { + %padded = scalar.mulf %vval7, %bv7 : f32 + scf.yield %padded : f32 + } + %dh7 = scalar.mulf %dl7, %ev7 : f32 + %grw7_0_a = vector.load %gr_flat[%k112] : view<256xf32> -> vector<4xf32> + %grw7_0_bi = index.constant 116 : index + %grw7_0_b = vector.load %gr_flat[%grw7_0_bi] : view<256xf32> -> vector<4xf32> + %grw7_1_a = vector.load %gr_flat[%k120] : view<256xf32> -> vector<4xf32> + %grw7_1_bi = index.constant 124 : index + %grw7_1_b = vector.load %gr_flat[%grw7_1_bi] : view<256xf32> -> vector<4xf32> + %qkw7_0_a = vector.load %qkt_flat[%k112] : view<256xf32> -> vector<4xf32> + %qkw7_0_bi = index.constant 116 : index + %qkw7_0_b = vector.load %qkt_flat[%qkw7_0_bi] : view<256xf32> -> vector<4xf32> + %qkw7_1_a = vector.load %qkt_flat[%k120] : view<256xf32> -> vector<4xf32> + %qkw7_1_bi = index.constant 124 : index + %qkw7_1_b = vector.load %qkt_flat[%qkw7_1_bi] : view<256xf32> -> vector<4xf32> + %qkd7_7 = vector.extract %qkw7_0_b[3] : vector<4xf32> -> f32 + %qsum7 = scalar.fmaf %qkd7_7, %dl7, %aq6_7 : f32 + %ao7_r = scalar.fmaf %bvv7, %cp7, %qsum7 : f32 + %ao7 = scf.if %active7 -> (f32) { + %value = scalar.mulf %ao7_r, %scale : f32 + scf.yield %value : f32 + } else { + %zero = scalar.constant 0.0 : f32 + scf.yield %zero : f32 + } + %aoi7_a = index.add %rec_col, %k896 : index + %aoi7 = index.assume %aoi7_a [range(%aoi7_a, 0, 2047)] : index + %ddh7 = scalar.fptrunc %dh7 : f32 to f16 + %publish_active7 = scalar.andi %lane_live, %active7 : i1 + scf.if %publish_active7 { + view.store %ao7, %aost_flat[%aoi7] : f32, view<2048xf32> + } + %gp7_8 = vector.extract %grw7_1_a[0] : vector<4xf32> -> f32 + %ag7_8 = scalar.fmaf %gp7_8, %dl7, %ag6_8 : f32 + %qp7_8 = vector.extract %qkw7_1_a[0] : vector<4xf32> -> f32 + %aq7_8 = scalar.fmaf %qp7_8, %dl7, %aq6_8 : f32 + %gp7_9 = vector.extract %grw7_1_a[1] : vector<4xf32> -> f32 + %ag7_9 = scalar.fmaf %gp7_9, %dl7, %ag6_9 : f32 + %qp7_9 = vector.extract %qkw7_1_a[1] : vector<4xf32> -> f32 + %aq7_9 = scalar.fmaf %qp7_9, %dl7, %aq6_9 : f32 + %gp7_10 = vector.extract %grw7_1_a[2] : vector<4xf32> -> f32 + %ag7_10 = scalar.fmaf %gp7_10, %dl7, %ag6_10 : f32 + %qp7_10 = vector.extract %qkw7_1_a[2] : vector<4xf32> -> f32 + %aq7_10 = scalar.fmaf %qp7_10, %dl7, %aq6_10 : f32 + %gp7_11 = vector.extract %grw7_1_a[3] : vector<4xf32> -> f32 + %ag7_11 = scalar.fmaf %gp7_11, %dl7, %ag6_11 : f32 + %qp7_11 = vector.extract %qkw7_1_a[3] : vector<4xf32> -> f32 + %aq7_11 = scalar.fmaf %qp7_11, %dl7, %aq6_11 : f32 + %gp7_12 = vector.extract %grw7_1_b[0] : vector<4xf32> -> f32 + %ag7_12 = scalar.fmaf %gp7_12, %dl7, %ag6_12 : f32 + %qp7_12 = vector.extract %qkw7_1_b[0] : vector<4xf32> -> f32 + %aq7_12 = scalar.fmaf %qp7_12, %dl7, %aq6_12 : f32 + %gp7_13 = vector.extract %grw7_1_b[1] : vector<4xf32> -> f32 + %ag7_13 = scalar.fmaf %gp7_13, %dl7, %ag6_13 : f32 + %qp7_13 = vector.extract %qkw7_1_b[1] : vector<4xf32> -> f32 + %aq7_13 = scalar.fmaf %qp7_13, %dl7, %aq6_13 : f32 + %gp7_14 = vector.extract %grw7_1_b[2] : vector<4xf32> -> f32 + %ag7_14 = scalar.fmaf %gp7_14, %dl7, %ag6_14 : f32 + %qp7_14 = vector.extract %qkw7_1_b[2] : vector<4xf32> -> f32 + %aq7_14 = scalar.fmaf %qp7_14, %dl7, %aq6_14 : f32 + %gp7_15 = vector.extract %grw7_1_b[3] : vector<4xf32> -> f32 + %ag7_15 = scalar.fmaf %gp7_15, %dl7, %ag6_15 : f32 + %qp7_15 = vector.extract %qkw7_1_b[3] : vector<4xf32> -> f32 + %aq7_15 = scalar.fmaf %qp7_15, %dl7, %aq6_15 : f32 + %gq8 = vector.load %gs_flat[%k32] : view<80xf32> -> vector<4xf32> + %bv8 = vector.extract %gq8[0] : vector<4xf32> -> f32 + %cp8 = vector.extract %gq8[1] : vector<4xf32> -> f32 + %ev8 = vector.extract %gq8[2] : vector<4xf32> -> f32 + %tk8_r = index.add %t0, %k8 : index + %tk8_ok = index.cmp ult, %tk8_r, %n_tokens : index + %tk8_s = scf.select %tk8_ok, %tk8_r, %c0 : index + %tk8 = index.assume %tk8_s [range(%tk8_s, 0, 1048575)] : index + %av8_i = index.add %tile_row, %k8 : index + %av8_b = index.assume %av8_i [range(%av8_i, 0, 271)] : index + %av8 = view.load %a_flat[%av8_b] : view<272xf32> -> f32 + %bvv8 = view.load %b_flat[%av8_b] : view<272xf32> -> f32 + %vrd8_a = index.add %rec_col, %k1024 : index + %vrd8 = index.assume %vrd8_a [range(%vrd8_a, 0, 2047)] : index + %vval8 = view.load %vst_flat[%vrd8] : view<2048xf32> -> f32 + %adec8 = scalar.fmaf %av8, %cp8, %ag7_8 : f32 + %vsub8 = scalar.subf %vval8, %adec8 : f32 + %active_index8 = index.add %t0, %k8 : index + %row_in_range8 = index.cmp ult, %active_index8, %n_tokens : index + %active8 = scalar.ori %row_in_range8, %all_recurrence_rows : i1 + %dl8 = scf.if %active8 -> (f32) { + %value = scalar.mulf %vsub8, %bv8 : f32 + scf.yield %value : f32 + } else { + %padded = scalar.mulf %vval8, %bv8 : f32 + scf.yield %padded : f32 + } + %dh8 = scalar.mulf %dl8, %ev8 : f32 + %grw8_0_a = vector.load %gr_flat[%k128] : view<256xf32> -> vector<4xf32> + %grw8_0_bi = index.constant 132 : index + %grw8_0_b = vector.load %gr_flat[%grw8_0_bi] : view<256xf32> -> vector<4xf32> + %grw8_1_a = vector.load %gr_flat[%k136] : view<256xf32> -> vector<4xf32> + %grw8_1_bi = index.constant 140 : index + %grw8_1_b = vector.load %gr_flat[%grw8_1_bi] : view<256xf32> -> vector<4xf32> + %qkw8_0_a = vector.load %qkt_flat[%k128] : view<256xf32> -> vector<4xf32> + %qkw8_0_bi = index.constant 132 : index + %qkw8_0_b = vector.load %qkt_flat[%qkw8_0_bi] : view<256xf32> -> vector<4xf32> + %qkw8_1_a = vector.load %qkt_flat[%k136] : view<256xf32> -> vector<4xf32> + %qkw8_1_bi = index.constant 140 : index + %qkw8_1_b = vector.load %qkt_flat[%qkw8_1_bi] : view<256xf32> -> vector<4xf32> + %qkd8_8 = vector.extract %qkw8_1_a[0] : vector<4xf32> -> f32 + %qsum8 = scalar.fmaf %qkd8_8, %dl8, %aq7_8 : f32 + %ao8_r = scalar.fmaf %bvv8, %cp8, %qsum8 : f32 + %ao8 = scf.if %active8 -> (f32) { + %value = scalar.mulf %ao8_r, %scale : f32 + scf.yield %value : f32 + } else { + %zero = scalar.constant 0.0 : f32 + scf.yield %zero : f32 + } + %aoi8_a = index.add %rec_col, %k1024 : index + %aoi8 = index.assume %aoi8_a [range(%aoi8_a, 0, 2047)] : index + %ddh8 = scalar.fptrunc %dh8 : f32 to f16 + %publish_active8 = scalar.andi %lane_live, %active8 : i1 + scf.if %publish_active8 { + view.store %ao8, %aost_flat[%aoi8] : f32, view<2048xf32> + } + %gp8_9 = vector.extract %grw8_1_a[1] : vector<4xf32> -> f32 + %ag8_9 = scalar.fmaf %gp8_9, %dl8, %ag7_9 : f32 + %qp8_9 = vector.extract %qkw8_1_a[1] : vector<4xf32> -> f32 + %aq8_9 = scalar.fmaf %qp8_9, %dl8, %aq7_9 : f32 + %gp8_10 = vector.extract %grw8_1_a[2] : vector<4xf32> -> f32 + %ag8_10 = scalar.fmaf %gp8_10, %dl8, %ag7_10 : f32 + %qp8_10 = vector.extract %qkw8_1_a[2] : vector<4xf32> -> f32 + %aq8_10 = scalar.fmaf %qp8_10, %dl8, %aq7_10 : f32 + %gp8_11 = vector.extract %grw8_1_a[3] : vector<4xf32> -> f32 + %ag8_11 = scalar.fmaf %gp8_11, %dl8, %ag7_11 : f32 + %qp8_11 = vector.extract %qkw8_1_a[3] : vector<4xf32> -> f32 + %aq8_11 = scalar.fmaf %qp8_11, %dl8, %aq7_11 : f32 + %gp8_12 = vector.extract %grw8_1_b[0] : vector<4xf32> -> f32 + %ag8_12 = scalar.fmaf %gp8_12, %dl8, %ag7_12 : f32 + %qp8_12 = vector.extract %qkw8_1_b[0] : vector<4xf32> -> f32 + %aq8_12 = scalar.fmaf %qp8_12, %dl8, %aq7_12 : f32 + %gp8_13 = vector.extract %grw8_1_b[1] : vector<4xf32> -> f32 + %ag8_13 = scalar.fmaf %gp8_13, %dl8, %ag7_13 : f32 + %qp8_13 = vector.extract %qkw8_1_b[1] : vector<4xf32> -> f32 + %aq8_13 = scalar.fmaf %qp8_13, %dl8, %aq7_13 : f32 + %gp8_14 = vector.extract %grw8_1_b[2] : vector<4xf32> -> f32 + %ag8_14 = scalar.fmaf %gp8_14, %dl8, %ag7_14 : f32 + %qp8_14 = vector.extract %qkw8_1_b[2] : vector<4xf32> -> f32 + %aq8_14 = scalar.fmaf %qp8_14, %dl8, %aq7_14 : f32 + %gp8_15 = vector.extract %grw8_1_b[3] : vector<4xf32> -> f32 + %ag8_15 = scalar.fmaf %gp8_15, %dl8, %ag7_15 : f32 + %qp8_15 = vector.extract %qkw8_1_b[3] : vector<4xf32> -> f32 + %aq8_15 = scalar.fmaf %qp8_15, %dl8, %aq7_15 : f32 + %gq9 = vector.load %gs_flat[%k36] : view<80xf32> -> vector<4xf32> + %bv9 = vector.extract %gq9[0] : vector<4xf32> -> f32 + %cp9 = vector.extract %gq9[1] : vector<4xf32> -> f32 + %ev9 = vector.extract %gq9[2] : vector<4xf32> -> f32 + %tk9_r = index.add %t0, %k9 : index + %tk9_ok = index.cmp ult, %tk9_r, %n_tokens : index + %tk9_s = scf.select %tk9_ok, %tk9_r, %c0 : index + %tk9 = index.assume %tk9_s [range(%tk9_s, 0, 1048575)] : index + %av9_i = index.add %tile_row, %k9 : index + %av9_b = index.assume %av9_i [range(%av9_i, 0, 271)] : index + %av9 = view.load %a_flat[%av9_b] : view<272xf32> -> f32 + %bvv9 = view.load %b_flat[%av9_b] : view<272xf32> -> f32 + %vrd9_a = index.add %rec_col, %k1152 : index + %vrd9 = index.assume %vrd9_a [range(%vrd9_a, 0, 2047)] : index + %vval9 = view.load %vst_flat[%vrd9] : view<2048xf32> -> f32 + %adec9 = scalar.fmaf %av9, %cp9, %ag8_9 : f32 + %vsub9 = scalar.subf %vval9, %adec9 : f32 + %active_index9 = index.add %t0, %k9 : index + %row_in_range9 = index.cmp ult, %active_index9, %n_tokens : index + %active9 = scalar.ori %row_in_range9, %all_recurrence_rows : i1 + %dl9 = scf.if %active9 -> (f32) { + %value = scalar.mulf %vsub9, %bv9 : f32 + scf.yield %value : f32 + } else { + %padded = scalar.mulf %vval9, %bv9 : f32 + scf.yield %padded : f32 + } + %dh9 = scalar.mulf %dl9, %ev9 : f32 + %grw9_0_a = vector.load %gr_flat[%k144] : view<256xf32> -> vector<4xf32> + %grw9_0_bi = index.constant 148 : index + %grw9_0_b = vector.load %gr_flat[%grw9_0_bi] : view<256xf32> -> vector<4xf32> + %grw9_1_a = vector.load %gr_flat[%k152] : view<256xf32> -> vector<4xf32> + %grw9_1_bi = index.constant 156 : index + %grw9_1_b = vector.load %gr_flat[%grw9_1_bi] : view<256xf32> -> vector<4xf32> + %qkw9_0_a = vector.load %qkt_flat[%k144] : view<256xf32> -> vector<4xf32> + %qkw9_0_bi = index.constant 148 : index + %qkw9_0_b = vector.load %qkt_flat[%qkw9_0_bi] : view<256xf32> -> vector<4xf32> + %qkw9_1_a = vector.load %qkt_flat[%k152] : view<256xf32> -> vector<4xf32> + %qkw9_1_bi = index.constant 156 : index + %qkw9_1_b = vector.load %qkt_flat[%qkw9_1_bi] : view<256xf32> -> vector<4xf32> + %qkd9_9 = vector.extract %qkw9_1_a[1] : vector<4xf32> -> f32 + %qsum9 = scalar.fmaf %qkd9_9, %dl9, %aq8_9 : f32 + %ao9_r = scalar.fmaf %bvv9, %cp9, %qsum9 : f32 + %ao9 = scf.if %active9 -> (f32) { + %value = scalar.mulf %ao9_r, %scale : f32 + scf.yield %value : f32 + } else { + %zero = scalar.constant 0.0 : f32 + scf.yield %zero : f32 + } + %aoi9_a = index.add %rec_col, %k1152 : index + %aoi9 = index.assume %aoi9_a [range(%aoi9_a, 0, 2047)] : index + %ddh9 = scalar.fptrunc %dh9 : f32 to f16 + %publish_active9 = scalar.andi %lane_live, %active9 : i1 + scf.if %publish_active9 { + view.store %ao9, %aost_flat[%aoi9] : f32, view<2048xf32> + } + %gp9_10 = vector.extract %grw9_1_a[2] : vector<4xf32> -> f32 + %ag9_10 = scalar.fmaf %gp9_10, %dl9, %ag8_10 : f32 + %qp9_10 = vector.extract %qkw9_1_a[2] : vector<4xf32> -> f32 + %aq9_10 = scalar.fmaf %qp9_10, %dl9, %aq8_10 : f32 + %gp9_11 = vector.extract %grw9_1_a[3] : vector<4xf32> -> f32 + %ag9_11 = scalar.fmaf %gp9_11, %dl9, %ag8_11 : f32 + %qp9_11 = vector.extract %qkw9_1_a[3] : vector<4xf32> -> f32 + %aq9_11 = scalar.fmaf %qp9_11, %dl9, %aq8_11 : f32 + %gp9_12 = vector.extract %grw9_1_b[0] : vector<4xf32> -> f32 + %ag9_12 = scalar.fmaf %gp9_12, %dl9, %ag8_12 : f32 + %qp9_12 = vector.extract %qkw9_1_b[0] : vector<4xf32> -> f32 + %aq9_12 = scalar.fmaf %qp9_12, %dl9, %aq8_12 : f32 + %gp9_13 = vector.extract %grw9_1_b[1] : vector<4xf32> -> f32 + %ag9_13 = scalar.fmaf %gp9_13, %dl9, %ag8_13 : f32 + %qp9_13 = vector.extract %qkw9_1_b[1] : vector<4xf32> -> f32 + %aq9_13 = scalar.fmaf %qp9_13, %dl9, %aq8_13 : f32 + %gp9_14 = vector.extract %grw9_1_b[2] : vector<4xf32> -> f32 + %ag9_14 = scalar.fmaf %gp9_14, %dl9, %ag8_14 : f32 + %qp9_14 = vector.extract %qkw9_1_b[2] : vector<4xf32> -> f32 + %aq9_14 = scalar.fmaf %qp9_14, %dl9, %aq8_14 : f32 + %gp9_15 = vector.extract %grw9_1_b[3] : vector<4xf32> -> f32 + %ag9_15 = scalar.fmaf %gp9_15, %dl9, %ag8_15 : f32 + %qp9_15 = vector.extract %qkw9_1_b[3] : vector<4xf32> -> f32 + %aq9_15 = scalar.fmaf %qp9_15, %dl9, %aq8_15 : f32 + %gq10 = vector.load %gs_flat[%k40] : view<80xf32> -> vector<4xf32> + %bv10 = vector.extract %gq10[0] : vector<4xf32> -> f32 + %cp10 = vector.extract %gq10[1] : vector<4xf32> -> f32 + %ev10 = vector.extract %gq10[2] : vector<4xf32> -> f32 + %tk10_r = index.add %t0, %k10 : index + %tk10_ok = index.cmp ult, %tk10_r, %n_tokens : index + %tk10_s = scf.select %tk10_ok, %tk10_r, %c0 : index + %tk10 = index.assume %tk10_s [range(%tk10_s, 0, 1048575)] : index + %av10_i = index.add %tile_row, %k10 : index + %av10_b = index.assume %av10_i [range(%av10_i, 0, 271)] : index + %av10 = view.load %a_flat[%av10_b] : view<272xf32> -> f32 + %bvv10 = view.load %b_flat[%av10_b] : view<272xf32> -> f32 + %vrd10_a = index.add %rec_col, %k1280 : index + %vrd10 = index.assume %vrd10_a [range(%vrd10_a, 0, 2047)] : index + %vval10 = view.load %vst_flat[%vrd10] : view<2048xf32> -> f32 + %adec10 = scalar.fmaf %av10, %cp10, %ag9_10 : f32 + %vsub10 = scalar.subf %vval10, %adec10 : f32 + %active_index10 = index.add %t0, %k10 : index + %row_in_range10 = index.cmp ult, %active_index10, %n_tokens : index + %active10 = scalar.ori %row_in_range10, %all_recurrence_rows : i1 + %dl10 = scf.if %active10 -> (f32) { + %value = scalar.mulf %vsub10, %bv10 : f32 + scf.yield %value : f32 + } else { + %padded = scalar.mulf %vval10, %bv10 : f32 + scf.yield %padded : f32 + } + %dh10 = scalar.mulf %dl10, %ev10 : f32 + %grw10_0_a = vector.load %gr_flat[%k160] : view<256xf32> -> vector<4xf32> + %grw10_0_bi = index.constant 164 : index + %grw10_0_b = vector.load %gr_flat[%grw10_0_bi] : view<256xf32> -> vector<4xf32> + %grw10_1_a = vector.load %gr_flat[%k168] : view<256xf32> -> vector<4xf32> + %grw10_1_bi = index.constant 172 : index + %grw10_1_b = vector.load %gr_flat[%grw10_1_bi] : view<256xf32> -> vector<4xf32> + %qkw10_0_a = vector.load %qkt_flat[%k160] : view<256xf32> -> vector<4xf32> + %qkw10_0_bi = index.constant 164 : index + %qkw10_0_b = vector.load %qkt_flat[%qkw10_0_bi] : view<256xf32> -> vector<4xf32> + %qkw10_1_a = vector.load %qkt_flat[%k168] : view<256xf32> -> vector<4xf32> + %qkw10_1_bi = index.constant 172 : index + %qkw10_1_b = vector.load %qkt_flat[%qkw10_1_bi] : view<256xf32> -> vector<4xf32> + %qkd10_10 = vector.extract %qkw10_1_a[2] : vector<4xf32> -> f32 + %qsum10 = scalar.fmaf %qkd10_10, %dl10, %aq9_10 : f32 + %ao10_r = scalar.fmaf %bvv10, %cp10, %qsum10 : f32 + %ao10 = scf.if %active10 -> (f32) { + %value = scalar.mulf %ao10_r, %scale : f32 + scf.yield %value : f32 + } else { + %zero = scalar.constant 0.0 : f32 + scf.yield %zero : f32 + } + %aoi10_a = index.add %rec_col, %k1280 : index + %aoi10 = index.assume %aoi10_a [range(%aoi10_a, 0, 2047)] : index + %ddh10 = scalar.fptrunc %dh10 : f32 to f16 + %publish_active10 = scalar.andi %lane_live, %active10 : i1 + scf.if %publish_active10 { + view.store %ao10, %aost_flat[%aoi10] : f32, view<2048xf32> + } + %gp10_11 = vector.extract %grw10_1_a[3] : vector<4xf32> -> f32 + %ag10_11 = scalar.fmaf %gp10_11, %dl10, %ag9_11 : f32 + %qp10_11 = vector.extract %qkw10_1_a[3] : vector<4xf32> -> f32 + %aq10_11 = scalar.fmaf %qp10_11, %dl10, %aq9_11 : f32 + %gp10_12 = vector.extract %grw10_1_b[0] : vector<4xf32> -> f32 + %ag10_12 = scalar.fmaf %gp10_12, %dl10, %ag9_12 : f32 + %qp10_12 = vector.extract %qkw10_1_b[0] : vector<4xf32> -> f32 + %aq10_12 = scalar.fmaf %qp10_12, %dl10, %aq9_12 : f32 + %gp10_13 = vector.extract %grw10_1_b[1] : vector<4xf32> -> f32 + %ag10_13 = scalar.fmaf %gp10_13, %dl10, %ag9_13 : f32 + %qp10_13 = vector.extract %qkw10_1_b[1] : vector<4xf32> -> f32 + %aq10_13 = scalar.fmaf %qp10_13, %dl10, %aq9_13 : f32 + %gp10_14 = vector.extract %grw10_1_b[2] : vector<4xf32> -> f32 + %ag10_14 = scalar.fmaf %gp10_14, %dl10, %ag9_14 : f32 + %qp10_14 = vector.extract %qkw10_1_b[2] : vector<4xf32> -> f32 + %aq10_14 = scalar.fmaf %qp10_14, %dl10, %aq9_14 : f32 + %gp10_15 = vector.extract %grw10_1_b[3] : vector<4xf32> -> f32 + %ag10_15 = scalar.fmaf %gp10_15, %dl10, %ag9_15 : f32 + %qp10_15 = vector.extract %qkw10_1_b[3] : vector<4xf32> -> f32 + %aq10_15 = scalar.fmaf %qp10_15, %dl10, %aq9_15 : f32 + %gq11 = vector.load %gs_flat[%k44] : view<80xf32> -> vector<4xf32> + %bv11 = vector.extract %gq11[0] : vector<4xf32> -> f32 + %cp11 = vector.extract %gq11[1] : vector<4xf32> -> f32 + %ev11 = vector.extract %gq11[2] : vector<4xf32> -> f32 + %tk11_r = index.add %t0, %k11 : index + %tk11_ok = index.cmp ult, %tk11_r, %n_tokens : index + %tk11_s = scf.select %tk11_ok, %tk11_r, %c0 : index + %tk11 = index.assume %tk11_s [range(%tk11_s, 0, 1048575)] : index + %av11_i = index.add %tile_row, %k11 : index + %av11_b = index.assume %av11_i [range(%av11_i, 0, 271)] : index + %av11 = view.load %a_flat[%av11_b] : view<272xf32> -> f32 + %bvv11 = view.load %b_flat[%av11_b] : view<272xf32> -> f32 + %vrd11_a = index.add %rec_col, %k1408 : index + %vrd11 = index.assume %vrd11_a [range(%vrd11_a, 0, 2047)] : index + %vval11 = view.load %vst_flat[%vrd11] : view<2048xf32> -> f32 + %adec11 = scalar.fmaf %av11, %cp11, %ag10_11 : f32 + %vsub11 = scalar.subf %vval11, %adec11 : f32 + %active_index11 = index.add %t0, %k11 : index + %row_in_range11 = index.cmp ult, %active_index11, %n_tokens : index + %active11 = scalar.ori %row_in_range11, %all_recurrence_rows : i1 + %dl11 = scf.if %active11 -> (f32) { + %value = scalar.mulf %vsub11, %bv11 : f32 + scf.yield %value : f32 + } else { + %padded = scalar.mulf %vval11, %bv11 : f32 + scf.yield %padded : f32 + } + %dh11 = scalar.mulf %dl11, %ev11 : f32 + %grw11_0_a = vector.load %gr_flat[%k176] : view<256xf32> -> vector<4xf32> + %grw11_0_bi = index.constant 180 : index + %grw11_0_b = vector.load %gr_flat[%grw11_0_bi] : view<256xf32> -> vector<4xf32> + %grw11_1_a = vector.load %gr_flat[%k184] : view<256xf32> -> vector<4xf32> + %grw11_1_bi = index.constant 188 : index + %grw11_1_b = vector.load %gr_flat[%grw11_1_bi] : view<256xf32> -> vector<4xf32> + %qkw11_0_a = vector.load %qkt_flat[%k176] : view<256xf32> -> vector<4xf32> + %qkw11_0_bi = index.constant 180 : index + %qkw11_0_b = vector.load %qkt_flat[%qkw11_0_bi] : view<256xf32> -> vector<4xf32> + %qkw11_1_a = vector.load %qkt_flat[%k184] : view<256xf32> -> vector<4xf32> + %qkw11_1_bi = index.constant 188 : index + %qkw11_1_b = vector.load %qkt_flat[%qkw11_1_bi] : view<256xf32> -> vector<4xf32> + %qkd11_11 = vector.extract %qkw11_1_a[3] : vector<4xf32> -> f32 + %qsum11 = scalar.fmaf %qkd11_11, %dl11, %aq10_11 : f32 + %ao11_r = scalar.fmaf %bvv11, %cp11, %qsum11 : f32 + %ao11 = scf.if %active11 -> (f32) { + %value = scalar.mulf %ao11_r, %scale : f32 + scf.yield %value : f32 + } else { + %zero = scalar.constant 0.0 : f32 + scf.yield %zero : f32 + } + %aoi11_a = index.add %rec_col, %k1408 : index + %aoi11 = index.assume %aoi11_a [range(%aoi11_a, 0, 2047)] : index + %ddh11 = scalar.fptrunc %dh11 : f32 to f16 + %publish_active11 = scalar.andi %lane_live, %active11 : i1 + scf.if %publish_active11 { + view.store %ao11, %aost_flat[%aoi11] : f32, view<2048xf32> + } + %gp11_12 = vector.extract %grw11_1_b[0] : vector<4xf32> -> f32 + %ag11_12 = scalar.fmaf %gp11_12, %dl11, %ag10_12 : f32 + %qp11_12 = vector.extract %qkw11_1_b[0] : vector<4xf32> -> f32 + %aq11_12 = scalar.fmaf %qp11_12, %dl11, %aq10_12 : f32 + %gp11_13 = vector.extract %grw11_1_b[1] : vector<4xf32> -> f32 + %ag11_13 = scalar.fmaf %gp11_13, %dl11, %ag10_13 : f32 + %qp11_13 = vector.extract %qkw11_1_b[1] : vector<4xf32> -> f32 + %aq11_13 = scalar.fmaf %qp11_13, %dl11, %aq10_13 : f32 + %gp11_14 = vector.extract %grw11_1_b[2] : vector<4xf32> -> f32 + %ag11_14 = scalar.fmaf %gp11_14, %dl11, %ag10_14 : f32 + %qp11_14 = vector.extract %qkw11_1_b[2] : vector<4xf32> -> f32 + %aq11_14 = scalar.fmaf %qp11_14, %dl11, %aq10_14 : f32 + %gp11_15 = vector.extract %grw11_1_b[3] : vector<4xf32> -> f32 + %ag11_15 = scalar.fmaf %gp11_15, %dl11, %ag10_15 : f32 + %qp11_15 = vector.extract %qkw11_1_b[3] : vector<4xf32> -> f32 + %aq11_15 = scalar.fmaf %qp11_15, %dl11, %aq10_15 : f32 + %gq12 = vector.load %gs_flat[%k48] : view<80xf32> -> vector<4xf32> + %bv12 = vector.extract %gq12[0] : vector<4xf32> -> f32 + %cp12 = vector.extract %gq12[1] : vector<4xf32> -> f32 + %ev12 = vector.extract %gq12[2] : vector<4xf32> -> f32 + %tk12_r = index.add %t0, %k12 : index + %tk12_ok = index.cmp ult, %tk12_r, %n_tokens : index + %tk12_s = scf.select %tk12_ok, %tk12_r, %c0 : index + %tk12 = index.assume %tk12_s [range(%tk12_s, 0, 1048575)] : index + %av12_i = index.add %tile_row, %k12 : index + %av12_b = index.assume %av12_i [range(%av12_i, 0, 271)] : index + %av12 = view.load %a_flat[%av12_b] : view<272xf32> -> f32 + %bvv12 = view.load %b_flat[%av12_b] : view<272xf32> -> f32 + %vrd12_a = index.add %rec_col, %k1536 : index + %vrd12 = index.assume %vrd12_a [range(%vrd12_a, 0, 2047)] : index + %vval12 = view.load %vst_flat[%vrd12] : view<2048xf32> -> f32 + %adec12 = scalar.fmaf %av12, %cp12, %ag11_12 : f32 + %vsub12 = scalar.subf %vval12, %adec12 : f32 + %active_index12 = index.add %t0, %k12 : index + %row_in_range12 = index.cmp ult, %active_index12, %n_tokens : index + %active12 = scalar.ori %row_in_range12, %all_recurrence_rows : i1 + %dl12 = scf.if %active12 -> (f32) { + %value = scalar.mulf %vsub12, %bv12 : f32 + scf.yield %value : f32 + } else { + %padded = scalar.mulf %vval12, %bv12 : f32 + scf.yield %padded : f32 + } + %dh12 = scalar.mulf %dl12, %ev12 : f32 + %grw12_0_a = vector.load %gr_flat[%k192] : view<256xf32> -> vector<4xf32> + %grw12_0_bi = index.constant 196 : index + %grw12_0_b = vector.load %gr_flat[%grw12_0_bi] : view<256xf32> -> vector<4xf32> + %grw12_1_a = vector.load %gr_flat[%k200] : view<256xf32> -> vector<4xf32> + %grw12_1_bi = index.constant 204 : index + %grw12_1_b = vector.load %gr_flat[%grw12_1_bi] : view<256xf32> -> vector<4xf32> + %qkw12_0_a = vector.load %qkt_flat[%k192] : view<256xf32> -> vector<4xf32> + %qkw12_0_bi = index.constant 196 : index + %qkw12_0_b = vector.load %qkt_flat[%qkw12_0_bi] : view<256xf32> -> vector<4xf32> + %qkw12_1_a = vector.load %qkt_flat[%k200] : view<256xf32> -> vector<4xf32> + %qkw12_1_bi = index.constant 204 : index + %qkw12_1_b = vector.load %qkt_flat[%qkw12_1_bi] : view<256xf32> -> vector<4xf32> + %qkd12_12 = vector.extract %qkw12_1_b[0] : vector<4xf32> -> f32 + %qsum12 = scalar.fmaf %qkd12_12, %dl12, %aq11_12 : f32 + %ao12_r = scalar.fmaf %bvv12, %cp12, %qsum12 : f32 + %ao12 = scf.if %active12 -> (f32) { + %value = scalar.mulf %ao12_r, %scale : f32 + scf.yield %value : f32 + } else { + %zero = scalar.constant 0.0 : f32 + scf.yield %zero : f32 + } + %aoi12_a = index.add %rec_col, %k1536 : index + %aoi12 = index.assume %aoi12_a [range(%aoi12_a, 0, 2047)] : index + %ddh12 = scalar.fptrunc %dh12 : f32 to f16 + %publish_active12 = scalar.andi %lane_live, %active12 : i1 + scf.if %publish_active12 { + view.store %ao12, %aost_flat[%aoi12] : f32, view<2048xf32> + } + %gp12_13 = vector.extract %grw12_1_b[1] : vector<4xf32> -> f32 + %ag12_13 = scalar.fmaf %gp12_13, %dl12, %ag11_13 : f32 + %qp12_13 = vector.extract %qkw12_1_b[1] : vector<4xf32> -> f32 + %aq12_13 = scalar.fmaf %qp12_13, %dl12, %aq11_13 : f32 + %gp12_14 = vector.extract %grw12_1_b[2] : vector<4xf32> -> f32 + %ag12_14 = scalar.fmaf %gp12_14, %dl12, %ag11_14 : f32 + %qp12_14 = vector.extract %qkw12_1_b[2] : vector<4xf32> -> f32 + %aq12_14 = scalar.fmaf %qp12_14, %dl12, %aq11_14 : f32 + %gp12_15 = vector.extract %grw12_1_b[3] : vector<4xf32> -> f32 + %ag12_15 = scalar.fmaf %gp12_15, %dl12, %ag11_15 : f32 + %qp12_15 = vector.extract %qkw12_1_b[3] : vector<4xf32> -> f32 + %aq12_15 = scalar.fmaf %qp12_15, %dl12, %aq11_15 : f32 + %gq13 = vector.load %gs_flat[%k52] : view<80xf32> -> vector<4xf32> + %bv13 = vector.extract %gq13[0] : vector<4xf32> -> f32 + %cp13 = vector.extract %gq13[1] : vector<4xf32> -> f32 + %ev13 = vector.extract %gq13[2] : vector<4xf32> -> f32 + %tk13_r = index.add %t0, %k13 : index + %tk13_ok = index.cmp ult, %tk13_r, %n_tokens : index + %tk13_s = scf.select %tk13_ok, %tk13_r, %c0 : index + %tk13 = index.assume %tk13_s [range(%tk13_s, 0, 1048575)] : index + %av13_i = index.add %tile_row, %k13 : index + %av13_b = index.assume %av13_i [range(%av13_i, 0, 271)] : index + %av13 = view.load %a_flat[%av13_b] : view<272xf32> -> f32 + %bvv13 = view.load %b_flat[%av13_b] : view<272xf32> -> f32 + %vrd13_a = index.add %rec_col, %k1664 : index + %vrd13 = index.assume %vrd13_a [range(%vrd13_a, 0, 2047)] : index + %vval13 = view.load %vst_flat[%vrd13] : view<2048xf32> -> f32 + %adec13 = scalar.fmaf %av13, %cp13, %ag12_13 : f32 + %vsub13 = scalar.subf %vval13, %adec13 : f32 + %active_index13 = index.add %t0, %k13 : index + %row_in_range13 = index.cmp ult, %active_index13, %n_tokens : index + %active13 = scalar.ori %row_in_range13, %all_recurrence_rows : i1 + %dl13 = scf.if %active13 -> (f32) { + %value = scalar.mulf %vsub13, %bv13 : f32 + scf.yield %value : f32 + } else { + %padded = scalar.mulf %vval13, %bv13 : f32 + scf.yield %padded : f32 + } + %dh13 = scalar.mulf %dl13, %ev13 : f32 + %grw13_0_a = vector.load %gr_flat[%k208] : view<256xf32> -> vector<4xf32> + %grw13_0_bi = index.constant 212 : index + %grw13_0_b = vector.load %gr_flat[%grw13_0_bi] : view<256xf32> -> vector<4xf32> + %grw13_1_a = vector.load %gr_flat[%k216] : view<256xf32> -> vector<4xf32> + %grw13_1_bi = index.constant 220 : index + %grw13_1_b = vector.load %gr_flat[%grw13_1_bi] : view<256xf32> -> vector<4xf32> + %qkw13_0_a = vector.load %qkt_flat[%k208] : view<256xf32> -> vector<4xf32> + %qkw13_0_bi = index.constant 212 : index + %qkw13_0_b = vector.load %qkt_flat[%qkw13_0_bi] : view<256xf32> -> vector<4xf32> + %qkw13_1_a = vector.load %qkt_flat[%k216] : view<256xf32> -> vector<4xf32> + %qkw13_1_bi = index.constant 220 : index + %qkw13_1_b = vector.load %qkt_flat[%qkw13_1_bi] : view<256xf32> -> vector<4xf32> + %qkd13_13 = vector.extract %qkw13_1_b[1] : vector<4xf32> -> f32 + %qsum13 = scalar.fmaf %qkd13_13, %dl13, %aq12_13 : f32 + %ao13_r = scalar.fmaf %bvv13, %cp13, %qsum13 : f32 + %ao13 = scf.if %active13 -> (f32) { + %value = scalar.mulf %ao13_r, %scale : f32 + scf.yield %value : f32 + } else { + %zero = scalar.constant 0.0 : f32 + scf.yield %zero : f32 + } + %aoi13_a = index.add %rec_col, %k1664 : index + %aoi13 = index.assume %aoi13_a [range(%aoi13_a, 0, 2047)] : index + %ddh13 = scalar.fptrunc %dh13 : f32 to f16 + %publish_active13 = scalar.andi %lane_live, %active13 : i1 + scf.if %publish_active13 { + view.store %ao13, %aost_flat[%aoi13] : f32, view<2048xf32> + } + %gp13_14 = vector.extract %grw13_1_b[2] : vector<4xf32> -> f32 + %ag13_14 = scalar.fmaf %gp13_14, %dl13, %ag12_14 : f32 + %qp13_14 = vector.extract %qkw13_1_b[2] : vector<4xf32> -> f32 + %aq13_14 = scalar.fmaf %qp13_14, %dl13, %aq12_14 : f32 + %gp13_15 = vector.extract %grw13_1_b[3] : vector<4xf32> -> f32 + %ag13_15 = scalar.fmaf %gp13_15, %dl13, %ag12_15 : f32 + %qp13_15 = vector.extract %qkw13_1_b[3] : vector<4xf32> -> f32 + %aq13_15 = scalar.fmaf %qp13_15, %dl13, %aq12_15 : f32 + %gq14 = vector.load %gs_flat[%k56] : view<80xf32> -> vector<4xf32> + %bv14 = vector.extract %gq14[0] : vector<4xf32> -> f32 + %cp14 = vector.extract %gq14[1] : vector<4xf32> -> f32 + %ev14 = vector.extract %gq14[2] : vector<4xf32> -> f32 + %tk14_r = index.add %t0, %k14 : index + %tk14_ok = index.cmp ult, %tk14_r, %n_tokens : index + %tk14_s = scf.select %tk14_ok, %tk14_r, %c0 : index + %tk14 = index.assume %tk14_s [range(%tk14_s, 0, 1048575)] : index + %av14_i = index.add %tile_row, %k14 : index + %av14_b = index.assume %av14_i [range(%av14_i, 0, 271)] : index + %av14 = view.load %a_flat[%av14_b] : view<272xf32> -> f32 + %bvv14 = view.load %b_flat[%av14_b] : view<272xf32> -> f32 + %vrd14_a = index.add %rec_col, %k1792 : index + %vrd14 = index.assume %vrd14_a [range(%vrd14_a, 0, 2047)] : index + %vval14 = view.load %vst_flat[%vrd14] : view<2048xf32> -> f32 + %adec14 = scalar.fmaf %av14, %cp14, %ag13_14 : f32 + %vsub14 = scalar.subf %vval14, %adec14 : f32 + %active_index14 = index.add %t0, %k14 : index + %row_in_range14 = index.cmp ult, %active_index14, %n_tokens : index + %active14 = scalar.ori %row_in_range14, %all_recurrence_rows : i1 + %dl14 = scf.if %active14 -> (f32) { + %value = scalar.mulf %vsub14, %bv14 : f32 + scf.yield %value : f32 + } else { + %padded = scalar.mulf %vval14, %bv14 : f32 + scf.yield %padded : f32 + } + %dh14 = scalar.mulf %dl14, %ev14 : f32 + %grw14_0_a = vector.load %gr_flat[%k224] : view<256xf32> -> vector<4xf32> + %grw14_0_bi = index.constant 228 : index + %grw14_0_b = vector.load %gr_flat[%grw14_0_bi] : view<256xf32> -> vector<4xf32> + %grw14_1_a = vector.load %gr_flat[%k232] : view<256xf32> -> vector<4xf32> + %grw14_1_bi = index.constant 236 : index + %grw14_1_b = vector.load %gr_flat[%grw14_1_bi] : view<256xf32> -> vector<4xf32> + %qkw14_0_a = vector.load %qkt_flat[%k224] : view<256xf32> -> vector<4xf32> + %qkw14_0_bi = index.constant 228 : index + %qkw14_0_b = vector.load %qkt_flat[%qkw14_0_bi] : view<256xf32> -> vector<4xf32> + %qkw14_1_a = vector.load %qkt_flat[%k232] : view<256xf32> -> vector<4xf32> + %qkw14_1_bi = index.constant 236 : index + %qkw14_1_b = vector.load %qkt_flat[%qkw14_1_bi] : view<256xf32> -> vector<4xf32> + %qkd14_14 = vector.extract %qkw14_1_b[2] : vector<4xf32> -> f32 + %qsum14 = scalar.fmaf %qkd14_14, %dl14, %aq13_14 : f32 + %ao14_r = scalar.fmaf %bvv14, %cp14, %qsum14 : f32 + %ao14 = scf.if %active14 -> (f32) { + %value = scalar.mulf %ao14_r, %scale : f32 + scf.yield %value : f32 + } else { + %zero = scalar.constant 0.0 : f32 + scf.yield %zero : f32 + } + %aoi14_a = index.add %rec_col, %k1792 : index + %aoi14 = index.assume %aoi14_a [range(%aoi14_a, 0, 2047)] : index + %ddh14 = scalar.fptrunc %dh14 : f32 to f16 + %publish_active14 = scalar.andi %lane_live, %active14 : i1 + scf.if %publish_active14 { + view.store %ao14, %aost_flat[%aoi14] : f32, view<2048xf32> + } + %gp14_15 = vector.extract %grw14_1_b[3] : vector<4xf32> -> f32 + %ag14_15 = scalar.fmaf %gp14_15, %dl14, %ag13_15 : f32 + %qp14_15 = vector.extract %qkw14_1_b[3] : vector<4xf32> -> f32 + %aq14_15 = scalar.fmaf %qp14_15, %dl14, %aq13_15 : f32 + %gq15 = vector.load %gs_flat[%k60] : view<80xf32> -> vector<4xf32> + %bv15 = vector.extract %gq15[0] : vector<4xf32> -> f32 + %cp15 = vector.extract %gq15[1] : vector<4xf32> -> f32 + %ev15 = vector.extract %gq15[2] : vector<4xf32> -> f32 + %tk15_r = index.add %t0, %k15 : index + %tk15_ok = index.cmp ult, %tk15_r, %n_tokens : index + %tk15_s = scf.select %tk15_ok, %tk15_r, %c0 : index + %tk15 = index.assume %tk15_s [range(%tk15_s, 0, 1048575)] : index + %av15_i = index.add %tile_row, %k15 : index + %av15_b = index.assume %av15_i [range(%av15_i, 0, 271)] : index + %av15 = view.load %a_flat[%av15_b] : view<272xf32> -> f32 + %bvv15 = view.load %b_flat[%av15_b] : view<272xf32> -> f32 + %vrd15_a = index.add %rec_col, %k1920 : index + %vrd15 = index.assume %vrd15_a [range(%vrd15_a, 0, 2047)] : index + %vval15 = view.load %vst_flat[%vrd15] : view<2048xf32> -> f32 + %adec15 = scalar.fmaf %av15, %cp15, %ag14_15 : f32 + %vsub15 = scalar.subf %vval15, %adec15 : f32 + %active_index15 = index.add %t0, %k15 : index + %row_in_range15 = index.cmp ult, %active_index15, %n_tokens : index + %active15 = scalar.ori %row_in_range15, %all_recurrence_rows : i1 + %dl15 = scf.if %active15 -> (f32) { + %value = scalar.mulf %vsub15, %bv15 : f32 + scf.yield %value : f32 + } else { + %padded = scalar.mulf %vval15, %bv15 : f32 + scf.yield %padded : f32 + } + %dh15 = scalar.mulf %dl15, %ev15 : f32 + %grw15_0_a = vector.load %gr_flat[%k240] : view<256xf32> -> vector<4xf32> + %grw15_0_bi = index.constant 244 : index + %grw15_0_b = vector.load %gr_flat[%grw15_0_bi] : view<256xf32> -> vector<4xf32> + %grw15_1_a = vector.load %gr_flat[%k248] : view<256xf32> -> vector<4xf32> + %grw15_1_bi = index.constant 252 : index + %grw15_1_b = vector.load %gr_flat[%grw15_1_bi] : view<256xf32> -> vector<4xf32> + %qkw15_0_a = vector.load %qkt_flat[%k240] : view<256xf32> -> vector<4xf32> + %qkw15_0_bi = index.constant 244 : index + %qkw15_0_b = vector.load %qkt_flat[%qkw15_0_bi] : view<256xf32> -> vector<4xf32> + %qkw15_1_a = vector.load %qkt_flat[%k248] : view<256xf32> -> vector<4xf32> + %qkw15_1_bi = index.constant 252 : index + %qkw15_1_b = vector.load %qkt_flat[%qkw15_1_bi] : view<256xf32> -> vector<4xf32> + %qkd15_15 = vector.extract %qkw15_1_b[3] : vector<4xf32> -> f32 + %qsum15 = scalar.fmaf %qkd15_15, %dl15, %aq14_15 : f32 + %ao15_r = scalar.fmaf %bvv15, %cp15, %qsum15 : f32 + %ao15 = scf.if %active15 -> (f32) { + %value = scalar.mulf %ao15_r, %scale : f32 + scf.yield %value : f32 + } else { + %zero = scalar.constant 0.0 : f32 + scf.yield %zero : f32 + } + %aoi15_a = index.add %rec_col, %k1920 : index + %aoi15 = index.assume %aoi15_a [range(%aoi15_a, 0, 2047)] : index + %ddh15 = scalar.fptrunc %dh15 : f32 to f16 + %publish_active15 = scalar.andi %lane_live, %active15 : i1 + scf.if %publish_active15 { + view.store %ao15, %aost_flat[%aoi15] : f32, view<2048xf32> + } + // A is now dead. Materialize the register-held deltas into its + // allocation so the full-head owner does not reserve a separate D tile. + %ddv0 = vector.from_elements %ddh0, %ddh1, %ddh2, %ddh3, %ddh4, %ddh5, %ddh6, %ddh7 : vector<8xf16> + %ddv1 = vector.from_elements %ddh8, %ddh9, %ddh10, %ddh11, %ddh12, %ddh13, %ddh14, %ddh15 : vector<8xf16> + %dds0 = index.assume %d_row [range(%d_row, 0, 376)] : index + %dds1_a = index.add %d_row, %k8 : index + %dds1 = index.assume %dds1_a [range(%dds1_a, 8, 376)] : index + scf.if %lane_live { + vector.store %ddv0, %d_flat[%dds0] : vector<8xf16>, view<384xf16> + vector.store %ddv1, %d_flat[%dds1] : vector<8xf16>, view<384xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + + scf.if %rmsnorm_gate { + scf.if %publish_q8 { + } else { + // One native subgroup handles each local attention row, using the + // standalone RMS reduction and unary implementation without a global read. + %post_rows0 = index.mul %n_tokens, %n_heads : index + %post_rows = index.mul %post_rows0, %n_seqs : index + %post_tokens = index.mul %n_tokens, %n_seqs : index + %post_groups = index.mul %n_heads, %k8 : index + %post_channel0 = index.mul %lane, %k4 : index + %post_channel = index.assume %post_channel0 [range(%post_channel0, 0, 124), mul(%post_channel0, 4)] : index + %post_attention = buffer.view %lds[%base] : buffer -> view<16x128xf32> + %post_weights = buffer.view %rms_weight[%base] : buffer -> view<128xf32> + %post_gates = buffer.view %raw_gate[%base] : buffer -> view<[%post_rows]x128xf32> + %post_f32 = buffer.view %norm_output[%base] : buffer -> view<[%post_rows]x128xf32> + %post_f16 = buffer.view %half_output[%base] : buffer -> view<[%post_groups]x[%post_tokens]x16xf16> + scf.for %post_local0 = [%wave to %k16 step %k8] { + %post_local = index.assume %post_local0 [range(%post_local0, 0, 15)] : index + %post_time0 = index.add %t0, %post_local : index + %post_valid = index.cmp ult, %post_time0, %n_tokens : index + %post_time_safe = scf.select %post_valid, %post_time0, %c0 : index + %post_time = index.assume %post_time_safe [lt(%post_time_safe, %n_tokens)] : index + %post_token0 = index.madd %seq, %n_tokens, %post_time : index + %post_token = index.assume %post_token0 [lt(%post_token0, %post_tokens)] : index + %post_row0 = index.madd %post_token, %n_heads, %h_idx : index + %post_row = index.assume %post_row0 [lt(%post_row0, %post_rows)] : index + %post_scale = template.apply<@ggml.rmsnorm_f32.subgroup_row_scale>(%k16, %post_local, %k128, %rms_epsilon, %lds) : (index, index, index, f32, buffer) -> (f32) + scf.if %post_valid { + %post_x = vector.load %post_attention[%post_local, %post_channel] : view<16x128xf32> -> vector<4xf32> + %post_w = vector.load %post_weights[%post_channel] : view<128xf32> -> vector<4xf32> + %post_g = vector.load %post_gates[%post_row, %post_channel] : view<[%post_rows]x128xf32> -> vector<4xf32> + %post_scale_v = vector.splat %post_scale : vector<4xf32> + %post_normalized = vector.mulf %post_x, %post_scale_v : vector<4xf32> + %post_weighted = vector.mulf %post_normalized, %post_w : vector<4xf32> + %post_activated = template.apply<@ggml.unary_f32.apply_vector4>(%rms_gate_op, %post_g) : (index, vector<4xf32>) -> (vector<4xf32>) + %post_result = vector.mulf %post_weighted, %post_activated : vector<4xf32> + vector.store %post_result, %post_f32[%post_row, %post_channel] : vector<4xf32>, view<[%post_rows]x128xf32> + scf.if %publish_q8 { + %post_flat_channel = index.madd %post_row, %k128, %post_channel : index + %post_scratch_values = buffer.view %lds[%base] : buffer -> view<256xf32> + %post_scratch_d = buffer.view %lds[%base] : buffer -> view<32xf32> + template.apply<@ggml.quantize_q8_1_x4.publish_vector4_strict>(%post_valid, %base, %post_flat_channel, %post_result, %post_scratch_values, %post_scratch_d, %half_output) : (i1, offset, index, vector<4xf32>, view<256xf32>, view<32xf32>, buffer) + } else { + %post_half = vector.fptrunc %post_result : vector<4xf32> to vector<4xf16> + %post_head_channel = index.madd %h_idx, %k128, %post_channel : index + %post_group = index.div %post_head_channel, %k16 : index + %post_group_lane = index.rem %post_head_channel, %k16 : index + vector.store %post_half, %post_f16[%post_group, %post_token, %post_group_lane] : vector<4xf16>, view<[%post_groups]x[%post_tokens]x16xf16> + } + } + } + + } + } else { + %aow0_r = index.add %t0, %k0 : index + %aow0_s = index.add %aow0_r, %vs_tok : index + %aow0_ok = index.cmp ult, %aow0_s, %n_tokens : index + %aow0_c = scf.select %aow0_ok, %aow0_s, %c0 : index + %aow0 = index.assume %aow0_c [range(%aow0_c, 0, 1048575)] : index + %aol0_m = index.mul %vs_tok, %k128 : index + %aol0_a = index.add %aol0_m, %vs_col : index + %aol0_b = index.add %aol0_a, %k0 : index + %aol0 = index.assume %aol0_b [range(%aol0_b, 0, 2047)] : index + %aov0 = view.load %aost_flat[%aol0] : view<2048xf32> -> f32 + %aog0_m = index.mul %aow0, %attn_stride : index + %aog0_a = index.add %attn_base, %aog0_m : index + %aog0_c = index.add %aog0_a, %vs_gcol : index + %aog0 = index.assume %aog0_c [range(%aog0_c, 0, 1073741823)] : index + scf.if %aow0_ok { + view.store %aov0, %dst_view[%aog0] : f32, view<1073741824xf32> + } + %aow1_r = index.add %t0, %k2 : index + %aow1_s = index.add %aow1_r, %vs_tok : index + %aow1_ok = index.cmp ult, %aow1_s, %n_tokens : index + %aow1_c = scf.select %aow1_ok, %aow1_s, %c0 : index + %aow1 = index.assume %aow1_c [range(%aow1_c, 0, 1048575)] : index + %aol1_m = index.mul %vs_tok, %k128 : index + %aol1_a = index.add %aol1_m, %vs_col : index + %aol1_b = index.add %aol1_a, %k256 : index + %aol1 = index.assume %aol1_b [range(%aol1_b, 0, 2047)] : index + %aov1 = view.load %aost_flat[%aol1] : view<2048xf32> -> f32 + %aog1_m = index.mul %aow1, %attn_stride : index + %aog1_a = index.add %attn_base, %aog1_m : index + %aog1_c = index.add %aog1_a, %vs_gcol : index + %aog1 = index.assume %aog1_c [range(%aog1_c, 0, 1073741823)] : index + scf.if %aow1_ok { + view.store %aov1, %dst_view[%aog1] : f32, view<1073741824xf32> + } + %aow2_r = index.add %t0, %k4 : index + %aow2_s = index.add %aow2_r, %vs_tok : index + %aow2_ok = index.cmp ult, %aow2_s, %n_tokens : index + %aow2_c = scf.select %aow2_ok, %aow2_s, %c0 : index + %aow2 = index.assume %aow2_c [range(%aow2_c, 0, 1048575)] : index + %aol2_m = index.mul %vs_tok, %k128 : index + %aol2_a = index.add %aol2_m, %vs_col : index + %aol2_b = index.add %aol2_a, %k512 : index + %aol2 = index.assume %aol2_b [range(%aol2_b, 0, 2047)] : index + %aov2 = view.load %aost_flat[%aol2] : view<2048xf32> -> f32 + %aog2_m = index.mul %aow2, %attn_stride : index + %aog2_a = index.add %attn_base, %aog2_m : index + %aog2_c = index.add %aog2_a, %vs_gcol : index + %aog2 = index.assume %aog2_c [range(%aog2_c, 0, 1073741823)] : index + scf.if %aow2_ok { + view.store %aov2, %dst_view[%aog2] : f32, view<1073741824xf32> + } + %aow3_r = index.add %t0, %k6 : index + %aow3_s = index.add %aow3_r, %vs_tok : index + %aow3_ok = index.cmp ult, %aow3_s, %n_tokens : index + %aow3_c = scf.select %aow3_ok, %aow3_s, %c0 : index + %aow3 = index.assume %aow3_c [range(%aow3_c, 0, 1048575)] : index + %aol3_m = index.mul %vs_tok, %k128 : index + %aol3_a = index.add %aol3_m, %vs_col : index + %aol3_b = index.add %aol3_a, %k768 : index + %aol3 = index.assume %aol3_b [range(%aol3_b, 0, 2047)] : index + %aov3 = view.load %aost_flat[%aol3] : view<2048xf32> -> f32 + %aog3_m = index.mul %aow3, %attn_stride : index + %aog3_a = index.add %attn_base, %aog3_m : index + %aog3_c = index.add %aog3_a, %vs_gcol : index + %aog3 = index.assume %aog3_c [range(%aog3_c, 0, 1073741823)] : index + scf.if %aow3_ok { + view.store %aov3, %dst_view[%aog3] : f32, view<1073741824xf32> + } + %aow4_r = index.add %t0, %k8 : index + %aow4_s = index.add %aow4_r, %vs_tok : index + %aow4_ok = index.cmp ult, %aow4_s, %n_tokens : index + %aow4_c = scf.select %aow4_ok, %aow4_s, %c0 : index + %aow4 = index.assume %aow4_c [range(%aow4_c, 0, 1048575)] : index + %aol4_m = index.mul %vs_tok, %k128 : index + %aol4_a = index.add %aol4_m, %vs_col : index + %aol4_b = index.add %aol4_a, %k1024 : index + %aol4 = index.assume %aol4_b [range(%aol4_b, 0, 2047)] : index + %aov4 = view.load %aost_flat[%aol4] : view<2048xf32> -> f32 + %aog4_m = index.mul %aow4, %attn_stride : index + %aog4_a = index.add %attn_base, %aog4_m : index + %aog4_c = index.add %aog4_a, %vs_gcol : index + %aog4 = index.assume %aog4_c [range(%aog4_c, 0, 1073741823)] : index + scf.if %aow4_ok { + view.store %aov4, %dst_view[%aog4] : f32, view<1073741824xf32> + } + %aow5_r = index.add %t0, %k10 : index + %aow5_s = index.add %aow5_r, %vs_tok : index + %aow5_ok = index.cmp ult, %aow5_s, %n_tokens : index + %aow5_c = scf.select %aow5_ok, %aow5_s, %c0 : index + %aow5 = index.assume %aow5_c [range(%aow5_c, 0, 1048575)] : index + %aol5_m = index.mul %vs_tok, %k128 : index + %aol5_a = index.add %aol5_m, %vs_col : index + %aol5_b = index.add %aol5_a, %k1280 : index + %aol5 = index.assume %aol5_b [range(%aol5_b, 0, 2047)] : index + %aov5 = view.load %aost_flat[%aol5] : view<2048xf32> -> f32 + %aog5_m = index.mul %aow5, %attn_stride : index + %aog5_a = index.add %attn_base, %aog5_m : index + %aog5_c = index.add %aog5_a, %vs_gcol : index + %aog5 = index.assume %aog5_c [range(%aog5_c, 0, 1073741823)] : index + scf.if %aow5_ok { + view.store %aov5, %dst_view[%aog5] : f32, view<1073741824xf32> + } + %aow6_r = index.add %t0, %k12 : index + %aow6_s = index.add %aow6_r, %vs_tok : index + %aow6_ok = index.cmp ult, %aow6_s, %n_tokens : index + %aow6_c = scf.select %aow6_ok, %aow6_s, %c0 : index + %aow6 = index.assume %aow6_c [range(%aow6_c, 0, 1048575)] : index + %aol6_m = index.mul %vs_tok, %k128 : index + %aol6_a = index.add %aol6_m, %vs_col : index + %aol6_b = index.add %aol6_a, %k1536 : index + %aol6 = index.assume %aol6_b [range(%aol6_b, 0, 2047)] : index + %aov6 = view.load %aost_flat[%aol6] : view<2048xf32> -> f32 + %aog6_m = index.mul %aow6, %attn_stride : index + %aog6_a = index.add %attn_base, %aog6_m : index + %aog6_c = index.add %aog6_a, %vs_gcol : index + %aog6 = index.assume %aog6_c [range(%aog6_c, 0, 1073741823)] : index + scf.if %aow6_ok { + view.store %aov6, %dst_view[%aog6] : f32, view<1073741824xf32> + } + %aow7_r = index.add %t0, %k14 : index + %aow7_s = index.add %aow7_r, %vs_tok : index + %aow7_ok = index.cmp ult, %aow7_s, %n_tokens : index + %aow7_c = scf.select %aow7_ok, %aow7_s, %c0 : index + %aow7 = index.assume %aow7_c [range(%aow7_c, 0, 1048575)] : index + %aol7_m = index.mul %vs_tok, %k128 : index + %aol7_a = index.add %aol7_m, %vs_col : index + %aol7_b = index.add %aol7_a, %k1792 : index + %aol7 = index.assume %aol7_b [range(%aol7_b, 0, 2047)] : index + %aov7 = view.load %aost_flat[%aol7] : view<2048xf32> -> f32 + %aog7_m = index.mul %aow7, %attn_stride : index + %aog7_a = index.add %attn_base, %aog7_m : index + %aog7_c = index.add %aog7_a, %vs_gcol : index + %aog7 = index.assume %aog7_c [range(%aog7_c, 0, 1073741823)] : index + scf.if %aow7_ok { + view.store %aov7, %dst_view[%aog7] : f32, view<1073741824xf32> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + // K' was formed from the exact normalized f16 K tile before it + // was overwritten. Only the first 128 threads own rows. + scf.if %kt_live { + %ktrow2 = index.assume %tid [range(%tid, 0, 127)] : index + %ktdst0 = index.mul %ktrow2, %k24 : index + %ktdst = index.assume %ktdst0 [range(%ktdst0, 0, 3064)] : index + vector.store %ktv0, %ka_flat[%ktdst] : vector<8xf16>, view<3072xf16> + %ktdst8_a = index.add %ktdst, %k8 : index + %ktdst8 = index.assume %ktdst8_a [range(%ktdst8_a, 8, 3064)] : index + vector.store %ktv1, %ka_flat[%ktdst8] : vector<8xf16>, view<3072xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %rowh_r = index.rem %row0, %k64 : index + %rowh = index.assume %rowh_r [range(%rowh_r, 0, 60)] : index + %snapshot_next0, %snapshot_next1, %snapshot_next2, %snapshot_next3, %snapshot_next4, %snapshot_next5, %snapshot_next6, %snapshot_next7, %snapshot_next8, %snapshot_next9, %snapshot_next10, %snapshot_next11, %snapshot_next12, %snapshot_next13, %snapshot_next14, %snapshot_next15 = scf.if %reuse_snapshot_fragments -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + scf.if %publish_snapshots { + %snapshot_c0o = index.constant 0 : offset + %snapshot_c0 = index.constant 0 : index + %snapshot_c4 = index.constant 4 : index + %snapshot_c16 = index.constant 16 : index + %snapshot_c32 = index.constant 32 : index + %snapshot_c48 = index.constant 48 : index + %snapshot_c64 = index.constant 64 : index + %snapshot_c68 = index.constant 68 : index + %snapshot_c80 = index.constant 80 : index + %snapshot_c96 = index.constant 96 : index + %snapshot_c112 = index.constant 112 : index + %snapshot_c128 = index.constant 128 : index + %snapshot_c4352 = index.constant 4352 : index + %snapshot_c6528 = index.constant 6528 : index + %snapshot_c11072 = index.constant 11072 : index + %snapshot_ka_o = index.constant 4352 : offset + %snapshot_lay_dcol = encoding.layout.strided [24, 1] : encoding + %snapshot_lay_katok = encoding.layout.strided [1, 24] : encoding + %snapshot_lay_upd = encoding.layout.strided [68, 1] : encoding + %snapshot_lay_state = encoding.layout.strided [128, 1] : encoding + %snapshot_wave_bytes = index.mul %wave, %snapshot_c6528 : index + %snapshot_wave_base = index.add %snapshot_wave_bytes, %snapshot_c11072 : index + %snapshot_wave_off = index.cast %snapshot_wave_base : index to offset + %snapshot_d_base_i = index.add %snapshot_wave_base, %snapshot_c4352 : index + %snapshot_d_base = index.cast %snapshot_d_base_i : index to offset + %snapshot_d_lhs = buffer.view %lds[%snapshot_d_base] : buffer -> view<16x16xf16, %snapshot_lay_dcol> + %snapshot_ka_rhs = buffer.view %lds[%snapshot_ka_o] : buffer -> view<16x128xf16, %snapshot_lay_katok> + %snapshot_upd_flat = buffer.view %lds[%snapshot_wave_off] : buffer -> view<1088xf32> + %snapshot_upd_tile = buffer.view %lds[%snapshot_wave_off] : buffer -> view<16x68xf32, %snapshot_lay_upd> + %snapshot_lane_low = index.cmp ult, %lane, %snapshot_c16 : index + %snapshot_lane_high = index.cmp uge, %lane, %snapshot_c16 : index + scf.if %snapshot_lane_low { + vector.store %sa0, %snapshot_upd_flat[%rowh] : vector<4xf32>, view<1088xf32> + %snapshot_lo_s1_off = index.constant 68 : index + %snapshot_lo_s1_raw = index.add %snapshot_lo_s1_off, %rowh : index + %snapshot_lo_s1_idx = index.assume %snapshot_lo_s1_raw [range(%snapshot_lo_s1_raw, 68, 128)] : index + vector.store %sa1, %snapshot_upd_flat[%snapshot_lo_s1_idx] : vector<4xf32>, view<1088xf32> + %snapshot_lo_s2_off = index.constant 136 : index + %snapshot_lo_s2_raw = index.add %snapshot_lo_s2_off, %rowh : index + %snapshot_lo_s2_idx = index.assume %snapshot_lo_s2_raw [range(%snapshot_lo_s2_raw, 136, 196)] : index + vector.store %sa2, %snapshot_upd_flat[%snapshot_lo_s2_idx] : vector<4xf32>, view<1088xf32> + %snapshot_lo_s3_off = index.constant 204 : index + %snapshot_lo_s3_raw = index.add %snapshot_lo_s3_off, %rowh : index + %snapshot_lo_s3_idx = index.assume %snapshot_lo_s3_raw [range(%snapshot_lo_s3_raw, 204, 264)] : index + vector.store %sa3, %snapshot_upd_flat[%snapshot_lo_s3_idx] : vector<4xf32>, view<1088xf32> + %snapshot_lo_s4_off = index.constant 272 : index + %snapshot_lo_s4_raw = index.add %snapshot_lo_s4_off, %rowh : index + %snapshot_lo_s4_idx = index.assume %snapshot_lo_s4_raw [range(%snapshot_lo_s4_raw, 272, 332)] : index + vector.store %sa4, %snapshot_upd_flat[%snapshot_lo_s4_idx] : vector<4xf32>, view<1088xf32> + %snapshot_lo_s5_off = index.constant 340 : index + %snapshot_lo_s5_raw = index.add %snapshot_lo_s5_off, %rowh : index + %snapshot_lo_s5_idx = index.assume %snapshot_lo_s5_raw [range(%snapshot_lo_s5_raw, 340, 400)] : index + vector.store %sa5, %snapshot_upd_flat[%snapshot_lo_s5_idx] : vector<4xf32>, view<1088xf32> + %snapshot_lo_s6_off = index.constant 408 : index + %snapshot_lo_s6_raw = index.add %snapshot_lo_s6_off, %rowh : index + %snapshot_lo_s6_idx = index.assume %snapshot_lo_s6_raw [range(%snapshot_lo_s6_raw, 408, 468)] : index + vector.store %sa6, %snapshot_upd_flat[%snapshot_lo_s6_idx] : vector<4xf32>, view<1088xf32> + %snapshot_lo_s7_off = index.constant 476 : index + %snapshot_lo_s7_raw = index.add %snapshot_lo_s7_off, %rowh : index + %snapshot_lo_s7_idx = index.assume %snapshot_lo_s7_raw [range(%snapshot_lo_s7_raw, 476, 536)] : index + vector.store %sa7, %snapshot_upd_flat[%snapshot_lo_s7_idx] : vector<4xf32>, view<1088xf32> + %snapshot_lo_s8_off = index.constant 544 : index + %snapshot_lo_s8_raw = index.add %snapshot_lo_s8_off, %rowh : index + %snapshot_lo_s8_idx = index.assume %snapshot_lo_s8_raw [range(%snapshot_lo_s8_raw, 544, 604)] : index + vector.store %sa8, %snapshot_upd_flat[%snapshot_lo_s8_idx] : vector<4xf32>, view<1088xf32> + %snapshot_lo_s9_off = index.constant 612 : index + %snapshot_lo_s9_raw = index.add %snapshot_lo_s9_off, %rowh : index + %snapshot_lo_s9_idx = index.assume %snapshot_lo_s9_raw [range(%snapshot_lo_s9_raw, 612, 672)] : index + vector.store %sa9, %snapshot_upd_flat[%snapshot_lo_s9_idx] : vector<4xf32>, view<1088xf32> + %snapshot_lo_s10_off = index.constant 680 : index + %snapshot_lo_s10_raw = index.add %snapshot_lo_s10_off, %rowh : index + %snapshot_lo_s10_idx = index.assume %snapshot_lo_s10_raw [range(%snapshot_lo_s10_raw, 680, 740)] : index + vector.store %sa10, %snapshot_upd_flat[%snapshot_lo_s10_idx] : vector<4xf32>, view<1088xf32> + %snapshot_lo_s11_off = index.constant 748 : index + %snapshot_lo_s11_raw = index.add %snapshot_lo_s11_off, %rowh : index + %snapshot_lo_s11_idx = index.assume %snapshot_lo_s11_raw [range(%snapshot_lo_s11_raw, 748, 808)] : index + vector.store %sa11, %snapshot_upd_flat[%snapshot_lo_s11_idx] : vector<4xf32>, view<1088xf32> + %snapshot_lo_s12_off = index.constant 816 : index + %snapshot_lo_s12_raw = index.add %snapshot_lo_s12_off, %rowh : index + %snapshot_lo_s12_idx = index.assume %snapshot_lo_s12_raw [range(%snapshot_lo_s12_raw, 816, 876)] : index + vector.store %sa12, %snapshot_upd_flat[%snapshot_lo_s12_idx] : vector<4xf32>, view<1088xf32> + %snapshot_lo_s13_off = index.constant 884 : index + %snapshot_lo_s13_raw = index.add %snapshot_lo_s13_off, %rowh : index + %snapshot_lo_s13_idx = index.assume %snapshot_lo_s13_raw [range(%snapshot_lo_s13_raw, 884, 944)] : index + vector.store %sa13, %snapshot_upd_flat[%snapshot_lo_s13_idx] : vector<4xf32>, view<1088xf32> + %snapshot_lo_s14_off = index.constant 952 : index + %snapshot_lo_s14_raw = index.add %snapshot_lo_s14_off, %rowh : index + %snapshot_lo_s14_idx = index.assume %snapshot_lo_s14_raw [range(%snapshot_lo_s14_raw, 952, 1012)] : index + vector.store %sa14, %snapshot_upd_flat[%snapshot_lo_s14_idx] : vector<4xf32>, view<1088xf32> + %snapshot_lo_s15_off = index.constant 1020 : index + %snapshot_lo_s15_raw = index.add %snapshot_lo_s15_off, %rowh : index + %snapshot_lo_s15_idx = index.assume %snapshot_lo_s15_raw [range(%snapshot_lo_s15_raw, 1020, 1080)] : index + vector.store %sa15, %snapshot_upd_flat[%snapshot_lo_s15_idx] : vector<4xf32>, view<1088xf32> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + %snapshot_lo0_state = vector.fragment.load %snapshot_upd_tile[%snapshot_c0, %snapshot_c0] shape [%snapshot_c16, %snapshot_c16] : view<16x68xf32, %snapshot_lay_upd> -> vector<8xf32> + %snapshot_lo1_state = vector.fragment.load %snapshot_upd_tile[%snapshot_c0, %snapshot_c16] shape [%snapshot_c16, %snapshot_c16] : view<16x68xf32, %snapshot_lay_upd> -> vector<8xf32> + %snapshot_lo2_state = vector.fragment.load %snapshot_upd_tile[%snapshot_c0, %snapshot_c32] shape [%snapshot_c16, %snapshot_c16] : view<16x68xf32, %snapshot_lay_upd> -> vector<8xf32> + %snapshot_lo3_state = vector.fragment.load %snapshot_upd_tile[%snapshot_c0, %snapshot_c48] shape [%snapshot_c16, %snapshot_c16] : view<16x68xf32, %snapshot_lay_upd> -> vector<8xf32> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %snapshot_lane_high { + vector.store %sa0, %snapshot_upd_flat[%rowh] : vector<4xf32>, view<1088xf32> + %snapshot_hi_s1_off = index.constant 68 : index + %snapshot_hi_s1_raw = index.add %snapshot_hi_s1_off, %rowh : index + %snapshot_hi_s1_idx = index.assume %snapshot_hi_s1_raw [range(%snapshot_hi_s1_raw, 68, 128)] : index + vector.store %sa1, %snapshot_upd_flat[%snapshot_hi_s1_idx] : vector<4xf32>, view<1088xf32> + %snapshot_hi_s2_off = index.constant 136 : index + %snapshot_hi_s2_raw = index.add %snapshot_hi_s2_off, %rowh : index + %snapshot_hi_s2_idx = index.assume %snapshot_hi_s2_raw [range(%snapshot_hi_s2_raw, 136, 196)] : index + vector.store %sa2, %snapshot_upd_flat[%snapshot_hi_s2_idx] : vector<4xf32>, view<1088xf32> + %snapshot_hi_s3_off = index.constant 204 : index + %snapshot_hi_s3_raw = index.add %snapshot_hi_s3_off, %rowh : index + %snapshot_hi_s3_idx = index.assume %snapshot_hi_s3_raw [range(%snapshot_hi_s3_raw, 204, 264)] : index + vector.store %sa3, %snapshot_upd_flat[%snapshot_hi_s3_idx] : vector<4xf32>, view<1088xf32> + %snapshot_hi_s4_off = index.constant 272 : index + %snapshot_hi_s4_raw = index.add %snapshot_hi_s4_off, %rowh : index + %snapshot_hi_s4_idx = index.assume %snapshot_hi_s4_raw [range(%snapshot_hi_s4_raw, 272, 332)] : index + vector.store %sa4, %snapshot_upd_flat[%snapshot_hi_s4_idx] : vector<4xf32>, view<1088xf32> + %snapshot_hi_s5_off = index.constant 340 : index + %snapshot_hi_s5_raw = index.add %snapshot_hi_s5_off, %rowh : index + %snapshot_hi_s5_idx = index.assume %snapshot_hi_s5_raw [range(%snapshot_hi_s5_raw, 340, 400)] : index + vector.store %sa5, %snapshot_upd_flat[%snapshot_hi_s5_idx] : vector<4xf32>, view<1088xf32> + %snapshot_hi_s6_off = index.constant 408 : index + %snapshot_hi_s6_raw = index.add %snapshot_hi_s6_off, %rowh : index + %snapshot_hi_s6_idx = index.assume %snapshot_hi_s6_raw [range(%snapshot_hi_s6_raw, 408, 468)] : index + vector.store %sa6, %snapshot_upd_flat[%snapshot_hi_s6_idx] : vector<4xf32>, view<1088xf32> + %snapshot_hi_s7_off = index.constant 476 : index + %snapshot_hi_s7_raw = index.add %snapshot_hi_s7_off, %rowh : index + %snapshot_hi_s7_idx = index.assume %snapshot_hi_s7_raw [range(%snapshot_hi_s7_raw, 476, 536)] : index + vector.store %sa7, %snapshot_upd_flat[%snapshot_hi_s7_idx] : vector<4xf32>, view<1088xf32> + %snapshot_hi_s8_off = index.constant 544 : index + %snapshot_hi_s8_raw = index.add %snapshot_hi_s8_off, %rowh : index + %snapshot_hi_s8_idx = index.assume %snapshot_hi_s8_raw [range(%snapshot_hi_s8_raw, 544, 604)] : index + vector.store %sa8, %snapshot_upd_flat[%snapshot_hi_s8_idx] : vector<4xf32>, view<1088xf32> + %snapshot_hi_s9_off = index.constant 612 : index + %snapshot_hi_s9_raw = index.add %snapshot_hi_s9_off, %rowh : index + %snapshot_hi_s9_idx = index.assume %snapshot_hi_s9_raw [range(%snapshot_hi_s9_raw, 612, 672)] : index + vector.store %sa9, %snapshot_upd_flat[%snapshot_hi_s9_idx] : vector<4xf32>, view<1088xf32> + %snapshot_hi_s10_off = index.constant 680 : index + %snapshot_hi_s10_raw = index.add %snapshot_hi_s10_off, %rowh : index + %snapshot_hi_s10_idx = index.assume %snapshot_hi_s10_raw [range(%snapshot_hi_s10_raw, 680, 740)] : index + vector.store %sa10, %snapshot_upd_flat[%snapshot_hi_s10_idx] : vector<4xf32>, view<1088xf32> + %snapshot_hi_s11_off = index.constant 748 : index + %snapshot_hi_s11_raw = index.add %snapshot_hi_s11_off, %rowh : index + %snapshot_hi_s11_idx = index.assume %snapshot_hi_s11_raw [range(%snapshot_hi_s11_raw, 748, 808)] : index + vector.store %sa11, %snapshot_upd_flat[%snapshot_hi_s11_idx] : vector<4xf32>, view<1088xf32> + %snapshot_hi_s12_off = index.constant 816 : index + %snapshot_hi_s12_raw = index.add %snapshot_hi_s12_off, %rowh : index + %snapshot_hi_s12_idx = index.assume %snapshot_hi_s12_raw [range(%snapshot_hi_s12_raw, 816, 876)] : index + vector.store %sa12, %snapshot_upd_flat[%snapshot_hi_s12_idx] : vector<4xf32>, view<1088xf32> + %snapshot_hi_s13_off = index.constant 884 : index + %snapshot_hi_s13_raw = index.add %snapshot_hi_s13_off, %rowh : index + %snapshot_hi_s13_idx = index.assume %snapshot_hi_s13_raw [range(%snapshot_hi_s13_raw, 884, 944)] : index + vector.store %sa13, %snapshot_upd_flat[%snapshot_hi_s13_idx] : vector<4xf32>, view<1088xf32> + %snapshot_hi_s14_off = index.constant 952 : index + %snapshot_hi_s14_raw = index.add %snapshot_hi_s14_off, %rowh : index + %snapshot_hi_s14_idx = index.assume %snapshot_hi_s14_raw [range(%snapshot_hi_s14_raw, 952, 1012)] : index + vector.store %sa14, %snapshot_upd_flat[%snapshot_hi_s14_idx] : vector<4xf32>, view<1088xf32> + %snapshot_hi_s15_off = index.constant 1020 : index + %snapshot_hi_s15_raw = index.add %snapshot_hi_s15_off, %rowh : index + %snapshot_hi_s15_idx = index.assume %snapshot_hi_s15_raw [range(%snapshot_hi_s15_raw, 1020, 1080)] : index + vector.store %sa15, %snapshot_upd_flat[%snapshot_hi_s15_idx] : vector<4xf32>, view<1088xf32> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + %snapshot_hi0_state = vector.fragment.load %snapshot_upd_tile[%snapshot_c0, %snapshot_c0] shape [%snapshot_c16, %snapshot_c16] : view<16x68xf32, %snapshot_lay_upd> -> vector<8xf32> + %snapshot_hi1_state = vector.fragment.load %snapshot_upd_tile[%snapshot_c0, %snapshot_c16] shape [%snapshot_c16, %snapshot_c16] : view<16x68xf32, %snapshot_lay_upd> -> vector<8xf32> + %snapshot_hi2_state = vector.fragment.load %snapshot_upd_tile[%snapshot_c0, %snapshot_c32] shape [%snapshot_c16, %snapshot_c16] : view<16x68xf32, %snapshot_lay_upd> -> vector<8xf32> + %snapshot_hi3_state = vector.fragment.load %snapshot_upd_tile[%snapshot_c0, %snapshot_c48] shape [%snapshot_c16, %snapshot_c16] : view<16x68xf32, %snapshot_lay_upd> -> vector<8xf32> + kernel.barrier scope(subgroup) ordering(acq_rel) + %snapshot_zero = scalar.constant 0.0 : f16 + %snapshot_last = index.sub %n_tokens, %k1 : index + %has_prefix0 = index.cmp ult, %k1, %n_tokens : index + %has_prefix1 = index.cmp ult, %k2, %n_tokens : index + %has_prefix2 = index.cmp ult, %k3, %n_tokens : index + %has_prefix3 = index.cmp ult, %k4, %n_tokens : index + scf.if %has_prefix0 { + %snapshot_slot0 = index.mul %snapshot_stride, %snapshot_last : index + %p0d0 = scalar.fptrunc %dl0 : f32 to f16 + %p0v0 = vector.from_elements %p0d0, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero : vector<8xf16> + %p0v1 = vector.from_elements %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero : vector<8xf16> + scf.if %lane_live { + vector.store %p0v0, %d_flat[%dds0] : vector<8xf16>, view<384xf16> + vector.store %p0v1, %d_flat[%dds1] : vector<8xf16>, view<384xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + template.apply<@llm.gated_delta_net.snapshot_fragment_prefix.body>(%cp0, %snapshot_slot0, %state_offset, %wcol, %row0, %rowh, %wave, %lane, %lds, %snapshot_cache, %snapshot_lo0_state, %snapshot_lo1_state, %snapshot_lo2_state, %snapshot_lo3_state, %snapshot_hi0_state, %snapshot_hi1_state, %snapshot_hi2_state, %snapshot_hi3_state, %fz8) : (f32, index, index, index, index, index, index, index, buffer, buffer, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) + kernel.barrier scope(workgroup) ordering(acq_rel) + } + scf.if %has_prefix1 { + %snapshot_index1 = index.sub %snapshot_last, %k1 : index + %snapshot_slot1 = index.mul %snapshot_stride, %snapshot_index1 : index + %p10_log = scalar.subf %lc1, %lc0 : f32 + %p10 = scalar.expf %p10_log : f32 + %p1d0f = scalar.mulf %dl0, %p10 : f32 + %p1d0 = scalar.fptrunc %p1d0f : f32 to f16 + %p1d1 = scalar.fptrunc %dl1 : f32 to f16 + %p1v0 = vector.from_elements %p1d0, %p1d1, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero : vector<8xf16> + %p1v1 = vector.from_elements %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero : vector<8xf16> + scf.if %lane_live { + vector.store %p1v0, %d_flat[%dds0] : vector<8xf16>, view<384xf16> + vector.store %p1v1, %d_flat[%dds1] : vector<8xf16>, view<384xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + template.apply<@llm.gated_delta_net.snapshot_fragment_prefix.body>(%cp1, %snapshot_slot1, %state_offset, %wcol, %row0, %rowh, %wave, %lane, %lds, %snapshot_cache, %snapshot_lo0_state, %snapshot_lo1_state, %snapshot_lo2_state, %snapshot_lo3_state, %snapshot_hi0_state, %snapshot_hi1_state, %snapshot_hi2_state, %snapshot_hi3_state, %fz8) : (f32, index, index, index, index, index, index, index, buffer, buffer, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) + kernel.barrier scope(workgroup) ordering(acq_rel) + } + scf.if %has_prefix2 { + %snapshot_index2 = index.sub %snapshot_last, %k2 : index + %snapshot_slot2 = index.mul %snapshot_stride, %snapshot_index2 : index + %p20_log = scalar.subf %lc2, %lc0 : f32 + %p20 = scalar.expf %p20_log : f32 + %p21_log = scalar.subf %lc2, %lc1 : f32 + %p21 = scalar.expf %p21_log : f32 + %p2d0f = scalar.mulf %dl0, %p20 : f32 + %p2d1f = scalar.mulf %dl1, %p21 : f32 + %p2d0 = scalar.fptrunc %p2d0f : f32 to f16 + %p2d1 = scalar.fptrunc %p2d1f : f32 to f16 + %p2d2 = scalar.fptrunc %dl2 : f32 to f16 + %p2v0 = vector.from_elements %p2d0, %p2d1, %p2d2, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero : vector<8xf16> + %p2v1 = vector.from_elements %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero : vector<8xf16> + scf.if %lane_live { + vector.store %p2v0, %d_flat[%dds0] : vector<8xf16>, view<384xf16> + vector.store %p2v1, %d_flat[%dds1] : vector<8xf16>, view<384xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + template.apply<@llm.gated_delta_net.snapshot_fragment_prefix.body>(%cp2, %snapshot_slot2, %state_offset, %wcol, %row0, %rowh, %wave, %lane, %lds, %snapshot_cache, %snapshot_lo0_state, %snapshot_lo1_state, %snapshot_lo2_state, %snapshot_lo3_state, %snapshot_hi0_state, %snapshot_hi1_state, %snapshot_hi2_state, %snapshot_hi3_state, %fz8) : (f32, index, index, index, index, index, index, index, buffer, buffer, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) + kernel.barrier scope(workgroup) ordering(acq_rel) + } + scf.if %has_prefix3 { + %snapshot_index3 = index.sub %snapshot_last, %k3 : index + %snapshot_slot3 = index.mul %snapshot_stride, %snapshot_index3 : index + %p30_log = scalar.subf %lc3, %lc0 : f32 + %p30 = scalar.expf %p30_log : f32 + %p31_log = scalar.subf %lc3, %lc1 : f32 + %p31 = scalar.expf %p31_log : f32 + %p32_log = scalar.subf %lc3, %lc2 : f32 + %p32 = scalar.expf %p32_log : f32 + %p3d0f = scalar.mulf %dl0, %p30 : f32 + %p3d1f = scalar.mulf %dl1, %p31 : f32 + %p3d2f = scalar.mulf %dl2, %p32 : f32 + %p3d0 = scalar.fptrunc %p3d0f : f32 to f16 + %p3d1 = scalar.fptrunc %p3d1f : f32 to f16 + %p3d2 = scalar.fptrunc %p3d2f : f32 to f16 + %p3d3 = scalar.fptrunc %dl3 : f32 to f16 + %p3v0 = vector.from_elements %p3d0, %p3d1, %p3d2, %p3d3, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero : vector<8xf16> + %p3v1 = vector.from_elements %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero : vector<8xf16> + scf.if %lane_live { + vector.store %p3v0, %d_flat[%dds0] : vector<8xf16>, view<384xf16> + vector.store %p3v1, %d_flat[%dds1] : vector<8xf16>, view<384xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + template.apply<@llm.gated_delta_net.snapshot_fragment_prefix.body>(%cp3, %snapshot_slot3, %state_offset, %wcol, %row0, %rowh, %wave, %lane, %lds, %snapshot_cache, %snapshot_lo0_state, %snapshot_lo1_state, %snapshot_lo2_state, %snapshot_lo3_state, %snapshot_hi0_state, %snapshot_hi1_state, %snapshot_hi2_state, %snapshot_hi3_state, %fz8) : (f32, index, index, index, index, index, index, index, buffer, buffer, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) + kernel.barrier scope(workgroup) ordering(acq_rel) + } + + scf.if %lane_live { + vector.store %ddv0, %d_flat[%dds0] : vector<8xf16>, view<384xf16> + vector.store %ddv1, %d_flat[%dds1] : vector<8xf16>, view<384xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %snapshot_clast = view.load %gs_flat[%k61] : view<80xf32> -> f32 + template.apply<@llm.gated_delta_net.snapshot_fragment_prefix.body>(%snapshot_clast, %k0, %state_offset, %wcol, %row0, %rowh, %wave, %lane, %lds, %snapshot_cache, %snapshot_lo0_state, %snapshot_lo1_state, %snapshot_lo2_state, %snapshot_lo3_state, %snapshot_hi0_state, %snapshot_hi1_state, %snapshot_hi2_state, %snapshot_hi3_state, %fz8) : (f32, index, index, index, index, index, index, index, buffer, buffer, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) + } + scf.yield %sa0, %sa1, %sa2, %sa3, %sa4, %sa5, %sa6, %sa7, %sa8, %sa9, %sa10, %sa11, %sa12, %sa13, %sa14, %sa15 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } else { + scf.if %publish_snapshots { + %snapshot_zero = scalar.constant 0.0 : f16 + %snapshot_last = index.sub %n_tokens, %k1 : index + %has_prefix0 = index.cmp ult, %k1, %n_tokens : index + %has_prefix1 = index.cmp ult, %k2, %n_tokens : index + %has_prefix2 = index.cmp ult, %k3, %n_tokens : index + %has_prefix3 = index.cmp ult, %k4, %n_tokens : index + scf.if %has_prefix0 { + %snapshot_slot0 = index.mul %snapshot_stride, %snapshot_last : index + %p0d0 = scalar.fptrunc %dl0 : f32 to f16 + %p0v0 = vector.from_elements %p0d0, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero : vector<8xf16> + %p0v1 = vector.from_elements %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero : vector<8xf16> + scf.if %lane_live { + vector.store %p0v0, %d_flat[%dds0] : vector<8xf16>, view<384xf16> + vector.store %p0v1, %d_flat[%dds1] : vector<8xf16>, view<384xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + template.apply<@llm.gated_delta_net.snapshot_prefix.body>(%cp0, %snapshot_slot0, %state_offset, %wcol, %row0, %rowh, %wave, %lane, %lds, %snapshot_cache, %sa0, %sa1, %sa2, %sa3, %sa4, %sa5, %sa6, %sa7, %sa8, %sa9, %sa10, %sa11, %sa12, %sa13, %sa14, %sa15, %fz8) : (f32, index, index, index, index, index, index, index, buffer, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<8xf32>) + kernel.barrier scope(workgroup) ordering(acq_rel) + } + scf.if %has_prefix1 { + %snapshot_index1 = index.sub %snapshot_last, %k1 : index + %snapshot_slot1 = index.mul %snapshot_stride, %snapshot_index1 : index + %p10_log = scalar.subf %lc1, %lc0 : f32 + %p10 = scalar.expf %p10_log : f32 + %p1d0f = scalar.mulf %dl0, %p10 : f32 + %p1d0 = scalar.fptrunc %p1d0f : f32 to f16 + %p1d1 = scalar.fptrunc %dl1 : f32 to f16 + %p1v0 = vector.from_elements %p1d0, %p1d1, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero : vector<8xf16> + %p1v1 = vector.from_elements %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero : vector<8xf16> + scf.if %lane_live { + vector.store %p1v0, %d_flat[%dds0] : vector<8xf16>, view<384xf16> + vector.store %p1v1, %d_flat[%dds1] : vector<8xf16>, view<384xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + template.apply<@llm.gated_delta_net.snapshot_prefix.body>(%cp1, %snapshot_slot1, %state_offset, %wcol, %row0, %rowh, %wave, %lane, %lds, %snapshot_cache, %sa0, %sa1, %sa2, %sa3, %sa4, %sa5, %sa6, %sa7, %sa8, %sa9, %sa10, %sa11, %sa12, %sa13, %sa14, %sa15, %fz8) : (f32, index, index, index, index, index, index, index, buffer, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<8xf32>) + kernel.barrier scope(workgroup) ordering(acq_rel) + } + scf.if %has_prefix2 { + %snapshot_index2 = index.sub %snapshot_last, %k2 : index + %snapshot_slot2 = index.mul %snapshot_stride, %snapshot_index2 : index + %p20_log = scalar.subf %lc2, %lc0 : f32 + %p20 = scalar.expf %p20_log : f32 + %p21_log = scalar.subf %lc2, %lc1 : f32 + %p21 = scalar.expf %p21_log : f32 + %p2d0f = scalar.mulf %dl0, %p20 : f32 + %p2d1f = scalar.mulf %dl1, %p21 : f32 + %p2d0 = scalar.fptrunc %p2d0f : f32 to f16 + %p2d1 = scalar.fptrunc %p2d1f : f32 to f16 + %p2d2 = scalar.fptrunc %dl2 : f32 to f16 + %p2v0 = vector.from_elements %p2d0, %p2d1, %p2d2, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero : vector<8xf16> + %p2v1 = vector.from_elements %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero : vector<8xf16> + scf.if %lane_live { + vector.store %p2v0, %d_flat[%dds0] : vector<8xf16>, view<384xf16> + vector.store %p2v1, %d_flat[%dds1] : vector<8xf16>, view<384xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + template.apply<@llm.gated_delta_net.snapshot_prefix.body>(%cp2, %snapshot_slot2, %state_offset, %wcol, %row0, %rowh, %wave, %lane, %lds, %snapshot_cache, %sa0, %sa1, %sa2, %sa3, %sa4, %sa5, %sa6, %sa7, %sa8, %sa9, %sa10, %sa11, %sa12, %sa13, %sa14, %sa15, %fz8) : (f32, index, index, index, index, index, index, index, buffer, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<8xf32>) + kernel.barrier scope(workgroup) ordering(acq_rel) + } + scf.if %has_prefix3 { + %snapshot_index3 = index.sub %snapshot_last, %k3 : index + %snapshot_slot3 = index.mul %snapshot_stride, %snapshot_index3 : index + %p30_log = scalar.subf %lc3, %lc0 : f32 + %p30 = scalar.expf %p30_log : f32 + %p31_log = scalar.subf %lc3, %lc1 : f32 + %p31 = scalar.expf %p31_log : f32 + %p32_log = scalar.subf %lc3, %lc2 : f32 + %p32 = scalar.expf %p32_log : f32 + %p3d0f = scalar.mulf %dl0, %p30 : f32 + %p3d1f = scalar.mulf %dl1, %p31 : f32 + %p3d2f = scalar.mulf %dl2, %p32 : f32 + %p3d0 = scalar.fptrunc %p3d0f : f32 to f16 + %p3d1 = scalar.fptrunc %p3d1f : f32 to f16 + %p3d2 = scalar.fptrunc %p3d2f : f32 to f16 + %p3d3 = scalar.fptrunc %dl3 : f32 to f16 + %p3v0 = vector.from_elements %p3d0, %p3d1, %p3d2, %p3d3, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero : vector<8xf16> + %p3v1 = vector.from_elements %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero, %snapshot_zero : vector<8xf16> + scf.if %lane_live { + vector.store %p3v0, %d_flat[%dds0] : vector<8xf16>, view<384xf16> + vector.store %p3v1, %d_flat[%dds1] : vector<8xf16>, view<384xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + template.apply<@llm.gated_delta_net.snapshot_prefix.body>(%cp3, %snapshot_slot3, %state_offset, %wcol, %row0, %rowh, %wave, %lane, %lds, %snapshot_cache, %sa0, %sa1, %sa2, %sa3, %sa4, %sa5, %sa6, %sa7, %sa8, %sa9, %sa10, %sa11, %sa12, %sa13, %sa14, %sa15, %fz8) : (f32, index, index, index, index, index, index, index, buffer, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<8xf32>) + kernel.barrier scope(workgroup) ordering(acq_rel) + } + + scf.if %lane_live { + vector.store %ddv0, %d_flat[%dds0] : vector<8xf16>, view<384xf16> + vector.store %ddv1, %d_flat[%dds1] : vector<8xf16>, view<384xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + } + %clast = view.load %gs_flat[%k61] : view<80xf32> -> f32 + %clv = vector.splat %clast : vector<4xf32> + + // Update S <- c_last S + D K' in two 64-row halves and overlay the dead F16 state copy. + // Lane L owns rows 4L..4L+3 in the half selected by its 16-lane group. + %up00_i = vector.fragment %fz8 shape [%k16, %k16] : vector<8xf32> + %up00_l = vector.fragment.load %d_lhs[%k0, %k0] shape [%k16, %k16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %up00_r = vector.fragment.load %ka_rhs[%k0, %k0] shape [%k16, %k16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %up00_m = vector.mma %up00_l, %up00_r, %up00_i : vector<16xf16>, vector<16xf16>, vector<8xf32> + %up01_i = vector.fragment %fz8 shape [%k16, %k16] : vector<8xf32> + %up01_l = vector.fragment.load %d_lhs[%k0, %k0] shape [%k16, %k16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %up01_r = vector.fragment.load %ka_rhs[%k0, %k16] shape [%k16, %k16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %up01_m = vector.mma %up01_l, %up01_r, %up01_i : vector<16xf16>, vector<16xf16>, vector<8xf32> + %up02_i = vector.fragment %fz8 shape [%k16, %k16] : vector<8xf32> + %up02_l = vector.fragment.load %d_lhs[%k0, %k0] shape [%k16, %k16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %up02_r = vector.fragment.load %ka_rhs[%k0, %k32] shape [%k16, %k16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %up02_m = vector.mma %up02_l, %up02_r, %up02_i : vector<16xf16>, vector<16xf16>, vector<8xf32> + %up03_i = vector.fragment %fz8 shape [%k16, %k16] : vector<8xf32> + %up03_l = vector.fragment.load %d_lhs[%k0, %k0] shape [%k16, %k16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %up03_r = vector.fragment.load %ka_rhs[%k0, %k48] shape [%k16, %k16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %up03_m = vector.mma %up03_l, %up03_r, %up03_i : vector<16xf16>, vector<16xf16>, vector<8xf32> + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %up00_m, %upd_tile[%k0, %k0] shape [%k16, %k16] : vector<8xf32>, view<16x68xf32, %lay_upd> + vector.fragment.store %up01_m, %upd_tile[%k0, %k16] shape [%k16, %k16] : vector<8xf32>, view<16x68xf32, %lay_upd> + vector.fragment.store %up02_m, %upd_tile[%k0, %k32] shape [%k16, %k16] : vector<8xf32>, view<16x68xf32, %lay_upd> + vector.fragment.store %up03_m, %upd_tile[%k0, %k48] shape [%k16, %k16] : vector<8xf32>, view<16x68xf32, %lay_upd> + kernel.barrier scope(subgroup) ordering(acq_rel) + %ur00_m = index.mul %k0, %k68 : index + %ur00_a = index.add %ur00_m, %rowh : index + %ur00_b = index.assume %ur00_a [range(%ur00_a, 0, 1084)] : index + %ur00 = vector.load %upd_flat[%ur00_b] : view<1088xf32> -> vector<4xf32> + %ur01_m = index.mul %k1, %k68 : index + %ur01_a = index.add %ur01_m, %rowh : index + %ur01_b = index.assume %ur01_a [range(%ur01_a, 0, 1084)] : index + %ur01 = vector.load %upd_flat[%ur01_b] : view<1088xf32> -> vector<4xf32> + %ur02_m = index.mul %k2, %k68 : index + %ur02_a = index.add %ur02_m, %rowh : index + %ur02_b = index.assume %ur02_a [range(%ur02_a, 0, 1084)] : index + %ur02 = vector.load %upd_flat[%ur02_b] : view<1088xf32> -> vector<4xf32> + %ur03_m = index.mul %k3, %k68 : index + %ur03_a = index.add %ur03_m, %rowh : index + %ur03_b = index.assume %ur03_a [range(%ur03_a, 0, 1084)] : index + %ur03 = vector.load %upd_flat[%ur03_b] : view<1088xf32> -> vector<4xf32> + %ur04_m = index.mul %k4, %k68 : index + %ur04_a = index.add %ur04_m, %rowh : index + %ur04_b = index.assume %ur04_a [range(%ur04_a, 0, 1084)] : index + %ur04 = vector.load %upd_flat[%ur04_b] : view<1088xf32> -> vector<4xf32> + %ur05_m = index.mul %k5, %k68 : index + %ur05_a = index.add %ur05_m, %rowh : index + %ur05_b = index.assume %ur05_a [range(%ur05_a, 0, 1084)] : index + %ur05 = vector.load %upd_flat[%ur05_b] : view<1088xf32> -> vector<4xf32> + %ur06_m = index.mul %k6, %k68 : index + %ur06_a = index.add %ur06_m, %rowh : index + %ur06_b = index.assume %ur06_a [range(%ur06_a, 0, 1084)] : index + %ur06 = vector.load %upd_flat[%ur06_b] : view<1088xf32> -> vector<4xf32> + %ur07_m = index.mul %k7, %k68 : index + %ur07_a = index.add %ur07_m, %rowh : index + %ur07_b = index.assume %ur07_a [range(%ur07_a, 0, 1084)] : index + %ur07 = vector.load %upd_flat[%ur07_b] : view<1088xf32> -> vector<4xf32> + %ur08_m = index.mul %k8, %k68 : index + %ur08_a = index.add %ur08_m, %rowh : index + %ur08_b = index.assume %ur08_a [range(%ur08_a, 0, 1084)] : index + %ur08 = vector.load %upd_flat[%ur08_b] : view<1088xf32> -> vector<4xf32> + %ur09_m = index.mul %k9, %k68 : index + %ur09_a = index.add %ur09_m, %rowh : index + %ur09_b = index.assume %ur09_a [range(%ur09_a, 0, 1084)] : index + %ur09 = vector.load %upd_flat[%ur09_b] : view<1088xf32> -> vector<4xf32> + %ur010_m = index.mul %k10, %k68 : index + %ur010_a = index.add %ur010_m, %rowh : index + %ur010_b = index.assume %ur010_a [range(%ur010_a, 0, 1084)] : index + %ur010 = vector.load %upd_flat[%ur010_b] : view<1088xf32> -> vector<4xf32> + %ur011_m = index.mul %k11, %k68 : index + %ur011_a = index.add %ur011_m, %rowh : index + %ur011_b = index.assume %ur011_a [range(%ur011_a, 0, 1084)] : index + %ur011 = vector.load %upd_flat[%ur011_b] : view<1088xf32> -> vector<4xf32> + %ur012_m = index.mul %k12, %k68 : index + %ur012_a = index.add %ur012_m, %rowh : index + %ur012_b = index.assume %ur012_a [range(%ur012_a, 0, 1084)] : index + %ur012 = vector.load %upd_flat[%ur012_b] : view<1088xf32> -> vector<4xf32> + %ur013_m = index.mul %k13, %k68 : index + %ur013_a = index.add %ur013_m, %rowh : index + %ur013_b = index.assume %ur013_a [range(%ur013_a, 0, 1084)] : index + %ur013 = vector.load %upd_flat[%ur013_b] : view<1088xf32> -> vector<4xf32> + %ur014_m = index.mul %k14, %k68 : index + %ur014_a = index.add %ur014_m, %rowh : index + %ur014_b = index.assume %ur014_a [range(%ur014_a, 0, 1084)] : index + %ur014 = vector.load %upd_flat[%ur014_b] : view<1088xf32> -> vector<4xf32> + %ur015_m = index.mul %k15, %k68 : index + %ur015_a = index.add %ur015_m, %rowh : index + %ur015_b = index.assume %ur015_a [range(%ur015_a, 0, 1084)] : index + %ur015 = vector.load %upd_flat[%ur015_b] : view<1088xf32> -> vector<4xf32> + %up10_i = vector.fragment %fz8 shape [%k16, %k16] : vector<8xf32> + %up10_l = vector.fragment.load %d_lhs[%k0, %k0] shape [%k16, %k16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %up10_r = vector.fragment.load %ka_rhs[%k0, %k64] shape [%k16, %k16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %up10_m = vector.mma %up10_l, %up10_r, %up10_i : vector<16xf16>, vector<16xf16>, vector<8xf32> + %up11_i = vector.fragment %fz8 shape [%k16, %k16] : vector<8xf32> + %up11_l = vector.fragment.load %d_lhs[%k0, %k0] shape [%k16, %k16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %up11_r = vector.fragment.load %ka_rhs[%k0, %k80] shape [%k16, %k16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %up11_m = vector.mma %up11_l, %up11_r, %up11_i : vector<16xf16>, vector<16xf16>, vector<8xf32> + %up12_i = vector.fragment %fz8 shape [%k16, %k16] : vector<8xf32> + %up12_l = vector.fragment.load %d_lhs[%k0, %k0] shape [%k16, %k16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %up12_r = vector.fragment.load %ka_rhs[%k0, %k96] shape [%k16, %k16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %up12_m = vector.mma %up12_l, %up12_r, %up12_i : vector<16xf16>, vector<16xf16>, vector<8xf32> + %up13_i = vector.fragment %fz8 shape [%k16, %k16] : vector<8xf32> + %up13_l = vector.fragment.load %d_lhs[%k0, %k0] shape [%k16, %k16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %up13_r = vector.fragment.load %ka_rhs[%k0, %k112] shape [%k16, %k16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %up13_m = vector.mma %up13_l, %up13_r, %up13_i : vector<16xf16>, vector<16xf16>, vector<8xf32> + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %up10_m, %upd_tile[%k0, %k0] shape [%k16, %k16] : vector<8xf32>, view<16x68xf32, %lay_upd> + vector.fragment.store %up11_m, %upd_tile[%k0, %k16] shape [%k16, %k16] : vector<8xf32>, view<16x68xf32, %lay_upd> + vector.fragment.store %up12_m, %upd_tile[%k0, %k32] shape [%k16, %k16] : vector<8xf32>, view<16x68xf32, %lay_upd> + vector.fragment.store %up13_m, %upd_tile[%k0, %k48] shape [%k16, %k16] : vector<8xf32>, view<16x68xf32, %lay_upd> + kernel.barrier scope(subgroup) ordering(acq_rel) + %ur10_m = index.mul %k0, %k68 : index + %ur10_a = index.add %ur10_m, %rowh : index + %ur10_b = index.assume %ur10_a [range(%ur10_a, 0, 1084)] : index + %ur10 = vector.load %upd_flat[%ur10_b] : view<1088xf32> -> vector<4xf32> + %ur11_m = index.mul %k1, %k68 : index + %ur11_a = index.add %ur11_m, %rowh : index + %ur11_b = index.assume %ur11_a [range(%ur11_a, 0, 1084)] : index + %ur11 = vector.load %upd_flat[%ur11_b] : view<1088xf32> -> vector<4xf32> + %ur12_m = index.mul %k2, %k68 : index + %ur12_a = index.add %ur12_m, %rowh : index + %ur12_b = index.assume %ur12_a [range(%ur12_a, 0, 1084)] : index + %ur12 = vector.load %upd_flat[%ur12_b] : view<1088xf32> -> vector<4xf32> + %ur13_m = index.mul %k3, %k68 : index + %ur13_a = index.add %ur13_m, %rowh : index + %ur13_b = index.assume %ur13_a [range(%ur13_a, 0, 1084)] : index + %ur13 = vector.load %upd_flat[%ur13_b] : view<1088xf32> -> vector<4xf32> + %ur14_m = index.mul %k4, %k68 : index + %ur14_a = index.add %ur14_m, %rowh : index + %ur14_b = index.assume %ur14_a [range(%ur14_a, 0, 1084)] : index + %ur14 = vector.load %upd_flat[%ur14_b] : view<1088xf32> -> vector<4xf32> + %ur15_m = index.mul %k5, %k68 : index + %ur15_a = index.add %ur15_m, %rowh : index + %ur15_b = index.assume %ur15_a [range(%ur15_a, 0, 1084)] : index + %ur15 = vector.load %upd_flat[%ur15_b] : view<1088xf32> -> vector<4xf32> + %ur16_m = index.mul %k6, %k68 : index + %ur16_a = index.add %ur16_m, %rowh : index + %ur16_b = index.assume %ur16_a [range(%ur16_a, 0, 1084)] : index + %ur16 = vector.load %upd_flat[%ur16_b] : view<1088xf32> -> vector<4xf32> + %ur17_m = index.mul %k7, %k68 : index + %ur17_a = index.add %ur17_m, %rowh : index + %ur17_b = index.assume %ur17_a [range(%ur17_a, 0, 1084)] : index + %ur17 = vector.load %upd_flat[%ur17_b] : view<1088xf32> -> vector<4xf32> + %ur18_m = index.mul %k8, %k68 : index + %ur18_a = index.add %ur18_m, %rowh : index + %ur18_b = index.assume %ur18_a [range(%ur18_a, 0, 1084)] : index + %ur18 = vector.load %upd_flat[%ur18_b] : view<1088xf32> -> vector<4xf32> + %ur19_m = index.mul %k9, %k68 : index + %ur19_a = index.add %ur19_m, %rowh : index + %ur19_b = index.assume %ur19_a [range(%ur19_a, 0, 1084)] : index + %ur19 = vector.load %upd_flat[%ur19_b] : view<1088xf32> -> vector<4xf32> + %ur110_m = index.mul %k10, %k68 : index + %ur110_a = index.add %ur110_m, %rowh : index + %ur110_b = index.assume %ur110_a [range(%ur110_a, 0, 1084)] : index + %ur110 = vector.load %upd_flat[%ur110_b] : view<1088xf32> -> vector<4xf32> + %ur111_m = index.mul %k11, %k68 : index + %ur111_a = index.add %ur111_m, %rowh : index + %ur111_b = index.assume %ur111_a [range(%ur111_a, 0, 1084)] : index + %ur111 = vector.load %upd_flat[%ur111_b] : view<1088xf32> -> vector<4xf32> + %ur112_m = index.mul %k12, %k68 : index + %ur112_a = index.add %ur112_m, %rowh : index + %ur112_b = index.assume %ur112_a [range(%ur112_a, 0, 1084)] : index + %ur112 = vector.load %upd_flat[%ur112_b] : view<1088xf32> -> vector<4xf32> + %ur113_m = index.mul %k13, %k68 : index + %ur113_a = index.add %ur113_m, %rowh : index + %ur113_b = index.assume %ur113_a [range(%ur113_a, 0, 1084)] : index + %ur113 = vector.load %upd_flat[%ur113_b] : view<1088xf32> -> vector<4xf32> + %ur114_m = index.mul %k14, %k68 : index + %ur114_a = index.add %ur114_m, %rowh : index + %ur114_b = index.assume %ur114_a [range(%ur114_a, 0, 1084)] : index + %ur114 = vector.load %upd_flat[%ur114_b] : view<1088xf32> -> vector<4xf32> + %ur115_m = index.mul %k15, %k68 : index + %ur115_a = index.add %ur115_m, %rowh : index + %ur115_b = index.assume %ur115_a [range(%ur115_a, 0, 1084)] : index + %ur115 = vector.load %upd_flat[%ur115_b] : view<1088xf32> -> vector<4xf32> + + %usel0 = scf.select %lane_lo, %ur00, %ur10 : vector<4xf32> + %sdec0 = vector.mulf %sa0, %clv : vector<4xf32> + %snew0 = vector.addf %sdec0, %usel0 : vector<4xf32> + %usel1 = scf.select %lane_lo, %ur01, %ur11 : vector<4xf32> + %sdec1 = vector.mulf %sa1, %clv : vector<4xf32> + %snew1 = vector.addf %sdec1, %usel1 : vector<4xf32> + %usel2 = scf.select %lane_lo, %ur02, %ur12 : vector<4xf32> + %sdec2 = vector.mulf %sa2, %clv : vector<4xf32> + %snew2 = vector.addf %sdec2, %usel2 : vector<4xf32> + %usel3 = scf.select %lane_lo, %ur03, %ur13 : vector<4xf32> + %sdec3 = vector.mulf %sa3, %clv : vector<4xf32> + %snew3 = vector.addf %sdec3, %usel3 : vector<4xf32> + %usel4 = scf.select %lane_lo, %ur04, %ur14 : vector<4xf32> + %sdec4 = vector.mulf %sa4, %clv : vector<4xf32> + %snew4 = vector.addf %sdec4, %usel4 : vector<4xf32> + %usel5 = scf.select %lane_lo, %ur05, %ur15 : vector<4xf32> + %sdec5 = vector.mulf %sa5, %clv : vector<4xf32> + %snew5 = vector.addf %sdec5, %usel5 : vector<4xf32> + %usel6 = scf.select %lane_lo, %ur06, %ur16 : vector<4xf32> + %sdec6 = vector.mulf %sa6, %clv : vector<4xf32> + %snew6 = vector.addf %sdec6, %usel6 : vector<4xf32> + %usel7 = scf.select %lane_lo, %ur07, %ur17 : vector<4xf32> + %sdec7 = vector.mulf %sa7, %clv : vector<4xf32> + %snew7 = vector.addf %sdec7, %usel7 : vector<4xf32> + %usel8 = scf.select %lane_lo, %ur08, %ur18 : vector<4xf32> + %sdec8 = vector.mulf %sa8, %clv : vector<4xf32> + %snew8 = vector.addf %sdec8, %usel8 : vector<4xf32> + %usel9 = scf.select %lane_lo, %ur09, %ur19 : vector<4xf32> + %sdec9 = vector.mulf %sa9, %clv : vector<4xf32> + %snew9 = vector.addf %sdec9, %usel9 : vector<4xf32> + %usel10 = scf.select %lane_lo, %ur010, %ur110 : vector<4xf32> + %sdec10 = vector.mulf %sa10, %clv : vector<4xf32> + %snew10 = vector.addf %sdec10, %usel10 : vector<4xf32> + %usel11 = scf.select %lane_lo, %ur011, %ur111 : vector<4xf32> + %sdec11 = vector.mulf %sa11, %clv : vector<4xf32> + %snew11 = vector.addf %sdec11, %usel11 : vector<4xf32> + %usel12 = scf.select %lane_lo, %ur012, %ur112 : vector<4xf32> + %sdec12 = vector.mulf %sa12, %clv : vector<4xf32> + %snew12 = vector.addf %sdec12, %usel12 : vector<4xf32> + %usel13 = scf.select %lane_lo, %ur013, %ur113 : vector<4xf32> + %sdec13 = vector.mulf %sa13, %clv : vector<4xf32> + %snew13 = vector.addf %sdec13, %usel13 : vector<4xf32> + %usel14 = scf.select %lane_lo, %ur014, %ur114 : vector<4xf32> + %sdec14 = vector.mulf %sa14, %clv : vector<4xf32> + %snew14 = vector.addf %sdec14, %usel14 : vector<4xf32> + %usel15 = scf.select %lane_lo, %ur015, %ur115 : vector<4xf32> + %sdec15 = vector.mulf %sa15, %clv : vector<4xf32> + %snew15 = vector.addf %sdec15, %usel15 : vector<4xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %snew0, %snew1, %snew2, %snew3, %snew4, %snew5, %snew6, %snew7, %snew8, %snew9, %snew10, %snew11, %snew12, %snew13, %snew14, %snew15 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + scf.yield %snapshot_next0, %snapshot_next1, %snapshot_next2, %snapshot_next3, %snapshot_next4, %snapshot_next5, %snapshot_next6, %snapshot_next7, %snapshot_next8, %snapshot_next9, %snapshot_next10, %snapshot_next11, %snapshot_next12, %snapshot_next13, %snapshot_next14, %snapshot_next15 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + + scf.if %reuse_snapshot_fragments { + } else { + scf.if %publish_snapshots { + vector.store %sf0, %snapshot_cache_view[%icb0] : vector<4xf32>, view<1073741824xf32> + vector.store %sf1, %snapshot_cache_view[%icb1] : vector<4xf32>, view<1073741824xf32> + vector.store %sf2, %snapshot_cache_view[%icb2] : vector<4xf32>, view<1073741824xf32> + vector.store %sf3, %snapshot_cache_view[%icb3] : vector<4xf32>, view<1073741824xf32> + vector.store %sf4, %snapshot_cache_view[%icb4] : vector<4xf32>, view<1073741824xf32> + vector.store %sf5, %snapshot_cache_view[%icb5] : vector<4xf32>, view<1073741824xf32> + vector.store %sf6, %snapshot_cache_view[%icb6] : vector<4xf32>, view<1073741824xf32> + vector.store %sf7, %snapshot_cache_view[%icb7] : vector<4xf32>, view<1073741824xf32> + vector.store %sf8, %snapshot_cache_view[%icb8] : vector<4xf32>, view<1073741824xf32> + vector.store %sf9, %snapshot_cache_view[%icb9] : vector<4xf32>, view<1073741824xf32> + vector.store %sf10, %snapshot_cache_view[%icb10] : vector<4xf32>, view<1073741824xf32> + vector.store %sf11, %snapshot_cache_view[%icb11] : vector<4xf32>, view<1073741824xf32> + vector.store %sf12, %snapshot_cache_view[%icb12] : vector<4xf32>, view<1073741824xf32> + vector.store %sf13, %snapshot_cache_view[%icb13] : vector<4xf32>, view<1073741824xf32> + vector.store %sf14, %snapshot_cache_view[%icb14] : vector<4xf32>, view<1073741824xf32> + vector.store %sf15, %snapshot_cache_view[%icb15] : vector<4xf32>, view<1073741824xf32> + } else { + scf.if %state_inplace { + // All state shards have already been consumed, and each workgroup owns a + // disjoint head/column range, so publication can safely reuse the cache row. + vector.store %sf0, %si_view[%icb0] : vector<4xf32>, view<1073741824xf32> + vector.store %sf1, %si_view[%icb1] : vector<4xf32>, view<1073741824xf32> + vector.store %sf2, %si_view[%icb2] : vector<4xf32>, view<1073741824xf32> + vector.store %sf3, %si_view[%icb3] : vector<4xf32>, view<1073741824xf32> + vector.store %sf4, %si_view[%icb4] : vector<4xf32>, view<1073741824xf32> + vector.store %sf5, %si_view[%icb5] : vector<4xf32>, view<1073741824xf32> + vector.store %sf6, %si_view[%icb6] : vector<4xf32>, view<1073741824xf32> + vector.store %sf7, %si_view[%icb7] : vector<4xf32>, view<1073741824xf32> + vector.store %sf8, %si_view[%icb8] : vector<4xf32>, view<1073741824xf32> + vector.store %sf9, %si_view[%icb9] : vector<4xf32>, view<1073741824xf32> + vector.store %sf10, %si_view[%icb10] : vector<4xf32>, view<1073741824xf32> + vector.store %sf11, %si_view[%icb11] : vector<4xf32>, view<1073741824xf32> + vector.store %sf12, %si_view[%icb12] : vector<4xf32>, view<1073741824xf32> + vector.store %sf13, %si_view[%icb13] : vector<4xf32>, view<1073741824xf32> + vector.store %sf14, %si_view[%icb14] : vector<4xf32>, view<1073741824xf32> + vector.store %sf15, %si_view[%icb15] : vector<4xf32>, view<1073741824xf32> + } else { + // The ordinary graph contract appends updated state after attention. + %woc0_m = index.mul %wcol, %s_v : index + %woc0_a = index.add %state_offset, %woc0_m : index + %woc0_e = index.add %woc0_a, %attn_elems : index + %woc0_r = index.add %woc0_e, %row0 : index + %woc0 = index.assume %woc0_r [range(%woc0_r, 0, 1073741820)] : index + vector.store %sf0, %dst_view[%woc0] : vector<4xf32>, view<1073741824xf32> + %woc1_r = index.add %woc0, %k128 : index + %woc1 = index.assume %woc1_r [range(%woc1_r, 0, 1073741820)] : index + vector.store %sf1, %dst_view[%woc1] : vector<4xf32>, view<1073741824xf32> + %woc2_r = index.add %woc1, %k128 : index + %woc2 = index.assume %woc2_r [range(%woc2_r, 0, 1073741820)] : index + vector.store %sf2, %dst_view[%woc2] : vector<4xf32>, view<1073741824xf32> + %woc3_r = index.add %woc2, %k128 : index + %woc3 = index.assume %woc3_r [range(%woc3_r, 0, 1073741820)] : index + vector.store %sf3, %dst_view[%woc3] : vector<4xf32>, view<1073741824xf32> + %woc4_r = index.add %woc3, %k128 : index + %woc4 = index.assume %woc4_r [range(%woc4_r, 0, 1073741820)] : index + vector.store %sf4, %dst_view[%woc4] : vector<4xf32>, view<1073741824xf32> + %woc5_r = index.add %woc4, %k128 : index + %woc5 = index.assume %woc5_r [range(%woc5_r, 0, 1073741820)] : index + vector.store %sf5, %dst_view[%woc5] : vector<4xf32>, view<1073741824xf32> + %woc6_r = index.add %woc5, %k128 : index + %woc6 = index.assume %woc6_r [range(%woc6_r, 0, 1073741820)] : index + vector.store %sf6, %dst_view[%woc6] : vector<4xf32>, view<1073741824xf32> + %woc7_r = index.add %woc6, %k128 : index + %woc7 = index.assume %woc7_r [range(%woc7_r, 0, 1073741820)] : index + vector.store %sf7, %dst_view[%woc7] : vector<4xf32>, view<1073741824xf32> + %woc8_r = index.add %woc7, %k128 : index + %woc8 = index.assume %woc8_r [range(%woc8_r, 0, 1073741820)] : index + vector.store %sf8, %dst_view[%woc8] : vector<4xf32>, view<1073741824xf32> + %woc9_r = index.add %woc8, %k128 : index + %woc9 = index.assume %woc9_r [range(%woc9_r, 0, 1073741820)] : index + vector.store %sf9, %dst_view[%woc9] : vector<4xf32>, view<1073741824xf32> + %woc10_r = index.add %woc9, %k128 : index + %woc10 = index.assume %woc10_r [range(%woc10_r, 0, 1073741820)] : index + vector.store %sf10, %dst_view[%woc10] : vector<4xf32>, view<1073741824xf32> + %woc11_r = index.add %woc10, %k128 : index + %woc11 = index.assume %woc11_r [range(%woc11_r, 0, 1073741820)] : index + vector.store %sf11, %dst_view[%woc11] : vector<4xf32>, view<1073741824xf32> + %woc12_r = index.add %woc11, %k128 : index + %woc12 = index.assume %woc12_r [range(%woc12_r, 0, 1073741820)] : index + vector.store %sf12, %dst_view[%woc12] : vector<4xf32>, view<1073741824xf32> + %woc13_r = index.add %woc12, %k128 : index + %woc13 = index.assume %woc13_r [range(%woc13_r, 0, 1073741820)] : index + vector.store %sf13, %dst_view[%woc13] : vector<4xf32>, view<1073741824xf32> + %woc14_r = index.add %woc13, %k128 : index + %woc14 = index.assume %woc14_r [range(%woc14_r, 0, 1073741820)] : index + vector.store %sf14, %dst_view[%woc14] : vector<4xf32>, view<1073741824xf32> + %woc15_r = index.add %woc14, %k128 : index + %woc15 = index.assume %woc15_r [range(%woc15_r, 0, 1073741820)] : index + vector.store %sf15, %dst_view[%woc15] : vector<4xf32>, view<1073741824xf32> + } + } + } + scf.if %publish_q8 { + %t0 = index.constant 0 : index + // One native subgroup handles each local attention row, using the + // standalone RMS reduction and unary implementation without a global read. + %post_rows0 = index.mul %n_tokens, %n_heads : index + %post_rows = index.mul %post_rows0, %n_seqs : index + %post_tokens = index.mul %n_tokens, %n_seqs : index + %post_groups = index.mul %n_heads, %k8 : index + %post_channel0 = index.mul %lane, %k4 : index + %post_channel = index.assume %post_channel0 [range(%post_channel0, 0, 124), mul(%post_channel0, 4)] : index + %post_attention = buffer.view %lds[%base] : buffer -> view<16x128xf32> + %post_weights = buffer.view %rms_weight[%base] : buffer -> view<128xf32> + %post_gates = buffer.view %raw_gate[%base] : buffer -> view<[%post_rows]x128xf32> + %post_f32 = buffer.view %norm_output[%base] : buffer -> view<[%post_rows]x128xf32> + %post_f16 = buffer.view %half_output[%base] : buffer -> view<[%post_groups]x[%post_tokens]x16xf16> + scf.for %post_local0 = [%wave to %k16 step %k8] { + %post_local = index.assume %post_local0 [range(%post_local0, 0, 15)] : index + %post_time0 = index.add %t0, %post_local : index + %post_valid = index.cmp ult, %post_time0, %n_tokens : index + %post_time_safe = scf.select %post_valid, %post_time0, %c0 : index + %post_time = index.assume %post_time_safe [lt(%post_time_safe, %n_tokens)] : index + %post_token0 = index.madd %seq, %n_tokens, %post_time : index + %post_token = index.assume %post_token0 [lt(%post_token0, %post_tokens)] : index + %post_row0 = index.madd %post_token, %n_heads, %h_idx : index + %post_row = index.assume %post_row0 [lt(%post_row0, %post_rows)] : index + %post_scale = template.apply<@ggml.rmsnorm_f32.subgroup_row_scale>(%k16, %post_local, %k128, %rms_epsilon, %lds) : (index, index, index, f32, buffer) -> (f32) + scf.if %post_valid { + %post_x = vector.load %post_attention[%post_local, %post_channel] : view<16x128xf32> -> vector<4xf32> + %post_w = vector.load %post_weights[%post_channel] : view<128xf32> -> vector<4xf32> + %post_g = vector.load %post_gates[%post_row, %post_channel] : view<[%post_rows]x128xf32> -> vector<4xf32> + %post_scale_v = vector.splat %post_scale : vector<4xf32> + %post_normalized = vector.mulf %post_x, %post_scale_v : vector<4xf32> + %post_weighted = vector.mulf %post_normalized, %post_w : vector<4xf32> + %post_activated = template.apply<@ggml.unary_f32.apply_vector4>(%rms_gate_op, %post_g) : (index, vector<4xf32>) -> (vector<4xf32>) + %post_result = vector.mulf %post_weighted, %post_activated : vector<4xf32> + vector.store %post_result, %post_f32[%post_row, %post_channel] : vector<4xf32>, view<[%post_rows]x128xf32> + scf.if %publish_q8 { + %post_flat_channel = index.madd %post_row, %k128, %post_channel : index + %post_scratch_values = buffer.view %lds[%base] : buffer -> view<256xf32> + %post_scratch_d = buffer.view %lds[%base] : buffer -> view<32xf32> + template.apply<@ggml.quantize_q8_1_x4.publish_vector4_strict>(%post_valid, %base, %post_flat_channel, %post_result, %post_scratch_values, %post_scratch_d, %half_output) : (i1, offset, index, vector<4xf32>, view<256xf32>, view<32xf32>, buffer) + } else { + %post_half = vector.fptrunc %post_result : vector<4xf32> to vector<4xf16> + %post_head_channel = index.madd %h_idx, %k128, %post_channel : index + %post_group = index.div %post_head_channel, %k16 : index + %post_group_lane = index.rem %post_head_channel, %k16 : index + vector.store %post_half, %post_f16[%post_group, %post_token, %post_group_lane] : vector<4xf16>, view<[%post_groups]x[%post_tokens]x16xf16> + } + } + } + + } + + template.return +} + +// Prefix-state publisher shared by the four-row rollback specialization. +// Stage each resident state half in logical result layout, then combine the +// exact WMMA result in registers and publish it directly to the rollback cache. +template.def<@llm.gated_delta_net.snapshot_prefix.body> device @llm_gated_delta_net_snapshot_prefix_body(%state_factor: f32, %cache_base: index, %state_offset: index, %wcol: index, %row0: index, %rowh: index, %wave: index, %lane: index, %lds: buffer, %snapshot_cache: buffer, %sa0: vector<4xf32>, %sa1: vector<4xf32>, %sa2: vector<4xf32>, %sa3: vector<4xf32>, %sa4: vector<4xf32>, %sa5: vector<4xf32>, %sa6: vector<4xf32>, %sa7: vector<4xf32>, %sa8: vector<4xf32>, %sa9: vector<4xf32>, %sa10: vector<4xf32>, %sa11: vector<4xf32>, %sa12: vector<4xf32>, %sa13: vector<4xf32>, %sa14: vector<4xf32>, %sa15: vector<4xf32>, %fz8: vector<8xf32>) { + %c0o = index.constant 0 : offset + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c48 = index.constant 48 : index + %c64 = index.constant 64 : index + %c68 = index.constant 68 : index + %c80 = index.constant 80 : index + %c96 = index.constant 96 : index + %c112 = index.constant 112 : index + %c128 = index.constant 128 : index + %c4352 = index.constant 4352 : index + %c6528 = index.constant 6528 : index + %c11072 = index.constant 11072 : index + %ka_o = index.constant 4352 : offset + %lay_dcol = encoding.layout.strided [24, 1] : encoding + %lay_katok = encoding.layout.strided [1, 24] : encoding + %lay_upd = encoding.layout.strided [68, 1] : encoding + %lay_state = encoding.layout.strided [128, 1] : encoding + %wave_bytes = index.mul %wave, %c6528 : index + %wave_base = index.add %wave_bytes, %c11072 : index + %wave_off = index.cast %wave_base : index to offset + %d_base_i = index.add %wave_base, %c4352 : index + %d_base = index.cast %d_base_i : index to offset + %d_lhs = buffer.view %lds[%d_base] : buffer -> view<16x16xf16, %lay_dcol> + %ka_rhs = buffer.view %lds[%ka_o] : buffer -> view<16x128xf16, %lay_katok> + %upd_flat = buffer.view %lds[%wave_off] : buffer -> view<1088xf32> + %upd_tile = buffer.view %lds[%wave_off] : buffer -> view<16x68xf32, %lay_upd> + %cache_g = buffer.assume.memory_space %snapshot_cache : buffer + %cache_state_base = index.add %cache_base, %state_offset : index + %wave_state = index.mul %wcol, %c128 : index + %cache_wave_base = index.add %cache_state_base, %wave_state : index + %cache_byte_i = index.mul %cache_wave_base, %c4 : index + %cache_byte = index.cast %cache_byte_i : index to offset + %cache_view = buffer.view %cache_g[%cache_byte] : buffer -> view<16x128xf32, %lay_state> + %state_factor_v = vector.splat %state_factor : vector<8xf32> + %lane_low = index.cmp ult, %lane, %c16 : index + %lane_high = index.cmp uge, %lane, %c16 : index + + scf.if %lane_low { + vector.store %sa0, %upd_flat[%rowh] : vector<4xf32>, view<1088xf32> + %lo_s1_off = index.constant 68 : index + %lo_s1_raw = index.add %lo_s1_off, %rowh : index + %lo_s1_idx = index.assume %lo_s1_raw [range(%lo_s1_raw, 68, 128)] : index + vector.store %sa1, %upd_flat[%lo_s1_idx] : vector<4xf32>, view<1088xf32> + %lo_s2_off = index.constant 136 : index + %lo_s2_raw = index.add %lo_s2_off, %rowh : index + %lo_s2_idx = index.assume %lo_s2_raw [range(%lo_s2_raw, 136, 196)] : index + vector.store %sa2, %upd_flat[%lo_s2_idx] : vector<4xf32>, view<1088xf32> + %lo_s3_off = index.constant 204 : index + %lo_s3_raw = index.add %lo_s3_off, %rowh : index + %lo_s3_idx = index.assume %lo_s3_raw [range(%lo_s3_raw, 204, 264)] : index + vector.store %sa3, %upd_flat[%lo_s3_idx] : vector<4xf32>, view<1088xf32> + %lo_s4_off = index.constant 272 : index + %lo_s4_raw = index.add %lo_s4_off, %rowh : index + %lo_s4_idx = index.assume %lo_s4_raw [range(%lo_s4_raw, 272, 332)] : index + vector.store %sa4, %upd_flat[%lo_s4_idx] : vector<4xf32>, view<1088xf32> + %lo_s5_off = index.constant 340 : index + %lo_s5_raw = index.add %lo_s5_off, %rowh : index + %lo_s5_idx = index.assume %lo_s5_raw [range(%lo_s5_raw, 340, 400)] : index + vector.store %sa5, %upd_flat[%lo_s5_idx] : vector<4xf32>, view<1088xf32> + %lo_s6_off = index.constant 408 : index + %lo_s6_raw = index.add %lo_s6_off, %rowh : index + %lo_s6_idx = index.assume %lo_s6_raw [range(%lo_s6_raw, 408, 468)] : index + vector.store %sa6, %upd_flat[%lo_s6_idx] : vector<4xf32>, view<1088xf32> + %lo_s7_off = index.constant 476 : index + %lo_s7_raw = index.add %lo_s7_off, %rowh : index + %lo_s7_idx = index.assume %lo_s7_raw [range(%lo_s7_raw, 476, 536)] : index + vector.store %sa7, %upd_flat[%lo_s7_idx] : vector<4xf32>, view<1088xf32> + %lo_s8_off = index.constant 544 : index + %lo_s8_raw = index.add %lo_s8_off, %rowh : index + %lo_s8_idx = index.assume %lo_s8_raw [range(%lo_s8_raw, 544, 604)] : index + vector.store %sa8, %upd_flat[%lo_s8_idx] : vector<4xf32>, view<1088xf32> + %lo_s9_off = index.constant 612 : index + %lo_s9_raw = index.add %lo_s9_off, %rowh : index + %lo_s9_idx = index.assume %lo_s9_raw [range(%lo_s9_raw, 612, 672)] : index + vector.store %sa9, %upd_flat[%lo_s9_idx] : vector<4xf32>, view<1088xf32> + %lo_s10_off = index.constant 680 : index + %lo_s10_raw = index.add %lo_s10_off, %rowh : index + %lo_s10_idx = index.assume %lo_s10_raw [range(%lo_s10_raw, 680, 740)] : index + vector.store %sa10, %upd_flat[%lo_s10_idx] : vector<4xf32>, view<1088xf32> + %lo_s11_off = index.constant 748 : index + %lo_s11_raw = index.add %lo_s11_off, %rowh : index + %lo_s11_idx = index.assume %lo_s11_raw [range(%lo_s11_raw, 748, 808)] : index + vector.store %sa11, %upd_flat[%lo_s11_idx] : vector<4xf32>, view<1088xf32> + %lo_s12_off = index.constant 816 : index + %lo_s12_raw = index.add %lo_s12_off, %rowh : index + %lo_s12_idx = index.assume %lo_s12_raw [range(%lo_s12_raw, 816, 876)] : index + vector.store %sa12, %upd_flat[%lo_s12_idx] : vector<4xf32>, view<1088xf32> + %lo_s13_off = index.constant 884 : index + %lo_s13_raw = index.add %lo_s13_off, %rowh : index + %lo_s13_idx = index.assume %lo_s13_raw [range(%lo_s13_raw, 884, 944)] : index + vector.store %sa13, %upd_flat[%lo_s13_idx] : vector<4xf32>, view<1088xf32> + %lo_s14_off = index.constant 952 : index + %lo_s14_raw = index.add %lo_s14_off, %rowh : index + %lo_s14_idx = index.assume %lo_s14_raw [range(%lo_s14_raw, 952, 1012)] : index + vector.store %sa14, %upd_flat[%lo_s14_idx] : vector<4xf32>, view<1088xf32> + %lo_s15_off = index.constant 1020 : index + %lo_s15_raw = index.add %lo_s15_off, %rowh : index + %lo_s15_idx = index.assume %lo_s15_raw [range(%lo_s15_raw, 1020, 1080)] : index + vector.store %sa15, %upd_flat[%lo_s15_idx] : vector<4xf32>, view<1088xf32> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + %lo0_state = vector.fragment.load %upd_tile[%c0, %c0] shape [%c16, %c16] : view<16x68xf32, %lay_upd> -> vector<8xf32> + %lo0_scaled = vector.mulf %lo0_state, %state_factor_v : vector<8xf32> + %lo0_init = vector.fragment %fz8 shape [%c16, %c16] : vector<8xf32> + %lo0_lhs = vector.fragment.load %d_lhs[%c0, %c0] shape [%c16, %c16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %lo0_rhs = vector.fragment.load %ka_rhs[%c0, %c0] shape [%c16, %c16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %lo0_mma = vector.mma %lo0_lhs, %lo0_rhs, %lo0_init : vector<16xf16>, vector<16xf16>, vector<8xf32> + %lo0_result = vector.addf %lo0_scaled, %lo0_mma : vector<8xf32> + vector.fragment.store %lo0_result, %cache_view[%c0, %c0] shape [%c16, %c16] : vector<8xf32>, view<16x128xf32, %lay_state> + %lo1_state = vector.fragment.load %upd_tile[%c0, %c16] shape [%c16, %c16] : view<16x68xf32, %lay_upd> -> vector<8xf32> + %lo1_scaled = vector.mulf %lo1_state, %state_factor_v : vector<8xf32> + %lo1_init = vector.fragment %fz8 shape [%c16, %c16] : vector<8xf32> + %lo1_lhs = vector.fragment.load %d_lhs[%c0, %c0] shape [%c16, %c16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %lo1_rhs = vector.fragment.load %ka_rhs[%c0, %c16] shape [%c16, %c16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %lo1_mma = vector.mma %lo1_lhs, %lo1_rhs, %lo1_init : vector<16xf16>, vector<16xf16>, vector<8xf32> + %lo1_result = vector.addf %lo1_scaled, %lo1_mma : vector<8xf32> + vector.fragment.store %lo1_result, %cache_view[%c0, %c16] shape [%c16, %c16] : vector<8xf32>, view<16x128xf32, %lay_state> + %lo2_state = vector.fragment.load %upd_tile[%c0, %c32] shape [%c16, %c16] : view<16x68xf32, %lay_upd> -> vector<8xf32> + %lo2_scaled = vector.mulf %lo2_state, %state_factor_v : vector<8xf32> + %lo2_init = vector.fragment %fz8 shape [%c16, %c16] : vector<8xf32> + %lo2_lhs = vector.fragment.load %d_lhs[%c0, %c0] shape [%c16, %c16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %lo2_rhs = vector.fragment.load %ka_rhs[%c0, %c32] shape [%c16, %c16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %lo2_mma = vector.mma %lo2_lhs, %lo2_rhs, %lo2_init : vector<16xf16>, vector<16xf16>, vector<8xf32> + %lo2_result = vector.addf %lo2_scaled, %lo2_mma : vector<8xf32> + vector.fragment.store %lo2_result, %cache_view[%c0, %c32] shape [%c16, %c16] : vector<8xf32>, view<16x128xf32, %lay_state> + %lo3_state = vector.fragment.load %upd_tile[%c0, %c48] shape [%c16, %c16] : view<16x68xf32, %lay_upd> -> vector<8xf32> + %lo3_scaled = vector.mulf %lo3_state, %state_factor_v : vector<8xf32> + %lo3_init = vector.fragment %fz8 shape [%c16, %c16] : vector<8xf32> + %lo3_lhs = vector.fragment.load %d_lhs[%c0, %c0] shape [%c16, %c16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %lo3_rhs = vector.fragment.load %ka_rhs[%c0, %c48] shape [%c16, %c16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %lo3_mma = vector.mma %lo3_lhs, %lo3_rhs, %lo3_init : vector<16xf16>, vector<16xf16>, vector<8xf32> + %lo3_result = vector.addf %lo3_scaled, %lo3_mma : vector<8xf32> + vector.fragment.store %lo3_result, %cache_view[%c0, %c48] shape [%c16, %c16] : vector<8xf32>, view<16x128xf32, %lay_state> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %lane_high { + vector.store %sa0, %upd_flat[%rowh] : vector<4xf32>, view<1088xf32> + %hi_s1_off = index.constant 68 : index + %hi_s1_raw = index.add %hi_s1_off, %rowh : index + %hi_s1_idx = index.assume %hi_s1_raw [range(%hi_s1_raw, 68, 128)] : index + vector.store %sa1, %upd_flat[%hi_s1_idx] : vector<4xf32>, view<1088xf32> + %hi_s2_off = index.constant 136 : index + %hi_s2_raw = index.add %hi_s2_off, %rowh : index + %hi_s2_idx = index.assume %hi_s2_raw [range(%hi_s2_raw, 136, 196)] : index + vector.store %sa2, %upd_flat[%hi_s2_idx] : vector<4xf32>, view<1088xf32> + %hi_s3_off = index.constant 204 : index + %hi_s3_raw = index.add %hi_s3_off, %rowh : index + %hi_s3_idx = index.assume %hi_s3_raw [range(%hi_s3_raw, 204, 264)] : index + vector.store %sa3, %upd_flat[%hi_s3_idx] : vector<4xf32>, view<1088xf32> + %hi_s4_off = index.constant 272 : index + %hi_s4_raw = index.add %hi_s4_off, %rowh : index + %hi_s4_idx = index.assume %hi_s4_raw [range(%hi_s4_raw, 272, 332)] : index + vector.store %sa4, %upd_flat[%hi_s4_idx] : vector<4xf32>, view<1088xf32> + %hi_s5_off = index.constant 340 : index + %hi_s5_raw = index.add %hi_s5_off, %rowh : index + %hi_s5_idx = index.assume %hi_s5_raw [range(%hi_s5_raw, 340, 400)] : index + vector.store %sa5, %upd_flat[%hi_s5_idx] : vector<4xf32>, view<1088xf32> + %hi_s6_off = index.constant 408 : index + %hi_s6_raw = index.add %hi_s6_off, %rowh : index + %hi_s6_idx = index.assume %hi_s6_raw [range(%hi_s6_raw, 408, 468)] : index + vector.store %sa6, %upd_flat[%hi_s6_idx] : vector<4xf32>, view<1088xf32> + %hi_s7_off = index.constant 476 : index + %hi_s7_raw = index.add %hi_s7_off, %rowh : index + %hi_s7_idx = index.assume %hi_s7_raw [range(%hi_s7_raw, 476, 536)] : index + vector.store %sa7, %upd_flat[%hi_s7_idx] : vector<4xf32>, view<1088xf32> + %hi_s8_off = index.constant 544 : index + %hi_s8_raw = index.add %hi_s8_off, %rowh : index + %hi_s8_idx = index.assume %hi_s8_raw [range(%hi_s8_raw, 544, 604)] : index + vector.store %sa8, %upd_flat[%hi_s8_idx] : vector<4xf32>, view<1088xf32> + %hi_s9_off = index.constant 612 : index + %hi_s9_raw = index.add %hi_s9_off, %rowh : index + %hi_s9_idx = index.assume %hi_s9_raw [range(%hi_s9_raw, 612, 672)] : index + vector.store %sa9, %upd_flat[%hi_s9_idx] : vector<4xf32>, view<1088xf32> + %hi_s10_off = index.constant 680 : index + %hi_s10_raw = index.add %hi_s10_off, %rowh : index + %hi_s10_idx = index.assume %hi_s10_raw [range(%hi_s10_raw, 680, 740)] : index + vector.store %sa10, %upd_flat[%hi_s10_idx] : vector<4xf32>, view<1088xf32> + %hi_s11_off = index.constant 748 : index + %hi_s11_raw = index.add %hi_s11_off, %rowh : index + %hi_s11_idx = index.assume %hi_s11_raw [range(%hi_s11_raw, 748, 808)] : index + vector.store %sa11, %upd_flat[%hi_s11_idx] : vector<4xf32>, view<1088xf32> + %hi_s12_off = index.constant 816 : index + %hi_s12_raw = index.add %hi_s12_off, %rowh : index + %hi_s12_idx = index.assume %hi_s12_raw [range(%hi_s12_raw, 816, 876)] : index + vector.store %sa12, %upd_flat[%hi_s12_idx] : vector<4xf32>, view<1088xf32> + %hi_s13_off = index.constant 884 : index + %hi_s13_raw = index.add %hi_s13_off, %rowh : index + %hi_s13_idx = index.assume %hi_s13_raw [range(%hi_s13_raw, 884, 944)] : index + vector.store %sa13, %upd_flat[%hi_s13_idx] : vector<4xf32>, view<1088xf32> + %hi_s14_off = index.constant 952 : index + %hi_s14_raw = index.add %hi_s14_off, %rowh : index + %hi_s14_idx = index.assume %hi_s14_raw [range(%hi_s14_raw, 952, 1012)] : index + vector.store %sa14, %upd_flat[%hi_s14_idx] : vector<4xf32>, view<1088xf32> + %hi_s15_off = index.constant 1020 : index + %hi_s15_raw = index.add %hi_s15_off, %rowh : index + %hi_s15_idx = index.assume %hi_s15_raw [range(%hi_s15_raw, 1020, 1080)] : index + vector.store %sa15, %upd_flat[%hi_s15_idx] : vector<4xf32>, view<1088xf32> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + %hi0_state = vector.fragment.load %upd_tile[%c0, %c0] shape [%c16, %c16] : view<16x68xf32, %lay_upd> -> vector<8xf32> + %hi0_scaled = vector.mulf %hi0_state, %state_factor_v : vector<8xf32> + %hi0_init = vector.fragment %fz8 shape [%c16, %c16] : vector<8xf32> + %hi0_lhs = vector.fragment.load %d_lhs[%c0, %c0] shape [%c16, %c16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %hi0_rhs = vector.fragment.load %ka_rhs[%c0, %c64] shape [%c16, %c16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %hi0_mma = vector.mma %hi0_lhs, %hi0_rhs, %hi0_init : vector<16xf16>, vector<16xf16>, vector<8xf32> + %hi0_result = vector.addf %hi0_scaled, %hi0_mma : vector<8xf32> + vector.fragment.store %hi0_result, %cache_view[%c0, %c64] shape [%c16, %c16] : vector<8xf32>, view<16x128xf32, %lay_state> + %hi1_state = vector.fragment.load %upd_tile[%c0, %c16] shape [%c16, %c16] : view<16x68xf32, %lay_upd> -> vector<8xf32> + %hi1_scaled = vector.mulf %hi1_state, %state_factor_v : vector<8xf32> + %hi1_init = vector.fragment %fz8 shape [%c16, %c16] : vector<8xf32> + %hi1_lhs = vector.fragment.load %d_lhs[%c0, %c0] shape [%c16, %c16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %hi1_rhs = vector.fragment.load %ka_rhs[%c0, %c80] shape [%c16, %c16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %hi1_mma = vector.mma %hi1_lhs, %hi1_rhs, %hi1_init : vector<16xf16>, vector<16xf16>, vector<8xf32> + %hi1_result = vector.addf %hi1_scaled, %hi1_mma : vector<8xf32> + vector.fragment.store %hi1_result, %cache_view[%c0, %c80] shape [%c16, %c16] : vector<8xf32>, view<16x128xf32, %lay_state> + %hi2_state = vector.fragment.load %upd_tile[%c0, %c32] shape [%c16, %c16] : view<16x68xf32, %lay_upd> -> vector<8xf32> + %hi2_scaled = vector.mulf %hi2_state, %state_factor_v : vector<8xf32> + %hi2_init = vector.fragment %fz8 shape [%c16, %c16] : vector<8xf32> + %hi2_lhs = vector.fragment.load %d_lhs[%c0, %c0] shape [%c16, %c16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %hi2_rhs = vector.fragment.load %ka_rhs[%c0, %c96] shape [%c16, %c16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %hi2_mma = vector.mma %hi2_lhs, %hi2_rhs, %hi2_init : vector<16xf16>, vector<16xf16>, vector<8xf32> + %hi2_result = vector.addf %hi2_scaled, %hi2_mma : vector<8xf32> + vector.fragment.store %hi2_result, %cache_view[%c0, %c96] shape [%c16, %c16] : vector<8xf32>, view<16x128xf32, %lay_state> + %hi3_state = vector.fragment.load %upd_tile[%c0, %c48] shape [%c16, %c16] : view<16x68xf32, %lay_upd> -> vector<8xf32> + %hi3_scaled = vector.mulf %hi3_state, %state_factor_v : vector<8xf32> + %hi3_init = vector.fragment %fz8 shape [%c16, %c16] : vector<8xf32> + %hi3_lhs = vector.fragment.load %d_lhs[%c0, %c0] shape [%c16, %c16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %hi3_rhs = vector.fragment.load %ka_rhs[%c0, %c112] shape [%c16, %c16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %hi3_mma = vector.mma %hi3_lhs, %hi3_rhs, %hi3_init : vector<16xf16>, vector<16xf16>, vector<8xf32> + %hi3_result = vector.addf %hi3_scaled, %hi3_mma : vector<8xf32> + vector.fragment.store %hi3_result, %cache_view[%c0, %c112] shape [%c16, %c16] : vector<8xf32>, view<16x128xf32, %lay_state> + template.return +} + +template.def<@llm.gated_delta_net.snapshot_fragment_prefix.body> device @llm_gated_delta_net_snapshot_fragment_prefix_body(%state_factor: f32, %cache_base: index, %state_offset: index, %wcol: index, %row0: index, %rowh: index, %wave: index, %lane: index, %lds: buffer, %snapshot_cache: buffer, %lo0_state: vector<8xf32>, %lo1_state: vector<8xf32>, %lo2_state: vector<8xf32>, %lo3_state: vector<8xf32>, %hi0_state: vector<8xf32>, %hi1_state: vector<8xf32>, %hi2_state: vector<8xf32>, %hi3_state: vector<8xf32>, %fz8: vector<8xf32>) { + %c0o = index.constant 0 : offset + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c48 = index.constant 48 : index + %c64 = index.constant 64 : index + %c68 = index.constant 68 : index + %c80 = index.constant 80 : index + %c96 = index.constant 96 : index + %c112 = index.constant 112 : index + %c128 = index.constant 128 : index + %c4352 = index.constant 4352 : index + %c6528 = index.constant 6528 : index + %c11072 = index.constant 11072 : index + %ka_o = index.constant 4352 : offset + %lay_dcol = encoding.layout.strided [24, 1] : encoding + %lay_katok = encoding.layout.strided [1, 24] : encoding + %lay_upd = encoding.layout.strided [68, 1] : encoding + %lay_state = encoding.layout.strided [128, 1] : encoding + %wave_bytes = index.mul %wave, %c6528 : index + %wave_base = index.add %wave_bytes, %c11072 : index + %wave_off = index.cast %wave_base : index to offset + %d_base_i = index.add %wave_base, %c4352 : index + %d_base = index.cast %d_base_i : index to offset + %d_lhs = buffer.view %lds[%d_base] : buffer -> view<16x16xf16, %lay_dcol> + %ka_rhs = buffer.view %lds[%ka_o] : buffer -> view<16x128xf16, %lay_katok> + %upd_flat = buffer.view %lds[%wave_off] : buffer -> view<1088xf32> + %upd_tile = buffer.view %lds[%wave_off] : buffer -> view<16x68xf32, %lay_upd> + %cache_g = buffer.assume.memory_space %snapshot_cache : buffer + %cache_state_base = index.add %cache_base, %state_offset : index + %wave_state = index.mul %wcol, %c128 : index + %cache_wave_base = index.add %cache_state_base, %wave_state : index + %cache_byte_i = index.mul %cache_wave_base, %c4 : index + %cache_byte = index.cast %cache_byte_i : index to offset + %cache_view = buffer.view %cache_g[%cache_byte] : buffer -> view<16x128xf32, %lay_state> + %state_factor_v = vector.splat %state_factor : vector<8xf32> + %lane_low = index.cmp ult, %lane, %c16 : index + %lane_high = index.cmp uge, %lane, %c16 : index + + %lo0_scaled = vector.mulf %lo0_state, %state_factor_v : vector<8xf32> + %lo0_init = vector.fragment %fz8 shape [%c16, %c16] : vector<8xf32> + %lo0_lhs = vector.fragment.load %d_lhs[%c0, %c0] shape [%c16, %c16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %lo0_rhs = vector.fragment.load %ka_rhs[%c0, %c0] shape [%c16, %c16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %lo0_mma = vector.mma %lo0_lhs, %lo0_rhs, %lo0_init : vector<16xf16>, vector<16xf16>, vector<8xf32> + %lo0_result = vector.addf %lo0_scaled, %lo0_mma : vector<8xf32> + vector.fragment.store %lo0_result, %cache_view[%c0, %c0] shape [%c16, %c16] : vector<8xf32>, view<16x128xf32, %lay_state> + %lo1_scaled = vector.mulf %lo1_state, %state_factor_v : vector<8xf32> + %lo1_init = vector.fragment %fz8 shape [%c16, %c16] : vector<8xf32> + %lo1_lhs = vector.fragment.load %d_lhs[%c0, %c0] shape [%c16, %c16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %lo1_rhs = vector.fragment.load %ka_rhs[%c0, %c16] shape [%c16, %c16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %lo1_mma = vector.mma %lo1_lhs, %lo1_rhs, %lo1_init : vector<16xf16>, vector<16xf16>, vector<8xf32> + %lo1_result = vector.addf %lo1_scaled, %lo1_mma : vector<8xf32> + vector.fragment.store %lo1_result, %cache_view[%c0, %c16] shape [%c16, %c16] : vector<8xf32>, view<16x128xf32, %lay_state> + %lo2_scaled = vector.mulf %lo2_state, %state_factor_v : vector<8xf32> + %lo2_init = vector.fragment %fz8 shape [%c16, %c16] : vector<8xf32> + %lo2_lhs = vector.fragment.load %d_lhs[%c0, %c0] shape [%c16, %c16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %lo2_rhs = vector.fragment.load %ka_rhs[%c0, %c32] shape [%c16, %c16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %lo2_mma = vector.mma %lo2_lhs, %lo2_rhs, %lo2_init : vector<16xf16>, vector<16xf16>, vector<8xf32> + %lo2_result = vector.addf %lo2_scaled, %lo2_mma : vector<8xf32> + vector.fragment.store %lo2_result, %cache_view[%c0, %c32] shape [%c16, %c16] : vector<8xf32>, view<16x128xf32, %lay_state> + %lo3_scaled = vector.mulf %lo3_state, %state_factor_v : vector<8xf32> + %lo3_init = vector.fragment %fz8 shape [%c16, %c16] : vector<8xf32> + %lo3_lhs = vector.fragment.load %d_lhs[%c0, %c0] shape [%c16, %c16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %lo3_rhs = vector.fragment.load %ka_rhs[%c0, %c48] shape [%c16, %c16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %lo3_mma = vector.mma %lo3_lhs, %lo3_rhs, %lo3_init : vector<16xf16>, vector<16xf16>, vector<8xf32> + %lo3_result = vector.addf %lo3_scaled, %lo3_mma : vector<8xf32> + vector.fragment.store %lo3_result, %cache_view[%c0, %c48] shape [%c16, %c16] : vector<8xf32>, view<16x128xf32, %lay_state> + kernel.barrier scope(subgroup) ordering(acq_rel) + %hi0_scaled = vector.mulf %hi0_state, %state_factor_v : vector<8xf32> + %hi0_init = vector.fragment %fz8 shape [%c16, %c16] : vector<8xf32> + %hi0_lhs = vector.fragment.load %d_lhs[%c0, %c0] shape [%c16, %c16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %hi0_rhs = vector.fragment.load %ka_rhs[%c0, %c64] shape [%c16, %c16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %hi0_mma = vector.mma %hi0_lhs, %hi0_rhs, %hi0_init : vector<16xf16>, vector<16xf16>, vector<8xf32> + %hi0_result = vector.addf %hi0_scaled, %hi0_mma : vector<8xf32> + vector.fragment.store %hi0_result, %cache_view[%c0, %c64] shape [%c16, %c16] : vector<8xf32>, view<16x128xf32, %lay_state> + %hi1_scaled = vector.mulf %hi1_state, %state_factor_v : vector<8xf32> + %hi1_init = vector.fragment %fz8 shape [%c16, %c16] : vector<8xf32> + %hi1_lhs = vector.fragment.load %d_lhs[%c0, %c0] shape [%c16, %c16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %hi1_rhs = vector.fragment.load %ka_rhs[%c0, %c80] shape [%c16, %c16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %hi1_mma = vector.mma %hi1_lhs, %hi1_rhs, %hi1_init : vector<16xf16>, vector<16xf16>, vector<8xf32> + %hi1_result = vector.addf %hi1_scaled, %hi1_mma : vector<8xf32> + vector.fragment.store %hi1_result, %cache_view[%c0, %c80] shape [%c16, %c16] : vector<8xf32>, view<16x128xf32, %lay_state> + %hi2_scaled = vector.mulf %hi2_state, %state_factor_v : vector<8xf32> + %hi2_init = vector.fragment %fz8 shape [%c16, %c16] : vector<8xf32> + %hi2_lhs = vector.fragment.load %d_lhs[%c0, %c0] shape [%c16, %c16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %hi2_rhs = vector.fragment.load %ka_rhs[%c0, %c96] shape [%c16, %c16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %hi2_mma = vector.mma %hi2_lhs, %hi2_rhs, %hi2_init : vector<16xf16>, vector<16xf16>, vector<8xf32> + %hi2_result = vector.addf %hi2_scaled, %hi2_mma : vector<8xf32> + vector.fragment.store %hi2_result, %cache_view[%c0, %c96] shape [%c16, %c16] : vector<8xf32>, view<16x128xf32, %lay_state> + %hi3_scaled = vector.mulf %hi3_state, %state_factor_v : vector<8xf32> + %hi3_init = vector.fragment %fz8 shape [%c16, %c16] : vector<8xf32> + %hi3_lhs = vector.fragment.load %d_lhs[%c0, %c0] shape [%c16, %c16] : view<16x16xf16, %lay_dcol> -> vector<16xf16> + %hi3_rhs = vector.fragment.load %ka_rhs[%c0, %c112] shape [%c16, %c16] : view<16x128xf16, %lay_katok> -> vector<16xf16> + %hi3_mma = vector.mma %hi3_lhs, %hi3_rhs, %hi3_init : vector<16xf16>, vector<16xf16>, vector<8xf32> + %hi3_result = vector.addf %hi3_scaled, %hi3_mma : vector<8xf32> + vector.fragment.store %hi3_result, %cache_view[%c0, %c112] shape [%c16, %c16] : vector<8xf32>, view<16x128xf32, %lay_state> + template.return +} + +// Quantized alpha/beta projection epilogue. +config.decl @llm.gated_delta_net.epilogue_head_count : %value: index where [range(%value, 1, 4096)] + +config.decl @llm.gated_delta_net.epilogue_element_count : %value: index where [range(%value, 1, 1048576)] + +config.decl @llm.gated_delta_net.epilogue_workgroup_size : %value: index where [range(%value, 32, 1024), mul(%value, 32)] + +kernel.def export("llm_gated_delta_net_projection_epilogue_f32") @llm_gated_delta_net_projection_epilogue_f32() { + %unit = index.constant 1 : index + %cneg1 = index.constant -1 : index + %elements = config.get @llm.gated_delta_net.epilogue_element_count : index + %workgroup_size = config.get @llm.gated_delta_net.epilogue_workgroup_size : index + %rounding = index.add %workgroup_size, %cneg1 : index + %rounded = index.add %elements, %rounding : index + %workgroups = index.div %rounded, %workgroup_size : index + kernel.launch.config workgroups(%workgroups, %unit, %unit) workgroup_size(%workgroup_size, %unit, %unit) : index +} launch(%alpha_raw: buffer, %beta_raw: buffer, %bias: buffer, %a_scale: buffer, %gate_dst: buffer, %beta_dst: buffer) { + %base = index.constant 0 : offset + %elements0 = config.get @llm.gated_delta_net.epilogue_element_count : index + %heads0 = config.get @llm.gated_delta_net.epilogue_head_count : index + %workgroup_size = config.get @llm.gated_delta_net.epilogue_workgroup_size : index + %elements = index.assume %elements0 [range(%elements0, 1, 1048576)] : index + %heads = index.assume %heads0 [range(%heads0, 1, 4096)] : index + %workgroup0 = kernel.workgroup.id : index + %workgroup = index.assume %workgroup0 [range(%workgroup0, 0, 32767)] : index + %lane0 = kernel.workitem.id : index + %lane = index.assume %lane0 [range(%lane0, 0, 1023)] : index + %linear_base = index.mul %workgroup, %workgroup_size : index + %linear0 = index.add %linear_base, %lane : index + %linear = index.assume %linear0 [range(%linear0, 0, 1048575)] : index + %in_bounds = index.cmp ult, %linear, %elements : index + + %alpha_global = buffer.assume.memory_space %alpha_raw : buffer + %beta_global = buffer.assume.memory_space %beta_raw : buffer + %bias_global = buffer.assume.memory_space %bias : buffer + %a_scale_global = buffer.assume.memory_space %a_scale : buffer + %gate_global = buffer.assume.memory_space %gate_dst : buffer + %beta_dst_global = buffer.assume.memory_space %beta_dst : buffer + %alpha_noalias, %beta_noalias, %bias_noalias, %a_scale_noalias, %gate_noalias, %beta_dst_noalias = buffer.assume.noalias %alpha_global, %beta_global, %bias_global, %a_scale_global, %gate_global, %beta_dst_global : buffer, buffer, buffer, buffer, buffer, buffer + %alpha_view = buffer.view %alpha_noalias[%base] : buffer -> view<1048576xf32> + %beta_view = buffer.view %beta_noalias[%base] : buffer -> view<1048576xf32> + %bias_view = buffer.view %bias_noalias[%base] : buffer -> view<4096xf32> + %a_scale_view = buffer.view %a_scale_noalias[%base] : buffer -> view<4096xf32> + %gate_view = buffer.view %gate_noalias[%base] : buffer -> view<1048576xf32> + %beta_dst_view = buffer.view %beta_dst_noalias[%base] : buffer -> view<1048576xf32> + + scf.if %in_bounds { + %head = index.rem %linear, %heads : index + %alpha = view.load %alpha_view[%linear] : view<1048576xf32> -> f32 + %beta = view.load %beta_view[%linear] : view<1048576xf32> -> f32 + %bias_value = view.load %bias_view[%head] : view<4096xf32> -> f32 + %a_value = view.load %a_scale_view[%head] : view<4096xf32> -> f32 + %c0 = scalar.constant 0.0 : f32 + %c1 = scalar.constant 1.0 : f32 + %biased = scalar.addf %alpha, %bias_value : f32 + %abs = scalar.absf %biased : f32 + %neg_abs = scalar.negf %abs : f32 + %exp = scalar.expf %neg_abs : f32 + %one_plus = scalar.addf %c1, %exp : f32 + %log = scalar.logf %one_plus : f32 + %positive = scalar.maxnumf %biased, %c0 : f32 + %softplus = scalar.addf %positive, %log : f32 + %gate = scalar.mulf %softplus, %a_value : f32 + %beta_result = scalar.logisticf %beta : f32 + view.store %gate, %gate_view[%linear] : f32, view<1048576xf32> + view.store %beta_result, %beta_dst_view[%linear] : f32, view<1048576xf32> + } + kernel.return +} + +kernel.def export("llm_gated_delta_net_f32_wmma_head128_rmsnorm_gate") @llm_gated_delta_net_f32_wmma_head128_rmsnorm_gate() { + %unit = index.constant 1 : index + %wg = index.constant 256 : index + %blk = index.constant 128 : index + %blk_m = index.constant 127 : index + %s_v = config.get @llm.gated_delta_net.head_width : index + %n_heads = config.get @llm.gated_delta_net.head_count : index + %n_seqs = config.get @llm.gated_delta_net.sequence_count : index + %cr = index.add %s_v, %blk_m : index + %col_blocks = index.div %cr, %blk : index + kernel.launch.config workgroups(%n_heads, %n_seqs, %col_blocks) workgroup_size(%wg, %unit, %unit) : index +} launch(%q: buffer, %k: buffer, %v: buffer, %g: buffer, %beta: buffer, %state_in: buffer, %dst: buffer, %rms_weight: buffer, %raw_gate: buffer, %norm_output: buffer, %half_output: buffer) { + %state_inplace = scalar.constant false : i1 + %publish_snapshots = scalar.constant false : i1 + %projection_epilogue = scalar.constant false : i1 + %snapshot_stride = index.constant 1 : index + %scale = scalar.constant 0.0883883461356163 : f32 + %rmsnorm_gate = scalar.constant true : i1 + %rms_epsilon = config.get @ggml.rmsnorm_gate_f32.rms_epsilon : f32 + %rms_gate_op = config.get @ggml.rmsnorm_gate_f32.gate_op : index + %select_state = scalar.constant false : i1 + %state_row_count = index.constant 1 : index + %publish_q8 = scalar.constant false : i1 + template.apply<@llm.gated_delta_net.f32_wmma_head128.body>(%state_inplace, %publish_snapshots, %projection_epilogue, %snapshot_stride, %scale, %q, %k, %v, %g, %beta, %g, %beta, %state_in, %dst, %state_in, %rmsnorm_gate, %rms_epsilon, %rms_gate_op, %rms_weight, %raw_gate, %norm_output, %half_output, %select_state, %state_row_count, %dst, %publish_q8) : (i1, i1, i1, index, f32, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, i1, f32, index, buffer, buffer, buffer, buffer, i1, index, buffer, i1) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/get_rows_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/get_rows_f32.loom new file mode 100644 index 000000000000..c2f12072f465 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/get_rows_f32.loom @@ -0,0 +1,189 @@ +// Generic GGML GET_ROWS for contiguous rows decoded through F32. +// +// One workitem publishes four adjacent output channels for one requested row. +// Storage-format-specific row decoding is delegated to the common dequant +// motif so Q1_0, Q3_K, Q4_K, Q5_K, Q6_K, IQ4_XS, Q8_0, Q8_1, F16, BF16, and F32 share the same gather shape. +template.decl @ggml.get_rows_f32.body(%token_count: index, %row_count: index, %hidden_size: index, %token_ids: buffer, %weight: buffer, %output: buffer) + +template.decl @ggml.get_rows_f32.launch(%hidden_capacity: index, %token_capacity: index) -> (index, index, index, index) + +template.decl @ggml.get_rows_f32.load_vector4(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight_format: index, %publish: i1, %valid_token_id: i1, %token: index, %row_count: index, %hidden_size: index, %channel: index, %weight: buffer) -> (vector<4xf32>) + +template.decl @ggml.get_rows_f32.next_body(%token_count: index, %row_count: index, %hidden_size: index, %token_ids: buffer, %weight: buffer, %output: buffer, %next_output: buffer) + +template.decl @ggml.publish_f32.f32_vector4(%arg0: i1, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: vector<4xf32>, %arg6: buffer) + +template.decl @ggml.publish_f32.next_vector4(%arg0: index, %arg1: i1, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: vector<4xf32>, %arg7: view<256xf32>, %arg8: view<32xf32>, %arg9: buffer) + +amdgpu.target @ggml_get_rows_f32_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.get_rows_f32.token_capacity : %value: index where [range(%value, 1, 2048)] + +config.decl @ggml.get_rows_f32.hidden_capacity : %value: index where [range(%value, 4, 1073741824), mul(%value, 4)] + +config.decl @ggml.get_rows_f32.weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.get_rows_f32.next_format : %value: index where [range(%value, 16, 81)] + +func.decl @ggml_dequant_weight_tile_bytes(%weight_format: index) -> (offset) + +func.decl @ggml_dequant_weight_row_bytes(%weight_format: index, %hidden_size: index) -> (offset) + +func.decl @ggml_iq4nl_table_i8() -> (vector<16xi8>) + +func.decl @ggml_iq3s_grid_f32_halves() -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + +func.decl @ggml_dequant_f32_vector4(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight_format: index, %weight: buffer, %row_byte_base: offset, %input_size: index, %quant_block: index, %quant_group: index, %packet: index, %k: index) -> (vector<4xf32>) + +template.def<@ggml.get_rows_f32.load_vector4> device @ggml_get_rows_f32_load_vector4(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight_format: index, %publish: i1, %valid_token_id: i1, %token: index, %row_count: index, %hidden_size: index, %channel: index, %weight: buffer) -> (vector<4xf32>) { + %c4 = index.constant 4 : index + %c32 = index.constant 32 : index + %c256 = index.constant 256 : index + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %token_id = index.assume %token [range(%token, 0, 262207)] : index + %bounded_token_id, %bounded_row_count = index.assume %token_id, %row_count [lt(%token_id, %row_count)] : index, index + %safe_values = scf.if %publish -> (vector<4xf32>) { + %row_bytes = func.call @ggml_dequant_weight_row_bytes(%weight_format, %hidden_size) : (index, index) -> (offset) + %row_byte_base = index.scale %bounded_token_id, %row_bytes : index, offset -> offset + %quant_block = index.div %channel, %c256 : index + %block_channel = index.rem %channel, %c256 : index + %quant_group = index.div %block_channel, %c32 : index + %group_channel = index.rem %block_channel, %c32 : index + %quant_packet = index.div %group_channel, %c4 : index + %values = func.call @ggml_dequant_f32_vector4(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight_format, %weight, %row_byte_base, %hidden_size, %quant_block, %quant_group, %quant_packet, %channel) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, buffer, offset, index, index, index, index, index) -> (vector<4xf32>) + %row_values = scf.select %valid_token_id, %values, %c0_f32x4 : vector<4xf32> + scf.yield %row_values : vector<4xf32> + } else { + scf.yield %c0_f32x4 : vector<4xf32> + } + template.return %safe_values : vector<4xf32> +} + +template.def<@ggml.get_rows_f32.body> device @ggml_get_rows_f32_body(%token_count: index, %row_count: index, %hidden_size: index, %token_ids: buffer, %weight: buffer, %output: buffer) { + %token_capacity = config.get @ggml.get_rows_f32.token_capacity : index + %hidden_capacity = config.get @ggml.get_rows_f32.hidden_capacity : index + %weight_format = config.get @ggml.get_rows_f32.weight_format : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %bounded_row_count = index.assume %row_count [range(%row_count, 1, 262208)] : index + %bounded_hidden_size = index.assume %hidden_size [range(%hidden_size, 4, 1073741824), mul(%hidden_size, 4), le(%hidden_size, %hidden_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %c256 = index.constant 256 : index + %c0_i32 = scalar.constant 0 : i32 + %c0_offset = index.constant 0 : offset + %workitem = index.assume %workitem0 [range(%workitem0, 0, 255)] : index + %packet_tile_base = index.mul %channel_tile, %c256 : index + %packet0 = index.add %packet_tile_base, %workitem : index + %packet_count = index.div %bounded_hidden_size, %c4 : index + %valid_token = index.cmp ult, %token0, %bounded_token_count : index + %valid_packet = index.cmp ult, %packet0, %packet_count : index + %publish = scalar.andi %valid_token, %valid_packet : i1 + %token = scf.select %valid_token, %token0, %c0 : index + %packet = scf.select %valid_packet, %packet0, %c0 : index + %token_ids_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %token_ids, %weight, %output : buffer, buffer, buffer + %token_ids_view = buffer.view %token_ids_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]xi32> + %token_id_raw = view.load %token_ids_view[%token] : view<[%bounded_token_count]xi32> -> i32 + %token_id_nonnegative = scalar.cmpi sge, %token_id_raw, %c0_i32 : i32 + %safe_token_id_i32 = scf.select %token_id_nonnegative, %token_id_raw, %c0_i32 : i32 + %token_id = index.cast %safe_token_id_i32 : i32 to index + %channel = index.mul %packet, %c4 : index + %iq4nl_table = func.call @ggml_iq4nl_table_i8() : () -> (vector<16xi8>) + %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15 = func.call @ggml_iq3s_grid_f32_halves() : () -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + %values = template.apply<@ggml.get_rows_f32.load_vector4>(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight_format, %publish, %token_id_nonnegative, %token_id, %bounded_row_count, %bounded_hidden_size, %channel, %weight_noalias) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, i1, i1, index, index, index, index, buffer) -> (vector<4xf32>) + scf.if %publish { + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_hidden_size]xf32> + vector.store %values, %output_view[%token, %channel] : vector<4xf32>, view<[%bounded_token_count]x[%bounded_hidden_size]xf32> + } + template.return +} + +template.def<@ggml.get_rows_f32.next_body> device @ggml_get_rows_f32_next_body(%token_count: index, %row_count: index, %hidden_size: index, %token_ids: buffer, %weight: buffer, %output: buffer, %next_output: buffer) { + %token_capacity = config.get @ggml.get_rows_f32.token_capacity : index + %hidden_capacity = config.get @ggml.get_rows_f32.hidden_capacity : index + %weight_format = config.get @ggml.get_rows_f32.weight_format : index + %next_format = config.get @ggml.get_rows_f32.next_format : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %bounded_row_count = index.assume %row_count [range(%row_count, 1, 262208)] : index + %bounded_hidden_size = index.assume %hidden_size [range(%hidden_size, 256, 1073741824), mul(%hidden_size, 256), le(%hidden_size, %hidden_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %c256 = index.constant 256 : index + %c0_i32 = scalar.constant 0 : i32 + %c0_offset = index.constant 0 : offset + %scratch_d_byte_add = index.constant 1024 : offset + %scratch_bytes = index.constant 1152 : offset + %workitem = index.assume %workitem0 [range(%workitem0, 0, 255)] : index + %packet_tile_base = index.mul %channel_tile, %c256 : index + %packet0 = index.add %packet_tile_base, %workitem : index + %packet_count = index.div %bounded_hidden_size, %c4 : index + %valid_token = index.cmp ult, %token0, %bounded_token_count : index + %valid_packet = index.cmp ult, %packet0, %packet_count : index + %publish = scalar.andi %valid_token, %valid_packet : i1 + %token = scf.select %valid_token, %token0, %c0 : index + %packet = scf.select %valid_packet, %packet0, %c0 : index + %token_ids_noalias, %weight_noalias, %output_noalias, %next_output_noalias = buffer.assume.noalias %token_ids, %weight, %output, %next_output : buffer, buffer, buffer, buffer + %token_ids_view = buffer.view %token_ids_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]xi32> + %token_id_raw = view.load %token_ids_view[%token] : view<[%bounded_token_count]xi32> -> i32 + %token_id_nonnegative = scalar.cmpi sge, %token_id_raw, %c0_i32 : i32 + %safe_token_id_i32 = scf.select %token_id_nonnegative, %token_id_raw, %c0_i32 : i32 + %token_id = index.cast %safe_token_id_i32 : i32 to index + %channel = index.mul %packet, %c4 : index + %iq4nl_table = func.call @ggml_iq4nl_table_i8() : () -> (vector<16xi8>) + %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15 = func.call @ggml_iq3s_grid_f32_halves() : () -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + %values = template.apply<@ggml.get_rows_f32.load_vector4>(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight_format, %publish, %token_id_nonnegative, %token_id, %bounded_row_count, %bounded_hidden_size, %channel, %weight_noalias) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, i1, i1, index, index, index, index, buffer) -> (vector<4xf32>) + template.apply<@ggml.publish_f32.f32_vector4>(%publish, %token, %channel, %bounded_token_count, %bounded_hidden_size, %values, %output_noalias) : (i1, index, index, index, index, vector<4xf32>, buffer) + %scratch = buffer.alloca align(16) %scratch_bytes : buffer + %scratch_values = buffer.view %scratch[%c0_offset] : buffer -> view<256xf32> + %scratch_d = buffer.view %scratch[%scratch_d_byte_add] : buffer -> view<32xf32> + template.apply<@ggml.publish_f32.next_vector4>(%next_format, %publish, %token, %channel, %bounded_token_count, %bounded_hidden_size, %values, %scratch_values, %scratch_d, %next_output_noalias) : (index, i1, index, index, index, index, vector<4xf32>, view<256xf32>, view<32xf32>, buffer) + template.return +} + +template.def<@ggml.get_rows_f32.launch> @ggml_get_rows_f32_launch(%hidden_capacity: index, %token_capacity: index) -> (index, index, index, index) { + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %packet_count = index.div %hidden_capacity, %c4 : index + %padded_packet_count = index.add %packet_count, %c255 : index + %packet_tiles = index.div %padded_packet_count, %c256 : index + template.return %packet_tiles, %token_capacity, %c1, %c256 : index, index, index, index +} + +kernel.def target(@ggml_get_rows_f32_gfx11_wave64) @ggml_get_rows_f32(%token_count: index, %row_count: index, %hidden_size: index) { + %token_capacity = config.get @ggml.get_rows_f32.token_capacity : index + %hidden_capacity = config.get @ggml.get_rows_f32.hidden_capacity : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c255 = index.constant 255 : index + %workgroup_size = index.constant 256 : index + %packet_count = index.div %hidden_capacity, %c4 : index + %padded_packet_count = index.add %packet_count, %c255 : index + %packet_tiles = index.div %padded_packet_count, %workgroup_size : index + kernel.launch.config workgroups(%packet_tiles, %token_capacity, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %row_count: index, %hidden_size: index, %token_ids: buffer, %weight: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + template.apply<@ggml.get_rows_f32.body>(%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : (index, index, index, buffer, buffer, buffer) + kernel.return +} + +kernel.def target(@ggml_get_rows_f32_gfx11_wave64) @ggml_get_rows_f32_next(%token_count: index, %row_count: index, %hidden_size: index) { + %token_capacity = config.get @ggml.get_rows_f32.token_capacity : index + %hidden_capacity = config.get @ggml.get_rows_f32.hidden_capacity : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c255 = index.constant 255 : index + %workgroup_size = index.constant 256 : index + %packet_count = index.div %hidden_capacity, %c4 : index + %padded_packet_count = index.add %packet_count, %c255 : index + %packet_tiles = index.div %padded_packet_count, %workgroup_size : index + kernel.launch.config workgroups(%packet_tiles, %token_capacity, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %row_count: index, %hidden_size: index, %token_ids: buffer, %weight: buffer, %output: buffer, %next_output: buffer) where [range(%token_count, 1, 2048)] { + template.apply<@ggml.get_rows_f32.next_body>(%token_count, %row_count, %hidden_size, %token_ids, %weight, %output, %next_output) : (index, index, index, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/grouped_mul_mat_f16_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/grouped_mul_mat_f16_f32.loom new file mode 100644 index 000000000000..3042d9e87b05 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/grouped_mul_mat_f16_f32.loom @@ -0,0 +1,73 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// Batched ("grouped") MUL_MAT with F16 weights: for every group g and token m, +// output[g][m][n] = sum_k weight[g][n][k] * input[g][m][k] (F32 accumulation) +// The weight has one [N x K] matrix per group and no broadcast. ZAYA's CCA grouped +// convolution is one of these per tap (10 groups of 128 x 128). +// One workitem computes one output element; workgroups tile (N / 64, M, G). + +amdgpu.target @ggml_grouped_mul_mat_f16_f32_gfx11_wave64 {subgroup_size = 64} + +kernel.def target(@ggml_grouped_mul_mat_f16_f32_gfx11_wave64) export("ggml_grouped_mul_mat_f16_f32") @ggml_grouped_mul_mat_f16_f32(%input_size: index, %output_size: index, %token_count: index, %group_count: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %rounded = index.add %output_size, %rounding : index + %tiles = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%tiles, %token_count, %group_count) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%input_size: index, %output_size: index, %token_count: index, %group_count: index, %weight: buffer, %input: buffer, %output: buffer) { + %k = index.assume %input_size [range(%input_size, 1, 65536)] : index + %n = index.assume %output_size [range(%output_size, 1, 65536)] : index + %m = index.assume %token_count [range(%token_count, 1, 65536)] : index + %g = index.assume %group_count [range(%group_count, 1, 4096)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %c0_f32 = scalar.constant 0.0 : f32 + %zero_offset = index.constant 0 : offset + %tile = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %group0 = kernel.workgroup.id : index + %lane0 = kernel.workitem.id : index + %lane = index.assume %lane0 [range(%lane0, 0, 63)] : index + %token = index.assume %token0 [range(%token0, 0, 65535)] : index + %group = index.assume %group0 [range(%group0, 0, 4095)] : index + %tile_base = index.mul %tile, %c64 : index + %column0 = index.add %tile_base, %lane : index + %valid = index.cmp ult, %column0, %n : index + %column = scf.select %valid, %column0, %c0 : index + %weight_rows = index.mul %g, %n : index + %input_rows = index.mul %g, %m : index + %weight_row_base = index.mul %group, %n : index + %weight_row = index.add %weight_row_base, %column : index + %input_row_base = index.mul %group, %m : index + %input_row = index.add %input_row_base, %token : index + %weight_noalias, %input_noalias, %output_noalias = buffer.assume.noalias %weight, %input, %output : buffer, buffer, buffer + %weight_view = buffer.view %weight_noalias[%zero_offset] : buffer -> view<[%weight_rows]x[%k]xf16> + %input_view = buffer.view %input_noalias[%zero_offset] : buffer -> view<[%input_rows]x[%k]xf32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%input_rows]x[%n]xf32> + %sum = scf.for %channel = [%c0 to %k step %c1](%accumulator = %c0_f32 : f32) -> (f32) { + %weight_f16 = view.load %weight_view[%weight_row, %channel] : view<[%weight_rows]x[%k]xf16> -> f16 + %weight_value = scalar.extf %weight_f16 : f16 to f32 + %input_value = view.load %input_view[%input_row, %channel] : view<[%input_rows]x[%k]xf32> -> f32 + %next_accumulator = scalar.fmaf %weight_value, %input_value, %accumulator : f32 + scf.yield %next_accumulator : f32 + } + scf.if %valid { + view.store %sum, %output_view[%input_row, %column] : f32, view<[%input_rows]x[%n]xf32> + } + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/hadamard_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/hadamard_f32.loom new file mode 100644 index 000000000000..6644b73e2da5 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/hadamard_f32.loom @@ -0,0 +1,339 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// The normalized Walsh-Hadamard transform of each row: out = H x / sqrt(n), n = 2^log2_block. +// This is what a MUL_MAT carrying GGML_HINT_SRC0_IS_HADAMARD computes (llama-hadamard: the +// rotation folded into PrismML Bonsai weights; the CPU, Vulkan, CUDA and Metal backends run the +// same transform for that hint). The matrix itself is never read: n log2 n adds per row instead +// of n^2 multiply-adds, and no limit on the number of rows (the dense prefill matmul stops at 2048, +// which sent Bonsai's 512-token rotations, 512 x 17 rows of 1024, to the CPU every layer). +// +// One workgroup per row: the row goes into workgroup memory, log2(n) butterfly stages run with a +// barrier between them, and the scaled row is written back. A stage pairs index i (bit s clear) +// with i | 2^s; pair p maps to i = ((p >> s) << (s + 1)) | (p & (2^s - 1)), all shifts and masks, +// so no thread divides by a runtime value. The scale is exactly 2^(-k/2) for even k and +// 2^(-(k-1)/2) / sqrt(2) for odd k, matching the CPU's 1 / sqrtf(n). + +amdgpu.target @ggml_hadamard_f32_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.hadamard_f32.log2_block : %value: index where [range(%value, 6, 12)] + +kernel.def target(@ggml_hadamard_f32_gfx11_wave64) export("ggml_hadamard_f32") @ggml_hadamard_f32(%row_count: index) { + %one = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%row_count, %one, %one) workgroup_size(%c256, %one, %one) : index +} launch(%row_count: index, %input: buffer, %output: buffer) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c256 = index.constant 256 : index + %c4096 = index.constant 4096 : index + %c0_offset = index.constant 0 : offset + %lds_bytes = index.constant 16384 : offset + %one_i32 = scalar.constant 1 : i32 + %one_f32 = scalar.constant 1.0 : f32 + %inv_sqrt2 = scalar.constant 0.70710678118654752 : f32 + + %rows = index.assume %row_count [range(%row_count, 1, 1048576)] : index + %log2 = config.get @ggml.hadamard_f32.log2_block : index + %log2_i32 = index.cast %log2 : index to i32 + %n_i32 = scalar.shli %one_i32, %log2_i32 : i32 + %n0 = index.cast %n_i32 : i32 to index + %n = index.assume %n0 [range(%n0, 64, 4096)] : index + %half = index.div %n, %c2 : index + %count0 = index.mul %rows, %n : index + %count = index.assume %count0 [range(%count0, 64, 4294967296)] : index + + // scale = 2^(-floor(k/2)) * (k odd ? 1/sqrt(2) : 1) + %k_half_i32 = scalar.shrui %log2_i32, %one_i32 : i32 + %pow_i32 = scalar.shli %one_i32, %k_half_i32 : i32 + %pow_f32 = scalar.sitofp %pow_i32 : i32 to f32 + %even_scale = scalar.divf %one_f32, %pow_f32 : f32 + %k_odd_i32 = scalar.andi %log2_i32, %one_i32 : i32 + %zero_i32 = scalar.constant 0 : i32 + %k_odd = scalar.cmpi ne, %k_odd_i32, %zero_i32 : i32 + %odd_scale = scalar.mulf %even_scale, %inv_sqrt2 : f32 + %scale = scf.select %k_odd, %odd_scale, %even_scale : f32 + + %row = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %row_base = index.mul %row, %n : index + %row_ok = index.cmp ult, %row, %rows : index + + %input_view = buffer.view %input[%c0_offset] : buffer -> view<[%count]xf32> + %output_view = buffer.view %output[%c0_offset] : buffer -> view<[%count]xf32> + %lds = buffer.alloca align(16) %lds_bytes : buffer + %values = buffer.view %lds[%c0_offset] : buffer -> view<4096xf32> + + scf.for %e = [%workitem to %n step %c256] { + %g = index.add %row_base, %e : index + %g_ok = index.cmp ult, %g, %count : index + %e_ok = index.cmp ult, %e, %c4096 : index + %ok0 = scalar.andi %g_ok, %e_ok : i1 + %ok = scalar.andi %ok0, %row_ok : i1 + scf.if %ok { + %v = view.load %input_view[%g] : view<[%count]xf32> -> f32 + view.store %v, %values[%e] : f32, view<4096xf32> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + + scf.for %s = [%c0 to %log2 step %c1] { + %s_i32 = index.cast %s : index to i32 + %s1_i32 = scalar.addi %s_i32, %one_i32 : i32 + %h_i32 = scalar.shli %one_i32, %s_i32 : i32 + %low_mask = scalar.subi %h_i32, %one_i32 : i32 + scf.for %p = [%workitem to %half step %c256] { + %p_i32 = index.cast %p : index to i32 + %high = scalar.shrui %p_i32, %s_i32 : i32 + %high_shifted = scalar.shli %high, %s1_i32 : i32 + %low = scalar.andi %p_i32, %low_mask : i32 + %i_i32 = scalar.ori %high_shifted, %low : i32 + %j_i32 = scalar.ori %i_i32, %h_i32 : i32 + %i = index.cast %i_i32 : i32 to index + %j = index.cast %j_i32 : i32 to index + %i_ok = index.cmp ult, %i, %c4096 : index + %j_ok = index.cmp ult, %j, %c4096 : index + %pair_ok = scalar.andi %i_ok, %j_ok : i1 + scf.if %pair_ok { + %a = view.load %values[%i] : view<4096xf32> -> f32 + %b = view.load %values[%j] : view<4096xf32> -> f32 + %sum = scalar.addf %a, %b : f32 + %diff = scalar.subf %a, %b : f32 + view.store %sum, %values[%i] : f32, view<4096xf32> + view.store %diff, %values[%j] : f32, view<4096xf32> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + } + + scf.for %e2 = [%workitem to %n step %c256] { + %g2 = index.add %row_base, %e2 : index + %g2_ok = index.cmp ult, %g2, %count : index + %e2_ok = index.cmp ult, %e2, %c4096 : index + %ok2a = scalar.andi %g2_ok, %e2_ok : i1 + %ok2 = scalar.andi %ok2a, %row_ok : i1 + scf.if %ok2 { + %v2 = view.load %values[%e2] : view<4096xf32> -> f32 + %scaled = scalar.mulf %v2, %scale : f32 + view.store %scaled, %output_view[%g2] : f32, view<[%count]xf32> + } + } + kernel.return +} + +// Reference for the cases: one thread per output element, out[j] = scale * sum_i (-1)^popcount(i & j) +// x[i], the dense Sylvester product the transform replaces (no butterflies, no workgroup memory). +kernel.def target(@ggml_hadamard_f32_gfx11_wave64) export("ggml_hadamard_reference_f32") @ggml_hadamard_reference_f32(%row_count: index) { + %one = index.constant 1 : index + %c64 = index.constant 64 : index + %log2 = config.get @ggml.hadamard_f32.log2_block : index + %log2_i32 = index.cast %log2 : index to i32 + %one_i32 = scalar.constant 1 : i32 + %n_i32 = scalar.shli %one_i32, %log2_i32 : i32 + %n = index.cast %n_i32 : i32 to index + %groups = index.div %n, %c64 : index + kernel.launch.config workgroups(%groups, %row_count, %one) workgroup_size(%c64, %one, %one) : index +} launch(%row_count: index, %input: buffer, %output: buffer) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %c0_offset = index.constant 0 : offset + %one_i32 = scalar.constant 1 : i32 + %zero_i32 = scalar.constant 0 : i32 + %zero_f32 = scalar.constant 0.0 : f32 + %one_f32 = scalar.constant 1.0 : f32 + %two_f32 = scalar.constant 2.0 : f32 + %inv_sqrt2 = scalar.constant 0.70710678118654752 : f32 + %rows = index.assume %row_count [range(%row_count, 1, 64)] : index + %log2 = config.get @ggml.hadamard_f32.log2_block : index + %log2_i32 = index.cast %log2 : index to i32 + %n_i32 = scalar.shli %one_i32, %log2_i32 : i32 + %n0 = index.cast %n_i32 : i32 to index + %n = index.assume %n0 [range(%n0, 64, 4096)] : index + %count0 = index.mul %rows, %n : index + %count = index.assume %count0 [range(%count0, 64, 262144)] : index + %k_half_i32 = scalar.shrui %log2_i32, %one_i32 : i32 + %pow_i32 = scalar.shli %one_i32, %k_half_i32 : i32 + %pow_f32 = scalar.sitofp %pow_i32 : i32 to f32 + %even_scale = scalar.divf %one_f32, %pow_f32 : f32 + %k_odd_i32 = scalar.andi %log2_i32, %one_i32 : i32 + %k_odd = scalar.cmpi ne, %k_odd_i32, %zero_i32 : i32 + %odd_scale = scalar.mulf %even_scale, %inv_sqrt2 : f32 + %scale = scf.select %k_odd, %odd_scale, %even_scale : f32 + %group = kernel.workgroup.id : index + %row = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %group_base = index.mul %group, %c64 : index + %j = index.add %group_base, %lane : index + %j_i32 = index.cast %j : index to i32 + %row_base = index.mul %row, %n : index + %input_view = buffer.view %input[%c0_offset] : buffer -> view<[%count]xf32> + %output_view = buffer.view %output[%c0_offset] : buffer -> view<[%count]xf32> + %sum = scf.for %i = [%c0 to %n step %c1](%acc = %zero_f32 : f32) -> (f32) { + %i_i32 = index.cast %i : index to i32 + %both = scalar.andi %i_i32, %j_i32 : i32 + %parity = scf.for %b = [%c0 to %log2 step %c1](%p = %zero_i32 : i32) -> (i32) { + %b_i32 = index.cast %b : index to i32 + %shifted = scalar.shrui %both, %b_i32 : i32 + %bit = scalar.andi %shifted, %one_i32 : i32 + %next = scalar.xori %p, %bit : i32 + scf.yield %next : i32 + } + %parity_f32 = scalar.sitofp %parity : i32 to f32 + %twice = scalar.mulf %parity_f32, %two_f32 : f32 + %sign = scalar.subf %one_f32, %twice : f32 + %g = index.add %row_base, %i : index + %g_ok = index.cmp ult, %g, %count : index + %x = scf.if %g_ok -> (f32) { + %v = view.load %input_view[%g] : view<[%count]xf32> -> f32 + scf.yield %v : f32 + } else { + scf.yield %zero_f32 : f32 + } + %term = scalar.mulf %sign, %x : f32 + %next_acc = scalar.addf %acc, %term : f32 + scf.yield %next_acc : f32 + } + %result = scalar.mulf %sum, %scale : f32 + %out_index = index.add %row_base, %j : index + %out_ok0 = index.cmp ult, %out_index, %count : index + %j_ok = index.cmp ult, %j, %n : index + %out_ok = scalar.andi %out_ok0, %j_ok : i1 + scf.if %out_ok { + view.store %result, %output_view[%out_index] : f32, view<[%count]xf32> + } + kernel.return +} + +// Cases: every launch is checked against the reference on random inputs. A run takes one block +// size, so run each group with its --config (n64 -> 6, n1024 -> 10, n4096 -> 12). + +// n=64, 1 row(s): random input vs the reference. +// Run with --config=ggml.hadamard_f32.log2_block=6. +check.case public @ggml_hadamard_f32_n64_one_row_case { + %rows = check.literal value(1) : index + %seed = check.param.seed base(7400000000001000001) count(1) : i64 + %x = check.generate.random.uniform seed(%seed) range(-1.0 to 1.0) : tensor<64xf32> + %x_ref = check.generate.random.uniform seed(%seed) range(-1.0 to 1.0) : tensor<64xf32> + %expected = check.generate.fill value(7.0) : tensor<64xf32> + kernel.launch @ggml_hadamard_reference_f32[%rows](%rows, %x_ref, %expected) : [index](index, tensor<64xf32>, tensor<64xf32>) + %output = check.generate.fill value(-7.0) : tensor<64xf32> + kernel.launch @ggml_hadamard_f32[%rows](%rows, %x, %output) : [index](index, tensor<64xf32>, tensor<64xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-4) rtol(1.0e-4) nan(same) : tensor<64xf32> + check.return +} + +// n=64, 3 row(s): random input vs the reference. +// Run with --config=ggml.hadamard_f32.log2_block=6. +check.case public @ggml_hadamard_f32_n64_rows_case { + %rows = check.literal value(3) : index + %seed = check.param.seed base(7400000000001000008) count(1) : i64 + %x = check.generate.random.uniform seed(%seed) range(-1.0 to 1.0) : tensor<192xf32> + %x_ref = check.generate.random.uniform seed(%seed) range(-1.0 to 1.0) : tensor<192xf32> + %expected = check.generate.fill value(7.0) : tensor<192xf32> + kernel.launch @ggml_hadamard_reference_f32[%rows](%rows, %x_ref, %expected) : [index](index, tensor<192xf32>, tensor<192xf32>) + %output = check.generate.fill value(-7.0) : tensor<192xf32> + kernel.launch @ggml_hadamard_f32[%rows](%rows, %x, %output) : [index](index, tensor<192xf32>, tensor<192xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-4) rtol(1.0e-4) nan(same) : tensor<192xf32> + check.return +} + +// n=64, 2 row(s), in place (input and output are one buffer): random input vs the reference. +// Run with --config=ggml.hadamard_f32.log2_block=6. +check.case public @ggml_hadamard_f32_n64_in_place_case { + %rows = check.literal value(2) : index + %seed = check.param.seed base(7400000000001000015) count(1) : i64 + %x = check.generate.random.uniform seed(%seed) range(-1.0 to 1.0) : tensor<128xf32> + %x_ref = check.generate.random.uniform seed(%seed) range(-1.0 to 1.0) : tensor<128xf32> + %expected = check.generate.fill value(7.0) : tensor<128xf32> + kernel.launch @ggml_hadamard_reference_f32[%rows](%rows, %x_ref, %expected) : [index](index, tensor<128xf32>, tensor<128xf32>) + kernel.launch @ggml_hadamard_f32[%rows](%rows, %x, %x) : [index](index, tensor<128xf32>, tensor<128xf32>) + check.expect.close actual(%x) expected(%expected) atol(1.0e-4) rtol(1.0e-4) nan(same) : tensor<128xf32> + check.return +} + +// n=1024, 1 row(s): random input vs the reference. +// Run with --config=ggml.hadamard_f32.log2_block=10. +check.case public @ggml_hadamard_f32_n1024_one_row_case { + %rows = check.literal value(1) : index + %seed = check.param.seed base(7400000000001000022) count(1) : i64 + %x = check.generate.random.uniform seed(%seed) range(-1.0 to 1.0) : tensor<1024xf32> + %x_ref = check.generate.random.uniform seed(%seed) range(-1.0 to 1.0) : tensor<1024xf32> + %expected = check.generate.fill value(7.0) : tensor<1024xf32> + kernel.launch @ggml_hadamard_reference_f32[%rows](%rows, %x_ref, %expected) : [index](index, tensor<1024xf32>, tensor<1024xf32>) + %output = check.generate.fill value(-7.0) : tensor<1024xf32> + kernel.launch @ggml_hadamard_f32[%rows](%rows, %x, %output) : [index](index, tensor<1024xf32>, tensor<1024xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-4) rtol(1.0e-4) nan(same) : tensor<1024xf32> + check.return +} + +// n=1024, 3 row(s): random input vs the reference. +// Run with --config=ggml.hadamard_f32.log2_block=10. +check.case public @ggml_hadamard_f32_n1024_rows_case { + %rows = check.literal value(3) : index + %seed = check.param.seed base(7400000000001000029) count(1) : i64 + %x = check.generate.random.uniform seed(%seed) range(-1.0 to 1.0) : tensor<3072xf32> + %x_ref = check.generate.random.uniform seed(%seed) range(-1.0 to 1.0) : tensor<3072xf32> + %expected = check.generate.fill value(7.0) : tensor<3072xf32> + kernel.launch @ggml_hadamard_reference_f32[%rows](%rows, %x_ref, %expected) : [index](index, tensor<3072xf32>, tensor<3072xf32>) + %output = check.generate.fill value(-7.0) : tensor<3072xf32> + kernel.launch @ggml_hadamard_f32[%rows](%rows, %x, %output) : [index](index, tensor<3072xf32>, tensor<3072xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-4) rtol(1.0e-4) nan(same) : tensor<3072xf32> + check.return +} + +// n=1024, 2 row(s), in place (input and output are one buffer): random input vs the reference. +// Run with --config=ggml.hadamard_f32.log2_block=10. +check.case public @ggml_hadamard_f32_n1024_in_place_case { + %rows = check.literal value(2) : index + %seed = check.param.seed base(7400000000001000036) count(1) : i64 + %x = check.generate.random.uniform seed(%seed) range(-1.0 to 1.0) : tensor<2048xf32> + %x_ref = check.generate.random.uniform seed(%seed) range(-1.0 to 1.0) : tensor<2048xf32> + %expected = check.generate.fill value(7.0) : tensor<2048xf32> + kernel.launch @ggml_hadamard_reference_f32[%rows](%rows, %x_ref, %expected) : [index](index, tensor<2048xf32>, tensor<2048xf32>) + kernel.launch @ggml_hadamard_f32[%rows](%rows, %x, %x) : [index](index, tensor<2048xf32>, tensor<2048xf32>) + check.expect.close actual(%x) expected(%expected) atol(1.0e-4) rtol(1.0e-4) nan(same) : tensor<2048xf32> + check.return +} + +// n=4096, 2 row(s): random input vs the reference. +// Run with --config=ggml.hadamard_f32.log2_block=12. +check.case public @ggml_hadamard_f32_n4096_rows_case { + %rows = check.literal value(2) : index + %seed = check.param.seed base(7400000000001000043) count(1) : i64 + %x = check.generate.random.uniform seed(%seed) range(-1.0 to 1.0) : tensor<8192xf32> + %x_ref = check.generate.random.uniform seed(%seed) range(-1.0 to 1.0) : tensor<8192xf32> + %expected = check.generate.fill value(7.0) : tensor<8192xf32> + kernel.launch @ggml_hadamard_reference_f32[%rows](%rows, %x_ref, %expected) : [index](index, tensor<8192xf32>, tensor<8192xf32>) + %output = check.generate.fill value(-7.0) : tensor<8192xf32> + kernel.launch @ggml_hadamard_f32[%rows](%rows, %x, %output) : [index](index, tensor<8192xf32>, tensor<8192xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-4) rtol(1.0e-4) nan(same) : tensor<8192xf32> + check.return +} + +// n=4096, 1 row(s), in place (input and output are one buffer): random input vs the reference. +// Run with --config=ggml.hadamard_f32.log2_block=12. +check.case public @ggml_hadamard_f32_n4096_in_place_case { + %rows = check.literal value(1) : index + %seed = check.param.seed base(7400000000001000050) count(1) : i64 + %x = check.generate.random.uniform seed(%seed) range(-1.0 to 1.0) : tensor<4096xf32> + %x_ref = check.generate.random.uniform seed(%seed) range(-1.0 to 1.0) : tensor<4096xf32> + %expected = check.generate.fill value(7.0) : tensor<4096xf32> + kernel.launch @ggml_hadamard_reference_f32[%rows](%rows, %x_ref, %expected) : [index](index, tensor<4096xf32>, tensor<4096xf32>) + kernel.launch @ggml_hadamard_f32[%rows](%rows, %x, %x) : [index](index, tensor<4096xf32>, tensor<4096xf32>) + check.expect.close actual(%x) expected(%expected) atol(1.0e-4) rtol(1.0e-4) nan(same) : tensor<4096xf32> + check.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/kquant_decode_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/kquant_decode_f32.loom new file mode 100644 index 000000000000..825495aef6ee --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/kquant_decode_f32.loom @@ -0,0 +1,11806 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// Decode (one token) projections reading Q2_K, Q3_K, Q4_K, Q5_K, Q6_K, IQ3_S, IQ4_NL, IQ4_XS and Q8_0 +// weights in their GGUF block layout (no repack, no activation quantization), in two kernels: +// ggml_kquant_swiglu_decode_f32: out[r] = silu(sum_k gate[r, k] * x[k]) * sum_k up[r, k] * x[k] +// ggml_kquant_mul_mat_decode_f32: out[r] = sum_k w[r, k] * x[k] (+ addend[r]) +// One wave64 per output row. A 256-value block is split over 16 lanes, 16 values each, so a wave +// covers four blocks per step and each lane reads its weights with 8-byte loads. +// Q4_K / Q5_K: lane l owns bytes 8(l%4)..+8 of the 32-byte code group l/4, i.e. 8 values of +// sub-block 2(l/4) (low nibbles) and the same 8 positions of sub-block 2(l/4)+1 (high +// nibbles); Q5_K adds bit 2(l/4) / 2(l/4)+1 of the matching qh bytes as the fifth bit. +// IQ4_XS: lane l owns bytes 8(l%2)..+8 of sub-block l/2: 8 values from the low nibbles and the +// 8 values 16 positions later from the high nibbles. The IQ4_NL table value of each code is +// read from lane (code) of the 16-lane group, which holds kvalues_iq4nl[code]; a 16-entry +// vector.table.lookup lowers to 15 selects per byte on gfx11 and made IQ4_XS compute-bound. +// Q2_K, Q3_K, Q6_K, IQ1_S, IQ1_M, IQ2_XXS, IQ2_XS, IQ2_S, IQ3_XXS, IQ3_S, IQ4_NL, Q8_0: see their lane +// functions (2-byte aligned blocks, 16-bit loads). IQ1_S, IQ1_M, IQ2_XXS, IQ2_XS, IQ2_S, IQ3_XXS and +// IQ3_S read their grids from workgroup memory (4 KiB: IQ1's 2048 entries, two per word), written +// once per workgroup; the SwiGLU kernels keep separate gate and up grids. +// Weight format config values follow dispatch-mul-mat-weight-format.h (Q2_K 12, Q3_K 11, Q4_K 4, Q5_K 5, +// Q6_K 6, IQ1_S 26, IQ1_M 27, IQ2_S 22, IQ2_XXS 24, IQ2_XS 25, IQ3_XXS 28, IQ3_S 21, IQ4_NL 20, IQ4_XS 23, +// Q8_0 80, packed ternary 90: see @ggml_kquant_t2_lane_parts; TQ1_0 34, TQ2_0 35, MXFP4 39: see +// @ggml_kquant_run16_lane_parts; PrismML PQ2_0 72 and PTQ1_0 73: see @ggml_kquant_pq2_lane_parts). + +amdgpu.target @ggml_kquant_decode_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.kquant_swiglu_decode.input_size : %value: index where [range(%value, 256, 65536), mul(%value, 256)] + +config.decl @ggml.kquant_swiglu_decode.output_size : %value: index where [range(%value, 1, 1048576)] + +config.decl @ggml.kquant_swiglu_decode.gate_weight_format : %value: index where [range(%value, 4, 90)] + +config.decl @ggml.kquant_swiglu_decode.up_weight_format : %value: index where [range(%value, 4, 90)] + +// get_scale_min_k4(sub, scales) from the three scale words of a Q4_K / Q5_K block. +func.def inline @ggml_kquant_scale_min(%s0: i32, %s1: i32, %s2: i32, %sub: index) -> (i32, i32) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c6_i32 = scalar.constant 6 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c63_i32 = scalar.constant 63 : i32 + %s = index.assume %sub [range(%sub, 0, 7)] : index + %low = index.cmp ult, %s, %c4 : index + %t = index.rem %s, %c4 : index + %t8 = index.mul %t, %c8 : index + %shift = index.cast %t8 : index to i32 + %shift4 = scalar.addi %shift, %c4_i32 : i32 + %shift6 = scalar.addi %shift, %c6_i32 : i32 + %sc_a0 = scalar.shrui %s0, %shift : i32 + %sc_a = scalar.andi %sc_a0, %c63_i32 : i32 + %m_a0 = scalar.shrui %s1, %shift : i32 + %m_a = scalar.andi %m_a0, %c63_i32 : i32 + %sc_b0 = scalar.shrui %s2, %shift : i32 + %sc_b1 = scalar.andi %sc_b0, %c15_i32 : i32 + %sc_bh0 = scalar.shrui %s0, %shift6 : i32 + %sc_bh1 = scalar.andi %sc_bh0, %c3_i32 : i32 + %sc_bh = scalar.shli %sc_bh1, %c4_i32 : i32 + %sc_b = scalar.ori %sc_b1, %sc_bh : i32 + %m_b0 = scalar.shrui %s2, %shift4 : i32 + %m_b1 = scalar.andi %m_b0, %c15_i32 : i32 + %m_bh0 = scalar.shrui %s1, %shift6 : i32 + %m_bh1 = scalar.andi %m_bh0, %c3_i32 : i32 + %m_bh = scalar.shli %m_bh1, %c4_i32 : i32 + %m_b = scalar.ori %m_b1, %m_bh : i32 + %sc = scf.select %low, %sc_a, %sc_b : i32 + %m = scf.select %low, %m_a, %m_b : i32 + func.return %sc, %m : i32, i32 +} + +// kvalues_iq4nl. +func.def inline @ggml_kquant_iq4nl_table() -> (vector<16xi8>) { + %v0 = scalar.constant -127 : i8 + %v1 = scalar.constant -104 : i8 + %v2 = scalar.constant -83 : i8 + %v3 = scalar.constant -65 : i8 + %v4 = scalar.constant -49 : i8 + %v5 = scalar.constant -35 : i8 + %v6 = scalar.constant -22 : i8 + %v7 = scalar.constant -10 : i8 + %v8 = scalar.constant 1 : i8 + %v9 = scalar.constant 13 : i8 + %v10 = scalar.constant 25 : i8 + %v11 = scalar.constant 38 : i8 + %v12 = scalar.constant 53 : i8 + %v13 = scalar.constant 69 : i8 + %v14 = scalar.constant 89 : i8 + %v15 = scalar.constant 113 : i8 + %table = vector.from_elements %v0, %v1, %v2, %v3, %v4, %v5, %v6, %v7, %v8, %v9, %v10, %v11, %v12, %v13, %v14, %v15 : vector<16xi8> + func.return %table : vector<16xi8> +} + +// 8 codes (two words of 4 bytes) as f32. +config.decl @ggml.kquant_mul_mat_decode.input_size : %value: index where [range(%value, 256, 65536), mul(%value, 256)] + +config.decl @ggml.kquant_mul_mat_decode.output_size : %value: index where [range(%value, 1, 1048576)] + +config.decl @ggml.kquant_mul_mat_decode.weight_format : %value: index where [range(%value, 4, 90)] + +config.decl @ggml.kquant_mul_mat_decode.add : %value: index where [range(%value, 0, 1)] + +func.def inline @ggml_kquant_codes_u8x8(%codes: vector<2xi32>) -> (vector<8xf32>) { + %bytes = vector.bitcast %codes : vector<2xi32> to vector<8xi8> + %values = vector.uitofp %bytes : vector<8xi8> to vector<8xf32> + func.return %values : vector<8xf32> +} + + +// Lane partial of one Q4_K / Q5_K block from its header, scale words and the lane's 5-bit (or +// 4-bit) codes: d * sc * sum(q * x) - dmin * m * sum(x) over the lane's two 8-value runs. +func.def inline @ggml_kquant_q45k_lane_finish(%dm: vector<2xf16>, %header: vector<4xi32>, %low: vector<2xi32>, %high: vector<2xi32>, %group: index, %packet: index, %input: buffer, %block: index) -> (f32) { + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c8 = index.constant 8 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %zero_scalar = scalar.constant 0.0 : f32 + %x_block_bytes = index.constant 1024 : offset + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %d_f16 = vector.extract %dm[0] : vector<2xf16> -> f16 + %dmin_f16 = vector.extract %dm[1] : vector<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %s0 = vector.extract %header[1] : vector<4xi32> -> i32 + %s1 = vector.extract %header[2] : vector<4xi32> -> i32 + %s2 = vector.extract %header[3] : vector<4xi32> -> i32 + %sub_lo = index.mul %group, %c2 : index + %sub_hi = index.add %sub_lo, %c1 : index + %sc_lo_i, %m_lo_i = func.call @ggml_kquant_scale_min(%s0, %s1, %s2, %sub_lo) : (i32, i32, i32, index) -> (i32, i32) + %sc_hi_i, %m_hi_i = func.call @ggml_kquant_scale_min(%s0, %s1, %s2, %sub_hi) : (i32, i32, i32, index) -> (i32, i32) + %q_lo = func.call @ggml_kquant_codes_u8x8(%low) : (vector<2xi32>) -> (vector<8xf32>) + %q_hi = func.call @ggml_kquant_codes_u8x8(%high) : (vector<2xi32>) -> (vector<8xf32>) + %sub_base = index.mul %group, %c64 : index + %pos = index.mul %packet, %c8 : index + %x_lo_k = index.add %sub_base, %pos : index + %x_hi_k = index.add %x_lo_k, %c32 : index + %x_lo = vector.load %xv[%x_lo_k] : view<256xf32> -> vector<8xf32> + %x_hi = vector.load %xv[%x_hi_k] : view<256xf32> -> vector<8xf32> + %qx_lo_v = vector.mulf %q_lo, %x_lo : vector<8xf32> + %qx_hi_v = vector.mulf %q_hi, %x_hi : vector<8xf32> + %qx_lo = vector.reduce %qx_lo_v, %zero_scalar : vector<8xf32>, f32 + %qx_hi = vector.reduce %qx_hi_v, %zero_scalar : vector<8xf32>, f32 + %sx_lo = vector.reduce %x_lo, %zero_scalar : vector<8xf32>, f32 + %sx_hi = vector.reduce %x_hi, %zero_scalar : vector<8xf32>, f32 + %sc_lo_f = scalar.uitofp %sc_lo_i : i32 to f32 + %sc_hi_f = scalar.uitofp %sc_hi_i : i32 to f32 + %m_lo_f = scalar.uitofp %m_lo_i : i32 to f32 + %m_hi_f = scalar.uitofp %m_hi_i : i32 to f32 + %d_lo = scalar.mulf %d, %sc_lo_f : f32 + %d_hi = scalar.mulf %d, %sc_hi_f : f32 + %mn_lo = scalar.mulf %dmin, %m_lo_f : f32 + %mn_hi = scalar.mulf %dmin, %m_hi_f : f32 + %t0 = scalar.mulf %d_lo, %qx_lo : f32 + %t1 = scalar.fmaf %d_hi, %qx_hi, %t0 : f32 + %n0 = scalar.mulf %mn_lo, %sx_lo : f32 + %n1 = scalar.fmaf %mn_hi, %sx_hi, %n0 : f32 + %result = scalar.subf %t1, %n1 : f32 + func.return %result : f32 +} + +// Q4_K: 144 bytes = d, dmin (f16), scales[12], qs[128]. +func.def inline @ggml_kquant_q4k_lane_dot(%weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %nibble_mask = vector.constant 252645135 : vector<2xi32> + %four = vector.constant 4 : vector<2xi32> + %block_bytes = index.constant 144 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %group = index.div %l, %c4 : index + %packet = index.rem %l, %c4 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %wv = buffer.view %weight[%block_base] : buffer -> view<36xi32> + %hv = buffer.view %weight[%block_base] : buffer -> view<2xf16> + %dm = vector.load %hv[%c0] : view<2xf16> -> vector<2xf16> + %header = vector.load %wv[%c0] : view<36xi32> -> vector<4xi32> + %group_words = index.mul %group, %c8 : index + %packet_words = index.mul %packet, %c2 : index + %code_rel = index.add %group_words, %packet_words : index + %code_word = index.add %c4, %code_rel : index + %codes = vector.load %wv[%code_word] : view<36xi32> -> vector<2xi32> + %low = vector.andi %codes, %nibble_mask : vector<2xi32> + %codes_shr = vector.shrui %codes, %four : vector<2xi32> + %high = vector.andi %codes_shr, %nibble_mask : vector<2xi32> + %r = func.call @ggml_kquant_q45k_lane_finish(%dm, %header, %low, %high, %group, %packet, %input, %block) : (vector<2xf16>, vector<4xi32>, vector<2xi32>, vector<2xi32>, index, index, buffer, index) -> (f32) + func.return %r : f32 +} + +// Q5_K: 176 bytes = d, dmin (f16), scales[12], qh[32], qs[128]. +func.def inline @ggml_kquant_q5k_lane_dot(%weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c12 = index.constant 12 : index + %nibble_mask = vector.constant 252645135 : vector<2xi32> + %bit_mask = vector.constant 16843009 : vector<2xi32> + %four = vector.constant 4 : vector<2xi32> + %block_bytes = index.constant 176 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %group = index.div %l, %c4 : index + %packet = index.rem %l, %c4 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %wv = buffer.view %weight[%block_base] : buffer -> view<44xi32> + %hv = buffer.view %weight[%block_base] : buffer -> view<2xf16> + %dm = vector.load %hv[%c0] : view<2xf16> -> vector<2xf16> + %header = vector.load %wv[%c0] : view<44xi32> -> vector<4xi32> + %group_words = index.mul %group, %c8 : index + %packet_words = index.mul %packet, %c2 : index + %code_rel = index.add %group_words, %packet_words : index + %code_word = index.add %c12, %code_rel : index + %codes = vector.load %wv[%code_word] : view<44xi32> -> vector<2xi32> + %qh_word = index.add %c4, %packet_words : index + %qh = vector.load %wv[%qh_word] : view<44xi32> -> vector<2xi32> + %sub_lo = index.mul %group, %c2 : index + %sub_hi = index.add %sub_lo, %c1 : index + %shift_lo_i = index.cast %sub_lo : index to i32 + %shift_hi_i = index.cast %sub_hi : index to i32 + %shift_lo = vector.splat %shift_lo_i : vector<2xi32> + %shift_hi = vector.splat %shift_hi_i : vector<2xi32> + %qh_lo0 = vector.shrui %qh, %shift_lo : vector<2xi32> + %qh_hi0 = vector.shrui %qh, %shift_hi : vector<2xi32> + %qh_lo1 = vector.andi %qh_lo0, %bit_mask : vector<2xi32> + %qh_hi1 = vector.andi %qh_hi0, %bit_mask : vector<2xi32> + %qh_lo = vector.shli %qh_lo1, %four : vector<2xi32> + %qh_hi = vector.shli %qh_hi1, %four : vector<2xi32> + %low4 = vector.andi %codes, %nibble_mask : vector<2xi32> + %codes_shr = vector.shrui %codes, %four : vector<2xi32> + %high4 = vector.andi %codes_shr, %nibble_mask : vector<2xi32> + %low = vector.ori %low4, %qh_lo : vector<2xi32> + %high = vector.ori %high4, %qh_hi : vector<2xi32> + %r = func.call @ggml_kquant_q45k_lane_finish(%dm, %header, %low, %high, %group, %packet, %input, %block) : (vector<2xf16>, vector<4xi32>, vector<2xi32>, vector<2xi32>, index, index, buffer, index) -> (f32) + func.return %r : f32 +} + +// sum over the lane's 16 IQ4 codes (two words: low nibbles at x_lo, high nibbles at x_hi) of +// kvalues_iq4nl[code] * x, the table value read from lane (code) of the 16-lane group. +func.def inline @ggml_kquant_iq4_codes_dot(%kv_lane: f32, %codes: vector<2xi32>, %x_lo: vector<8xf32>, %x_hi: vector<8xf32>) -> (f32) { + %zero_scalar = scalar.constant 0.0 : f32 + %c15_i32 = scalar.constant 15 : i32 + %c16w = scalar.constant 16 : i32 + %w0 = vector.extract %codes[0] : vector<2xi32> -> i32 + %w1 = vector.extract %codes[1] : vector<2xi32> -> i32 + %sh0 = scalar.constant 0 : i32 + %sh4 = scalar.constant 4 : i32 + %sh8 = scalar.constant 8 : i32 + %sh12 = scalar.constant 12 : i32 + %sh16 = scalar.constant 16 : i32 + %sh20 = scalar.constant 20 : i32 + %sh24 = scalar.constant 24 : i32 + %sh28 = scalar.constant 28 : i32 + %c_lo00 = scalar.shrui %w0, %sh0 : i32 + %c_lo0 = scalar.andi %c_lo00, %c15_i32 : i32 + %k_lo0, %k_lo0_ok = kernel.subgroup.shuffle %kv_lane, %c_lo0, %c16w : f32, i32, i32 + %x_lo0 = vector.extract %x_lo[0] : vector<8xf32> -> f32 + %acc0 = scalar.fmaf %k_lo0, %x_lo0, %zero_scalar : f32 + %c_lo10 = scalar.shrui %w0, %sh8 : i32 + %c_lo1 = scalar.andi %c_lo10, %c15_i32 : i32 + %k_lo1, %k_lo1_ok = kernel.subgroup.shuffle %kv_lane, %c_lo1, %c16w : f32, i32, i32 + %x_lo1 = vector.extract %x_lo[1] : vector<8xf32> -> f32 + %acc1 = scalar.fmaf %k_lo1, %x_lo1, %acc0 : f32 + %c_lo20 = scalar.shrui %w0, %sh16 : i32 + %c_lo2 = scalar.andi %c_lo20, %c15_i32 : i32 + %k_lo2, %k_lo2_ok = kernel.subgroup.shuffle %kv_lane, %c_lo2, %c16w : f32, i32, i32 + %x_lo2 = vector.extract %x_lo[2] : vector<8xf32> -> f32 + %acc2 = scalar.fmaf %k_lo2, %x_lo2, %acc1 : f32 + %c_lo30 = scalar.shrui %w0, %sh24 : i32 + %c_lo3 = scalar.andi %c_lo30, %c15_i32 : i32 + %k_lo3, %k_lo3_ok = kernel.subgroup.shuffle %kv_lane, %c_lo3, %c16w : f32, i32, i32 + %x_lo3 = vector.extract %x_lo[3] : vector<8xf32> -> f32 + %acc3 = scalar.fmaf %k_lo3, %x_lo3, %acc2 : f32 + %c_lo40 = scalar.shrui %w1, %sh0 : i32 + %c_lo4 = scalar.andi %c_lo40, %c15_i32 : i32 + %k_lo4, %k_lo4_ok = kernel.subgroup.shuffle %kv_lane, %c_lo4, %c16w : f32, i32, i32 + %x_lo4 = vector.extract %x_lo[4] : vector<8xf32> -> f32 + %acc4 = scalar.fmaf %k_lo4, %x_lo4, %acc3 : f32 + %c_lo50 = scalar.shrui %w1, %sh8 : i32 + %c_lo5 = scalar.andi %c_lo50, %c15_i32 : i32 + %k_lo5, %k_lo5_ok = kernel.subgroup.shuffle %kv_lane, %c_lo5, %c16w : f32, i32, i32 + %x_lo5 = vector.extract %x_lo[5] : vector<8xf32> -> f32 + %acc5 = scalar.fmaf %k_lo5, %x_lo5, %acc4 : f32 + %c_lo60 = scalar.shrui %w1, %sh16 : i32 + %c_lo6 = scalar.andi %c_lo60, %c15_i32 : i32 + %k_lo6, %k_lo6_ok = kernel.subgroup.shuffle %kv_lane, %c_lo6, %c16w : f32, i32, i32 + %x_lo6 = vector.extract %x_lo[6] : vector<8xf32> -> f32 + %acc6 = scalar.fmaf %k_lo6, %x_lo6, %acc5 : f32 + %c_lo70 = scalar.shrui %w1, %sh24 : i32 + %c_lo7 = scalar.andi %c_lo70, %c15_i32 : i32 + %k_lo7, %k_lo7_ok = kernel.subgroup.shuffle %kv_lane, %c_lo7, %c16w : f32, i32, i32 + %x_lo7 = vector.extract %x_lo[7] : vector<8xf32> -> f32 + %acc7 = scalar.fmaf %k_lo7, %x_lo7, %acc6 : f32 + %c_hi00 = scalar.shrui %w0, %sh4 : i32 + %c_hi0 = scalar.andi %c_hi00, %c15_i32 : i32 + %k_hi0, %k_hi0_ok = kernel.subgroup.shuffle %kv_lane, %c_hi0, %c16w : f32, i32, i32 + %x_hi0 = vector.extract %x_hi[0] : vector<8xf32> -> f32 + %acc8 = scalar.fmaf %k_hi0, %x_hi0, %acc7 : f32 + %c_hi10 = scalar.shrui %w0, %sh12 : i32 + %c_hi1 = scalar.andi %c_hi10, %c15_i32 : i32 + %k_hi1, %k_hi1_ok = kernel.subgroup.shuffle %kv_lane, %c_hi1, %c16w : f32, i32, i32 + %x_hi1 = vector.extract %x_hi[1] : vector<8xf32> -> f32 + %acc9 = scalar.fmaf %k_hi1, %x_hi1, %acc8 : f32 + %c_hi20 = scalar.shrui %w0, %sh20 : i32 + %c_hi2 = scalar.andi %c_hi20, %c15_i32 : i32 + %k_hi2, %k_hi2_ok = kernel.subgroup.shuffle %kv_lane, %c_hi2, %c16w : f32, i32, i32 + %x_hi2 = vector.extract %x_hi[2] : vector<8xf32> -> f32 + %acc10 = scalar.fmaf %k_hi2, %x_hi2, %acc9 : f32 + %c_hi30 = scalar.shrui %w0, %sh28 : i32 + %c_hi3 = scalar.andi %c_hi30, %c15_i32 : i32 + %k_hi3, %k_hi3_ok = kernel.subgroup.shuffle %kv_lane, %c_hi3, %c16w : f32, i32, i32 + %x_hi3 = vector.extract %x_hi[3] : vector<8xf32> -> f32 + %acc11 = scalar.fmaf %k_hi3, %x_hi3, %acc10 : f32 + %c_hi40 = scalar.shrui %w1, %sh4 : i32 + %c_hi4 = scalar.andi %c_hi40, %c15_i32 : i32 + %k_hi4, %k_hi4_ok = kernel.subgroup.shuffle %kv_lane, %c_hi4, %c16w : f32, i32, i32 + %x_hi4 = vector.extract %x_hi[4] : vector<8xf32> -> f32 + %acc12 = scalar.fmaf %k_hi4, %x_hi4, %acc11 : f32 + %c_hi50 = scalar.shrui %w1, %sh12 : i32 + %c_hi5 = scalar.andi %c_hi50, %c15_i32 : i32 + %k_hi5, %k_hi5_ok = kernel.subgroup.shuffle %kv_lane, %c_hi5, %c16w : f32, i32, i32 + %x_hi5 = vector.extract %x_hi[5] : vector<8xf32> -> f32 + %acc13 = scalar.fmaf %k_hi5, %x_hi5, %acc12 : f32 + %c_hi60 = scalar.shrui %w1, %sh20 : i32 + %c_hi6 = scalar.andi %c_hi60, %c15_i32 : i32 + %k_hi6, %k_hi6_ok = kernel.subgroup.shuffle %kv_lane, %c_hi6, %c16w : f32, i32, i32 + %x_hi6 = vector.extract %x_hi[6] : vector<8xf32> -> f32 + %acc14 = scalar.fmaf %k_hi6, %x_hi6, %acc13 : f32 + %c_hi70 = scalar.shrui %w1, %sh28 : i32 + %c_hi7 = scalar.andi %c_hi70, %c15_i32 : i32 + %k_hi7, %k_hi7_ok = kernel.subgroup.shuffle %kv_lane, %c_hi7, %c16w : f32, i32, i32 + %x_hi7 = vector.extract %x_hi[7] : vector<8xf32> -> f32 + %acc15 = scalar.fmaf %k_hi7, %x_hi7, %acc14 : f32 + func.return %acc15 : f32 +} + +// IQ4_XS: 136 bytes = d (f16), scales_h (u16), scales_l[4], qs[128]. +func.def inline @ggml_kquant_iq4xs_lane_dot(%kv_lane: f32, %weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c16_i32 = scalar.constant 16 : i32 + %c32_i32 = scalar.constant 32 : i32 + %block_bytes = index.constant 136 : offset + %x_block_bytes = index.constant 1024 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %sub = index.div %l, %c2 : index + %half = index.rem %l, %c2 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %wv = buffer.view %weight[%block_base] : buffer -> view<34xi32> + %hv = buffer.view %weight[%block_base] : buffer -> view<2xf16> + %dh = vector.load %hv[%c0] : view<2xf16> -> vector<2xf16> + %d_f16 = vector.extract %dh[0] : vector<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %header = vector.load %wv[%c0] : view<34xi32> -> vector<2xi32> + %word0 = vector.extract %header[0] : vector<2xi32> -> i32 + %scales_l = vector.extract %header[1] : vector<2xi32> -> i32 + %sub_i32 = index.cast %sub : index to i32 + %low_shift = scalar.muli %sub_i32, %c4_i32 : i32 + %high_shift0 = scalar.muli %sub_i32, %c2_i32 : i32 + %high_shift = scalar.addi %high_shift0, %c16_i32 : i32 + %ls_low0 = scalar.shrui %scales_l, %low_shift : i32 + %ls_low = scalar.andi %ls_low0, %c15_i32 : i32 + %ls_high0 = scalar.shrui %word0, %high_shift : i32 + %ls_high1 = scalar.andi %ls_high0, %c3_i32 : i32 + %ls_high = scalar.shli %ls_high1, %c4_i32 : i32 + %ls = scalar.ori %ls_low, %ls_high : i32 + %ls_centered = scalar.subi %ls, %c32_i32 : i32 + %ls_f = scalar.sitofp %ls_centered : i32 to f32 + %dl = scalar.mulf %d, %ls_f : f32 + %sub_words = index.mul %sub, %c4 : index + %half_words = index.mul %half, %c2 : index + %code_rel = index.add %sub_words, %half_words : index + %code_word = index.add %c2, %code_rel : index + %codes = vector.load %wv[%code_word] : view<34xi32> -> vector<2xi32> + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %sub_k = index.mul %sub, %c32 : index + %half_k = index.mul %half, %c8 : index + %x_lo_k = index.add %sub_k, %half_k : index + %x_hi_k = index.add %x_lo_k, %c16 : index + %x_lo = vector.load %xv[%x_lo_k] : view<256xf32> -> vector<8xf32> + %x_hi = vector.load %xv[%x_hi_k] : view<256xf32> -> vector<8xf32> + %sum = func.call @ggml_kquant_iq4_codes_dot(%kv_lane, %codes, %x_lo, %x_hi) : (f32, vector<2xi32>, vector<8xf32>, vector<8xf32>) -> (f32) + %result = scalar.mulf %dl, %sum : f32 + func.return %result : f32 +} + +// IQ4_NL: 18 bytes per 32 values = d (f16), qs[16]; eight blocks per 256 values. Lane l owns +// bytes 8(l%2)..+8 of block l/2 (only 2-byte aligned, so 16-bit loads). +func.def inline @ggml_kquant_iq4nl_lane_dot(%kv_lane: f32, %weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %super_bytes = index.constant 144 : offset + %qblock_bytes = index.constant 18 : offset + %x_block_bytes = index.constant 1024 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %qb = index.div %l, %c2 : index + %half = index.rem %l, %c2 : index + %super_add = index.scale %block, %super_bytes : index, offset -> offset + %qb_add = index.scale %qb, %qblock_bytes : index, offset -> offset + %super_base = index.add %row_base, %super_add : offset + %qb_base = index.add %super_base, %qb_add : offset + %hv = buffer.view %weight[%qb_base] : buffer -> view<9xf16> + %iv = buffer.view %weight[%qb_base] : buffer -> view<9xi16> + %d_f16 = view.load %hv[%c0] : view<9xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %half_words = index.mul %half, %c4 : index + %code_at = index.add %c1, %half_words : index + %codes16 = vector.load %iv[%code_at] : view<9xi16> -> vector<4xi16> + %codes = vector.bitcast %codes16 : vector<4xi16> to vector<2xi32> + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %qb_k = index.mul %qb, %c32 : index + %half_k = index.mul %half, %c8 : index + %x_lo_k = index.add %qb_k, %half_k : index + %x_hi_k = index.add %x_lo_k, %c16 : index + %x_lo = vector.load %xv[%x_lo_k] : view<256xf32> -> vector<8xf32> + %x_hi = vector.load %xv[%x_hi_k] : view<256xf32> -> vector<8xf32> + %sum = func.call @ggml_kquant_iq4_codes_dot(%kv_lane, %codes, %x_lo, %x_hi) : (f32, vector<2xi32>, vector<8xf32>, vector<8xf32>) -> (f32) + %result = scalar.mulf %d, %sum : f32 + func.return %result : f32 +} + +// Q8_0: 34 bytes per 32 values = d (f16), qs[32] (int8); eight blocks per 256 values. Lane l owns +// bytes 16(l%2)..+16 of block l/2 (2-byte aligned). +func.def inline @ggml_kquant_q8_0_lane_dot(%weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %zero_scalar = scalar.constant 0.0 : f32 + %super_bytes = index.constant 272 : offset + %qblock_bytes = index.constant 34 : offset + %x_block_bytes = index.constant 1024 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %qb = index.div %l, %c2 : index + %half = index.rem %l, %c2 : index + %super_add = index.scale %block, %super_bytes : index, offset -> offset + %qb_add = index.scale %qb, %qblock_bytes : index, offset -> offset + %super_base = index.add %row_base, %super_add : offset + %qb_base = index.add %super_base, %qb_add : offset + %hv = buffer.view %weight[%qb_base] : buffer -> view<17xf16> + %iv = buffer.view %weight[%qb_base] : buffer -> view<17xi16> + %d_f16 = view.load %hv[%c0] : view<17xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %half_words = index.mul %half, %c8 : index + %code_at = index.add %c1, %half_words : index + %codes16 = vector.load %iv[%code_at] : view<17xi16> -> vector<8xi16> + %codes = vector.bitcast %codes16 : vector<8xi16> to vector<16xi8> + %q = vector.sitofp %codes : vector<16xi8> to vector<16xf32> + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %qb_k = index.mul %qb, %c32 : index + %half_k = index.mul %half, %c16 : index + %x_k = index.add %qb_k, %half_k : index + %x = vector.load %xv[%x_k] : view<256xf32> -> vector<16xf32> + %qx = vector.mulf %q, %x : vector<16xf32> + %sum = vector.reduce %qx, %zero_scalar : vector<16xf32>, f32 + %result = scalar.mulf %d, %sum : f32 + func.return %result : f32 +} + +// Q6_K: 210 bytes = ql[128], qh[64], scales[16] (int8), d (f16); 2-byte aligned. Lane l takes half +// n = l/8 of the block and positions p = 4(l%8)..+4 of each of its four 32-value quarters: +// ql[64n + p] (low nibble: quarter 0, high: quarter 2), ql[64n + 32 + p] (quarters 1 and 3), +// qh[32n + p] (two bits per quarter), scale sc[8n + (l%8)/4 + 2 quarter]. +func.def inline @ggml_kquant_q6k_quarter(%ql: i32, %qh: i32, %nibble_shift: i32, %high_shift: i32, %x: vector<4xf32>) -> (f32) { + %c4_i32 = scalar.constant 4 : i32 + %mask4 = scalar.constant 252645135 : i32 + %mask2 = scalar.constant 50529027 : i32 + %zero_scalar = scalar.constant 0.0 : f32 + %c32 = vector.constant 32.0 : vector<4xf32> + %low0 = scalar.shrui %ql, %nibble_shift : i32 + %low = scalar.andi %low0, %mask4 : i32 + %high0 = scalar.shrui %qh, %high_shift : i32 + %high1 = scalar.andi %high0, %mask2 : i32 + %high = scalar.shli %high1, %c4_i32 : i32 + %q_word = scalar.ori %low, %high : i32 + %q_vec = vector.from_elements %q_word : vector<1xi32> + %q_bytes = vector.bitcast %q_vec : vector<1xi32> to vector<4xi8> + %q_u = vector.uitofp %q_bytes : vector<4xi8> to vector<4xf32> + %q = vector.subf %q_u, %c32 : vector<4xf32> + %qx = vector.mulf %q, %x : vector<4xf32> + %sum = vector.reduce %qx, %zero_scalar : vector<4xf32>, f32 + func.return %sum : f32 +} + +func.def inline @ggml_kquant_q6k_lane_dot(%weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c96 = index.constant 96 : index + %c104 = index.constant 104 : index + %c128 = index.constant 128 : index + %s0 = scalar.constant 0 : i32 + %s2 = scalar.constant 2 : i32 + %s4 = scalar.constant 4 : i32 + %s6 = scalar.constant 6 : i32 + %block_bytes = index.constant 210 : offset + %x_block_bytes = index.constant 1024 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %n = index.div %l, %c8 : index + %j = index.rem %l, %c8 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %iv = buffer.view %weight[%block_base] : buffer -> view<105xi16> + %hv = buffer.view %weight[%block_base] : buffer -> view<105xf16> + %n32 = index.mul %n, %c32 : index + %j2 = index.mul %j, %c2 : index + %ql_a_at = index.add %n32, %j2 : index + %ql_b_at = index.add %ql_a_at, %c16 : index + %n16 = index.mul %n, %c16 : index + %qh_rel = index.add %n16, %j2 : index + %qh_at = index.add %c64, %qh_rel : index + %n4 = index.mul %n, %c4 : index + %sc_at = index.add %c96, %n4 : index + %ql_a2 = vector.load %iv[%ql_a_at] : view<105xi16> -> vector<2xi16> + %ql_b2 = vector.load %iv[%ql_b_at] : view<105xi16> -> vector<2xi16> + %qh2 = vector.load %iv[%qh_at] : view<105xi16> -> vector<2xi16> + %sc4 = vector.load %iv[%sc_at] : view<105xi16> -> vector<4xi16> + %ql_a1 = vector.bitcast %ql_a2 : vector<2xi16> to vector<1xi32> + %ql_b1 = vector.bitcast %ql_b2 : vector<2xi16> to vector<1xi32> + %qh1 = vector.bitcast %qh2 : vector<2xi16> to vector<1xi32> + %ql_a = vector.extract %ql_a1[0] : vector<1xi32> -> i32 + %ql_b = vector.extract %ql_b1[0] : vector<1xi32> -> i32 + %qh = vector.extract %qh1[0] : vector<1xi32> -> i32 + %sc8 = vector.bitcast %sc4 : vector<4xi16> to vector<8xi8> + %d_f16 = view.load %hv[%c104] : view<105xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %second = index.cmp uge, %j, %c4 : index + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %n128 = index.mul %n, %c128 : index + %j4 = index.mul %j, %c4 : index + %p0 = index.add %n128, %j4 : index + %p1 = index.add %p0, %c32 : index + %p2 = index.add %p0, %c64 : index + %p3 = index.add %p0, %c96 : index + %x0 = vector.load %xv[%p0] : view<256xf32> -> vector<4xf32> + %x1 = vector.load %xv[%p1] : view<256xf32> -> vector<4xf32> + %x2 = vector.load %xv[%p2] : view<256xf32> -> vector<4xf32> + %x3 = vector.load %xv[%p3] : view<256xf32> -> vector<4xf32> + %t0 = func.call @ggml_kquant_q6k_quarter(%ql_a, %qh, %s0, %s0, %x0) : (i32, i32, i32, i32, vector<4xf32>) -> (f32) + %t1 = func.call @ggml_kquant_q6k_quarter(%ql_b, %qh, %s0, %s2, %x1) : (i32, i32, i32, i32, vector<4xf32>) -> (f32) + %t2 = func.call @ggml_kquant_q6k_quarter(%ql_a, %qh, %s4, %s4, %x2) : (i32, i32, i32, i32, vector<4xf32>) -> (f32) + %t3 = func.call @ggml_kquant_q6k_quarter(%ql_b, %qh, %s4, %s6, %x3) : (i32, i32, i32, i32, vector<4xf32>) -> (f32) + %sc0_a = vector.extract %sc8[0] : vector<8xi8> -> i8 + %sc0_b = vector.extract %sc8[1] : vector<8xi8> -> i8 + %sc0_i8 = scf.select %second, %sc0_b, %sc0_a : i8 + %sc0_i32 = scalar.extsi %sc0_i8 : i8 to i32 + %sc0 = scalar.sitofp %sc0_i32 : i32 to f32 + %sc1_a = vector.extract %sc8[2] : vector<8xi8> -> i8 + %sc1_b = vector.extract %sc8[3] : vector<8xi8> -> i8 + %sc1_i8 = scf.select %second, %sc1_b, %sc1_a : i8 + %sc1_i32 = scalar.extsi %sc1_i8 : i8 to i32 + %sc1 = scalar.sitofp %sc1_i32 : i32 to f32 + %sc2_a = vector.extract %sc8[4] : vector<8xi8> -> i8 + %sc2_b = vector.extract %sc8[5] : vector<8xi8> -> i8 + %sc2_i8 = scf.select %second, %sc2_b, %sc2_a : i8 + %sc2_i32 = scalar.extsi %sc2_i8 : i8 to i32 + %sc2 = scalar.sitofp %sc2_i32 : i32 to f32 + %sc3_a = vector.extract %sc8[6] : vector<8xi8> -> i8 + %sc3_b = vector.extract %sc8[7] : vector<8xi8> -> i8 + %sc3_i8 = scf.select %second, %sc3_b, %sc3_a : i8 + %sc3_i32 = scalar.extsi %sc3_i8 : i8 to i32 + %sc3 = scalar.sitofp %sc3_i32 : i32 to f32 + %u0 = scalar.mulf %sc0, %t0 : f32 + %u1 = scalar.fmaf %sc1, %t1, %u0 : f32 + %u2 = scalar.fmaf %sc2, %t2, %u1 : f32 + %u3 = scalar.fmaf %sc3, %t3, %u2 : f32 + %result = scalar.mulf %d, %u3 : f32 + func.return %result : f32 +} + +// Q3_K: 110 bytes = hmask[32], qs[64], scales[12] (16 6-bit scales), d (f16); 2-byte aligned. +// Value 128n + 32j + 16h + l (l < 16) is d * (scale[8n + 2j + h] - 32) * (q - 4 * !hbit) with +// q = (qs[32n + 16h + l] >> 2j) & 3 and hbit = bit 4n + j of hmask[16h + l], so lane l16 = 8n + +// 2j + h owns one run of 16 values with a single scale. +func.def inline @ggml_kquant_q3k_lane_dot(%weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c48 = index.constant 48 : index + %c54 = index.constant 54 : index + %c128 = index.constant 128 : index + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c32_i32 = scalar.constant 32 : i32 + %zero_scalar = scalar.constant 0.0 : f32 + %mask2 = vector.constant 50529027 : vector<4xi32> + %mask1 = vector.constant 16843009 : vector<4xi32> + %two = vector.constant 2 : vector<4xi32> + %four_f = vector.constant 4.0 : vector<16xf32> + %block_bytes = index.constant 110 : offset + %x_block_bytes = index.constant 1024 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %n = index.div %l, %c8 : index + %jh = index.rem %l, %c8 : index + %j = index.div %jh, %c2 : index + %h = index.rem %jh, %c2 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %iv = buffer.view %weight[%block_base] : buffer -> view<55xi16> + %hv = buffer.view %weight[%block_base] : buffer -> view<55xf16> + %h8 = index.mul %h, %c8 : index + %n16 = index.mul %n, %c16 : index + %qs_rel = index.add %n16, %h8 : index + %qs_at = index.add %c16, %qs_rel : index + %qs16 = vector.load %iv[%qs_at] : view<55xi16> -> vector<8xi16> + %hm16 = vector.load %iv[%h8] : view<55xi16> -> vector<8xi16> + %sc16 = vector.load %iv[%c48] : view<55xi16> -> vector<6xi16> + %d_f16 = view.load %hv[%c54] : view<55xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %qs = vector.bitcast %qs16 : vector<8xi16> to vector<4xi32> + %hm = vector.bitcast %hm16 : vector<8xi16> to vector<4xi32> + %scw = vector.bitcast %sc16 : vector<6xi16> to vector<3xi32> + %j_i32 = index.cast %j : index to i32 + %n_i32 = index.cast %n : index to i32 + %q_shift_i = scalar.muli %j_i32, %c2_i32 : i32 + %n4 = scalar.muli %n_i32, %c4_i32 : i32 + %h_shift_i = scalar.addi %n4, %j_i32 : i32 + %q_shift = vector.splat %q_shift_i : vector<4xi32> + %h_shift = vector.splat %h_shift_i : vector<4xi32> + %q0 = vector.shrui %qs, %q_shift : vector<4xi32> + %q = vector.andi %q0, %mask2 : vector<4xi32> + %hb0 = vector.shrui %hm, %h_shift : vector<4xi32> + %hb1 = vector.andi %hb0, %mask1 : vector<4xi32> + %hb = vector.shli %hb1, %two : vector<4xi32> + %u = vector.ori %q, %hb : vector<4xi32> + %u8 = vector.bitcast %u : vector<4xi32> to vector<16xi8> + %uf = vector.uitofp %u8 : vector<16xi8> to vector<16xf32> + %vf = vector.subf %uf, %four_f : vector<16xf32> + // scale s = 8n + 2j + h: word w = s / 4 = 2n + j / 2, byte b = s % 4 = 2 (j % 2) + h. + %j2 = index.div %j, %c2 : index + %jodd = index.rem %j, %c2 : index + %n2 = index.mul %n, %c2 : index + %w = index.add %n2, %j2 : index + %jodd2 = index.mul %jodd, %c2 : index + %b = index.add %jodd2, %h : index + %w_odd = index.rem %w, %c2 : index + %w_hi = index.div %w, %c2 : index + %sw0 = vector.extract %scw[0] : vector<3xi32> -> i32 + %sw1 = vector.extract %scw[1] : vector<3xi32> -> i32 + %sw2 = vector.extract %scw[2] : vector<3xi32> -> i32 + %c1 = index.constant 1 : index + %odd = index.cmp eq, %w_odd, %c1 : index + %low_word = scf.select %odd, %sw1, %sw0 : i32 + %b_i32 = index.cast %b : index to i32 + %w_hi_i32 = index.cast %w_hi : index to i32 + %w_i32 = index.cast %w : index to i32 + %b8 = scalar.muli %b_i32, %c8_i32 : i32 + %wh4 = scalar.muli %w_hi_i32, %c4_i32 : i32 + %low_shift = scalar.addi %b8, %wh4 : i32 + %w2 = scalar.muli %w_i32, %c2_i32 : i32 + %high_shift = scalar.addi %b8, %w2 : i32 + %sl0 = scalar.shrui %low_word, %low_shift : i32 + %sl = scalar.andi %sl0, %c15_i32 : i32 + %sh0 = scalar.shrui %sw2, %high_shift : i32 + %sh1 = scalar.andi %sh0, %c3_i32 : i32 + %sh = scalar.shli %sh1, %c4_i32 : i32 + %sc0 = scalar.ori %sl, %sh : i32 + %sc = scalar.subi %sc0, %c32_i32 : i32 + %sc_f = scalar.sitofp %sc : i32 to f32 + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %n128 = index.mul %n, %c128 : index + %j32 = index.mul %j, %c32 : index + %h16 = index.mul %h, %c16 : index + %p0 = index.add %n128, %j32 : index + %p = index.add %p0, %h16 : index + %x = vector.load %xv[%p] : view<256xf32> -> vector<16xf32> + %vx = vector.mulf %vf, %x : vector<16xf32> + %sum = vector.reduce %vx, %zero_scalar : vector<16xf32>, f32 + %dl = scalar.mulf %d, %sc_f : f32 + %result = scalar.mulf %dl, %sum : f32 + func.return %result : f32 +} + +// Q2_K: 84 bytes = scales[16] (low nibble scale, high nibble min), qs[64], d, dmin (f16); 2-byte +// aligned. Value 128n + 32j + 16h + l (l < 16) is d * (sc & 15) * q - dmin * (sc >> 4) with +// sc = scales[8n + 2j + h] and q = (qs[32n + 16h + l] >> 2j) & 3, so lane l16 = 8n + 2j + h owns one +// run of 16 values with one scale and one min (ggml-quants.c dequantize_row_q2_K). +func.def inline @ggml_kquant_q2k_lane_parts(%weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, f32, index) { + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c40 = index.constant 40 : index + %c41 = index.constant 41 : index + %c128 = index.constant 128 : index + %c1 = index.constant 1 : index + %c3 = index.constant 3 : index + %c0 = index.constant 0 : index + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c15_i32 = scalar.constant 15 : i32 + %mask2 = vector.constant 50529027 : vector<4xi32> + %block_bytes = index.constant 84 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %n = index.div %l, %c8 : index + %jh = index.rem %l, %c8 : index + %j = index.div %jh, %c2 : index + %h = index.rem %jh, %c2 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %iv = buffer.view %weight[%block_base] : buffer -> view<42xi16> + %hv = buffer.view %weight[%block_base] : buffer -> view<42xf16> + // qs bytes 16 + 32n + 16h .. +16 = i16 words 8 + 16n + 8h .. +8 + %h8 = index.mul %h, %c8 : index + %n16 = index.mul %n, %c16 : index + %qs_rel = index.add %n16, %h8 : index + %qs_at = index.add %c8, %qs_rel : index + %qs16 = vector.load %iv[%qs_at] : view<42xi16> -> vector<8xi16> + %sc16 = vector.load %iv[%c0] : view<42xi16> -> vector<8xi16> + %d_f16 = view.load %hv[%c40] : view<42xf16> -> f16 + %dmin_f16 = view.load %hv[%c41] : view<42xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %qs = vector.bitcast %qs16 : vector<8xi16> to vector<4xi32> + %j_i32 = index.cast %j : index to i32 + %q_shift_i = scalar.muli %j_i32, %c2_i32 : i32 + %q_shift = vector.splat %q_shift_i : vector<4xi32> + %q0 = vector.shrui %qs, %q_shift : vector<4xi32> + %q = vector.andi %q0, %mask2 : vector<4xi32> + %q8 = vector.bitcast %q : vector<4xi32> to vector<16xi8> + %qf = vector.uitofp %q8 : vector<16xi8> to vector<16xf32> + // scale byte l: word l / 4 of the 16 scale bytes, byte l % 4 of that word + %scw = vector.bitcast %sc16 : vector<8xi16> to vector<4xi32> + %sw0 = vector.extract %scw[0] : vector<4xi32> -> i32 + %sw1 = vector.extract %scw[1] : vector<4xi32> -> i32 + %sw2 = vector.extract %scw[2] : vector<4xi32> -> i32 + %sw3 = vector.extract %scw[3] : vector<4xi32> -> i32 + %w = index.div %l, %c4 : index + %b = index.rem %l, %c4 : index + %is1 = index.cmp eq, %w, %c1 : index + %is2 = index.cmp eq, %w, %c2 : index + %is3 = index.cmp eq, %w, %c3 : index + %s01 = scf.select %is1, %sw1, %sw0 : i32 + %s012 = scf.select %is2, %sw2, %s01 : i32 + %word = scf.select %is3, %sw3, %s012 : i32 + %b_i32 = index.cast %b : index to i32 + %b8 = scalar.muli %b_i32, %c8_i32 : i32 + %sc_byte0 = scalar.shrui %word, %b8 : i32 + %scl_i = scalar.andi %sc_byte0, %c15_i32 : i32 + %mn0 = scalar.shrui %sc_byte0, %c4_i32 : i32 + %mn_i = scalar.andi %mn0, %c15_i32 : i32 + %scl = scalar.uitofp %scl_i : i32 to f32 + %mn = scalar.uitofp %mn_i : i32 to f32 + %dl = scalar.mulf %d, %scl : f32 + %ml = scalar.mulf %dmin, %mn : f32 + %n128 = index.mul %n, %c128 : index + %j32 = index.mul %j, %c32 : index + %h16 = index.mul %h, %c16 : index + %pa = index.add %n128, %j32 : index + %p0 = index.add %pa, %h16 : index + func.return %qf, %dl, %ml, %p0 : vector<16xf32>, f32, f32, index +} + +func.def inline @ggml_kquant_q2k_lane_dot(%weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %zero_scalar = scalar.constant 0.0 : f32 + %x_block_bytes = index.constant 1024 : offset + %qf, %dl, %ml, %p = func.call @ggml_kquant_q2k_lane_parts(%weight, %row_base, %block, %lane16) : (buffer, offset, index, index) -> (vector<16xf32>, f32, f32, index) + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %x = vector.load %xv[%p] : view<256xf32> -> vector<16xf32> + %qx = vector.mulf %qf, %x : vector<16xf32> + %sum_qx = vector.reduce %qx, %zero_scalar : vector<16xf32>, f32 + %sum_x = vector.reduce %x, %zero_scalar : vector<16xf32>, f32 + %a = scalar.mulf %dl, %sum_qx : f32 + %m = scalar.mulf %ml, %sum_x : f32 + %result = scalar.subf %a, %m : f32 + func.return %result : f32 +} + +func.def inline @ggml_kquant_q2k_lane_weights(%weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c12 = index.constant 12 : index + %qf, %dl, %ml, %p0 = func.call @ggml_kquant_q2k_lane_parts(%weight, %row_base, %block, %lane16) : (buffer, offset, index, index) -> (vector<16xf32>, f32, f32, index) + %dl_v = vector.splat %dl : vector<16xf32> + %ml_v = vector.splat %ml : vector<16xf32> + %w0 = vector.mulf %qf, %dl_v : vector<16xf32> + %w_out = vector.subf %w0, %ml_v : vector<16xf32> + %p1 = index.add %p0, %c4 : index + %p2 = index.add %p0, %c8 : index + %p3 = index.add %p0, %c12 : index + func.return %w_out, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +// iq3s_grid from ggml-common.h (MIT, the ggml authors). Kernels cannot read Loom rodata on amdgpu, +// so each workgroup writes the 2 KiB table to workgroup memory: chunk %chunk (0..3) = entries +// 128 chunk .. +128, one chunk per wave, then a workgroup barrier. +func.def inline @ggml_kquant_iq3s_grid_fill(%grid: buffer, %chunk: index) { + %zero_offset = index.constant 0 : offset + %gv = buffer.view %grid[%zero_offset] : buffer -> view<512xi32> + %k0 = index.constant 0 : index + %is0 = index.cmp eq, %chunk, %k0 : index + scf.if %is0 { + %v0 = scalar.constant 16843009 : i32 + %v1 = scalar.constant 16843011 : i32 + %v2 = scalar.constant 16843013 : i32 + %v3 = scalar.constant 16843019 : i32 + %w0 = vector.from_elements %v0, %v1, %v2, %v3 : vector<4xi32> + %o0 = index.constant 0 : index + vector.store %w0, %gv[%o0] : vector<4xi32>, view<512xi32> + %v4 = scalar.constant 16843023 : i32 + %v5 = scalar.constant 16843521 : i32 + %v6 = scalar.constant 16843523 : i32 + %v7 = scalar.constant 16843525 : i32 + %w4 = vector.from_elements %v4, %v5, %v6, %v7 : vector<4xi32> + %o4 = index.constant 4 : index + vector.store %w4, %gv[%o4] : vector<4xi32>, view<512xi32> + %v8 = scalar.constant 16843529 : i32 + %v9 = scalar.constant 16843533 : i32 + %v10 = scalar.constant 16844033 : i32 + %v11 = scalar.constant 16844035 : i32 + %w8 = vector.from_elements %v8, %v9, %v10, %v11 : vector<4xi32> + %o8 = index.constant 8 : index + vector.store %w8, %gv[%o8] : vector<4xi32>, view<512xi32> + %v12 = scalar.constant 16844043 : i32 + %v13 = scalar.constant 16844551 : i32 + %v14 = scalar.constant 16845057 : i32 + %v15 = scalar.constant 16845061 : i32 + %w12 = vector.from_elements %v12, %v13, %v14, %v15 : vector<4xi32> + %o12 = index.constant 12 : index + vector.store %w12, %gv[%o12] : vector<4xi32>, view<512xi32> + %v16 = scalar.constant 16845067 : i32 + %v17 = scalar.constant 16845071 : i32 + %v18 = scalar.constant 16845571 : i32 + %v19 = scalar.constant 16845575 : i32 + %w16 = vector.from_elements %v16, %v17, %v18, %v19 : vector<4xi32> + %o16 = index.constant 16 : index + vector.store %w16, %gv[%o16] : vector<4xi32>, view<512xi32> + %v20 = scalar.constant 16846081 : i32 + %v21 = scalar.constant 16846085 : i32 + %v22 = scalar.constant 16846595 : i32 + %v23 = scalar.constant 16846601 : i32 + %w20 = vector.from_elements %v20, %v21, %v22, %v23 : vector<4xi32> + %o20 = index.constant 20 : index + vector.store %w20, %gv[%o20] : vector<4xi32>, view<512xi32> + %v24 = scalar.constant 16846607 : i32 + %v25 = scalar.constant 16974081 : i32 + %v26 = scalar.constant 16974083 : i32 + %v27 = scalar.constant 16974085 : i32 + %w24 = vector.from_elements %v24, %v25, %v26, %v27 : vector<4xi32> + %o24 = index.constant 24 : index + vector.store %w24, %gv[%o24] : vector<4xi32>, view<512xi32> + %v28 = scalar.constant 16974089 : i32 + %v29 = scalar.constant 16974593 : i32 + %v30 = scalar.constant 16974595 : i32 + %v31 = scalar.constant 16974603 : i32 + %w28 = vector.from_elements %v28, %v29, %v30, %v31 : vector<4xi32> + %o28 = index.constant 28 : index + vector.store %w28, %gv[%o28] : vector<4xi32>, view<512xi32> + %v32 = scalar.constant 16975105 : i32 + %v33 = scalar.constant 16975111 : i32 + %v34 = scalar.constant 16975119 : i32 + %v35 = scalar.constant 16975619 : i32 + %w32 = vector.from_elements %v32, %v33, %v34, %v35 : vector<4xi32> + %o32 = index.constant 32 : index + vector.store %w32, %gv[%o32] : vector<4xi32>, view<512xi32> + %v36 = scalar.constant 16975627 : i32 + %v37 = scalar.constant 16976137 : i32 + %v38 = scalar.constant 16977155 : i32 + %v39 = scalar.constant 16977163 : i32 + %w36 = vector.from_elements %v36, %v37, %v38, %v39 : vector<4xi32> + %o36 = index.constant 36 : index + vector.store %w36, %gv[%o36] : vector<4xi32>, view<512xi32> + %v40 = scalar.constant 16977669 : i32 + %v41 = scalar.constant 17105153 : i32 + %v42 = scalar.constant 17105155 : i32 + %v43 = scalar.constant 17105163 : i32 + %w40 = vector.from_elements %v40, %v41, %v42, %v43 : vector<4xi32> + %o40 = index.constant 40 : index + vector.store %w40, %gv[%o40] : vector<4xi32>, view<512xi32> + %v44 = scalar.constant 17105167 : i32 + %v45 = scalar.constant 17105665 : i32 + %v46 = scalar.constant 17105671 : i32 + %v47 = scalar.constant 17105677 : i32 + %w44 = vector.from_elements %v44, %v45, %v46, %v47 : vector<4xi32> + %o44 = index.constant 44 : index + vector.store %w44, %gv[%o44] : vector<4xi32>, view<512xi32> + %v48 = scalar.constant 17106179 : i32 + %v49 = scalar.constant 17106187 : i32 + %v50 = scalar.constant 17106689 : i32 + %v51 = scalar.constant 17106697 : i32 + %w48 = vector.from_elements %v48, %v49, %v50, %v51 : vector<4xi32> + %o48 = index.constant 48 : index + vector.store %w48, %gv[%o48] : vector<4xi32>, view<512xi32> + %v52 = scalar.constant 17107205 : i32 + %v53 = scalar.constant 17107211 : i32 + %v54 = scalar.constant 17107215 : i32 + %v55 = scalar.constant 17107715 : i32 + %w52 = vector.from_elements %v52, %v53, %v54, %v55 : vector<4xi32> + %o52 = index.constant 52 : index + vector.store %w52, %gv[%o52] : vector<4xi32>, view<512xi32> + %v56 = scalar.constant 17107719 : i32 + %v57 = scalar.constant 17108737 : i32 + %v58 = scalar.constant 17108743 : i32 + %v59 = scalar.constant 17236231 : i32 + %w56 = vector.from_elements %v56, %v57, %v58, %v59 : vector<4xi32> + %o56 = index.constant 56 : index + vector.store %w56, %gv[%o56] : vector<4xi32>, view<512xi32> + %v60 = scalar.constant 17236739 : i32 + %v61 = scalar.constant 17236747 : i32 + %v62 = scalar.constant 17237249 : i32 + %v63 = scalar.constant 17237253 : i32 + %w60 = vector.from_elements %v60, %v61, %v62, %v63 : vector<4xi32> + %o60 = index.constant 60 : index + vector.store %w60, %gv[%o60] : vector<4xi32>, view<512xi32> + %v64 = scalar.constant 17237763 : i32 + %v65 = scalar.constant 17237767 : i32 + %v66 = scalar.constant 17237773 : i32 + %v67 = scalar.constant 17238281 : i32 + %w64 = vector.from_elements %v64, %v65, %v66, %v67 : vector<4xi32> + %o64 = index.constant 64 : index + vector.store %w64, %gv[%o64] : vector<4xi32>, view<512xi32> + %v68 = scalar.constant 17238785 : i32 + %v69 = scalar.constant 17238789 : i32 + %v70 = scalar.constant 17239311 : i32 + %v71 = scalar.constant 17239811 : i32 + %w68 = vector.from_elements %v68, %v69, %v70, %v71 : vector<4xi32> + %o68 = index.constant 68 : index + vector.store %w68, %gv[%o68] : vector<4xi32>, view<512xi32> + %v72 = scalar.constant 17239819 : i32 + %v73 = scalar.constant 17367297 : i32 + %v74 = scalar.constant 17367815 : i32 + %v75 = scalar.constant 17367823 : i32 + %w72 = vector.from_elements %v72, %v73, %v74, %v75 : vector<4xi32> + %o72 = index.constant 72 : index + vector.store %w72, %gv[%o72] : vector<4xi32>, view<512xi32> + %v76 = scalar.constant 17368323 : i32 + %v77 = scalar.constant 17368329 : i32 + %v78 = scalar.constant 17368837 : i32 + %v79 = scalar.constant 17369345 : i32 + %w76 = vector.from_elements %v76, %v77, %v78, %v79 : vector<4xi32> + %o76 = index.constant 76 : index + vector.store %w76, %gv[%o76] : vector<4xi32>, view<512xi32> + %v80 = scalar.constant 17369351 : i32 + %v81 = scalar.constant 17369859 : i32 + %v82 = scalar.constant 17370881 : i32 + %v83 = scalar.constant 17498373 : i32 + %w80 = vector.from_elements %v80, %v81, %v82, %v83 : vector<4xi32> + %o80 = index.constant 80 : index + vector.store %w80, %gv[%o80] : vector<4xi32>, view<512xi32> + %v84 = scalar.constant 17498377 : i32 + %v85 = scalar.constant 17499393 : i32 + %v86 = scalar.constant 17499397 : i32 + %v87 = scalar.constant 17499405 : i32 + %w84 = vector.from_elements %v84, %v85, %v86, %v87 : vector<4xi32> + %o84 = index.constant 84 : index + vector.store %w84, %gv[%o84] : vector<4xi32>, view<512xi32> + %v88 = scalar.constant 17499911 : i32 + %v89 = scalar.constant 17500419 : i32 + %v90 = scalar.constant 17500427 : i32 + %v91 = scalar.constant 17500431 : i32 + %w88 = vector.from_elements %v88, %v89, %v90, %v91 : vector<4xi32> + %o88 = index.constant 88 : index + vector.store %w88, %gv[%o88] : vector<4xi32>, view<512xi32> + %v92 = scalar.constant 17501453 : i32 + %v93 = scalar.constant 17501959 : i32 + %v94 = scalar.constant 17629453 : i32 + %v95 = scalar.constant 17629955 : i32 + %w92 = vector.from_elements %v92, %v93, %v94, %v95 : vector<4xi32> + %o92 = index.constant 92 : index + vector.store %w92, %gv[%o92] : vector<4xi32>, view<512xi32> + %v96 = scalar.constant 17629959 : i32 + %v97 = scalar.constant 17630979 : i32 + %v98 = scalar.constant 17632005 : i32 + %v99 = scalar.constant 17633027 : i32 + %w96 = vector.from_elements %v96, %v97, %v98, %v99 : vector<4xi32> + %o96 = index.constant 96 : index + vector.store %w96, %gv[%o96] : vector<4xi32>, view<512xi32> + %v100 = scalar.constant 17760513 : i32 + %v101 = scalar.constant 17760517 : i32 + %v102 = scalar.constant 17760521 : i32 + %v103 = scalar.constant 17761537 : i32 + %w100 = vector.from_elements %v100, %v101, %v102, %v103 : vector<4xi32> + %o100 = index.constant 100 : index + vector.store %w100, %gv[%o100] : vector<4xi32>, view<512xi32> + %v104 = scalar.constant 17761541 : i32 + %v105 = scalar.constant 17761549 : i32 + %v106 = scalar.constant 17762055 : i32 + %v107 = scalar.constant 17763073 : i32 + %w104 = vector.from_elements %v104, %v105, %v106, %v107 : vector<4xi32> + %o104 = index.constant 104 : index + vector.store %w104, %gv[%o104] : vector<4xi32>, view<512xi32> + %v108 = scalar.constant 17763081 : i32 + %v109 = scalar.constant 50397441 : i32 + %v110 = scalar.constant 50397443 : i32 + %v111 = scalar.constant 50397445 : i32 + %w108 = vector.from_elements %v108, %v109, %v110, %v111 : vector<4xi32> + %o108 = index.constant 108 : index + vector.store %w108, %gv[%o108] : vector<4xi32>, view<512xi32> + %v112 = scalar.constant 50397449 : i32 + %v113 = scalar.constant 50397953 : i32 + %v114 = scalar.constant 50397955 : i32 + %v115 = scalar.constant 50397959 : i32 + %w112 = vector.from_elements %v112, %v113, %v114, %v115 : vector<4xi32> + %o112 = index.constant 112 : index + vector.store %w112, %gv[%o112] : vector<4xi32>, view<512xi32> + %v116 = scalar.constant 50397963 : i32 + %v117 = scalar.constant 50397967 : i32 + %v118 = scalar.constant 50398465 : i32 + %v119 = scalar.constant 50398469 : i32 + %w116 = vector.from_elements %v116, %v117, %v118, %v119 : vector<4xi32> + %o116 = index.constant 116 : index + vector.store %w116, %gv[%o116] : vector<4xi32>, view<512xi32> + %v120 = scalar.constant 50398979 : i32 + %v121 = scalar.constant 50398985 : i32 + %v122 = scalar.constant 50398989 : i32 + %v123 = scalar.constant 50400009 : i32 + %w120 = vector.from_elements %v120, %v121, %v122, %v123 : vector<4xi32> + %o120 = index.constant 120 : index + vector.store %w120, %gv[%o120] : vector<4xi32>, view<512xi32> + %v124 = scalar.constant 50400013 : i32 + %v125 = scalar.constant 50400515 : i32 + %v126 = scalar.constant 50401029 : i32 + %v127 = scalar.constant 50528513 : i32 + %w124 = vector.from_elements %v124, %v125, %v126, %v127 : vector<4xi32> + %o124 = index.constant 124 : index + vector.store %w124, %gv[%o124] : vector<4xi32>, view<512xi32> + } + %k1 = index.constant 1 : index + %is1 = index.cmp eq, %chunk, %k1 : index + scf.if %is1 { + %v128 = scalar.constant 50528515 : i32 + %v129 = scalar.constant 50528519 : i32 + %v130 = scalar.constant 50528525 : i32 + %v131 = scalar.constant 50529025 : i32 + %w128 = vector.from_elements %v128, %v129, %v130, %v131 : vector<4xi32> + %o128 = index.constant 128 : index + vector.store %w128, %gv[%o128] : vector<4xi32>, view<512xi32> + %v132 = scalar.constant 50529033 : i32 + %v133 = scalar.constant 50529539 : i32 + %v134 = scalar.constant 50530049 : i32 + %v135 = scalar.constant 50530055 : i32 + %w132 = vector.from_elements %v132, %v133, %v134, %v135 : vector<4xi32> + %o132 = index.constant 132 : index + vector.store %w132, %gv[%o132] : vector<4xi32>, view<512xi32> + %v136 = scalar.constant 50530563 : i32 + %v137 = scalar.constant 50531073 : i32 + %v138 = scalar.constant 50531077 : i32 + %v139 = scalar.constant 50532097 : i32 + %w136 = vector.from_elements %v136, %v137, %v138, %v139 : vector<4xi32> + %o136 = index.constant 136 : index + vector.store %w136, %gv[%o136] : vector<4xi32>, view<512xi32> + %v140 = scalar.constant 50532109 : i32 + %v141 = scalar.constant 50659585 : i32 + %v142 = scalar.constant 50660101 : i32 + %v143 = scalar.constant 50660107 : i32 + %w140 = vector.from_elements %v140, %v141, %v142, %v143 : vector<4xi32> + %o140 = index.constant 140 : index + vector.store %w140, %gv[%o140] : vector<4xi32>, view<512xi32> + %v144 = scalar.constant 50660111 : i32 + %v145 = scalar.constant 50660609 : i32 + %v146 = scalar.constant 50660617 : i32 + %v147 = scalar.constant 50661125 : i32 + %w144 = vector.from_elements %v144, %v145, %v146, %v147 : vector<4xi32> + %o144 = index.constant 144 : index + vector.store %w144, %gv[%o144] : vector<4xi32>, view<512xi32> + %v148 = scalar.constant 50661633 : i32 + %v149 = scalar.constant 50661639 : i32 + %v150 = scalar.constant 50662155 : i32 + %v151 = scalar.constant 50662657 : i32 + %w148 = vector.from_elements %v148, %v149, %v150, %v151 : vector<4xi32> + %o148 = index.constant 148 : index + vector.store %w148, %gv[%o148] : vector<4xi32>, view<512xi32> + %v152 = scalar.constant 50663173 : i32 + %v153 = scalar.constant 50790659 : i32 + %v154 = scalar.constant 50790665 : i32 + %v155 = scalar.constant 50790671 : i32 + %w152 = vector.from_elements %v152, %v153, %v154, %v155 : vector<4xi32> + %o152 = index.constant 152 : index + vector.store %w152, %gv[%o152] : vector<4xi32>, view<512xi32> + %v156 = scalar.constant 50791169 : i32 + %v157 = scalar.constant 50791175 : i32 + %v158 = scalar.constant 50791683 : i32 + %v159 = scalar.constant 50791695 : i32 + %w156 = vector.from_elements %v156, %v157, %v158, %v159 : vector<4xi32> + %o156 = index.constant 156 : index + vector.store %w156, %gv[%o156] : vector<4xi32>, view<512xi32> + %v160 = scalar.constant 50792193 : i32 + %v161 = scalar.constant 50792201 : i32 + %v162 = scalar.constant 50792707 : i32 + %v163 = scalar.constant 50793733 : i32 + %w160 = vector.from_elements %v160, %v161, %v162, %v163 : vector<4xi32> + %o160 = index.constant 160 : index + vector.store %w160, %gv[%o160] : vector<4xi32>, view<512xi32> + %v164 = scalar.constant 50794241 : i32 + %v165 = scalar.constant 50921735 : i32 + %v166 = scalar.constant 50921739 : i32 + %v167 = scalar.constant 50922245 : i32 + %w164 = vector.from_elements %v164, %v165, %v166, %v167 : vector<4xi32> + %o164 = index.constant 164 : index + vector.store %w164, %gv[%o164] : vector<4xi32>, view<512xi32> + %v168 = scalar.constant 50922249 : i32 + %v169 = scalar.constant 50923267 : i32 + %v170 = scalar.constant 50923271 : i32 + %v171 = scalar.constant 50923781 : i32 + %w168 = vector.from_elements %v168, %v169, %v170, %v171 : vector<4xi32> + %o168 = index.constant 168 : index + vector.store %w168, %gv[%o168] : vector<4xi32>, view<512xi32> + %v172 = scalar.constant 50923789 : i32 + %v173 = scalar.constant 50924289 : i32 + %v174 = scalar.constant 50924297 : i32 + %v175 = scalar.constant 51052803 : i32 + %w172 = vector.from_elements %v172, %v173, %v174, %v175 : vector<4xi32> + %o172 = index.constant 172 : index + vector.store %w172, %gv[%o172] : vector<4xi32>, view<512xi32> + %v176 = scalar.constant 51053313 : i32 + %v177 = scalar.constant 51053319 : i32 + %v178 = scalar.constant 51053827 : i32 + %v179 = scalar.constant 51054337 : i32 + %w176 = vector.from_elements %v176, %v177, %v178, %v179 : vector<4xi32> + %o176 = index.constant 176 : index + vector.store %w176, %gv[%o176] : vector<4xi32>, view<512xi32> + %v180 = scalar.constant 51054341 : i32 + %v181 = scalar.constant 51055363 : i32 + %v182 = scalar.constant 51184897 : i32 + %v183 = scalar.constant 51184905 : i32 + %w180 = vector.from_elements %v180, %v181, %v182, %v183 : vector<4xi32> + %o180 = index.constant 180 : index + vector.store %w180, %gv[%o180] : vector<4xi32>, view<512xi32> + %v184 = scalar.constant 51184911 : i32 + %v185 = scalar.constant 51185929 : i32 + %v186 = scalar.constant 51185933 : i32 + %v187 = scalar.constant 51314947 : i32 + %w184 = vector.from_elements %v184, %v185, %v186, %v187 : vector<4xi32> + %o184 = index.constant 184 : index + vector.store %w184, %gv[%o184] : vector<4xi32>, view<512xi32> + %v188 = scalar.constant 51314951 : i32 + %v189 = scalar.constant 51315457 : i32 + %v190 = scalar.constant 51315461 : i32 + %v191 = scalar.constant 51315971 : i32 + %w188 = vector.from_elements %v188, %v189, %v190, %v191 : vector<4xi32> + %o188 = index.constant 188 : index + vector.store %w188, %gv[%o188] : vector<4xi32>, view<512xi32> + %v192 = scalar.constant 51316491 : i32 + %v193 = scalar.constant 51316995 : i32 + %v194 = scalar.constant 51318021 : i32 + %v195 = scalar.constant 51318529 : i32 + %w192 = vector.from_elements %v192, %v193, %v194, %v195 : vector<4xi32> + %o192 = index.constant 192 : index + vector.store %w192, %gv[%o192] : vector<4xi32>, view<512xi32> + %v196 = scalar.constant 83951873 : i32 + %v197 = scalar.constant 83951875 : i32 + %v198 = scalar.constant 83951879 : i32 + %v199 = scalar.constant 83951883 : i32 + %w196 = vector.from_elements %v196, %v197, %v198, %v199 : vector<4xi32> + %o196 = index.constant 196 : index + vector.store %w196, %gv[%o196] : vector<4xi32>, view<512xi32> + %v200 = scalar.constant 83951887 : i32 + %v201 = scalar.constant 83952385 : i32 + %v202 = scalar.constant 83952389 : i32 + %v203 = scalar.constant 83952393 : i32 + %w200 = vector.from_elements %v200, %v201, %v202, %v203 : vector<4xi32> + %o200 = index.constant 200 : index + vector.store %w200, %gv[%o200] : vector<4xi32>, view<512xi32> + %v204 = scalar.constant 83952397 : i32 + %v205 = scalar.constant 83952899 : i32 + %v206 = scalar.constant 83952903 : i32 + %v207 = scalar.constant 83952911 : i32 + %w204 = vector.from_elements %v204, %v205, %v206, %v207 : vector<4xi32> + %o204 = index.constant 204 : index + vector.store %w204, %gv[%o204] : vector<4xi32>, view<512xi32> + %v208 = scalar.constant 83953409 : i32 + %v209 = scalar.constant 83953413 : i32 + %v210 = scalar.constant 83953923 : i32 + %v211 = scalar.constant 83953927 : i32 + %w208 = vector.from_elements %v208, %v209, %v210, %v211 : vector<4xi32> + %o208 = index.constant 208 : index + vector.store %w208, %gv[%o208] : vector<4xi32>, view<512xi32> + %v212 = scalar.constant 83953931 : i32 + %v213 = scalar.constant 83954433 : i32 + %v214 = scalar.constant 83954437 : i32 + %v215 = scalar.constant 83954959 : i32 + %w212 = vector.from_elements %v212, %v213, %v214, %v215 : vector<4xi32> + %o212 = index.constant 212 : index + vector.store %w212, %gv[%o212] : vector<4xi32>, view<512xi32> + %v216 = scalar.constant 83955457 : i32 + %v217 = scalar.constant 83955463 : i32 + %v218 = scalar.constant 83955467 : i32 + %v219 = scalar.constant 84082945 : i32 + %w216 = vector.from_elements %v216, %v217, %v218, %v219 : vector<4xi32> + %o216 = index.constant 216 : index + vector.store %w216, %gv[%o216] : vector<4xi32>, view<512xi32> + %v220 = scalar.constant 84082949 : i32 + %v221 = scalar.constant 84083457 : i32 + %v222 = scalar.constant 84083463 : i32 + %v223 = scalar.constant 84083471 : i32 + %w220 = vector.from_elements %v220, %v221, %v222, %v223 : vector<4xi32> + %o220 = index.constant 220 : index + vector.store %w220, %gv[%o220] : vector<4xi32>, view<512xi32> + %v224 = scalar.constant 84083973 : i32 + %v225 = scalar.constant 84083979 : i32 + %v226 = scalar.constant 84084483 : i32 + %v227 = scalar.constant 84084489 : i32 + %w224 = vector.from_elements %v224, %v225, %v226, %v227 : vector<4xi32> + %o224 = index.constant 224 : index + vector.store %w224, %gv[%o224] : vector<4xi32>, view<512xi32> + %v228 = scalar.constant 84084997 : i32 + %v229 = scalar.constant 84085507 : i32 + %v230 = scalar.constant 84214019 : i32 + %v231 = scalar.constant 84214025 : i32 + %w228 = vector.from_elements %v228, %v229, %v230, %v231 : vector<4xi32> + %o228 = index.constant 228 : index + vector.store %w228, %gv[%o228] : vector<4xi32>, view<512xi32> + %v232 = scalar.constant 84214031 : i32 + %v233 = scalar.constant 84215043 : i32 + %v234 = scalar.constant 84215047 : i32 + %v235 = scalar.constant 84215553 : i32 + %w232 = vector.from_elements %v232, %v233, %v234, %v235 : vector<4xi32> + %o232 = index.constant 232 : index + vector.store %w232, %gv[%o232] : vector<4xi32>, view<512xi32> + %v236 = scalar.constant 84215567 : i32 + %v237 = scalar.constant 84216067 : i32 + %v238 = scalar.constant 84216583 : i32 + %v239 = scalar.constant 84216591 : i32 + %w236 = vector.from_elements %v236, %v237, %v238, %v239 : vector<4xi32> + %o236 = index.constant 236 : index + vector.store %w236, %gv[%o236] : vector<4xi32>, view<512xi32> + %v240 = scalar.constant 84217603 : i32 + %v241 = scalar.constant 84217609 : i32 + %v242 = scalar.constant 84345089 : i32 + %v243 = scalar.constant 84345093 : i32 + %w240 = vector.from_elements %v240, %v241, %v242, %v243 : vector<4xi32> + %o240 = index.constant 240 : index + vector.store %w240, %gv[%o240] : vector<4xi32>, view<512xi32> + %v244 = scalar.constant 84345099 : i32 + %v245 = scalar.constant 84345603 : i32 + %v246 = scalar.constant 84346117 : i32 + %v247 = scalar.constant 84346121 : i32 + %w244 = vector.from_elements %v244, %v245, %v246, %v247 : vector<4xi32> + %o244 = index.constant 244 : index + vector.store %w244, %gv[%o244] : vector<4xi32>, view<512xi32> + %v248 = scalar.constant 84346627 : i32 + %v249 = scalar.constant 84346631 : i32 + %v250 = scalar.constant 84347141 : i32 + %v251 = scalar.constant 84347649 : i32 + %w248 = vector.from_elements %v248, %v249, %v250, %v251 : vector<4xi32> + %o248 = index.constant 248 : index + vector.store %w248, %gv[%o248] : vector<4xi32>, view<512xi32> + %v252 = scalar.constant 84348173 : i32 + %v253 = scalar.constant 84476163 : i32 + %v254 = scalar.constant 84476175 : i32 + %v255 = scalar.constant 84477185 : i32 + %w252 = vector.from_elements %v252, %v253, %v254, %v255 : vector<4xi32> + %o252 = index.constant 252 : index + vector.store %w252, %gv[%o252] : vector<4xi32>, view<512xi32> + } + %k2 = index.constant 2 : index + %is2 = index.cmp eq, %chunk, %k2 : index + scf.if %is2 { + %v256 = scalar.constant 84477191 : i32 + %v257 = scalar.constant 84477701 : i32 + %v258 = scalar.constant 84477707 : i32 + %v259 = scalar.constant 84478211 : i32 + %w256 = vector.from_elements %v256, %v257, %v258, %v259 : vector<4xi32> + %o256 = index.constant 256 : index + vector.store %w256, %gv[%o256] : vector<4xi32>, view<512xi32> + %v260 = scalar.constant 84479749 : i32 + %v261 = scalar.constant 84479755 : i32 + %v262 = scalar.constant 84607241 : i32 + %v263 = scalar.constant 84607747 : i32 + %w260 = vector.from_elements %v260, %v261, %v262, %v263 : vector<4xi32> + %o260 = index.constant 260 : index + vector.store %w260, %gv[%o260] : vector<4xi32>, view<512xi32> + %v264 = scalar.constant 84608261 : i32 + %v265 = scalar.constant 84608783 : i32 + %v266 = scalar.constant 84609281 : i32 + %v267 = scalar.constant 84609799 : i32 + %w264 = vector.from_elements %v264, %v265, %v266, %v267 : vector<4xi32> + %o264 = index.constant 264 : index + vector.store %w264, %gv[%o264] : vector<4xi32>, view<512xi32> + %v268 = scalar.constant 84610817 : i32 + %v269 = scalar.constant 84738305 : i32 + %v270 = scalar.constant 84738309 : i32 + %v271 = scalar.constant 84738319 : i32 + %w268 = vector.from_elements %v268, %v269, %v270, %v271 : vector<4xi32> + %o268 = index.constant 268 : index + vector.store %w268, %gv[%o268] : vector<4xi32>, view<512xi32> + %v272 = scalar.constant 84739331 : i32 + %v273 = scalar.constant 84740875 : i32 + %v274 = scalar.constant 84741379 : i32 + %v275 = scalar.constant 84869387 : i32 + %w272 = vector.from_elements %v272, %v273, %v274, %v275 : vector<4xi32> + %o272 = index.constant 272 : index + vector.store %w272, %gv[%o272] : vector<4xi32>, view<512xi32> + %v276 = scalar.constant 84869891 : i32 + %v277 = scalar.constant 84870413 : i32 + %v278 = scalar.constant 84870913 : i32 + %v279 = scalar.constant 84871431 : i32 + %w276 = vector.from_elements %v276, %v277, %v278, %v279 : vector<4xi32> + %o276 = index.constant 276 : index + vector.store %w276, %gv[%o276] : vector<4xi32>, view<512xi32> + %v280 = scalar.constant 84871937 : i32 + %v281 = scalar.constant 117506309 : i32 + %v282 = scalar.constant 117506819 : i32 + %v283 = scalar.constant 117506823 : i32 + %w280 = vector.from_elements %v280, %v281, %v282, %v283 : vector<4xi32> + %o280 = index.constant 280 : index + vector.store %w280, %gv[%o280] : vector<4xi32>, view<512xi32> + %v284 = scalar.constant 117506827 : i32 + %v285 = scalar.constant 117506831 : i32 + %v286 = scalar.constant 117507333 : i32 + %v287 = scalar.constant 117507843 : i32 + %w284 = vector.from_elements %v284, %v285, %v286, %v287 : vector<4xi32> + %o284 = index.constant 284 : index + vector.store %w284, %gv[%o284] : vector<4xi32>, view<512xi32> + %v288 = scalar.constant 117507847 : i32 + %v289 = scalar.constant 117507851 : i32 + %v290 = scalar.constant 117508357 : i32 + %v291 = scalar.constant 117508361 : i32 + %w288 = vector.from_elements %v288, %v289, %v290, %v291 : vector<4xi32> + %o288 = index.constant 288 : index + vector.store %w288, %gv[%o288] : vector<4xi32>, view<512xi32> + %v292 = scalar.constant 117508367 : i32 + %v293 = scalar.constant 117508867 : i32 + %v294 = scalar.constant 117509383 : i32 + %v295 = scalar.constant 117509891 : i32 + %w292 = vector.from_elements %v292, %v293, %v294, %v295 : vector<4xi32> + %o292 = index.constant 292 : index + vector.store %w292, %gv[%o292] : vector<4xi32>, view<512xi32> + %v296 = scalar.constant 117637379 : i32 + %v297 = scalar.constant 117637383 : i32 + %v298 = scalar.constant 117637387 : i32 + %v299 = scalar.constant 117637897 : i32 + %w296 = vector.from_elements %v296, %v297, %v298, %v299 : vector<4xi32> + %o296 = index.constant 296 : index + vector.store %w296, %gv[%o296] : vector<4xi32>, view<512xi32> + %v300 = scalar.constant 117638403 : i32 + %v301 = scalar.constant 117638407 : i32 + %v302 = scalar.constant 117639425 : i32 + %v303 = scalar.constant 117640449 : i32 + %w300 = vector.from_elements %v300, %v301, %v302, %v303 : vector<4xi32> + %o300 = index.constant 300 : index + vector.store %w300, %gv[%o300] : vector<4xi32>, view<512xi32> + %v304 = scalar.constant 117640965 : i32 + %v305 = scalar.constant 117640973 : i32 + %v306 = scalar.constant 117768449 : i32 + %v307 = scalar.constant 117768965 : i32 + %w304 = vector.from_elements %v304, %v305, %v306, %v307 : vector<4xi32> + %o304 = index.constant 304 : index + vector.store %w304, %gv[%o304] : vector<4xi32>, view<512xi32> + %v308 = scalar.constant 117769473 : i32 + %v309 = scalar.constant 117769989 : i32 + %v310 = scalar.constant 117769993 : i32 + %v311 = scalar.constant 117771009 : i32 + %w308 = vector.from_elements %v308, %v309, %v310, %v311 : vector<4xi32> + %o308 = index.constant 308 : index + vector.store %w308, %gv[%o308] : vector<4xi32>, view<512xi32> + %v312 = scalar.constant 117899523 : i32 + %v313 = scalar.constant 117900033 : i32 + %v314 = scalar.constant 117900041 : i32 + %v315 = scalar.constant 117900547 : i32 + %w312 = vector.from_elements %v312, %v313, %v314, %v315 : vector<4xi32> + %o312 = index.constant 312 : index + vector.store %w312, %gv[%o312] : vector<4xi32>, view<512xi32> + %v316 = scalar.constant 117900551 : i32 + %v317 = scalar.constant 117900559 : i32 + %v318 = scalar.constant 117901057 : i32 + %v319 = scalar.constant 117901571 : i32 + %w316 = vector.from_elements %v316, %v317, %v318, %v319 : vector<4xi32> + %o316 = index.constant 316 : index + vector.store %w316, %gv[%o316] : vector<4xi32>, view<512xi32> + %v320 = scalar.constant 117901575 : i32 + %v321 = scalar.constant 117901583 : i32 + %v322 = scalar.constant 117902091 : i32 + %v323 = scalar.constant 117903111 : i32 + %w320 = vector.from_elements %v320, %v321, %v322, %v323 : vector<4xi32> + %o320 = index.constant 320 : index + vector.store %w320, %gv[%o320] : vector<4xi32>, view<512xi32> + %v324 = scalar.constant 118030599 : i32 + %v325 = scalar.constant 118031107 : i32 + %v326 = scalar.constant 118031117 : i32 + %v327 = scalar.constant 118031621 : i32 + %w324 = vector.from_elements %v324, %v325, %v326, %v327 : vector<4xi32> + %o324 = index.constant 324 : index + vector.store %w324, %gv[%o324] : vector<4xi32>, view<512xi32> + %v328 = scalar.constant 118032131 : i32 + %v329 = scalar.constant 118033157 : i32 + %v330 = scalar.constant 118033665 : i32 + %v331 = scalar.constant 118033673 : i32 + %w328 = vector.from_elements %v328, %v329, %v330, %v331 : vector<4xi32> + %o328 = index.constant 328 : index + vector.store %w328, %gv[%o328] : vector<4xi32>, view<512xi32> + %v332 = scalar.constant 118161667 : i32 + %v333 = scalar.constant 118162177 : i32 + %v334 = scalar.constant 118162181 : i32 + %v335 = scalar.constant 118162699 : i32 + %w332 = vector.from_elements %v332, %v333, %v334, %v335 : vector<4xi32> + %o332 = index.constant 332 : index + vector.store %w332, %gv[%o332] : vector<4xi32>, view<512xi32> + %v336 = scalar.constant 118163205 : i32 + %v337 = scalar.constant 118163721 : i32 + %v338 = scalar.constant 118164237 : i32 + %v339 = scalar.constant 118165255 : i32 + %w336 = vector.from_elements %v336, %v337, %v338, %v339 : vector<4xi32> + %o336 = index.constant 336 : index + vector.store %w336, %gv[%o336] : vector<4xi32>, view<512xi32> + %v340 = scalar.constant 118293261 : i32 + %v341 = scalar.constant 118294787 : i32 + %v342 = scalar.constant 118423811 : i32 + %v343 = scalar.constant 118423815 : i32 + %w340 = vector.from_elements %v340, %v341, %v342, %v343 : vector<4xi32> + %o340 = index.constant 340 : index + vector.store %w340, %gv[%o340] : vector<4xi32>, view<512xi32> + %v344 = scalar.constant 118424833 : i32 + %v345 = scalar.constant 118424837 : i32 + %v346 = scalar.constant 118425355 : i32 + %v347 = scalar.constant 151060737 : i32 + %w344 = vector.from_elements %v344, %v345, %v346, %v347 : vector<4xi32> + %o344 = index.constant 344 : index + vector.store %w344, %gv[%o344] : vector<4xi32>, view<512xi32> + %v348 = scalar.constant 151060745 : i32 + %v349 = scalar.constant 151061253 : i32 + %v350 = scalar.constant 151061761 : i32 + %v351 = scalar.constant 151061769 : i32 + %w348 = vector.from_elements %v348, %v349, %v350, %v351 : vector<4xi32> + %o348 = index.constant 348 : index + vector.store %w348, %gv[%o348] : vector<4xi32>, view<512xi32> + %v352 = scalar.constant 151061775 : i32 + %v353 = scalar.constant 151062277 : i32 + %v354 = scalar.constant 151062787 : i32 + %v355 = scalar.constant 151063297 : i32 + %w352 = vector.from_elements %v352, %v353, %v354, %v355 : vector<4xi32> + %o352 = index.constant 352 : index + vector.store %w352, %gv[%o352] : vector<4xi32>, view<512xi32> + %v356 = scalar.constant 151064321 : i32 + %v357 = scalar.constant 151191813 : i32 + %v358 = scalar.constant 151191823 : i32 + %v359 = scalar.constant 151192323 : i32 + %w356 = vector.from_elements %v356, %v357, %v358, %v359 : vector<4xi32> + %o356 = index.constant 356 : index + vector.store %w356, %gv[%o356] : vector<4xi32>, view<512xi32> + %v360 = scalar.constant 151192327 : i32 + %v361 = scalar.constant 151192837 : i32 + %v362 = scalar.constant 151193345 : i32 + %v363 = scalar.constant 151193355 : i32 + %w360 = vector.from_elements %v360, %v361, %v362, %v363 : vector<4xi32> + %o360 = index.constant 360 : index + vector.store %w360, %gv[%o360] : vector<4xi32>, view<512xi32> + %v364 = scalar.constant 151193863 : i32 + %v365 = scalar.constant 151194371 : i32 + %v366 = scalar.constant 151194379 : i32 + %v367 = scalar.constant 151322883 : i32 + %w364 = vector.from_elements %v364, %v365, %v366, %v367 : vector<4xi32> + %o364 = index.constant 364 : index + vector.store %w364, %gv[%o364] : vector<4xi32>, view<512xi32> + %v368 = scalar.constant 151322887 : i32 + %v369 = scalar.constant 151323393 : i32 + %v370 = scalar.constant 151323403 : i32 + %v371 = scalar.constant 151323907 : i32 + %w368 = vector.from_elements %v368, %v369, %v370, %v371 : vector<4xi32> + %o368 = index.constant 368 : index + vector.store %w368, %gv[%o368] : vector<4xi32>, view<512xi32> + %v372 = scalar.constant 151324423 : i32 + %v373 = scalar.constant 151324929 : i32 + %v374 = scalar.constant 151325455 : i32 + %v375 = scalar.constant 151325957 : i32 + %w372 = vector.from_elements %v372, %v373, %v374, %v375 : vector<4xi32> + %o372 = index.constant 372 : index + vector.store %w372, %gv[%o372] : vector<4xi32>, view<512xi32> + %v376 = scalar.constant 151326465 : i32 + %v377 = scalar.constant 151453961 : i32 + %v378 = scalar.constant 151454467 : i32 + %v379 = scalar.constant 151454471 : i32 + %w376 = vector.from_elements %v376, %v377, %v378, %v379 : vector<4xi32> + %o376 = index.constant 376 : index + vector.store %w376, %gv[%o376] : vector<4xi32>, view<512xi32> + %v380 = scalar.constant 151454977 : i32 + %v381 = scalar.constant 151454981 : i32 + %v382 = scalar.constant 151455491 : i32 + %v383 = scalar.constant 151455499 : i32 + %w380 = vector.from_elements %v380, %v381, %v382, %v383 : vector<4xi32> + %o380 = index.constant 380 : index + vector.store %w380, %gv[%o380] : vector<4xi32>, view<512xi32> + } + %k3 = index.constant 3 : index + %is3 = index.cmp eq, %chunk, %k3 : index + scf.if %is3 { + %v384 = scalar.constant 151585025 : i32 + %v385 = scalar.constant 151585029 : i32 + %v386 = scalar.constant 151586057 : i32 + %v387 = scalar.constant 151586575 : i32 + %w384 = vector.from_elements %v384, %v385, %v386, %v387 : vector<4xi32> + %o384 = index.constant 384 : index + vector.store %w384, %gv[%o384] : vector<4xi32>, view<512xi32> + %v388 = scalar.constant 151587073 : i32 + %v389 = scalar.constant 151588611 : i32 + %v390 = scalar.constant 151716107 : i32 + %v391 = scalar.constant 151716111 : i32 + %w388 = vector.from_elements %v388, %v389, %v390, %v391 : vector<4xi32> + %o388 = index.constant 388 : index + vector.store %w388, %gv[%o388] : vector<4xi32>, view<512xi32> + %v392 = scalar.constant 151717123 : i32 + %v393 = scalar.constant 151719173 : i32 + %v394 = scalar.constant 151847687 : i32 + %v395 = scalar.constant 151848713 : i32 + %w392 = vector.from_elements %v392, %v393, %v394, %v395 : vector<4xi32> + %o392 = index.constant 392 : index + vector.store %w392, %gv[%o392] : vector<4xi32>, view<512xi32> + %v396 = scalar.constant 151850241 : i32 + %v397 = scalar.constant 151978753 : i32 + %v398 = scalar.constant 151978763 : i32 + %v399 = scalar.constant 151979777 : i32 + %w396 = vector.from_elements %v396, %v397, %v398, %v399 : vector<4xi32> + %o396 = index.constant 396 : index + vector.store %w396, %gv[%o396] : vector<4xi32>, view<512xi32> + %v400 = scalar.constant 151980295 : i32 + %v401 = scalar.constant 151980803 : i32 + %v402 = scalar.constant 184615173 : i32 + %v403 = scalar.constant 184615681 : i32 + %w400 = vector.from_elements %v400, %v401, %v402, %v403 : vector<4xi32> + %o400 = index.constant 400 : index + vector.store %w400, %gv[%o400] : vector<4xi32>, view<512xi32> + %v404 = scalar.constant 184615689 : i32 + %v405 = scalar.constant 184616197 : i32 + %v406 = scalar.constant 184617217 : i32 + %v407 = scalar.constant 184617225 : i32 + %w404 = vector.from_elements %v404, %v405, %v406, %v407 : vector<4xi32> + %o404 = index.constant 404 : index + vector.store %w404, %gv[%o404] : vector<4xi32>, view<512xi32> + %v408 = scalar.constant 184617231 : i32 + %v409 = scalar.constant 184617733 : i32 + %v410 = scalar.constant 184618253 : i32 + %v411 = scalar.constant 184618761 : i32 + %w408 = vector.from_elements %v408, %v409, %v410, %v411 : vector<4xi32> + %o408 = index.constant 408 : index + vector.store %w408, %gv[%o408] : vector<4xi32>, view<512xi32> + %v412 = scalar.constant 184746243 : i32 + %v413 = scalar.constant 184746247 : i32 + %v414 = scalar.constant 184746251 : i32 + %v415 = scalar.constant 184746757 : i32 + %w412 = vector.from_elements %v412, %v413, %v414, %v415 : vector<4xi32> + %o412 = index.constant 412 : index + vector.store %w412, %gv[%o412] : vector<4xi32>, view<512xi32> + %v416 = scalar.constant 184747267 : i32 + %v417 = scalar.constant 184747781 : i32 + %v418 = scalar.constant 184749829 : i32 + %v419 = scalar.constant 184877313 : i32 + %w416 = vector.from_elements %v416, %v417, %v418, %v419 : vector<4xi32> + %o416 = index.constant 416 : index + vector.store %w416, %gv[%o416] : vector<4xi32>, view<512xi32> + %v420 = scalar.constant 184877827 : i32 + %v421 = scalar.constant 184878343 : i32 + %v422 = scalar.constant 184878849 : i32 + %v423 = scalar.constant 184878861 : i32 + %w420 = vector.from_elements %v420, %v421, %v422, %v423 : vector<4xi32> + %o420 = index.constant 420 : index + vector.store %w420, %gv[%o420] : vector<4xi32>, view<512xi32> + %v424 = scalar.constant 184879879 : i32 + %v425 = scalar.constant 185008389 : i32 + %v426 = scalar.constant 185008399 : i32 + %v427 = scalar.constant 185008897 : i32 + %w424 = vector.from_elements %v424, %v425, %v426, %v427 : vector<4xi32> + %o424 = index.constant 424 : index + vector.store %w424, %gv[%o424] : vector<4xi32>, view<512xi32> + %v428 = scalar.constant 185009423 : i32 + %v429 = scalar.constant 185010441 : i32 + %v430 = scalar.constant 185010947 : i32 + %v431 = scalar.constant 185011467 : i32 + %w428 = vector.from_elements %v428, %v429, %v430, %v431 : vector<4xi32> + %o428 = index.constant 428 : index + vector.store %w428, %gv[%o428] : vector<4xi32>, view<512xi32> + %v432 = scalar.constant 185011975 : i32 + %v433 = scalar.constant 185139459 : i32 + %v434 = scalar.constant 185139465 : i32 + %v435 = scalar.constant 185140481 : i32 + %w432 = vector.from_elements %v432, %v433, %v434, %v435 : vector<4xi32> + %o432 = index.constant 432 : index + vector.store %w432, %gv[%o432] : vector<4xi32>, view<512xi32> + %v436 = scalar.constant 185140997 : i32 + %v437 = scalar.constant 185141517 : i32 + %v438 = scalar.constant 185271045 : i32 + %v439 = scalar.constant 185271565 : i32 + %w436 = vector.from_elements %v436, %v437, %v438, %v439 : vector<4xi32> + %o436 = index.constant 436 : index + vector.store %w436, %gv[%o436] : vector<4xi32>, view<512xi32> + %v440 = scalar.constant 185273091 : i32 + %v441 = scalar.constant 185273095 : i32 + %v442 = scalar.constant 185403653 : i32 + %v443 = scalar.constant 185532677 : i32 + %w440 = vector.from_elements %v440, %v441, %v442, %v443 : vector<4xi32> + %o440 = index.constant 440 : index + vector.store %w440, %gv[%o440] : vector<4xi32>, view<512xi32> + %v444 = scalar.constant 185532681 : i32 + %v445 = scalar.constant 185533701 : i32 + %v446 = scalar.constant 218170115 : i32 + %v447 = scalar.constant 218170119 : i32 + %w444 = vector.from_elements %v444, %v445, %v446, %v447 : vector<4xi32> + %o444 = index.constant 444 : index + vector.store %w444, %gv[%o444] : vector<4xi32>, view<512xi32> + %v448 = scalar.constant 218170123 : i32 + %v449 = scalar.constant 218171139 : i32 + %v450 = scalar.constant 218171143 : i32 + %v451 = scalar.constant 218172673 : i32 + %w448 = vector.from_elements %v448, %v449, %v450, %v451 : vector<4xi32> + %o448 = index.constant 448 : index + vector.store %w448, %gv[%o448] : vector<4xi32>, view<512xi32> + %v452 = scalar.constant 218300673 : i32 + %v453 = scalar.constant 218301697 : i32 + %v454 = scalar.constant 218301711 : i32 + %v455 = scalar.constant 218303753 : i32 + %w452 = vector.from_elements %v452, %v453, %v454, %v455 : vector<4xi32> + %o452 = index.constant 452 : index + vector.store %w452, %gv[%o452] : vector<4xi32>, view<512xi32> + %v456 = scalar.constant 218432261 : i32 + %v457 = scalar.constant 218433289 : i32 + %v458 = scalar.constant 218433797 : i32 + %v459 = scalar.constant 218434315 : i32 + %w456 = vector.from_elements %v456, %v457, %v458, %v459 : vector<4xi32> + %o456 = index.constant 456 : index + vector.store %w456, %gv[%o456] : vector<4xi32>, view<512xi32> + %v460 = scalar.constant 218434821 : i32 + %v461 = scalar.constant 218435329 : i32 + %v462 = scalar.constant 218562817 : i32 + %v463 = scalar.constant 218563337 : i32 + %w460 = vector.from_elements %v460, %v461, %v462, %v463 : vector<4xi32> + %o460 = index.constant 460 : index + vector.store %w460, %gv[%o460] : vector<4xi32>, view<512xi32> + %v464 = scalar.constant 218563843 : i32 + %v465 = scalar.constant 218564865 : i32 + %v466 = scalar.constant 218694923 : i32 + %v467 = scalar.constant 218695943 : i32 + %w464 = vector.from_elements %v464, %v465, %v466, %v467 : vector<4xi32> + %o464 = index.constant 464 : index + vector.store %w464, %gv[%o464] : vector<4xi32>, view<512xi32> + %v468 = scalar.constant 218696965 : i32 + %v469 = scalar.constant 218824961 : i32 + %v470 = scalar.constant 218824967 : i32 + %v471 = scalar.constant 218826505 : i32 + %w468 = vector.from_elements %v468, %v469, %v470, %v471 : vector<4xi32> + %o468 = index.constant 468 : index + vector.store %w468, %gv[%o468] : vector<4xi32>, view<512xi32> + %v472 = scalar.constant 218828033 : i32 + %v473 = scalar.constant 218956043 : i32 + %v474 = scalar.constant 218958081 : i32 + %v475 = scalar.constant 219087619 : i32 + %w472 = vector.from_elements %v472, %v473, %v474, %v475 : vector<4xi32> + %o472 = index.constant 472 : index + vector.store %w472, %gv[%o472] : vector<4xi32>, view<512xi32> + %v476 = scalar.constant 219087623 : i32 + %v477 = scalar.constant 251724033 : i32 + %v478 = scalar.constant 251724041 : i32 + %v479 = scalar.constant 251724047 : i32 + %w476 = vector.from_elements %v476, %v477, %v478, %v479 : vector<4xi32> + %o476 = index.constant 476 : index + vector.store %w476, %gv[%o476] : vector<4xi32>, view<512xi32> + %v480 = scalar.constant 251725057 : i32 + %v481 = scalar.constant 251725061 : i32 + %v482 = scalar.constant 251725581 : i32 + %v483 = scalar.constant 251726081 : i32 + %w480 = vector.from_elements %v480, %v481, %v482, %v483 : vector<4xi32> + %o480 = index.constant 480 : index + vector.store %w480, %gv[%o480] : vector<4xi32>, view<512xi32> + %v484 = scalar.constant 251726601 : i32 + %v485 = scalar.constant 251727109 : i32 + %v486 = scalar.constant 251855109 : i32 + %v487 = scalar.constant 251855619 : i32 + %w484 = vector.from_elements %v484, %v485, %v486, %v487 : vector<4xi32> + %o484 = index.constant 484 : index + vector.store %w484, %gv[%o484] : vector<4xi32>, view<512xi32> + %v488 = scalar.constant 251856137 : i32 + %v489 = scalar.constant 251857159 : i32 + %v490 = scalar.constant 251857163 : i32 + %v491 = scalar.constant 251986179 : i32 + %w488 = vector.from_elements %v488, %v489, %v490, %v491 : vector<4xi32> + %o488 = index.constant 488 : index + vector.store %w488, %gv[%o488] : vector<4xi32>, view<512xi32> + %v492 = scalar.constant 251986185 : i32 + %v493 = scalar.constant 251986689 : i32 + %v494 = scalar.constant 251986701 : i32 + %v495 = scalar.constant 251987203 : i32 + %w492 = vector.from_elements %v492, %v493, %v494, %v495 : vector<4xi32> + %o492 = index.constant 492 : index + vector.store %w492, %gv[%o492] : vector<4xi32>, view<512xi32> + %v496 = scalar.constant 251987713 : i32 + %v497 = scalar.constant 251988739 : i32 + %v498 = scalar.constant 252117253 : i32 + %v499 = scalar.constant 252118789 : i32 + %w496 = vector.from_elements %v496, %v497, %v498, %v499 : vector<4xi32> + %o496 = index.constant 496 : index + vector.store %w496, %gv[%o496] : vector<4xi32>, view<512xi32> + %v500 = scalar.constant 252118795 : i32 + %v501 = scalar.constant 252119815 : i32 + %v502 = scalar.constant 252248323 : i32 + %v503 = scalar.constant 252248331 : i32 + %w500 = vector.from_elements %v500, %v501, %v502, %v503 : vector<4xi32> + %o500 = index.constant 500 : index + vector.store %w500, %gv[%o500] : vector<4xi32>, view<512xi32> + %v504 = scalar.constant 252248839 : i32 + %v505 = scalar.constant 252249345 : i32 + %v506 = scalar.constant 252250881 : i32 + %v507 = scalar.constant 252380421 : i32 + %w504 = vector.from_elements %v504, %v505, %v506, %v507 : vector<4xi32> + %o504 = index.constant 504 : index + vector.store %w504, %gv[%o504] : vector<4xi32>, view<512xi32> + %v508 = scalar.constant 252381445 : i32 + %v509 = scalar.constant 252510469 : i32 + %v510 = scalar.constant 252512003 : i32 + %v511 = scalar.constant 252641537 : i32 + %w508 = vector.from_elements %v508, %v509, %v510, %v511 : vector<4xi32> + %o508 = index.constant 508 : index + vector.store %w508, %gv[%o508] : vector<4xi32>, view<512xi32> + } + func.return +} + +// IQ3_S: 110 bytes = d (f16), qs[64], qh[8], signs[32], scales[4]; 2-byte aligned. Sub-block ib (32 +// values) has scale d * (1 + 2 * nibble ib of scales), grid indices qs[8 ib + e] | bit e of qh[ib] +// << 8 (e = 0..7, four values each from iq3s_grid) and signs[4 ib ..]. Lane l16 owns entries +// 4 (l16 % 2) .. +4 of sub-block l16 / 2: 16 consecutive values. +func.def inline @ggml_kquant_iq3s_lane_dot(%grid_buf: buffer, %weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c33 = index.constant 33 : index + %c37 = index.constant 37 : index + %c53 = index.constant 53 : index + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c255_i32 = scalar.constant 255 : i32 + %spread = scalar.constant 2113665 : i32 + %lsb = scalar.constant 16843009 : i32 + %one4 = vector.constant 1.0 : vector<4xf32> + %two4 = vector.constant 2.0 : vector<4xf32> + %zero_scalar = scalar.constant 0.0 : f32 + %zero_offset = index.constant 0 : offset + %block_bytes = index.constant 110 : offset + %x_block_bytes = index.constant 1024 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %ib = index.div %l, %c2 : index + %half = index.rem %l, %c2 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %iv = buffer.view %weight[%block_base] : buffer -> view<55xi16> + %hv = buffer.view %weight[%block_base] : buffer -> view<55xf16> + %d_f16 = view.load %hv[%c0] : view<55xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %ib4 = index.mul %ib, %c4 : index + %half2 = index.mul %half, %c2 : index + %qs_rel = index.add %ib4, %half2 : index + %qs_at = index.add %c1, %qs_rel : index + %qs2 = vector.load %iv[%qs_at] : view<55xi16> -> vector<2xi16> + %qs1 = vector.bitcast %qs2 : vector<2xi16> to vector<1xi32> + %qs = vector.extract %qs1[0] : vector<1xi32> -> i32 + %ib_2 = index.div %ib, %c2 : index + %qh_at = index.add %c33, %ib_2 : index + %qh16 = view.load %iv[%qh_at] : view<55xi16> -> i16 + %qh16u = scalar.extui %qh16 : i16 to i32 + %ib_odd = index.rem %ib, %c2 : index + %ib_odd_i32 = index.cast %ib_odd : index to i32 + %qh_shift = scalar.muli %ib_odd_i32, %c8_i32 : i32 + %qh0 = scalar.shrui %qh16u, %qh_shift : i32 + %qh = scalar.andi %qh0, %c255_i32 : i32 + %ib2 = index.mul %ib, %c2 : index + %sg_rel = index.add %ib2, %half : index + %sg_at = index.add %c37, %sg_rel : index + %sg16 = view.load %iv[%sg_at] : view<55xi16> -> i16 + %sg = scalar.extui %sg16 : i16 to i32 + %ib_4 = index.div %ib, %c4 : index + %sc_at = index.add %c53, %ib_4 : index + %sc16 = view.load %iv[%sc_at] : view<55xi16> -> i16 + %sc16u = scalar.extui %sc16 : i16 to i32 + %ib_i32 = index.cast %ib : index to i32 + %sc_shift = scalar.muli %ib_i32, %c4_i32 : i32 + %sc_shift_w = scalar.andi %sc_shift, %c15_i32 : i32 + %sc0 = scalar.shrui %sc16u, %sc_shift_w : i32 + %sc = scalar.andi %sc0, %c15_i32 : i32 + %sc2 = scalar.muli %sc, %c2_i32 : i32 + %sc21 = scalar.addi %sc2, %c1_i32 : i32 + %sc_f = scalar.sitofp %sc21 : i32 to f32 + %db = scalar.mulf %d, %sc_f : f32 + %half_i32 = index.cast %half : index to i32 + %qh_base = scalar.muli %half_i32, %c4_i32 : i32 + %grid = buffer.view %grid_buf[%zero_offset] : buffer -> view<512xi32> + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %ib32 = index.mul %ib, %c32 : index + %half16 = index.mul %half, %c16 : index + %xk = index.add %ib32, %half16 : index + %x = vector.load %xv[%xk] : view<256xf32> -> vector<16xf32> + %e0_sh = scalar.constant 0 : i32 + %e0_b0 = scalar.shrui %qs, %e0_sh : i32 + %e0_b = scalar.andi %e0_b0, %c255_i32 : i32 + %e0_hs0 = scalar.constant 0 : i32 + %e0_hs = scalar.addi %qh_base, %e0_hs0 : i32 + %e0_h0 = scalar.shrui %qh, %e0_hs : i32 + %e0_h1 = scalar.andi %e0_h0, %c1_i32 : i32 + %e0_h = scalar.shli %e0_h1, %c8_i32 : i32 + %e0_i32 = scalar.ori %e0_b, %e0_h : i32 + %e0_idx0 = index.cast %e0_i32 : i32 to index + %e0_idx = index.assume %e0_idx0 [range(%e0_idx0, 0, 511)] : index + %e0_g = view.load %grid[%e0_idx] : view<512xi32> -> i32 + %e0_gv = vector.from_elements %e0_g : vector<1xi32> + %e0_gb = vector.bitcast %e0_gv : vector<1xi32> to vector<4xi8> + %e0_gf = vector.uitofp %e0_gb : vector<4xi8> to vector<4xf32> + %e0_ss = scalar.constant 0 : i32 + %e0_s0 = scalar.shrui %sg, %e0_ss : i32 + %e0_s1 = scalar.andi %e0_s0, %c15_i32 : i32 + %e0_s2 = scalar.muli %e0_s1, %spread : i32 + %e0_s3 = scalar.andi %e0_s2, %lsb : i32 + %e0_sv = vector.from_elements %e0_s3 : vector<1xi32> + %e0_sb = vector.bitcast %e0_sv : vector<1xi32> to vector<4xi8> + %e0_sf = vector.uitofp %e0_sb : vector<4xi8> to vector<4xf32> + %e0_s2f = vector.mulf %e0_sf, %two4 : vector<4xf32> + %e0_sign = vector.subf %one4, %e0_s2f : vector<4xf32> + %e0_w = vector.mulf %e0_gf, %e0_sign : vector<4xf32> + %e0_w0 = vector.extract %e0_w[0] : vector<4xf32> -> f32 + %e0_x0 = vector.extract %x[0] : vector<16xf32> -> f32 + %e0_a0 = scalar.fmaf %e0_w0, %e0_x0, %zero_scalar : f32 + %e0_w1 = vector.extract %e0_w[1] : vector<4xf32> -> f32 + %e0_x1 = vector.extract %x[1] : vector<16xf32> -> f32 + %e0_a1 = scalar.fmaf %e0_w1, %e0_x1, %e0_a0 : f32 + %e0_w2 = vector.extract %e0_w[2] : vector<4xf32> -> f32 + %e0_x2 = vector.extract %x[2] : vector<16xf32> -> f32 + %e0_a2 = scalar.fmaf %e0_w2, %e0_x2, %e0_a1 : f32 + %e0_w3 = vector.extract %e0_w[3] : vector<4xf32> -> f32 + %e0_x3 = vector.extract %x[3] : vector<16xf32> -> f32 + %e0_a3 = scalar.fmaf %e0_w3, %e0_x3, %e0_a2 : f32 + %e1_sh = scalar.constant 8 : i32 + %e1_b0 = scalar.shrui %qs, %e1_sh : i32 + %e1_b = scalar.andi %e1_b0, %c255_i32 : i32 + %e1_hs0 = scalar.constant 1 : i32 + %e1_hs = scalar.addi %qh_base, %e1_hs0 : i32 + %e1_h0 = scalar.shrui %qh, %e1_hs : i32 + %e1_h1 = scalar.andi %e1_h0, %c1_i32 : i32 + %e1_h = scalar.shli %e1_h1, %c8_i32 : i32 + %e1_i32 = scalar.ori %e1_b, %e1_h : i32 + %e1_idx0 = index.cast %e1_i32 : i32 to index + %e1_idx = index.assume %e1_idx0 [range(%e1_idx0, 0, 511)] : index + %e1_g = view.load %grid[%e1_idx] : view<512xi32> -> i32 + %e1_gv = vector.from_elements %e1_g : vector<1xi32> + %e1_gb = vector.bitcast %e1_gv : vector<1xi32> to vector<4xi8> + %e1_gf = vector.uitofp %e1_gb : vector<4xi8> to vector<4xf32> + %e1_ss = scalar.constant 4 : i32 + %e1_s0 = scalar.shrui %sg, %e1_ss : i32 + %e1_s1 = scalar.andi %e1_s0, %c15_i32 : i32 + %e1_s2 = scalar.muli %e1_s1, %spread : i32 + %e1_s3 = scalar.andi %e1_s2, %lsb : i32 + %e1_sv = vector.from_elements %e1_s3 : vector<1xi32> + %e1_sb = vector.bitcast %e1_sv : vector<1xi32> to vector<4xi8> + %e1_sf = vector.uitofp %e1_sb : vector<4xi8> to vector<4xf32> + %e1_s2f = vector.mulf %e1_sf, %two4 : vector<4xf32> + %e1_sign = vector.subf %one4, %e1_s2f : vector<4xf32> + %e1_w = vector.mulf %e1_gf, %e1_sign : vector<4xf32> + %e1_w0 = vector.extract %e1_w[0] : vector<4xf32> -> f32 + %e1_x0 = vector.extract %x[4] : vector<16xf32> -> f32 + %e1_a0 = scalar.fmaf %e1_w0, %e1_x0, %e0_a3 : f32 + %e1_w1 = vector.extract %e1_w[1] : vector<4xf32> -> f32 + %e1_x1 = vector.extract %x[5] : vector<16xf32> -> f32 + %e1_a1 = scalar.fmaf %e1_w1, %e1_x1, %e1_a0 : f32 + %e1_w2 = vector.extract %e1_w[2] : vector<4xf32> -> f32 + %e1_x2 = vector.extract %x[6] : vector<16xf32> -> f32 + %e1_a2 = scalar.fmaf %e1_w2, %e1_x2, %e1_a1 : f32 + %e1_w3 = vector.extract %e1_w[3] : vector<4xf32> -> f32 + %e1_x3 = vector.extract %x[7] : vector<16xf32> -> f32 + %e1_a3 = scalar.fmaf %e1_w3, %e1_x3, %e1_a2 : f32 + %e2_sh = scalar.constant 16 : i32 + %e2_b0 = scalar.shrui %qs, %e2_sh : i32 + %e2_b = scalar.andi %e2_b0, %c255_i32 : i32 + %e2_hs0 = scalar.constant 2 : i32 + %e2_hs = scalar.addi %qh_base, %e2_hs0 : i32 + %e2_h0 = scalar.shrui %qh, %e2_hs : i32 + %e2_h1 = scalar.andi %e2_h0, %c1_i32 : i32 + %e2_h = scalar.shli %e2_h1, %c8_i32 : i32 + %e2_i32 = scalar.ori %e2_b, %e2_h : i32 + %e2_idx0 = index.cast %e2_i32 : i32 to index + %e2_idx = index.assume %e2_idx0 [range(%e2_idx0, 0, 511)] : index + %e2_g = view.load %grid[%e2_idx] : view<512xi32> -> i32 + %e2_gv = vector.from_elements %e2_g : vector<1xi32> + %e2_gb = vector.bitcast %e2_gv : vector<1xi32> to vector<4xi8> + %e2_gf = vector.uitofp %e2_gb : vector<4xi8> to vector<4xf32> + %e2_ss = scalar.constant 8 : i32 + %e2_s0 = scalar.shrui %sg, %e2_ss : i32 + %e2_s1 = scalar.andi %e2_s0, %c15_i32 : i32 + %e2_s2 = scalar.muli %e2_s1, %spread : i32 + %e2_s3 = scalar.andi %e2_s2, %lsb : i32 + %e2_sv = vector.from_elements %e2_s3 : vector<1xi32> + %e2_sb = vector.bitcast %e2_sv : vector<1xi32> to vector<4xi8> + %e2_sf = vector.uitofp %e2_sb : vector<4xi8> to vector<4xf32> + %e2_s2f = vector.mulf %e2_sf, %two4 : vector<4xf32> + %e2_sign = vector.subf %one4, %e2_s2f : vector<4xf32> + %e2_w = vector.mulf %e2_gf, %e2_sign : vector<4xf32> + %e2_w0 = vector.extract %e2_w[0] : vector<4xf32> -> f32 + %e2_x0 = vector.extract %x[8] : vector<16xf32> -> f32 + %e2_a0 = scalar.fmaf %e2_w0, %e2_x0, %e1_a3 : f32 + %e2_w1 = vector.extract %e2_w[1] : vector<4xf32> -> f32 + %e2_x1 = vector.extract %x[9] : vector<16xf32> -> f32 + %e2_a1 = scalar.fmaf %e2_w1, %e2_x1, %e2_a0 : f32 + %e2_w2 = vector.extract %e2_w[2] : vector<4xf32> -> f32 + %e2_x2 = vector.extract %x[10] : vector<16xf32> -> f32 + %e2_a2 = scalar.fmaf %e2_w2, %e2_x2, %e2_a1 : f32 + %e2_w3 = vector.extract %e2_w[3] : vector<4xf32> -> f32 + %e2_x3 = vector.extract %x[11] : vector<16xf32> -> f32 + %e2_a3 = scalar.fmaf %e2_w3, %e2_x3, %e2_a2 : f32 + %e3_sh = scalar.constant 24 : i32 + %e3_b0 = scalar.shrui %qs, %e3_sh : i32 + %e3_b = scalar.andi %e3_b0, %c255_i32 : i32 + %e3_hs0 = scalar.constant 3 : i32 + %e3_hs = scalar.addi %qh_base, %e3_hs0 : i32 + %e3_h0 = scalar.shrui %qh, %e3_hs : i32 + %e3_h1 = scalar.andi %e3_h0, %c1_i32 : i32 + %e3_h = scalar.shli %e3_h1, %c8_i32 : i32 + %e3_i32 = scalar.ori %e3_b, %e3_h : i32 + %e3_idx0 = index.cast %e3_i32 : i32 to index + %e3_idx = index.assume %e3_idx0 [range(%e3_idx0, 0, 511)] : index + %e3_g = view.load %grid[%e3_idx] : view<512xi32> -> i32 + %e3_gv = vector.from_elements %e3_g : vector<1xi32> + %e3_gb = vector.bitcast %e3_gv : vector<1xi32> to vector<4xi8> + %e3_gf = vector.uitofp %e3_gb : vector<4xi8> to vector<4xf32> + %e3_ss = scalar.constant 12 : i32 + %e3_s0 = scalar.shrui %sg, %e3_ss : i32 + %e3_s1 = scalar.andi %e3_s0, %c15_i32 : i32 + %e3_s2 = scalar.muli %e3_s1, %spread : i32 + %e3_s3 = scalar.andi %e3_s2, %lsb : i32 + %e3_sv = vector.from_elements %e3_s3 : vector<1xi32> + %e3_sb = vector.bitcast %e3_sv : vector<1xi32> to vector<4xi8> + %e3_sf = vector.uitofp %e3_sb : vector<4xi8> to vector<4xf32> + %e3_s2f = vector.mulf %e3_sf, %two4 : vector<4xf32> + %e3_sign = vector.subf %one4, %e3_s2f : vector<4xf32> + %e3_w = vector.mulf %e3_gf, %e3_sign : vector<4xf32> + %e3_w0 = vector.extract %e3_w[0] : vector<4xf32> -> f32 + %e3_x0 = vector.extract %x[12] : vector<16xf32> -> f32 + %e3_a0 = scalar.fmaf %e3_w0, %e3_x0, %e2_a3 : f32 + %e3_w1 = vector.extract %e3_w[1] : vector<4xf32> -> f32 + %e3_x1 = vector.extract %x[13] : vector<16xf32> -> f32 + %e3_a1 = scalar.fmaf %e3_w1, %e3_x1, %e3_a0 : f32 + %e3_w2 = vector.extract %e3_w[2] : vector<4xf32> -> f32 + %e3_x2 = vector.extract %x[14] : vector<16xf32> -> f32 + %e3_a2 = scalar.fmaf %e3_w2, %e3_x2, %e3_a1 : f32 + %e3_w3 = vector.extract %e3_w[3] : vector<4xf32> -> f32 + %e3_x3 = vector.extract %x[15] : vector<16xf32> -> f32 + %e3_a3 = scalar.fmaf %e3_w3, %e3_x3, %e3_a2 : f32 + %result = scalar.mulf %db, %e3_a3 : f32 + func.return %result : f32 +} + +func.def inline @ggml_kquant_iq2xxs_grid_fill(%grid: buffer, %chunk: index) { + %zero_offset = index.constant 0 : offset + %gv = buffer.view %grid[%zero_offset] : buffer -> view<512xi32> + %k0 = index.constant 0 : index + %is0 = index.cmp eq, %chunk, %k0 : index + scf.if %is0 { + %v0_0_0 = scalar.constant 0 : i32 + %v0_0_1 = scalar.constant 2 : i32 + %v0_0_2 = scalar.constant 5 : i32 + %v0_0_3 = scalar.constant 8 : i32 + %w0_0 = vector.from_elements %v0_0_0, %v0_0_1, %v0_0_2, %v0_0_3 : vector<4xi32> + %o0_0 = index.constant 0 : index + vector.store %w0_0, %gv[%o0_0] : vector<4xi32>, view<512xi32> + %v0_1_0 = scalar.constant 10 : i32 + %v0_1_1 = scalar.constant 17 : i32 + %v0_1_2 = scalar.constant 20 : i32 + %v0_1_3 = scalar.constant 32 : i32 + %w0_1 = vector.from_elements %v0_1_0, %v0_1_1, %v0_1_2, %v0_1_3 : vector<4xi32> + %o0_1 = index.constant 4 : index + vector.store %w0_1, %gv[%o0_1] : vector<4xi32>, view<512xi32> + %v0_2_0 = scalar.constant 34 : i32 + %v0_2_1 = scalar.constant 40 : i32 + %v0_2_2 = scalar.constant 42 : i32 + %v0_2_3 = scalar.constant 65 : i32 + %w0_2 = vector.from_elements %v0_2_0, %v0_2_1, %v0_2_2, %v0_2_3 : vector<4xi32> + %o0_2 = index.constant 8 : index + vector.store %w0_2, %gv[%o0_2] : vector<4xi32>, view<512xi32> + %v0_3_0 = scalar.constant 68 : i32 + %v0_3_1 = scalar.constant 80 : i32 + %v0_3_2 = scalar.constant 88 : i32 + %v0_3_3 = scalar.constant 97 : i32 + %w0_3 = vector.from_elements %v0_3_0, %v0_3_1, %v0_3_2, %v0_3_3 : vector<4xi32> + %o0_3 = index.constant 12 : index + vector.store %w0_3, %gv[%o0_3] : vector<4xi32>, view<512xi32> + %v0_4_0 = scalar.constant 100 : i32 + %v0_4_1 = scalar.constant 128 : i32 + %v0_4_2 = scalar.constant 130 : i32 + %v0_4_3 = scalar.constant 138 : i32 + %w0_4 = vector.from_elements %v0_4_0, %v0_4_1, %v0_4_2, %v0_4_3 : vector<4xi32> + %o0_4 = index.constant 16 : index + vector.store %w0_4, %gv[%o0_4] : vector<4xi32>, view<512xi32> + %v0_5_0 = scalar.constant 162 : i32 + %v0_5_1 = scalar.constant 257 : i32 + %v0_5_2 = scalar.constant 260 : i32 + %v0_5_3 = scalar.constant 272 : i32 + %w0_5 = vector.from_elements %v0_5_0, %v0_5_1, %v0_5_2, %v0_5_3 : vector<4xi32> + %o0_5 = index.constant 20 : index + vector.store %w0_5, %gv[%o0_5] : vector<4xi32>, view<512xi32> + %v0_6_0 = scalar.constant 277 : i32 + %v0_6_1 = scalar.constant 320 : i32 + %v0_6_2 = scalar.constant 388 : i32 + %v0_6_3 = scalar.constant 408 : i32 + %w0_6 = vector.from_elements %v0_6_0, %v0_6_1, %v0_6_2, %v0_6_3 : vector<4xi32> + %o0_6 = index.constant 24 : index + vector.store %w0_6, %gv[%o0_6] : vector<4xi32>, view<512xi32> + %v0_7_0 = scalar.constant 512 : i32 + %v0_7_1 = scalar.constant 514 : i32 + %v0_7_2 = scalar.constant 546 : i32 + %v0_7_3 = scalar.constant 642 : i32 + %w0_7 = vector.from_elements %v0_7_0, %v0_7_1, %v0_7_2, %v0_7_3 : vector<4xi32> + %o0_7 = index.constant 28 : index + vector.store %w0_7, %gv[%o0_7] : vector<4xi32>, view<512xi32> + %v0_8_0 = scalar.constant 1025 : i32 + %v0_8_1 = scalar.constant 1028 : i32 + %v0_8_2 = scalar.constant 1040 : i32 + %v0_8_3 = scalar.constant 1057 : i32 + %w0_8 = vector.from_elements %v0_8_0, %v0_8_1, %v0_8_2, %v0_8_3 : vector<4xi32> + %o0_8 = index.constant 32 : index + vector.store %w0_8, %gv[%o0_8] : vector<4xi32>, view<512xi32> + %v0_9_0 = scalar.constant 1060 : i32 + %v0_9_1 = scalar.constant 1088 : i32 + %v0_9_2 = scalar.constant 1090 : i32 + %v0_9_3 = scalar.constant 1096 : i32 + %w0_9 = vector.from_elements %v0_9_0, %v0_9_1, %v0_9_2, %v0_9_3 : vector<4xi32> + %o0_9 = index.constant 36 : index + vector.store %w0_9, %gv[%o0_9] : vector<4xi32>, view<512xi32> + %v0_10_0 = scalar.constant 1120 : i32 + %v0_10_1 = scalar.constant 1153 : i32 + %v0_10_2 = scalar.constant 1156 : i32 + %v0_10_3 = scalar.constant 1168 : i32 + %w0_10 = vector.from_elements %v0_10_0, %v0_10_1, %v0_10_2, %v0_10_3 : vector<4xi32> + %o0_10 = index.constant 40 : index + vector.store %w0_10, %gv[%o0_10] : vector<4xi32>, view<512xi32> + %v0_11_0 = scalar.constant 1188 : i32 + %v0_11_1 = scalar.constant 1280 : i32 + %v0_11_2 = scalar.constant 1282 : i32 + %v0_11_3 = scalar.constant 1288 : i32 + %w0_11 = vector.from_elements %v0_11_0, %v0_11_1, %v0_11_2, %v0_11_3 : vector<4xi32> + %o0_11 = index.constant 44 : index + vector.store %w0_11, %gv[%o0_11] : vector<4xi32>, view<512xi32> + %v0_12_0 = scalar.constant 1312 : i32 + %v0_12_1 = scalar.constant 1350 : i32 + %v0_12_2 = scalar.constant 1385 : i32 + %v0_12_3 = scalar.constant 1408 : i32 + %w0_12 = vector.from_elements %v0_12_0, %v0_12_1, %v0_12_2, %v0_12_3 : vector<4xi32> + %o0_12 = index.constant 48 : index + vector.store %w0_12, %gv[%o0_12] : vector<4xi32>, view<512xi32> + %v0_13_0 = scalar.constant 1425 : i32 + %v0_13_1 = scalar.constant 1545 : i32 + %v0_13_2 = scalar.constant 1552 : i32 + %v0_13_3 = scalar.constant 1600 : i32 + %w0_13 = vector.from_elements %v0_13_0, %v0_13_1, %v0_13_2, %v0_13_3 : vector<4xi32> + %o0_13 = index.constant 52 : index + vector.store %w0_13, %gv[%o0_13] : vector<4xi32>, view<512xi32> + %v0_14_0 = scalar.constant 1668 : i32 + %v0_14_1 = scalar.constant 1700 : i32 + %v0_14_2 = scalar.constant 2048 : i32 + %v0_14_3 = scalar.constant 2053 : i32 + %w0_14 = vector.from_elements %v0_14_0, %v0_14_1, %v0_14_2, %v0_14_3 : vector<4xi32> + %o0_14 = index.constant 56 : index + vector.store %w0_14, %gv[%o0_14] : vector<4xi32>, view<512xi32> + %v0_15_0 = scalar.constant 2056 : i32 + %v0_15_1 = scalar.constant 2068 : i32 + %v0_15_2 = scalar.constant 2088 : i32 + %v0_15_3 = scalar.constant 2113 : i32 + %w0_15 = vector.from_elements %v0_15_0, %v0_15_1, %v0_15_2, %v0_15_3 : vector<4xi32> + %o0_15 = index.constant 60 : index + vector.store %w0_15, %gv[%o0_15] : vector<4xi32>, view<512xi32> + %v0_16_0 = scalar.constant 2116 : i32 + %v0_16_1 = scalar.constant 2128 : i32 + %v0_16_2 = scalar.constant 2130 : i32 + %v0_16_3 = scalar.constant 2184 : i32 + %w0_16 = vector.from_elements %v0_16_0, %v0_16_1, %v0_16_2, %v0_16_3 : vector<4xi32> + %o0_16 = index.constant 64 : index + vector.store %w0_16, %gv[%o0_16] : vector<4xi32>, view<512xi32> + %v0_17_0 = scalar.constant 2308 : i32 + %v0_17_1 = scalar.constant 2368 : i32 + %v0_17_2 = scalar.constant 2562 : i32 + %v0_17_3 = scalar.constant 2580 : i32 + %w0_17 = vector.from_elements %v0_17_0, %v0_17_1, %v0_17_2, %v0_17_3 : vector<4xi32> + %o0_17 = index.constant 68 : index + vector.store %w0_17, %gv[%o0_17] : vector<4xi32>, view<512xi32> + %v0_18_0 = scalar.constant 4097 : i32 + %v0_18_1 = scalar.constant 4100 : i32 + %v0_18_2 = scalar.constant 4112 : i32 + %v0_18_3 = scalar.constant 4129 : i32 + %w0_18 = vector.from_elements %v0_18_0, %v0_18_1, %v0_18_2, %v0_18_3 : vector<4xi32> + %o0_18 = index.constant 72 : index + vector.store %w0_18, %gv[%o0_18] : vector<4xi32>, view<512xi32> + %v0_19_0 = scalar.constant 4160 : i32 + %v0_19_1 = scalar.constant 4192 : i32 + %v0_19_2 = scalar.constant 4228 : i32 + %v0_19_3 = scalar.constant 4240 : i32 + %w0_19 = vector.from_elements %v0_19_0, %v0_19_1, %v0_19_2, %v0_19_3 : vector<4xi32> + %o0_19 = index.constant 76 : index + vector.store %w0_19, %gv[%o0_19] : vector<4xi32>, view<512xi32> + %v0_20_0 = scalar.constant 4245 : i32 + %v0_20_1 = scalar.constant 4352 : i32 + %v0_20_2 = scalar.constant 4360 : i32 + %v0_20_3 = scalar.constant 4384 : i32 + %w0_20 = vector.from_elements %v0_20_0, %v0_20_1, %v0_20_2, %v0_20_3 : vector<4xi32> + %o0_20 = index.constant 80 : index + vector.store %w0_20, %gv[%o0_20] : vector<4xi32>, view<512xi32> + %v0_21_0 = scalar.constant 4432 : i32 + %v0_21_1 = scalar.constant 4442 : i32 + %v0_21_2 = scalar.constant 4480 : i32 + %v0_21_3 = scalar.constant 4644 : i32 + %w0_21 = vector.from_elements %v0_21_0, %v0_21_1, %v0_21_2, %v0_21_3 : vector<4xi32> + %o0_21 = index.constant 84 : index + vector.store %w0_21, %gv[%o0_21] : vector<4xi32>, view<512xi32> + %v0_22_0 = scalar.constant 4677 : i32 + %v0_22_1 = scalar.constant 5120 : i32 + %v0_22_2 = scalar.constant 5128 : i32 + %v0_22_3 = scalar.constant 5152 : i32 + %w0_22 = vector.from_elements %v0_22_0, %v0_22_1, %v0_22_2, %v0_22_3 : vector<4xi32> + %o0_22 = index.constant 88 : index + vector.store %w0_22, %gv[%o0_22] : vector<4xi32>, view<512xi32> + %v0_23_0 = scalar.constant 5157 : i32 + %v0_23_1 = scalar.constant 5193 : i32 + %v0_23_2 = scalar.constant 5248 : i32 + %v0_23_3 = scalar.constant 5400 : i32 + %w0_23 = vector.from_elements %v0_23_0, %v0_23_1, %v0_23_2, %v0_23_3 : vector<4xi32> + %o0_23 = index.constant 92 : index + vector.store %w0_23, %gv[%o0_23] : vector<4xi32>, view<512xi32> + %v0_24_0 = scalar.constant 5474 : i32 + %v0_24_1 = scalar.constant 5632 : i32 + %v0_24_2 = scalar.constant 5654 : i32 + %v0_24_3 = scalar.constant 6145 : i32 + %w0_24 = vector.from_elements %v0_24_0, %v0_24_1, %v0_24_2, %v0_24_3 : vector<4xi32> + %o0_24 = index.constant 96 : index + vector.store %w0_24, %gv[%o0_24] : vector<4xi32>, view<512xi32> + %v0_25_0 = scalar.constant 6148 : i32 + %v0_25_1 = scalar.constant 6160 : i32 + %v0_25_2 = scalar.constant 6208 : i32 + %v0_25_3 = scalar.constant 6273 : i32 + %w0_25 = vector.from_elements %v0_25_0, %v0_25_1, %v0_25_2, %v0_25_3 : vector<4xi32> + %o0_25 = index.constant 100 : index + vector.store %w0_25, %gv[%o0_25] : vector<4xi32>, view<512xi32> + %v0_26_0 = scalar.constant 6400 : i32 + %v0_26_1 = scalar.constant 6405 : i32 + %v0_26_2 = scalar.constant 6560 : i32 + %v0_26_3 = scalar.constant 6737 : i32 + %w0_26 = vector.from_elements %v0_26_0, %v0_26_1, %v0_26_2, %v0_26_3 : vector<4xi32> + %o0_26 = index.constant 104 : index + vector.store %w0_26, %gv[%o0_26] : vector<4xi32>, view<512xi32> + %v0_27_0 = scalar.constant 8192 : i32 + %v0_27_1 = scalar.constant 8194 : i32 + %v0_27_2 = scalar.constant 8202 : i32 + %v0_27_3 = scalar.constant 8260 : i32 + %w0_27 = vector.from_elements %v0_27_0, %v0_27_1, %v0_27_2, %v0_27_3 : vector<4xi32> + %o0_27 = index.constant 108 : index + vector.store %w0_27, %gv[%o0_27] : vector<4xi32>, view<512xi32> + %v0_28_0 = scalar.constant 8289 : i32 + %v0_28_1 = scalar.constant 8320 : i32 + %v0_28_2 = scalar.constant 8322 : i32 + %v0_28_3 = scalar.constant 8489 : i32 + %w0_28 = vector.from_elements %v0_28_0, %v0_28_1, %v0_28_2, %v0_28_3 : vector<4xi32> + %o0_28 = index.constant 112 : index + vector.store %w0_28, %gv[%o0_28] : vector<4xi32>, view<512xi32> + %v0_29_0 = scalar.constant 8520 : i32 + %v0_29_1 = scalar.constant 8704 : i32 + %v0_29_2 = scalar.constant 8706 : i32 + %v0_29_3 = scalar.constant 9217 : i32 + %w0_29 = vector.from_elements %v0_29_0, %v0_29_1, %v0_29_2, %v0_29_3 : vector<4xi32> + %o0_29 = index.constant 116 : index + vector.store %w0_29, %gv[%o0_29] : vector<4xi32>, view<512xi32> + %v0_30_0 = scalar.constant 9220 : i32 + %v0_30_1 = scalar.constant 9232 : i32 + %v0_30_2 = scalar.constant 9280 : i32 + %v0_30_3 = scalar.constant 9302 : i32 + %w0_30 = vector.from_elements %v0_30_0, %v0_30_1, %v0_30_2, %v0_30_3 : vector<4xi32> + %o0_30 = index.constant 120 : index + vector.store %w0_30, %gv[%o0_30] : vector<4xi32>, view<512xi32> + %v0_31_0 = scalar.constant 9472 : i32 + %v0_31_1 = scalar.constant 9537 : i32 + %v0_31_2 = scalar.constant 9572 : i32 + %v0_31_3 = scalar.constant 9872 : i32 + %w0_31 = vector.from_elements %v0_31_0, %v0_31_1, %v0_31_2, %v0_31_3 : vector<4xi32> + %o0_31 = index.constant 124 : index + vector.store %w0_31, %gv[%o0_31] : vector<4xi32>, view<512xi32> + } + %k1 = index.constant 1 : index + %is1 = index.cmp eq, %chunk, %k1 : index + scf.if %is1 { + %v1_0_0 = scalar.constant 10248 : i32 + %v1_0_1 = scalar.constant 10272 : i32 + %v1_0_2 = scalar.constant 10388 : i32 + %v1_0_3 = scalar.constant 10820 : i32 + %w1_0 = vector.from_elements %v1_0_0, %v1_0_1, %v1_0_2, %v1_0_3 : vector<4xi32> + %o1_0 = index.constant 128 : index + vector.store %w1_0, %gv[%o1_0] : vector<4xi32>, view<512xi32> + %v1_1_0 = scalar.constant 16385 : i32 + %v1_1_1 = scalar.constant 16388 : i32 + %v1_1_2 = scalar.constant 16400 : i32 + %v1_1_3 = scalar.constant 16408 : i32 + %w1_1 = vector.from_elements %v1_1_0, %v1_1_1, %v1_1_2, %v1_1_3 : vector<4xi32> + %o1_1 = index.constant 132 : index + vector.store %w1_1, %gv[%o1_1] : vector<4xi32>, view<512xi32> + %v1_2_0 = scalar.constant 16417 : i32 + %v1_2_1 = scalar.constant 16420 : i32 + %v1_2_2 = scalar.constant 16448 : i32 + %v1_2_3 = scalar.constant 16456 : i32 + %w1_2 = vector.from_elements %v1_2_0, %v1_2_1, %v1_2_2, %v1_2_3 : vector<4xi32> + %o1_2 = index.constant 136 : index + vector.store %w1_2, %gv[%o1_2] : vector<4xi32>, view<512xi32> + %v1_3_0 = scalar.constant 16470 : i32 + %v1_3_1 = scalar.constant 16480 : i32 + %v1_3_2 = scalar.constant 16513 : i32 + %v1_3_3 = scalar.constant 16516 : i32 + %w1_3 = vector.from_elements %v1_3_0, %v1_3_1, %v1_3_2, %v1_3_3 : vector<4xi32> + %o1_3 = index.constant 140 : index + vector.store %w1_3, %gv[%o1_3] : vector<4xi32>, view<512xi32> + %v1_4_0 = scalar.constant 16528 : i32 + %v1_4_1 = scalar.constant 16640 : i32 + %v1_4_2 = scalar.constant 16672 : i32 + %v1_4_3 = scalar.constant 16737 : i32 + %w1_4 = vector.from_elements %v1_4_0, %v1_4_1, %v1_4_2, %v1_4_3 : vector<4xi32> + %o1_4 = index.constant 144 : index + vector.store %w1_4, %gv[%o1_4] : vector<4xi32>, view<512xi32> + %v1_5_0 = scalar.constant 16768 : i32 + %v1_5_1 = scalar.constant 16773 : i32 + %v1_5_2 = scalar.constant 16897 : i32 + %v1_5_3 = scalar.constant 16912 : i32 + %w1_5 = vector.from_elements %v1_5_0, %v1_5_1, %v1_5_2, %v1_5_3 : vector<4xi32> + %o1_5 = index.constant 148 : index + vector.store %w1_5, %gv[%o1_5] : vector<4xi32>, view<512xi32> + %v1_6_0 = scalar.constant 16968 : i32 + %v1_6_1 = scalar.constant 16982 : i32 + %v1_6_2 = scalar.constant 17000 : i32 + %v1_6_3 = scalar.constant 17408 : i32 + %w1_6 = vector.from_elements %v1_6_0, %v1_6_1, %v1_6_2, %v1_6_3 : vector<4xi32> + %o1_6 = index.constant 152 : index + vector.store %w1_6, %gv[%o1_6] : vector<4xi32>, view<512xi32> + %v1_7_0 = scalar.constant 17416 : i32 + %v1_7_1 = scalar.constant 17440 : i32 + %v1_7_2 = scalar.constant 17536 : i32 + %v1_7_3 = scalar.constant 17561 : i32 + %w1_7 = vector.from_elements %v1_7_0, %v1_7_1, %v1_7_2, %v1_7_3 : vector<4xi32> + %o1_7 = index.constant 156 : index + vector.store %w1_7, %gv[%o1_7] : vector<4xi32>, view<512xi32> + %v1_8_0 = scalar.constant 17682 : i32 + %v1_8_1 = scalar.constant 17700 : i32 + %v1_8_2 = scalar.constant 17920 : i32 + %v1_8_3 = scalar.constant 18433 : i32 + %w1_8 = vector.from_elements %v1_8_0, %v1_8_1, %v1_8_2, %v1_8_3 : vector<4xi32> + %o1_8 = index.constant 160 : index + vector.store %w1_8, %gv[%o1_8] : vector<4xi32>, view<512xi32> + %v1_9_0 = scalar.constant 18436 : i32 + %v1_9_1 = scalar.constant 18448 : i32 + %v1_9_2 = scalar.constant 18496 : i32 + %v1_9_3 = scalar.constant 18501 : i32 + %w1_9 = vector.from_elements %v1_9_0, %v1_9_1, %v1_9_2, %v1_9_3 : vector<4xi32> + %o1_9 = index.constant 164 : index + vector.store %w1_9, %gv[%o1_9] : vector<4xi32>, view<512xi32> + %v1_10_0 = scalar.constant 18688 : i32 + %v1_10_1 = scalar.constant 18776 : i32 + %v1_10_2 = scalar.constant 18785 : i32 + %v1_10_3 = scalar.constant 18818 : i32 + %w1_10 = vector.from_elements %v1_10_0, %v1_10_1, %v1_10_2, %v1_10_3 : vector<4xi32> + %o1_10 = index.constant 168 : index + vector.store %w1_10, %gv[%o1_10] : vector<4xi32>, view<512xi32> + %v1_11_0 = scalar.constant 19013 : i32 + %v1_11_1 = scalar.constant 19088 : i32 + %v1_11_2 = scalar.constant 20480 : i32 + %v1_11_3 = scalar.constant 20488 : i32 + %w1_11 = vector.from_elements %v1_11_0, %v1_11_1, %v1_11_2, %v1_11_3 : vector<4xi32> + %o1_11 = index.constant 172 : index + vector.store %w1_11, %gv[%o1_11] : vector<4xi32>, view<512xi32> + %v1_12_0 = scalar.constant 20497 : i32 + %v1_12_1 = scalar.constant 20505 : i32 + %v1_12_2 = scalar.constant 20512 : i32 + %v1_12_3 = scalar.constant 20608 : i32 + %w1_12 = vector.from_elements %v1_12_0, %v1_12_1, %v1_12_2, %v1_12_3 : vector<4xi32> + %o1_12 = index.constant 176 : index + vector.store %w1_12, %gv[%o1_12] : vector<4xi32>, view<512xi32> + %v1_13_0 = scalar.constant 20616 : i32 + %v1_13_1 = scalar.constant 20740 : i32 + %v1_13_2 = scalar.constant 20802 : i32 + %v1_13_3 = scalar.constant 20900 : i32 + %w1_13 = vector.from_elements %v1_13_0, %v1_13_1, %v1_13_2, %v1_13_3 : vector<4xi32> + %o1_13 = index.constant 180 : index + vector.store %w1_13, %gv[%o1_13] : vector<4xi32>, view<512xi32> + %v1_14_0 = scalar.constant 21137 : i32 + %v1_14_1 = scalar.constant 21648 : i32 + %v1_14_2 = scalar.constant 21650 : i32 + %v1_14_3 = scalar.constant 21770 : i32 + %w1_14 = vector.from_elements %v1_14_0, %v1_14_1, %v1_14_2, %v1_14_3 : vector<4xi32> + %o1_14 = index.constant 184 : index + vector.store %w1_14, %gv[%o1_14] : vector<4xi32>, view<512xi32> + %v1_15_0 = scalar.constant 22017 : i32 + %v1_15_1 = scalar.constant 22100 : i32 + %v1_15_2 = scalar.constant 22528 : i32 + %v1_15_3 = scalar.constant 22545 : i32 + %w1_15 = vector.from_elements %v1_15_0, %v1_15_1, %v1_15_2, %v1_15_3 : vector<4xi32> + %o1_15 = index.constant 188 : index + vector.store %w1_15, %gv[%o1_15] : vector<4xi32>, view<512xi32> + %v1_16_0 = scalar.constant 22553 : i32 + %v1_16_1 = scalar.constant 22628 : i32 + %v1_16_2 = scalar.constant 22848 : i32 + %v1_16_3 = scalar.constant 23048 : i32 + %w1_16 = vector.from_elements %v1_16_0, %v1_16_1, %v1_16_2, %v1_16_3 : vector<4xi32> + %o1_16 = index.constant 192 : index + vector.store %w1_16, %gv[%o1_16] : vector<4xi32>, view<512xi32> + %v1_17_0 = scalar.constant 24580 : i32 + %v1_17_1 = scalar.constant 24592 : i32 + %v1_17_2 = scalar.constant 24640 : i32 + %v1_17_3 = scalar.constant 24680 : i32 + %w1_17 = vector.from_elements %v1_17_0, %v1_17_1, %v1_17_2, %v1_17_3 : vector<4xi32> + %o1_17 = index.constant 196 : index + vector.store %w1_17, %gv[%o1_17] : vector<4xi32>, view<512xi32> + %v1_18_0 = scalar.constant 24832 : i32 + %v1_18_1 = scalar.constant 24917 : i32 + %v1_18_2 = scalar.constant 25112 : i32 + %v1_18_3 = scalar.constant 25184 : i32 + %w1_18 = vector.from_elements %v1_18_0, %v1_18_1, %v1_18_2, %v1_18_3 : vector<4xi32> + %o1_18 = index.constant 200 : index + vector.store %w1_18, %gv[%o1_18] : vector<4xi32>, view<512xi32> + %v1_19_0 = scalar.constant 25600 : i32 + %v1_19_1 = scalar.constant 25605 : i32 + %v1_19_2 = scalar.constant 25872 : i32 + %v1_19_3 = scalar.constant 25874 : i32 + %w1_19 = vector.from_elements %v1_19_0, %v1_19_1, %v1_19_2, %v1_19_3 : vector<4xi32> + %o1_19 = index.constant 204 : index + vector.store %w1_19, %gv[%o1_19] : vector<4xi32>, view<512xi32> + %v1_20_0 = scalar.constant 25988 : i32 + %v1_20_1 = scalar.constant 26690 : i32 + %v1_20_2 = scalar.constant 32768 : i32 + %v1_20_3 = scalar.constant 32770 : i32 + %w1_20 = vector.from_elements %v1_20_0, %v1_20_1, %v1_20_2, %v1_20_3 : vector<4xi32> + %o1_20 = index.constant 208 : index + vector.store %w1_20, %gv[%o1_20] : vector<4xi32>, view<512xi32> + %v1_21_0 = scalar.constant 32778 : i32 + %v1_21_1 = scalar.constant 32833 : i32 + %v1_21_2 = scalar.constant 32898 : i32 + %v1_21_3 = scalar.constant 33028 : i32 + %w1_21 = vector.from_elements %v1_21_0, %v1_21_1, %v1_21_2, %v1_21_3 : vector<4xi32> + %o1_21 = index.constant 212 : index + vector.store %w1_21, %gv[%o1_21] : vector<4xi32>, view<512xi32> + %v1_22_0 = scalar.constant 33048 : i32 + %v1_22_1 = scalar.constant 33088 : i32 + %v1_22_2 = scalar.constant 33297 : i32 + %v1_22_3 = scalar.constant 33793 : i32 + %w1_22 = vector.from_elements %v1_22_0, %v1_22_1, %v1_22_2, %v1_22_3 : vector<4xi32> + %o1_22 = index.constant 216 : index + vector.store %w1_22, %gv[%o1_22] : vector<4xi32>, view<512xi32> + %v1_23_0 = scalar.constant 33796 : i32 + %v1_23_1 = scalar.constant 33808 : i32 + %v1_23_2 = scalar.constant 33813 : i32 + %v1_23_3 = scalar.constant 33856 : i32 + %w1_23 = vector.from_elements %v1_23_0, %v1_23_1, %v1_23_2, %v1_23_3 : vector<4xi32> + %o1_23 = index.constant 220 : index + vector.store %w1_23, %gv[%o1_23] : vector<4xi32>, view<512xi32> + %v1_24_0 = scalar.constant 33888 : i32 + %v1_24_1 = scalar.constant 34048 : i32 + %v1_24_2 = scalar.constant 34118 : i32 + %v1_24_3 = scalar.constant 34196 : i32 + %w1_24 = vector.from_elements %v1_24_0, %v1_24_1, %v1_24_2, %v1_24_3 : vector<4xi32> + %o1_24 = index.constant 224 : index + vector.store %w1_24, %gv[%o1_24] : vector<4xi32>, view<512xi32> + %v1_25_0 = scalar.constant 34313 : i32 + %v1_25_1 = scalar.constant 34368 : i32 + %v1_25_2 = scalar.constant 34400 : i32 + %v1_25_3 = scalar.constant 34818 : i32 + %w1_25 = vector.from_elements %v1_25_0, %v1_25_1, %v1_25_2, %v1_25_3 : vector<4xi32> + %o1_25 = index.constant 228 : index + vector.store %w1_25, %gv[%o1_25] : vector<4xi32>, view<512xi32> + %v1_26_0 = scalar.constant 35076 : i32 + %v1_26_1 = scalar.constant 35345 : i32 + %v1_26_2 = scalar.constant 36868 : i32 + %v1_26_3 = scalar.constant 36880 : i32 + %w1_26 = vector.from_elements %v1_26_0, %v1_26_1, %v1_26_2, %v1_26_3 : vector<4xi32> + %o1_26 = index.constant 232 : index + vector.store %w1_26, %gv[%o1_26] : vector<4xi32>, view<512xi32> + %v1_27_0 = scalar.constant 36900 : i32 + %v1_27_1 = scalar.constant 36928 : i32 + %v1_27_2 = scalar.constant 37025 : i32 + %v1_27_3 = scalar.constant 37142 : i32 + %w1_27 = vector.from_elements %v1_27_0, %v1_27_1, %v1_27_2, %v1_27_3 : vector<4xi32> + %o1_27 = index.constant 236 : index + vector.store %w1_27, %gv[%o1_27] : vector<4xi32>, view<512xi32> + %v1_28_0 = scalar.constant 37248 : i32 + %v1_28_1 = scalar.constant 37445 : i32 + %v1_28_2 = scalar.constant 37888 : i32 + %v1_28_3 = scalar.constant 37922 : i32 + %w1_28 = vector.from_elements %v1_28_0, %v1_28_1, %v1_28_2, %v1_28_3 : vector<4xi32> + %o1_28 = index.constant 240 : index + vector.store %w1_28, %gv[%o1_28] : vector<4xi32>, view<512xi32> + %v1_29_0 = scalar.constant 37956 : i32 + %v1_29_1 = scalar.constant 38225 : i32 + %v1_29_2 = scalar.constant 39041 : i32 + %v1_29_3 = scalar.constant 39200 : i32 + %w1_29 = vector.from_elements %v1_29_0, %v1_29_1, %v1_29_2, %v1_29_3 : vector<4xi32> + %o1_29 = index.constant 244 : index + vector.store %w1_29, %gv[%o1_29] : vector<4xi32>, view<512xi32> + %v1_30_0 = scalar.constant 40962 : i32 + %v1_30_1 = scalar.constant 41040 : i32 + %v1_30_2 = scalar.constant 41093 : i32 + %v1_30_3 = scalar.constant 41225 : i32 + %w1_30 = vector.from_elements %v1_30_0, %v1_30_1, %v1_30_2, %v1_30_3 : vector<4xi32> + %o1_30 = index.constant 248 : index + vector.store %w1_30, %gv[%o1_30] : vector<4xi32>, view<512xi32> + %v1_31_0 = scalar.constant 41472 : i32 + %v1_31_1 = scalar.constant 42008 : i32 + %v1_31_2 = scalar.constant 43088 : i32 + %v1_31_3 = scalar.constant 43268 : i32 + %w1_31 = vector.from_elements %v1_31_0, %v1_31_1, %v1_31_2, %v1_31_3 : vector<4xi32> + %o1_31 = index.constant 252 : index + vector.store %w1_31, %gv[%o1_31] : vector<4xi32>, view<512xi32> + } + func.return +} +func.def inline @ggml_kquant_iq2xs_grid_fill(%grid: buffer, %chunk: index) { + %zero_offset = index.constant 0 : offset + %gv = buffer.view %grid[%zero_offset] : buffer -> view<512xi32> + %k0 = index.constant 0 : index + %is0 = index.cmp eq, %chunk, %k0 : index + scf.if %is0 { + %v0_0_0 = scalar.constant 0 : i32 + %v0_0_1 = scalar.constant 2 : i32 + %v0_0_2 = scalar.constant 5 : i32 + %v0_0_3 = scalar.constant 8 : i32 + %w0_0 = vector.from_elements %v0_0_0, %v0_0_1, %v0_0_2, %v0_0_3 : vector<4xi32> + %o0_0 = index.constant 0 : index + vector.store %w0_0, %gv[%o0_0] : vector<4xi32>, view<512xi32> + %v0_1_0 = scalar.constant 10 : i32 + %v0_1_1 = scalar.constant 17 : i32 + %v0_1_2 = scalar.constant 20 : i32 + %v0_1_3 = scalar.constant 22 : i32 + %w0_1 = vector.from_elements %v0_1_0, %v0_1_1, %v0_1_2, %v0_1_3 : vector<4xi32> + %o0_1 = index.constant 4 : index + vector.store %w0_1, %gv[%o0_1] : vector<4xi32>, view<512xi32> + %v0_2_0 = scalar.constant 25 : i32 + %v0_2_1 = scalar.constant 32 : i32 + %v0_2_2 = scalar.constant 34 : i32 + %v0_2_3 = scalar.constant 37 : i32 + %w0_2 = vector.from_elements %v0_2_0, %v0_2_1, %v0_2_2, %v0_2_3 : vector<4xi32> + %o0_2 = index.constant 8 : index + vector.store %w0_2, %gv[%o0_2] : vector<4xi32>, view<512xi32> + %v0_3_0 = scalar.constant 40 : i32 + %v0_3_1 = scalar.constant 65 : i32 + %v0_3_2 = scalar.constant 68 : i32 + %v0_3_3 = scalar.constant 70 : i32 + %w0_3 = vector.from_elements %v0_3_0, %v0_3_1, %v0_3_2, %v0_3_3 : vector<4xi32> + %o0_3 = index.constant 12 : index + vector.store %w0_3, %gv[%o0_3] : vector<4xi32>, view<512xi32> + %v0_4_0 = scalar.constant 73 : i32 + %v0_4_1 = scalar.constant 80 : i32 + %v0_4_2 = scalar.constant 82 : i32 + %v0_4_3 = scalar.constant 85 : i32 + %w0_4 = vector.from_elements %v0_4_0, %v0_4_1, %v0_4_2, %v0_4_3 : vector<4xi32> + %o0_4 = index.constant 16 : index + vector.store %w0_4, %gv[%o0_4] : vector<4xi32>, view<512xi32> + %v0_5_0 = scalar.constant 88 : i32 + %v0_5_1 = scalar.constant 97 : i32 + %v0_5_2 = scalar.constant 100 : i32 + %v0_5_3 = scalar.constant 128 : i32 + %w0_5 = vector.from_elements %v0_5_0, %v0_5_1, %v0_5_2, %v0_5_3 : vector<4xi32> + %o0_5 = index.constant 20 : index + vector.store %w0_5, %gv[%o0_5] : vector<4xi32>, view<512xi32> + %v0_6_0 = scalar.constant 130 : i32 + %v0_6_1 = scalar.constant 133 : i32 + %v0_6_2 = scalar.constant 136 : i32 + %v0_6_3 = scalar.constant 145 : i32 + %w0_6 = vector.from_elements %v0_6_0, %v0_6_1, %v0_6_2, %v0_6_3 : vector<4xi32> + %o0_6 = index.constant 24 : index + vector.store %w0_6, %gv[%o0_6] : vector<4xi32>, view<512xi32> + %v0_7_0 = scalar.constant 148 : i32 + %v0_7_1 = scalar.constant 153 : i32 + %v0_7_2 = scalar.constant 160 : i32 + %v0_7_3 = scalar.constant 257 : i32 + %w0_7 = vector.from_elements %v0_7_0, %v0_7_1, %v0_7_2, %v0_7_3 : vector<4xi32> + %o0_7 = index.constant 28 : index + vector.store %w0_7, %gv[%o0_7] : vector<4xi32>, view<512xi32> + %v0_8_0 = scalar.constant 260 : i32 + %v0_8_1 = scalar.constant 262 : i32 + %v0_8_2 = scalar.constant 265 : i32 + %v0_8_3 = scalar.constant 272 : i32 + %w0_8 = vector.from_elements %v0_8_0, %v0_8_1, %v0_8_2, %v0_8_3 : vector<4xi32> + %o0_8 = index.constant 32 : index + vector.store %w0_8, %gv[%o0_8] : vector<4xi32>, view<512xi32> + %v0_9_0 = scalar.constant 274 : i32 + %v0_9_1 = scalar.constant 277 : i32 + %v0_9_2 = scalar.constant 280 : i32 + %v0_9_3 = scalar.constant 282 : i32 + %w0_9 = vector.from_elements %v0_9_0, %v0_9_1, %v0_9_2, %v0_9_3 : vector<4xi32> + %o0_9 = index.constant 36 : index + vector.store %w0_9, %gv[%o0_9] : vector<4xi32>, view<512xi32> + %v0_10_0 = scalar.constant 289 : i32 + %v0_10_1 = scalar.constant 292 : i32 + %v0_10_2 = scalar.constant 320 : i32 + %v0_10_3 = scalar.constant 322 : i32 + %w0_10 = vector.from_elements %v0_10_0, %v0_10_1, %v0_10_2, %v0_10_3 : vector<4xi32> + %o0_10 = index.constant 40 : index + vector.store %w0_10, %gv[%o0_10] : vector<4xi32>, view<512xi32> + %v0_11_0 = scalar.constant 325 : i32 + %v0_11_1 = scalar.constant 328 : i32 + %v0_11_2 = scalar.constant 337 : i32 + %v0_11_3 = scalar.constant 340 : i32 + %w0_11 = vector.from_elements %v0_11_0, %v0_11_1, %v0_11_2, %v0_11_3 : vector<4xi32> + %o0_11 = index.constant 44 : index + vector.store %w0_11, %gv[%o0_11] : vector<4xi32>, view<512xi32> + %v0_12_0 = scalar.constant 352 : i32 + %v0_12_1 = scalar.constant 360 : i32 + %v0_12_2 = scalar.constant 385 : i32 + %v0_12_3 = scalar.constant 388 : i32 + %w0_12 = vector.from_elements %v0_12_0, %v0_12_1, %v0_12_2, %v0_12_3 : vector<4xi32> + %o0_12 = index.constant 48 : index + vector.store %w0_12, %gv[%o0_12] : vector<4xi32>, view<512xi32> + %v0_13_0 = scalar.constant 400 : i32 + %v0_13_1 = scalar.constant 512 : i32 + %v0_13_2 = scalar.constant 514 : i32 + %v0_13_3 = scalar.constant 517 : i32 + %w0_13 = vector.from_elements %v0_13_0, %v0_13_1, %v0_13_2, %v0_13_3 : vector<4xi32> + %o0_13 = index.constant 52 : index + vector.store %w0_13, %gv[%o0_13] : vector<4xi32>, view<512xi32> + %v0_14_0 = scalar.constant 520 : i32 + %v0_14_1 = scalar.constant 529 : i32 + %v0_14_2 = scalar.constant 532 : i32 + %v0_14_3 = scalar.constant 544 : i32 + %w0_14 = vector.from_elements %v0_14_0, %v0_14_1, %v0_14_2, %v0_14_3 : vector<4xi32> + %o0_14 = index.constant 56 : index + vector.store %w0_14, %gv[%o0_14] : vector<4xi32>, view<512xi32> + %v0_15_0 = scalar.constant 577 : i32 + %v0_15_1 = scalar.constant 580 : i32 + %v0_15_2 = scalar.constant 592 : i32 + %v0_15_3 = scalar.constant 597 : i32 + %w0_15 = vector.from_elements %v0_15_0, %v0_15_1, %v0_15_2, %v0_15_3 : vector<4xi32> + %o0_15 = index.constant 60 : index + vector.store %w0_15, %gv[%o0_15] : vector<4xi32>, view<512xi32> + %v0_16_0 = scalar.constant 640 : i32 + %v0_16_1 = scalar.constant 650 : i32 + %v0_16_2 = scalar.constant 1025 : i32 + %v0_16_3 = scalar.constant 1028 : i32 + %w0_16 = vector.from_elements %v0_16_0, %v0_16_1, %v0_16_2, %v0_16_3 : vector<4xi32> + %o0_16 = index.constant 64 : index + vector.store %w0_16, %gv[%o0_16] : vector<4xi32>, view<512xi32> + %v0_17_0 = scalar.constant 1030 : i32 + %v0_17_1 = scalar.constant 1033 : i32 + %v0_17_2 = scalar.constant 1040 : i32 + %v0_17_3 = scalar.constant 1042 : i32 + %w0_17 = vector.from_elements %v0_17_0, %v0_17_1, %v0_17_2, %v0_17_3 : vector<4xi32> + %o0_17 = index.constant 68 : index + vector.store %w0_17, %gv[%o0_17] : vector<4xi32>, view<512xi32> + %v0_18_0 = scalar.constant 1045 : i32 + %v0_18_1 = scalar.constant 1048 : i32 + %v0_18_2 = scalar.constant 1057 : i32 + %v0_18_3 = scalar.constant 1060 : i32 + %w0_18 = vector.from_elements %v0_18_0, %v0_18_1, %v0_18_2, %v0_18_3 : vector<4xi32> + %o0_18 = index.constant 72 : index + vector.store %w0_18, %gv[%o0_18] : vector<4xi32>, view<512xi32> + %v0_19_0 = scalar.constant 1088 : i32 + %v0_19_1 = scalar.constant 1090 : i32 + %v0_19_2 = scalar.constant 1093 : i32 + %v0_19_3 = scalar.constant 1096 : i32 + %w0_19 = vector.from_elements %v0_19_0, %v0_19_1, %v0_19_2, %v0_19_3 : vector<4xi32> + %o0_19 = index.constant 76 : index + vector.store %w0_19, %gv[%o0_19] : vector<4xi32>, view<512xi32> + %v0_20_0 = scalar.constant 1105 : i32 + %v0_20_1 = scalar.constant 1108 : i32 + %v0_20_2 = scalar.constant 1110 : i32 + %v0_20_3 = scalar.constant 1120 : i32 + %w0_20 = vector.from_elements %v0_20_0, %v0_20_1, %v0_20_2, %v0_20_3 : vector<4xi32> + %o0_20 = index.constant 80 : index + vector.store %w0_20, %gv[%o0_20] : vector<4xi32>, view<512xi32> + %v0_21_0 = scalar.constant 1153 : i32 + %v0_21_1 = scalar.constant 1156 : i32 + %v0_21_2 = scalar.constant 1168 : i32 + %v0_21_3 = scalar.constant 1280 : i32 + %w0_21 = vector.from_elements %v0_21_0, %v0_21_1, %v0_21_2, %v0_21_3 : vector<4xi32> + %o0_21 = index.constant 84 : index + vector.store %w0_21, %gv[%o0_21] : vector<4xi32>, view<512xi32> + %v0_22_0 = scalar.constant 1282 : i32 + %v0_22_1 = scalar.constant 1285 : i32 + %v0_22_2 = scalar.constant 1288 : i32 + %v0_22_3 = scalar.constant 1297 : i32 + %w0_22 = vector.from_elements %v0_22_0, %v0_22_1, %v0_22_2, %v0_22_3 : vector<4xi32> + %o0_22 = index.constant 88 : index + vector.store %w0_22, %gv[%o0_22] : vector<4xi32>, view<512xi32> + %v0_23_0 = scalar.constant 1300 : i32 + %v0_23_1 = scalar.constant 1312 : i32 + %v0_23_2 = scalar.constant 1345 : i32 + %v0_23_3 = scalar.constant 1348 : i32 + %w0_23 = vector.from_elements %v0_23_0, %v0_23_1, %v0_23_2, %v0_23_3 : vector<4xi32> + %o0_23 = index.constant 92 : index + vector.store %w0_23, %gv[%o0_23] : vector<4xi32>, view<512xi32> + %v0_24_0 = scalar.constant 1360 : i32 + %v0_24_1 = scalar.constant 1377 : i32 + %v0_24_2 = scalar.constant 1408 : i32 + %v0_24_3 = scalar.constant 1537 : i32 + %w0_24 = vector.from_elements %v0_24_0, %v0_24_1, %v0_24_2, %v0_24_3 : vector<4xi32> + %o0_24 = index.constant 96 : index + vector.store %w0_24, %gv[%o0_24] : vector<4xi32>, view<512xi32> + %v0_25_0 = scalar.constant 1540 : i32 + %v0_25_1 = scalar.constant 1552 : i32 + %v0_25_2 = scalar.constant 1574 : i32 + %v0_25_3 = scalar.constant 1600 : i32 + %w0_25 = vector.from_elements %v0_25_0, %v0_25_1, %v0_25_2, %v0_25_3 : vector<4xi32> + %o0_25 = index.constant 100 : index + vector.store %w0_25, %gv[%o0_25] : vector<4xi32>, view<512xi32> + %v0_26_0 = scalar.constant 1602 : i32 + %v0_26_1 = scalar.constant 1668 : i32 + %v0_26_2 = scalar.constant 2048 : i32 + %v0_26_3 = scalar.constant 2050 : i32 + %w0_26 = vector.from_elements %v0_26_0, %v0_26_1, %v0_26_2, %v0_26_3 : vector<4xi32> + %o0_26 = index.constant 104 : index + vector.store %w0_26, %gv[%o0_26] : vector<4xi32>, view<512xi32> + %v0_27_0 = scalar.constant 2053 : i32 + %v0_27_1 = scalar.constant 2056 : i32 + %v0_27_2 = scalar.constant 2058 : i32 + %v0_27_3 = scalar.constant 2065 : i32 + %w0_27 = vector.from_elements %v0_27_0, %v0_27_1, %v0_27_2, %v0_27_3 : vector<4xi32> + %o0_27 = index.constant 108 : index + vector.store %w0_27, %gv[%o0_27] : vector<4xi32>, view<512xi32> + %v0_28_0 = scalar.constant 2068 : i32 + %v0_28_1 = scalar.constant 2080 : i32 + %v0_28_2 = scalar.constant 2085 : i32 + %v0_28_3 = scalar.constant 2113 : i32 + %w0_28 = vector.from_elements %v0_28_0, %v0_28_1, %v0_28_2, %v0_28_3 : vector<4xi32> + %o0_28 = index.constant 112 : index + vector.store %w0_28, %gv[%o0_28] : vector<4xi32>, view<512xi32> + %v0_29_0 = scalar.constant 2116 : i32 + %v0_29_1 = scalar.constant 2128 : i32 + %v0_29_2 = scalar.constant 2136 : i32 + %v0_29_3 = scalar.constant 2176 : i32 + %w0_29 = vector.from_elements %v0_29_0, %v0_29_1, %v0_29_2, %v0_29_3 : vector<4xi32> + %o0_29 = index.constant 116 : index + vector.store %w0_29, %gv[%o0_29] : vector<4xi32>, view<512xi32> + %v0_30_0 = scalar.constant 2208 : i32 + %v0_30_1 = scalar.constant 2218 : i32 + %v0_30_2 = scalar.constant 2305 : i32 + %v0_30_3 = scalar.constant 2308 : i32 + %w0_30 = vector.from_elements %v0_30_0, %v0_30_1, %v0_30_2, %v0_30_3 : vector<4xi32> + %o0_30 = index.constant 120 : index + vector.store %w0_30, %gv[%o0_30] : vector<4xi32>, view<512xi32> + %v0_31_0 = scalar.constant 2320 : i32 + %v0_31_1 = scalar.constant 2368 : i32 + %v0_31_2 = scalar.constant 2433 : i32 + %v0_31_3 = scalar.constant 2441 : i32 + %w0_31 = vector.from_elements %v0_31_0, %v0_31_1, %v0_31_2, %v0_31_3 : vector<4xi32> + %o0_31 = index.constant 124 : index + vector.store %w0_31, %gv[%o0_31] : vector<4xi32>, view<512xi32> + } + %k1 = index.constant 1 : index + %is1 = index.cmp eq, %chunk, %k1 : index + scf.if %is1 { + %v1_0_0 = scalar.constant 2560 : i32 + %v1_0_1 = scalar.constant 2592 : i32 + %v1_0_2 = scalar.constant 2600 : i32 + %v1_0_3 = scalar.constant 2710 : i32 + %w1_0 = vector.from_elements %v1_0_0, %v1_0_1, %v1_0_2, %v1_0_3 : vector<4xi32> + %o1_0 = index.constant 128 : index + vector.store %w1_0, %gv[%o1_0] : vector<4xi32>, view<512xi32> + %v1_1_0 = scalar.constant 2720 : i32 + %v1_1_1 = scalar.constant 4097 : i32 + %v1_1_2 = scalar.constant 4100 : i32 + %v1_1_3 = scalar.constant 4102 : i32 + %w1_1 = vector.from_elements %v1_1_0, %v1_1_1, %v1_1_2, %v1_1_3 : vector<4xi32> + %o1_1 = index.constant 132 : index + vector.store %w1_1, %gv[%o1_1] : vector<4xi32>, view<512xi32> + %v1_2_0 = scalar.constant 4105 : i32 + %v1_2_1 = scalar.constant 4112 : i32 + %v1_2_2 = scalar.constant 4114 : i32 + %v1_2_3 = scalar.constant 4117 : i32 + %w1_2 = vector.from_elements %v1_2_0, %v1_2_1, %v1_2_2, %v1_2_3 : vector<4xi32> + %o1_2 = index.constant 136 : index + vector.store %w1_2, %gv[%o1_2] : vector<4xi32>, view<512xi32> + %v1_3_0 = scalar.constant 4120 : i32 + %v1_3_1 = scalar.constant 4129 : i32 + %v1_3_2 = scalar.constant 4132 : i32 + %v1_3_3 = scalar.constant 4160 : i32 + %w1_3 = vector.from_elements %v1_3_0, %v1_3_1, %v1_3_2, %v1_3_3 : vector<4xi32> + %o1_3 = index.constant 140 : index + vector.store %w1_3, %gv[%o1_3] : vector<4xi32>, view<512xi32> + %v1_4_0 = scalar.constant 4162 : i32 + %v1_4_1 = scalar.constant 4165 : i32 + %v1_4_2 = scalar.constant 4168 : i32 + %v1_4_3 = scalar.constant 4177 : i32 + %w1_4 = vector.from_elements %v1_4_0, %v1_4_1, %v1_4_2, %v1_4_3 : vector<4xi32> + %o1_4 = index.constant 144 : index + vector.store %w1_4, %gv[%o1_4] : vector<4xi32>, view<512xi32> + %v1_5_0 = scalar.constant 4180 : i32 + %v1_5_1 = scalar.constant 4192 : i32 + %v1_5_2 = scalar.constant 4202 : i32 + %v1_5_3 = scalar.constant 4225 : i32 + %w1_5 = vector.from_elements %v1_5_0, %v1_5_1, %v1_5_2, %v1_5_3 : vector<4xi32> + %o1_5 = index.constant 148 : index + vector.store %w1_5, %gv[%o1_5] : vector<4xi32>, view<512xi32> + %v1_6_0 = scalar.constant 4228 : i32 + %v1_6_1 = scalar.constant 4240 : i32 + %v1_6_2 = scalar.constant 4352 : i32 + %v1_6_3 = scalar.constant 4354 : i32 + %w1_6 = vector.from_elements %v1_6_0, %v1_6_1, %v1_6_2, %v1_6_3 : vector<4xi32> + %o1_6 = index.constant 152 : index + vector.store %w1_6, %gv[%o1_6] : vector<4xi32>, view<512xi32> + %v1_7_0 = scalar.constant 4357 : i32 + %v1_7_1 = scalar.constant 4360 : i32 + %v1_7_2 = scalar.constant 4369 : i32 + %v1_7_3 = scalar.constant 4372 : i32 + %w1_7 = vector.from_elements %v1_7_0, %v1_7_1, %v1_7_2, %v1_7_3 : vector<4xi32> + %o1_7 = index.constant 156 : index + vector.store %w1_7, %gv[%o1_7] : vector<4xi32>, view<512xi32> + %v1_8_0 = scalar.constant 4384 : i32 + %v1_8_1 = scalar.constant 4417 : i32 + %v1_8_2 = scalar.constant 4420 : i32 + %v1_8_3 = scalar.constant 4432 : i32 + %w1_8 = vector.from_elements %v1_8_0, %v1_8_1, %v1_8_2, %v1_8_3 : vector<4xi32> + %o1_8 = index.constant 160 : index + vector.store %w1_8, %gv[%o1_8] : vector<4xi32>, view<512xi32> + %v1_9_0 = scalar.constant 4480 : i32 + %v1_9_1 = scalar.constant 4500 : i32 + %v1_9_2 = scalar.constant 4502 : i32 + %v1_9_3 = scalar.constant 4609 : i32 + %w1_9 = vector.from_elements %v1_9_0, %v1_9_1, %v1_9_2, %v1_9_3 : vector<4xi32> + %o1_9 = index.constant 164 : index + vector.store %w1_9, %gv[%o1_9] : vector<4xi32>, view<512xi32> + %v1_10_0 = scalar.constant 4612 : i32 + %v1_10_1 = scalar.constant 4614 : i32 + %v1_10_2 = scalar.constant 4624 : i32 + %v1_10_3 = scalar.constant 4672 : i32 + %w1_10 = vector.from_elements %v1_10_0, %v1_10_1, %v1_10_2, %v1_10_3 : vector<4xi32> + %o1_10 = index.constant 168 : index + vector.store %w1_10, %gv[%o1_10] : vector<4xi32>, view<512xi32> + %v1_11_0 = scalar.constant 4704 : i32 + %v1_11_1 = scalar.constant 5120 : i32 + %v1_11_2 = scalar.constant 5122 : i32 + %v1_11_3 = scalar.constant 5125 : i32 + %w1_11 = vector.from_elements %v1_11_0, %v1_11_1, %v1_11_2, %v1_11_3 : vector<4xi32> + %o1_11 = index.constant 172 : index + vector.store %w1_11, %gv[%o1_11] : vector<4xi32>, view<512xi32> + %v1_12_0 = scalar.constant 5128 : i32 + %v1_12_1 = scalar.constant 5137 : i32 + %v1_12_2 = scalar.constant 5140 : i32 + %v1_12_3 = scalar.constant 5152 : i32 + %w1_12 = vector.from_elements %v1_12_0, %v1_12_1, %v1_12_2, %v1_12_3 : vector<4xi32> + %o1_12 = index.constant 176 : index + vector.store %w1_12, %gv[%o1_12] : vector<4xi32>, view<512xi32> + %v1_13_0 = scalar.constant 5185 : i32 + %v1_13_1 = scalar.constant 5188 : i32 + %v1_13_2 = scalar.constant 5193 : i32 + %v1_13_3 = scalar.constant 5200 : i32 + %w1_13 = vector.from_elements %v1_13_0, %v1_13_1, %v1_13_2, %v1_13_3 : vector<4xi32> + %o1_13 = index.constant 180 : index + vector.store %w1_13, %gv[%o1_13] : vector<4xi32>, view<512xi32> + %v1_14_0 = scalar.constant 5220 : i32 + %v1_14_1 = scalar.constant 5248 : i32 + %v1_14_2 = scalar.constant 5377 : i32 + %v1_14_3 = scalar.constant 5380 : i32 + %w1_14 = vector.from_elements %v1_14_0, %v1_14_1, %v1_14_2, %v1_14_3 : vector<4xi32> + %o1_14 = index.constant 184 : index + vector.store %w1_14, %gv[%o1_14] : vector<4xi32>, view<512xi32> + %v1_15_0 = scalar.constant 5392 : i32 + %v1_15_1 = scalar.constant 5440 : i32 + %v1_15_2 = scalar.constant 5632 : i32 + %v1_15_3 = scalar.constant 5652 : i32 + %w1_15 = vector.from_elements %v1_15_0, %v1_15_1, %v1_15_2, %v1_15_3 : vector<4xi32> + %o1_15 = index.constant 188 : index + vector.store %w1_15, %gv[%o1_15] : vector<4xi32>, view<512xi32> + %v1_16_0 = scalar.constant 5705 : i32 + %v1_16_1 = scalar.constant 6145 : i32 + %v1_16_2 = scalar.constant 6148 : i32 + %v1_16_3 = scalar.constant 6160 : i32 + %w1_16 = vector.from_elements %v1_16_0, %v1_16_1, %v1_16_2, %v1_16_3 : vector<4xi32> + %o1_16 = index.constant 192 : index + vector.store %w1_16, %gv[%o1_16] : vector<4xi32>, view<512xi32> + %v1_17_0 = scalar.constant 6162 : i32 + %v1_17_1 = scalar.constant 6208 : i32 + %v1_17_2 = scalar.constant 6228 : i32 + %v1_17_3 = scalar.constant 6278 : i32 + %w1_17 = vector.from_elements %v1_17_0, %v1_17_1, %v1_17_2, %v1_17_3 : vector<4xi32> + %o1_17 = index.constant 196 : index + vector.store %w1_17, %gv[%o1_17] : vector<4xi32>, view<512xi32> + %v1_18_0 = scalar.constant 6400 : i32 + %v1_18_1 = scalar.constant 6405 : i32 + %v1_18_2 = scalar.constant 6502 : i32 + %v1_18_3 = scalar.constant 6737 : i32 + %w1_18 = vector.from_elements %v1_18_0, %v1_18_1, %v1_18_2, %v1_18_3 : vector<4xi32> + %o1_18 = index.constant 200 : index + vector.store %w1_18, %gv[%o1_18] : vector<4xi32>, view<512xi32> + %v1_19_0 = scalar.constant 6825 : i32 + %v1_19_1 = scalar.constant 8192 : i32 + %v1_19_2 = scalar.constant 8194 : i32 + %v1_19_3 = scalar.constant 8197 : i32 + %w1_19 = vector.from_elements %v1_19_0, %v1_19_1, %v1_19_2, %v1_19_3 : vector<4xi32> + %o1_19 = index.constant 204 : index + vector.store %w1_19, %gv[%o1_19] : vector<4xi32>, view<512xi32> + %v1_20_0 = scalar.constant 8200 : i32 + %v1_20_1 = scalar.constant 8202 : i32 + %v1_20_2 = scalar.constant 8209 : i32 + %v1_20_3 = scalar.constant 8212 : i32 + %w1_20 = vector.from_elements %v1_20_0, %v1_20_1, %v1_20_2, %v1_20_3 : vector<4xi32> + %o1_20 = index.constant 208 : index + vector.store %w1_20, %gv[%o1_20] : vector<4xi32>, view<512xi32> + %v1_21_0 = scalar.constant 8224 : i32 + %v1_21_1 = scalar.constant 8257 : i32 + %v1_21_2 = scalar.constant 8260 : i32 + %v1_21_3 = scalar.constant 8272 : i32 + %w1_21 = vector.from_elements %v1_21_0, %v1_21_1, %v1_21_2, %v1_21_3 : vector<4xi32> + %o1_21 = index.constant 212 : index + vector.store %w1_21, %gv[%o1_21] : vector<4xi32>, view<512xi32> + %v1_22_0 = scalar.constant 8320 : i32 + %v1_22_1 = scalar.constant 8352 : i32 + %v1_22_2 = scalar.constant 8449 : i32 + %v1_22_3 = scalar.constant 8452 : i32 + %w1_22 = vector.from_elements %v1_22_0, %v1_22_1, %v1_22_2, %v1_22_3 : vector<4xi32> + %o1_22 = index.constant 216 : index + vector.store %w1_22, %gv[%o1_22] : vector<4xi32>, view<512xi32> + %v1_23_0 = scalar.constant 8464 : i32 + %v1_23_1 = scalar.constant 8512 : i32 + %v1_23_2 = scalar.constant 8520 : i32 + %v1_23_3 = scalar.constant 8549 : i32 + %w1_23 = vector.from_elements %v1_23_0, %v1_23_1, %v1_23_2, %v1_23_3 : vector<4xi32> + %o1_23 = index.constant 220 : index + vector.store %w1_23, %gv[%o1_23] : vector<4xi32>, view<512xi32> + %v1_24_0 = scalar.constant 8704 : i32 + %v1_24_1 = scalar.constant 8738 : i32 + %v1_24_2 = scalar.constant 8832 : i32 + %v1_24_3 = scalar.constant 8872 : i32 + %w1_24 = vector.from_elements %v1_24_0, %v1_24_1, %v1_24_2, %v1_24_3 : vector<4xi32> + %o1_24 = index.constant 224 : index + vector.store %w1_24, %gv[%o1_24] : vector<4xi32>, view<512xi32> + %v1_25_0 = scalar.constant 9217 : i32 + %v1_25_1 = scalar.constant 9220 : i32 + %v1_25_2 = scalar.constant 9232 : i32 + %v1_25_3 = scalar.constant 9257 : i32 + %w1_25 = vector.from_elements %v1_25_0, %v1_25_1, %v1_25_2, %v1_25_3 : vector<4xi32> + %o1_25 = index.constant 228 : index + vector.store %w1_25, %gv[%o1_25] : vector<4xi32>, view<512xi32> + %v1_26_0 = scalar.constant 9280 : i32 + %v1_26_1 = scalar.constant 9472 : i32 + %v1_26_2 = scalar.constant 9537 : i32 + %v1_26_3 = scalar.constant 9554 : i32 + %w1_26 = vector.from_elements %v1_26_0, %v1_26_1, %v1_26_2, %v1_26_3 : vector<4xi32> + %o1_26 = index.constant 232 : index + vector.store %w1_26, %gv[%o1_26] : vector<4xi32>, view<512xi32> + %v1_27_0 = scalar.constant 9625 : i32 + %v1_27_1 = scalar.constant 9729 : i32 + %v1_27_2 = scalar.constant 9754 : i32 + %v1_27_3 = scalar.constant 9894 : i32 + %w1_27 = vector.from_elements %v1_27_0, %v1_27_1, %v1_27_2, %v1_27_3 : vector<4xi32> + %o1_27 = index.constant 236 : index + vector.store %w1_27, %gv[%o1_27] : vector<4xi32>, view<512xi32> + %v1_28_0 = scalar.constant 10240 : i32 + %v1_28_1 = scalar.constant 10248 : i32 + %v1_28_2 = scalar.constant 10250 : i32 + %v1_28_3 = scalar.constant 10272 : i32 + %w1_28 = vector.from_elements %v1_28_0, %v1_28_1, %v1_28_2, %v1_28_3 : vector<4xi32> + %o1_28 = index.constant 240 : index + vector.store %w1_28, %gv[%o1_28] : vector<4xi32>, view<512xi32> + %v1_29_0 = scalar.constant 10325 : i32 + %v1_29_1 = scalar.constant 10376 : i32 + %v1_29_2 = scalar.constant 10402 : i32 + %v1_29_3 = scalar.constant 10600 : i32 + %w1_29 = vector.from_elements %v1_29_0, %v1_29_1, %v1_29_2, %v1_29_3 : vector<4xi32> + %o1_29 = index.constant 244 : index + vector.store %w1_29, %gv[%o1_29] : vector<4xi32>, view<512xi32> + %v1_30_0 = scalar.constant 10640 : i32 + %v1_30_1 = scalar.constant 10760 : i32 + %v1_30_2 = scalar.constant 10784 : i32 + %v1_30_3 = scalar.constant 10882 : i32 + %w1_30 = vector.from_elements %v1_30_0, %v1_30_1, %v1_30_2, %v1_30_3 : vector<4xi32> + %o1_30 = index.constant 248 : index + vector.store %w1_30, %gv[%o1_30] : vector<4xi32>, view<512xi32> + %v1_31_0 = scalar.constant 10888 : i32 + %v1_31_1 = scalar.constant 10890 : i32 + %v1_31_2 = scalar.constant 16385 : i32 + %v1_31_3 = scalar.constant 16388 : i32 + %w1_31 = vector.from_elements %v1_31_0, %v1_31_1, %v1_31_2, %v1_31_3 : vector<4xi32> + %o1_31 = index.constant 252 : index + vector.store %w1_31, %gv[%o1_31] : vector<4xi32>, view<512xi32> + } + %k2 = index.constant 2 : index + %is2 = index.cmp eq, %chunk, %k2 : index + scf.if %is2 { + %v2_0_0 = scalar.constant 16390 : i32 + %v2_0_1 = scalar.constant 16393 : i32 + %v2_0_2 = scalar.constant 16400 : i32 + %v2_0_3 = scalar.constant 16402 : i32 + %w2_0 = vector.from_elements %v2_0_0, %v2_0_1, %v2_0_2, %v2_0_3 : vector<4xi32> + %o2_0 = index.constant 256 : index + vector.store %w2_0, %gv[%o2_0] : vector<4xi32>, view<512xi32> + %v2_1_0 = scalar.constant 16405 : i32 + %v2_1_1 = scalar.constant 16408 : i32 + %v2_1_2 = scalar.constant 16417 : i32 + %v2_1_3 = scalar.constant 16420 : i32 + %w2_1 = vector.from_elements %v2_1_0, %v2_1_1, %v2_1_2, %v2_1_3 : vector<4xi32> + %o2_1 = index.constant 260 : index + vector.store %w2_1, %gv[%o2_1] : vector<4xi32>, view<512xi32> + %v2_2_0 = scalar.constant 16448 : i32 + %v2_2_1 = scalar.constant 16450 : i32 + %v2_2_2 = scalar.constant 16453 : i32 + %v2_2_3 = scalar.constant 16456 : i32 + %w2_2 = vector.from_elements %v2_2_0, %v2_2_1, %v2_2_2, %v2_2_3 : vector<4xi32> + %o2_2 = index.constant 264 : index + vector.store %w2_2, %gv[%o2_2] : vector<4xi32>, view<512xi32> + %v2_3_0 = scalar.constant 16458 : i32 + %v2_3_1 = scalar.constant 16465 : i32 + %v2_3_2 = scalar.constant 16468 : i32 + %v2_3_3 = scalar.constant 16480 : i32 + %w2_3 = vector.from_elements %v2_3_0, %v2_3_1, %v2_3_2, %v2_3_3 : vector<4xi32> + %o2_3 = index.constant 268 : index + vector.store %w2_3, %gv[%o2_3] : vector<4xi32>, view<512xi32> + %v2_4_0 = scalar.constant 16485 : i32 + %v2_4_1 = scalar.constant 16513 : i32 + %v2_4_2 = scalar.constant 16516 : i32 + %v2_4_3 = scalar.constant 16528 : i32 + %w2_4 = vector.from_elements %v2_4_0, %v2_4_1, %v2_4_2, %v2_4_3 : vector<4xi32> + %o2_4 = index.constant 272 : index + vector.store %w2_4, %gv[%o2_4] : vector<4xi32>, view<512xi32> + %v2_5_0 = scalar.constant 16640 : i32 + %v2_5_1 = scalar.constant 16642 : i32 + %v2_5_2 = scalar.constant 16645 : i32 + %v2_5_3 = scalar.constant 16648 : i32 + %w2_5 = vector.from_elements %v2_5_0, %v2_5_1, %v2_5_2, %v2_5_3 : vector<4xi32> + %o2_5 = index.constant 276 : index + vector.store %w2_5, %gv[%o2_5] : vector<4xi32>, view<512xi32> + %v2_6_0 = scalar.constant 16657 : i32 + %v2_6_1 = scalar.constant 16660 : i32 + %v2_6_2 = scalar.constant 16672 : i32 + %v2_6_3 = scalar.constant 16705 : i32 + %w2_6 = vector.from_elements %v2_6_0, %v2_6_1, %v2_6_2, %v2_6_3 : vector<4xi32> + %o2_6 = index.constant 280 : index + vector.store %w2_6, %gv[%o2_6] : vector<4xi32>, view<512xi32> + %v2_7_0 = scalar.constant 16708 : i32 + %v2_7_1 = scalar.constant 16720 : i32 + %v2_7_2 = scalar.constant 16768 : i32 + %v2_7_3 = scalar.constant 16773 : i32 + %w2_7 = vector.from_elements %v2_7_0, %v2_7_1, %v2_7_2, %v2_7_3 : vector<4xi32> + %o2_7 = index.constant 284 : index + vector.store %w2_7, %gv[%o2_7] : vector<4xi32>, view<512xi32> + %v2_8_0 = scalar.constant 16802 : i32 + %v2_8_1 = scalar.constant 16897 : i32 + %v2_8_2 = scalar.constant 16900 : i32 + %v2_8_3 = scalar.constant 16912 : i32 + %w2_8 = vector.from_elements %v2_8_0, %v2_8_1, %v2_8_2, %v2_8_3 : vector<4xi32> + %o2_8 = index.constant 288 : index + vector.store %w2_8, %gv[%o2_8] : vector<4xi32>, view<512xi32> + %v2_9_0 = scalar.constant 16914 : i32 + %v2_9_1 = scalar.constant 16937 : i32 + %v2_9_2 = scalar.constant 16960 : i32 + %v2_9_3 = scalar.constant 17408 : i32 + %w2_9 = vector.from_elements %v2_9_0, %v2_9_1, %v2_9_2, %v2_9_3 : vector<4xi32> + %o2_9 = index.constant 292 : index + vector.store %w2_9, %gv[%o2_9] : vector<4xi32>, view<512xi32> + %v2_10_0 = scalar.constant 17410 : i32 + %v2_10_1 = scalar.constant 17413 : i32 + %v2_10_2 = scalar.constant 17416 : i32 + %v2_10_3 = scalar.constant 17425 : i32 + %w2_10 = vector.from_elements %v2_10_0, %v2_10_1, %v2_10_2, %v2_10_3 : vector<4xi32> + %o2_10 = index.constant 296 : index + vector.store %w2_10, %gv[%o2_10] : vector<4xi32>, view<512xi32> + %v2_11_0 = scalar.constant 17428 : i32 + %v2_11_1 = scalar.constant 17433 : i32 + %v2_11_2 = scalar.constant 17440 : i32 + %v2_11_3 = scalar.constant 17473 : i32 + %w2_11 = vector.from_elements %v2_11_0, %v2_11_1, %v2_11_2, %v2_11_3 : vector<4xi32> + %o2_11 = index.constant 300 : index + vector.store %w2_11, %gv[%o2_11] : vector<4xi32>, view<512xi32> + %v2_12_0 = scalar.constant 17476 : i32 + %v2_12_1 = scalar.constant 17488 : i32 + %v2_12_2 = scalar.constant 17536 : i32 + %v2_12_3 = scalar.constant 17556 : i32 + %w2_12 = vector.from_elements %v2_12_0, %v2_12_1, %v2_12_2, %v2_12_3 : vector<4xi32> + %o2_12 = index.constant 304 : index + vector.store %w2_12, %gv[%o2_12] : vector<4xi32>, view<512xi32> + %v2_13_0 = scalar.constant 17665 : i32 + %v2_13_1 = scalar.constant 17668 : i32 + %v2_13_2 = scalar.constant 17680 : i32 + %v2_13_3 = scalar.constant 17700 : i32 + %w2_13 = vector.from_elements %v2_13_0, %v2_13_1, %v2_13_2, %v2_13_3 : vector<4xi32> + %o2_13 = index.constant 308 : index + vector.store %w2_13, %gv[%o2_13] : vector<4xi32>, view<512xi32> + %v2_14_0 = scalar.constant 17728 : i32 + %v2_14_1 = scalar.constant 17818 : i32 + %v2_14_2 = scalar.constant 17920 : i32 + %v2_14_3 = scalar.constant 17930 : i32 + %w2_14 = vector.from_elements %v2_14_0, %v2_14_1, %v2_14_2, %v2_14_3 : vector<4xi32> + %o2_14 = index.constant 312 : index + vector.store %w2_14, %gv[%o2_14] : vector<4xi32>, view<512xi32> + %v2_15_0 = scalar.constant 17988 : i32 + %v2_15_1 = scalar.constant 18000 : i32 + %v2_15_2 = scalar.constant 18433 : i32 + %v2_15_3 = scalar.constant 18436 : i32 + %w2_15 = vector.from_elements %v2_15_0, %v2_15_1, %v2_15_2, %v2_15_3 : vector<4xi32> + %o2_15 = index.constant 316 : index + vector.store %w2_15, %gv[%o2_15] : vector<4xi32>, view<512xi32> + %v2_16_0 = scalar.constant 18448 : i32 + %v2_16_1 = scalar.constant 18496 : i32 + %v2_16_2 = scalar.constant 18501 : i32 + %v2_16_3 = scalar.constant 18516 : i32 + %w2_16 = vector.from_elements %v2_16_0, %v2_16_1, %v2_16_2, %v2_16_3 : vector<4xi32> + %o2_16 = index.constant 320 : index + vector.store %w2_16, %gv[%o2_16] : vector<4xi32>, view<512xi32> + %v2_17_0 = scalar.constant 18530 : i32 + %v2_17_1 = scalar.constant 18688 : i32 + %v2_17_2 = scalar.constant 18705 : i32 + %v2_17_3 = scalar.constant 18756 : i32 + %w2_17 = vector.from_elements %v2_17_0, %v2_17_1, %v2_17_2, %v2_17_3 : vector<4xi32> + %o2_17 = index.constant 324 : index + vector.store %w2_17, %gv[%o2_17] : vector<4xi32>, view<512xi32> + %v2_18_0 = scalar.constant 18768 : i32 + %v2_18_1 = scalar.constant 18793 : i32 + %v2_18_2 = scalar.constant 18948 : i32 + %v2_18_3 = scalar.constant 20480 : i32 + %w2_18 = vector.from_elements %v2_18_0, %v2_18_1, %v2_18_2, %v2_18_3 : vector<4xi32> + %o2_18 = index.constant 328 : index + vector.store %w2_18, %gv[%o2_18] : vector<4xi32>, view<512xi32> + %v2_19_0 = scalar.constant 20482 : i32 + %v2_19_1 = scalar.constant 20485 : i32 + %v2_19_2 = scalar.constant 20488 : i32 + %v2_19_3 = scalar.constant 20497 : i32 + %w2_19 = vector.from_elements %v2_19_0, %v2_19_1, %v2_19_2, %v2_19_3 : vector<4xi32> + %o2_19 = index.constant 332 : index + vector.store %w2_19, %gv[%o2_19] : vector<4xi32>, view<512xi32> + %v2_20_0 = scalar.constant 20500 : i32 + %v2_20_1 = scalar.constant 20512 : i32 + %v2_20_2 = scalar.constant 20520 : i32 + %v2_20_3 = scalar.constant 20545 : i32 + %w2_20 = vector.from_elements %v2_20_0, %v2_20_1, %v2_20_2, %v2_20_3 : vector<4xi32> + %o2_20 = index.constant 336 : index + vector.store %w2_20, %gv[%o2_20] : vector<4xi32>, view<512xi32> + %v2_21_0 = scalar.constant 20548 : i32 + %v2_21_1 = scalar.constant 20560 : i32 + %v2_21_2 = scalar.constant 20608 : i32 + %v2_21_3 = scalar.constant 20737 : i32 + %w2_21 = vector.from_elements %v2_21_0, %v2_21_1, %v2_21_2, %v2_21_3 : vector<4xi32> + %o2_21 = index.constant 340 : index + vector.store %w2_21, %gv[%o2_21] : vector<4xi32>, view<512xi32> + %v2_22_0 = scalar.constant 20740 : i32 + %v2_22_1 = scalar.constant 20752 : i32 + %v2_22_2 = scalar.constant 20757 : i32 + %v2_22_3 = scalar.constant 20800 : i32 + %w2_22 = vector.from_elements %v2_22_0, %v2_22_1, %v2_22_2, %v2_22_3 : vector<4xi32> + %o2_22 = index.constant 344 : index + vector.store %w2_22, %gv[%o2_22] : vector<4xi32>, view<512xi32> + %v2_23_0 = scalar.constant 20802 : i32 + %v2_23_1 = scalar.constant 20992 : i32 + %v2_23_2 = scalar.constant 21060 : i32 + %v2_23_3 = scalar.constant 21162 : i32 + %w2_23 = vector.from_elements %v2_23_0, %v2_23_1, %v2_23_2, %v2_23_3 : vector<4xi32> + %o2_23 = index.constant 348 : index + vector.store %w2_23, %gv[%o2_23] : vector<4xi32>, view<512xi32> + %v2_24_0 = scalar.constant 21505 : i32 + %v2_24_1 = scalar.constant 21508 : i32 + %v2_24_2 = scalar.constant 21520 : i32 + %v2_24_3 = scalar.constant 21537 : i32 + %w2_24 = vector.from_elements %v2_24_0, %v2_24_1, %v2_24_2, %v2_24_3 : vector<4xi32> + %o2_24 = index.constant 352 : index + vector.store %w2_24, %gv[%o2_24] : vector<4xi32>, view<512xi32> + %v2_25_0 = scalar.constant 21568 : i32 + %v2_25_1 = scalar.constant 21600 : i32 + %v2_25_2 = scalar.constant 21633 : i32 + %v2_25_3 = scalar.constant 21665 : i32 + %w2_25 = vector.from_elements %v2_25_0, %v2_25_1, %v2_25_2, %v2_25_3 : vector<4xi32> + %o2_25 = index.constant 356 : index + vector.store %w2_25, %gv[%o2_25] : vector<4xi32>, view<512xi32> + %v2_26_0 = scalar.constant 21760 : i32 + %v2_26_1 = scalar.constant 21768 : i32 + %v2_26_2 = scalar.constant 21888 : i32 + %v2_26_3 = scalar.constant 21896 : i32 + %w2_26 = vector.from_elements %v2_26_0, %v2_26_1, %v2_26_2, %v2_26_3 : vector<4xi32> + %o2_26 = index.constant 360 : index + vector.store %w2_26, %gv[%o2_26] : vector<4xi32>, view<512xi32> + %v2_27_0 = scalar.constant 22049 : i32 + %v2_27_1 = scalar.constant 22120 : i32 + %v2_27_2 = scalar.constant 22177 : i32 + %v2_27_3 = scalar.constant 22528 : i32 + %w2_27 = vector.from_elements %v2_27_0, %v2_27_1, %v2_27_2, %v2_27_3 : vector<4xi32> + %o2_27 = index.constant 364 : index + vector.store %w2_27, %gv[%o2_27] : vector<4xi32>, view<512xi32> + %v2_28_0 = scalar.constant 22548 : i32 + %v2_28_1 = scalar.constant 22593 : i32 + %v2_28_2 = scalar.constant 22608 : i32 + %v2_28_3 = scalar.constant 22681 : i32 + %w2_28 = vector.from_elements %v2_28_0, %v2_28_1, %v2_28_2, %v2_28_3 : vector<4xi32> + %o2_28 = index.constant 368 : index + vector.store %w2_28, %gv[%o2_28] : vector<4xi32>, view<512xi32> + %v2_29_0 = scalar.constant 22810 : i32 + %v2_29_1 = scalar.constant 22848 : i32 + %v2_29_2 = scalar.constant 22850 : i32 + %v2_29_3 = scalar.constant 23173 : i32 + %w2_29 = vector.from_elements %v2_29_0, %v2_29_1, %v2_29_2, %v2_29_3 : vector<4xi32> + %o2_29 = index.constant 372 : index + vector.store %w2_29, %gv[%o2_29] : vector<4xi32>, view<512xi32> + %v2_30_0 = scalar.constant 24577 : i32 + %v2_30_1 = scalar.constant 24580 : i32 + %v2_30_2 = scalar.constant 24592 : i32 + %v2_30_3 = scalar.constant 24640 : i32 + %w2_30 = vector.from_elements %v2_30_0, %v2_30_1, %v2_30_2, %v2_30_3 : vector<4xi32> + %o2_30 = index.constant 376 : index + vector.store %w2_30, %gv[%o2_30] : vector<4xi32>, view<512xi32> + %v2_31_0 = scalar.constant 24660 : i32 + %v2_31_1 = scalar.constant 24674 : i32 + %v2_31_2 = scalar.constant 24710 : i32 + %v2_31_3 = scalar.constant 24745 : i32 + %w2_31 = vector.from_elements %v2_31_0, %v2_31_1, %v2_31_2, %v2_31_3 : vector<4xi32> + %o2_31 = index.constant 380 : index + vector.store %w2_31, %gv[%o2_31] : vector<4xi32>, view<512xi32> + } + %k3 = index.constant 3 : index + %is3 = index.cmp eq, %chunk, %k3 : index + scf.if %is3 { + %v3_0_0 = scalar.constant 24832 : i32 + %v3_0_1 = scalar.constant 25124 : i32 + %v3_0_2 = scalar.constant 25162 : i32 + %v3_0_3 = scalar.constant 25234 : i32 + %w3_0 = vector.from_elements %v3_0_0, %v3_0_1, %v3_0_2, %v3_0_3 : vector<4xi32> + %o3_0 = index.constant 384 : index + vector.store %w3_0, %gv[%o3_0] : vector<4xi32>, view<512xi32> + %v3_1_0 = scalar.constant 25600 : i32 + %v3_1_1 = scalar.constant 25622 : i32 + %v3_1_2 = scalar.constant 25872 : i32 + %v3_1_3 = scalar.constant 25920 : i32 + %w3_1 = vector.from_elements %v3_1_0, %v3_1_1, %v3_1_2, %v3_1_3 : vector<4xi32> + %o3_1 = index.constant 388 : index + vector.store %w3_1, %gv[%o3_1] : vector<4xi32>, view<512xi32> + %v3_2_0 = scalar.constant 25925 : i32 + %v3_2_1 = scalar.constant 26020 : i32 + %v3_2_2 = scalar.constant 26625 : i32 + %v3_2_3 = scalar.constant 26730 : i32 + %w3_2 = vector.from_elements %v3_2_0, %v3_2_1, %v3_2_2, %v3_2_3 : vector<4xi32> + %o3_2 = index.constant 392 : index + vector.store %w3_2, %gv[%o3_2] : vector<4xi32>, view<512xi32> + %v3_3_0 = scalar.constant 26917 : i32 + %v3_3_1 = scalar.constant 27142 : i32 + %v3_3_2 = scalar.constant 27220 : i32 + %v3_3_3 = scalar.constant 27234 : i32 + %w3_3 = vector.from_elements %v3_3_0, %v3_3_1, %v3_3_2, %v3_3_3 : vector<4xi32> + %o3_3 = index.constant 396 : index + vector.store %w3_3, %gv[%o3_3] : vector<4xi32>, view<512xi32> + %v3_4_0 = scalar.constant 32768 : i32 + %v3_4_1 = scalar.constant 32770 : i32 + %v3_4_2 = scalar.constant 32773 : i32 + %v3_4_3 = scalar.constant 32776 : i32 + %w3_4 = vector.from_elements %v3_4_0, %v3_4_1, %v3_4_2, %v3_4_3 : vector<4xi32> + %o3_4 = index.constant 400 : index + vector.store %w3_4, %gv[%o3_4] : vector<4xi32>, view<512xi32> + %v3_5_0 = scalar.constant 32785 : i32 + %v3_5_1 = scalar.constant 32788 : i32 + %v3_5_2 = scalar.constant 32800 : i32 + %v3_5_3 = scalar.constant 32810 : i32 + %w3_5 = vector.from_elements %v3_5_0, %v3_5_1, %v3_5_2, %v3_5_3 : vector<4xi32> + %o3_5 = index.constant 404 : index + vector.store %w3_5, %gv[%o3_5] : vector<4xi32>, view<512xi32> + %v3_6_0 = scalar.constant 32833 : i32 + %v3_6_1 = scalar.constant 32836 : i32 + %v3_6_2 = scalar.constant 32848 : i32 + %v3_6_3 = scalar.constant 32896 : i32 + %w3_6 = vector.from_elements %v3_6_0, %v3_6_1, %v3_6_2, %v3_6_3 : vector<4xi32> + %o3_6 = index.constant 408 : index + vector.store %w3_6, %gv[%o3_6] : vector<4xi32>, view<512xi32> + %v3_7_0 = scalar.constant 32898 : i32 + %v3_7_1 = scalar.constant 32936 : i32 + %v3_7_2 = scalar.constant 32938 : i32 + %v3_7_3 = scalar.constant 33025 : i32 + %w3_7 = vector.from_elements %v3_7_0, %v3_7_1, %v3_7_2, %v3_7_3 : vector<4xi32> + %o3_7 = index.constant 412 : index + vector.store %w3_7, %gv[%o3_7] : vector<4xi32>, view<512xi32> + %v3_8_0 = scalar.constant 33028 : i32 + %v3_8_1 = scalar.constant 33030 : i32 + %v3_8_2 = scalar.constant 33040 : i32 + %v3_8_3 = scalar.constant 33088 : i32 + %w3_8 = vector.from_elements %v3_8_0, %v3_8_1, %v3_8_2, %v3_8_3 : vector<4xi32> + %o3_8 = index.constant 416 : index + vector.store %w3_8, %gv[%o3_8] : vector<4xi32>, view<512xi32> + %v3_9_0 = scalar.constant 33105 : i32 + %v3_9_1 = scalar.constant 33113 : i32 + %v3_9_2 = scalar.constant 33280 : i32 + %v3_9_3 = scalar.constant 33312 : i32 + %w3_9 = vector.from_elements %v3_9_0, %v3_9_1, %v3_9_2, %v3_9_3 : vector<4xi32> + %o3_9 = index.constant 420 : index + vector.store %w3_9, %gv[%o3_9] : vector<4xi32>, view<512xi32> + %v3_10_0 = scalar.constant 33408 : i32 + %v3_10_1 = scalar.constant 33410 : i32 + %v3_10_2 = scalar.constant 33440 : i32 + %v3_10_3 = scalar.constant 33448 : i32 + %w3_10 = vector.from_elements %v3_10_0, %v3_10_1, %v3_10_2, %v3_10_3 : vector<4xi32> + %o3_10 = index.constant 424 : index + vector.store %w3_10, %gv[%o3_10] : vector<4xi32>, view<512xi32> + %v3_11_0 = scalar.constant 33793 : i32 + %v3_11_1 = scalar.constant 33796 : i32 + %v3_11_2 = scalar.constant 33808 : i32 + %v3_11_3 = scalar.constant 33810 : i32 + %w3_11 = vector.from_elements %v3_11_0, %v3_11_1, %v3_11_2, %v3_11_3 : vector<4xi32> + %o3_11 = index.constant 428 : index + vector.store %w3_11, %gv[%o3_11] : vector<4xi32>, view<512xi32> + %v3_12_0 = scalar.constant 33813 : i32 + %v3_12_1 = scalar.constant 33856 : i32 + %v3_12_2 = scalar.constant 33888 : i32 + %v3_12_3 = scalar.constant 33929 : i32 + %w3_12 = vector.from_elements %v3_12_0, %v3_12_1, %v3_12_2, %v3_12_3 : vector<4xi32> + %o3_12 = index.constant 432 : index + vector.store %w3_12, %gv[%o3_12] : vector<4xi32>, view<512xi32> + %v3_13_0 = scalar.constant 34048 : i32 + %v3_13_1 = scalar.constant 34116 : i32 + %v3_13_2 = scalar.constant 34213 : i32 + %v3_13_3 = scalar.constant 34328 : i32 + %w3_13 = vector.from_elements %v3_13_0, %v3_13_1, %v3_13_2, %v3_13_3 : vector<4xi32> + %o3_13 = index.constant 436 : index + vector.store %w3_13, %gv[%o3_13] : vector<4xi32>, view<512xi32> + %v3_14_0 = scalar.constant 34410 : i32 + %v3_14_1 = scalar.constant 34816 : i32 + %v3_14_2 = scalar.constant 34824 : i32 + %v3_14_3 = scalar.constant 34853 : i32 + %w3_14 = vector.from_elements %v3_14_0, %v3_14_1, %v3_14_2, %v3_14_3 : vector<4xi32> + %o3_14 = index.constant 440 : index + vector.store %w3_14, %gv[%o3_14] : vector<4xi32>, view<512xi32> + %v3_15_0 = scalar.constant 34906 : i32 + %v3_15_1 = scalar.constant 34944 : i32 + %v3_15_2 = scalar.constant 34946 : i32 + %v3_15_3 = scalar.constant 34984 : i32 + %w3_15 = vector.from_elements %v3_15_0, %v3_15_1, %v3_15_2, %v3_15_3 : vector<4xi32> + %o3_15 = index.constant 444 : index + vector.store %w3_15, %gv[%o3_15] : vector<4xi32>, view<512xi32> + %v3_16_0 = scalar.constant 35078 : i32 + %v3_16_1 = scalar.constant 35362 : i32 + %v3_16_2 = scalar.constant 35456 : i32 + %v3_16_3 = scalar.constant 35464 : i32 + %w3_16 = vector.from_elements %v3_16_0, %v3_16_1, %v3_16_2, %v3_16_3 : vector<4xi32> + %o3_16 = index.constant 448 : index + vector.store %w3_16, %gv[%o3_16] : vector<4xi32>, view<512xi32> + %v3_17_0 = scalar.constant 35478 : i32 + %v3_17_1 = scalar.constant 35496 : i32 + %v3_17_2 = scalar.constant 36865 : i32 + %v3_17_3 = scalar.constant 36868 : i32 + %w3_17 = vector.from_elements %v3_17_0, %v3_17_1, %v3_17_2, %v3_17_3 : vector<4xi32> + %o3_17 = index.constant 452 : index + vector.store %w3_17, %gv[%o3_17] : vector<4xi32>, view<512xi32> + %v3_18_0 = scalar.constant 36880 : i32 + %v3_18_1 = scalar.constant 36928 : i32 + %v3_18_2 = scalar.constant 36950 : i32 + %v3_18_3 = scalar.constant 36996 : i32 + %w3_18 = vector.from_elements %v3_18_0, %v3_18_1, %v3_18_2, %v3_18_3 : vector<4xi32> + %o3_18 = index.constant 456 : index + vector.store %w3_18, %gv[%o3_18] : vector<4xi32>, view<512xi32> + %v3_19_0 = scalar.constant 37120 : i32 + %v3_19_1 = scalar.constant 37154 : i32 + %v3_19_2 = scalar.constant 37220 : i32 + %v3_19_3 = scalar.constant 37462 : i32 + %w3_19 = vector.from_elements %v3_19_0, %v3_19_1, %v3_19_2, %v3_19_3 : vector<4xi32> + %o3_19 = index.constant 460 : index + vector.store %w3_19, %gv[%o3_19] : vector<4xi32>, view<512xi32> + %v3_20_0 = scalar.constant 37513 : i32 + %v3_20_1 = scalar.constant 37888 : i32 + %v3_20_2 = scalar.constant 37893 : i32 + %v3_20_3 = scalar.constant 37956 : i32 + %w3_20 = vector.from_elements %v3_20_0, %v3_20_1, %v3_20_2, %v3_20_3 : vector<4xi32> + %o3_20 = index.constant 464 : index + vector.store %w3_20, %gv[%o3_20] : vector<4xi32>, view<512xi32> + %v3_21_0 = scalar.constant 37968 : i32 + %v3_21_1 = scalar.constant 37976 : i32 + %v3_21_2 = scalar.constant 38185 : i32 + %v3_21_3 = scalar.constant 38288 : i32 + %w3_21 = vector.from_elements %v3_21_0, %v3_21_1, %v3_21_2, %v3_21_3 : vector<4xi32> + %o3_21 = index.constant 468 : index + vector.store %w3_21, %gv[%o3_21] : vector<4xi32>, view<512xi32> + %v3_22_0 = scalar.constant 38290 : i32 + %v3_22_1 = scalar.constant 38465 : i32 + %v3_22_2 = scalar.constant 38993 : i32 + %v3_22_3 = scalar.constant 39078 : i32 + %w3_22 = vector.from_elements %v3_22_0, %v3_22_1, %v3_22_2, %v3_22_3 : vector<4xi32> + %o3_22 = index.constant 472 : index + vector.store %w3_22, %gv[%o3_22] : vector<4xi32>, view<512xi32> + %v3_23_0 = scalar.constant 39241 : i32 + %v3_23_1 = scalar.constant 39445 : i32 + %v3_23_2 = scalar.constant 39520 : i32 + %v3_23_3 = scalar.constant 40960 : i32 + %w3_23 = vector.from_elements %v3_23_0, %v3_23_1, %v3_23_2, %v3_23_3 : vector<4xi32> + %o3_23 = index.constant 476 : index + vector.store %w3_23, %gv[%o3_23] : vector<4xi32>, view<512xi32> + %v3_24_0 = scalar.constant 40962 : i32 + %v3_24_1 = scalar.constant 40968 : i32 + %v3_24_2 = scalar.constant 40970 : i32 + %v3_24_3 = scalar.constant 40992 : i32 + %w3_24 = vector.from_elements %v3_24_0, %v3_24_1, %v3_24_2, %v3_24_3 : vector<4xi32> + %o3_24 = index.constant 480 : index + vector.store %w3_24, %gv[%o3_24] : vector<4xi32>, view<512xi32> + %v3_25_0 = scalar.constant 41002 : i32 + %v3_25_1 = scalar.constant 41120 : i32 + %v3_25_2 = scalar.constant 41297 : i32 + %v3_25_3 = scalar.constant 41305 : i32 + %w3_25 = vector.from_elements %v3_25_0, %v3_25_1, %v3_25_2, %v3_25_3 : vector<4xi32> + %o3_25 = index.constant 484 : index + vector.store %w3_25, %gv[%o3_25] : vector<4xi32>, view<512xi32> + %v3_26_0 = scalar.constant 41382 : i32 + %v3_26_1 = scalar.constant 41472 : i32 + %v3_26_2 = scalar.constant 41474 : i32 + %v3_26_3 = scalar.constant 41480 : i32 + %w3_26 = vector.from_elements %v3_26_0, %v3_26_1, %v3_26_2, %v3_26_3 : vector<4xi32> + %o3_26 = index.constant 488 : index + vector.store %w3_26, %gv[%o3_26] : vector<4xi32>, view<512xi32> + %v3_27_0 = scalar.constant 41514 : i32 + %v3_27_1 = scalar.constant 41600 : i32 + %v3_27_2 = scalar.constant 41632 : i32 + %v3_27_3 = scalar.constant 42048 : i32 + %w3_27 = vector.from_elements %v3_27_0, %v3_27_1, %v3_27_2, %v3_27_3 : vector<4xi32> + %o3_27 = index.constant 492 : index + vector.store %w3_27, %gv[%o3_27] : vector<4xi32>, view<512xi32> + %v3_28_0 = scalar.constant 42133 : i32 + %v3_28_1 = scalar.constant 42597 : i32 + %v3_28_2 = scalar.constant 42648 : i32 + %v3_28_3 = scalar.constant 43018 : i32 + %w3_28 = vector.from_elements %v3_28_0, %v3_28_1, %v3_28_2, %v3_28_3 : vector<4xi32> + %o3_28 = index.constant 496 : index + vector.store %w3_28, %gv[%o3_28] : vector<4xi32>, view<512xi32> + %v3_29_0 = scalar.constant 43040 : i32 + %v3_29_1 = scalar.constant 43042 : i32 + %v3_29_2 = scalar.constant 43048 : i32 + %v3_29_3 = scalar.constant 43168 : i32 + %w3_29 = vector.from_elements %v3_29_0, %v3_29_1, %v3_29_2, %v3_29_3 : vector<4xi32> + %o3_29 = index.constant 500 : index + vector.store %w3_29, %gv[%o3_29] : vector<4xi32>, view<512xi32> + %v3_30_0 = scalar.constant 43176 : i32 + %v3_30_1 = scalar.constant 43268 : i32 + %v3_30_2 = scalar.constant 43396 : i32 + %v3_30_3 = scalar.constant 43398 : i32 + %w3_30 = vector.from_elements %v3_30_0, %v3_30_1, %v3_30_2, %v3_30_3 : vector<4xi32> + %o3_30 = index.constant 504 : index + vector.store %w3_30, %gv[%o3_30] : vector<4xi32>, view<512xi32> + %v3_31_0 = scalar.constant 43560 : i32 + %v3_31_1 = scalar.constant 43562 : i32 + %v3_31_2 = scalar.constant 43665 : i32 + %v3_31_3 = scalar.constant 43690 : i32 + %w3_31 = vector.from_elements %v3_31_0, %v3_31_1, %v3_31_2, %v3_31_3 : vector<4xi32> + %o3_31 = index.constant 508 : index + vector.store %w3_31, %gv[%o3_31] : vector<4xi32>, view<512xi32> + } + func.return +} + +// IQ2_XXS / IQ2_XS: a 256-value block is 8 groups of 32, each 4 grid slots of 8 values; lane l16 +// owns group l16 / 2, slots 2 (l16 % 2) and +1: values 32 (l16 / 2) + 16 (l16 % 2) .. +15, one +// scale. The grid (16-bit codes, 2 bits per value: level 8 + 17c + (c >> 1) = 8, 25, 43) is +// staged in workgroup memory like IQ3_S's; the sign byte is ksigns_iq2xs = 7 bits + parity. +func.def inline @ggml_kquant_iq2_slot_values(%code: i32, %signs7: i32) -> (vector<8xf32>) { + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c7_i32 = scalar.constant 7 : i32 + %p4 = scalar.shrui %signs7, %c4_i32 : i32 + %x4 = scalar.xori %signs7, %p4 : i32 + %p2 = scalar.shrui %x4, %c2_i32 : i32 + %x2 = scalar.xori %x4, %p2 : i32 + %p1 = scalar.shrui %x2, %c1_i32 : i32 + %x1 = scalar.xori %x2, %p1 : i32 + %parity = scalar.andi %x1, %c1_i32 : i32 + %high = scalar.shli %parity, %c7_i32 : i32 + %signs8 = scalar.ori %signs7, %high : i32 + %s0 = scalar.constant 0 : i32 + %s2 = scalar.constant 2 : i32 + %s4 = scalar.constant 4 : i32 + %s6 = scalar.constant 6 : i32 + %s8 = scalar.constant 8 : i32 + %s10 = scalar.constant 10 : i32 + %s12 = scalar.constant 12 : i32 + %s14 = scalar.constant 14 : i32 + %s3 = scalar.constant 3 : i32 + %s5 = scalar.constant 5 : i32 + %shift2 = vector.from_elements %s0, %s2, %s4, %s6, %s8, %s10, %s12, %s14 : vector<8xi32> + %shift1 = vector.from_elements %s0, %c1_i32, %s2, %s3, %s4, %s5, %s6, %c7_i32 : vector<8xi32> + %three = vector.splat %s3 : vector<8xi32> + %one = vector.splat %c1_i32 : vector<8xi32> + %c8v = vector.splat %s8 : vector<8xi32> + %c17_i32 = scalar.constant 17 : i32 + %c17v = vector.splat %c17_i32 : vector<8xi32> + %cv = vector.splat %code : vector<8xi32> + %cs = vector.shrui %cv, %shift2 : vector<8xi32> + %c = vector.andi %cs, %three : vector<8xi32> + %c17 = vector.muli %c, %c17v : vector<8xi32> + %chi = vector.shrui %c, %one : vector<8xi32> + %lv0 = vector.addi %c17, %chi : vector<8xi32> + %lv = vector.addi %lv0, %c8v : vector<8xi32> + %sv = vector.splat %signs8 : vector<8xi32> + %sb0 = vector.shrui %sv, %shift1 : vector<8xi32> + %sb = vector.andi %sb0, %one : vector<8xi32> + %sb2 = vector.shli %sb, %one : vector<8xi32> + %sgn = vector.subi %one, %sb2 : vector<8xi32> + %v = vector.muli %lv, %sgn : vector<8xi32> + %vf = vector.sitofp %v : vector<8xi32> to vector<8xf32> + func.return %vf : vector<8xf32> +} + +func.def inline @ggml_kquant_iq2_join16(%a: vector<8xf32>, %b: vector<8xf32>) -> (vector<16xf32>) { + %a0 = vector.extract %a[0] : vector<8xf32> -> f32 + %a1 = vector.extract %a[1] : vector<8xf32> -> f32 + %a2 = vector.extract %a[2] : vector<8xf32> -> f32 + %a3 = vector.extract %a[3] : vector<8xf32> -> f32 + %a4 = vector.extract %a[4] : vector<8xf32> -> f32 + %a5 = vector.extract %a[5] : vector<8xf32> -> f32 + %a6 = vector.extract %a[6] : vector<8xf32> -> f32 + %a7 = vector.extract %a[7] : vector<8xf32> -> f32 + %b0 = vector.extract %b[0] : vector<8xf32> -> f32 + %b1 = vector.extract %b[1] : vector<8xf32> -> f32 + %b2 = vector.extract %b[2] : vector<8xf32> -> f32 + %b3 = vector.extract %b[3] : vector<8xf32> -> f32 + %b4 = vector.extract %b[4] : vector<8xf32> -> f32 + %b5 = vector.extract %b[5] : vector<8xf32> -> f32 + %b6 = vector.extract %b[6] : vector<8xf32> -> f32 + %b7 = vector.extract %b[7] : vector<8xf32> -> f32 + %r = vector.from_elements %a0, %a1, %a2, %a3, %a4, %a5, %a6, %a7, %b0, %b1, %b2, %b3, %b4, %b5, %b6, %b7 : vector<16xf32> + func.return %r : vector<16xf32> +} + +// IQ2_XXS (66 bytes: d, qs[32] u16): group g = 8 bytes at 2 + 8g: slot indices (bytes 0..3), +// then u32 aux: four 7-bit sign groups, scale nibble in bits 28..31 (d * (0.5 + s) * 0.25). +func.def inline @ggml_kquant_iq2xxs_lane_parts(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c7_i32 = scalar.constant 7 : i32 + %c16_i32 = scalar.constant 16 : i32 + %c28_i32 = scalar.constant 28 : i32 + %c127_i32 = scalar.constant 127 : i32 + %c05 = scalar.constant 0.5 : f32 + %c025 = scalar.constant 0.25 : f32 + %block_bytes = index.constant 66 : offset + %zero_offset = index.constant 0 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %g = index.div %l, %c2 : index + %h = index.rem %l, %c2 : index + %s0 = index.mul %h, %c2 : index + %s1 = index.add %s0, %c1 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %hv = buffer.view %weight[%block_base] : buffer -> view<33xf16> + %wv = buffer.view %weight[%block_base] : buffer -> view<33xi16> + %bv = buffer.view %weight[%block_base] : buffer -> view<66xi8> + %gv = buffer.view %grid[%zero_offset] : buffer -> view<512xi32> + %d_f16 = view.load %hv[%c0] : view<33xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %g8 = index.mul %g, %c8 : index + %ib0 = index.add %g8, %c2 : index + %ia = index.add %ib0, %s0 : index + %ib = index.add %ib0, %s1 : index + %ia_i8 = view.load %bv[%ia] : view<66xi8> -> i8 + %ib_i8 = view.load %bv[%ib] : view<66xi8> -> i8 + %ga = scalar.extui %ia_i8 : i8 to i32 + %gb = scalar.extui %ib_i8 : i8 to i32 + %ga_x = index.cast %ga : i32 to index + %gb_x = index.cast %gb : i32 to index + %ga_b = index.assume %ga_x [range(%ga_x, 0, 255)] : index + %gb_b = index.assume %gb_x [range(%gb_x, 0, 255)] : index + %code_a = view.load %gv[%ga_b] : view<512xi32> -> i32 + %code_b = view.load %gv[%gb_b] : view<512xi32> -> i32 + %g4 = index.mul %g, %c4 : index + %w0_at = index.add %g4, %c3 : index + %w1_at = index.add %g4, %c4 : index + %w0_i16 = view.load %wv[%w0_at] : view<33xi16> -> i16 + %w1_i16 = view.load %wv[%w1_at] : view<33xi16> -> i16 + %w0 = scalar.extui %w0_i16 : i16 to i32 + %w1 = scalar.extui %w1_i16 : i16 to i32 + %w1s = scalar.shli %w1, %c16_i32 : i32 + %aux = scalar.ori %w0, %w1s : i32 + %sc4 = scalar.shrui %aux, %c28_i32 : i32 + %sc_f = scalar.uitofp %sc4 : i32 to f32 + %sc_p = scalar.addf %sc_f, %c05 : f32 + %ds = scalar.mulf %d, %sc_p : f32 + %scale = scalar.mulf %ds, %c025 : f32 + %s0_i32 = index.cast %s0 : index to i32 + %s1_i32 = index.cast %s1 : index to i32 + %sha = scalar.muli %s0_i32, %c7_i32 : i32 + %shb = scalar.muli %s1_i32, %c7_i32 : i32 + %sa0 = scalar.shrui %aux, %sha : i32 + %sb0 = scalar.shrui %aux, %shb : i32 + %sa = scalar.andi %sa0, %c127_i32 : i32 + %sb = scalar.andi %sb0, %c127_i32 : i32 + %va = func.call @ggml_kquant_iq2_slot_values(%code_a, %sa) : (i32, i32) -> (vector<8xf32>) + %vb = func.call @ggml_kquant_iq2_slot_values(%code_b, %sb) : (i32, i32) -> (vector<8xf32>) + %v = func.call @ggml_kquant_iq2_join16(%va, %vb) : (vector<8xf32>, vector<8xf32>) -> (vector<16xf32>) + %g32 = index.mul %g, %c32 : index + %h16 = index.mul %h, %c16 : index + %p0 = index.add %g32, %h16 : index + func.return %v, %scale, %p0 : vector<16xf32>, f32, index +} + +// IQ2_XS (74 bytes: d, qs[32] u16, scales[8]): slot q = qs[4g + s]: grid index q & 511, 7 sign +// bits q >> 9; scale nibble h of scales[g] (slots 2h, 2h + 1). +func.def inline @ggml_kquant_iq2xs_lane_parts(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c66 = index.constant 66 : index + %c4_i32 = scalar.constant 4 : i32 + %c9_i32 = scalar.constant 9 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c127_i32 = scalar.constant 127 : i32 + %c511_i32 = scalar.constant 511 : i32 + %c05 = scalar.constant 0.5 : f32 + %c025 = scalar.constant 0.25 : f32 + %block_bytes = index.constant 74 : offset + %zero_offset = index.constant 0 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %g = index.div %l, %c2 : index + %h = index.rem %l, %c2 : index + %s0 = index.mul %h, %c2 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %hv = buffer.view %weight[%block_base] : buffer -> view<37xf16> + %wv = buffer.view %weight[%block_base] : buffer -> view<37xi16> + %bv = buffer.view %weight[%block_base] : buffer -> view<74xi8> + %gv = buffer.view %grid[%zero_offset] : buffer -> view<512xi32> + %d_f16 = view.load %hv[%c0] : view<37xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %g4 = index.mul %g, %c4 : index + %qa_at0 = index.add %g4, %s0 : index + %qa_at = index.add %qa_at0, %c1 : index + %qb_at = index.add %qa_at, %c1 : index + %qa_i16 = view.load %wv[%qa_at] : view<37xi16> -> i16 + %qb_i16 = view.load %wv[%qb_at] : view<37xi16> -> i16 + %qa = scalar.extui %qa_i16 : i16 to i32 + %qb = scalar.extui %qb_i16 : i16 to i32 + %ga = scalar.andi %qa, %c511_i32 : i32 + %gb = scalar.andi %qb, %c511_i32 : i32 + %sa0 = scalar.shrui %qa, %c9_i32 : i32 + %sb0 = scalar.shrui %qb, %c9_i32 : i32 + %sa = scalar.andi %sa0, %c127_i32 : i32 + %sb = scalar.andi %sb0, %c127_i32 : i32 + %ga_x = index.cast %ga : i32 to index + %gb_x = index.cast %gb : i32 to index + %ga_b = index.assume %ga_x [range(%ga_x, 0, 511)] : index + %gb_b = index.assume %gb_x [range(%gb_x, 0, 511)] : index + %code_a = view.load %gv[%ga_b] : view<512xi32> -> i32 + %code_b = view.load %gv[%gb_b] : view<512xi32> -> i32 + %sc_at = index.add %c66, %g : index + %sc_i8 = view.load %bv[%sc_at] : view<74xi8> -> i8 + %sc = scalar.extui %sc_i8 : i8 to i32 + %h_i32 = index.cast %h : index to i32 + %nsh = scalar.muli %h_i32, %c4_i32 : i32 + %nib0 = scalar.shrui %sc, %nsh : i32 + %nib = scalar.andi %nib0, %c15_i32 : i32 + %nib_f = scalar.uitofp %nib : i32 to f32 + %nib_p = scalar.addf %nib_f, %c05 : f32 + %ds = scalar.mulf %d, %nib_p : f32 + %scale = scalar.mulf %ds, %c025 : f32 + %va = func.call @ggml_kquant_iq2_slot_values(%code_a, %sa) : (i32, i32) -> (vector<8xf32>) + %vb = func.call @ggml_kquant_iq2_slot_values(%code_b, %sb) : (i32, i32) -> (vector<8xf32>) + %v = func.call @ggml_kquant_iq2_join16(%va, %vb) : (vector<8xf32>, vector<8xf32>) -> (vector<16xf32>) + %g32 = index.mul %g, %c32 : index + %h16 = index.mul %h, %c16 : index + %p0 = index.add %g32, %h16 : index + func.return %v, %scale, %p0 : vector<16xf32>, f32, index +} + +func.def inline @ggml_kquant_iq2xxs_lane_dot(%grid: buffer, %weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %zero_scalar = scalar.constant 0.0 : f32 + %x_block_bytes = index.constant 1024 : offset + %v, %scale, %p = func.call @ggml_kquant_iq2xxs_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %x = vector.load %xv[%p] : view<256xf32> -> vector<16xf32> + %vx = vector.mulf %v, %x : vector<16xf32> + %sum = vector.reduce %vx, %zero_scalar : vector<16xf32>, f32 + %result = scalar.mulf %scale, %sum : f32 + func.return %result : f32 +} + +func.def inline @ggml_kquant_iq2xxs_lane_weights(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c12 = index.constant 12 : index + %v, %scale, %p0 = func.call @ggml_kquant_iq2xxs_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %sv = vector.splat %scale : vector<16xf32> + %w = vector.mulf %v, %sv : vector<16xf32> + %p1 = index.add %p0, %c4 : index + %p2 = index.add %p0, %c8 : index + %p3 = index.add %p0, %c12 : index + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +func.def inline @ggml_kquant_iq2xs_lane_dot(%grid: buffer, %weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %zero_scalar = scalar.constant 0.0 : f32 + %x_block_bytes = index.constant 1024 : offset + %v, %scale, %p = func.call @ggml_kquant_iq2xs_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %x = vector.load %xv[%p] : view<256xf32> -> vector<16xf32> + %vx = vector.mulf %v, %x : vector<16xf32> + %sum = vector.reduce %vx, %zero_scalar : vector<16xf32>, f32 + %result = scalar.mulf %scale, %sum : f32 + func.return %result : f32 +} + +func.def inline @ggml_kquant_iq2xs_lane_weights(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c12 = index.constant 12 : index + %v, %scale, %p0 = func.call @ggml_kquant_iq2xs_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %sv = vector.splat %scale : vector<16xf32> + %w = vector.mulf %v, %sv : vector<16xf32> + %p1 = index.add %p0, %c4 : index + %p2 = index.add %p0, %c8 : index + %p3 = index.add %p0, %c12 : index + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +// Stages the IQ3_XXS grid, one 12-bit code per word; subgroup %chunk (0..3) writes words 64 %chunk .. +63. +func.def inline @ggml_kquant_iq3xxs_grid_fill(%grid: buffer, %chunk: index) { + %zero_offset = index.constant 0 : offset + %gv = buffer.view %grid[%zero_offset] : buffer -> view<512xi32> + %k0 = index.constant 0 : index + %is0 = index.cmp eq, %chunk, %k0 : index + scf.if %is0 { + %v0_0_0 = scalar.constant 0 : i32 + %v0_0_1 = scalar.constant 2 : i32 + %v0_0_2 = scalar.constant 4 : i32 + %v0_0_3 = scalar.constant 9 : i32 + %w0_0 = vector.from_elements %v0_0_0, %v0_0_1, %v0_0_2, %v0_0_3 : vector<4xi32> + %o0_0 = index.constant 0 : index + vector.store %w0_0, %gv[%o0_0] : vector<4xi32>, view<512xi32> + %v0_1_0 = scalar.constant 11 : i32 + %v0_1_1 = scalar.constant 15 : i32 + %v0_1_2 = scalar.constant 16 : i32 + %v0_1_3 = scalar.constant 18 : i32 + %w0_1 = vector.from_elements %v0_1_0, %v0_1_1, %v0_1_2, %v0_1_3 : vector<4xi32> + %o0_1 = index.constant 4 : index + vector.store %w0_1, %gv[%o0_1] : vector<4xi32>, view<512xi32> + %v0_2_0 = scalar.constant 25 : i32 + %v0_2_1 = scalar.constant 34 : i32 + %v0_2_2 = scalar.constant 59 : i32 + %v0_2_3 = scalar.constant 61 : i32 + %w0_2 = vector.from_elements %v0_2_0, %v0_2_1, %v0_2_2, %v0_2_3 : vector<4xi32> + %o0_2 = index.constant 8 : index + vector.store %w0_2, %gv[%o0_2] : vector<4xi32>, view<512xi32> + %v0_3_0 = scalar.constant 65 : i32 + %v0_3_1 = scalar.constant 67 : i32 + %v0_3_2 = scalar.constant 72 : i32 + %v0_3_3 = scalar.constant 74 : i32 + %w0_3 = vector.from_elements %v0_3_0, %v0_3_1, %v0_3_2, %v0_3_3 : vector<4xi32> + %o0_3 = index.constant 12 : index + vector.store %w0_3, %gv[%o0_3] : vector<4xi32>, view<512xi32> + %v0_4_0 = scalar.constant 81 : i32 + %v0_4_1 = scalar.constant 85 : i32 + %v0_4_2 = scalar.constant 88 : i32 + %v0_4_3 = scalar.constant 90 : i32 + %w0_4 = vector.from_elements %v0_4_0, %v0_4_1, %v0_4_2, %v0_4_3 : vector<4xi32> + %o0_4 = index.constant 16 : index + vector.store %w0_4, %gv[%o0_4] : vector<4xi32>, view<512xi32> + %v0_5_0 = scalar.constant 97 : i32 + %v0_5_1 = scalar.constant 108 : i32 + %v0_5_2 = scalar.constant 120 : i32 + %v0_5_3 = scalar.constant 128 : i32 + %w0_5 = vector.from_elements %v0_5_0, %v0_5_1, %v0_5_2, %v0_5_3 : vector<4xi32> + %o0_5 = index.constant 20 : index + vector.store %w0_5, %gv[%o0_5] : vector<4xi32>, view<512xi32> + %v0_6_0 = scalar.constant 130 : i32 + %v0_6_1 = scalar.constant 132 : i32 + %v0_6_2 = scalar.constant 137 : i32 + %v0_6_3 = scalar.constant 144 : i32 + %w0_6 = vector.from_elements %v0_6_0, %v0_6_1, %v0_6_2, %v0_6_3 : vector<4xi32> + %o0_6 = index.constant 24 : index + vector.store %w0_6, %gv[%o0_6] : vector<4xi32>, view<512xi32> + %v0_7_0 = scalar.constant 146 : i32 + %v0_7_1 = scalar.constant 153 : i32 + %v0_7_2 = scalar.constant 155 : i32 + %v0_7_3 = scalar.constant 159 : i32 + %w0_7 = vector.from_elements %v0_7_0, %v0_7_1, %v0_7_2, %v0_7_3 : vector<4xi32> + %o0_7 = index.constant 28 : index + vector.store %w0_7, %gv[%o0_7] : vector<4xi32>, view<512xi32> + %v0_8_0 = scalar.constant 169 : i32 + %v0_8_1 = scalar.constant 175 : i32 + %v0_8_2 = scalar.constant 189 : i32 + %v0_8_3 = scalar.constant 193 : i32 + %w0_8 = vector.from_elements %v0_8_0, %v0_8_1, %v0_8_2, %v0_8_3 : vector<4xi32> + %o0_8 = index.constant 32 : index + vector.store %w0_8, %gv[%o0_8] : vector<4xi32>, view<512xi32> + %v0_9_0 = scalar.constant 199 : i32 + %v0_9_1 = scalar.constant 200 : i32 + %v0_9_2 = scalar.constant 202 : i32 + %v0_9_3 = scalar.constant 213 : i32 + %w0_9 = vector.from_elements %v0_9_0, %v0_9_1, %v0_9_2, %v0_9_3 : vector<4xi32> + %o0_9 = index.constant 36 : index + vector.store %w0_9, %gv[%o0_9] : vector<4xi32>, view<512xi32> + %v0_10_0 = scalar.constant 248 : i32 + %v0_10_1 = scalar.constant 267 : i32 + %v0_10_2 = scalar.constant 287 : i32 + %v0_10_3 = scalar.constant 292 : i32 + %w0_10 = vector.from_elements %v0_10_0, %v0_10_1, %v0_10_2, %v0_10_3 : vector<4xi32> + %o0_10 = index.constant 40 : index + vector.store %w0_10, %gv[%o0_10] : vector<4xi32>, view<512xi32> + %v0_11_0 = scalar.constant 303 : i32 + %v0_11_1 = scalar.constant 315 : i32 + %v0_11_2 = scalar.constant 317 : i32 + %v0_11_3 = scalar.constant 321 : i32 + %w0_11 = vector.from_elements %v0_11_0, %v0_11_1, %v0_11_2, %v0_11_3 : vector<4xi32> + %o0_11 = index.constant 44 : index + vector.store %w0_11, %gv[%o0_11] : vector<4xi32>, view<512xi32> + %v0_12_0 = scalar.constant 327 : i32 + %v0_12_1 = scalar.constant 346 : i32 + %v0_12_2 = scalar.constant 362 : i32 + %v0_12_3 = scalar.constant 413 : i32 + %w0_12 = vector.from_elements %v0_12_0, %v0_12_1, %v0_12_2, %v0_12_3 : vector<4xi32> + %o0_12 = index.constant 48 : index + vector.store %w0_12, %gv[%o0_12] : vector<4xi32>, view<512xi32> + %v0_13_0 = scalar.constant 436 : i32 + %v0_13_1 = scalar.constant 456 : i32 + %v0_13_2 = scalar.constant 460 : i32 + %v0_13_3 = scalar.constant 462 : i32 + %w0_13 = vector.from_elements %v0_13_0, %v0_13_1, %v0_13_2, %v0_13_3 : vector<4xi32> + %o0_13 = index.constant 52 : index + vector.store %w0_13, %gv[%o0_13] : vector<4xi32>, view<512xi32> + %v0_14_0 = scalar.constant 483 : i32 + %v0_14_1 = scalar.constant 497 : i32 + %v0_14_2 = scalar.constant 513 : i32 + %v0_14_3 = scalar.constant 515 : i32 + %w0_14 = vector.from_elements %v0_14_0, %v0_14_1, %v0_14_2, %v0_14_3 : vector<4xi32> + %o0_14 = index.constant 56 : index + vector.store %w0_14, %gv[%o0_14] : vector<4xi32>, view<512xi32> + %v0_15_0 = scalar.constant 520 : i32 + %v0_15_1 = scalar.constant 522 : i32 + %v0_15_2 = scalar.constant 529 : i32 + %v0_15_3 = scalar.constant 531 : i32 + %w0_15 = vector.from_elements %v0_15_0, %v0_15_1, %v0_15_2, %v0_15_3 : vector<4xi32> + %o0_15 = index.constant 60 : index + vector.store %w0_15, %gv[%o0_15] : vector<4xi32>, view<512xi32> + } + %k1 = index.constant 1 : index + %is1 = index.cmp eq, %chunk, %k1 : index + scf.if %is1 { + %v1_0_0 = scalar.constant 536 : i32 + %v1_0_1 = scalar.constant 538 : i32 + %v1_0_2 = scalar.constant 540 : i32 + %v1_0_3 = scalar.constant 551 : i32 + %w1_0 = vector.from_elements %v1_0_0, %v1_0_1, %v1_0_2, %v1_0_3 : vector<4xi32> + %o1_0 = index.constant 64 : index + vector.store %w1_0, %gv[%o1_0] : vector<4xi32>, view<512xi32> + %v1_1_0 = scalar.constant 552 : i32 + %v1_1_1 = scalar.constant 576 : i32 + %v1_1_2 = scalar.constant 578 : i32 + %v1_1_3 = scalar.constant 585 : i32 + %w1_1 = vector.from_elements %v1_1_0, %v1_1_1, %v1_1_2, %v1_1_3 : vector<4xi32> + %o1_1 = index.constant 68 : index + vector.store %w1_1, %gv[%o1_1] : vector<4xi32>, view<512xi32> + %v1_2_0 = scalar.constant 592 : i32 + %v1_2_1 = scalar.constant 594 : i32 + %v1_2_2 = scalar.constant 641 : i32 + %v1_2_3 = scalar.constant 643 : i32 + %w1_2 = vector.from_elements %v1_2_0, %v1_2_1, %v1_2_2, %v1_2_3 : vector<4xi32> + %o1_2 = index.constant 72 : index + vector.store %w1_2, %gv[%o1_2] : vector<4xi32>, view<512xi32> + %v1_3_0 = scalar.constant 648 : i32 + %v1_3_1 = scalar.constant 650 : i32 + %v1_3_2 = scalar.constant 657 : i32 + %v1_3_3 = scalar.constant 664 : i32 + %w1_3 = vector.from_elements %v1_3_0, %v1_3_1, %v1_3_2, %v1_3_3 : vector<4xi32> + %o1_3 = index.constant 76 : index + vector.store %w1_3, %gv[%o1_3] : vector<4xi32>, view<512xi32> + %v1_4_0 = scalar.constant 698 : i32 + %v1_4_1 = scalar.constant 704 : i32 + %v1_4_2 = scalar.constant 706 : i32 + %v1_4_3 = scalar.constant 720 : i32 + %w1_4 = vector.from_elements %v1_4_0, %v1_4_1, %v1_4_2, %v1_4_3 : vector<4xi32> + %o1_4 = index.constant 80 : index + vector.store %w1_4, %gv[%o1_4] : vector<4xi32>, view<512xi32> + %v1_5_0 = scalar.constant 729 : i32 + %v1_5_1 = scalar.constant 742 : i32 + %v1_5_2 = scalar.constant 758 : i32 + %v1_5_3 = scalar.constant 769 : i32 + %w1_5 = vector.from_elements %v1_5_0, %v1_5_1, %v1_5_2, %v1_5_3 : vector<4xi32> + %o1_5 = index.constant 84 : index + vector.store %w1_5, %gv[%o1_5] : vector<4xi32>, view<512xi32> + %v1_6_0 = scalar.constant 773 : i32 + %v1_6_1 = scalar.constant 808 : i32 + %v1_6_2 = scalar.constant 848 : i32 + %v1_6_3 = scalar.constant 852 : i32 + %w1_6 = vector.from_elements %v1_6_0, %v1_6_1, %v1_6_2, %v1_6_3 : vector<4xi32> + %o1_6 = index.constant 88 : index + vector.store %w1_6, %gv[%o1_6] : vector<4xi32>, view<512xi32> + %v1_7_0 = scalar.constant 870 : i32 + %v1_7_1 = scalar.constant 889 : i32 + %v1_7_2 = scalar.constant 901 : i32 + %v1_7_3 = scalar.constant 978 : i32 + %w1_7 = vector.from_elements %v1_7_0, %v1_7_1, %v1_7_2, %v1_7_3 : vector<4xi32> + %o1_7 = index.constant 92 : index + vector.store %w1_7, %gv[%o1_7] : vector<4xi32>, view<512xi32> + %v1_8_0 = scalar.constant 992 : i32 + %v1_8_1 = scalar.constant 1024 : i32 + %v1_8_2 = scalar.constant 1026 : i32 + %v1_8_3 = scalar.constant 1033 : i32 + %w1_8 = vector.from_elements %v1_8_0, %v1_8_1, %v1_8_2, %v1_8_3 : vector<4xi32> + %o1_8 = index.constant 96 : index + vector.store %w1_8, %gv[%o1_8] : vector<4xi32>, view<512xi32> + %v1_9_0 = scalar.constant 1035 : i32 + %v1_9_1 = scalar.constant 1040 : i32 + %v1_9_2 = scalar.constant 1042 : i32 + %v1_9_3 = scalar.constant 1046 : i32 + %w1_9 = vector.from_elements %v1_9_0, %v1_9_1, %v1_9_2, %v1_9_3 : vector<4xi32> + %o1_9 = index.constant 100 : index + vector.store %w1_9, %gv[%o1_9] : vector<4xi32>, view<512xi32> + %v1_10_0 = scalar.constant 1049 : i32 + %v1_10_1 = scalar.constant 1058 : i32 + %v1_10_2 = scalar.constant 1089 : i32 + %v1_10_3 = scalar.constant 1091 : i32 + %w1_10 = vector.from_elements %v1_10_0, %v1_10_1, %v1_10_2, %v1_10_3 : vector<4xi32> + %o1_10 = index.constant 104 : index + vector.store %w1_10, %gv[%o1_10] : vector<4xi32>, view<512xi32> + %v1_11_0 = scalar.constant 1093 : i32 + %v1_11_1 = scalar.constant 1096 : i32 + %v1_11_2 = scalar.constant 1098 : i32 + %v1_11_3 = scalar.constant 1105 : i32 + %w1_11 = vector.from_elements %v1_11_0, %v1_11_1, %v1_11_2, %v1_11_3 : vector<4xi32> + %o1_11 = index.constant 108 : index + vector.store %w1_11, %gv[%o1_11] : vector<4xi32>, view<512xi32> + %v1_12_0 = scalar.constant 1112 : i32 + %v1_12_1 = scalar.constant 1139 : i32 + %v1_12_2 = scalar.constant 1143 : i32 + %v1_12_3 = scalar.constant 1144 : i32 + %w1_12 = vector.from_elements %v1_12_0, %v1_12_1, %v1_12_2, %v1_12_3 : vector<4xi32> + %o1_12 = index.constant 112 : index + vector.store %w1_12, %gv[%o1_12] : vector<4xi32>, view<512xi32> + %v1_13_0 = scalar.constant 1152 : i32 + %v1_13_1 = scalar.constant 1154 : i32 + %v1_13_2 = scalar.constant 1161 : i32 + %v1_13_3 = scalar.constant 1167 : i32 + %w1_13 = vector.from_elements %v1_13_0, %v1_13_1, %v1_13_2, %v1_13_3 : vector<4xi32> + %o1_13 = index.constant 116 : index + vector.store %w1_13, %gv[%o1_13] : vector<4xi32>, view<512xi32> + %v1_14_0 = scalar.constant 1168 : i32 + %v1_14_1 = scalar.constant 1170 : i32 + %v1_14_2 = scalar.constant 1183 : i32 + %v1_14_3 = scalar.constant 1184 : i32 + %w1_14 = vector.from_elements %v1_14_0, %v1_14_1, %v1_14_2, %v1_14_3 : vector<4xi32> + %o1_14 = index.constant 120 : index + vector.store %w1_14, %gv[%o1_14] : vector<4xi32>, view<512xi32> + %v1_15_0 = scalar.constant 1197 : i32 + %v1_15_1 = scalar.constant 1217 : i32 + %v1_15_2 = scalar.constant 1224 : i32 + %v1_15_3 = scalar.constant 1228 : i32 + %w1_15 = vector.from_elements %v1_15_0, %v1_15_1, %v1_15_2, %v1_15_3 : vector<4xi32> + %o1_15 = index.constant 124 : index + vector.store %w1_15, %gv[%o1_15] : vector<4xi32>, view<512xi32> + } + %k2 = index.constant 2 : index + %is2 = index.cmp eq, %chunk, %k2 : index + scf.if %is2 { + %v2_0_0 = scalar.constant 1272 : i32 + %v2_0_1 = scalar.constant 1276 : i32 + %v2_0_2 = scalar.constant 1309 : i32 + %v2_0_3 = scalar.constant 1323 : i32 + %w2_0 = vector.from_elements %v2_0_0, %v2_0_1, %v2_0_2, %v2_0_3 : vector<4xi32> + %o2_0 = index.constant 128 : index + vector.store %w2_0, %gv[%o2_0] : vector<4xi32>, view<512xi32> + %v2_1_0 = scalar.constant 1347 : i32 + %v2_1_1 = scalar.constant 1367 : i32 + %v2_1_2 = scalar.constant 1377 : i32 + %v2_1_3 = scalar.constant 1404 : i32 + %w2_1 = vector.from_elements %v2_1_0, %v2_1_1, %v2_1_2, %v2_1_3 : vector<4xi32> + %o2_1 = index.constant 132 : index + vector.store %w2_1, %gv[%o2_1] : vector<4xi32>, view<512xi32> + %v2_2_0 = scalar.constant 1473 : i32 + %v2_2_1 = scalar.constant 1475 : i32 + %v2_2_2 = scalar.constant 1486 : i32 + %v2_2_3 = scalar.constant 1509 : i32 + %w2_2 = vector.from_elements %v2_2_0, %v2_2_1, %v2_2_2, %v2_2_3 : vector<4xi32> + %o2_2 = index.constant 136 : index + vector.store %w2_2, %gv[%o2_2] : vector<4xi32>, view<512xi32> + %v2_3_0 = scalar.constant 1537 : i32 + %v2_3_1 = scalar.constant 1544 : i32 + %v2_3_2 = scalar.constant 1546 : i32 + %v2_3_3 = scalar.constant 1553 : i32 + %w2_3 = vector.from_elements %v2_3_0, %v2_3_1, %v2_3_2, %v2_3_3 : vector<4xi32> + %o2_3 = index.constant 140 : index + vector.store %w2_3, %gv[%o2_3] : vector<4xi32>, view<512xi32> + %v2_4_0 = scalar.constant 1555 : i32 + %v2_4_1 = scalar.constant 1576 : i32 + %v2_4_2 = scalar.constant 1589 : i32 + %v2_4_3 = scalar.constant 1594 : i32 + %w2_4 = vector.from_elements %v2_4_0, %v2_4_1, %v2_4_2, %v2_4_3 : vector<4xi32> + %o2_4 = index.constant 144 : index + vector.store %w2_4, %gv[%o2_4] : vector<4xi32>, view<512xi32> + %v2_5_0 = scalar.constant 1600 : i32 + %v2_5_1 = scalar.constant 1602 : i32 + %v2_5_2 = scalar.constant 1616 : i32 + %v2_5_3 = scalar.constant 1625 : i32 + %w2_5 = vector.from_elements %v2_5_0, %v2_5_1, %v2_5_2, %v2_5_3 : vector<4xi32> + %o2_5 = index.constant 148 : index + vector.store %w2_5, %gv[%o2_5] : vector<4xi32>, view<512xi32> + %v2_6_0 = scalar.constant 1636 : i32 + %v2_6_1 = scalar.constant 1638 : i32 + %v2_6_2 = scalar.constant 1665 : i32 + %v2_6_3 = scalar.constant 1667 : i32 + %w2_6 = vector.from_elements %v2_6_0, %v2_6_1, %v2_6_2, %v2_6_3 : vector<4xi32> + %o2_6 = index.constant 152 : index + vector.store %w2_6, %gv[%o2_6] : vector<4xi32>, view<512xi32> + %v2_7_0 = scalar.constant 1672 : i32 + %v2_7_1 = scalar.constant 1685 : i32 + %v2_7_2 = scalar.constant 1706 : i32 + %v2_7_3 = scalar.constant 1722 : i32 + %w2_7 = vector.from_elements %v2_7_0, %v2_7_1, %v2_7_2, %v2_7_3 : vector<4xi32> + %o2_7 = index.constant 156 : index + vector.store %w2_7, %gv[%o2_7] : vector<4xi32>, view<512xi32> + %v2_8_0 = scalar.constant 1737 : i32 + %v2_8_1 = scalar.constant 1755 : i32 + %v2_8_2 = scalar.constant 1816 : i32 + %v2_8_3 = scalar.constant 1831 : i32 + %w2_8 = vector.from_elements %v2_8_0, %v2_8_1, %v2_8_2, %v2_8_3 : vector<4xi32> + %o2_8 = index.constant 160 : index + vector.store %w2_8, %gv[%o2_8] : vector<4xi32>, view<512xi32> + %v2_9_0 = scalar.constant 1850 : i32 + %v2_9_1 = scalar.constant 1856 : i32 + %v2_9_2 = scalar.constant 1862 : i32 + %v2_9_3 = scalar.constant 1874 : i32 + %w2_9 = vector.from_elements %v2_9_0, %v2_9_1, %v2_9_2, %v2_9_3 : vector<4xi32> + %o2_9 = index.constant 164 : index + vector.store %w2_9, %gv[%o2_9] : vector<4xi32>, view<512xi32> + %v2_10_0 = scalar.constant 1901 : i32 + %v2_10_1 = scalar.constant 1932 : i32 + %v2_10_2 = scalar.constant 1950 : i32 + %v2_10_3 = scalar.constant 1971 : i32 + %w2_10 = vector.from_elements %v2_10_0, %v2_10_1, %v2_10_2, %v2_10_3 : vector<4xi32> + %o2_10 = index.constant 168 : index + vector.store %w2_10, %gv[%o2_10] : vector<4xi32>, view<512xi32> + %v2_11_0 = scalar.constant 2011 : i32 + %v2_11_1 = scalar.constant 2032 : i32 + %v2_11_2 = scalar.constant 2052 : i32 + %v2_11_3 = scalar.constant 2063 : i32 + %w2_11 = vector.from_elements %v2_11_0, %v2_11_1, %v2_11_2, %v2_11_3 : vector<4xi32> + %o2_11 = index.constant 172 : index + vector.store %w2_11, %gv[%o2_11] : vector<4xi32>, view<512xi32> + %v2_12_0 = scalar.constant 2077 : i32 + %v2_12_1 = scalar.constant 2079 : i32 + %v2_12_2 = scalar.constant 2091 : i32 + %v2_12_3 = scalar.constant 2095 : i32 + %w2_12 = vector.from_elements %v2_12_0, %v2_12_1, %v2_12_2, %v2_12_3 : vector<4xi32> + %o2_12 = index.constant 176 : index + vector.store %w2_12, %gv[%o2_12] : vector<4xi32>, view<512xi32> + %v2_13_0 = scalar.constant 2172 : i32 + %v2_13_1 = scalar.constant 2192 : i32 + %v2_13_2 = scalar.constant 2207 : i32 + %v2_13_3 = scalar.constant 2208 : i32 + %w2_13 = vector.from_elements %v2_13_0, %v2_13_1, %v2_13_2, %v2_13_3 : vector<4xi32> + %o2_13 = index.constant 180 : index + vector.store %w2_13, %gv[%o2_13] : vector<4xi32>, view<512xi32> + %v2_14_0 = scalar.constant 2224 : i32 + %v2_14_1 = scalar.constant 2230 : i32 + %v2_14_2 = scalar.constant 2247 : i32 + %v2_14_3 = scalar.constant 2277 : i32 + %w2_14 = vector.from_elements %v2_14_0, %v2_14_1, %v2_14_2, %v2_14_3 : vector<4xi32> + %o2_14 = index.constant 184 : index + vector.store %w2_14, %gv[%o2_14] : vector<4xi32>, view<512xi32> + %v2_15_0 = scalar.constant 2308 : i32 + %v2_15_1 = scalar.constant 2345 : i32 + %v2_15_2 = scalar.constant 2356 : i32 + %v2_15_3 = scalar.constant 2389 : i32 + %w2_15 = vector.from_elements %v2_15_0, %v2_15_1, %v2_15_2, %v2_15_3 : vector<4xi32> + %o2_15 = index.constant 188 : index + vector.store %w2_15, %gv[%o2_15] : vector<4xi32>, view<512xi32> + } + %k3 = index.constant 3 : index + %is3 = index.cmp eq, %chunk, %k3 : index + scf.if %is3 { + %v3_0_0 = scalar.constant 2403 : i32 + %v3_0_1 = scalar.constant 2424 : i32 + %v3_0_2 = scalar.constant 2501 : i32 + %v3_0_3 = scalar.constant 2504 : i32 + %w3_0 = vector.from_elements %v3_0_0, %v3_0_1, %v3_0_2, %v3_0_3 : vector<4xi32> + %o3_0 = index.constant 192 : index + vector.store %w3_0, %gv[%o3_0] : vector<4xi32>, view<512xi32> + %v3_1_0 = scalar.constant 2506 : i32 + %v3_1_1 = scalar.constant 2520 : i32 + %v3_1_2 = scalar.constant 2570 : i32 + %v3_1_3 = scalar.constant 2593 : i32 + %w3_1 = vector.from_elements %v3_1_0, %v3_1_1, %v3_1_2, %v3_1_3 : vector<4xi32> + %o3_1 = index.constant 196 : index + vector.store %w3_1, %gv[%o3_1] : vector<4xi32>, view<512xi32> + %v3_2_0 = scalar.constant 2616 : i32 + %v3_2_1 = scalar.constant 2624 : i32 + %v3_2_2 = scalar.constant 2630 : i32 + %v3_2_3 = scalar.constant 2646 : i32 + %w3_2 = vector.from_elements %v3_2_0, %v3_2_1, %v3_2_2, %v3_2_3 : vector<4xi32> + %o3_2 = index.constant 200 : index + vector.store %w3_2, %gv[%o3_2] : vector<4xi32>, view<512xi32> + %v3_3_0 = scalar.constant 2669 : i32 + %v3_3_1 = scalar.constant 2700 : i32 + %v3_3_2 = scalar.constant 2714 : i32 + %v3_3_3 = scalar.constant 2746 : i32 + %w3_3 = vector.from_elements %v3_3_0, %v3_3_1, %v3_3_2, %v3_3_3 : vector<4xi32> + %o3_3 = index.constant 204 : index + vector.store %w3_3, %gv[%o3_3] : vector<4xi32>, view<512xi32> + %v3_4_0 = scalar.constant 2754 : i32 + %v3_4_1 = scalar.constant 2795 : i32 + %v3_4_2 = scalar.constant 2824 : i32 + %v3_4_3 = scalar.constant 2835 : i32 + %w3_4 = vector.from_elements %v3_4_0, %v3_4_1, %v3_4_2, %v3_4_3 : vector<4xi32> + %o3_4 = index.constant 208 : index + vector.store %w3_4, %gv[%o3_4] : vector<4xi32>, view<512xi32> + %v3_5_0 = scalar.constant 2839 : i32 + %v3_5_1 = scalar.constant 2874 : i32 + %v3_5_2 = scalar.constant 2882 : i32 + %v3_5_3 = scalar.constant 2905 : i32 + %w3_5 = vector.from_elements %v3_5_0, %v3_5_1, %v3_5_2, %v3_5_3 : vector<4xi32> + %o3_5 = index.constant 212 : index + vector.store %w3_5, %gv[%o3_5] : vector<4xi32>, view<512xi32> + %v3_6_0 = scalar.constant 2984 : i32 + %v3_6_1 = scalar.constant 3028 : i32 + %v3_6_2 = scalar.constant 3042 : i32 + %v3_6_3 = scalar.constant 3092 : i32 + %w3_6 = vector.from_elements %v3_6_0, %v3_6_1, %v3_6_2, %v3_6_3 : vector<4xi32> + %o3_6 = index.constant 216 : index + vector.store %w3_6, %gv[%o3_6] : vector<4xi32>, view<512xi32> + %v3_7_0 = scalar.constant 3108 : i32 + %v3_7_1 = scalar.constant 3110 : i32 + %v3_7_2 = scalar.constant 3124 : i32 + %v3_7_3 = scalar.constant 3153 : i32 + %w3_7 = vector.from_elements %v3_7_0, %v3_7_1, %v3_7_2, %v3_7_3 : vector<4xi32> + %o3_7 = index.constant 220 : index + vector.store %w3_7, %gv[%o3_7] : vector<4xi32>, view<512xi32> + %v3_8_0 = scalar.constant 3185 : i32 + %v3_8_1 = scalar.constant 3215 : i32 + %v3_8_2 = scalar.constant 3252 : i32 + %v3_8_3 = scalar.constant 3288 : i32 + %w3_8 = vector.from_elements %v3_8_0, %v3_8_1, %v3_8_2, %v3_8_3 : vector<4xi32> + %o3_8 = index.constant 224 : index + vector.store %w3_8, %gv[%o3_8] : vector<4xi32>, view<512xi32> + %v3_9_0 = scalar.constant 3294 : i32 + %v3_9_1 = scalar.constant 3364 : i32 + %v3_9_2 = scalar.constant 3397 : i32 + %v3_9_3 = scalar.constant 3434 : i32 + %w3_9 = vector.from_elements %v3_9_0, %v3_9_1, %v3_9_2, %v3_9_3 : vector<4xi32> + %o3_9 = index.constant 228 : index + vector.store %w3_9, %gv[%o3_9] : vector<4xi32>, view<512xi32> + %v3_10_0 = scalar.constant 3483 : i32 + %v3_10_1 = scalar.constant 3523 : i32 + %v3_10_2 = scalar.constant 3537 : i32 + %v3_10_3 = scalar.constant 3587 : i32 + %w3_10 = vector.from_elements %v3_10_0, %v3_10_1, %v3_10_2, %v3_10_3 : vector<4xi32> + %o3_10 = index.constant 232 : index + vector.store %w3_10, %gv[%o3_10] : vector<4xi32>, view<512xi32> + %v3_11_0 = scalar.constant 3589 : i32 + %v3_11_1 = scalar.constant 3591 : i32 + %v3_11_2 = scalar.constant 3592 : i32 + %v3_11_3 = scalar.constant 3610 : i32 + %w3_11 = vector.from_elements %v3_11_0, %v3_11_1, %v3_11_2, %v3_11_3 : vector<4xi32> + %o3_11 = index.constant 236 : index + vector.store %w3_11, %gv[%o3_11] : vector<4xi32>, view<512xi32> + %v3_12_0 = scalar.constant 3626 : i32 + %v3_12_1 = scalar.constant 3670 : i32 + %v3_12_2 = scalar.constant 3680 : i32 + %v3_12_3 = scalar.constant 3722 : i32 + %w3_12 = vector.from_elements %v3_12_0, %v3_12_1, %v3_12_2, %v3_12_3 : vector<4xi32> + %o3_12 = index.constant 240 : index + vector.store %w3_12, %gv[%o3_12] : vector<4xi32>, view<512xi32> + %v3_13_0 = scalar.constant 3749 : i32 + %v3_13_1 = scalar.constant 3754 : i32 + %v3_13_2 = scalar.constant 3776 : i32 + %v3_13_3 = scalar.constant 3789 : i32 + %w3_13 = vector.from_elements %v3_13_0, %v3_13_1, %v3_13_2, %v3_13_3 : vector<4xi32> + %o3_13 = index.constant 244 : index + vector.store %w3_13, %gv[%o3_13] : vector<4xi32>, view<512xi32> + %v3_14_0 = scalar.constant 3803 : i32 + %v3_14_1 = scalar.constant 3824 : i32 + %v3_14_2 = scalar.constant 3857 : i32 + %v3_14_3 = scalar.constant 3873 : i32 + %w3_14 = vector.from_elements %v3_14_0, %v3_14_1, %v3_14_2, %v3_14_3 : vector<4xi32> + %o3_14 = index.constant 248 : index + vector.store %w3_14, %gv[%o3_14] : vector<4xi32>, view<512xi32> + %v3_15_0 = scalar.constant 3904 : i32 + %v3_15_1 = scalar.constant 3906 : i32 + %v3_15_2 = scalar.constant 3924 : i32 + %v3_15_3 = scalar.constant 3992 : i32 + %w3_15 = vector.from_elements %v3_15_0, %v3_15_1, %v3_15_2, %v3_15_3 : vector<4xi32> + %o3_15 = index.constant 252 : index + vector.store %w3_15, %gv[%o3_15] : vector<4xi32>, view<512xi32> + } + func.return +} + +// IQ3_XXS: a 256-value block is 8 groups of 32, each 4 slots of 8 values (two grid entries of 4); lane l16 +// owns group l16 / 2, slots 2 (l16 % 2) and +1: values 32 (l16 / 2) + 16 (l16 % 2) .. +15, one scale. +// The grid (12-bit codes) is staged in workgroup memory like IQ2's; signs are ksigns_iq2xs (7 bits + parity). +func.def inline @ggml_kquant_iq3xxs_slot_values(%code_a: i32, %code_b: i32, %signs7: i32) -> (vector<8xf32>) { + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c12_i32 = scalar.constant 12 : i32 + %p4 = scalar.shrui %signs7, %c4_i32 : i32 + %x4 = scalar.xori %signs7, %p4 : i32 + %p2 = scalar.shrui %x4, %c2_i32 : i32 + %x2 = scalar.xori %x4, %p2 : i32 + %p1 = scalar.shrui %x2, %c1_i32 : i32 + %x1 = scalar.xori %x2, %p1 : i32 + %parity = scalar.andi %x1, %c1_i32 : i32 + %high = scalar.shli %parity, %c7_i32 : i32 + %signs8 = scalar.ori %signs7, %high : i32 + %b_hi = scalar.shli %code_b, %c12_i32 : i32 + %both = scalar.ori %code_a, %b_hi : i32 + %s0 = scalar.constant 0 : i32 + %s3 = scalar.constant 3 : i32 + %s5 = scalar.constant 5 : i32 + %s6 = scalar.constant 6 : i32 + %s8 = scalar.constant 8 : i32 + %s9 = scalar.constant 9 : i32 + %s12 = scalar.constant 12 : i32 + %s15 = scalar.constant 15 : i32 + %s18 = scalar.constant 18 : i32 + %s21 = scalar.constant 21 : i32 + %shift3 = vector.from_elements %s0, %s3, %s6, %s9, %s12, %s15, %s18, %s21 : vector<8xi32> + %shift1 = vector.from_elements %s0, %c1_i32, %c2_i32, %s3, %c4_i32, %s5, %s6, %c7_i32 : vector<8xi32> + %cv = vector.splat %both : vector<8xi32> + %lvs = vector.shrui %cv, %shift3 : vector<8xi32> + %seven = vector.splat %c7_i32 : vector<8xi32> + %lv = vector.andi %lvs, %seven : vector<8xi32> + %eight = vector.splat %s8 : vector<8xi32> + %four = vector.splat %c4_i32 : vector<8xi32> + %lv8 = vector.muli %lv, %eight : vector<8xi32> + %base = vector.addi %lv8, %four : vector<8xi32> + // level 7 is the only one with all three bits set: bump = 2 (L & L >> 1 & L >> 2 & 1) + %one7 = vector.splat %c1_i32 : vector<8xi32> + %two7 = vector.splat %c2_i32 : vector<8xi32> + %l1 = vector.shrui %lv, %one7 : vector<8xi32> + %l2 = vector.shrui %lv, %two7 : vector<8xi32> + %a01 = vector.andi %lv, %l1 : vector<8xi32> + %a012 = vector.andi %a01, %l2 : vector<8xi32> + %is7 = vector.andi %a012, %one7 : vector<8xi32> + %bump = vector.shli %is7, %one7 : vector<8xi32> + %mag = vector.addi %base, %bump : vector<8xi32> + %sv = vector.splat %signs8 : vector<8xi32> + %sb0 = vector.shrui %sv, %shift1 : vector<8xi32> + %one = vector.splat %c1_i32 : vector<8xi32> + %sb = vector.andi %sb0, %one : vector<8xi32> + %sb2 = vector.shli %sb, %one : vector<8xi32> + %sgn = vector.subi %one, %sb2 : vector<8xi32> + %v = vector.muli %mag, %sgn : vector<8xi32> + %vf = vector.sitofp %v : vector<8xi32> to vector<8xf32> + func.return %vf : vector<8xf32> +} + +// IQ3_XXS (98 bytes): lane slots s = 2h, 2h + 1 of group g use grid indices at bytes 2 + 8 g + 2 s (+1) +// and sign group s of aux = u32 at byte 66 + 4 g, scale d * (0.5 + (aux >> 28)) * 0.5. +func.def inline @ggml_kquant_iq3xxs_lane_parts(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c33 = index.constant 33 : index + %c7_i32 = scalar.constant 7 : i32 + %c16_i32 = scalar.constant 16 : i32 + %c28_i32 = scalar.constant 28 : i32 + %c127_i32 = scalar.constant 127 : i32 + %c05 = scalar.constant 0.5 : f32 + %block_bytes = index.constant 98 : offset + %zero_offset = index.constant 0 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %g = index.div %l, %c2 : index + %h = index.rem %l, %c2 : index + %s0 = index.mul %h, %c2 : index + %s1 = index.add %s0, %c1 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %hv = buffer.view %weight[%block_base] : buffer -> view<49xf16> + %wv = buffer.view %weight[%block_base] : buffer -> view<49xi16> + %bv = buffer.view %weight[%block_base] : buffer -> view<98xi8> + %gv = buffer.view %grid[%zero_offset] : buffer -> view<512xi32> + %d_f16 = view.load %hv[%c0] : view<49xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %g8 = index.mul %g, %c8 : index + %i0 = index.add %g8, %c2 : index + %h4 = index.mul %h, %c4 : index + %ia = index.add %i0, %h4 : index + %ib = index.add %ia, %c1 : index + %ic = index.add %ia, %c2 : index + %id = index.add %ia, %c3 : index + %ia_i8 = view.load %bv[%ia] : view<98xi8> -> i8 + %ib_i8 = view.load %bv[%ib] : view<98xi8> -> i8 + %ic_i8 = view.load %bv[%ic] : view<98xi8> -> i8 + %id_i8 = view.load %bv[%id] : view<98xi8> -> i8 + %ga = scalar.extui %ia_i8 : i8 to i32 + %gb = scalar.extui %ib_i8 : i8 to i32 + %gc = scalar.extui %ic_i8 : i8 to i32 + %gd = scalar.extui %id_i8 : i8 to i32 + %ga_x = index.cast %ga : i32 to index + %gb_x = index.cast %gb : i32 to index + %gc_x = index.cast %gc : i32 to index + %gd_x = index.cast %gd : i32 to index + %ga_b = index.assume %ga_x [range(%ga_x, 0, 255)] : index + %gb_b = index.assume %gb_x [range(%gb_x, 0, 255)] : index + %gc_b = index.assume %gc_x [range(%gc_x, 0, 255)] : index + %gd_b = index.assume %gd_x [range(%gd_x, 0, 255)] : index + %code_a = view.load %gv[%ga_b] : view<512xi32> -> i32 + %code_b = view.load %gv[%gb_b] : view<512xi32> -> i32 + %code_c = view.load %gv[%gc_b] : view<512xi32> -> i32 + %code_d = view.load %gv[%gd_b] : view<512xi32> -> i32 + %g2 = index.add %g, %g : index + %w0_at = index.add %c33, %g2 : index + %w1_at = index.add %w0_at, %c1 : index + %w0_i16 = view.load %wv[%w0_at] : view<49xi16> -> i16 + %w1_i16 = view.load %wv[%w1_at] : view<49xi16> -> i16 + %w0 = scalar.extui %w0_i16 : i16 to i32 + %w1 = scalar.extui %w1_i16 : i16 to i32 + %w1s = scalar.shli %w1, %c16_i32 : i32 + %aux = scalar.ori %w0, %w1s : i32 + %sc4 = scalar.shrui %aux, %c28_i32 : i32 + %sc_f = scalar.uitofp %sc4 : i32 to f32 + %sc_p = scalar.addf %sc_f, %c05 : f32 + %ds = scalar.mulf %d, %sc_p : f32 + %scale = scalar.mulf %ds, %c05 : f32 + %s0_i32 = index.cast %s0 : index to i32 + %s1_i32 = index.cast %s1 : index to i32 + %sha = scalar.muli %s0_i32, %c7_i32 : i32 + %shb = scalar.muli %s1_i32, %c7_i32 : i32 + %sa0 = scalar.shrui %aux, %sha : i32 + %sb0 = scalar.shrui %aux, %shb : i32 + %sa = scalar.andi %sa0, %c127_i32 : i32 + %sb = scalar.andi %sb0, %c127_i32 : i32 + %va = func.call @ggml_kquant_iq3xxs_slot_values(%code_a, %code_b, %sa) : (i32, i32, i32) -> (vector<8xf32>) + %vb = func.call @ggml_kquant_iq3xxs_slot_values(%code_c, %code_d, %sb) : (i32, i32, i32) -> (vector<8xf32>) + %v = func.call @ggml_kquant_iq2_join16(%va, %vb) : (vector<8xf32>, vector<8xf32>) -> (vector<16xf32>) + %g32 = index.mul %g, %c32 : index + %h16 = index.mul %h, %c16 : index + %p0 = index.add %g32, %h16 : index + func.return %v, %scale, %p0 : vector<16xf32>, f32, index +} + +func.def inline @ggml_kquant_iq3xxs_lane_dot(%grid: buffer, %weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %zero_scalar = scalar.constant 0.0 : f32 + %x_block_bytes = index.constant 1024 : offset + %v, %scale, %p = func.call @ggml_kquant_iq3xxs_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %x = vector.load %xv[%p] : view<256xf32> -> vector<16xf32> + %vx = vector.mulf %v, %x : vector<16xf32> + %sum = vector.reduce %vx, %zero_scalar : vector<16xf32>, f32 + %result = scalar.mulf %scale, %sum : f32 + func.return %result : f32 +} + +func.def inline @ggml_kquant_iq3xxs_lane_weights(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c12 = index.constant 12 : index + %v, %scale, %p0 = func.call @ggml_kquant_iq3xxs_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %sv = vector.splat %scale : vector<16xf32> + %w = vector.mulf %v, %sv : vector<16xf32> + %p1 = index.add %p0, %c4 : index + %p2 = index.add %p0, %c8 : index + %p3 = index.add %p0, %c12 : index + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +// Stages the IQ2_S grid as 512 words (two 16-bit codes each); subgroup %chunk (0..3) writes +// words 128 %chunk .. +127. +func.def inline @ggml_kquant_iq2s_grid_fill(%grid: buffer, %chunk: index) { + %zero_offset = index.constant 0 : offset + %gv = buffer.view %grid[%zero_offset] : buffer -> view<512xi32> + %k0 = index.constant 0 : index + %is0 = index.cmp eq, %chunk, %k0 : index + scf.if %is0 { + %v0_0_0 = scalar.constant 131072 : i32 + %v0_0_1 = scalar.constant 524293 : i32 + %v0_0_2 = scalar.constant 1114122 : i32 + %v0_0_3 = scalar.constant 1441812 : i32 + %w0_0 = vector.from_elements %v0_0_0, %v0_0_1, %v0_0_2, %v0_0_3 : vector<4xi32> + %o0_0 = index.constant 0 : index + vector.store %w0_0, %gv[%o0_0] : vector<4xi32>, view<512xi32> + %v0_1_0 = scalar.constant 2097177 : i32 + %v0_1_1 = scalar.constant 2424866 : i32 + %v0_1_2 = scalar.constant 4259880 : i32 + %v0_1_3 = scalar.constant 4587588 : i32 + %w0_1 = vector.from_elements %v0_1_0, %v0_1_1, %v0_1_2, %v0_1_3 : vector<4xi32> + %o0_1 = index.constant 4 : index + vector.store %w0_1, %gv[%o0_1] : vector<4xi32>, view<512xi32> + %v0_2_0 = scalar.constant 5242953 : i32 + %v0_2_1 = scalar.constant 5570642 : i32 + %v0_2_2 = scalar.constant 6357080 : i32 + %v0_2_3 = scalar.constant 6684772 : i32 + %w0_2 = vector.from_elements %v0_2_0, %v0_2_1, %v0_2_2, %v0_2_3 : vector<4xi32> + %o0_2 = index.constant 8 : index + vector.store %w0_2, %gv[%o0_2] : vector<4xi32>, view<512xi32> + %v0_3_0 = scalar.constant 8388713 : i32 + %v0_3_1 = scalar.constant 8716418 : i32 + %v0_3_2 = scalar.constant 9502856 : i32 + %v0_3_3 = scalar.constant 10485908 : i32 + %w0_3 = vector.from_elements %v0_3_0, %v0_3_1, %v0_3_2, %v0_3_3 : vector<4xi32> + %o0_3 = index.constant 12 : index + vector.store %w0_3, %gv[%o0_3] : vector<4xi32>, view<512xi32> + %v0_4_0 = scalar.constant 11141285 : i32 + %v0_4_1 = scalar.constant 17039617 : i32 + %v0_4_2 = scalar.constant 17367302 : i32 + %v0_4_3 = scalar.constant 17957136 : i32 + %w0_4 = vector.from_elements %v0_4_0, %v0_4_1, %v0_4_2, %v0_4_3 : vector<4xi32> + %o0_4 = index.constant 16 : index + vector.store %w0_4, %gv[%o0_4] : vector<4xi32>, view<512xi32> + %v0_5_0 = scalar.constant 18350357 : i32 + %v0_5_1 = scalar.constant 19136801 : i32 + %v0_5_2 = scalar.constant 21102912 : i32 + %v0_5_3 = scalar.constant 21496133 : i32 + %w0_5 = vector.from_elements %v0_5_0, %v0_5_1, %v0_5_2, %v0_5_3 : vector<4xi32> + %o0_5 = index.constant 20 : index + vector.store %w0_5, %gv[%o0_5] : vector<4xi32>, view<512xi32> + %v0_6_0 = scalar.constant 22282577 : i32 + %v0_6_1 = scalar.constant 22610262 : i32 + %v0_6_2 = scalar.constant 23396704 : i32 + %v0_6_3 = scalar.constant 25231720 : i32 + %w0_6 = vector.from_elements %v0_6_0, %v0_6_1, %v0_6_2, %v0_6_3 : vector<4xi32> + %o0_6 = index.constant 24 : index + vector.store %w0_6, %gv[%o0_6] : vector<4xi32>, view<512xi32> + %v0_7_0 = scalar.constant 26214788 : i32 + %v0_7_1 = scalar.constant 26542482 : i32 + %v0_7_2 = scalar.constant 27525537 : i32 + %v0_7_3 = scalar.constant 33686016 : i32 + %w0_7 = vector.from_elements %v0_7_0, %v0_7_1, %v0_7_2, %v0_7_3 : vector<4xi32> + %o0_7 = index.constant 28 : index + vector.store %w0_7, %gv[%o0_7] : vector<4xi32>, view<512xi32> + %v0_8_0 = scalar.constant 34079237 : i32 + %v0_8_1 = scalar.constant 34865681 : i32 + %v0_8_2 = scalar.constant 36307488 : i32 + %v0_8_3 = scalar.constant 38011457 : i32 + %w0_8 = vector.from_elements %v0_8_0, %v0_8_1, %v0_8_2, %v0_8_3 : vector<4xi32> + %o0_8 = index.constant 32 : index + vector.store %w0_8, %gv[%o0_8] : vector<4xi32>, view<512xi32> + %v0_9_0 = scalar.constant 38339142 : i32 + %v0_9_1 = scalar.constant 39125584 : i32 + %v0_9_2 = scalar.constant 42271360 : i32 + %v0_9_3 = scalar.constant 43254410 : i32 + %w0_9 = vector.from_elements %v0_9_0, %v0_9_1, %v0_9_2, %v0_9_3 : vector<4xi32> + %o0_9 = index.constant 36 : index + vector.store %w0_9, %gv[%o0_9] : vector<4xi32>, view<512xi32> + %v0_10_0 = scalar.constant 67175074 : i32 + %v0_10_1 = scalar.constant 67503108 : i32 + %v0_10_2 = scalar.constant 68158473 : i32 + %v0_10_3 = scalar.constant 68486162 : i32 + %w0_10 = vector.from_elements %v0_10_0, %v0_10_1, %v0_10_2, %v0_10_3 : vector<4xi32> + %o0_10 = index.constant 40 : index + vector.store %w0_10, %gv[%o0_10] : vector<4xi32>, view<512xi32> + %v0_11_0 = scalar.constant 69272600 : i32 + %v0_11_1 = scalar.constant 69600292 : i32 + %v0_11_2 = scalar.constant 71304233 : i32 + %v0_11_3 = scalar.constant 71631938 : i32 + %w0_11 = vector.from_elements %v0_11_0, %v0_11_1, %v0_11_2, %v0_11_3 : vector<4xi32> + %o0_11 = index.constant 44 : index + vector.store %w0_11, %gv[%o0_11] : vector<4xi32>, view<512xi32> + %v0_12_0 = scalar.constant 71959624 : i32 + %v0_12_1 = scalar.constant 72614993 : i32 + %v0_12_2 = scalar.constant 72942678 : i32 + %v0_12_3 = scalar.constant 73532512 : i32 + %w0_12 = vector.from_elements %v0_12_0, %v0_12_1, %v0_12_2, %v0_12_3 : vector<4xi32> + %o0_12 = index.constant 48 : index + vector.store %w0_12, %gv[%o0_12] : vector<4xi32>, view<512xi32> + %v0_13_0 = scalar.constant 75564133 : i32 + %v0_13_1 = scalar.constant 75891844 : i32 + %v0_13_2 = scalar.constant 76547209 : i32 + %v0_13_3 = scalar.constant 77071509 : i32 + %w0_13 = vector.from_elements %v0_13_0, %v0_13_1, %v0_13_2, %v0_13_3 : vector<4xi32> + %o0_13 = index.constant 52 : index + vector.store %w0_13, %gv[%o0_13] : vector<4xi32>, view<512xi32> + %v0_14_0 = scalar.constant 77857953 : i32 + %v0_14_1 = scalar.constant 84018432 : i32 + %v0_14_2 = scalar.constant 84411653 : i32 + %v0_14_3 = scalar.constant 85001482 : i32 + %w0_14 = vector.from_elements %v0_14_0, %v0_14_1, %v0_14_2, %v0_14_3 : vector<4xi32> + %o0_14 = index.constant 56 : index + vector.store %w0_14, %gv[%o0_14] : vector<4xi32>, view<512xi32> + %v0_15_0 = scalar.constant 85329172 : i32 + %v0_15_1 = scalar.constant 85984537 : i32 + %v0_15_2 = scalar.constant 86508837 : i32 + %v0_15_3 = scalar.constant 88343873 : i32 + %w0_15 = vector.from_elements %v0_15_0, %v0_15_1, %v0_15_2, %v0_15_3 : vector<4xi32> + %o0_15 = index.constant 60 : index + vector.store %w0_15, %gv[%o0_15] : vector<4xi32>, view<512xi32> + %v0_16_0 = scalar.constant 88671558 : i32 + %v0_16_1 = scalar.constant 89261392 : i32 + %v0_16_2 = scalar.constant 89654613 : i32 + %v0_16_3 = scalar.constant 90441057 : i32 + %w0_16 = vector.from_elements %v0_16_0, %v0_16_1, %v0_16_2, %v0_16_3 : vector<4xi32> + %o0_16 = index.constant 64 : index + vector.store %w0_16, %gv[%o0_16] : vector<4xi32>, view<512xi32> + %v0_17_0 = scalar.constant 92407168 : i32 + %v0_17_1 = scalar.constant 92800389 : i32 + %v0_17_2 = scalar.constant 93586833 : i32 + %v0_17_3 = scalar.constant 100730272 : i32 + %w0_17 = vector.from_elements %v0_17_0, %v0_17_1, %v0_17_2, %v0_17_3 : vector<4xi32> + %o0_17 = index.constant 68 : index + vector.store %w0_17, %gv[%o0_17] : vector<4xi32>, view<512xi32> + %v0_18_0 = scalar.constant 101058052 : i32 + %v0_18_1 = scalar.constant 101713417 : i32 + %v0_18_2 = scalar.constant 104859157 : i32 + %v0_18_3 = scalar.constant 105383493 : i32 + %w0_18 = vector.from_elements %v0_18_0, %v0_18_1, %v0_18_2, %v0_18_3 : vector<4xi32> + %o0_18 = index.constant 72 : index + vector.store %w0_18, %gv[%o0_18] : vector<4xi32>, view<512xi32> + %v0_19_0 = scalar.constant 106169937 : i32 + %v0_19_1 = scalar.constant 109119072 : i32 + %v0_19_2 = scalar.constant 110102148 : i32 + %v0_19_3 = scalar.constant 134350848 : i32 + %w0_19 = vector.from_elements %v0_19_0, %v0_19_1, %v0_19_2, %v0_19_3 : vector<4xi32> + %o0_19 = index.constant 76 : index + vector.store %w0_19, %gv[%o0_19] : vector<4xi32>, view<512xi32> + %v0_20_0 = scalar.constant 134744069 : i32 + %v0_20_1 = scalar.constant 135530513 : i32 + %v0_20_2 = scalar.constant 135858198 : i32 + %v0_20_3 = scalar.constant 136644640 : i32 + %w0_20 = vector.from_elements %v0_20_0, %v0_20_1, %v0_20_2, %v0_20_3 : vector<4xi32> + %o0_20 = index.constant 80 : index + vector.store %w0_20, %gv[%o0_20] : vector<4xi32>, view<512xi32> + %v0_21_0 = scalar.constant 138479658 : i32 + %v0_21_1 = scalar.constant 138807364 : i32 + %v0_21_2 = scalar.constant 139462729 : i32 + %v0_21_3 = scalar.constant 139790418 : i32 + %w0_21 = vector.from_elements %v0_21_0, %v0_21_1, %v0_21_2, %v0_21_3 : vector<4xi32> + %o0_21 = index.constant 84 : index + vector.store %w0_21, %gv[%o0_21] : vector<4xi32>, view<512xi32> + %v0_22_0 = scalar.constant 140576856 : i32 + %v0_22_1 = scalar.constant 142608484 : i32 + %v0_22_2 = scalar.constant 143919237 : i32 + %v0_22_3 = scalar.constant 151062698 : i32 + %w0_22 = vector.from_elements %v0_22_0, %v0_22_1, %v0_22_2, %v0_22_3 : vector<4xi32> + %o0_22 = index.constant 88 : index + vector.store %w0_22, %gv[%o0_22] : vector<4xi32>, view<512xi32> + %v0_23_0 = scalar.constant 152045828 : i32 + %v0_23_1 = scalar.constant 152373522 : i32 + %v0_23_2 = scalar.constant 153159960 : i32 + %v0_23_3 = scalar.constant 155519296 : i32 + %w0_23 = vector.from_elements %v0_23_0, %v0_23_1, %v0_23_2, %v0_23_3 : vector<4xi32> + %o0_23 = index.constant 92 : index + vector.store %w0_23, %gv[%o0_23] : vector<4xi32>, view<512xi32> + %v0_24_0 = scalar.constant 156305736 : i32 + %v0_24_1 = scalar.constant 157288788 : i32 + %v0_24_2 = scalar.constant 160434561 : i32 + %v0_24_3 = scalar.constant 168888832 : i32 + %w0_24 = vector.from_elements %v0_24_0, %v0_24_1, %v0_24_2, %v0_24_3 : vector<4xi32> + %o0_24 = index.constant 96 : index + vector.store %w0_24, %gv[%o0_24] : vector<4xi32>, view<512xi32> + %v0_25_0 = scalar.constant 170002964 : i32 + %v0_25_1 = scalar.constant 170527272 : i32 + %v0_25_2 = scalar.constant 177801808 : i32 + %v0_25_3 = scalar.constant 268701697 : i32 + %w0_25 = vector.from_elements %v0_25_0, %v0_25_1, %v0_25_2, %v0_25_3 : vector<4xi32> + %o0_25 = index.constant 100 : index + vector.store %w0_25, %gv[%o0_25] : vector<4xi32>, view<512xi32> + %v0_26_0 = scalar.constant 269029382 : i32 + %v0_26_1 = scalar.constant 269619216 : i32 + %v0_26_2 = scalar.constant 270012437 : i32 + %v0_26_3 = scalar.constant 270798881 : i32 + %w0_26 = vector.from_elements %v0_26_0, %v0_26_1, %v0_26_2, %v0_26_3 : vector<4xi32> + %o0_26 = index.constant 104 : index + vector.store %w0_26, %gv[%o0_26] : vector<4xi32>, view<512xi32> + %v0_27_0 = scalar.constant 272633894 : i32 + %v0_27_1 = scalar.constant 272961602 : i32 + %v0_27_2 = scalar.constant 273748040 : i32 + %v0_27_3 = scalar.constant 274075732 : i32 + %w0_27 = vector.from_elements %v0_27_0, %v0_27_1, %v0_27_2, %v0_27_3 : vector<4xi32> + %o0_27 = index.constant 108 : index + vector.store %w0_27, %gv[%o0_27] : vector<4xi32>, view<512xi32> + %v0_28_0 = scalar.constant 274731097 : i32 + %v0_28_1 = scalar.constant 275058786 : i32 + %v0_28_2 = scalar.constant 276893800 : i32 + %v0_28_3 = scalar.constant 277221508 : i32 + %w0_28 = vector.from_elements %v0_28_0, %v0_28_1, %v0_28_2, %v0_28_3 : vector<4xi32> + %o0_28 = index.constant 112 : index + vector.store %w0_28, %gv[%o0_28] : vector<4xi32>, view<512xi32> + %v0_29_0 = scalar.constant 278204560 : i32 + %v0_29_1 = scalar.constant 278991000 : i32 + %v0_29_2 = scalar.constant 285216932 : i32 + %v0_29_3 = scalar.constant 285544706 : i32 + %w0_29 = vector.from_elements %v0_29_0, %v0_29_1, %v0_29_2, %v0_29_3 : vector<4xi32> + %o0_29 = index.constant 116 : index + vector.store %w0_29, %gv[%o0_29] : vector<4xi32>, view<512xi32> + %v0_30_0 = scalar.constant 285872392 : i32 + %v0_30_1 = scalar.constant 286527761 : i32 + %v0_30_2 = scalar.constant 286855446 : i32 + %v0_30_3 = scalar.constant 287445280 : i32 + %w0_30 = vector.from_elements %v0_30_0, %v0_30_1, %v0_30_2, %v0_30_3 : vector<4xi32> + %o0_30 = index.constant 120 : index + vector.store %w0_30, %gv[%o0_30] : vector<4xi32>, view<512xi32> + %v0_31_0 = scalar.constant 287838501 : i32 + %v0_31_1 = scalar.constant 289673537 : i32 + %v0_31_2 = scalar.constant 290001222 : i32 + %v0_31_3 = scalar.constant 290591056 : i32 + %w0_31 = vector.from_elements %v0_31_0, %v0_31_1, %v0_31_2, %v0_31_3 : vector<4xi32> + %o0_31 = index.constant 124 : index + vector.store %w0_31, %gv[%o0_31] : vector<4xi32>, view<512xi32> + } + %k1 = index.constant 1 : index + %is1 = index.cmp eq, %chunk, %k1 : index + scf.if %is1 { + %v1_0_0 = scalar.constant 290984277 : i32 + %v1_0_1 = scalar.constant 291770721 : i32 + %v1_0_2 = scalar.constant 293736832 : i32 + %v1_0_3 = scalar.constant 294130053 : i32 + %w1_0 = vector.from_elements %v1_0_0, %v1_0_1, %v1_0_2, %v1_0_3 : vector<4xi32> + %o1_0 = index.constant 128 : index + vector.store %w1_0, %gv[%o1_0] : vector<4xi32>, view<512xi32> + %v1_1_0 = scalar.constant 294916497 : i32 + %v1_1_1 = scalar.constant 302256641 : i32 + %v1_1_2 = scalar.constant 303043081 : i32 + %v1_1_3 = scalar.constant 304157205 : i32 + %w1_1 = vector.from_elements %v1_1_0, %v1_1_1, %v1_1_2, %v1_1_3 : vector<4xi32> + %o1_1 = index.constant 132 : index + vector.store %w1_1, %gv[%o1_1] : vector<4xi32>, view<512xi32> + %v1_2_0 = scalar.constant 306188836 : i32 + %v1_2_1 = scalar.constant 307302981 : i32 + %v1_2_2 = scalar.constant 310448724 : i32 + %v1_2_3 = scalar.constant 311431812 : i32 + %w1_2 = vector.from_elements %v1_2_0, %v1_2_1, %v1_2_2, %v1_2_3 : vector<4xi32> + %o1_2 = index.constant 136 : index + vector.store %w1_2, %gv[%o1_2] : vector<4xi32>, view<512xi32> + %v1_3_0 = scalar.constant 335680512 : i32 + %v1_3_1 = scalar.constant 336073733 : i32 + %v1_3_2 = scalar.constant 336860177 : i32 + %v1_3_3 = scalar.constant 337187862 : i32 + %w1_3 = vector.from_elements %v1_3_0, %v1_3_1, %v1_3_2, %v1_3_3 : vector<4xi32> + %o1_3 = index.constant 140 : index + vector.store %w1_3, %gv[%o1_3] : vector<4xi32>, view<512xi32> + %v1_4_0 = scalar.constant 337974304 : i32 + %v1_4_1 = scalar.constant 339809320 : i32 + %v1_4_2 = scalar.constant 340137028 : i32 + %v1_4_3 = scalar.constant 340792393 : i32 + %w1_4 = vector.from_elements %v1_4_0, %v1_4_1, %v1_4_2, %v1_4_3 : vector<4xi32> + %o1_4 = index.constant 144 : index + vector.store %w1_4, %gv[%o1_4] : vector<4xi32>, view<512xi32> + %v1_5_0 = scalar.constant 341120082 : i32 + %v1_5_1 = scalar.constant 341906520 : i32 + %v1_5_2 = scalar.constant 343938148 : i32 + %v1_5_3 = scalar.constant 344265858 : i32 + %w1_5 = vector.from_elements %v1_5_0, %v1_5_1, %v1_5_2, %v1_5_3 : vector<4xi32> + %o1_5 = index.constant 148 : index + vector.store %w1_5, %gv[%o1_5] : vector<4xi32>, view<512xi32> + %v1_6_0 = scalar.constant 345052296 : i32 + %v1_6_1 = scalar.constant 346035348 : i32 + %v1_6_2 = scalar.constant 352589057 : i32 + %v1_6_3 = scalar.constant 352916742 : i32 + %w1_6 = vector.from_elements %v1_6_0, %v1_6_1, %v1_6_2, %v1_6_3 : vector<4xi32> + %o1_6 = index.constant 152 : index + vector.store %w1_6, %gv[%o1_6] : vector<4xi32>, view<512xi32> + %v1_7_0 = scalar.constant 353506576 : i32 + %v1_7_1 = scalar.constant 353899797 : i32 + %v1_7_2 = scalar.constant 354686241 : i32 + %v1_7_3 = scalar.constant 356652352 : i32 + %w1_7 = vector.from_elements %v1_7_0, %v1_7_1, %v1_7_2, %v1_7_3 : vector<4xi32> + %o1_7 = index.constant 156 : index + vector.store %w1_7, %gv[%o1_7] : vector<4xi32>, view<512xi32> + %v1_8_0 = scalar.constant 357045573 : i32 + %v1_8_1 = scalar.constant 357832017 : i32 + %v1_8_2 = scalar.constant 360781152 : i32 + %v1_8_3 = scalar.constant 361764228 : i32 + %w1_8 = vector.from_elements %v1_8_0, %v1_8_1, %v1_8_2, %v1_8_3 : vector<4xi32> + %o1_8 = index.constant 160 : index + vector.store %w1_8, %gv[%o1_8] : vector<4xi32>, view<512xi32> + %v1_9_0 = scalar.constant 369432064 : i32 + %v1_9_1 = scalar.constant 370218504 : i32 + %v1_9_2 = scalar.constant 371201556 : i32 + %v1_9_3 = scalar.constant 373560897 : i32 + %w1_9 = vector.from_elements %v1_9_0, %v1_9_1, %v1_9_2, %v1_9_3 : vector<4xi32> + %o1_9 = index.constant 164 : index + vector.store %w1_9, %gv[%o1_9] : vector<4xi32>, view<512xi32> + %v1_10_0 = scalar.constant 377493072 : i32 + %v1_10_1 = scalar.constant 402724522 : i32 + %v1_10_2 = scalar.constant 403052548 : i32 + %v1_10_3 = scalar.constant 403707913 : i32 + %w1_10 = vector.from_elements %v1_10_0, %v1_10_1, %v1_10_2, %v1_10_3 : vector<4xi32> + %o1_10 = index.constant 168 : index + vector.store %w1_10, %gv[%o1_10] : vector<4xi32>, view<512xi32> + %v1_11_0 = scalar.constant 404232213 : i32 + %v1_11_1 = scalar.constant 406853665 : i32 + %v1_11_2 = scalar.constant 407181378 : i32 + %v1_11_3 = scalar.constant 407967816 : i32 + %w1_11 = vector.from_elements %v1_11_0, %v1_11_1, %v1_11_2, %v1_11_3 : vector<4xi32> + %o1_11 = index.constant 172 : index + vector.store %w1_11, %gv[%o1_11] : vector<4xi32>, view<512xi32> + %v1_12_0 = scalar.constant 408950868 : i32 + %v1_12_1 = scalar.constant 411310209 : i32 + %v1_12_2 = scalar.constant 419567872 : i32 + %v1_12_3 = scalar.constant 419961093 : i32 + %w1_12 = vector.from_elements %v1_12_0, %v1_12_1, %v1_12_2, %v1_12_3 : vector<4xi32> + %o1_12 = index.constant 176 : index + vector.store %w1_12, %gv[%o1_12] : vector<4xi32>, view<512xi32> + %v1_13_0 = scalar.constant 420747537 : i32 + %v1_13_1 = scalar.constant 423696672 : i32 + %v1_13_2 = scalar.constant 424679748 : i32 + %v1_13_3 = scalar.constant 430053737 : i32 + %w1_13 = vector.from_elements %v1_13_0, %v1_13_1, %v1_13_2, %v1_13_3 : vector<4xi32> + %o1_13 = index.constant 180 : index + vector.store %w1_13, %gv[%o1_13] : vector<4xi32>, view<512xi32> + %v1_14_0 = scalar.constant 437262852 : i32 + %v1_14_1 = scalar.constant 441850432 : i32 + %v1_14_2 = scalar.constant 537010176 : i32 + %v1_14_3 = scalar.constant 537403397 : i32 + %w1_14 = vector.from_elements %v1_14_0, %v1_14_1, %v1_14_2, %v1_14_3 : vector<4xi32> + %o1_14 = index.constant 184 : index + vector.store %w1_14, %gv[%o1_14] : vector<4xi32>, view<512xi32> + %v1_15_0 = scalar.constant 538189841 : i32 + %v1_15_1 = scalar.constant 538517526 : i32 + %v1_15_2 = scalar.constant 539303968 : i32 + %v1_15_3 = scalar.constant 541138986 : i32 + %w1_15 = vector.from_elements %v1_15_0, %v1_15_1, %v1_15_2, %v1_15_3 : vector<4xi32> + %o1_15 = index.constant 188 : index + vector.store %w1_15, %gv[%o1_15] : vector<4xi32>, view<512xi32> + %v1_16_0 = scalar.constant 542122052 : i32 + %v1_16_1 = scalar.constant 542449746 : i32 + %v1_16_2 = scalar.constant 545267812 : i32 + %v1_16_3 = scalar.constant 546578570 : i32 + %w1_16 = vector.from_elements %v1_16_0, %v1_16_1, %v1_16_2, %v1_16_3 : vector<4xi32> + %o1_16 = index.constant 192 : index + vector.store %w1_16, %gv[%o1_16] : vector<4xi32>, view<512xi32> + %v1_17_0 = scalar.constant 553722026 : i32 + %v1_17_1 = scalar.constant 554705156 : i32 + %v1_17_2 = scalar.constant 555032850 : i32 + %v1_17_3 = scalar.constant 557850913 : i32 + %w1_17 = vector.from_elements %v1_17_0, %v1_17_1, %v1_17_2, %v1_17_3 : vector<4xi32> + %o1_17 = index.constant 196 : index + vector.store %w1_17, %gv[%o1_17] : vector<4xi32>, view<512xi32> + %v1_18_0 = scalar.constant 558178626 : i32 + %v1_18_1 = scalar.constant 559161681 : i32 + %v1_18_2 = scalar.constant 562110816 : i32 + %v1_18_3 = scalar.constant 563093892 : i32 + %w1_18 = vector.from_elements %v1_18_0, %v1_18_1, %v1_18_2, %v1_18_3 : vector<4xi32> + %o1_18 = index.constant 200 : index + vector.store %w1_18, %gv[%o1_18] : vector<4xi32>, view<512xi32> + %v1_19_0 = scalar.constant 571089408 : i32 + %v1_19_1 = scalar.constant 573055522 : i32 + %v1_19_2 = scalar.constant 574890538 : i32 + %v1_19_3 = scalar.constant 579347024 : i32 + %w1_19 = vector.from_elements %v1_19_0, %v1_19_1, %v1_19_2, %v1_19_3 : vector<4xi32> + %o1_19 = index.constant 204 : index + vector.store %w1_19, %gv[%o1_19] : vector<4xi32>, view<512xi32> + %v1_20_0 = scalar.constant 581444234 : i32 + %v1_20_1 = scalar.constant 604251137 : i32 + %v1_20_2 = scalar.constant 604578822 : i32 + %v1_20_3 = scalar.constant 605365264 : i32 + %w1_20 = vector.from_elements %v1_20_0, %v1_20_1, %v1_20_2, %v1_20_3 : vector<4xi32> + %o1_20 = index.constant 208 : index + vector.store %w1_20, %gv[%o1_20] : vector<4xi32>, view<512xi32> + %v1_21_0 = scalar.constant 606151704 : i32 + %v1_21_1 = scalar.constant 608183332 : i32 + %v1_21_2 = scalar.constant 608511042 : i32 + %v1_21_3 = scalar.constant 609297480 : i32 + %w1_21 = vector.from_elements %v1_21_0, %v1_21_1, %v1_21_2, %v1_21_3 : vector<4xi32> + %o1_21 = index.constant 212 : index + vector.store %w1_21, %gv[%o1_21] : vector<4xi32>, view<512xi32> + %v1_22_0 = scalar.constant 610280532 : i32 + %v1_22_1 = scalar.constant 612639873 : i32 + %v1_22_2 = scalar.constant 620766352 : i32 + %v1_22_3 = scalar.constant 621290757 : i32 + %w1_22 = vector.from_elements %v1_22_0, %v1_22_1, %v1_22_2, %v1_22_3 : vector<4xi32> + %o1_22 = index.constant 216 : index + vector.store %w1_22, %gv[%o1_22] : vector<4xi32>, view<512xi32> + %v1_23_0 = scalar.constant 622077201 : i32 + %v1_23_1 = scalar.constant 625026336 : i32 + %v1_23_2 = scalar.constant 626009412 : i32 + %v1_23_3 = scalar.constant 629155174 : i32 + %w1_23 = vector.from_elements %v1_23_0, %v1_23_1, %v1_23_2, %v1_23_3 : vector<4xi32> + %o1_23 = index.constant 220 : index + vector.store %w1_23, %gv[%o1_23] : vector<4xi32>, view<512xi32> + %v1_24_0 = scalar.constant 637806081 : i32 + %v1_24_1 = scalar.constant 641738256 : i32 + %v1_24_2 = scalar.constant 671098457 : i32 + %v1_24_3 = scalar.constant 672212997 : i32 + %w1_24 = vector.from_elements %v1_24_0, %v1_24_1, %v1_24_2, %v1_24_3 : vector<4xi32> + %o1_24 = index.constant 224 : index + vector.store %w1_24, %gv[%o1_24] : vector<4xi32>, view<512xi32> + %v1_25_0 = scalar.constant 675358740 : i32 + %v1_25_1 = scalar.constant 676341828 : i32 + %v1_25_2 = scalar.constant 682240138 : i32 + %v1_25_3 = scalar.constant 688138497 : i32 + %w1_25 = vector.from_elements %v1_25_0, %v1_25_1, %v1_25_2, %v1_25_3 : vector<4xi32> + %o1_25 = index.constant 228 : index + vector.store %w1_25, %gv[%o1_25] : vector<4xi32>, view<512xi32> + %v1_26_0 = scalar.constant 697641232 : i32 + %v1_26_1 = scalar.constant 706882058 : i32 + %v1_26_2 = scalar.constant 713566820 : i32 + %v1_26_3 = scalar.constant 1073818250 : i32 + %w1_26 = vector.from_elements %v1_26_0, %v1_26_1, %v1_26_2, %v1_26_3 : vector<4xi32> + %o1_26 = index.constant 232 : index + vector.store %w1_26, %gv[%o1_26] : vector<4xi32>, view<512xi32> + %v1_27_0 = scalar.constant 1074151428 : i32 + %v1_27_1 = scalar.constant 1074806793 : i32 + %v1_27_2 = scalar.constant 1075134482 : i32 + %v1_27_3 = scalar.constant 1075462168 : i32 + %w1_27 = vector.from_elements %v1_27_0, %v1_27_1, %v1_27_2, %v1_27_3 : vector<4xi32> + %o1_27 = index.constant 236 : index + vector.store %w1_27, %gv[%o1_27] : vector<4xi32>, view<512xi32> + %v1_28_0 = scalar.constant 1076117537 : i32 + %v1_28_1 = scalar.constant 1077952550 : i32 + %v1_28_2 = scalar.constant 1078280258 : i32 + %v1_28_3 = scalar.constant 1078607944 : i32 + %w1_28 = vector.from_elements %v1_28_0, %v1_28_1, %v1_28_2, %v1_28_3 : vector<4xi32> + %o1_28 = index.constant 240 : index + vector.store %w1_28, %gv[%o1_28] : vector<4xi32>, view<512xi32> + %v1_29_0 = scalar.constant 1079263313 : i32 + %v1_29_1 = scalar.constant 1079590998 : i32 + %v1_29_2 = scalar.constant 1080180832 : i32 + %v1_29_3 = scalar.constant 1082212453 : i32 + %w1_29 = vector.from_elements %v1_29_0, %v1_29_1, %v1_29_2, %v1_29_3 : vector<4xi32> + %o1_29 = index.constant 244 : index + vector.store %w1_29, %gv[%o1_29] : vector<4xi32>, view<512xi32> + %v1_30_0 = scalar.constant 1083195524 : i32 + %v1_30_1 = scalar.constant 1083719829 : i32 + %v1_30_2 = scalar.constant 1084506273 : i32 + %v1_30_3 = scalar.constant 1090666752 : i32 + %w1_30 = vector.from_elements %v1_30_0, %v1_30_1, %v1_30_2, %v1_30_3 : vector<4xi32> + %o1_30 = index.constant 248 : index + vector.store %w1_30, %gv[%o1_30] : vector<4xi32>, view<512xi32> + %v1_31_0 = scalar.constant 1091059973 : i32 + %v1_31_1 = scalar.constant 1091846417 : i32 + %v1_31_2 = scalar.constant 1092174102 : i32 + %v1_31_3 = scalar.constant 1092763936 : i32 + %w1_31 = vector.from_elements %v1_31_0, %v1_31_1, %v1_31_2, %v1_31_3 : vector<4xi32> + %o1_31 = index.constant 252 : index + vector.store %w1_31, %gv[%o1_31] : vector<4xi32>, view<512xi32> + } + %k2 = index.constant 2 : index + %is2 = index.cmp eq, %chunk, %k2 : index + scf.if %is2 { + %v2_0_0 = scalar.constant 1094795557 : i32 + %v2_0_1 = scalar.constant 1095123268 : i32 + %v2_0_2 = scalar.constant 1095778633 : i32 + %v2_0_3 = scalar.constant 1096106322 : i32 + %w2_0 = vector.from_elements %v2_0_0, %v2_0_1, %v2_0_2, %v2_0_3 : vector<4xi32> + %o2_0 = index.constant 256 : index + vector.store %w2_0, %gv[%o2_0] : vector<4xi32>, view<512xi32> + %v2_1_0 = scalar.constant 1096892760 : i32 + %v2_1_1 = scalar.constant 1098924388 : i32 + %v2_1_2 = scalar.constant 1099252098 : i32 + %v2_1_3 = scalar.constant 1100038536 : i32 + %w2_1 = vector.from_elements %v2_1_0, %v2_1_1, %v2_1_2, %v2_1_3 : vector<4xi32> + %o2_1 = index.constant 260 : index + vector.store %w2_1, %gv[%o2_1] : vector<4xi32>, view<512xi32> + %v2_2_0 = scalar.constant 1101021588 : i32 + %v2_2_1 = scalar.constant 1107575297 : i32 + %v2_2_2 = scalar.constant 1108492816 : i32 + %v2_2_3 = scalar.constant 1108886037 : i32 + %w2_2 = vector.from_elements %v2_2_0, %v2_2_1, %v2_2_2, %v2_2_3 : vector<4xi32> + %o2_2 = index.constant 264 : index + vector.store %w2_2, %gv[%o2_2] : vector<4xi32>, view<512xi32> + %v2_3_0 = scalar.constant 1111507492 : i32 + %v2_3_1 = scalar.constant 1112031813 : i32 + %v2_3_2 = scalar.constant 1112818257 : i32 + %v2_3_3 = scalar.constant 1115767392 : i32 + %w2_3 = vector.from_elements %v2_3_0, %v2_3_1, %v2_3_2, %v2_3_3 : vector<4xi32> + %o2_3 = index.constant 268 : index + vector.store %w2_3, %gv[%o2_3] : vector<4xi32>, view<512xi32> + %v2_4_0 = scalar.constant 1140867716 : i32 + %v2_4_1 = scalar.constant 1141195778 : i32 + %v2_4_2 = scalar.constant 1141523464 : i32 + %v2_4_3 = scalar.constant 1142178833 : i32 + %w2_4 = vector.from_elements %v2_4_0, %v2_4_1, %v2_4_2, %v2_4_3 : vector<4xi32> + %o2_4 = index.constant 272 : index + vector.store %w2_4, %gv[%o2_4] : vector<4xi32>, view<512xi32> + %v2_5_0 = scalar.constant 1142506518 : i32 + %v2_5_1 = scalar.constant 1143096352 : i32 + %v2_5_2 = scalar.constant 1143489573 : i32 + %v2_5_3 = scalar.constant 1145324609 : i32 + %w2_5 = vector.from_elements %v2_5_0, %v2_5_1, %v2_5_2, %v2_5_3 : vector<4xi32> + %o2_5 = index.constant 276 : index + vector.store %w2_5, %gv[%o2_5] : vector<4xi32>, view<512xi32> + %v2_6_0 = scalar.constant 1145652294 : i32 + %v2_6_1 = scalar.constant 1146242128 : i32 + %v2_6_2 = scalar.constant 1146635349 : i32 + %v2_6_3 = scalar.constant 1147421793 : i32 + %w2_6 = vector.from_elements %v2_6_0, %v2_6_1, %v2_6_2, %v2_6_3 : vector<4xi32> + %o2_6 = index.constant 280 : index + vector.store %w2_6, %gv[%o2_6] : vector<4xi32>, view<512xi32> + %v2_7_0 = scalar.constant 1149387904 : i32 + %v2_7_1 = scalar.constant 1149781125 : i32 + %v2_7_2 = scalar.constant 1150567569 : i32 + %v2_7_3 = scalar.constant 1157711008 : i32 + %w2_7 = vector.from_elements %v2_7_0, %v2_7_1, %v2_7_2, %v2_7_3 : vector<4xi32> + %o2_7 = index.constant 284 : index + vector.store %w2_7, %gv[%o2_7] : vector<4xi32>, view<512xi32> + %v2_8_0 = scalar.constant 1158038788 : i32 + %v2_8_1 = scalar.constant 1158694153 : i32 + %v2_8_2 = scalar.constant 1159021842 : i32 + %v2_8_3 = scalar.constant 1159808280 : i32 + %w2_8 = vector.from_elements %v2_8_0, %v2_8_1, %v2_8_2, %v2_8_3 : vector<4xi32> + %o2_8 = index.constant 288 : index + vector.store %w2_8, %gv[%o2_8] : vector<4xi32>, view<512xi32> + %v2_9_0 = scalar.constant 1161839908 : i32 + %v2_9_1 = scalar.constant 1162167618 : i32 + %v2_9_2 = scalar.constant 1162954056 : i32 + %v2_9_3 = scalar.constant 1163937108 : i32 + %w2_9 = vector.from_elements %v2_9_0, %v2_9_1, %v2_9_2, %v2_9_3 : vector<4xi32> + %o2_9 = index.constant 292 : index + vector.store %w2_9, %gv[%o2_9] : vector<4xi32>, view<512xi32> + %v2_10_0 = scalar.constant 1166099818 : i32 + %v2_10_1 = scalar.constant 1167082884 : i32 + %v2_10_2 = scalar.constant 1174554112 : i32 + %v2_10_3 = scalar.constant 1174947333 : i32 + %w2_10 = vector.from_elements %v2_10_0, %v2_10_1, %v2_10_2, %v2_10_3 : vector<4xi32> + %o2_10 = index.constant 296 : index + vector.store %w2_10, %gv[%o2_10] : vector<4xi32>, view<512xi32> + %v2_11_0 = scalar.constant 1175733777 : i32 + %v2_11_1 = scalar.constant 1178682912 : i32 + %v2_11_2 = scalar.constant 1179665988 : i32 + %v2_11_3 = scalar.constant 1185236608 : i32 + %w2_11 = vector.from_elements %v2_11_0, %v2_11_1, %v2_11_2, %v2_11_3 : vector<4xi32> + %o2_11 = index.constant 300 : index + vector.store %w2_11, %gv[%o2_11] : vector<4xi32>, view<512xi32> + %v2_12_0 = scalar.constant 1208240129 : i32 + %v2_12_1 = scalar.constant 1209026569 : i32 + %v2_12_2 = scalar.constant 1209354258 : i32 + %v2_12_3 = scalar.constant 1210140696 : i32 + %w2_12 = vector.from_elements %v2_12_0, %v2_12_1, %v2_12_2, %v2_12_3 : vector<4xi32> + %o2_12 = index.constant 304 : index + vector.store %w2_12, %gv[%o2_12] : vector<4xi32>, view<512xi32> + %v2_13_0 = scalar.constant 1212172324 : i32 + %v2_13_1 = scalar.constant 1212500034 : i32 + %v2_13_2 = scalar.constant 1213286472 : i32 + %v2_13_3 = scalar.constant 1214269524 : i32 + %w2_13 = vector.from_elements %v2_13_0, %v2_13_1, %v2_13_2, %v2_13_3 : vector<4xi32> + %o2_13 = index.constant 308 : index + vector.store %w2_13, %gv[%o2_13] : vector<4xi32>, view<512xi32> + %v2_14_0 = scalar.constant 1217415300 : i32 + %v2_14_1 = scalar.constant 1224886528 : i32 + %v2_14_2 = scalar.constant 1225279749 : i32 + %v2_14_3 = scalar.constant 1226066193 : i32 + %w2_14 = vector.from_elements %v2_14_0, %v2_14_1, %v2_14_2, %v2_14_3 : vector<4xi32> + %o2_14 = index.constant 312 : index + vector.store %w2_14, %gv[%o2_14] : vector<4xi32>, view<512xi32> + %v2_15_0 = scalar.constant 1229015328 : i32 + %v2_15_1 = scalar.constant 1229998404 : i32 + %v2_15_2 = scalar.constant 1234585984 : i32 + %v2_15_3 = scalar.constant 1241795073 : i32 + %w2_15 = vector.from_elements %v2_15_0, %v2_15_1, %v2_15_2, %v2_15_3 : vector<4xi32> + %o2_15 = index.constant 316 : index + vector.store %w2_15, %gv[%o2_15] : vector<4xi32>, view<512xi32> + %v2_16_0 = scalar.constant 1245727248 : i32 + %v2_16_1 = scalar.constant 1342328832 : i32 + %v2_16_2 = scalar.constant 1342722053 : i32 + %v2_16_3 = scalar.constant 1343508497 : i32 + %w2_16 = vector.from_elements %v2_16_0, %v2_16_1, %v2_16_2, %v2_16_3 : vector<4xi32> + %o2_16 = index.constant 320 : index + vector.store %w2_16, %gv[%o2_16] : vector<4xi32>, view<512xi32> + %v2_17_0 = scalar.constant 1343836182 : i32 + %v2_17_1 = scalar.constant 1344426016 : i32 + %v2_17_2 = scalar.constant 1344819237 : i32 + %v2_17_3 = scalar.constant 1346654273 : i32 + %w2_17 = vector.from_elements %v2_17_0, %v2_17_1, %v2_17_2, %v2_17_3 : vector<4xi32> + %o2_17 = index.constant 324 : index + vector.store %w2_17, %gv[%o2_17] : vector<4xi32>, view<512xi32> + %v2_18_0 = scalar.constant 1346981958 : i32 + %v2_18_1 = scalar.constant 1347571792 : i32 + %v2_18_2 = scalar.constant 1347965013 : i32 + %v2_18_3 = scalar.constant 1348751457 : i32 + %w2_18 = vector.from_elements %v2_18_0, %v2_18_1, %v2_18_2, %v2_18_3 : vector<4xi32> + %o2_18 = index.constant 328 : index + vector.store %w2_18, %gv[%o2_18] : vector<4xi32>, view<512xi32> + %v2_19_0 = scalar.constant 1350717568 : i32 + %v2_19_1 = scalar.constant 1351110789 : i32 + %v2_19_2 = scalar.constant 1351897233 : i32 + %v2_19_3 = scalar.constant 1359237377 : i32 + %w2_19 = vector.from_elements %v2_19_0, %v2_19_1, %v2_19_2, %v2_19_3 : vector<4xi32> + %o2_19 = index.constant 332 : index + vector.store %w2_19, %gv[%o2_19] : vector<4xi32>, view<512xi32> + %v2_20_0 = scalar.constant 1359565062 : i32 + %v2_20_1 = scalar.constant 1360154896 : i32 + %v2_20_2 = scalar.constant 1360548117 : i32 + %v2_20_3 = scalar.constant 1361334561 : i32 + %w2_20 = vector.from_elements %v2_20_0, %v2_20_1, %v2_20_2, %v2_20_3 : vector<4xi32> + %o2_20 = index.constant 336 : index + vector.store %w2_20, %gv[%o2_20] : vector<4xi32>, view<512xi32> + %v2_21_0 = scalar.constant 1363300672 : i32 + %v2_21_1 = scalar.constant 1363693893 : i32 + %v2_21_2 = scalar.constant 1364480337 : i32 + %v2_21_3 = scalar.constant 1367429472 : i32 + %w2_21 = vector.from_elements %v2_21_0, %v2_21_1, %v2_21_2, %v2_21_3 : vector<4xi32> + %o2_21 = index.constant 340 : index + vector.store %w2_21, %gv[%o2_21] : vector<4xi32>, view<512xi32> + %v2_22_0 = scalar.constant 1368412548 : i32 + %v2_22_1 = scalar.constant 1376080384 : i32 + %v2_22_2 = scalar.constant 1376866824 : i32 + %v2_22_3 = scalar.constant 1377849876 : i32 + %w2_22 = vector.from_elements %v2_22_0, %v2_22_1, %v2_22_2, %v2_22_3 : vector<4xi32> + %o2_22 = index.constant 344 : index + vector.store %w2_22, %gv[%o2_22] : vector<4xi32>, view<512xi32> + %v2_23_0 = scalar.constant 1380209217 : i32 + %v2_23_1 = scalar.constant 1382634064 : i32 + %v2_23_2 = scalar.constant 1409372800 : i32 + %v2_23_3 = scalar.constant 1409700868 : i32 + %w2_23 = vector.from_elements %v2_23_0, %v2_23_1, %v2_23_2, %v2_23_3 : vector<4xi32> + %o2_23 = index.constant 348 : index + vector.store %w2_23, %gv[%o2_23] : vector<4xi32>, view<512xi32> + %v2_24_0 = scalar.constant 1410356233 : i32 + %v2_24_1 = scalar.constant 1410683922 : i32 + %v2_24_2 = scalar.constant 1411470360 : i32 + %v2_24_3 = scalar.constant 1413501988 : i32 + %w2_24 = vector.from_elements %v2_24_0, %v2_24_1, %v2_24_2, %v2_24_3 : vector<4xi32> + %o2_24 = index.constant 352 : index + vector.store %w2_24, %gv[%o2_24] : vector<4xi32>, view<512xi32> + %v2_25_0 = scalar.constant 1413829698 : i32 + %v2_25_1 = scalar.constant 1414616136 : i32 + %v2_25_2 = scalar.constant 1415599188 : i32 + %v2_25_3 = scalar.constant 1417958529 : i32 + %w2_25 = vector.from_elements %v2_25_0, %v2_25_1, %v2_25_2, %v2_25_3 : vector<4xi32> + %o2_25 = index.constant 356 : index + vector.store %w2_25, %gv[%o2_25] : vector<4xi32>, view<512xi32> + %v2_26_0 = scalar.constant 1426085008 : i32 + %v2_26_1 = scalar.constant 1426412802 : i32 + %v2_26_2 = scalar.constant 1427199240 : i32 + %v2_26_3 = scalar.constant 1428182292 : i32 + %w2_26 = vector.from_elements %v2_26_0, %v2_26_1, %v2_26_2, %v2_26_3 : vector<4xi32> + %o2_26 = index.constant 360 : index + vector.store %w2_26, %gv[%o2_26] : vector<4xi32>, view<512xi32> + %v2_27_0 = scalar.constant 1430541633 : i32 + %v2_27_1 = scalar.constant 1434473808 : i32 + %v2_27_2 = scalar.constant 1443124737 : i32 + %v2_27_3 = scalar.constant 1445352976 : i32 + %w2_27 = vector.from_elements %v2_27_0, %v2_27_1, %v2_27_2, %v2_27_3 : vector<4xi32> + %o2_27 = index.constant 364 : index + vector.store %w2_27, %gv[%o2_27] : vector<4xi32>, view<512xi32> + %v2_28_0 = scalar.constant 1476417088 : i32 + %v2_28_1 = scalar.constant 1476745218 : i32 + %v2_28_2 = scalar.constant 1477531656 : i32 + %v2_28_3 = scalar.constant 1478514708 : i32 + %w2_28 = vector.from_elements %v2_28_0, %v2_28_1, %v2_28_2, %v2_28_3 : vector<4xi32> + %o2_28 = index.constant 368 : index + vector.store %w2_28, %gv[%o2_28] : vector<4xi32>, view<512xi32> + %v2_29_0 = scalar.constant 1480874049 : i32 + %v2_29_1 = scalar.constant 1482315856 : i32 + %v2_29_2 = scalar.constant 1493260416 : i32 + %v2_29_3 = scalar.constant 1494243588 : i32 + %w2_29 = vector.from_elements %v2_29_0, %v2_29_1, %v2_29_2, %v2_29_3 : vector<4xi32> + %o2_29 = index.constant 372 : index + vector.store %w2_29, %gv[%o2_29] : vector<4xi32>, view<512xi32> + %v2_30_0 = scalar.constant 1509972288 : i32 + %v2_30_1 = scalar.constant 1518688793 : i32 + %v2_30_2 = scalar.constant 1610701480 : i32 + %v2_30_3 = scalar.constant 1611030532 : i32 + %w2_30 = vector.from_elements %v2_30_0, %v2_30_1, %v2_30_2, %v2_30_3 : vector<4xi32> + %o2_30 = index.constant 376 : index + vector.store %w2_30, %gv[%o2_30] : vector<4xi32>, view<512xi32> + %v2_31_0 = scalar.constant 1611816976 : i32 + %v2_31_1 = scalar.constant 1612210197 : i32 + %v2_31_2 = scalar.constant 1612996641 : i32 + %v2_31_3 = scalar.constant 1615159360 : i32 + %w2_31 = vector.from_elements %v2_31_0, %v2_31_1, %v2_31_2, %v2_31_3 : vector<4xi32> + %o2_31 = index.constant 380 : index + vector.store %w2_31, %gv[%o2_31] : vector<4xi32>, view<512xi32> + } + %k3 = index.constant 3 : index + %is3 = index.cmp eq, %chunk, %k3 : index + scf.if %is3 { + %v3_0_0 = scalar.constant 1615945800 : i32 + %v3_0_1 = scalar.constant 1616928852 : i32 + %v3_0_2 = scalar.constant 1620074628 : i32 + %v3_0_3 = scalar.constant 1627545856 : i32 + %w3_0 = vector.from_elements %v3_0_0, %v3_0_1, %v3_0_2, %v3_0_3 : vector<4xi32> + %o3_0 = index.constant 384 : index + vector.store %w3_0, %gv[%o3_0] : vector<4xi32>, view<512xi32> + %v3_1_0 = scalar.constant 1627939077 : i32 + %v3_1_1 = scalar.constant 1628725521 : i32 + %v3_1_2 = scalar.constant 1631674656 : i32 + %v3_1_3 = scalar.constant 1632657732 : i32 + %w3_1 = vector.from_elements %v3_1_0, %v3_1_1, %v3_1_2, %v3_1_3 : vector<4xi32> + %o3_1 = index.constant 388 : index + vector.store %w3_1, %gv[%o3_1] : vector<4xi32>, view<512xi32> + %v3_2_0 = scalar.constant 1637441920 : i32 + %v3_2_1 = scalar.constant 1645240836 : i32 + %v3_2_2 = scalar.constant 1649828416 : i32 + %v3_2_3 = scalar.constant 1677746849 : i32 + %w3_2 = vector.from_elements %v3_2_0, %v3_2_1, %v3_2_2, %v3_2_3 : vector<4xi32> + %o3_2 = index.constant 392 : index + vector.store %w3_2, %gv[%o3_2] : vector<4xi32>, view<512xi32> + %v3_3_0 = scalar.constant 1678271493 : i32 + %v3_3_1 = scalar.constant 1679057937 : i32 + %v3_3_2 = scalar.constant 1682007072 : i32 + %v3_3_3 = scalar.constant 1682990148 : i32 + %w3_3 = vector.from_elements %v3_3_0, %v3_3_1, %v3_3_2, %v3_3_3 : vector<4xi32> + %o3_3 = index.constant 396 : index + vector.store %w3_3, %gv[%o3_3] : vector<4xi32>, view<512xi32> + %v3_4_0 = scalar.constant 1694590080 : i32 + %v3_4_1 = scalar.constant 1695573252 : i32 + %v3_4_2 = scalar.constant 1699374400 : i32 + %v3_4_3 = scalar.constant 1704093032 : i32 + %w3_4 = vector.from_elements %v3_4_0, %v3_4_1, %v3_4_2, %v3_4_3 : vector<4xi32> + %o3_4 = index.constant 400 : index + vector.store %w3_4, %gv[%o3_4] : vector<4xi32>, view<512xi32> + %v3_5_0 = scalar.constant 1721001472 : i32 + %v3_5_1 = scalar.constant 1745119233 : i32 + %v3_5_2 = scalar.constant 1751476240 : i32 + %v3_5_3 = scalar.constant 1761634456 : i32 + %w3_5 = vector.from_elements %v3_5_0, %v3_5_1, %v3_5_2, %v3_5_3 : vector<4xi32> + %o3_5 = index.constant 404 : index + vector.store %w3_5, %gv[%o3_5] : vector<4xi32>, view<512xi32> + %v3_6_0 = scalar.constant 1782737194 : i32 + %v3_6_1 = scalar.constant -2147456351 : i32 + %v3_6_2 = scalar.constant -2147123198 : i32 + %v3_6_3 = scalar.constant -2146336760 : i32 + %w3_6 = vector.from_elements %v3_6_0, %v3_6_1, %v3_6_2, %v3_6_3 : vector<4xi32> + %o3_6 = index.constant 408 : index + vector.store %w3_6, %gv[%o3_6] : vector<4xi32>, view<512xi32> + %v3_7_0 = scalar.constant -2145812460 : i32 + %v3_7_1 = scalar.constant -2145026016 : i32 + %v3_7_2 = scalar.constant -2142994367 : i32 + %v3_7_3 = scalar.constant -2142076848 : i32 + %w3_7 = vector.from_elements %v3_7_0, %v3_7_1, %v3_7_2, %v3_7_3 : vector<4xi32> + %o3_7 = index.constant 412 : index + vector.store %w3_7, %gv[%o3_7] : vector<4xi32>, view<512xi32> + %v3_8_0 = scalar.constant -2141683627 : i32 + %v3_8_1 = scalar.constant -2139062175 : i32 + %v3_8_2 = scalar.constant -2137948027 : i32 + %v3_8_3 = scalar.constant -2130607980 : i32 + %w3_8 = vector.from_elements %v3_8_0, %v3_8_1, %v3_8_2, %v3_8_3 : vector<4xi32> + %o3_8 = index.constant 416 : index + vector.store %w3_8, %gv[%o3_8] : vector<4xi32>, view<512xi32> + %v3_9_0 = scalar.constant -2130083580 : i32 + %v3_9_1 = scalar.constant -2129493744 : i32 + %v3_9_2 = scalar.constant -2129100523 : i32 + %v3_9_3 = scalar.constant -2128314079 : i32 + %w3_9 = vector.from_elements %v3_9_0, %v3_9_1, %v3_9_2, %v3_9_3 : vector<4xi32> + %o3_9 = index.constant 420 : index + vector.store %w3_9, %gv[%o3_9] : vector<4xi32>, view<512xi32> + %v3_10_0 = scalar.constant -2126347968 : i32 + %v3_10_1 = scalar.constant -2125954747 : i32 + %v3_10_2 = scalar.constant -2125168303 : i32 + %v3_10_3 = scalar.constant -2122022527 : i32 + %w3_10 = vector.from_elements %v3_10_0, %v3_10_1, %v3_10_2, %v3_10_3 : vector<4xi32> + %o3_10 = index.constant 424 : index + vector.store %w3_10, %gv[%o3_10] : vector<4xi32>, view<512xi32> + %v3_11_0 = scalar.constant -2119597680 : i32 + %v3_11_1 = scalar.constant -2113568256 : i32 + %v3_11_2 = scalar.constant -2112781814 : i32 + %v3_11_3 = scalar.constant -2109636076 : i32 + %w3_11 = vector.from_elements %v3_11_0, %v3_11_1, %v3_11_2, %v3_11_3 : vector<4xi32> + %o3_11 = index.constant 428 : index + vector.store %w3_11, %gv[%o3_11] : vector<4xi32>, view<512xi32> + %v3_12_0 = scalar.constant -2108652988 : i32 + %v3_12_1 = scalar.constant -2080078847 : i32 + %v3_12_2 = scalar.constant -2079751162 : i32 + %v3_12_3 = scalar.constant -2079161328 : i32 + %w3_12 = vector.from_elements %v3_12_0, %v3_12_1, %v3_12_2, %v3_12_3 : vector<4xi32> + %o3_12 = index.constant 432 : index + vector.store %w3_12, %gv[%o3_12] : vector<4xi32>, view<512xi32> + %v3_13_0 = scalar.constant -2078768107 : i32 + %v3_13_1 = scalar.constant -2076146655 : i32 + %v3_13_2 = scalar.constant -2075818942 : i32 + %v3_13_3 = scalar.constant -2075032504 : i32 + %w3_13 = vector.from_elements %v3_13_0, %v3_13_1, %v3_13_2, %v3_13_3 : vector<4xi32> + %o3_13 = index.constant 436 : index + vector.store %w3_13, %gv[%o3_13] : vector<4xi32>, view<512xi32> + %v3_14_0 = scalar.constant -2074049452 : i32 + %v3_14_1 = scalar.constant -2071690111 : i32 + %v3_14_2 = scalar.constant -2063563632 : i32 + %v3_14_3 = scalar.constant -2063235838 : i32 + %w3_14 = vector.from_elements %v3_14_0, %v3_14_1, %v3_14_2, %v3_14_3 : vector<4xi32> + %o3_14 = index.constant 440 : index + vector.store %w3_14, %gv[%o3_14] : vector<4xi32>, view<512xi32> + %v3_15_0 = scalar.constant -2062449400 : i32 + %v3_15_1 = scalar.constant -2061466348 : i32 + %v3_15_2 = scalar.constant -2059107007 : i32 + %v3_15_3 = scalar.constant -2055174832 : i32 + %w3_15 = vector.from_elements %v3_15_0, %v3_15_1, %v3_15_2, %v3_15_3 : vector<4xi32> + %o3_15 = index.constant 444 : index + vector.store %w3_15, %gv[%o3_15] : vector<4xi32>, view<512xi32> + %v3_16_0 = scalar.constant -2046720630 : i32 + %v3_16_1 = scalar.constant -2045737468 : i32 + %v3_16_2 = scalar.constant -2042591703 : i32 + %v3_16_3 = scalar.constant -2012903424 : i32 + %w3_16 = vector.from_elements %v3_16_0, %v3_16_1, %v3_16_2, %v3_16_3 : vector<4xi32> + %o3_16 = index.constant 448 : index + vector.store %w3_16, %gv[%o3_16] : vector<4xi32>, view<512xi32> + %v3_17_0 = scalar.constant -2011920367 : i32 + %v3_17_1 = scalar.constant -2008774591 : i32 + %v3_17_2 = scalar.constant -2002614192 : i32 + %v3_17_3 = scalar.constant -1996191487 : i32 + %w3_17 = vector.from_elements %v3_17_0, %v3_17_1, %v3_17_2, %v3_17_3 : vector<4xi32> + %o3_17 = index.constant 452 : index + vector.store %w3_17, %gv[%o3_17] : vector<4xi32>, view<512xi32> + %v3_18_0 = scalar.constant -1989834432 : i32 + %v3_18_1 = scalar.constant -1973908958 : i32 + %v3_18_2 = scalar.constant -1971156390 : i32 + %v3_18_3 = scalar.constant -1878947166 : i32 + %w3_18 = vector.from_elements %v3_18_0, %v3_18_1, %v3_18_2, %v3_18_3 : vector<4xi32> + %o3_18 = index.constant 456 : index + vector.store %w3_18, %gv[%o3_18] : vector<4xi32>, view<512xi32> + %v3_19_0 = scalar.constant -1878421500 : i32 + %v3_19_1 = scalar.constant -1877831664 : i32 + %v3_19_2 = scalar.constant -1877438443 : i32 + %v3_19_3 = scalar.constant -1874816988 : i32 + %w3_19 = vector.from_elements %v3_19_0, %v3_19_1, %v3_19_2, %v3_19_3 : vector<4xi32> + %o3_19 = index.constant 460 : index + vector.store %w3_19, %gv[%o3_19] : vector<4xi32>, view<512xi32> + %v3_20_0 = scalar.constant -1874489278 : i32 + %v3_20_1 = scalar.constant -1873702840 : i32 + %v3_20_2 = scalar.constant -1872719788 : i32 + %v3_20_3 = scalar.constant -1870360447 : i32 + %w3_20 = vector.from_elements %v3_20_0, %v3_20_1, %v3_20_2, %v3_20_3 : vector<4xi32> + %o3_20 = index.constant 464 : index + vector.store %w3_20, %gv[%o3_20] : vector<4xi32>, view<512xi32> + %v3_21_0 = scalar.constant -1862233968 : i32 + %v3_21_1 = scalar.constant -1861119739 : i32 + %v3_21_2 = scalar.constant -1857973996 : i32 + %v3_21_3 = scalar.constant -1856990908 : i32 + %w3_21 = vector.from_elements %v3_21_0, %v3_21_1, %v3_21_2, %v3_21_3 : vector<4xi32> + %o3_21 = index.constant 468 : index + vector.store %w3_21, %gv[%o3_21] : vector<4xi32>, view<512xi32> + %v3_22_0 = scalar.constant -1845391014 : i32 + %v3_22_1 = scalar.constant -1844407804 : i32 + %v3_22_2 = scalar.constant -1834577344 : i32 + %v3_22_3 = scalar.constant -1811770368 : i32 + %w3_22 = vector.from_elements %v3_22_0, %v3_22_1, %v3_22_2, %v3_22_3 : vector<4xi32> + %o3_22 = index.constant 472 : index + vector.store %w3_22, %gv[%o3_22] : vector<4xi32>, view<512xi32> + %v3_23_0 = scalar.constant -1811377147 : i32 + %v3_23_1 = scalar.constant -1810590703 : i32 + %v3_23_2 = scalar.constant -1807641568 : i32 + %v3_23_3 = scalar.constant -1806658492 : i32 + %w3_23 = vector.from_elements %v3_23_0, %v3_23_1, %v3_23_2, %v3_23_3 : vector<4xi32> + %o3_23 = index.constant 476 : index + vector.store %w3_23, %gv[%o3_23] : vector<4xi32>, view<512xi32> + %v3_24_0 = scalar.constant -1802070912 : i32 + %v3_24_1 = scalar.constant -1794861823 : i32 + %v3_24_2 = scalar.constant -1790929648 : i32 + %v3_24_3 = scalar.constant -1784572520 : i32 + %w3_24 = vector.from_elements %v3_24_0, %v3_24_1, %v3_24_2, %v3_24_3 : vector<4xi32> + %o3_24 = index.constant 480 : index + vector.store %w3_24, %gv[%o3_24] : vector<4xi32>, view<512xi32> + %v3_25_0 = scalar.constant -1773758976 : i32 + %v3_25_1 = scalar.constant -1744726428 : i32 + %v3_25_2 = scalar.constant -1743742972 : i32 + %v3_25_3 = scalar.constant -1740597210 : i32 + %w3_25 = vector.from_elements %v3_25_0, %v3_25_1, %v3_25_2, %v3_25_3 : vector<4xi32> + %o3_25 = index.constant 484 : index + vector.store %w3_25, %gv[%o3_25] : vector<4xi32>, view<512xi32> + %v3_26_0 = scalar.constant -1728014167 : i32 + %v3_26_1 = scalar.constant -1722640055 : i32 + %v3_26_2 = scalar.constant -1610573168 : i32 + %v3_26_3 = scalar.constant -1609916411 : i32 + %w3_26 = vector.from_elements %v3_26_0, %v3_26_1, %v3_26_2, %v3_26_3 : vector<4xi32> + %o3_26 = index.constant 488 : index + vector.store %w3_26, %gv[%o3_26] : vector<4xi32>, view<512xi32> + %v3_27_0 = scalar.constant -1608343532 : i32 + %v3_27_1 = scalar.constant -1606311894 : i32 + %v3_27_2 = scalar.constant -1605328828 : i32 + %v3_27_3 = scalar.constant -1599430494 : i32 + %w3_27 = vector.from_elements %v3_27_0, %v3_27_1, %v3_27_2, %v3_27_3 : vector<4xi32> + %o3_27 = index.constant 492 : index + vector.store %w3_27, %gv[%o3_27] : vector<4xi32>, view<512xi32> + %v3_28_0 = scalar.constant -1587175104 : i32 + %v3_28_1 = scalar.constant -1576361470 : i32 + %v3_28_2 = scalar.constant -1574395358 : i32 + %v3_28_3 = scalar.constant -1568497110 : i32 + %w3_28 = vector.from_elements %v3_28_0, %v3_28_1, %v3_28_2, %v3_28_3 : vector<4xi32> + %o3_28 = index.constant 496 : index + vector.store %w3_28, %gv[%o3_28] : vector<4xi32>, view<512xi32> + %v3_29_0 = scalar.constant -1567972728 : i32 + %v3_29_1 = scalar.constant -1543396696 : i32 + %v3_29_2 = scalar.constant -1542413308 : i32 + %v3_29_3 = scalar.constant -1534483392 : i32 + %w3_29 = vector.from_elements %v3_29_0, %v3_29_1, %v3_29_2, %v3_29_3 : vector<4xi32> + %o3_29 = index.constant 500 : index + vector.store %w3_29, %gv[%o3_29] : vector<4xi32>, view<512xi32> + %v3_30_0 = scalar.constant -1526684508 : i32 + %v3_30_1 = scalar.constant -1504598759 : i32 + %v3_30_2 = scalar.constant -1473730550 : i32 + %v3_30_3 = scalar.constant -1454069598 : i32 + %w3_30 = vector.from_elements %v3_30_0, %v3_30_1, %v3_30_2, %v3_30_3 : vector<4xi32> + %o3_30 = index.constant 504 : index + vector.store %w3_30, %gv[%o3_30] : vector<4xi32>, view<512xi32> + %v3_31_0 = scalar.constant -1442272890 : i32 + %v3_31_1 = scalar.constant -1440699894 : i32 + %v3_31_2 = scalar.constant -1440175582 : i32 + %v3_31_3 = scalar.constant -1431655800 : i32 + %w3_31 = vector.from_elements %v3_31_0, %v3_31_1, %v3_31_2, %v3_31_3 : vector<4xi32> + %o3_31 = index.constant 508 : index + vector.store %w3_31, %gv[%o3_31] : vector<4xi32>, view<512xi32> + } + func.return +} + +// IQ2_S: lane l16 owns group l16 / 2, slots 2 (l16 % 2) and +1 (values 32 (l16 / 2) + 16 (l16 % 2) .. +// +15), which share scale nibble l16 % 2 of scales[g]. Slot s: grid index qs[4 g + s] | (qh[g] << (8 - 2 s) +// & 0x300), a full sign byte signs[4 g + s] (no parity bit, unlike IQ2_XXS / IQ2_XS). +func.def inline @ggml_kquant_iq2s_code(%grid: buffer, %index: i32) -> (i32) { + %zero_offset = index.constant 0 : offset + %gv = buffer.view %grid[%zero_offset] : buffer -> view<512xi32> + %c1_i32 = scalar.constant 1 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c65535_i32 = scalar.constant 65535 : i32 + %w_i32 = scalar.shrui %index, %c1_i32 : i32 + %w_x = index.cast %w_i32 : i32 to index + %w_b = index.assume %w_x [range(%w_x, 0, 511)] : index + %word = view.load %gv[%w_b] : view<512xi32> -> i32 + %odd = scalar.andi %index, %c1_i32 : i32 + %sh = scalar.shli %odd, %c4_i32 : i32 + %shifted = scalar.shrui %word, %sh : i32 + %code = scalar.andi %shifted, %c65535_i32 : i32 + func.return %code : i32 +} + +func.def inline @ggml_kquant_iq2s_slot_values(%code: i32, %signs8: i32) -> (vector<8xf32>) { + %c1_i32 = scalar.constant 1 : i32 + %c7_i32 = scalar.constant 7 : i32 + %s0 = scalar.constant 0 : i32 + %s2 = scalar.constant 2 : i32 + %s3 = scalar.constant 3 : i32 + %s4 = scalar.constant 4 : i32 + %s5 = scalar.constant 5 : i32 + %s6 = scalar.constant 6 : i32 + %s8 = scalar.constant 8 : i32 + %s10 = scalar.constant 10 : i32 + %s12 = scalar.constant 12 : i32 + %s14 = scalar.constant 14 : i32 + %shift2 = vector.from_elements %s0, %s2, %s4, %s6, %s8, %s10, %s12, %s14 : vector<8xi32> + %shift1 = vector.from_elements %s0, %c1_i32, %s2, %s3, %s4, %s5, %s6, %c7_i32 : vector<8xi32> + %three = vector.splat %s3 : vector<8xi32> + %one = vector.splat %c1_i32 : vector<8xi32> + %c8v = vector.splat %s8 : vector<8xi32> + %c17_i32 = scalar.constant 17 : i32 + %c17v = vector.splat %c17_i32 : vector<8xi32> + %cv = vector.splat %code : vector<8xi32> + %cs = vector.shrui %cv, %shift2 : vector<8xi32> + %c = vector.andi %cs, %three : vector<8xi32> + %c17 = vector.muli %c, %c17v : vector<8xi32> + %chi = vector.shrui %c, %one : vector<8xi32> + %lv0 = vector.addi %c17, %chi : vector<8xi32> + %lv = vector.addi %lv0, %c8v : vector<8xi32> + %sv = vector.splat %signs8 : vector<8xi32> + %sb0 = vector.shrui %sv, %shift1 : vector<8xi32> + %sb = vector.andi %sb0, %one : vector<8xi32> + %sb2 = vector.shli %sb, %one : vector<8xi32> + %sgn = vector.subi %one, %sb2 : vector<8xi32> + %v = vector.muli %lv, %sgn : vector<8xi32> + %vf = vector.sitofp %v : vector<8xi32> to vector<8xf32> + func.return %vf : vector<8xf32> +} + +// IQ2_S (82 bytes: d, qs[32] grid low bits, signs[32], qh[8], scales[8]): value = d * (0.5 + nibble) * 0.25 +// * grid * sign. +func.def inline @ggml_kquant_iq2s_lane_parts(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c34 = index.constant 34 : index + %c66 = index.constant 66 : index + %c74 = index.constant 74 : index + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c768_i32 = scalar.constant 768 : i32 + %c05 = scalar.constant 0.5 : f32 + %c025 = scalar.constant 0.25 : f32 + %block_bytes = index.constant 82 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %g = index.div %l, %c2 : index + %h = index.rem %l, %c2 : index + %sa = index.mul %h, %c2 : index + %sb = index.add %sa, %c1 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %hv = buffer.view %weight[%block_base] : buffer -> view<41xf16> + %bv = buffer.view %weight[%block_base] : buffer -> view<82xi8> + %d_f16 = view.load %hv[%c0] : view<41xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %g4 = index.mul %g, %c4 : index + %qs0 = index.add %g4, %c2 : index + %qa_at = index.add %qs0, %sa : index + %qb_at = index.add %qs0, %sb : index + %sg0 = index.add %g4, %c34 : index + %sa_at = index.add %sg0, %sa : index + %sb_at = index.add %sg0, %sb : index + %qh_at = index.add %c66, %g : index + %sc_at = index.add %c74, %g : index + %qa_i8 = view.load %bv[%qa_at] : view<82xi8> -> i8 + %qb_i8 = view.load %bv[%qb_at] : view<82xi8> -> i8 + %sga_i8 = view.load %bv[%sa_at] : view<82xi8> -> i8 + %sgb_i8 = view.load %bv[%sb_at] : view<82xi8> -> i8 + %qh_i8 = view.load %bv[%qh_at] : view<82xi8> -> i8 + %sc_i8 = view.load %bv[%sc_at] : view<82xi8> -> i8 + %qa = scalar.extui %qa_i8 : i8 to i32 + %qb = scalar.extui %qb_i8 : i8 to i32 + %sga = scalar.extui %sga_i8 : i8 to i32 + %sgb = scalar.extui %sgb_i8 : i8 to i32 + %qh = scalar.extui %qh_i8 : i8 to i32 + %sc = scalar.extui %sc_i8 : i8 to i32 + %sa_i32 = index.cast %sa : index to i32 + %sb_i32 = index.cast %sb : index to i32 + %sa2 = scalar.muli %sa_i32, %c2_i32 : i32 + %sb2 = scalar.muli %sb_i32, %c2_i32 : i32 + %sha = scalar.subi %c8_i32, %sa2 : i32 + %shb = scalar.subi %c8_i32, %sb2 : i32 + %ha0 = scalar.shli %qh, %sha : i32 + %hb0 = scalar.shli %qh, %shb : i32 + %ha = scalar.andi %ha0, %c768_i32 : i32 + %hb = scalar.andi %hb0, %c768_i32 : i32 + %ia = scalar.ori %qa, %ha : i32 + %ib = scalar.ori %qb, %hb : i32 + %code_a = func.call @ggml_kquant_iq2s_code(%grid, %ia) : (buffer, i32) -> (i32) + %code_b = func.call @ggml_kquant_iq2s_code(%grid, %ib) : (buffer, i32) -> (i32) + %h_i32 = index.cast %h : index to i32 + %nsh = scalar.muli %h_i32, %c4_i32 : i32 + %nib0 = scalar.shrui %sc, %nsh : i32 + %nib = scalar.andi %nib0, %c15_i32 : i32 + %nib_f = scalar.uitofp %nib : i32 to f32 + %nib_p = scalar.addf %nib_f, %c05 : f32 + %ds = scalar.mulf %d, %nib_p : f32 + %scale = scalar.mulf %ds, %c025 : f32 + %va = func.call @ggml_kquant_iq2s_slot_values(%code_a, %sga) : (i32, i32) -> (vector<8xf32>) + %vb = func.call @ggml_kquant_iq2s_slot_values(%code_b, %sgb) : (i32, i32) -> (vector<8xf32>) + %v = func.call @ggml_kquant_iq2_join16(%va, %vb) : (vector<8xf32>, vector<8xf32>) -> (vector<16xf32>) + %g32 = index.mul %g, %c32 : index + %h16 = index.mul %h, %c16 : index + %p0 = index.add %g32, %h16 : index + func.return %v, %scale, %p0 : vector<16xf32>, f32, index +} + +func.def inline @ggml_kquant_iq2s_lane_dot(%grid: buffer, %weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %zero_scalar = scalar.constant 0.0 : f32 + %x_block_bytes = index.constant 1024 : offset + %v, %scale, %p = func.call @ggml_kquant_iq2s_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %x = vector.load %xv[%p] : view<256xf32> -> vector<16xf32> + %vx = vector.mulf %v, %x : vector<16xf32> + %sum = vector.reduce %vx, %zero_scalar : vector<16xf32>, f32 + %result = scalar.mulf %scale, %sum : f32 + func.return %result : f32 +} + +func.def inline @ggml_kquant_iq2s_lane_weights(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c12 = index.constant 12 : index + %v, %scale, %p0 = func.call @ggml_kquant_iq2s_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %sv = vector.splat %scale : vector<16xf32> + %w = vector.mulf %v, %sv : vector<16xf32> + %p1 = index.add %p0, %c4 : index + %p2 = index.add %p0, %c8 : index + %p3 = index.add %p0, %c12 : index + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +// Stages the IQ1 grid as 1024 words (two 16-bit codes each); subgroup %chunk (0..3) writes +// words 256 %chunk .. +255. +func.def inline @ggml_kquant_iq1s_grid_fill(%grid: buffer, %chunk: index) { + %zero_offset = index.constant 0 : offset + %gv = buffer.view %grid[%zero_offset] : buffer -> view<1024xi32> + %k0 = index.constant 0 : index + %is0 = index.cmp eq, %chunk, %k0 : index + scf.if %is0 { + %v0_0_0 = scalar.constant 131072 : i32 + %v0_0_1 = scalar.constant 524293 : i32 + %v0_0_2 = scalar.constant 1114122 : i32 + %v0_0_3 = scalar.constant 2097173 : i32 + %w0_0 = vector.from_elements %v0_0_0, %v0_0_1, %v0_0_2, %v0_0_3 : vector<4xi32> + %o0_0 = index.constant 0 : index + vector.store %w0_0, %gv[%o0_0] : vector<4xi32>, view<1024xi32> + %v0_1_0 = scalar.constant 2621474 : i32 + %v0_1_1 = scalar.constant 4522026 : i32 + %v0_1_2 = scalar.constant 5505105 : i32 + %v0_1_3 = scalar.constant 6619222 : i32 + %w0_1 = vector.from_elements %v0_1_0, %v0_1_1, %v0_1_2, %v0_1_3 : vector<4xi32> + %o0_1 = index.constant 4 : index + vector.store %w0_1, %gv[%o0_1] : vector<4xi32>, view<1024xi32> + %v0_2_0 = scalar.constant 8519808 : i32 + %v0_2_1 = scalar.constant 9044104 : i32 + %v0_2_2 = scalar.constant 10485909 : i32 + %v0_2_3 = scalar.constant 11010210 : i32 + %w0_2 = vector.from_elements %v0_2_0, %v0_2_1, %v0_2_2, %v0_2_3 : vector<4xi32> + %o0_2 = index.constant 8 : index + vector.store %w0_2, %gv[%o0_2] : vector<4xi32>, view<1024xi32> + %v0_3_0 = scalar.constant 17039530 : i32 + %v0_3_1 = scalar.constant 17891589 : i32 + %v0_3_2 = scalar.constant 18219284 : i32 + %v0_3_3 = scalar.constant 18481433 : i32 + %w0_3 = vector.from_elements %v0_3_0, %v0_3_1, %v0_3_2, %v0_3_3 : vector<4xi32> + %o0_3 = index.constant 12 : index + vector.store %w0_3, %gv[%o0_3] : vector<4xi32>, view<1024xi32> + %v0_4_0 = scalar.constant 21037349 : i32 + %v0_4_1 = scalar.constant 21561670 : i32 + %v0_4_2 = scalar.constant 22348114 : i32 + %v0_4_3 = scalar.constant 23134554 : i32 + %w0_4 = vector.from_elements %v0_4_0, %v0_4_1, %v0_4_2, %v0_4_3 : vector<4xi32> + %o0_4 = index.constant 16 : index + vector.store %w0_4, %gv[%o0_4] : vector<4xi32>, view<1024xi32> + %v0_5_0 = scalar.constant 23462244 : i32 + %v0_5_1 = scalar.constant 25493864 : i32 + %v0_5_2 = scalar.constant 26476945 : i32 + %v0_5_3 = scalar.constant 27591062 : i32 + %w0_5 = vector.from_elements %v0_5_0, %v0_5_1, %v0_5_2, %v0_5_3 : vector<4xi32> + %o0_5 = index.constant 20 : index + vector.store %w0_5, %gv[%o0_5] : vector<4xi32>, view<1024xi32> + %v0_6_0 = scalar.constant 33686016 : i32 + %v0_6_1 = scalar.constant 34210312 : i32 + %v0_6_2 = scalar.constant 35652117 : i32 + %v0_6_3 = scalar.constant 36176418 : i32 + %w0_6 = vector.from_elements %v0_6_0, %v0_6_1, %v0_6_2, %v0_6_3 : vector<4xi32> + %o0_6 = index.constant 24 : index + vector.store %w0_6, %gv[%o0_6] : vector<4xi32>, view<1024xi32> + %v0_7_0 = scalar.constant 38076970 : i32 + %v0_7_1 = scalar.constant 39387729 : i32 + %v0_7_2 = scalar.constant 40436324 : i32 + %v0_7_3 = scalar.constant 42074752 : i32 + %w0_7 = vector.from_elements %v0_7_0, %v0_7_1, %v0_7_2, %v0_7_3 : vector<4xi32> + %o0_7 = index.constant 28 : index + vector.store %w0_7, %gv[%o0_7] : vector<4xi32>, view<1024xi32> + %v0_8_0 = scalar.constant 42599048 : i32 + %v0_8_1 = scalar.constant 43319953 : i32 + %v0_8_2 = scalar.constant 44040857 : i32 + %v0_8_3 = scalar.constant 44565154 : i32 + %w0_8 = vector.from_elements %v0_8_0, %v0_8_1, %v0_8_2, %v0_8_3 : vector<4xi32> + %o0_8 = index.constant 32 : index + vector.store %w0_8, %gv[%o0_8] : vector<4xi32>, view<1024xi32> + %v0_9_0 = scalar.constant 68223658 : i32 + %v0_9_1 = scalar.constant 68551700 : i32 + %v0_9_2 = scalar.constant 71369765 : i32 + %v0_9_3 = scalar.constant 72680521 : i32 + %w0_9 = vector.from_elements %v0_9_0, %v0_9_1, %v0_9_2, %v0_9_3 : vector<4xi32> + %o0_9 = index.constant 36 : index + vector.store %w0_9, %gv[%o0_9] : vector<4xi32>, view<1024xi32> + %v0_10_0 = scalar.constant 73663578 : i32 + %v0_10_1 = scalar.constant 76612709 : i32 + %v0_10_2 = scalar.constant 77923481 : i32 + %v0_10_3 = scalar.constant 84149505 : i32 + %w0_10 = vector.from_elements %v0_10_0, %v0_10_1, %v0_10_2, %v0_10_3 : vector<4xi32> + %o0_10 = index.constant 40 : index + vector.store %w0_10, %gv[%o0_10] : vector<4xi32>, view<1024xi32> + %v0_11_0 = scalar.constant 84280581 : i32 + %v0_11_1 = scalar.constant 85460245 : i32 + %v0_11_2 = scalar.constant 86574362 : i32 + %v0_11_3 = scalar.constant 88409408 : i32 + %w0_11 = vector.from_elements %v0_11_0, %v0_11_1, %v0_11_2, %v0_11_3 : vector<4xi32> + %o0_11 = index.constant 44 : index + vector.store %w0_11, %gv[%o0_11] : vector<4xi32>, view<1024xi32> + %v0_12_0 = scalar.constant 89130314 : i32 + %v0_12_1 = scalar.constant 89392465 : i32 + %v0_12_2 = scalar.constant 89523541 : i32 + %v0_12_3 = scalar.constant 90178905 : i32 + %w0_12 = vector.from_elements %v0_12_0, %v0_12_1, %v0_12_2, %v0_12_3 : vector<4xi32> + %o0_12 = index.constant 48 : index + vector.store %w0_12, %gv[%o0_12] : vector<4xi32>, view<1024xi32> + %v0_13_0 = scalar.constant 90506594 : i32 + %v0_13_1 = scalar.constant 90834280 : i32 + %v0_13_2 = scalar.constant 93390209 : i32 + %v0_13_3 = scalar.constant 93848981 : i32 + %w0_13 = vector.from_elements %v0_13_0, %v0_13_1, %v0_13_2, %v0_13_3 : vector<4xi32> + %o0_13 = index.constant 52 : index + vector.store %w0_13, %gv[%o0_13] : vector<4xi32>, view<1024xi32> + %v0_14_0 = scalar.constant 94438810 : i32 + %v0_14_1 = scalar.constant 94700964 : i32 + %v0_14_2 = scalar.constant 94963110 : i32 + %v0_14_3 = scalar.constant 102303252 : i32 + %w0_14 = vector.from_elements %v0_14_0, %v0_14_1, %v0_14_2, %v0_14_3 : vector<4xi32> + %o0_14 = index.constant 56 : index + vector.store %w0_14, %gv[%o0_14] : vector<4xi32>, view<1024xi32> + %v0_15_0 = scalar.constant 105121345 : i32 + %v0_15_1 = scalar.constant 106038864 : i32 + %v0_15_2 = scalar.constant 106432085 : i32 + %v0_15_3 = scalar.constant 107021920 : i32 + %w0_15 = vector.from_elements %v0_15_0, %v0_15_1, %v0_15_2, %v0_15_3 : vector<4xi32> + %o0_15 = index.constant 60 : index + vector.store %w0_15, %gv[%o0_15] : vector<4xi32>, view<1024xi32> + %v0_16_0 = scalar.constant 107546214 : i32 + %v0_16_1 = scalar.constant 110167685 : i32 + %v0_16_2 = scalar.constant 110691988 : i32 + %v0_16_3 = scalar.constant 134350848 : i32 + %w0_16 = vector.from_elements %v0_16_0, %v0_16_1, %v0_16_2, %v0_16_3 : vector<4xi32> + %o0_16 = index.constant 64 : index + vector.store %w0_16, %gv[%o0_16] : vector<4xi32>, view<1024xi32> + %v0_17_0 = scalar.constant 134875144 : i32 + %v0_17_1 = scalar.constant 136316949 : i32 + %v0_17_2 = scalar.constant 136841250 : i32 + %v0_17_3 = scalar.constant 138741802 : i32 + %w0_17 = vector.from_elements %v0_17_0, %v0_17_1, %v0_17_2, %v0_17_3 : vector<4xi32> + %o0_17 = index.constant 68 : index + vector.store %w0_17, %gv[%o0_17] : vector<4xi32>, view<1024xi32> + %v0_18_0 = scalar.constant 139855953 : i32 + %v0_18_1 = scalar.constant 142608485 : i32 + %v0_18_2 = scalar.constant 143132802 : i32 + %v0_18_3 = scalar.constant 143984778 : i32 + %w0_18 = vector.from_elements %v0_18_0, %v0_18_1, %v0_18_2, %v0_18_3 : vector<4xi32> + %o0_18 = index.constant 72 : index + vector.store %w0_18, %gv[%o0_18] : vector<4xi32>, view<1024xi32> + %v0_19_0 = scalar.constant 144836768 : i32 + %v0_19_1 = scalar.constant 145361064 : i32 + %v0_19_2 = scalar.constant 152111365 : i32 + %v0_19_3 = scalar.constant 152635668 : i32 + %w0_19 = vector.from_elements %v0_19_0, %v0_19_1, %v0_19_2, %v0_19_3 : vector<4xi32> + %o0_19 = index.constant 76 : index + vector.store %w0_19, %gv[%o0_19] : vector<4xi32>, view<1024xi32> + %v0_20_0 = scalar.constant 153422116 : i32 + %v0_20_1 = scalar.constant 156240193 : i32 + %v0_20_2 = scalar.constant 156567889 : i32 + %v0_20_3 = scalar.constant 157550945 : i32 + %w0_20 = vector.from_elements %v0_20_0, %v0_20_1, %v0_20_2, %v0_20_3 : vector<4xi32> + %o0_20 = index.constant 80 : index + vector.store %w0_20, %gv[%o0_20] : vector<4xi32>, view<1024xi32> + %v0_21_0 = scalar.constant 160500073 : i32 + %v0_21_1 = scalar.constant 160827796 : i32 + %v0_21_2 = scalar.constant 161810841 : i32 + %v0_21_3 = scalar.constant 167905792 : i32 + %w0_21 = vector.from_elements %v0_21_0, %v0_21_1, %v0_21_2, %v0_21_3 : vector<4xi32> + %o0_21 = index.constant 84 : index + vector.store %w0_21, %gv[%o0_21] : vector<4xi32>, view<1024xi32> + %v0_22_0 = scalar.constant 168430088 : i32 + %v0_22_1 = scalar.constant 169871893 : i32 + %v0_22_2 = scalar.constant 170396194 : i32 + %v0_22_3 = scalar.constant 172296746 : i32 + %w0_22 = vector.from_elements %v0_22_0, %v0_22_1, %v0_22_2, %v0_22_3 : vector<4xi32> + %o0_22 = index.constant 88 : index + vector.store %w0_22, %gv[%o0_22] : vector<4xi32>, view<1024xi32> + %v0_23_0 = scalar.constant 173607505 : i32 + %v0_23_1 = scalar.constant 174393953 : i32 + %v0_23_2 = scalar.constant 176294528 : i32 + %v0_23_3 = scalar.constant 176687749 : i32 + %w0_23 = vector.from_elements %v0_23_0, %v0_23_1, %v0_23_2, %v0_23_3 : vector<4xi32> + %o0_23 = index.constant 92 : index + vector.store %w0_23, %gv[%o0_23] : vector<4xi32>, view<1024xi32> + %v0_24_0 = scalar.constant 177539722 : i32 + %v0_24_1 = scalar.constant 178391712 : i32 + %v0_24_2 = scalar.constant 178916008 : i32 + %v0_24_3 = scalar.constant 269553680 : i32 + %w0_24 = vector.from_elements %v0_24_0, %v0_24_1, %v0_24_2, %v0_24_3 : vector<4xi32> + %o0_24 = index.constant 96 : index + vector.store %w0_24, %gv[%o0_24] : vector<4xi32>, view<1024xi32> + %v0_25_0 = scalar.constant 270077972 : i32 + %v0_25_1 = scalar.constant 270864420 : i32 + %v0_25_2 = scalar.constant 272896065 : i32 + %v0_25_3 = scalar.constant 274010192 : i32 + %w0_25 = vector.from_elements %v0_25_0, %v0_25_1, %v0_25_2, %v0_25_3 : vector<4xi32> + %o0_25 = index.constant 100 : index + vector.store %w0_25, %gv[%o0_25] : vector<4xi32>, view<1024xi32> + %v0_26_0 = scalar.constant 274796632 : i32 + %v0_26_1 = scalar.constant 275058788 : i32 + %v0_26_2 = scalar.constant 277942377 : i32 + %v0_26_3 = scalar.constant 278270100 : i32 + %w0_26 = vector.from_elements %v0_26_0, %v0_26_1, %v0_26_2, %v0_26_3 : vector<4xi32> + %o0_26 = index.constant 104 : index + vector.store %w0_26, %gv[%o0_26] : vector<4xi32>, view<1024xi32> + %v0_27_0 = scalar.constant 279253153 : i32 + %v0_27_1 = scalar.constant 285479169 : i32 + %v0_27_2 = scalar.constant 285806854 : i32 + %v0_27_3 = scalar.constant 286396688 : i32 + %w0_27 = vector.from_elements %v0_27_0, %v0_27_1, %v0_27_2, %v0_27_3 : vector<4xi32> + %o0_27 = index.constant 108 : index + vector.store %w0_27, %gv[%o0_27] : vector<4xi32>, view<1024xi32> + %v0_28_0 = scalar.constant 286789909 : i32 + %v0_28_1 = scalar.constant 287576353 : i32 + %v0_28_2 = scalar.constant 289739049 : i32 + %v0_28_3 = scalar.constant 290459978 : i32 + %w0_28 = vector.from_elements %v0_28_0, %v0_28_1, %v0_28_2, %v0_28_3 : vector<4xi32> + %o0_28 = index.constant 112 : index + vector.store %w0_28, %gv[%o0_28] : vector<4xi32>, view<1024xi32> + %v0_29_0 = scalar.constant 290591057 : i32 + %v0_29_1 = scalar.constant 290787668 : i32 + %v0_29_2 = scalar.constant 291049814 : i32 + %v0_29_3 = scalar.constant 291836256 : i32 + %w0_29 = vector.from_elements %v0_29_0, %v0_29_1, %v0_29_2, %v0_29_3 : vector<4xi32> + %o0_29 = index.constant 116 : index + vector.store %w0_29, %gv[%o0_29] : vector<4xi32>, view<1024xi32> + %v0_30_0 = scalar.constant 294785412 : i32 + %v0_30_1 = scalar.constant 295768469 : i32 + %v0_30_2 = scalar.constant 303108516 : i32 + %v0_30_3 = scalar.constant 303436308 : i32 + %w0_30 = vector.from_elements %v0_30_0, %v0_30_1, %v0_30_2, %v0_30_3 : vector<4xi32> + %o0_30 = index.constant 120 : index + vector.store %w0_30, %gv[%o0_30] : vector<4xi32>, view<1024xi32> + %v0_31_0 = scalar.constant 306188837 : i32 + %v0_31_1 = scalar.constant 306778694 : i32 + %v0_31_2 = scalar.constant 307565138 : i32 + %v0_31_3 = scalar.constant 307892824 : i32 + %w0_31 = vector.from_elements %v0_31_0, %v0_31_1, %v0_31_2, %v0_31_3 : vector<4xi32> + %o0_31 = index.constant 124 : index + vector.store %w0_31, %gv[%o0_31] : vector<4xi32>, view<1024xi32> + %v0_32_0 = scalar.constant 308679268 : i32 + %v0_32_1 = scalar.constant 311497349 : i32 + %v0_32_2 = scalar.constant 311825044 : i32 + %v0_32_3 = scalar.constant 335614629 : i32 + %w0_32 = vector.from_elements %v0_32_0, %v0_32_1, %v0_32_2, %v0_32_3 : vector<4xi32> + %o0_32 = index.constant 128 : index + vector.store %w0_32, %gv[%o0_32] : vector<4xi32>, view<1024xi32> + %v0_33_0 = scalar.constant 336139270 : i32 + %v0_33_1 = scalar.constant 336925716 : i32 + %v0_33_2 = scalar.constant 337187864 : i32 + %v0_33_3 = scalar.constant 338039841 : i32 + %w0_33 = vector.from_elements %v0_33_0, %v0_33_1, %v0_33_2, %v0_33_3 : vector<4xi32> + %o0_33 = index.constant 132 : index + vector.store %w0_33, %gv[%o0_33] : vector<4xi32>, view<1024xi32> + %v0_34_0 = scalar.constant 340071489 : i32 + %v0_34_1 = scalar.constant 340268102 : i32 + %v0_34_2 = scalar.constant 340857930 : i32 + %v0_34_3 = scalar.constant 341120084 : i32 + %w0_34 = vector.from_elements %v0_34_0, %v0_34_1, %v0_34_2, %v0_34_3 : vector<4xi32> + %o0_34 = index.constant 136 : index + vector.store %w0_34, %gv[%o0_34] : vector<4xi32>, view<1024xi32> + %v0_35_0 = scalar.constant 341382230 : i32 + %v0_35_1 = scalar.constant 342168674 : i32 + %v0_35_2 = scalar.constant 344200296 : i32 + %v0_35_3 = scalar.constant 344986761 : i32 + %w0_35 = vector.from_elements %v0_35_0, %v0_35_1, %v0_35_2, %v0_35_3 : vector<4xi32> + %o0_35 = index.constant 140 : index + vector.store %w0_35, %gv[%o0_35] : vector<4xi32>, view<1024xi32> + %v0_36_0 = scalar.constant 345314452 : i32 + %v0_36_1 = scalar.constant 345576600 : i32 + %v0_36_2 = scalar.constant 346100890 : i32 + %v0_36_3 = scalar.constant 346363044 : i32 + %w0_36 = vector.from_elements %v0_36_0, %v0_36_1, %v0_36_2, %v0_36_3 : vector<4xi32> + %o0_36 = index.constant 144 : index + vector.store %w0_36, %gv[%o0_36] : vector<4xi32>, view<1024xi32> + %v0_37_0 = scalar.constant 352457897 : i32 + %v0_37_1 = scalar.constant 352982277 : i32 + %v0_37_2 = scalar.constant 353637649 : i32 + %v0_37_3 = scalar.constant 353768725 : i32 + %w0_37 = vector.from_elements %v0_37_0, %v0_37_1, %v0_37_2, %v0_37_3 : vector<4xi32> + %o0_37 = index.constant 148 : index + vector.store %w0_37, %gv[%o0_37] : vector<4xi32>, view<1024xi32> + %v0_38_0 = scalar.constant 354424089 : i32 + %v0_38_1 = scalar.constant 354751778 : i32 + %v0_38_2 = scalar.constant 355079464 : i32 + %v0_38_3 = scalar.constant 356783425 : i32 + %w0_38 = vector.from_elements %v0_38_0, %v0_38_1, %v0_38_2, %v0_38_3 : vector<4xi32> + %o0_38 = index.constant 152 : index + vector.store %w0_38, %gv[%o0_38] : vector<4xi32>, view<1024xi32> + %v0_39_0 = scalar.constant 356914501 : i32 + %v0_39_1 = scalar.constant 357700945 : i32 + %v0_39_2 = scalar.constant 357897556 : i32 + %v0_39_3 = scalar.constant 358159702 : i32 + %w0_39 = vector.from_elements %v0_39_0, %v0_39_1, %v0_39_2, %v0_39_3 : vector<4xi32> + %o0_39 = index.constant 156 : index + vector.store %w0_39, %gv[%o0_39] : vector<4xi32>, view<1024xi32> + %v0_40_0 = scalar.constant 358683994 : i32 + %v0_40_1 = scalar.constant 358946148 : i32 + %v0_40_2 = scalar.constant 359208294 : i32 + %v0_40_3 = scalar.constant 360846720 : i32 + %w0_40 = vector.from_elements %v0_40_0, %v0_40_1, %v0_40_2, %v0_40_3 : vector<4xi32> + %o0_40 = index.constant 160 : index + vector.store %w0_40, %gv[%o0_40] : vector<4xi32>, view<1024xi32> + %v0_41_0 = scalar.constant 361043332 : i32 + %v0_41_1 = scalar.constant 361371016 : i32 + %v0_41_2 = scalar.constant 361829776 : i32 + %v0_41_3 = scalar.constant 362091924 : i32 + %w0_41 = vector.from_elements %v0_41_0, %v0_41_1, %v0_41_2, %v0_41_3 : vector<4xi32> + %o0_41 = index.constant 164 : index + vector.store %w0_41, %gv[%o0_41] : vector<4xi32>, view<1024xi32> + %v0_42_0 = scalar.constant 362354070 : i32 + %v0_42_1 = scalar.constant 362812826 : i32 + %v0_42_2 = scalar.constant 363140514 : i32 + %v0_42_3 = scalar.constant 369366529 : i32 + %w0_42 = vector.from_elements %v0_42_0, %v0_42_1, %v0_42_2, %v0_42_3 : vector<4xi32> + %o0_42 = index.constant 168 : index + vector.store %w0_42, %gv[%o0_42] : vector<4xi32>, view<1024xi32> + %v0_43_0 = scalar.constant 369497605 : i32 + %v0_43_1 = scalar.constant 370546197 : i32 + %v0_43_2 = scalar.constant 370808344 : i32 + %v0_43_3 = scalar.constant 371594785 : i32 + %w0_43 = vector.from_elements %v0_43_0, %v0_43_1, %v0_43_2, %v0_43_3 : vector<4xi32> + %o0_43 = index.constant 172 : index + vector.store %w0_43, %gv[%o0_43] : vector<4xi32>, view<1024xi32> + %v0_44_0 = scalar.constant 373429824 : i32 + %v0_44_1 = scalar.constant 373626436 : i32 + %v0_44_2 = scalar.constant 373954120 : i32 + %v0_44_3 = scalar.constant 374675025 : i32 + %w0_44 = vector.from_elements %v0_44_0, %v0_44_1, %v0_44_2, %v0_44_3 : vector<4xi32> + %o0_44 = index.constant 176 : index + vector.store %w0_44, %gv[%o0_44] : vector<4xi32>, view<1024xi32> + %v0_45_0 = scalar.constant 374871638 : i32 + %v0_45_1 = scalar.constant 375461465 : i32 + %v0_45_2 = scalar.constant 375723620 : i32 + %v0_45_3 = scalar.constant 375985768 : i32 + %w0_45 = vector.from_elements %v0_45_0, %v0_45_1, %v0_45_2, %v0_45_3 : vector<4xi32> + %o0_45 = index.constant 180 : index + vector.store %w0_45, %gv[%o0_45] : vector<4xi32>, view<1024xi32> + %v0_46_0 = scalar.constant 377886314 : i32 + %v0_46_1 = scalar.constant 378672778 : i32 + %v0_46_2 = scalar.constant 379852437 : i32 + %v0_46_3 = scalar.constant 403773097 : i32 + %w0_46 = vector.from_elements %v0_46_0, %v0_46_1, %v0_46_2, %v0_46_3 : vector<4xi32> + %o0_46 = index.constant 184 : index + vector.store %w0_46, %gv[%o0_46] : vector<4xi32>, view<1024xi32> + %v0_47_0 = scalar.constant 405084182 : i32 + %v0_47_1 = scalar.constant 407115841 : i32 + %v0_47_2 = scalar.constant 407443526 : i32 + %v0_47_3 = scalar.constant 408229968 : i32 + %w0_47 = vector.from_elements %v0_47_0, %v0_47_1, %v0_47_2, %v0_47_3 : vector<4xi32> + %o0_47 = index.constant 188 : index + vector.store %w0_47, %gv[%o0_47] : vector<4xi32>, view<1024xi32> + %v0_48_0 = scalar.constant 408557656 : i32 + %v0_48_1 = scalar.constant 409016416 : i32 + %v0_48_2 = scalar.constant 409344100 : i32 + %v0_48_3 = scalar.constant 411375721 : i32 + %w0_48 = vector.from_elements %v0_48_0, %v0_48_1, %v0_48_2, %v0_48_3 : vector<4xi32> + %o0_48 = index.constant 192 : index + vector.store %w0_48, %gv[%o0_48] : vector<4xi32>, view<1024xi32> + %v0_49_0 = scalar.constant 412358801 : i32 + %v0_49_1 = scalar.constant 420485285 : i32 + %v0_49_2 = scalar.constant 420813074 : i32 + %v0_49_3 = scalar.constant 421599514 : i32 + %w0_49 = vector.from_elements %v0_49_0, %v0_49_1, %v0_49_2, %v0_49_3 : vector<4xi32> + %o0_49 = index.constant 196 : index + vector.store %w0_49, %gv[%o0_49] : vector<4xi32>, view<1024xi32> + %v0_50_0 = scalar.constant 423762213 : i32 + %v0_50_1 = scalar.constant 423958852 : i32 + %v0_50_2 = scalar.constant 424745288 : i32 + %v0_50_3 = scalar.constant 425007444 : i32 + %w0_50 = vector.from_elements %v0_50_0, %v0_50_1, %v0_50_2, %v0_50_3 : vector<4xi32> + %o0_50 = index.constant 200 : index + vector.store %w0_50, %gv[%o0_50] : vector<4xi32>, view<1024xi32> + %v0_51_0 = scalar.constant 425269590 : i32 + %v0_51_1 = scalar.constant 425728346 : i32 + %v0_51_2 = scalar.constant 426383717 : i32 + %v0_51_3 = scalar.constant 428939657 : i32 + %w0_51 = vector.from_elements %v0_51_0, %v0_51_1, %v0_51_2, %v0_51_3 : vector<4xi32> + %o0_51 = index.constant 204 : index + vector.store %w0_51, %gv[%o0_51] : vector<4xi32>, view<1024xi32> + %v0_52_0 = scalar.constant 429201810 : i32 + %v0_52_1 = scalar.constant 429988248 : i32 + %v0_52_2 = scalar.constant 430512550 : i32 + %v0_52_3 = scalar.constant 437656073 : i32 + %w0_52 = vector.from_elements %v0_52_0, %v0_52_1, %v0_52_2, %v0_52_3 : vector<4xi32> + %o0_52 = index.constant 208 : index + vector.store %w0_52, %gv[%o0_52] : vector<4xi32>, view<1024xi32> + %v0_53_0 = scalar.constant 438704676 : i32 + %v0_53_1 = scalar.constant 440801860 : i32 + %v0_53_2 = scalar.constant 441457225 : i32 + %v0_53_3 = scalar.constant 441784914 : i32 + %w0_53 = vector.from_elements %v0_53_0, %v0_53_1, %v0_53_2, %v0_53_3 : vector<4xi32> + %o0_53 = index.constant 212 : index + vector.store %w0_53, %gv[%o0_53] : vector<4xi32>, view<1024xi32> + %v0_54_0 = scalar.constant 442571352 : i32 + %v0_54_1 = scalar.constant 443095654 : i32 + %v0_54_2 = scalar.constant 445717125 : i32 + %v0_54_3 = scalar.constant 446306966 : i32 + %w0_54 = vector.from_elements %v0_54_0, %v0_54_1, %v0_54_2, %v0_54_3 : vector<4xi32> + %o0_54 = index.constant 216 : index + vector.store %w0_54, %gv[%o0_54] : vector<4xi32>, view<1024xi32> + %v0_55_0 = scalar.constant 537010176 : i32 + %v0_55_1 = scalar.constant 537534472 : i32 + %v0_55_2 = scalar.constant 538976277 : i32 + %v0_55_3 = scalar.constant 539303970 : i32 + %w0_55 = vector.from_elements %v0_55_0, %v0_55_1, %v0_55_2, %v0_55_3 : vector<4xi32> + %o0_55 = index.constant 220 : index + vector.store %w0_55, %gv[%o0_55] : vector<4xi32>, view<1024xi32> + %v0_56_0 = scalar.constant 539631656 : i32 + %v0_56_1 = scalar.constant 542187589 : i32 + %v0_56_2 = scalar.constant 543236185 : i32 + %v0_56_3 = scalar.constant 545267813 : i32 + %w0_56 = vector.from_elements %v0_56_0, %v0_56_1, %v0_56_2, %v0_56_3 : vector<4xi32> + %o0_56 = index.constant 224 : index + vector.store %w0_56, %gv[%o0_56] : vector<4xi32>, view<1024xi32> + %v0_57_0 = scalar.constant 545792130 : i32 + %v0_57_1 = scalar.constant 546644106 : i32 + %v0_57_2 = scalar.constant 547496096 : i32 + %v0_57_3 = scalar.constant 547889317 : i32 + %w0_57 = vector.from_elements %v0_57_0, %v0_57_1, %v0_57_2, %v0_57_3 : vector<4xi32> + %o0_57 = index.constant 228 : index + vector.store %w0_57, %gv[%o0_57] : vector<4xi32>, view<1024xi32> + %v0_58_0 = scalar.constant 553984170 : i32 + %v0_58_1 = scalar.constant 554967313 : i32 + %v0_58_2 = scalar.constant 556081433 : i32 + %v0_58_3 = scalar.constant 558113090 : i32 + %w0_58 = vector.from_elements %v0_58_0, %v0_58_1, %v0_58_2, %v0_58_3 : vector<4xi32> + %o0_58 = index.constant 232 : index + vector.store %w0_58, %gv[%o0_58] : vector<4xi32>, view<1024xi32> + %v0_59_0 = scalar.constant 559227209 : i32 + %v0_59_1 = scalar.constant 559554904 : i32 + %v0_59_2 = scalar.constant 560210273 : i32 + %v0_59_3 = scalar.constant 560341349 : i32 + %w0_59 = vector.from_elements %v0_59_0, %v0_59_1, %v0_59_2, %v0_59_3 : vector<4xi32> + %o0_59 = index.constant 236 : index + vector.store %w0_59, %gv[%o0_59] : vector<4xi32>, view<1024xi32> + %v0_60_0 = scalar.constant 563093893 : i32 + %v0_60_1 = scalar.constant 563683734 : i32 + %v0_60_2 = scalar.constant 570499493 : i32 + %v0_60_3 = scalar.constant 571089416 : i32 + %w0_60 = vector.from_elements %v0_60_0, %v0_60_1, %v0_60_2, %v0_60_3 : vector<4xi32> + %o0_60 = index.constant 240 : index + vector.store %w0_60, %gv[%o0_60] : vector<4xi32>, view<1024xi32> + %v0_61_0 = scalar.constant 571810321 : i32 + %v0_61_1 = scalar.constant 572662304 : i32 + %v0_61_2 = scalar.constant 573186600 : i32 + %v0_61_3 = scalar.constant 575742533 : i32 + %w0_61 = vector.from_elements %v0_61_0, %v0_61_1, %v0_61_2, %v0_61_3 : vector<4xi32> + %o0_61 = index.constant 244 : index + vector.store %w0_61, %gv[%o0_61] : vector<4xi32>, view<1024xi32> + %v0_62_0 = scalar.constant 576266838 : i32 + %v0_62_1 = scalar.constant 578888293 : i32 + %v0_62_2 = scalar.constant 579478152 : i32 + %v0_62_3 = scalar.constant 580199057 : i32 + %w0_62 = vector.from_elements %v0_62_0, %v0_62_1, %v0_62_2, %v0_62_3 : vector<4xi32> + %o0_62 = index.constant 248 : index + vector.store %w0_62, %gv[%o0_62] : vector<4xi32>, view<1024xi32> + %v0_63_0 = scalar.constant 581051040 : i32 + %v0_63_1 = scalar.constant 581575336 : i32 + %v0_63_2 = scalar.constant 605299717 : i32 + %v0_63_3 = scalar.constant 605627414 : i32 + %w0_63 = vector.from_elements %v0_63_0, %v0_63_1, %v0_63_2, %v0_63_3 : vector<4xi32> + %o0_63 = index.constant 252 : index + vector.store %w0_63, %gv[%o0_63] : vector<4xi32>, view<1024xi32> + } + %k1 = index.constant 1 : index + %is1 = index.cmp eq, %chunk, %k1 : index + scf.if %is1 { + %v1_0_0 = scalar.constant 608445477 : i32 + %v1_0_1 = scalar.constant 608576581 : i32 + %v1_0_2 = scalar.constant 609363017 : i32 + %v1_0_3 = scalar.constant 609756245 : i32 + %w1_0 = vector.from_elements %v1_0_0, %v1_0_1, %v1_0_2, %v1_0_3 : vector<4xi32> + %o1_0 = index.constant 256 : index + vector.store %w1_0, %gv[%o1_0] : vector<4xi32>, view<1024xi32> + %v1_1_0 = scalar.constant 610673754 : i32 + %v1_1_1 = scalar.constant 613491845 : i32 + %v1_1_2 = scalar.constant 614016148 : i32 + %v1_1_3 = scalar.constant 614802593 : i32 + %w1_1 = vector.from_elements %v1_1_0, %v1_1_1, %v1_1_2, %v1_1_3 : vector<4xi32> + %o1_1 = index.constant 260 : index + vector.store %w1_1, %gv[%o1_1] : vector<4xi32>, view<1024xi32> + %v1_2_0 = scalar.constant 622142729 : i32 + %v1_2_1 = scalar.constant 623453473 : i32 + %v1_2_2 = scalar.constant 625288512 : i32 + %v1_2_3 = scalar.constant 626074952 : i32 + %w1_2 = vector.from_elements %v1_2_0, %v1_2_1, %v1_2_2, %v1_2_3 : vector<4xi32> + %o1_2 = index.constant 264 : index + vector.store %w1_2, %gv[%o1_2] : vector<4xi32>, view<1024xi32> + %v1_3_0 = scalar.constant 626337108 : i32 + %v1_3_1 = scalar.constant 627189081 : i32 + %v1_3_2 = scalar.constant 627582309 : i32 + %v1_3_3 = scalar.constant 630203785 : i32 + %w1_3 = vector.from_elements %v1_3_0, %v1_3_1, %v1_3_2, %v1_3_3 : vector<4xi32> + %o1_3 = index.constant 268 : index + vector.store %w1_3, %gv[%o1_3] : vector<4xi32>, view<1024xi32> + %v1_4_0 = scalar.constant 630531476 : i32 + %v1_4_1 = scalar.constant 630859160 : i32 + %v1_4_2 = scalar.constant 631514529 : i32 + %v1_4_3 = scalar.constant 631842214 : i32 + %w1_4 = vector.from_elements %v1_4_0, %v1_4_1, %v1_4_2, %v1_4_3 : vector<4xi32> + %o1_4 = index.constant 272 : index + vector.store %w1_4, %gv[%o1_4] : vector<4xi32>, view<1024xi32> + %v1_5_0 = scalar.constant 638592517 : i32 + %v1_5_1 = scalar.constant 639182354 : i32 + %v1_5_2 = scalar.constant 641803813 : i32 + %v1_5_3 = scalar.constant 643114569 : i32 + %w1_5 = vector.from_elements %v1_5_0, %v1_5_1, %v1_5_2, %v1_5_3 : vector<4xi32> + %o1_5 = index.constant 276 : index + vector.store %w1_5, %gv[%o1_5] : vector<4xi32>, view<1024xi32> + %v1_6_0 = scalar.constant 643901024 : i32 + %v1_6_1 = scalar.constant 646194793 : i32 + %v1_6_2 = scalar.constant 646981254 : i32 + %v1_6_3 = scalar.constant 671098522 : i32 + %w1_6 = vector.from_elements %v1_6_0, %v1_6_1, %v1_6_2, %v1_6_3 : vector<4xi32> + %o1_6 = index.constant 280 : index + vector.store %w1_6, %gv[%o1_6] : vector<4xi32>, view<1024xi32> + %v1_7_0 = scalar.constant 671623170 : i32 + %v1_7_1 = scalar.constant 672475146 : i32 + %v1_7_2 = scalar.constant 673327136 : i32 + %v1_7_3 = scalar.constant 673851432 : i32 + %w1_7 = vector.from_elements %v1_7_0, %v1_7_1, %v1_7_2, %v1_7_3 : vector<4xi32> + %o1_7 = index.constant 284 : index + vector.store %w1_7, %gv[%o1_7] : vector<4xi32>, view<1024xi32> + %v1_8_0 = scalar.constant 676407365 : i32 + %v1_8_1 = scalar.constant 677718100 : i32 + %v1_8_2 = scalar.constant 679618688 : i32 + %v1_8_3 = scalar.constant 680142984 : i32 + %w1_8 = vector.from_elements %v1_8_0, %v1_8_1, %v1_8_2, %v1_8_3 : vector<4xi32> + %o1_8 = index.constant 288 : index + vector.store %w1_8, %gv[%o1_8] : vector<4xi32>, view<1024xi32> + %v1_9_0 = scalar.constant 681715872 : i32 + %v1_9_1 = scalar.constant 682240168 : i32 + %v1_9_2 = scalar.constant 688990473 : i32 + %v1_9_3 = scalar.constant 689514772 : i32 + %w1_9 = vector.from_elements %v1_9_0, %v1_9_1, %v1_9_2, %v1_9_3 : vector<4xi32> + %o1_9 = index.constant 292 : index + vector.store %w1_9, %gv[%o1_9] : vector<4xi32>, view<1024xi32> + %v1_10_0 = scalar.constant 692463909 : i32 + %v1_10_1 = scalar.constant 693250377 : i32 + %v1_10_2 = scalar.constant 694233429 : i32 + %v1_10_3 = scalar.constant 694561124 : i32 + %w1_10 = vector.from_elements %v1_10_0, %v1_10_1, %v1_10_2, %v1_10_3 : vector<4xi32> + %o1_10 = index.constant 296 : index + vector.store %w1_10, %gv[%o1_10] : vector<4xi32>, view<1024xi32> + %v1_11_0 = scalar.constant 696592745 : i32 + %v1_11_1 = scalar.constant 697706896 : i32 + %v1_11_2 = scalar.constant 698624409 : i32 + %v1_11_3 = scalar.constant 704653733 : i32 + %w1_11 = vector.from_elements %v1_11_0, %v1_11_1, %v1_11_2, %v1_11_3 : vector<4xi32> + %o1_11 = index.constant 300 : index + vector.store %w1_11, %gv[%o1_11] : vector<4xi32>, view<1024xi32> + %v1_12_0 = scalar.constant 705178114 : i32 + %v1_12_1 = scalar.constant 706750986 : i32 + %v1_12_2 = scalar.constant 707275298 : i32 + %v1_12_3 = scalar.constant 709175850 : i32 + %w1_12 = vector.from_elements %v1_12_0, %v1_12_1, %v1_12_2, %v1_12_3 : vector<4xi32> + %o1_12 = index.constant 304 : index + vector.store %w1_12, %gv[%o1_12] : vector<4xi32>, view<1024xi32> + %v1_13_0 = scalar.constant 710290001 : i32 + %v1_13_1 = scalar.constant 711273049 : i32 + %v1_13_2 = scalar.constant 713173632 : i32 + %v1_13_3 = scalar.constant 713697928 : i32 + %w1_13 = vector.from_elements %v1_13_0, %v1_13_1, %v1_13_2, %v1_13_3 : vector<4xi32> + %o1_13 = index.constant 308 : index + vector.store %w1_13, %gv[%o1_13] : vector<4xi32>, view<1024xi32> + %v1_14_0 = scalar.constant 715139733 : i32 + %v1_14_1 = scalar.constant 715664034 : i32 + %v1_14_2 = scalar.constant 1074080426 : i32 + %v1_14_3 = scalar.constant 1075200017 : i32 + %w1_14 = vector.from_elements %v1_14_0, %v1_14_1, %v1_14_2, %v1_14_3 : vector<4xi32> + %o1_14 = index.constant 312 : index + vector.store %w1_14, %gv[%o1_14] : vector<4xi32>, view<1024xi32> + %v1_15_0 = scalar.constant 1078542373 : i32 + %v1_15_1 = scalar.constant 1079328850 : i32 + %v1_15_2 = scalar.constant 1079656536 : i32 + %v1_15_3 = scalar.constant 1080311905 : i32 + %w1_15 = vector.from_elements %v1_15_0, %v1_15_1, %v1_15_2, %v1_15_3 : vector<4xi32> + %o1_15 = index.constant 316 : index + vector.store %w1_15, %gv[%o1_15] : vector<4xi32>, view<1024xi32> + %v1_16_0 = scalar.constant 1083457638 : i32 + %v1_16_1 = scalar.constant 1084309657 : i32 + %v1_16_2 = scalar.constant 1090535590 : i32 + %v1_16_3 = scalar.constant 1090797825 : i32 + %w1_16 = vector.from_elements %v1_16_0, %v1_16_1, %v1_16_2, %v1_16_3 : vector<4xi32> + %o1_16 = index.constant 320 : index + vector.store %w1_16, %gv[%o1_16] : vector<4xi32>, view<1024xi32> + %v1_17_0 = scalar.constant 1091125510 : i32 + %v1_17_1 = scalar.constant 1091911954 : i32 + %v1_17_2 = scalar.constant 1092108566 : i32 + %v1_17_3 = scalar.constant 1092698394 : i32 + %w1_17 = vector.from_elements %v1_17_0, %v1_17_1, %v1_17_2, %v1_17_3 : vector<4xi32> + %o1_17 = index.constant 324 : index + vector.store %w1_17, %gv[%o1_17] : vector<4xi32>, view<1024xi32> + %v1_18_0 = scalar.constant 1093222694 : i32 + %v1_18_1 = scalar.constant 1095254341 : i32 + %v1_18_2 = scalar.constant 1095844170 : i32 + %v1_18_3 = scalar.constant 1096106324 : i32 + %w1_18 = vector.from_elements %v1_18_0, %v1_18_1, %v1_18_2, %v1_18_3 : vector<4xi32> + %o1_18 = index.constant 328 : index + vector.store %w1_18, %gv[%o1_18] : vector<4xi32>, view<1024xi32> + %v1_19_0 = scalar.constant 1096368470 : i32 + %v1_19_1 = scalar.constant 1097154906 : i32 + %v1_19_2 = scalar.constant 1097482600 : i32 + %v1_19_3 = scalar.constant 1099186561 : i32 + %w1_19 = vector.from_elements %v1_19_0, %v1_19_1, %v1_19_2, %v1_19_3 : vector<4xi32> + %o1_19 = index.constant 332 : index + vector.store %w1_19, %gv[%o1_19] : vector<4xi32>, view<1024xi32> + %v1_20_0 = scalar.constant 1099972998 : i32 + %v1_20_1 = scalar.constant 1100300690 : i32 + %v1_20_2 = scalar.constant 1101087136 : i32 + %v1_20_3 = scalar.constant 1107640738 : i32 + %w1_20 = vector.from_elements %v1_20_0, %v1_20_1, %v1_20_2, %v1_20_3 : vector<4xi32> + %o1_20 = index.constant 336 : index + vector.store %w1_20, %gv[%o1_20] : vector<4xi32>, view<1024xi32> + %v1_21_0 = scalar.constant 1108623889 : i32 + %v1_21_1 = scalar.constant 1109738006 : i32 + %v1_21_2 = scalar.constant 1112687169 : i32 + %v1_21_3 = scalar.constant 1113211477 : i32 + %w1_21 = vector.from_elements %v1_21_0, %v1_21_1, %v1_21_2, %v1_21_3 : vector<4xi32> + %o1_21 = index.constant 340 : index + vector.store %w1_21, %gv[%o1_21] : vector<4xi32>, view<1024xi32> + %v1_22_0 = scalar.constant 1114194532 : i32 + %v1_22_1 = scalar.constant 1117012617 : i32 + %v1_22_2 = scalar.constant 1140933285 : i32 + %v1_22_3 = scalar.constant 1142506517 : i32 + %w1_22 = vector.from_elements %v1_22_0, %v1_22_1, %v1_22_2, %v1_22_3 : vector<4xi32> + %o1_22 = index.constant 344 : index + vector.store %w1_22, %gv[%o1_22] : vector<4xi32>, view<1024xi32> + %v1_23_0 = scalar.constant 1145390121 : i32 + %v1_23_1 = scalar.constant 1145717832 : i32 + %v1_23_2 = scalar.constant 1146373201 : i32 + %v1_23_3 = scalar.constant 1146504277 : i32 + %w1_23 = vector.from_elements %v1_23_0, %v1_23_1, %v1_23_2, %v1_23_3 : vector<4xi32> + %o1_23 = index.constant 348 : index + vector.store %w1_23, %gv[%o1_23] : vector<4xi32>, view<1024xi32> + %v1_24_0 = scalar.constant 1147290721 : i32 + %v1_24_1 = scalar.constant 1147683941 : i32 + %v1_24_2 = scalar.constant 1149322346 : i32 + %v1_24_3 = scalar.constant 1149846662 : i32 + %w1_24 = vector.from_elements %v1_24_0, %v1_24_1, %v1_24_2, %v1_24_3 : vector<4xi32> + %o1_24 = index.constant 352 : index + vector.store %w1_24, %gv[%o1_24] : vector<4xi32>, view<1024xi32> + %v1_25_0 = scalar.constant 1150436496 : i32 + %v1_25_1 = scalar.constant 1151354005 : i32 + %v1_25_2 = scalar.constant 1151943841 : i32 + %v1_25_3 = scalar.constant 1157776641 : i32 + %w1_25 = vector.from_elements %v1_25_0, %v1_25_1, %v1_25_2, %v1_25_3 : vector<4xi32> + %o1_25 = index.constant 356 : index + vector.store %w1_25, %gv[%o1_25] : vector<4xi32>, view<1024xi32> + %v1_26_0 = scalar.constant 1158300933 : i32 + %v1_26_1 = scalar.constant 1158956305 : i32 + %v1_26_2 = scalar.constant 1159087381 : i32 + %v1_26_3 = scalar.constant 1159742745 : i32 + %w1_26 = vector.from_elements %v1_26_0, %v1_26_1, %v1_26_2, %v1_26_3 : vector<4xi32> + %o1_26 = index.constant 360 : index + vector.store %w1_26, %gv[%o1_26] : vector<4xi32>, view<1024xi32> + %v1_27_0 = scalar.constant 1160398117 : i32 + %v1_27_1 = scalar.constant 1162102081 : i32 + %v1_27_2 = scalar.constant 1162233157 : i32 + %v1_27_3 = scalar.constant 1162888521 : i32 + %w1_27 = vector.from_elements %v1_27_0, %v1_27_1, %v1_27_2, %v1_27_3 : vector<4xi32> + %o1_27 = index.constant 364 : index + vector.store %w1_27, %gv[%o1_27] : vector<4xi32>, view<1024xi32> + %v1_28_0 = scalar.constant 1163150673 : i32 + %v1_28_1 = scalar.constant 1163281749 : i32 + %v1_28_2 = scalar.constant 1163478360 : i32 + %v1_28_3 = scalar.constant 1164199265 : i32 + %w1_28 = vector.from_elements %v1_28_0, %v1_28_1, %v1_28_2, %v1_28_3 : vector<4xi32> + %o1_28 = index.constant 368 : index + vector.store %w1_28, %gv[%o1_28] : vector<4xi32>, view<1024xi32> + %v1_29_0 = scalar.constant 1164330341 : i32 + %v1_29_1 = scalar.constant 1166165353 : i32 + %v1_29_2 = scalar.constant 1166361988 : i32 + %v1_29_3 = scalar.constant 1167148424 : i32 + %w1_29 = vector.from_elements %v1_29_0, %v1_29_1, %v1_29_2, %v1_29_3 : vector<4xi32> + %o1_29 = index.constant 372 : index + vector.store %w1_29, %gv[%o1_29] : vector<4xi32>, view<1024xi32> + %v1_30_0 = scalar.constant 1167410580 : i32 + %v1_30_1 = scalar.constant 1167672726 : i32 + %v1_30_2 = scalar.constant 1168459162 : i32 + %v1_30_3 = scalar.constant 1168786856 : i32 + %w1_30 = vector.from_elements %v1_30_0, %v1_30_1, %v1_30_2, %v1_30_3 : vector<4xi32> + %o1_30 = index.constant 376 : index + vector.store %w1_30, %gv[%o1_30] : vector<4xi32>, view<1024xi32> + %v1_31_0 = scalar.constant 1174750721 : i32 + %v1_31_1 = scalar.constant 1175733769 : i32 + %v1_31_2 = scalar.constant 1175995925 : i32 + %v1_31_3 = scalar.constant 1176585754 : i32 + %w1_31 = vector.from_elements %v1_31_0, %v1_31_1, %v1_31_2, %v1_31_3 : vector<4xi32> + %o1_31 = index.constant 380 : index + vector.store %w1_31, %gv[%o1_31] : vector<4xi32>, view<1024xi32> + %v1_32_0 = scalar.constant 1177110052 : i32 + %v1_32_1 = scalar.constant 1178748480 : i32 + %v1_32_2 = scalar.constant 1179141701 : i32 + %v1_32_3 = scalar.constant 1179731536 : i32 + %w1_32 = vector.from_elements %v1_32_0, %v1_32_1, %v1_32_2, %v1_32_3 : vector<4xi32> + %o1_32 = index.constant 384 : index + vector.store %w1_32, %gv[%o1_32] : vector<4xi32>, view<1024xi32> + %v1_33_0 = scalar.constant 1179993682 : i32 + %v1_33_1 = scalar.constant 1180255830 : i32 + %v1_33_2 = scalar.constant 1181042274 : i32 + %v1_33_3 = scalar.constant 1182877288 : i32 + %w1_33 = vector.from_elements %v1_33_0, %v1_33_1, %v1_33_2, %v1_33_3 : vector<4xi32> + %o1_33 = index.constant 388 : index + vector.store %w1_33, %gv[%o1_33] : vector<4xi32>, view<1024xi32> + %v1_34_0 = scalar.constant 1183467141 : i32 + %v1_34_1 = scalar.constant 1184188052 : i32 + %v1_34_2 = scalar.constant 1185171105 : i32 + %v1_34_3 = scalar.constant 1208305318 : i32 + %w1_34 = vector.from_elements %v1_34_0, %v1_34_1, %v1_34_2, %v1_34_3 : vector<4xi32> + %o1_34 = index.constant 392 : index + vector.store %w1_34, %gv[%o1_34] : vector<4xi32>, view<1024xi32> + %v1_35_0 = scalar.constant 1209354257 : i32 + %v1_35_1 = scalar.constant 1210402842 : i32 + %v1_35_2 = scalar.constant 1212762178 : i32 + %v1_35_3 = scalar.constant 1213548624 : i32 + %w1_35 = vector.from_elements %v1_35_0, %v1_35_1, %v1_35_2, %v1_35_3 : vector<4xi32> + %o1_35 = index.constant 396 : index + vector.store %w1_35, %gv[%o1_35] : vector<4xi32>, view<1024xi32> + %v1_36_0 = scalar.constant 1214335064 : i32 + %v1_36_1 = scalar.constant 1214662756 : i32 + %v1_36_2 = scalar.constant 1216694377 : i32 + %v1_36_3 = scalar.constant 1217677457 : i32 + %w1_36 = vector.from_elements %v1_36_0, %v1_36_1, %v1_36_2, %v1_36_3 : vector<4xi32> + %o1_36 = index.constant 400 : index + vector.store %w1_36, %gv[%o1_36] : vector<4xi32>, view<1024xi32> + %v1_37_0 = scalar.constant 1218005142 : i32 + %v1_37_1 = scalar.constant 1224820901 : i32 + %v1_37_2 = scalar.constant 1225148677 : i32 + %v1_37_3 = scalar.constant 1225804042 : i32 + %w1_37 = vector.from_elements %v1_37_0, %v1_37_1, %v1_37_2, %v1_37_3 : vector<4xi32> + %o1_37 = index.constant 404 : index + vector.store %w1_37, %gv[%o1_37] : vector<4xi32>, view<1024xi32> + %v1_38_0 = scalar.constant 1226131732 : i32 + %v1_38_1 = scalar.constant 1226918168 : i32 + %v1_38_2 = scalar.constant 1227245860 : i32 + %v1_38_3 = scalar.constant 1229277504 : i32 + %w1_38 = vector.from_elements %v1_38_0, %v1_38_1, %v1_38_2, %v1_38_3 : vector<4xi32> + %o1_38 = index.constant 408 : index + vector.store %w1_38, %gv[%o1_38] : vector<4xi32>, view<1024xi32> + %v1_39_0 = scalar.constant 1230063946 : i32 + %v1_39_1 = scalar.constant 1230260562 : i32 + %v1_39_2 = scalar.constant 1230391637 : i32 + %v1_39_3 = scalar.constant 1231047001 : i32 + %w1_39 = vector.from_elements %v1_39_0, %v1_39_1, %v1_39_2, %v1_39_3 : vector<4xi32> + %o1_39 = index.constant 412 : index + vector.store %w1_39, %gv[%o1_39] : vector<4xi32>, view<1024xi32> + %v1_40_0 = scalar.constant 1231374690 : i32 + %v1_40_1 = scalar.constant 1231702374 : i32 + %v1_40_2 = scalar.constant 1233734022 : i32 + %v1_40_3 = scalar.constant 1234520466 : i32 + %w1_40 = vector.from_elements %v1_40_0, %v1_40_1, %v1_40_2, %v1_40_3 : vector<4xi32> + %o1_40 = index.constant 416 : index + vector.store %w1_40, %gv[%o1_40] : vector<4xi32>, view<1024xi32> + %v1_41_0 = scalar.constant 1234717078 : i32 + %v1_41_1 = scalar.constant 1235503521 : i32 + %v1_41_2 = scalar.constant 1235831206 : i32 + %v1_41_3 = scalar.constant 1245989398 : i32 + %w1_41 = vector.from_elements %v1_41_0, %v1_41_1, %v1_41_2, %v1_41_3 : vector<4xi32> + %o1_41 = index.constant 420 : index + vector.store %w1_41, %gv[%o1_41] : vector<4xi32>, view<1024xi32> + %v1_42_0 = scalar.constant 1246317126 : i32 + %v1_42_1 = scalar.constant 1247300181 : i32 + %v1_42_2 = scalar.constant 1248086618 : i32 + %v1_42_3 = scalar.constant 1251232361 : i32 + %w1_42 = vector.from_elements %v1_42_0, %v1_42_1, %v1_42_2, %v1_42_3 : vector<4xi32> + %o1_42 = index.constant 424 : index + vector.store %w1_42, %gv[%o1_42] : vector<4xi32>, view<1024xi32> + %v1_43_0 = scalar.constant 1342261925 : i32 + %v1_43_1 = scalar.constant 1342525444 : i32 + %v1_43_2 = scalar.constant 1342787590 : i32 + %v1_43_3 = scalar.constant 1343574034 : i32 + %w1_43 = vector.from_elements %v1_43_0, %v1_43_1, %v1_43_2, %v1_43_3 : vector<4xi32> + %o1_43 = index.constant 428 : index + vector.store %w1_43, %gv[%o1_43] : vector<4xi32>, view<1024xi32> + %v1_44_0 = scalar.constant 1344360474 : i32 + %v1_44_1 = scalar.constant 1344884772 : i32 + %v1_44_2 = scalar.constant 1346719808 : i32 + %v1_44_3 = scalar.constant 1347506248 : i32 + %w1_44 = vector.from_elements %v1_44_0, %v1_44_1, %v1_44_2, %v1_44_3 : vector<4xi32> + %o1_44 = index.constant 432 : index + vector.store %w1_44, %gv[%o1_44] : vector<4xi32>, view<1024xi32> + %v1_45_0 = scalar.constant 1347768404 : i32 + %v1_45_1 = scalar.constant 1348030550 : i32 + %v1_45_2 = scalar.constant 1349013605 : i32 + %v1_45_3 = scalar.constant 1351176326 : i32 + %w1_45 = vector.from_elements %v1_45_0, %v1_45_1, %v1_45_2, %v1_45_3 : vector<4xi32> + %o1_45 = index.constant 436 : index + vector.store %w1_45, %gv[%o1_45] : vector<4xi32>, view<1024xi32> + %v1_46_0 = scalar.constant 1352159381 : i32 + %v1_46_1 = scalar.constant 1352749216 : i32 + %v1_46_2 = scalar.constant 1353273510 : i32 + %v1_46_3 = scalar.constant 1359499525 : i32 + %w1_46 = vector.from_elements %v1_46_0, %v1_46_1, %v1_46_2, %v1_46_3 : vector<4xi32> + %o1_46 = index.constant 440 : index + vector.store %w1_46, %gv[%o1_46] : vector<4xi32>, view<1024xi32> + %v1_47_0 = scalar.constant 1359630601 : i32 + %v1_47_1 = scalar.constant 1360285969 : i32 + %v1_47_2 = scalar.constant 1360417045 : i32 + %v1_47_3 = scalar.constant 1360613656 : i32 + %w1_47 = vector.from_elements %v1_47_0, %v1_47_1, %v1_47_2, %v1_47_3 : vector<4xi32> + %o1_47 = index.constant 444 : index + vector.store %w1_47, %gv[%o1_47] : vector<4xi32>, view<1024xi32> + %v1_48_0 = scalar.constant 1361400096 : i32 + %v1_48_1 = scalar.constant 1361596710 : i32 + %v1_48_2 = scalar.constant 1363235114 : i32 + %v1_48_3 = scalar.constant 1363497284 : i32 + %w1_48 = vector.from_elements %v1_48_0, %v1_48_1, %v1_48_2, %v1_48_3 : vector<4xi32> + %o1_48 = index.constant 448 : index + vector.store %w1_48, %gv[%o1_48] : vector<4xi32>, view<1024xi32> + %v1_49_0 = scalar.constant 1363759430 : i32 + %v1_49_1 = scalar.constant 1364283728 : i32 + %v1_49_2 = scalar.constant 1364480338 : i32 + %v1_49_3 = scalar.constant 1364611413 : i32 + %w1_49 = vector.from_elements %v1_49_0, %v1_49_1, %v1_49_2, %v1_49_3 : vector<4xi32> + %o1_49 = index.constant 452 : index + vector.store %w1_49, %gv[%o1_49] : vector<4xi32>, view<1024xi32> + %v1_50_0 = scalar.constant 1364808024 : i32 + %v1_50_1 = scalar.constant 1365332314 : i32 + %v1_50_2 = scalar.constant 1365594468 : i32 + %v1_50_3 = scalar.constant 1365856614 : i32 + %w1_50 = vector.from_elements %v1_50_0, %v1_50_1, %v1_50_2, %v1_50_3 : vector<4xi32> + %o1_50 = index.constant 456 : index + vector.store %w1_50, %gv[%o1_50] : vector<4xi32>, view<1024xi32> + %v1_51_0 = scalar.constant 1367691650 : i32 + %v1_51_1 = scalar.constant 1368674705 : i32 + %v1_51_2 = scalar.constant 1368805781 : i32 + %v1_51_3 = scalar.constant 1369461145 : i32 + %w1_51 = vector.from_elements %v1_51_0, %v1_51_1, %v1_51_2, %v1_51_3 : vector<4xi32> + %o1_51 = index.constant 460 : index + vector.store %w1_51, %gv[%o1_51] : vector<4xi32>, view<1024xi32> + %v1_52_0 = scalar.constant 1370116517 : i32 + %v1_52_1 = scalar.constant 1376145921 : i32 + %v1_52_2 = scalar.constant 1377128978 : i32 + %v1_52_3 = scalar.constant 1377915418 : i32 + %w1_52 = vector.from_elements %v1_52_0, %v1_52_1, %v1_52_2, %v1_52_3 : vector<4xi32> + %o1_52 = index.constant 464 : index + vector.store %w1_52, %gv[%o1_52] : vector<4xi32>, view<1024xi32> + %v1_53_0 = scalar.constant 1380078116 : i32 + %v1_53_1 = scalar.constant 1380602437 : i32 + %v1_53_2 = scalar.constant 1381257809 : i32 + %v1_53_3 = scalar.constant 1381388885 : i32 + %w1_53 = vector.from_elements %v1_53_0, %v1_53_1, %v1_53_2, %v1_53_3 : vector<4xi32> + %o1_53 = index.constant 468 : index + vector.store %w1_53, %gv[%o1_53] : vector<4xi32>, view<1024xi32> + %v1_54_0 = scalar.constant 1382175321 : i32 + %v1_54_1 = scalar.constant 1384469093 : i32 + %v1_54_2 = scalar.constant 1385321104 : i32 + %v1_54_3 = scalar.constant 1385779861 : i32 + %w1_54 = vector.from_elements %v1_54_0, %v1_54_1, %v1_54_2, %v1_54_3 : vector<4xi32> + %o1_54 = index.constant 472 : index + vector.store %w1_54, %gv[%o1_54] : vector<4xi32>, view<1024xi32> + %v1_55_0 = scalar.constant 1386500762 : i32 + %v1_55_1 = scalar.constant 1409635332 : i32 + %v1_55_2 = scalar.constant 1410618385 : i32 + %v1_55_3 = scalar.constant 1410749461 : i32 + %w1_55 = vector.from_elements %v1_55_0, %v1_55_1, %v1_55_2, %v1_55_3 : vector<4xi32> + %o1_55 = index.constant 476 : index + vector.store %w1_55, %gv[%o1_55] : vector<4xi32>, view<1024xi32> + %v1_56_0 = scalar.constant 1410946072 : i32 + %v1_56_1 = scalar.constant 1411732513 : i32 + %v1_56_2 = scalar.constant 1412060200 : i32 + %v1_56_3 = scalar.constant 1413764161 : i32 + %w1_56 = vector.from_elements %v1_56_0, %v1_56_1, %v1_56_2, %v1_56_3 : vector<4xi32> + %o1_56 = index.constant 480 : index + vector.store %w1_56, %gv[%o1_56] : vector<4xi32>, view<1024xi32> + %v1_57_0 = scalar.constant 1413895237 : i32 + %v1_57_1 = scalar.constant 1414157385 : i32 + %v1_57_2 = scalar.constant 1414616144 : i32 + %v1_57_3 = scalar.constant 1414878292 : i32 + %w1_57 = vector.from_elements %v1_57_0, %v1_57_1, %v1_57_2, %v1_57_3 : vector<4xi32> + %o1_57 = index.constant 484 : index + vector.store %w1_57, %gv[%o1_57] : vector<4xi32>, view<1024xi32> + %v1_58_0 = scalar.constant 1415074902 : i32 + %v1_58_1 = scalar.constant 1415205977 : i32 + %v1_58_2 = scalar.constant 1415730273 : i32 + %v1_58_3 = scalar.constant 1415926884 : i32 + %w1_58 = vector.from_elements %v1_58_0, %v1_58_1, %v1_58_2, %v1_58_3 : vector<4xi32> + %o1_58 = index.constant 488 : index + vector.store %w1_58, %gv[%o1_58] : vector<4xi32>, view<1024xi32> + %v1_59_0 = scalar.constant 1416189030 : i32 + %v1_59_1 = scalar.constant 1418220672 : i32 + %v1_59_2 = scalar.constant 1418810506 : i32 + %v1_59_3 = scalar.constant 1419072660 : i32 + %w1_59 = vector.from_elements %v1_59_0, %v1_59_1, %v1_59_2, %v1_59_3 : vector<4xi32> + %o1_59 = index.constant 492 : index + vector.store %w1_59, %gv[%o1_59] : vector<4xi32>, view<1024xi32> + %v1_60_0 = scalar.constant 1419334806 : i32 + %v1_60_1 = scalar.constant 1420055713 : i32 + %v1_60_2 = scalar.constant 1420448933 : i32 + %v1_60_3 = scalar.constant 1426216193 : i32 + %w1_60 = vector.from_elements %v1_60_0, %v1_60_1, %v1_60_2, %v1_60_3 : vector<4xi32> + %o1_60 = index.constant 496 : index + vector.store %w1_60, %gv[%o1_60] : vector<4xi32>, view<1024xi32> + %v1_61_0 = scalar.constant 1426412804 : i32 + %v1_61_1 = scalar.constant 1426674950 : i32 + %v1_61_2 = scalar.constant 1427199248 : i32 + %v1_61_3 = scalar.constant 1427395858 : i32 + %w1_61 = vector.from_elements %v1_61_0, %v1_61_1, %v1_61_2, %v1_61_3 : vector<4xi32> + %o1_61 = index.constant 500 : index + vector.store %w1_61, %gv[%o1_61] : vector<4xi32>, view<1024xi32> + %v1_62_0 = scalar.constant 1427526933 : i32 + %v1_62_1 = scalar.constant 1427789081 : i32 + %v1_62_2 = scalar.constant 1428444449 : i32 + %v1_62_3 = scalar.constant 1428575525 : i32 + %w1_62 = vector.from_elements %v1_62_0, %v1_62_1, %v1_62_2, %v1_62_3 : vector<4xi32> + %o1_62 = index.constant 504 : index + vector.store %w1_62, %gv[%o1_62] : vector<4xi32>, view<1024xi32> + %v1_63_0 = scalar.constant 1430279465 : i32 + %v1_63_1 = scalar.constant 1430410561 : i32 + %v1_63_2 = scalar.constant 1430607172 : i32 + %v1_63_3 = scalar.constant 1430803782 : i32 + %w1_63 = vector.from_elements %v1_63_0, %v1_63_1, %v1_63_2, %v1_63_3 : vector<4xi32> + %o1_63 = index.constant 508 : index + vector.store %w1_63, %gv[%o1_63] : vector<4xi32>, view<1024xi32> + } + %k2 = index.constant 2 : index + %is2 = index.cmp eq, %chunk, %k2 : index + scf.if %is2 { + %v2_0_0 = scalar.constant 1431328073 : i32 + %v2_0_1 = scalar.constant 1431459153 : i32 + %v2_0_2 = scalar.constant 1431655764 : i32 + %v2_0_3 = scalar.constant 1431852374 : i32 + %w2_0 = vector.from_elements %v2_0_0, %v2_0_1, %v2_0_2, %v2_0_3 : vector<4xi32> + %o2_0 = index.constant 512 : index + vector.store %w2_0, %gv[%o2_0] : vector<4xi32>, view<1024xi32> + %v2_1_0 = scalar.constant 1431983449 : i32 + %v2_1_1 = scalar.constant 1432442208 : i32 + %v2_1_2 = scalar.constant 1432704356 : i32 + %v2_1_3 = scalar.constant 1432900966 : i32 + %w2_1 = vector.from_elements %v2_1_0, %v2_1_1, %v2_1_2, %v2_1_3 : vector<4xi32> + %o2_1 = index.constant 516 : index + vector.store %w2_1, %gv[%o2_1] : vector<4xi32>, view<1024xi32> + %v2_2_0 = scalar.constant 1433032041 : i32 + %v2_2_1 = scalar.constant 1434736001 : i32 + %v2_2_2 = scalar.constant 1435063685 : i32 + %v2_2_3 = scalar.constant 1435522442 : i32 + %w2_2 = vector.from_elements %v2_2_0, %v2_2_1, %v2_2_2, %v2_2_3 : vector<4xi32> + %o2_2 = index.constant 520 : index + vector.store %w2_2, %gv[%o2_2] : vector<4xi32>, view<1024xi32> + %v2_3_0 = scalar.constant 1435784593 : i32 + %v2_3_1 = scalar.constant 1435915669 : i32 + %v2_3_2 = scalar.constant 1436112280 : i32 + %v2_3_3 = scalar.constant 1436833185 : i32 + %w2_3 = vector.from_elements %v2_3_0, %v2_3_1, %v2_3_2, %v2_3_3 : vector<4xi32> + %o2_3 = index.constant 524 : index + vector.store %w2_3, %gv[%o2_3] : vector<4xi32>, view<1024xi32> + %v2_4_0 = scalar.constant 1436964261 : i32 + %v2_4_1 = scalar.constant 1442862505 : i32 + %v2_4_2 = scalar.constant 1442993665 : i32 + %v2_4_3 = scalar.constant 1443255812 : i32 + %w2_4 = vector.from_elements %v2_4_0, %v2_4_1, %v2_4_2, %v2_4_3 : vector<4xi32> + %o2_4 = index.constant 528 : index + vector.store %w2_4, %gv[%o2_4] : vector<4xi32>, view<1024xi32> + %v2_5_0 = scalar.constant 1443452424 : i32 + %v2_5_1 = scalar.constant 1444173329 : i32 + %v2_5_2 = scalar.constant 1444435477 : i32 + %v2_5_3 = scalar.constant 1444959769 : i32 + %w2_5 = vector.from_elements %v2_5_0, %v2_5_1, %v2_5_2, %v2_5_3 : vector<4xi32> + %o2_5 = index.constant 532 : index + vector.store %w2_5, %gv[%o2_5] : vector<4xi32>, view<1024xi32> + %v2_6_0 = scalar.constant 1445090849 : i32 + %v2_6_1 = scalar.constant 1445287460 : i32 + %v2_6_2 = scalar.constant 1445484070 : i32 + %v2_6_3 = scalar.constant 1447122473 : i32 + %w2_6 = vector.from_elements %v2_6_0, %v2_6_1, %v2_6_2, %v2_6_3 : vector<4xi32> + %o2_6 = index.constant 536 : index + vector.store %w2_6, %gv[%o2_6] : vector<4xi32>, view<1024xi32> + %v2_7_0 = scalar.constant 1447450181 : i32 + %v2_7_1 = scalar.constant 1447646792 : i32 + %v2_7_2 = scalar.constant 1448105546 : i32 + %v2_7_3 = scalar.constant 1448236625 : i32 + %w2_7 = vector.from_elements %v2_7_0, %v2_7_1, %v2_7_2, %v2_7_3 : vector<4xi32> + %o2_7 = index.constant 540 : index + vector.store %w2_7, %gv[%o2_7] : vector<4xi32>, view<1024xi32> + %v2_8_0 = scalar.constant 1448433236 : i32 + %v2_8_1 = scalar.constant 1448629846 : i32 + %v2_8_2 = scalar.constant 1448760921 : i32 + %v2_8_3 = scalar.constant 1449416289 : i32 + %w2_8 = vector.from_elements %v2_8_0, %v2_8_1, %v2_8_2, %v2_8_3 : vector<4xi32> + %o2_8 = index.constant 544 : index + vector.store %w2_8, %gv[%o2_8] : vector<4xi32>, view<1024xi32> + %v2_9_0 = scalar.constant 1449743973 : i32 + %v2_9_1 = scalar.constant 1451579010 : i32 + %v2_9_2 = scalar.constant 1451775622 : i32 + %v2_9_3 = scalar.constant 1451906697 : i32 + %w2_9 = vector.from_elements %v2_9_0, %v2_9_1, %v2_9_2, %v2_9_3 : vector<4xi32> + %o2_9 = index.constant 548 : index + vector.store %w2_9, %gv[%o2_9] : vector<4xi32>, view<1024xi32> + %v2_10_0 = scalar.constant 1452627601 : i32 + %v2_10_1 = scalar.constant 1453479578 : i32 + %v2_10_2 = scalar.constant 1453741733 : i32 + %v2_10_3 = scalar.constant 1453938344 : i32 + %w2_10 = vector.from_elements %v2_10_0, %v2_10_1, %v2_10_2, %v2_10_3 : vector<4xi32> + %o2_10 = index.constant 552 : index + vector.store %w2_10, %gv[%o2_10] : vector<4xi32>, view<1024xi32> + %v2_11_0 = scalar.constant 1476745220 : i32 + %v2_11_1 = scalar.constant 1477007366 : i32 + %v2_11_2 = scalar.constant 1477793808 : i32 + %v2_11_3 = scalar.constant 1478580248 : i32 + %w2_11 = vector.from_elements %v2_11_0, %v2_11_1, %v2_11_2, %v2_11_3 : vector<4xi32> + %o2_11 = index.constant 556 : index + vector.store %w2_11, %gv[%o2_11] : vector<4xi32>, view<1024xi32> + %v2_12_0 = scalar.constant 1480939562 : i32 + %v2_12_1 = scalar.constant 1481267272 : i32 + %v2_12_2 = scalar.constant 1481922641 : i32 + %v2_12_3 = scalar.constant 1482053717 : i32 + %w2_12 = vector.from_elements %v2_12_0, %v2_12_1, %v2_12_2, %v2_12_3 : vector<4xi32> + %o2_12 = index.constant 560 : index + vector.store %w2_12, %gv[%o2_12] : vector<4xi32>, view<1024xi32> + %v2_13_0 = scalar.constant 1482250328 : i32 + %v2_13_1 = scalar.constant 1482840160 : i32 + %v2_13_2 = scalar.constant 1483036772 : i32 + %v2_13_3 = scalar.constant 1485396098 : i32 + %w2_13 = vector.from_elements %v2_13_0, %v2_13_1, %v2_13_2, %v2_13_3 : vector<4xi32> + %o2_13 = index.constant 564 : index + vector.store %w2_13, %gv[%o2_13] : vector<4xi32>, view<1024xi32> + %v2_14_0 = scalar.constant 1485985936 : i32 + %v2_14_1 = scalar.constant 1486379157 : i32 + %v2_14_2 = scalar.constant 1487493281 : i32 + %v2_14_3 = scalar.constant 1493326081 : i32 + %w2_14 = vector.from_elements %v2_14_0, %v2_14_1, %v2_14_2, %v2_14_3 : vector<4xi32> + %o2_14 = index.constant 568 : index + vector.store %w2_14, %gv[%o2_14] : vector<4xi32>, view<1024xi32> + %v2_15_0 = scalar.constant 1493850373 : i32 + %v2_15_1 = scalar.constant 1494505745 : i32 + %v2_15_2 = scalar.constant 1494636821 : i32 + %v2_15_3 = scalar.constant 1495619865 : i32 + %w2_15 = vector.from_elements %v2_15_0, %v2_15_1, %v2_15_2, %v2_15_3 : vector<4xi32> + %o2_15 = index.constant 572 : index + vector.store %w2_15, %gv[%o2_15] : vector<4xi32>, view<1024xi32> + %v2_16_0 = scalar.constant 1497651521 : i32 + %v2_16_1 = scalar.constant 1497782597 : i32 + %v2_16_2 = scalar.constant 1498437961 : i32 + %v2_16_3 = scalar.constant 1498569041 : i32 + %w2_16 = vector.from_elements %v2_16_0, %v2_16_1, %v2_16_2, %v2_16_3 : vector<4xi32> + %o2_16 = index.constant 576 : index + vector.store %w2_16, %gv[%o2_16] : vector<4xi32>, view<1024xi32> + %v2_17_0 = scalar.constant 1498765652 : i32 + %v2_17_1 = scalar.constant 1498962262 : i32 + %v2_17_2 = scalar.constant 1499093337 : i32 + %v2_17_3 = scalar.constant 1499748705 : i32 + %w2_17 = vector.from_elements %v2_17_0, %v2_17_1, %v2_17_2, %v2_17_3 : vector<4xi32> + %o2_17 = index.constant 580 : index + vector.store %w2_17, %gv[%o2_17] : vector<4xi32>, view<1024xi32> + %v2_18_0 = scalar.constant 1499879781 : i32 + %v2_18_1 = scalar.constant 1501649257 : i32 + %v2_18_2 = scalar.constant 1502173573 : i32 + %v2_18_3 = scalar.constant 1502894481 : i32 + %w2_18 = vector.from_elements %v2_18_0, %v2_18_1, %v2_18_2, %v2_18_3 : vector<4xi32> + %o2_18 = index.constant 584 : index + vector.store %w2_18, %gv[%o2_18] : vector<4xi32>, view<1024xi32> + %v2_19_0 = scalar.constant 1503025557 : i32 + %v2_19_1 = scalar.constant 1503222168 : i32 + %v2_19_2 = scalar.constant 1510234533 : i32 + %v2_19_3 = scalar.constant 1511348744 : i32 + %w2_19 = vector.from_elements %v2_19_0, %v2_19_1, %v2_19_2, %v2_19_3 : vector<4xi32> + %o2_19 = index.constant 588 : index + vector.store %w2_19, %gv[%o2_19] : vector<4xi32>, view<1024xi32> + %v2_20_0 = scalar.constant 1512069658 : i32 + %v2_20_1 = scalar.constant 1512462885 : i32 + %v2_20_2 = scalar.constant 1514494505 : i32 + %v2_20_3 = scalar.constant 1514756680 : i32 + %w2_20 = vector.from_elements %v2_20_0, %v2_20_1, %v2_20_2, %v2_20_3 : vector<4xi32> + %o2_20 = index.constant 592 : index + vector.store %w2_20, %gv[%o2_20] : vector<4xi32>, view<1024xi32> + %v2_21_0 = scalar.constant 1515543121 : i32 + %v2_21_1 = scalar.constant 1515739734 : i32 + %v2_21_2 = scalar.constant 1516395097 : i32 + %v2_21_3 = scalar.constant 1516788325 : i32 + %w2_21 = vector.from_elements %v2_21_0, %v2_21_1, %v2_21_2, %v2_21_3 : vector<4xi32> + %o2_21 = index.constant 596 : index + vector.store %w2_21, %gv[%o2_21] : vector<4xi32>, view<1024xi32> + %v2_22_0 = scalar.constant 1518426730 : i32 + %v2_22_1 = scalar.constant 1519540874 : i32 + %v2_22_2 = scalar.constant 1519803029 : i32 + %v2_22_3 = scalar.constant 1520065176 : i32 + %w2_22 = vector.from_elements %v2_22_0, %v2_22_1, %v2_22_2, %v2_22_3 : vector<4xi32> + %o2_22 = index.constant 600 : index + vector.store %w2_22, %gv[%o2_22] : vector<4xi32>, view<1024xi32> + %v2_23_0 = scalar.constant 1610963617 : i32 + %v2_23_1 = scalar.constant 1612079124 : i32 + %v2_23_2 = scalar.constant 1613062169 : i32 + %v2_23_3 = scalar.constant 1615880260 : i32 + %w2_23 = vector.from_elements %v2_23_0, %v2_23_1, %v2_23_2, %v2_23_3 : vector<4xi32> + %o2_23 = index.constant 604 : index + vector.store %w2_23, %gv[%o2_23] : vector<4xi32>, view<1024xi32> + %v2_24_0 = scalar.constant 1616273493 : i32 + %v2_24_1 = scalar.constant 1616535640 : i32 + %v2_24_2 = scalar.constant 1617191009 : i32 + %v2_24_3 = scalar.constant 1617518694 : i32 + %w2_24 = vector.from_elements %v2_24_0, %v2_24_1, %v2_24_2, %v2_24_3 : vector<4xi32> + %o2_24 = index.constant 608 : index + vector.store %w2_24, %gv[%o2_24] : vector<4xi32>, view<1024xi32> + %v2_25_0 = scalar.constant 1620467841 : i32 + %v2_25_1 = scalar.constant 1627480229 : i32 + %v2_25_2 = scalar.constant 1627808004 : i32 + %v2_25_3 = scalar.constant 1628594441 : i32 + %w2_25 = vector.from_elements %v2_25_0, %v2_25_1, %v2_25_2, %v2_25_3 : vector<4xi32> + %o2_25 = index.constant 612 : index + vector.store %w2_25, %gv[%o2_25] : vector<4xi32>, view<1024xi32> + %v2_26_0 = scalar.constant 1629577493 : i32 + %v2_26_1 = scalar.constant 1629905186 : i32 + %v2_26_2 = scalar.constant 1631936809 : i32 + %v2_26_3 = scalar.constant 1632723273 : i32 + %w2_26 = vector.from_elements %v2_26_0, %v2_26_1, %v2_26_2, %v2_26_3 : vector<4xi32> + %o2_26 = index.constant 616 : index + vector.store %w2_26, %gv[%o2_26] : vector<4xi32>, view<1024xi32> + %v2_27_0 = scalar.constant 1633050965 : i32 + %v2_27_1 = scalar.constant 1634034009 : i32 + %v2_27_2 = scalar.constant 1634361702 : i32 + %v2_27_3 = scalar.constant 1636458884 : i32 + %w2_27 = vector.from_elements %v2_27_0, %v2_27_1, %v2_27_2, %v2_27_3 : vector<4xi32> + %o2_27 = index.constant 620 : index + vector.store %w2_27, %gv[%o2_27] : vector<4xi32>, view<1024xi32> + %v2_28_0 = scalar.constant 1637179794 : i32 + %v2_28_1 = scalar.constant 1638293921 : i32 + %v2_28_2 = scalar.constant 1645306281 : i32 + %v2_28_3 = scalar.constant 1645830678 : i32 + %w2_28 = vector.from_elements %v2_28_0, %v2_28_1, %v2_28_2, %v2_28_3 : vector<4xi32> + %o2_28 = index.constant 624 : index + vector.store %w2_28, %gv[%o2_28] : vector<4xi32>, view<1024xi32> + %v2_29_0 = scalar.constant 1648452160 : i32 + %v2_29_1 = scalar.constant 1649762886 : i32 + %v2_29_2 = scalar.constant 1649959510 : i32 + %v2_29_3 = scalar.constant 1652908640 : i32 + %w2_29 = vector.from_elements %v2_29_0, %v2_29_1, %v2_29_2, %v2_29_3 : vector<4xi32> + %o2_29 = index.constant 628 : index + vector.store %w2_29, %gv[%o2_29] : vector<4xi32>, view<1024xi32> + %v2_30_0 = scalar.constant 1654022801 : i32 + %v2_30_1 = scalar.constant 1678860965 : i32 + %v2_30_2 = scalar.constant 1679123474 : i32 + %v2_30_3 = scalar.constant 1679451158 : i32 + %w2_30 = vector.from_elements %v2_30_0, %v2_30_1, %v2_30_2, %v2_30_3 : vector<4xi32> + %o2_30 = index.constant 632 : index + vector.store %w2_30, %gv[%o2_30] : vector<4xi32>, view<1024xi32> + %v2_31_0 = scalar.constant 1680237601 : i32 + %v2_31_1 = scalar.constant 1681941545 : i32 + %v2_31_2 = scalar.constant 1682269250 : i32 + %v2_31_3 = scalar.constant 1682596936 : i32 + %w2_31 = vector.from_elements %v2_31_0, %v2_31_1, %v2_31_2, %v2_31_3 : vector<4xi32> + %o2_31 = index.constant 636 : index + vector.store %w2_31, %gv[%o2_31] : vector<4xi32>, view<1024xi32> + %v2_32_0 = scalar.constant 1683252305 : i32 + %v2_32_1 = scalar.constant 1683383381 : i32 + %v2_32_2 = scalar.constant 1683645529 : i32 + %v2_32_3 = scalar.constant 1684169824 : i32 + %w2_32 = vector.from_elements %v2_32_0, %v2_32_1, %v2_32_2, %v2_32_3 : vector<4xi32> + %o2_32 = index.constant 640 : index + vector.store %w2_32, %gv[%o2_32] : vector<4xi32>, view<1024xi32> + %v2_33_0 = scalar.constant 1686398053 : i32 + %v2_33_1 = scalar.constant 1686725765 : i32 + %v2_33_2 = scalar.constant 1687315600 : i32 + %v2_33_3 = scalar.constant 1687512212 : i32 + %w2_33 = vector.from_elements %v2_33_0, %v2_33_1, %v2_33_2, %v2_33_3 : vector<4xi32> + %o2_33 = index.constant 644 : index + vector.store %w2_33, %gv[%o2_33] : vector<4xi32>, view<1024xi32> + %v2_34_0 = scalar.constant 1687708822 : i32 + %v2_34_1 = scalar.constant 1688298650 : i32 + %v2_34_2 = scalar.constant 1688822948 : i32 + %v2_34_3 = scalar.constant 1695048965 : i32 + %w2_34 = vector.from_elements %v2_34_0, %v2_34_1, %v2_34_2, %v2_34_3 : vector<4xi32> + %o2_34 = index.constant 648 : index + vector.store %w2_34, %gv[%o2_34] : vector<4xi32>, view<1024xi32> + %v2_35_0 = scalar.constant 1695638794 : i32 + %v2_35_1 = scalar.constant 1695966485 : i32 + %v2_35_2 = scalar.constant 1698981145 : i32 + %v2_35_3 = scalar.constant 1699112261 : i32 + %w2_35 = vector.from_elements %v2_35_0, %v2_35_1, %v2_35_2, %v2_35_3 : vector<4xi32> + %o2_35 = index.constant 652 : index + vector.store %w2_35, %gv[%o2_35] : vector<4xi32>, view<1024xi32> + %v2_36_0 = scalar.constant 1699767625 : i32 + %v2_36_1 = scalar.constant 1700029777 : i32 + %v2_36_2 = scalar.constant 1700160853 : i32 + %v2_36_3 = scalar.constant 1700881753 : i32 + %w2_36 = vector.from_elements %v2_36_0, %v2_36_1, %v2_36_2, %v2_36_3 : vector<4xi32> + %o2_36 = index.constant 656 : index + vector.store %w2_36, %gv[%o2_36] : vector<4xi32>, view<1024xi32> + %v2_37_0 = scalar.constant 1701143908 : i32 + %v2_37_1 = scalar.constant 1701406054 : i32 + %v2_37_2 = scalar.constant 1703503238 : i32 + %v2_37_3 = scalar.constant 1704027530 : i32 + %w2_37 = vector.from_elements %v2_37_0, %v2_37_1, %v2_37_2, %v2_37_3 : vector<4xi32> + %o2_37 = index.constant 660 : index + vector.store %w2_37, %gv[%o2_37] : vector<4xi32>, view<1024xi32> + %v2_38_0 = scalar.constant 1704355221 : i32 + %v2_38_1 = scalar.constant 1704617369 : i32 + %v2_38_2 = scalar.constant 1705338274 : i32 + %v2_38_3 = scalar.constant 1705534886 : i32 + %w2_38 = vector.from_elements %v2_38_0, %v2_38_1, %v2_38_2, %v2_38_3 : vector<4xi32> + %o2_38 = index.constant 664 : index + vector.store %w2_38, %gv[%o2_38] : vector<4xi32>, view<1024xi32> + %v2_39_0 = scalar.constant 1711891970 : i32 + %v2_39_1 = scalar.constant 1713399317 : i32 + %v2_39_2 = scalar.constant 1713923622 : i32 + %v2_39_3 = scalar.constant 1715496489 : i32 + %w2_39 = vector.from_elements %v2_39_0, %v2_39_1, %v2_39_2, %v2_39_3 : vector<4xi32> + %o2_39 = index.constant 668 : index + vector.store %w2_39, %gv[%o2_39] : vector<4xi32>, view<1024xi32> + %v2_40_0 = scalar.constant 1716020805 : i32 + %v2_40_1 = scalar.constant 1716610634 : i32 + %v2_40_2 = scalar.constant 1716872788 : i32 + %v2_40_3 = scalar.constant 1717069398 : i32 + %w2_40 = vector.from_elements %v2_40_0, %v2_40_1, %v2_40_2, %v2_40_3 : vector<4xi32> + %o2_40 = index.constant 672 : index + vector.store %w2_40, %gv[%o2_40] : vector<4xi32>, view<1024xi32> + %v2_41_0 = scalar.constant 1717593690 : i32 + %v2_41_1 = scalar.constant 1718117989 : i32 + %v2_41_2 = scalar.constant 1719821952 : i32 + %v2_41_3 = scalar.constant 1720346245 : i32 + %w2_41 = vector.from_elements %v2_41_0, %v2_41_1, %v2_41_2, %v2_41_3 : vector<4xi32> + %o2_41 = index.constant 676 : index + vector.store %w2_41, %gv[%o2_41] : vector<4xi32>, view<1024xi32> + %v2_42_0 = scalar.constant 1721132692 : i32 + %v2_42_1 = scalar.constant 1721329304 : i32 + %v2_42_2 = scalar.constant 1722050208 : i32 + %v2_42_3 = scalar.constant 1722443430 : i32 + %w2_42 = vector.from_elements %v2_42_0, %v2_42_1, %v2_42_2, %v2_42_3 : vector<4xi32> + %o2_42 = index.constant 680 : index + vector.store %w2_42, %gv[%o2_42] : vector<4xi32>, view<1024xi32> + %v2_43_0 = scalar.constant 1746495510 : i32 + %v2_43_1 = scalar.constant 1749116965 : i32 + %v2_43_2 = scalar.constant 1750427730 : i32 + %v2_43_3 = scalar.constant 1751214170 : i32 + %w2_43 = vector.from_elements %v2_43_0, %v2_43_1, %v2_43_2, %v2_43_3 : vector<4xi32> + %o2_43 = index.constant 684 : index + vector.store %w2_43, %gv[%o2_43] : vector<4xi32>, view<1024xi32> + %v2_44_0 = scalar.constant 1753573481 : i32 + %v2_44_1 = scalar.constant 1754818705 : i32 + %v2_44_2 = scalar.constant 1761700006 : i32 + %v2_44_3 = scalar.constant 1762683140 : i32 + %w2_44 = vector.from_elements %v2_44_0, %v2_44_1, %v2_44_2, %v2_44_3 : vector<4xi32> + %o2_44 = index.constant 688 : index + vector.store %w2_44, %gv[%o2_44] : vector<4xi32>, view<1024xi32> + %v2_45_0 = scalar.constant 1763797269 : i32 + %v2_45_1 = scalar.constant 1764124964 : i32 + %v2_45_2 = scalar.constant 1765828905 : i32 + %v2_45_3 = scalar.constant 1766156609 : i32 + %w2_45 = vector.from_elements %v2_45_0, %v2_45_1, %v2_45_2, %v2_45_3 : vector<4xi32> + %o2_45 = index.constant 692 : index + vector.store %w2_45, %gv[%o2_45] : vector<4xi32>, view<1024xi32> + %v2_46_0 = scalar.constant 1766353222 : i32 + %v2_46_1 = scalar.constant 1767139665 : i32 + %v2_46_2 = scalar.constant 1767270741 : i32 + %v2_46_3 = scalar.constant 1767926105 : i32 + %w2_46 = vector.from_elements %v2_46_0, %v2_46_1, %v2_46_2, %v2_46_3 : vector<4xi32> + %o2_46 = index.constant 696 : index + vector.store %w2_46, %gv[%o2_46] : vector<4xi32>, view<1024xi32> + %v2_47_0 = scalar.constant 1768581477 : i32 + %v2_47_1 = scalar.constant 1770285442 : i32 + %v2_47_2 = scalar.constant 1771399562 : i32 + %v2_47_3 = scalar.constant 1772382625 : i32 + %w2_47 = vector.from_elements %v2_47_0, %v2_47_1, %v2_47_2, %v2_47_3 : vector<4xi32> + %o2_47 = index.constant 700 : index + vector.store %w2_47, %gv[%o2_47] : vector<4xi32>, view<1024xi32> + %v2_48_0 = scalar.constant 1772710309 : i32 + %v2_48_1 = scalar.constant 1779853841 : i32 + %v2_48_2 = scalar.constant 1782671896 : i32 + %v2_48_3 = scalar.constant 1783196228 : i32 + %w2_48 = vector.from_elements %v2_48_0, %v2_48_1, %v2_48_2, %v2_48_3 : vector<4xi32> + %o2_48 = index.constant 704 : index + vector.store %w2_48, %gv[%o2_48] : vector<4xi32>, view<1024xi32> + %v2_49_0 = scalar.constant 1783982672 : i32 + %v2_49_1 = scalar.constant 1784310360 : i32 + %v2_49_2 = scalar.constant 1785031268 : i32 + %v2_49_3 = scalar.constant 1787193961 : i32 + %w2_49 = vector.from_elements %v2_49_0, %v2_49_1, %v2_49_2, %v2_49_3 : vector<4xi32> + %o2_49 = index.constant 708 : index + vector.store %w2_49, %gv[%o2_49] : vector<4xi32>, view<1024xi32> + %v2_50_0 = scalar.constant 1788373652 : i32 + %v2_50_1 = scalar.constant 1789291162 : i32 + %v2_50_2 = scalar.constant -2147319808 : i32 + %v2_50_3 = scalar.constant -2146795512 : i32 + %w2_50 = vector.from_elements %v2_50_0, %v2_50_1, %v2_50_2, %v2_50_3 : vector<4xi32> + %o2_50 = index.constant 712 : index + vector.store %w2_50, %gv[%o2_50] : vector<4xi32>, view<1024xi32> + %v2_51_0 = scalar.constant -2145222624 : i32 + %v2_51_1 = scalar.constant -2144698328 : i32 + %v2_51_2 = scalar.constant -2142207931 : i32 + %v2_51_3 = scalar.constant -2141945775 : i32 + %w2_51 = vector.from_elements %v2_51_0, %v2_51_1, %v2_51_2, %v2_51_3 : vector<4xi32> + %o2_51 = index.constant 716 : index + vector.store %w2_51, %gv[%o2_51] : vector<4xi32>, view<1024xi32> + %v2_52_0 = scalar.constant -2141618090 : i32 + %v2_52_1 = scalar.constant -2139062171 : i32 + %v2_52_2 = scalar.constant -2138537854 : i32 + %v2_52_3 = scalar.constant -2137685878 : i32 + %w2_52 = vector.from_elements %v2_52_0, %v2_52_1, %v2_52_2, %v2_52_3 : vector<4xi32> + %o2_52 = index.constant 720 : index + vector.store %w2_52, %gv[%o2_52] : vector<4xi32>, view<1024xi32> + %v2_53_0 = scalar.constant -2136833888 : i32 + %v2_53_1 = scalar.constant -2136309592 : i32 + %v2_53_2 = scalar.constant -2129559291 : i32 + %v2_53_3 = scalar.constant -2129231596 : i32 + %w2_53 = vector.from_elements %v2_53_0, %v2_53_1, %v2_53_2, %v2_53_3 : vector<4xi32> + %o2_53 = index.constant 724 : index + vector.store %w2_53, %gv[%o2_53] : vector<4xi32>, view<1024xi32> + %v2_54_0 = scalar.constant -2128248551 : i32 + %v2_54_1 = scalar.constant -2126216895 : i32 + %v2_54_2 = scalar.constant -2125430455 : i32 + %v2_54_3 = scalar.constant -2125102766 : i32 + %w2_54 = vector.from_elements %v2_54_0, %v2_54_1, %v2_54_2, %v2_54_3 : vector<4xi32> + %o2_54 = index.constant 728 : index + vector.store %w2_54, %gv[%o2_54] : vector<4xi32>, view<1024xi32> + %v2_55_0 = scalar.constant -2124906154 : i32 + %v2_55_1 = scalar.constant -2124119719 : i32 + %v2_55_2 = scalar.constant -2123792026 : i32 + %v2_55_3 = scalar.constant -2121694843 : i32 + %w2_55 = vector.from_elements %v2_55_0, %v2_55_1, %v2_55_2, %v2_55_3 : vector<4xi32> + %o2_55 = index.constant 732 : index + vector.store %w2_55, %gv[%o2_55] : vector<4xi32>, view<1024xi32> + %v2_56_0 = scalar.constant -2120842860 : i32 + %v2_56_1 = scalar.constant -2119859815 : i32 + %v2_56_2 = scalar.constant -2113764864 : i32 + %v2_56_3 = scalar.constant -2113240568 : i32 + %w2_56 = vector.from_elements %v2_56_0, %v2_56_1, %v2_56_2, %v2_56_3 : vector<4xi32> + %o2_56 = index.constant 736 : index + vector.store %w2_56, %gv[%o2_56] : vector<4xi32>, view<1024xi32> + %v2_57_0 = scalar.constant -2111798763 : i32 + %v2_57_1 = scalar.constant -2111274462 : i32 + %v2_57_2 = scalar.constant -2108587478 : i32 + %v2_57_3 = scalar.constant -2108063148 : i32 + %w2_57 = vector.from_elements %v2_57_0, %v2_57_1, %v2_57_2, %v2_57_3 : vector<4xi32> + %o2_57 = index.constant 740 : index + vector.store %w2_57, %gv[%o2_57] : vector<4xi32>, view<1024xi32> + %v2_58_0 = scalar.constant -2105507227 : i32 + %v2_58_1 = scalar.constant -2104982910 : i32 + %v2_58_2 = scalar.constant -2104130934 : i32 + %v2_58_3 = scalar.constant -2103278944 : i32 + %w2_58 = vector.from_elements %v2_58_0, %v2_58_1, %v2_58_2, %v2_58_3 : vector<4xi32> + %o2_58 = index.constant 744 : index + vector.store %w2_58, %gv[%o2_58] : vector<4xi32>, view<1024xi32> + %v2_59_0 = scalar.constant -2102754648 : i32 + %v2_59_1 = scalar.constant -2078702572 : i32 + %v2_59_2 = scalar.constant -2075884479 : i32 + %v2_59_3 = scalar.constant -2074770351 : i32 + %w2_59 = vector.from_elements %v2_59_0, %v2_59_1, %v2_59_2, %v2_59_3 : vector<4xi32> + %o2_59 = index.constant 748 : index + vector.store %w2_59, %gv[%o2_59] : vector<4xi32>, view<1024xi32> + %v2_60_0 = scalar.constant -2073983910 : i32 + %v2_60_1 = scalar.constant -2073459612 : i32 + %v2_60_2 = scalar.constant -2070313836 : i32 + %v2_60_3 = scalar.constant -2062973695 : i32 + %w2_60 = vector.from_elements %v2_60_0, %v2_60_1, %v2_60_2, %v2_60_3 : vector<4xi32> + %o2_60 = index.constant 752 : index + vector.store %w2_60, %gv[%o2_60] : vector<4xi32>, view<1024xi32> + %v2_61_0 = scalar.constant -2062187246 : i32 + %v2_61_1 = scalar.constant -2061073126 : i32 + %v2_61_2 = scalar.constant -2059369175 : i32 + %v2_61_3 = scalar.constant -2059041471 : i32 + %w2_61 = vector.from_elements %v2_61_0, %v2_61_1, %v2_61_2, %v2_61_3 : vector<4xi32> + %o2_61 = index.constant 756 : index + vector.store %w2_61, %gv[%o2_61] : vector<4xi32>, view<1024xi32> + %v2_62_0 = scalar.constant -2058255032 : i32 + %v2_62_1 = scalar.constant -2057992876 : i32 + %v2_62_2 = scalar.constant -2057730730 : i32 + %v2_62_3 = scalar.constant -2056944294 : i32 + %w2_62 = vector.from_elements %v2_62_0, %v2_62_1, %v2_62_2, %v2_62_3 : vector<4xi32> + %o2_62 = index.constant 760 : index + vector.store %w2_62, %gv[%o2_62] : vector<4xi32>, view<1024xi32> + %v2_63_0 = scalar.constant -2056747674 : i32 + %v2_63_1 = scalar.constant -2055109270 : i32 + %v2_63_2 = scalar.constant -2054781564 : i32 + %v2_63_3 = scalar.constant -2054126199 : i32 + %w2_63 = vector.from_elements %v2_63_0, %v2_63_1, %v2_63_2, %v2_63_3 : vector<4xi32> + %o2_63 = index.constant 764 : index + vector.store %w2_63, %gv[%o2_63] : vector<4xi32>, view<1024xi32> + } + %k3 = index.constant 3 : index + %is3 = index.cmp eq, %chunk, %k3 : index + scf.if %is3 { + %v3_0_0 = scalar.constant -2053798510 : i32 + %v3_0_1 = scalar.constant -2052684392 : i32 + %v3_0_2 = scalar.constant -2045344239 : i32 + %v3_0_3 = scalar.constant -2044361191 : i32 + %w3_0 = vector.from_elements %v3_0_0, %v3_0_1, %v3_0_2, %v3_0_3 : vector<4xi32> + %o3_0 = index.constant 768 : index + vector.store %w3_0, %gv[%o3_0] : vector<4xi32>, view<1024xi32> + %v3_1_0 = scalar.constant -2042329535 : i32 + %v3_1_1 = scalar.constant -2041936311 : i32 + %v3_1_2 = scalar.constant -2041215408 : i32 + %v3_1_3 = scalar.constant -2040887719 : i32 + %w3_1 = vector.from_elements %v3_1_0, %v3_1_1, %v3_1_2, %v3_1_3 : vector<4xi32> + %o3_1 = index.constant 772 : index + vector.store %w3_1, %gv[%o3_1] : vector<4xi32>, view<1024xi32> + %v3_2_0 = scalar.constant -2040101279 : i32 + %v3_2_1 = scalar.constant -2038069654 : i32 + %v3_2_2 = scalar.constant -2036693359 : i32 + %v3_2_3 = scalar.constant -2013231452 : i32 + %w3_2 = vector.from_elements %v3_2_0, %v3_2_1, %v3_2_2, %v3_2_3 : vector<4xi32> + %o3_2 = index.constant 776 : index + vector.store %w3_2, %gv[%o3_2] : vector<4xi32>, view<1024xi32> + %v3_3_0 = scalar.constant -2012706814 : i32 + %v3_3_1 = scalar.constant -2011854838 : i32 + %v3_3_2 = scalar.constant -2011002848 : i32 + %v3_3_3 = scalar.constant -2010478552 : i32 + %w3_3 = vector.from_elements %v3_3_0, %v3_3_1, %v3_3_2, %v3_3_3 : vector<4xi32> + %o3_3 = index.constant 780 : index + vector.store %w3_3, %gv[%o3_3] : vector<4xi32>, view<1024xi32> + %v3_4_0 = scalar.constant -2008709055 : i32 + %v3_4_1 = scalar.constant -2007725999 : i32 + %v3_4_2 = scalar.constant -2006611879 : i32 + %v3_4_3 = scalar.constant -2004842391 : i32 + %w3_4 = vector.from_elements %v3_4_0, %v3_4_1, %v3_4_2, %v3_4_3 : vector<4xi32> + %o3_4 = index.constant 784 : index + vector.store %w3_4, %gv[%o3_4] : vector<4xi32>, view<1024xi32> + %v3_5_0 = scalar.constant -2004318078 : i32 + %v3_5_1 = scalar.constant -2003466102 : i32 + %v3_5_2 = scalar.constant -2002614112 : i32 + %v3_5_3 = scalar.constant -2002089816 : i32 + %w3_5 = vector.from_elements %v3_5_0, %v3_5_1, %v3_5_2, %v3_5_3 : vector<4xi32> + %o3_5 = index.constant 788 : index + vector.store %w3_5, %gv[%o3_5] : vector<4xi32>, view<1024xi32> + %v3_6_0 = scalar.constant -1996060411 : i32 + %v3_6_1 = scalar.constant -1995142895 : i32 + %v3_6_2 = scalar.constant -1994028778 : i32 + %v3_6_3 = scalar.constant -1991997119 : i32 + %w3_6 = vector.from_elements %v3_6_0, %v3_6_1, %v3_6_2, %v3_6_3 : vector<4xi32> + %o3_6 = index.constant 792 : index + vector.store %w3_6, %gv[%o3_6] : vector<4xi32>, view<1024xi32> + %v3_7_0 = scalar.constant -1991669434 : i32 + %v3_7_1 = scalar.constant -1991079600 : i32 + %v3_7_2 = scalar.constant -1990555307 : i32 + %v3_7_3 = scalar.constant -1989899935 : i32 + %w3_7 = vector.from_elements %v3_7_0, %v3_7_1, %v3_7_2, %v3_7_3 : vector<4xi32> + %o3_7 = index.constant 796 : index + vector.store %w3_7, %gv[%o3_7] : vector<4xi32>, view<1024xi32> + %v3_8_0 = scalar.constant -1986623099 : i32 + %v3_8_1 = scalar.constant -1985640039 : i32 + %v3_8_2 = scalar.constant -1979545088 : i32 + %v3_8_3 = scalar.constant -1979020792 : i32 + %w3_8 = vector.from_elements %v3_8_0, %v3_8_1, %v3_8_2, %v3_8_3 : vector<4xi32> + %o3_8 = index.constant 800 : index + vector.store %w3_8, %gv[%o3_8] : vector<4xi32>, view<1024xi32> + %v3_9_0 = scalar.constant -1977578987 : i32 + %v3_9_1 = scalar.constant -1977054686 : i32 + %v3_9_2 = scalar.constant -1975154134 : i32 + %v3_9_3 = scalar.constant -1974171055 : i32 + %w3_9 = vector.from_elements %v3_9_0, %v3_9_1, %v3_9_2, %v3_9_3 : vector<4xi32> + %o3_9 = index.constant 804 : index + vector.store %w3_9, %gv[%o3_9] : vector<4xi32>, view<1024xi32> + %v3_10_0 = scalar.constant -1971287466 : i32 + %v3_10_1 = scalar.constant -1970763134 : i32 + %v3_10_2 = scalar.constant -1969911158 : i32 + %v3_10_3 = scalar.constant -1969059168 : i32 + %w3_10 = vector.from_elements %v3_10_0, %v3_10_1, %v3_10_2, %v3_10_3 : vector<4xi32> + %o3_10 = index.constant 808 : index + vector.store %w3_10, %gv[%o3_10] : vector<4xi32>, view<1024xi32> + %v3_11_0 = scalar.constant -1968534872 : i32 + %v3_11_1 = scalar.constant -1877897211 : i32 + %v3_11_2 = scalar.constant -1877438442 : i32 + %v3_11_3 = scalar.constant -1876586471 : i32 + %w3_11 = vector.from_elements %v3_11_0, %v3_11_1, %v3_11_2, %v3_11_3 : vector<4xi32> + %o3_11 = index.constant 812 : index + vector.store %w3_11, %gv[%o3_11] : vector<4xi32>, view<1024xi32> + %v3_12_0 = scalar.constant -1874423743 : i32 + %v3_12_1 = scalar.constant -1873440695 : i32 + %v3_12_2 = scalar.constant -1873113000 : i32 + %v3_12_3 = scalar.constant -1872064407 : i32 + %w3_12 = vector.from_elements %v3_12_0, %v3_12_1, %v3_12_2, %v3_12_3 : vector<4xi32> + %o3_12 = index.constant 816 : index + vector.store %w3_12, %gv[%o3_12] : vector<4xi32>, view<1024xi32> + %v3_13_0 = scalar.constant -1869508475 : i32 + %v3_13_1 = scalar.constant -1869180780 : i32 + %v3_13_2 = scalar.constant -1868197735 : i32 + %v3_13_3 = scalar.constant -1861971711 : i32 + %w3_13 = vector.from_elements %v3_13_0, %v3_13_1, %v3_13_2, %v3_13_3 : vector<4xi32> + %o3_13 = index.constant 820 : index + vector.store %w3_13, %gv[%o3_13] : vector<4xi32>, view<1024xi32> + %v3_14_0 = scalar.constant -1861644026 : i32 + %v3_14_1 = scalar.constant -1860857584 : i32 + %v3_14_2 = scalar.constant -1860529896 : i32 + %v3_14_3 = scalar.constant -1859874527 : i32 + %w3_14 = vector.from_elements %v3_14_0, %v3_14_1, %v3_14_2, %v3_14_3 : vector<4xi32> + %o3_14 = index.constant 824 : index + vector.store %w3_14, %gv[%o3_14] : vector<4xi32>, view<1024xi32> + %v3_15_0 = scalar.constant -1859546842 : i32 + %v3_15_1 = scalar.constant -1857711808 : i32 + %v3_15_2 = scalar.constant -1856925360 : i32 + %v3_15_3 = scalar.constant -1856663212 : i32 + %w3_15 = vector.from_elements %v3_15_0, %v3_15_1, %v3_15_2, %v3_15_3 : vector<4xi32> + %o3_15 = index.constant 828 : index + vector.store %w3_15, %gv[%o3_15] : vector<4xi32>, view<1024xi32> + %v3_16_0 = scalar.constant -1856401066 : i32 + %v3_16_1 = scalar.constant -1855614622 : i32 + %v3_16_2 = scalar.constant -1853451900 : i32 + %v3_16_3 = scalar.constant -1852468846 : i32 + %w3_16 = vector.from_elements %v3_16_0, %v3_16_1, %v3_16_2, %v3_16_3 : vector<4xi32> + %o3_16 = index.constant 832 : index + vector.store %w3_16, %gv[%o3_16] : vector<4xi32>, view<1024xi32> + %v3_17_0 = scalar.constant -1851682408 : i32 + %v3_17_1 = scalar.constant -1851354716 : i32 + %v3_17_2 = scalar.constant -1845128791 : i32 + %v3_17_3 = scalar.constant -1844145647 : i32 + %w3_17 = vector.from_elements %v3_17_0, %v3_17_1, %v3_17_2, %v3_17_3 : vector<4xi32> + %o3_17 = index.constant 836 : index + vector.store %w3_17, %gv[%o3_17] : vector<4xi32>, view<1024xi32> + %v3_18_0 = scalar.constant -1843031527 : i32 + %v3_18_1 = scalar.constant -1840868796 : i32 + %v3_18_2 = scalar.constant -1840213431 : i32 + %v3_18_3 = scalar.constant -1839885742 : i32 + %w3_18 = vector.from_elements %v3_18_0, %v3_18_1, %v3_18_2, %v3_18_3 : vector<4xi32> + %o3_18 = index.constant 840 : index + vector.store %w3_18, %gv[%o3_18] : vector<4xi32>, view<1024xi32> + %v3_19_0 = scalar.constant -1838771624 : i32 + %v3_19_1 = scalar.constant -1836739991 : i32 + %v3_19_2 = scalar.constant -1835625836 : i32 + %v3_19_3 = scalar.constant -1811836247 : i32 + %w3_19 = vector.from_elements %v3_19_0, %v3_19_1, %v3_19_2, %v3_19_3 : vector<4xi32> + %o3_19 = index.constant 844 : index + vector.store %w3_19, %gv[%o3_19] : vector<4xi32>, view<1024xi32> + %v3_20_0 = scalar.constant -1811508220 : i32 + %v3_20_1 = scalar.constant -1810525168 : i32 + %v3_20_2 = scalar.constant -1809411048 : i32 + %v3_20_3 = scalar.constant -1807051712 : i32 + %w3_20 = vector.from_elements %v3_20_0, %v3_20_1, %v3_20_2, %v3_20_3 : vector<4xi32> + %o3_20 = index.constant 848 : index + vector.store %w3_20, %gv[%o3_20] : vector<4xi32>, view<1024xi32> + %v3_21_0 = scalar.constant -1806396335 : i32 + %v3_21_1 = scalar.constant -1806265259 : i32 + %v3_21_2 = scalar.constant -1806068648 : i32 + %v3_21_3 = scalar.constant -1805544352 : i32 + %w3_21 = vector.from_elements %v3_21_0, %v3_21_1, %v3_21_2, %v3_21_3 : vector<4xi32> + %o3_21 = index.constant 852 : index + vector.store %w3_21, %gv[%o3_21] : vector<4xi32>, view<1024xi32> + %v3_22_0 = scalar.constant -1805282206 : i32 + %v3_22_1 = scalar.constant -1803119484 : i32 + %v3_22_2 = scalar.constant -1802201966 : i32 + %v3_22_3 = scalar.constant -1801939819 : i32 + %w3_22 = vector.from_elements %v3_22_0, %v3_22_1, %v3_22_2, %v3_22_3 : vector<4xi32> + %o3_22 = index.constant 856 : index + vector.store %w3_22, %gv[%o3_22] : vector<4xi32>, view<1024xi32> + %v3_23_0 = scalar.constant -1800825695 : i32 + %v3_23_1 = scalar.constant -1794796288 : i32 + %v3_23_2 = scalar.constant -1794468600 : i32 + %v3_23_3 = scalar.constant -1794009840 : i32 + %w3_23 = vector.from_elements %v3_23_0, %v3_23_1, %v3_23_2, %v3_23_3 : vector<4xi32> + %o3_23 = index.constant 860 : index + vector.store %w3_23, %gv[%o3_23] : vector<4xi32>, view<1024xi32> + %v3_24_0 = scalar.constant -1793747692 : i32 + %v3_24_1 = scalar.constant -1793485546 : i32 + %v3_24_2 = scalar.constant -1792699103 : i32 + %v3_24_3 = scalar.constant -1792371415 : i32 + %w3_24 = vector.from_elements %v3_24_0, %v3_24_1, %v3_24_2, %v3_24_3 : vector<4xi32> + %o3_24 = index.constant 864 : index + vector.store %w3_24, %gv[%o3_24] : vector<4xi32>, view<1024xi32> + %v3_25_0 = scalar.constant -1790667455 : i32 + %v3_25_1 = scalar.constant -1790536379 : i32 + %v3_25_2 = scalar.constant -1789881015 : i32 + %v3_25_3 = scalar.constant -1789749935 : i32 + %w3_25 = vector.from_elements %v3_25_0, %v3_25_1, %v3_25_2, %v3_25_3 : vector<4xi32> + %o3_25 = index.constant 868 : index + vector.store %w3_25, %gv[%o3_25] : vector<4xi32>, view<1024xi32> + %v3_26_0 = scalar.constant -1789553324 : i32 + %v3_26_1 = scalar.constant -1789356714 : i32 + %v3_26_2 = scalar.constant -1789225639 : i32 + %v3_26_3 = scalar.constant -1788570271 : i32 + %w3_26 = vector.from_elements %v3_26_0, %v3_26_1, %v3_26_2, %v3_26_3 : vector<4xi32> + %o3_26 = index.constant 872 : index + vector.store %w3_26, %gv[%o3_26] : vector<4xi32>, view<1024xi32> + %v3_27_0 = scalar.constant -1788439195 : i32 + %v3_27_1 = scalar.constant -1786669719 : i32 + %v3_27_2 = scalar.constant -1786210939 : i32 + %v3_27_3 = scalar.constant -1785555567 : i32 + %w3_27 = vector.from_elements %v3_27_0, %v3_27_1, %v3_27_2, %v3_27_3 : vector<4xi32> + %o3_27 = index.constant 876 : index + vector.store %w3_27, %gv[%o3_27] : vector<4xi32>, view<1024xi32> + %v3_28_0 = scalar.constant -1785358956 : i32 + %v3_28_1 = scalar.constant -1785096810 : i32 + %v3_28_2 = scalar.constant -1784638054 : i32 + %v3_28_3 = scalar.constant -1784310366 : i32 + %w3_28 = vector.from_elements %v3_28_0, %v3_28_1, %v3_28_2, %v3_28_3 : vector<4xi32> + %o3_28 = index.constant 880 : index + vector.store %w3_28, %gv[%o3_28] : vector<4xi32>, view<1024xi32> + %v3_29_0 = scalar.constant -1783982680 : i32 + %v3_29_1 = scalar.constant -1778084351 : i32 + %v3_29_2 = scalar.constant -1776970224 : i32 + %v3_29_3 = scalar.constant -1776249319 : i32 + %w3_29 = vector.from_elements %v3_29_0, %v3_29_1, %v3_29_2, %v3_29_3 : vector<4xi32> + %o3_29 = index.constant 884 : index + vector.store %w3_29, %gv[%o3_29] : vector<4xi32>, view<1024xi32> + %v3_30_0 = scalar.constant -1775659482 : i32 + %v3_30_1 = scalar.constant -1773627835 : i32 + %v3_30_2 = scalar.constant -1773038007 : i32 + %v3_30_3 = scalar.constant -1772775854 : i32 + %w3_30 = vector.from_elements %v3_30_0, %v3_30_1, %v3_30_2, %v3_30_3 : vector<4xi32> + %o3_30 = index.constant 888 : index + vector.store %w3_30, %gv[%o3_30] : vector<4xi32>, view<1024xi32> + %v3_31_0 = scalar.constant -1772513706 : i32 + %v3_31_1 = scalar.constant -1771530651 : i32 + %v3_31_2 = scalar.constant -1769695614 : i32 + %v3_31_3 = scalar.constant -1769302391 : i32 + %w3_31 = vector.from_elements %v3_31_0, %v3_31_1, %v3_31_2, %v3_31_3 : vector<4xi32> + %o3_31 = index.constant 892 : index + vector.store %w3_31, %gv[%o3_31] : vector<4xi32>, view<1024xi32> + %v3_32_0 = scalar.constant -1768647022 : i32 + %v3_32_1 = scalar.constant -1767598443 : i32 + %v3_32_2 = scalar.constant -1767270746 : i32 + %v3_32_3 = scalar.constant -1743349755 : i32 + %w3_32 = vector.from_elements %v3_32_0, %v3_32_1, %v3_32_2, %v3_32_3 : vector<4xi32> + %o3_32 = index.constant 896 : index + vector.store %w3_32, %gv[%o3_32] : vector<4xi32>, view<1024xi32> + %v3_33_0 = scalar.constant -1742366695 : i32 + %v3_33_1 = scalar.constant -1740203967 : i32 + %v3_33_2 = scalar.constant -1739417520 : i32 + %v3_33_3 = scalar.constant -1739155371 : i32 + %w3_33 = vector.from_elements %v3_33_0, %v3_33_1, %v3_33_2, %v3_33_3 : vector<4xi32> + %o3_33 = index.constant 900 : index + vector.store %w3_33, %gv[%o3_33] : vector<4xi32>, view<1024xi32> + %v3_34_0 = scalar.constant -1738237862 : i32 + %v3_34_1 = scalar.constant -1736075163 : i32 + %v3_34_2 = scalar.constant -1734961007 : i32 + %v3_34_3 = scalar.constant -1733977959 : i32 + %w3_34 = vector.from_elements %v3_34_0, %v3_34_1, %v3_34_2, %v3_34_3 : vector<4xi32> + %o3_34 = index.constant 904 : index + vector.store %w3_34, %gv[%o3_34] : vector<4xi32>, view<1024xi32> + %v3_35_0 = scalar.constant -1727620860 : i32 + %v3_35_1 = scalar.constant -1726965495 : i32 + %v3_35_2 = scalar.constant -1726637806 : i32 + %v3_35_3 = scalar.constant -1726310120 : i32 + %w3_35 = vector.from_elements %v3_35_0, %v3_35_1, %v3_35_2, %v3_35_3 : vector<4xi32> + %o3_35 = index.constant 908 : index + vector.store %w3_35, %gv[%o3_35] : vector<4xi32>, view<1024xi32> + %v3_36_0 = scalar.constant -1725851360 : i32 + %v3_36_1 = scalar.constant -1725523676 : i32 + %v3_36_2 = scalar.constant -1723688640 : i32 + %v3_36_3 = scalar.constant -1723295419 : i32 + %w3_36 = vector.from_elements %v3_36_0, %v3_36_1, %v3_36_2, %v3_36_3 : vector<4xi32> + %o3_36 = index.constant 912 : index + vector.store %w3_36, %gv[%o3_36] : vector<4xi32>, view<1024xi32> + %v3_37_0 = scalar.constant -1722705590 : i32 + %v3_37_1 = scalar.constant -1722443436 : i32 + %v3_37_2 = scalar.constant -1722181290 : i32 + %v3_37_3 = scalar.constant -1721394846 : i32 + %w3_37 = vector.from_elements %v3_37_0, %v3_37_1, %v3_37_2, %v3_37_3 : vector<4xi32> + %o3_37 = index.constant 916 : index + vector.store %w3_37, %gv[%o3_37] : vector<4xi32>, view<1024xi32> + %v3_38_0 = scalar.constant -1721067162 : i32 + %v3_38_1 = scalar.constant -1719363199 : i32 + %v3_38_2 = scalar.constant -1718445680 : i32 + %v3_38_3 = scalar.constant -1717921387 : i32 + %w3_38 = vector.from_elements %v3_38_0, %v3_38_1, %v3_38_2, %v3_38_3 : vector<4xi32> + %o3_38 = index.constant 920 : index + vector.store %w3_38, %gv[%o3_38] : vector<4xi32>, view<1024xi32> + %v3_39_0 = scalar.constant -1717134943 : i32 + %v3_39_1 = scalar.constant -1709860347 : i32 + %v3_39_2 = scalar.constant -1706780123 : i32 + %v3_39_3 = scalar.constant -1706452410 : i32 + %w3_39 = vector.from_elements %v3_39_0, %v3_39_1, %v3_39_2, %v3_39_3 : vector<4xi32> + %o3_39 = index.constant 924 : index + vector.store %w3_39, %gv[%o3_39] : vector<4xi32>, view<1024xi32> + %v3_40_0 = scalar.constant -1705665968 : i32 + %v3_40_1 = scalar.constant -1704879528 : i32 + %v3_40_2 = scalar.constant -1701733755 : i32 + %v3_40_3 = scalar.constant -1701471596 : i32 + %w3_40 = vector.from_elements %v3_40_0, %v3_40_1, %v3_40_2, %v3_40_3 : vector<4xi32> + %o3_40 = index.constant 928 : index + vector.store %w3_40, %gv[%o3_40] : vector<4xi32>, view<1024xi32> + %v3_41_0 = scalar.constant -1610573162 : i32 + %v3_41_1 = scalar.constant -1610047486 : i32 + %v3_41_2 = scalar.constant -1609195510 : i32 + %v3_41_3 = scalar.constant -1608343520 : i32 + %w3_41 = vector.from_elements %v3_41_0, %v3_41_1, %v3_41_2, %v3_41_3 : vector<4xi32> + %o3_41 = index.constant 932 : index + vector.store %w3_41, %gv[%o3_41] : vector<4xi32>, view<1024xi32> + %v3_42_0 = scalar.constant -1607819224 : i32 + %v3_42_1 = scalar.constant -1605263291 : i32 + %v3_42_2 = scalar.constant -1604935596 : i32 + %v3_42_3 = scalar.constant -1602183079 : i32 + %w3_42 = vector.from_elements %v3_42_0, %v3_42_1, %v3_42_2, %v3_42_3 : vector<4xi32> + %o3_42 = index.constant 936 : index + vector.store %w3_42, %gv[%o3_42] : vector<4xi32>, view<1024xi32> + %v3_43_0 = scalar.constant -1601658750 : i32 + %v3_43_1 = scalar.constant -1600806774 : i32 + %v3_43_2 = scalar.constant -1599954784 : i32 + %v3_43_3 = scalar.constant -1599430488 : i32 + %w3_43 = vector.from_elements %v3_43_0, %v3_43_1, %v3_43_2, %v3_43_3 : vector<4xi32> + %o3_43 = index.constant 940 : index + vector.store %w3_43, %gv[%o3_43] : vector<4xi32>, view<1024xi32> + %v3_44_0 = scalar.constant -1593204475 : i32 + %v3_44_1 = scalar.constant -1592483567 : i32 + %v3_44_2 = scalar.constant -1592155882 : i32 + %v3_44_3 = scalar.constant -1589206758 : i32 + %w3_44 = vector.from_elements %v3_44_0, %v3_44_1, %v3_44_2, %v3_44_3 : vector<4xi32> + %o3_44 = index.constant 944 : index + vector.store %w3_44, %gv[%o3_44] : vector<4xi32>, view<1024xi32> + %v3_45_0 = scalar.constant -1588485815 : i32 + %v3_45_1 = scalar.constant -1588027051 : i32 + %v3_45_2 = scalar.constant -1587437222 : i32 + %v3_45_3 = scalar.constant -1585077916 : i32 + %w3_45 = vector.from_elements %v3_45_0, %v3_45_1, %v3_45_2, %v3_45_3 : vector<4xi32> + %o3_45 = index.constant 948 : index + vector.store %w3_45, %gv[%o3_45] : vector<4xi32>, view<1024xi32> + %v3_46_0 = scalar.constant -1584225904 : i32 + %v3_46_1 = scalar.constant -1583767146 : i32 + %v3_46_2 = scalar.constant -1576492542 : i32 + %v3_46_3 = scalar.constant -1575968246 : i32 + %w3_46 = vector.from_elements %v3_46_0, %v3_46_1, %v3_46_2, %v3_46_3 : vector<4xi32> + %o3_46 = index.constant 952 : index + vector.store %w3_46, %gv[%o3_46] : vector<4xi32>, view<1024xi32> + %v3_47_0 = scalar.constant -1574788583 : i32 + %v3_47_1 = scalar.constant -1574264280 : i32 + %v3_47_2 = scalar.constant -1571708347 : i32 + %v3_47_3 = scalar.constant -1571184042 : i32 + %w3_47 = vector.from_elements %v3_47_0, %v3_47_1, %v3_47_2, %v3_47_3 : vector<4xi32> + %o3_47 = index.constant 956 : index + vector.store %w3_47, %gv[%o3_47] : vector<4xi32>, view<1024xi32> + %v3_48_0 = scalar.constant -1568628123 : i32 + %v3_48_1 = scalar.constant -1568103806 : i32 + %v3_48_2 = scalar.constant -1567251830 : i32 + %v3_48_3 = scalar.constant -1566399840 : i32 + %w3_48 = vector.from_elements %v3_48_0, %v3_48_1, %v3_48_2, %v3_48_3 : vector<4xi32> + %o3_48 = index.constant 960 : index + vector.store %w3_48, %gv[%o3_48] : vector<4xi32>, view<1024xi32> + %v3_49_0 = scalar.constant -1565875544 : i32 + %v3_49_1 = scalar.constant -1541037031 : i32 + %v3_49_2 = scalar.constant -1539005375 : i32 + %v3_49_3 = scalar.constant -1537956784 : i32 + %w3_49 = vector.from_elements %v3_49_0, %v3_49_1, %v3_49_2, %v3_49_3 : vector<4xi32> + %o3_49 = index.constant 964 : index + vector.store %w3_49, %gv[%o3_49] : vector<4xi32>, view<1024xi32> + %v3_50_0 = scalar.constant -1537694635 : i32 + %v3_50_1 = scalar.constant -1537104806 : i32 + %v3_50_2 = scalar.constant -1536777115 : i32 + %v3_50_3 = scalar.constant -1536580504 : i32 + %w3_50 = vector.from_elements %v3_50_0, %v3_50_1, %v3_50_2, %v3_50_3 : vector<4xi32> + %o3_50 = index.constant 968 : index + vector.store %w3_50, %gv[%o3_50] : vector<4xi32>, view<1024xi32> + %v3_51_0 = scalar.constant -1526291323 : i32 + %v3_51_1 = scalar.constant -1525635831 : i32 + %v3_51_2 = scalar.constant -1525308142 : i32 + %v3_51_3 = scalar.constant -1524194024 : i32 + %w3_51 = vector.from_elements %v3_51_0, %v3_51_1, %v3_51_2, %v3_51_3 : vector<4xi32> + %o3_51 = index.constant 972 : index + vector.store %w3_51, %gv[%o3_51] : vector<4xi32>, view<1024xi32> + %v3_52_0 = scalar.constant -1522358999 : i32 + %v3_52_1 = scalar.constant -1521375931 : i32 + %v3_52_2 = scalar.constant -1521113772 : i32 + %v3_52_3 = scalar.constant -1520851626 : i32 + %w3_52 = vector.from_elements %v3_52_0, %v3_52_1, %v3_52_2, %v3_52_3 : vector<4xi32> + %o3_52 = index.constant 976 : index + vector.store %w3_52, %gv[%o3_52] : vector<4xi32>, view<1024xi32> + %v3_53_0 = scalar.constant -1519737499 : i32 + %v3_53_1 = scalar.constant -1518033535 : i32 + %v3_53_2 = scalar.constant -1517902459 : i32 + %v3_53_3 = scalar.constant -1517116023 : i32 + %w3_53 = vector.from_elements %v3_53_0, %v3_53_1, %v3_53_2, %v3_53_3 : vector<4xi32> + %o3_53 = index.constant 980 : index + vector.store %w3_53, %gv[%o3_53] : vector<4xi32>, view<1024xi32> + %v3_54_0 = scalar.constant -1516722795 : i32 + %v3_54_1 = scalar.constant -1508792827 : i32 + %v3_54_2 = scalar.constant -1508202986 : i32 + %v3_54_3 = scalar.constant -1507482079 : i32 + %w3_54 = vector.from_elements %v3_54_0, %v3_54_1, %v3_54_2, %v3_54_3 : vector<4xi32> + %o3_54 = index.constant 984 : index + vector.store %w3_54, %gv[%o3_54] : vector<4xi32>, view<1024xi32> + %v3_55_0 = scalar.constant -1505319356 : i32 + %v3_55_1 = scalar.constant -1504532918 : i32 + %v3_55_2 = scalar.constant -1504270763 : i32 + %v3_55_3 = scalar.constant -1503615400 : i32 + %w3_55 = vector.from_elements %v3_55_0, %v3_55_1, %v3_55_2, %v3_55_3 : vector<4xi32> + %o3_55 = index.constant 988 : index + vector.store %w3_55, %gv[%o3_55] : vector<4xi32>, view<1024xi32> + %v3_56_0 = scalar.constant -1501125022 : i32 + %v3_56_1 = scalar.constant -1500141936 : i32 + %v3_56_2 = scalar.constant -1499879786 : i32 + %v3_56_3 = scalar.constant -1499158879 : i32 + %w3_56 = vector.from_elements %v3_56_0, %v3_56_1, %v3_56_2, %v3_56_3 : vector<4xi32> + %o3_56 = index.constant 992 : index + vector.store %w3_56, %gv[%o3_56] : vector<4xi32>, view<1024xi32> + %v3_57_0 = scalar.constant -1476352346 : i32 + %v3_57_1 = scalar.constant -1475827710 : i32 + %v3_57_2 = scalar.constant -1474254838 : i32 + %v3_57_3 = scalar.constant -1473730526 : i32 + %w3_57 = vector.from_elements %v3_57_0, %v3_57_1, %v3_57_2, %v3_57_3 : vector<4xi32> + %o3_57 = index.constant 996 : index + vector.store %w3_57, %gv[%o3_57] : vector<4xi32>, view<1024xi32> + %v3_58_0 = scalar.constant -1471043542 : i32 + %v3_58_1 = scalar.constant -1470715820 : i32 + %v3_58_2 = scalar.constant -1467963303 : i32 + %v3_58_3 = scalar.constant -1467438974 : i32 + %w3_58 = vector.from_elements %v3_58_0, %v3_58_1, %v3_58_2, %v3_58_3 : vector<4xi32> + %o3_58 = index.constant 1000 : index + vector.store %w3_58, %gv[%o3_58] : vector<4xi32>, view<1024xi32> + %v3_59_0 = scalar.constant -1466586998 : i32 + %v3_59_1 = scalar.constant -1465735008 : i32 + %v3_59_2 = scalar.constant -1465210712 : i32 + %v3_59_3 = scalar.constant -1458263803 : i32 + %w3_59 = vector.from_elements %v3_59_0, %v3_59_1, %v3_59_2, %v3_59_3 : vector<4xi32> + %o3_59 = index.constant 1004 : index + vector.store %w3_59, %gv[%o3_59] : vector<4xi32>, view<1024xi32> + %v3_60_0 = scalar.constant -1457411815 : i32 + %v3_60_1 = scalar.constant -1455314651 : i32 + %v3_60_2 = scalar.constant -1454003888 : i32 + %v3_60_3 = scalar.constant -1453217446 : i32 + %w3_60 = vector.from_elements %v3_60_0, %v3_60_1, %v3_60_2, %v3_60_3 : vector<4xi32> + %o3_60 = index.constant 1008 : index + vector.store %w3_60, %gv[%o3_60] : vector<4xi32>, view<1024xi32> + %v3_61_0 = scalar.constant -1452693146 : i32 + %v3_61_1 = scalar.constant -1449743984 : i32 + %v3_61_2 = scalar.constant -1442665984 : i32 + %v3_61_3 = scalar.constant -1442141688 : i32 + %w3_61 = vector.from_elements %v3_61_0, %v3_61_1, %v3_61_2, %v3_61_3 : vector<4xi32> + %o3_61 = index.constant 1012 : index + vector.store %w3_61, %gv[%o3_61] : vector<4xi32>, view<1024xi32> + %v3_62_0 = scalar.constant -1440568800 : i32 + %v3_62_1 = scalar.constant -1440044504 : i32 + %v3_62_2 = scalar.constant -1437291951 : i32 + %v3_62_3 = scalar.constant -1434408362 : i32 + %w3_62 = vector.from_elements %v3_62_0, %v3_62_1, %v3_62_2, %v3_62_3 : vector<4xi32> + %o3_62 = index.constant 1016 : index + vector.store %w3_62, %gv[%o3_62] : vector<4xi32>, view<1024xi32> + %v3_63_0 = scalar.constant -1433884030 : i32 + %v3_63_1 = scalar.constant -1433032054 : i32 + %v3_63_2 = scalar.constant -1432180064 : i32 + %v3_63_3 = scalar.constant -1431655768 : i32 + %w3_63 = vector.from_elements %v3_63_0, %v3_63_1, %v3_63_2, %v3_63_3 : vector<4xi32> + %o3_63 = index.constant 1020 : index + vector.store %w3_63, %gv[%o3_63] : vector<4xi32>, view<1024xi32> + } + func.return +} + +// IQ1_M's fp16 block scale (as motifs/dequant.loom's ggml_iq1m_block_scale). +func.def inline @ggml_kquant_iq1m_block_scale(%weight: buffer, %block_byte_base: offset) -> (f32) { + %wv = buffer.view %weight[%block_byte_base] : buffer -> view<28xi16> + %c24 = index.constant 24 : index + %c25 = index.constant 25 : index + %c26 = index.constant 26 : index + %c27 = index.constant 27 : index + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c12_i32 = scalar.constant 12 : i32 + %c240_i32 = scalar.constant 240 : i32 + %c3840_i32 = scalar.constant 3840 : i32 + %c61440_i32 = scalar.constant 61440 : i32 + %c255_i32 = scalar.constant 255 : i32 + %w0_i16 = view.load %wv[%c24] : view<28xi16> -> i16 + %w1_i16 = view.load %wv[%c25] : view<28xi16> -> i16 + %w2_i16 = view.load %wv[%c26] : view<28xi16> -> i16 + %w3_i16 = view.load %wv[%c27] : view<28xi16> -> i16 + %w0 = scalar.extui %w0_i16 : i16 to i32 + %w1 = scalar.extui %w1_i16 : i16 to i32 + %w2 = scalar.extui %w2_i16 : i16 to i32 + %w3 = scalar.extui %w3_i16 : i16 to i32 + %n0 = scalar.shrui %w0, %c12_i32 : i32 + %n1a = scalar.shrui %w1, %c8_i32 : i32 + %n1 = scalar.andi %n1a, %c240_i32 : i32 + %n2a = scalar.shrui %w2, %c4_i32 : i32 + %n2 = scalar.andi %n2a, %c3840_i32 : i32 + %n3 = scalar.andi %w3, %c61440_i32 : i32 + %u01 = scalar.ori %n0, %n1 : i32 + %u012 = scalar.ori %u01, %n2 : i32 + %u = scalar.ori %u012, %n3 : i32 + %lo = scalar.andi %u, %c255_i32 : i32 + %hi = scalar.shrui %u, %c8_i32 : i32 + %lo8 = scalar.trunci %lo : i32 to i8 + %hi8 = scalar.trunci %hi : i32 to i8 + %bytes = vector.from_elements %lo8, %hi8 : vector<2xi8> + %h = vector.bitcast %bytes : vector<2xi8> to vector<1xf16> + %h0 = vector.extract %h[0] : vector<1xf16> -> f16 + %f = scalar.extf %h0 : f16 to f32 + func.return %f : f32 +} + +// IQ1_S / IQ1_M: lane l16 owns group l16 / 2, slots 2 (l16 % 2) and +1: values 32 (l16 / 2) + +// 16 (l16 % 2) .. +15, one scale. The grid (2048 16-bit codes, 2 bits per value: value + 1) is +// staged in workgroup memory two codes per word. +func.def inline @ggml_kquant_iq1_code(%grid: buffer, %index: i32) -> (i32) { + %zero_offset = index.constant 0 : offset + %gv = buffer.view %grid[%zero_offset] : buffer -> view<1024xi32> + %c1_i32 = scalar.constant 1 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c65535_i32 = scalar.constant 65535 : i32 + %w_i32 = scalar.shrui %index, %c1_i32 : i32 + %w_x = index.cast %w_i32 : i32 to index + %w_b = index.assume %w_x [range(%w_x, 0, 1023)] : index + %word = view.load %gv[%w_b] : view<1024xi32> -> i32 + %odd = scalar.andi %index, %c1_i32 : i32 + %sh = scalar.shli %odd, %c4_i32 : i32 + %shifted = scalar.shrui %word, %sh : i32 + %code = scalar.andi %shifted, %c65535_i32 : i32 + func.return %code : i32 +} + +func.def inline @ggml_kquant_iq1_slot_values(%code: i32, %delta: f32) -> (vector<8xf32>) { + %c1_i32 = scalar.constant 1 : i32 + %s0 = scalar.constant 0 : i32 + %s2 = scalar.constant 2 : i32 + %s3 = scalar.constant 3 : i32 + %s4 = scalar.constant 4 : i32 + %s6 = scalar.constant 6 : i32 + %s8 = scalar.constant 8 : i32 + %s10 = scalar.constant 10 : i32 + %s12 = scalar.constant 12 : i32 + %s14 = scalar.constant 14 : i32 + %shift2 = vector.from_elements %s0, %s2, %s4, %s6, %s8, %s10, %s12, %s14 : vector<8xi32> + %three = vector.splat %s3 : vector<8xi32> + %one = vector.splat %c1_i32 : vector<8xi32> + %cv = vector.splat %code : vector<8xi32> + %cs = vector.shrui %cv, %shift2 : vector<8xi32> + %c = vector.andi %cs, %three : vector<8xi32> + %t = vector.subi %c, %one : vector<8xi32> + %tf = vector.sitofp %t : vector<8xi32> to vector<8xf32> + %dv = vector.splat %delta : vector<8xf32> + %v = vector.addf %tf, %dv : vector<8xf32> + func.return %v : vector<8xf32> +} + +// IQ1_S (50 bytes: d, qs[32], qh[8] u16): qh[g] = high index bits (slot l at bit 3l), scale s at +// bit 12, delta sign at bit 15: value = d * (2 s + 1) * (grid + delta). +func.def inline @ggml_kquant_iq1s_lane_parts(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c17 = index.constant 17 : index + %c32 = index.constant 32 : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c12_i32 = scalar.constant 12 : i32 + %c32768_i32 = scalar.constant 32768 : i32 + %pos_delta = scalar.constant 0.125 : f32 + %neg_delta = scalar.constant -0.125 : f32 + %block_bytes = index.constant 50 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %g = index.div %l, %c2 : index + %h = index.rem %l, %c2 : index + %sa = index.mul %h, %c2 : index + %sb = index.add %sa, %c1 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %hv = buffer.view %weight[%block_base] : buffer -> view<25xf16> + %wv = buffer.view %weight[%block_base] : buffer -> view<25xi16> + %bv = buffer.view %weight[%block_base] : buffer -> view<50xi8> + %d_f16 = view.load %hv[%c0] : view<25xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %qh_at = index.add %c17, %g : index + %qh_i16 = view.load %wv[%qh_at] : view<25xi16> -> i16 + %qh = scalar.extui %qh_i16 : i16 to i32 + %g4 = index.mul %g, %c4 : index + %qs0 = index.add %g4, %c2 : index + %qa_at = index.add %qs0, %sa : index + %qb_at = index.add %qs0, %sb : index + %qa_i8 = view.load %bv[%qa_at] : view<50xi8> -> i8 + %qb_i8 = view.load %bv[%qb_at] : view<50xi8> -> i8 + %qa = scalar.extui %qa_i8 : i8 to i32 + %qb = scalar.extui %qb_i8 : i8 to i32 + %sa_i32 = index.cast %sa : index to i32 + %sb_i32 = index.cast %sb : index to i32 + %sha = scalar.muli %sa_i32, %c3_i32 : i32 + %shb = scalar.muli %sb_i32, %c3_i32 : i32 + %ha0 = scalar.shrui %qh, %sha : i32 + %hb0 = scalar.shrui %qh, %shb : i32 + %ha1 = scalar.andi %ha0, %c7_i32 : i32 + %hb1 = scalar.andi %hb0, %c7_i32 : i32 + %ha = scalar.shli %ha1, %c8_i32 : i32 + %hb = scalar.shli %hb1, %c8_i32 : i32 + %ia = scalar.ori %qa, %ha : i32 + %ib = scalar.ori %qb, %hb : i32 + %code_a = func.call @ggml_kquant_iq1_code(%grid, %ia) : (buffer, i32) -> (i32) + %code_b = func.call @ggml_kquant_iq1_code(%grid, %ib) : (buffer, i32) -> (i32) + %s0 = scalar.shrui %qh, %c12_i32 : i32 + %s = scalar.andi %s0, %c7_i32 : i32 + %s2 = scalar.addi %s, %s : i32 + %s21 = scalar.addi %s2, %c1_i32 : i32 + %sf = scalar.uitofp %s21 : i32 to f32 + %scale = scalar.mulf %d, %sf : f32 + %neg_bit = scalar.andi %qh, %c32768_i32 : i32 + %neg = scalar.cmpi ne, %neg_bit, %c0_i32 : i32 + %delta = scf.select %neg, %neg_delta, %pos_delta : f32 + %va = func.call @ggml_kquant_iq1_slot_values(%code_a, %delta) : (i32, f32) -> (vector<8xf32>) + %vb = func.call @ggml_kquant_iq1_slot_values(%code_b, %delta) : (i32, f32) -> (vector<8xf32>) + %v = func.call @ggml_kquant_iq2_join16(%va, %vb) : (vector<8xf32>, vector<8xf32>) -> (vector<16xf32>) + %g32 = index.mul %g, %c32 : index + %h16 = index.mul %h, %c16 : index + %p0 = index.add %g32, %h16 : index + func.return %v, %scale, %p0 : vector<16xf32>, f32, index +} + +// IQ1_M (56 bytes: qs[32], qh[16], scales[8]): the lane's two slots share qh byte 2g + h (index +// high bits 0..2 / 4..6, delta signs bit 3 / 7) and the 3-bit scale at bit 6 (g % 2) + 3h of u16 +// scale word g / 2; d is the fp16 made of the four scale words' top nibbles. +func.def inline @ggml_kquant_iq1m_lane_parts(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c24 = index.constant 24 : index + %c32 = index.constant 32 : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c6_i32 = scalar.constant 6 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c128_i32 = scalar.constant 128 : i32 + %c1792_i32 = scalar.constant 1792 : i32 + %pos_delta = scalar.constant 0.125 : f32 + %neg_delta = scalar.constant -0.125 : f32 + %block_bytes = index.constant 56 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %g = index.div %l, %c2 : index + %h = index.rem %l, %c2 : index + %sa = index.mul %h, %c2 : index + %sb = index.add %sa, %c1 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %wv = buffer.view %weight[%block_base] : buffer -> view<28xi16> + %bv = buffer.view %weight[%block_base] : buffer -> view<56xi8> + %d = func.call @ggml_kquant_iq1m_block_scale(%weight, %block_base) : (buffer, offset) -> (f32) + %g4 = index.mul %g, %c4 : index + %qa_at = index.add %g4, %sa : index + %qb_at = index.add %g4, %sb : index + %qa_i8 = view.load %bv[%qa_at] : view<56xi8> -> i8 + %qb_i8 = view.load %bv[%qb_at] : view<56xi8> -> i8 + %qa = scalar.extui %qa_i8 : i8 to i32 + %qb = scalar.extui %qb_i8 : i8 to i32 + %g2 = index.add %g, %g : index + %qh_at0 = index.add %c32, %g2 : index + %qh_at = index.add %qh_at0, %h : index + %qh_i8 = view.load %bv[%qh_at] : view<56xi8> -> i8 + %qh = scalar.extui %qh_i8 : i8 to i32 + %ha0 = scalar.shli %qh, %c8_i32 : i32 + %ha = scalar.andi %ha0, %c1792_i32 : i32 + %hb0 = scalar.shli %qh, %c4_i32 : i32 + %hb = scalar.andi %hb0, %c1792_i32 : i32 + %ia = scalar.ori %qa, %ha : i32 + %ib = scalar.ori %qb, %hb : i32 + %code_a = func.call @ggml_kquant_iq1_code(%grid, %ia) : (buffer, i32) -> (i32) + %code_b = func.call @ggml_kquant_iq1_code(%grid, %ib) : (buffer, i32) -> (i32) + %na0 = scalar.andi %qh, %c8_i32 : i32 + %nb0 = scalar.andi %qh, %c128_i32 : i32 + %na = scalar.cmpi ne, %na0, %c0_i32 : i32 + %nb = scalar.cmpi ne, %nb0, %c0_i32 : i32 + %delta_a = scf.select %na, %neg_delta, %pos_delta : f32 + %delta_b = scf.select %nb, %neg_delta, %pos_delta : f32 + %gh = index.div %g, %c2 : index + %gp = index.rem %g, %c2 : index + %sc_at = index.add %c24, %gh : index + %sc_i16 = view.load %wv[%sc_at] : view<28xi16> -> i16 + %sc = scalar.extui %sc_i16 : i16 to i32 + %gp_i32 = index.cast %gp : index to i32 + %h_i32 = index.cast %h : index to i32 + %ssh0 = scalar.muli %gp_i32, %c6_i32 : i32 + %ssh1 = scalar.muli %h_i32, %c3_i32 : i32 + %ssh = scalar.addi %ssh0, %ssh1 : i32 + %s0 = scalar.shrui %sc, %ssh : i32 + %s = scalar.andi %s0, %c7_i32 : i32 + %s2 = scalar.addi %s, %s : i32 + %s21 = scalar.addi %s2, %c1_i32 : i32 + %sf = scalar.uitofp %s21 : i32 to f32 + %scale = scalar.mulf %d, %sf : f32 + %va = func.call @ggml_kquant_iq1_slot_values(%code_a, %delta_a) : (i32, f32) -> (vector<8xf32>) + %vb = func.call @ggml_kquant_iq1_slot_values(%code_b, %delta_b) : (i32, f32) -> (vector<8xf32>) + %v = func.call @ggml_kquant_iq2_join16(%va, %vb) : (vector<8xf32>, vector<8xf32>) -> (vector<16xf32>) + %g32 = index.mul %g, %c32 : index + %h16 = index.mul %h, %c16 : index + %p0 = index.add %g32, %h16 : index + func.return %v, %scale, %p0 : vector<16xf32>, f32, index +} + +func.def inline @ggml_kquant_iq1s_lane_dot(%grid: buffer, %weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %zero_scalar = scalar.constant 0.0 : f32 + %x_block_bytes = index.constant 1024 : offset + %v, %scale, %p = func.call @ggml_kquant_iq1s_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %x = vector.load %xv[%p] : view<256xf32> -> vector<16xf32> + %vx = vector.mulf %v, %x : vector<16xf32> + %sum = vector.reduce %vx, %zero_scalar : vector<16xf32>, f32 + %result = scalar.mulf %scale, %sum : f32 + func.return %result : f32 +} + +func.def inline @ggml_kquant_iq1s_lane_weights(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c12 = index.constant 12 : index + %v, %scale, %p0 = func.call @ggml_kquant_iq1s_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %sv = vector.splat %scale : vector<16xf32> + %w = vector.mulf %v, %sv : vector<16xf32> + %p1 = index.add %p0, %c4 : index + %p2 = index.add %p0, %c8 : index + %p3 = index.add %p0, %c12 : index + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +func.def inline @ggml_kquant_iq1m_lane_dot(%grid: buffer, %weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %zero_scalar = scalar.constant 0.0 : f32 + %x_block_bytes = index.constant 1024 : offset + %v, %scale, %p = func.call @ggml_kquant_iq1m_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %x = vector.load %xv[%p] : view<256xf32> -> vector<16xf32> + %vx = vector.mulf %v, %x : vector<16xf32> + %sum = vector.reduce %vx, %zero_scalar : vector<16xf32>, f32 + %result = scalar.mulf %scale, %sum : f32 + func.return %result : f32 +} + +func.def inline @ggml_kquant_iq1m_lane_weights(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c12 = index.constant 12 : index + %v, %scale, %p0 = func.call @ggml_kquant_iq1m_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %sv = vector.splat %scale : vector<16xf32> + %w = vector.mulf %v, %sv : vector<16xf32> + %p1 = index.add %p0, %c4 : index + %p2 = index.add %p0, %c8 : index + %p3 = index.add %p0, %c12 : index + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +// Packed ternary (format 90): a Q4_0 tensor whose values are trits t * d (Bonsai's ternary weights written as +// exact Q4_0 by tools/ternary_to_q4_0.py), repacked at upload by runtime/host-memory-ternary.cpp into 68 bytes +// per 256 values: d0, d1 (f16, one scale per 128-value group), then 64 bytes of 2-bit codes c = t + 1, value i +// at bits 2(i % 4) of byte 4 + i / 4. Lane l owns values 16l..+16: the 32-bit word at byte 4 + 4l (read as two +// 16-bit words; blocks are 2-byte aligned), value j at bits 2j, and the scale of group l / 8. +func.def inline @ggml_kquant_t2_lane_parts(%weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c2 = index.constant 2 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %block_bytes = index.constant 68 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %hv = buffer.view %weight[%block_base] : buffer -> view<34xf16> + %iv = buffer.view %weight[%block_base] : buffer -> view<34xi16> + %g = index.div %l, %c8 : index + %d_f16 = view.load %hv[%g] : view<34xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %l2 = index.mul %l, %c2 : index + %w_at = index.add %l2, %c2 : index + %words16 = vector.load %iv[%w_at] : view<34xi16> -> vector<2xi16> + %words32 = vector.bitcast %words16 : vector<2xi16> to vector<1xi32> + %word = vector.extract %words32[0] : vector<1xi32> -> i32 + %s0 = scalar.constant 0 : i32 + %s1 = scalar.constant 2 : i32 + %s2 = scalar.constant 4 : i32 + %s3 = scalar.constant 6 : i32 + %s4 = scalar.constant 8 : i32 + %s5 = scalar.constant 10 : i32 + %s6 = scalar.constant 12 : i32 + %s7 = scalar.constant 14 : i32 + %s8 = scalar.constant 16 : i32 + %s9 = scalar.constant 18 : i32 + %s10 = scalar.constant 20 : i32 + %s11 = scalar.constant 22 : i32 + %s12 = scalar.constant 24 : i32 + %s13 = scalar.constant 26 : i32 + %s14 = scalar.constant 28 : i32 + %s15 = scalar.constant 30 : i32 + %shift = vector.from_elements %s0, %s1, %s2, %s3, %s4, %s5, %s6, %s7, %s8, %s9, %s10, %s11, %s12, %s13, %s14, %s15 : vector<16xi32> + %three = vector.splat %c3_i32 : vector<16xi32> + %one = vector.splat %c1_i32 : vector<16xi32> + %wv = vector.splat %word : vector<16xi32> + %ws = vector.shrui %wv, %shift : vector<16xi32> + %c = vector.andi %ws, %three : vector<16xi32> + %t = vector.subi %c, %one : vector<16xi32> + %tf = vector.sitofp %t : vector<16xi32> to vector<16xf32> + %p = index.mul %l, %c16 : index + func.return %tf, %d, %p : vector<16xf32>, f32, index +} + +func.def inline @ggml_kquant_t2_lane_dot(%weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %zero_scalar = scalar.constant 0.0 : f32 + %x_block_bytes = index.constant 1024 : offset + %v, %scale, %p = func.call @ggml_kquant_t2_lane_parts(%weight, %row_base, %block, %lane16) : (buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %x = vector.load %xv[%p] : view<256xf32> -> vector<16xf32> + %vx = vector.mulf %v, %x : vector<16xf32> + %sum = vector.reduce %vx, %zero_scalar : vector<16xf32>, f32 + %result = scalar.mulf %scale, %sum : f32 + func.return %result : f32 +} + +func.def inline @ggml_kquant_t2_lane_weights(%weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c12 = index.constant 12 : index + %v, %scale, %p0 = func.call @ggml_kquant_t2_lane_parts(%weight, %row_base, %block, %lane16) : (buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %sv = vector.splat %scale : vector<16xf32> + %w = vector.mulf %v, %sv : vector<16xf32> + %p1 = index.add %p0, %c4 : index + %p2 = index.add %p0, %c8 : index + %p3 = index.add %p0, %c12 : index + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +// Group-128 types read in their GGUF layout: PrismML's PQ2_0 / PTQ1_0 (ggml-prism.h) and ggml's Q1_0 (below). +// A 256-value block is two 128-value blocks; lane l owns values 16l..+16, all in block l / 8, whose fp16 scale it reads. +// PQ2_0 (format 72): 68 bytes per 256 values = two [d, qs[32]]; lane l's codes are the 32-bit word at byte +// 34 (l / 8) + 2 + 4 (l % 8) (read as two 16-bit words; blocks are 2-byte aligned), value j at bits 2j, +// w = (code - 1) * d (codes 0..3 = -1, 0, +1, +2), as packed ternary (format 90) with the scale beside +// its codes. +// PTQ1_0 (format 73): 56 bytes per 256 values = two [qs[24], qh[2], d], 4-byte aligned, read as seven 32-bit +// words per block. With r = l % 8: r < 5 takes trit r of qs[0..15] (words 0..3); r = 5, 6 take trits +// 2(r - 5), 2(r - 5) + 1 of qs[16..23] (words 4, 5, 4, 5); r = 7 takes trit 4 of qs[16..23] and then +// qh0, qh1, qh0, qh1, ... with trits 0, 0, 1, 1, 2, 2, 3, 3 (word 6). Trit n of byte b is +// ((uint8_t)(b * 3^n) * 3) >> 8, minus 1. +func.def pure inline @ggml_kquant_g128_format(%format: index) -> (i1) { + %f10 = index.constant 10 : index + %f72 = index.constant 72 : index + %f73 = index.constant 73 : index + %is10 = index.cmp eq, %format, %f10 : index + %is72 = index.cmp eq, %format, %f72 : index + %is73 = index.cmp eq, %format, %f73 : index + %is7273 = scalar.ori %is72, %is73 : i1 + %is = scalar.ori %is7273, %is10 : i1 + func.return %is : i1 +} + +// Bytes per 256 values: 36 (Q1_0), 68 (PQ2_0), 56 (PTQ1_0), else %other. +func.def pure inline @ggml_kquant_g128_block_bytes(%format: index, %other: offset) -> (offset) { + %f72 = index.constant 72 : index + %f73 = index.constant 73 : index + %b68 = index.constant 68 : offset + %b56 = index.constant 56 : offset + %is72 = index.cmp eq, %format, %f72 : index + %is73 = index.cmp eq, %format, %f73 : index + %b0 = scf.select %is72, %b68, %other : offset + %b1 = scf.select %is73, %b56, %b0 : offset + %f10 = index.constant 10 : index + %b36 = index.constant 36 : offset + %is10 = index.cmp eq, %format, %f10 : index + %b2 = scf.select %is10, %b36, %b1 : offset + func.return %b2 : offset +} + +func.def inline @ggml_kquant_pq2_lane_parts(%weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c17 = index.constant 17 : index + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %block_bytes = index.constant 68 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %hv = buffer.view %weight[%block_base] : buffer -> view<34xf16> + %iv = buffer.view %weight[%block_base] : buffer -> view<34xi16> + %g = index.div %l, %c8 : index + %r = index.rem %l, %c8 : index + %d_at = index.mul %g, %c17 : index + %d_f16 = view.load %hv[%d_at] : view<34xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %r2 = index.mul %r, %c2 : index + %w_at0 = index.add %d_at, %c1 : index + %w_at = index.add %w_at0, %r2 : index + %words16 = vector.load %iv[%w_at] : view<34xi16> -> vector<2xi16> + %words32 = vector.bitcast %words16 : vector<2xi16> to vector<1xi32> + %word = vector.extract %words32[0] : vector<1xi32> -> i32 + %s0 = scalar.constant 0 : i32 + %s1 = scalar.constant 2 : i32 + %s2 = scalar.constant 4 : i32 + %s3 = scalar.constant 6 : i32 + %s4 = scalar.constant 8 : i32 + %s5 = scalar.constant 10 : i32 + %s6 = scalar.constant 12 : i32 + %s7 = scalar.constant 14 : i32 + %s8 = scalar.constant 16 : i32 + %s9 = scalar.constant 18 : i32 + %s10 = scalar.constant 20 : i32 + %s11 = scalar.constant 22 : i32 + %s12 = scalar.constant 24 : i32 + %s13 = scalar.constant 26 : i32 + %s14 = scalar.constant 28 : i32 + %s15 = scalar.constant 30 : i32 + %shift = vector.from_elements %s0, %s1, %s2, %s3, %s4, %s5, %s6, %s7, %s8, %s9, %s10, %s11, %s12, %s13, %s14, %s15 : vector<16xi32> + %three = vector.splat %c3_i32 : vector<16xi32> + %one = vector.splat %c1_i32 : vector<16xi32> + %wv = vector.splat %word : vector<16xi32> + %ws = vector.shrui %wv, %shift : vector<16xi32> + %c = vector.andi %ws, %three : vector<16xi32> + %t = vector.subi %c, %one : vector<16xi32> + %tf = vector.sitofp %t : vector<16xi32> to vector<16xf32> + %p = index.mul %l, %c16 : index + func.return %tf, %d, %p : vector<16xf32>, f32, index +} + +func.def inline @ggml_kquant_ptq1_lane_parts(%weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c5 = index.constant 5 : index + %c6 = index.constant 6 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c13 = index.constant 13 : index + %c14 = index.constant 14 : index + %c16 = index.constant 16 : index + %block_bytes = index.constant 56 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %wv = buffer.view %weight[%block_base] : buffer -> view<14xi32> + %hv = buffer.view %weight[%block_base] : buffer -> view<28xf16> + %g = index.div %l, %c8 : index + %r = index.rem %l, %c8 : index + %d_at0 = index.mul %g, %c14 : index + %d_at = index.add %d_at0, %c13 : index + %d_f16 = view.load %hv[%d_at] : view<28xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %is_a = index.cmp ult, %r, %c5 : index + %is_c = index.cmp eq, %r, %c7 : index + // the four words: 0, 1, 2, 3 (r < 5); 4, 5, 4, 5 (r = 5, 6); 4, 5, 6, 6 (r = 7) + %i0 = scf.select %is_a, %c0, %c4 : index + %i1 = scf.select %is_a, %c1, %c5 : index + %i2b = scf.select %is_c, %c6, %c4 : index + %i3b = scf.select %is_c, %c6, %c5 : index + %i2 = scf.select %is_a, %c2, %i2b : index + %i3 = scf.select %is_a, %c3, %i3b : index + %wbase = index.mul %g, %c7 : index + %j0 = index.add %wbase, %i0 : index + %j1 = index.add %wbase, %i1 : index + %j2 = index.add %wbase, %i2 : index + %j3 = index.add %wbase, %i3 : index + %w0 = view.load %wv[%j0] : view<14xi32> -> i32 + %w1 = view.load %wv[%j1] : view<14xi32> -> i32 + %w2 = view.load %wv[%j2] : view<14xi32> -> i32 + %w3 = view.load %wv[%j3] : view<14xi32> -> i32 + %words = vector.from_elements %w0, %w0, %w0, %w0, %w1, %w1, %w1, %w1, %w2, %w2, %w2, %w2, %w3, %w3, %w3, %w3 : vector<16xi32> + %s0 = scalar.constant 0 : i32 + %s8 = scalar.constant 8 : i32 + %s16 = scalar.constant 16 : i32 + %s24 = scalar.constant 24 : i32 + %shift_n = vector.from_elements %s0, %s8, %s16, %s24, %s0, %s8, %s16, %s24, %s0, %s8, %s16, %s24, %s0, %s8, %s16, %s24 : vector<16xi32> + %shift_c = vector.from_elements %s0, %s8, %s16, %s24, %s0, %s8, %s16, %s24, %s0, %s8, %s0, %s8, %s0, %s8, %s0, %s8 : vector<16xi32> + %shift = scf.select %is_c, %shift_c, %shift_n : vector<16xi32> + // powers of three: 3^r (r < 5); 3^(2(r - 5)) then 3^(2(r - 5) + 1) (r = 5, 6); 81 then 1, 1, 3, 3, 9, 9, 27, 27 (r = 7) + %p1 = scalar.constant 1 : i32 + %p3 = scalar.constant 3 : i32 + %p9 = scalar.constant 9 : i32 + %p27 = scalar.constant 27 : i32 + %p81 = scalar.constant 81 : i32 + %is_r1 = index.cmp eq, %r, %c1 : index + %is_r2 = index.cmp eq, %r, %c2 : index + %is_r3 = index.cmp eq, %r, %c3 : index + %is_r4 = index.cmp eq, %r, %c4 : index + %is_r6 = index.cmp eq, %r, %c6 : index + %pa1 = scf.select %is_r1, %p3, %p1 : i32 + %pa2 = scf.select %is_r2, %p9, %pa1 : i32 + %pa3 = scf.select %is_r3, %p27, %pa2 : i32 + %pa = scf.select %is_r4, %p81, %pa3 : i32 + %pb = scf.select %is_r6, %p9, %p1 : i32 + %pb_hi = scalar.muli %pb, %p3 : i32 + %plo0 = scf.select %is_a, %pa, %pb : i32 + %plo = scf.select %is_c, %p81, %plo0 : i32 + %phi = scf.select %is_a, %pa, %pb_hi : i32 + %pow_n = vector.from_elements %plo, %plo, %plo, %plo, %plo, %plo, %plo, %plo, %phi, %phi, %phi, %phi, %phi, %phi, %phi, %phi : vector<16xi32> + %pow_c = vector.from_elements %p81, %p81, %p81, %p81, %p81, %p81, %p81, %p81, %p1, %p1, %p3, %p3, %p9, %p9, %p27, %p27 : vector<16xi32> + %pow = scf.select %is_c, %pow_c, %pow_n : vector<16xi32> + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c255_i32 = scalar.constant 255 : i32 + %ff = vector.splat %c255_i32 : vector<16xi32> + %three = vector.splat %c3_i32 : vector<16xi32> + %eight = vector.splat %c8_i32 : vector<16xi32> + %one = vector.splat %c1_i32 : vector<16xi32> + %shifted = vector.shrui %words, %shift : vector<16xi32> + %bytes = vector.andi %shifted, %ff : vector<16xi32> + %m = vector.muli %bytes, %pow : vector<16xi32> + %q = vector.andi %m, %ff : vector<16xi32> + %q3 = vector.muli %q, %three : vector<16xi32> + %u = vector.shrui %q3, %eight : vector<16xi32> + %t = vector.subi %u, %one : vector<16xi32> + %tf = vector.sitofp %t : vector<16xi32> to vector<16xf32> + %p = index.mul %l, %c16 : index + func.return %tf, %d, %p : vector<16xf32>, f32, index +} + + +// Q1_0 (format 10, ggml's 1-bit type): 36 bytes per 256 values = two [d (fp16), qs[16]], 2-byte aligned; value j of a +// block is +d if bit j % 8 of qs[j / 8] is set, else -d. Lane l reads the 16 bits of its values, the 16-bit word at +// byte 18 (l / 8) + 2 + 2 (l % 8). +func.def inline @ggml_kquant_q1_0_lane_parts(%weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c1 = index.constant 1 : index + %c8 = index.constant 8 : index + %c9 = index.constant 9 : index + %c16 = index.constant 16 : index + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %block_bytes = index.constant 36 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %hv = buffer.view %weight[%block_base] : buffer -> view<18xf16> + %iv = buffer.view %weight[%block_base] : buffer -> view<18xi16> + %g = index.div %l, %c8 : index + %r = index.rem %l, %c8 : index + %d_at = index.mul %g, %c9 : index + %d_f16 = view.load %hv[%d_at] : view<18xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %w_at0 = index.add %d_at, %c1 : index + %w_at = index.add %w_at0, %r : index + %bits16 = view.load %iv[%w_at] : view<18xi16> -> i16 + %bits = scalar.extui %bits16 : i16 to i32 + %s0 = scalar.constant 0 : i32 + %s1 = scalar.constant 1 : i32 + %s2 = scalar.constant 2 : i32 + %s3 = scalar.constant 3 : i32 + %s4 = scalar.constant 4 : i32 + %s5 = scalar.constant 5 : i32 + %s6 = scalar.constant 6 : i32 + %s7 = scalar.constant 7 : i32 + %s8 = scalar.constant 8 : i32 + %s9 = scalar.constant 9 : i32 + %s10 = scalar.constant 10 : i32 + %s11 = scalar.constant 11 : i32 + %s12 = scalar.constant 12 : i32 + %s13 = scalar.constant 13 : i32 + %s14 = scalar.constant 14 : i32 + %s15 = scalar.constant 15 : i32 + %shift = vector.from_elements %s0, %s1, %s2, %s3, %s4, %s5, %s6, %s7, %s8, %s9, %s10, %s11, %s12, %s13, %s14, %s15 : vector<16xi32> + %one = vector.splat %c1_i32 : vector<16xi32> + %two = vector.splat %c2_i32 : vector<16xi32> + %bv = vector.splat %bits : vector<16xi32> + %bs = vector.shrui %bv, %shift : vector<16xi32> + %b = vector.andi %bs, %one : vector<16xi32> + %b2 = vector.muli %b, %two : vector<16xi32> + %t = vector.subi %b2, %one : vector<16xi32> + %tf = vector.sitofp %t : vector<16xi32> to vector<16xf32> + %p = index.mul %l, %c16 : index + func.return %tf, %d, %p : vector<16xf32>, f32, index +} + +// PTQ1_0 single-token lane dot with byte-uniform ownership: in 128-value block g = l / 8, lane r = l % 8 decodes +// qs[2r] and qs[2r + 1] (values 16n + 2r, 16n + 2r + 1, n = 0..4), qs[16 + r] (values 80 + 8n + r) and trit r / 2 of +// qh[r % 2] (value 120 + r), so every lane runs the same code: the trits of a byte come out most significant first +// as q <- q * 3, t = q >> 8, q &= 255, and the lane gathers its 16 inputs. The 2..8-token kernels keep +// @ggml_kquant_ptq1_lane_parts (16 contiguous values per lane). +func.def inline @ggml_kquant_ptq1_trit_step(%q: i32) -> (i32, i32) { + %c3_i32 = scalar.constant 3 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c255_i32 = scalar.constant 255 : i32 + %u = scalar.muli %q, %c3_i32 : i32 + %t = scalar.shrui %u, %c8_i32 : i32 + %next = scalar.andi %u, %c255_i32 : i32 + func.return %t, %next : i32, i32 +} + +func.def inline @ggml_kquant_ptq1_tf(%t: i32) -> (f32) { + %c1_i32 = scalar.constant 1 : i32 + %s = scalar.subi %t, %c1_i32 : i32 + %f = scalar.sitofp %s : i32 to f32 + func.return %f : f32 +} + +func.def inline @ggml_kquant_ptq1_lane_dot(%weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c8 = index.constant 8 : index + %c14 = index.constant 14 : index + %c13 = index.constant 13 : index + %c16 = index.constant 16 : index + %c24 = index.constant 24 : index + %c28 = index.constant 28 : index + %c80 = index.constant 80 : index + %c120 = index.constant 120 : index + %c128 = index.constant 128 : index + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c9_i32 = scalar.constant 9 : i32 + %c27_i32 = scalar.constant 27 : i32 + %c255_i32 = scalar.constant 255 : i32 + %block_bytes = index.constant 56 : offset + %x_block_bytes = index.constant 1024 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %bv = buffer.view %weight[%block_base] : buffer -> view<56xi8> + %iv = buffer.view %weight[%block_base] : buffer -> view<28xi16> + %hv = buffer.view %weight[%block_base] : buffer -> view<28xf16> + %g = index.div %l, %c8 : index + %r = index.rem %l, %c8 : index + %d_at0 = index.mul %g, %c14 : index + %d_at = index.add %d_at0, %c13 : index + %d_f16 = view.load %hv[%d_at] : view<28xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + // qs[2r], qs[2r + 1] + %a_at = index.add %d_at0, %r : index + %a16 = view.load %iv[%a_at] : view<28xi16> -> i16 + %a = scalar.extui %a16 : i16 to i32 + %qa0 = scalar.andi %a, %c255_i32 : i32 + %qa1 = scalar.shrui %a, %c8_i32 : i32 + // qs[16 + r] + %gb = index.mul %g, %c28 : index + %b_at0 = index.add %gb, %c16 : index + %b_at1 = index.add %b_at0, %r : index + %b_at = index.assume %b_at1 [range(%b_at1, 0, 55)] : index + %b8 = view.load %bv[%b_at] : view<56xi8> -> i8 + %qb = scalar.extui %b8 : i8 to i32 + // trit r / 2 of qh[r % 2] + %h = index.rem %r, %c2 : index + %nc = index.div %r, %c2 : index + %c_at0 = index.add %gb, %c24 : index + %c_at1 = index.add %c_at0, %h : index + %c_at = index.assume %c_at1 [range(%c_at1, 0, 55)] : index + %c8v = view.load %bv[%c_at] : view<56xi8> -> i8 + %qc0 = scalar.extui %c8v : i8 to i32 + %is1 = index.cmp eq, %nc, %c1 : index + %is2 = index.cmp eq, %nc, %c2 : index + %is3 = index.cmp eq, %nc, %c3 : index + %pw1 = scf.select %is1, %c3_i32, %c1_i32 : i32 + %pw2 = scf.select %is2, %c9_i32, %pw1 : i32 + %pw = scf.select %is3, %c27_i32, %pw2 : i32 + %qc1 = scalar.muli %qc0, %pw : i32 + %qc = scalar.andi %qc1, %c255_i32 : i32 + %tc, %qc_n = func.call @ggml_kquant_ptq1_trit_step(%qc) : (i32) -> (i32, i32) + // trits n = 0..4 of the three bytes + %ta0, %qa0_1 = func.call @ggml_kquant_ptq1_trit_step(%qa0) : (i32) -> (i32, i32) + %ta1, %qa0_2 = func.call @ggml_kquant_ptq1_trit_step(%qa0_1) : (i32) -> (i32, i32) + %ta2, %qa0_3 = func.call @ggml_kquant_ptq1_trit_step(%qa0_2) : (i32) -> (i32, i32) + %ta3, %qa0_4 = func.call @ggml_kquant_ptq1_trit_step(%qa0_3) : (i32) -> (i32, i32) + %ta4, %qa0_5 = func.call @ggml_kquant_ptq1_trit_step(%qa0_4) : (i32) -> (i32, i32) + %tb0, %qa1_1 = func.call @ggml_kquant_ptq1_trit_step(%qa1) : (i32) -> (i32, i32) + %tb1, %qa1_2 = func.call @ggml_kquant_ptq1_trit_step(%qa1_1) : (i32) -> (i32, i32) + %tb2, %qa1_3 = func.call @ggml_kquant_ptq1_trit_step(%qa1_2) : (i32) -> (i32, i32) + %tb3, %qa1_4 = func.call @ggml_kquant_ptq1_trit_step(%qa1_3) : (i32) -> (i32, i32) + %tb4, %qa1_5 = func.call @ggml_kquant_ptq1_trit_step(%qa1_4) : (i32) -> (i32, i32) + %tq0, %qb_1 = func.call @ggml_kquant_ptq1_trit_step(%qb) : (i32) -> (i32, i32) + %tq1, %qb_2 = func.call @ggml_kquant_ptq1_trit_step(%qb_1) : (i32) -> (i32, i32) + %tq2, %qb_3 = func.call @ggml_kquant_ptq1_trit_step(%qb_2) : (i32) -> (i32, i32) + %tq3, %qb_4 = func.call @ggml_kquant_ptq1_trit_step(%qb_3) : (i32) -> (i32, i32) + %tq4, %qb_5 = func.call @ggml_kquant_ptq1_trit_step(%qb_4) : (i32) -> (i32, i32) + // inputs: 128g + 16n + 2r (+1), 128g + 80 + 8n + r, 128g + 120 + r + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %p = index.mul %g, %c128 : index + %r2 = index.mul %r, %c2 : index + %pa = index.add %p, %r2 : index + %pa1 = index.add %pa, %c16 : index + %pa2 = index.add %pa1, %c16 : index + %pa3 = index.add %pa2, %c16 : index + %pa4 = index.add %pa3, %c16 : index + %xa0 = vector.load %xv[%pa] : view<256xf32> -> vector<2xf32> + %xa1 = vector.load %xv[%pa1] : view<256xf32> -> vector<2xf32> + %xa2 = vector.load %xv[%pa2] : view<256xf32> -> vector<2xf32> + %xa3 = vector.load %xv[%pa3] : view<256xf32> -> vector<2xf32> + %xa4 = vector.load %xv[%pa4] : view<256xf32> -> vector<2xf32> + %pb0a = index.add %p, %c80 : index + %pb0 = index.add %pb0a, %r : index + %pb1 = index.add %pb0, %c8 : index + %pb2 = index.add %pb1, %c8 : index + %pb3 = index.add %pb2, %c8 : index + %pb4 = index.add %pb3, %c8 : index + %xb0 = view.load %xv[%pb0] : view<256xf32> -> f32 + %xb1 = view.load %xv[%pb1] : view<256xf32> -> f32 + %xb2 = view.load %xv[%pb2] : view<256xf32> -> f32 + %xb3 = view.load %xv[%pb3] : view<256xf32> -> f32 + %xb4 = view.load %xv[%pb4] : view<256xf32> -> f32 + %pca = index.add %p, %c120 : index + %pc = index.add %pca, %r : index + %xc = view.load %xv[%pc] : view<256xf32> -> f32 + %wa0 = func.call @ggml_kquant_ptq1_tf(%ta0) : (i32) -> (f32) + %wa1 = func.call @ggml_kquant_ptq1_tf(%ta1) : (i32) -> (f32) + %wa2 = func.call @ggml_kquant_ptq1_tf(%ta2) : (i32) -> (f32) + %wa3 = func.call @ggml_kquant_ptq1_tf(%ta3) : (i32) -> (f32) + %wa4 = func.call @ggml_kquant_ptq1_tf(%ta4) : (i32) -> (f32) + %wb0 = func.call @ggml_kquant_ptq1_tf(%tb0) : (i32) -> (f32) + %wb1 = func.call @ggml_kquant_ptq1_tf(%tb1) : (i32) -> (f32) + %wb2 = func.call @ggml_kquant_ptq1_tf(%tb2) : (i32) -> (f32) + %wb3 = func.call @ggml_kquant_ptq1_tf(%tb3) : (i32) -> (f32) + %wb4 = func.call @ggml_kquant_ptq1_tf(%tb4) : (i32) -> (f32) + %wq0 = func.call @ggml_kquant_ptq1_tf(%tq0) : (i32) -> (f32) + %wq1 = func.call @ggml_kquant_ptq1_tf(%tq1) : (i32) -> (f32) + %wq2 = func.call @ggml_kquant_ptq1_tf(%tq2) : (i32) -> (f32) + %wq3 = func.call @ggml_kquant_ptq1_tf(%tq3) : (i32) -> (f32) + %wq4 = func.call @ggml_kquant_ptq1_tf(%tq4) : (i32) -> (f32) + %wc = func.call @ggml_kquant_ptq1_tf(%tc) : (i32) -> (f32) + %xa01 = vector.concat<0> %xa0, %xa1 : vector<2xf32>, vector<2xf32> -> vector<4xf32> + %xa23 = vector.concat<0> %xa2, %xa3 : vector<2xf32>, vector<2xf32> -> vector<4xf32> + %xa0123 = vector.concat<0> %xa01, %xa23 : vector<4xf32>, vector<4xf32> -> vector<8xf32> + %xa4b = vector.from_elements %xb0, %xb1, %xb2, %xb3, %xb4, %xc : vector<6xf32> + %xa4full = vector.concat<0> %xa4, %xa4b : vector<2xf32>, vector<6xf32> -> vector<8xf32> + %x16 = vector.concat<0> %xa0123, %xa4full : vector<8xf32>, vector<8xf32> -> vector<16xf32> + // weight order matching %x16: a0[n=0], a1[0], a0[1], a1[1], ... a0[4], a1[4], b[0..4], c + %w16 = vector.from_elements %wa0, %wb0, %wa1, %wb1, %wa2, %wb2, %wa3, %wb3, %wa4, %wb4, %wq0, %wq1, %wq2, %wq3, %wq4, %wc : vector<16xf32> + %zero_scalar = scalar.constant 0.0 : f32 + %vx = vector.mulf %w16, %x16 : vector<16xf32> + %sum = vector.reduce %vx, %zero_scalar : vector<16xf32>, f32 + %result = scalar.mulf %d, %sum : f32 + func.return %result : f32 +} + +func.def inline @ggml_kquant_g128_lane_parts(%format: index, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %f72 = index.constant 72 : index + %is72 = index.cmp eq, %format, %f72 : index + %v, %d, %p = scf.if %is72 -> (vector<16xf32>, f32, index) { + %v0, %d0, %p0 = func.call @ggml_kquant_pq2_lane_parts(%weight, %row_base, %block, %lane16) : (buffer, offset, index, index) -> (vector<16xf32>, f32, index) + scf.yield %v0, %d0, %p0 : vector<16xf32>, f32, index + } else { + %f10 = index.constant 10 : index + %is10 = index.cmp eq, %format, %f10 : index + %v1, %d1, %p1 = scf.if %is10 -> (vector<16xf32>, f32, index) { + %v2, %d2, %p2 = func.call @ggml_kquant_q1_0_lane_parts(%weight, %row_base, %block, %lane16) : (buffer, offset, index, index) -> (vector<16xf32>, f32, index) + scf.yield %v2, %d2, %p2 : vector<16xf32>, f32, index + } else { + %v3, %d3, %p3 = func.call @ggml_kquant_ptq1_lane_parts(%weight, %row_base, %block, %lane16) : (buffer, offset, index, index) -> (vector<16xf32>, f32, index) + scf.yield %v3, %d3, %p3 : vector<16xf32>, f32, index + } + scf.yield %v1, %d1, %p1 : vector<16xf32>, f32, index + } + func.return %v, %d, %p : vector<16xf32>, f32, index +} + +// The lane dot of formats 72 / 73, else Q8_0 (the lane-dot switch's last branch). +func.def inline @ggml_kquant_g128_or_q8_0_lane_dot(%format: index, %weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %is_prism = func.call pure @ggml_kquant_g128_format(%format) : (index) -> (i1) + %f73d = index.constant 73 : index + %is73d = index.cmp eq, %format, %f73d : index + %r = scf.if %is73d -> (f32) { + %v73 = func.call @ggml_kquant_ptq1_lane_dot(%weight, %input, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (f32) + scf.yield %v73 : f32 + } else { + %r0 = scf.if %is_prism -> (f32) { + %zero_scalar = scalar.constant 0.0 : f32 + %x_block_bytes = index.constant 1024 : offset + %v, %scale, %p = func.call @ggml_kquant_g128_lane_parts(%format, %weight, %row_base, %block, %lane16) : (index, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %x = vector.load %xv[%p] : view<256xf32> -> vector<16xf32> + %vx = vector.mulf %v, %x : vector<16xf32> + %sum = vector.reduce %vx, %zero_scalar : vector<16xf32>, f32 + %result = scalar.mulf %scale, %sum : f32 + scf.yield %result : f32 + } else { + %q8 = func.call @ggml_kquant_q8_0_lane_dot(%weight, %input, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (f32) + scf.yield %q8 : f32 + } + scf.yield %r0 : f32 + } + func.return %r : f32 +} + +// The lane weights of formats 72 / 73, else Q8_0 (the lane-weights switch's last branch). +func.def inline @ggml_kquant_g128_or_q8_0_lane_weights(%format: index, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %is_prism = func.call pure @ggml_kquant_g128_format(%format) : (index) -> (i1) + %w, %p0, %p1, %p2, %p3 = scf.if %is_prism -> (vector<16xf32>, index, index, index, index) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c12 = index.constant 12 : index + %v, %scale, %q0 = func.call @ggml_kquant_g128_lane_parts(%format, %weight, %row_base, %block, %lane16) : (index, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %sv = vector.splat %scale : vector<16xf32> + %wv = vector.mulf %v, %sv : vector<16xf32> + %q1 = index.add %q0, %c4 : index + %q2 = index.add %q0, %c8 : index + %q3 = index.add %q0, %c12 : index + scf.yield %wv, %q0, %q1, %q2, %q3 : vector<16xf32>, index, index, index, index + } else { + %v, %v0, %v1, %v2, %v3 = func.call @ggml_kquant_q8_0_lane_weights(%weight, %row_base, %block, %lane16) : (buffer, offset, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %v, %v0, %v1, %v2, %v3 : vector<16xf32>, index, index, index, index + } + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +// Stage the grid of %format (IQ3_S 21, IQ2_S 22, IQ2_XXS 24, IQ2_XS 25, IQ3_XXS 28, and IQ1_S 26 / IQ1_M 27: +// one grid) into workgroup memory, chunk %chunk. +// Callers test the format inline: the barrier after the fill needs a condition Loom can prove +// workgroup-uniform, which a call result is not (STRUCTURE/038). +func.def inline @ggml_kquant_grid_fill_for(%format: index, %grid: buffer, %chunk: index) { + %f21 = index.constant 21 : index + %f24 = index.constant 24 : index + %f25 = index.constant 25 : index + %is21 = index.cmp eq, %format, %f21 : index + %is24 = index.cmp eq, %format, %f24 : index + %is25 = index.cmp eq, %format, %f25 : index + %f28 = index.constant 28 : index + %is28 = index.cmp eq, %format, %f28 : index + %f22 = index.constant 22 : index + %is22 = index.cmp eq, %format, %f22 : index + %f26 = index.constant 26 : index + %f27 = index.constant 27 : index + %is26 = index.cmp eq, %format, %f26 : index + %is27 = index.cmp eq, %format, %f27 : index + %is_iq1 = scalar.ori %is26, %is27 : i1 + scf.if %is21 { + func.call @ggml_kquant_iq3s_grid_fill(%grid, %chunk) : (buffer, index) -> () + } + scf.if %is24 { + func.call @ggml_kquant_iq2xxs_grid_fill(%grid, %chunk) : (buffer, index) -> () + } + scf.if %is25 { + func.call @ggml_kquant_iq2xs_grid_fill(%grid, %chunk) : (buffer, index) -> () + } + scf.if %is28 { + func.call @ggml_kquant_iq3xxs_grid_fill(%grid, %chunk) : (buffer, index) -> () + } + scf.if %is22 { + func.call @ggml_kquant_iq2s_grid_fill(%grid, %chunk) : (buffer, index) -> () + } + scf.if %is_iq1 { + func.call @ggml_kquant_iq1s_grid_fill(%grid, %chunk) : (buffer, index) -> () + } + func.return +} + +// Four little-endian 32-bit words as their 16 bytes, one per element (0..255). +func.def inline @ggml_kquant_run16_bytes(%words: vector<4xi32>) -> (vector<16xi32>) { + %w0 = vector.extract %words[0] : vector<4xi32> -> i32 + %w1 = vector.extract %words[1] : vector<4xi32> -> i32 + %w2 = vector.extract %words[2] : vector<4xi32> -> i32 + %w3 = vector.extract %words[3] : vector<4xi32> -> i32 + %wv = vector.from_elements %w0, %w0, %w0, %w0, %w1, %w1, %w1, %w1, %w2, %w2, %w2, %w2, %w3, %w3, %w3, %w3 : vector<16xi32> + %b0 = scalar.constant 0 : i32 + %b8 = scalar.constant 8 : i32 + %b16 = scalar.constant 16 : i32 + %b24 = scalar.constant 24 : i32 + %shift = vector.from_elements %b0, %b8, %b16, %b24, %b0, %b8, %b16, %b24, %b0, %b8, %b16, %b24, %b0, %b8, %b16, %b24 : vector<16xi32> + %c255 = scalar.constant 255 : i32 + %mask = vector.splat %c255 : vector<16xi32> + %ws = vector.shrui %wv, %shift : vector<16xi32> + %bytes = vector.andi %ws, %mask : vector<16xi32> + func.return %bytes : vector<16xi32> +} + +// TQ2_0 (format 35): 66 bytes = qs[64], d (f16); value 128 a + 32 l + m is ((qs[32 a + m] >> 2 l) & 3) - 1 +// (dequantize_row_tq2_0). Lane l16 owns values 16 l16..+16: a = l16 / 8, l = (l16 % 8) / 2, bytes 32 a + 16 (l16 % 2)..+16. +func.def inline @ggml_kquant_tq2_lane_parts(%weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c2 = index.constant 2 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %block_bytes = index.constant 66 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %hv = buffer.view %weight[%block_base] : buffer -> view<33xf16> + %iv = buffer.view %weight[%block_base] : buffer -> view<33xi16> + %d_f16 = view.load %hv[%c32] : view<33xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %a = index.div %l, %c8 : index + %l8 = index.rem %l, %c8 : index + %ll = index.div %l8, %c2 : index + %h = index.rem %l, %c2 : index + %a16 = index.mul %a, %c16 : index + %h8 = index.mul %h, %c8 : index + %at0 = index.add %a16, %h8 : index + %at = index.assume %at0 [range(%at0, 0, 24), mul(%at0, 8)] : index + %w16 = vector.load %iv[%at] : view<33xi16> -> vector<8xi16> + %words = vector.bitcast %w16 : vector<8xi16> to vector<4xi32> + %bytes = func.call @ggml_kquant_run16_bytes(%words) : (vector<4xi32>) -> (vector<16xi32>) + %ll_i32 = index.cast %ll : index to i32 + %sh = scalar.muli %ll_i32, %c2_i32 : i32 + %shv = vector.splat %sh : vector<16xi32> + %three = vector.splat %c3_i32 : vector<16xi32> + %one = vector.splat %c1_i32 : vector<16xi32> + %bs = vector.shrui %bytes, %shv : vector<16xi32> + %q = vector.andi %bs, %three : vector<16xi32> + %t = vector.subi %q, %one : vector<16xi32> + %tf = vector.sitofp %t : vector<16xi32> to vector<16xf32> + %p = index.mul %l, %c16 : index + func.return %tf, %d, %p : vector<16xf32>, f32, index +} + +// TQ1_0 (format 34): 54 bytes = qs[48], qh[4], d (f16); five base-3 trits per byte (four in qh), trit n of +// byte b = ((b 3^n mod 256) 3) >> 8, minus 1 (dequantize_row_tq1_0). Lane l16 owns values 16 l16..+16: +// l16 < 10: trit l16 / 2 of qs[16 (l16 % 2)..+16]; 10..14: trit l16 - 10 of qs[32..47]; 15: values 240..255, +// element 4 n + j = trit n of qh[j]. +func.def inline @ggml_kquant_tq1_lane_parts(%weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c2 = index.constant 2 : index + %c8 = index.constant 8 : index + %c10 = index.constant 10 : index + %c15 = index.constant 15 : index + %c16 = index.constant 16 : index + %c24 = index.constant 24 : index + %c26 = index.constant 26 : index + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c255_i32 = scalar.constant 255 : i32 + %block_bytes = index.constant 54 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %hv = buffer.view %weight[%block_base] : buffer -> view<27xf16> + %iv = buffer.view %weight[%block_base] : buffer -> view<27xi16> + %d_f16 = view.load %hv[%c26] : view<27xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %low = index.cmp ult, %l, %c10 : index + %h = index.rem %l, %c2 : index + %h8 = index.mul %h, %c8 : index + %at0 = scf.select %low, %h8, %c16 : index + %at = index.assume %at0 [range(%at0, 0, 16), mul(%at0, 8)] : index + %w16 = vector.load %iv[%at] : view<27xi16> -> vector<8xi16> + %words = vector.bitcast %w16 : vector<8xi16> to vector<4xi32> + %qs_bytes = func.call @ggml_kquant_run16_bytes(%words) : (vector<4xi32>) -> (vector<16xi32>) + %qh16 = vector.load %iv[%c24] : view<27xi16> -> vector<2xi16> + %qh32 = vector.bitcast %qh16 : vector<2xi16> to vector<1xi32> + %qh = vector.extract %qh32[0] : vector<1xi32> -> i32 + %qh_words = vector.splat %qh : vector<4xi32> + %qh_bytes = func.call @ggml_kquant_run16_bytes(%qh_words) : (vector<4xi32>) -> (vector<16xi32>) + %n_low = index.div %l, %c2 : index + %l_mid = scf.select %low, %c10, %l : index + %n_mid = index.sub %l_mid, %c10 : index + %n = scf.select %low, %n_low, %n_mid : index + %n_i32 = index.cast %n : index to i32 + %p1 = scalar.constant 1 : i32 + %p3 = scalar.constant 3 : i32 + %p9 = scalar.constant 9 : i32 + %p27 = scalar.constant 27 : i32 + %p81 = scalar.constant 81 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %is1 = scalar.cmpi eq, %n_i32, %c1_i32 : i32 + %is2 = scalar.cmpi eq, %n_i32, %c2_i32 : i32 + %is3 = scalar.cmpi eq, %n_i32, %c3_i32 : i32 + %is4 = scalar.cmpi eq, %n_i32, %c4_i32 : i32 + %s1 = scf.select %is1, %p3, %p1 : i32 + %s2 = scf.select %is2, %p9, %s1 : i32 + %s3 = scf.select %is3, %p27, %s2 : i32 + %pow = scf.select %is4, %p81, %s3 : i32 + %pow_qs = vector.splat %pow : vector<16xi32> + %pow_qh = vector.from_elements %p1, %p1, %p1, %p1, %p3, %p3, %p3, %p3, %p9, %p9, %p9, %p9, %p27, %p27, %p27, %p27 : vector<16xi32> + %last = index.cmp eq, %l, %c15 : index + %b = scf.select %last, %qh_bytes, %qs_bytes : vector<16xi32> + %powv = scf.select %last, %pow_qh, %pow_qs : vector<16xi32> + %mask = vector.splat %c255_i32 : vector<16xi32> + %three = vector.splat %c3_i32 : vector<16xi32> + %eight = vector.splat %c8_i32 : vector<16xi32> + %one = vector.splat %c1_i32 : vector<16xi32> + %bp = vector.muli %b, %powv : vector<16xi32> + %q = vector.andi %bp, %mask : vector<16xi32> + %q3 = vector.muli %q, %three : vector<16xi32> + %xi = vector.shrui %q3, %eight : vector<16xi32> + %t = vector.subi %xi, %one : vector<16xi32> + %tf = vector.sitofp %t : vector<16xi32> to vector<16xf32> + %p = index.mul %l, %c16 : index + func.return %tf, %d, %p : vector<16xf32>, f32, index +} + +// MXFP4 (format 39): eight 17-byte blocks per 256 values, each e (E8M0) then qs[16]; value j of a block is the +// low nibble of qs[j] and value j + 16 the high nibble. Lane l16 owns block l16 / 2, nibbles (l16 % 2): 16 +// consecutive values. E2M1 code c (kvalues_mxfp4 = twice the FP4 value): m = c & 7 is 0, 1 for m < 2 and +// ((2 + (m & 1)) << (m >> 1)) >> 1 otherwise, negated when c >= 8; the scale is 2^(e - 128) (GGML_E8M0_TO_FP32_HALF). +func.def inline @ggml_kquant_mxfp4_lane_parts(%weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c16 = index.constant 16 : index + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c23_i32 = scalar.constant 23 : i32 + %sub_bytes = index.constant 17 : offset + %one_byte = index.constant 1 : offset + %block_bytes = index.constant 136 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %sub = index.div %l, %c2 : index + %h = index.rem %l, %c2 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %sub_add = index.scale %sub, %sub_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %sub_base = index.add %block_base, %sub_add : offset + %qs_base = index.add %sub_base, %one_byte : offset + %ev = buffer.view %weight[%sub_base] : buffer -> view<1xi8> + %qv = buffer.view %weight[%qs_base] : buffer -> view<16xi8> + %e_i8 = view.load %ev[%c0] : view<1xi8> -> i8 + %e = scalar.extui %e_i8 : i8 to i32 + %q8 = vector.load %qv[%c0] : view<16xi8> -> vector<16xi8> + %words = vector.bitcast %q8 : vector<16xi8> to vector<4xi32> + %bytes = func.call @ggml_kquant_run16_bytes(%words) : (vector<4xi32>) -> (vector<16xi32>) + %h_i32 = index.cast %h : index to i32 + %nsh = scalar.muli %h_i32, %c4_i32 : i32 + %nshv = vector.splat %nsh : vector<16xi32> + %fifteen = vector.splat %c15_i32 : vector<16xi32> + %seven = vector.splat %c7_i32 : vector<16xi32> + %onev = vector.splat %c1_i32 : vector<16xi32> + %twov = vector.splat %c2_i32 : vector<16xi32> + %threev = vector.splat %c3_i32 : vector<16xi32> + %c = vector.shrui %bytes, %nshv : vector<16xi32> + %code = vector.andi %c, %fifteen : vector<16xi32> + %m = vector.andi %code, %seven : vector<16xi32> + %m_odd = vector.andi %m, %onev : vector<16xi32> + %m_exp = vector.shrui %m, %onev : vector<16xi32> + %base = vector.addi %twov, %m_odd : vector<16xi32> + %scaled = vector.shli %base, %m_exp : vector<16xi32> + %mag_any = vector.shrui %scaled, %onev : vector<16xi32> + %m7 = vector.addi %m, %seven : vector<16xi32> + %nonzero_any = vector.shrui %m7, %threev : vector<16xi32> + %mag = vector.muli %mag_any, %nonzero_any : vector<16xi32> + %sign = vector.shrui %code, %threev : vector<16xi32> + %sign2 = vector.muli %sign, %twov : vector<16xi32> + %factor = vector.subi %onev, %sign2 : vector<16xi32> + %val = vector.muli %mag, %factor : vector<16xi32> + %vf = vector.sitofp %val : vector<16xi32> to vector<16xf32> + %is_small = scalar.cmpi ult, %e, %c2_i32 : i32 + %e_m1 = scalar.subi %e, %c1_i32 : i32 + %normal_bits = scalar.shli %e_m1, %c23_i32 : i32 + %sub_unit = scalar.constant 2097152 : i32 + %sub_bits = scalar.shli %sub_unit, %e : i32 + %bits = scf.select %is_small, %sub_bits, %normal_bits : i32 + %bits_v = vector.from_elements %bits : vector<1xi32> + %scale_v = vector.bitcast %bits_v : vector<1xi32> to vector<1xf32> + %scale = vector.extract %scale_v[0] : vector<1xf32> -> f32 + %p = index.mul %l, %c16 : index + func.return %vf, %scale, %p : vector<16xf32>, f32, index +} + +// 16-consecutive-value lanes with one scale: packed ternary 90, TQ1_0 34, TQ2_0 35, MXFP4 39. +func.def inline @ggml_kquant_run16_lane_parts(%format: index, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %f34 = index.constant 34 : index + %f35 = index.constant 35 : index + %f39 = index.constant 39 : index + %is34 = index.cmp eq, %format, %f34 : index + %is35 = index.cmp eq, %format, %f35 : index + %is39 = index.cmp eq, %format, %f39 : index + %v, %s, %p = scf.if %is35 -> (vector<16xf32>, f32, index) { + %a, %b, %c = func.call @ggml_kquant_tq2_lane_parts(%weight, %row_base, %block, %lane16) : (buffer, offset, index, index) -> (vector<16xf32>, f32, index) + scf.yield %a, %b, %c : vector<16xf32>, f32, index + } else { + %v34, %s34, %p34 = scf.if %is34 -> (vector<16xf32>, f32, index) { + %a, %b, %c = func.call @ggml_kquant_tq1_lane_parts(%weight, %row_base, %block, %lane16) : (buffer, offset, index, index) -> (vector<16xf32>, f32, index) + scf.yield %a, %b, %c : vector<16xf32>, f32, index + } else { + %v39, %s39, %p39 = scf.if %is39 -> (vector<16xf32>, f32, index) { + %a, %b, %c = func.call @ggml_kquant_mxfp4_lane_parts(%weight, %row_base, %block, %lane16) : (buffer, offset, index, index) -> (vector<16xf32>, f32, index) + scf.yield %a, %b, %c : vector<16xf32>, f32, index + } else { + %a, %b, %c = func.call @ggml_kquant_t2_lane_parts(%weight, %row_base, %block, %lane16) : (buffer, offset, index, index) -> (vector<16xf32>, f32, index) + scf.yield %a, %b, %c : vector<16xf32>, f32, index + } + scf.yield %v39, %s39, %p39 : vector<16xf32>, f32, index + } + scf.yield %v34, %s34, %p34 : vector<16xf32>, f32, index + } + func.return %v, %s, %p : vector<16xf32>, f32, index +} + +func.def inline @ggml_kquant_run16_lane_dot(%format: index, %weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %zero_scalar = scalar.constant 0.0 : f32 + %x_block_bytes = index.constant 1024 : offset + %v, %scale, %p = func.call @ggml_kquant_run16_lane_parts(%format, %weight, %row_base, %block, %lane16) : (index, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %x = vector.load %xv[%p] : view<256xf32> -> vector<16xf32> + %vx = vector.mulf %v, %x : vector<16xf32> + %sum = vector.reduce %vx, %zero_scalar : vector<16xf32>, f32 + %result = scalar.mulf %scale, %sum : f32 + func.return %result : f32 +} + +func.def inline @ggml_kquant_run16_lane_weights(%format: index, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c12 = index.constant 12 : index + %v, %scale, %p0 = func.call @ggml_kquant_run16_lane_parts(%format, %weight, %row_base, %block, %lane16) : (index, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %sv = vector.splat %scale : vector<16xf32> + %w = vector.mulf %v, %sv : vector<16xf32> + %p1 = index.add %p0, %c4 : index + %p2 = index.add %p0, %c8 : index + %p3 = index.add %p0, %c12 : index + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +// Bytes per 256 values of a weight format. +func.def pure inline @ggml_kquant_block_bytes(%format: index) -> (offset) { + %f5 = index.constant 5 : index + %f6 = index.constant 6 : index + %f23 = index.constant 23 : index + %f80 = index.constant 80 : index + %b144 = index.constant 144 : offset + %b176 = index.constant 176 : offset + %b210 = index.constant 210 : offset + %b136 = index.constant 136 : offset + %b272 = index.constant 272 : offset + %f11 = index.constant 11 : index + %b110 = index.constant 110 : offset + %is11 = index.cmp eq, %format, %f11 : index + %is5 = index.cmp eq, %format, %f5 : index + %is6 = index.cmp eq, %format, %f6 : index + %is23 = index.cmp eq, %format, %f23 : index + %is80 = index.cmp eq, %format, %f80 : index + %b0 = scf.select %is5, %b176, %b144 : offset + %b1 = scf.select %is6, %b210, %b0 : offset + %b2 = scf.select %is23, %b136, %b1 : offset + %b3 = scf.select %is80, %b272, %b2 : offset + %f21 = index.constant 21 : index + %is21 = index.cmp eq, %format, %f21 : index + %is110 = scalar.ori %is11, %is21 : i1 + %bytes0 = scf.select %is110, %b110, %b3 : offset + %f12 = index.constant 12 : index + %b84 = index.constant 84 : offset + %is12 = index.cmp eq, %format, %f12 : index + %bytes1 = scf.select %is12, %b84, %bytes0 : offset + %f24 = index.constant 24 : index + %f25 = index.constant 25 : index + %b66 = index.constant 66 : offset + %b74 = index.constant 74 : offset + %is24 = index.cmp eq, %format, %f24 : index + %is25 = index.cmp eq, %format, %f25 : index + %bytes2 = scf.select %is24, %b66, %bytes1 : offset + %bytes3 = scf.select %is25, %b74, %bytes2 : offset + %f28 = index.constant 28 : index + %b98 = index.constant 98 : offset + %is28 = index.cmp eq, %format, %f28 : index + %bytes4 = scf.select %is28, %b98, %bytes3 : offset + %f22 = index.constant 22 : index + %b82 = index.constant 82 : offset + %is22 = index.cmp eq, %format, %f22 : index + %bytes5 = scf.select %is22, %b82, %bytes4 : offset + %f26 = index.constant 26 : index + %f27 = index.constant 27 : index + %b50 = index.constant 50 : offset + %b56 = index.constant 56 : offset + %is26 = index.cmp eq, %format, %f26 : index + %is27 = index.cmp eq, %format, %f27 : index + %bytes6 = scf.select %is26, %b50, %bytes5 : offset + %bytes7 = scf.select %is27, %b56, %bytes6 : offset + %f90 = index.constant 90 : index + %b68 = index.constant 68 : offset + %is90 = index.cmp eq, %format, %f90 : index + %bytes8 = scf.select %is90, %b68, %bytes7 : offset + %f34 = index.constant 34 : index + %f35 = index.constant 35 : index + %f39 = index.constant 39 : index + %b54 = index.constant 54 : offset + %is34 = index.cmp eq, %format, %f34 : index + %is35 = index.cmp eq, %format, %f35 : index + %is39 = index.cmp eq, %format, %f39 : index + %bytes9 = scf.select %is34, %b54, %bytes8 : offset + %bytes10 = scf.select %is35, %b66, %bytes9 : offset + %bytes11 = scf.select %is39, %b136, %bytes10 : offset + %bytes = func.call pure @ggml_kquant_g128_block_bytes(%format, %bytes11) : (index, offset) -> (offset) + func.return %bytes : offset +} + +// The lane's partial dot product of 256-value block %block of weight row %row with the input. +// Formats: Q2_K 12, IQ1_S 26, IQ1_M 27, packed ternary 90, TQ1_0 34, TQ2_0 35, MXFP4 39, IQ2_S 22, IQ2_XXS 24, IQ2_XS 25, IQ3_XXS 28, Q3_K 11, IQ3_S 21, Q4_K 4, Q5_K 5, Q6_K 6, IQ4_NL 20, IQ4_XS 23, Q8_0 80. +func.def inline @ggml_kquant_lane_dot(%format: index, %kv_lane: f32, %grid: buffer, %weight: buffer, %input: buffer, %row: index, %blocks: index, %block: index, %lane16: index) -> (f32) { + %f4 = index.constant 4 : index + %f5 = index.constant 5 : index + %f6 = index.constant 6 : index + %f20 = index.constant 20 : index + %f23 = index.constant 23 : index + %block_bytes = func.call pure @ggml_kquant_block_bytes(%format) : (index) -> (offset) + %row_bytes = index.scale %blocks, %block_bytes : index, offset -> offset + %row_base = index.scale %row, %row_bytes : index, offset -> offset + %is4 = index.cmp eq, %format, %f4 : index + %is5 = index.cmp eq, %format, %f5 : index + %is6 = index.cmp eq, %format, %f6 : index + %is20 = index.cmp eq, %format, %f20 : index + %is23 = index.cmp eq, %format, %f23 : index + %r = scf.if %is5 -> (f32) { + %v = func.call @ggml_kquant_q5k_lane_dot(%weight, %input, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (f32) + scf.yield %v : f32 + } else { + %r4 = scf.if %is4 -> (f32) { + %v = func.call @ggml_kquant_q4k_lane_dot(%weight, %input, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (f32) + scf.yield %v : f32 + } else { + %r23 = scf.if %is23 -> (f32) { + %v = func.call @ggml_kquant_iq4xs_lane_dot(%kv_lane, %weight, %input, %row_base, %block, %lane16) : (f32, buffer, buffer, offset, index, index) -> (f32) + scf.yield %v : f32 + } else { + %r6 = scf.if %is6 -> (f32) { + %v = func.call @ggml_kquant_q6k_lane_dot(%weight, %input, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (f32) + scf.yield %v : f32 + } else { + %r20 = scf.if %is20 -> (f32) { + %v = func.call @ggml_kquant_iq4nl_lane_dot(%kv_lane, %weight, %input, %row_base, %block, %lane16) : (f32, buffer, buffer, offset, index, index) -> (f32) + scf.yield %v : f32 + } else { + %f11 = index.constant 11 : index + %is11 = index.cmp eq, %format, %f11 : index + %r11 = scf.if %is11 -> (f32) { + %v = func.call @ggml_kquant_q3k_lane_dot(%weight, %input, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (f32) + scf.yield %v : f32 + } else { + %f21 = index.constant 21 : index + %is21 = index.cmp eq, %format, %f21 : index + %r21 = scf.if %is21 -> (f32) { + %v = func.call @ggml_kquant_iq3s_lane_dot(%grid, %weight, %input, %row_base, %block, %lane16) : (buffer, buffer, buffer, offset, index, index) -> (f32) + scf.yield %v : f32 + } else { + %f12 = index.constant 12 : index + %is12 = index.cmp eq, %format, %f12 : index + %f24 = index.constant 24 : index + %f25 = index.constant 25 : index + %is24 = index.cmp eq, %format, %f24 : index + %is25 = index.cmp eq, %format, %f25 : index + %r12 = scf.if %is12 -> (f32) { + %v = func.call @ggml_kquant_q2k_lane_dot(%weight, %input, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (f32) + scf.yield %v : f32 + } else { + %r24 = scf.if %is24 -> (f32) { + %v = func.call @ggml_kquant_iq2xxs_lane_dot(%grid, %weight, %input, %row_base, %block, %lane16) : (buffer, buffer, buffer, offset, index, index) -> (f32) + scf.yield %v : f32 + } else { + %r25 = scf.if %is25 -> (f32) { + %v = func.call @ggml_kquant_iq2xs_lane_dot(%grid, %weight, %input, %row_base, %block, %lane16) : (buffer, buffer, buffer, offset, index, index) -> (f32) + scf.yield %v : f32 + } else { + %f28 = index.constant 28 : index + %is28 = index.cmp eq, %format, %f28 : index + %r28 = scf.if %is28 -> (f32) { + %v = func.call @ggml_kquant_iq3xxs_lane_dot(%grid, %weight, %input, %row_base, %block, %lane16) : (buffer, buffer, buffer, offset, index, index) -> (f32) + scf.yield %v : f32 + } else { + %f22 = index.constant 22 : index + %is22 = index.cmp eq, %format, %f22 : index + %r22 = scf.if %is22 -> (f32) { + %v = func.call @ggml_kquant_iq2s_lane_dot(%grid, %weight, %input, %row_base, %block, %lane16) : (buffer, buffer, buffer, offset, index, index) -> (f32) + scf.yield %v : f32 + } else { + %f26 = index.constant 26 : index + %f27 = index.constant 27 : index + %is26 = index.cmp eq, %format, %f26 : index + %is27 = index.cmp eq, %format, %f27 : index + %r26 = scf.if %is26 -> (f32) { + %v = func.call @ggml_kquant_iq1s_lane_dot(%grid, %weight, %input, %row_base, %block, %lane16) : (buffer, buffer, buffer, offset, index, index) -> (f32) + scf.yield %v : f32 + } else { + %r27 = scf.if %is27 -> (f32) { + %v = func.call @ggml_kquant_iq1m_lane_dot(%grid, %weight, %input, %row_base, %block, %lane16) : (buffer, buffer, buffer, offset, index, index) -> (f32) + scf.yield %v : f32 + } else { + %f90 = index.constant 90 : index + %f34 = index.constant 34 : index + %f35 = index.constant 35 : index + %f39 = index.constant 39 : index + %is34 = index.cmp eq, %format, %f34 : index + %is35 = index.cmp eq, %format, %f35 : index + %is39 = index.cmp eq, %format, %f39 : index + %is90_only = index.cmp eq, %format, %f90 : index + %is_tq = scalar.ori %is34, %is35 : i1 + %is_tq_mx = scalar.ori %is_tq, %is39 : i1 + %is90 = scalar.ori %is90_only, %is_tq_mx : i1 + %r90 = scf.if %is90 -> (f32) { + %v = func.call @ggml_kquant_run16_lane_dot(%format, %weight, %input, %row_base, %block, %lane16) : (index, buffer, buffer, offset, index, index) -> (f32) + scf.yield %v : f32 + } else { + %v = func.call @ggml_kquant_g128_or_q8_0_lane_dot(%format, %weight, %input, %row_base, %block, %lane16) : (index, buffer, buffer, offset, index, index) -> (f32) + scf.yield %v : f32 + } + scf.yield %r90 : f32 + } + scf.yield %r27 : f32 + } + scf.yield %r26 : f32 + } + scf.yield %r22 : f32 + } + scf.yield %r28 : f32 + } + scf.yield %r25 : f32 + } + scf.yield %r24 : f32 + } + scf.yield %r12 : f32 + } + scf.yield %r21 : f32 + } + scf.yield %r11 : f32 + } + scf.yield %r20 : f32 + } + scf.yield %r6 : f32 + } + scf.yield %r23 : f32 + } + scf.yield %r4 : f32 + } + func.return %r : f32 +} + +func.def inline @ggml_kquant_silu_mul(%gate: f32, %up: f32) -> (f32) { + %one = scalar.constant 1.0 : f32 + %neg = scalar.negf %gate : f32 + %e = scalar.expf %neg : f32 + %den = scalar.addf %one, %e : f32 + %silu = scalar.divf %gate, %den : f32 + %r = scalar.mulf %silu, %up : f32 + func.return %r : f32 +} + +kernel.def target(@ggml_kquant_decode_gfx11_wave64) export("ggml_kquant_swiglu_decode_f32") @ggml_kquant_swiglu_decode_f32() { + %output_size = config.get @ggml.kquant_swiglu_decode.output_size : index + %c1 = index.constant 1 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c256 = index.constant 256 : index + %last = index.add %output_size, %c3 : index + %groups = index.div %last, %c4 : index + kernel.launch.config workgroups(%groups, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%input: buffer, %gate: buffer, %up: buffer, %output: buffer) { + %input_size = config.get @ggml.kquant_swiglu_decode.input_size : index + %output_size = config.get @ggml.kquant_swiglu_decode.output_size : index + %gate_format = config.get @ggml.kquant_swiglu_decode.gate_weight_format : index + %up_format = config.get @ggml.kquant_swiglu_decode.up_weight_format : index + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c256 = index.constant 256 : index + %zero_offset = index.constant 0 : offset + %zero = scalar.constant 0.0 : f32 + %table = func.call @ggml_kquant_iq4nl_table() : () -> (vector<16xi8>) + %blocks = index.div %input_size, %c256 : index + %wg = kernel.workgroup.id : index + %subgroup = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %wg_row = index.mul %wg, %c4 : index + %row = index.add %wg_row, %subgroup : index + %first_block = index.div %lane, %c16 : index + %lane16 = index.rem %lane, %c16 : index + // Lane l of every 16-lane group holds kvalues_iq4nl[l]; IQ4_XS codes read it across lanes. + %lane16_i32 = index.cast %lane16 : index to i32 + %lane16_i8 = scalar.trunci %lane16_i32 : i32 to i8 + %lane16_v = vector.splat %lane16_i8 : vector<4xi8> + %kv_v = vector.table.lookup %table[%lane16_v] : vector<16xi8>, vector<4xi8> -> vector<4xi8> + %kv_i8 = vector.extract %kv_v[0] : vector<4xi8> -> i8 + %kv_i32 = scalar.extsi %kv_i8 : i8 to i32 + %kv_lane = scalar.sitofp %kv_i32 : i32 to f32 + %g_grid_f21 = index.constant 21 : index + %g_grid_f24 = index.constant 24 : index + %g_grid_f25 = index.constant 25 : index + %g_grid_is21 = index.cmp eq, %gate_format, %g_grid_f21 : index + %g_grid_is24 = index.cmp eq, %gate_format, %g_grid_f24 : index + %g_grid_is25 = index.cmp eq, %gate_format, %g_grid_f25 : index + %g_grid_a = scalar.ori %g_grid_is21, %g_grid_is24 : i1 + %g_grid_f28 = index.constant 28 : index + %g_grid_is28 = index.cmp eq, %gate_format, %g_grid_f28 : index + %g_grid_b = scalar.ori %g_grid_a, %g_grid_is25 : i1 + %g_grid_f22 = index.constant 22 : index + %g_grid_is22 = index.cmp eq, %gate_format, %g_grid_f22 : index + %g_grid_c = scalar.ori %g_grid_b, %g_grid_is28 : i1 + %g_grid_d = scalar.ori %g_grid_c, %g_grid_is22 : i1 + %g_grid_f26 = index.constant 26 : index + %g_grid_f27 = index.constant 27 : index + %g_grid_is26 = index.cmp eq, %gate_format, %g_grid_f26 : index + %g_grid_is27 = index.cmp eq, %gate_format, %g_grid_f27 : index + %g_grid_e = scalar.ori %g_grid_d, %g_grid_is26 : i1 + %g_grid = scalar.ori %g_grid_e, %g_grid_is27 : i1 + %u_grid_f21 = index.constant 21 : index + %u_grid_f24 = index.constant 24 : index + %u_grid_f25 = index.constant 25 : index + %u_grid_is21 = index.cmp eq, %up_format, %u_grid_f21 : index + %u_grid_is24 = index.cmp eq, %up_format, %u_grid_f24 : index + %u_grid_is25 = index.cmp eq, %up_format, %u_grid_f25 : index + %u_grid_a = scalar.ori %u_grid_is21, %u_grid_is24 : i1 + %u_grid_f28 = index.constant 28 : index + %u_grid_is28 = index.cmp eq, %up_format, %u_grid_f28 : index + %u_grid_b = scalar.ori %u_grid_a, %u_grid_is25 : i1 + %u_grid_f22 = index.constant 22 : index + %u_grid_is22 = index.cmp eq, %up_format, %u_grid_f22 : index + %u_grid_c = scalar.ori %u_grid_b, %u_grid_is28 : i1 + %u_grid_d = scalar.ori %u_grid_c, %u_grid_is22 : i1 + %u_grid_f26 = index.constant 26 : index + %u_grid_f27 = index.constant 27 : index + %u_grid_is26 = index.cmp eq, %up_format, %u_grid_f26 : index + %u_grid_is27 = index.cmp eq, %up_format, %u_grid_f27 : index + %u_grid_e = scalar.ori %u_grid_d, %u_grid_is26 : i1 + %u_grid = scalar.ori %u_grid_e, %u_grid_is27 : i1 + %needs_grid = scalar.ori %g_grid, %u_grid : i1 + // Gate and up each get their own codebook buffer, so a pair whose two formats need different grids runs here. + // A pair sharing one codebook (same format, or IQ1_S with IQ1_M) fills it once: the up lanes read the gate grid. + %same_format = index.cmp eq, %gate_format, %up_format : index + %iq1_g26 = index.constant 26 : index + %iq1_g27 = index.constant 27 : index + %gate_iq1_s = index.cmp eq, %gate_format, %iq1_g26 : index + %gate_iq1_m = index.cmp eq, %gate_format, %iq1_g27 : index + %up_iq1_s = index.cmp eq, %up_format, %iq1_g26 : index + %up_iq1_m = index.cmp eq, %up_format, %iq1_g27 : index + %gate_iq1 = scalar.ori %gate_iq1_s, %gate_iq1_m : i1 + %up_iq1 = scalar.ori %up_iq1_s, %up_iq1_m : i1 + %both_iq1 = scalar.andi %gate_iq1, %up_iq1 : i1 + %same_grid0 = scalar.ori %same_format, %both_iq1 : i1 + %same_grid = scalar.andi %same_grid0, %g_grid : i1 + %true_u = scalar.constant true : i1 + %not_same_grid = scalar.xori %same_grid, %true_u : i1 + %u_own_grid = scalar.andi %u_grid, %not_same_grid : i1 + %grid_bytes = index.constant 4096 : offset + %grid = buffer.alloca align(16) %grid_bytes : buffer + %grid_u = buffer.alloca align(16) %grid_bytes : buffer + scf.if %needs_grid { + scf.if %g_grid { + func.call @ggml_kquant_grid_fill_for(%gate_format, %grid, %subgroup) : (index, buffer, index) -> () + } + scf.if %u_own_grid { + func.call @ggml_kquant_grid_fill_for(%up_format, %grid_u, %subgroup) : (index, buffer, index) -> () + } + kernel.barrier scope(workgroup) ordering(acq_rel) + } + %valid = index.cmp ult, %row, %output_size : index + scf.if %valid { + %g_part, %u_part = scf.for %block = [%first_block to %blocks step %c4](%g_acc = %zero : f32, %u_acc = %zero : f32) -> (f32, f32) { + %g = func.call @ggml_kquant_lane_dot(%gate_format, %kv_lane, %grid, %gate, %input, %row, %blocks, %block, %lane16) : (index, f32, buffer, buffer, buffer, index, index, index, index) -> (f32) + %u = scf.if %same_grid -> (f32) { + %u_s = func.call @ggml_kquant_lane_dot(%up_format, %kv_lane, %grid, %up, %input, %row, %blocks, %block, %lane16) : (index, f32, buffer, buffer, buffer, index, index, index, index) -> (f32) + scf.yield %u_s : f32 + } else { + %u_o = func.call @ggml_kquant_lane_dot(%up_format, %kv_lane, %grid_u, %up, %input, %row, %blocks, %block, %lane16) : (index, f32, buffer, buffer, buffer, index, index, index, index) -> (f32) + scf.yield %u_o : f32 + } + %g_next = scalar.addf %g_acc, %g : f32 + %u_next = scalar.addf %u_acc, %u : f32 + scf.yield %g_next, %u_next : f32, f32 + } + %g_sum = kernel.subgroup.reduce %g_part : f32 + %u_sum = kernel.subgroup.reduce %u_part : f32 + %leader = index.cmp eq, %lane, %c0 : index + scf.if %leader { + %value = func.call @ggml_kquant_silu_mul(%g_sum, %u_sum) : (f32, f32) -> (f32) + %r, %r_bound = index.assume %row, %output_size [lt(%row, %output_size)] : index, index + %ov = buffer.view %output[%zero_offset] : buffer -> view<[%r_bound]xf32> + view.store %value, %ov[%r] : f32, view<[%r_bound]xf32> + } + } + kernel.return +} + +// Plain projection for one token on the same formats: out[r] = sum_k w[r, k] * x[k] (+ addend[r] +// when ggml.kquant_mul_mat_decode.add is 1, the residual ADD that follows most projections). +kernel.def target(@ggml_kquant_decode_gfx11_wave64) export("ggml_kquant_mul_mat_decode_f32") @ggml_kquant_mul_mat_decode_f32() { + %output_size = config.get @ggml.kquant_mul_mat_decode.output_size : index + %c1 = index.constant 1 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c256 = index.constant 256 : index + %last = index.add %output_size, %c3 : index + %groups = index.div %last, %c4 : index + kernel.launch.config workgroups(%groups, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%input: buffer, %weight: buffer, %addend: buffer, %output: buffer) { + %input_size = config.get @ggml.kquant_mul_mat_decode.input_size : index + %output_size = config.get @ggml.kquant_mul_mat_decode.output_size : index + %format = config.get @ggml.kquant_mul_mat_decode.weight_format : index + %add = config.get @ggml.kquant_mul_mat_decode.add : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c256 = index.constant 256 : index + %zero_offset = index.constant 0 : offset + %zero = scalar.constant 0.0 : f32 + %zero4 = vector.constant 0.0 : vector<4xf32> + %table = func.call @ggml_kquant_iq4nl_table() : () -> (vector<16xi8>) + %blocks = index.div %input_size, %c256 : index + %wg = kernel.workgroup.id : index + %subgroup = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %wg_row = index.mul %wg, %c4 : index + %row = index.add %wg_row, %subgroup : index + %first_block = index.div %lane, %c16 : index + %lane16 = index.rem %lane, %c16 : index + %lane16_i32 = index.cast %lane16 : index to i32 + %lane16_i8 = scalar.trunci %lane16_i32 : i32 to i8 + %lane16_v = vector.splat %lane16_i8 : vector<4xi8> + %kv_v = vector.table.lookup %table[%lane16_v] : vector<16xi8>, vector<4xi8> -> vector<4xi8> + %kv_i8 = vector.extract %kv_v[0] : vector<4xi8> -> i8 + %kv_i32 = scalar.extsi %kv_i8 : i8 to i32 + %kv_lane = scalar.sitofp %kv_i32 : i32 to f32 + %needs_grid_f21 = index.constant 21 : index + %needs_grid_f24 = index.constant 24 : index + %needs_grid_f25 = index.constant 25 : index + %needs_grid_is21 = index.cmp eq, %format, %needs_grid_f21 : index + %needs_grid_is24 = index.cmp eq, %format, %needs_grid_f24 : index + %needs_grid_is25 = index.cmp eq, %format, %needs_grid_f25 : index + %needs_grid_a = scalar.ori %needs_grid_is21, %needs_grid_is24 : i1 + %needs_grid_f28 = index.constant 28 : index + %needs_grid_is28 = index.cmp eq, %format, %needs_grid_f28 : index + %needs_grid_b = scalar.ori %needs_grid_a, %needs_grid_is25 : i1 + %needs_grid_f22 = index.constant 22 : index + %needs_grid_is22 = index.cmp eq, %format, %needs_grid_f22 : index + %needs_grid_c = scalar.ori %needs_grid_b, %needs_grid_is28 : i1 + %needs_grid_d = scalar.ori %needs_grid_c, %needs_grid_is22 : i1 + %needs_grid_f26 = index.constant 26 : index + %needs_grid_f27 = index.constant 27 : index + %needs_grid_is26 = index.cmp eq, %format, %needs_grid_f26 : index + %needs_grid_is27 = index.cmp eq, %format, %needs_grid_f27 : index + %needs_grid_e = scalar.ori %needs_grid_d, %needs_grid_is26 : i1 + %needs_grid = scalar.ori %needs_grid_e, %needs_grid_is27 : i1 + %grid_bytes = index.constant 4096 : offset + %grid = buffer.alloca align(16) %grid_bytes : buffer + scf.if %needs_grid { + func.call @ggml_kquant_grid_fill_for(%format, %grid, %subgroup) : (index, buffer, index) -> () + kernel.barrier scope(workgroup) ordering(acq_rel) + } + %valid = index.cmp ult, %row, %output_size : index + scf.if %valid { + %part = scf.for %block = [%first_block to %blocks step %c4](%acc = %zero : f32) -> (f32) { + %v = func.call @ggml_kquant_lane_dot(%format, %kv_lane, %grid, %weight, %input, %row, %blocks, %block, %lane16) : (index, f32, buffer, buffer, buffer, index, index, index, index) -> (f32) + %next = scalar.addf %acc, %v : f32 + scf.yield %next : f32 + } + %sum = kernel.subgroup.reduce %part : f32 + %leader = index.cmp eq, %lane, %c0 : index + scf.if %leader { + %r, %r_bound = index.assume %row, %output_size [lt(%row, %output_size)] : index, index + %has_add = index.cmp eq, %add, %c1 : index + %value = scf.if %has_add -> (f32) { + %av = buffer.view %addend[%zero_offset] : buffer -> view<[%r_bound]xf32> + %a = view.load %av[%r] : view<[%r_bound]xf32> -> f32 + %s = scalar.addf %sum, %a : f32 + scf.yield %s : f32 + } else { + scf.yield %sum : f32 + } + %ov = buffer.view %output[%zero_offset] : buffer -> view<[%r_bound]xf32> + view.store %value, %ov[%r] : f32, view<[%r_bound]xf32> + } + } + kernel.return +} + + +// Multi-token decode (2..8 tokens, e.g. MTP / speculative verify batches): each lane dequantizes +// its 16 weights once per block and applies them to every token. +config.decl @ggml.kquant_decode.token_count : %value: index where [range(%value, 1, 8)] + +// Lane weights for multi-token decode: the lane's 16 dequantized weights and the positions (within +// the 256-value block) of their four runs of 4 values, so each token only needs its dot product. +func.def inline @ggml_kquant_q45k_lane_weights_finish(%dm: vector<2xf16>, %header: vector<4xi32>, %low: vector<2xi32>, %high: vector<2xi32>, %group: index, %packet: index) -> (vector<16xf32>, index, index, index, index) { + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c8 = index.constant 8 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %d_f16 = vector.extract %dm[0] : vector<2xf16> -> f16 + %dmin_f16 = vector.extract %dm[1] : vector<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %s0 = vector.extract %header[1] : vector<4xi32> -> i32 + %s1 = vector.extract %header[2] : vector<4xi32> -> i32 + %s2 = vector.extract %header[3] : vector<4xi32> -> i32 + %sub_lo = index.mul %group, %c2 : index + %sub_hi = index.add %sub_lo, %c1 : index + %sc_lo_i, %m_lo_i = func.call @ggml_kquant_scale_min(%s0, %s1, %s2, %sub_lo) : (i32, i32, i32, index) -> (i32, i32) + %sc_hi_i, %m_hi_i = func.call @ggml_kquant_scale_min(%s0, %s1, %s2, %sub_hi) : (i32, i32, i32, index) -> (i32, i32) + %q_lo = func.call @ggml_kquant_codes_u8x8(%low) : (vector<2xi32>) -> (vector<8xf32>) + %q_hi = func.call @ggml_kquant_codes_u8x8(%high) : (vector<2xi32>) -> (vector<8xf32>) + + %c4w = index.constant 4 : index + %c36w = index.constant 36 : index + %sub_base = index.mul %group, %c64 : index + %pos = index.mul %packet, %c8 : index + %p0 = index.add %sub_base, %pos : index + %p1 = index.add %p0, %c4w : index + %p2 = index.add %p0, %c32 : index + %p3 = index.add %p0, %c36w : index + %sc_lo_f = scalar.uitofp %sc_lo_i : i32 to f32 + %sc_hi_f = scalar.uitofp %sc_hi_i : i32 to f32 + %m_lo_f = scalar.uitofp %m_lo_i : i32 to f32 + %m_hi_f = scalar.uitofp %m_hi_i : i32 to f32 + %d_lo = scalar.mulf %d, %sc_lo_f : f32 + %d_hi = scalar.mulf %d, %sc_hi_f : f32 + %mn_lo = scalar.mulf %dmin, %m_lo_f : f32 + %mn_hi = scalar.mulf %dmin, %m_hi_f : f32 + %nmn_lo = scalar.negf %mn_lo : f32 + %nmn_hi = scalar.negf %mn_hi : f32 + %d_lo_v = vector.splat %d_lo : vector<8xf32> + %d_hi_v = vector.splat %d_hi : vector<8xf32> + %m_lo_v = vector.splat %nmn_lo : vector<8xf32> + %m_hi_v = vector.splat %nmn_hi : vector<8xf32> + %w_lo = vector.fmaf %q_lo, %d_lo_v, %m_lo_v : vector<8xf32> + %w_hi = vector.fmaf %q_hi, %d_hi_v, %m_hi_v : vector<8xf32> + %w = vector.concat<0> %w_lo, %w_hi : vector<8xf32>, vector<8xf32> -> vector<16xf32> + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +func.def inline @ggml_kquant_q4k_lane_weights(%weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %nibble_mask = vector.constant 252645135 : vector<2xi32> + %four = vector.constant 4 : vector<2xi32> + %block_bytes = index.constant 144 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %group = index.div %l, %c4 : index + %packet = index.rem %l, %c4 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %wv = buffer.view %weight[%block_base] : buffer -> view<36xi32> + %hv = buffer.view %weight[%block_base] : buffer -> view<2xf16> + %dm = vector.load %hv[%c0] : view<2xf16> -> vector<2xf16> + %header = vector.load %wv[%c0] : view<36xi32> -> vector<4xi32> + %group_words = index.mul %group, %c8 : index + %packet_words = index.mul %packet, %c2 : index + %code_rel = index.add %group_words, %packet_words : index + %code_word = index.add %c4, %code_rel : index + %codes = vector.load %wv[%code_word] : view<36xi32> -> vector<2xi32> + %low = vector.andi %codes, %nibble_mask : vector<2xi32> + %codes_shr = vector.shrui %codes, %four : vector<2xi32> + %high = vector.andi %codes_shr, %nibble_mask : vector<2xi32> + %w, %p0, %p1, %p2, %p3 = func.call @ggml_kquant_q45k_lane_weights_finish(%dm, %header, %low, %high, %group, %packet) : (vector<2xf16>, vector<4xi32>, vector<2xi32>, vector<2xi32>, index, index) -> (vector<16xf32>, index, index, index, index) + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +func.def inline @ggml_kquant_q5k_lane_weights(%weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c12 = index.constant 12 : index + %nibble_mask = vector.constant 252645135 : vector<2xi32> + %bit_mask = vector.constant 16843009 : vector<2xi32> + %four = vector.constant 4 : vector<2xi32> + %block_bytes = index.constant 176 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %group = index.div %l, %c4 : index + %packet = index.rem %l, %c4 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %wv = buffer.view %weight[%block_base] : buffer -> view<44xi32> + %hv = buffer.view %weight[%block_base] : buffer -> view<2xf16> + %dm = vector.load %hv[%c0] : view<2xf16> -> vector<2xf16> + %header = vector.load %wv[%c0] : view<44xi32> -> vector<4xi32> + %group_words = index.mul %group, %c8 : index + %packet_words = index.mul %packet, %c2 : index + %code_rel = index.add %group_words, %packet_words : index + %code_word = index.add %c12, %code_rel : index + %codes = vector.load %wv[%code_word] : view<44xi32> -> vector<2xi32> + %qh_word = index.add %c4, %packet_words : index + %qh = vector.load %wv[%qh_word] : view<44xi32> -> vector<2xi32> + %sub_lo = index.mul %group, %c2 : index + %sub_hi = index.add %sub_lo, %c1 : index + %shift_lo_i = index.cast %sub_lo : index to i32 + %shift_hi_i = index.cast %sub_hi : index to i32 + %shift_lo = vector.splat %shift_lo_i : vector<2xi32> + %shift_hi = vector.splat %shift_hi_i : vector<2xi32> + %qh_lo0 = vector.shrui %qh, %shift_lo : vector<2xi32> + %qh_hi0 = vector.shrui %qh, %shift_hi : vector<2xi32> + %qh_lo1 = vector.andi %qh_lo0, %bit_mask : vector<2xi32> + %qh_hi1 = vector.andi %qh_hi0, %bit_mask : vector<2xi32> + %qh_lo = vector.shli %qh_lo1, %four : vector<2xi32> + %qh_hi = vector.shli %qh_hi1, %four : vector<2xi32> + %low4 = vector.andi %codes, %nibble_mask : vector<2xi32> + %codes_shr = vector.shrui %codes, %four : vector<2xi32> + %high4 = vector.andi %codes_shr, %nibble_mask : vector<2xi32> + %low = vector.ori %low4, %qh_lo : vector<2xi32> + %high = vector.ori %high4, %qh_hi : vector<2xi32> + %w, %p0, %p1, %p2, %p3 = func.call @ggml_kquant_q45k_lane_weights_finish(%dm, %header, %low, %high, %group, %packet) : (vector<2xf16>, vector<4xi32>, vector<2xi32>, vector<2xi32>, index, index) -> (vector<16xf32>, index, index, index, index) + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +func.def inline @ggml_kquant_iq4_codes_weights(%kv_lane: f32, %codes: vector<2xi32>) -> (vector<16xf32>) { + %c15_i32 = scalar.constant 15 : i32 + %c16w = scalar.constant 16 : i32 + %w0 = vector.extract %codes[0] : vector<2xi32> -> i32 + %w1 = vector.extract %codes[1] : vector<2xi32> -> i32 + %sh_lo0 = scalar.constant 0 : i32 + %c_lo00 = scalar.shrui %w0, %sh_lo0 : i32 + %c_lo0 = scalar.andi %c_lo00, %c15_i32 : i32 + %k_lo0, %k_lo0_ok = kernel.subgroup.shuffle %kv_lane, %c_lo0, %c16w : f32, i32, i32 + %sh_lo1 = scalar.constant 8 : i32 + %c_lo10 = scalar.shrui %w0, %sh_lo1 : i32 + %c_lo1 = scalar.andi %c_lo10, %c15_i32 : i32 + %k_lo1, %k_lo1_ok = kernel.subgroup.shuffle %kv_lane, %c_lo1, %c16w : f32, i32, i32 + %sh_lo2 = scalar.constant 16 : i32 + %c_lo20 = scalar.shrui %w0, %sh_lo2 : i32 + %c_lo2 = scalar.andi %c_lo20, %c15_i32 : i32 + %k_lo2, %k_lo2_ok = kernel.subgroup.shuffle %kv_lane, %c_lo2, %c16w : f32, i32, i32 + %sh_lo3 = scalar.constant 24 : i32 + %c_lo30 = scalar.shrui %w0, %sh_lo3 : i32 + %c_lo3 = scalar.andi %c_lo30, %c15_i32 : i32 + %k_lo3, %k_lo3_ok = kernel.subgroup.shuffle %kv_lane, %c_lo3, %c16w : f32, i32, i32 + %sh_lo4 = scalar.constant 0 : i32 + %c_lo40 = scalar.shrui %w1, %sh_lo4 : i32 + %c_lo4 = scalar.andi %c_lo40, %c15_i32 : i32 + %k_lo4, %k_lo4_ok = kernel.subgroup.shuffle %kv_lane, %c_lo4, %c16w : f32, i32, i32 + %sh_lo5 = scalar.constant 8 : i32 + %c_lo50 = scalar.shrui %w1, %sh_lo5 : i32 + %c_lo5 = scalar.andi %c_lo50, %c15_i32 : i32 + %k_lo5, %k_lo5_ok = kernel.subgroup.shuffle %kv_lane, %c_lo5, %c16w : f32, i32, i32 + %sh_lo6 = scalar.constant 16 : i32 + %c_lo60 = scalar.shrui %w1, %sh_lo6 : i32 + %c_lo6 = scalar.andi %c_lo60, %c15_i32 : i32 + %k_lo6, %k_lo6_ok = kernel.subgroup.shuffle %kv_lane, %c_lo6, %c16w : f32, i32, i32 + %sh_lo7 = scalar.constant 24 : i32 + %c_lo70 = scalar.shrui %w1, %sh_lo7 : i32 + %c_lo7 = scalar.andi %c_lo70, %c15_i32 : i32 + %k_lo7, %k_lo7_ok = kernel.subgroup.shuffle %kv_lane, %c_lo7, %c16w : f32, i32, i32 + %sh_hi0 = scalar.constant 4 : i32 + %c_hi00 = scalar.shrui %w0, %sh_hi0 : i32 + %c_hi0 = scalar.andi %c_hi00, %c15_i32 : i32 + %k_hi0, %k_hi0_ok = kernel.subgroup.shuffle %kv_lane, %c_hi0, %c16w : f32, i32, i32 + %sh_hi1 = scalar.constant 12 : i32 + %c_hi10 = scalar.shrui %w0, %sh_hi1 : i32 + %c_hi1 = scalar.andi %c_hi10, %c15_i32 : i32 + %k_hi1, %k_hi1_ok = kernel.subgroup.shuffle %kv_lane, %c_hi1, %c16w : f32, i32, i32 + %sh_hi2 = scalar.constant 20 : i32 + %c_hi20 = scalar.shrui %w0, %sh_hi2 : i32 + %c_hi2 = scalar.andi %c_hi20, %c15_i32 : i32 + %k_hi2, %k_hi2_ok = kernel.subgroup.shuffle %kv_lane, %c_hi2, %c16w : f32, i32, i32 + %sh_hi3 = scalar.constant 28 : i32 + %c_hi30 = scalar.shrui %w0, %sh_hi3 : i32 + %c_hi3 = scalar.andi %c_hi30, %c15_i32 : i32 + %k_hi3, %k_hi3_ok = kernel.subgroup.shuffle %kv_lane, %c_hi3, %c16w : f32, i32, i32 + %sh_hi4 = scalar.constant 4 : i32 + %c_hi40 = scalar.shrui %w1, %sh_hi4 : i32 + %c_hi4 = scalar.andi %c_hi40, %c15_i32 : i32 + %k_hi4, %k_hi4_ok = kernel.subgroup.shuffle %kv_lane, %c_hi4, %c16w : f32, i32, i32 + %sh_hi5 = scalar.constant 12 : i32 + %c_hi50 = scalar.shrui %w1, %sh_hi5 : i32 + %c_hi5 = scalar.andi %c_hi50, %c15_i32 : i32 + %k_hi5, %k_hi5_ok = kernel.subgroup.shuffle %kv_lane, %c_hi5, %c16w : f32, i32, i32 + %sh_hi6 = scalar.constant 20 : i32 + %c_hi60 = scalar.shrui %w1, %sh_hi6 : i32 + %c_hi6 = scalar.andi %c_hi60, %c15_i32 : i32 + %k_hi6, %k_hi6_ok = kernel.subgroup.shuffle %kv_lane, %c_hi6, %c16w : f32, i32, i32 + %sh_hi7 = scalar.constant 28 : i32 + %c_hi70 = scalar.shrui %w1, %sh_hi7 : i32 + %c_hi7 = scalar.andi %c_hi70, %c15_i32 : i32 + %k_hi7, %k_hi7_ok = kernel.subgroup.shuffle %kv_lane, %c_hi7, %c16w : f32, i32, i32 + %values = vector.from_elements %k_lo0, %k_lo1, %k_lo2, %k_lo3, %k_lo4, %k_lo5, %k_lo6, %k_lo7, %k_hi0, %k_hi1, %k_hi2, %k_hi3, %k_hi4, %k_hi5, %k_hi6, %k_hi7 : vector<16xf32> + func.return %values : vector<16xf32> +} + +func.def inline @ggml_kquant_iq4xs_lane_weights(%kv_lane: f32, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c16_i32 = scalar.constant 16 : i32 + %c32_i32 = scalar.constant 32 : i32 + %block_bytes = index.constant 136 : offset + %x_block_bytes = index.constant 1024 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %sub = index.div %l, %c2 : index + %half = index.rem %l, %c2 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %wv = buffer.view %weight[%block_base] : buffer -> view<34xi32> + %hv = buffer.view %weight[%block_base] : buffer -> view<2xf16> + %dh = vector.load %hv[%c0] : view<2xf16> -> vector<2xf16> + %d_f16 = vector.extract %dh[0] : vector<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %header = vector.load %wv[%c0] : view<34xi32> -> vector<2xi32> + %word0 = vector.extract %header[0] : vector<2xi32> -> i32 + %scales_l = vector.extract %header[1] : vector<2xi32> -> i32 + %sub_i32 = index.cast %sub : index to i32 + %low_shift = scalar.muli %sub_i32, %c4_i32 : i32 + %high_shift0 = scalar.muli %sub_i32, %c2_i32 : i32 + %high_shift = scalar.addi %high_shift0, %c16_i32 : i32 + %ls_low0 = scalar.shrui %scales_l, %low_shift : i32 + %ls_low = scalar.andi %ls_low0, %c15_i32 : i32 + %ls_high0 = scalar.shrui %word0, %high_shift : i32 + %ls_high1 = scalar.andi %ls_high0, %c3_i32 : i32 + %ls_high = scalar.shli %ls_high1, %c4_i32 : i32 + %ls = scalar.ori %ls_low, %ls_high : i32 + %ls_centered = scalar.subi %ls, %c32_i32 : i32 + %ls_f = scalar.sitofp %ls_centered : i32 to f32 + %dl = scalar.mulf %d, %ls_f : f32 + %sub_words = index.mul %sub, %c4 : index + %half_words = index.mul %half, %c2 : index + %code_rel = index.add %sub_words, %half_words : index + %code_word = index.add %c2, %code_rel : index + %codes = vector.load %wv[%code_word] : view<34xi32> -> vector<2xi32> + %c4w = index.constant 4 : index + %c20w = index.constant 20 : index + %blk_k = index.mul %sub, %c32 : index + %half_k = index.mul %half, %c8 : index + %p0 = index.add %blk_k, %half_k : index + %p1 = index.add %p0, %c4w : index + %p2 = index.add %p0, %c16 : index + %p3 = index.add %p0, %c20w : index + %kv = func.call @ggml_kquant_iq4_codes_weights(%kv_lane, %codes) : (f32, vector<2xi32>) -> (vector<16xf32>) + %scale_v = vector.splat %dl : vector<16xf32> + %w = vector.mulf %kv, %scale_v : vector<16xf32> + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +func.def inline @ggml_kquant_iq4nl_lane_weights(%kv_lane: f32, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %super_bytes = index.constant 144 : offset + %qblock_bytes = index.constant 18 : offset + %x_block_bytes = index.constant 1024 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %qb = index.div %l, %c2 : index + %half = index.rem %l, %c2 : index + %super_add = index.scale %block, %super_bytes : index, offset -> offset + %qb_add = index.scale %qb, %qblock_bytes : index, offset -> offset + %super_base = index.add %row_base, %super_add : offset + %qb_base = index.add %super_base, %qb_add : offset + %hv = buffer.view %weight[%qb_base] : buffer -> view<9xf16> + %iv = buffer.view %weight[%qb_base] : buffer -> view<9xi16> + %d_f16 = view.load %hv[%c0] : view<9xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %half_words = index.mul %half, %c4 : index + %code_at = index.add %c1, %half_words : index + %codes16 = vector.load %iv[%code_at] : view<9xi16> -> vector<4xi16> + %codes = vector.bitcast %codes16 : vector<4xi16> to vector<2xi32> + %c4w = index.constant 4 : index + %c20w = index.constant 20 : index + %blk_k = index.mul %qb, %c32 : index + %half_k = index.mul %half, %c8 : index + %p0 = index.add %blk_k, %half_k : index + %p1 = index.add %p0, %c4w : index + %p2 = index.add %p0, %c16 : index + %p3 = index.add %p0, %c20w : index + %kv = func.call @ggml_kquant_iq4_codes_weights(%kv_lane, %codes) : (f32, vector<2xi32>) -> (vector<16xf32>) + %scale_v = vector.splat %d : vector<16xf32> + %w = vector.mulf %kv, %scale_v : vector<16xf32> + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +func.def inline @ggml_kquant_q8_0_lane_weights(%weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %zero_scalar = scalar.constant 0.0 : f32 + %super_bytes = index.constant 272 : offset + %qblock_bytes = index.constant 34 : offset + %x_block_bytes = index.constant 1024 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %qb = index.div %l, %c2 : index + %half = index.rem %l, %c2 : index + %super_add = index.scale %block, %super_bytes : index, offset -> offset + %qb_add = index.scale %qb, %qblock_bytes : index, offset -> offset + %super_base = index.add %row_base, %super_add : offset + %qb_base = index.add %super_base, %qb_add : offset + %hv = buffer.view %weight[%qb_base] : buffer -> view<17xf16> + %iv = buffer.view %weight[%qb_base] : buffer -> view<17xi16> + %d_f16 = view.load %hv[%c0] : view<17xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %half_words = index.mul %half, %c8 : index + %code_at = index.add %c1, %half_words : index + %codes16 = vector.load %iv[%code_at] : view<17xi16> -> vector<8xi16> + %codes = vector.bitcast %codes16 : vector<8xi16> to vector<16xi8> + %q = vector.sitofp %codes : vector<16xi8> to vector<16xf32> + %c4w = index.constant 4 : index + %c12w = index.constant 12 : index + %qb_k = index.mul %qb, %c32 : index + %half_k = index.mul %half, %c16 : index + %p0 = index.add %qb_k, %half_k : index + %p1 = index.add %p0, %c4w : index + %p2 = index.add %p0, %c8 : index + %p3 = index.add %p0, %c12w : index + %d_v = vector.splat %d : vector<16xf32> + %w = vector.mulf %q, %d_v : vector<16xf32> + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +func.def inline @ggml_kquant_q6k_quarter_values(%ql: i32, %qh: i32, %nibble_shift: i32, %high_shift: i32) -> (vector<4xf32>) { + %c4_i32 = scalar.constant 4 : i32 + %mask4 = scalar.constant 252645135 : i32 + %mask2 = scalar.constant 50529027 : i32 + %c32 = vector.constant 32.0 : vector<4xf32> + %low0 = scalar.shrui %ql, %nibble_shift : i32 + %low = scalar.andi %low0, %mask4 : i32 + %high0 = scalar.shrui %qh, %high_shift : i32 + %high1 = scalar.andi %high0, %mask2 : i32 + %high = scalar.shli %high1, %c4_i32 : i32 + %q_word = scalar.ori %low, %high : i32 + %q_vec = vector.from_elements %q_word : vector<1xi32> + %q_bytes = vector.bitcast %q_vec : vector<1xi32> to vector<4xi8> + %q_u = vector.uitofp %q_bytes : vector<4xi8> to vector<4xf32> + %q = vector.subf %q_u, %c32 : vector<4xf32> + func.return %q : vector<4xf32> +} + +func.def inline @ggml_kquant_q6k_lane_weights(%weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c96 = index.constant 96 : index + %c104 = index.constant 104 : index + %c128 = index.constant 128 : index + %s0 = scalar.constant 0 : i32 + %s2 = scalar.constant 2 : i32 + %s4 = scalar.constant 4 : i32 + %s6 = scalar.constant 6 : i32 + %block_bytes = index.constant 210 : offset + %x_block_bytes = index.constant 1024 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %n = index.div %l, %c8 : index + %j = index.rem %l, %c8 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %iv = buffer.view %weight[%block_base] : buffer -> view<105xi16> + %hv = buffer.view %weight[%block_base] : buffer -> view<105xf16> + %n32 = index.mul %n, %c32 : index + %j2 = index.mul %j, %c2 : index + %ql_a_at = index.add %n32, %j2 : index + %ql_b_at = index.add %ql_a_at, %c16 : index + %n16 = index.mul %n, %c16 : index + %qh_rel = index.add %n16, %j2 : index + %qh_at = index.add %c64, %qh_rel : index + %n4 = index.mul %n, %c4 : index + %sc_at = index.add %c96, %n4 : index + %ql_a2 = vector.load %iv[%ql_a_at] : view<105xi16> -> vector<2xi16> + %ql_b2 = vector.load %iv[%ql_b_at] : view<105xi16> -> vector<2xi16> + %qh2 = vector.load %iv[%qh_at] : view<105xi16> -> vector<2xi16> + %sc4 = vector.load %iv[%sc_at] : view<105xi16> -> vector<4xi16> + %ql_a1 = vector.bitcast %ql_a2 : vector<2xi16> to vector<1xi32> + %ql_b1 = vector.bitcast %ql_b2 : vector<2xi16> to vector<1xi32> + %qh1 = vector.bitcast %qh2 : vector<2xi16> to vector<1xi32> + %ql_a = vector.extract %ql_a1[0] : vector<1xi32> -> i32 + %ql_b = vector.extract %ql_b1[0] : vector<1xi32> -> i32 + %qh = vector.extract %qh1[0] : vector<1xi32> -> i32 + %sc8 = vector.bitcast %sc4 : vector<4xi16> to vector<8xi8> + %d_f16 = view.load %hv[%c104] : view<105xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %second = index.cmp uge, %j, %c4 : index + %n128 = index.mul %n, %c128 : index + %j4 = index.mul %j, %c4 : index + %p0 = index.add %n128, %j4 : index + %p1 = index.add %p0, %c32 : index + %p2 = index.add %p0, %c64 : index + %p3 = index.add %p0, %c96 : index + %v0 = func.call @ggml_kquant_q6k_quarter_values(%ql_a, %qh, %s0, %s0) : (i32, i32, i32, i32) -> (vector<4xf32>) + %v1 = func.call @ggml_kquant_q6k_quarter_values(%ql_b, %qh, %s0, %s2) : (i32, i32, i32, i32) -> (vector<4xf32>) + %v2 = func.call @ggml_kquant_q6k_quarter_values(%ql_a, %qh, %s4, %s4) : (i32, i32, i32, i32) -> (vector<4xf32>) + %v3 = func.call @ggml_kquant_q6k_quarter_values(%ql_b, %qh, %s4, %s6) : (i32, i32, i32, i32) -> (vector<4xf32>) + %sc0_a = vector.extract %sc8[0] : vector<8xi8> -> i8 + %sc0_b = vector.extract %sc8[1] : vector<8xi8> -> i8 + %sc0_i8 = scf.select %second, %sc0_b, %sc0_a : i8 + %sc0_i32 = scalar.extsi %sc0_i8 : i8 to i32 + %sc0 = scalar.sitofp %sc0_i32 : i32 to f32 + %ds0 = scalar.mulf %d, %sc0 : f32 + %ds0_v = vector.splat %ds0 : vector<4xf32> + %w0 = vector.mulf %v0, %ds0_v : vector<4xf32> + %sc1_a = vector.extract %sc8[2] : vector<8xi8> -> i8 + %sc1_b = vector.extract %sc8[3] : vector<8xi8> -> i8 + %sc1_i8 = scf.select %second, %sc1_b, %sc1_a : i8 + %sc1_i32 = scalar.extsi %sc1_i8 : i8 to i32 + %sc1 = scalar.sitofp %sc1_i32 : i32 to f32 + %ds1 = scalar.mulf %d, %sc1 : f32 + %ds1_v = vector.splat %ds1 : vector<4xf32> + %w1 = vector.mulf %v1, %ds1_v : vector<4xf32> + %sc2_a = vector.extract %sc8[4] : vector<8xi8> -> i8 + %sc2_b = vector.extract %sc8[5] : vector<8xi8> -> i8 + %sc2_i8 = scf.select %second, %sc2_b, %sc2_a : i8 + %sc2_i32 = scalar.extsi %sc2_i8 : i8 to i32 + %sc2 = scalar.sitofp %sc2_i32 : i32 to f32 + %ds2 = scalar.mulf %d, %sc2 : f32 + %ds2_v = vector.splat %ds2 : vector<4xf32> + %w2 = vector.mulf %v2, %ds2_v : vector<4xf32> + %sc3_a = vector.extract %sc8[6] : vector<8xi8> -> i8 + %sc3_b = vector.extract %sc8[7] : vector<8xi8> -> i8 + %sc3_i8 = scf.select %second, %sc3_b, %sc3_a : i8 + %sc3_i32 = scalar.extsi %sc3_i8 : i8 to i32 + %sc3 = scalar.sitofp %sc3_i32 : i32 to f32 + %ds3 = scalar.mulf %d, %sc3 : f32 + %ds3_v = vector.splat %ds3 : vector<4xf32> + %w3 = vector.mulf %v3, %ds3_v : vector<4xf32> + %w01 = vector.concat<0> %w0, %w1 : vector<4xf32>, vector<4xf32> -> vector<8xf32> + %w23 = vector.concat<0> %w2, %w3 : vector<4xf32>, vector<4xf32> -> vector<8xf32> + %w = vector.concat<0> %w01, %w23 : vector<8xf32>, vector<8xf32> -> vector<16xf32> + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +func.def inline @ggml_kquant_q3k_lane_weights(%weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c48 = index.constant 48 : index + %c54 = index.constant 54 : index + %c128 = index.constant 128 : index + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c32_i32 = scalar.constant 32 : i32 + %zero_scalar = scalar.constant 0.0 : f32 + %mask2 = vector.constant 50529027 : vector<4xi32> + %mask1 = vector.constant 16843009 : vector<4xi32> + %two = vector.constant 2 : vector<4xi32> + %four_f = vector.constant 4.0 : vector<16xf32> + %block_bytes = index.constant 110 : offset + %x_block_bytes = index.constant 1024 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %n = index.div %l, %c8 : index + %jh = index.rem %l, %c8 : index + %j = index.div %jh, %c2 : index + %h = index.rem %jh, %c2 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %iv = buffer.view %weight[%block_base] : buffer -> view<55xi16> + %hv = buffer.view %weight[%block_base] : buffer -> view<55xf16> + %h8 = index.mul %h, %c8 : index + %n16 = index.mul %n, %c16 : index + %qs_rel = index.add %n16, %h8 : index + %qs_at = index.add %c16, %qs_rel : index + %qs16 = vector.load %iv[%qs_at] : view<55xi16> -> vector<8xi16> + %hm16 = vector.load %iv[%h8] : view<55xi16> -> vector<8xi16> + %sc16 = vector.load %iv[%c48] : view<55xi16> -> vector<6xi16> + %d_f16 = view.load %hv[%c54] : view<55xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %qs = vector.bitcast %qs16 : vector<8xi16> to vector<4xi32> + %hm = vector.bitcast %hm16 : vector<8xi16> to vector<4xi32> + %scw = vector.bitcast %sc16 : vector<6xi16> to vector<3xi32> + %j_i32 = index.cast %j : index to i32 + %n_i32 = index.cast %n : index to i32 + %q_shift_i = scalar.muli %j_i32, %c2_i32 : i32 + %n4 = scalar.muli %n_i32, %c4_i32 : i32 + %h_shift_i = scalar.addi %n4, %j_i32 : i32 + %q_shift = vector.splat %q_shift_i : vector<4xi32> + %h_shift = vector.splat %h_shift_i : vector<4xi32> + %q0 = vector.shrui %qs, %q_shift : vector<4xi32> + %q = vector.andi %q0, %mask2 : vector<4xi32> + %hb0 = vector.shrui %hm, %h_shift : vector<4xi32> + %hb1 = vector.andi %hb0, %mask1 : vector<4xi32> + %hb = vector.shli %hb1, %two : vector<4xi32> + %u = vector.ori %q, %hb : vector<4xi32> + %u8 = vector.bitcast %u : vector<4xi32> to vector<16xi8> + %uf = vector.uitofp %u8 : vector<16xi8> to vector<16xf32> + %vf = vector.subf %uf, %four_f : vector<16xf32> + // scale s = 8n + 2j + h: word w = s / 4 = 2n + j / 2, byte b = s % 4 = 2 (j % 2) + h. + %j2 = index.div %j, %c2 : index + %jodd = index.rem %j, %c2 : index + %n2 = index.mul %n, %c2 : index + %w = index.add %n2, %j2 : index + %jodd2 = index.mul %jodd, %c2 : index + %b = index.add %jodd2, %h : index + %w_odd = index.rem %w, %c2 : index + %w_hi = index.div %w, %c2 : index + %sw0 = vector.extract %scw[0] : vector<3xi32> -> i32 + %sw1 = vector.extract %scw[1] : vector<3xi32> -> i32 + %sw2 = vector.extract %scw[2] : vector<3xi32> -> i32 + %c1 = index.constant 1 : index + %odd = index.cmp eq, %w_odd, %c1 : index + %low_word = scf.select %odd, %sw1, %sw0 : i32 + %b_i32 = index.cast %b : index to i32 + %w_hi_i32 = index.cast %w_hi : index to i32 + %w_i32 = index.cast %w : index to i32 + %b8 = scalar.muli %b_i32, %c8_i32 : i32 + %wh4 = scalar.muli %w_hi_i32, %c4_i32 : i32 + %low_shift = scalar.addi %b8, %wh4 : i32 + %w2 = scalar.muli %w_i32, %c2_i32 : i32 + %high_shift = scalar.addi %b8, %w2 : i32 + %sl0 = scalar.shrui %low_word, %low_shift : i32 + %sl = scalar.andi %sl0, %c15_i32 : i32 + %sh0 = scalar.shrui %sw2, %high_shift : i32 + %sh1 = scalar.andi %sh0, %c3_i32 : i32 + %sh = scalar.shli %sh1, %c4_i32 : i32 + %sc0 = scalar.ori %sl, %sh : i32 + %sc = scalar.subi %sc0, %c32_i32 : i32 + %sc_f = scalar.sitofp %sc : i32 to f32 + %c12w = index.constant 12 : index + %n128 = index.mul %n, %c128 : index + %j32 = index.mul %j, %c32 : index + %h16 = index.mul %h, %c16 : index + %pa = index.add %n128, %j32 : index + %p0 = index.add %pa, %h16 : index + %p1 = index.add %p0, %c4 : index + %p2 = index.add %p0, %c8 : index + %p3 = index.add %p0, %c12w : index + %dl = scalar.mulf %d, %sc_f : f32 + %dl_v = vector.splat %dl : vector<16xf32> + %w_out = vector.mulf %vf, %dl_v : vector<16xf32> + func.return %w_out, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +func.def inline @ggml_kquant_iq3s_lane_weights(%grid_buf: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c33 = index.constant 33 : index + %c37 = index.constant 37 : index + %c53 = index.constant 53 : index + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c255_i32 = scalar.constant 255 : i32 + %spread = scalar.constant 2113665 : i32 + %lsb = scalar.constant 16843009 : i32 + %one4 = vector.constant 1.0 : vector<4xf32> + %two4 = vector.constant 2.0 : vector<4xf32> + %zero_scalar = scalar.constant 0.0 : f32 + %zero_offset = index.constant 0 : offset + %block_bytes = index.constant 110 : offset + %x_block_bytes = index.constant 1024 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %ib = index.div %l, %c2 : index + %half = index.rem %l, %c2 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %iv = buffer.view %weight[%block_base] : buffer -> view<55xi16> + %hv = buffer.view %weight[%block_base] : buffer -> view<55xf16> + %d_f16 = view.load %hv[%c0] : view<55xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %ib4 = index.mul %ib, %c4 : index + %half2 = index.mul %half, %c2 : index + %qs_rel = index.add %ib4, %half2 : index + %qs_at = index.add %c1, %qs_rel : index + %qs2 = vector.load %iv[%qs_at] : view<55xi16> -> vector<2xi16> + %qs1 = vector.bitcast %qs2 : vector<2xi16> to vector<1xi32> + %qs = vector.extract %qs1[0] : vector<1xi32> -> i32 + %ib_2 = index.div %ib, %c2 : index + %qh_at = index.add %c33, %ib_2 : index + %qh16 = view.load %iv[%qh_at] : view<55xi16> -> i16 + %qh16u = scalar.extui %qh16 : i16 to i32 + %ib_odd = index.rem %ib, %c2 : index + %ib_odd_i32 = index.cast %ib_odd : index to i32 + %qh_shift = scalar.muli %ib_odd_i32, %c8_i32 : i32 + %qh0 = scalar.shrui %qh16u, %qh_shift : i32 + %qh = scalar.andi %qh0, %c255_i32 : i32 + %ib2 = index.mul %ib, %c2 : index + %sg_rel = index.add %ib2, %half : index + %sg_at = index.add %c37, %sg_rel : index + %sg16 = view.load %iv[%sg_at] : view<55xi16> -> i16 + %sg = scalar.extui %sg16 : i16 to i32 + %ib_4 = index.div %ib, %c4 : index + %sc_at = index.add %c53, %ib_4 : index + %sc16 = view.load %iv[%sc_at] : view<55xi16> -> i16 + %sc16u = scalar.extui %sc16 : i16 to i32 + %ib_i32 = index.cast %ib : index to i32 + %sc_shift = scalar.muli %ib_i32, %c4_i32 : i32 + %sc_shift_w = scalar.andi %sc_shift, %c15_i32 : i32 + %sc0 = scalar.shrui %sc16u, %sc_shift_w : i32 + %sc = scalar.andi %sc0, %c15_i32 : i32 + %sc2 = scalar.muli %sc, %c2_i32 : i32 + %sc21 = scalar.addi %sc2, %c1_i32 : i32 + %sc_f = scalar.sitofp %sc21 : i32 to f32 + %db = scalar.mulf %d, %sc_f : f32 + %half_i32 = index.cast %half : index to i32 + %qh_base = scalar.muli %half_i32, %c4_i32 : i32 + %grid = buffer.view %grid_buf[%zero_offset] : buffer -> view<512xi32> + %ib32 = index.mul %ib, %c32 : index + %half16 = index.mul %half, %c16 : index + %xk = index.add %ib32, %half16 : index + %e0_sh = scalar.constant 0 : i32 + %e0_b0 = scalar.shrui %qs, %e0_sh : i32 + %e0_b = scalar.andi %e0_b0, %c255_i32 : i32 + %e0_hs0 = scalar.constant 0 : i32 + %e0_hs = scalar.addi %qh_base, %e0_hs0 : i32 + %e0_h0 = scalar.shrui %qh, %e0_hs : i32 + %e0_h1 = scalar.andi %e0_h0, %c1_i32 : i32 + %e0_h = scalar.shli %e0_h1, %c8_i32 : i32 + %e0_i32 = scalar.ori %e0_b, %e0_h : i32 + %e0_idx0 = index.cast %e0_i32 : i32 to index + %e0_idx = index.assume %e0_idx0 [range(%e0_idx0, 0, 511)] : index + %e0_g = view.load %grid[%e0_idx] : view<512xi32> -> i32 + %e0_gv = vector.from_elements %e0_g : vector<1xi32> + %e0_gb = vector.bitcast %e0_gv : vector<1xi32> to vector<4xi8> + %e0_gf = vector.uitofp %e0_gb : vector<4xi8> to vector<4xf32> + %e0_ss = scalar.constant 0 : i32 + %e0_s0 = scalar.shrui %sg, %e0_ss : i32 + %e0_s1 = scalar.andi %e0_s0, %c15_i32 : i32 + %e0_s2 = scalar.muli %e0_s1, %spread : i32 + %e0_s3 = scalar.andi %e0_s2, %lsb : i32 + %e0_sv = vector.from_elements %e0_s3 : vector<1xi32> + %e0_sb = vector.bitcast %e0_sv : vector<1xi32> to vector<4xi8> + %e0_sf = vector.uitofp %e0_sb : vector<4xi8> to vector<4xf32> + %e0_s2f = vector.mulf %e0_sf, %two4 : vector<4xf32> + %e0_sign = vector.subf %one4, %e0_s2f : vector<4xf32> + %e0_w = vector.mulf %e0_gf, %e0_sign : vector<4xf32> + %e1_sh = scalar.constant 8 : i32 + %e1_b0 = scalar.shrui %qs, %e1_sh : i32 + %e1_b = scalar.andi %e1_b0, %c255_i32 : i32 + %e1_hs0 = scalar.constant 1 : i32 + %e1_hs = scalar.addi %qh_base, %e1_hs0 : i32 + %e1_h0 = scalar.shrui %qh, %e1_hs : i32 + %e1_h1 = scalar.andi %e1_h0, %c1_i32 : i32 + %e1_h = scalar.shli %e1_h1, %c8_i32 : i32 + %e1_i32 = scalar.ori %e1_b, %e1_h : i32 + %e1_idx0 = index.cast %e1_i32 : i32 to index + %e1_idx = index.assume %e1_idx0 [range(%e1_idx0, 0, 511)] : index + %e1_g = view.load %grid[%e1_idx] : view<512xi32> -> i32 + %e1_gv = vector.from_elements %e1_g : vector<1xi32> + %e1_gb = vector.bitcast %e1_gv : vector<1xi32> to vector<4xi8> + %e1_gf = vector.uitofp %e1_gb : vector<4xi8> to vector<4xf32> + %e1_ss = scalar.constant 4 : i32 + %e1_s0 = scalar.shrui %sg, %e1_ss : i32 + %e1_s1 = scalar.andi %e1_s0, %c15_i32 : i32 + %e1_s2 = scalar.muli %e1_s1, %spread : i32 + %e1_s3 = scalar.andi %e1_s2, %lsb : i32 + %e1_sv = vector.from_elements %e1_s3 : vector<1xi32> + %e1_sb = vector.bitcast %e1_sv : vector<1xi32> to vector<4xi8> + %e1_sf = vector.uitofp %e1_sb : vector<4xi8> to vector<4xf32> + %e1_s2f = vector.mulf %e1_sf, %two4 : vector<4xf32> + %e1_sign = vector.subf %one4, %e1_s2f : vector<4xf32> + %e1_w = vector.mulf %e1_gf, %e1_sign : vector<4xf32> + %e2_sh = scalar.constant 16 : i32 + %e2_b0 = scalar.shrui %qs, %e2_sh : i32 + %e2_b = scalar.andi %e2_b0, %c255_i32 : i32 + %e2_hs0 = scalar.constant 2 : i32 + %e2_hs = scalar.addi %qh_base, %e2_hs0 : i32 + %e2_h0 = scalar.shrui %qh, %e2_hs : i32 + %e2_h1 = scalar.andi %e2_h0, %c1_i32 : i32 + %e2_h = scalar.shli %e2_h1, %c8_i32 : i32 + %e2_i32 = scalar.ori %e2_b, %e2_h : i32 + %e2_idx0 = index.cast %e2_i32 : i32 to index + %e2_idx = index.assume %e2_idx0 [range(%e2_idx0, 0, 511)] : index + %e2_g = view.load %grid[%e2_idx] : view<512xi32> -> i32 + %e2_gv = vector.from_elements %e2_g : vector<1xi32> + %e2_gb = vector.bitcast %e2_gv : vector<1xi32> to vector<4xi8> + %e2_gf = vector.uitofp %e2_gb : vector<4xi8> to vector<4xf32> + %e2_ss = scalar.constant 8 : i32 + %e2_s0 = scalar.shrui %sg, %e2_ss : i32 + %e2_s1 = scalar.andi %e2_s0, %c15_i32 : i32 + %e2_s2 = scalar.muli %e2_s1, %spread : i32 + %e2_s3 = scalar.andi %e2_s2, %lsb : i32 + %e2_sv = vector.from_elements %e2_s3 : vector<1xi32> + %e2_sb = vector.bitcast %e2_sv : vector<1xi32> to vector<4xi8> + %e2_sf = vector.uitofp %e2_sb : vector<4xi8> to vector<4xf32> + %e2_s2f = vector.mulf %e2_sf, %two4 : vector<4xf32> + %e2_sign = vector.subf %one4, %e2_s2f : vector<4xf32> + %e2_w = vector.mulf %e2_gf, %e2_sign : vector<4xf32> + %e3_sh = scalar.constant 24 : i32 + %e3_b0 = scalar.shrui %qs, %e3_sh : i32 + %e3_b = scalar.andi %e3_b0, %c255_i32 : i32 + %e3_hs0 = scalar.constant 3 : i32 + %e3_hs = scalar.addi %qh_base, %e3_hs0 : i32 + %e3_h0 = scalar.shrui %qh, %e3_hs : i32 + %e3_h1 = scalar.andi %e3_h0, %c1_i32 : i32 + %e3_h = scalar.shli %e3_h1, %c8_i32 : i32 + %e3_i32 = scalar.ori %e3_b, %e3_h : i32 + %e3_idx0 = index.cast %e3_i32 : i32 to index + %e3_idx = index.assume %e3_idx0 [range(%e3_idx0, 0, 511)] : index + %e3_g = view.load %grid[%e3_idx] : view<512xi32> -> i32 + %e3_gv = vector.from_elements %e3_g : vector<1xi32> + %e3_gb = vector.bitcast %e3_gv : vector<1xi32> to vector<4xi8> + %e3_gf = vector.uitofp %e3_gb : vector<4xi8> to vector<4xf32> + %e3_ss = scalar.constant 12 : i32 + %e3_s0 = scalar.shrui %sg, %e3_ss : i32 + %e3_s1 = scalar.andi %e3_s0, %c15_i32 : i32 + %e3_s2 = scalar.muli %e3_s1, %spread : i32 + %e3_s3 = scalar.andi %e3_s2, %lsb : i32 + %e3_sv = vector.from_elements %e3_s3 : vector<1xi32> + %e3_sb = vector.bitcast %e3_sv : vector<1xi32> to vector<4xi8> + %e3_sf = vector.uitofp %e3_sb : vector<4xi8> to vector<4xf32> + %e3_s2f = vector.mulf %e3_sf, %two4 : vector<4xf32> + %e3_sign = vector.subf %one4, %e3_s2f : vector<4xf32> + %e3_w = vector.mulf %e3_gf, %e3_sign : vector<4xf32> + %c4w = index.constant 4 : index + %c8w = index.constant 8 : index + %c12w = index.constant 12 : index + %p0 = index.add %xk, %c0 : index + %p1 = index.add %xk, %c4w : index + %p2 = index.add %xk, %c8w : index + %p3 = index.add %xk, %c12w : index + %w01 = vector.concat<0> %e0_w, %e1_w : vector<4xf32>, vector<4xf32> -> vector<8xf32> + %w23 = vector.concat<0> %e2_w, %e3_w : vector<4xf32>, vector<4xf32> -> vector<8xf32> + %wu = vector.concat<0> %w01, %w23 : vector<8xf32>, vector<8xf32> -> vector<16xf32> + %db_v = vector.splat %db : vector<16xf32> + %w = vector.mulf %wu, %db_v : vector<16xf32> + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} + +// The lane's dequantized weights of block %block of row %row (multi-token decode). +func.def inline @ggml_kquant_lane_weights(%format: index, %kv_lane: f32, %grid: buffer, %weight: buffer, %row: index, %blocks: index, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %f4 = index.constant 4 : index + %f5 = index.constant 5 : index + %f6 = index.constant 6 : index + %f20 = index.constant 20 : index + %f23 = index.constant 23 : index + %block_bytes = func.call pure @ggml_kquant_block_bytes(%format) : (index) -> (offset) + %row_bytes = index.scale %blocks, %block_bytes : index, offset -> offset + %row_base = index.scale %row, %row_bytes : index, offset -> offset + %is4 = index.cmp eq, %format, %f4 : index + %is5 = index.cmp eq, %format, %f5 : index + %is6 = index.cmp eq, %format, %f6 : index + %is20 = index.cmp eq, %format, %f20 : index + %is23 = index.cmp eq, %format, %f23 : index + %r, %r_0, %r_1, %r_2, %r_3 = scf.if %is5 -> (vector<16xf32>, index, index, index, index) { + %v, %v0, %v1, %v2, %v3 = func.call @ggml_kquant_q5k_lane_weights(%weight, %row_base, %block, %lane16) : (buffer, offset, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %v, %v0, %v1, %v2, %v3 : vector<16xf32>, index, index, index, index + } else { + %r4, %r4_0, %r4_1, %r4_2, %r4_3 = scf.if %is4 -> (vector<16xf32>, index, index, index, index) { + %v, %v0, %v1, %v2, %v3 = func.call @ggml_kquant_q4k_lane_weights(%weight, %row_base, %block, %lane16) : (buffer, offset, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %v, %v0, %v1, %v2, %v3 : vector<16xf32>, index, index, index, index + } else { + %r23, %r23_0, %r23_1, %r23_2, %r23_3 = scf.if %is23 -> (vector<16xf32>, index, index, index, index) { + %v, %v0, %v1, %v2, %v3 = func.call @ggml_kquant_iq4xs_lane_weights(%kv_lane, %weight, %row_base, %block, %lane16) : (f32, buffer, offset, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %v, %v0, %v1, %v2, %v3 : vector<16xf32>, index, index, index, index + } else { + %r6, %r6_0, %r6_1, %r6_2, %r6_3 = scf.if %is6 -> (vector<16xf32>, index, index, index, index) { + %v, %v0, %v1, %v2, %v3 = func.call @ggml_kquant_q6k_lane_weights(%weight, %row_base, %block, %lane16) : (buffer, offset, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %v, %v0, %v1, %v2, %v3 : vector<16xf32>, index, index, index, index + } else { + %r20, %r20_0, %r20_1, %r20_2, %r20_3 = scf.if %is20 -> (vector<16xf32>, index, index, index, index) { + %v, %v0, %v1, %v2, %v3 = func.call @ggml_kquant_iq4nl_lane_weights(%kv_lane, %weight, %row_base, %block, %lane16) : (f32, buffer, offset, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %v, %v0, %v1, %v2, %v3 : vector<16xf32>, index, index, index, index + } else { + %f11 = index.constant 11 : index + %is11 = index.cmp eq, %format, %f11 : index + %r11, %r11_0, %r11_1, %r11_2, %r11_3 = scf.if %is11 -> (vector<16xf32>, index, index, index, index) { + %v, %v0, %v1, %v2, %v3 = func.call @ggml_kquant_q3k_lane_weights(%weight, %row_base, %block, %lane16) : (buffer, offset, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %v, %v0, %v1, %v2, %v3 : vector<16xf32>, index, index, index, index + } else { + %f21 = index.constant 21 : index + %is21 = index.cmp eq, %format, %f21 : index + %r21, %r21_0, %r21_1, %r21_2, %r21_3 = scf.if %is21 -> (vector<16xf32>, index, index, index, index) { + %v, %v0, %v1, %v2, %v3 = func.call @ggml_kquant_iq3s_lane_weights(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %v, %v0, %v1, %v2, %v3 : vector<16xf32>, index, index, index, index + } else { + %f12 = index.constant 12 : index + %is12 = index.cmp eq, %format, %f12 : index + %f24 = index.constant 24 : index + %f25 = index.constant 25 : index + %is24 = index.cmp eq, %format, %f24 : index + %is25 = index.cmp eq, %format, %f25 : index + %r12, %r12_0, %r12_1, %r12_2, %r12_3 = scf.if %is12 -> (vector<16xf32>, index, index, index, index) { + %v, %v0, %v1, %v2, %v3 = func.call @ggml_kquant_q2k_lane_weights(%weight, %row_base, %block, %lane16) : (buffer, offset, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %v, %v0, %v1, %v2, %v3 : vector<16xf32>, index, index, index, index + } else { + %r24, %r24_0, %r24_1, %r24_2, %r24_3 = scf.if %is24 -> (vector<16xf32>, index, index, index, index) { + %v, %v0, %v1, %v2, %v3 = func.call @ggml_kquant_iq2xxs_lane_weights(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %v, %v0, %v1, %v2, %v3 : vector<16xf32>, index, index, index, index + } else { + %r25, %r25_0, %r25_1, %r25_2, %r25_3 = scf.if %is25 -> (vector<16xf32>, index, index, index, index) { + %v, %v0, %v1, %v2, %v3 = func.call @ggml_kquant_iq2xs_lane_weights(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %v, %v0, %v1, %v2, %v3 : vector<16xf32>, index, index, index, index + } else { + %f28 = index.constant 28 : index + %is28 = index.cmp eq, %format, %f28 : index + %r28, %r28_0, %r28_1, %r28_2, %r28_3 = scf.if %is28 -> (vector<16xf32>, index, index, index, index) { + %v, %v0, %v1, %v2, %v3 = func.call @ggml_kquant_iq3xxs_lane_weights(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %v, %v0, %v1, %v2, %v3 : vector<16xf32>, index, index, index, index + } else { + %f22 = index.constant 22 : index + %is22 = index.cmp eq, %format, %f22 : index + %r22, %r22_0, %r22_1, %r22_2, %r22_3 = scf.if %is22 -> (vector<16xf32>, index, index, index, index) { + %v, %v0, %v1, %v2, %v3 = func.call @ggml_kquant_iq2s_lane_weights(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %v, %v0, %v1, %v2, %v3 : vector<16xf32>, index, index, index, index + } else { + %f26 = index.constant 26 : index + %f27 = index.constant 27 : index + %is26 = index.cmp eq, %format, %f26 : index + %is27 = index.cmp eq, %format, %f27 : index + %r26, %r26_0, %r26_1, %r26_2, %r26_3 = scf.if %is26 -> (vector<16xf32>, index, index, index, index) { + %v, %v0, %v1, %v2, %v3 = func.call @ggml_kquant_iq1s_lane_weights(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %v, %v0, %v1, %v2, %v3 : vector<16xf32>, index, index, index, index + } else { + %r27, %r27_0, %r27_1, %r27_2, %r27_3 = scf.if %is27 -> (vector<16xf32>, index, index, index, index) { + %v, %v0, %v1, %v2, %v3 = func.call @ggml_kquant_iq1m_lane_weights(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %v, %v0, %v1, %v2, %v3 : vector<16xf32>, index, index, index, index + } else { + %f90 = index.constant 90 : index + %f34 = index.constant 34 : index + %f35 = index.constant 35 : index + %f39 = index.constant 39 : index + %is34 = index.cmp eq, %format, %f34 : index + %is35 = index.cmp eq, %format, %f35 : index + %is39 = index.cmp eq, %format, %f39 : index + %is90_only = index.cmp eq, %format, %f90 : index + %is_tq = scalar.ori %is34, %is35 : i1 + %is_tq_mx = scalar.ori %is_tq, %is39 : i1 + %is90 = scalar.ori %is90_only, %is_tq_mx : i1 + %r90, %r90_0, %r90_1, %r90_2, %r90_3 = scf.if %is90 -> (vector<16xf32>, index, index, index, index) { + %v, %v0, %v1, %v2, %v3 = func.call @ggml_kquant_run16_lane_weights(%format, %weight, %row_base, %block, %lane16) : (index, buffer, offset, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %v, %v0, %v1, %v2, %v3 : vector<16xf32>, index, index, index, index + } else { + %v, %v0, %v1, %v2, %v3 = func.call @ggml_kquant_g128_or_q8_0_lane_weights(%format, %weight, %row_base, %block, %lane16) : (index, buffer, offset, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %v, %v0, %v1, %v2, %v3 : vector<16xf32>, index, index, index, index + } + scf.yield %r90, %r90_0, %r90_1, %r90_2, %r90_3 : vector<16xf32>, index, index, index, index + } + scf.yield %r27, %r27_0, %r27_1, %r27_2, %r27_3 : vector<16xf32>, index, index, index, index + } + scf.yield %r26, %r26_0, %r26_1, %r26_2, %r26_3 : vector<16xf32>, index, index, index, index + } + scf.yield %r22, %r22_0, %r22_1, %r22_2, %r22_3 : vector<16xf32>, index, index, index, index + } + scf.yield %r28, %r28_0, %r28_1, %r28_2, %r28_3 : vector<16xf32>, index, index, index, index + } + scf.yield %r25, %r25_0, %r25_1, %r25_2, %r25_3 : vector<16xf32>, index, index, index, index + } + scf.yield %r24, %r24_0, %r24_1, %r24_2, %r24_3 : vector<16xf32>, index, index, index, index + } + scf.yield %r12, %r12_0, %r12_1, %r12_2, %r12_3 : vector<16xf32>, index, index, index, index + } + scf.yield %r21, %r21_0, %r21_1, %r21_2, %r21_3 : vector<16xf32>, index, index, index, index + } + scf.yield %r11, %r11_0, %r11_1, %r11_2, %r11_3 : vector<16xf32>, index, index, index, index + } + scf.yield %r20, %r20_0, %r20_1, %r20_2, %r20_3 : vector<16xf32>, index, index, index, index + } + scf.yield %r6, %r6_0, %r6_1, %r6_2, %r6_3 : vector<16xf32>, index, index, index, index + } + scf.yield %r23, %r23_0, %r23_1, %r23_2, %r23_3 : vector<16xf32>, index, index, index, index + } + scf.yield %r4, %r4_0, %r4_1, %r4_2, %r4_3 : vector<16xf32>, index, index, index, index + } + func.return %r, %r_0, %r_1, %r_2, %r_3 : vector<16xf32>, index, index, index, index +} + +// Accumulate lane weights times token %t's input into %acc (four partial sums; reduced once after +// the block loop). %shape (folded per format): 0 = four runs of 4 (Q6_K), 1 = two runs of 8 at %p0 +// and %p2 (Q4_K, Q5_K, IQ4_XS, IQ4_NL), 2 = 16 contiguous values at %p0 (Q2_K, Q3_K, IQ3_S, Q8_0). +func.def inline @ggml_kquant_token_fma(%shape: index, %w: vector<16xf32>, %p0: index, %p1: index, %p2: index, %p3: index, %input: buffer, %t: index, %blocks: index, %block: index, %acc: vector<4xf32>) -> (vector<4xf32>) { + %s1 = index.constant 1 : index + %s2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c12 = index.constant 12 : index + %x_block_bytes = index.constant 1024 : offset + %row_blocks = index.mul %t, %blocks : index + %x_block = index.add %row_blocks, %block : index + %x_base = index.scale %x_block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %is_run16 = index.cmp eq, %shape, %s2 : index + %is_run8 = index.cmp eq, %shape, %s1 : index + %xa, %xb, %xc, %xd = scf.if %is_run16 -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + %q0 = index.assume %p0 [range(%p0, 0, 240)] : index + %x = vector.load %xv[%q0] : view<256xf32> -> vector<16xf32> + %x0 = vector.slice %x[0] : vector<16xf32> -> vector<4xf32> + %x1 = vector.slice %x[4] : vector<16xf32> -> vector<4xf32> + %x2 = vector.slice %x[8] : vector<16xf32> -> vector<4xf32> + %x3 = vector.slice %x[12] : vector<16xf32> -> vector<4xf32> + scf.yield %x0, %x1, %x2, %x3 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } else { + %y0, %y1, %y2, %y3 = scf.if %is_run8 -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + %q0 = index.assume %p0 [range(%p0, 0, 248)] : index + %q2 = index.assume %p2 [range(%p2, 0, 248)] : index + %xl = vector.load %xv[%q0] : view<256xf32> -> vector<8xf32> + %xh = vector.load %xv[%q2] : view<256xf32> -> vector<8xf32> + %x0 = vector.slice %xl[0] : vector<8xf32> -> vector<4xf32> + %x1 = vector.slice %xl[4] : vector<8xf32> -> vector<4xf32> + %x2 = vector.slice %xh[0] : vector<8xf32> -> vector<4xf32> + %x3 = vector.slice %xh[4] : vector<8xf32> -> vector<4xf32> + scf.yield %x0, %x1, %x2, %x3 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } else { + %q0 = index.assume %p0 [range(%p0, 0, 252)] : index + %q1 = index.assume %p1 [range(%p1, 0, 252)] : index + %q2 = index.assume %p2 [range(%p2, 0, 252)] : index + %q3 = index.assume %p3 [range(%p3, 0, 252)] : index + %x0 = vector.load %xv[%q0] : view<256xf32> -> vector<4xf32> + %x1 = vector.load %xv[%q1] : view<256xf32> -> vector<4xf32> + %x2 = vector.load %xv[%q2] : view<256xf32> -> vector<4xf32> + %x3 = vector.load %xv[%q3] : view<256xf32> -> vector<4xf32> + scf.yield %x0, %x1, %x2, %x3 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + scf.yield %y0, %y1, %y2, %y3 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + %w0 = vector.slice %w[0] : vector<16xf32> -> vector<4xf32> + %w1 = vector.slice %w[4] : vector<16xf32> -> vector<4xf32> + %w2 = vector.slice %w[8] : vector<16xf32> -> vector<4xf32> + %w3 = vector.slice %w[12] : vector<16xf32> -> vector<4xf32> + %a0 = vector.fmaf %w0, %xa, %acc : vector<4xf32> + %a1 = vector.fmaf %w1, %xb, %a0 : vector<4xf32> + %a2 = vector.fmaf %w2, %xc, %a1 : vector<4xf32> + %a3 = vector.fmaf %w3, %xd, %a2 : vector<4xf32> + func.return %a3 : vector<4xf32> +} + +// Run shape of a format's lane weights, for ggml_kquant_token_fma. +func.def pure inline @ggml_kquant_run_shape(%format: index) -> (index) { + %f6 = index.constant 6 : index + %f11 = index.constant 11 : index + %f21 = index.constant 21 : index + %f80 = index.constant 80 : index + %s0 = index.constant 0 : index + %s1 = index.constant 1 : index + %s2 = index.constant 2 : index + %is6 = index.cmp eq, %format, %f6 : index + %is11 = index.cmp eq, %format, %f11 : index + %is21 = index.cmp eq, %format, %f21 : index + %is80 = index.cmp eq, %format, %f80 : index + %f12 = index.constant 12 : index + %is12 = index.cmp eq, %format, %f12 : index + %a0 = scalar.ori %is11, %is21 : i1 + %f24 = index.constant 24 : index + %f25 = index.constant 25 : index + %is24 = index.cmp eq, %format, %f24 : index + %is25 = index.cmp eq, %format, %f25 : index + %a1 = scalar.ori %a0, %is12 : i1 + %a2 = scalar.ori %a1, %is24 : i1 + %a3 = scalar.ori %a2, %is25 : i1 + %f28 = index.constant 28 : index + %is28 = index.cmp eq, %format, %f28 : index + %a4 = scalar.ori %a3, %is28 : i1 + %f22 = index.constant 22 : index + %is22 = index.cmp eq, %format, %f22 : index + %a5 = scalar.ori %a4, %is22 : i1 + %f26 = index.constant 26 : index + %f27 = index.constant 27 : index + %is26 = index.cmp eq, %format, %f26 : index + %is27 = index.cmp eq, %format, %f27 : index + %a6 = scalar.ori %a5, %is26 : i1 + %a7 = scalar.ori %a6, %is27 : i1 + %f90 = index.constant 90 : index + %is90 = index.cmp eq, %format, %f90 : index + %a8 = scalar.ori %a7, %is90 : i1 + %f34 = index.constant 34 : index + %f35 = index.constant 35 : index + %f39 = index.constant 39 : index + %is34 = index.cmp eq, %format, %f34 : index + %is35 = index.cmp eq, %format, %f35 : index + %is39 = index.cmp eq, %format, %f39 : index + %a9 = scalar.ori %a8, %is34 : i1 + %a10 = scalar.ori %a9, %is35 : i1 + %a11 = scalar.ori %a10, %is39 : i1 + %is_prism = func.call pure @ggml_kquant_g128_format(%format) : (index) -> (i1) + %a = scalar.ori %a11, %is_prism : i1 + %run16 = scalar.ori %a, %is80 : i1 + %shape0 = scf.select %run16, %s2, %s1 : index + %shape = scf.select %is6, %s0, %shape0 : index + func.return %shape : index +} + +kernel.def target(@ggml_kquant_decode_gfx11_wave64) export("ggml_kquant_swiglu_decode_tokens_f32") @ggml_kquant_swiglu_decode_tokens_f32() { + %output_size = config.get @ggml.kquant_swiglu_decode.output_size : index + %c1 = index.constant 1 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c256 = index.constant 256 : index + %last = index.add %output_size, %c3 : index + %groups = index.div %last, %c4 : index + kernel.launch.config workgroups(%groups, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%input: buffer, %gate: buffer, %up: buffer, %output: buffer) { + %input_size = config.get @ggml.kquant_swiglu_decode.input_size : index + %output_size = config.get @ggml.kquant_swiglu_decode.output_size : index + %tokens = config.get @ggml.kquant_decode.token_count : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c256 = index.constant 256 : index + %gate_format = config.get @ggml.kquant_swiglu_decode.gate_weight_format : index + %up_format = config.get @ggml.kquant_swiglu_decode.up_weight_format : index + %c0 = index.constant 0 : index + %c16 = index.constant 16 : index + %zero_offset = index.constant 0 : offset + %zero = scalar.constant 0.0 : f32 + %zero4 = vector.constant 0.0 : vector<4xf32> + %table = func.call @ggml_kquant_iq4nl_table() : () -> (vector<16xi8>) + %blocks = index.div %input_size, %c256 : index + %wg = kernel.workgroup.id : index + %subgroup = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %wg_row = index.mul %wg, %c4 : index + %row = index.add %wg_row, %subgroup : index + %first_block = index.div %lane, %c16 : index + %lane16 = index.rem %lane, %c16 : index + %lane16_i32 = index.cast %lane16 : index to i32 + %lane16_i8 = scalar.trunci %lane16_i32 : i32 to i8 + %lane16_v = vector.splat %lane16_i8 : vector<4xi8> + %kv_v = vector.table.lookup %table[%lane16_v] : vector<16xi8>, vector<4xi8> -> vector<4xi8> + %kv_i8 = vector.extract %kv_v[0] : vector<4xi8> -> i8 + %kv_i32 = scalar.extsi %kv_i8 : i8 to i32 + %kv_lane = scalar.sitofp %kv_i32 : i32 to f32 + %g_grid_f21 = index.constant 21 : index + %g_grid_f24 = index.constant 24 : index + %g_grid_f25 = index.constant 25 : index + %g_grid_is21 = index.cmp eq, %gate_format, %g_grid_f21 : index + %g_grid_is24 = index.cmp eq, %gate_format, %g_grid_f24 : index + %g_grid_is25 = index.cmp eq, %gate_format, %g_grid_f25 : index + %g_grid_a = scalar.ori %g_grid_is21, %g_grid_is24 : i1 + %g_grid_f28 = index.constant 28 : index + %g_grid_is28 = index.cmp eq, %gate_format, %g_grid_f28 : index + %g_grid_b = scalar.ori %g_grid_a, %g_grid_is25 : i1 + %g_grid_f22 = index.constant 22 : index + %g_grid_is22 = index.cmp eq, %gate_format, %g_grid_f22 : index + %g_grid_c = scalar.ori %g_grid_b, %g_grid_is28 : i1 + %g_grid_d = scalar.ori %g_grid_c, %g_grid_is22 : i1 + %g_grid_f26 = index.constant 26 : index + %g_grid_f27 = index.constant 27 : index + %g_grid_is26 = index.cmp eq, %gate_format, %g_grid_f26 : index + %g_grid_is27 = index.cmp eq, %gate_format, %g_grid_f27 : index + %g_grid_e = scalar.ori %g_grid_d, %g_grid_is26 : i1 + %g_grid = scalar.ori %g_grid_e, %g_grid_is27 : i1 + %u_grid_f21 = index.constant 21 : index + %u_grid_f24 = index.constant 24 : index + %u_grid_f25 = index.constant 25 : index + %u_grid_is21 = index.cmp eq, %up_format, %u_grid_f21 : index + %u_grid_is24 = index.cmp eq, %up_format, %u_grid_f24 : index + %u_grid_is25 = index.cmp eq, %up_format, %u_grid_f25 : index + %u_grid_a = scalar.ori %u_grid_is21, %u_grid_is24 : i1 + %u_grid_f28 = index.constant 28 : index + %u_grid_is28 = index.cmp eq, %up_format, %u_grid_f28 : index + %u_grid_b = scalar.ori %u_grid_a, %u_grid_is25 : i1 + %u_grid_f22 = index.constant 22 : index + %u_grid_is22 = index.cmp eq, %up_format, %u_grid_f22 : index + %u_grid_c = scalar.ori %u_grid_b, %u_grid_is28 : i1 + %u_grid_d = scalar.ori %u_grid_c, %u_grid_is22 : i1 + %u_grid_f26 = index.constant 26 : index + %u_grid_f27 = index.constant 27 : index + %u_grid_is26 = index.cmp eq, %up_format, %u_grid_f26 : index + %u_grid_is27 = index.cmp eq, %up_format, %u_grid_f27 : index + %u_grid_e = scalar.ori %u_grid_d, %u_grid_is26 : i1 + %u_grid = scalar.ori %u_grid_e, %u_grid_is27 : i1 + %needs_grid = scalar.ori %g_grid, %u_grid : i1 + // Gate and up each get their own codebook buffer, so a pair whose two formats need different grids runs here. + // A pair sharing one codebook (same format, or IQ1_S with IQ1_M) fills it once: the up lanes read the gate grid. + %same_format = index.cmp eq, %gate_format, %up_format : index + %iq1_g26 = index.constant 26 : index + %iq1_g27 = index.constant 27 : index + %gate_iq1_s = index.cmp eq, %gate_format, %iq1_g26 : index + %gate_iq1_m = index.cmp eq, %gate_format, %iq1_g27 : index + %up_iq1_s = index.cmp eq, %up_format, %iq1_g26 : index + %up_iq1_m = index.cmp eq, %up_format, %iq1_g27 : index + %gate_iq1 = scalar.ori %gate_iq1_s, %gate_iq1_m : i1 + %up_iq1 = scalar.ori %up_iq1_s, %up_iq1_m : i1 + %both_iq1 = scalar.andi %gate_iq1, %up_iq1 : i1 + %same_grid0 = scalar.ori %same_format, %both_iq1 : i1 + %same_grid = scalar.andi %same_grid0, %g_grid : i1 + %true_u = scalar.constant true : i1 + %not_same_grid = scalar.xori %same_grid, %true_u : i1 + %u_own_grid = scalar.andi %u_grid, %not_same_grid : i1 + %grid_bytes = index.constant 4096 : offset + %grid = buffer.alloca align(16) %grid_bytes : buffer + %grid_u = buffer.alloca align(16) %grid_bytes : buffer + scf.if %needs_grid { + scf.if %g_grid { + func.call @ggml_kquant_grid_fill_for(%gate_format, %grid, %subgroup) : (index, buffer, index) -> () + } + scf.if %u_own_grid { + func.call @ggml_kquant_grid_fill_for(%up_format, %grid_u, %subgroup) : (index, buffer, index) -> () + } + kernel.barrier scope(workgroup) ordering(acq_rel) + } + %c0t = index.constant 0 : index + %tv0 = index.cmp ult, %c0t, %tokens : index + %c1t = index.constant 1 : index + %tv1 = index.cmp ult, %c1t, %tokens : index + %c2t = index.constant 2 : index + %tv2 = index.cmp ult, %c2t, %tokens : index + %c3t = index.constant 3 : index + %tv3 = index.cmp ult, %c3t, %tokens : index + %c4t = index.constant 4 : index + %tv4 = index.cmp ult, %c4t, %tokens : index + %c5t = index.constant 5 : index + %tv5 = index.cmp ult, %c5t, %tokens : index + %c6t = index.constant 6 : index + %tv6 = index.cmp ult, %c6t, %tokens : index + %c7t = index.constant 7 : index + %tv7 = index.cmp ult, %c7t, %tokens : index + %valid = index.cmp ult, %row, %output_size : index + scf.if %valid { + %g0_part, %u0_part, %g1_part, %u1_part, %g2_part, %u2_part, %g3_part, %u3_part, %g4_part, %u4_part, %g5_part, %u5_part, %g6_part, %u6_part, %g7_part, %u7_part = scf.for %block = [%first_block to %blocks step %c4](%g0_acc = %zero4 : vector<4xf32>, %u0_acc = %zero4 : vector<4xf32>, %g1_acc = %zero4 : vector<4xf32>, %u1_acc = %zero4 : vector<4xf32>, %g2_acc = %zero4 : vector<4xf32>, %u2_acc = %zero4 : vector<4xf32>, %g3_acc = %zero4 : vector<4xf32>, %u3_acc = %zero4 : vector<4xf32>, %g4_acc = %zero4 : vector<4xf32>, %u4_acc = %zero4 : vector<4xf32>, %g5_acc = %zero4 : vector<4xf32>, %u5_acc = %zero4 : vector<4xf32>, %g6_acc = %zero4 : vector<4xf32>, %u6_acc = %zero4 : vector<4xf32>, %g7_acc = %zero4 : vector<4xf32>, %u7_acc = %zero4 : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + %gshape = func.call pure @ggml_kquant_run_shape(%gate_format) : (index) -> (index) + %gw, %gp0, %gp1, %gp2, %gp3 = func.call @ggml_kquant_lane_weights(%gate_format, %kv_lane, %grid, %gate, %row, %blocks, %block, %lane16) : (index, f32, buffer, buffer, index, index, index, index) -> (vector<16xf32>, index, index, index, index) + %ushape = func.call pure @ggml_kquant_run_shape(%up_format) : (index) -> (index) + %uw, %up0, %up1, %up2, %up3 = scf.if %same_grid -> (vector<16xf32>, index, index, index, index) { + %uws, %ups0, %ups1, %ups2, %ups3 = func.call @ggml_kquant_lane_weights(%up_format, %kv_lane, %grid, %up, %row, %blocks, %block, %lane16) : (index, f32, buffer, buffer, index, index, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %uws, %ups0, %ups1, %ups2, %ups3 : vector<16xf32>, index, index, index, index + } else { + %uwo, %upo0, %upo1, %upo2, %upo3 = func.call @ggml_kquant_lane_weights(%up_format, %kv_lane, %grid_u, %up, %row, %blocks, %block, %lane16) : (index, f32, buffer, buffer, index, index, index, index) -> (vector<16xf32>, index, index, index, index) + scf.yield %uwo, %upo0, %upo1, %upo2, %upo3 : vector<16xf32>, index, index, index, index + } + %g0_next, %u0_next = scf.if %tv0 -> (vector<4xf32>, vector<4xf32>) { + %g0_s = func.call @ggml_kquant_token_fma(%gshape, %gw, %gp0, %gp1, %gp2, %gp3, %input, %c0t, %blocks, %block, %g0_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + %u0_s = func.call @ggml_kquant_token_fma(%ushape, %uw, %up0, %up1, %up2, %up3, %input, %c0t, %blocks, %block, %u0_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + scf.yield %g0_s, %u0_s : vector<4xf32>, vector<4xf32> + } else { + scf.yield %g0_acc, %u0_acc : vector<4xf32>, vector<4xf32> + } + %g1_next, %u1_next = scf.if %tv1 -> (vector<4xf32>, vector<4xf32>) { + %g1_s = func.call @ggml_kquant_token_fma(%gshape, %gw, %gp0, %gp1, %gp2, %gp3, %input, %c1t, %blocks, %block, %g1_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + %u1_s = func.call @ggml_kquant_token_fma(%ushape, %uw, %up0, %up1, %up2, %up3, %input, %c1t, %blocks, %block, %u1_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + scf.yield %g1_s, %u1_s : vector<4xf32>, vector<4xf32> + } else { + scf.yield %g1_acc, %u1_acc : vector<4xf32>, vector<4xf32> + } + %g2_next, %u2_next = scf.if %tv2 -> (vector<4xf32>, vector<4xf32>) { + %g2_s = func.call @ggml_kquant_token_fma(%gshape, %gw, %gp0, %gp1, %gp2, %gp3, %input, %c2t, %blocks, %block, %g2_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + %u2_s = func.call @ggml_kquant_token_fma(%ushape, %uw, %up0, %up1, %up2, %up3, %input, %c2t, %blocks, %block, %u2_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + scf.yield %g2_s, %u2_s : vector<4xf32>, vector<4xf32> + } else { + scf.yield %g2_acc, %u2_acc : vector<4xf32>, vector<4xf32> + } + %g3_next, %u3_next = scf.if %tv3 -> (vector<4xf32>, vector<4xf32>) { + %g3_s = func.call @ggml_kquant_token_fma(%gshape, %gw, %gp0, %gp1, %gp2, %gp3, %input, %c3t, %blocks, %block, %g3_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + %u3_s = func.call @ggml_kquant_token_fma(%ushape, %uw, %up0, %up1, %up2, %up3, %input, %c3t, %blocks, %block, %u3_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + scf.yield %g3_s, %u3_s : vector<4xf32>, vector<4xf32> + } else { + scf.yield %g3_acc, %u3_acc : vector<4xf32>, vector<4xf32> + } + %g4_next, %u4_next = scf.if %tv4 -> (vector<4xf32>, vector<4xf32>) { + %g4_s = func.call @ggml_kquant_token_fma(%gshape, %gw, %gp0, %gp1, %gp2, %gp3, %input, %c4t, %blocks, %block, %g4_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + %u4_s = func.call @ggml_kquant_token_fma(%ushape, %uw, %up0, %up1, %up2, %up3, %input, %c4t, %blocks, %block, %u4_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + scf.yield %g4_s, %u4_s : vector<4xf32>, vector<4xf32> + } else { + scf.yield %g4_acc, %u4_acc : vector<4xf32>, vector<4xf32> + } + %g5_next, %u5_next = scf.if %tv5 -> (vector<4xf32>, vector<4xf32>) { + %g5_s = func.call @ggml_kquant_token_fma(%gshape, %gw, %gp0, %gp1, %gp2, %gp3, %input, %c5t, %blocks, %block, %g5_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + %u5_s = func.call @ggml_kquant_token_fma(%ushape, %uw, %up0, %up1, %up2, %up3, %input, %c5t, %blocks, %block, %u5_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + scf.yield %g5_s, %u5_s : vector<4xf32>, vector<4xf32> + } else { + scf.yield %g5_acc, %u5_acc : vector<4xf32>, vector<4xf32> + } + %g6_next, %u6_next = scf.if %tv6 -> (vector<4xf32>, vector<4xf32>) { + %g6_s = func.call @ggml_kquant_token_fma(%gshape, %gw, %gp0, %gp1, %gp2, %gp3, %input, %c6t, %blocks, %block, %g6_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + %u6_s = func.call @ggml_kquant_token_fma(%ushape, %uw, %up0, %up1, %up2, %up3, %input, %c6t, %blocks, %block, %u6_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + scf.yield %g6_s, %u6_s : vector<4xf32>, vector<4xf32> + } else { + scf.yield %g6_acc, %u6_acc : vector<4xf32>, vector<4xf32> + } + %g7_next, %u7_next = scf.if %tv7 -> (vector<4xf32>, vector<4xf32>) { + %g7_s = func.call @ggml_kquant_token_fma(%gshape, %gw, %gp0, %gp1, %gp2, %gp3, %input, %c7t, %blocks, %block, %g7_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + %u7_s = func.call @ggml_kquant_token_fma(%ushape, %uw, %up0, %up1, %up2, %up3, %input, %c7t, %blocks, %block, %u7_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + scf.yield %g7_s, %u7_s : vector<4xf32>, vector<4xf32> + } else { + scf.yield %g7_acc, %u7_acc : vector<4xf32>, vector<4xf32> + } + scf.yield %g0_next, %u0_next, %g1_next, %u1_next, %g2_next, %u2_next, %g3_next, %u3_next, %g4_next, %u4_next, %g5_next, %u5_next, %g6_next, %u6_next, %g7_next, %u7_next : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + %leader = index.cmp eq, %lane, %c0 : index + %n_total = index.mul %tokens, %output_size : index + scf.if %tv0 { + %g0_lane = vector.reduce %g0_part, %zero : vector<4xf32>, f32 + %g0_sum = kernel.subgroup.reduce %g0_lane : f32 + %u0_lane = vector.reduce %u0_part, %zero : vector<4xf32>, f32 + %u0_sum = kernel.subgroup.reduce %u0_lane : f32 + scf.if %leader { + %o0_row = index.mul %c0t, %output_size : index + %o0_raw = index.add %o0_row, %row : index + %o0, %o0_bound = index.assume %o0_raw, %n_total [lt(%o0_raw, %n_total)] : index, index + %v0 = func.call @ggml_kquant_silu_mul(%g0_sum, %u0_sum) : (f32, f32) -> (f32) + %ov0 = buffer.view %output[%zero_offset] : buffer -> view<[%o0_bound]xf32> + view.store %v0, %ov0[%o0] : f32, view<[%o0_bound]xf32> + } + } + scf.if %tv1 { + %g1_lane = vector.reduce %g1_part, %zero : vector<4xf32>, f32 + %g1_sum = kernel.subgroup.reduce %g1_lane : f32 + %u1_lane = vector.reduce %u1_part, %zero : vector<4xf32>, f32 + %u1_sum = kernel.subgroup.reduce %u1_lane : f32 + scf.if %leader { + %o1_row = index.mul %c1t, %output_size : index + %o1_raw = index.add %o1_row, %row : index + %o1, %o1_bound = index.assume %o1_raw, %n_total [lt(%o1_raw, %n_total)] : index, index + %v1 = func.call @ggml_kquant_silu_mul(%g1_sum, %u1_sum) : (f32, f32) -> (f32) + %ov1 = buffer.view %output[%zero_offset] : buffer -> view<[%o1_bound]xf32> + view.store %v1, %ov1[%o1] : f32, view<[%o1_bound]xf32> + } + } + scf.if %tv2 { + %g2_lane = vector.reduce %g2_part, %zero : vector<4xf32>, f32 + %g2_sum = kernel.subgroup.reduce %g2_lane : f32 + %u2_lane = vector.reduce %u2_part, %zero : vector<4xf32>, f32 + %u2_sum = kernel.subgroup.reduce %u2_lane : f32 + scf.if %leader { + %o2_row = index.mul %c2t, %output_size : index + %o2_raw = index.add %o2_row, %row : index + %o2, %o2_bound = index.assume %o2_raw, %n_total [lt(%o2_raw, %n_total)] : index, index + %v2 = func.call @ggml_kquant_silu_mul(%g2_sum, %u2_sum) : (f32, f32) -> (f32) + %ov2 = buffer.view %output[%zero_offset] : buffer -> view<[%o2_bound]xf32> + view.store %v2, %ov2[%o2] : f32, view<[%o2_bound]xf32> + } + } + scf.if %tv3 { + %g3_lane = vector.reduce %g3_part, %zero : vector<4xf32>, f32 + %g3_sum = kernel.subgroup.reduce %g3_lane : f32 + %u3_lane = vector.reduce %u3_part, %zero : vector<4xf32>, f32 + %u3_sum = kernel.subgroup.reduce %u3_lane : f32 + scf.if %leader { + %o3_row = index.mul %c3t, %output_size : index + %o3_raw = index.add %o3_row, %row : index + %o3, %o3_bound = index.assume %o3_raw, %n_total [lt(%o3_raw, %n_total)] : index, index + %v3 = func.call @ggml_kquant_silu_mul(%g3_sum, %u3_sum) : (f32, f32) -> (f32) + %ov3 = buffer.view %output[%zero_offset] : buffer -> view<[%o3_bound]xf32> + view.store %v3, %ov3[%o3] : f32, view<[%o3_bound]xf32> + } + } + scf.if %tv4 { + %g4_lane = vector.reduce %g4_part, %zero : vector<4xf32>, f32 + %g4_sum = kernel.subgroup.reduce %g4_lane : f32 + %u4_lane = vector.reduce %u4_part, %zero : vector<4xf32>, f32 + %u4_sum = kernel.subgroup.reduce %u4_lane : f32 + scf.if %leader { + %o4_row = index.mul %c4t, %output_size : index + %o4_raw = index.add %o4_row, %row : index + %o4, %o4_bound = index.assume %o4_raw, %n_total [lt(%o4_raw, %n_total)] : index, index + %v4 = func.call @ggml_kquant_silu_mul(%g4_sum, %u4_sum) : (f32, f32) -> (f32) + %ov4 = buffer.view %output[%zero_offset] : buffer -> view<[%o4_bound]xf32> + view.store %v4, %ov4[%o4] : f32, view<[%o4_bound]xf32> + } + } + scf.if %tv5 { + %g5_lane = vector.reduce %g5_part, %zero : vector<4xf32>, f32 + %g5_sum = kernel.subgroup.reduce %g5_lane : f32 + %u5_lane = vector.reduce %u5_part, %zero : vector<4xf32>, f32 + %u5_sum = kernel.subgroup.reduce %u5_lane : f32 + scf.if %leader { + %o5_row = index.mul %c5t, %output_size : index + %o5_raw = index.add %o5_row, %row : index + %o5, %o5_bound = index.assume %o5_raw, %n_total [lt(%o5_raw, %n_total)] : index, index + %v5 = func.call @ggml_kquant_silu_mul(%g5_sum, %u5_sum) : (f32, f32) -> (f32) + %ov5 = buffer.view %output[%zero_offset] : buffer -> view<[%o5_bound]xf32> + view.store %v5, %ov5[%o5] : f32, view<[%o5_bound]xf32> + } + } + scf.if %tv6 { + %g6_lane = vector.reduce %g6_part, %zero : vector<4xf32>, f32 + %g6_sum = kernel.subgroup.reduce %g6_lane : f32 + %u6_lane = vector.reduce %u6_part, %zero : vector<4xf32>, f32 + %u6_sum = kernel.subgroup.reduce %u6_lane : f32 + scf.if %leader { + %o6_row = index.mul %c6t, %output_size : index + %o6_raw = index.add %o6_row, %row : index + %o6, %o6_bound = index.assume %o6_raw, %n_total [lt(%o6_raw, %n_total)] : index, index + %v6 = func.call @ggml_kquant_silu_mul(%g6_sum, %u6_sum) : (f32, f32) -> (f32) + %ov6 = buffer.view %output[%zero_offset] : buffer -> view<[%o6_bound]xf32> + view.store %v6, %ov6[%o6] : f32, view<[%o6_bound]xf32> + } + } + scf.if %tv7 { + %g7_lane = vector.reduce %g7_part, %zero : vector<4xf32>, f32 + %g7_sum = kernel.subgroup.reduce %g7_lane : f32 + %u7_lane = vector.reduce %u7_part, %zero : vector<4xf32>, f32 + %u7_sum = kernel.subgroup.reduce %u7_lane : f32 + scf.if %leader { + %o7_row = index.mul %c7t, %output_size : index + %o7_raw = index.add %o7_row, %row : index + %o7, %o7_bound = index.assume %o7_raw, %n_total [lt(%o7_raw, %n_total)] : index, index + %v7 = func.call @ggml_kquant_silu_mul(%g7_sum, %u7_sum) : (f32, f32) -> (f32) + %ov7 = buffer.view %output[%zero_offset] : buffer -> view<[%o7_bound]xf32> + view.store %v7, %ov7[%o7] : f32, view<[%o7_bound]xf32> + } + } + } + kernel.return +} + +kernel.def target(@ggml_kquant_decode_gfx11_wave64) export("ggml_kquant_mul_mat_decode_tokens_f32") @ggml_kquant_mul_mat_decode_tokens_f32() { + %output_size = config.get @ggml.kquant_mul_mat_decode.output_size : index + %c1 = index.constant 1 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c256 = index.constant 256 : index + %last = index.add %output_size, %c3 : index + %groups = index.div %last, %c4 : index + kernel.launch.config workgroups(%groups, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%input: buffer, %weight: buffer, %addend: buffer, %output: buffer) { + %input_size = config.get @ggml.kquant_mul_mat_decode.input_size : index + %output_size = config.get @ggml.kquant_mul_mat_decode.output_size : index + %tokens = config.get @ggml.kquant_decode.token_count : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c256 = index.constant 256 : index + %format = config.get @ggml.kquant_mul_mat_decode.weight_format : index + %add = config.get @ggml.kquant_mul_mat_decode.add : index + %c0 = index.constant 0 : index + %c16 = index.constant 16 : index + %zero_offset = index.constant 0 : offset + %zero = scalar.constant 0.0 : f32 + %zero4 = vector.constant 0.0 : vector<4xf32> + %table = func.call @ggml_kquant_iq4nl_table() : () -> (vector<16xi8>) + %blocks = index.div %input_size, %c256 : index + %wg = kernel.workgroup.id : index + %subgroup = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %wg_row = index.mul %wg, %c4 : index + %row = index.add %wg_row, %subgroup : index + %first_block = index.div %lane, %c16 : index + %lane16 = index.rem %lane, %c16 : index + %lane16_i32 = index.cast %lane16 : index to i32 + %lane16_i8 = scalar.trunci %lane16_i32 : i32 to i8 + %lane16_v = vector.splat %lane16_i8 : vector<4xi8> + %kv_v = vector.table.lookup %table[%lane16_v] : vector<16xi8>, vector<4xi8> -> vector<4xi8> + %kv_i8 = vector.extract %kv_v[0] : vector<4xi8> -> i8 + %kv_i32 = scalar.extsi %kv_i8 : i8 to i32 + %kv_lane = scalar.sitofp %kv_i32 : i32 to f32 + %needs_grid_f21 = index.constant 21 : index + %needs_grid_f24 = index.constant 24 : index + %needs_grid_f25 = index.constant 25 : index + %needs_grid_is21 = index.cmp eq, %format, %needs_grid_f21 : index + %needs_grid_is24 = index.cmp eq, %format, %needs_grid_f24 : index + %needs_grid_is25 = index.cmp eq, %format, %needs_grid_f25 : index + %needs_grid_a = scalar.ori %needs_grid_is21, %needs_grid_is24 : i1 + %needs_grid_f28 = index.constant 28 : index + %needs_grid_is28 = index.cmp eq, %format, %needs_grid_f28 : index + %needs_grid_b = scalar.ori %needs_grid_a, %needs_grid_is25 : i1 + %needs_grid_f22 = index.constant 22 : index + %needs_grid_is22 = index.cmp eq, %format, %needs_grid_f22 : index + %needs_grid_c = scalar.ori %needs_grid_b, %needs_grid_is28 : i1 + %needs_grid_d = scalar.ori %needs_grid_c, %needs_grid_is22 : i1 + %needs_grid_f26 = index.constant 26 : index + %needs_grid_f27 = index.constant 27 : index + %needs_grid_is26 = index.cmp eq, %format, %needs_grid_f26 : index + %needs_grid_is27 = index.cmp eq, %format, %needs_grid_f27 : index + %needs_grid_e = scalar.ori %needs_grid_d, %needs_grid_is26 : i1 + %needs_grid = scalar.ori %needs_grid_e, %needs_grid_is27 : i1 + %grid_bytes = index.constant 4096 : offset + %grid = buffer.alloca align(16) %grid_bytes : buffer + scf.if %needs_grid { + func.call @ggml_kquant_grid_fill_for(%format, %grid, %subgroup) : (index, buffer, index) -> () + kernel.barrier scope(workgroup) ordering(acq_rel) + } + %c0t = index.constant 0 : index + %tv0 = index.cmp ult, %c0t, %tokens : index + %c1t = index.constant 1 : index + %tv1 = index.cmp ult, %c1t, %tokens : index + %c2t = index.constant 2 : index + %tv2 = index.cmp ult, %c2t, %tokens : index + %c3t = index.constant 3 : index + %tv3 = index.cmp ult, %c3t, %tokens : index + %c4t = index.constant 4 : index + %tv4 = index.cmp ult, %c4t, %tokens : index + %c5t = index.constant 5 : index + %tv5 = index.cmp ult, %c5t, %tokens : index + %c6t = index.constant 6 : index + %tv6 = index.cmp ult, %c6t, %tokens : index + %c7t = index.constant 7 : index + %tv7 = index.cmp ult, %c7t, %tokens : index + %valid = index.cmp ult, %row, %output_size : index + scf.if %valid { + %m0_part, %m1_part, %m2_part, %m3_part, %m4_part, %m5_part, %m6_part, %m7_part = scf.for %block = [%first_block to %blocks step %c4](%m0_acc = %zero4 : vector<4xf32>, %m1_acc = %zero4 : vector<4xf32>, %m2_acc = %zero4 : vector<4xf32>, %m3_acc = %zero4 : vector<4xf32>, %m4_acc = %zero4 : vector<4xf32>, %m5_acc = %zero4 : vector<4xf32>, %m6_acc = %zero4 : vector<4xf32>, %m7_acc = %zero4 : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + %mshape = func.call pure @ggml_kquant_run_shape(%format) : (index) -> (index) + %mw, %mp0, %mp1, %mp2, %mp3 = func.call @ggml_kquant_lane_weights(%format, %kv_lane, %grid, %weight, %row, %blocks, %block, %lane16) : (index, f32, buffer, buffer, index, index, index, index) -> (vector<16xf32>, index, index, index, index) + %m0_next = scf.if %tv0 -> (vector<4xf32>) { + %m0_s = func.call @ggml_kquant_token_fma(%mshape, %mw, %mp0, %mp1, %mp2, %mp3, %input, %c0t, %blocks, %block, %m0_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + scf.yield %m0_s : vector<4xf32> + } else { + scf.yield %m0_acc : vector<4xf32> + } + %m1_next = scf.if %tv1 -> (vector<4xf32>) { + %m1_s = func.call @ggml_kquant_token_fma(%mshape, %mw, %mp0, %mp1, %mp2, %mp3, %input, %c1t, %blocks, %block, %m1_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + scf.yield %m1_s : vector<4xf32> + } else { + scf.yield %m1_acc : vector<4xf32> + } + %m2_next = scf.if %tv2 -> (vector<4xf32>) { + %m2_s = func.call @ggml_kquant_token_fma(%mshape, %mw, %mp0, %mp1, %mp2, %mp3, %input, %c2t, %blocks, %block, %m2_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + scf.yield %m2_s : vector<4xf32> + } else { + scf.yield %m2_acc : vector<4xf32> + } + %m3_next = scf.if %tv3 -> (vector<4xf32>) { + %m3_s = func.call @ggml_kquant_token_fma(%mshape, %mw, %mp0, %mp1, %mp2, %mp3, %input, %c3t, %blocks, %block, %m3_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + scf.yield %m3_s : vector<4xf32> + } else { + scf.yield %m3_acc : vector<4xf32> + } + %m4_next = scf.if %tv4 -> (vector<4xf32>) { + %m4_s = func.call @ggml_kquant_token_fma(%mshape, %mw, %mp0, %mp1, %mp2, %mp3, %input, %c4t, %blocks, %block, %m4_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + scf.yield %m4_s : vector<4xf32> + } else { + scf.yield %m4_acc : vector<4xf32> + } + %m5_next = scf.if %tv5 -> (vector<4xf32>) { + %m5_s = func.call @ggml_kquant_token_fma(%mshape, %mw, %mp0, %mp1, %mp2, %mp3, %input, %c5t, %blocks, %block, %m5_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + scf.yield %m5_s : vector<4xf32> + } else { + scf.yield %m5_acc : vector<4xf32> + } + %m6_next = scf.if %tv6 -> (vector<4xf32>) { + %m6_s = func.call @ggml_kquant_token_fma(%mshape, %mw, %mp0, %mp1, %mp2, %mp3, %input, %c6t, %blocks, %block, %m6_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + scf.yield %m6_s : vector<4xf32> + } else { + scf.yield %m6_acc : vector<4xf32> + } + %m7_next = scf.if %tv7 -> (vector<4xf32>) { + %m7_s = func.call @ggml_kquant_token_fma(%mshape, %mw, %mp0, %mp1, %mp2, %mp3, %input, %c7t, %blocks, %block, %m7_acc) : (index, vector<16xf32>, index, index, index, index, buffer, index, index, index, vector<4xf32>) -> (vector<4xf32>) + scf.yield %m7_s : vector<4xf32> + } else { + scf.yield %m7_acc : vector<4xf32> + } + scf.yield %m0_next, %m1_next, %m2_next, %m3_next, %m4_next, %m5_next, %m6_next, %m7_next : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + %leader = index.cmp eq, %lane, %c0 : index + %n_total = index.mul %tokens, %output_size : index + scf.if %tv0 { + %m0_lane = vector.reduce %m0_part, %zero : vector<4xf32>, f32 + %m0_sum = kernel.subgroup.reduce %m0_lane : f32 + scf.if %leader { + %o0_row = index.mul %c0t, %output_size : index + %o0_raw = index.add %o0_row, %row : index + %o0, %o0_bound = index.assume %o0_raw, %n_total [lt(%o0_raw, %n_total)] : index, index + %has_add0 = index.cmp eq, %add, %c1 : index + %v0 = scf.if %has_add0 -> (f32) { + %av0 = buffer.view %addend[%zero_offset] : buffer -> view<[%o0_bound]xf32> + %a0 = view.load %av0[%o0] : view<[%o0_bound]xf32> -> f32 + %s0 = scalar.addf %m0_sum, %a0 : f32 + scf.yield %s0 : f32 + } else { + scf.yield %m0_sum : f32 + } + %ov0 = buffer.view %output[%zero_offset] : buffer -> view<[%o0_bound]xf32> + view.store %v0, %ov0[%o0] : f32, view<[%o0_bound]xf32> + } + } + scf.if %tv1 { + %m1_lane = vector.reduce %m1_part, %zero : vector<4xf32>, f32 + %m1_sum = kernel.subgroup.reduce %m1_lane : f32 + scf.if %leader { + %o1_row = index.mul %c1t, %output_size : index + %o1_raw = index.add %o1_row, %row : index + %o1, %o1_bound = index.assume %o1_raw, %n_total [lt(%o1_raw, %n_total)] : index, index + %has_add1 = index.cmp eq, %add, %c1 : index + %v1 = scf.if %has_add1 -> (f32) { + %av1 = buffer.view %addend[%zero_offset] : buffer -> view<[%o1_bound]xf32> + %a1 = view.load %av1[%o1] : view<[%o1_bound]xf32> -> f32 + %s1 = scalar.addf %m1_sum, %a1 : f32 + scf.yield %s1 : f32 + } else { + scf.yield %m1_sum : f32 + } + %ov1 = buffer.view %output[%zero_offset] : buffer -> view<[%o1_bound]xf32> + view.store %v1, %ov1[%o1] : f32, view<[%o1_bound]xf32> + } + } + scf.if %tv2 { + %m2_lane = vector.reduce %m2_part, %zero : vector<4xf32>, f32 + %m2_sum = kernel.subgroup.reduce %m2_lane : f32 + scf.if %leader { + %o2_row = index.mul %c2t, %output_size : index + %o2_raw = index.add %o2_row, %row : index + %o2, %o2_bound = index.assume %o2_raw, %n_total [lt(%o2_raw, %n_total)] : index, index + %has_add2 = index.cmp eq, %add, %c1 : index + %v2 = scf.if %has_add2 -> (f32) { + %av2 = buffer.view %addend[%zero_offset] : buffer -> view<[%o2_bound]xf32> + %a2 = view.load %av2[%o2] : view<[%o2_bound]xf32> -> f32 + %s2 = scalar.addf %m2_sum, %a2 : f32 + scf.yield %s2 : f32 + } else { + scf.yield %m2_sum : f32 + } + %ov2 = buffer.view %output[%zero_offset] : buffer -> view<[%o2_bound]xf32> + view.store %v2, %ov2[%o2] : f32, view<[%o2_bound]xf32> + } + } + scf.if %tv3 { + %m3_lane = vector.reduce %m3_part, %zero : vector<4xf32>, f32 + %m3_sum = kernel.subgroup.reduce %m3_lane : f32 + scf.if %leader { + %o3_row = index.mul %c3t, %output_size : index + %o3_raw = index.add %o3_row, %row : index + %o3, %o3_bound = index.assume %o3_raw, %n_total [lt(%o3_raw, %n_total)] : index, index + %has_add3 = index.cmp eq, %add, %c1 : index + %v3 = scf.if %has_add3 -> (f32) { + %av3 = buffer.view %addend[%zero_offset] : buffer -> view<[%o3_bound]xf32> + %a3 = view.load %av3[%o3] : view<[%o3_bound]xf32> -> f32 + %s3 = scalar.addf %m3_sum, %a3 : f32 + scf.yield %s3 : f32 + } else { + scf.yield %m3_sum : f32 + } + %ov3 = buffer.view %output[%zero_offset] : buffer -> view<[%o3_bound]xf32> + view.store %v3, %ov3[%o3] : f32, view<[%o3_bound]xf32> + } + } + scf.if %tv4 { + %m4_lane = vector.reduce %m4_part, %zero : vector<4xf32>, f32 + %m4_sum = kernel.subgroup.reduce %m4_lane : f32 + scf.if %leader { + %o4_row = index.mul %c4t, %output_size : index + %o4_raw = index.add %o4_row, %row : index + %o4, %o4_bound = index.assume %o4_raw, %n_total [lt(%o4_raw, %n_total)] : index, index + %has_add4 = index.cmp eq, %add, %c1 : index + %v4 = scf.if %has_add4 -> (f32) { + %av4 = buffer.view %addend[%zero_offset] : buffer -> view<[%o4_bound]xf32> + %a4 = view.load %av4[%o4] : view<[%o4_bound]xf32> -> f32 + %s4 = scalar.addf %m4_sum, %a4 : f32 + scf.yield %s4 : f32 + } else { + scf.yield %m4_sum : f32 + } + %ov4 = buffer.view %output[%zero_offset] : buffer -> view<[%o4_bound]xf32> + view.store %v4, %ov4[%o4] : f32, view<[%o4_bound]xf32> + } + } + scf.if %tv5 { + %m5_lane = vector.reduce %m5_part, %zero : vector<4xf32>, f32 + %m5_sum = kernel.subgroup.reduce %m5_lane : f32 + scf.if %leader { + %o5_row = index.mul %c5t, %output_size : index + %o5_raw = index.add %o5_row, %row : index + %o5, %o5_bound = index.assume %o5_raw, %n_total [lt(%o5_raw, %n_total)] : index, index + %has_add5 = index.cmp eq, %add, %c1 : index + %v5 = scf.if %has_add5 -> (f32) { + %av5 = buffer.view %addend[%zero_offset] : buffer -> view<[%o5_bound]xf32> + %a5 = view.load %av5[%o5] : view<[%o5_bound]xf32> -> f32 + %s5 = scalar.addf %m5_sum, %a5 : f32 + scf.yield %s5 : f32 + } else { + scf.yield %m5_sum : f32 + } + %ov5 = buffer.view %output[%zero_offset] : buffer -> view<[%o5_bound]xf32> + view.store %v5, %ov5[%o5] : f32, view<[%o5_bound]xf32> + } + } + scf.if %tv6 { + %m6_lane = vector.reduce %m6_part, %zero : vector<4xf32>, f32 + %m6_sum = kernel.subgroup.reduce %m6_lane : f32 + scf.if %leader { + %o6_row = index.mul %c6t, %output_size : index + %o6_raw = index.add %o6_row, %row : index + %o6, %o6_bound = index.assume %o6_raw, %n_total [lt(%o6_raw, %n_total)] : index, index + %has_add6 = index.cmp eq, %add, %c1 : index + %v6 = scf.if %has_add6 -> (f32) { + %av6 = buffer.view %addend[%zero_offset] : buffer -> view<[%o6_bound]xf32> + %a6 = view.load %av6[%o6] : view<[%o6_bound]xf32> -> f32 + %s6 = scalar.addf %m6_sum, %a6 : f32 + scf.yield %s6 : f32 + } else { + scf.yield %m6_sum : f32 + } + %ov6 = buffer.view %output[%zero_offset] : buffer -> view<[%o6_bound]xf32> + view.store %v6, %ov6[%o6] : f32, view<[%o6_bound]xf32> + } + } + scf.if %tv7 { + %m7_lane = vector.reduce %m7_part, %zero : vector<4xf32>, f32 + %m7_sum = kernel.subgroup.reduce %m7_lane : f32 + scf.if %leader { + %o7_row = index.mul %c7t, %output_size : index + %o7_raw = index.add %o7_row, %row : index + %o7, %o7_bound = index.assume %o7_raw, %n_total [lt(%o7_raw, %n_total)] : index, index + %has_add7 = index.cmp eq, %add, %c1 : index + %v7 = scf.if %has_add7 -> (f32) { + %av7 = buffer.view %addend[%zero_offset] : buffer -> view<[%o7_bound]xf32> + %a7 = view.load %av7[%o7] : view<[%o7_bound]xf32> -> f32 + %s7 = scalar.addf %m7_sum, %a7 : f32 + scf.yield %s7 : f32 + } else { + scf.yield %m7_sum : f32 + } + %ov7 = buffer.view %output[%zero_offset] : buffer -> view<[%o7_bound]xf32> + view.store %v7, %ov7[%o7] : f32, view<[%o7_bound]xf32> + } + } + } + kernel.return +} + +// Reference for the cases: one thread per row, one weight at a time from single bytes, following +// ggml's dequantize_row_q4_K / q5_K / iq4_xs. +func.def inline @ggml_kquant_ref_byte(%format: index, %weight: buffer, %block_base: offset, %index: index) -> (i32) { + %f11 = index.constant 11 : index + %f21 = index.constant 21 : index + %is11a = index.cmp eq, %format, %f11 : index + %is21 = index.cmp eq, %format, %f21 : index + %is11 = scalar.ori %is11a, %is21 : i1 + %r = scf.if %is11 -> (i32) { + %i = index.assume %index [range(%index, 0, 109)] : index + %v = buffer.view %weight[%block_base] : buffer -> view<110xi8> + %b = view.load %v[%i] : view<110xi8> -> i8 + %u = scalar.extui %b : i8 to i32 + scf.yield %u : i32 + } else { + %u = func.call @ggml_kquant_ref_byte_sized(%format, %weight, %block_base, %index) : (index, buffer, offset, index) -> (i32) + scf.yield %u : i32 + } + func.return %r : i32 +} + +func.def inline @ggml_kquant_ref_byte_sized(%format: index, %weight: buffer, %block_base: offset, %index: index) -> (i32) { + %f5 = index.constant 5 : index + %f6 = index.constant 6 : index + %f23 = index.constant 23 : index + %f80 = index.constant 80 : index + %is5 = index.cmp eq, %format, %f5 : index + %is6 = index.cmp eq, %format, %f6 : index + %is23 = index.cmp eq, %format, %f23 : index + %is80 = index.cmp eq, %format, %f80 : index + %byte = scf.if %is5 -> (i8) { + %i = index.assume %index [range(%index, 0, 175)] : index + %v = buffer.view %weight[%block_base] : buffer -> view<176xi8> + %b = view.load %v[%i] : view<176xi8> -> i8 + scf.yield %b : i8 + } else { + %b6 = scf.if %is6 -> (i8) { + %i = index.assume %index [range(%index, 0, 209)] : index + %v = buffer.view %weight[%block_base] : buffer -> view<210xi8> + %b = view.load %v[%i] : view<210xi8> -> i8 + scf.yield %b : i8 + } else { + %b23 = scf.if %is23 -> (i8) { + %i = index.assume %index [range(%index, 0, 135)] : index + %v = buffer.view %weight[%block_base] : buffer -> view<136xi8> + %b = view.load %v[%i] : view<136xi8> -> i8 + scf.yield %b : i8 + } else { + %b80 = scf.if %is80 -> (i8) { + %i = index.assume %index [range(%index, 0, 271)] : index + %v = buffer.view %weight[%block_base] : buffer -> view<272xi8> + %b = view.load %v[%i] : view<272xi8> -> i8 + scf.yield %b : i8 + } else { + %i = index.assume %index [range(%index, 0, 143)] : index + %v = buffer.view %weight[%block_base] : buffer -> view<144xi8> + %b = view.load %v[%i] : view<144xi8> -> i8 + scf.yield %b : i8 + } + scf.yield %b80 : i8 + } + scf.yield %b23 : i8 + } + scf.yield %b6 : i8 + } + %r = scalar.extui %byte : i8 to i32 + func.return %r : i32 +} + +// The f16 at bytes %index, %index + 1 of a block. +func.def inline @ggml_kquant_ref_f16(%format: index, %weight: buffer, %block_base: offset, %index: index) -> (f32) { + %c1 = index.constant 1 : index + %lo = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %index) : (index, buffer, offset, index) -> (i32) + %index1 = index.add %index, %c1 : index + %hi = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %index1) : (index, buffer, offset, index) -> (i32) + %lo8 = scalar.trunci %lo : i32 to i8 + %hi8 = scalar.trunci %hi : i32 to i8 + %bytes = vector.from_elements %lo8, %hi8 : vector<2xi8> + %h = vector.bitcast %bytes : vector<2xi8> to vector<1xf16> + %h0 = vector.extract %h[0] : vector<1xf16> -> f16 + %f = scalar.extf %h0 : f16 to f32 + func.return %f : f32 +} + +func.def inline @ggml_kquant_ref_value_k(%format: index, %table: vector<16xi8>, %weight: buffer, %row_base: offset, %k: index) -> (f32) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c48 = index.constant 48 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %f5 = index.constant 5 : index + %f23 = index.constant 23 : index + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c6_i32 = scalar.constant 6 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c16_i32 = scalar.constant 16 : i32 + %c32_i32 = scalar.constant 32 : i32 + %c63_i32 = scalar.constant 63 : i32 + %block = index.div %k, %c256 : index + %i = index.rem %k, %c256 : index + %block_bytes = func.call pure @ggml_kquant_block_bytes(%format) : (index) -> (offset) + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %hv = buffer.view %weight[%block_base] : buffer -> view<2xf16> + %dm = vector.load %hv[%c0] : view<2xf16> -> vector<2xf16> + %d_f16 = vector.extract %dm[0] : vector<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %is_iq4 = index.cmp eq, %format, %f23 : index + %value = scf.if %is_iq4 -> (f32) { + %ib = index.div %i, %c32 : index + %r = index.rem %i, %c32 : index + %r16 = index.rem %r, %c16 : index + %qs_ib = index.mul %ib, %c16 : index + %qs_rel = index.add %qs_ib, %r16 : index + %qs_at = index.add %c8, %qs_rel : index + %byte = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %qs_at) : (index, buffer, offset, index) -> (i32) + %is_high = index.cmp uge, %r, %c16 : index + %byte_hi = scalar.shrui %byte, %c4_i32 : i32 + %byte_lo = scalar.andi %byte, %c15_i32 : i32 + %nibble = scf.select %is_high, %byte_hi, %byte_lo : i32 + %nibble_i8 = scalar.trunci %nibble : i32 to i8 + %codes = vector.splat %nibble_i8 : vector<4xi8> + %looked = vector.table.lookup %table[%codes] : vector<16xi8>, vector<4xi8> -> vector<4xi8> + %kv_i8 = vector.extract %looked[0] : vector<4xi8> -> i8 + %kv_i32 = scalar.extsi %kv_i8 : i8 to i32 + %kv = scalar.sitofp %kv_i32 : i32 to f32 + %sl_at0 = index.div %ib, %c2 : index + %sl_at = index.add %c4, %sl_at0 : index + %sl_byte = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %sl_at) : (index, buffer, offset, index) -> (i32) + %ib_odd = index.rem %ib, %c2 : index + %ib_odd_i32 = index.cast %ib_odd : index to i32 + %sl_shift = scalar.muli %ib_odd_i32, %c4_i32 : i32 + %sl0 = scalar.shrui %sl_byte, %sl_shift : i32 + %sl = scalar.andi %sl0, %c15_i32 : i32 + %sh_lo = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %c2) : (index, buffer, offset, index) -> (i32) + %sh_hi = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %c3) : (index, buffer, offset, index) -> (i32) + %sh_hi_s = scalar.shli %sh_hi, %c8_i32 : i32 + %scales_h = scalar.ori %sh_lo, %sh_hi_s : i32 + %ib_i32 = index.cast %ib : index to i32 + %sh_shift = scalar.muli %ib_i32, %c2_i32 : i32 + %sh0 = scalar.shrui %scales_h, %sh_shift : i32 + %sh1 = scalar.andi %sh0, %c3_i32 : i32 + %sh = scalar.shli %sh1, %c4_i32 : i32 + %ls = scalar.ori %sl, %sh : i32 + %ls_c = scalar.subi %ls, %c32_i32 : i32 + %ls_f = scalar.sitofp %ls_c : i32 to f32 + %dl = scalar.mulf %d, %ls_f : f32 + %v = scalar.mulf %dl, %kv : f32 + scf.yield %v : f32 + } else { + %is5 = index.cmp eq, %format, %f5 : index + %dmin_f16 = vector.extract %dm[1] : vector<2xf16> -> f16 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %j64 = index.div %i, %c64 : index + %r = index.rem %i, %c64 : index + %is_high = index.cmp uge, %r, %c32 : index + %l = index.rem %r, %c32 : index + %sub0 = index.mul %j64, %c2 : index + %sub1 = index.add %sub0, %c1 : index + %sub = scf.select %is_high, %sub1, %sub0 : index + %qs_off = scf.select %is5, %c48, %c16 : index + %qs_j = index.mul %j64, %c32 : index + %qs_rel = index.add %qs_j, %l : index + %qs_at = index.add %qs_off, %qs_rel : index + %byte = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %qs_at) : (index, buffer, offset, index) -> (i32) + %byte_hi = scalar.shrui %byte, %c4_i32 : i32 + %byte_lo = scalar.andi %byte, %c15_i32 : i32 + %nibble = scf.select %is_high, %byte_hi, %byte_lo : i32 + %q = scf.if %is5 -> (i32) { + %qh_at = index.add %c16, %l : index + %qh = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %qh_at) : (index, buffer, offset, index) -> (i32) + %sub_i32 = index.cast %sub : index to i32 + %bit0 = scalar.shrui %qh, %sub_i32 : i32 + %bit = scalar.andi %bit0, %c1_i32 : i32 + %add = scalar.shli %bit, %c4_i32 : i32 + %q5 = scalar.addi %nibble, %add : i32 + scf.yield %q5 : i32 + } else { + scf.yield %nibble : i32 + } + // get_scale_min_k4(sub, scales), scales at byte 4. + %low_sub = index.cmp ult, %sub, %c4 : index + %sub_m4_raw = index.add %sub, %c4 : index + %at_j = index.add %c4, %sub : index + %at_j4 = index.add %c4, %sub_m4_raw : index + %b_j = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %at_j) : (index, buffer, offset, index) -> (i32) + %sc, %m = scf.if %low_sub -> (i32, i32) { + %b_j4 = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %at_j4) : (index, buffer, offset, index) -> (i32) + %s = scalar.andi %b_j, %c63_i32 : i32 + %mm = scalar.andi %b_j4, %c63_i32 : i32 + scf.yield %s, %mm : i32, i32 + } else { + %b_j4 = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %at_j4) : (index, buffer, offset, index) -> (i32) + %at_jm4 = index.sub %at_j, %c4 : index + %b_jm4 = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %at_jm4) : (index, buffer, offset, index) -> (i32) + %s_lo = scalar.andi %b_j4, %c15_i32 : i32 + %s_hi0 = scalar.shrui %b_jm4, %c6_i32 : i32 + %s_hi = scalar.shli %s_hi0, %c4_i32 : i32 + %s = scalar.ori %s_lo, %s_hi : i32 + %m_lo = scalar.shrui %b_j4, %c4_i32 : i32 + %m_hi0 = scalar.shrui %b_j, %c6_i32 : i32 + %m_hi = scalar.shli %m_hi0, %c4_i32 : i32 + %mm = scalar.ori %m_lo, %m_hi : i32 + scf.yield %s, %mm : i32, i32 + } + %q_f = scalar.sitofp %q : i32 to f32 + %sc_f = scalar.sitofp %sc : i32 to f32 + %m_f = scalar.sitofp %m : i32 to f32 + %dsc = scalar.mulf %d, %sc_f : f32 + %dm_m = scalar.mulf %dmin, %m_f : f32 + %w = scalar.mulf %dsc, %q_f : f32 + %v = scalar.subf %w, %dm_m : f32 + scf.yield %v : f32 + } + func.return %value : f32 +} + +// Q3_K weight for the reference, from single bytes (dequantize_row_q3_K). +func.def inline @ggml_kquant_ref_value_q3k(%format: index, %weight: buffer, %block_base: offset, %i: index) -> (f32) { + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c96 = index.constant 96 : index + %c104 = index.constant 104 : index + %c108 = index.constant 108 : index + %c128 = index.constant 128 : index + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c32_i32 = scalar.constant 32 : i32 + %zero_i32 = scalar.constant 0 : i32 + %n = index.div %i, %c128 : index + %r = index.rem %i, %c128 : index + %j = index.div %r, %c32 : index + %r2 = index.rem %r, %c32 : index + %h = index.div %r2, %c16 : index + %l = index.rem %r2, %c16 : index + %n32 = index.mul %n, %c32 : index + %h16 = index.mul %h, %c16 : index + %qs_rel0 = index.add %n32, %h16 : index + %qs_rel = index.add %qs_rel0, %l : index + %qs_at = index.add %c32, %qs_rel : index + %qb = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %qs_at) : (index, buffer, offset, index) -> (i32) + %j_i32 = index.cast %j : index to i32 + %n_i32 = index.cast %n : index to i32 + %qshift = scalar.muli %j_i32, %c2_i32 : i32 + %q0 = scalar.shrui %qb, %qshift : i32 + %q = scalar.andi %q0, %c3_i32 : i32 + %hm_at = index.add %h16, %l : index + %hmb = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %hm_at) : (index, buffer, offset, index) -> (i32) + %n4 = scalar.muli %n_i32, %c4_i32 : i32 + %hshift = scalar.addi %n4, %j_i32 : i32 + %hb0 = scalar.shrui %hmb, %hshift : i32 + %hb = scalar.andi %hb0, %c1_i32 : i32 + %has = scalar.cmpi ne, %hb, %zero_i32 : i32 + %sub4 = scf.select %has, %zero_i32, %c4_i32 : i32 + %qv = scalar.subi %q, %sub4 : i32 + %n8 = index.mul %n, %c8 : index + %j2 = index.mul %j, %c2 : index + %s0 = index.add %n8, %j2 : index + %s = index.add %s0, %h : index + %w = index.div %s, %c4 : index + %b = index.rem %s, %c4 : index + %w_odd = index.rem %w, %c2 : index + %w_hi = index.div %w, %c2 : index + %lo_rel0 = index.mul %w_odd, %c4 : index + %lo_rel = index.add %lo_rel0, %b : index + %lo_at = index.add %c96, %lo_rel : index + %hi_at = index.add %c104, %b : index + %lob = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %lo_at) : (index, buffer, offset, index) -> (i32) + %hib = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %hi_at) : (index, buffer, offset, index) -> (i32) + %w_hi_i32 = index.cast %w_hi : index to i32 + %w_i32 = index.cast %w : index to i32 + %lshift = scalar.muli %w_hi_i32, %c4_i32 : i32 + %hshift2 = scalar.muli %w_i32, %c2_i32 : i32 + %sl0 = scalar.shrui %lob, %lshift : i32 + %sl = scalar.andi %sl0, %c15_i32 : i32 + %sh0 = scalar.shrui %hib, %hshift2 : i32 + %sh1 = scalar.andi %sh0, %c3_i32 : i32 + %sh = scalar.shli %sh1, %c4_i32 : i32 + %sc0 = scalar.ori %sl, %sh : i32 + %sc = scalar.subi %sc0, %c32_i32 : i32 + %d = func.call @ggml_kquant_ref_f16(%format, %weight, %block_base, %c108) : (index, buffer, offset, index) -> (f32) + %qv_f = scalar.sitofp %qv : i32 to f32 + %sc_f = scalar.sitofp %sc : i32 to f32 + %dl = scalar.mulf %d, %sc_f : f32 + %v = scalar.mulf %dl, %qv_f : f32 + func.return %v : f32 +} + +// IQ3_S weight for the reference, from single bytes (dequantize_row_iq3_s). +func.def inline @ggml_kquant_ref_value_iq3s(%grid_buf: buffer, %format: index, %weight: buffer, %block_base: offset, %i: index) -> (f32) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c32 = index.constant 32 : index + %c66 = index.constant 66 : index + %c74 = index.constant 74 : index + %c106 = index.constant 106 : index + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c255_i32 = scalar.constant 255 : i32 + %zero_i32 = scalar.constant 0 : i32 + %zero_offset = index.constant 0 : offset + %ib = index.div %i, %c32 : index + %r = index.rem %i, %c32 : index + %l = index.div %r, %c8 : index + %r8 = index.rem %r, %c8 : index + %k = index.div %r8, %c4 : index + %j = index.rem %r8, %c4 : index + %l2 = index.mul %l, %c2 : index + %e = index.add %l2, %k : index + %ib8 = index.mul %ib, %c8 : index + %qs_rel = index.add %ib8, %e : index + %qs_at = index.add %c2, %qs_rel : index + %qs = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %qs_at) : (index, buffer, offset, index) -> (i32) + %qh_at = index.add %c66, %ib : index + %qh = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %qh_at) : (index, buffer, offset, index) -> (i32) + %e_i32 = index.cast %e : index to i32 + %hb0 = scalar.shrui %qh, %e_i32 : i32 + %hb1 = scalar.andi %hb0, %c1_i32 : i32 + %hb = scalar.shli %hb1, %c8_i32 : i32 + %gi_i32 = scalar.ori %qs, %hb : i32 + %gi0 = index.cast %gi_i32 : i32 to index + %gi = index.assume %gi0 [range(%gi0, 0, 511)] : index + %grid = buffer.view %grid_buf[%zero_offset] : buffer -> view<512xi32> + %g = view.load %grid[%gi] : view<512xi32> -> i32 + %j_i32 = index.cast %j : index to i32 + %j8 = scalar.muli %j_i32, %c8_i32 : i32 + %gb0 = scalar.shrui %g, %j8 : i32 + %gb = scalar.andi %gb0, %c255_i32 : i32 + %ib4 = index.mul %ib, %c4 : index + %sg_rel = index.add %ib4, %l : index + %sg_at = index.add %c74, %sg_rel : index + %sg = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %sg_at) : (index, buffer, offset, index) -> (i32) + %k_i32 = index.cast %k : index to i32 + %k4 = scalar.muli %k_i32, %c4_i32 : i32 + %sbit_at = scalar.addi %k4, %j_i32 : i32 + %sb0 = scalar.shrui %sg, %sbit_at : i32 + %sb = scalar.andi %sb0, %c1_i32 : i32 + %neg = scalar.cmpi ne, %sb, %zero_i32 : i32 + %ib_2 = index.div %ib, %c2 : index + %sc_at = index.add %c106, %ib_2 : index + %scb = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %sc_at) : (index, buffer, offset, index) -> (i32) + %ib_odd = index.rem %ib, %c2 : index + %ib_odd_i32 = index.cast %ib_odd : index to i32 + %sc_shift = scalar.muli %ib_odd_i32, %c4_i32 : i32 + %sc0 = scalar.shrui %scb, %sc_shift : i32 + %sc = scalar.andi %sc0, %c15_i32 : i32 + %sc2 = scalar.muli %sc, %c2_i32 : i32 + %sc21 = scalar.addi %sc2, %c1_i32 : i32 + %d = func.call @ggml_kquant_ref_f16(%format, %weight, %block_base, %c0) : (index, buffer, offset, index) -> (f32) + %sc_f = scalar.sitofp %sc21 : i32 to f32 + %g_f = scalar.sitofp %gb : i32 to f32 + %db = scalar.mulf %d, %sc_f : f32 + %v0 = scalar.mulf %db, %g_f : f32 + %nv = scalar.negf %v0 : f32 + %v = scf.select %neg, %nv, %v0 : f32 + func.return %v : f32 +} + +// Q6_K, IQ4_NL and Q8_0 weights for the reference, from single bytes. +// Scalar reference for TQ1_0 (34), TQ2_0 (35) and MXFP4 (39), written from ggml's dequantize_row_tq1_0, +// dequantize_row_tq2_0 and dequantize_row_mxfp4, independent of motifs/dequant_1bit.loom. +// %block_base is the first byte of the 256-value group; %i is 0..255 within it. +func.def pure inline @ggml_kquant_ref_1bit_format(%format: index) -> (i1) { + %f34 = index.constant 34 : index + %f35 = index.constant 35 : index + %f39 = index.constant 39 : index + %is34 = index.cmp eq, %format, %f34 : index + %is35 = index.cmp eq, %format, %f35 : index + %is39 = index.cmp eq, %format, %f39 : index + %is3435 = scalar.ori %is34, %is35 : i1 + %is = scalar.ori %is3435, %is39 : i1 + func.return %is : i1 +} + +func.def inline @ggml_kquant_ref_value_1bit(%format: index, %weight: buffer, %block_base: offset, %i: index) -> (f32) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c17 = index.constant 17 : index + %c32 = index.constant 32 : index + %c128 = index.constant 128 : index + %c160 = index.constant 160 : index + %c240 = index.constant 240 : index + %f35 = index.constant 35 : index + %f39 = index.constant 39 : index + %is35 = index.cmp eq, %format, %f35 : index + %is39 = index.cmp eq, %format, %f39 : index + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c5_i32 = scalar.constant 5 : i32 + %c6_i32 = scalar.constant 6 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c9_i32 = scalar.constant 9 : i32 + %c12_i32 = scalar.constant 12 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c23_i32 = scalar.constant 23 : i32 + %c27_i32 = scalar.constant 27 : i32 + %c81_i32 = scalar.constant 81 : i32 + %c255_i32 = scalar.constant 255 : i32 + %c0_i32 = scalar.constant 0 : i32 + %value = scf.if %is39 -> (f32) { + // MXFP4: eight 17-byte blocks of 32 values (E8M0 e, then qs[16]). Value j is the low nibble of + // qs[j], value j + 16 the high nibble; kvalues {0,1,2,3,4,6,8,12} with sign bit 8, times the + // E8M0 half scale 2^(e - 127) / 2 built from its f32 bits ((e - 1) << 23, or 0x00200000 << e). + %mb = index.div %i, %c32 : index + %j = index.rem %i, %c32 : index + %bv = buffer.view %weight[%block_base] : buffer -> view<136xi8> + %e_at0 = index.mul %mb, %c17 : index + %e_at = index.assume %e_at0 [range(%e_at0, 0, 119)] : index + %e_i8 = view.load %bv[%e_at] : view<136xi8> -> i8 + %e = scalar.extui %e_i8 : i8 to i32 + %hi = index.cmp uge, %j, %c16 : index + %jq = index.rem %j, %c16 : index + %q_at0 = index.add %e_at0, %c1 : index + %q_at1 = index.add %q_at0, %jq : index + %q_at = index.assume %q_at1 [range(%q_at1, 1, 135)] : index + %q_i8 = view.load %bv[%q_at] : view<136xi8> -> i8 + %q = scalar.extui %q_i8 : i8 to i32 + %q_high = scalar.shrui %q, %c4_i32 : i32 + %q_sel = scf.select %hi, %q_high, %q : i32 + %code = scalar.andi %q_sel, %c15_i32 : i32 + %mag3 = scalar.andi %code, %c7_i32 : i32 + %is5 = scalar.cmpi eq, %mag3, %c5_i32 : i32 + %is6 = scalar.cmpi eq, %mag3, %c6_i32 : i32 + %is7 = scalar.cmpi eq, %mag3, %c7_i32 : i32 + %mag_a = scf.select %is5, %c6_i32, %mag3 : i32 + %mag_b = scf.select %is6, %c8_i32, %mag_a : i32 + %mag = scf.select %is7, %c12_i32, %mag_b : i32 + %sign_bit = scalar.andi %code, %c8_i32 : i32 + %negative = scalar.cmpi ne, %sign_bit, %c0_i32 : i32 + %mag_f = scalar.sitofp %mag : i32 to f32 + %neg_mag_f = scalar.negf %mag_f : f32 + %kv = scf.select %negative, %neg_mag_f, %mag_f : f32 + %is_small = scalar.cmpi ult, %e, %c2_i32 : i32 + %e_m1 = scalar.subi %e, %c1_i32 : i32 + %normal_bits = scalar.shli %e_m1, %c23_i32 : i32 + %sub_unit = scalar.constant 2097152 : i32 + %sub_bits = scalar.shli %sub_unit, %e : i32 + %bits = scf.select %is_small, %sub_bits, %normal_bits : i32 + %bits_v = vector.from_elements %bits : vector<1xi32> + %scale_v = vector.bitcast %bits_v : vector<1xi32> to vector<1xf32> + %scale = vector.extract %scale_v[0] : vector<1xf32> -> f32 + %r = scalar.mulf %kv, %scale : f32 + scf.yield %r : f32 + } else { + %value_tq = scf.if %is35 -> (f32) { + // TQ2_0: qs[64] then d (f16, byte 64). Value i: byte (i / 128) * 32 + i % 32, 2-bit field (i % 128) / 32, + // (q - 1) * d. + %bv = buffer.view %weight[%block_base] : buffer -> view<66xi8> + %hv = buffer.view %weight[%block_base] : buffer -> view<33xf16> + %d_at = index.constant 32 : index + %d_f16 = view.load %hv[%d_at] : view<33xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %half = index.div %i, %c128 : index + %r128 = index.rem %i, %c128 : index + %field = index.div %r128, %c32 : index + %m = index.rem %i, %c32 : index + %b0 = index.mul %half, %c32 : index + %b1 = index.add %b0, %m : index + %b = index.assume %b1 [range(%b1, 0, 63)] : index + %byte_i8 = view.load %bv[%b] : view<66xi8> -> i8 + %byte = scalar.extui %byte_i8 : i8 to i32 + %field_i32 = index.cast %field : index to i32 + %shift = scalar.muli %field_i32, %c2_i32 : i32 + %shifted = scalar.shrui %byte, %shift : i32 + %q = scalar.andi %shifted, %c3_i32 : i32 + %t = scalar.subi %q, %c1_i32 : i32 + %t_f = scalar.sitofp %t : i32 to f32 + %r = scalar.mulf %t_f, %d : f32 + scf.yield %r : f32 + } else { + // TQ1_0: qs[48], qh[4], then d (f16, byte 52). Values 0..159 use qs[i % 32] with n = i / 32, values + // 160..239 qs[32 + (i - 160) % 16] with n = (i - 160) / 16, values 240..255 qh[(i - 240) % 4] with + // n = (i - 240) / 4; q = (byte * 3^n) mod 256, xi = (q * 3) >> 8, value (xi - 1) * d. + %bv = buffer.view %weight[%block_base] : buffer -> view<54xi8> + %hv = buffer.view %weight[%block_base] : buffer -> view<27xf16> + %d_at = index.constant 26 : index + %d_f16 = view.load %hv[%d_at] : view<27xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %in_a = index.cmp ult, %i, %c160 : index + %in_b = index.cmp ult, %i, %c240 : index + %a_byte = index.rem %i, %c32 : index + %a_n = index.div %i, %c32 : index + // Select before subtracting: index arithmetic is no-wrap, so an unconditional i - 160 would let the + // compiler assume i >= 160 and fold the first region away. + %i_b_in = scf.select %in_a, %c160, %i : index + %i_b = index.sub %i_b_in, %c160 : index + %b_r = index.rem %i_b, %c16 : index + %b_byte = index.add %b_r, %c32 : index + %b_n = index.div %i_b, %c16 : index + %i_c_in = scf.select %in_b, %c240, %i : index + %i_c = index.sub %i_c_in, %c240 : index + %c_r = index.rem %i_c, %c4 : index + %c48 = index.constant 48 : index + %c_byte = index.add %c_r, %c48 : index + %c_n = index.div %i_c, %c4 : index + %bc_byte = scf.select %in_b, %b_byte, %c_byte : index + %bc_n = scf.select %in_b, %b_n, %c_n : index + %byte_at0 = scf.select %in_a, %a_byte, %bc_byte : index + %n = scf.select %in_a, %a_n, %bc_n : index + %byte_at = index.assume %byte_at0 [range(%byte_at0, 0, 51)] : index + %byte_i8 = view.load %bv[%byte_at] : view<54xi8> -> i8 + %byte = scalar.extui %byte_i8 : i8 to i32 + %n_i32 = index.cast %n : index to i32 + %n1 = scalar.cmpi eq, %n_i32, %c1_i32 : i32 + %n2 = scalar.cmpi eq, %n_i32, %c2_i32 : i32 + %n3 = scalar.cmpi eq, %n_i32, %c3_i32 : i32 + %n4 = scalar.cmpi eq, %n_i32, %c4_i32 : i32 + %p1 = scf.select %n1, %c3_i32, %c1_i32 : i32 + %p2 = scf.select %n2, %c9_i32, %p1 : i32 + %p3 = scf.select %n3, %c27_i32, %p2 : i32 + %pow3 = scf.select %n4, %c81_i32, %p3 : i32 + %prod = scalar.muli %byte, %pow3 : i32 + %q = scalar.andi %prod, %c255_i32 : i32 + %q3 = scalar.muli %q, %c3_i32 : i32 + %xi = scalar.shrui %q3, %c8_i32 : i32 + %t = scalar.subi %xi, %c1_i32 : i32 + %t_f = scalar.sitofp %t : i32 to f32 + %r = scalar.mulf %t_f, %d : f32 + scf.yield %r : f32 + } + scf.yield %value_tq : f32 + } + func.return %value : f32 +} + +// The reference values of the formats above go to ggml_kquant_ref_value_1bit, the rest to the PrismML +// g128 reference. +func.def inline @ggml_kquant_ref_value_alt(%format: index, %weight: buffer, %block_base: offset, %i: index) -> (f32) { + %is_1bit = func.call pure @ggml_kquant_ref_1bit_format(%format) : (index) -> (i1) + %v = scf.if %is_1bit -> (f32) { + %t = func.call @ggml_kquant_ref_value_1bit(%format, %weight, %block_base, %i) : (index, buffer, offset, index) -> (f32) + scf.yield %t : f32 + } else { + %g = func.call @ggml_kquant_g128_ref_value(%format, %weight, %block_base, %i) : (index, buffer, offset, index) -> (f32) + scf.yield %g : f32 + } + func.return %v : f32 +} + +func.def inline @ggml_kquant_ref_value(%grid: buffer, %format: index, %table: vector<16xi8>, %weight: buffer, %row_base: offset, %k: index) -> (f32) { + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %c192 = index.constant 192 : index + %c208 = index.constant 208 : index + %c256 = index.constant 256 : index + %f6 = index.constant 6 : index + %f20 = index.constant 20 : index + %f80 = index.constant 80 : index + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c32_i32 = scalar.constant 32 : i32 + %is6 = index.cmp eq, %format, %f6 : index + %is20 = index.cmp eq, %format, %f20 : index + %is80 = index.cmp eq, %format, %f80 : index + %block = index.div %k, %c256 : index + %i = index.rem %k, %c256 : index + %block_bytes = func.call pure @ggml_kquant_block_bytes(%format) : (index) -> (offset) + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %f11 = index.constant 11 : index + %is11 = index.cmp eq, %format, %f11 : index + %f21r = index.constant 21 : index + %is21r = index.cmp eq, %format, %f21r : index + %is_g128_ref = func.call pure @ggml_kquant_g128_format(%format) : (index) -> (i1) + %is_1bit_ref = func.call pure @ggml_kquant_ref_1bit_format(%format) : (index) -> (i1) + %is_prism_ref = scalar.ori %is_g128_ref, %is_1bit_ref : i1 + %value = scf.if %is_prism_ref -> (f32) { + %pv = func.call @ggml_kquant_ref_value_alt(%format, %weight, %block_base, %i) : (index, buffer, offset, index) -> (f32) + scf.yield %pv : f32 + } else { + %value_k = scf.if %is11 -> (f32) { + %v = func.call @ggml_kquant_ref_value_q3k(%format, %weight, %block_base, %i) : (index, buffer, offset, index) -> (f32) + scf.yield %v : f32 + } else { + %value21 = scf.if %is21r -> (f32) { + %v = func.call @ggml_kquant_ref_value_iq3s(%grid, %format, %weight, %block_base, %i) : (buffer, index, buffer, offset, index) -> (f32) + scf.yield %v : f32 + } else { + %value6 = scf.if %is6 -> (f32) { + %n = index.div %i, %c128 : index + %r = index.rem %i, %c128 : index + %quarter = index.div %r, %c32 : index + %l = index.rem %r, %c32 : index + %is = index.div %l, %c16 : index + %q_odd = index.rem %quarter, %c2 : index + %q_hi = index.cmp uge, %quarter, %c2 : index + %n64 = index.mul %n, %c64 : index + %odd32 = index.mul %q_odd, %c32 : index + %ql_rel = index.add %n64, %odd32 : index + %ql_at = index.add %ql_rel, %l : index + %ql = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %ql_at) : (index, buffer, offset, index) -> (i32) + %ql_hi = scalar.shrui %ql, %c4_i32 : i32 + %ql_lo = scalar.andi %ql, %c15_i32 : i32 + %nibble = scf.select %q_hi, %ql_hi, %ql_lo : i32 + %n32 = index.mul %n, %c32 : index + %qh_rel = index.add %n32, %l : index + %qh_at = index.add %c128, %qh_rel : index + %qh = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %qh_at) : (index, buffer, offset, index) -> (i32) + %quarter2 = index.mul %quarter, %c2 : index + %qh_shift = index.cast %quarter2 : index to i32 + %hb0 = scalar.shrui %qh, %qh_shift : i32 + %hb1 = scalar.andi %hb0, %c3_i32 : i32 + %hb = scalar.shli %hb1, %c4_i32 : i32 + %q0 = scalar.ori %nibble, %hb : i32 + %q = scalar.subi %q0, %c32_i32 : i32 + %n8 = index.mul %n, %c4 : index + %n8b = index.mul %n8, %c2 : index + %sc_rel0 = index.add %n8b, %is : index + %sc_rel = index.add %sc_rel0, %quarter2 : index + %sc_at = index.add %c192, %sc_rel : index + %sc_u = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %sc_at) : (index, buffer, offset, index) -> (i32) + %sc8 = scalar.trunci %sc_u : i32 to i8 + %sc = scalar.extsi %sc8 : i8 to i32 + %d = func.call @ggml_kquant_ref_f16(%format, %weight, %block_base, %c208) : (index, buffer, offset, index) -> (f32) + %q_f = scalar.sitofp %q : i32 to f32 + %sc_f = scalar.sitofp %sc : i32 to f32 + %dsc = scalar.mulf %d, %sc_f : f32 + %v = scalar.mulf %dsc, %q_f : f32 + scf.yield %v : f32 + } else { + %v2 = scf.if %is80 -> (f32) { + %qb = index.div %i, %c32 : index + %r = index.rem %i, %c32 : index + %qb_base = index.mul %qb, %c32 : index + %qb_2 = index.mul %qb, %c2 : index + %qb_at = index.add %qb_base, %qb_2 : index + %d = func.call @ggml_kquant_ref_f16(%format, %weight, %block_base, %qb_at) : (index, buffer, offset, index) -> (f32) + %q_rel = index.add %qb_at, %c2 : index + %q_at = index.add %q_rel, %r : index + %q_u = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %q_at) : (index, buffer, offset, index) -> (i32) + %q8 = scalar.trunci %q_u : i32 to i8 + %q = scalar.extsi %q8 : i8 to i32 + %q_f = scalar.sitofp %q : i32 to f32 + %v = scalar.mulf %d, %q_f : f32 + scf.yield %v : f32 + } else { + %v20 = scf.if %is20 -> (f32) { + %qb = index.div %i, %c32 : index + %r = index.rem %i, %c32 : index + %qb_16 = index.mul %qb, %c16 : index + %qb_2 = index.mul %qb, %c2 : index + %qb_at = index.add %qb_16, %qb_2 : index + %d = func.call @ggml_kquant_ref_f16(%format, %weight, %block_base, %qb_at) : (index, buffer, offset, index) -> (f32) + %r16 = index.rem %r, %c16 : index + %q_rel = index.add %qb_at, %c2 : index + %q_at = index.add %q_rel, %r16 : index + %byte = func.call @ggml_kquant_ref_byte(%format, %weight, %block_base, %q_at) : (index, buffer, offset, index) -> (i32) + %is_high = index.cmp uge, %r, %c16 : index + %byte_hi = scalar.shrui %byte, %c4_i32 : i32 + %byte_lo = scalar.andi %byte, %c15_i32 : i32 + %nibble = scf.select %is_high, %byte_hi, %byte_lo : i32 + %nibble_i8 = scalar.trunci %nibble : i32 to i8 + %codes = vector.splat %nibble_i8 : vector<4xi8> + %looked = vector.table.lookup %table[%codes] : vector<16xi8>, vector<4xi8> -> vector<4xi8> + %kv_i8 = vector.extract %looked[0] : vector<4xi8> -> i8 + %kv_i32 = scalar.extsi %kv_i8 : i8 to i32 + %kv = scalar.sitofp %kv_i32 : i32 to f32 + %v = scalar.mulf %d, %kv : f32 + scf.yield %v : f32 + } else { + %v = func.call @ggml_kquant_ref_value_k(%format, %table, %weight, %row_base, %k) : (index, vector<16xi8>, buffer, offset, index) -> (f32) + scf.yield %v : f32 + } + scf.yield %v20 : f32 + } + scf.yield %v2 : f32 + } + scf.yield %value6 : f32 + } + scf.yield %value21 : f32 + } + scf.yield %value_k : f32 + } + func.return %value : f32 +} + +kernel.def target(@ggml_kquant_decode_gfx11_wave64) export("ggml_kquant_swiglu_decode_reference_f32") @ggml_kquant_swiglu_decode_reference_f32() { + %output_size = config.get @ggml.kquant_swiglu_decode.output_size : index + %c1 = index.constant 1 : index + kernel.launch.config workgroups(%output_size, %c1, %c1) workgroup_size(%c1, %c1, %c1) : index +} launch(%input: buffer, %gate: buffer, %up: buffer, %output: buffer) { + %input_size = config.get @ggml.kquant_swiglu_decode.input_size : index + %output_size = config.get @ggml.kquant_swiglu_decode.output_size : index + %gate_format = config.get @ggml.kquant_swiglu_decode.gate_weight_format : index + %up_format = config.get @ggml.kquant_swiglu_decode.up_weight_format : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + %zero_offset = index.constant 0 : offset + %zero = scalar.constant 0.0 : f32 + %grid_bytes = index.constant 2048 : offset + %grid = buffer.alloca align(16) %grid_bytes : buffer + %ch0 = index.constant 0 : index + %ch1 = index.constant 1 : index + %ch2 = index.constant 2 : index + %ch3 = index.constant 3 : index + func.call @ggml_kquant_iq3s_grid_fill(%grid, %ch0) : (buffer, index) -> () + func.call @ggml_kquant_iq3s_grid_fill(%grid, %ch1) : (buffer, index) -> () + func.call @ggml_kquant_iq3s_grid_fill(%grid, %ch2) : (buffer, index) -> () + func.call @ggml_kquant_iq3s_grid_fill(%grid, %ch3) : (buffer, index) -> () + %table = func.call @ggml_kquant_iq4nl_table() : () -> (vector<16xi8>) + %blocks = index.div %input_size, %c256 : index + %row0 = kernel.workgroup.id : index + %row, %r_bound = index.assume %row0, %output_size [lt(%row0, %output_size)] : index, index + %g_bytes = func.call pure @ggml_kquant_block_bytes(%gate_format) : (index) -> (offset) + %u_bytes = func.call pure @ggml_kquant_block_bytes(%up_format) : (index) -> (offset) + %g_row_bytes = index.scale %blocks, %g_bytes : index, offset -> offset + %u_row_bytes = index.scale %blocks, %u_bytes : index, offset -> offset + %g_base = index.scale %row, %g_row_bytes : index, offset -> offset + %u_base = index.scale %row, %u_row_bytes : index, offset -> offset + %tokens = config.get @ggml.kquant_decode.token_count : index + %n_total = index.mul %tokens, %output_size : index + %k_total = index.mul %tokens, %input_size : index + %xv = buffer.view %input[%zero_offset] : buffer -> view<[%k_total]xf32> + scf.for %t = [%c0 to %tokens step %c1] { + %x_row = index.mul %t, %input_size : index + %g_sum, %u_sum = scf.for %k = [%c0 to %input_size step %c1](%g_acc = %zero : f32, %u_acc = %zero : f32) -> (f32, f32) { + %xk_raw = index.add %x_row, %k : index + %xk, %xk_bound = index.assume %xk_raw, %k_total [lt(%xk_raw, %k_total)] : index, index + %xvt = buffer.view %input[%zero_offset] : buffer -> view<[%xk_bound]xf32> + %x = view.load %xvt[%xk] : view<[%xk_bound]xf32> -> f32 + %gw = func.call @ggml_kquant_ref_value(%grid, %gate_format, %table, %gate, %g_base, %k) : (buffer, index, vector<16xi8>, buffer, offset, index) -> (f32) + %uw = func.call @ggml_kquant_ref_value(%grid, %up_format, %table, %up, %u_base, %k) : (buffer, index, vector<16xi8>, buffer, offset, index) -> (f32) + %g_next = scalar.fmaf %gw, %x, %g_acc : f32 + %u_next = scalar.fmaf %uw, %x, %u_acc : f32 + scf.yield %g_next, %u_next : f32, f32 + } + %value = func.call @ggml_kquant_silu_mul(%g_sum, %u_sum) : (f32, f32) -> (f32) + %o_row = index.mul %t, %output_size : index + %o_raw = index.add %o_row, %row : index + %o, %o_bound = index.assume %o_raw, %n_total [lt(%o_raw, %n_total)] : index, index + %ov = buffer.view %output[%zero_offset] : buffer -> view<[%o_bound]xf32> + view.store %value, %ov[%o] : f32, view<[%o_bound]xf32> + } + kernel.return +} + +// Cases: 512 inputs (two blocks per row), 6 rows (the last workgroup half full). Weights are an +// integer ramp over the bytes (every byte of word w is w % 64), so scales, codes and the f16 +// d/dmin headers (0x0000..0x3f3f) are all finite and varied. Run with +// --config=ggml.kquant_swiglu_decode.input_size=512 --config=ggml.kquant_swiglu_decode.output_size=6 +// and the formats named in each case. + +// --config=ggml.kquant_swiglu_decode.gate_weight_format=5 --config=ggml.kquant_swiglu_decode.up_weight_format=23 +check.case public @ggml_kquant_swiglu_decode_q5k_iq4xs_case { + %x_seed = check.param.seed base(7300000000000031001) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<512xf32> + %gate = check.generate.iota offset(0) step(16843009) period(64) : tensor<528xi32> + %up = check.generate.iota offset(0) step(16843009) period(64) : tensor<408xi32> + %output = check.generate.fill value(-7.0) : tensor<6xf32> + %expected = check.generate.fill value(7.0) : tensor<6xf32> + kernel.launch @ggml_kquant_swiglu_decode_reference_f32[](%input, %gate, %up, %expected) : [](tensor<512xf32>, tensor<528xi32>, tensor<408xi32>, tensor<6xf32>) + kernel.launch @ggml_kquant_swiglu_decode_f32[](%input, %gate, %up, %output) : [](tensor<512xf32>, tensor<528xi32>, tensor<408xi32>, tensor<6xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<6xf32> + check.return +} + +// --config=ggml.kquant_swiglu_decode.gate_weight_format=4 --config=ggml.kquant_swiglu_decode.up_weight_format=5 +check.case public @ggml_kquant_swiglu_decode_q4k_q5k_case { + %x_seed = check.param.seed base(7300000000000031002) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<512xf32> + %gate = check.generate.iota offset(0) step(16843009) period(64) : tensor<432xi32> + %up = check.generate.iota offset(0) step(16843009) period(64) : tensor<528xi32> + %output = check.generate.fill value(-7.0) : tensor<6xf32> + %expected = check.generate.fill value(7.0) : tensor<6xf32> + kernel.launch @ggml_kquant_swiglu_decode_reference_f32[](%input, %gate, %up, %expected) : [](tensor<512xf32>, tensor<432xi32>, tensor<528xi32>, tensor<6xf32>) + kernel.launch @ggml_kquant_swiglu_decode_f32[](%input, %gate, %up, %output) : [](tensor<512xf32>, tensor<432xi32>, tensor<528xi32>, tensor<6xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<6xf32> + check.return +} + +kernel.def target(@ggml_kquant_decode_gfx11_wave64) export("ggml_kquant_mul_mat_decode_reference_f32") @ggml_kquant_mul_mat_decode_reference_f32() { + %output_size = config.get @ggml.kquant_mul_mat_decode.output_size : index + %c1 = index.constant 1 : index + kernel.launch.config workgroups(%output_size, %c1, %c1) workgroup_size(%c1, %c1, %c1) : index +} launch(%input: buffer, %weight: buffer, %addend: buffer, %output: buffer) { + %input_size = config.get @ggml.kquant_mul_mat_decode.input_size : index + %output_size = config.get @ggml.kquant_mul_mat_decode.output_size : index + %format = config.get @ggml.kquant_mul_mat_decode.weight_format : index + %add = config.get @ggml.kquant_mul_mat_decode.add : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + %zero_offset = index.constant 0 : offset + %zero = scalar.constant 0.0 : f32 + %grid_bytes = index.constant 2048 : offset + %grid = buffer.alloca align(16) %grid_bytes : buffer + %ch0 = index.constant 0 : index + %ch1 = index.constant 1 : index + %ch2 = index.constant 2 : index + %ch3 = index.constant 3 : index + func.call @ggml_kquant_iq3s_grid_fill(%grid, %ch0) : (buffer, index) -> () + func.call @ggml_kquant_iq3s_grid_fill(%grid, %ch1) : (buffer, index) -> () + func.call @ggml_kquant_iq3s_grid_fill(%grid, %ch2) : (buffer, index) -> () + func.call @ggml_kquant_iq3s_grid_fill(%grid, %ch3) : (buffer, index) -> () + %table = func.call @ggml_kquant_iq4nl_table() : () -> (vector<16xi8>) + %blocks = index.div %input_size, %c256 : index + %row0 = kernel.workgroup.id : index + %row, %r_bound = index.assume %row0, %output_size [lt(%row0, %output_size)] : index, index + %bytes = func.call pure @ggml_kquant_block_bytes(%format) : (index) -> (offset) + %row_bytes = index.scale %blocks, %bytes : index, offset -> offset + %base = index.scale %row, %row_bytes : index, offset -> offset + %tokens = config.get @ggml.kquant_decode.token_count : index + %n_total = index.mul %tokens, %output_size : index + %k_total = index.mul %tokens, %input_size : index + scf.for %t = [%c0 to %tokens step %c1] { + %x_row = index.mul %t, %input_size : index + %sum = scf.for %k = [%c0 to %input_size step %c1](%acc = %zero : f32) -> (f32) { + %xk_raw = index.add %x_row, %k : index + %xk, %xk_bound = index.assume %xk_raw, %k_total [lt(%xk_raw, %k_total)] : index, index + %xvt = buffer.view %input[%zero_offset] : buffer -> view<[%xk_bound]xf32> + %x = view.load %xvt[%xk] : view<[%xk_bound]xf32> -> f32 + %w = func.call @ggml_kquant_ref_value(%grid, %format, %table, %weight, %base, %k) : (buffer, index, vector<16xi8>, buffer, offset, index) -> (f32) + %next = scalar.fmaf %w, %x, %acc : f32 + scf.yield %next : f32 + } + %o_row = index.mul %t, %output_size : index + %o_raw = index.add %o_row, %row : index + %o, %o_bound = index.assume %o_raw, %n_total [lt(%o_raw, %n_total)] : index, index + %has_add = index.cmp eq, %add, %c1 : index + %value = scf.if %has_add -> (f32) { + %av = buffer.view %addend[%zero_offset] : buffer -> view<[%o_bound]xf32> + %a = view.load %av[%o] : view<[%o_bound]xf32> -> f32 + %s = scalar.addf %sum, %a : f32 + scf.yield %s : f32 + } else { + scf.yield %sum : f32 + } + %ov = buffer.view %output[%zero_offset] : buffer -> view<[%o_bound]xf32> + view.store %value, %ov[%o] : f32, view<[%o_bound]xf32> + } + kernel.return +} + +// Projection case: 512 inputs, 6 IQ4_XS rows plus an addend. Run with +// --config=ggml.kquant_mul_mat_decode.input_size=512 --config=ggml.kquant_mul_mat_decode.output_size=6 +// --config=ggml.kquant_mul_mat_decode.weight_format=23 --config=ggml.kquant_mul_mat_decode.add=1 +check.case public @ggml_kquant_mul_mat_decode_iq4xs_add_case { + %x_seed = check.param.seed base(7300000000000031003) count(1) : i64 + %a_seed = check.param.seed base(7300000000000031004) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<512xf32> + %weight = check.generate.iota offset(0) step(16843009) period(64) : tensor<408xi32> + %addend = check.generate.random.uniform seed(%a_seed) range(-1.0 to 1.0) : tensor<6xf32> + %output = check.generate.fill value(-7.0) : tensor<6xf32> + %expected = check.generate.fill value(7.0) : tensor<6xf32> + kernel.launch @ggml_kquant_mul_mat_decode_reference_f32[](%input, %weight, %addend, %expected) : [](tensor<512xf32>, tensor<408xi32>, tensor<6xf32>, tensor<6xf32>) + kernel.launch @ggml_kquant_mul_mat_decode_f32[](%input, %weight, %addend, %output) : [](tensor<512xf32>, tensor<408xi32>, tensor<6xf32>, tensor<6xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<6xf32> + check.return +} + +// --config ... input_size=512 output_size=6 gate_weight_format=6 up_weight_format=80 +check.case public @ggml_kquant_swiglu_decode_q6k_q8_0_case { + %x_seed = check.param.seed base(7300000000000031005) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<512xf32> + %gate = check.generate.iota offset(0) step(16843009) period(64) : tensor<630xi32> + %up = check.generate.iota offset(0) step(16843009) period(64) : tensor<816xi32> + %output = check.generate.fill value(-7.0) : tensor<6xf32> + %expected = check.generate.fill value(7.0) : tensor<6xf32> + kernel.launch @ggml_kquant_swiglu_decode_reference_f32[](%input, %gate, %up, %expected) : [](tensor<512xf32>, tensor<630xi32>, tensor<816xi32>, tensor<6xf32>) + kernel.launch @ggml_kquant_swiglu_decode_f32[](%input, %gate, %up, %output) : [](tensor<512xf32>, tensor<630xi32>, tensor<816xi32>, tensor<6xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<6xf32> + check.return +} + +// --config ... ggml.kquant_mul_mat_decode.weight_format=20 add=0 +check.case public @ggml_kquant_mul_mat_decode_iq4nl_case { + %x_seed = check.param.seed base(7300000000000031006) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<512xf32> + %weight = check.generate.iota offset(0) step(16843009) period(64) : tensor<432xi32> + %addend = check.generate.fill value(100.0) : tensor<6xf32> + %output = check.generate.fill value(-7.0) : tensor<6xf32> + %expected = check.generate.fill value(7.0) : tensor<6xf32> + kernel.launch @ggml_kquant_mul_mat_decode_reference_f32[](%input, %weight, %addend, %expected) : [](tensor<512xf32>, tensor<432xi32>, tensor<6xf32>, tensor<6xf32>) + kernel.launch @ggml_kquant_mul_mat_decode_f32[](%input, %weight, %addend, %output) : [](tensor<512xf32>, tensor<432xi32>, tensor<6xf32>, tensor<6xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<6xf32> + check.return +} + +// --config ... input_size=512 output_size=6 gate_weight_format=11 up_weight_format=23 +check.case public @ggml_kquant_swiglu_decode_q3k_iq4xs_case { + %x_seed = check.param.seed base(7300000000000031007) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<512xf32> + %gate = check.generate.iota offset(0) step(16843009) period(64) : tensor<330xi32> + %up = check.generate.iota offset(0) step(16843009) period(64) : tensor<408xi32> + %output = check.generate.fill value(-7.0) : tensor<6xf32> + %expected = check.generate.fill value(7.0) : tensor<6xf32> + kernel.launch @ggml_kquant_swiglu_decode_reference_f32[](%input, %gate, %up, %expected) : [](tensor<512xf32>, tensor<330xi32>, tensor<408xi32>, tensor<6xf32>) + kernel.launch @ggml_kquant_swiglu_decode_f32[](%input, %gate, %up, %output) : [](tensor<512xf32>, tensor<330xi32>, tensor<408xi32>, tensor<6xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<6xf32> + check.return +} + +// --config ... ggml.kquant_mul_mat_decode.weight_format=21 add=1 (IQ3_S: 6 rows x 2 x 110 B) +check.case public @ggml_kquant_mul_mat_decode_iq3s_case { + %x_seed = check.param.seed base(7300000000000031008) count(1) : i64 + %a_seed = check.param.seed base(7300000000000031009) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<512xf32> + %weight = check.generate.iota offset(0) step(16843009) period(64) : tensor<330xi32> + %addend = check.generate.random.uniform seed(%a_seed) range(-1.0 to 1.0) : tensor<6xf32> + %output = check.generate.fill value(-7.0) : tensor<6xf32> + %expected = check.generate.fill value(7.0) : tensor<6xf32> + kernel.launch @ggml_kquant_mul_mat_decode_reference_f32[](%input, %weight, %addend, %expected) : [](tensor<512xf32>, tensor<330xi32>, tensor<6xf32>, tensor<6xf32>) + kernel.launch @ggml_kquant_mul_mat_decode_f32[](%input, %weight, %addend, %output) : [](tensor<512xf32>, tensor<330xi32>, tensor<6xf32>, tensor<6xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<6xf32> + check.return +} + +// Multi-token cases (3 tokens). Run with --config=ggml.kquant_decode.token_count=3 and the formats +// named in each case. + +// gate_weight_format=5 up_weight_format=23 +check.case public @ggml_kquant_swiglu_decode_tokens_q5k_iq4xs_case { + %x_seed = check.param.seed base(7300000000000031101) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<1536xf32> + %gate = check.generate.iota offset(0) step(16843009) period(64) : tensor<528xi32> + %up = check.generate.iota offset(0) step(16843009) period(64) : tensor<408xi32> + %output = check.generate.fill value(-7.0) : tensor<18xf32> + %expected = check.generate.fill value(7.0) : tensor<18xf32> + kernel.launch @ggml_kquant_swiglu_decode_reference_f32[](%input, %gate, %up, %expected) : [](tensor<1536xf32>, tensor<528xi32>, tensor<408xi32>, tensor<18xf32>) + kernel.launch @ggml_kquant_swiglu_decode_tokens_f32[](%input, %gate, %up, %output) : [](tensor<1536xf32>, tensor<528xi32>, tensor<408xi32>, tensor<18xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<18xf32> + check.return +} + +// gate_weight_format=6 up_weight_format=80 +check.case public @ggml_kquant_swiglu_decode_tokens_q6k_q8_0_case { + %x_seed = check.param.seed base(7300000000000031102) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<1536xf32> + %gate = check.generate.iota offset(0) step(16843009) period(64) : tensor<630xi32> + %up = check.generate.iota offset(0) step(16843009) period(64) : tensor<816xi32> + %output = check.generate.fill value(-7.0) : tensor<18xf32> + %expected = check.generate.fill value(7.0) : tensor<18xf32> + kernel.launch @ggml_kquant_swiglu_decode_reference_f32[](%input, %gate, %up, %expected) : [](tensor<1536xf32>, tensor<630xi32>, tensor<816xi32>, tensor<18xf32>) + kernel.launch @ggml_kquant_swiglu_decode_tokens_f32[](%input, %gate, %up, %output) : [](tensor<1536xf32>, tensor<630xi32>, tensor<816xi32>, tensor<18xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<18xf32> + check.return +} + +// gate_weight_format=11 up_weight_format=4 +check.case public @ggml_kquant_swiglu_decode_tokens_q3k_q4k_case { + %x_seed = check.param.seed base(7300000000000031103) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<1536xf32> + %gate = check.generate.iota offset(0) step(16843009) period(64) : tensor<330xi32> + %up = check.generate.iota offset(0) step(16843009) period(64) : tensor<432xi32> + %output = check.generate.fill value(-7.0) : tensor<18xf32> + %expected = check.generate.fill value(7.0) : tensor<18xf32> + kernel.launch @ggml_kquant_swiglu_decode_reference_f32[](%input, %gate, %up, %expected) : [](tensor<1536xf32>, tensor<330xi32>, tensor<432xi32>, tensor<18xf32>) + kernel.launch @ggml_kquant_swiglu_decode_tokens_f32[](%input, %gate, %up, %output) : [](tensor<1536xf32>, tensor<330xi32>, tensor<432xi32>, tensor<18xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<18xf32> + check.return +} + +// ggml.kquant_mul_mat_decode.weight_format=21 add=1 (IQ3_S) +check.case public @ggml_kquant_mul_mat_decode_tokens_iq3s_case { + %x_seed = check.param.seed base(7300000000000031104) count(1) : i64 + %a_seed = check.param.seed base(7300000000000031105) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<1536xf32> + %weight = check.generate.iota offset(0) step(16843009) period(64) : tensor<330xi32> + %addend = check.generate.random.uniform seed(%a_seed) range(-1.0 to 1.0) : tensor<18xf32> + %output = check.generate.fill value(-7.0) : tensor<18xf32> + %expected = check.generate.fill value(7.0) : tensor<18xf32> + kernel.launch @ggml_kquant_mul_mat_decode_reference_f32[](%input, %weight, %addend, %expected) : [](tensor<1536xf32>, tensor<330xi32>, tensor<18xf32>, tensor<18xf32>) + kernel.launch @ggml_kquant_mul_mat_decode_tokens_f32[](%input, %weight, %addend, %output) : [](tensor<1536xf32>, tensor<330xi32>, tensor<18xf32>, tensor<18xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<18xf32> + check.return +} + +// ggml.kquant_mul_mat_decode.weight_format=20 add=0 (IQ4_NL) +check.case public @ggml_kquant_mul_mat_decode_tokens_iq4nl_case { + %x_seed = check.param.seed base(7300000000000031106) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<1536xf32> + %weight = check.generate.iota offset(0) step(16843009) period(64) : tensor<432xi32> + %addend = check.generate.fill value(100.0) : tensor<18xf32> + %output = check.generate.fill value(-7.0) : tensor<18xf32> + %expected = check.generate.fill value(7.0) : tensor<18xf32> + kernel.launch @ggml_kquant_mul_mat_decode_reference_f32[](%input, %weight, %addend, %expected) : [](tensor<1536xf32>, tensor<432xi32>, tensor<18xf32>, tensor<18xf32>) + kernel.launch @ggml_kquant_mul_mat_decode_tokens_f32[](%input, %weight, %addend, %output) : [](tensor<1536xf32>, tensor<432xi32>, tensor<18xf32>, tensor<18xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<18xf32> + check.return +} + +// PrismML formats for the reference kernels: value %i (0..255) of the 256-value block at %block_base, one byte at +// a time in dequantize_row_pq2_0 / dequantize_row_ptq1_0 order (ggml-prism-quants.c). +func.def inline @ggml_kquant_g128_ref_value(%format: index, %weight: buffer, %block_base: offset, %i: index) -> (f32) { + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c13 = index.constant 13 : index + %c14 = index.constant 14 : index + %c16 = index.constant 16 : index + %c17 = index.constant 17 : index + %c24 = index.constant 24 : index + %c28 = index.constant 28 : index + %c34 = index.constant 34 : index + %c80 = index.constant 80 : index + %c120 = index.constant 120 : index + %c128 = index.constant 128 : index + %f72 = index.constant 72 : index + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c255_i32 = scalar.constant 255 : i32 + %c0_i32r = scalar.constant 0 : i32 + %g = index.div %i, %c128 : index + %j = index.rem %i, %c128 : index + %f10r = index.constant 10 : index + %is10r = index.cmp eq, %format, %f10r : index + %is72 = index.cmp eq, %format, %f72 : index + %v = scf.if %is10r -> (f32) { + %c9 = index.constant 9 : index + %c18 = index.constant 18 : index + %bv = buffer.view %weight[%block_base] : buffer -> view<36xi8> + %hv = buffer.view %weight[%block_base] : buffer -> view<18xf16> + %d_at = index.mul %g, %c9 : index + %d_f16 = view.load %hv[%d_at] : view<18xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %gb = index.mul %g, %c18 : index + %jb = index.div %j, %c8 : index + %b_at0 = index.add %gb, %c2 : index + %b_at1 = index.add %b_at0, %jb : index + %b_at = index.assume %b_at1 [range(%b_at1, 0, 35)] : index + %byte_i8 = view.load %bv[%b_at] : view<36xi8> -> i8 + %byte = scalar.extui %byte_i8 : i8 to i32 + %jm = index.rem %j, %c8 : index + %jm_i32 = index.cast %jm : index to i32 + %s = scalar.shrui %byte, %jm_i32 : i32 + %bit = scalar.andi %s, %c1_i32 : i32 + %set = scalar.cmpi ne, %bit, %c0_i32r : i32 + %neg_d = scalar.negf %d : f32 + %r = scf.select %set, %d, %neg_d : f32 + scf.yield %r : f32 + } else { + %v72 = scf.if %is72 -> (f32) { + %bv = buffer.view %weight[%block_base] : buffer -> view<68xi8> + %hv = buffer.view %weight[%block_base] : buffer -> view<34xf16> + %d_at = index.mul %g, %c17 : index + %d_f16 = view.load %hv[%d_at] : view<34xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %gb = index.mul %g, %c34 : index + %jb = index.div %j, %c4 : index + %b_at0 = index.add %gb, %c2 : index + %b_at1 = index.add %b_at0, %jb : index + %b_at = index.assume %b_at1 [range(%b_at1, 0, 67)] : index + %byte_i8 = view.load %bv[%b_at] : view<68xi8> -> i8 + %byte = scalar.extui %byte_i8 : i8 to i32 + %jm = index.rem %j, %c4 : index + %jm_i32 = index.cast %jm : index to i32 + %sh = scalar.muli %jm_i32, %c2_i32 : i32 + %s = scalar.shrui %byte, %sh : i32 + %code = scalar.andi %s, %c3_i32 : i32 + %t = scalar.subi %code, %c1_i32 : i32 + %tf = scalar.sitofp %t : i32 to f32 + %r = scalar.mulf %tf, %d : f32 + scf.yield %r : f32 + } else { + %bv = buffer.view %weight[%block_base] : buffer -> view<56xi8> + %hv = buffer.view %weight[%block_base] : buffer -> view<28xf16> + %d_at0 = index.mul %g, %c14 : index + %d_at = index.add %d_at0, %c13 : index + %d_f16 = view.load %hv[%d_at] : view<28xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + // byte and trit of value j: j < 80: (j % 16, j / 16); j < 120: (16 + (j - 80) % 8, (j - 80) / 8); + // else (24 + (j - 120) % 2, (j - 120) / 2) + %in_a = index.cmp ult, %j, %c80 : index + %in_c = index.cmp uge, %j, %c120 : index + %jb0 = scf.select %in_a, %c80, %j : index + %jb = index.sub %jb0, %c80 : index + %jc0 = scf.select %in_c, %j, %c120 : index + %jc = index.sub %jc0, %c120 : index + %byte_a = index.rem %j, %c16 : index + %n_a = index.div %j, %c16 : index + %byte_b0 = index.rem %jb, %c8 : index + %byte_b = index.add %byte_b0, %c16 : index + %n_b = index.div %jb, %c8 : index + %byte_c0 = index.rem %jc, %c2 : index + %byte_c = index.add %byte_c0, %c24 : index + %n_c = index.div %jc, %c2 : index + %byte_ab = scf.select %in_a, %byte_a, %byte_b : index + %byte_in = scf.select %in_c, %byte_c, %byte_ab : index + %n_ab = scf.select %in_a, %n_a, %n_b : index + %n = scf.select %in_c, %n_c, %n_ab : index + %gb = index.mul %g, %c28 : index + %b_at0 = index.add %gb, %byte_in : index + %b_at = index.assume %b_at0 [range(%b_at0, 0, 55)] : index + %byte_i8 = view.load %bv[%b_at] : view<56xi8> -> i8 + %byte = scalar.extui %byte_i8 : i8 to i32 + // 3^n, n = 0..4 + %c1 = index.constant 1 : index + %c3 = index.constant 3 : index + %c9_i32 = scalar.constant 9 : i32 + %c27_i32 = scalar.constant 27 : i32 + %c81_i32 = scalar.constant 81 : i32 + %n1 = index.cmp eq, %n, %c1 : index + %n2 = index.cmp eq, %n, %c2 : index + %n3 = index.cmp eq, %n, %c3 : index + %n4 = index.cmp eq, %n, %c4 : index + %pw1 = scf.select %n1, %c3_i32, %c1_i32 : i32 + %pw2 = scf.select %n2, %c9_i32, %pw1 : i32 + %pw3 = scf.select %n3, %c27_i32, %pw2 : i32 + %pow = scf.select %n4, %c81_i32, %pw3 : i32 + %m = scalar.muli %byte, %pow : i32 + %q = scalar.andi %m, %c255_i32 : i32 + %q3 = scalar.muli %q, %c3_i32 : i32 + %u = scalar.shrui %q3, %c8_i32 : i32 + %t = scalar.subi %u, %c1_i32 : i32 + %tf = scalar.sitofp %t : i32 to f32 + %r = scalar.mulf %tf, %d : f32 + scf.yield %r : f32 + } + scf.yield %v72 : f32 + } + func.return %v : f32 +} + +// PrismML cases: 512 inputs (two 256-value blocks per row), 6 rows, the integer-ramp weights of the cases above +// (every byte of word w is w % 64: finite fp16 scales; PQ2_0 codes include 3, i.e. +2). PQ2_0 is 68 bytes per 256 +// values (6 x 2 x 68 B = 204 words), PTQ1_0 56 (168 words). + +// --config=ggml.kquant_mul_mat_decode.input_size=512 --config=ggml.kquant_mul_mat_decode.output_size=6 +// --config=ggml.kquant_mul_mat_decode.weight_format=73 --config=ggml.kquant_mul_mat_decode.add=1 +check.case public @ggml_kquant_mul_mat_decode_ptq1_add_case { + %x_seed = check.param.seed base(7300000000000031201) count(1) : i64 + %a_seed = check.param.seed base(7300000000000031202) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<512xf32> + %weight = check.generate.iota offset(0) step(16843009) period(64) : tensor<168xi32> + %addend = check.generate.random.uniform seed(%a_seed) range(-1.0 to 1.0) : tensor<6xf32> + %output = check.generate.fill value(-7.0) : tensor<6xf32> + %expected = check.generate.fill value(7.0) : tensor<6xf32> + kernel.launch @ggml_kquant_mul_mat_decode_reference_f32[](%input, %weight, %addend, %expected) : [](tensor<512xf32>, tensor<168xi32>, tensor<6xf32>, tensor<6xf32>) + kernel.launch @ggml_kquant_mul_mat_decode_f32[](%input, %weight, %addend, %output) : [](tensor<512xf32>, tensor<168xi32>, tensor<6xf32>, tensor<6xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<6xf32> + check.return +} + +// --config ... ggml.kquant_mul_mat_decode.weight_format=72 add=0 +check.case public @ggml_kquant_mul_mat_decode_pq2_case { + %x_seed = check.param.seed base(7300000000000031203) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<512xf32> + %weight = check.generate.iota offset(0) step(16843009) period(64) : tensor<204xi32> + %addend = check.generate.fill value(100.0) : tensor<6xf32> + %output = check.generate.fill value(-7.0) : tensor<6xf32> + %expected = check.generate.fill value(7.0) : tensor<6xf32> + kernel.launch @ggml_kquant_mul_mat_decode_reference_f32[](%input, %weight, %addend, %expected) : [](tensor<512xf32>, tensor<204xi32>, tensor<6xf32>, tensor<6xf32>) + kernel.launch @ggml_kquant_mul_mat_decode_f32[](%input, %weight, %addend, %output) : [](tensor<512xf32>, tensor<204xi32>, tensor<6xf32>, tensor<6xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<6xf32> + check.return +} + +// --config ... input_size=512 output_size=6 gate_weight_format=73 up_weight_format=72 +check.case public @ggml_kquant_swiglu_decode_ptq1_pq2_case { + %x_seed = check.param.seed base(7300000000000031204) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<512xf32> + %gate = check.generate.iota offset(0) step(16843009) period(64) : tensor<168xi32> + %up = check.generate.iota offset(0) step(16843009) period(64) : tensor<204xi32> + %output = check.generate.fill value(-7.0) : tensor<6xf32> + %expected = check.generate.fill value(7.0) : tensor<6xf32> + kernel.launch @ggml_kquant_swiglu_decode_reference_f32[](%input, %gate, %up, %expected) : [](tensor<512xf32>, tensor<168xi32>, tensor<204xi32>, tensor<6xf32>) + kernel.launch @ggml_kquant_swiglu_decode_f32[](%input, %gate, %up, %output) : [](tensor<512xf32>, tensor<168xi32>, tensor<204xi32>, tensor<6xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<6xf32> + check.return +} + +// Multi-token (--config=ggml.kquant_decode.token_count=3): gate_weight_format=72 up_weight_format=73 +check.case public @ggml_kquant_swiglu_decode_tokens_pq2_ptq1_case { + %x_seed = check.param.seed base(7300000000000031205) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<1536xf32> + %gate = check.generate.iota offset(0) step(16843009) period(64) : tensor<204xi32> + %up = check.generate.iota offset(0) step(16843009) period(64) : tensor<168xi32> + %output = check.generate.fill value(-7.0) : tensor<18xf32> + %expected = check.generate.fill value(7.0) : tensor<18xf32> + kernel.launch @ggml_kquant_swiglu_decode_reference_f32[](%input, %gate, %up, %expected) : [](tensor<1536xf32>, tensor<204xi32>, tensor<168xi32>, tensor<18xf32>) + kernel.launch @ggml_kquant_swiglu_decode_tokens_f32[](%input, %gate, %up, %output) : [](tensor<1536xf32>, tensor<204xi32>, tensor<168xi32>, tensor<18xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<18xf32> + check.return +} + +// Multi-token: ggml.kquant_mul_mat_decode.weight_format=73 add=1 +check.case public @ggml_kquant_mul_mat_decode_tokens_ptq1_case { + %x_seed = check.param.seed base(7300000000000031206) count(1) : i64 + %a_seed = check.param.seed base(7300000000000031207) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<1536xf32> + %weight = check.generate.iota offset(0) step(16843009) period(64) : tensor<168xi32> + %addend = check.generate.random.uniform seed(%a_seed) range(-1.0 to 1.0) : tensor<18xf32> + %output = check.generate.fill value(-7.0) : tensor<18xf32> + %expected = check.generate.fill value(7.0) : tensor<18xf32> + kernel.launch @ggml_kquant_mul_mat_decode_reference_f32[](%input, %weight, %addend, %expected) : [](tensor<1536xf32>, tensor<168xi32>, tensor<18xf32>, tensor<18xf32>) + kernel.launch @ggml_kquant_mul_mat_decode_tokens_f32[](%input, %weight, %addend, %output) : [](tensor<1536xf32>, tensor<168xi32>, tensor<18xf32>, tensor<18xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<18xf32> + check.return +} + +// --config ... ggml.kquant_mul_mat_decode.weight_format=10 add=1 (Q1_0: 6 x 2 x 36 B = 108 words) +check.case public @ggml_kquant_mul_mat_decode_q1_0_add_case { + %x_seed = check.param.seed base(7300000000000031208) count(1) : i64 + %a_seed = check.param.seed base(7300000000000031209) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<512xf32> + %weight = check.generate.iota offset(0) step(16843009) period(64) : tensor<108xi32> + %addend = check.generate.random.uniform seed(%a_seed) range(-1.0 to 1.0) : tensor<6xf32> + %output = check.generate.fill value(-7.0) : tensor<6xf32> + %expected = check.generate.fill value(7.0) : tensor<6xf32> + kernel.launch @ggml_kquant_mul_mat_decode_reference_f32[](%input, %weight, %addend, %expected) : [](tensor<512xf32>, tensor<108xi32>, tensor<6xf32>, tensor<6xf32>) + kernel.launch @ggml_kquant_mul_mat_decode_f32[](%input, %weight, %addend, %output) : [](tensor<512xf32>, tensor<108xi32>, tensor<6xf32>, tensor<6xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<6xf32> + check.return +} + +// Multi-token (token_count=3): gate_weight_format=10 up_weight_format=73 +check.case public @ggml_kquant_swiglu_decode_tokens_q1_0_ptq1_case { + %x_seed = check.param.seed base(7300000000000031210) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<1536xf32> + %gate = check.generate.iota offset(0) step(16843009) period(64) : tensor<108xi32> + %up = check.generate.iota offset(0) step(16843009) period(64) : tensor<168xi32> + %output = check.generate.fill value(-7.0) : tensor<18xf32> + %expected = check.generate.fill value(7.0) : tensor<18xf32> + kernel.launch @ggml_kquant_swiglu_decode_reference_f32[](%input, %gate, %up, %expected) : [](tensor<1536xf32>, tensor<108xi32>, tensor<168xi32>, tensor<18xf32>) + kernel.launch @ggml_kquant_swiglu_decode_tokens_f32[](%input, %gate, %up, %output) : [](tensor<1536xf32>, tensor<108xi32>, tensor<168xi32>, tensor<18xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<18xf32> + check.return +} + +// Mixed pairs with a codebook format (two-buffer staging of #60): input_size=512 output_size=6. +// gate_weight_format=21 up_weight_format=73 (IQ3_S gate grid; the PTQ1_0 up lanes read no grid) +check.case public @ggml_kquant_swiglu_decode_iq3s_ptq1_case { + %x_seed = check.param.seed base(7300000000000031211) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<512xf32> + %gate = check.generate.iota offset(0) step(16843009) period(64) : tensor<330xi32> + %up = check.generate.iota offset(0) step(16843009) period(64) : tensor<168xi32> + %output = check.generate.fill value(-7.0) : tensor<6xf32> + %expected = check.generate.fill value(7.0) : tensor<6xf32> + kernel.launch @ggml_kquant_swiglu_decode_reference_f32[](%input, %gate, %up, %expected) : [](tensor<512xf32>, tensor<330xi32>, tensor<168xi32>, tensor<6xf32>) + kernel.launch @ggml_kquant_swiglu_decode_f32[](%input, %gate, %up, %output) : [](tensor<512xf32>, tensor<330xi32>, tensor<168xi32>, tensor<6xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<6xf32> + check.return +} + +// gate_weight_format=73 up_weight_format=21 (the IQ3_S up fills its own grid buffer); token_count=3 +check.case public @ggml_kquant_swiglu_decode_tokens_ptq1_iq3s_case { + %x_seed = check.param.seed base(7300000000000031212) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<1536xf32> + %gate = check.generate.iota offset(0) step(16843009) period(64) : tensor<168xi32> + %up = check.generate.iota offset(0) step(16843009) period(64) : tensor<330xi32> + %output = check.generate.fill value(-7.0) : tensor<18xf32> + %expected = check.generate.fill value(7.0) : tensor<18xf32> + kernel.launch @ggml_kquant_swiglu_decode_reference_f32[](%input, %gate, %up, %expected) : [](tensor<1536xf32>, tensor<168xi32>, tensor<330xi32>, tensor<18xf32>) + kernel.launch @ggml_kquant_swiglu_decode_tokens_f32[](%input, %gate, %up, %output) : [](tensor<1536xf32>, tensor<168xi32>, tensor<330xi32>, tensor<18xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<18xf32> + check.return +} + +// gate_weight_format=73 up_weight_format=21, single token +check.case public @ggml_kquant_swiglu_decode_ptq1_iq3s_case { + %x_seed = check.param.seed base(7300000000000031213) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<512xf32> + %gate = check.generate.iota offset(0) step(16843009) period(64) : tensor<168xi32> + %up = check.generate.iota offset(0) step(16843009) period(64) : tensor<330xi32> + %output = check.generate.fill value(-7.0) : tensor<6xf32> + %expected = check.generate.fill value(7.0) : tensor<6xf32> + kernel.launch @ggml_kquant_swiglu_decode_reference_f32[](%input, %gate, %up, %expected) : [](tensor<512xf32>, tensor<168xi32>, tensor<330xi32>, tensor<6xf32>) + kernel.launch @ggml_kquant_swiglu_decode_f32[](%input, %gate, %up, %output) : [](tensor<512xf32>, tensor<168xi32>, tensor<330xi32>, tensor<6xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<6xf32> + check.return +} + +// TQ1_0 (34), TQ2_0 (35), MXFP4 (39) against ggml_kquant_ref_value_1bit. Weight words per case at +// input_size 512, output_size 6: TQ1_0 6 x 2 x 54 B = 162, TQ2_0 6 x 2 x 66 B = 198, MXFP4 6 x 2 x 136 B +// = 408. The MXFP4 cases fill bytes 121..127 (0x79797979 + k * 0x01010101 stays a positive i32) so every +// E8M0 scale is normal and between 2^-7 and 2^-1 (the GPU flushes subnormals; tiny scales would also pass +// the tolerance whatever the codes decode to). + +// --config=ggml.kquant_mul_mat_decode.input_size=512 --config=ggml.kquant_mul_mat_decode.output_size=6 +// --config=ggml.kquant_mul_mat_decode.weight_format=34 --config=ggml.kquant_mul_mat_decode.add=1 +check.case public @ggml_kquant_mul_mat_decode_tq1_0_add_case { + %x_seed = check.param.seed base(7300000000000031301) count(1) : i64 + %a_seed = check.param.seed base(7300000000000031302) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<512xf32> + %weight = check.generate.iota offset(0) step(16843009) period(64) : tensor<162xi32> + %addend = check.generate.random.uniform seed(%a_seed) range(-1.0 to 1.0) : tensor<6xf32> + %output = check.generate.fill value(-7.0) : tensor<6xf32> + %expected = check.generate.fill value(7.0) : tensor<6xf32> + kernel.launch @ggml_kquant_mul_mat_decode_reference_f32[](%input, %weight, %addend, %expected) : [](tensor<512xf32>, tensor<162xi32>, tensor<6xf32>, tensor<6xf32>) + kernel.launch @ggml_kquant_mul_mat_decode_f32[](%input, %weight, %addend, %output) : [](tensor<512xf32>, tensor<162xi32>, tensor<6xf32>, tensor<6xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<6xf32> + check.return +} + +// --config ... ggml.kquant_mul_mat_decode.weight_format=35 add=0 +check.case public @ggml_kquant_mul_mat_decode_tq2_0_case { + %x_seed = check.param.seed base(7300000000000031303) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<512xf32> + %weight = check.generate.iota offset(0) step(16843009) period(64) : tensor<198xi32> + %addend = check.generate.fill value(100.0) : tensor<6xf32> + %output = check.generate.fill value(-7.0) : tensor<6xf32> + %expected = check.generate.fill value(7.0) : tensor<6xf32> + kernel.launch @ggml_kquant_mul_mat_decode_reference_f32[](%input, %weight, %addend, %expected) : [](tensor<512xf32>, tensor<198xi32>, tensor<6xf32>, tensor<6xf32>) + kernel.launch @ggml_kquant_mul_mat_decode_f32[](%input, %weight, %addend, %output) : [](tensor<512xf32>, tensor<198xi32>, tensor<6xf32>, tensor<6xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<6xf32> + check.return +} + +// --config ... ggml.kquant_mul_mat_decode.weight_format=39 add=0 +check.case public @ggml_kquant_mul_mat_decode_mxfp4_case { + %x_seed = check.param.seed base(7300000000000031304) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<512xf32> + %weight = check.generate.iota offset(2038003065) step(16843009) period(7) : tensor<408xi32> + %addend = check.generate.fill value(100.0) : tensor<6xf32> + %output = check.generate.fill value(-7.0) : tensor<6xf32> + %expected = check.generate.fill value(7.0) : tensor<6xf32> + kernel.launch @ggml_kquant_mul_mat_decode_reference_f32[](%input, %weight, %addend, %expected) : [](tensor<512xf32>, tensor<408xi32>, tensor<6xf32>, tensor<6xf32>) + kernel.launch @ggml_kquant_mul_mat_decode_f32[](%input, %weight, %addend, %output) : [](tensor<512xf32>, tensor<408xi32>, tensor<6xf32>, tensor<6xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<6xf32> + check.return +} + +// Multi-token: ggml.kquant_mul_mat_decode.weight_format=34 add=1 +check.case public @ggml_kquant_mul_mat_decode_tokens_tq1_0_case { + %x_seed = check.param.seed base(7300000000000031305) count(1) : i64 + %a_seed = check.param.seed base(7300000000000031306) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<1536xf32> + %weight = check.generate.iota offset(0) step(16843009) period(64) : tensor<162xi32> + %addend = check.generate.random.uniform seed(%a_seed) range(-1.0 to 1.0) : tensor<18xf32> + %output = check.generate.fill value(-7.0) : tensor<18xf32> + %expected = check.generate.fill value(7.0) : tensor<18xf32> + kernel.launch @ggml_kquant_mul_mat_decode_reference_f32[](%input, %weight, %addend, %expected) : [](tensor<1536xf32>, tensor<162xi32>, tensor<18xf32>, tensor<18xf32>) + kernel.launch @ggml_kquant_mul_mat_decode_tokens_f32[](%input, %weight, %addend, %output) : [](tensor<1536xf32>, tensor<162xi32>, tensor<18xf32>, tensor<18xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<18xf32> + check.return +} + +// Multi-token: ggml.kquant_mul_mat_decode.weight_format=35 add=0 +check.case public @ggml_kquant_mul_mat_decode_tokens_tq2_0_case { + %x_seed = check.param.seed base(7300000000000031307) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<1536xf32> + %weight = check.generate.iota offset(0) step(16843009) period(64) : tensor<198xi32> + %addend = check.generate.fill value(100.0) : tensor<18xf32> + %output = check.generate.fill value(-7.0) : tensor<18xf32> + %expected = check.generate.fill value(7.0) : tensor<18xf32> + kernel.launch @ggml_kquant_mul_mat_decode_reference_f32[](%input, %weight, %addend, %expected) : [](tensor<1536xf32>, tensor<198xi32>, tensor<18xf32>, tensor<18xf32>) + kernel.launch @ggml_kquant_mul_mat_decode_tokens_f32[](%input, %weight, %addend, %output) : [](tensor<1536xf32>, tensor<198xi32>, tensor<18xf32>, tensor<18xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<18xf32> + check.return +} + +// Multi-token: ggml.kquant_mul_mat_decode.weight_format=39 add=0 +check.case public @ggml_kquant_mul_mat_decode_tokens_mxfp4_case { + %x_seed = check.param.seed base(7300000000000031308) count(1) : i64 + %input = check.generate.random.uniform seed(%x_seed) range(-1.0 to 1.0) : tensor<1536xf32> + %weight = check.generate.iota offset(2038003065) step(16843009) period(7) : tensor<408xi32> + %addend = check.generate.fill value(100.0) : tensor<18xf32> + %output = check.generate.fill value(-7.0) : tensor<18xf32> + %expected = check.generate.fill value(7.0) : tensor<18xf32> + kernel.launch @ggml_kquant_mul_mat_decode_reference_f32[](%input, %weight, %addend, %expected) : [](tensor<1536xf32>, tensor<408xi32>, tensor<18xf32>, tensor<18xf32>) + kernel.launch @ggml_kquant_mul_mat_decode_tokens_f32[](%input, %weight, %addend, %output) : [](tensor<1536xf32>, tensor<408xi32>, tensor<18xf32>, tensor<18xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-3) rtol(1.0e-4) nan(same) : tensor<18xf32> + check.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_k_matmul_rope_set_rows_decode_f32_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_k_matmul_rope_set_rows_decode_f32_f32.loom new file mode 100644 index 000000000000..aa2591431cfe --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_k_matmul_rope_set_rows_decode_f32_f32.loom @@ -0,0 +1,110 @@ +template.decl @ggml.mul_mat_f32_f32_decode.dispatch(%weight_format: index, %token_count: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %output: buffer, %aux0: buffer, %aux1: buffer, %aux2: buffer, %aux3: buffer) + +template.decl @ggml.mul_mat_f32_f32_decode.publish_pair2(%publish_pair: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %row1_valid: i1, %value0: f32, %value1: f32, %cache: buffer, %positions: buffer, %indices: buffer, %theta: buffer, %freq_factors: buffer) + +func.decl @ggml_rope_f32_pair_packet(%position: f32, %theta: vector<2xf32>, %freq_factors: vector<2xf32>, %x_values: vector<2xf32>, %y_values: vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) + +amdgpu.target @llm_attention_k_matmul_rope_set_rows_decode_gfx11_wave64 {subgroup_size = 64} + +config.decl @llm.attention_qkv.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 32)] + +config.decl @llm.attention_qkv.output_size : %value: index where [range(%value, 4, 32768), mul(%value, 4)] + +config.decl @llm.attention_qkv.weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @llm.attention_qkv.head_size : %value: index where [range(%value, 4, 1024), mul(%value, 4)] + +config.decl @llm.attention_qkv.head_count : %value: index where [range(%value, 1, 64)] + +config.decl @llm.attention_qkv.cache_row_count : %value: index where [range(%value, 1, 1048576)] + +config.decl @llm.attention_qkv.cache_output_format : %value: index where [range(%value, 16, 32)] + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +template.def<@ggml.mul_mat_f32_f32_decode.publish_pair2> device priority(20) @llm_attention_k_matmul_rope_set_rows_decode_publish_pair2(%publish_pair: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %row1_valid: i1, %value0: f32, %value1: f32, %cache: buffer, %positions: buffer, %indices: buffer, %theta: buffer, %freq_factors: buffer) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 4, 32768), mul(%output_size0, 4)] : index + %head_size0 = config.get @llm.attention_qkv.head_size : index + %head_size = index.assume %head_size0 [range(%head_size0, 4, 1024), mul(%head_size0, 4)] : index + %cache_row_count0 = config.get @llm.attention_qkv.cache_row_count : index + %cache_row_count = index.assume %cache_row_count0 [range(%cache_row_count0, 1, 1048576)] : index + %cache_output_format = config.get @llm.attention_qkv.cache_output_format : index + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c0_i64 = scalar.constant 0 : i64 + %c1048575_i64 = scalar.constant 1048575 : i64 + %c0_offset = index.constant 0 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %publish0 = scalar.andi %publish_pair, %row1_valid : i1 + %token, %launch_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %head_channel = index.rem %channel, %head_size : index + %theta_channel = index.div %head_channel, %c2 : index + %half_head_size = index.div %head_size, %c2 : index + %is_f16 = index.cmp eq, %cache_output_format, %c16 : index + %is_f32 = index.cmp eq, %cache_output_format, %c32 : index + %positions_noalias, %indices_noalias, %theta_noalias, %freq_factors_noalias, %cache_noalias = buffer.assume.noalias %positions, %indices, %theta, %freq_factors, %cache : buffer, buffer, buffer, buffer, buffer + %indices_view = buffer.view %indices_noalias[%c0_offset] : buffer -> view<[%launch_token_count]xi64> + %index_raw = view.load %indices_view[%token] : view<[%launch_token_count]xi64> -> i64 + %index_nonnegative = scalar.cmpi sge, %index_raw, %c0_i64 : i64 + %index_in_cast_range = scalar.cmpi sle, %index_raw, %c1048575_i64 : i64 + %valid_index = scalar.andi %index_nonnegative, %index_in_cast_range : i1 + %safe_index0_i64 = scf.select %valid_index, %index_raw, %c0_i64 : i64 + %safe_index_i64 = scalar.assume %safe_index0_i64 [range(%safe_index0_i64, 0, 1048575)] : i64 + %cache_row0 = index.cast %safe_index_i64 : i64 to index + %valid_row = index.cmp ult, %cache_row0, %cache_row_count : index + %cache_row = scf.select %valid_row, %cache_row0, %c0 : index + %publish1 = scalar.andi %publish0, %valid_index : i1 + %publish = scalar.andi %publish1, %valid_row : i1 + scf.if %publish { + %positions_view = buffer.view %positions_noalias[%c0_offset] : buffer -> view<[%launch_token_count]xi32> + %theta_view = buffer.view %theta_noalias[%c0_offset] : buffer -> view<[%half_head_size]xf32> + %freq_factors_view = buffer.view %freq_factors_noalias[%c0_offset] : buffer -> view<[%half_head_size]xf32> + %position_i32 = view.load %positions_view[%token] : view<[%launch_token_count]xi32> -> i32 + %position = scalar.sitofp %position_i32 : i32 to f32 + %theta_scalar = view.load %theta_view[%theta_channel] : view<[%half_head_size]xf32> -> f32 + %freq_factors_scalar = view.load %freq_factors_view[%theta_channel] : view<[%half_head_size]xf32> -> f32 + %theta_packet = vector.splat %theta_scalar : vector<2xf32> + %freq_factors_packet = vector.splat %freq_factors_scalar : vector<2xf32> + %x_values = vector.from_elements %value0, %c0_f32 : vector<2xf32> + %y_values = vector.from_elements %value1, %c0_f32 : vector<2xf32> + %rotated_x, %rotated_y = func.call @ggml_rope_f32_pair_packet(%position, %theta_packet, %freq_factors_packet, %x_values, %y_values) : (f32, vector<2xf32>, vector<2xf32>, vector<2xf32>, vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) + %rx0 = vector.extract %rotated_x[0] : vector<2xf32> -> f32 + %ry0 = vector.extract %rotated_y[0] : vector<2xf32> -> f32 + %values = vector.from_elements %rx0, %ry0 : vector<2xf32> + %publish_f16 = scalar.andi %publish, %is_f16 : i1 + %publish_f32 = scalar.andi %publish, %is_f32 : i1 + scf.if %publish_f16 { + %cache_view = buffer.view %cache_noalias[%c0_offset] : buffer -> view<[%cache_row_count]x[%output_size]xf16> + %truncated = vector.fptrunc %values : vector<2xf32> to vector<2xf16> + vector.store %truncated, %cache_view[%cache_row, %channel] : vector<2xf16>, view<[%cache_row_count]x[%output_size]xf16> + } + scf.if %publish_f32 { + %cache_view = buffer.view %cache_noalias[%c0_offset] : buffer -> view<[%cache_row_count]x[%output_size]xf32> + vector.store %values, %cache_view[%cache_row, %channel] : vector<2xf32>, view<[%cache_row_count]x[%output_size]xf32> + } + } + template.return +} + +kernel.def target(@llm_attention_k_matmul_rope_set_rows_decode_gfx11_wave64) @llm_attention_k_matmul_rope_set_rows_decode_f32_f32(%token_count: index) { + %token_capacity = config.get @ggml.workload.token_capacity : index + %output_capacity = config.get @llm.attention_qkv.output_size : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %padded_output_size = index.add %output_capacity, %c1 : index + %output_pairs = index.div %padded_output_size, %c2 : index + %padded_token_count = index.add %token_capacity, %c63 : index + %launch_tokens = index.div %padded_token_count, %c64 : index + kernel.launch.config workgroups(%output_pairs, %launch_tokens, %c1) workgroup_size(%c64, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %positions: buffer, %indices: buffer, %theta: buffer, %freq_factors: buffer, %cache: buffer) { + %input_size = config.get @llm.attention_qkv.input_size : index + %output_size = config.get @llm.attention_qkv.output_size : index + %weight_format = config.get @llm.attention_qkv.weight_format : index + template.apply<@ggml.mul_mat_f32_f32_decode.dispatch>(%weight_format, %token_count, %input_size, %output_size, %input, %weight, %cache, %positions, %indices, %theta, %freq_factors) : (index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_k_matmul_rope_set_rows_f32_f32_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_k_matmul_rope_set_rows_f32_f32_wmma.loom new file mode 100644 index 000000000000..c53008364e4b --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_k_matmul_rope_set_rows_f32_f32_wmma.loom @@ -0,0 +1,67 @@ +template.decl @ggml.mul_mat_f32_f32_wmma.core(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: f32, %arg7: index, %arg8: index, %arg9: buffer, %arg10: buffer, %arg11: buffer, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer, %arg16: buffer, %arg17: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.launch(%token_capacity: index, %input_size: index, %output_size: index, %weight_format: index) -> (index, index, index, index) + +template.decl @ggml.mul_mat_f32_f32_wmma.publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %positions: buffer, %cache: buffer, %indices: buffer, %theta: buffer, %freq_factors: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @llm.attention_qkv.rope_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: buffer, %arg6: buffer, %arg7: buffer, %arg8: vector<4xf32>) -> (vector<4xf32>) + +template.decl @llm.attention_qkv.store_cache_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: buffer, %arg7: vector<4xf32>, %arg8: buffer) + +amdgpu.target @llm_attention_k_matmul_rope_set_rows_gfx11_wave64 {subgroup_size = 64} + +config.decl @llm.attention_qkv.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +config.decl @llm.attention_qkv.output_size : %value: index where [range(%value, 4, 32768), mul(%value, 4)] + +config.decl @llm.attention_qkv.weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @llm.attention_qkv.head_size : %value: index where [range(%value, 4, 1024), mul(%value, 4)] + +config.decl @llm.attention_qkv.head_count : %value: index where [range(%value, 1, 64)] + +config.decl @llm.attention_qkv.cache_row_count : %value: index where [range(%value, 1, 1048576)] + +config.decl @llm.attention_qkv.cache_output_format : %value: index where [range(%value, 16, 32)] + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +template.def<@ggml.mul_mat_f32_f32_wmma.publish_vector4> device @llm_attention_k_matmul_rope_set_rows_publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %positions: buffer, %cache: buffer, %indices: buffer, %theta: buffer, %freq_factors: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %head_size = config.get @llm.attention_qkv.head_size : index + %head_count = config.get @llm.attention_qkv.head_count : index + %cache_row_count = config.get @llm.attention_qkv.cache_row_count : index + %cache_output_format = config.get @llm.attention_qkv.cache_output_format : index + scf.if %publish_word { + %rotated = template.apply<@llm.attention_qkv.rope_vector4>(%token0, %channel, %token_count0, %head_size, %head_count, %positions, %theta, %freq_factors, %values) : (index, index, index, index, index, buffer, buffer, buffer, vector<4xf32>) -> (vector<4xf32>) + template.apply<@llm.attention_qkv.store_cache_vector4>(%cache_output_format, %token0, %channel, %token_count0, %cache_row_count, %output_size0, %indices, %rotated, %cache) : (index, index, index, index, index, index, buffer, vector<4xf32>, buffer) + } + template.return +} + +template.def<@ggml.mul_mat_f32_f32_wmma.finish_tile> device @llm_attention_k_matmul_rope_set_rows_finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + template.return +} + +kernel.def target(@llm_attention_k_matmul_rope_set_rows_gfx11_wave64) @llm_attention_k_matmul_rope_set_rows_f32_f32_wmma(%token_count: index) { + %output_size = config.get @llm.attention_qkv.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %c1 = index.constant 1 : index + %input_size = config.get @llm.attention_qkv.input_size : index + %weight_format = config.get @llm.attention_qkv.weight_format : index + %launch_x, %launch_y, %launch_z, %workgroup_size = template.apply<@ggml.mul_mat_f32_f32_wmma.launch>(%token_capacity, %input_size, %output_size, %weight_format) pure : (index, index, index, index) -> (index, index, index, index) + kernel.launch.config workgroups(%launch_x, %launch_y, %launch_z) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %positions: buffer, %indices: buffer, %theta: buffer, %freq_factors: buffer, %cache: buffer) { + %input_size = config.get @llm.attention_qkv.input_size : index + %output_size = config.get @llm.attention_qkv.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %weight_format = config.get @llm.attention_qkv.weight_format : index + %c0 = index.constant 0 : index + %c0_f32 = scalar.constant 0.0 : f32 + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_f32_f32_wmma.core>(%weight_format, %bounded_token_count, %input_size, %output_size, %c0, %c0, %c0_f32, %channel_tile, %token_tile, %input, %weight, %positions, %cache, %indices, %theta, %freq_factors, %cache, %cache) : (index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_q_matmul_rope_decode_f32_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_q_matmul_rope_decode_f32_f32.loom new file mode 100644 index 000000000000..a4f90842af5e --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_q_matmul_rope_decode_f32_f32.loom @@ -0,0 +1,78 @@ +template.decl @ggml.mul_mat_f32_f32_decode.dispatch(%weight_format: index, %token_count: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %output: buffer, %aux0: buffer, %aux1: buffer, %aux2: buffer, %aux3: buffer) + +template.decl @ggml.mul_mat_f32_f32_decode.publish_pair2(%publish_pair: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %row1_valid: i1, %value0: f32, %value1: f32, %output: buffer, %positions: buffer, %theta: buffer, %freq_factors: buffer, %aux3: buffer) + +func.decl @ggml_rope_f32_pair_packet(%position: f32, %theta: vector<2xf32>, %freq_factors: vector<2xf32>, %x_values: vector<2xf32>, %y_values: vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) + +amdgpu.target @llm_attention_q_matmul_rope_decode_gfx11_wave64 {subgroup_size = 64} + +config.decl @llm.attention_qkv.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 32)] + +config.decl @llm.attention_qkv.output_size : %value: index where [range(%value, 4, 32768), mul(%value, 4)] + +config.decl @llm.attention_qkv.weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @llm.attention_qkv.head_size : %value: index where [range(%value, 4, 1024), mul(%value, 4)] + +config.decl @llm.attention_qkv.head_count : %value: index where [range(%value, 1, 64)] + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +template.def<@ggml.mul_mat_f32_f32_decode.publish_pair2> device priority(20) @llm_attention_q_matmul_rope_decode_publish_pair2(%publish_pair: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %row1_valid: i1, %value0: f32, %value1: f32, %output: buffer, %positions: buffer, %theta: buffer, %freq_factors: buffer, %aux3: buffer) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 4, 32768), mul(%output_size0, 4)] : index + %head_size0 = config.get @llm.attention_qkv.head_size : index + %head_count0 = config.get @llm.attention_qkv.head_count : index + %head_size, %head_count = index.assume %head_size0, %head_count0 [range(%head_size0, 4, 1024), mul(%head_size0, 4), range(%head_count0, 1, 64)] : index, index + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c0_offset = index.constant 0 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %publish = scalar.andi %publish_pair, %row1_valid : i1 + %token, %launch_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %head = index.div %channel, %head_size : index + %head_channel = index.rem %channel, %head_size : index + %theta_channel = index.div %head_channel, %c2 : index + %half_head_size = index.div %head_size, %c2 : index + %positions_noalias, %theta_noalias, %freq_factors_noalias, %output_noalias = buffer.assume.noalias %positions, %theta, %freq_factors, %output : buffer, buffer, buffer, buffer + scf.if %publish { + %positions_view = buffer.view %positions_noalias[%c0_offset] : buffer -> view<[%launch_token_count]xi32> + %theta_view = buffer.view %theta_noalias[%c0_offset] : buffer -> view<[%half_head_size]xf32> + %freq_factors_view = buffer.view %freq_factors_noalias[%c0_offset] : buffer -> view<[%half_head_size]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%head_count]x[%head_size]xf32> + %position_i32 = view.load %positions_view[%token] : view<[%launch_token_count]xi32> -> i32 + %position = scalar.sitofp %position_i32 : i32 to f32 + %theta_scalar = view.load %theta_view[%theta_channel] : view<[%half_head_size]xf32> -> f32 + %freq_factors_scalar = view.load %freq_factors_view[%theta_channel] : view<[%half_head_size]xf32> -> f32 + %theta_packet = vector.splat %theta_scalar : vector<2xf32> + %freq_factors_packet = vector.splat %freq_factors_scalar : vector<2xf32> + %x_values = vector.from_elements %value0, %c0_f32 : vector<2xf32> + %y_values = vector.from_elements %value1, %c0_f32 : vector<2xf32> + %rotated_x, %rotated_y = func.call @ggml_rope_f32_pair_packet(%position, %theta_packet, %freq_factors_packet, %x_values, %y_values) : (f32, vector<2xf32>, vector<2xf32>, vector<2xf32>, vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) + %rx0 = vector.extract %rotated_x[0] : vector<2xf32> -> f32 + %ry0 = vector.extract %rotated_y[0] : vector<2xf32> -> f32 + %values = vector.from_elements %rx0, %ry0 : vector<2xf32> + vector.store %values, %output_view[%token, %head, %head_channel] : vector<2xf32>, view<[%launch_token_count]x[%head_count]x[%head_size]xf32> + } + template.return +} + +kernel.def target(@llm_attention_q_matmul_rope_decode_gfx11_wave64) @llm_attention_q_matmul_rope_decode_f32_f32(%token_count: index) { + %token_capacity = config.get @ggml.workload.token_capacity : index + %output_capacity = config.get @llm.attention_qkv.output_size : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %padded_output_size = index.add %output_capacity, %c1 : index + %output_pairs = index.div %padded_output_size, %c2 : index + %padded_token_count = index.add %token_capacity, %c63 : index + %launch_tokens = index.div %padded_token_count, %c64 : index + kernel.launch.config workgroups(%output_pairs, %launch_tokens, %c1) workgroup_size(%c64, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %positions: buffer, %theta: buffer, %freq_factors: buffer, %output: buffer) { + %input_size = config.get @llm.attention_qkv.input_size : index + %output_size = config.get @llm.attention_qkv.output_size : index + %weight_format = config.get @llm.attention_qkv.weight_format : index + template.apply<@ggml.mul_mat_f32_f32_decode.dispatch>(%weight_format, %token_count, %input_size, %output_size, %input, %weight, %output, %positions, %theta, %freq_factors, %output) : (index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_q_matmul_rope_f32_f32_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_q_matmul_rope_f32_f32_wmma.loom new file mode 100644 index 000000000000..21db9448a88c --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_q_matmul_rope_f32_f32_wmma.loom @@ -0,0 +1,61 @@ +template.decl @ggml.mul_mat_f32_f32_wmma.core(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: f32, %arg7: index, %arg8: index, %arg9: buffer, %arg10: buffer, %arg11: buffer, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer, %arg16: buffer, %arg17: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.launch(%token_capacity: index, %input_size: index, %output_size: index, %weight_format: index) -> (index, index, index, index) + +template.decl @ggml.mul_mat_f32_f32_wmma.publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %positions: buffer, %output: buffer, %theta: buffer, %freq_factors: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @llm.attention_qkv.rope_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: buffer, %arg6: buffer, %arg7: buffer, %arg8: vector<4xf32>) -> (vector<4xf32>) + +template.decl @llm.attention_qkv.store_query_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: vector<4xf32>, %arg6: buffer) + +amdgpu.target @llm_attention_q_matmul_rope_gfx11_wave64 {subgroup_size = 64} + +config.decl @llm.attention_qkv.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +config.decl @llm.attention_qkv.output_size : %value: index where [range(%value, 4, 32768), mul(%value, 4)] + +config.decl @llm.attention_qkv.weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @llm.attention_qkv.head_size : %value: index where [range(%value, 4, 1024), mul(%value, 4)] + +config.decl @llm.attention_qkv.head_count : %value: index where [range(%value, 1, 64)] + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +template.def<@ggml.mul_mat_f32_f32_wmma.publish_vector4> device @llm_attention_q_matmul_rope_publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %positions: buffer, %output: buffer, %theta: buffer, %freq_factors: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %head_size = config.get @llm.attention_qkv.head_size : index + %head_count = config.get @llm.attention_qkv.head_count : index + scf.if %publish_word { + %rotated = template.apply<@llm.attention_qkv.rope_vector4>(%token0, %channel, %token_count0, %head_size, %head_count, %positions, %theta, %freq_factors, %values) : (index, index, index, index, index, buffer, buffer, buffer, vector<4xf32>) -> (vector<4xf32>) + template.apply<@llm.attention_qkv.store_query_vector4>(%token0, %channel, %token_count0, %head_size, %head_count, %rotated, %output) : (index, index, index, index, index, vector<4xf32>, buffer) + } + template.return +} + +template.def<@ggml.mul_mat_f32_f32_wmma.finish_tile> device @llm_attention_q_matmul_rope_finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + template.return +} + +kernel.def target(@llm_attention_q_matmul_rope_gfx11_wave64) @llm_attention_q_matmul_rope_f32_f32_wmma(%token_count: index) { + %output_size = config.get @llm.attention_qkv.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %c1 = index.constant 1 : index + %input_size = config.get @llm.attention_qkv.input_size : index + %weight_format = config.get @llm.attention_qkv.weight_format : index + %launch_x, %launch_y, %launch_z, %workgroup_size = template.apply<@ggml.mul_mat_f32_f32_wmma.launch>(%token_capacity, %input_size, %output_size, %weight_format) pure : (index, index, index, index) -> (index, index, index, index) + kernel.launch.config workgroups(%launch_x, %launch_y, %launch_z) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %positions: buffer, %theta: buffer, %freq_factors: buffer, %output: buffer) { + %input_size = config.get @llm.attention_qkv.input_size : index + %output_size = config.get @llm.attention_qkv.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %weight_format = config.get @llm.attention_qkv.weight_format : index + %c0 = index.constant 0 : index + %c0_f32 = scalar.constant 0.0 : f32 + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_f32_f32_wmma.core>(%weight_format, %bounded_token_count, %input_size, %output_size, %c0, %c0, %c0_f32, %channel_tile, %token_tile, %input, %weight, %positions, %output, %theta, %freq_factors, %output, %output, %output) : (index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_v_matmul_set_rows_decode_f32_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_v_matmul_set_rows_decode_f32_f32.loom new file mode 100644 index 000000000000..496e9fbe0282 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_v_matmul_set_rows_decode_f32_f32.loom @@ -0,0 +1,94 @@ +template.decl @ggml.mul_mat_f32_f32_decode.dispatch(%weight_format: index, %token_count: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %output: buffer, %aux0: buffer, %aux1: buffer, %aux2: buffer, %aux3: buffer) + +template.decl @ggml.mul_mat_f32_f32_decode.publish_pair2(%publish_pair: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %row1_valid: i1, %value0: f32, %value1: f32, %cache: buffer, %indices: buffer, %aux1: buffer, %aux2: buffer, %aux3: buffer) + +amdgpu.target @llm_attention_v_matmul_set_rows_decode_gfx11_wave64 {subgroup_size = 64} + +config.decl @llm.attention_qkv.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 32)] + +config.decl @llm.attention_qkv.output_size : %value: index where [range(%value, 1, 32768)] + +config.decl @llm.attention_qkv.weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @llm.attention_qkv.cache_row_count : %value: index where [range(%value, 1, 1048576)] + +config.decl @llm.attention_qkv.cache_output_format : %value: index where [range(%value, 16, 32)] + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +template.def<@ggml.mul_mat_f32_f32_decode.publish_pair2> device priority(20) @llm_attention_v_matmul_set_rows_decode_publish_pair2(%publish_pair: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %row1_valid: i1, %value0: f32, %value1: f32, %cache: buffer, %indices: buffer, %aux1: buffer, %aux2: buffer, %aux3: buffer) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 1, 32768)] : index + %cache_row_count0 = config.get @llm.attention_qkv.cache_row_count : index + %cache_row_count = index.assume %cache_row_count0 [range(%cache_row_count0, 1, 1048576)] : index + %cache_output_format = config.get @llm.attention_qkv.cache_output_format : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c0_i64 = scalar.constant 0 : i64 + %c1048575_i64 = scalar.constant 1048575 : i64 + %c0_offset = index.constant 0 : offset + %token, %launch_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %is_f16 = index.cmp eq, %cache_output_format, %c16 : index + %is_f32 = index.cmp eq, %cache_output_format, %c32 : index + %indices_noalias, %cache_noalias = buffer.assume.noalias %indices, %cache : buffer, buffer + %indices_view = buffer.view %indices_noalias[%c0_offset] : buffer -> view<[%launch_token_count]xi64> + %index_raw = view.load %indices_view[%token] : view<[%launch_token_count]xi64> -> i64 + %index_nonnegative = scalar.cmpi sge, %index_raw, %c0_i64 : i64 + %index_in_cast_range = scalar.cmpi sle, %index_raw, %c1048575_i64 : i64 + %valid_index = scalar.andi %index_nonnegative, %index_in_cast_range : i1 + %safe_index0_i64 = scf.select %valid_index, %index_raw, %c0_i64 : i64 + %safe_index_i64 = scalar.assume %safe_index0_i64 [range(%safe_index0_i64, 0, 1048575)] : i64 + %cache_row0 = index.cast %safe_index_i64 : i64 to index + %valid_row = index.cmp ult, %cache_row0, %cache_row_count : index + %cache_row = scf.select %valid_row, %cache_row0, %c0 : index + %publish0 = scalar.andi %publish_pair, %valid_index : i1 + %publish1 = scalar.andi %publish0, %valid_row : i1 + %publish_first_f16 = scalar.andi %publish1, %is_f16 : i1 + %publish_first_f32 = scalar.andi %publish1, %is_f32 : i1 + %publish_second = scalar.andi %publish1, %row1_valid : i1 + %publish_second_f16 = scalar.andi %publish_second, %is_f16 : i1 + %publish_second_f32 = scalar.andi %publish_second, %is_f32 : i1 + scf.if %publish_first_f16 { + %cache_view = buffer.view %cache_noalias[%c0_offset] : buffer -> view<[%cache_row_count]x[%output_size]xf16> + %truncated = scalar.fptrunc %value0 : f32 to f16 + view.store %truncated, %cache_view[%cache_row, %channel] : f16, view<[%cache_row_count]x[%output_size]xf16> + } + scf.if %publish_first_f32 { + %cache_view = buffer.view %cache_noalias[%c0_offset] : buffer -> view<[%cache_row_count]x[%output_size]xf32> + view.store %value0, %cache_view[%cache_row, %channel] : f32, view<[%cache_row_count]x[%output_size]xf32> + } + scf.if %publish_second_f16 { + %cache_view = buffer.view %cache_noalias[%c0_offset] : buffer -> view<[%cache_row_count]x[%output_size]xf16> + %channel1 = index.add %channel, %c1 : index + %truncated = scalar.fptrunc %value1 : f32 to f16 + view.store %truncated, %cache_view[%cache_row, %channel1] : f16, view<[%cache_row_count]x[%output_size]xf16> + } + scf.if %publish_second_f32 { + %cache_view = buffer.view %cache_noalias[%c0_offset] : buffer -> view<[%cache_row_count]x[%output_size]xf32> + %channel1 = index.add %channel, %c1 : index + view.store %value1, %cache_view[%cache_row, %channel1] : f32, view<[%cache_row_count]x[%output_size]xf32> + } + template.return +} + +kernel.def target(@llm_attention_v_matmul_set_rows_decode_gfx11_wave64) @llm_attention_v_matmul_set_rows_decode_f32_f32(%token_count: index) { + %token_capacity = config.get @ggml.workload.token_capacity : index + %output_capacity = config.get @llm.attention_qkv.output_size : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %padded_output_size = index.add %output_capacity, %c1 : index + %output_pairs = index.div %padded_output_size, %c2 : index + %padded_token_count = index.add %token_capacity, %c63 : index + %launch_tokens = index.div %padded_token_count, %c64 : index + kernel.launch.config workgroups(%output_pairs, %launch_tokens, %c1) workgroup_size(%c64, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %indices: buffer, %cache: buffer) { + %input_size = config.get @llm.attention_qkv.input_size : index + %output_size = config.get @llm.attention_qkv.output_size : index + %weight_format = config.get @llm.attention_qkv.weight_format : index + template.apply<@ggml.mul_mat_f32_f32_decode.dispatch>(%weight_format, %token_count, %input_size, %output_size, %input, %weight, %cache, %indices, %cache, %cache, %cache) : (index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_v_matmul_set_rows_f32_f32_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_v_matmul_set_rows_f32_f32_wmma.loom new file mode 100644 index 000000000000..2cd1e8d4e7fc --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/llm_attention_v_matmul_set_rows_f32_f32_wmma.loom @@ -0,0 +1,58 @@ +template.decl @ggml.mul_mat_f32_f32_wmma.core(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: f32, %arg7: index, %arg8: index, %arg9: buffer, %arg10: buffer, %arg11: buffer, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer, %arg16: buffer, %arg17: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.launch(%token_capacity: index, %input_size: index, %output_size: index, %weight_format: index) -> (index, index, index, index) + +template.decl @ggml.mul_mat_f32_f32_wmma.publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %indices: buffer, %cache: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @llm.attention_qkv.store_cache_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: buffer, %arg7: vector<4xf32>, %arg8: buffer) + +amdgpu.target @llm_attention_v_matmul_set_rows_gfx11_wave64 {subgroup_size = 64} + +config.decl @llm.attention_qkv.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +config.decl @llm.attention_qkv.output_size : %value: index where [range(%value, 4, 32768), mul(%value, 4)] + +config.decl @llm.attention_qkv.weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @llm.attention_qkv.cache_row_count : %value: index where [range(%value, 1, 1048576)] + +config.decl @llm.attention_qkv.cache_output_format : %value: index where [range(%value, 16, 32)] + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +template.def<@ggml.mul_mat_f32_f32_wmma.publish_vector4> device @llm_attention_v_matmul_set_rows_publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %indices: buffer, %cache: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %cache_row_count = config.get @llm.attention_qkv.cache_row_count : index + %cache_output_format = config.get @llm.attention_qkv.cache_output_format : index + scf.if %publish_word { + template.apply<@llm.attention_qkv.store_cache_vector4>(%cache_output_format, %token0, %channel, %token_count0, %cache_row_count, %output_size0, %indices, %values, %cache) : (index, index, index, index, index, index, buffer, vector<4xf32>, buffer) + } + template.return +} + +template.def<@ggml.mul_mat_f32_f32_wmma.finish_tile> device @llm_attention_v_matmul_set_rows_finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + template.return +} + +kernel.def target(@llm_attention_v_matmul_set_rows_gfx11_wave64) @llm_attention_v_matmul_set_rows_f32_f32_wmma(%token_count: index) { + %output_size = config.get @llm.attention_qkv.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %input_size = config.get @llm.attention_qkv.input_size : index + %weight_format = config.get @llm.attention_qkv.weight_format : index + %c1 = index.constant 1 : index + %launch_x, %launch_y, %launch_z, %workgroup_size = template.apply<@ggml.mul_mat_f32_f32_wmma.launch>(%token_capacity, %input_size, %output_size, %weight_format) pure : (index, index, index, index) -> (index, index, index, index) + kernel.launch.config workgroups(%launch_x, %launch_y, %launch_z) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %indices: buffer, %cache: buffer) { + %input_size = config.get @llm.attention_qkv.input_size : index + %output_size = config.get @llm.attention_qkv.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %weight_format = config.get @llm.attention_qkv.weight_format : index + %c0 = index.constant 0 : index + %c0_f32 = scalar.constant 0.0 : f32 + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_f32_f32_wmma.core>(%weight_format, %bounded_token_count, %input_size, %output_size, %c0, %c0, %c0_f32, %channel_tile, %token_tile, %input, %weight, %indices, %cache, %cache, %cache, %cache, %cache, %cache) : (index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/moe_routing_tables.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/moe_routing_tables.loom new file mode 100644 index 000000000000..525081faef51 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/moe_routing_tables.loom @@ -0,0 +1,178 @@ +// Common MoE routing table preparation. +// +// The expert table stores [expert_count] assignment counts followed by +// [expert_count][token_count] compact assignment ordinals. The partition table +// stores a descriptor count followed by compact 32-row expert partitions. +amdgpu.target @ggml_moe_routing_gfx11_wave32 {subgroup_size = 32} + +config.decl @ggml.moe_routing.route_count : %value: index where [range(%value, 1, 32)] + +config.decl @ggml.moe_routing.expert_count : %value: index where [range(%value, 1, 512)] + +config.decl @ggml.moe_routing.descriptor_expert_mask : %value: i32 + +config.decl @ggml.moe_routing.descriptor_partition_shift : %value: i32 + +config.decl @ggml.moe_routing.descriptor_row_count_shift : %value: i32 + +config.decl @ggml.moe_routing.partition_workgroup_size : %value: index where [range(%value, 1, 1024)] + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +kernel.def target(@ggml_moe_routing_gfx11_wave32) @ggml_moe_build_expert_table(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index) { + %configured_expert_count = config.get @ggml.moe_routing.expert_count : index + %c1 = index.constant 1 : index + %workgroup_size = index.constant 256 : index + kernel.launch.config workgroups(%configured_expert_count, %c1, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index, %route_ids: buffer, %expert_table: buffer) where [range(%token_count, 1, 2048)] { + %configured_route_count0 = config.get @ggml.moe_routing.route_count : index + %configured_expert_count0 = config.get @ggml.moe_routing.expert_count : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %bounded_route_count, %configured_route_count = index.assume %route_count, %configured_route_count0 [range(%route_count, 1, 32), eq(%route_count, %configured_route_count0)] : index, index + %bounded_route_stride = index.assume %route_stride [range(%route_stride, 1, 512), ge(%route_stride, %bounded_route_count)] : index + %bounded_expert_count, %configured_expert_count = index.assume %expert_count, %configured_expert_count0 [range(%expert_count, 1, 512), eq(%expert_count, %configured_expert_count0)] : index, index + %expert0 = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %c0 = index.constant 0 : index + %workgroup_size = index.constant 256 : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %cn1_i32 = scalar.constant -1 : i32 + %c0_offset = index.constant 0 : offset + %c4_bytes = index.constant 4 : offset + %assignment_count = index.mul %bounded_token_count, %bounded_route_count : index + %expert, %table_expert_count = index.assume %expert0, %bounded_expert_count [lt(%expert0, %bounded_expert_count)] : index, index + %assignment_table_byte_base = index.scale %table_expert_count, %c4_bytes : index, offset -> offset + %route_view = buffer.view %route_ids[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_route_stride]xi32> + %count_view = buffer.view %expert_table[%c0_offset] : buffer -> view<[%table_expert_count]xi32> + %assignment_view = buffer.view %expert_table[%assignment_table_byte_base] : buffer -> view<[%table_expert_count]x[%bounded_token_count]xi32> + %expert_route_count = scf.for %block_base = [%c0 to %assignment_count step %workgroup_size](%matched_base = %c0_i32 : i32) -> (i32) { + %assignment = index.add %block_base, %lane : index + %in_range = index.cmp ult, %assignment, %assignment_count : index + %route_expert_i32 = scf.if %in_range -> (i32) { + %token0 = index.div %assignment, %configured_route_count : index + %route0 = index.rem %assignment, %configured_route_count : index + %token, %route_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %route, %route_row_stride = index.assume %route0, %bounded_route_stride [lt(%route0, %bounded_route_stride)] : index, index + %loaded = view.load %route_view[%token, %route] : view<[%bounded_token_count]x[%bounded_route_stride]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %cn1_i32 : i32 + } + %route_expert0 = index.cast %route_expert_i32 : i32 to index + %route_expert = index.assume %route_expert0 [range(%route_expert0, -1, 511)] : index + %matches = index.cmp eq, %route_expert, %expert : index + %match_i32 = scf.if %matches -> (i32) { + scf.yield %c1_i32 : i32 + } else { + scf.yield %c0_i32 : i32 + } + %block_prefix = kernel.workgroup.scan %match_i32 {direction = forward, mode = exclusive} : i32 + %block_match_count_reduced = kernel.workgroup.reduce %match_i32 : i32 + %block_match_count = kernel.subgroup.broadcast.first %block_match_count_reduced : i32 + scf.if %matches { + %match_ordinal_i32 = scalar.addi %matched_base, %block_prefix : i32 + %match_ordinal0 = index.cast %match_ordinal_i32 : i32 to index + %match_ordinal = index.assume %match_ordinal0 [range(%match_ordinal0, 0, 2047)] : index + %bounded_match_ordinal, %table_token_count = index.assume %match_ordinal, %bounded_token_count [lt(%match_ordinal, %bounded_token_count)] : index, index + %assignment_i32 = index.cast %assignment : index to i32 + view.store %assignment_i32, %assignment_view[%expert, %bounded_match_ordinal] : i32, view<[%table_expert_count]x[%bounded_token_count]xi32> + } + %next_matched_base = scalar.addi %matched_base, %block_match_count : i32 + scf.yield %next_matched_base : i32 + } + %is_lane_zero = index.cmp eq, %lane, %c0 : index + scf.if %is_lane_zero { + view.store %expert_route_count, %count_view[%expert] : i32, view<[%table_expert_count]xi32> + } + kernel.return +} + +kernel.def target(@ggml_moe_routing_gfx11_wave32) @ggml_moe_build_expert_partition_table(%token_count: index, %route_count: index, %expert_count: index) { + %c1 = index.constant 1 : index + %workgroup_size = config.get @ggml.moe_routing.partition_workgroup_size : index + kernel.launch.config workgroups(%c1, %c1, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %route_count: index, %expert_count: index, %expert_table: buffer, %partition_table: buffer) where [range(%token_count, 1, 2048)] { + %configured_route_count0 = config.get @ggml.moe_routing.route_count : index + %configured_expert_count0 = config.get @ggml.moe_routing.expert_count : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %bounded_route_count, %configured_route_count = index.assume %route_count, %configured_route_count0 [range(%route_count, 1, 32), eq(%route_count, %configured_route_count0)] : index, index + %bounded_expert_count, %configured_expert_count = index.assume %expert_count, %configured_expert_count0 [range(%expert_count, 1, 512), eq(%expert_count, %configured_expert_count0)] : index, index + %lane = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %descriptor_expert_mask = config.get @ggml.moe_routing.descriptor_expert_mask : i32 + %descriptor_partition_shift = config.get @ggml.moe_routing.descriptor_partition_shift : i32 + %descriptor_row_count_shift = config.get @ggml.moe_routing.descriptor_row_count_shift : i32 + %c0_offset = index.constant 0 : offset + %c4_bytes = index.constant 4 : offset + %assignment_count = index.mul %bounded_token_count, %bounded_route_count : index + %rounded_assignment_count = index.add %assignment_count, %c31 : index + %assignment_partition_count = index.div %rounded_assignment_count, %c32 : index + %maximum_partition_count = index.add %assignment_partition_count, %bounded_expert_count : index + %count_view = buffer.view %expert_table[%c0_offset] : buffer -> view<[%bounded_expert_count]xi32> + %partition_count_view = buffer.view %partition_table[%c0_offset] : buffer -> view<1xi32> + %partition_descriptor_view = buffer.view %partition_table[%c4_bytes] : buffer -> view<[%maximum_partition_count]xi32> + %has_expert = index.cmp ult, %lane, %bounded_expert_count : index + %expert_assignment_count_i32 = scf.if %has_expert -> (i32) { + %expert, %table_expert_count = index.assume %lane, %bounded_expert_count [lt(%lane, %bounded_expert_count)] : index, index + %loaded = view.load %count_view[%expert] : view<[%bounded_expert_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %c0_i32 : i32 + } + %expert_assignment_count0 = index.cast %expert_assignment_count_i32 : i32 to index + %expert_assignment_count = index.assume %expert_assignment_count0 [range(%expert_assignment_count0, 0, 2048)] : index + %rounded_expert_assignment_count = index.add %expert_assignment_count, %c31 : index + %expert_partition_count = index.div %rounded_expert_assignment_count, %c32 : index + %expert_partition_count_i32 = index.cast %expert_partition_count : index to i32 + %expert_partition_base_i32 = kernel.workgroup.scan %expert_partition_count_i32 {direction = forward, mode = exclusive} : i32 + %partition_count_i32 = kernel.workgroup.reduce %expert_partition_count_i32 : i32 + %partition_count0 = index.cast %partition_count_i32 : i32 to index + %partition_count, %table_partition_capacity = index.assume %partition_count0, %maximum_partition_count [le(%partition_count0, %maximum_partition_count)] : index, index + %expert_partition_base0 = index.cast %expert_partition_base_i32 : i32 to index + %expert_partition_base = index.assume %expert_partition_base0 [range(%expert_partition_base0, 0, 2559)] : index + scf.if %has_expert { + %expert, %table_expert_count = index.assume %lane, %bounded_expert_count [lt(%lane, %bounded_expert_count)] : index, index + %expert_i32 = index.cast %expert : index to i32 + scf.for %partition = [%c0 to %expert_partition_count step %c1] { + %descriptor_ordinal0 = index.add %expert_partition_base, %partition : index + %descriptor_ordinal, %descriptor_count = index.assume %descriptor_ordinal0, %partition_count [lt(%descriptor_ordinal0, %partition_count)] : index, index + %table_descriptor_ordinal, %table_descriptor_capacity = index.assume %descriptor_ordinal, %maximum_partition_count [lt(%descriptor_ordinal, %maximum_partition_count)] : index, index + %partition_remainder = index.rem %expert_assignment_count, %c32 : index + %has_partial_tail = index.cmp ne, %partition_remainder, %c0 : index + %partition_row_count = scf.if %has_partial_tail -> (index) { + %next_partition = index.add %partition, %c1 : index + %is_tail_partition = index.cmp eq, %next_partition, %expert_partition_count : index + %tail_row_count = scf.if %is_tail_partition -> (index) { + scf.yield %partition_remainder : index + } else { + scf.yield %c32 : index + } + scf.yield %tail_row_count : index + } else { + scf.yield %c32 : index + } + %partition_i32 = index.cast %partition : index to i32 + %partition_row_count_i32 = index.cast %partition_row_count : index to i32 + %bounded_expert_i32 = scalar.andi %expert_i32, %descriptor_expert_mask : i32 + %packed_partition = scalar.shli %partition_i32, %descriptor_partition_shift : i32 + %partition_row_count_minus_one = scalar.subi %partition_row_count_i32, %c1_i32 : i32 + %packed_row_count = scalar.shli %partition_row_count_minus_one, %descriptor_row_count_shift : i32 + %packed_expert_partition = scalar.ori %bounded_expert_i32, %packed_partition : i32 + %packed_descriptor = scalar.ori %packed_expert_partition, %packed_row_count : i32 + view.store %packed_descriptor, %partition_descriptor_view[%table_descriptor_ordinal] : i32, view<[%maximum_partition_count]xi32> + } + } + %is_lane_zero = index.cmp eq, %lane, %c0 : index + scf.if %is_lane_zero { + view.store %partition_count_i32, %partition_count_view[%c0] : i32, view<1xi32> + } + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_add_f32_f32_decode.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_add_f32_f32_decode.loom new file mode 100644 index 000000000000..22e7cc97c168 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_add_f32_f32_decode.loom @@ -0,0 +1,55 @@ +template.decl @ggml.mul_mat_f32_f32_decode.dispatch(%weight_format: index, %token_count: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %output: buffer, %aux0: buffer, %aux1: buffer, %aux2: buffer, %aux3: buffer) + +template.decl @ggml.mul_mat_f32_f32_decode.publish_pair2(%publish_pair: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %row1_valid: i1, %value0: f32, %value1: f32, %output: buffer, %residual_input: buffer, %aux1: buffer, %aux2: buffer, %aux3: buffer) + +amdgpu.target @ggml_mul_mat_add_f32_f32_decode_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.mul_mat_f32_f32_decode.token_capacity : %value: index where [range(%value, 1, 2048)] + +config.decl @ggml.mul_mat_f32_f32_decode.output_capacity : %value: index where [range(%value, 1, 1048576)] + +config.decl @ggml.mul_mat_f32_f32_decode.weight_format : %value: index where [range(%value, 4, 81)] + +template.def<@ggml.mul_mat_f32_f32_decode.publish_pair2> device priority(20) @ggml_mul_mat_add_f32_f32_decode_publish_pair2(%publish_pair: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %row1_valid: i1, %value0: f32, %value1: f32, %output: buffer, %residual_input: buffer, %aux1: buffer, %aux2: buffer, %aux3: buffer) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 1, 1048576)] : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %token, %launch_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %output_noalias = buffer.assume.noalias %output : buffer + %residual_input_noalias = buffer.assume.noalias %residual_input : buffer + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%output_size]xf32> + %residual_input_view = buffer.view %residual_input_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%output_size]xf32> + scf.if %publish_pair { + %residual0 = view.load %residual_input_view[%token, %channel] : view<[%launch_token_count]x[%output_size]xf32> -> f32 + %sum0 = scalar.addf %value0, %residual0 : f32 + view.store %sum0, %output_view[%token, %channel] : f32, view<[%launch_token_count]x[%output_size]xf32> + } + scf.if %row1_valid { + scf.if %publish_pair { + %row1 = index.add %channel, %c1 : index + %residual1 = view.load %residual_input_view[%token, %row1] : view<[%launch_token_count]x[%output_size]xf32> -> f32 + %sum1 = scalar.addf %value1, %residual1 : f32 + view.store %sum1, %output_view[%token, %row1] : f32, view<[%launch_token_count]x[%output_size]xf32> + } + } + template.return +} + +kernel.def target(@ggml_mul_mat_add_f32_f32_decode_gfx11_wave64) @ggml_mul_mat_add_f32_f32_decode_wave64(%token_count: index, %input_size: index, %output_size: index) { + %token_capacity = config.get @ggml.mul_mat_f32_f32_decode.token_capacity : index + %output_capacity = config.get @ggml.mul_mat_f32_f32_decode.output_capacity : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %padded_output_size = index.add %output_capacity, %c1 : index + %output_pairs = index.div %padded_output_size, %c2 : index + %padded_token_count = index.add %token_capacity, %c63 : index + %launch_tokens = index.div %padded_token_count, %c64 : index + kernel.launch.config workgroups(%output_pairs, %launch_tokens, %c1) workgroup_size(%c64, %c1, %c1) : index +} launch(%token_count: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %residual_input: buffer, %residual_output: buffer) { + %weight_format = config.get @ggml.mul_mat_f32_f32_decode.weight_format : index + template.apply<@ggml.mul_mat_f32_f32_decode.dispatch>(%weight_format, %token_count, %input_size, %output_size, %input, %weight, %residual_output, %residual_input, %residual_output, %residual_output, %residual_output) : (index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_add_f32_f32_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_add_f32_f32_wmma.loom new file mode 100644 index 000000000000..95cd5db63032 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_add_f32_f32_wmma.loom @@ -0,0 +1,55 @@ +template.decl @ggml.mul_mat_f32_f32_wmma.add_residual_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: vector<4xf32>, %arg5: buffer) -> (vector<4xf32>) + +template.decl @ggml.mul_mat_f32_f32_wmma.core(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: f32, %arg7: index, %arg8: index, %arg9: buffer, %arg10: buffer, %arg11: buffer, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer, %arg16: buffer, %arg17: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.launch(%token_capacity: index, %input_size: index, %output_size: index, %weight_format: index) -> (index, index, index, index) + +template.decl @ggml.mul_mat_f32_f32_wmma.publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.store_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: vector<4xf32>, %arg5: buffer) + +amdgpu.target @ggml_mul_mat_add_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.mul_mat_postops.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +config.decl @ggml.mul_mat_postops.output_size : %value: index where [range(%value, 1, 262144)] + +config.decl @ggml.mul_mat_postops.weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +template.def<@ggml.mul_mat_f32_f32_wmma.publish_vector4> device @ggml_mul_mat_add_publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + scf.if %publish_word { + %published = template.apply<@ggml.mul_mat_f32_f32_wmma.add_residual_vector4>(%token_count0, %output_size0, %token0, %channel, %values, %residual_input) : (index, index, index, index, vector<4xf32>, buffer) -> (vector<4xf32>) + template.apply<@ggml.mul_mat_f32_f32_wmma.store_vector4>(%token_count0, %output_size0, %token0, %channel, %published, %residual_output) : (index, index, index, index, vector<4xf32>, buffer) + } + template.return +} + +template.def<@ggml.mul_mat_f32_f32_wmma.finish_tile> device @ggml_mul_mat_add_finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + template.return +} + +kernel.def target(@ggml_mul_mat_add_gfx11_wave64) @ggml_mul_mat_add_f32_f32_wmma(%token_count: index) { + %output_size = config.get @ggml.mul_mat_postops.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %c1 = index.constant 1 : index + %input_size = config.get @ggml.mul_mat_postops.input_size : index + %weight_format = config.get @ggml.mul_mat_postops.weight_format : index + %launch_x, %launch_y, %launch_z, %workgroup_size = template.apply<@ggml.mul_mat_f32_f32_wmma.launch>(%token_capacity, %input_size, %output_size, %weight_format) pure : (index, index, index, index) -> (index, index, index, index) + kernel.launch.config workgroups(%launch_x, %launch_y, %launch_z) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %residual_input: buffer, %residual_output: buffer) { + %input_size = config.get @ggml.mul_mat_postops.input_size : index + %output_size = config.get @ggml.mul_mat_postops.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %weight_format = config.get @ggml.mul_mat_postops.weight_format : index + %c0 = index.constant 0 : index + %c0_f32 = scalar.constant 0.0 : f32 + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_f32_f32_wmma.core>(%weight_format, %bounded_token_count, %input_size, %output_size, %c0, %c0, %c0_f32, %channel_tile, %token_tile, %input, %weight, %residual_output, %residual_output, %residual_input, %residual_output, %residual_output, %residual_output, %residual_output) : (index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_add_next_rmsnorm_f32_f32_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_add_next_rmsnorm_f32_f32_wmma.loom new file mode 100644 index 000000000000..cd630d060222 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_add_next_rmsnorm_f32_f32_wmma.loom @@ -0,0 +1,68 @@ +// Dense GGML matmul fused with the following residual add and RMSNorm-weight +// multiply: +// MUL_MAT -> ADD(residual) -> RMS_NORM -> MUL(norm_weight) +// +// The residual-add output is explicitly materialized for later graph users. +// The raw matmul and raw RMSNorm outputs are internal to this fused dispatch. +template.decl @ggml.mul_mat_f32_f32_wmma.add_residual_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: vector<4xf32>, %arg5: buffer) -> (vector<4xf32>) + +template.decl @ggml.mul_mat_f32_f32_wmma.core(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: f32, %arg7: index, %arg8: index, %arg9: buffer, %arg10: buffer, %arg11: buffer, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer, %arg16: buffer, %arg17: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.finish_rmsnorm_weight(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: f32, %arg5: index, %arg6: buffer, %arg7: buffer, %arg8: buffer, %arg9: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.launch(%token_capacity: index, %input_size: index, %output_size: index, %weight_format: index) -> (index, index, index, index) + +template.decl @ggml.mul_mat_f32_f32_wmma.publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.store_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: vector<4xf32>, %arg5: buffer) + +amdgpu.target @ggml_mul_mat_add_next_rmsnorm_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.mul_mat_postops.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +config.decl @ggml.mul_mat_postops.output_size : %value: index where [range(%value, 128, 32768), mul(%value, 128)] + +config.decl @ggml.mul_mat_postops.weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.mul_mat_postops.rms_epsilon : f32 + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +template.def<@ggml.mul_mat_f32_f32_wmma.publish_vector4> device @ggml_mul_mat_add_next_rmsnorm_publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 128, 32768), mul(%output_size0, 128)] : index + scf.if %publish_word { + %published = template.apply<@ggml.mul_mat_f32_f32_wmma.add_residual_vector4>(%token_count, %output_size, %token0, %channel, %values, %residual_input) : (index, index, index, index, vector<4xf32>, buffer) -> (vector<4xf32>) + template.apply<@ggml.mul_mat_f32_f32_wmma.store_vector4>(%token_count, %output_size, %token0, %channel, %published, %residual_output) : (index, index, index, index, vector<4xf32>, buffer) + } + template.return +} + +template.def<@ggml.mul_mat_f32_f32_wmma.finish_tile> device @ggml_mul_mat_add_next_rmsnorm_finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + template.apply<@ggml.mul_mat_f32_f32_wmma.finish_rmsnorm_weight>(%token_count, %output_size, %output_tile_count, %token_tile_count, %epsilon, %token_tile, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (index, index, index, index, f32, index, buffer, buffer, buffer, buffer) + template.return +} + +kernel.def target(@ggml_mul_mat_add_next_rmsnorm_gfx11_wave64) @ggml_mul_mat_add_next_rmsnorm_f32_f32_wmma(%token_count: index) { + %output_size = config.get @ggml.mul_mat_postops.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %c1 = index.constant 1 : index + %input_size = config.get @ggml.mul_mat_postops.input_size : index + %weight_format = config.get @ggml.mul_mat_postops.weight_format : index + %launch_x, %launch_y, %launch_z, %workgroup_size = template.apply<@ggml.mul_mat_f32_f32_wmma.launch>(%token_capacity, %input_size, %output_size, %weight_format) pure : (index, index, index, index) -> (index, index, index, index) + kernel.launch.config workgroups(%launch_x, %launch_y, %launch_z) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %input_size = config.get @ggml.mul_mat_postops.input_size : index + %output_size = config.get @ggml.mul_mat_postops.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %weight_format = config.get @ggml.mul_mat_postops.weight_format : index + %epsilon = config.get @ggml.mul_mat_postops.rms_epsilon : f32 + %c0 = index.constant 0 : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_f32_f32_wmma.core>(%weight_format, %bounded_token_count, %input_size, %output_size, %c0, %c0, %epsilon, %channel_tile, %token_tile, %input, %weight, %residual_output, %residual_output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_bias_add_f32_f32_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_bias_add_f32_f32_wmma.loom new file mode 100644 index 000000000000..27c9fef60c0e --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_bias_add_f32_f32_wmma.loom @@ -0,0 +1,58 @@ +template.decl @ggml.mul_mat_f32_f32_wmma.add_bias_vector4(%arg0: index, %arg1: index, %arg2: vector<4xf32>, %arg3: buffer) -> (vector<4xf32>) + +template.decl @ggml.mul_mat_f32_f32_wmma.add_residual_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: vector<4xf32>, %arg5: buffer) -> (vector<4xf32>) + +template.decl @ggml.mul_mat_f32_f32_wmma.core(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: f32, %arg7: index, %arg8: index, %arg9: buffer, %arg10: buffer, %arg11: buffer, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer, %arg16: buffer, %arg17: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.launch(%token_capacity: index, %input_size: index, %output_size: index, %weight_format: index) -> (index, index, index, index) + +template.decl @ggml.mul_mat_f32_f32_wmma.publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.store_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: vector<4xf32>, %arg5: buffer) + +amdgpu.target @ggml_mul_mat_bias_add_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.mul_mat_postops.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +config.decl @ggml.mul_mat_postops.output_size : %value: index where [range(%value, 1, 262144)] + +config.decl @ggml.mul_mat_postops.weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +template.def<@ggml.mul_mat_f32_f32_wmma.publish_vector4> device @ggml_mul_mat_bias_add_publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + scf.if %publish_word { + %biased = template.apply<@ggml.mul_mat_f32_f32_wmma.add_bias_vector4>(%output_size0, %channel, %values, %bias) : (index, index, vector<4xf32>, buffer) -> (vector<4xf32>) + %published = template.apply<@ggml.mul_mat_f32_f32_wmma.add_residual_vector4>(%token_count0, %output_size0, %token0, %channel, %biased, %residual_input) : (index, index, index, index, vector<4xf32>, buffer) -> (vector<4xf32>) + template.apply<@ggml.mul_mat_f32_f32_wmma.store_vector4>(%token_count0, %output_size0, %token0, %channel, %published, %residual_output) : (index, index, index, index, vector<4xf32>, buffer) + } + template.return +} + +template.def<@ggml.mul_mat_f32_f32_wmma.finish_tile> device @ggml_mul_mat_bias_add_finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + template.return +} + +kernel.def target(@ggml_mul_mat_bias_add_gfx11_wave64) @ggml_mul_mat_bias_add_f32_f32_wmma(%token_count: index) { + %output_size = config.get @ggml.mul_mat_postops.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %c1 = index.constant 1 : index + %input_size = config.get @ggml.mul_mat_postops.input_size : index + %weight_format = config.get @ggml.mul_mat_postops.weight_format : index + %launch_x, %launch_y, %launch_z, %workgroup_size = template.apply<@ggml.mul_mat_f32_f32_wmma.launch>(%token_capacity, %input_size, %output_size, %weight_format) pure : (index, index, index, index) -> (index, index, index, index) + kernel.launch.config workgroups(%launch_x, %launch_y, %launch_z) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %bias: buffer, %residual_input: buffer, %residual_output: buffer) { + %input_size = config.get @ggml.mul_mat_postops.input_size : index + %output_size = config.get @ggml.mul_mat_postops.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %weight_format = config.get @ggml.mul_mat_postops.weight_format : index + %c0 = index.constant 0 : index + %c0_f32 = scalar.constant 0.0 : f32 + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_f32_f32_wmma.core>(%weight_format, %bounded_token_count, %input_size, %output_size, %c0, %c0, %c0_f32, %channel_tile, %token_tile, %input, %weight, %bias, %residual_output, %residual_input, %residual_output, %residual_output, %residual_output, %residual_output) : (index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_bias_add_next_rmsnorm_f32_f32_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_bias_add_next_rmsnorm_f32_f32_wmma.loom new file mode 100644 index 000000000000..aa6b4bfd4148 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_bias_add_next_rmsnorm_f32_f32_wmma.loom @@ -0,0 +1,63 @@ +template.decl @ggml.mul_mat_f32_f32_wmma.add_bias_vector4(%arg0: index, %arg1: index, %arg2: vector<4xf32>, %arg3: buffer) -> (vector<4xf32>) + +template.decl @ggml.mul_mat_f32_f32_wmma.add_residual_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: vector<4xf32>, %arg5: buffer) -> (vector<4xf32>) + +template.decl @ggml.mul_mat_f32_f32_wmma.core(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: f32, %arg7: index, %arg8: index, %arg9: buffer, %arg10: buffer, %arg11: buffer, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer, %arg16: buffer, %arg17: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.finish_rmsnorm_weight(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: f32, %arg5: index, %arg6: buffer, %arg7: buffer, %arg8: buffer, %arg9: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.launch(%token_capacity: index, %input_size: index, %output_size: index, %weight_format: index) -> (index, index, index, index) + +template.decl @ggml.mul_mat_f32_f32_wmma.publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.store_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: vector<4xf32>, %arg5: buffer) + +amdgpu.target @ggml_mul_mat_bias_add_next_rmsnorm_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.mul_mat_postops.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +config.decl @ggml.mul_mat_postops.output_size : %value: index where [range(%value, 128, 32768), mul(%value, 128)] + +config.decl @ggml.mul_mat_postops.weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.mul_mat_postops.rms_epsilon : f32 + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +template.def<@ggml.mul_mat_f32_f32_wmma.publish_vector4> device @ggml_mul_mat_bias_add_next_rmsnorm_publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + scf.if %publish_word { + %biased = template.apply<@ggml.mul_mat_f32_f32_wmma.add_bias_vector4>(%output_size0, %channel, %values, %bias) : (index, index, vector<4xf32>, buffer) -> (vector<4xf32>) + %published = template.apply<@ggml.mul_mat_f32_f32_wmma.add_residual_vector4>(%token_count0, %output_size0, %token0, %channel, %biased, %residual_input) : (index, index, index, index, vector<4xf32>, buffer) -> (vector<4xf32>) + template.apply<@ggml.mul_mat_f32_f32_wmma.store_vector4>(%token_count0, %output_size0, %token0, %channel, %published, %residual_output) : (index, index, index, index, vector<4xf32>, buffer) + } + template.return +} + +template.def<@ggml.mul_mat_f32_f32_wmma.finish_tile> device @ggml_mul_mat_bias_add_next_rmsnorm_finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + template.apply<@ggml.mul_mat_f32_f32_wmma.finish_rmsnorm_weight>(%token_count, %output_size, %output_tile_count, %token_tile_count, %epsilon, %token_tile, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (index, index, index, index, f32, index, buffer, buffer, buffer, buffer) + template.return +} + +kernel.def target(@ggml_mul_mat_bias_add_next_rmsnorm_gfx11_wave64) @ggml_mul_mat_bias_add_next_rmsnorm_f32_f32_wmma(%token_count: index) { + %output_size = config.get @ggml.mul_mat_postops.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %c1 = index.constant 1 : index + %input_size = config.get @ggml.mul_mat_postops.input_size : index + %weight_format = config.get @ggml.mul_mat_postops.weight_format : index + %launch_x, %launch_y, %launch_z, %workgroup_size = template.apply<@ggml.mul_mat_f32_f32_wmma.launch>(%token_capacity, %input_size, %output_size, %weight_format) pure : (index, index, index, index) -> (index, index, index, index) + kernel.launch.config workgroups(%launch_x, %launch_y, %launch_z) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %bias: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %input_size = config.get @ggml.mul_mat_postops.input_size : index + %output_size = config.get @ggml.mul_mat_postops.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %weight_format = config.get @ggml.mul_mat_postops.weight_format : index + %epsilon = config.get @ggml.mul_mat_postops.rms_epsilon : f32 + %c0 = index.constant 0 : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_f32_f32_wmma.core>(%weight_format, %bounded_token_count, %input_size, %output_size, %c0, %c0, %epsilon, %channel_tile, %token_tile, %input, %weight, %bias, %residual_output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_bias_f32_f32_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_bias_f32_f32_wmma.loom new file mode 100644 index 000000000000..cac1acce67bf --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_bias_f32_f32_wmma.loom @@ -0,0 +1,55 @@ +template.decl @ggml.mul_mat_f32_f32_wmma.add_bias_vector4(%arg0: index, %arg1: index, %arg2: vector<4xf32>, %arg3: buffer) -> (vector<4xf32>) + +template.decl @ggml.mul_mat_f32_f32_wmma.core(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: f32, %arg7: index, %arg8: index, %arg9: buffer, %arg10: buffer, %arg11: buffer, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer, %arg16: buffer, %arg17: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.launch(%token_capacity: index, %input_size: index, %output_size: index, %weight_format: index) -> (index, index, index, index) + +template.decl @ggml.mul_mat_f32_f32_wmma.publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.store_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: vector<4xf32>, %arg5: buffer) + +amdgpu.target @ggml_mul_mat_bias_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.mul_mat_postops.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +config.decl @ggml.mul_mat_postops.output_size : %value: index where [range(%value, 1, 262144)] + +config.decl @ggml.mul_mat_postops.weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +template.def<@ggml.mul_mat_f32_f32_wmma.publish_vector4> device @ggml_mul_mat_bias_publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + scf.if %publish_word { + %biased = template.apply<@ggml.mul_mat_f32_f32_wmma.add_bias_vector4>(%output_size0, %channel, %values, %bias) : (index, index, vector<4xf32>, buffer) -> (vector<4xf32>) + template.apply<@ggml.mul_mat_f32_f32_wmma.store_vector4>(%token_count0, %output_size0, %token0, %channel, %biased, %output) : (index, index, index, index, vector<4xf32>, buffer) + } + template.return +} + +template.def<@ggml.mul_mat_f32_f32_wmma.finish_tile> device @ggml_mul_mat_bias_finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + template.return +} + +kernel.def target(@ggml_mul_mat_bias_gfx11_wave64) @ggml_mul_mat_bias_f32_f32_wmma(%token_count: index) { + %output_size = config.get @ggml.mul_mat_postops.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %c1 = index.constant 1 : index + %input_size = config.get @ggml.mul_mat_postops.input_size : index + %weight_format = config.get @ggml.mul_mat_postops.weight_format : index + %launch_x, %launch_y, %launch_z, %workgroup_size = template.apply<@ggml.mul_mat_f32_f32_wmma.launch>(%token_capacity, %input_size, %output_size, %weight_format) pure : (index, index, index, index) -> (index, index, index, index) + kernel.launch.config workgroups(%launch_x, %launch_y, %launch_z) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %bias: buffer, %output: buffer) { + %input_size = config.get @ggml.mul_mat_postops.input_size : index + %output_size = config.get @ggml.mul_mat_postops.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %weight_format = config.get @ggml.mul_mat_postops.weight_format : index + %c0 = index.constant 0 : index + %c0_f32 = scalar.constant 0.0 : f32 + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_f32_f32_wmma.core>(%weight_format, %bounded_token_count, %input_size, %output_size, %c0, %c0, %c0_f32, %channel_tile, %token_tile, %input, %weight, %bias, %output, %output, %output, %output, %output, %output) : (index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_dual_q4_f32_decode.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_dual_q4_f32_decode.loom new file mode 100644 index 000000000000..8bbc43af5d39 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_dual_q4_f32_decode.loom @@ -0,0 +1,88 @@ +// Two raw-Q4_K one-token F32 projections sharing activation loads. +// Each wave64 computes the same row of two distinct weights. The original +// public decoder retains F16 weight rounding, four strided K cohorts, dot4 +// contraction and reduction order. Both ordinary F32 outputs are published. +config.decl @ggml.mul_mat_dual_q4_f32_c1.input_size : %value: index where [range(%value, 4096, 32768), mul(%value, 1024)] +config.decl @ggml.mul_mat_dual_q4_f32_c1.output_size : %value: index where [range(%value, 24, 63)] + +amdgpu.target @ggml_mul_mat_dual_q4_f32_decode_gfx11_wave64 {subgroup_size = 64} + +func.decl @ggml_iq4nl_table_i8() -> (vector<16xi8>) + +func.decl @ggml_iq3s_grid_f32_halves() -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + +func.decl @ggml_mul_mat_decode_load_f32_block(%token_count: index, %input_size: index, %token: index, %block: index, %lane: index, %input: buffer) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) + +func.decl @ggml_mul_mat_decode_f32_block_row(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight_format: index, %input_size: index, %row: index, %block: index, %lane: index, %weight: buffer, %input0: vector<4xf32>, %input1: vector<4xf32>, %input2: vector<4xf32>, %input3: vector<4xf32>) -> (f32) + +func.def inline @ggml_mul_mat_dual_q4_f32_decode_body(%weight_format: index, %publish_output: i1, %token_count: index, %token0: index, %pair: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %output: buffer, %aux0: buffer, %aux1: buffer, %aux2: buffer, %aux3: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 32)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 1048576)] : index + %lane0 = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %c0_f32 = scalar.constant 0.0 : f32 + %c0_offset = index.constant 0 : offset + %token, %launch_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %lane = index.assume %lane0 [range(%lane0, 0, 63)] : index + %row00 = index.add %pair, %c0 : index + %row0, %launch_output_size = index.assume %row00, %bounded_output_size [lt(%row00, %bounded_output_size)] : index, index + %row1 = index.add %row0, %c0 : index + %row1_valid = index.cmp ult, %row1, %launch_output_size : index + %padded_input_size = index.add %bounded_input_size, %c255 : index + %block_count = index.div %padded_input_size, %c256 : index + %cohort = index.div %lane, %c16 : index + scf.if %publish_output { + %input_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input, %weight, %output : buffer, buffer, buffer + %iq4nl_table = func.call @ggml_iq4nl_table_i8() : () -> (vector<16xi8>) + %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15 = func.call @ggml_iq3s_grid_f32_halves() : () -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + %acc0, %acc1 = scf.for %block_base = [%c0 to %block_count step %c4](%row_acc0 = %c0_f32 : f32, %row_acc1 = %c0_f32 : f32) -> (f32, f32) { + %block = index.add %block_base, %cohort : index + %input0, %input1, %input2, %input3 = func.call @ggml_mul_mat_decode_load_f32_block(%launch_token_count, %bounded_input_size, %token, %block, %lane, %input_noalias) : (index, index, index, index, index, buffer) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) + %contribution0 = func.call @ggml_mul_mat_decode_f32_block_row(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight_format, %bounded_input_size, %row0, %block, %lane, %weight_noalias, %input0, %input1, %input2, %input3) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, index, index, index, index, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) -> (f32) + %contribution1 = scf.if %row1_valid -> (f32) { + %row1_contribution = func.call @ggml_mul_mat_decode_f32_block_row(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight_format, %bounded_input_size, %row1, %block, %lane, %aux0, %input0, %input1, %input2, %input3) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, index, index, index, index, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) -> (f32) + scf.yield %row1_contribution : f32 + } else { + scf.yield %c0_f32 : f32 + } + %next0 = scalar.addf %row_acc0, %contribution0 : f32 + %next1 = scalar.addf %row_acc1, %contribution1 : f32 + scf.yield %next0, %next1 : f32, f32 + } + %sum0 = kernel.workgroup.reduce %acc0 : f32 + %sum1 = kernel.workgroup.reduce %acc1 : f32 + %is_lane_zero = index.cmp eq, %lane, %c0 : index + %publish_pair = scalar.andi %publish_output, %is_lane_zero : i1 + %alpha_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_output_size]xf32> + %beta_view = buffer.view %aux1[%c0_offset] : buffer -> view<[%launch_output_size]xf32> + scf.if %publish_pair { + view.store %sum0, %alpha_view[%row0] : f32, view<[%launch_output_size]xf32> + view.store %sum1, %beta_view[%row0] : f32, view<[%launch_output_size]xf32> + } + } + func.return +} + +kernel.def target(@ggml_mul_mat_dual_q4_f32_decode_gfx11_wave64) export("ggml_mul_mat_dual_q4_f32_decode") @ggml_mul_mat_dual_q4_f32_decode() { + %heads = config.get @ggml.mul_mat_dual_q4_f32_c1.output_size : index + %threads = index.constant 64 : index + %one = index.constant 1 : index + kernel.launch.config workgroups(%heads, %one, %one) workgroup_size(%threads, %one, %one) : index +} launch(%input: buffer, %alpha: buffer, %beta: buffer, %output_alpha: buffer, %output_beta: buffer) { + %format = index.constant 4 : index + %one = index.constant 1 : index + %zero = index.constant 0 : index + %yes = scalar.constant true : i1 + %k = config.get @ggml.mul_mat_dual_q4_f32_c1.input_size : index + %n = config.get @ggml.mul_mat_dual_q4_f32_c1.output_size : index + %row = kernel.workgroup.id : index + func.call @ggml_mul_mat_dual_q4_f32_decode_body(%format, %yes, %one, %zero, %row, %k, %n, %input, %alpha, %output_alpha, %beta, %output_beta, %output_alpha, %output_alpha) : (index, i1, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_f32_f32_decode.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_f32_f32_decode.loom new file mode 100644 index 000000000000..84b8a284d4f6 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_f32_f32_decode.loom @@ -0,0 +1,400 @@ +// Dense GGML decode matmul for one contiguous F32 activation row. +// +// Pipeline: +// 1. Load four masked F32 activation packets from each 256-wide K block. +// 2. Decode four matching weight packets through @ggml_dequant_f16_vector4. +// 3. Dot the activation/weight packets and reduce one contribution per output row. +// +// One 64-workitem workgroup computes two adjacent output rows. Four cohorts +// of 16 lanes each process four 256-value weight blocks in parallel, while +// each lane contracts four packets from the 0, 32, 64, and 96 element +// quarters of its block. Storage-format-specific weight decoding stays in +// dequant.loom, so decode scheduling remains format-agnostic. +template.decl @ggml.mul_mat_f32_f32_decode.body(%weight_format: index, %publish_output: i1, %token_count: index, %token0: index, %pair: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %output: buffer, %aux0: buffer, %aux1: buffer, %aux2: buffer, %aux3: buffer) + +template.decl @ggml.mul_mat_f32_f32_decode.dispatch(%weight_format: index, %token_count: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %output: buffer, %aux0: buffer, %aux1: buffer, %aux2: buffer, %aux3: buffer) + +template.decl @ggml.mul_mat_f32_f32_decode.launch(%token_capacity: index, %output_capacity: index) -> (index, index, index, index) + +template.decl @ggml.mul_mat_f32_f32_decode.publish_pair2(%publish_pair: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %row1_valid: i1, %value0: f32, %value1: f32, %output: buffer, %aux0: buffer, %aux1: buffer, %aux2: buffer, %aux3: buffer) + +template.decl @ggml.mul_mat_f32_f32_decode.dual_body(%op: index, %lhs_weight_format: index, %rhs_weight_format: index, %publish_output: i1, %token_count: index, %token0: index, %pair: index, %input_size: index, %output_size: index, %input: buffer, %lhs_weight: buffer, %rhs_weight: buffer, %output: buffer) + +template.decl @ggml.mul_mat_f32_f32_decode.dual_dispatch(%op: index, %lhs_weight_format: index, %rhs_weight_format: index, %token_capacity: index, %token_count: index, %input_size: index, %output_size: index, %input: buffer, %lhs_weight: buffer, %rhs_weight: buffer, %output: buffer) + +template.decl @ggml.mul_mat_f32_f32_decode.publish_dual_pair2(%op: index, %publish_pair: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %row1_valid: i1, %lhs0: f32, %lhs1: f32, %rhs0: f32, %rhs1: f32, %output: buffer) + +amdgpu.target @ggml_mul_mat_f32_f32_decode_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.mul_mat_f32_f32_decode.token_capacity : %value: index where [range(%value, 1, 2048)] + +config.decl @ggml.mul_mat_f32_f32_decode.output_capacity : %value: index where [range(%value, 1, 1048576)] + +config.decl @ggml.mul_mat_f32_f32_decode.weight_format : %value: index where [range(%value, 4, 81)] + +func.decl @ggml_dequant_weight_tile_bytes(%weight_format: index) -> (offset) + +func.decl @ggml_dequant_weight_row_bytes(%weight_format: index, %hidden_size: index) -> (offset) + +func.decl @ggml_iq4nl_table_i8() -> (vector<16xi8>) + +func.decl @ggml_iq3s_grid_f32_halves() -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + +func.decl @ggml_dequant_f16_vector4(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight_format: index, %weight: buffer, %row_byte_base: offset, %input_size: index, %quant_block: index, %quant_group: index, %packet: index, %k: index) -> (vector<4xf16>) + +func.def inline @ggml_mul_mat_decode_dot4_f32(%lhs: vector<4xf32>, %rhs: vector<4xf32>) -> (f32) { + %c0 = scalar.constant 0.0 : f32 + %lhs0 = vector.extract %lhs[0] : vector<4xf32> -> f32 + %lhs1 = vector.extract %lhs[1] : vector<4xf32> -> f32 + %lhs2 = vector.extract %lhs[2] : vector<4xf32> -> f32 + %lhs3 = vector.extract %lhs[3] : vector<4xf32> -> f32 + %rhs0 = vector.extract %rhs[0] : vector<4xf32> -> f32 + %rhs1 = vector.extract %rhs[1] : vector<4xf32> -> f32 + %rhs2 = vector.extract %rhs[2] : vector<4xf32> -> f32 + %rhs3 = vector.extract %rhs[3] : vector<4xf32> -> f32 + %sum0 = scalar.fmaf %lhs0, %rhs0, %c0 : f32 + %sum1 = scalar.fmaf %lhs1, %rhs1, %sum0 : f32 + %sum2 = scalar.fmaf %lhs2, %rhs2, %sum1 : f32 + %sum3 = scalar.fmaf %lhs3, %rhs3, %sum2 : f32 + func.return %sum3 : f32 +} + +// Load activation packets once per K block. The masks zero K-tail lanes for +// input sizes that are not exact 256 multiples. +func.def inline @ggml_mul_mat_decode_load_f32_block(%token_count: index, %input_size: index, %token: index, %block: index, %lane: index, %input: buffer) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c96 = index.constant 96 : index + %c128 = index.constant 128 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %bounded_lane = index.assume %lane [range(%lane, 0, 63)] : index + %padded_input_size = index.add %input_size, %c255 : index + %block_count = index.div %padded_input_size, %c256 : index + %valid_block = index.cmp ult, %block, %block_count : index + %safe_block0 = scf.select %valid_block, %block, %c0 : index + %safe_block, %launch_block_count = index.assume %safe_block0, %block_count [lt(%safe_block0, %block_count)] : index, index + %itid = index.rem %bounded_lane, %c16 : index + %vector_half = index.div %itid, %c8 : index + %vector_index = index.rem %itid, %c8 : index + %input_view = buffer.view %input[%c0_offset] : buffer -> view<[%token_count]x[%input_size]xf32> + %vector_half_base = index.mul %vector_half, %c128 : index + %vector_offset = index.mul %vector_index, %c4 : index + %input_index0 = index.add %vector_half_base, %vector_offset : index + %input_index1 = index.add %input_index0, %c32 : index + %input_index2 = index.add %input_index0, %c64 : index + %input_index3 = index.add %input_index0, %c96 : index + %block_k_base = index.mul %safe_block, %c256 : index + %k0 = index.add %block_k_base, %input_index0 : index + %k1 = index.add %block_k_base, %input_index1 : index + %k2 = index.add %block_k_base, %input_index2 : index + %k3 = index.add %block_k_base, %input_index3 : index + %mask0 = vector.mask.range [%k0 to %input_size step %c1] : index -> vector<4xi1> + %mask1 = vector.mask.range [%k1 to %input_size step %c1] : index -> vector<4xi1> + %mask2 = vector.mask.range [%k2 to %input_size step %c1] : index -> vector<4xi1> + %mask3 = vector.mask.range [%k3 to %input_size step %c1] : index -> vector<4xi1> + %input0 = vector.load.mask %input_view[%token, %k0], %mask0, %c0_f32x4 : view<[%token_count]x[%input_size]xf32>, vector<4xi1>, vector<4xf32> + %input1 = vector.load.mask %input_view[%token, %k1], %mask1, %c0_f32x4 : view<[%token_count]x[%input_size]xf32>, vector<4xi1>, vector<4xf32> + %input2 = vector.load.mask %input_view[%token, %k2], %mask2, %c0_f32x4 : view<[%token_count]x[%input_size]xf32>, vector<4xi1>, vector<4xf32> + %input3 = vector.load.mask %input_view[%token, %k3], %mask3, %c0_f32x4 : view<[%token_count]x[%input_size]xf32>, vector<4xi1>, vector<4xf32> + func.return %input0, %input1, %input2, %input3 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> +} + +// Decode one output-row contribution for the current K block. The four weight +// packets mirror the activation packet offsets: 0, 32, 64, and 96. +func.def inline @ggml_mul_mat_decode_f32_block_row(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight_format: index, %input_size: index, %row: index, %block: index, %lane: index, %weight: buffer, %input0: vector<4xf32>, %input1: vector<4xf32>, %input2: vector<4xf32>, %input3: vector<4xf32>) -> (f32) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c96 = index.constant 96 : index + %c128 = index.constant 128 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %c0_f16x4 = vector.constant 0.0 : vector<4xf16> + %c0_f32 = scalar.constant 0.0 : f32 + %bounded_lane = index.assume %lane [range(%lane, 0, 63)] : index + %padded_input_size = index.add %input_size, %c255 : index + %block_count = index.div %padded_input_size, %c256 : index + %valid_block = index.cmp ult, %block, %block_count : index + %safe_block = scf.select %valid_block, %block, %c0 : index + %itid = index.rem %bounded_lane, %c16 : index + %vector_half = index.div %itid, %c8 : index + %packet = index.rem %itid, %c8 : index + %weight_row_bytes = func.call @ggml_dequant_weight_row_bytes(%weight_format, %input_size) : (index, index) -> (offset) + %row_byte_base = index.scale %row, %weight_row_bytes : index, offset -> offset + %group_base = index.mul %vector_half, %c4 : index + %group0 = index.add %group_base, %c0 : index + %group1 = index.add %group_base, %c1 : index + %group2 = index.add %group_base, %c2 : index + %group3 = index.add %group_base, %c3 : index + %block_k_base = index.mul %safe_block, %c256 : index + %vector_half_base = index.mul %vector_half, %c128 : index + %packet_k = index.mul %packet, %c4 : index + %k_half = index.add %block_k_base, %vector_half_base : index + %k0 = index.add %k_half, %packet_k : index + %k1_add = index.add %k0, %c32 : index + %k2_add = index.add %k0, %c64 : index + %k3_add = index.add %k0, %c96 : index + %valid_k0 = index.cmp ult, %k0, %input_size : index + %valid_k1 = index.cmp ult, %k1_add, %input_size : index + %valid_k2 = index.cmp ult, %k2_add, %input_size : index + %valid_k3 = index.cmp ult, %k3_add, %input_size : index + %weight0_f16 = scf.if %valid_k0 -> (vector<4xf16>) { + %decoded = func.call @ggml_dequant_f16_vector4(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight_format, %weight, %row_byte_base, %input_size, %safe_block, %group0, %packet, %k0) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, buffer, offset, index, index, index, index, index) -> (vector<4xf16>) + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + %weight1_f16 = scf.if %valid_k1 -> (vector<4xf16>) { + %decoded = func.call @ggml_dequant_f16_vector4(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight_format, %weight, %row_byte_base, %input_size, %safe_block, %group1, %packet, %k1_add) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, buffer, offset, index, index, index, index, index) -> (vector<4xf16>) + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + %weight2_f16 = scf.if %valid_k2 -> (vector<4xf16>) { + %decoded = func.call @ggml_dequant_f16_vector4(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight_format, %weight, %row_byte_base, %input_size, %safe_block, %group2, %packet, %k2_add) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, buffer, offset, index, index, index, index, index) -> (vector<4xf16>) + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + %weight3_f16 = scf.if %valid_k3 -> (vector<4xf16>) { + %decoded = func.call @ggml_dequant_f16_vector4(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight_format, %weight, %row_byte_base, %input_size, %safe_block, %group3, %packet, %k3_add) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, buffer, offset, index, index, index, index, index) -> (vector<4xf16>) + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + %weight0 = vector.extf %weight0_f16 : vector<4xf16> to vector<4xf32> + %weight1 = vector.extf %weight1_f16 : vector<4xf16> to vector<4xf32> + %weight2 = vector.extf %weight2_f16 : vector<4xf16> to vector<4xf32> + %weight3 = vector.extf %weight3_f16 : vector<4xf16> to vector<4xf32> + %dot0 = func.call @ggml_mul_mat_decode_dot4_f32(%input0, %weight0) : (vector<4xf32>, vector<4xf32>) -> (f32) + %dot1 = func.call @ggml_mul_mat_decode_dot4_f32(%input1, %weight1) : (vector<4xf32>, vector<4xf32>) -> (f32) + %dot2 = func.call @ggml_mul_mat_decode_dot4_f32(%input2, %weight2) : (vector<4xf32>, vector<4xf32>) -> (f32) + %dot3 = func.call @ggml_mul_mat_decode_dot4_f32(%input3, %weight3) : (vector<4xf32>, vector<4xf32>) -> (f32) + %sum3 = scalar.addf %dot2, %dot3 : f32 + %sum2 = scalar.addf %dot1, %sum3 : f32 + %sum = scalar.addf %dot0, %sum2 : f32 + %contribution = scf.select %valid_block, %sum, %c0_f32 : f32 + func.return %contribution : f32 +} + +template.def<@ggml.mul_mat_f32_f32_decode.publish_pair2> device priority(1) @ggml_mul_mat_f32_f32_decode_publish_pair2(%publish_pair: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %row1_valid: i1, %value0: f32, %value1: f32, %output: buffer, %aux0: buffer, %aux1: buffer, %aux2: buffer, %aux3: buffer) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 1, 1048576)] : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %token, %launch_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %output_noalias = buffer.assume.noalias %output : buffer + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%output_size]xf32> + scf.if %publish_pair { + view.store %value0, %output_view[%token, %channel] : f32, view<[%launch_token_count]x[%output_size]xf32> + } + scf.if %row1_valid { + scf.if %publish_pair { + %row1 = index.add %channel, %c1 : index + view.store %value1, %output_view[%token, %row1] : f32, view<[%launch_token_count]x[%output_size]xf32> + } + } + template.return +} + +template.def<@ggml.mul_mat_f32_f32_decode.publish_dual_pair2> device priority(1) @ggml_mul_mat_f32_f32_decode_publish_dual_pair2(%op: index, %publish_pair: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %row1_valid: i1, %lhs0: f32, %lhs1: f32, %rhs0: f32, %rhs1: f32, %output: buffer) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 1, 1048576)] : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %token, %launch_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %output_noalias = buffer.assume.noalias %output : buffer + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%output_size]xf32> + scf.if %publish_pair { + view.store %lhs0, %output_view[%token, %channel] : f32, view<[%launch_token_count]x[%output_size]xf32> + } + scf.if %row1_valid { + scf.if %publish_pair { + %row1 = index.add %channel, %c1 : index + view.store %lhs1, %output_view[%token, %row1] : f32, view<[%launch_token_count]x[%output_size]xf32> + } + } + template.return +} + +template.def<@ggml.mul_mat_f32_f32_decode.body> device @ggml_mul_mat_f32_f32_decode_body(%weight_format: index, %publish_output: i1, %token_count: index, %token0: index, %pair: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %output: buffer, %aux0: buffer, %aux1: buffer, %aux2: buffer, %aux3: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 32)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 1048576)] : index + %lane0 = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %c0_f32 = scalar.constant 0.0 : f32 + %c0_offset = index.constant 0 : offset + %token, %launch_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %lane = index.assume %lane0 [range(%lane0, 0, 63)] : index + %row00 = index.mul %pair, %c2 : index + %row0, %launch_output_size = index.assume %row00, %bounded_output_size [lt(%row00, %bounded_output_size)] : index, index + %row1 = index.add %row0, %c1 : index + %row1_valid = index.cmp ult, %row1, %launch_output_size : index + %padded_input_size = index.add %bounded_input_size, %c255 : index + %block_count = index.div %padded_input_size, %c256 : index + %cohort = index.div %lane, %c16 : index + scf.if %publish_output { + %input_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input, %weight, %output : buffer, buffer, buffer + %iq4nl_table = func.call @ggml_iq4nl_table_i8() : () -> (vector<16xi8>) + %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15 = func.call @ggml_iq3s_grid_f32_halves() : () -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + %acc0, %acc1 = scf.for %block_base = [%c0 to %block_count step %c4](%row_acc0 = %c0_f32 : f32, %row_acc1 = %c0_f32 : f32) -> (f32, f32) { + %block = index.add %block_base, %cohort : index + %input0, %input1, %input2, %input3 = func.call @ggml_mul_mat_decode_load_f32_block(%launch_token_count, %bounded_input_size, %token, %block, %lane, %input_noalias) : (index, index, index, index, index, buffer) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) + %contribution0 = func.call @ggml_mul_mat_decode_f32_block_row(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight_format, %bounded_input_size, %row0, %block, %lane, %weight_noalias, %input0, %input1, %input2, %input3) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, index, index, index, index, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) -> (f32) + %contribution1 = scf.if %row1_valid -> (f32) { + %row1_contribution = func.call @ggml_mul_mat_decode_f32_block_row(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight_format, %bounded_input_size, %row1, %block, %lane, %weight_noalias, %input0, %input1, %input2, %input3) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, index, index, index, index, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) -> (f32) + scf.yield %row1_contribution : f32 + } else { + scf.yield %c0_f32 : f32 + } + %next0 = scalar.addf %row_acc0, %contribution0 : f32 + %next1 = scalar.addf %row_acc1, %contribution1 : f32 + scf.yield %next0, %next1 : f32, f32 + } + %sum0 = kernel.workgroup.reduce %acc0 : f32 + %sum1 = kernel.workgroup.reduce %acc1 : f32 + %is_lane_zero = index.cmp eq, %lane, %c0 : index + %publish_pair = scalar.andi %publish_output, %is_lane_zero : i1 + template.apply<@ggml.mul_mat_f32_f32_decode.publish_pair2>(%publish_pair, %token, %row0, %launch_token_count, %launch_output_size, %row1_valid, %sum0, %sum1, %output_noalias, %aux0, %aux1, %aux2, %aux3) : (i1, index, index, index, index, i1, f32, f32, buffer, buffer, buffer, buffer, buffer) + } + template.return +} + +template.def<@ggml.mul_mat_f32_f32_decode.dual_body> device @ggml_mul_mat_f32_f32_decode_dual_body(%op: index, %lhs_weight_format: index, %rhs_weight_format: index, %publish_output: i1, %token_count: index, %token0: index, %pair: index, %input_size: index, %output_size: index, %input: buffer, %lhs_weight: buffer, %rhs_weight: buffer, %output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 32)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 1048576)] : index + %lane0 = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %c0_f32 = scalar.constant 0.0 : f32 + %token, %launch_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %lane = index.assume %lane0 [range(%lane0, 0, 63)] : index + %row00 = index.mul %pair, %c2 : index + %row0, %launch_output_size = index.assume %row00, %bounded_output_size [lt(%row00, %bounded_output_size)] : index, index + %row1 = index.add %row0, %c1 : index + %row1_valid = index.cmp ult, %row1, %launch_output_size : index + %padded_input_size = index.add %bounded_input_size, %c255 : index + %block_count = index.div %padded_input_size, %c256 : index + %cohort = index.div %lane, %c16 : index + scf.if %publish_output { + %input_noalias, %lhs_weight_noalias, %rhs_weight_noalias, %output_noalias = buffer.assume.noalias %input, %lhs_weight, %rhs_weight, %output : buffer, buffer, buffer, buffer + %iq4nl_table = func.call @ggml_iq4nl_table_i8() : () -> (vector<16xi8>) + %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15 = func.call @ggml_iq3s_grid_f32_halves() : () -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + %lhs_acc0, %lhs_acc1, %rhs_acc0, %rhs_acc1 = scf.for %block_base = [%c0 to %block_count step %c4](%lhs_row_acc0 = %c0_f32 : f32, %lhs_row_acc1 = %c0_f32 : f32, %rhs_row_acc0 = %c0_f32 : f32, %rhs_row_acc1 = %c0_f32 : f32) -> (f32, f32, f32, f32) { + %block = index.add %block_base, %cohort : index + %input0, %input1, %input2, %input3 = func.call @ggml_mul_mat_decode_load_f32_block(%launch_token_count, %bounded_input_size, %token, %block, %lane, %input_noalias) : (index, index, index, index, index, buffer) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) + %lhs_contribution0 = func.call @ggml_mul_mat_decode_f32_block_row(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %lhs_weight_format, %bounded_input_size, %row0, %block, %lane, %lhs_weight_noalias, %input0, %input1, %input2, %input3) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, index, index, index, index, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) -> (f32) + %rhs_contribution0 = func.call @ggml_mul_mat_decode_f32_block_row(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %rhs_weight_format, %bounded_input_size, %row0, %block, %lane, %rhs_weight_noalias, %input0, %input1, %input2, %input3) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, index, index, index, index, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) -> (f32) + %lhs_contribution1, %rhs_contribution1 = scf.if %row1_valid -> (f32, f32) { + %lhs_row1_contribution = func.call @ggml_mul_mat_decode_f32_block_row(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %lhs_weight_format, %bounded_input_size, %row1, %block, %lane, %lhs_weight_noalias, %input0, %input1, %input2, %input3) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, index, index, index, index, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) -> (f32) + %rhs_row1_contribution = func.call @ggml_mul_mat_decode_f32_block_row(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %rhs_weight_format, %bounded_input_size, %row1, %block, %lane, %rhs_weight_noalias, %input0, %input1, %input2, %input3) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, index, index, index, index, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) -> (f32) + scf.yield %lhs_row1_contribution, %rhs_row1_contribution : f32, f32 + } else { + scf.yield %c0_f32, %c0_f32 : f32, f32 + } + %lhs_next0 = scalar.addf %lhs_row_acc0, %lhs_contribution0 : f32 + %lhs_next1 = scalar.addf %lhs_row_acc1, %lhs_contribution1 : f32 + %rhs_next0 = scalar.addf %rhs_row_acc0, %rhs_contribution0 : f32 + %rhs_next1 = scalar.addf %rhs_row_acc1, %rhs_contribution1 : f32 + scf.yield %lhs_next0, %lhs_next1, %rhs_next0, %rhs_next1 : f32, f32, f32, f32 + } + %lhs_sum0 = kernel.workgroup.reduce %lhs_acc0 : f32 + %lhs_sum1 = kernel.workgroup.reduce %lhs_acc1 : f32 + %rhs_sum0 = kernel.workgroup.reduce %rhs_acc0 : f32 + %rhs_sum1 = kernel.workgroup.reduce %rhs_acc1 : f32 + %is_lane_zero = index.cmp eq, %lane, %c0 : index + %publish_pair = scalar.andi %publish_output, %is_lane_zero : i1 + template.apply<@ggml.mul_mat_f32_f32_decode.publish_dual_pair2>(%op, %publish_pair, %token, %row0, %launch_token_count, %launch_output_size, %row1_valid, %lhs_sum0, %lhs_sum1, %rhs_sum0, %rhs_sum1, %output_noalias) : (index, i1, index, index, index, index, i1, f32, f32, f32, f32, buffer) + } + template.return +} + +template.def<@ggml.mul_mat_f32_f32_decode.launch> @ggml_mul_mat_f32_f32_decode_launch(%token_capacity: index, %output_capacity: index) -> (index, index, index, index) { + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c64 = index.constant 64 : index + %padded_output_size = index.add %output_capacity, %c1 : index + %output_pairs = index.div %padded_output_size, %c2 : index + template.return %output_pairs, %token_capacity, %c1, %c64 : index, index, index, index +} + +template.def<@ggml.mul_mat_f32_f32_decode.dual_dispatch> device @ggml_mul_mat_f32_f32_decode_dual_dispatch(%op: index, %lhs_weight_format: index, %rhs_weight_format: index, %token_capacity: index, %token_count: index, %input_size: index, %output_size: index, %input: buffer, %lhs_weight: buffer, %rhs_weight: buffer, %output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 1048576)] : index + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %pair = kernel.workgroup.id : index + %token = kernel.workgroup.id : index + %row0 = index.mul %pair, %c2 : index + %valid_token = index.cmp ult, %token, %bounded_token_count : index + %valid_row = index.cmp ult, %row0, %bounded_output_size : index + %publish_output = scalar.andi %valid_token, %valid_row : i1 + %safe_token = scf.select %valid_token, %token, %c0 : index + %safe_pair = scf.select %valid_row, %pair, %c0 : index + template.apply<@ggml.mul_mat_f32_f32_decode.dual_body>(%op, %lhs_weight_format, %rhs_weight_format, %publish_output, %bounded_token_count, %safe_token, %safe_pair, %input_size, %bounded_output_size, %input, %lhs_weight, %rhs_weight, %output) : (index, index, index, i1, index, index, index, index, index, buffer, buffer, buffer, buffer) + template.return +} + +template.def<@ggml.mul_mat_f32_f32_decode.dispatch> device @ggml_mul_mat_f32_f32_decode_dispatch(%weight_format: index, %token_count: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %output: buffer, %aux0: buffer, %aux1: buffer, %aux2: buffer, %aux3: buffer) { + %token_capacity = config.get @ggml.mul_mat_f32_f32_decode.token_capacity : index + %output_capacity = config.get @ggml.mul_mat_f32_f32_decode.output_capacity : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 1048576), le(%output_size, %output_capacity)] : index + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %pair = kernel.workgroup.id : index + %token = kernel.workgroup.id : index + %row0 = index.mul %pair, %c2 : index + %valid_token = index.cmp ult, %token, %bounded_token_count : index + %valid_row = index.cmp ult, %row0, %bounded_output_size : index + %publish_output = scalar.andi %valid_token, %valid_row : i1 + %safe_token = scf.select %valid_token, %token, %c0 : index + %safe_pair = scf.select %valid_row, %pair, %c0 : index + template.apply<@ggml.mul_mat_f32_f32_decode.body>(%weight_format, %publish_output, %bounded_token_count, %safe_token, %safe_pair, %input_size, %bounded_output_size, %input, %weight, %output, %aux0, %aux1, %aux2, %aux3) : (index, i1, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + template.return +} + +kernel.def target(@ggml_mul_mat_f32_f32_decode_gfx11_wave64) @ggml_mul_mat_f32_f32_decode_wave64(%token_count: index, %input_size: index, %output_size: index) { + %token_capacity = config.get @ggml.mul_mat_f32_f32_decode.token_capacity : index + %output_capacity = config.get @ggml.mul_mat_f32_f32_decode.output_capacity : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %padded_output_size = index.add %output_capacity, %c1 : index + %output_pairs = index.div %padded_output_size, %c2 : index + %padded_token_count = index.add %token_capacity, %c63 : index + %launch_tokens = index.div %padded_token_count, %c64 : index + kernel.launch.config workgroups(%output_pairs, %launch_tokens, %c1) workgroup_size(%c64, %c1, %c1) : index +} launch(%token_count: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %output: buffer) { + %weight_format = config.get @ggml.mul_mat_f32_f32_decode.weight_format : index + template.apply<@ggml.mul_mat_f32_f32_decode.dispatch>(%weight_format, %token_count, %input_size, %output_size, %input, %weight, %output, %output, %output, %output, %output) : (index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_f32_f32_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_f32_f32_wmma.loom new file mode 100644 index 000000000000..209d2650725a --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_f32_f32_wmma.loom @@ -0,0 +1,299 @@ +func.decl @ggml_mul_mat_quantized_f16_prefill_conv4_interior(%weight_format: index, %token_count: index, %input_size: index, %output_size: index, %channel_tile: index, %input: buffer, %weight: buffer, %filter: buffer, %output: buffer, %edges: buffer) + +func.decl @ggml_mul_mat_quantized_f16_prefill_wave32(%weight_format: index, %binary_op: index, %paired: i1, %packed_input: i1, %packed_output: i1, %token_count: index, %input_size: index, %output_size: index, %channel_tile: index, %token_tile: index, %input: buffer, %weight: buffer, %up_weight: buffer, %output: buffer, %f16_output: buffer) + +// Dense GGML matmul for contiguous 2D F32 activations. +// +// Each two-wave workgroup computes 64 output channels for 32 contiguous +// tokens. GGUF weight rows are decoded or converted directly into a padded +// FP16 LDS tile, while the same workgroup converts the corresponding F32 +// activation tile to FP16. Four wave64 WMMA accumulators cover the 32x32 result +// owned by each wave. A wave-private LDS slice transposes each accumulator for +// coalesced F32 publication. +// +// Weight contracts use the unmodified GGUF layout: +// Q3_K: [output channel][input size / 256][110 bytes] +// Q4_K: [output channel][input size / 256][144 bytes] +// Q5_K: [output channel][input size / 256][176 bytes] +// Q6_K: [output channel][input size / 256][210 bytes] +// IQ4_NL: [output channel][input size / 32][18 bytes] +// IQ3_S: [output channel][input size / 256][110 bytes] +// IQ4_XS: [output channel][input size / 256][136 bytes] +// Q8_0: [output channel][input size / 32][34 bytes] +// Q8_1: [output channel][input size / 32][36 bytes] +// F16: [output channel][input size][2 bytes] +// BF16: [output channel][input size][2 bytes] +// F32: [output channel][input size][4 bytes] +// The configured weight format selects the storage reader before the shared +// device template is instantiated, so inactive decode logic is absent from the +// emitted kernel. No persistent repacking or expanded-weight allocation is +// required. +template.decl @ggml.mul_mat_f32_f32_wmma.core(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: f32, %arg7: index, %arg8: index, %arg9: buffer, %arg10: buffer, %arg11: buffer, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer, %arg16: buffer, %arg17: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_f32_f32_wmma.launch(%token_capacity: index, %input_size: index, %output_size: index, %weight_format: index) -> (index, index, index, index) + +template.decl @ggml.mul_mat_f32_f32_wmma.publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.unary_f32.apply_vector4(%arg0: index, %arg1: vector<4xf32>) -> (vector<4xf32>) + +amdgpu.target @ggml_mul_mat_gfx11_wave64 {subgroup_size = 64} +amdgpu.target @ggml_mul_mat_gfx11_wave32 {subgroup_size = 32} + +config.decl @ggml.mul_mat.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 32)] + +config.decl @ggml.mul_mat.output_size : %value: index where [range(%value, 1, 262144)] + +// Private F16 activation storage for the Q4 prefill entry: row-major (0) or +// K16-major [K/16, tokens, 16] (1). The latter uses full 512-token tiles. +config.decl @ggml.mul_mat.f16_input_layout : %value: index where [range(%value, 0, 1)] + +// Selects whether publication overwrites the output (0) or accumulates the +// projection into a caller-provided residual (1). +config.decl @ggml.mul_mat.output_accumulation : %value: index where [range(%value, 0, 1)] + +// Unary op applied to the published F32 output. Identity is op 23. +config.decl @ggml.mul_mat.output_unary_op : %value: index where [range(%value, 0, 23)] + +config.decl @ggml.mul_mat.weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +template.def<@ggml.mul_mat_f32_f32_wmma.publish_vector4> device @ggml_mul_mat_f32_f32_wmma_publish_vector4(%publish_word: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %output_accumulation: index, %output_unary_op: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 1, 262144)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %output_noalias = buffer.assume.noalias %output : buffer + %partition = kernel.workgroup.id : index + %element_bytes = index.constant 4 : offset + %plane_elements = index.mul %token_count, %output_size : index + %plane_bytes = index.scale %plane_elements, %element_bytes : index, offset -> offset + %plane_offset = index.scale %partition, %plane_bytes : index, offset -> offset + %output_view = buffer.view %output_noalias[%plane_offset] : buffer -> view<[%token_count]x[%output_size]xf32> + %accumulates_output = index.cmp eq, %output_accumulation, %c1 : index + scf.if %publish_word { + %token, %launch_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %mask = vector.mask.range [%channel to %output_size step %c1] : index -> vector<4xi1> + %published = scf.if %accumulates_output -> (vector<4xf32>) { + %residual = vector.load.mask %output_view[%token, %channel], %mask, %c0_f32x4 : view<[%token_count]x[%output_size]xf32>, vector<4xi1>, vector<4xf32> + %sum = vector.addf %residual, %values : vector<4xf32> + scf.yield %sum : vector<4xf32> + } else { + scf.yield %values : vector<4xf32> + } + %activated = template.apply<@ggml.unary_f32.apply_vector4>(%output_unary_op, %published) : (index, vector<4xf32>) -> (vector<4xf32>) + vector.store.mask %activated, %output_view[%token, %channel], %mask : vector<4xf32>, view<[%token_count]x[%output_size]xf32>, vector<4xi1> + } + template.return +} + +template.def<@ggml.mul_mat_f32_f32_wmma.finish_tile> device @ggml_mul_mat_f32_f32_wmma_finish_tile(%token_count: index, %output_size: index, %output_tile_count: index, %token_tile_count: index, %output_accumulation: index, %output_unary_op: index, %epsilon: f32, %channel_tile: index, %token_tile: index, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + template.return +} + +// Shape and format configs specialize the native code once while token count +// remains the only per-dispatch scalar. +kernel.def target(@ggml_mul_mat_gfx11_wave64) @ggml_mul_mat_f32_f32_wmma(%token_count: index) { + %output_size = config.get @ggml.mul_mat.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %c1 = index.constant 1 : index + %input_size = config.get @ggml.mul_mat.input_size : index + %weight_format = config.get @ggml.mul_mat.weight_format : index + %launch_x, %launch_y, %launch_z, %workgroup_size = template.apply<@ggml.mul_mat_f32_f32_wmma.launch>(%token_capacity, %input_size, %output_size, %weight_format) pure : (index, index, index, index) -> (index, index, index, index) + kernel.launch.config workgroups(%launch_x, %launch_y, %launch_z) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) { + %input_size = config.get @ggml.mul_mat.input_size : index + %output_size = config.get @ggml.mul_mat.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %output_accumulation = config.get @ggml.mul_mat.output_accumulation : index + %output_unary_op = config.get @ggml.mul_mat.output_unary_op : index + %weight_format = config.get @ggml.mul_mat.weight_format : index + %c0_f32 = scalar.constant 0.0 : f32 + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_f32_f32_wmma.core>(%weight_format, %bounded_token_count, %input_size, %output_size, %output_accumulation, %output_unary_op, %c0_f32, %channel_tile, %token_tile, %input, %weight, %output, %output, %output, %output, %output, %output, %output) : (index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +kernel.def target(@ggml_mul_mat_gfx11_wave32) @ggml_mul_mat_q4_k_f16_wmma_prefill_wave32(%token_count: index) { + %input_size = config.get @ggml.mul_mat.input_size : index + %output_size = config.get @ggml.mul_mat.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %c512 = index.constant 512 : index + %c1024 = index.constant 1024 : index + %minimum_workgroups = index.constant 32 : index + %twice_input_size = index.mul %input_size, %c2 : index + %contracts = index.cmp ule, %output_size, %input_size : index + %large_expansion = index.cmp uge, %output_size, %twice_input_size : index + %wide_shape = scalar.ori %contracts, %large_expansion : i1 + %output_tail = index.rem %output_size, %c128 : index + %full_wide_tile = index.cmp eq, %output_tail, %c0 : index + %wide_tile = scalar.andi %wide_shape, %full_wide_tile : i1 + %ordinary_tile_channels = scf.select %wide_tile, %c128, %c64 : index + %ordinary_workgroup_size = scf.select %wide_tile, %c512, %c256 : index + %active_token_tiles = index.div %token_count, %c256 : index + %larger_output_tiles = index.div %output_size, %c256 : index + %larger_grid = index.mul %active_token_tiles, %larger_output_tiles : index + %enough_workgroups = index.cmp uge, %larger_grid, %minimum_workgroups : index + %larger_channel_tail = index.rem %output_size, %c256 : index + %full_larger_channels = index.cmp eq, %larger_channel_tail, %c0 : index + %multiple_token_tiles = index.cmp ugt, %active_token_tiles, %c1 : index + %larger_shape = scalar.andi %contracts, %full_larger_channels : i1 + %larger_grid_tile = scalar.andi %larger_shape, %enough_workgroups : i1 + %larger_tile = scalar.andi %larger_grid_tile, %multiple_token_tiles : i1 + %staged_tile_channels = scf.select %larger_tile, %c256, %ordinary_tile_channels : index + %staged_workgroup_size = scf.select %larger_tile, %c1024, %ordinary_workgroup_size : index + %staged_output_tiles = index.div %output_size, %staged_tile_channels : index + %staged_token_tiles = index.div %token_capacity, %c256 : index + %input_layout = config.get @ggml.mul_mat.f16_input_layout : index + %packed = index.cmp eq, %input_layout, %c1 : index + %workgroup_size = scf.select %packed, %c512, %staged_workgroup_size : index + %token_tiles = scf.select %packed, %c1, %staged_token_tiles : index + %packed_output_tiles = index.div %output_size, %c64 : index + %output_tiles = scf.select %packed, %packed_output_tiles, %staged_output_tiles : index + kernel.launch.config workgroups(%token_tiles, %output_tiles, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) { + %input_size = config.get @ggml.mul_mat.input_size : index + %output_size = config.get @ggml.mul_mat.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %bounded_token_count = index.assume %token_count [range(%token_count, 256, 2048), mul(%token_count, 256), le(%token_count, %token_capacity)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 64, 262144), mul(%output_size, 64)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + %paired = scalar.constant false : i1 + %input_layout = config.get @ggml.mul_mat.f16_input_layout : index + %packed_layout = index.constant 1 : index + %packed_input = index.cmp eq, %input_layout, %packed_layout : index + %prefill_weight_format = index.constant 44 : index + %prefill_binary_op = index.constant 0 : index + func.call @ggml_mul_mat_quantized_f16_prefill_wave32(%prefill_weight_format, %prefill_binary_op, %paired, %packed_input, %paired, %bounded_token_count, %input_size, %bounded_output_size, %channel_tile, %token_tile, %input, %weight, %weight, %output, %output) : (index, index, i1, i1, i1, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +kernel.def target(@ggml_mul_mat_gfx11_wave64) @ggml_mul_mat_f32_f32_narrow_split_k4(%token_count: index) { + %token_capacity = config.get @ggml.workload.token_capacity : index + %one = index.constant 1 : index + %split_count = index.constant 4 : index + %tile_tokens = index.constant 32 : index + %workgroup_size = index.constant 128 : index + %token_tiles = index.div %token_capacity, %tile_tokens : index + kernel.launch.config workgroups(%one, %token_tiles, %split_count) workgroup_size(%workgroup_size, %one, %one) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer, %final_output: buffer, %counters: buffer) { + %input_size0 = config.get @ggml.mul_mat.input_size : index + %input_size = index.assume %input_size0 [range(%input_size0, 4096, 32768), mul(%input_size0, 1024)] : index + %output_size0 = config.get @ggml.mul_mat.output_size : index + %output_size = index.assume %output_size0 [range(%output_size0, 4, 64), mul(%output_size0, 4)] : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %output_accumulation = index.constant 0 : index + %output_unary_op = index.constant 23 : index + %weight_format = config.get @ggml.mul_mat.weight_format : index + %c0_f32 = scalar.constant 0.0 : f32 + %bounded_token_count = index.assume %token_count [range(%token_count, 32, 2048), mul(%token_count, 32), le(%token_count, %token_capacity), le(%token_capacity, %token_count)] : index + %channel_tile = kernel.workgroup.id : index + %split_count = index.constant 4 : index + %token_tile = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_f32_f32_wmma.core>(%weight_format, %bounded_token_count, %input_size, %output_size, %output_accumulation, %output_unary_op, %c0_f32, %channel_tile, %token_tile, %input, %weight, %output, %output, %output, %output, %output, %output, %output) : (index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + %base = index.constant 0 : offset + %zero = index.constant 0 : index + %four = index.constant 4 : index + %tile_m = index.constant 32 : index + %step = index.constant 512 : index + %zero_i32 = scalar.constant 0 : i32 + %one_i32 = scalar.constant 1 : i32 + %three_i32 = scalar.constant 3 : i32 + %negative_four_i32 = scalar.constant -4 : i32 + %tid = kernel.workitem.id : index + %first_thread = index.cmp eq, %tid, %zero : index + %counter_count = index.div %token_capacity, %tile_m : index + %counter_view = buffer.view %counters[%base] : buffer -> view<[%counter_count]xi32> + kernel.barrier scope(workgroup) ordering(release) + %arrival_local = scf.if %first_thread -> (i32) { + %previous = view.atomic.rmw %one_i32, %counter_view[%token_tile] {ordering = acq_rel, scope = device} : i32, view<[%counter_count]xi32> -> i32 + scf.yield %previous : i32 + } else { + scf.yield %zero_i32 : i32 + } + %arrival = kernel.workgroup.reduce %arrival_local : i32 + %last = scalar.cmpi eq, %arrival, %three_i32 : i32 + scf.if %last { + kernel.barrier scope(workgroup) ordering(acquire) + %byte_size = index.constant 4 : offset + %plane_elements = index.mul %bounded_token_count, %output_size : index + %plane_bytes = index.scale %plane_elements, %byte_size : index, offset -> offset + %plane2 = index.add %plane_bytes, %plane_bytes : offset + %plane3 = index.add %plane2, %plane_bytes : offset + %partial0 = buffer.view %output[%base] : buffer -> view<[%bounded_token_count]x[%output_size]xf32> + %partial1 = buffer.view %output[%plane_bytes] : buffer -> view<[%bounded_token_count]x[%output_size]xf32> + %partial2 = buffer.view %output[%plane2] : buffer -> view<[%bounded_token_count]x[%output_size]xf32> + %partial3 = buffer.view %output[%plane3] : buffer -> view<[%bounded_token_count]x[%output_size]xf32> + %final_view = buffer.view %final_output[%base] : buffer -> view<[%bounded_token_count]x[%output_size]xf32> + %token_origin = index.mul %token_tile, %tile_m : index + %elements = index.mul %tile_m, %output_size : index + %first_element = index.mul %tid, %four : index + scf.for %linear = [%first_element to %elements step %step] { + %row = index.div %linear, %output_size : index + %channel = index.rem %linear, %output_size : index + %token = index.add %token_origin, %row : index + %v0 = vector.load %partial0[%token, %channel] : view<[%bounded_token_count]x[%output_size]xf32> -> vector<4xf32> + %v1 = vector.load %partial1[%token, %channel] : view<[%bounded_token_count]x[%output_size]xf32> -> vector<4xf32> + %v2 = vector.load %partial2[%token, %channel] : view<[%bounded_token_count]x[%output_size]xf32> -> vector<4xf32> + %v3 = vector.load %partial3[%token, %channel] : view<[%bounded_token_count]x[%output_size]xf32> -> vector<4xf32> + %sum01 = vector.addf %v0, %v1 : vector<4xf32> + %sum012 = vector.addf %sum01, %v2 : vector<4xf32> + %sum = vector.addf %sum012, %v3 : vector<4xf32> + vector.store %sum, %final_view[%token, %channel] : vector<4xf32>, view<[%bounded_token_count]x[%output_size]xf32> + scf.yield + } + kernel.barrier scope(workgroup) ordering(release) + scf.if %first_thread { + view.atomic.reduce %negative_four_i32, %counter_view[%token_tile] {ordering = release, scope = device} : i32, view<[%counter_count]xi32> + } + } + kernel.return +} + +func.decl @ggml_mul_mat_quantized_f16_prefill_conv4(%weight_format: index, %token_count: index, %input_size: index, %output_size: index, %channel_tile: index, %input: buffer, %weight: buffer, %state: buffer, %filter: buffer, %output: buffer, %cache: buffer) + +kernel.def target(@ggml_mul_mat_gfx11_wave32) @ggml_mul_mat_quantized_f16_wmma_prefill_conv4(%token_count: index) { + %output_size = config.get @ggml.mul_mat.output_size : index + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %threads = index.constant 512 : index + %groups = index.div %output_size, %sixtyfour : index + kernel.launch.config workgroups(%one, %groups, %one) workgroup_size(%threads, %one, %one) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %state: buffer, %filter: buffer, %output: buffer, %cache: buffer) { + %input_size = config.get @ggml.mul_mat.input_size : index + %output_size = config.get @ggml.mul_mat.output_size : index + %weight_format = config.get @ggml.mul_mat.weight_format : index + %channel_tile = kernel.workgroup.id : index + func.call @ggml_mul_mat_quantized_f16_prefill_conv4(%weight_format, %token_count, %input_size, %output_size, %channel_tile, %input, %weight, %state, %filter, %output, %cache) : (index, index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +// Recurrent state is consumed by the later finish dispatch. +kernel.def target(@ggml_mul_mat_gfx11_wave32) @ggml_mul_mat_quantized_f16_wmma_prefill_conv4_interior(%token_count: index) { + %output_size = config.get @ggml.mul_mat.output_size : index + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %threads = index.constant 512 : index + %groups = index.div %output_size, %sixtyfour : index + kernel.launch.config workgroups(%one, %groups, %one) workgroup_size(%threads, %one, %one) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %filter: buffer, %output: buffer, %edges: buffer) { + %input_size = config.get @ggml.mul_mat.input_size : index + %output_size = config.get @ggml.mul_mat.output_size : index + %weight_format = config.get @ggml.mul_mat.weight_format : index + %channel_tile = kernel.workgroup.id : index + func.call @ggml_mul_mat_quantized_f16_prefill_conv4_interior(%weight_format, %token_count, %input_size, %output_size, %channel_tile, %input, %weight, %filter, %output, %edges) : (index, index, index, index, index, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_decode_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_decode_f32.loom new file mode 100644 index 000000000000..28eb3388734d --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_decode_f32.loom @@ -0,0 +1,144 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// One-token MUL_MAT_ID (a MoE expert matmul at decode) as a GEMV: each wave64 reads only the +// selected expert's rows. The per-block dequant/dot loop is the one in +// ops/mul_mat_f32_f32_decode.loom (Copyright The HRX Authors, Apache-2.0), linked as a library +// and copied here unchanged apart from the weight row (expert * output_size + row) and the input +// row. The WMMA mul_mat_id kernel does the same work for one token at ~45 GB/s on ZAYA1-8B. + +config.decl @ggml.mul_mat_id_decode.weight_format : %value: index where [range(%value, 4, 81)] +config.decl @ggml.mul_mat_id_decode.row_capacity : %value: index where [range(%value, 1, 2048)] +config.decl @ggml.mul_mat_id_decode.output_capacity : %value: index where [range(%value, 1, 1048576)] + +amdgpu.target @ggml_mul_mat_id_decode_gfx11_wave64 {subgroup_size = 64} + +func.decl @ggml_iq4nl_table_i8() -> (vector<16xi8>) + +func.decl @ggml_iq3s_grid_f32_halves() -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + +func.decl @ggml_mul_mat_decode_load_f32_block(%token_count: index, %input_size: index, %token: index, %block: index, %lane: index, %input: buffer) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) + +func.decl @ggml_mul_mat_decode_f32_block_row(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight_format: index, %input_size: index, %row: index, %block: index, %lane: index, %weight: buffer, %input0: vector<4xf32>, %input1: vector<4xf32>, %input2: vector<4xf32>, %input3: vector<4xf32>) -> (f32) + +// One workgroup (one wave64) per (output-row pair, slot, token): row pair x, slot y, token z, so +// no thread divides by a runtime count (AMDGPU Loom lowers only constant divisors). The slot's +// expert comes from ids[token][slot]; the input row is token * input_rows + (input_rows == 1 ? +// 0 : slot): input_rows is 1 when every slot reads the same activation, slot_count when each +// slot has its own (the matcher admits only those two). ids are [token][route_stride] i32 (llama.cpp's +// route ids are a strided view of the argsort). Output layout [token][slot][output_size]. +kernel.def target(@ggml_mul_mat_id_decode_gfx11_wave64) export("ggml_mul_mat_id_decode_f32_wave64") @ggml_mul_mat_id_decode_f32_wave64(%token_count: index, %slot_count: index, %input_rows: index, %input_size: index, %output_size: index, %expert_count: index, %route_stride: index) { + %row_capacity = config.get @ggml.mul_mat_id_decode.row_capacity : index + %output_capacity = config.get @ggml.mul_mat_id_decode.output_capacity : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c64 = index.constant 64 : index + %padded_output_size = index.add %output_capacity, %c1 : index + %output_pairs = index.div %padded_output_size, %c2 : index + kernel.launch.config workgroups(%output_pairs, %slot_count, %token_count) workgroup_size(%c64, %c1, %c1) : index +} launch(%token_count: index, %slot_count: index, %input_rows: index, %input_size: index, %output_size: index, %expert_count: index, %route_stride: index, %input: buffer, %weight: buffer, %ids: buffer, %output: buffer) { + %weight_format = config.get @ggml.mul_mat_id_decode.weight_format : index + %row_capacity = config.get @ggml.mul_mat_id_decode.row_capacity : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %c0_f32 = scalar.constant 0.0 : f32 + %c0_offset = index.constant 0 : offset + %bounded_tokens = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_slots = index.assume %slot_count [range(%slot_count, 1, 64)] : index + %bounded_input_rows = index.assume %input_rows [range(%input_rows, 1, 64)] : index + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 32)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 1048576)] : index + %flat_count0 = index.mul %bounded_tokens, %bounded_slots : index + %flat_count = index.assume %flat_count0 [range(%flat_count0, 1, 2048), le(%flat_count0, %row_capacity)] : index + %input_token_count0 = index.mul %bounded_tokens, %bounded_input_rows : index + %input_token_count = index.assume %input_token_count0 [range(%input_token_count0, 1, 131072)] : index + %lane0 = kernel.workitem.id : index + %lane = index.assume %lane0 [range(%lane0, 0, 63)] : index + %pair = kernel.workgroup.id : index + %slot_id = kernel.workgroup.id : index + %token_id = kernel.workgroup.id : index + %flat_base = index.mul %token_id, %bounded_slots : index + %flat0 = index.add %flat_base, %slot_id : index + %row00 = index.mul %pair, %c2 : index + %valid_flat = index.cmp ult, %flat0, %flat_count : index + %valid_row = index.cmp ult, %row00, %bounded_output_size : index + %publish_output = scalar.andi %valid_flat, %valid_row : i1 + %safe_flat = scf.select %valid_flat, %flat0, %c0 : index + %safe_row = scf.select %valid_row, %row00, %c0 : index + %flat, %launch_flat_count = index.assume %safe_flat, %flat_count [lt(%safe_flat, %flat_count)] : index, index + %row0, %launch_output_size = index.assume %safe_row, %bounded_output_size [lt(%safe_row, %bounded_output_size)] : index, index + %row1 = index.add %row0, %c1 : index + %row1_valid = index.cmp ult, %row1, %launch_output_size : index + %token = scf.select %valid_flat, %token_id, %c0 : index + %slot = scf.select %valid_flat, %slot_id, %c0 : index + %shared_input = index.cmp eq, %bounded_input_rows, %c1 : index + %input_slot = scf.select %shared_input, %c0, %slot : index + %input_row_base = index.mul %token, %bounded_input_rows : index + %input_row0 = index.add %input_row_base, %input_slot : index + %input_row = index.assume %input_row0 [lt(%input_row0, %input_token_count)] : index + %padded_input_size = index.add %bounded_input_size, %c255 : index + %block_count = index.div %padded_input_size, %c256 : index + %cohort = index.div %lane, %c16 : index + %ids_noalias = buffer.assume.noalias %ids : buffer + %bounded_route_stride = index.assume %route_stride [range(%route_stride, 1, 4096)] : index + %ids_view = buffer.view %ids_noalias[%c0_offset] : buffer -> view<[%bounded_tokens]x[%bounded_route_stride]xi32> + %ids_token, %ids_token_count = index.assume %token, %bounded_tokens [lt(%token, %bounded_tokens)] : index, index + %ids_slot, %ids_stride = index.assume %slot, %bounded_route_stride [lt(%slot, %bounded_route_stride)] : index, index + %expert_i32 = view.load %ids_view[%ids_token, %ids_slot] : view<[%bounded_tokens]x[%bounded_route_stride]xi32> -> i32 + %expert_raw = index.cast %expert_i32 : i32 to index + %bounded_experts = index.assume %expert_count [range(%expert_count, 1, 4096)] : index + // A bad router id (negative or past the expert count) reads expert 0 instead of past the weights. + %expert_in_range = index.cmp ult, %expert_raw, %bounded_experts : index + %expert0 = scf.select %expert_in_range, %expert_raw, %c0 : index + %expert = index.assume %expert0 [lt(%expert0, %bounded_experts)] : index + %expert_row_base = index.mul %expert, %bounded_output_size : index + %wrow0 = index.add %expert_row_base, %row0 : index + %wrow1 = index.add %wrow0, %c1 : index + scf.if %publish_output { + %input_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input, %weight, %output : buffer, buffer, buffer + %iq4nl_table = func.call @ggml_iq4nl_table_i8() : () -> (vector<16xi8>) + %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15 = func.call @ggml_iq3s_grid_f32_halves() : () -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + %acc0, %acc1 = scf.for %block_base = [%c0 to %block_count step %c4](%row_acc0 = %c0_f32 : f32, %row_acc1 = %c0_f32 : f32) -> (f32, f32) { + %block = index.add %block_base, %cohort : index + %input0, %input1, %input2, %input3 = func.call @ggml_mul_mat_decode_load_f32_block(%input_token_count, %bounded_input_size, %input_row, %block, %lane, %input_noalias) : (index, index, index, index, index, buffer) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) + %contribution0 = func.call @ggml_mul_mat_decode_f32_block_row(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight_format, %bounded_input_size, %wrow0, %block, %lane, %weight_noalias, %input0, %input1, %input2, %input3) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, index, index, index, index, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) -> (f32) + %contribution1 = scf.if %row1_valid -> (f32) { + %row1_contribution = func.call @ggml_mul_mat_decode_f32_block_row(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %weight_format, %bounded_input_size, %wrow1, %block, %lane, %weight_noalias, %input0, %input1, %input2, %input3) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, index, index, index, index, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) -> (f32) + scf.yield %row1_contribution : f32 + } else { + scf.yield %c0_f32 : f32 + } + %next0 = scalar.addf %row_acc0, %contribution0 : f32 + %next1 = scalar.addf %row_acc1, %contribution1 : f32 + scf.yield %next0, %next1 : f32, f32 + } + %sum0 = kernel.workgroup.reduce %acc0 : f32 + %sum1 = kernel.workgroup.reduce %acc1 : f32 + %is_lane_zero = index.cmp eq, %lane, %c0 : index + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_flat_count]x[%launch_output_size]xf32> + scf.if %is_lane_zero { + view.store %sum0, %output_view[%flat, %row0] : f32, view<[%launch_flat_count]x[%launch_output_size]xf32> + scf.if %row1_valid { + view.store %sum1, %output_view[%flat, %row1] : f32, view<[%launch_flat_count]x[%launch_output_size]xf32> + } + } + } + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_f16_f16_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_f16_f16_wmma.loom new file mode 100644 index 000000000000..f300aba28975 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_f16_f16_wmma.loom @@ -0,0 +1,71 @@ +// Standalone F16 MUL_MAT_ID fallback for routed expert projections. + +template.decl @ggml.mul_mat_id_f16_f16_wmma.core(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: index, %arg7: index, %arg8: index, %arg9: buffer, %arg10: buffer, %arg11: buffer, %arg12: buffer) + +template.decl @ggml.mul_mat_id_f16_f16_wmma.finish_tile(%token_count: index, %route_count: index, %output_size: index, %expert_count: index, %channel_tile: index, %expert: index, %route_tile_base: index, %expert_route_count: index, %expert_table: buffer, %output: buffer) + +template.decl @ggml.mul_mat_id_f16_f16_wmma.publish_vector4(%publish_word: i1, %assignment: index, %channel: index, %token_count: index, %route_count: index, %output_size: index, %values: vector<4xf16>, %output: buffer) + +amdgpu.target @ggml_mul_mat_id_f16_f16_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.mul_mat_id_f16_f16.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +config.decl @ggml.mul_mat_id_f16_f16.route_count : %value: index where [range(%value, 1, 8)] + +config.decl @ggml.mul_mat_id_f16_f16.expert_count : %value: index where [range(%value, 1, 512)] + +config.decl @ggml.mul_mat_id_f16_f16.output_size : %value: index where [range(%value, 1, 4096)] + +config.decl @ggml.mul_mat_id_f16_f16.weight_format : %value: index where [range(%value, 4, 6)] + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +template.def<@ggml.mul_mat_id_f16_f16_wmma.publish_vector4> device @ggml_mul_mat_id_f16_f16_wmma_publish_vector4(%publish_word: i1, %assignment: index, %channel: index, %token_count0: index, %route_count0: index, %output_size0: index, %values: vector<4xf16>, %output: buffer) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %route_count = index.assume %route_count0 [range(%route_count0, 1, 8)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 1, 4096)] : index + %c0_offset = index.constant 0 : offset + %c1 = index.constant 1 : index + %assignment_count = index.mul %token_count, %route_count : index + %output_noalias = buffer.assume.noalias %output : buffer + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%assignment_count]x[%output_size]xf16> + scf.if %publish_word { + %bounded_assignment, %bounded_assignment_count = index.assume %assignment, %assignment_count [lt(%assignment, %assignment_count)] : index, index + %mask = vector.mask.range [%channel to %output_size step %c1] : index -> vector<4xi1> + vector.store.mask %values, %output_view[%bounded_assignment, %channel], %mask : vector<4xf16>, view<[%assignment_count]x[%output_size]xf16>, vector<4xi1> + } + template.return +} + +template.def<@ggml.mul_mat_id_f16_f16_wmma.finish_tile> device @ggml_mul_mat_id_f16_f16_wmma_finish_tile(%token_count: index, %route_count: index, %output_size: index, %expert_count: index, %channel_tile: index, %expert: index, %route_tile_base: index, %expert_route_count: index, %expert_table: buffer, %output: buffer) { + template.return +} + +// Q4_K and Q6_K share one entry point. The model dispatcher supplies a +// compile-time weight format so specialization erases the inactive decoder. +kernel.def target(@ggml_mul_mat_id_f16_f16_gfx11_wave64) @ggml_mul_mat_id_f16_f16_wmma(%token_count: index) { + %expert_count = config.get @ggml.mul_mat_id_f16_f16.expert_count : index + %output_size = config.get @ggml.mul_mat_id_f16_f16.output_size : index + %c1 = index.constant 1 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %padded_token_count = index.add %token_count, %c63 : index + %route_tiles = index.div %padded_token_count, %c64 : index + kernel.launch.config workgroups(%output_tiles, %route_tiles, %expert_count) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %expert_table: buffer, %weight: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %input_size = config.get @ggml.mul_mat_id_f16_f16.input_size : index + %route_count = config.get @ggml.mul_mat_id_f16_f16.route_count : index + %expert_count = config.get @ggml.mul_mat_id_f16_f16.expert_count : index + %output_size = config.get @ggml.mul_mat_id_f16_f16.output_size : index + %weight_format = config.get @ggml.mul_mat_id_f16_f16.weight_format : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %route_tile = kernel.workgroup.id : index + %expert = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_id_f16_f16_wmma.core>(%weight_format, %bounded_token_count, %input_size, %route_count, %expert_count, %output_size, %channel_tile, %route_tile, %expert, %input, %expert_table, %weight, %output) : (index, index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_f32_f32_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_f32_f32_wmma.loom new file mode 100644 index 000000000000..e1f86b47fbaa --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_f32_f32_wmma.loom @@ -0,0 +1,81 @@ +// Standalone GGML MUL_MAT_ID fallback for routed expert projections. +// +// The kernel consumes a precomputed expert table. Each table row stores compact +// assignment indices encoded as token * route_count + route. Input can either +// be broadcast across all routes or contain one input row per route. +template.decl @ggml.mul_mat_id_f32_f32_wmma.core(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: index, %arg7: f32, %arg8: index, %arg9: index, %arg10: buffer, %arg11: buffer, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer, %arg16: buffer, %arg17: buffer, %arg18: buffer, %arg19: buffer, %arg20: buffer) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.finish_tile(%token_count: index, %route_count: index, %output_size: index, %expert_count: index, %output_tile_count: index, %maximum_partition_count: index, %epsilon: f32, %channel_tile: index, %descriptor_ordinal: index, %expert: index, %route_tile_base: index, %partition_row_count: index, %expert_table: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.publish_vector4(%publish_word: i1, %assignment: index, %token: index, %route: index, %channel: index, %token_count0: index, %route_count0: index, %output_size0: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +amdgpu.target @ggml_mul_mat_id_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.mul_mat_id.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 32)] + +config.decl @ggml.mul_mat_id.output_size : %value: index where [range(%value, 1, 262144)] + +config.decl @ggml.mul_mat_id.expert_count : %value: index where [range(%value, 1, 512)] + +config.decl @ggml.mul_mat_id.route_count : %value: index where [range(%value, 1, 32)] + +config.decl @ggml.mul_mat_id.input_route_count : %value: index where [range(%value, 1, 32)] + +config.decl @ggml.mul_mat_id.weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +template.def<@ggml.mul_mat_id_f32_f32_wmma.publish_vector4> device @ggml_mul_mat_id_f32_f32_wmma_publish_vector4(%publish_word: i1, %assignment: index, %token: index, %route: index, %channel: index, %token_count0: index, %route_count0: index, %output_size0: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %route_count = index.assume %route_count0 [range(%route_count0, 1, 32)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 1, 262144)] : index + %c0_offset = index.constant 0 : offset + %c1 = index.constant 1 : index + %output_row_count = index.mul %token_count, %route_count : index + %output_noalias = buffer.assume.noalias %output : buffer + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%output_row_count]x[%output_size]xf32> + scf.if %publish_word { + %bounded_assignment, %assignment_count = index.assume %assignment, %output_row_count [lt(%assignment, %output_row_count)] : index, index + %mask = vector.mask.range [%channel to %output_size step %c1] : index -> vector<4xi1> + vector.store.mask %values, %output_view[%bounded_assignment, %channel], %mask : vector<4xf32>, view<[%output_row_count]x[%output_size]xf32>, vector<4xi1> + } + template.return +} + +template.def<@ggml.mul_mat_id_f32_f32_wmma.finish_tile> device @ggml_mul_mat_id_f32_f32_wmma_finish_tile(%token_count: index, %route_count: index, %output_size: index, %expert_count: index, %output_tile_count: index, %maximum_partition_count: index, %epsilon: f32, %channel_tile: index, %descriptor_ordinal: index, %expert: index, %route_tile_base: index, %partition_row_count: index, %expert_table: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + template.return +} + +kernel.def target(@ggml_mul_mat_id_gfx11_wave64) @ggml_mul_mat_id_f32_f32_wmma(%token_count: index) { + %expert_count = config.get @ggml.mul_mat_id.expert_count : index + %route_count = config.get @ggml.mul_mat_id.route_count : index + %output_size = config.get @ggml.mul_mat_id.output_size : index + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %assignment_count = index.mul %token_count, %route_count : index + %rounded_assignment_count = index.add %assignment_count, %c31 : index + %assignment_partition_count = index.div %rounded_assignment_count, %c32 : index + %has_more_assignment_partitions = index.cmp ugt, %assignment_partition_count, %expert_count : index + %launch_partition_count = scf.select %has_more_assignment_partitions, %assignment_partition_count, %expert_count : index + kernel.launch.config workgroups(%output_tiles, %launch_partition_count, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %expert_table: buffer, %partition_table: buffer, %weight: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %input_size = config.get @ggml.mul_mat_id.input_size : index + %route_count = config.get @ggml.mul_mat_id.route_count : index + %input_route_count = config.get @ggml.mul_mat_id.input_route_count : index + %expert_count = config.get @ggml.mul_mat_id.expert_count : index + %output_size = config.get @ggml.mul_mat_id.output_size : index + %weight_format = config.get @ggml.mul_mat_id.weight_format : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %c0_f32 = scalar.constant 0.0 : f32 + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %partition_ordinal = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_id_f32_f32_wmma.core>(%weight_format, %bounded_token_count, %input_size, %output_size, %route_count, %input_route_count, %expert_count, %c0_f32, %channel_tile, %partition_ordinal, %input, %expert_table, %partition_table, %weight, %output, %output, %output, %output, %output, %output, %output) : (index, index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_postops_f32_f32_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_postops_f32_f32_wmma.loom new file mode 100644 index 000000000000..194e0d9d297f --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_postops_f32_f32_wmma.loom @@ -0,0 +1,93 @@ +template.decl @ggml.mul_mat_id_f32_f32_wmma.add_bias_vector4(%arg0: index, %arg1: index, %arg2: vector<4xf32>, %arg3: buffer) -> (vector<4xf32>) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.add_residual_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: vector<4xf32>, %arg6: buffer) -> (vector<4xf32>) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.core(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: index, %arg7: f32, %arg8: index, %arg9: index, %arg10: buffer, %arg11: buffer, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer, %arg16: buffer, %arg17: buffer, %arg18: buffer, %arg19: buffer, %arg20: buffer) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.finish_tile(%token_count: index, %route_count: index, %output_size: index, %expert_count: index, %output_tile_count: index, %maximum_partition_count: index, %epsilon: f32, %channel_tile: index, %descriptor_ordinal: index, %expert: index, %route_tile_base: index, %partition_row_count: index, %expert_table: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.publish_vector4(%publish_word: i1, %assignment: index, %token: index, %route: index, %channel: index, %token_count0: index, %route_count0: index, %output_size0: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.store_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: vector<4xf32>, %arg6: buffer) + +amdgpu.target @ggml_mul_mat_id_postops_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.mul_mat_id_postops.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 32)] + +config.decl @ggml.mul_mat_id_postops.output_size : %value: index where [range(%value, 1, 262144)] + +config.decl @ggml.mul_mat_id_postops.expert_count : %value: index where [range(%value, 1, 512)] + +config.decl @ggml.mul_mat_id_postops.route_count : %value: index where [range(%value, 1, 32)] + +config.decl @ggml.mul_mat_id_postops.input_route_count : %value: index where [range(%value, 1, 32)] + +config.decl @ggml.mul_mat_id_postops.weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.mul_mat_id_postops.has_bias : %value: index where [range(%value, 0, 1)] + +config.decl @ggml.mul_mat_id_postops.has_residual : %value: index where [range(%value, 0, 1)] + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +template.def<@ggml.mul_mat_id_f32_f32_wmma.publish_vector4> device @ggml_mul_mat_id_postops_publish_vector4(%publish_word: i1, %assignment: index, %token: index, %route: index, %channel: index, %token_count0: index, %route_count0: index, %output_size0: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %has_bias_value = config.get @ggml.mul_mat_id_postops.has_bias : index + %has_residual_value = config.get @ggml.mul_mat_id_postops.has_residual : index + %c0 = index.constant 0 : index + %has_bias = index.cmp ne, %has_bias_value, %c0 : index + %has_residual = index.cmp ne, %has_residual_value, %c0 : index + scf.if %publish_word { + %with_bias = scf.if %has_bias -> (vector<4xf32>) { + %biased = template.apply<@ggml.mul_mat_id_f32_f32_wmma.add_bias_vector4>(%output_size0, %channel, %values, %bias) : (index, index, vector<4xf32>, buffer) -> (vector<4xf32>) + scf.yield %biased : vector<4xf32> + } else { + scf.yield %values : vector<4xf32> + } + %published = scf.if %has_residual -> (vector<4xf32>) { + %added = template.apply<@ggml.mul_mat_id_f32_f32_wmma.add_residual_vector4>(%token_count0, %route_count0, %output_size0, %assignment, %channel, %with_bias, %residual_input) : (index, index, index, index, index, vector<4xf32>, buffer) -> (vector<4xf32>) + scf.yield %added : vector<4xf32> + } else { + scf.yield %with_bias : vector<4xf32> + } + template.apply<@ggml.mul_mat_id_f32_f32_wmma.store_vector4>(%token_count0, %route_count0, %output_size0, %assignment, %channel, %published, %output) : (index, index, index, index, index, vector<4xf32>, buffer) + } + template.return +} + +template.def<@ggml.mul_mat_id_f32_f32_wmma.finish_tile> device @ggml_mul_mat_id_postops_finish_tile(%token_count: index, %route_count: index, %output_size: index, %expert_count: index, %output_tile_count: index, %maximum_partition_count: index, %epsilon: f32, %channel_tile: index, %descriptor_ordinal: index, %expert: index, %route_tile_base: index, %partition_row_count: index, %expert_table: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + template.return +} + +kernel.def target(@ggml_mul_mat_id_postops_gfx11_wave64) @ggml_mul_mat_id_postops_f32_f32_wmma(%token_count: index) { + %expert_count = config.get @ggml.mul_mat_id_postops.expert_count : index + %route_count = config.get @ggml.mul_mat_id_postops.route_count : index + %output_size = config.get @ggml.mul_mat_id_postops.output_size : index + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %assignment_count = index.mul %token_count, %route_count : index + %rounded_assignment_count = index.add %assignment_count, %c31 : index + %assignment_partition_count = index.div %rounded_assignment_count, %c32 : index + %has_more_assignment_partitions = index.cmp ugt, %assignment_partition_count, %expert_count : index + %launch_partition_count = scf.select %has_more_assignment_partitions, %assignment_partition_count, %expert_count : index + kernel.launch.config workgroups(%output_tiles, %launch_partition_count, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %expert_table: buffer, %partition_table: buffer, %weight: buffer, %bias: buffer, %residual_input: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %input_size = config.get @ggml.mul_mat_id_postops.input_size : index + %route_count = config.get @ggml.mul_mat_id_postops.route_count : index + %input_route_count = config.get @ggml.mul_mat_id_postops.input_route_count : index + %expert_count = config.get @ggml.mul_mat_id_postops.expert_count : index + %output_size = config.get @ggml.mul_mat_id_postops.output_size : index + %weight_format = config.get @ggml.mul_mat_id_postops.weight_format : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %c0_f32 = scalar.constant 0.0 : f32 + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %partition_ordinal = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_id_f32_f32_wmma.core>(%weight_format, %bounded_token_count, %input_size, %output_size, %route_count, %input_route_count, %expert_count, %c0_f32, %channel_tile, %partition_ordinal, %input, %expert_table, %partition_table, %weight, %bias, %output, %residual_input, %output, %output, %output, %output) : (index, index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_postops_next_rmsnorm_f32_f32_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_postops_next_rmsnorm_f32_f32_wmma.loom new file mode 100644 index 000000000000..123fea1658eb --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_postops_next_rmsnorm_f32_f32_wmma.loom @@ -0,0 +1,98 @@ +template.decl @ggml.mul_mat_id_f32_f32_wmma.add_bias_vector4(%arg0: index, %arg1: index, %arg2: vector<4xf32>, %arg3: buffer) -> (vector<4xf32>) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.add_residual_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: vector<4xf32>, %arg6: buffer) -> (vector<4xf32>) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.core(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: index, %arg7: f32, %arg8: index, %arg9: index, %arg10: buffer, %arg11: buffer, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer, %arg16: buffer, %arg17: buffer, %arg18: buffer, %arg19: buffer, %arg20: buffer) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.finish_routed_rmsnorm_weight(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: f32, %arg7: index, %arg8: index, %arg9: index, %arg10: index, %arg11: buffer, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.finish_tile(%token_count: index, %route_count: index, %output_size: index, %expert_count: index, %output_tile_count: index, %maximum_partition_count: index, %epsilon: f32, %channel_tile: index, %descriptor_ordinal: index, %expert: index, %route_tile_base: index, %partition_row_count: index, %expert_table: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.publish_vector4(%publish_word: i1, %assignment: index, %token: index, %route: index, %channel: index, %token_count0: index, %route_count0: index, %output_size0: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) + +template.decl @ggml.mul_mat_id_f32_f32_wmma.store_vector4(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: vector<4xf32>, %arg6: buffer) + +amdgpu.target @ggml_mul_mat_id_postops_next_rmsnorm_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.mul_mat_id_postops.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 32)] + +config.decl @ggml.mul_mat_id_postops.output_size : %value: index where [range(%value, 128, 32768), mul(%value, 128)] + +config.decl @ggml.mul_mat_id_postops.expert_count : %value: index where [range(%value, 1, 512)] + +config.decl @ggml.mul_mat_id_postops.route_count : %value: index where [range(%value, 1, 32)] + +config.decl @ggml.mul_mat_id_postops.input_route_count : %value: index where [range(%value, 1, 32)] + +config.decl @ggml.mul_mat_id_postops.weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.mul_mat_id_postops.has_bias : %value: index where [range(%value, 0, 1)] + +config.decl @ggml.mul_mat_id_postops.has_residual : %value: index where [range(%value, 0, 1)] + +config.decl @ggml.mul_mat_id_postops.rms_epsilon : f32 + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +template.def<@ggml.mul_mat_id_f32_f32_wmma.publish_vector4> device @ggml_mul_mat_id_postops_next_rmsnorm_publish_vector4(%publish_word: i1, %assignment: index, %token: index, %route: index, %channel: index, %token_count0: index, %route_count0: index, %output_size0: index, %values: vector<4xf32>, %bias: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + %has_bias_value = config.get @ggml.mul_mat_id_postops.has_bias : index + %has_residual_value = config.get @ggml.mul_mat_id_postops.has_residual : index + %c0 = index.constant 0 : index + %has_bias = index.cmp ne, %has_bias_value, %c0 : index + %has_residual = index.cmp ne, %has_residual_value, %c0 : index + scf.if %publish_word { + %with_bias = scf.if %has_bias -> (vector<4xf32>) { + %biased = template.apply<@ggml.mul_mat_id_f32_f32_wmma.add_bias_vector4>(%output_size0, %channel, %values, %bias) : (index, index, vector<4xf32>, buffer) -> (vector<4xf32>) + scf.yield %biased : vector<4xf32> + } else { + scf.yield %values : vector<4xf32> + } + %published = scf.if %has_residual -> (vector<4xf32>) { + %added = template.apply<@ggml.mul_mat_id_f32_f32_wmma.add_residual_vector4>(%token_count0, %route_count0, %output_size0, %assignment, %channel, %with_bias, %residual_input) : (index, index, index, index, index, vector<4xf32>, buffer) -> (vector<4xf32>) + scf.yield %added : vector<4xf32> + } else { + scf.yield %with_bias : vector<4xf32> + } + template.apply<@ggml.mul_mat_id_f32_f32_wmma.store_vector4>(%token_count0, %route_count0, %output_size0, %assignment, %channel, %published, %residual_output) : (index, index, index, index, index, vector<4xf32>, buffer) + } + template.return +} + +template.def<@ggml.mul_mat_id_f32_f32_wmma.finish_tile> device @ggml_mul_mat_id_postops_next_rmsnorm_finish_tile(%token_count: index, %route_count: index, %output_size: index, %expert_count: index, %output_tile_count: index, %maximum_partition_count: index, %epsilon: f32, %channel_tile: index, %descriptor_ordinal: index, %expert: index, %route_tile_base: index, %partition_row_count: index, %expert_table: buffer, %output: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) { + template.apply<@ggml.mul_mat_id_f32_f32_wmma.finish_routed_rmsnorm_weight>(%token_count, %route_count, %output_size, %expert_count, %output_tile_count, %maximum_partition_count, %epsilon, %descriptor_ordinal, %expert, %route_tile_base, %partition_row_count, %expert_table, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (index, index, index, index, index, index, f32, index, index, index, index, buffer, buffer, buffer, buffer, buffer) + template.return +} + +kernel.def target(@ggml_mul_mat_id_postops_next_rmsnorm_gfx11_wave64) @ggml_mul_mat_id_postops_next_rmsnorm_f32_f32_wmma(%token_count: index) { + %expert_count = config.get @ggml.mul_mat_id_postops.expert_count : index + %route_count = config.get @ggml.mul_mat_id_postops.route_count : index + %output_size = config.get @ggml.mul_mat_id_postops.output_size : index + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %assignment_count = index.mul %token_count, %route_count : index + %rounded_assignment_count = index.add %assignment_count, %c31 : index + %assignment_partition_count = index.div %rounded_assignment_count, %c32 : index + %has_more_assignment_partitions = index.cmp ugt, %assignment_partition_count, %expert_count : index + %launch_partition_count = scf.select %has_more_assignment_partitions, %assignment_partition_count, %expert_count : index + kernel.launch.config workgroups(%output_tiles, %launch_partition_count, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %expert_table: buffer, %partition_table: buffer, %weight: buffer, %bias: buffer, %residual_input: buffer, %residual_output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counters: buffer) where [range(%token_count, 1, 2048)] { + %input_size = config.get @ggml.mul_mat_id_postops.input_size : index + %route_count = config.get @ggml.mul_mat_id_postops.route_count : index + %input_route_count = config.get @ggml.mul_mat_id_postops.input_route_count : index + %expert_count = config.get @ggml.mul_mat_id_postops.expert_count : index + %output_size = config.get @ggml.mul_mat_id_postops.output_size : index + %weight_format = config.get @ggml.mul_mat_id_postops.weight_format : index + %epsilon = config.get @ggml.mul_mat_id_postops.rms_epsilon : f32 + %token_capacity = config.get @ggml.workload.token_capacity : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %partition_ordinal = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_id_f32_f32_wmma.core>(%weight_format, %bounded_token_count, %input_size, %output_size, %route_count, %input_route_count, %expert_count, %epsilon, %channel_tile, %partition_ordinal, %input, %expert_table, %partition_table, %weight, %bias, %residual_output, %residual_input, %residual_output, %norm_weight, %normalized_output, %completion_counters) : (index, index, index, index, index, index, index, f32, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_swiglu_f16_f16_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_swiglu_f16_f16_wmma.loom new file mode 100644 index 000000000000..c352f7169f92 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_swiglu_f16_f16_wmma.loom @@ -0,0 +1,57 @@ +// Routed gate/up SwiGLU f16-output path with compile-time weight-format selection. + +template.decl @ggml.mul_mat_id_swiglu_f16_f16.body(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: index, %arg7: i32, %arg8: i32, %arg9: i32, %arg10: index, %arg11: index, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer, %arg16: buffer, %arg17: buffer) + +amdgpu.target @ggml_mul_mat_id_swiglu_f16_f16_gfx11_wave32 {subgroup_size = 32} + +config.decl @ggml.mul_mat_id_swiglu_f16_f16.input_size : %value: index where [range(%value, 512, 32768), mul(%value, 512)] + +config.decl @ggml.mul_mat_id_swiglu_f16_f16.route_count : %value: index where [range(%value, 1, 8)] + +config.decl @ggml.mul_mat_id_swiglu_f16_f16.expert_count : %value: index where [range(%value, 1, 512)] + +config.decl @ggml.mul_mat_id_swiglu_f16_f16.output_size : %value: index where [range(%value, 1, 4096)] + +config.decl @ggml.mul_mat_id_swiglu_f16_f16.gate_weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.mul_mat_id_swiglu_f16_f16.up_weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.mul_mat_id_swiglu_f16_f16.descriptor_expert_mask : %value: i32 + +config.decl @ggml.mul_mat_id_swiglu_f16_f16.descriptor_partition_shift : %value: i32 + +config.decl @ggml.mul_mat_id_swiglu_f16_f16.descriptor_row_count_shift : %value: i32 + +kernel.def target(@ggml_mul_mat_id_swiglu_f16_f16_gfx11_wave32) @ggml_mul_mat_id_swiglu_f16_f16_wmma(%token_count: index) { + %route_count = config.get @ggml.mul_mat_id_swiglu_f16_f16.route_count : index + %expert_count = config.get @ggml.mul_mat_id_swiglu_f16_f16.expert_count : index + %output_size = config.get @ggml.mul_mat_id_swiglu_f16_f16.output_size : index + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %assignment_count = index.mul %token_count, %route_count : index + %rounded_assignment_count = index.add %assignment_count, %c31 : index + %assignment_partition_count = index.div %rounded_assignment_count, %c32 : index + %has_more_assignment_partitions = index.cmp ugt, %assignment_partition_count, %expert_count : index + %launch_partition_count = scf.select %has_more_assignment_partitions, %assignment_partition_count, %expert_count : index + kernel.launch.config workgroups(%output_tiles, %launch_partition_count, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %expert_table: buffer, %partition_table: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %input_size = config.get @ggml.mul_mat_id_swiglu_f16_f16.input_size : index + %route_count = config.get @ggml.mul_mat_id_swiglu_f16_f16.route_count : index + %expert_count = config.get @ggml.mul_mat_id_swiglu_f16_f16.expert_count : index + %output_size = config.get @ggml.mul_mat_id_swiglu_f16_f16.output_size : index + %gate_weight_format = config.get @ggml.mul_mat_id_swiglu_f16_f16.gate_weight_format : index + %up_weight_format = config.get @ggml.mul_mat_id_swiglu_f16_f16.up_weight_format : index + %descriptor_expert_mask = config.get @ggml.mul_mat_id_swiglu_f16_f16.descriptor_expert_mask : i32 + %descriptor_partition_shift = config.get @ggml.mul_mat_id_swiglu_f16_f16.descriptor_partition_shift : i32 + %descriptor_row_count_shift = config.get @ggml.mul_mat_id_swiglu_f16_f16.descriptor_row_count_shift : i32 + %channel_tile = kernel.workgroup.id : index + %partition_ordinal = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_id_swiglu_f16_f16.body>(%gate_weight_format, %up_weight_format, %token_count, %input_size, %route_count, %expert_count, %output_size, %descriptor_expert_mask, %descriptor_partition_shift, %descriptor_row_count_shift, %channel_tile, %partition_ordinal, %input, %expert_table, %partition_table, %gate_weight, %up_weight, %output) : (index, index, index, index, index, index, index, i32, i32, i32, index, index, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_swiglu_f32_f32_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_swiglu_f32_f32_wmma.loom new file mode 100644 index 000000000000..79b16cc1f01d --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_id_swiglu_f32_f32_wmma.loom @@ -0,0 +1,397 @@ +// Routed GGML gate/up MUL_MAT_ID fused with SwiGLU for compact expert tables. +template.decl @ggml.mul_mat_id_swiglu.apply_vector4(%gate: vector<4xf32>, %up: vector<4xf32>) -> (vector<4xf32>) + +template.decl @ggml.mul_mat_id_swiglu.body(%gate_weight_format: index, %up_weight_format: index, %token_count: index, %input_size: index, %output_size: index, %route_count: index, %input_route_count: index, %expert_count: index, %channel_tile: index, %partition_ordinal: index, %input: buffer, %expert_table: buffer, %partition_table: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) + +amdgpu.target @ggml_mul_mat_id_swiglu_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.mul_mat_id_swiglu.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 32)] + +config.decl @ggml.mul_mat_id_swiglu.output_size : %value: index where [range(%value, 1, 262144)] + +config.decl @ggml.mul_mat_id_swiglu.expert_count : %value: index where [range(%value, 1, 512)] + +config.decl @ggml.mul_mat_id_swiglu.route_count : %value: index where [range(%value, 1, 32)] + +config.decl @ggml.mul_mat_id_swiglu.input_route_count : %value: index where [range(%value, 1, 32)] + +config.decl @ggml.mul_mat_id_swiglu.gate_weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.mul_mat_id_swiglu.up_weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +func.decl @ggml_dequant_weight_row_bytes(%weight_format: index, %hidden_size: index) -> (offset) + +func.decl @ggml_iq4nl_table_i8() -> (vector<16xi8>) + +func.decl @ggml_iq3s_grid_f32_halves() -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + +func.decl @ggml_dequant_f16_vector4(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight_format: index, %weight: buffer, %row_byte_base: offset, %input_size: index, %quant_block: index, %quant_group: index, %packet: index, %k: index) -> (vector<4xf16>) + +func.def inline @ggml_moe_unpack_expert_partition_descriptor(%descriptor: i32) -> (index, index, index) { + %c1_i32 = scalar.constant 1 : i32 + %c5_i32 = scalar.constant 5 : i32 + %c9_i32 = scalar.constant 9 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c31_i32 = scalar.constant 31 : i32 + %c63_i32 = scalar.constant 63 : i32 + %c511_i32 = scalar.constant 511 : i32 + %expert_i32 = scalar.andi %descriptor, %c511_i32 : i32 + %partition_shifted_i32 = scalar.shrui %descriptor, %c9_i32 : i32 + %partition_i32 = scalar.andi %partition_shifted_i32, %c63_i32 : i32 + %route_tile_base_i32 = scalar.shli %partition_i32, %c5_i32 : i32 + %row_count_shifted_i32 = scalar.shrui %descriptor, %c15_i32 : i32 + %row_count_minus_one_i32 = scalar.andi %row_count_shifted_i32, %c31_i32 : i32 + %partition_row_count_i32 = scalar.addi %row_count_minus_one_i32, %c1_i32 : i32 + %expert0 = index.cast %expert_i32 : i32 to index + %expert = index.assume %expert0 [range(%expert0, 0, 511)] : index + %route_tile_base0 = index.cast %route_tile_base_i32 : i32 to index + %route_tile_base = index.assume %route_tile_base0 [range(%route_tile_base0, 0, 2016)] : index + %partition_row_count0 = index.cast %partition_row_count_i32 : i32 to index + %partition_row_count = index.assume %partition_row_count0 [range(%partition_row_count0, 1, 32)] : index + func.return %expert, %route_tile_base, %partition_row_count : index, index, index +} + +template.def<@ggml.mul_mat_id_swiglu.apply_vector4> device @ggml_mul_mat_id_swiglu_apply_vector4(%gate: vector<4xf32>, %up: vector<4xf32>) -> (vector<4xf32>) { + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %c1_f32x4 = vector.constant 1.0 : vector<4xf32> + %negative_gate = vector.subf %c0_f32x4, %gate : vector<4xf32> + %gate_exp = vector.expf %negative_gate : vector<4xf32> + %denominator = vector.addf %c1_f32x4, %gate_exp : vector<4xf32> + %sigmoid = vector.divf %c1_f32x4, %denominator : vector<4xf32> + %silu = vector.mulf %gate, %sigmoid : vector<4xf32> + %result = vector.mulf %silu, %up : vector<4xf32> + template.return %result : vector<4xf32> +} + +template.def<@ggml.mul_mat_id_swiglu.body> device @ggml_mul_mat_id_swiglu_f32_f32_wmma_body(%gate_weight_format: index, %up_weight_format: index, %token_count: index, %input_size: index, %output_size: index, %route_count: index, %input_route_count: index, %expert_count: index, %channel_tile: index, %partition_ordinal: index, %input: buffer, %expert_table: buffer, %partition_table: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 32)] : index + %bounded_route_count = index.assume %route_count [range(%route_count, 1, 32)] : index + %bounded_input_route_count = index.assume %input_route_count [range(%input_route_count, 1, 32), le(%input_route_count, %bounded_route_count)] : index + %bounded_expert_count = index.assume %expert_count [range(%expert_count, 1, 512)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 262144)] : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 1)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c40 = index.constant 40 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %c4_bytes = index.constant 4 : offset + %weight_stage_bytes = index.constant 5120 : offset + %activation_stage_bytes = index.constant 2560 : offset + %route_stage_bytes = index.constant 128 : offset + %wave_result_stage_bytes = index.constant 1024 : offset + %result_stage_bytes = index.constant 2048 : offset + %c0_i32 = scalar.constant 0 : i32 + %cn1_i32 = scalar.constant -1 : i32 + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %c0_f16x4 = vector.constant 0.0 : vector<4xf16> + %zero_accumulator = vector.constant 0.0 : vector<4xf32> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %assignment_count = index.mul %bounded_token_count, %bounded_route_count : index + %output_row_count = index.mul %bounded_token_count, %bounded_route_count : index + %padded_output_size = index.add %bounded_output_size, %c63 : index + %output_tile_count = index.div %padded_output_size, %c64 : index + %rounded_assignment_count = index.add %assignment_count, %c31 : index + %assignment_partition_count = index.div %rounded_assignment_count, %c32 : index + %maximum_partition_count = index.add %assignment_partition_count, %bounded_expert_count : index + %has_more_assignment_partitions = index.cmp ugt, %assignment_partition_count, %bounded_expert_count : index + %launch_partition_count = scf.select %has_more_assignment_partitions, %assignment_partition_count, %bounded_expert_count : index + %assignment_table_byte_base = index.scale %bounded_expert_count, %c4_bytes : index, offset -> offset + %padded_input_size = index.add %bounded_input_size, %c255 : index + %quant_block_count = index.div %padded_input_size, %c256 : index + %gate_weight_row_bytes = func.call @ggml_dequant_weight_row_bytes(%gate_weight_format, %bounded_input_size) : (index, index) -> (offset) + %up_weight_row_bytes = func.call @ggml_dequant_weight_row_bytes(%up_weight_format, %bounded_input_size) : (index, index) -> (offset) + %gate_weight_expert_bytes = index.scale %bounded_output_size, %gate_weight_row_bytes : index, offset -> offset + %up_weight_expert_bytes = index.scale %bounded_output_size, %up_weight_row_bytes : index, offset -> offset + %input_noalias, %expert_table_noalias, %partition_table_noalias, %gate_weight_noalias, %up_weight_noalias, %output_noalias = buffer.assume.noalias %input, %expert_table, %partition_table, %gate_weight, %up_weight, %output : buffer, buffer, buffer, buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_input_route_count]x[%bounded_input_size]xf32> + %assignment_view = buffer.view %expert_table_noalias[%assignment_table_byte_base] : buffer -> view<[%bounded_expert_count]x[%bounded_token_count]xi32> + %partition_count_view = buffer.view %partition_table_noalias[%c0_offset] : buffer -> view<1xi32> + %partition_descriptor_view = buffer.view %partition_table_noalias[%c4_bytes] : buffer -> view<[%maximum_partition_count]xi32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%output_row_count]x[%bounded_output_size]xf32> + %gate_weight_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %up_weight_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %activation_stage = buffer.alloca align(16) %activation_stage_bytes : buffer + %route_stage = buffer.alloca align(16) %route_stage_bytes : buffer + %gate_result_stage = buffer.alloca align(16) %result_stage_bytes : buffer + %up_result_stage = buffer.alloca align(16) %result_stage_bytes : buffer + %gate_weight_stage_view = buffer.view %gate_weight_stage[%c0_offset] : buffer -> view<64x40xf16> + %up_weight_stage_view = buffer.view %up_weight_stage[%c0_offset] : buffer -> view<64x40xf16> + %activation_stage_physical_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x40xf16> + %activation_fragment_layout = encoding.layout.strided [1, %c40] : encoding + %activation_fragment_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x32xf16, %activation_fragment_layout> + %route_stage_view = buffer.view %route_stage[%c0_offset] : buffer -> view<32xi32> + %wave_result_stage_offset = index.scale %subgroup, %wave_result_stage_bytes : index, offset -> offset + %result_fragment_layout = encoding.layout.strided [1, %c16] : encoding + %gate_result_fragment_view = buffer.view %gate_result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf32, %result_fragment_layout> + %up_result_fragment_view = buffer.view %up_result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf32, %result_fragment_layout> + %gate_result_physical_view = buffer.view %gate_result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf32> + %up_result_physical_view = buffer.view %up_result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf32> + %channel_tile_base = index.mul %channel_tile, %c64 : index + %is_workitem_zero = index.cmp eq, %workitem, %c0 : index + %lane_partition_count_i32 = scf.if %is_workitem_zero -> (i32) { + %loaded = view.load %partition_count_view[%c0] : view<1xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %c0_i32 : i32 + } + %partition_count_reduced = kernel.workgroup.reduce %lane_partition_count_i32 : i32 + %partition_count_i32 = kernel.subgroup.broadcast.first %partition_count_reduced : i32 + %partition_count0 = index.cast %partition_count_i32 : i32 to index + %partition_count, %partition_capacity = index.assume %partition_count0, %maximum_partition_count [le(%partition_count0, %maximum_partition_count)] : index, index + scf.for %active_partition = [%partition_ordinal to %partition_count step %launch_partition_count] { + %descriptor_ordinal, %descriptor_count = index.assume %active_partition, %partition_count [lt(%active_partition, %partition_count)] : index, index + %table_descriptor_ordinal, %table_descriptor_capacity = index.assume %descriptor_ordinal, %maximum_partition_count [lt(%descriptor_ordinal, %maximum_partition_count)] : index, index + %lane_descriptor_i32 = scf.if %is_workitem_zero -> (i32) { + %loaded = view.load %partition_descriptor_view[%table_descriptor_ordinal] : view<[%maximum_partition_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %c0_i32 : i32 + } + %descriptor_reduced_i32 = kernel.workgroup.reduce %lane_descriptor_i32 : i32 + %descriptor_i32 = kernel.subgroup.broadcast.first %descriptor_reduced_i32 : i32 + %expert, %route_tile_base, %partition_row_count = func.call @ggml_moe_unpack_expert_partition_descriptor(%descriptor_i32) : (i32) -> (index, index, index) + %bounded_expert, %table_expert_count = index.assume %expert, %bounded_expert_count [lt(%expert, %bounded_expert_count)] : index, index + %loads_route = index.cmp ult, %workitem, %c32 : index + scf.if %loads_route { + %local_route = index.assume %workitem [range(%workitem, 0, 31)] : index + %assignment_ordinal = index.add %route_tile_base, %local_route : index + %valid_row = index.cmp ult, %local_route, %partition_row_count : index + %assignment_i32 = scf.if %valid_row -> (i32) { + %bounded_assignment_ordinal, %table_token_count = index.assume %assignment_ordinal, %bounded_token_count [lt(%assignment_ordinal, %bounded_token_count)] : index, index + %loaded = view.load %assignment_view[%bounded_expert, %bounded_assignment_ordinal] : view<[%bounded_expert_count]x[%bounded_token_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %cn1_i32 : i32 + } + view.store %assignment_i32, %route_stage_view[%local_route] : i32, view<32xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %load_packet = index.rem %workitem, %c8 : index + %load_k = index.mul %load_packet, %c4 : index + %load_row0 = index.div %workitem, %c8 : index + %load_row = index.assume %load_row0 [range(%load_row0, 0, 15)] : index + %gate_expert_byte_base = index.scale %bounded_expert, %gate_weight_expert_bytes : index, offset -> offset + %up_expert_byte_base = index.scale %bounded_expert, %up_weight_expert_bytes : index, offset -> offset + %subgroup_channel_add = index.mul %subgroup, %c32 : index + %subgroup_channel1 = index.add %subgroup_channel_add, %c16 : index + %init00 = vector.fragment %zero_accumulator shape [%m, %n] : vector<4xf32> + %init01 = vector.fragment %zero_accumulator shape [%m, %n] : vector<4xf32> + %init10 = vector.fragment %zero_accumulator shape [%m, %n] : vector<4xf32> + %init11 = vector.fragment %zero_accumulator shape [%m, %n] : vector<4xf32> + %iq4nl_table = func.call @ggml_iq4nl_table_i8() : () -> (vector<16xi8>) + %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15 = func.call @ggml_iq3s_grid_f32_halves() : () -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + %gate_result00, %gate_result01, %gate_result10, %gate_result11, %up_result00, %up_result01, %up_result10, %up_result11 = scf.for %quant_block = [%c0 to %quant_block_count step %c1](%gate_block_acc00 = %init00 : vector<4xf32>, %gate_block_acc01 = %init01 : vector<4xf32>, %gate_block_acc10 = %init10 : vector<4xf32>, %gate_block_acc11 = %init11 : vector<4xf32>, %up_block_acc00 = %init00 : vector<4xf32>, %up_block_acc01 = %init01 : vector<4xf32>, %up_block_acc10 = %init10 : vector<4xf32>, %up_block_acc11 = %init11 : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + %gate_block_result00, %gate_block_result01, %gate_block_result10, %gate_block_result11, %up_block_result00, %up_block_result01, %up_block_result10, %up_block_result11 = scf.for %quant_group = [%c0 to %c8 step %c1](%gate_acc00 = %gate_block_acc00 : vector<4xf32>, %gate_acc01 = %gate_block_acc01 : vector<4xf32>, %gate_acc10 = %gate_block_acc10 : vector<4xf32>, %gate_acc11 = %gate_block_acc11 : vector<4xf32>, %up_acc00 = %up_block_acc00 : vector<4xf32>, %up_acc01 = %up_block_acc01 : vector<4xf32>, %up_acc10 = %up_block_acc10 : vector<4xf32>, %up_acc11 = %up_block_acc11 : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + %block_k_base = index.mul %quant_block, %c256 : index + %group_k_add = index.mul %quant_group, %c32 : index + %k_origin = index.add %block_k_base, %group_k_add : index + scf.for %row_offset = [%c0 to %c64 step %c16] unroll { + %local_row0 = index.add %load_row, %row_offset : index + %local_row = index.assume %local_row0 [range(%local_row0, 0, 63)] : index + %channel = index.add %channel_tile_base, %local_row : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %weight_k = index.add %k_origin, %load_k : index + %valid_k = index.cmp ult, %weight_k, %bounded_input_size : index + %valid_weight = scalar.andi %valid_channel, %valid_k : i1 + %gate_weight_values = scf.if %valid_weight -> (vector<4xf16>) { + %channel_byte_add = index.scale %channel, %gate_weight_row_bytes : index, offset -> offset + %row_byte_base = index.add %gate_expert_byte_base, %channel_byte_add : offset + %decoded = func.call @ggml_dequant_f16_vector4(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %gate_weight_format, %gate_weight_noalias, %row_byte_base, %bounded_input_size, %quant_block, %quant_group, %load_packet, %weight_k) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, buffer, offset, index, index, index, index, index) -> (vector<4xf16>) + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + %up_weight_values = scf.if %valid_weight -> (vector<4xf16>) { + %channel_byte_add = index.scale %channel, %up_weight_row_bytes : index, offset -> offset + %row_byte_base = index.add %up_expert_byte_base, %channel_byte_add : offset + %decoded = func.call @ggml_dequant_f16_vector4(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %up_weight_format, %up_weight_noalias, %row_byte_base, %bounded_input_size, %quant_block, %quant_group, %load_packet, %weight_k) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, buffer, offset, index, index, index, index, index) -> (vector<4xf16>) + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + %is_activation_row = index.cmp ult, %local_row, %c32 : index + vector.store %gate_weight_values, %gate_weight_stage_view[%local_row, %load_k] : vector<4xf16>, view<64x40xf16> + vector.store %up_weight_values, %up_weight_stage_view[%local_row, %load_k] : vector<4xf16>, view<64x40xf16> + scf.if %is_activation_row { + %activation_row = index.assume %local_row [range(%local_row, 0, 31)] : index + %assignment_i32 = view.load %route_stage_view[%activation_row] : view<32xi32> -> i32 + %valid_assignment = scalar.cmpi sge, %assignment_i32, %c0_i32 : i32 + %input_k = index.add %k_origin, %load_k : index + %valid_input = index.cmp ult, %input_k, %bounded_input_size : index + %valid_activation = scalar.andi %valid_assignment, %valid_input : i1 + %activation_values = scf.if %valid_activation -> (vector<4xf16>) { + %assignment0 = index.cast %assignment_i32 : i32 to index + %assignment = index.assume %assignment0 [range(%assignment0, 0, 65535)] : index + %bounded_assignment, %bounded_assignment_count = index.assume %assignment, %assignment_count [lt(%assignment, %assignment_count)] : index, index + %token0 = index.div %bounded_assignment, %bounded_route_count : index + %route0 = index.rem %bounded_assignment, %bounded_route_count : index + %token, %input_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %input_route0 = index.rem %route0, %bounded_input_route_count : index + %input_route = index.assume %input_route0 [range(%input_route0, 0, 31), lt(%input_route0, %bounded_input_route_count)] : index + %mask = vector.mask.range [%input_k to %bounded_input_size step %c1] : index -> vector<4xi1> + %loaded = vector.load.mask %input_view[%token, %input_route, %input_k], %mask, %c0_f32x4 : view<[%bounded_token_count]x[%bounded_input_route_count]x[%bounded_input_size]xf32>, vector<4xi1>, vector<4xf32> + %converted = vector.fptrunc %loaded : vector<4xf32> to vector<4xf16> + scf.yield %converted : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %activation_values, %activation_stage_physical_view[%activation_row, %load_k] : vector<4xf16>, view<32x40xf16> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %gate_next00, %gate_next01, %gate_next10, %gate_next11, %up_next00, %up_next01, %up_next10, %up_next11 = scf.for %k_half = [%c0 to %c32 step %c16](%gate_half_acc00 = %gate_acc00 : vector<4xf32>, %gate_half_acc01 = %gate_acc01 : vector<4xf32>, %gate_half_acc10 = %gate_acc10 : vector<4xf32>, %gate_half_acc11 = %gate_acc11 : vector<4xf32>, %up_half_acc00 = %up_acc00 : vector<4xf32>, %up_half_acc01 = %up_acc01 : vector<4xf32>, %up_half_acc10 = %up_acc10 : vector<4xf32>, %up_half_acc11 = %up_acc11 : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) unroll { + %gate_lhs0 = vector.fragment.load %gate_weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %gate_lhs1 = vector.fragment.load %gate_weight_stage_view[%subgroup_channel1, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %up_lhs0 = vector.fragment.load %up_weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %up_lhs1 = vector.fragment.load %up_weight_stage_view[%subgroup_channel1, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %rhs0 = vector.fragment.load %activation_fragment_view[%k_half, %c0] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %rhs1 = vector.fragment.load %activation_fragment_view[%k_half, %c16] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %gate_half_next00 = vector.mma %gate_lhs0, %rhs0, %gate_half_acc00 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %gate_half_next01 = vector.mma %gate_lhs0, %rhs1, %gate_half_acc01 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %gate_half_next10 = vector.mma %gate_lhs1, %rhs0, %gate_half_acc10 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %gate_half_next11 = vector.mma %gate_lhs1, %rhs1, %gate_half_acc11 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %up_half_next00 = vector.mma %up_lhs0, %rhs0, %up_half_acc00 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %up_half_next01 = vector.mma %up_lhs0, %rhs1, %up_half_acc01 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %up_half_next10 = vector.mma %up_lhs1, %rhs0, %up_half_acc10 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %up_half_next11 = vector.mma %up_lhs1, %rhs1, %up_half_acc11 : vector<16xf16>, vector<16xf16>, vector<4xf32> + scf.yield %gate_half_next00, %gate_half_next01, %gate_half_next10, %gate_half_next11, %up_half_next00, %up_half_next01, %up_half_next10, %up_half_next11 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %gate_next00, %gate_next01, %gate_next10, %gate_next11, %up_next00, %up_next01, %up_next10, %up_next11 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + scf.yield %gate_block_result00, %gate_block_result01, %gate_block_result10, %gate_block_result11, %up_block_result00, %up_block_result01, %up_block_result10, %up_block_result11 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + %publish_route0 = index.div %lane, %c4 : index + %publish_route = index.assume %publish_route0 [range(%publish_route0, 0, 15)] : index + %publish_packet0 = index.rem %lane, %c4 : index + %publish_packet = index.assume %publish_packet0 [range(%publish_packet0, 0, 3)] : index + %publish_channel_add = index.mul %publish_packet, %c4 : index + %local_route1 = index.add %c16, %publish_route : index + %assignment0_i32 = view.load %route_stage_view[%publish_route] : view<32xi32> -> i32 + %assignment1_i32 = view.load %route_stage_view[%local_route1] : view<32xi32> -> i32 + %assignment0_nonnegative = scalar.cmpi sge, %assignment0_i32, %c0_i32 : i32 + %assignment1_nonnegative = scalar.cmpi sge, %assignment1_i32, %c0_i32 : i32 + %safe_assignment0_i32 = scf.if %assignment0_nonnegative -> (i32) { + scf.yield %assignment0_i32 : i32 + } else { + scf.yield %c0_i32 : i32 + } + %safe_assignment1_i32 = scf.if %assignment1_nonnegative -> (i32) { + scf.yield %assignment1_i32 : i32 + } else { + scf.yield %c0_i32 : i32 + } + %safe_assignment0_0 = index.cast %safe_assignment0_i32 : i32 to index + %safe_assignment1_0 = index.cast %safe_assignment1_i32 : i32 to index + %safe_assignment0 = index.assume %safe_assignment0_0 [range(%safe_assignment0_0, 0, 65535)] : index + %safe_assignment1 = index.assume %safe_assignment1_0 [range(%safe_assignment1_0, 0, 65535)] : index + %bounded_assignment0, %bounded_assignment_count0 = index.assume %safe_assignment0, %assignment_count [lt(%safe_assignment0, %assignment_count)] : index, index + %bounded_assignment1, %bounded_assignment_count1 = index.assume %safe_assignment1, %assignment_count [lt(%safe_assignment1, %assignment_count)] : index, index + %subgroup_channel_base = index.add %channel_tile_base, %subgroup_channel_add : index + %channel0 = index.add %subgroup_channel_base, %publish_channel_add : index + %channel1_base = index.add %subgroup_channel_base, %c16 : index + %channel1 = index.add %channel1_base, %publish_channel_add : index + %valid_channel0 = index.cmp ult, %channel0, %bounded_output_size : index + %valid_channel1 = index.cmp ult, %channel1, %bounded_output_size : index + %writes00 = scalar.andi %assignment0_nonnegative, %valid_channel0 : i1 + %writes01 = scalar.andi %assignment1_nonnegative, %valid_channel0 : i1 + %writes10 = scalar.andi %assignment0_nonnegative, %valid_channel1 : i1 + %writes11 = scalar.andi %assignment1_nonnegative, %valid_channel1 : i1 + vector.fragment.store %gate_result00, %gate_result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + vector.fragment.store %up_result00, %up_result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes00 { + %gate_values = vector.load %gate_result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + %up_values = vector.load %up_result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + %activated = template.apply<@ggml.mul_mat_id_swiglu.apply_vector4>(%gate_values, %up_values) : (vector<4xf32>, vector<4xf32>) -> (vector<4xf32>) + %mask = vector.mask.range [%channel0 to %bounded_output_size step %c1] : index -> vector<4xi1> + vector.store.mask %activated, %output_view[%bounded_assignment0, %channel0], %mask : vector<4xf32>, view<[%output_row_count]x[%bounded_output_size]xf32>, vector<4xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %gate_result01, %gate_result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + vector.fragment.store %up_result01, %up_result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes01 { + %gate_values = vector.load %gate_result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + %up_values = vector.load %up_result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + %activated = template.apply<@ggml.mul_mat_id_swiglu.apply_vector4>(%gate_values, %up_values) : (vector<4xf32>, vector<4xf32>) -> (vector<4xf32>) + %mask = vector.mask.range [%channel0 to %bounded_output_size step %c1] : index -> vector<4xi1> + vector.store.mask %activated, %output_view[%bounded_assignment1, %channel0], %mask : vector<4xf32>, view<[%output_row_count]x[%bounded_output_size]xf32>, vector<4xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %gate_result10, %gate_result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + vector.fragment.store %up_result10, %up_result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes10 { + %gate_values = vector.load %gate_result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + %up_values = vector.load %up_result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + %activated = template.apply<@ggml.mul_mat_id_swiglu.apply_vector4>(%gate_values, %up_values) : (vector<4xf32>, vector<4xf32>) -> (vector<4xf32>) + %mask = vector.mask.range [%channel1 to %bounded_output_size step %c1] : index -> vector<4xi1> + vector.store.mask %activated, %output_view[%bounded_assignment0, %channel1], %mask : vector<4xf32>, view<[%output_row_count]x[%bounded_output_size]xf32>, vector<4xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %gate_result11, %gate_result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + vector.fragment.store %up_result11, %up_result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes11 { + %gate_values = vector.load %gate_result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + %up_values = vector.load %up_result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + %activated = template.apply<@ggml.mul_mat_id_swiglu.apply_vector4>(%gate_values, %up_values) : (vector<4xf32>, vector<4xf32>) -> (vector<4xf32>) + %mask = vector.mask.range [%channel1 to %bounded_output_size step %c1] : index -> vector<4xi1> + vector.store.mask %activated, %output_view[%bounded_assignment1, %channel1], %mask : vector<4xf32>, view<[%output_row_count]x[%bounded_output_size]xf32>, vector<4xi1> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + } + template.return +} + +kernel.def target(@ggml_mul_mat_id_swiglu_gfx11_wave64) @ggml_mul_mat_id_swiglu_f32_f32_wmma(%token_count: index) { + %expert_count = config.get @ggml.mul_mat_id_swiglu.expert_count : index + %route_count = config.get @ggml.mul_mat_id_swiglu.route_count : index + %output_size = config.get @ggml.mul_mat_id_swiglu.output_size : index + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %assignment_count = index.mul %token_count, %route_count : index + %rounded_assignment_count = index.add %assignment_count, %c31 : index + %assignment_partition_count = index.div %rounded_assignment_count, %c32 : index + %has_more_assignment_partitions = index.cmp ugt, %assignment_partition_count, %expert_count : index + %launch_partition_count = scf.select %has_more_assignment_partitions, %assignment_partition_count, %expert_count : index + kernel.launch.config workgroups(%output_tiles, %launch_partition_count, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %expert_table: buffer, %partition_table: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %input_size = config.get @ggml.mul_mat_id_swiglu.input_size : index + %route_count = config.get @ggml.mul_mat_id_swiglu.route_count : index + %input_route_count = config.get @ggml.mul_mat_id_swiglu.input_route_count : index + %expert_count = config.get @ggml.mul_mat_id_swiglu.expert_count : index + %output_size = config.get @ggml.mul_mat_id_swiglu.output_size : index + %gate_weight_format = config.get @ggml.mul_mat_id_swiglu.gate_weight_format : index + %up_weight_format = config.get @ggml.mul_mat_id_swiglu.up_weight_format : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %partition_ordinal = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_id_swiglu.body>(%gate_weight_format, %up_weight_format, %bounded_token_count, %input_size, %output_size, %route_count, %input_route_count, %expert_count, %channel_tile, %partition_ordinal, %input, %expert_table, %partition_table, %gate_weight, %up_weight, %output) : (index, index, index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_q5_k_q8_plane_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_q5_k_q8_plane_wmma.loom new file mode 100644 index 000000000000..7ea86aff87fb --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_q5_k_q8_plane_wmma.loom @@ -0,0 +1,1409 @@ +template.decl @ggml.mul_mat_q8_1_x4.entry(%token_count: index, %paired_weights: i1, %q8_input: buffer, %weight: buffer, %peer_weight: buffer, %output: buffer) + +// Generic Q5_K projection consuming the Q8_1 plane published by the fused +// gate/up path. Weights use the backend-resident symmetric-I5/K32 carrier; +// input, output, and token dimensions are all JIT configuration values. + +template.decl @ggml.mul_mat_q5_k_q8_plane.wmmai8.body(%is_iq4xs: i1, %token256: i1, %q8_plane: i1, %paired_weights: i1, %peer_weight: buffer, %src0_na: buffer, %dst_na: buffer, %q8_na: buffer, %base: offset, %lds_bytes: offset, %w_off: offset, %as_off: offset, %ws_off: offset, %asum_off: offset, %wc_off: offset, %c0: index, %c1: index, %c2: index, %c4: index, %c8: index, %c16: index, %c32: index, %c17: index, %c34: index, %k_kc: index, %k_astride: index, %k_wstride: index, %k_m: index, %k_n: index, %k_blocks: index, %k_wave: index, %k_wm: index, %fzero: vector<8xf32>, %izero: vector<8xi32>, %k_b: index, %rows_b: index, %cols_b: index, %nchunks: index, %kblocks: index, %col_tile: index, %row_tile: index, %tid: index, %wave: index, %is_symi5: i1, %is_q4: i1, %is_q4_packed: i1) + +amdgpu.target @ggml_q5_plane_dense_gfx11_wave32 {subgroup_size = 32} + +config.decl @ggml.mul_mat_q5_k_q8_plane.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +config.decl @ggml.mul_mat_q5_k_q8_plane.output_size : %value: index where [range(%value, 1, 262144)] + +config.decl @ggml.mul_mat_q5_k_q8_plane.token_capacity : %value: index where [range(%value, 1, 2048), range(%value, 1, 2048)] + +config.decl @ggml.mul_mat_q8_1_x4.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +config.decl @ggml.mul_mat_q8_1_x4.output_size : %value: index where [range(%value, 64, 262144), mul(%value, 64)] + +config.decl @ggml.mul_mat_q8_1_x4.token_capacity : %value: index where [range(%value, 256, 2048), mul(%value, 256)] + +config.decl @ggml.mul_mat_q8_1_x4.weight_format : %value: index where [range(%value, 4, 44)] + +// Paired projections share a tile without persisting gate/up intermediates. +// Select loaded words rather than opaque buffer identities in device code. +func.def inline @ggml_q8_1_x4_paired_weight_words(%paired: i1, %load_up: i1, %weight: buffer, %peer: buffer, %offset: offset) -> (vector<4xi32>) { + %zero = index.constant 0 : index + %weight_view = buffer.view %weight[%offset] : buffer -> view<4xi32> + %weight_words = vector.load %weight_view[%zero] : view<4xi32> -> vector<4xi32> + %peer_words = scf.if %paired -> (vector<4xi32>) { + %peer_view = buffer.view %peer[%offset] : buffer -> view<4xi32> + %words = vector.load %peer_view[%zero] : view<4xi32> -> vector<4xi32> + scf.yield %words : vector<4xi32> + } else { + scf.yield %weight_words : vector<4xi32> + } + %words = scf.select %load_up, %peer_words, %weight_words : vector<4xi32> + func.return %words : vector<4xi32> +} + +// Decodes one Q4_K scale/minimum pair from a header loaded as three packed +// scale words. Keeping the complete header in registers lets paired-group +// contractions share both the header and packed-code load. +func.def inline @ggml_q5_plane_q4k_scale_from_header(%scale0: i32, %scale1: i32, %scale2: i32, %q4_group: index) -> (i32, i32) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c48_i32 = scalar.constant 48 : i32 + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + %is_low_group = index.cmp ult, %bounded_group, %c4 : index + %scale_lane = index.rem %bounded_group, %c4 : index + %scale_shift_index = index.mul %scale_lane, %c8 : index + %scale_shift = index.cast %scale_shift_index : index to i32 + %high_shift = scalar.addi %scale_shift, %c2_i32 : i32 + %minimum_shift = scalar.addi %scale_shift, %c4_i32 : i32 + %selected_scale_source = scf.select %is_low_group, %scale0, %scale2 : i32 + %selected_minimum_source = scf.select %is_low_group, %scale1, %scale2 : i32 + %selected_scale_high_shift = scf.select %is_low_group, %scale_shift, %high_shift : i32 + %selected_minimum_low_shift = scf.select %is_low_group, %scale_shift, %minimum_shift : i32 + %scale_low0 = scalar.shrui %selected_scale_source, %scale_shift : i32 + %scale_low = scalar.andi %scale_low0, %c15_i32 : i32 + %scale_high0 = scalar.shrui %scale0, %selected_scale_high_shift : i32 + %scale_high = scalar.andi %scale_high0, %c48_i32 : i32 + %scale = scalar.ori %scale_low, %scale_high : i32 + %minimum_low0 = scalar.shrui %selected_minimum_source, %selected_minimum_low_shift : i32 + %minimum_low = scalar.andi %minimum_low0, %c15_i32 : i32 + %minimum_high0 = scalar.shrui %scale1, %selected_scale_high_shift : i32 + %minimum_high = scalar.andi %minimum_high0, %c48_i32 : i32 + %minimum = scalar.ori %minimum_low, %minimum_high : i32 + func.return %scale, %minimum : i32, i32 +} + +// Packs either nibble from four adjacent GGUF IQ4_XS bytes for the gfx11 table lookup. +// The nibble is the table index as it is: vector.table.lookup indexes the table in order on gfx1151 +// (the XOR 12 remap made every IQ4_XS weight of this prefill path wrong; same fix as motifs/dequant.loom). +func.def inline @ggml_q5_plane_dense_iq4xs_wmma_table_codes4(%code_word: vector<1xi32>, %uses_high: i1) -> (vector<4xi8>) { + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c12_i32 = scalar.constant 12 : i32 + %c16_i32 = scalar.constant 16 : i32 + %mask0_i32 = scalar.constant 15 : i32 + %mask1_i32 = scalar.constant 240 : i32 + %mask2_i32 = scalar.constant 3840 : i32 + %mask3_i32 = scalar.constant 61440 : i32 + %word = vector.extract %code_word[0] : vector<1xi32> -> i32 + %word_shr4 = scalar.shrui %word, %c4_i32 : i32 + %word_shr8 = scalar.shrui %word, %c8_i32 : i32 + %word_shr12 = scalar.shrui %word, %c12_i32 : i32 + %word_shr16 = scalar.shrui %word, %c16_i32 : i32 + %low0 = scalar.andi %word, %mask0_i32 : i32 + %low1 = scalar.andi %word_shr4, %mask1_i32 : i32 + %low2 = scalar.andi %word_shr8, %mask2_i32 : i32 + %low3 = scalar.andi %word_shr12, %mask3_i32 : i32 + %low01 = scalar.ori %low0, %low1 : i32 + %low23 = scalar.ori %low2, %low3 : i32 + %low = scalar.ori %low01, %low23 : i32 + %high0 = scalar.andi %word_shr4, %mask0_i32 : i32 + %high1 = scalar.andi %word_shr8, %mask1_i32 : i32 + %high2 = scalar.andi %word_shr12, %mask2_i32 : i32 + %high3 = scalar.andi %word_shr16, %mask3_i32 : i32 + %high01 = scalar.ori %high0, %high1 : i32 + %high23 = scalar.ori %high2, %high3 : i32 + %high = scalar.ori %high01, %high23 : i32 + %selected = scf.select %uses_high, %high, %low : i32 + %target_selected_shr8 = scalar.shrui %selected, %c8_i32 : i32 + %packed0 = scalar.trunci %selected : i32 to i8 + %packed1 = scalar.trunci %target_selected_shr8 : i32 to i8 + %packed = vector.from_elements %packed0, %packed1 : vector<2xi8> + %codes = vector.bitunpacku<4> %packed : vector<2xi8> -> vector<4xi8> + func.return %codes : vector<4xi8> +} + +func.def inline @ggml_q5_plane_dense_iq4xs_wmmai8_signed_values16(%code_words: vector<4xi32>, %uses_high: i1) -> (vector<4xi8>, vector<4xi8>, vector<4xi8>, vector<4xi8>) { + %v0 = scalar.constant -127 : i8 + %v1 = scalar.constant -104 : i8 + %v2 = scalar.constant -83 : i8 + %v3 = scalar.constant -65 : i8 + %v4 = scalar.constant -49 : i8 + %v5 = scalar.constant -35 : i8 + %v6 = scalar.constant -22 : i8 + %v7 = scalar.constant -10 : i8 + %v8 = scalar.constant 1 : i8 + %v9 = scalar.constant 13 : i8 + %v10 = scalar.constant 25 : i8 + %v11 = scalar.constant 38 : i8 + %v12 = scalar.constant 53 : i8 + %v13 = scalar.constant 69 : i8 + %v14 = scalar.constant 89 : i8 + %v15 = scalar.constant 113 : i8 + %value_table = vector.from_elements %v0, %v1, %v2, %v3, %v4, %v5, %v6, %v7, %v8, %v9, %v10, %v11, %v12, %v13, %v14, %v15 : vector<16xi8> + %word0_i32 = vector.extract %code_words[0] : vector<4xi32> -> i32 + %word1_i32 = vector.extract %code_words[1] : vector<4xi32> -> i32 + %word2_i32 = vector.extract %code_words[2] : vector<4xi32> -> i32 + %word3_i32 = vector.extract %code_words[3] : vector<4xi32> -> i32 + %word0 = vector.from_elements %word0_i32 : vector<1xi32> + %word1 = vector.from_elements %word1_i32 : vector<1xi32> + %word2 = vector.from_elements %word2_i32 : vector<1xi32> + %word3 = vector.from_elements %word3_i32 : vector<1xi32> + %codes0 = func.call @ggml_q5_plane_dense_iq4xs_wmma_table_codes4(%word0, %uses_high) : (vector<1xi32>, i1) -> (vector<4xi8>) + %codes1 = func.call @ggml_q5_plane_dense_iq4xs_wmma_table_codes4(%word1, %uses_high) : (vector<1xi32>, i1) -> (vector<4xi8>) + %codes2 = func.call @ggml_q5_plane_dense_iq4xs_wmma_table_codes4(%word2, %uses_high) : (vector<1xi32>, i1) -> (vector<4xi8>) + %codes3 = func.call @ggml_q5_plane_dense_iq4xs_wmma_table_codes4(%word3, %uses_high) : (vector<1xi32>, i1) -> (vector<4xi8>) + %values0_i8 = vector.table.lookup %value_table[%codes0] : vector<16xi8>, vector<4xi8> -> vector<4xi8> + %values1_i8 = vector.table.lookup %value_table[%codes1] : vector<16xi8>, vector<4xi8> -> vector<4xi8> + %values2_i8 = vector.table.lookup %value_table[%codes2] : vector<16xi8>, vector<4xi8> -> vector<4xi8> + %values3_i8 = vector.table.lookup %value_table[%codes3] : vector<16xi8>, vector<4xi8> -> vector<4xi8> + func.return %values0_i8, %values1_i8, %values2_i8, %values3_i8 : vector<4xi8>, vector<4xi8>, vector<4xi8>, vector<4xi8> +} + +func.def inline @ggml_q5_plane_dense_iq4xs_wmma_scale(%weight: buffer, %row_byte_base: offset, %iq4_block: index, %iq4_group: index) -> (f32) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %block_bytes = index.constant 136 : offset + %scale_high_offset = index.constant 2 : offset + %scale_low_offset = index.constant 4 : offset + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c32_i32 = scalar.constant 32 : i32 + %bounded_group = index.assume %iq4_group [range(%iq4_group, 0, 7)] : index + %block_byte_add = index.scale %iq4_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %scale_high_byte_base = index.add %block_byte_base, %scale_high_offset : offset + %scale_low_byte_base = index.add %block_byte_base, %scale_low_offset : offset + %d_view = buffer.view %weight[%block_byte_base] : buffer -> view<1xf16> + %scale_high_view = buffer.view %weight[%scale_high_byte_base] : buffer -> view<1xi16> + %scale_low_view = buffer.view %weight[%scale_low_byte_base] : buffer -> view<4xi8> + %d_f16 = view.load %d_view[%c0] : view<1xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %scale_low_index = index.div %bounded_group, %c2 : index + %scale_low_i8 = view.load %scale_low_view[%scale_low_index] : view<4xi8> -> i8 + %scale_low_i32 = scalar.extui %scale_low_i8 : i8 to i32 + %scale_low_half = index.rem %bounded_group, %c2 : index + %scale_low_shift_index = index.mul %scale_low_half, %c4 : index + %scale_low_shift = index.cast %scale_low_shift_index : index to i32 + %scale_low_shifted = scalar.shrui %scale_low_i32, %scale_low_shift : i32 + %scale_low = scalar.andi %scale_low_shifted, %c15_i32 : i32 + %scale_high_i16 = view.load %scale_high_view[%c0] : view<1xi16> -> i16 + %scale_high_i32 = scalar.extui %scale_high_i16 : i16 to i32 + %scale_high_shift_index = index.mul %bounded_group, %c2 : index + %scale_high_shift = index.cast %scale_high_shift_index : index to i32 + %scale_high_shifted = scalar.shrui %scale_high_i32, %scale_high_shift : i32 + %scale_high_two_bits = scalar.andi %scale_high_shifted, %c3_i32 : i32 + %scale_high = scalar.shli %scale_high_two_bits, %c4_i32 : i32 + %scale_unsigned = scalar.ori %scale_low, %scale_high : i32 + %scale_signed = scalar.subi %scale_unsigned, %c32_i32 : i32 + %scale_f32 = scalar.sitofp %scale_signed : i32 to f32 + %d_scale = scalar.mulf %d, %scale_f32 : f32 + func.return %d_scale : f32 +} + +// Token-128 and token-256 share one body; a compile-time selector removes the unused statically bounded LDS view. +func.def inline @ggml_q5_plane_dense_q8_wmmai8_stage_values(%token256: i1, %scratch: buffer, %base: offset, %adbase: index, %ad1: index, %values0: vector<32xi8>, %values1: vector<32xi8>) { + scf.if %token256 { + %al_flat = buffer.view %scratch[%base] : buffer -> view<20480xi8> + %ad0_bounded = index.assume %adbase [range(%adbase, 0, 20400)] : index + %ad1_bounded = index.assume %ad1 [range(%ad1, 32, 20432)] : index + vector.store %values0, %al_flat[%ad0_bounded] : vector<32xi8>, view<20480xi8> + vector.store %values1, %al_flat[%ad1_bounded] : vector<32xi8>, view<20480xi8> + } else { + %al_flat = buffer.view %scratch[%base] : buffer -> view<10240xi8> + %ad0_bounded = index.assume %adbase [range(%adbase, 0, 10160)] : index + %ad1_bounded = index.assume %ad1 [range(%ad1, 32, 10192)] : index + vector.store %values0, %al_flat[%ad0_bounded] : vector<32xi8>, view<10240xi8> + vector.store %values1, %al_flat[%ad1_bounded] : vector<32xi8>, view<10240xi8> + } + func.return +} + +func.def inline @ggml_q5_plane_dense_q8_wmmai8_stage_metadata(%token256: i1, %scratch: buffer, %as_off: offset, %asum_off: offset, %dst0: index, %dst1: index, %d0: f32, %d1: f32, %sum0: f32, %sum1: f32) { + scf.if %token256 { + %scale_view = buffer.view %scratch[%as_off] : buffer -> view<512xf32> + %sum_view = buffer.view %scratch[%asum_off] : buffer -> view<512xf32> + %dst0_bounded = index.assume %dst0 [range(%dst0, 0, 495)] : index + %dst1_bounded = index.assume %dst1 [range(%dst1, 16, 511)] : index + view.store %d0, %scale_view[%dst0_bounded] : f32, view<512xf32> + view.store %d1, %scale_view[%dst1_bounded] : f32, view<512xf32> + view.store %sum0, %sum_view[%dst0_bounded] : f32, view<512xf32> + view.store %sum1, %sum_view[%dst1_bounded] : f32, view<512xf32> + } else { + %scale_view = buffer.view %scratch[%as_off] : buffer -> view<256xf32> + %sum_view = buffer.view %scratch[%asum_off] : buffer -> view<256xf32> + %dst0_bounded = index.assume %dst0 [range(%dst0, 0, 239)] : index + %dst1_bounded = index.assume %dst1 [range(%dst1, 16, 255)] : index + view.store %d0, %scale_view[%dst0_bounded] : f32, view<256xf32> + view.store %d1, %scale_view[%dst1_bounded] : f32, view<256xf32> + view.store %sum0, %sum_view[%dst0_bounded] : f32, view<256xf32> + view.store %sum1, %sum_view[%dst1_bounded] : f32, view<256xf32> + } + func.return +} + +func.def inline @ggml_q5_plane_dense_q8_wmmai8_load_lhs(%token256: i1, %scratch: buffer, %base: offset, %row: index, %column: index, %c16: index) -> (vector<16xi8>) { + %lhs_layout = encoding.layout.strided [80, 1] : encoding + %result = scf.if %token256 -> (vector<16xi8>) { + %al_view = buffer.view %scratch[%base] : buffer -> view<256x64xi8, %lhs_layout> + %values = vector.fragment.load %al_view[%row, %column] shape [%c16, %c16] : view<256x64xi8, %lhs_layout> -> vector<16xi8> + scf.yield %values : vector<16xi8> + } else { + %al_view = buffer.view %scratch[%base] : buffer -> view<128x64xi8, %lhs_layout> + %values = vector.fragment.load %al_view[%row, %column] shape [%c16, %c16] : view<128x64xi8, %lhs_layout> -> vector<16xi8> + scf.yield %values : vector<16xi8> + } + func.return %result : vector<16xi8> +} + +func.def inline @ggml_q5_plane_dense_q8_wmmai8_load_metadata(%token256: i1, %scratch: buffer, %as_off: offset, %asum_off: offset, %index: index) -> (vector<8xf32>, vector<8xf32>) { + %scale, %sum = scf.if %token256 -> (vector<8xf32>, vector<8xf32>) { + %scale_view = buffer.view %scratch[%as_off] : buffer -> view<512xf32> + %sum_view = buffer.view %scratch[%asum_off] : buffer -> view<512xf32> + %bounded_index = index.assume %index [range(%index, 0, 504)] : index + %scale_values = vector.load %scale_view[%bounded_index] : view<512xf32> -> vector<8xf32> + %sum_values = vector.load %sum_view[%bounded_index] : view<512xf32> -> vector<8xf32> + scf.yield %scale_values, %sum_values : vector<8xf32>, vector<8xf32> + } else { + %scale_view = buffer.view %scratch[%as_off] : buffer -> view<256xf32> + %sum_view = buffer.view %scratch[%asum_off] : buffer -> view<256xf32> + %bounded_index = index.assume %index [range(%index, 0, 120)] : index + %scale_values = vector.load %scale_view[%bounded_index] : view<256xf32> -> vector<8xf32> + %sum_values = vector.load %sum_view[%bounded_index] : view<256xf32> -> vector<8xf32> + scf.yield %scale_values, %sum_values : vector<8xf32>, vector<8xf32> + } + func.return %scale, %sum : vector<8xf32>, vector<8xf32> +} + +// Raw Q5_K x packed-Q8_1_x4 IU8 WMMA using the existing Loom activation layout. +// Keep one weighted K32 fragment live while updating the paired projection tile. +func.def inline @ggml_q8_1_x4_affine_rhs_tile(%scratch: buffer, %weight_offset: offset, %scale_offset: offset, %correction_offset: offset, %block: index, %column: index, %lane: index, %lhs0: vector<4xi32>, %lhs1: vector<4xi32>, %zero: vector<8xi32>, %activation_scale: vector<8xf32>, %activation_sum: vector<8xf32>, %accumulator: vector<8xf32>) -> (vector<8xf32>) { + %c2 = index.constant 2 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %row_layout = encoding.layout.strided [1, 80] : encoding + %rhs_schema = encoding.define #encoding.operand : encoding + %weight = buffer.view %scratch[%weight_offset] : buffer -> view<64x128xi8, %row_layout> + %scales = buffer.view %scratch[%scale_offset] : buffer -> view<256xf32> + %corrections = buffer.view %scratch[%correction_offset] : buffer -> view<256xf32> + %row = index.add %column, %lane : index + %row_pair = index.mul %row, %c2 : index + %metadata0 = index.add %row_pair, %block : index + %metadata = index.assume %metadata0 [range(%metadata0, 0, 255)] : index + %scale = view.load %scales[%metadata] : view<256xf32> -> f32 + %scale_vector = vector.splat %scale : vector<8xf32> + %correction = view.load %corrections[%metadata] : view<256xf32> -> f32 + %correction_vector = vector.splat %correction : vector<8xf32> + %k0 = index.mul %block, %c32 : index + %k1 = index.add %k0, %c16 : index + %rhs0_raw = vector.fragment.load %weight[%k0, %column] shape [%c16, %c16] : view<64x128xi8, %row_layout> -> vector<16xi8> + %rhs0_words = vector.bitcast %rhs0_raw : vector<16xi8> to vector<4xi32> + %rhs0 = vector.fragment %rhs0_words shape [%c16, %c16] using {schema = %rhs_schema : encoding} : vector<4xi32> + %rhs1_raw = vector.fragment.load %weight[%k1, %column] shape [%c16, %c16] : view<64x128xi8, %row_layout> -> vector<16xi8> + %rhs1_words = vector.bitcast %rhs1_raw : vector<16xi8> to vector<4xi32> + %rhs1 = vector.fragment %rhs1_words shape [%c16, %c16] using {schema = %rhs_schema : encoding} : vector<4xi32> + %dot0 = vector.mma %lhs0, %rhs0, %zero : vector<4xi32>, vector<4xi32>, vector<8xi32> + %dot = vector.mma %lhs1, %rhs1, %dot0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %dot_f32 = vector.sitofp %dot : vector<8xi32> to vector<8xf32> + %product_scale = vector.mulf %activation_scale, %scale_vector : vector<8xf32> + %base = vector.fmaf %activation_sum, %correction_vector, %accumulator : vector<8xf32> + %result = vector.fmaf %dot_f32, %product_scale, %base : vector<8xf32> + func.return %result : vector<8xf32> +} + +template.def<@ggml.mul_mat_q5_k_q8_plane.wmmai8.body> device @ggml_mul_mat_q5_k_q8_plane_wmmai8_body(%is_iq4xs: i1, %token256: i1, %q8_plane: i1, %paired_weights: i1, %peer_weight: buffer, %src0_na: buffer, %dst_na: buffer, %q8_na: buffer, %base: offset, %lds_bytes: offset, %w_off: offset, %as_off: offset, %ws_off: offset, %asum_off: offset, %wc_off: offset, %c0: index, %c1: index, %c2: index, %c4: index, %c8: index, %c16: index, %c32: index, %c17: index, %c34: index, %k_kc: index, %k_astride: index, %k_wstride: index, %k_m: index, %k_n: index, %k_blocks: index, %k_wave: index, %k_wm: index, %fzero: vector<8xf32>, %izero: vector<8xi32>, %k_b: index, %rows_b: index, %cols_b: index, %nchunks: index, %kblocks: index, %col_tile: index, %row_tile: index, %tid: index, %wave: index, %is_symi5: i1, %is_q4: i1, %is_q4_packed: i1) { + %c64_wide = index.constant 64 : index + %weight_elements = index.mul %k_n, %k_wstride : index + %weight_metadata = index.mul %k_n, %c2 : index + %q4_k_block = index.constant 256 : index + %q5_native_block_bytes = index.constant 176 : offset + %q4_block_bytes = index.constant 144 : offset + %q5_block_bytes = scf.select %is_q4, %q4_block_bytes, %q5_native_block_bytes : offset + %q5_high_offset = index.constant 16 : offset + %q5_native_code_offset = index.constant 48 : offset + %q4_code_offset = index.constant 16 : offset + %q5_code_offset = scf.select %is_q4, %q4_code_offset, %q5_native_code_offset : offset + %q8_group_bytes = index.constant 144 : offset + %q8_payload_offset = index.constant 16 : offset + %q8_elements_per_group = index.constant 128 : index + // dst is [rows, cols] with rows contiguous, and C is [cols, rows]. + %dst_view = buffer.view %dst_na[%base] : buffer -> view<[%cols_b]x[%rows_b]xf32> + + // Waves share the staged weight tile across their token rows. + %scratch = buffer.alloca align(16) %lds_bytes : buffer + %i8_schema = encoding.define #encoding.operand : encoding + // Compile-time formats select unsigned Q5_K or signed IQ4_XS/I5 operands. + %signed_rhs_schema = encoding.define #encoding.operand : encoding + %unsigned_rhs_schema = encoding.define #encoding.operand : encoding + %is_signed_rhs = scalar.ori %is_iq4xs, %is_symi5 : i1 + %u8_schema = scf.select %is_signed_rhs, %signed_rhs_schema, %unsigned_rhs_schema : encoding + %rhs_layout = encoding.layout.strided [1, 80] : encoding + %wl_view = buffer.view %scratch[%w_off] : buffer -> view<64x[%k_n]xi8, %rhs_layout> + %wl_flat = buffer.view %scratch[%w_off] : buffer -> view<[%weight_elements]xi8> + %wsl_view = buffer.view %scratch[%ws_off] : buffer -> view<[%weight_metadata]xf32> + %wcl_view = buffer.view %scratch[%wc_off] : buffer -> view<[%weight_metadata]xf32> + + %col_base = index.mul %col_tile, %k_m : index + %output_tile_rows = scf.select %paired_weights, %c64_wide, %k_n : index + %row_base = index.mul %row_tile, %output_tile_rows : index + %up_half = index.cmp uge, %tid, %c64_wide : index + %load_up = scalar.andi %paired_weights, %up_half : i1 + + // One thread stages both K32 blocks for a token from the 144-byte Q8_1_x4 group. + %ascol_g = index.add %col_base, %tid : index + %adbase = index.mul %tid, %k_astride : index + %q8_groups_per_row = index.div %k_b, %q8_elements_per_group : index + %q8_row_bytes = index.scale %q8_groups_per_row, %q8_group_bytes : index, offset -> offset + %q8_row_byte_base = index.scale %ascol_g, %q8_row_bytes : index, offset -> offset + %q8_plane_payload_elements = index.mul %k_b, %cols_b : index + %q8_plane_payload_words = index.div %q8_plane_payload_elements, %c4 : index + %q8_plane_metadata_offset = index.cast %q8_plane_payload_elements : index to offset + %q8_plane_metadata_count = index.mul %kblocks, %cols_b : index + %q8_plane_payload = buffer.view %q8_na[%base] : buffer -> view<[%q8_plane_payload_words]xi32> + %q8_plane_metadata = buffer.view %q8_na[%q8_plane_metadata_offset] : buffer -> view<[%q8_plane_metadata_count]xi32> + + // Paired projections use fewer token rows and a wider shared weight tile. + %wcol = index.rem %wave, %k_wm : index + %wrow = index.div %wave, %k_wm : index + %k_mspan = scf.select %paired_weights, %c16, %c32 : index + %k_nspan = index.add %k_n, %c0 : index + %wm_off = index.mul %wcol, %k_mspan : index + %wn_off = index.mul %wrow, %k_nspan : index + %m_out = index.add %col_base, %wm_off : index + %n_out = index.add %row_base, %wn_off : index + %lm0 = index.add %wm_off, %c0 : index + %gm0 = index.add %m_out, %c0 : index + %k_ma1 = index.constant 16 : index + %lm1 = index.add %wm_off, %k_ma1 : index + %gm1 = index.add %m_out, %k_ma1 : index + %ln0 = index.add %wn_off, %c0 : index + %gn0 = index.add %n_out, %c0 : index + %k_nb1 = index.constant 16 : index + %ln1 = index.add %wn_off, %k_nb1 : index + %gn1 = index.add %n_out, %k_nb1 : index + %k_nb2 = index.constant 32 : index + %ln2 = index.add %wn_off, %k_nb2 : index + %gn2 = index.add %n_out, %k_nb2 : index + %k_nb3 = index.constant 48 : index + %ln3 = index.add %wn_off, %k_nb3 : index + %gn3 = index.add %n_out, %k_nb3 : index + %lane = index.rem %tid, %k_wave : index + %lane_lo = index.rem %lane, %c16 : index + %lane_hi = index.div %lane, %c16 : index + %wsb0_0 = index.add %ln0, %lane_lo : index + %wsb0 = index.mul %wsb0_0, %k_blocks : index + %wsb1_0 = index.add %ln1, %lane_lo : index + %wsb1 = index.mul %wsb1_0, %k_blocks : index + %wsb2_0 = index.add %ln2, %lane_lo : index + %wsb2 = index.mul %wsb2_0, %k_blocks : index + %wsb3_0 = index.add %ln3, %lane_lo : index + %wsb3 = index.mul %wsb3_0, %k_blocks : index + // Activation scale metadata is [tile][block][lane-half][register]. + %as_tile0 = index.div %lm0, %c16 : index + %asb0 = index.mul %as_tile0, %k_blocks : index + %as_tile1 = index.div %lm1, %c16 : index + %asb1 = index.mul %as_tile1, %k_blocks : index + + %f0_0, %f0_1, %f0_2, %f0_3, %f1_0, %f1_1, %f1_2, %f1_3 = scf.for %chunk = [%c0 to %nchunks step %c1](%fc0_0 = %fzero : vector<8xf32>, %fc0_1 = %fzero : vector<8xf32>, %fc0_2 = %fzero : vector<8xf32>, %fc0_3 = %fzero : vector<8xf32>, %fc1_0 = %fzero : vector<8xf32>, %fc1_1 = %fzero : vector<8xf32>, %fc1_2 = %fzero : vector<8xf32>, %fc1_3 = %fzero : vector<8xf32>) -> (vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) { + %activation_active = index.cmp ult, %tid, %k_m : index + scf.if %activation_active { + %q8_words0, %q8_words1, %q8_d0_f16, %q8_s0_f16, %q8_d1_f16, %q8_s1_f16 = scf.if %q8_plane -> (vector<8xi32>, vector<8xi32>, f16, f16, f16, f16) { + %q8_block0 = index.mul %chunk, %c2 : index + %q8_block1 = index.add %q8_block0, %c1 : index + %q8_block0_token_base = index.mul %q8_block0, %cols_b : index + %q8_block1_token_base = index.mul %q8_block1, %cols_b : index + %q8_block0_token0 = index.add %q8_block0_token_base, %ascol_g : index + %q8_block1_token0 = index.add %q8_block1_token_base, %ascol_g : index + %q8_block0_token = index.assume %q8_block0_token0 [range(%q8_block0_token0, 0, 33554431)] : index + %q8_block1_token = index.assume %q8_block1_token0 [range(%q8_block1_token0, 1, 33554431)] : index + %q8_payload_word0_0 = index.mul %q8_block0_token, %c8 : index + %q8_payload_word1_0 = index.mul %q8_block1_token, %c8 : index + %q8_payload_word0 = index.assume %q8_payload_word0_0 [range(%q8_payload_word0_0, 0, 268435448), mul(%q8_payload_word0_0, 8)] : index + %q8_payload_word1 = index.assume %q8_payload_word1_0 [range(%q8_payload_word1_0, 8, 268435448), mul(%q8_payload_word1_0, 8)] : index + %q8_plane_words0 = vector.load %q8_plane_payload[%q8_payload_word0] : view<[%q8_plane_payload_words]xi32> -> vector<8xi32> + %q8_plane_words1 = vector.load %q8_plane_payload[%q8_payload_word1] : view<[%q8_plane_payload_words]xi32> -> vector<8xi32> + %q8_meta0_word = view.load %q8_plane_metadata[%q8_block0_token] : view<[%q8_plane_metadata_count]xi32> -> i32 + %q8_meta1_word = view.load %q8_plane_metadata[%q8_block1_token] : view<[%q8_plane_metadata_count]xi32> -> i32 + %q8_meta0_vector = vector.from_elements %q8_meta0_word : vector<1xi32> + %q8_meta1_vector = vector.from_elements %q8_meta1_word : vector<1xi32> + %q8_ds0_vector = vector.bitcast %q8_meta0_vector : vector<1xi32> to vector<2xf16> + %q8_ds1_vector = vector.bitcast %q8_meta1_vector : vector<1xi32> to vector<2xf16> + %q8_plane_d0_f16 = vector.extract %q8_ds0_vector[0] : vector<2xf16> -> f16 + %q8_plane_s0_f16 = vector.extract %q8_ds0_vector[1] : vector<2xf16> -> f16 + %q8_plane_d1_f16 = vector.extract %q8_ds1_vector[0] : vector<2xf16> -> f16 + %q8_plane_s1_f16 = vector.extract %q8_ds1_vector[1] : vector<2xf16> -> f16 + scf.yield %q8_plane_words0, %q8_plane_words1, %q8_plane_d0_f16, %q8_plane_s0_f16, %q8_plane_d1_f16, %q8_plane_s1_f16 : vector<8xi32>, vector<8xi32>, f16, f16, f16, f16 + } else { + %q8_group = index.div %chunk, %c2 : index + %q8_pair = index.rem %chunk, %c2 : index + %q8_inner0 = index.mul %q8_pair, %c2 : index + %q8_inner1 = index.add %q8_inner0, %c1 : index + %q8_group_byte_add = index.scale %q8_group, %q8_group_bytes : index, offset -> offset + %q8_group_byte_base = index.add %q8_row_byte_base, %q8_group_byte_add : offset + %q8_payload_byte_base = index.add %q8_group_byte_base, %q8_payload_offset : offset + %q8_ds = buffer.view %q8_na[%q8_group_byte_base] : buffer -> view<8xf16> + %q8_payload = buffer.view %q8_na[%q8_payload_byte_base] : buffer -> view<32xi32> + %q8_word0 = index.mul %q8_inner0, %c8 : index + %q8_word1 = index.mul %q8_inner1, %c8 : index + %q8_standard_words0 = vector.load %q8_payload[%q8_word0] : view<32xi32> -> vector<8xi32> + %q8_standard_words1 = vector.load %q8_payload[%q8_word1] : view<32xi32> -> vector<8xi32> + %q8_ds0 = index.mul %q8_inner0, %c2 : index + %q8_ds1 = index.mul %q8_inner1, %c2 : index + %q8_s0 = index.add %q8_ds0, %c1 : index + %q8_s1 = index.add %q8_ds1, %c1 : index + %q8_standard_d0_f16 = view.load %q8_ds[%q8_ds0] : view<8xf16> -> f16 + %q8_standard_s0_f16 = view.load %q8_ds[%q8_s0] : view<8xf16> -> f16 + %q8_standard_d1_f16 = view.load %q8_ds[%q8_ds1] : view<8xf16> -> f16 + %q8_standard_s1_f16 = view.load %q8_ds[%q8_s1] : view<8xf16> -> f16 + scf.yield %q8_standard_words0, %q8_standard_words1, %q8_standard_d0_f16, %q8_standard_s0_f16, %q8_standard_d1_f16, %q8_standard_s1_f16 : vector<8xi32>, vector<8xi32>, f16, f16, f16, f16 + } + %q8_values0 = vector.bitcast %q8_words0 : vector<8xi32> to vector<32xi8> + %q8_values1 = vector.bitcast %q8_words1 : vector<8xi32> to vector<32xi8> + %ad1_0 = index.add %adbase, %c32 : index + func.call @ggml_q5_plane_dense_q8_wmmai8_stage_values(%token256, %scratch, %base, %adbase, %ad1_0, %q8_values0, %q8_values1) : (i1, buffer, offset, index, index, vector<32xi8>, vector<32xi8>) + + // Match the scale-plane permutation consumed by the WMMA fragments. + %q8_token_tile = index.div %tid, %c16 : index + %q8_token_inner = index.rem %tid, %c16 : index + %q8_token_parity = index.rem %q8_token_inner, %c2 : index + %q8_token_vector = index.div %q8_token_inner, %c2 : index + %q8_scale_tile0 = index.mul %q8_token_tile, %c4 : index + %q8_scale_tile1 = index.add %q8_scale_tile0, %q8_token_parity : index + %q8_scale_tile2 = index.mul %q8_scale_tile1, %c8 : index + %q8_scale_dst0_0 = index.add %q8_scale_tile2, %q8_token_vector : index + %q8_scale_dst1_0 = index.add %q8_scale_dst0_0, %c16 : index + %q8_d0 = scalar.extf %q8_d0_f16 : f16 to f32 + %q8_scaled_sum0 = scalar.extf %q8_s0_f16 : f16 to f32 + %q8_d1 = scalar.extf %q8_d1_f16 : f16 to f32 + %q8_scaled_sum1 = scalar.extf %q8_s1_f16 : f16 to f32 + func.call @ggml_q5_plane_dense_q8_wmmai8_stage_metadata(%token256, %scratch, %as_off, %asum_off, %q8_scale_dst0_0, %q8_scale_dst1_0, %q8_d0, %q8_d1, %q8_scaled_sum0, %q8_scaled_sum1) : (i1, buffer, offset, offset, index, index, f32, f32, f32, f32) + } + // One lane decodes two adjacent Q5_K groups; each selects its bits from the shared 32-byte fifth-bit plane. + %q5_weight_lane_count = index.add %k_n, %c0 : index + %q5_weight_lane_active = index.cmp ult, %tid, %q5_weight_lane_count : index + scf.if %q5_weight_lane_active { + scf.if %is_iq4xs { + %iq4_block_bytes = index.constant 136 : offset + %iq4_code_offset = index.constant 8 : offset + %iq4_weight_tid0 = index.rem %tid, %q5_weight_lane_count : index + %iq4_weight_tid = index.assume %iq4_weight_tid0 [range(%iq4_weight_tid0, 0, 63)] : index + %iq4_block_count = index.div %k_b, %q4_k_block : index + %iq4_block = index.div %chunk, %c4 : index + %iq4_pair = index.rem %chunk, %c4 : index + %iq4_group0 = index.mul %iq4_pair, %c2 : index + %iq4_group1 = index.add %iq4_group0, %c1 : index + %iq4_row0 = index.add %row_base, %iq4_weight_tid : index + %iq4_row = index.assume %iq4_row0 [range(%iq4_row0, 0, 262143)] : index + %iq4_row_bytes = index.scale %iq4_block_count, %iq4_block_bytes : index, offset -> offset + %iq4_row_byte_base = index.scale %iq4_row, %iq4_row_bytes : index, offset -> offset + %iq4_block_byte_add = index.scale %iq4_block, %iq4_block_bytes : index, offset -> offset + %iq4_block_byte_base = index.add %iq4_row_byte_base, %iq4_block_byte_add : offset + %iq4_code_byte_base = index.add %iq4_block_byte_base, %iq4_code_offset : offset + %iq4_code_view = buffer.view %src0_na[%iq4_code_byte_base] : buffer -> view<32xi32> + %iq4_scale0 = func.call @ggml_q5_plane_dense_iq4xs_wmma_scale(%src0_na, %iq4_row_byte_base, %iq4_block, %iq4_group0) : (buffer, offset, index, index) -> (f32) + %iq4_scale1 = func.call @ggml_q5_plane_dense_iq4xs_wmma_scale(%src0_na, %iq4_row_byte_base, %iq4_block, %iq4_group1) : (buffer, offset, index, index) -> (f32) + %iq4_meta0_0 = index.mul %iq4_weight_tid, %c2 : index + %iq4_meta0 = index.assume %iq4_meta0_0 [range(%iq4_meta0_0, 0, 126)] : index + %iq4_meta1_0 = index.add %iq4_meta0, %c1 : index + %iq4_meta1 = index.assume %iq4_meta1_0 [range(%iq4_meta1_0, 1, 127)] : index + %iq4_bias = scalar.constant -128.0 : f32 + %iq4_correction0 = scalar.mulf %iq4_scale0, %iq4_bias : f32 + %iq4_correction1 = scalar.mulf %iq4_scale1, %iq4_bias : f32 + view.store %iq4_scale0, %wsl_view[%iq4_meta0] : f32, view<[%weight_metadata]xf32> + view.store %iq4_scale1, %wsl_view[%iq4_meta1] : f32, view<[%weight_metadata]xf32> + view.store %iq4_correction0, %wcl_view[%iq4_meta0] : f32, view<[%weight_metadata]xf32> + view.store %iq4_correction1, %wcl_view[%iq4_meta1] : f32, view<[%weight_metadata]xf32> + + %iq4_group0_word_base0 = index.mul %iq4_group0, %c4 : index + %iq4_group0_word_base = index.assume %iq4_group0_word_base0 [range(%iq4_group0_word_base0, 0, 28)] : index + %iq4_group1_word_base0 = index.mul %iq4_group1, %c4 : index + %iq4_group1_word_base = index.assume %iq4_group1_word_base0 [range(%iq4_group1_word_base0, 4, 28)] : index + %iq4_weight_row_base = index.mul %iq4_weight_tid, %k_wstride : index + %iq4_w0 = index.assume %iq4_weight_row_base [range(%iq4_weight_row_base, 0, 5040)] : index + %iq4_payload16 = index.constant 16 : index + %iq4_payload48 = index.constant 48 : index + %iq4_w0_hi0 = index.add %iq4_weight_row_base, %iq4_payload16 : index + %iq4_w0_hi = index.assume %iq4_w0_hi0 [range(%iq4_w0_hi0, 16, 5056)] : index + %iq4_w1_0 = index.add %iq4_weight_row_base, %c32 : index + %iq4_w1 = index.assume %iq4_w1_0 [range(%iq4_w1_0, 32, 5072)] : index + %iq4_w1_hi0 = index.add %iq4_weight_row_base, %iq4_payload48 : index + %iq4_w1_hi = index.assume %iq4_w1_hi0 [range(%iq4_w1_hi0, 48, 5088)] : index + %iq4_false = scalar.constant false : i1 + %iq4_true = scalar.constant true : i1 + + %iq4_words0 = vector.load %iq4_code_view[%iq4_group0_word_base] : view<32xi32> -> vector<4xi32> + %iq4_values00, %iq4_values01, %iq4_values02, %iq4_values03 = func.call @ggml_q5_plane_dense_iq4xs_wmmai8_signed_values16(%iq4_words0, %iq4_false) : (vector<4xi32>, i1) -> (vector<4xi8>, vector<4xi8>, vector<4xi8>, vector<4xi8>) + vector.store %iq4_values00, %wl_flat[%iq4_w0] : vector<4xi8>, view<[%weight_elements]xi8> + %iq4_w0_4_0 = index.add %iq4_w0, %c4 : index + %iq4_w0_4 = index.assume %iq4_w0_4_0 [range(%iq4_w0_4_0, 4, 5044)] : index + vector.store %iq4_values01, %wl_flat[%iq4_w0_4] : vector<4xi8>, view<[%weight_elements]xi8> + %iq4_w0_8_0 = index.add %iq4_w0, %c8 : index + %iq4_w0_8 = index.assume %iq4_w0_8_0 [range(%iq4_w0_8_0, 8, 5048)] : index + vector.store %iq4_values02, %wl_flat[%iq4_w0_8] : vector<4xi8>, view<[%weight_elements]xi8> + %iq4_w0_12_0 = index.add %iq4_w0_8, %c4 : index + %iq4_w0_12 = index.assume %iq4_w0_12_0 [range(%iq4_w0_12_0, 12, 5052)] : index + vector.store %iq4_values03, %wl_flat[%iq4_w0_12] : vector<4xi8>, view<[%weight_elements]xi8> + %iq4_values04, %iq4_values05, %iq4_values06, %iq4_values07 = func.call @ggml_q5_plane_dense_iq4xs_wmmai8_signed_values16(%iq4_words0, %iq4_true) : (vector<4xi32>, i1) -> (vector<4xi8>, vector<4xi8>, vector<4xi8>, vector<4xi8>) + vector.store %iq4_values04, %wl_flat[%iq4_w0_hi] : vector<4xi8>, view<[%weight_elements]xi8> + %iq4_w0_hi4_0 = index.add %iq4_w0_hi, %c4 : index + %iq4_w0_hi4 = index.assume %iq4_w0_hi4_0 [range(%iq4_w0_hi4_0, 20, 5060)] : index + vector.store %iq4_values05, %wl_flat[%iq4_w0_hi4] : vector<4xi8>, view<[%weight_elements]xi8> + %iq4_w0_hi8_0 = index.add %iq4_w0_hi, %c8 : index + %iq4_w0_hi8 = index.assume %iq4_w0_hi8_0 [range(%iq4_w0_hi8_0, 24, 5064)] : index + vector.store %iq4_values06, %wl_flat[%iq4_w0_hi8] : vector<4xi8>, view<[%weight_elements]xi8> + %iq4_w0_hi12_0 = index.add %iq4_w0_hi8, %c4 : index + %iq4_w0_hi12 = index.assume %iq4_w0_hi12_0 [range(%iq4_w0_hi12_0, 28, 5068)] : index + vector.store %iq4_values07, %wl_flat[%iq4_w0_hi12] : vector<4xi8>, view<[%weight_elements]xi8> + + %iq4_words1 = vector.load %iq4_code_view[%iq4_group1_word_base] : view<32xi32> -> vector<4xi32> + %iq4_values10, %iq4_values11, %iq4_values12, %iq4_values13 = func.call @ggml_q5_plane_dense_iq4xs_wmmai8_signed_values16(%iq4_words1, %iq4_false) : (vector<4xi32>, i1) -> (vector<4xi8>, vector<4xi8>, vector<4xi8>, vector<4xi8>) + vector.store %iq4_values10, %wl_flat[%iq4_w1] : vector<4xi8>, view<[%weight_elements]xi8> + %iq4_w1_4_0 = index.add %iq4_w1, %c4 : index + %iq4_w1_4 = index.assume %iq4_w1_4_0 [range(%iq4_w1_4_0, 36, 5076)] : index + vector.store %iq4_values11, %wl_flat[%iq4_w1_4] : vector<4xi8>, view<[%weight_elements]xi8> + %iq4_w1_8_0 = index.add %iq4_w1, %c8 : index + %iq4_w1_8 = index.assume %iq4_w1_8_0 [range(%iq4_w1_8_0, 40, 5080)] : index + vector.store %iq4_values12, %wl_flat[%iq4_w1_8] : vector<4xi8>, view<[%weight_elements]xi8> + %iq4_w1_12_0 = index.add %iq4_w1_8, %c4 : index + %iq4_w1_12 = index.assume %iq4_w1_12_0 [range(%iq4_w1_12_0, 44, 5084)] : index + vector.store %iq4_values13, %wl_flat[%iq4_w1_12] : vector<4xi8>, view<[%weight_elements]xi8> + %iq4_values14, %iq4_values15, %iq4_values16, %iq4_values17 = func.call @ggml_q5_plane_dense_iq4xs_wmmai8_signed_values16(%iq4_words1, %iq4_true) : (vector<4xi32>, i1) -> (vector<4xi8>, vector<4xi8>, vector<4xi8>, vector<4xi8>) + vector.store %iq4_values14, %wl_flat[%iq4_w1_hi] : vector<4xi8>, view<[%weight_elements]xi8> + %iq4_w1_hi4_0 = index.add %iq4_w1_hi, %c4 : index + %iq4_w1_hi4 = index.assume %iq4_w1_hi4_0 [range(%iq4_w1_hi4_0, 52, 5092)] : index + vector.store %iq4_values15, %wl_flat[%iq4_w1_hi4] : vector<4xi8>, view<[%weight_elements]xi8> + %iq4_w1_hi8_0 = index.add %iq4_w1_hi, %c8 : index + %iq4_w1_hi8 = index.assume %iq4_w1_hi8_0 [range(%iq4_w1_hi8_0, 56, 5096)] : index + vector.store %iq4_values16, %wl_flat[%iq4_w1_hi8] : vector<4xi8>, view<[%weight_elements]xi8> + %iq4_w1_hi12_0 = index.add %iq4_w1_hi8, %c4 : index + %iq4_w1_hi12 = index.assume %iq4_w1_hi12_0 [range(%iq4_w1_hi12_0, 60, 5100)] : index + vector.store %iq4_values17, %wl_flat[%iq4_w1_hi12] : vector<4xi8>, view<[%weight_elements]xi8> + } else { + %q5_weight_tid0 = index.rem %tid, %q5_weight_lane_count : index + %q5_weight_tid = index.assume %q5_weight_tid0 [range(%q5_weight_tid0, 0, 127)] : index + %q5_block_count = index.div %k_b, %q4_k_block : index + %q5_block = index.div %chunk, %c4 : index + %q5_pair = index.rem %chunk, %c4 : index + %q5_group0 = index.mul %q5_pair, %c2 : index + %q5_group1 = index.add %q5_group0, %c1 : index + %paired_row_lane = index.rem %q5_weight_tid, %c64_wide : index + %weight_row_lane = scf.select %paired_weights, %paired_row_lane, %q5_weight_tid : index + %q5_row0 = index.add %row_base, %weight_row_lane : index + %q5_row = index.assume %q5_row0 [range(%q5_row0, 0, 262143)] : index + %q5_row_bytes = index.scale %q5_block_count, %q5_block_bytes : index, offset -> offset + %q5_row_byte_base = index.scale %q5_row, %q5_row_bytes : index, offset -> offset + %q5_block_byte_add = index.scale %q5_block, %q5_block_bytes : index, offset -> offset + %q5_native_block_byte_base = index.add %q5_row_byte_base, %q5_block_byte_add : offset + %q4_row64 = index.constant 64 : index + %q4_field_bytes = index.constant 16 : offset + %q4_group_bytes = index.constant 9216 : offset + %q4_row_group = index.div %q5_row, %q4_row64 : index + %q4_record = index.madd %q4_row_group, %q5_block_count, %q5_block : index + %q4_record_base = index.scale %q4_record, %q4_group_bytes : index, offset -> offset + %packed_row_lane = index.rem %q5_row, %q4_row64 : index + %q4_row_offset = index.scale %packed_row_lane, %q4_field_bytes : index, offset -> offset + %q4_packed_block_byte_base = index.add %q4_record_base, %q4_row_offset : offset + %q5_block_byte_base = scf.select %is_q4_packed, %q4_packed_block_byte_base, %q5_native_block_byte_base : offset + %q5_high_byte_base = index.add %q5_block_byte_base, %q5_high_offset : offset + %q5_code_byte_base = index.add %q5_block_byte_base, %q5_code_offset : offset + %q5_high_view = buffer.view %src0_na[%q5_high_byte_base] : buffer -> view<8xi32> + %q5_code_view = buffer.view %src0_na[%q5_code_byte_base] : buffer -> view<32xi32> + %q5_d0, %q5_d1, %q5_correction0, %q5_correction1 = scf.if %is_symi5 -> (f32, f32, f32, f32) { + // The capacity-neutral signed-I5 carrier reuses the Q5 bit planes + // and stores one adjacent F16 scale pair for these two K32 groups. + %q5_s5_header_view = buffer.view %src0_na[%q5_block_byte_base] : buffer -> view<16xi8> + %q5_s5_scale_byte0 = index.mul %q5_group0, %c2 : index + %q5_s5_scale_byte = index.assume %q5_s5_scale_byte0 [range(%q5_s5_scale_byte0, 0, 12), mul(%q5_s5_scale_byte0, 4)] : index + %q5_s5_scale_raw = vector.load %q5_s5_header_view[%q5_s5_scale_byte] : view<16xi8> -> vector<4xi8> + %q5_s5_scale_halves = vector.bitcast %q5_s5_scale_raw : vector<4xi8> to vector<2xf16> + %q5_s5_scale0_f16 = vector.extract %q5_s5_scale_halves[0] : vector<2xf16> -> f16 + %q5_s5_scale1_f16 = vector.extract %q5_s5_scale_halves[1] : vector<2xf16> -> f16 + %q5_s5_scale0 = scalar.extf %q5_s5_scale0_f16 : f16 to f32 + %q5_s5_scale1 = scalar.extf %q5_s5_scale1_f16 : f16 to f32 + %q5_s5_zero = scalar.constant 0.0 : f32 + scf.yield %q5_s5_scale0, %q5_s5_scale1, %q5_s5_zero, %q5_s5_zero : f32, f32, f32, f32 + } else { + %q5_header = func.call @ggml_q8_1_x4_paired_weight_words(%paired_weights, %load_up, %src0_na, %peer_weight, %q5_block_byte_base) : (i1, i1, buffer, buffer, offset) -> (vector<4xi32>) + %q5_header_halves = vector.bitcast %q5_header : vector<4xi32> to vector<8xf16> + %q5_d_f16 = vector.extract %q5_header_halves[0] : vector<8xf16> -> f16 + %q5_dmin_f16 = vector.extract %q5_header_halves[1] : vector<8xf16> -> f16 + %q5_scale0 = vector.extract %q5_header[1] : vector<4xi32> -> i32 + %q5_scale1 = vector.extract %q5_header[2] : vector<4xi32> -> i32 + %q5_scale2 = vector.extract %q5_header[3] : vector<4xi32> -> i32 + %q5_d = scalar.extf %q5_d_f16 : f16 to f32 + %q5_dmin = scalar.extf %q5_dmin_f16 : f16 to f32 + %q5_scale_group0, %q5_min_group0 = func.call @ggml_q5_plane_q4k_scale_from_header(%q5_scale0, %q5_scale1, %q5_scale2, %q5_group0) : (i32, i32, i32, index) -> (i32, i32) + %q5_scale_group1, %q5_min_group1 = func.call @ggml_q5_plane_q4k_scale_from_header(%q5_scale0, %q5_scale1, %q5_scale2, %q5_group1) : (i32, i32, i32, index) -> (i32, i32) + %q5_scale_group0_f32 = scalar.uitofp %q5_scale_group0 : i32 to f32 + %q5_scale_group1_f32 = scalar.uitofp %q5_scale_group1 : i32 to f32 + %q5_min_group0_f32 = scalar.uitofp %q5_min_group0 : i32 to f32 + %q5_min_group1_f32 = scalar.uitofp %q5_min_group1 : i32 to f32 + %q5_exact_d0 = scalar.mulf %q5_d, %q5_scale_group0_f32 : f32 + %q5_exact_d1 = scalar.mulf %q5_d, %q5_scale_group1_f32 : f32 + %q5_m0 = scalar.mulf %q5_dmin, %q5_min_group0_f32 : f32 + %q5_m1 = scalar.mulf %q5_dmin, %q5_min_group1_f32 : f32 + %q5_neg_m0 = scalar.negf %q5_m0 : f32 + %q5_neg_m1 = scalar.negf %q5_m1 : f32 + scf.yield %q5_exact_d0, %q5_exact_d1, %q5_neg_m0, %q5_neg_m1 : f32, f32, f32, f32 + } + %q5_meta0_0 = index.mul %q5_weight_tid, %c2 : index + %q5_meta0 = index.assume %q5_meta0_0 [range(%q5_meta0_0, 0, 254)] : index + %q5_meta1_0 = index.add %q5_meta0, %c1 : index + %q5_meta1 = index.assume %q5_meta1_0 [range(%q5_meta1_0, 1, 255)] : index + view.store %q5_d0, %wsl_view[%q5_meta0] : f32, view<[%weight_metadata]xf32> + view.store %q5_d1, %wsl_view[%q5_meta1] : f32, view<[%weight_metadata]xf32> + view.store %q5_correction0, %wcl_view[%q5_meta0] : f32, view<[%weight_metadata]xf32> + view.store %q5_correction1, %wcl_view[%q5_meta1] : f32, view<[%weight_metadata]xf32> + + %q5_pair_word_base = index.mul %q5_pair, %c8 : index + %q5_group0_i32 = index.cast %q5_group0 : index to i32 + %q5_group1_i32 = index.cast %q5_group1 : index to i32 + %q5_group0_shift = vector.splat %q5_group0_i32 : vector<4xi32> + %q5_group1_shift = vector.splat %q5_group1_i32 : vector<4xi32> + %q5_nibble_mask = vector.constant 252645135 : vector<4xi32> + %q5_byte_ones = vector.constant 16843009 : vector<4xi32> + %q5_sign_extend = vector.constant 240 : vector<4xi32> + %q5_shift4 = vector.constant 4 : vector<4xi32> + %q5_weight_row_base = index.mul %q5_weight_tid, %k_wstride : index + %q5_w0 = index.assume %q5_weight_row_base [range(%q5_weight_row_base, 0, 10160)] : index + %q5_payload16 = index.constant 16 : index + %q5_payload48 = index.constant 48 : index + %q5_w0_hi0 = index.add %q5_weight_row_base, %q5_payload16 : index + %q5_w0_hi = index.assume %q5_w0_hi0 [range(%q5_w0_hi0, 16, 10176)] : index + %q5_w1_0 = index.add %q5_weight_row_base, %c32 : index + %q5_w1 = index.assume %q5_w1_0 [range(%q5_w1_0, 32, 10192)] : index + %q5_w1_hi0 = index.add %q5_weight_row_base, %q5_payload48 : index + %q5_w1_hi = index.assume %q5_w1_hi0 [range(%q5_w1_hi0, 48, 10208)] : index + + // First 16-byte payload half. + %q5_code_words0 = scf.if %is_q4_packed -> (vector<4xi32>) { + %field_add = index.constant 1 : index + %field_bytes = index.constant 1024 : offset + %field_pair = index.mul %q5_pair, %c2 : index + %field = index.add %field_pair, %field_add : index + %field_offset = index.scale %field, %field_bytes : index, offset -> offset + %payload_base = index.add %q5_block_byte_base, %field_offset : offset + %words = func.call @ggml_q8_1_x4_paired_weight_words(%paired_weights, %load_up, %src0_na, %peer_weight, %payload_base) : (i1, i1, buffer, buffer, offset) -> (vector<4xi32>) + scf.yield %words : vector<4xi32> + } else { + %words = vector.load %q5_code_view[%q5_pair_word_base] : view<32xi32> -> vector<4xi32> + scf.yield %words : vector<4xi32> + } + %q5_high_words0 = scf.if %is_q4 -> (vector<4xi32>) { + %zero = vector.constant 0 : vector<4xi32> + scf.yield %zero : vector<4xi32> + } else { + %high = vector.load %q5_high_view[%c0] : view<8xi32> -> vector<4xi32> + scf.yield %high : vector<4xi32> + } + %q5_low_codes0 = vector.andi %q5_code_words0, %q5_nibble_mask : vector<4xi32> + %q5_fifth00_shifted = vector.shrui %q5_high_words0, %q5_group0_shift : vector<4xi32> + %q5_fifth00_bits = vector.andi %q5_fifth00_shifted, %q5_byte_ones : vector<4xi32> + %q5_fifth00_unsigned = vector.shli %q5_fifth00_bits, %q5_shift4 : vector<4xi32> + %q5_fifth00_signed = vector.muli %q5_fifth00_bits, %q5_sign_extend : vector<4xi32> + %q5_fifth00 = scf.select %is_symi5, %q5_fifth00_signed, %q5_fifth00_unsigned : vector<4xi32> + %q5_values00_words = vector.ori %q5_low_codes0, %q5_fifth00 : vector<4xi32> + %q5_values00 = vector.bitcast %q5_values00_words : vector<4xi32> to vector<16xi8> + vector.store %q5_values00, %wl_flat[%q5_w0] : vector<16xi8>, view<[%weight_elements]xi8> + %q5_high_codes00 = vector.shrui %q5_code_words0, %q5_shift4 : vector<4xi32> + %q5_high_codes0 = vector.andi %q5_high_codes00, %q5_nibble_mask : vector<4xi32> + %q5_fifth10_shifted = vector.shrui %q5_high_words0, %q5_group1_shift : vector<4xi32> + %q5_fifth10_bits = vector.andi %q5_fifth10_shifted, %q5_byte_ones : vector<4xi32> + %q5_fifth10_unsigned = vector.shli %q5_fifth10_bits, %q5_shift4 : vector<4xi32> + %q5_fifth10_signed = vector.muli %q5_fifth10_bits, %q5_sign_extend : vector<4xi32> + %q5_fifth10 = scf.select %is_symi5, %q5_fifth10_signed, %q5_fifth10_unsigned : vector<4xi32> + %q5_values10_words = vector.ori %q5_high_codes0, %q5_fifth10 : vector<4xi32> + %q5_values10 = vector.bitcast %q5_values10_words : vector<4xi32> to vector<16xi8> + vector.store %q5_values10, %wl_flat[%q5_w1] : vector<16xi8>, view<[%weight_elements]xi8> + + // Second 16-byte payload half. Recompute its addresses after the first + // stores so only one half's decode payload remains live at a time. + %q5_pair_word_hi0 = index.add %q5_pair_word_base, %c4 : index + %q5_pair_word_hi = index.assume %q5_pair_word_hi0 [range(%q5_pair_word_hi0, 4, 28)] : index + %q5_code_words1 = scf.if %is_q4_packed -> (vector<4xi32>) { + %field_add = index.constant 2 : index + %field_bytes = index.constant 1024 : offset + %field_pair = index.mul %q5_pair, %c2 : index + %field = index.add %field_pair, %field_add : index + %field_offset = index.scale %field, %field_bytes : index, offset -> offset + %payload_base = index.add %q5_block_byte_base, %field_offset : offset + %words = func.call @ggml_q8_1_x4_paired_weight_words(%paired_weights, %load_up, %src0_na, %peer_weight, %payload_base) : (i1, i1, buffer, buffer, offset) -> (vector<4xi32>) + scf.yield %words : vector<4xi32> + } else { + %words = vector.load %q5_code_view[%q5_pair_word_hi] : view<32xi32> -> vector<4xi32> + scf.yield %words : vector<4xi32> + } + %q5_high_words1 = scf.if %is_q4 -> (vector<4xi32>) { + %zero = vector.constant 0 : vector<4xi32> + scf.yield %zero : vector<4xi32> + } else { + %high = vector.load %q5_high_view[%c4] : view<8xi32> -> vector<4xi32> + scf.yield %high : vector<4xi32> + } + %q5_low_codes1 = vector.andi %q5_code_words1, %q5_nibble_mask : vector<4xi32> + %q5_fifth01_shifted = vector.shrui %q5_high_words1, %q5_group0_shift : vector<4xi32> + %q5_fifth01_bits = vector.andi %q5_fifth01_shifted, %q5_byte_ones : vector<4xi32> + %q5_fifth01_unsigned = vector.shli %q5_fifth01_bits, %q5_shift4 : vector<4xi32> + %q5_fifth01_signed = vector.muli %q5_fifth01_bits, %q5_sign_extend : vector<4xi32> + %q5_fifth01 = scf.select %is_symi5, %q5_fifth01_signed, %q5_fifth01_unsigned : vector<4xi32> + %q5_values01_words = vector.ori %q5_low_codes1, %q5_fifth01 : vector<4xi32> + %q5_values01 = vector.bitcast %q5_values01_words : vector<4xi32> to vector<16xi8> + vector.store %q5_values01, %wl_flat[%q5_w0_hi] : vector<16xi8>, view<[%weight_elements]xi8> + %q5_high_codes10 = vector.shrui %q5_code_words1, %q5_shift4 : vector<4xi32> + %q5_high_codes1 = vector.andi %q5_high_codes10, %q5_nibble_mask : vector<4xi32> + %q5_fifth11_shifted = vector.shrui %q5_high_words1, %q5_group1_shift : vector<4xi32> + %q5_fifth11_bits = vector.andi %q5_fifth11_shifted, %q5_byte_ones : vector<4xi32> + %q5_fifth11_unsigned = vector.shli %q5_fifth11_bits, %q5_shift4 : vector<4xi32> + %q5_fifth11_signed = vector.muli %q5_fifth11_bits, %q5_sign_extend : vector<4xi32> + %q5_fifth11 = scf.select %is_symi5, %q5_fifth11_signed, %q5_fifth11_unsigned : vector<4xi32> + %q5_values11_words = vector.ori %q5_high_codes1, %q5_fifth11 : vector<4xi32> + %q5_values11 = vector.bitcast %q5_values11_words : vector<4xi32> to vector<16xi8> + vector.store %q5_values11, %wl_flat[%q5_w1_hi] : vector<16xi8>, view<[%weight_elements]xi8> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + + %step0_0, %step0_1, %step0_2, %step0_3, %step1_0, %step1_1, %step1_2, %step1_3 = scf.if %paired_weights -> (vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) { + %bk0 = index.mul %c0, %c32 : index + %bkh0 = index.add %bk0, %c16 : index + %iz0 = vector.fragment %izero shape [%c16, %c16] : vector<8xi32> + %lhs0_0_raw = func.call @ggml_q5_plane_dense_q8_wmmai8_load_lhs(%token256, %scratch, %base, %lm0, %bk0, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %lhs0_0_words = vector.bitcast %lhs0_0_raw : vector<16xi8> to vector<4xi32> + %lhs0_0 = vector.fragment %lhs0_0_words shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %lhs0_1_raw = func.call @ggml_q5_plane_dense_q8_wmmai8_load_lhs(%token256, %scratch, %base, %lm0, %bkh0, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %lhs0_1_words = vector.bitcast %lhs0_1_raw : vector<16xi8> to vector<4xi32> + %lhs0_1 = vector.fragment %lhs0_1_words shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %asi0_0 = index.add %asb0, %c0 : index + %asi0_1 = index.mul %asi0_0, %c2 : index + %asi0_2 = index.add %asi0_1, %lane_hi : index + %asi0_3 = index.mul %asi0_2, %c8 : index + %asv0, %auv0 = func.call @ggml_q5_plane_dense_q8_wmmai8_load_metadata(%token256, %scratch, %as_off, %asum_off, %asi0_3) : (i1, buffer, offset, offset, index) -> (vector<8xf32>, vector<8xf32>) + %n0_0 = index.constant 0 : index + %b0n0 = func.call @ggml_q8_1_x4_affine_rhs_tile(%scratch, %w_off, %ws_off, %wc_off, %c0, %n0_0, %lane_lo, %lhs0_0, %lhs0_1, %iz0, %asv0, %auv0, %fc0_0) : (buffer, offset, offset, offset, index, index, index, vector<4xi32>, vector<4xi32>, vector<8xi32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + scf.schedule.fence + %n0_1 = index.constant 16 : index + %b0n1 = func.call @ggml_q8_1_x4_affine_rhs_tile(%scratch, %w_off, %ws_off, %wc_off, %c0, %n0_1, %lane_lo, %lhs0_0, %lhs0_1, %iz0, %asv0, %auv0, %fc0_1) : (buffer, offset, offset, offset, index, index, index, vector<4xi32>, vector<4xi32>, vector<8xi32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + scf.schedule.fence + %n0_2 = index.constant 32 : index + %b0n2 = func.call @ggml_q8_1_x4_affine_rhs_tile(%scratch, %w_off, %ws_off, %wc_off, %c0, %n0_2, %lane_lo, %lhs0_0, %lhs0_1, %iz0, %asv0, %auv0, %fc0_2) : (buffer, offset, offset, offset, index, index, index, vector<4xi32>, vector<4xi32>, vector<8xi32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + scf.schedule.fence + %n0_3 = index.constant 48 : index + %b0n3 = func.call @ggml_q8_1_x4_affine_rhs_tile(%scratch, %w_off, %ws_off, %wc_off, %c0, %n0_3, %lane_lo, %lhs0_0, %lhs0_1, %iz0, %asv0, %auv0, %fc0_3) : (buffer, offset, offset, offset, index, index, index, vector<4xi32>, vector<4xi32>, vector<8xi32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + scf.schedule.fence + %n0_4 = index.constant 64 : index + %b0n4 = func.call @ggml_q8_1_x4_affine_rhs_tile(%scratch, %w_off, %ws_off, %wc_off, %c0, %n0_4, %lane_lo, %lhs0_0, %lhs0_1, %iz0, %asv0, %auv0, %fc1_0) : (buffer, offset, offset, offset, index, index, index, vector<4xi32>, vector<4xi32>, vector<8xi32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + scf.schedule.fence + %n0_5 = index.constant 80 : index + %b0n5 = func.call @ggml_q8_1_x4_affine_rhs_tile(%scratch, %w_off, %ws_off, %wc_off, %c0, %n0_5, %lane_lo, %lhs0_0, %lhs0_1, %iz0, %asv0, %auv0, %fc1_1) : (buffer, offset, offset, offset, index, index, index, vector<4xi32>, vector<4xi32>, vector<8xi32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + scf.schedule.fence + %n0_6 = index.constant 96 : index + %b0n6 = func.call @ggml_q8_1_x4_affine_rhs_tile(%scratch, %w_off, %ws_off, %wc_off, %c0, %n0_6, %lane_lo, %lhs0_0, %lhs0_1, %iz0, %asv0, %auv0, %fc1_2) : (buffer, offset, offset, offset, index, index, index, vector<4xi32>, vector<4xi32>, vector<8xi32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + scf.schedule.fence + %n0_7 = index.constant 112 : index + %b0n7 = func.call @ggml_q8_1_x4_affine_rhs_tile(%scratch, %w_off, %ws_off, %wc_off, %c0, %n0_7, %lane_lo, %lhs0_0, %lhs0_1, %iz0, %asv0, %auv0, %fc1_3) : (buffer, offset, offset, offset, index, index, index, vector<4xi32>, vector<4xi32>, vector<8xi32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + scf.schedule.fence + %bk1 = index.mul %c1, %c32 : index + %bkh1 = index.add %bk1, %c16 : index + %iz1 = vector.fragment %izero shape [%c16, %c16] : vector<8xi32> + %lhs1_0_raw = func.call @ggml_q5_plane_dense_q8_wmmai8_load_lhs(%token256, %scratch, %base, %lm0, %bk1, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %lhs1_0_words = vector.bitcast %lhs1_0_raw : vector<16xi8> to vector<4xi32> + %lhs1_0 = vector.fragment %lhs1_0_words shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %lhs1_1_raw = func.call @ggml_q5_plane_dense_q8_wmmai8_load_lhs(%token256, %scratch, %base, %lm0, %bkh1, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %lhs1_1_words = vector.bitcast %lhs1_1_raw : vector<16xi8> to vector<4xi32> + %lhs1_1 = vector.fragment %lhs1_1_words shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %asi1_0 = index.add %asb0, %c1 : index + %asi1_1 = index.mul %asi1_0, %c2 : index + %asi1_2 = index.add %asi1_1, %lane_hi : index + %asi1_3 = index.mul %asi1_2, %c8 : index + %asv1, %auv1 = func.call @ggml_q5_plane_dense_q8_wmmai8_load_metadata(%token256, %scratch, %as_off, %asum_off, %asi1_3) : (i1, buffer, offset, offset, index) -> (vector<8xf32>, vector<8xf32>) + %n1_0 = index.constant 0 : index + %b1n0 = func.call @ggml_q8_1_x4_affine_rhs_tile(%scratch, %w_off, %ws_off, %wc_off, %c1, %n1_0, %lane_lo, %lhs1_0, %lhs1_1, %iz1, %asv1, %auv1, %b0n0) : (buffer, offset, offset, offset, index, index, index, vector<4xi32>, vector<4xi32>, vector<8xi32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + scf.schedule.fence + %n1_1 = index.constant 16 : index + %b1n1 = func.call @ggml_q8_1_x4_affine_rhs_tile(%scratch, %w_off, %ws_off, %wc_off, %c1, %n1_1, %lane_lo, %lhs1_0, %lhs1_1, %iz1, %asv1, %auv1, %b0n1) : (buffer, offset, offset, offset, index, index, index, vector<4xi32>, vector<4xi32>, vector<8xi32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + scf.schedule.fence + %n1_2 = index.constant 32 : index + %b1n2 = func.call @ggml_q8_1_x4_affine_rhs_tile(%scratch, %w_off, %ws_off, %wc_off, %c1, %n1_2, %lane_lo, %lhs1_0, %lhs1_1, %iz1, %asv1, %auv1, %b0n2) : (buffer, offset, offset, offset, index, index, index, vector<4xi32>, vector<4xi32>, vector<8xi32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + scf.schedule.fence + %n1_3 = index.constant 48 : index + %b1n3 = func.call @ggml_q8_1_x4_affine_rhs_tile(%scratch, %w_off, %ws_off, %wc_off, %c1, %n1_3, %lane_lo, %lhs1_0, %lhs1_1, %iz1, %asv1, %auv1, %b0n3) : (buffer, offset, offset, offset, index, index, index, vector<4xi32>, vector<4xi32>, vector<8xi32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + scf.schedule.fence + %n1_4 = index.constant 64 : index + %b1n4 = func.call @ggml_q8_1_x4_affine_rhs_tile(%scratch, %w_off, %ws_off, %wc_off, %c1, %n1_4, %lane_lo, %lhs1_0, %lhs1_1, %iz1, %asv1, %auv1, %b0n4) : (buffer, offset, offset, offset, index, index, index, vector<4xi32>, vector<4xi32>, vector<8xi32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + scf.schedule.fence + %n1_5 = index.constant 80 : index + %b1n5 = func.call @ggml_q8_1_x4_affine_rhs_tile(%scratch, %w_off, %ws_off, %wc_off, %c1, %n1_5, %lane_lo, %lhs1_0, %lhs1_1, %iz1, %asv1, %auv1, %b0n5) : (buffer, offset, offset, offset, index, index, index, vector<4xi32>, vector<4xi32>, vector<8xi32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + scf.schedule.fence + %n1_6 = index.constant 96 : index + %b1n6 = func.call @ggml_q8_1_x4_affine_rhs_tile(%scratch, %w_off, %ws_off, %wc_off, %c1, %n1_6, %lane_lo, %lhs1_0, %lhs1_1, %iz1, %asv1, %auv1, %b0n6) : (buffer, offset, offset, offset, index, index, index, vector<4xi32>, vector<4xi32>, vector<8xi32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + scf.schedule.fence + %n1_7 = index.constant 112 : index + %b1n7 = func.call @ggml_q8_1_x4_affine_rhs_tile(%scratch, %w_off, %ws_off, %wc_off, %c1, %n1_7, %lane_lo, %lhs1_0, %lhs1_1, %iz1, %asv1, %auv1, %b0n7) : (buffer, offset, offset, offset, index, index, index, vector<4xi32>, vector<4xi32>, vector<8xi32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) -> (vector<8xf32>) + scf.schedule.fence + scf.yield %b1n0, %b1n1, %b1n2, %b1n3, %b1n4, %b1n5, %b1n6, %b1n7 : vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32> + } else { + // Two K32 blocks explicitly unrolled. Arithmetic order and the + // enclosing K64 staging/barrier cadence are unchanged. + %blk_k_u0 = index.mul %c0, %c32 : index + %blk_k16_u0 = index.add %blk_k_u0, %c16 : index + %iz_u0 = vector.fragment %izero shape [%c16, %c16] : vector<8xi32> + %wsi0_2_u0 = index.add %wsb0, %c0 : index + %wsi0_u0 = index.assume %wsi0_2_u0 [range(%wsi0_2_u0, 0, 127)] : index + %wsc0_u0 = view.load %wsl_view[%wsi0_u0] : view<[%weight_metadata]xf32> -> f32 + %wsv0_u0 = vector.splat %wsc0_u0 : vector<8xf32> + %wc0_u0 = view.load %wcl_view[%wsi0_u0] : view<[%weight_metadata]xf32> -> f32 + %wcv0_u0 = vector.splat %wc0_u0 : vector<8xf32> + %wsi1_2_u0 = index.add %wsb1, %c0 : index + %wsi1_u0 = index.assume %wsi1_2_u0 [range(%wsi1_2_u0, 0, 127)] : index + %wsc1_u0 = view.load %wsl_view[%wsi1_u0] : view<[%weight_metadata]xf32> -> f32 + %wsv1_u0 = vector.splat %wsc1_u0 : vector<8xf32> + %wc1_u0 = view.load %wcl_view[%wsi1_u0] : view<[%weight_metadata]xf32> -> f32 + %wcv1_u0 = vector.splat %wc1_u0 : vector<8xf32> + %wsi2_2_u0 = index.add %wsb2, %c0 : index + %wsi2_u0 = index.assume %wsi2_2_u0 [range(%wsi2_2_u0, 0, 127)] : index + %wsc2_u0 = view.load %wsl_view[%wsi2_u0] : view<[%weight_metadata]xf32> -> f32 + %wsv2_u0 = vector.splat %wsc2_u0 : vector<8xf32> + %wc2_u0 = view.load %wcl_view[%wsi2_u0] : view<[%weight_metadata]xf32> -> f32 + %wcv2_u0 = vector.splat %wc2_u0 : vector<8xf32> + %wsi3_2_u0 = index.add %wsb3, %c0 : index + %wsi3_u0 = index.assume %wsi3_2_u0 [range(%wsi3_2_u0, 0, 127)] : index + %wsc3_u0 = view.load %wsl_view[%wsi3_u0] : view<[%weight_metadata]xf32> -> f32 + %wsv3_u0 = vector.splat %wsc3_u0 : vector<8xf32> + %wc3_u0 = view.load %wcl_view[%wsi3_u0] : view<[%weight_metadata]xf32> -> f32 + %wcv3_u0 = vector.splat %wc3_u0 : vector<8xf32> + %lf0_0_raw_u0 = func.call @ggml_q5_plane_dense_q8_wmmai8_load_lhs(%token256, %scratch, %base, %lm0, %blk_k_u0, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %lf0_0_words_u0 = vector.bitcast %lf0_0_raw_u0 : vector<16xi8> to vector<4xi32> + %lf0_0_u0 = vector.fragment %lf0_0_words_u0 shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %lf0_1_raw_u0 = func.call @ggml_q5_plane_dense_q8_wmmai8_load_lhs(%token256, %scratch, %base, %lm0, %blk_k16_u0, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %lf0_1_words_u0 = vector.bitcast %lf0_1_raw_u0 : vector<16xi8> to vector<4xi32> + %lf0_1_u0 = vector.fragment %lf0_1_words_u0 shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %asi0_0_u0 = index.add %asb0, %c0 : index + %asi0_1_u0 = index.mul %asi0_0_u0, %c2 : index + %asi0_2_u0 = index.add %asi0_1_u0, %lane_hi : index + %asi0_3_u0 = index.mul %asi0_2_u0, %c8 : index + %asv0_u0, %auv0_u0 = func.call @ggml_q5_plane_dense_q8_wmmai8_load_metadata(%token256, %scratch, %as_off, %asum_off, %asi0_3_u0) : (i1, buffer, offset, offset, index) -> (vector<8xf32>, vector<8xf32>) + // Issue independent RHS fragments before consuming the first pair. + %rfs0_0_0_raw_u0 = vector.fragment.load %wl_view[%blk_k_u0, %ln0] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs0_0_0_words_u0 = vector.bitcast %rfs0_0_0_raw_u0 : vector<16xi8> to vector<4xi32> + %rfs0_0_0_u0 = vector.fragment %rfs0_0_0_words_u0 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs0_0_1_raw_u0 = vector.fragment.load %wl_view[%blk_k16_u0, %ln0] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs0_0_1_words_u0 = vector.bitcast %rfs0_0_1_raw_u0 : vector<16xi8> to vector<4xi32> + %rfs0_0_1_u0 = vector.fragment %rfs0_0_1_words_u0 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs0_1_0_raw_u0 = vector.fragment.load %wl_view[%blk_k_u0, %ln1] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs0_1_0_words_u0 = vector.bitcast %rfs0_1_0_raw_u0 : vector<16xi8> to vector<4xi32> + %rfs0_1_0_u0 = vector.fragment %rfs0_1_0_words_u0 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs0_1_1_raw_u0 = vector.fragment.load %wl_view[%blk_k16_u0, %ln1] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs0_1_1_words_u0 = vector.bitcast %rfs0_1_1_raw_u0 : vector<16xi8> to vector<4xi32> + %rfs0_1_1_u0 = vector.fragment %rfs0_1_1_words_u0 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs0_2_0_raw_u0 = vector.fragment.load %wl_view[%blk_k_u0, %ln2] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs0_2_0_words_u0 = vector.bitcast %rfs0_2_0_raw_u0 : vector<16xi8> to vector<4xi32> + %rfs0_2_0_u0 = vector.fragment %rfs0_2_0_words_u0 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs0_2_1_raw_u0 = vector.fragment.load %wl_view[%blk_k16_u0, %ln2] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs0_2_1_words_u0 = vector.bitcast %rfs0_2_1_raw_u0 : vector<16xi8> to vector<4xi32> + %rfs0_2_1_u0 = vector.fragment %rfs0_2_1_words_u0 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs0_3_0_raw_u0 = vector.fragment.load %wl_view[%blk_k_u0, %ln3] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs0_3_0_words_u0 = vector.bitcast %rfs0_3_0_raw_u0 : vector<16xi8> to vector<4xi32> + %rfs0_3_0_u0 = vector.fragment %rfs0_3_0_words_u0 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs0_3_1_raw_u0 = vector.fragment.load %wl_view[%blk_k16_u0, %ln3] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs0_3_1_words_u0 = vector.bitcast %rfs0_3_1_raw_u0 : vector<16xi8> to vector<4xi32> + %rfs0_3_1_u0 = vector.fragment %rfs0_3_1_words_u0 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %i0_0_0_u0 = vector.mma %lf0_0_u0, %rfs0_0_0_u0, %iz_u0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i1_0_0_u0 = vector.mma %lf0_1_u0, %rfs0_0_1_u0, %i0_0_0_u0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %ff0_0_u0 = vector.sitofp %i1_0_0_u0 : vector<8xi32> to vector<8xf32> + %sv0_0_u0 = vector.mulf %asv0_u0, %wsv0_u0 : vector<8xf32> + %base_fn0_0_u0 = vector.fmaf %auv0_u0, %wcv0_u0, %fc0_0 : vector<8xf32> + %selected_base_fn0_0_u0 = scf.select %is_signed_rhs, %fc0_0, %base_fn0_0_u0 : vector<8xf32> + + %fn0_0_u0 = vector.fmaf %ff0_0_u0, %sv0_0_u0, %selected_base_fn0_0_u0 : vector<8xf32> + scf.schedule.fence + %i0_0_1_u0 = vector.mma %lf0_0_u0, %rfs0_1_0_u0, %iz_u0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i1_0_1_u0 = vector.mma %lf0_1_u0, %rfs0_1_1_u0, %i0_0_1_u0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %ff0_1_u0 = vector.sitofp %i1_0_1_u0 : vector<8xi32> to vector<8xf32> + %sv0_1_u0 = vector.mulf %asv0_u0, %wsv1_u0 : vector<8xf32> + %base_fn0_1_u0 = vector.fmaf %auv0_u0, %wcv1_u0, %fc0_1 : vector<8xf32> + %selected_base_fn0_1_u0 = scf.select %is_signed_rhs, %fc0_1, %base_fn0_1_u0 : vector<8xf32> + + %fn0_1_u0 = vector.fmaf %ff0_1_u0, %sv0_1_u0, %selected_base_fn0_1_u0 : vector<8xf32> + scf.schedule.fence + %i0_0_2_u0 = vector.mma %lf0_0_u0, %rfs0_2_0_u0, %iz_u0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i1_0_2_u0 = vector.mma %lf0_1_u0, %rfs0_2_1_u0, %i0_0_2_u0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %ff0_2_u0 = vector.sitofp %i1_0_2_u0 : vector<8xi32> to vector<8xf32> + %sv0_2_u0 = vector.mulf %asv0_u0, %wsv2_u0 : vector<8xf32> + %base_fn0_2_u0 = vector.fmaf %auv0_u0, %wcv2_u0, %fc0_2 : vector<8xf32> + %selected_base_fn0_2_u0 = scf.select %is_signed_rhs, %fc0_2, %base_fn0_2_u0 : vector<8xf32> + + %fn0_2_u0 = vector.fmaf %ff0_2_u0, %sv0_2_u0, %selected_base_fn0_2_u0 : vector<8xf32> + scf.schedule.fence + %i0_0_3_u0 = vector.mma %lf0_0_u0, %rfs0_3_0_u0, %iz_u0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i1_0_3_u0 = vector.mma %lf0_1_u0, %rfs0_3_1_u0, %i0_0_3_u0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %ff0_3_u0 = vector.sitofp %i1_0_3_u0 : vector<8xi32> to vector<8xf32> + %sv0_3_u0 = vector.mulf %asv0_u0, %wsv3_u0 : vector<8xf32> + %base_fn0_3_u0 = vector.fmaf %auv0_u0, %wcv3_u0, %fc0_3 : vector<8xf32> + %selected_base_fn0_3_u0 = scf.select %is_signed_rhs, %fc0_3, %base_fn0_3_u0 : vector<8xf32> + + %fn0_3_u0 = vector.fmaf %ff0_3_u0, %sv0_3_u0, %selected_base_fn0_3_u0 : vector<8xf32> + scf.schedule.fence + %lf1_0_raw_u0 = func.call @ggml_q5_plane_dense_q8_wmmai8_load_lhs(%token256, %scratch, %base, %lm1, %blk_k_u0, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %lf1_0_words_u0 = vector.bitcast %lf1_0_raw_u0 : vector<16xi8> to vector<4xi32> + %lf1_0_u0 = vector.fragment %lf1_0_words_u0 shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %lf1_1_raw_u0 = func.call @ggml_q5_plane_dense_q8_wmmai8_load_lhs(%token256, %scratch, %base, %lm1, %blk_k16_u0, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %lf1_1_words_u0 = vector.bitcast %lf1_1_raw_u0 : vector<16xi8> to vector<4xi32> + %lf1_1_u0 = vector.fragment %lf1_1_words_u0 shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %asi1_0_u0 = index.add %asb1, %c0 : index + %asi1_1_u0 = index.mul %asi1_0_u0, %c2 : index + %asi1_2_u0 = index.add %asi1_1_u0, %lane_hi : index + %asi1_3_u0 = index.mul %asi1_2_u0, %c8 : index + %asv1_u0, %auv1_u0 = func.call @ggml_q5_plane_dense_q8_wmmai8_load_metadata(%token256, %scratch, %as_off, %asum_off, %asi1_3_u0) : (i1, buffer, offset, offset, index) -> (vector<8xf32>, vector<8xf32>) + // Issue independent RHS fragments before consuming the first pair. + %rfs1_0_0_raw_u0 = vector.fragment.load %wl_view[%blk_k_u0, %ln0] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs1_0_0_words_u0 = vector.bitcast %rfs1_0_0_raw_u0 : vector<16xi8> to vector<4xi32> + %rfs1_0_0_u0 = vector.fragment %rfs1_0_0_words_u0 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs1_0_1_raw_u0 = vector.fragment.load %wl_view[%blk_k16_u0, %ln0] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs1_0_1_words_u0 = vector.bitcast %rfs1_0_1_raw_u0 : vector<16xi8> to vector<4xi32> + %rfs1_0_1_u0 = vector.fragment %rfs1_0_1_words_u0 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs1_1_0_raw_u0 = vector.fragment.load %wl_view[%blk_k_u0, %ln1] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs1_1_0_words_u0 = vector.bitcast %rfs1_1_0_raw_u0 : vector<16xi8> to vector<4xi32> + %rfs1_1_0_u0 = vector.fragment %rfs1_1_0_words_u0 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs1_1_1_raw_u0 = vector.fragment.load %wl_view[%blk_k16_u0, %ln1] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs1_1_1_words_u0 = vector.bitcast %rfs1_1_1_raw_u0 : vector<16xi8> to vector<4xi32> + %rfs1_1_1_u0 = vector.fragment %rfs1_1_1_words_u0 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs1_2_0_raw_u0 = vector.fragment.load %wl_view[%blk_k_u0, %ln2] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs1_2_0_words_u0 = vector.bitcast %rfs1_2_0_raw_u0 : vector<16xi8> to vector<4xi32> + %rfs1_2_0_u0 = vector.fragment %rfs1_2_0_words_u0 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs1_2_1_raw_u0 = vector.fragment.load %wl_view[%blk_k16_u0, %ln2] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs1_2_1_words_u0 = vector.bitcast %rfs1_2_1_raw_u0 : vector<16xi8> to vector<4xi32> + %rfs1_2_1_u0 = vector.fragment %rfs1_2_1_words_u0 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs1_3_0_raw_u0 = vector.fragment.load %wl_view[%blk_k_u0, %ln3] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs1_3_0_words_u0 = vector.bitcast %rfs1_3_0_raw_u0 : vector<16xi8> to vector<4xi32> + %rfs1_3_0_u0 = vector.fragment %rfs1_3_0_words_u0 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs1_3_1_raw_u0 = vector.fragment.load %wl_view[%blk_k16_u0, %ln3] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs1_3_1_words_u0 = vector.bitcast %rfs1_3_1_raw_u0 : vector<16xi8> to vector<4xi32> + %rfs1_3_1_u0 = vector.fragment %rfs1_3_1_words_u0 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %i0_1_0_u0 = vector.mma %lf1_0_u0, %rfs1_0_0_u0, %iz_u0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i1_1_0_u0 = vector.mma %lf1_1_u0, %rfs1_0_1_u0, %i0_1_0_u0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %ff1_0_u0 = vector.sitofp %i1_1_0_u0 : vector<8xi32> to vector<8xf32> + %sv1_0_u0 = vector.mulf %asv1_u0, %wsv0_u0 : vector<8xf32> + %base_fn1_0_u0 = vector.fmaf %auv1_u0, %wcv0_u0, %fc1_0 : vector<8xf32> + %selected_base_fn1_0_u0 = scf.select %is_signed_rhs, %fc1_0, %base_fn1_0_u0 : vector<8xf32> + + %fn1_0_u0 = vector.fmaf %ff1_0_u0, %sv1_0_u0, %selected_base_fn1_0_u0 : vector<8xf32> + scf.schedule.fence + %i0_1_1_u0 = vector.mma %lf1_0_u0, %rfs1_1_0_u0, %iz_u0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i1_1_1_u0 = vector.mma %lf1_1_u0, %rfs1_1_1_u0, %i0_1_1_u0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %ff1_1_u0 = vector.sitofp %i1_1_1_u0 : vector<8xi32> to vector<8xf32> + %sv1_1_u0 = vector.mulf %asv1_u0, %wsv1_u0 : vector<8xf32> + %base_fn1_1_u0 = vector.fmaf %auv1_u0, %wcv1_u0, %fc1_1 : vector<8xf32> + %selected_base_fn1_1_u0 = scf.select %is_signed_rhs, %fc1_1, %base_fn1_1_u0 : vector<8xf32> + + %fn1_1_u0 = vector.fmaf %ff1_1_u0, %sv1_1_u0, %selected_base_fn1_1_u0 : vector<8xf32> + scf.schedule.fence + %i0_1_2_u0 = vector.mma %lf1_0_u0, %rfs1_2_0_u0, %iz_u0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i1_1_2_u0 = vector.mma %lf1_1_u0, %rfs1_2_1_u0, %i0_1_2_u0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %ff1_2_u0 = vector.sitofp %i1_1_2_u0 : vector<8xi32> to vector<8xf32> + %sv1_2_u0 = vector.mulf %asv1_u0, %wsv2_u0 : vector<8xf32> + %base_fn1_2_u0 = vector.fmaf %auv1_u0, %wcv2_u0, %fc1_2 : vector<8xf32> + %selected_base_fn1_2_u0 = scf.select %is_signed_rhs, %fc1_2, %base_fn1_2_u0 : vector<8xf32> + + %fn1_2_u0 = vector.fmaf %ff1_2_u0, %sv1_2_u0, %selected_base_fn1_2_u0 : vector<8xf32> + scf.schedule.fence + %i0_1_3_u0 = vector.mma %lf1_0_u0, %rfs1_3_0_u0, %iz_u0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i1_1_3_u0 = vector.mma %lf1_1_u0, %rfs1_3_1_u0, %i0_1_3_u0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %ff1_3_u0 = vector.sitofp %i1_1_3_u0 : vector<8xi32> to vector<8xf32> + %sv1_3_u0 = vector.mulf %asv1_u0, %wsv3_u0 : vector<8xf32> + %base_fn1_3_u0 = vector.fmaf %auv1_u0, %wcv3_u0, %fc1_3 : vector<8xf32> + %selected_base_fn1_3_u0 = scf.select %is_signed_rhs, %fc1_3, %base_fn1_3_u0 : vector<8xf32> + + %fn1_3_u0 = vector.fmaf %ff1_3_u0, %sv1_3_u0, %selected_base_fn1_3_u0 : vector<8xf32> + scf.schedule.fence + %blk_k_u1 = index.mul %c1, %c32 : index + %blk_k16_u1 = index.add %blk_k_u1, %c16 : index + %iz_u1 = vector.fragment %izero shape [%c16, %c16] : vector<8xi32> + %wsi0_2_u1 = index.add %wsb0, %c1 : index + %wsi0_u1 = index.assume %wsi0_2_u1 [range(%wsi0_2_u1, 0, 127)] : index + %wsc0_u1 = view.load %wsl_view[%wsi0_u1] : view<[%weight_metadata]xf32> -> f32 + %wsv0_u1 = vector.splat %wsc0_u1 : vector<8xf32> + %wc0_u1 = view.load %wcl_view[%wsi0_u1] : view<[%weight_metadata]xf32> -> f32 + %wcv0_u1 = vector.splat %wc0_u1 : vector<8xf32> + %wsi1_2_u1 = index.add %wsb1, %c1 : index + %wsi1_u1 = index.assume %wsi1_2_u1 [range(%wsi1_2_u1, 0, 127)] : index + %wsc1_u1 = view.load %wsl_view[%wsi1_u1] : view<[%weight_metadata]xf32> -> f32 + %wsv1_u1 = vector.splat %wsc1_u1 : vector<8xf32> + %wc1_u1 = view.load %wcl_view[%wsi1_u1] : view<[%weight_metadata]xf32> -> f32 + %wcv1_u1 = vector.splat %wc1_u1 : vector<8xf32> + %wsi2_2_u1 = index.add %wsb2, %c1 : index + %wsi2_u1 = index.assume %wsi2_2_u1 [range(%wsi2_2_u1, 0, 127)] : index + %wsc2_u1 = view.load %wsl_view[%wsi2_u1] : view<[%weight_metadata]xf32> -> f32 + %wsv2_u1 = vector.splat %wsc2_u1 : vector<8xf32> + %wc2_u1 = view.load %wcl_view[%wsi2_u1] : view<[%weight_metadata]xf32> -> f32 + %wcv2_u1 = vector.splat %wc2_u1 : vector<8xf32> + %wsi3_2_u1 = index.add %wsb3, %c1 : index + %wsi3_u1 = index.assume %wsi3_2_u1 [range(%wsi3_2_u1, 0, 127)] : index + %wsc3_u1 = view.load %wsl_view[%wsi3_u1] : view<[%weight_metadata]xf32> -> f32 + %wsv3_u1 = vector.splat %wsc3_u1 : vector<8xf32> + %wc3_u1 = view.load %wcl_view[%wsi3_u1] : view<[%weight_metadata]xf32> -> f32 + %wcv3_u1 = vector.splat %wc3_u1 : vector<8xf32> + %lf0_0_raw_u1 = func.call @ggml_q5_plane_dense_q8_wmmai8_load_lhs(%token256, %scratch, %base, %lm0, %blk_k_u1, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %lf0_0_words_u1 = vector.bitcast %lf0_0_raw_u1 : vector<16xi8> to vector<4xi32> + %lf0_0_u1 = vector.fragment %lf0_0_words_u1 shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %lf0_1_raw_u1 = func.call @ggml_q5_plane_dense_q8_wmmai8_load_lhs(%token256, %scratch, %base, %lm0, %blk_k16_u1, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %lf0_1_words_u1 = vector.bitcast %lf0_1_raw_u1 : vector<16xi8> to vector<4xi32> + %lf0_1_u1 = vector.fragment %lf0_1_words_u1 shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %asi0_0_u1 = index.add %asb0, %c1 : index + %asi0_1_u1 = index.mul %asi0_0_u1, %c2 : index + %asi0_2_u1 = index.add %asi0_1_u1, %lane_hi : index + %asi0_3_u1 = index.mul %asi0_2_u1, %c8 : index + %asv0_u1, %auv0_u1 = func.call @ggml_q5_plane_dense_q8_wmmai8_load_metadata(%token256, %scratch, %as_off, %asum_off, %asi0_3_u1) : (i1, buffer, offset, offset, index) -> (vector<8xf32>, vector<8xf32>) + // Issue independent RHS fragments before consuming the first pair. + %rfs0_0_0_raw_u1 = vector.fragment.load %wl_view[%blk_k_u1, %ln0] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs0_0_0_words_u1 = vector.bitcast %rfs0_0_0_raw_u1 : vector<16xi8> to vector<4xi32> + %rfs0_0_0_u1 = vector.fragment %rfs0_0_0_words_u1 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs0_0_1_raw_u1 = vector.fragment.load %wl_view[%blk_k16_u1, %ln0] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs0_0_1_words_u1 = vector.bitcast %rfs0_0_1_raw_u1 : vector<16xi8> to vector<4xi32> + %rfs0_0_1_u1 = vector.fragment %rfs0_0_1_words_u1 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs0_1_0_raw_u1 = vector.fragment.load %wl_view[%blk_k_u1, %ln1] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs0_1_0_words_u1 = vector.bitcast %rfs0_1_0_raw_u1 : vector<16xi8> to vector<4xi32> + %rfs0_1_0_u1 = vector.fragment %rfs0_1_0_words_u1 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs0_1_1_raw_u1 = vector.fragment.load %wl_view[%blk_k16_u1, %ln1] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs0_1_1_words_u1 = vector.bitcast %rfs0_1_1_raw_u1 : vector<16xi8> to vector<4xi32> + %rfs0_1_1_u1 = vector.fragment %rfs0_1_1_words_u1 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs0_2_0_raw_u1 = vector.fragment.load %wl_view[%blk_k_u1, %ln2] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs0_2_0_words_u1 = vector.bitcast %rfs0_2_0_raw_u1 : vector<16xi8> to vector<4xi32> + %rfs0_2_0_u1 = vector.fragment %rfs0_2_0_words_u1 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs0_2_1_raw_u1 = vector.fragment.load %wl_view[%blk_k16_u1, %ln2] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs0_2_1_words_u1 = vector.bitcast %rfs0_2_1_raw_u1 : vector<16xi8> to vector<4xi32> + %rfs0_2_1_u1 = vector.fragment %rfs0_2_1_words_u1 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs0_3_0_raw_u1 = vector.fragment.load %wl_view[%blk_k_u1, %ln3] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs0_3_0_words_u1 = vector.bitcast %rfs0_3_0_raw_u1 : vector<16xi8> to vector<4xi32> + %rfs0_3_0_u1 = vector.fragment %rfs0_3_0_words_u1 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs0_3_1_raw_u1 = vector.fragment.load %wl_view[%blk_k16_u1, %ln3] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs0_3_1_words_u1 = vector.bitcast %rfs0_3_1_raw_u1 : vector<16xi8> to vector<4xi32> + %rfs0_3_1_u1 = vector.fragment %rfs0_3_1_words_u1 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %i0_0_0_u1 = vector.mma %lf0_0_u1, %rfs0_0_0_u1, %iz_u1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i1_0_0_u1 = vector.mma %lf0_1_u1, %rfs0_0_1_u1, %i0_0_0_u1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %ff0_0_u1 = vector.sitofp %i1_0_0_u1 : vector<8xi32> to vector<8xf32> + %sv0_0_u1 = vector.mulf %asv0_u1, %wsv0_u1 : vector<8xf32> + %base_fn0_0 = vector.fmaf %auv0_u1, %wcv0_u1, %fn0_0_u0 : vector<8xf32> + %selected_base_fn0_0 = scf.select %is_signed_rhs, %fn0_0_u0, %base_fn0_0 : vector<8xf32> + + %fn0_0 = vector.fmaf %ff0_0_u1, %sv0_0_u1, %selected_base_fn0_0 : vector<8xf32> + scf.schedule.fence + %i0_0_1_u1 = vector.mma %lf0_0_u1, %rfs0_1_0_u1, %iz_u1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i1_0_1_u1 = vector.mma %lf0_1_u1, %rfs0_1_1_u1, %i0_0_1_u1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %ff0_1_u1 = vector.sitofp %i1_0_1_u1 : vector<8xi32> to vector<8xf32> + %sv0_1_u1 = vector.mulf %asv0_u1, %wsv1_u1 : vector<8xf32> + %base_fn0_1 = vector.fmaf %auv0_u1, %wcv1_u1, %fn0_1_u0 : vector<8xf32> + %selected_base_fn0_1 = scf.select %is_signed_rhs, %fn0_1_u0, %base_fn0_1 : vector<8xf32> + + %fn0_1 = vector.fmaf %ff0_1_u1, %sv0_1_u1, %selected_base_fn0_1 : vector<8xf32> + scf.schedule.fence + %i0_0_2_u1 = vector.mma %lf0_0_u1, %rfs0_2_0_u1, %iz_u1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i1_0_2_u1 = vector.mma %lf0_1_u1, %rfs0_2_1_u1, %i0_0_2_u1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %ff0_2_u1 = vector.sitofp %i1_0_2_u1 : vector<8xi32> to vector<8xf32> + %sv0_2_u1 = vector.mulf %asv0_u1, %wsv2_u1 : vector<8xf32> + %base_fn0_2 = vector.fmaf %auv0_u1, %wcv2_u1, %fn0_2_u0 : vector<8xf32> + %selected_base_fn0_2 = scf.select %is_signed_rhs, %fn0_2_u0, %base_fn0_2 : vector<8xf32> + + %fn0_2 = vector.fmaf %ff0_2_u1, %sv0_2_u1, %selected_base_fn0_2 : vector<8xf32> + scf.schedule.fence + %i0_0_3_u1 = vector.mma %lf0_0_u1, %rfs0_3_0_u1, %iz_u1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i1_0_3_u1 = vector.mma %lf0_1_u1, %rfs0_3_1_u1, %i0_0_3_u1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %ff0_3_u1 = vector.sitofp %i1_0_3_u1 : vector<8xi32> to vector<8xf32> + %sv0_3_u1 = vector.mulf %asv0_u1, %wsv3_u1 : vector<8xf32> + %base_fn0_3 = vector.fmaf %auv0_u1, %wcv3_u1, %fn0_3_u0 : vector<8xf32> + %selected_base_fn0_3 = scf.select %is_signed_rhs, %fn0_3_u0, %base_fn0_3 : vector<8xf32> + + %fn0_3 = vector.fmaf %ff0_3_u1, %sv0_3_u1, %selected_base_fn0_3 : vector<8xf32> + scf.schedule.fence + %lf1_0_raw_u1 = func.call @ggml_q5_plane_dense_q8_wmmai8_load_lhs(%token256, %scratch, %base, %lm1, %blk_k_u1, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %lf1_0_words_u1 = vector.bitcast %lf1_0_raw_u1 : vector<16xi8> to vector<4xi32> + %lf1_0_u1 = vector.fragment %lf1_0_words_u1 shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %lf1_1_raw_u1 = func.call @ggml_q5_plane_dense_q8_wmmai8_load_lhs(%token256, %scratch, %base, %lm1, %blk_k16_u1, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %lf1_1_words_u1 = vector.bitcast %lf1_1_raw_u1 : vector<16xi8> to vector<4xi32> + %lf1_1_u1 = vector.fragment %lf1_1_words_u1 shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %asi1_0_u1 = index.add %asb1, %c1 : index + %asi1_1_u1 = index.mul %asi1_0_u1, %c2 : index + %asi1_2_u1 = index.add %asi1_1_u1, %lane_hi : index + %asi1_3_u1 = index.mul %asi1_2_u1, %c8 : index + %asv1_u1, %auv1_u1 = func.call @ggml_q5_plane_dense_q8_wmmai8_load_metadata(%token256, %scratch, %as_off, %asum_off, %asi1_3_u1) : (i1, buffer, offset, offset, index) -> (vector<8xf32>, vector<8xf32>) + // Issue independent RHS fragments before consuming the first pair. + %rfs1_0_0_raw_u1 = vector.fragment.load %wl_view[%blk_k_u1, %ln0] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs1_0_0_words_u1 = vector.bitcast %rfs1_0_0_raw_u1 : vector<16xi8> to vector<4xi32> + %rfs1_0_0_u1 = vector.fragment %rfs1_0_0_words_u1 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs1_0_1_raw_u1 = vector.fragment.load %wl_view[%blk_k16_u1, %ln0] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs1_0_1_words_u1 = vector.bitcast %rfs1_0_1_raw_u1 : vector<16xi8> to vector<4xi32> + %rfs1_0_1_u1 = vector.fragment %rfs1_0_1_words_u1 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs1_1_0_raw_u1 = vector.fragment.load %wl_view[%blk_k_u1, %ln1] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs1_1_0_words_u1 = vector.bitcast %rfs1_1_0_raw_u1 : vector<16xi8> to vector<4xi32> + %rfs1_1_0_u1 = vector.fragment %rfs1_1_0_words_u1 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs1_1_1_raw_u1 = vector.fragment.load %wl_view[%blk_k16_u1, %ln1] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs1_1_1_words_u1 = vector.bitcast %rfs1_1_1_raw_u1 : vector<16xi8> to vector<4xi32> + %rfs1_1_1_u1 = vector.fragment %rfs1_1_1_words_u1 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs1_2_0_raw_u1 = vector.fragment.load %wl_view[%blk_k_u1, %ln2] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs1_2_0_words_u1 = vector.bitcast %rfs1_2_0_raw_u1 : vector<16xi8> to vector<4xi32> + %rfs1_2_0_u1 = vector.fragment %rfs1_2_0_words_u1 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs1_2_1_raw_u1 = vector.fragment.load %wl_view[%blk_k16_u1, %ln2] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs1_2_1_words_u1 = vector.bitcast %rfs1_2_1_raw_u1 : vector<16xi8> to vector<4xi32> + %rfs1_2_1_u1 = vector.fragment %rfs1_2_1_words_u1 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs1_3_0_raw_u1 = vector.fragment.load %wl_view[%blk_k_u1, %ln3] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs1_3_0_words_u1 = vector.bitcast %rfs1_3_0_raw_u1 : vector<16xi8> to vector<4xi32> + %rfs1_3_0_u1 = vector.fragment %rfs1_3_0_words_u1 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %rfs1_3_1_raw_u1 = vector.fragment.load %wl_view[%blk_k16_u1, %ln3] shape [%c16, %c16] : view<64x[%k_n]xi8, %rhs_layout> -> vector<16xi8> + %rfs1_3_1_words_u1 = vector.bitcast %rfs1_3_1_raw_u1 : vector<16xi8> to vector<4xi32> + %rfs1_3_1_u1 = vector.fragment %rfs1_3_1_words_u1 shape [%c16, %c16] using {schema = %u8_schema : encoding} : vector<4xi32> + %i0_1_0_u1 = vector.mma %lf1_0_u1, %rfs1_0_0_u1, %iz_u1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i1_1_0_u1 = vector.mma %lf1_1_u1, %rfs1_0_1_u1, %i0_1_0_u1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %ff1_0_u1 = vector.sitofp %i1_1_0_u1 : vector<8xi32> to vector<8xf32> + %sv1_0_u1 = vector.mulf %asv1_u1, %wsv0_u1 : vector<8xf32> + %base_fn1_0 = vector.fmaf %auv1_u1, %wcv0_u1, %fn1_0_u0 : vector<8xf32> + %selected_base_fn1_0 = scf.select %is_signed_rhs, %fn1_0_u0, %base_fn1_0 : vector<8xf32> + + %fn1_0 = vector.fmaf %ff1_0_u1, %sv1_0_u1, %selected_base_fn1_0 : vector<8xf32> + scf.schedule.fence + %i0_1_1_u1 = vector.mma %lf1_0_u1, %rfs1_1_0_u1, %iz_u1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i1_1_1_u1 = vector.mma %lf1_1_u1, %rfs1_1_1_u1, %i0_1_1_u1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %ff1_1_u1 = vector.sitofp %i1_1_1_u1 : vector<8xi32> to vector<8xf32> + %sv1_1_u1 = vector.mulf %asv1_u1, %wsv1_u1 : vector<8xf32> + %base_fn1_1 = vector.fmaf %auv1_u1, %wcv1_u1, %fn1_1_u0 : vector<8xf32> + %selected_base_fn1_1 = scf.select %is_signed_rhs, %fn1_1_u0, %base_fn1_1 : vector<8xf32> + + %fn1_1 = vector.fmaf %ff1_1_u1, %sv1_1_u1, %selected_base_fn1_1 : vector<8xf32> + scf.schedule.fence + %i0_1_2_u1 = vector.mma %lf1_0_u1, %rfs1_2_0_u1, %iz_u1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i1_1_2_u1 = vector.mma %lf1_1_u1, %rfs1_2_1_u1, %i0_1_2_u1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %ff1_2_u1 = vector.sitofp %i1_1_2_u1 : vector<8xi32> to vector<8xf32> + %sv1_2_u1 = vector.mulf %asv1_u1, %wsv2_u1 : vector<8xf32> + %base_fn1_2 = vector.fmaf %auv1_u1, %wcv2_u1, %fn1_2_u0 : vector<8xf32> + %selected_base_fn1_2 = scf.select %is_signed_rhs, %fn1_2_u0, %base_fn1_2 : vector<8xf32> + + %fn1_2 = vector.fmaf %ff1_2_u1, %sv1_2_u1, %selected_base_fn1_2 : vector<8xf32> + scf.schedule.fence + %i0_1_3_u1 = vector.mma %lf1_0_u1, %rfs1_3_0_u1, %iz_u1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i1_1_3_u1 = vector.mma %lf1_1_u1, %rfs1_3_1_u1, %i0_1_3_u1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %ff1_3_u1 = vector.sitofp %i1_1_3_u1 : vector<8xi32> to vector<8xf32> + %sv1_3_u1 = vector.mulf %asv1_u1, %wsv3_u1 : vector<8xf32> + %base_fn1_3 = vector.fmaf %auv1_u1, %wcv3_u1, %fn1_3_u0 : vector<8xf32> + %selected_base_fn1_3 = scf.select %is_signed_rhs, %fn1_3_u0, %base_fn1_3 : vector<8xf32> + + %fn1_3 = vector.fmaf %ff1_3_u1, %sv1_3_u1, %selected_base_fn1_3 : vector<8xf32> + scf.yield %fn0_0, %fn0_1, %fn0_2, %fn0_3, %fn1_0, %fn1_1, %fn1_2, %fn1_3 : vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %step0_0, %step0_1, %step0_2, %step0_3, %step1_0, %step1_1, %step1_2, %step1_3 : vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32> + } + + scf.if %paired_weights { + %silu0 = vector.siluf %f0_0 : vector<8xf32> + %result0 = vector.mulf %silu0, %f1_0 : vector<8xf32> + vector.fragment.store %result0, %dst_view[%gm0, %gn0] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + %silu1 = vector.siluf %f0_1 : vector<8xf32> + %result1 = vector.mulf %silu1, %f1_1 : vector<8xf32> + vector.fragment.store %result1, %dst_view[%gm0, %gn1] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + %silu2 = vector.siluf %f0_2 : vector<8xf32> + %result2 = vector.mulf %silu2, %f1_2 : vector<8xf32> + vector.fragment.store %result2, %dst_view[%gm0, %gn2] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + %silu3 = vector.siluf %f0_3 : vector<8xf32> + %result3 = vector.mulf %silu3, %f1_3 : vector<8xf32> + vector.fragment.store %result3, %dst_view[%gm0, %gn3] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + } else { + vector.fragment.store %f0_0, %dst_view[%gm0, %gn0] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %f0_1, %dst_view[%gm0, %gn1] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %f0_2, %dst_view[%gm0, %gn2] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %f0_3, %dst_view[%gm0, %gn3] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %f1_0, %dst_view[%gm1, %gn0] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %f1_1, %dst_view[%gm1, %gn1] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %f1_2, %dst_view[%gm1, %gn2] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %f1_3, %dst_view[%gm1, %gn3] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + } + template.return +} + +kernel.def target(@ggml_q5_plane_dense_gfx11_wave32) @ggml_mul_mat_q5_k_q8_plane_wmmai8_token256(%token_count: index) { + %unit = index.constant 1 : index + %tile_rows = index.constant 64 : index + %tile_tokens = index.constant 256 : index + %workgroup_size = index.constant 256 : index + %output_size = config.get @ggml.mul_mat_q5_k_q8_plane.output_size : index + %token_capacity = config.get @ggml.mul_mat_q5_k_q8_plane.token_capacity : index + %row_tiles = index.div %output_size, %tile_rows : index + %token_tiles = index.div %token_capacity, %tile_tokens : index + kernel.launch.config workgroups(%token_tiles, %row_tiles, %unit) workgroup_size(%workgroup_size, %unit, %unit) : index +} launch(%token_count: index, %q8_input: buffer, %weight: buffer, %output: buffer) { + %base = index.constant 0 : offset + %lds_bytes = index.constant 30720 : offset + %w_off = index.constant 20480 : offset + %as_off = index.constant 25600 : offset + %ws_off = index.constant 27648 : offset + %asum_off = index.constant 28160 : offset + %wc_off = index.constant 30208 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c17 = index.constant 17 : index + %c34 = index.constant 34 : index + %k_kc = index.constant 64 : index + %k_astride = index.constant 80 : index + %k_wstride = index.constant 80 : index + %k_m = index.constant 256 : index + %k_n = index.constant 64 : index + %k_blocks = index.constant 2 : index + %k_wave = index.constant 32 : index + %k_wm = index.constant 8 : index + %fzero = vector.constant 0.0 : vector<8xf32> + %izero = vector.constant 0 : vector<8xi32> + %is_iq4xs = scalar.constant false : i1 + %is_symi5 = scalar.constant true : i1 + %is_q4 = scalar.constant false : i1 + %is_q4_packed = scalar.constant false : i1 + %token256 = scalar.constant true : i1 + %q8_plane = scalar.constant true : i1 + %paired_weights = scalar.constant false : i1 + + %input_size0 = config.get @ggml.mul_mat_q5_k_q8_plane.input_size : index + %output_size0 = config.get @ggml.mul_mat_q5_k_q8_plane.output_size : index + %token_capacity = config.get @ggml.mul_mat_q5_k_q8_plane.token_capacity : index + %input_size = index.assume %input_size0 [range(%input_size0, 256, 32768), mul(%input_size0, 256)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 64, 262144), mul(%output_size0, 64)] : index + %bounded_token_count, %configured_token_capacity = index.assume %token_count, %token_capacity [range(%token_count, 256, 2048), mul(%token_count, 256), eq(%token_count, %token_capacity)] : index, index + %nchunks = index.div %input_size, %k_kc : index + %kblocks = index.div %input_size, %c32 : index + + %col_tile0 = kernel.workgroup.id : index + %row_tile0 = kernel.workgroup.id : index + %tid0 = kernel.workitem.id : index + %col_tile = index.assume %col_tile0 [range(%col_tile0, 0, 7)] : index + %row_tile = index.assume %row_tile0 [range(%row_tile0, 0, 4095)] : index + %tid = index.assume %tid0 [range(%tid0, 0, 255)] : index + %wave = index.div %tid, %k_wave : index + + %weight_g = buffer.assume.memory_space %weight : buffer + %output_g = buffer.assume.memory_space %output : buffer + %q8_g = buffer.assume.memory_space %q8_input : buffer + %weight_noalias, %output_noalias, %q8_noalias = buffer.assume.noalias %weight_g, %output_g, %q8_g : buffer, buffer, buffer + + template.apply<@ggml.mul_mat_q5_k_q8_plane.wmmai8.body>(%is_iq4xs, %token256, %q8_plane, %paired_weights, %weight_noalias, %weight_noalias, %output_noalias, %q8_noalias, %base, %lds_bytes, %w_off, %as_off, %ws_off, %asum_off, %wc_off, %c0, %c1, %c2, %c4, %c8, %c16, %c32, %c17, %c34, %k_kc, %k_astride, %k_wstride, %k_m, %k_n, %k_blocks, %k_wave, %k_wm, %fzero, %izero, %input_size, %output_size, %bounded_token_count, %nchunks, %kblocks, %col_tile, %row_tile, %tid, %wave, %is_symi5, %is_q4, %is_q4_packed) : (i1, i1, i1, i1, buffer, buffer, buffer, buffer, offset, offset, offset, offset, offset, offset, offset, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, vector<8xf32>, vector<8xi32>, index, index, index, index, index, index, index, index, index, i1, i1, i1) + kernel.return +} + +// The JIT format and paired-projection constants remove inactive paths. +template.def<@ggml.mul_mat_q8_1_x4.entry> device @ggml_mul_mat_q8_1_x4_entry(%token_count: index, %paired_weights: i1, %q8_input: buffer, %weight: buffer, %peer_weight: buffer, %output: buffer) { + %base = index.constant 0 : offset + %lds_bytes_ordinary = index.constant 30720 : offset + %lds_bytes_paired = index.constant 24576 : offset + %lds_bytes = scf.select %paired_weights, %lds_bytes_paired, %lds_bytes_ordinary : offset + %w_off_ordinary = index.constant 20480 : offset + %w_off_paired = index.constant 10240 : offset + %w_off = scf.select %paired_weights, %w_off_paired, %w_off_ordinary : offset + %as_off_ordinary = index.constant 25600 : offset + %as_off_paired = index.constant 20480 : offset + %as_off = scf.select %paired_weights, %as_off_paired, %as_off_ordinary : offset + %ws_off_ordinary = index.constant 27648 : offset + %ws_off_paired = index.constant 21504 : offset + %ws_off = scf.select %paired_weights, %ws_off_paired, %ws_off_ordinary : offset + %asum_off_ordinary = index.constant 28160 : offset + %asum_off_paired = index.constant 22528 : offset + %asum_off = scf.select %paired_weights, %asum_off_paired, %asum_off_ordinary : offset + %wc_off_ordinary = index.constant 30208 : offset + %wc_off_paired = index.constant 23552 : offset + %wc_off = scf.select %paired_weights, %wc_off_paired, %wc_off_ordinary : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c23 = index.constant 23 : index + %c32 = index.constant 32 : index + %c17 = index.constant 17 : index + %c34 = index.constant 34 : index + %k_kc = index.constant 64 : index + %k_astride = index.constant 80 : index + %k_wstride = index.constant 80 : index + %k_m_ordinary = index.constant 256 : index + %k_m_paired = index.constant 128 : index + %k_m = scf.select %paired_weights, %k_m_paired, %k_m_ordinary : index + %k_n_ordinary = index.constant 64 : index + %k_n_paired = index.constant 128 : index + %k_n = scf.select %paired_weights, %k_n_paired, %k_n_ordinary : index + %k_blocks = index.constant 2 : index + %k_wave = index.constant 32 : index + %k_wm = index.constant 8 : index + %fzero = vector.constant 0.0 : vector<8xf32> + %izero = vector.constant 0 : vector<8xi32> + %weight_format = config.get @ggml.mul_mat_q8_1_x4.weight_format : index + %is_iq4xs = index.cmp eq, %weight_format, %c23 : index + %true = scalar.constant true : i1 + %token256 = scalar.xori %paired_weights, %true : i1 + %q8_plane = scalar.constant false : i1 + %is_symi5 = scalar.constant false : i1 + %q4_native = index.cmp eq, %weight_format, %c4 : index + %q4_packed_format = index.constant 44 : index + %is_q4_packed = index.cmp eq, %weight_format, %q4_packed_format : index + %is_q4 = scalar.ori %q4_native, %is_q4_packed : i1 + + %input_size0 = config.get @ggml.mul_mat_q8_1_x4.input_size : index + %output_size0 = config.get @ggml.mul_mat_q8_1_x4.output_size : index + %token_capacity = config.get @ggml.mul_mat_q8_1_x4.token_capacity : index + %input_size = index.assume %input_size0 [range(%input_size0, 256, 32768), mul(%input_size0, 256)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 64, 262144), mul(%output_size0, 64)] : index + %bounded_token_count, %configured_token_capacity = index.assume %token_count, %token_capacity [range(%token_count, 256, 2048), mul(%token_count, 256), eq(%token_count, %token_capacity)] : index, index + %nchunks = index.div %input_size, %k_kc : index + %kblocks = index.div %input_size, %c32 : index + + %col_tile0 = kernel.workgroup.id : index + %row_tile0 = kernel.workgroup.id : index + %tid0 = kernel.workitem.id : index + %col_tile = scf.if %paired_weights -> (index) { + %col = index.assume %col_tile0 [range(%col_tile0, 0, 15)] : index + scf.yield %col : index + } else { + %col = index.assume %col_tile0 [range(%col_tile0, 0, 7)] : index + scf.yield %col : index + } + %row_tile = scf.if %paired_weights -> (index) { + %row = index.assume %row_tile0 [range(%row_tile0, 0, 8191)] : index + scf.yield %row : index + } else { + %row = index.assume %row_tile0 [range(%row_tile0, 0, 4095)] : index + scf.yield %row : index + } + %tid = index.assume %tid0 [range(%tid0, 0, 255)] : index + %wave = index.div %tid, %k_wave : index + + %peer_global = buffer.assume.memory_space %peer_weight : buffer + %weight_g = buffer.assume.memory_space %weight : buffer + %output_g = buffer.assume.memory_space %output : buffer + %q8_g = buffer.assume.memory_space %q8_input : buffer + %weight_noalias, %output_noalias, %q8_noalias = buffer.assume.noalias %weight_g, %output_g, %q8_g : buffer, buffer, buffer + + template.apply<@ggml.mul_mat_q5_k_q8_plane.wmmai8.body>(%is_iq4xs, %token256, %q8_plane, %paired_weights, %peer_global, %weight_noalias, %output_noalias, %q8_noalias, %base, %lds_bytes, %w_off, %as_off, %ws_off, %asum_off, %wc_off, %c0, %c1, %c2, %c4, %c8, %c16, %c32, %c17, %c34, %k_kc, %k_astride, %k_wstride, %k_m, %k_n, %k_blocks, %k_wave, %k_wm, %fzero, %izero, %input_size, %output_size, %bounded_token_count, %nchunks, %kblocks, %col_tile, %row_tile, %tid, %wave, %is_symi5, %is_q4, %is_q4_packed) : (i1, i1, i1, i1, buffer, buffer, buffer, buffer, offset, offset, offset, offset, offset, offset, offset, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, vector<8xf32>, vector<8xi32>, index, index, index, index, index, index, index, index, index, i1, i1, i1) + template.return +} + +kernel.def target(@ggml_q5_plane_dense_gfx11_wave32) export("ggml_mul_mat_q5_k_iq4_xs_q8_1_x4_wmma_token256") @ggml_mul_mat_q5_k_iq4_xs_q8_1_x4_wmma_token256(%token_count: index) { + %unit = index.constant 1 : index + %tile_rows = index.constant 64 : index + %tile_tokens = index.constant 256 : index + %workgroup_size = index.constant 256 : index + %output_size = config.get @ggml.mul_mat_q8_1_x4.output_size : index + %token_capacity = config.get @ggml.mul_mat_q8_1_x4.token_capacity : index + %row_tiles = index.div %output_size, %tile_rows : index + %token_tiles = index.div %token_capacity, %tile_tokens : index + kernel.launch.config workgroups(%token_tiles, %row_tiles, %unit) workgroup_size(%workgroup_size, %unit, %unit) : index +} launch(%token_count: index, %q8_input: buffer, %weight: buffer, %output: buffer) { + %paired_weights = scalar.constant false : i1 + template.apply<@ggml.mul_mat_q8_1_x4.entry>(%token_count, %paired_weights, %q8_input, %weight, %weight, %output) : (index, i1, buffer, buffer, buffer, buffer) + kernel.return +} + +kernel.def target(@ggml_q5_plane_dense_gfx11_wave32) export("ggml_mul_mat_q4_k_q8_1_x4_swiglu_wmma_token256") @ggml_mul_mat_q4_k_q8_1_x4_swiglu_wmma_token256(%token_count: index) { + %unit = index.constant 1 : index + %tile_rows = index.constant 64 : index + %tile_tokens = index.constant 128 : index + %workgroup_size = index.constant 256 : index + %output_size = config.get @ggml.mul_mat_q8_1_x4.output_size : index + %token_capacity = config.get @ggml.mul_mat_q8_1_x4.token_capacity : index + %row_tiles = index.div %output_size, %tile_rows : index + %token_tiles = index.div %token_capacity, %tile_tokens : index + kernel.launch.config workgroups(%token_tiles, %row_tiles, %unit) workgroup_size(%workgroup_size, %unit, %unit) : index +} launch(%token_count: index, %q8_input: buffer, %weight: buffer, %up_weight: buffer, %output: buffer) { + %paired_weights = scalar.constant true : i1 + template.apply<@ggml.mul_mat_q8_1_x4.entry>(%token_count, %paired_weights, %q8_input, %weight, %up_weight, %output) : (index, i1, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_q6_k_packed_f16_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_q6_k_packed_f16_wmma.loom new file mode 100644 index 000000000000..52341d38eeb9 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_q6_k_packed_f16_wmma.loom @@ -0,0 +1,1580 @@ +func.decl @ggml_mul_mat_quantized_f16_prefill_wave32(%weight_format: index, %binary_op: index, %paired: i1, %packed_input: i1, %packed_output: i1, %token_count: index, %input_size: index, %output_size: index, %channel_tile: index, %token_tile: index, %input: buffer, %weight: buffer, %up_weight: buffer, %output: buffer, %f16_output: buffer) + +config.decl @ggml.mul_mat_q6_k_packed.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +config.decl @ggml.mul_mat_q6_k_packed.output_size : %value: index where [range(%value, 64, 262144), mul(%value, 64)] + +config.decl @ggml.mul_mat_q6_k_packed.output_accumulation : %value: index where [range(%value, 0, 1)] + +config.decl @ggml.mul_mat_q6_k_packed.weight_offset : %value: index where [range(%value, 0, 2147483647), mul(%value, 256)] + +config.decl @ggml.mul_mat_q6_k_packed.token_capacity : %value: index where [range(%value, 128, 2048), mul(%value, 128)] + +config.decl @ggml.mul_mat_q6_k_shortlist.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +config.decl @ggml.mul_mat_q6_k_shortlist.output_size : %value: index where [range(%value, 64, 262144), mul(%value, 64)] + +config.decl @ggml.mul_mat_q6_k_shortlist.selected_group_count : %value: index where [range(%value, 1, 1024)] + +template.decl @ggml.mul_mat_q6_k_packed.publish_f16_wave32(%result: vector<16xf16>, %token_offset: index, %token_tile_base: index, %bounded_token_count: index, %bounded_output_size: index, %channel_tile_base: index, %subgroup_channel_add: index, %result_stage: buffer, %output_noalias: buffer, %output_accumulation: index) + +template.decl @ggml.mul_mat_q6_k_packed.publish_token1_overwrite_f32_wave32(%result: vector<8xf32>, %output_size: index, %channel_tile_base: index, %subgroup_channel_add: index, %output: buffer) + +template.decl @ggml.mul_mat_q6_k_packed.token1.body(%token_count: index, %input_size0: index, %output_size0: index, %channel_tile: index, %scale_row_layout: index, %input: buffer, %weight: buffer, %output: buffer) + +template.decl @ggml.mul_mat_q6_k_i8_prepacked.body(%token_count: index, %input_size0: index, %output_size0: index, %channel_tile: index, %input: buffer, %weight: buffer, %output: buffer) + +template.decl @ggml.mul_mat_q6_k_packed.selected_refine.body(%token_count: index, %input_size0: index, %output_size0: index, %scale_row_layout: index, %candidate_count0: index, %input: buffer, %weight: buffer, %candidates: buffer, %output: buffer) + +template.decl @ggml.mul_mat_q6_k_packed.prefill_wave32.body(%input_is_f16: i1, %token_count: index, %input_size0: index, %output_size0: index, %channel_tile: index, %token_tile: index, %input: buffer, %weight: buffer, %output: buffer) + +amdgpu.target @ggml_mul_mat_q6_k_prefill_gfx11_wave32 {subgroup_size = 32} + +func.decl @ggml_q6k_f16_pair(%half_packet: i1, %weight: buffer, %weight_row_byte_base: offset, %q6_block: index, %q6_group: index, %packet: index) -> (vector<16xf16>, vector<16xf16>) + +template.def<@ggml.mul_mat_q6_k_packed.publish_f16_wave32> device @ggml_mul_mat_q6_k_packed_publish_f16_wave32(%result: vector<16xf16>, %token_offset: index, %token_tile_base: index, %bounded_token_count: index, %bounded_output_size: index, %channel_tile_base: index, %subgroup_channel_add: index, %result_stage: buffer, %output_noalias: buffer, %output_accumulation: index) { + %lane = kernel.subgroup.lane.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c0_offset = index.constant 0 : offset + %wave_result_stage_bytes = index.constant 512 : offset + %c0_f32x8 = vector.constant 0.0 : vector<8xf32> + %m = index.constant 16 : index + %n = index.constant 16 : index + %wave_result_stage_offset = index.scale %subgroup, %wave_result_stage_bytes : index, offset -> offset + %result_fragment_layout = encoding.layout.strided [1, %c16] : encoding + %result_fragment_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf16, %result_fragment_layout> + %result_physical_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf16> + %publish_token0 = index.div %lane, %c2 : index + %publish_token = index.assume %publish_token0 [range(%publish_token0, 0, 15)] : index + %publish_packet0 = index.rem %lane, %c2 : index + %publish_packet = index.assume %publish_packet0 [range(%publish_packet0, 0, 1)] : index + %publish_channel_add = index.mul %publish_packet, %c8 : index + %token_base = index.add %token_tile_base, %token_offset : index + %token = index.add %token_base, %publish_token : index + %subgroup_channel_base = index.add %channel_tile_base, %subgroup_channel_add : index + %channel = index.add %subgroup_channel_base, %publish_channel_add : index + %valid_token = index.cmp ult, %token, %bounded_token_count : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %writes = scalar.andi %valid_token, %valid_channel : i1 + %accumulates_output = index.cmp eq, %output_accumulation, %c1 : index + vector.fragment.store %result, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<16xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes { + %bounded_token, %output_token_count = index.assume %token, %bounded_token_count [lt(%token, %bounded_token_count)] : index, index + %bounded_output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%output_token_count]x[%bounded_output_size]xf32> + %values = vector.load %result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf16> -> vector<8xf16> + %wide = vector.extf %values : vector<8xf16> to vector<8xf32> + %mask = vector.mask.range [%channel to %bounded_output_size step %c1] : index -> vector<8xi1> + %published = scf.if %accumulates_output -> (vector<8xf32>) { + %residual = vector.load.mask %bounded_output_view[%bounded_token, %channel], %mask, %c0_f32x8 : view<[%output_token_count]x[%bounded_output_size]xf32>, vector<8xi1>, vector<8xf32> + %sum = vector.addf %residual, %wide : vector<8xf32> + scf.yield %sum : vector<8xf32> + } else { + scf.yield %wide : vector<8xf32> + } + vector.store.mask %published, %bounded_output_view[%bounded_token, %channel], %mask : vector<8xf32>, view<[%output_token_count]x[%bounded_output_size]xf32>, vector<8xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + template.return +} + +template.def<@ggml.mul_mat_q6_k_packed.token1.body> device @ggml_mul_mat_q6_k_packed_token1_body(%token_count: index, %input_size0: index, %output_size0: index, %channel_tile: index, %scale_row_layout: index, %input: buffer, %weight: buffer, %output: buffer) { + %output_accumulation = config.get @ggml.mul_mat_q6_k_packed.output_accumulation : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 16)] : index + %bounded_input_size = index.assume %input_size0 [range(%input_size0, 256, 32768), mul(%input_size0, 256)] : index + %bounded_output_size = index.assume %output_size0 [range(%output_size0, 64, 262144), mul(%output_size0, 64)] : index + %workitem0 = kernel.workitem.id : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 127)] : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 3)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c72 = index.constant 72 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %c128_bytes = index.constant 128 : offset + %c1152_bytes = index.constant 1152 : offset + %c13440_bytes = index.constant 13440 : offset + %weight_stage_bytes = index.constant 9216 : offset + %activation_stage_bytes = index.constant 2304 : offset + %zero_accumulator = vector.constant 0.0 : vector<8xf32> + %empty_words = vector.constant 0 : vector<4xi32> + %zero_scale_bytes = vector.constant 0 : vector<8xi8> + %nibble_mask = vector.constant 252645135 : vector<4xi32> + %high0_mask = vector.constant 50529027 : vector<4xi32> + %high1_mask = vector.constant 202116108 : vector<4xi32> + %high2_mask = vector.constant 808464432 : vector<4xi32> + %high3_mask = vector.constant -1061109568 : vector<4xi32> + %shift2 = vector.constant 2 : vector<4xi32> + %shift4 = vector.constant 4 : vector<4xi32> + %negative32_f32 = vector.constant -32.0 : vector<16xf32> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + + %quant_block_count = index.div %bounded_input_size, %c256 : index + %input_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input, %weight, %output : buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_input_size]xf32> + %weight_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %activation_stage = buffer.alloca align(16) %activation_stage_bytes : buffer + %weight_stage_view = buffer.view %weight_stage[%c0_offset] : buffer -> view<64x72xf16> + %activation_stage_physical_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<16x72xf16> + %activation_fragment_layout = encoding.layout.strided [1, %c72] : encoding + %activation_fragment_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<64x16xf16, %activation_fragment_layout> + + %channel_tile_base = index.mul %channel_tile, %c64 : index + %weight_half = index.rem %workitem, %c2 : index + %weight_local_row0 = index.div %workitem, %c2 : index + %weight_local_row = index.assume %weight_local_row0 [range(%weight_local_row0, 0, 63)] : index + %weight_k = index.mul %weight_half, %c16 : index + %weight_k_next = index.add %weight_k, %c32 : index + %channel = index.add %channel_tile_base, %weight_local_row : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %use_scale_row_layout = index.cmp eq, %scale_row_layout, %c1 : index + %activation_row0 = index.div %workitem, %c8 : index + %activation_row = index.assume %activation_row0 [range(%activation_row0, 0, 15)] : index + %activation_packet = index.rem %workitem, %c8 : index + %activation_k = index.mul %activation_packet, %c4 : index + %activation_k_next = index.add %activation_k, %c32 : index + %valid_token = index.cmp ult, %activation_row, %bounded_token_count : index + %subgroup_channel_add = index.mul %subgroup, %c16 : index + %init = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + + %result = scf.for %quant_block = [%c0 to %quant_block_count step %c1](%block_acc = %init : vector<8xf32>) -> (vector<8xf32>) { + %row_group_block0 = index.mul %channel_tile, %quant_block_count : index + %row_group_block = index.add %row_group_block0, %quant_block : index + %tile_byte_base = index.scale %row_group_block, %c13440_bytes : index, offset -> offset + %scale_byte_base = index.add %tile_byte_base, %c128_bytes : offset + %raw_byte_base = index.add %tile_byte_base, %c1152_bytes : offset + %d_view = buffer.view %weight_noalias[%tile_byte_base] : buffer -> view<64xf16> + %scale_view = buffer.view %weight_noalias[%scale_byte_base] : buffer -> view<8x64x2xi8> + %scale_row_view = buffer.view %weight_noalias[%scale_byte_base] : buffer -> view<64x2x8xi8> + %raw_view = buffer.view %weight_noalias[%raw_byte_base] : buffer -> view<2x3x64x32xi8> + %d = scf.if %valid_channel -> (f32) { + %d_f16 = view.load %d_view[%weight_local_row] : view<64xf16> -> f16 + %d_f32 = scalar.extf %d_f16 : f16 to f32 + scf.yield %d_f32 : f32 + } else { + %c0_f32 = scalar.constant 0.0 : f32 + scf.yield %c0_f32 : f32 + } + %load_scale_row = scalar.andi %valid_channel, %use_scale_row_layout : i1 + %scale_bytes = scf.if %load_scale_row -> (vector<8xi8>) { + %loaded_scale_bytes = vector.load %scale_row_view[%weight_local_row, %weight_half, %c0] : view<64x2x8xi8> -> vector<8xi8> + scf.yield %loaded_scale_bytes : vector<8xi8> + } else { + scf.yield %zero_scale_bytes : vector<8xi8> + } + %block_result, %final_ql0_words, %final_ql1_words, %final_qh_words = scf.for %quant_group = [%c0 to %c8 step %c2](%acc = %block_acc : vector<8xf32>, %cached_ql0_words = %empty_words : vector<4xi32>, %cached_ql1_words = %empty_words : vector<4xi32>, %cached_qh_words = %empty_words : vector<4xi32>) -> (vector<8xf32>, vector<4xi32>, vector<4xi32>, vector<4xi32>) unroll { + %quant_group_next = index.add %quant_group, %c1 : index + %half128 = index.div %quant_group, %c4 : index + %group_in_half = index.rem %quant_group, %c4 : index + %is_low_pair = index.cmp eq, %group_in_half, %c0 : index + %is_first_group = index.cmp eq, %quant_group, %c0 : index + %prefetch_second_half = index.cmp eq, %quant_group, %c2 : index + %weight_values0, %weight_values1, %next_ql0_words, %next_ql1_words, %next_qh_words = scf.if %valid_channel -> (vector<16xf16>, vector<16xf16>, vector<4xi32>, vector<4xi32>, vector<4xi32>) { + %ql0_words, %ql1_words, %qh_words = scf.if %is_first_group -> (vector<4xi32>, vector<4xi32>, vector<4xi32>) { + %ql0_bytes = vector.load %raw_view[%half128, %c0, %weight_local_row, %weight_k] : view<2x3x64x32xi8> -> vector<16xi8> + %ql1_bytes = vector.load %raw_view[%half128, %c1, %weight_local_row, %weight_k] : view<2x3x64x32xi8> -> vector<16xi8> + %qh_bytes = vector.load %raw_view[%half128, %c2, %weight_local_row, %weight_k] : view<2x3x64x32xi8> -> vector<16xi8> + %loaded_ql0_words = vector.bitcast %ql0_bytes : vector<16xi8> to vector<4xi32> + %loaded_ql1_words = vector.bitcast %ql1_bytes : vector<16xi8> to vector<4xi32> + %loaded_qh_words = vector.bitcast %qh_bytes : vector<16xi8> to vector<4xi32> + scf.yield %loaded_ql0_words, %loaded_ql1_words, %loaded_qh_words : vector<4xi32>, vector<4xi32>, vector<4xi32> + } else { + scf.yield %cached_ql0_words, %cached_ql1_words, %cached_qh_words : vector<4xi32>, vector<4xi32>, vector<4xi32> + } + %codes0_words, %codes1_words = scf.if %is_low_pair -> (vector<4xi32>, vector<4xi32>) { + %ql0_low = vector.andi %ql0_words, %nibble_mask : vector<4xi32> + %ql1_low = vector.andi %ql1_words, %nibble_mask : vector<4xi32> + %qh0_bits0 = vector.andi %qh_words, %high0_mask : vector<4xi32> + %qh1_bits0 = vector.andi %qh_words, %high1_mask : vector<4xi32> + %qh0_bits = vector.shli %qh0_bits0, %shift4 : vector<4xi32> + %qh1_bits = vector.shli %qh1_bits0, %shift2 : vector<4xi32> + %codes0 = vector.ori %ql0_low, %qh0_bits : vector<4xi32> + %codes1 = vector.ori %ql1_low, %qh1_bits : vector<4xi32> + scf.yield %codes0, %codes1 : vector<4xi32>, vector<4xi32> + } else { + %ql0_high0 = vector.shrui %ql0_words, %shift4 : vector<4xi32> + %ql1_high0 = vector.shrui %ql1_words, %shift4 : vector<4xi32> + %ql0_high = vector.andi %ql0_high0, %nibble_mask : vector<4xi32> + %ql1_high = vector.andi %ql1_high0, %nibble_mask : vector<4xi32> + %qh2_bits = vector.andi %qh_words, %high2_mask : vector<4xi32> + %qh3_bits0 = vector.andi %qh_words, %high3_mask : vector<4xi32> + %qh3_bits = vector.shrui %qh3_bits0, %shift2 : vector<4xi32> + %codes0 = vector.ori %ql0_high, %qh2_bits : vector<4xi32> + %codes1 = vector.ori %ql1_high, %qh3_bits : vector<4xi32> + scf.yield %codes0, %codes1 : vector<4xi32>, vector<4xi32> + } + %codes0_u8 = vector.bitcast %codes0_words : vector<4xi32> to vector<16xi8> + %codes1_u8 = vector.bitcast %codes1_words : vector<4xi32> to vector<16xi8> + %codes0_unsigned_f32 = vector.uitofp %codes0_u8 : vector<16xi8> to vector<16xf32> + %codes1_unsigned_f32 = vector.uitofp %codes1_u8 : vector<16xi8> to vector<16xf32> + %codes0_f32 = vector.addf %codes0_unsigned_f32, %negative32_f32 : vector<16xf32> + %codes1_f32 = vector.addf %codes1_unsigned_f32, %negative32_f32 : vector<16xf32> + %scale0_i8, %scale1_i8 = scf.if %use_scale_row_layout -> (i8, i8) { + %row_scale0_i8 = vector.extract %scale_bytes[%quant_group] : vector<8xi8> -> i8 + %row_scale1_i8 = vector.extract %scale_bytes[%quant_group_next] : vector<8xi8> -> i8 + scf.yield %row_scale0_i8, %row_scale1_i8 : i8, i8 + } else { + %group_scale0_i8 = view.load %scale_view[%quant_group, %weight_local_row, %weight_half] : view<8x64x2xi8> -> i8 + %group_scale1_i8 = view.load %scale_view[%quant_group_next, %weight_local_row, %weight_half] : view<8x64x2xi8> -> i8 + scf.yield %group_scale0_i8, %group_scale1_i8 : i8, i8 + } + %scale0 = scalar.sitofp %scale0_i8 : i8 to f32 + %scale1 = scalar.sitofp %scale1_i8 : i8 to f32 + %combined0 = scalar.mulf %scale0, %d : f32 + %combined1 = scalar.mulf %scale1, %d : f32 + %combined0_vector = vector.splat %combined0 : vector<16xf32> + %combined1_vector = vector.splat %combined1 : vector<16xf32> + %values0_f32 = vector.mulf %codes0_f32, %combined0_vector : vector<16xf32> + %values1_f32 = vector.mulf %codes1_f32, %combined1_vector : vector<16xf32> + %values0 = vector.fptrunc %values0_f32 : vector<16xf32> to vector<16xf16> + %values1 = vector.fptrunc %values1_f32 : vector<16xf32> to vector<16xf16> + %prefetched_ql0_words, %prefetched_ql1_words, %prefetched_qh_words = scf.if %prefetch_second_half -> (vector<4xi32>, vector<4xi32>, vector<4xi32>) { + %next_ql0_bytes = vector.load %raw_view[%c1, %c0, %weight_local_row, %weight_k] : view<2x3x64x32xi8> -> vector<16xi8> + %next_ql1_bytes = vector.load %raw_view[%c1, %c1, %weight_local_row, %weight_k] : view<2x3x64x32xi8> -> vector<16xi8> + %next_qh_bytes = vector.load %raw_view[%c1, %c2, %weight_local_row, %weight_k] : view<2x3x64x32xi8> -> vector<16xi8> + %next_ql0 = vector.bitcast %next_ql0_bytes : vector<16xi8> to vector<4xi32> + %next_ql1 = vector.bitcast %next_ql1_bytes : vector<16xi8> to vector<4xi32> + %next_qh = vector.bitcast %next_qh_bytes : vector<16xi8> to vector<4xi32> + scf.yield %next_ql0, %next_ql1, %next_qh : vector<4xi32>, vector<4xi32>, vector<4xi32> + } else { + scf.yield %ql0_words, %ql1_words, %qh_words : vector<4xi32>, vector<4xi32>, vector<4xi32> + } + scf.yield %values0, %values1, %prefetched_ql0_words, %prefetched_ql1_words, %prefetched_qh_words : vector<16xf16>, vector<16xf16>, vector<4xi32>, vector<4xi32>, vector<4xi32> + } else { + %zeros0 = vector.constant 0.0 : vector<16xf16> + %zeros1 = vector.constant 0.0 : vector<16xf16> + scf.yield %zeros0, %zeros1, %cached_ql0_words, %cached_ql1_words, %cached_qh_words : vector<16xf16>, vector<16xf16>, vector<4xi32>, vector<4xi32>, vector<4xi32> + } + vector.store %weight_values0, %weight_stage_view[%weight_local_row, %weight_k] : vector<16xf16>, view<64x72xf16> + vector.store %weight_values1, %weight_stage_view[%weight_local_row, %weight_k_next] : vector<16xf16>, view<64x72xf16> + + %activation_values0, %activation_values1 = scf.if %valid_token -> (vector<4xf16>, vector<4xf16>) { + %bounded_token, %input_token_count = index.assume %activation_row, %bounded_token_count [lt(%activation_row, %bounded_token_count)] : index, index + %block_k_base = index.mul %quant_block, %c256 : index + %group_k_add = index.mul %quant_group, %c32 : index + %input_k0 = index.add %block_k_base, %group_k_add : index + %input_k = index.add %input_k0, %activation_k : index + %input_k_next = index.add %input_k0, %activation_k_next : index + %loaded0 = vector.load %input_view[%bounded_token, %input_k] : view<[%bounded_token_count]x[%bounded_input_size]xf32> -> vector<4xf32> + %loaded1 = vector.load %input_view[%bounded_token, %input_k_next] : view<[%bounded_token_count]x[%bounded_input_size]xf32> -> vector<4xf32> + %narrow0 = vector.fptrunc %loaded0 : vector<4xf32> to vector<4xf16> + %narrow1 = vector.fptrunc %loaded1 : vector<4xf32> to vector<4xf16> + scf.yield %narrow0, %narrow1 : vector<4xf16>, vector<4xf16> + } else { + %zeros0 = vector.constant 0.0 : vector<4xf16> + %zeros1 = vector.constant 0.0 : vector<4xf16> + scf.yield %zeros0, %zeros1 : vector<4xf16>, vector<4xf16> + } + vector.store %activation_values0, %activation_stage_physical_view[%activation_row, %activation_k] : vector<4xf16>, view<16x72xf16> + vector.store %activation_values1, %activation_stage_physical_view[%activation_row, %activation_k_next] : vector<4xf16>, view<16x72xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + %next = scf.for %k_half = [%c0 to %c64 step %c16](%half_acc = %acc : vector<8xf32>) -> (vector<8xf32>) unroll { + %lhs = vector.fragment.load %weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x72xf16> -> vector<16xf16> + %rhs = vector.fragment.load %activation_fragment_view[%k_half, %c0] shape [%k, %n] : view<64x16xf16, %activation_fragment_layout> -> vector<16xf16> + %half_next = vector.mma %lhs, %rhs, %half_acc : vector<16xf16>, vector<16xf16>, vector<8xf32> + scf.yield %half_next : vector<8xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %next, %next_ql0_words, %next_ql1_words, %next_qh_words : vector<8xf32>, vector<4xi32>, vector<4xi32>, vector<4xi32> + } + scf.yield %block_result : vector<8xf32> + } + + %overwrites_output = index.cmp eq, %output_accumulation, %c0 : index + %has_one_token = index.cmp eq, %bounded_token_count, %c1 : index + %use_token1_overwrite = scalar.andi %overwrites_output, %has_one_token : i1 + scf.if %use_token1_overwrite { + template.apply<@ggml.mul_mat_q6_k_packed.publish_token1_overwrite_f32_wave32>(%result, %bounded_output_size, %channel_tile_base, %subgroup_channel_add, %output_noalias) : (vector<8xf32>, index, index, index, buffer) + } else { + template.apply<@ggml.mul_mat_q6_k_packed.publish_f32_wave32>(%result, %c0, %c0, %bounded_token_count, %bounded_output_size, %channel_tile_base, %subgroup_channel_add, %output_noalias, %output_accumulation) : (vector<8xf32>, index, index, index, index, index, index, buffer, index) + } + template.return +} + +template.def<@ggml.mul_mat_q6_k_packed.publish_token1_overwrite_f32_wave32> device @ggml_mul_mat_q6_k_packed_publish_token1_overwrite_f32_wave32(%result: vector<8xf32>, %output_size: index, %channel_tile_base: index, %subgroup_channel_add: index, %output: buffer) { + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %channel = index.add %channel_tile_base, %subgroup_channel_add : index + %lane = kernel.subgroup.lane.id : index + %lane_token = index.rem %lane, %c16 : index + %lane_channel = index.div %lane, %c16 : index + %partial_channel_base = index.add %channel, %lane_channel : index + %stores_token0 = index.cmp eq, %lane_token, %c0 : index + scf.if %stores_token0 { + %dst = buffer.view %output[%base] : buffer -> view<1x[%output_size]xf32> + scf.for %element = [%c0 to %c8 step %c1] unroll { + %channel_step = index.mul %element, %c2 : index + %partial_channel = index.add %partial_channel_base, %channel_step : index + %channel_valid = index.cmp ult, %partial_channel, %output_size : index + scf.if %channel_valid { + %bounded_channel, %output_channels = index.assume %partial_channel, %output_size [lt(%partial_channel, %output_size)] : index, index + %value = vector.extract %result[%element] : vector<8xf32> -> f32 + view.store %value, %dst[%c0, %bounded_channel] : f32, view<1x[%output_size]xf32> + } + } + } + template.return +} + +kernel.def export("ggml_mul_mat_q6_k_packed_token1_f16_wmma") @ggml_mul_mat_q6_k_packed_token1_f16_wmma(%token_count: index) { + %output_size = config.get @ggml.mul_mat_q6_k_packed.output_size : index + %c1 = index.constant 1 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + kernel.launch.config workgroups(%output_tiles, %c1, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) { + %input_size = config.get @ggml.mul_mat_q6_k_packed.input_size : index + %output_size = config.get @ggml.mul_mat_q6_k_packed.output_size : index + %scale_row_layout = index.constant 1 : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 16)] : index + %channel_tile = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_q6_k_packed.token1.body>(%bounded_token_count, %input_size, %output_size, %channel_tile, %scale_row_layout, %input, %weight, %output) : (index, index, index, index, index, buffer, buffer, buffer) + kernel.return +} + +template.def<@ggml.mul_mat_q6_k_i8_prepacked.body> device @ggml_mul_mat_q6_k_i8_prepacked_body(%token_count: index, %input_size0: index, %output_size0: index, %channel_tile: index, %input: buffer, %weight: buffer, %output: buffer) { + %output_accumulation = config.get @ggml.mul_mat_q6_k_packed.output_accumulation : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 16)] : index + %bounded_input_size = index.assume %input_size0 [range(%input_size0, 256, 32768), mul(%input_size0, 256)] : index + %bounded_output_size = index.assume %output_size0 [range(%output_size0, 64, 262144), mul(%output_size0, 64)] : index + %workitem0 = kernel.workitem.id : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 127)] : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 3)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c72 = index.constant 72 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %c128_bytes = index.constant 128 : offset + %c1152_bytes = index.constant 1152 : offset + %c17536_bytes = index.constant 17536 : offset + %weight_stage_bytes = index.constant 9216 : offset + %activation_stage_bytes = index.constant 2304 : offset + %zero_accumulator = vector.constant 0.0 : vector<8xf32> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + + %quant_block_count = index.div %bounded_input_size, %c256 : index + %input_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input, %weight, %output : buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_input_size]xf32> + %weight_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %activation_stage = buffer.alloca align(16) %activation_stage_bytes : buffer + %weight_stage_view = buffer.view %weight_stage[%c0_offset] : buffer -> view<64x72xf16> + %activation_stage_physical_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<16x72xf16> + %activation_fragment_layout = encoding.layout.strided [1, %c72] : encoding + %activation_fragment_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<64x16xf16, %activation_fragment_layout> + + %channel_tile_base = index.mul %channel_tile, %c64 : index + %weight_half = index.rem %workitem, %c2 : index + %weight_local_row0 = index.div %workitem, %c2 : index + %weight_local_row = index.assume %weight_local_row0 [range(%weight_local_row0, 0, 63)] : index + %weight_k = index.mul %weight_half, %c16 : index + %weight_k_next = index.add %weight_k, %c32 : index + %channel = index.add %channel_tile_base, %weight_local_row : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %activation_row0 = index.div %workitem, %c8 : index + %activation_row = index.assume %activation_row0 [range(%activation_row0, 0, 15)] : index + %activation_packet = index.rem %workitem, %c8 : index + %activation_k = index.mul %activation_packet, %c4 : index + %activation_k_next = index.add %activation_k, %c32 : index + %valid_token = index.cmp ult, %activation_row, %bounded_token_count : index + %subgroup_channel_add = index.mul %subgroup, %c16 : index + %init = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + + %result = scf.for %quant_block = [%c0 to %quant_block_count step %c1](%block_acc = %init : vector<8xf32>) -> (vector<8xf32>) { + %row_group_block0 = index.mul %channel_tile, %quant_block_count : index + %row_group_block = index.add %row_group_block0, %quant_block : index + %tile_byte_base = index.scale %row_group_block, %c17536_bytes : index, offset -> offset + %scale_byte_base = index.add %tile_byte_base, %c128_bytes : offset + %code_byte_base = index.add %tile_byte_base, %c1152_bytes : offset + %d_view = buffer.view %weight_noalias[%tile_byte_base] : buffer -> view<64xf16> + %scale_view = buffer.view %weight_noalias[%scale_byte_base] : buffer -> view<8x64x2xi8> + %code_view = buffer.view %weight_noalias[%code_byte_base] : buffer -> view<8x64x32xi8> + %d = scf.if %valid_channel -> (f32) { + %d_f16 = view.load %d_view[%weight_local_row] : view<64xf16> -> f16 + %d_f32 = scalar.extf %d_f16 : f16 to f32 + scf.yield %d_f32 : f32 + } else { + %c0_f32 = scalar.constant 0.0 : f32 + scf.yield %c0_f32 : f32 + } + %block_result = scf.for %quant_group = [%c0 to %c8 step %c2](%acc = %block_acc : vector<8xf32>) -> (vector<8xf32>) { + %quant_group_next = index.add %quant_group, %c1 : index + %weight_values0, %weight_values1 = scf.if %valid_channel -> (vector<16xf16>, vector<16xf16>) { + %codes0 = vector.load %code_view[%quant_group, %weight_local_row, %weight_k] : view<8x64x32xi8> -> vector<16xi8> + %codes1 = vector.load %code_view[%quant_group_next, %weight_local_row, %weight_k] : view<8x64x32xi8> -> vector<16xi8> + %codes0_f32 = vector.sitofp %codes0 : vector<16xi8> to vector<16xf32> + %codes1_f32 = vector.sitofp %codes1 : vector<16xi8> to vector<16xf32> + %scale0_i8 = view.load %scale_view[%quant_group, %weight_local_row, %weight_half] : view<8x64x2xi8> -> i8 + %scale1_i8 = view.load %scale_view[%quant_group_next, %weight_local_row, %weight_half] : view<8x64x2xi8> -> i8 + %scale0 = scalar.sitofp %scale0_i8 : i8 to f32 + %scale1 = scalar.sitofp %scale1_i8 : i8 to f32 + %combined0 = scalar.mulf %scale0, %d : f32 + %combined1 = scalar.mulf %scale1, %d : f32 + %combined0_vector = vector.splat %combined0 : vector<16xf32> + %combined1_vector = vector.splat %combined1 : vector<16xf32> + %values0_f32 = vector.mulf %codes0_f32, %combined0_vector : vector<16xf32> + %values1_f32 = vector.mulf %codes1_f32, %combined1_vector : vector<16xf32> + %values0 = vector.fptrunc %values0_f32 : vector<16xf32> to vector<16xf16> + %values1 = vector.fptrunc %values1_f32 : vector<16xf32> to vector<16xf16> + scf.yield %values0, %values1 : vector<16xf16>, vector<16xf16> + } else { + %zeros0 = vector.constant 0.0 : vector<16xf16> + %zeros1 = vector.constant 0.0 : vector<16xf16> + scf.yield %zeros0, %zeros1 : vector<16xf16>, vector<16xf16> + } + vector.store %weight_values0, %weight_stage_view[%weight_local_row, %weight_k] : vector<16xf16>, view<64x72xf16> + vector.store %weight_values1, %weight_stage_view[%weight_local_row, %weight_k_next] : vector<16xf16>, view<64x72xf16> + + %activation_values0, %activation_values1 = scf.if %valid_token -> (vector<4xf16>, vector<4xf16>) { + %bounded_token, %input_token_count = index.assume %activation_row, %bounded_token_count [lt(%activation_row, %bounded_token_count)] : index, index + %block_k_base = index.mul %quant_block, %c256 : index + %group_k_add = index.mul %quant_group, %c32 : index + %input_k0 = index.add %block_k_base, %group_k_add : index + %input_k = index.add %input_k0, %activation_k : index + %input_k_next = index.add %input_k0, %activation_k_next : index + %loaded0 = vector.load %input_view[%bounded_token, %input_k] : view<[%bounded_token_count]x[%bounded_input_size]xf32> -> vector<4xf32> + %loaded1 = vector.load %input_view[%bounded_token, %input_k_next] : view<[%bounded_token_count]x[%bounded_input_size]xf32> -> vector<4xf32> + %narrow0 = vector.fptrunc %loaded0 : vector<4xf32> to vector<4xf16> + %narrow1 = vector.fptrunc %loaded1 : vector<4xf32> to vector<4xf16> + scf.yield %narrow0, %narrow1 : vector<4xf16>, vector<4xf16> + } else { + %zeros0 = vector.constant 0.0 : vector<4xf16> + %zeros1 = vector.constant 0.0 : vector<4xf16> + scf.yield %zeros0, %zeros1 : vector<4xf16>, vector<4xf16> + } + vector.store %activation_values0, %activation_stage_physical_view[%activation_row, %activation_k] : vector<4xf16>, view<16x72xf16> + vector.store %activation_values1, %activation_stage_physical_view[%activation_row, %activation_k_next] : vector<4xf16>, view<16x72xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + %next = scf.for %k_half = [%c0 to %c64 step %c16](%half_acc = %acc : vector<8xf32>) -> (vector<8xf32>) unroll { + %lhs_fragment = vector.fragment.load %weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x72xf16> -> vector<16xf16> + %rhs_fragment = vector.fragment.load %activation_fragment_view[%k_half, %c0] shape [%k, %n] : view<64x16xf16, %activation_fragment_layout> -> vector<16xf16> + %half_next = vector.mma %lhs_fragment, %rhs_fragment, %half_acc : vector<16xf16>, vector<16xf16>, vector<8xf32> + scf.yield %half_next : vector<8xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %next : vector<8xf32> + } + scf.yield %block_result : vector<8xf32> + } + + template.apply<@ggml.mul_mat_q6_k_packed.publish_f32_wave32>(%result, %c0, %c0, %bounded_token_count, %bounded_output_size, %channel_tile_base, %subgroup_channel_add, %output_noalias, %output_accumulation) : (vector<8xf32>, index, index, index, index, index, index, buffer, index) + template.return +} + +kernel.def export("ggml_mul_mat_q6_k_i8_prepacked_f16_wmma") @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma(%token_count: index) { + %output_size = config.get @ggml.mul_mat_q6_k_packed.output_size : index + %c1 = index.constant 1 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + kernel.launch.config workgroups(%output_tiles, %c1, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) { + %input_size = config.get @ggml.mul_mat_q6_k_packed.input_size : index + %output_size = config.get @ggml.mul_mat_q6_k_packed.output_size : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 16)] : index + %channel_tile = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_q6_k_i8_prepacked.body>(%bounded_token_count, %input_size, %output_size, %channel_tile, %input, %weight, %output) : (index, index, index, index, buffer, buffer, buffer) + kernel.return +} + +func.def inline @ggml_expand_symmetric_i2_words(%q2_packed: vector<2xi32>) -> (vector<4xi32>) { + %q2_mask16 = vector.constant 65535 : vector<2xi32> + %q2_mask8 = vector.constant 16711935 : vector<2xi32> + %q2_mask4 = vector.constant 252645135 : vector<2xi32> + %q2_mask2 = vector.constant 858993459 : vector<2xi32> + %q2_sign = vector.constant 572662306 : vector<2xi32> + %q2_shift16 = vector.constant 16 : vector<2xi32> + %q2_shift8 = vector.constant 8 : vector<2xi32> + %q2_shift4 = vector.constant 4 : vector<2xi32> + %q2_shift2 = vector.constant 2 : vector<2xi32> + %q2_shift1 = vector.constant 1 : vector<2xi32> + %q2_lo0 = vector.andi %q2_packed, %q2_mask16 : vector<2xi32> + %q2_hi0 = vector.shrui %q2_packed, %q2_shift16 : vector<2xi32> + %q2_lo1s = vector.shli %q2_lo0, %q2_shift8 : vector<2xi32> + %q2_lo1o = vector.ori %q2_lo0, %q2_lo1s : vector<2xi32> + %q2_lo1 = vector.andi %q2_lo1o, %q2_mask8 : vector<2xi32> + %q2_lo2s = vector.shli %q2_lo1, %q2_shift4 : vector<2xi32> + %q2_lo2o = vector.ori %q2_lo1, %q2_lo2s : vector<2xi32> + %q2_lo2 = vector.andi %q2_lo2o, %q2_mask4 : vector<2xi32> + %q2_lo3s = vector.shli %q2_lo2, %q2_shift2 : vector<2xi32> + %q2_lo3o = vector.ori %q2_lo2, %q2_lo3s : vector<2xi32> + %q2_lo3 = vector.andi %q2_lo3o, %q2_mask2 : vector<2xi32> + %q2_losign = vector.andi %q2_lo3, %q2_sign : vector<2xi32> + %q2_losign2 = vector.shli %q2_losign, %q2_shift1 : vector<2xi32> + %q2_losign3 = vector.shli %q2_losign, %q2_shift2 : vector<2xi32> + %q2_lo4 = vector.ori %q2_lo3, %q2_losign2 : vector<2xi32> + %q2_lo = vector.ori %q2_lo4, %q2_losign3 : vector<2xi32> + %q2_hi1s = vector.shli %q2_hi0, %q2_shift8 : vector<2xi32> + %q2_hi1o = vector.ori %q2_hi0, %q2_hi1s : vector<2xi32> + %q2_hi1 = vector.andi %q2_hi1o, %q2_mask8 : vector<2xi32> + %q2_hi2s = vector.shli %q2_hi1, %q2_shift4 : vector<2xi32> + %q2_hi2o = vector.ori %q2_hi1, %q2_hi2s : vector<2xi32> + %q2_hi2 = vector.andi %q2_hi2o, %q2_mask4 : vector<2xi32> + %q2_hi3s = vector.shli %q2_hi2, %q2_shift2 : vector<2xi32> + %q2_hi3o = vector.ori %q2_hi2, %q2_hi3s : vector<2xi32> + %q2_hi3 = vector.andi %q2_hi3o, %q2_mask2 : vector<2xi32> + %q2_hisign = vector.andi %q2_hi3, %q2_sign : vector<2xi32> + %q2_hisign2 = vector.shli %q2_hisign, %q2_shift1 : vector<2xi32> + %q2_hisign3 = vector.shli %q2_hisign, %q2_shift2 : vector<2xi32> + %q2_hi4 = vector.ori %q2_hi3, %q2_hisign2 : vector<2xi32> + %q2_hi = vector.ori %q2_hi4, %q2_hisign3 : vector<2xi32> + %weight_word0 = vector.extract %q2_lo[0] : vector<2xi32> -> i32 + %weight_word1 = vector.extract %q2_hi[0] : vector<2xi32> -> i32 + %weight_word2 = vector.extract %q2_lo[1] : vector<2xi32> -> i32 + %weight_word3 = vector.extract %q2_hi[1] : vector<2xi32> -> i32 + %weight_words = vector.from_elements %weight_word0, %weight_word1, %weight_word2, %weight_word3 : vector<4xi32> + func.return %weight_words : vector<4xi32> +} + +func.def inline @ggml_symmetric_i2_dot_preloaded(%acc: f32, %weight_words: vector<4xi32>, %activation_words: vector<4xi32>, %combined_scale: f32) -> (f32) { + %zero_i32 = scalar.constant 0 : i32 + %zero_i32x4 = vector.constant 0 : vector<4xi32> + %dot_parts = vector.dot8i4 %weight_words, %activation_words, %zero_i32x4 : vector<4xi32> + %dot_i32 = vector.reduce %dot_parts, %zero_i32 : vector<4xi32>, i32 + %dot = scalar.sitofp %dot_i32 : i32 to f32 + %next = scalar.fmaf %dot, %combined_scale, %acc : f32 + func.return %next : f32 +} + +kernel.def export("ggml_select_symmetric_i4_k32_groups") @ggml_select_symmetric_i4_k32_groups() { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%c1, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%input: buffer, %selected_groups: buffer) { + %input_size0 = config.get @ggml.mul_mat_q6_k_shortlist.input_size : index + %input_size = index.assume %input_size0 [range(%input_size0, 256, 32768), mul(%input_size0, 256)] : index + %selected_count = config.get @ggml.mul_mat_q6_k_shortlist.selected_group_count : index + %tid0 = kernel.workitem.id : index + %tid = index.assume %tid0 [range(%tid0, 0, 255)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c32 = index.constant 32 : index + %base = index.constant 0 : offset + %zero_f32 = scalar.constant 0.0 : f32 + %negative_large = scalar.constant -3.4028234663852886e+38 : f32 + %zero_i32 = scalar.constant 0 : i32 + %one_i32 = scalar.constant 1 : i32 + %largest_i32 = scalar.constant 2147483647 : i32 + %group_count0 = index.div %input_size, %c32 : index + %group_count = index.assume %group_count0 [range(%group_count0, 8, 256)] : index + %valid = index.cmp ult, %tid, %group_count : index + %input_global = buffer.assume.memory_space %input : buffer + %groups_global = buffer.assume.memory_space %selected_groups : buffer + %input_noalias, %groups_noalias = buffer.assume.noalias %input_global, %groups_global : buffer, buffer + %input_view = buffer.view %input_noalias[%base] : buffer -> view<[%input_size]xf32> + %groups_view = buffer.view %groups_noalias[%base] : buffer -> view<[%selected_count]xi32> + %score = scf.if %valid -> (f32) { + %group_base = index.mul %tid, %c32 : index + %sum = scf.for %offset = [%c0 to %c32 step %c4](%iter = %zero_f32 : f32) -> (f32) unroll { + %element = index.add %group_base, %offset : index + %values = vector.load %input_view[%element] : view<[%input_size]xf32> -> vector<4xf32> + %absolute = vector.absf %values : vector<4xf32> + %partial = vector.reduce %absolute, %zero_f32 : vector<4xf32>, f32 + %next = scalar.addf %iter, %partial : f32 + scf.yield %next : f32 + } + scf.yield %sum : f32 + } else { + scf.yield %negative_large : f32 + } + %tid_i32 = index.cast %tid : index to i32 + %final = scf.for %rank = [%c0 to %selected_count step %c1](%remaining = %score : f32) -> (f32) { + %winner_score = kernel.workgroup.reduce %remaining : f32 + %matches = scalar.cmpf oeq, %remaining, %winner_score : f32 + %winner_candidate = scf.select %matches, %tid_i32, %largest_i32 : i32 + %winner_id = kernel.workgroup.reduce %winner_candidate : i32 + %is_winner = scalar.cmpi eq, %tid_i32, %winner_id : i32 + %next = scf.select %is_winner, %negative_large, %remaining : f32 + scf.yield %next : f32 + } + %matches_selected_sentinel = scalar.cmpf oeq, %final, %negative_large : f32 + %is_selected = scalar.andi %valid, %matches_selected_sentinel : i1 + %selected_i32 = scf.select %is_selected, %one_i32, %zero_i32 : i32 + %rank_i32 = kernel.workgroup.scan %selected_i32 {direction = forward, mode = exclusive} : i32 + scf.if %is_selected { + %rank0 = index.cast %rank_i32 : i32 to index + %bounded_rank = index.assume %rank0 [range(%rank0, 0, 255)] : index + %rank, %bounded_selected_count = index.assume %bounded_rank, %selected_count [lt(%bounded_rank, %selected_count)] : index, index + view.store %tid_i32, %groups_view[%rank] : i32, view<[%selected_count]xi32> + } + kernel.return +} + +kernel.def export("ggml_mul_mat_q6_k_symmetric_i2_scan_token1") @ggml_mul_mat_q6_k_symmetric_i2_scan_token1() { + %rows = config.get @ggml.mul_mat_q6_k_shortlist.output_size : index + %c1 = index.constant 1 : index + %c128 = index.constant 128 : index + %c511 = index.constant 511 : index + %c512 = index.constant 512 : index + %rows_up = index.add %rows, %c511 : index + %workgroups = index.div %rows_up, %c512 : index + kernel.launch.config workgroups(%workgroups, %c1, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%weight: buffer, %output: buffer, %qact: buffer, %scales: buffer, %selected_groups: buffer) { + %input_size0 = config.get @ggml.mul_mat_q6_k_shortlist.input_size : index + %output_size0 = config.get @ggml.mul_mat_q6_k_shortlist.output_size : index + %selected_count0 = config.get @ggml.mul_mat_q6_k_shortlist.selected_group_count : index + %input_size = index.assume %input_size0 [range(%input_size0, 256, 32768), mul(%input_size0, 256)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 64, 262144), mul(%output_size0, 64)] : index + %selected_count = index.assume %selected_count0 [range(%selected_count0, 1, 1024)] : index + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c128 = index.constant 128 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %zero_f32 = scalar.constant 0.0 : f32 + %workgroup = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 127)] : index + %cohort_base = index.mul %workgroup, %c128 : index + %cohort = index.add %cohort_base, %workitem : index + %row = index.mul %cohort, %c4 : index + %valid_cohort = index.cmp ult, %row, %output_size : index + %block_count = index.div %input_size, %c256 : index + %group_count = index.div %input_size, %c32 : index + %rows_plus = index.add %output_size, %c255 : index + %row_group_units = index.div %rows_plus, %c256 : index + %row_group = index.mul %row_group_units, %c32 : index + %row_group_id = index.div %row, %row_group : index + %row_lane = index.rem %row, %row_group : index + %cohort_lane = index.div %row_lane, %c4 : index + %payload_plane_bytes = index.mul %row_group, %c8 : index + %payload_block_bytes = index.mul %payload_plane_bytes, %c8 : index + %scale_cohorts = index.div %row_group, %c4 : index + %scale_plane_bytes = index.mul %scale_cohorts, %c16 : index + %block_bytes = index.add %payload_block_bytes, %scale_plane_bytes : index + %row_group_bytes = index.mul %block_count, %block_bytes : index + %group_base = index.mul %row_group_id, %row_group_bytes : index + %scale_region_add = index.mul %block_count, %payload_block_bytes : index + %scale_region = index.add %group_base, %scale_region_add : index + %weight_global = buffer.assume.memory_space %weight : buffer + %output_global = buffer.assume.memory_space %output : buffer + %qact_global = buffer.assume.memory_space %qact : buffer + %scales_global = buffer.assume.memory_space %scales : buffer + %groups_global = buffer.assume.memory_space %selected_groups : buffer + %weight_na, %output_na, %qact_na, %scales_na, %groups_na = buffer.assume.noalias %weight_global, %output_global, %qact_global, %scales_global, %groups_global : buffer, buffer, buffer, buffer, buffer + %groups_view = buffer.view %groups_na[%base] : buffer -> view<[%selected_count]xi32> + %qact_flat = buffer.view %qact_na[%base] : buffer -> view<1073741824xi8> + %activation_scale_view = buffer.view %scales_na[%base] : buffer -> view<33554432xf32> + %result0, %result1, %result2, %result3 = scf.if %valid_cohort -> (f32, f32, f32, f32) { + %acc0, %acc1, %acc2, %acc3 = scf.for %selected_index = [%c0 to %selected_count step %c1](%iter0 = %zero_f32 : f32, %iter1 = %zero_f32 : f32, %iter2 = %zero_f32 : f32, %iter3 = %zero_f32 : f32) -> (f32, f32, f32, f32) { + %group_i32 = view.load %groups_view[%selected_index] : view<[%selected_count]xi32> -> i32 + %group0 = index.cast %group_i32 : i32 to index + %group = index.assume %group0 [range(%group0, 0, 1023)] : index + %block = index.div %group, %c8 : index + %group_in_block = index.rem %group, %c8 : index + %block_payload_add = index.mul %block, %payload_block_bytes : index + %group_payload_add = index.mul %group_in_block, %payload_plane_bytes : index + %row_payload_add = index.mul %row_lane, %c8 : index + %payload_index0 = index.add %group_base, %block_payload_add : index + %payload_index1 = index.add %payload_index0, %group_payload_add : index + %payload_index = index.add %payload_index1, %row_payload_add : index + %payload_base = index.cast %payload_index : index to offset + %weight_view = buffer.view %weight_na[%payload_base] : buffer -> view<8xi32> + %q2_payload = vector.load %weight_view[%c0] : view<8xi32> -> vector<8xi32> + %q2_00 = vector.extract %q2_payload[0] : vector<8xi32> -> i32 + %q2_01 = vector.extract %q2_payload[1] : vector<8xi32> -> i32 + %q2_10 = vector.extract %q2_payload[2] : vector<8xi32> -> i32 + %q2_11 = vector.extract %q2_payload[3] : vector<8xi32> -> i32 + %q2_20 = vector.extract %q2_payload[4] : vector<8xi32> -> i32 + %q2_21 = vector.extract %q2_payload[5] : vector<8xi32> -> i32 + %q2_30 = vector.extract %q2_payload[6] : vector<8xi32> -> i32 + %q2_31 = vector.extract %q2_payload[7] : vector<8xi32> -> i32 + %q2_packed0 = vector.from_elements %q2_00, %q2_01 : vector<2xi32> + %q2_packed1 = vector.from_elements %q2_10, %q2_11 : vector<2xi32> + %q2_packed2 = vector.from_elements %q2_20, %q2_21 : vector<2xi32> + %q2_packed3 = vector.from_elements %q2_30, %q2_31 : vector<2xi32> + %weight_words0 = func.call @ggml_expand_symmetric_i2_words(%q2_packed0) : (vector<2xi32>) -> (vector<4xi32>) + %weight_words1 = func.call @ggml_expand_symmetric_i2_words(%q2_packed1) : (vector<2xi32>) -> (vector<4xi32>) + %weight_words2 = func.call @ggml_expand_symmetric_i2_words(%q2_packed2) : (vector<2xi32>) -> (vector<4xi32>) + %weight_words3 = func.call @ggml_expand_symmetric_i2_words(%q2_packed3) : (vector<2xi32>) -> (vector<4xi32>) + %block_scale_add = index.mul %block, %scale_plane_bytes : index + %cohort_scale_add = index.mul %cohort_lane, %c16 : index + %group_scale_add = index.mul %group_in_block, %c2 : index + %scale_index0 = index.add %scale_region, %block_scale_add : index + %scale_index1 = index.add %scale_index0, %cohort_scale_add : index + %scale_byte_index = index.add %scale_index1, %group_scale_add : index + %scale_byte_base = index.cast %scale_byte_index : index to offset + %weight_scale_view = buffer.view %weight_na[%scale_byte_base] : buffer -> view<1xf16> + %weight_scale_f16 = view.load %weight_scale_view[%c0] : view<1xf16> -> f16 + %weight_scale = scalar.extf %weight_scale_f16 : f16 to f32 + %activation_payload_add = index.mul %group, %c16 : index + %activation_payload_index = index.assume %activation_payload_add [range(%activation_payload_add, 0, 1073741808)] : index + %activation_bytes = vector.load %qact_flat[%activation_payload_index] : view<1073741824xi8> -> vector<16xi8> + %activation_words = vector.bitcast %activation_bytes : vector<16xi8> to vector<4xi32> + %activation_scale_index = index.assume %group [range(%group, 0, 33554431)] : index + %activation_scale = view.load %activation_scale_view[%activation_scale_index] : view<33554432xf32> -> f32 + %combined_scale = scalar.mulf %weight_scale, %activation_scale : f32 + %next0 = func.call @ggml_symmetric_i2_dot_preloaded(%iter0, %weight_words0, %activation_words, %combined_scale) : (f32, vector<4xi32>, vector<4xi32>, f32) -> (f32) + %next1 = func.call @ggml_symmetric_i2_dot_preloaded(%iter1, %weight_words1, %activation_words, %combined_scale) : (f32, vector<4xi32>, vector<4xi32>, f32) -> (f32) + %next2 = func.call @ggml_symmetric_i2_dot_preloaded(%iter2, %weight_words2, %activation_words, %combined_scale) : (f32, vector<4xi32>, vector<4xi32>, f32) -> (f32) + %next3 = func.call @ggml_symmetric_i2_dot_preloaded(%iter3, %weight_words3, %activation_words, %combined_scale) : (f32, vector<4xi32>, vector<4xi32>, f32) -> (f32) + scf.yield %next0, %next1, %next2, %next3 : f32, f32, f32, f32 + } + scf.yield %acc0, %acc1, %acc2, %acc3 : f32, f32, f32, f32 + } else { + scf.yield %zero_f32, %zero_f32, %zero_f32, %zero_f32 : f32, f32, f32, f32 + } + scf.if %valid_cohort { + %output_view = buffer.view %output_na[%base] : buffer -> view<[%output_size]xf32> + %results = vector.from_elements %result0, %result1, %result2, %result3 : vector<4xf32> + vector.store %results, %output_view[%row] : vector<4xf32>, view<[%output_size]xf32> + } + kernel.return +} + +kernel.def export("ggml_top_k8_f32_partitions_register") @ggml_top_k8_f32_partitions_register(%element_count: index) { + %c1 = index.constant 1 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%c128, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%element_count: index, %values: buffer, %partial_values: buffer, %partial_ids: buffer) { + %count = index.assume %element_count [range(%element_count, 64, 262144)] : index + %partition = kernel.workgroup.id : index + %tid = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %candidate_count = index.constant 8 : index + %c8 = index.constant 8 : index + %c127 = index.constant 127 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %base = index.constant 0 : offset + %negative_large = scalar.constant -3.4028234663852886e+38 : f32 + %largest_i32 = scalar.constant 2147483647 : i32 + %negative_values = vector.splat %negative_large : vector<8xf32> + %invalid_ids = vector.splat %largest_i32 : vector<8xi32> + %values_noalias, %partial_values_noalias, %partial_ids_noalias = buffer.assume.noalias %values, %partial_values, %partial_ids : buffer, buffer, buffer + %values_view = buffer.view %values_noalias[%base] : buffer -> view<[%count]xf32> + %partial_values_view = buffer.view %partial_values_noalias[%base] : buffer -> view<1024xf32> + %partial_ids_view = buffer.view %partial_ids_noalias[%base] : buffer -> view<1024xi32> + %rounded_count = index.add %count, %c127 : index + %partition_span = index.div %rounded_count, %c128 : index + %partition_begin = index.mul %partition, %partition_span : index + %partition_end_unbounded = index.add %partition_begin, %partition_span : index + %end_exceeds_count = index.cmp ugt, %partition_end_unbounded, %count : index + %partition_end = scf.select %end_exceeds_count, %count, %partition_end_unbounded : index + %thread_begin = index.add %partition_begin, %tid : index + %partial_base = index.mul %partition, %candidate_count : index + %loaded_values, %loaded_ids = scf.for %slot = [%c0 to %c8 step %c1](%local_values = %negative_values : vector<8xf32>, %local_ids = %invalid_ids : vector<8xi32>) -> (vector<8xf32>, vector<8xi32>) unroll { + %slot_stride = index.mul %slot, %c256 : index + %candidate_index = index.add %thread_begin, %slot_stride : index + %valid = index.cmp ult, %candidate_index, %partition_end : index + %candidate_value, %candidate_id = scf.if %valid -> (f32, i32) { + %bounded_index = index.assume %candidate_index [lt(%candidate_index, %count)] : index + %value = view.load %values_view[%bounded_index] : view<[%count]xf32> -> f32 + %id = index.cast %bounded_index : index to i32 + scf.yield %value, %id : f32, i32 + } else { + scf.yield %negative_large, %largest_i32 : f32, i32 + } + %next_values = vector.insert %candidate_value into %local_values[%slot] : f32, vector<8xf32> + %next_ids = vector.insert %candidate_id into %local_ids[%slot] : i32, vector<8xi32> + scf.yield %next_values, %next_ids : vector<8xf32>, vector<8xi32> + } + %remaining_final = scf.for %rank = [%c0 to %candidate_count step %c1](%remaining = %loaded_values : vector<8xf32>) -> (vector<8xf32>) unroll { + %local_value, %local_id = scf.for %slot = [%c0 to %c8 step %c1](%best_value = %negative_large : f32, %best_id = %largest_i32 : i32) -> (f32, i32) unroll { + %candidate_value = vector.extract %remaining[%slot] : vector<8xf32> -> f32 + %candidate_id = vector.extract %loaded_ids[%slot] : vector<8xi32> -> i32 + %is_greater = scalar.cmpf ogt, %candidate_value, %best_value : f32 + %is_equal = scalar.cmpf oeq, %candidate_value, %best_value : f32 + %is_lower_id = scalar.cmpi ult, %candidate_id, %best_id : i32 + %is_lower_tie = scalar.andi %is_equal, %is_lower_id : i1 + %is_better = scalar.ori %is_greater, %is_lower_tie : i1 + %next_value = scf.select %is_better, %candidate_value, %best_value : f32 + %next_id = scf.select %is_better, %candidate_id, %best_id : i32 + scf.yield %next_value, %next_id : f32, i32 + } + %winner_value = kernel.workgroup.reduce %local_value : f32 + %matches_winner = scalar.cmpf oeq, %local_value, %winner_value : f32 + %winner_id_candidate = scf.select %matches_winner, %local_id, %largest_i32 : i32 + %winner_id = kernel.workgroup.reduce %winner_id_candidate : i32 + %is_tid_zero = index.cmp eq, %tid, %c0 : index + scf.if %is_tid_zero { + %partial_index = index.add %partial_base, %rank : index + view.store %winner_value, %partial_values_view[%partial_index] : f32, view<1024xf32> + view.store %winner_id, %partial_ids_view[%partial_index] : i32, view<1024xi32> + } + %next_remaining = scf.for %slot = [%c0 to %c8 step %c1](%current = %remaining : vector<8xf32>) -> (vector<8xf32>) unroll { + %candidate_id = vector.extract %loaded_ids[%slot] : vector<8xi32> -> i32 + %is_winner = scalar.cmpi eq, %candidate_id, %winner_id : i32 + %old_value = vector.extract %current[%slot] : vector<8xf32> -> f32 + %next_value = scf.select %is_winner, %negative_large, %old_value : f32 + %updated = vector.insert %next_value into %current[%slot] : f32, vector<8xf32> + scf.yield %updated : vector<8xf32> + } + scf.yield %next_remaining : vector<8xf32> + } + kernel.return +} + +kernel.def export("ggml_top_k128_f32_reduce_gather_register") @ggml_top_k128_f32_reduce_gather_register(%element_count: index) { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%c1, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%element_count: index, %partial_values: buffer, %partial_ids: buffer, %candidate_output: buffer, %value_output: buffer) { + %count = index.assume %element_count [range(%element_count, 64, 262144)] : index + %tid = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %candidate_count = index.constant 128 : index + %c256 = index.constant 256 : index + %c1024 = index.constant 1024 : index + %base = index.constant 0 : offset + %negative_large = scalar.constant -3.4028234663852886e+38 : f32 + %largest_i32 = scalar.constant 2147483647 : i32 + %negative_values = vector.splat %negative_large : vector<4xf32> + %invalid_ids = vector.splat %largest_i32 : vector<4xi32> + %partial_values_noalias, %partial_ids_noalias, %candidate_output_noalias, %value_output_noalias = buffer.assume.noalias %partial_values, %partial_ids, %candidate_output, %value_output : buffer, buffer, buffer, buffer + %partial_values_view = buffer.view %partial_values_noalias[%base] : buffer -> view<1024xf32> + %partial_ids_view = buffer.view %partial_ids_noalias[%base] : buffer -> view<1024xi32> + %candidate_output_view = buffer.view %candidate_output_noalias[%base] : buffer -> view<128xi32> + %value_output_view = buffer.view %value_output_noalias[%base] : buffer -> view<128xf32> + %loaded_values, %loaded_ids = scf.for %slot = [%c0 to %c4 step %c1](%values = %negative_values : vector<4xf32>, %ids = %invalid_ids : vector<4xi32>) -> (vector<4xf32>, vector<4xi32>) unroll { + %slot_stride = index.mul %slot, %c256 : index + %partial_index = index.add %tid, %slot_stride : index + %valid = index.cmp ult, %partial_index, %c1024 : index + %candidate_value, %candidate_id = scf.if %valid -> (f32, i32) { + %value = view.load %partial_values_view[%partial_index] : view<1024xf32> -> f32 + %id = view.load %partial_ids_view[%partial_index] : view<1024xi32> -> i32 + scf.yield %value, %id : f32, i32 + } else { + scf.yield %negative_large, %largest_i32 : f32, i32 + } + %next_values = vector.insert %candidate_value into %values[%slot] : f32, vector<4xf32> + %next_ids = vector.insert %candidate_id into %ids[%slot] : i32, vector<4xi32> + scf.yield %next_values, %next_ids : vector<4xf32>, vector<4xi32> + } + %remaining_final = scf.for %rank = [%c0 to %candidate_count step %c1](%remaining = %loaded_values : vector<4xf32>) -> (vector<4xf32>) { + %local_value, %local_id = scf.for %slot = [%c0 to %c4 step %c1](%best_value = %negative_large : f32, %best_id = %largest_i32 : i32) -> (f32, i32) unroll { + %candidate_value = vector.extract %remaining[%slot] : vector<4xf32> -> f32 + %candidate_id = vector.extract %loaded_ids[%slot] : vector<4xi32> -> i32 + %is_greater = scalar.cmpf ogt, %candidate_value, %best_value : f32 + %is_equal = scalar.cmpf oeq, %candidate_value, %best_value : f32 + %is_lower_id = scalar.cmpi ult, %candidate_id, %best_id : i32 + %is_lower_tie = scalar.andi %is_equal, %is_lower_id : i1 + %is_better = scalar.ori %is_greater, %is_lower_tie : i1 + %next_value = scf.select %is_better, %candidate_value, %best_value : f32 + %next_id = scf.select %is_better, %candidate_id, %best_id : i32 + scf.yield %next_value, %next_id : f32, i32 + } + %winner_value = kernel.workgroup.reduce %local_value : f32 + %matches_winner = scalar.cmpf oeq, %local_value, %winner_value : f32 + %winner_id_candidate = scf.select %matches_winner, %local_id, %largest_i32 : i32 + %winner_id = kernel.workgroup.reduce %winner_id_candidate : i32 + %is_tid_zero = index.cmp eq, %tid, %c0 : index + scf.if %is_tid_zero { + view.store %winner_id, %candidate_output_view[%rank] : i32, view<128xi32> + view.store %winner_value, %value_output_view[%rank] : f32, view<128xf32> + } + %next_remaining = scf.for %slot = [%c0 to %c4 step %c1](%values = %remaining : vector<4xf32>) -> (vector<4xf32>) unroll { + %candidate_id = vector.extract %loaded_ids[%slot] : vector<4xi32> -> i32 + %is_winner = scalar.cmpi eq, %candidate_id, %winner_id : i32 + %old_value = vector.extract %values[%slot] : vector<4xf32> -> f32 + %next_value = scf.select %is_winner, %negative_large, %old_value : f32 + %updated = vector.insert %next_value into %values[%slot] : f32, vector<4xf32> + scf.yield %updated : vector<4xf32> + } + scf.yield %next_remaining : vector<4xf32> + } + kernel.return +} + +kernel.def export("ggml_fill_negative_f32") @ggml_fill_negative_f32(%element_count: index) { + %c1 = index.constant 1 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %rounded_count = index.add %element_count, %c255 : index + %workgroups = index.div %rounded_count, %c256 : index + kernel.launch.config workgroups(%workgroups, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%element_count: index, %output: buffer) { + %count = index.assume %element_count [range(%element_count, 1, 262144)] : index + %workgroup = kernel.workgroup.id : index + %tid0 = kernel.workitem.id : index + %tid = index.assume %tid0 [range(%tid0, 0, 255)] : index + %c256 = index.constant 256 : index + %base = index.constant 0 : offset + %negative_large = scalar.constant -3.4028234663852886e+38 : f32 + %workgroup_base = index.mul %workgroup, %c256 : index + %index0 = index.add %workgroup_base, %tid : index + %index = index.assume %index0 [range(%index0, 0, 262399)] : index + %valid = index.cmp ult, %index, %count : index + %output_view = buffer.view %output[%base] : buffer -> view<[%count]xf32> + scf.if %valid { + %bounded_index = index.assume %index [lt(%index, %count)] : index + view.store %negative_large, %output_view[%bounded_index] : f32, view<[%count]xf32> + } + kernel.return +} + +template.def<@ggml.mul_mat_q6_k_packed.selected_refine.body> device @ggml_mul_mat_q6_k_packed_selected_refine_body(%token_count: index, %input_size0: index, %output_size0: index, %scale_row_layout: index, %candidate_count0: index, %input: buffer, %weight: buffer, %candidates: buffer, %output: buffer) { + %output_accumulation = config.get @ggml.mul_mat_q6_k_packed.output_accumulation : index + %weight_offset_index = config.get @ggml.mul_mat_q6_k_packed.weight_offset : index + %weight_offset = index.cast %weight_offset_index : index to offset + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 16)] : index + %bounded_input_size = index.assume %input_size0 [range(%input_size0, 256, 32768), mul(%input_size0, 256)] : index + %bounded_output_size = index.assume %output_size0 [range(%output_size0, 64, 262144), mul(%output_size0, 64)] : index + %bounded_candidate_count = index.assume %candidate_count0 [range(%candidate_count0, 1, 64)] : index + %workitem0 = kernel.workitem.id : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 127)] : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 3)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c72 = index.constant 72 : index + %c256 = index.constant 256 : index + %base = index.constant 0 : offset + %c128_bytes = index.constant 128 : offset + %c1152_bytes = index.constant 1152 : offset + %c13440_bytes = index.constant 13440 : offset + %weight_stage_bytes = index.constant 9216 : offset + %activation_stage_bytes = index.constant 2304 : offset + %result_stage_bytes = index.constant 2048 : offset + %zero_accumulator = vector.constant 0.0 : vector<16xf16> + %empty_words = vector.constant 0 : vector<4xi32> + %zero_scale_bytes = vector.constant 0 : vector<8xi8> + %nibble_mask = vector.constant 252645135 : vector<4xi32> + %high0_mask = vector.constant 50529027 : vector<4xi32> + %high1_mask = vector.constant 202116108 : vector<4xi32> + %high2_mask = vector.constant 808464432 : vector<4xi32> + %high3_mask = vector.constant -1061109568 : vector<4xi32> + %shift2 = vector.constant 2 : vector<4xi32> + %shift4 = vector.constant 4 : vector<4xi32> + %negative32_f32 = vector.constant -32.0 : vector<16xf32> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %quant_block_count = index.div %bounded_input_size, %c256 : index + %input_noalias, %weight_noalias, %candidates_noalias, %output_noalias = buffer.assume.noalias %input, %weight, %candidates, %output : buffer, buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%base] : buffer -> view<[%bounded_token_count]x[%bounded_input_size]xf32> + %candidates_view = buffer.view %candidates_noalias[%base] : buffer -> view<64xi32> + %weight_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %activation_stage = buffer.alloca align(16) %activation_stage_bytes : buffer + %result_stage = buffer.alloca align(16) %result_stage_bytes : buffer + %weight_stage_view = buffer.view %weight_stage[%base] : buffer -> view<64x72xf16> + %activation_stage_physical_view = buffer.view %activation_stage[%base] : buffer -> view<16x72xf16> + %activation_fragment_layout = encoding.layout.strided [1, %c72] : encoding + %activation_fragment_view = buffer.view %activation_stage[%base] : buffer -> view<64x16xf16, %activation_fragment_layout> + %weight_half = index.rem %workitem, %c2 : index + %weight_local_row0 = index.div %workitem, %c2 : index + %weight_local_row = index.assume %weight_local_row0 [range(%weight_local_row0, 0, 63)] : index + %weight_k = index.mul %weight_half, %c16 : index + %weight_k_next = index.add %weight_k, %c32 : index + %valid_candidate = index.cmp ult, %weight_local_row, %bounded_candidate_count : index + %channel0 = scf.if %valid_candidate -> (index) { + %candidate_i32 = view.load %candidates_view[%weight_local_row] : view<64xi32> -> i32 + %candidate_index0 = index.cast %candidate_i32 : i32 to index + %candidate_index = index.assume %candidate_index0 [range(%candidate_index0, 0, 262143)] : index + scf.yield %candidate_index : index + } else { + scf.yield %c0 : index + } + %channel = index.assume %channel0 [range(%channel0, 0, 262143)] : index + %channel_in_range = index.cmp ult, %channel, %bounded_output_size : index + %valid_channel = scalar.andi %valid_candidate, %channel_in_range : i1 + %selected_channel_tile = index.div %channel, %c64 : index + %selected_weight_local_row = index.rem %channel, %c64 : index + %use_scale_row_layout = index.cmp eq, %scale_row_layout, %c1 : index + %activation_row0 = index.div %workitem, %c8 : index + %activation_row = index.assume %activation_row0 [range(%activation_row0, 0, 15)] : index + %activation_packet = index.rem %workitem, %c8 : index + %activation_k = index.mul %activation_packet, %c4 : index + %activation_k_next = index.add %activation_k, %c32 : index + %valid_token = index.cmp ult, %activation_row, %bounded_token_count : index + %subgroup_channel_add = index.mul %subgroup, %c16 : index + %init = vector.fragment %zero_accumulator shape [%m, %n] : vector<16xf16> + %result = scf.for %quant_block = [%c0 to %quant_block_count step %c1](%block_acc = %init : vector<16xf16>) -> (vector<16xf16>) { + %row_group_block0 = index.mul %selected_channel_tile, %quant_block_count : index + %row_group_block = index.add %row_group_block0, %quant_block : index + %tile_byte_relative = index.scale %row_group_block, %c13440_bytes : index, offset -> offset + %tile_byte_base = index.add %weight_offset, %tile_byte_relative : offset + %scale_byte_base = index.add %tile_byte_base, %c128_bytes : offset + %raw_byte_base = index.add %tile_byte_base, %c1152_bytes : offset + %d_view = buffer.view %weight_noalias[%tile_byte_base] : buffer -> view<64xf16> + %scale_view = buffer.view %weight_noalias[%scale_byte_base] : buffer -> view<8x64x2xi8> + %scale_row_view = buffer.view %weight_noalias[%scale_byte_base] : buffer -> view<64x2x8xi8> + %raw_view = buffer.view %weight_noalias[%raw_byte_base] : buffer -> view<2x3x64x32xi8> + %d = scf.if %valid_channel -> (f32) { + %d_f16 = view.load %d_view[%selected_weight_local_row] : view<64xf16> -> f16 + %d_f32 = scalar.extf %d_f16 : f16 to f32 + scf.yield %d_f32 : f32 + } else { + %zero_f32 = scalar.constant 0.0 : f32 + scf.yield %zero_f32 : f32 + } + %load_scale_row = scalar.andi %valid_channel, %use_scale_row_layout : i1 + %scale_bytes = scf.if %load_scale_row -> (vector<8xi8>) { + %loaded_scale_bytes = vector.load %scale_row_view[%selected_weight_local_row, %weight_half, %c0] : view<64x2x8xi8> -> vector<8xi8> + scf.yield %loaded_scale_bytes : vector<8xi8> + } else { + scf.yield %zero_scale_bytes : vector<8xi8> + } + %block_result, %final_ql0_words, %final_ql1_words, %final_qh_words = scf.for %quant_group = [%c0 to %c8 step %c2](%acc = %block_acc : vector<16xf16>, %cached_ql0_words = %empty_words : vector<4xi32>, %cached_ql1_words = %empty_words : vector<4xi32>, %cached_qh_words = %empty_words : vector<4xi32>) -> (vector<16xf16>, vector<4xi32>, vector<4xi32>, vector<4xi32>) unroll { + %quant_group_next = index.add %quant_group, %c1 : index + %half128 = index.div %quant_group, %c4 : index + %group_in_half = index.rem %quant_group, %c4 : index + %is_low_pair = index.cmp eq, %group_in_half, %c0 : index + %is_first_group = index.cmp eq, %quant_group, %c0 : index + %prefetch_second_half = index.cmp eq, %quant_group, %c2 : index + %weight_values0, %weight_values1, %next_ql0_words, %next_ql1_words, %next_qh_words = scf.if %valid_channel -> (vector<16xf16>, vector<16xf16>, vector<4xi32>, vector<4xi32>, vector<4xi32>) { + %ql0_words, %ql1_words, %qh_words = scf.if %is_first_group -> (vector<4xi32>, vector<4xi32>, vector<4xi32>) { + %ql0_bytes = vector.load %raw_view[%half128, %c0, %selected_weight_local_row, %weight_k] : view<2x3x64x32xi8> -> vector<16xi8> + %ql1_bytes = vector.load %raw_view[%half128, %c1, %selected_weight_local_row, %weight_k] : view<2x3x64x32xi8> -> vector<16xi8> + %qh_bytes = vector.load %raw_view[%half128, %c2, %selected_weight_local_row, %weight_k] : view<2x3x64x32xi8> -> vector<16xi8> + %loaded_ql0_words = vector.bitcast %ql0_bytes : vector<16xi8> to vector<4xi32> + %loaded_ql1_words = vector.bitcast %ql1_bytes : vector<16xi8> to vector<4xi32> + %loaded_qh_words = vector.bitcast %qh_bytes : vector<16xi8> to vector<4xi32> + scf.yield %loaded_ql0_words, %loaded_ql1_words, %loaded_qh_words : vector<4xi32>, vector<4xi32>, vector<4xi32> + } else { + scf.yield %cached_ql0_words, %cached_ql1_words, %cached_qh_words : vector<4xi32>, vector<4xi32>, vector<4xi32> + } + %codes0_words, %codes1_words = scf.if %is_low_pair -> (vector<4xi32>, vector<4xi32>) { + %ql0_low = vector.andi %ql0_words, %nibble_mask : vector<4xi32> + %ql1_low = vector.andi %ql1_words, %nibble_mask : vector<4xi32> + %qh0_bits0 = vector.andi %qh_words, %high0_mask : vector<4xi32> + %qh1_bits0 = vector.andi %qh_words, %high1_mask : vector<4xi32> + %qh0_bits = vector.shli %qh0_bits0, %shift4 : vector<4xi32> + %qh1_bits = vector.shli %qh1_bits0, %shift2 : vector<4xi32> + %codes0 = vector.ori %ql0_low, %qh0_bits : vector<4xi32> + %codes1 = vector.ori %ql1_low, %qh1_bits : vector<4xi32> + scf.yield %codes0, %codes1 : vector<4xi32>, vector<4xi32> + } else { + %ql0_high0 = vector.shrui %ql0_words, %shift4 : vector<4xi32> + %ql1_high0 = vector.shrui %ql1_words, %shift4 : vector<4xi32> + %ql0_high = vector.andi %ql0_high0, %nibble_mask : vector<4xi32> + %ql1_high = vector.andi %ql1_high0, %nibble_mask : vector<4xi32> + %qh2_bits = vector.andi %qh_words, %high2_mask : vector<4xi32> + %qh3_bits0 = vector.andi %qh_words, %high3_mask : vector<4xi32> + %qh3_bits = vector.shrui %qh3_bits0, %shift2 : vector<4xi32> + %codes0 = vector.ori %ql0_high, %qh2_bits : vector<4xi32> + %codes1 = vector.ori %ql1_high, %qh3_bits : vector<4xi32> + scf.yield %codes0, %codes1 : vector<4xi32>, vector<4xi32> + } + %codes0_u8 = vector.bitcast %codes0_words : vector<4xi32> to vector<16xi8> + %codes1_u8 = vector.bitcast %codes1_words : vector<4xi32> to vector<16xi8> + %codes0_unsigned_f32 = vector.uitofp %codes0_u8 : vector<16xi8> to vector<16xf32> + %codes1_unsigned_f32 = vector.uitofp %codes1_u8 : vector<16xi8> to vector<16xf32> + %codes0_f32 = vector.addf %codes0_unsigned_f32, %negative32_f32 : vector<16xf32> + %codes1_f32 = vector.addf %codes1_unsigned_f32, %negative32_f32 : vector<16xf32> + %scale0_i8, %scale1_i8 = scf.if %use_scale_row_layout -> (i8, i8) { + %row_scale0_i8 = vector.extract %scale_bytes[%quant_group] : vector<8xi8> -> i8 + %row_scale1_i8 = vector.extract %scale_bytes[%quant_group_next] : vector<8xi8> -> i8 + scf.yield %row_scale0_i8, %row_scale1_i8 : i8, i8 + } else { + %group_scale0_i8 = view.load %scale_view[%quant_group, %selected_weight_local_row, %weight_half] : view<8x64x2xi8> -> i8 + %group_scale1_i8 = view.load %scale_view[%quant_group_next, %selected_weight_local_row, %weight_half] : view<8x64x2xi8> -> i8 + scf.yield %group_scale0_i8, %group_scale1_i8 : i8, i8 + } + %scale0 = scalar.sitofp %scale0_i8 : i8 to f32 + %scale1 = scalar.sitofp %scale1_i8 : i8 to f32 + %combined0 = scalar.mulf %scale0, %d : f32 + %combined1 = scalar.mulf %scale1, %d : f32 + %combined0_vector = vector.splat %combined0 : vector<16xf32> + %combined1_vector = vector.splat %combined1 : vector<16xf32> + %values0_f32 = vector.mulf %codes0_f32, %combined0_vector : vector<16xf32> + %values1_f32 = vector.mulf %codes1_f32, %combined1_vector : vector<16xf32> + %values0 = vector.fptrunc %values0_f32 : vector<16xf32> to vector<16xf16> + %values1 = vector.fptrunc %values1_f32 : vector<16xf32> to vector<16xf16> + %prefetched_ql0_words, %prefetched_ql1_words, %prefetched_qh_words = scf.if %prefetch_second_half -> (vector<4xi32>, vector<4xi32>, vector<4xi32>) { + %next_ql0_bytes = vector.load %raw_view[%c1, %c0, %selected_weight_local_row, %weight_k] : view<2x3x64x32xi8> -> vector<16xi8> + %next_ql1_bytes = vector.load %raw_view[%c1, %c1, %selected_weight_local_row, %weight_k] : view<2x3x64x32xi8> -> vector<16xi8> + %next_qh_bytes = vector.load %raw_view[%c1, %c2, %selected_weight_local_row, %weight_k] : view<2x3x64x32xi8> -> vector<16xi8> + %next_ql0 = vector.bitcast %next_ql0_bytes : vector<16xi8> to vector<4xi32> + %next_ql1 = vector.bitcast %next_ql1_bytes : vector<16xi8> to vector<4xi32> + %next_qh = vector.bitcast %next_qh_bytes : vector<16xi8> to vector<4xi32> + scf.yield %next_ql0, %next_ql1, %next_qh : vector<4xi32>, vector<4xi32>, vector<4xi32> + } else { + scf.yield %ql0_words, %ql1_words, %qh_words : vector<4xi32>, vector<4xi32>, vector<4xi32> + } + scf.yield %values0, %values1, %prefetched_ql0_words, %prefetched_ql1_words, %prefetched_qh_words : vector<16xf16>, vector<16xf16>, vector<4xi32>, vector<4xi32>, vector<4xi32> + } else { + %zeros0 = vector.constant 0.0 : vector<16xf16> + %zeros1 = vector.constant 0.0 : vector<16xf16> + scf.yield %zeros0, %zeros1, %cached_ql0_words, %cached_ql1_words, %cached_qh_words : vector<16xf16>, vector<16xf16>, vector<4xi32>, vector<4xi32>, vector<4xi32> + } + vector.store %weight_values0, %weight_stage_view[%weight_local_row, %weight_k] : vector<16xf16>, view<64x72xf16> + vector.store %weight_values1, %weight_stage_view[%weight_local_row, %weight_k_next] : vector<16xf16>, view<64x72xf16> + %activation_values0, %activation_values1 = scf.if %valid_token -> (vector<4xf16>, vector<4xf16>) { + %bounded_token, %input_token_count = index.assume %activation_row, %bounded_token_count [lt(%activation_row, %bounded_token_count)] : index, index + %block_k_base = index.mul %quant_block, %c256 : index + %group_k_add = index.mul %quant_group, %c32 : index + %input_k0 = index.add %block_k_base, %group_k_add : index + %input_k = index.add %input_k0, %activation_k : index + %input_k_next = index.add %input_k0, %activation_k_next : index + %loaded0 = vector.load %input_view[%bounded_token, %input_k] : view<[%bounded_token_count]x[%bounded_input_size]xf32> -> vector<4xf32> + %loaded1 = vector.load %input_view[%bounded_token, %input_k_next] : view<[%bounded_token_count]x[%bounded_input_size]xf32> -> vector<4xf32> + %narrow0 = vector.fptrunc %loaded0 : vector<4xf32> to vector<4xf16> + %narrow1 = vector.fptrunc %loaded1 : vector<4xf32> to vector<4xf16> + scf.yield %narrow0, %narrow1 : vector<4xf16>, vector<4xf16> + } else { + %zeros0 = vector.constant 0.0 : vector<4xf16> + %zeros1 = vector.constant 0.0 : vector<4xf16> + scf.yield %zeros0, %zeros1 : vector<4xf16>, vector<4xf16> + } + vector.store %activation_values0, %activation_stage_physical_view[%activation_row, %activation_k] : vector<4xf16>, view<16x72xf16> + vector.store %activation_values1, %activation_stage_physical_view[%activation_row, %activation_k_next] : vector<4xf16>, view<16x72xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + %next = scf.for %k_half = [%c0 to %c64 step %c16](%half_acc = %acc : vector<16xf16>) -> (vector<16xf16>) unroll { + %lhs = vector.fragment.load %weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x72xf16> -> vector<16xf16> + %rhs = vector.fragment.load %activation_fragment_view[%k_half, %c0] shape [%k, %n] : view<64x16xf16, %activation_fragment_layout> -> vector<16xf16> + %half_next = vector.mma %lhs, %rhs, %half_acc : vector<16xf16>, vector<16xf16>, vector<16xf16> + scf.yield %half_next : vector<16xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %next, %next_ql0_words, %next_ql1_words, %next_qh_words : vector<16xf16>, vector<4xi32>, vector<4xi32>, vector<4xi32> + } + scf.yield %block_result : vector<16xf16> + } + template.apply<@ggml.mul_mat_q6_k_packed.publish_f16_wave32>(%result, %c0, %c0, %bounded_token_count, %c64, %c0, %subgroup_channel_add, %result_stage, %output_noalias, %output_accumulation) : (vector<16xf16>, index, index, index, index, index, index, buffer, buffer, index) + template.return +} + +kernel.def export("ggml_mul_mat_q6_k_packed_selected_refine_token1") @ggml_mul_mat_q6_k_packed_selected_refine_token1(%token_count: index, %candidate_count: index) { + %c1 = index.constant 1 : index + %c128 = index.constant 128 : index + kernel.launch.config workgroups(%c1, %c1, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %candidate_count: index, %input: buffer, %weight: buffer, %candidates: buffer, %exact_output: buffer, %output: buffer) { + %input_size = config.get @ggml.mul_mat_q6_k_packed.input_size : index + %output_size0 = config.get @ggml.mul_mat_q6_k_packed.output_size : index + %output_size = index.assume %output_size0 [range(%output_size0, 64, 262144), mul(%output_size0, 64)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %base = index.constant 0 : offset + %scale_row_layout = index.constant 1 : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 1)] : index + %bounded_candidate_count = index.assume %candidate_count [range(%candidate_count, 1, 64)] : index + template.apply<@ggml.mul_mat_q6_k_packed.selected_refine.body>(%bounded_token_count, %input_size, %output_size, %scale_row_layout, %bounded_candidate_count, %input, %weight, %candidates, %exact_output) : (index, index, index, index, index, buffer, buffer, buffer, buffer) + kernel.barrier scope(workgroup) ordering(acq_rel) + %tid0 = kernel.workitem.id : index + %tid = index.assume %tid0 [range(%tid0, 0, 127)] : index + %publishes = index.cmp ult, %tid, %bounded_candidate_count : index + %candidates_noalias, %exact_noalias, %output_noalias = buffer.assume.noalias %candidates, %exact_output, %output : buffer, buffer, buffer + %candidate_view = buffer.view %candidates_noalias[%base] : buffer -> view<64xi32> + %exact_view = buffer.view %exact_noalias[%base] : buffer -> view<64xf32> + scf.if %publishes { + %candidate_i32 = view.load %candidate_view[%tid] : view<64xi32> -> i32 + %candidate0 = index.cast %candidate_i32 : i32 to index + %candidate1 = index.assume %candidate0 [range(%candidate0, 0, 262143)] : index + %candidate, %bounded_output_size = index.assume %candidate1, %output_size [lt(%candidate1, %output_size)] : index, index + %output_view = buffer.view %output_noalias[%base] : buffer -> view<[%bounded_output_size]xf32> + %value = view.load %exact_view[%tid] : view<64xf32> -> f32 + view.store %value, %output_view[%candidate] : f32, view<[%bounded_output_size]xf32> + } + kernel.return +} + +template.decl @ggml.mul_mat_q6_k_packed.publish_f32_wave32(%result: vector<8xf32>, %token_offset: index, %token_tile_base: index, %token_count: index, %output_size: index, %channel_tile_base: index, %subgroup_channel_add: index, %output: buffer, %output_accumulation: index) + +template.def<@ggml.mul_mat_q6_k_packed.publish_f32_wave32> device @ggml_mul_mat_q6_k_packed_publish_f32_wave32(%result: vector<8xf32>, %token_offset: index, %token_tile_base: index, %token_count: index, %output_size: index, %channel_tile_base: index, %subgroup_channel_add: index, %output: buffer, %output_accumulation: index) { + %base = index.constant 0 : offset + %c1 = index.constant 1 : index + %c16 = index.constant 16 : index + %token = index.add %token_tile_base, %token_offset : index + %channel = index.add %channel_tile_base, %subgroup_channel_add : index + %token_end = index.add %token, %c16 : index + %channel_end = index.add %channel, %c16 : index + %valid_token = index.cmp ule, %token_end, %token_count : index + %valid_channel = index.cmp ule, %channel_end, %output_size : index + %valid = scalar.andi %valid_token, %valid_channel : i1 + scf.if %valid { + %layout = encoding.layout.strided [1, %output_size] : encoding + %dst = buffer.view %output[%base] : buffer -> view<[%output_size]x[%token_count]xf32, %layout> + %accumulates = index.cmp eq, %output_accumulation, %c1 : index + %published = scf.if %accumulates -> (vector<8xf32>) { + %old = vector.fragment.load %dst[%channel, %token] shape [%c16, %c16] : view<[%output_size]x[%token_count]xf32, %layout> -> vector<8xf32> + %sum = vector.addf %old, %result : vector<8xf32> + scf.yield %sum : vector<8xf32> + } else { + scf.yield %result : vector<8xf32> + } + vector.fragment.store %published, %dst[%channel, %token] shape [%c16, %c16] : vector<8xf32>, view<[%output_size]x[%token_count]xf32, %layout> + } else { + // Native wave32 F32 result: lane % 16 selects the token; the high + // lane bit selects even/odd channels in the eight-value fragment. + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c8 = index.constant 8 : index + %lane = kernel.subgroup.lane.id : index + %lane_token = index.rem %lane, %c16 : index + %lane_channel = index.div %lane, %c16 : index + %partial_token = index.add %token, %lane_token : index + %partial_channel_base = index.add %channel, %lane_channel : index + %token_valid = index.cmp ult, %partial_token, %token_count : index + %accumulates = index.cmp eq, %output_accumulation, %c1 : index + scf.if %token_valid { + %bounded_token, %output_tokens = index.assume %partial_token, %token_count [lt(%partial_token, %token_count)] : index, index + %dst = buffer.view %output[%base] : buffer -> view<[%output_tokens]x[%output_size]xf32> + scf.for %element = [%c0 to %c8 step %c1] unroll { + %channel_step = index.mul %element, %c2 : index + %partial_channel = index.add %partial_channel_base, %channel_step : index + %channel_valid = index.cmp ult, %partial_channel, %output_size : index + scf.if %channel_valid { + %bounded_channel, %output_channels = index.assume %partial_channel, %output_size [lt(%partial_channel, %output_size)] : index, index + %value = vector.extract %result[%element] : vector<8xf32> -> f32 + %published = scf.if %accumulates -> (f32) { + %old = view.load %dst[%bounded_token, %bounded_channel] : view<[%output_tokens]x[%output_size]xf32> -> f32 + %sum = scalar.addf %value, %old : f32 + scf.yield %sum : f32 + } else { + scf.yield %value : f32 + } + view.store %published, %dst[%bounded_token, %bounded_channel] : f32, view<[%output_tokens]x[%output_size]xf32> + } + } + } + } + template.return +} + +template.def<@ggml.mul_mat_q6_k_packed.prefill_wave32.body> device @ggml_mul_mat_q6_k_packed_prefill_wave32_body(%input_is_f16: i1, %token_count: index, %input_size0: index, %output_size0: index, %channel_tile: index, %token_tile: index, %input: buffer, %weight: buffer, %output: buffer) { + %output_accumulation = config.get @ggml.mul_mat_q6_k_packed.output_accumulation : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_input_size = index.assume %input_size0 [range(%input_size0, 256, 32768), mul(%input_size0, 256)] : index + %bounded_output_size = index.assume %output_size0 [range(%output_size0, 1, 262144)] : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c72 = index.constant 72 : index + %c48 = index.constant 48 : index + %c64 = index.constant 64 : index + %c80 = index.constant 80 : index + %c96 = index.constant 96 : index + %c112 = index.constant 112 : index + %c128 = index.constant 128 : index + %c210 = index.constant 210 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %weight_stage_bytes = index.constant 18432 : offset + %activation_stage_bytes = index.constant 9216 : offset + %c0_f16x4 = vector.constant 0.0 : vector<16xf16> + %c0_f16x8 = vector.constant 0.0 : vector<8xf16> + %zero_accumulator = vector.constant 0.0 : vector<8xf32> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %quant_block_count = index.div %bounded_input_size, %c256 : index + %weight_row_bytes = index.mul %quant_block_count, %c210 : index + %input_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input, %weight, %output : buffer, buffer, buffer + %input_f32_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_input_size]xf32> + %input_f16_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_input_size]xf16> + %weight_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %activation_stage = buffer.alloca align(16) %activation_stage_bytes : buffer + %weight_stage_view = buffer.view %weight_stage[%c0_offset] : buffer -> view<128x72xf16> + %activation_stage_physical_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<64x72xf16> + %activation_fragment_layout = encoding.layout.strided [1, %c72] : encoding + %activation_fragment_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<64x64xf16, %activation_fragment_layout> + %channel_tile_base = index.mul %channel_tile, %c128 : index + %token_tile_base = index.mul %token_tile, %c128 : index + %second_token_tile_base = index.add %token_tile_base, %c64 : index + %load_half = index.rem %workitem, %c2 : index + %load_packet = index.mul %load_half, %c4 : index + %load_k = index.mul %load_packet, %c4 : index + %load_row0 = index.div %workitem, %c2 : index + %load_row = index.assume %load_row0 [range(%load_row0, 0, 127)] : index + %activation_load_packet = index.rem %workitem, %c4 : index + %activation_load_k = index.mul %activation_load_packet, %c8 : index + %activation_load_row0 = index.div %workitem, %c4 : index + %activation_load_row = index.assume %activation_load_row0 [range(%activation_load_row0, 0, 63)] : index + %subgroup_channel_add = index.mul %subgroup, %c16 : index + %rhs_column0 = index.add %c0, %c0 : index + %rhs_column1 = index.add %c16, %c0 : index + %rhs_column2 = index.add %c32, %c0 : index + %rhs_column3 = index.add %c48, %c0 : index + %init00 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + %init01 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + %init02 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + %init03 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + %init10 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + %init11 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + %init12 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + %init13 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf32> + %result00, %result01, %result02, %result03, %result10, %result11, %result12, %result13 = scf.for %quant_block = [%c0 to %quant_block_count step %c1](%block_acc00 = %init00 : vector<8xf32>, %block_acc01 = %init01 : vector<8xf32>, %block_acc02 = %init02 : vector<8xf32>, %block_acc03 = %init03 : vector<8xf32>, %block_acc10 = %init10 : vector<8xf32>, %block_acc11 = %init11 : vector<8xf32>, %block_acc12 = %init12 : vector<8xf32>, %block_acc13 = %init13 : vector<8xf32>) -> (vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) { + %block_result00, %block_result01, %block_result02, %block_result03, %block_result10, %block_result11, %block_result12, %block_result13 = scf.for %quant_group = [%c0 to %c8 step %c2](%acc00 = %block_acc00 : vector<8xf32>, %acc01 = %block_acc01 : vector<8xf32>, %acc02 = %block_acc02 : vector<8xf32>, %acc03 = %block_acc03 : vector<8xf32>, %acc10 = %block_acc10 : vector<8xf32>, %acc11 = %block_acc11 : vector<8xf32>, %acc12 = %block_acc12 : vector<8xf32>, %acc13 = %block_acc13 : vector<8xf32>) -> (vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) { + %block_k_base = index.mul %quant_block, %c256 : index + %group_k_add = index.mul %quant_group, %c32 : index + %k_origin = index.add %block_k_base, %group_k_add : index + %weight_local_row = index.assume %load_row [range(%load_row, 0, 127)] : index + %channel = index.add %channel_tile_base, %weight_local_row : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %weight_values0, %weight_values1 = scf.if %valid_channel -> (vector<16xf16>, vector<16xf16>) { + %channel_byte_index = index.mul %channel, %weight_row_bytes : index + %row_byte_base = index.cast %channel_byte_index : index to offset + %half_packet = scalar.constant false : i1 + %decoded0, %decoded1 = func.call @ggml_q6k_f16_pair(%half_packet, %weight_noalias, %row_byte_base, %quant_block, %quant_group, %load_packet) : (i1, buffer, offset, index, index, index) -> (vector<16xf16>, vector<16xf16>) + scf.yield %decoded0, %decoded1 : vector<16xf16>, vector<16xf16> + } else { + scf.yield %c0_f16x4, %c0_f16x4 : vector<16xf16>, vector<16xf16> + } + %weight_k1 = index.add %load_k, %c32 : index + vector.store %weight_values0, %weight_stage_view[%weight_local_row, %load_k] : vector<16xf16>, view<128x72xf16> + vector.store %weight_values1, %weight_stage_view[%weight_local_row, %weight_k1] : vector<16xf16>, view<128x72xf16> + scf.for %stage_half0 = [%c0 to %c2 step %c1] unroll { + %stage_k_base0 = index.mul %stage_half0, %c32 : index + %token0 = index.add %token_tile_base, %activation_load_row : index + %valid_token0 = index.cmp ult, %token0, %bounded_token_count : index + %activation_k0 = index.add %stage_k_base0, %activation_load_k : index + %activation_values0 = scf.if %valid_token0 -> (vector<8xf16>) { + %bounded_token0, %input_token_count0 = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %input_k0 = index.add %k_origin, %activation_k0 : index + %loaded0 = scf.if %input_is_f16 -> (vector<8xf16>) { + %narrow0 = vector.load %input_f16_view[%bounded_token0, %input_k0] : view<[%bounded_token_count]x[%bounded_input_size]xf16> -> vector<8xf16> + scf.yield %narrow0 : vector<8xf16> + } else { + %wide0 = vector.load %input_f32_view[%bounded_token0, %input_k0] : view<[%bounded_token_count]x[%bounded_input_size]xf32> -> vector<8xf32> + %narrow0 = vector.fptrunc %wide0 : vector<8xf32> to vector<8xf16> + scf.yield %narrow0 : vector<8xf16> + } + scf.yield %loaded0 : vector<8xf16> + } else { + scf.yield %c0_f16x8 : vector<8xf16> + } + vector.store %activation_values0, %activation_stage_physical_view[%activation_load_row, %activation_k0] : vector<8xf16>, view<64x72xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %next00, %next01, %next02, %next03 = scf.for %k_half0 = [%c0 to %c64 step %c16](%half_acc00 = %acc00 : vector<8xf32>, %half_acc01 = %acc01 : vector<8xf32>, %half_acc02 = %acc02 : vector<8xf32>, %half_acc03 = %acc03 : vector<8xf32>) -> (vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) unroll { + %lhs0 = vector.fragment.load %weight_stage_view[%subgroup_channel_add, %k_half0] shape [%m, %k] : view<128x72xf16> -> vector<16xf16> + %rhs00 = vector.fragment.load %activation_fragment_view[%k_half0, %c0] shape [%k, %n] : view<64x64xf16, %activation_fragment_layout> -> vector<16xf16> + %rhs01 = vector.fragment.load %activation_fragment_view[%k_half0, %c16] shape [%k, %n] : view<64x64xf16, %activation_fragment_layout> -> vector<16xf16> + %half_next00 = vector.mma %lhs0, %rhs00, %half_acc00 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %half_next01 = vector.mma %lhs0, %rhs01, %half_acc01 : vector<16xf16>, vector<16xf16>, vector<8xf32> + scf.schedule.fence + %rhs02 = vector.fragment.load %activation_fragment_view[%k_half0, %c32] shape [%k, %n] : view<64x64xf16, %activation_fragment_layout> -> vector<16xf16> + %rhs03 = vector.fragment.load %activation_fragment_view[%k_half0, %c48] shape [%k, %n] : view<64x64xf16, %activation_fragment_layout> -> vector<16xf16> + %half_next02 = vector.mma %lhs0, %rhs02, %half_acc02 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %half_next03 = vector.mma %lhs0, %rhs03, %half_acc03 : vector<16xf16>, vector<16xf16>, vector<8xf32> + scf.schedule.fence + scf.yield %half_next00, %half_next01, %half_next02, %half_next03 : vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.for %stage_half1 = [%c0 to %c2 step %c1] unroll { + %stage_k_base1 = index.mul %stage_half1, %c32 : index + %token1 = index.add %second_token_tile_base, %activation_load_row : index + %valid_token1 = index.cmp ult, %token1, %bounded_token_count : index + %activation_k1 = index.add %stage_k_base1, %activation_load_k : index + %activation_values1 = scf.if %valid_token1 -> (vector<8xf16>) { + %bounded_token1, %input_token_count1 = index.assume %token1, %bounded_token_count [lt(%token1, %bounded_token_count)] : index, index + %input_k1 = index.add %k_origin, %activation_k1 : index + %loaded1 = scf.if %input_is_f16 -> (vector<8xf16>) { + %narrow1 = vector.load %input_f16_view[%bounded_token1, %input_k1] : view<[%bounded_token_count]x[%bounded_input_size]xf16> -> vector<8xf16> + scf.yield %narrow1 : vector<8xf16> + } else { + %wide1 = vector.load %input_f32_view[%bounded_token1, %input_k1] : view<[%bounded_token_count]x[%bounded_input_size]xf32> -> vector<8xf32> + %narrow1 = vector.fptrunc %wide1 : vector<8xf32> to vector<8xf16> + scf.yield %narrow1 : vector<8xf16> + } + scf.yield %loaded1 : vector<8xf16> + } else { + scf.yield %c0_f16x8 : vector<8xf16> + } + vector.store %activation_values1, %activation_stage_physical_view[%activation_load_row, %activation_k1] : vector<8xf16>, view<64x72xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %next10, %next11, %next12, %next13 = scf.for %k_half1 = [%c0 to %c64 step %c16](%half_acc10 = %acc10 : vector<8xf32>, %half_acc11 = %acc11 : vector<8xf32>, %half_acc12 = %acc12 : vector<8xf32>, %half_acc13 = %acc13 : vector<8xf32>) -> (vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) unroll { + %lhs1 = vector.fragment.load %weight_stage_view[%subgroup_channel_add, %k_half1] shape [%m, %k] : view<128x72xf16> -> vector<16xf16> + %rhs10 = vector.fragment.load %activation_fragment_view[%k_half1, %c0] shape [%k, %n] : view<64x64xf16, %activation_fragment_layout> -> vector<16xf16> + %rhs11 = vector.fragment.load %activation_fragment_view[%k_half1, %c16] shape [%k, %n] : view<64x64xf16, %activation_fragment_layout> -> vector<16xf16> + %half_next10 = vector.mma %lhs1, %rhs10, %half_acc10 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %half_next11 = vector.mma %lhs1, %rhs11, %half_acc11 : vector<16xf16>, vector<16xf16>, vector<8xf32> + scf.schedule.fence + %rhs12 = vector.fragment.load %activation_fragment_view[%k_half1, %c32] shape [%k, %n] : view<64x64xf16, %activation_fragment_layout> -> vector<16xf16> + %rhs13 = vector.fragment.load %activation_fragment_view[%k_half1, %c48] shape [%k, %n] : view<64x64xf16, %activation_fragment_layout> -> vector<16xf16> + %half_next12 = vector.mma %lhs1, %rhs12, %half_acc12 : vector<16xf16>, vector<16xf16>, vector<8xf32> + %half_next13 = vector.mma %lhs1, %rhs13, %half_acc13 : vector<16xf16>, vector<16xf16>, vector<8xf32> + scf.schedule.fence + scf.yield %half_next10, %half_next11, %half_next12, %half_next13 : vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %next00, %next01, %next02, %next03, %next10, %next11, %next12, %next13 : vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32> + } + scf.yield %block_result00, %block_result01, %block_result02, %block_result03, %block_result10, %block_result11, %block_result12, %block_result13 : vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32> + } + template.apply<@ggml.mul_mat_q6_k_packed.publish_f32_wave32>(%result00, %c0, %token_tile_base, %bounded_token_count, %bounded_output_size, %channel_tile_base, %subgroup_channel_add, %output_noalias, %output_accumulation) : (vector<8xf32>, index, index, index, index, index, index, buffer, index) + template.apply<@ggml.mul_mat_q6_k_packed.publish_f32_wave32>(%result01, %c16, %token_tile_base, %bounded_token_count, %bounded_output_size, %channel_tile_base, %subgroup_channel_add, %output_noalias, %output_accumulation) : (vector<8xf32>, index, index, index, index, index, index, buffer, index) + template.apply<@ggml.mul_mat_q6_k_packed.publish_f32_wave32>(%result02, %c32, %token_tile_base, %bounded_token_count, %bounded_output_size, %channel_tile_base, %subgroup_channel_add, %output_noalias, %output_accumulation) : (vector<8xf32>, index, index, index, index, index, index, buffer, index) + template.apply<@ggml.mul_mat_q6_k_packed.publish_f32_wave32>(%result03, %c48, %token_tile_base, %bounded_token_count, %bounded_output_size, %channel_tile_base, %subgroup_channel_add, %output_noalias, %output_accumulation) : (vector<8xf32>, index, index, index, index, index, index, buffer, index) + template.apply<@ggml.mul_mat_q6_k_packed.publish_f32_wave32>(%result10, %c64, %token_tile_base, %bounded_token_count, %bounded_output_size, %channel_tile_base, %subgroup_channel_add, %output_noalias, %output_accumulation) : (vector<8xf32>, index, index, index, index, index, index, buffer, index) + template.apply<@ggml.mul_mat_q6_k_packed.publish_f32_wave32>(%result11, %c80, %token_tile_base, %bounded_token_count, %bounded_output_size, %channel_tile_base, %subgroup_channel_add, %output_noalias, %output_accumulation) : (vector<8xf32>, index, index, index, index, index, index, buffer, index) + template.apply<@ggml.mul_mat_q6_k_packed.publish_f32_wave32>(%result12, %c96, %token_tile_base, %bounded_token_count, %bounded_output_size, %channel_tile_base, %subgroup_channel_add, %output_noalias, %output_accumulation) : (vector<8xf32>, index, index, index, index, index, index, buffer, index) + template.apply<@ggml.mul_mat_q6_k_packed.publish_f32_wave32>(%result13, %c112, %token_tile_base, %bounded_token_count, %bounded_output_size, %channel_tile_base, %subgroup_channel_add, %output_noalias, %output_accumulation) : (vector<8xf32>, index, index, index, index, index, index, buffer, index) + template.return +} + +// PP128+ wave32 specialization for ordinary F32 activations. Keeping this as +// a separate provider leaves the wave64 decode schedule and ABI unchanged. +kernel.def target(@ggml_mul_mat_q6_k_prefill_gfx11_wave32) export("ggml_mul_mat_q6_k_f32_wmma_prefill_wave32") @ggml_mul_mat_q6_k_f32_wmma_prefill_wave32(%token_count: index) { + %output_size = config.get @ggml.mul_mat_q6_k_packed.output_size : index + %token_capacity = config.get @ggml.mul_mat_q6_k_packed.token_capacity : index + %c1 = index.constant 1 : index + %c127 = index.constant 127 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %padded_output_size = index.add %output_size, %c127 : index + %output_tiles = index.div %padded_output_size, %c128 : index + %padded_token_count = index.add %token_capacity, %c127 : index + %token_tiles = index.div %padded_token_count, %c128 : index + kernel.launch.config workgroups(%token_tiles, %output_tiles, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) { + %input_size = config.get @ggml.mul_mat_q6_k_packed.input_size : index + %output_size = config.get @ggml.mul_mat_q6_k_packed.output_size : index + %token_capacity = config.get @ggml.mul_mat_q6_k_packed.token_capacity : index + %bounded_token_count = index.assume %token_count [range(%token_count, 128, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + %input_is_f16 = scalar.constant false : i1 + template.apply<@ggml.mul_mat_q6_k_packed.prefill_wave32.body>(%input_is_f16, %bounded_token_count, %input_size, %output_size, %channel_tile, %token_tile, %input, %weight, %output) : (i1, index, index, index, index, index, buffer, buffer, buffer) + kernel.return +} + +// PP128+ wave32 specialization for the F16 transient produced by the fused +// gate/up/SwiGLU recipe. +// Complete M512 tiles amortize the private F16 copy on sufficiently large grids. +func.def pure inline @ggml_q6k_prefill_packed_tile(%tokens: index, %inputs: index, %outputs: index, %accumulation: index) -> (i1, index) { + %zero = index.constant 0 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %c512 = index.constant 512 : index + %contracting = index.cmp ugt, %inputs, %outputs : index + %n_tile = scf.select %contracting, %c64, %c128 : index + %m_tail = index.rem %tokens, %c512 : index + %n_tail = index.rem %outputs, %n_tile : index + %full_m = index.cmp eq, %m_tail, %zero : index + %full_n = index.cmp eq, %n_tail, %zero : index + %overwrite = index.cmp eq, %accumulation, %zero : index + %full_tile = scalar.andi %full_m, %full_n : i1 + %full_output = scalar.andi %full_tile, %overwrite : i1 + %m_tiles = index.div %tokens, %c512 : index + %n_tiles = index.div %outputs, %n_tile : index + %groups = index.mul %m_tiles, %n_tiles : index + %enough_groups = index.cmp uge, %groups, %c32 : index + %packed = scalar.andi %full_output, %enough_groups : i1 + func.return %packed, %n_tile : i1, index +} + +kernel.def target(@ggml_mul_mat_q6_k_prefill_gfx11_wave32) export("ggml_mul_mat_q6_k_f16_wmma_prefill_wave32") @ggml_mul_mat_q6_k_f16_wmma_prefill_wave32(%token_count: index) { + %input_size = config.get @ggml.mul_mat_q6_k_packed.input_size : index + %output_size = config.get @ggml.mul_mat_q6_k_packed.output_size : index + %token_capacity = config.get @ggml.mul_mat_q6_k_packed.token_capacity : index + %output_accumulation = config.get @ggml.mul_mat_q6_k_packed.output_accumulation : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c127 = index.constant 127 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %c512 = index.constant 512 : index + %c1024 = index.constant 1024 : index + %minimum_workgroups = index.constant 32 : index + %token_tail = index.rem %token_count, %c256 : index + %channel_tail = index.rem %output_size, %c128 : index + %full_tokens = index.cmp eq, %token_tail, %c0 : index + %full_channels = index.cmp eq, %channel_tail, %c0 : index + %overwrite = index.cmp eq, %output_accumulation, %c0 : index + %full_tile = scalar.andi %full_tokens, %full_channels : i1 + %wide_tile = scalar.andi %full_tile, %overwrite : i1 + %padded_output_size = index.add %output_size, %c127 : index + %ordinary_output_tiles = index.div %padded_output_size, %c128 : index + %padded_token_count = index.add %token_capacity, %c127 : index + %small_token_tiles = index.div %padded_token_count, %c128 : index + %wide_token_tiles = index.div %token_count, %c256 : index + %staged_token_tiles = scf.select %wide_tile, %wide_token_tiles, %small_token_tiles : index + %wide_workgroup_size = scf.select %wide_tile, %c512, %c256 : index + %larger_output_tiles = index.div %output_size, %c256 : index + %larger_grid = index.mul %wide_token_tiles, %larger_output_tiles : index + %enough_workgroups = index.cmp uge, %larger_grid, %minimum_workgroups : index + %larger_channel_tail = index.rem %output_size, %c256 : index + %larger_full_channels = index.cmp eq, %larger_channel_tail, %c0 : index + %larger_full_tile = scalar.andi %wide_tile, %larger_full_channels : i1 + %multiple_token_tiles = index.cmp ugt, %wide_token_tiles, %c1 : index + %larger_grid_tile = scalar.andi %larger_full_tile, %enough_workgroups : i1 + %larger_tile = scalar.andi %larger_grid_tile, %multiple_token_tiles : i1 + %staged_workgroup_size = scf.select %larger_tile, %c1024, %wide_workgroup_size : index + %staged_output_tiles = scf.select %larger_tile, %larger_output_tiles, %ordinary_output_tiles : index + %packed, %packed_n_tile = func.call pure @ggml_q6k_prefill_packed_tile(%token_count, %input_size, %output_size, %output_accumulation) : (index, index, index, index) -> (i1, index) + %c8 = index.constant 8 : index + %packed_workgroup_size = index.mul %packed_n_tile, %c8 : index + %packed_token_tiles = index.div %token_count, %c512 : index + %packed_output_tiles = index.div %output_size, %packed_n_tile : index + %workgroup_size = scf.select %packed, %packed_workgroup_size, %staged_workgroup_size : index + %token_tiles = scf.select %packed, %packed_token_tiles, %staged_token_tiles : index + %output_tiles = scf.select %packed, %packed_output_tiles, %staged_output_tiles : index + kernel.launch.config workgroups(%token_tiles, %output_tiles, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) { + %input_size = config.get @ggml.mul_mat_q6_k_packed.input_size : index + %output_size = config.get @ggml.mul_mat_q6_k_packed.output_size : index + %token_capacity = config.get @ggml.mul_mat_q6_k_packed.token_capacity : index + %output_accumulation = config.get @ggml.mul_mat_q6_k_packed.output_accumulation : index + %packed_input, %packed_n_tile = func.call pure @ggml_q6k_prefill_packed_tile(%token_capacity, %input_size, %output_size, %output_accumulation) : (index, index, index, index) -> (i1, index) + %bounded_token_count = index.assume %token_count [range(%token_count, 128, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + %workgroup_size = kernel.workgroup.size : index + %c512 = index.constant 512 : index + %wide_tile = index.cmp uge, %workgroup_size, %c512 : index + scf.if %wide_tile { + %weight_format = index.constant 6 : index + %paired = scalar.constant false : i1 + %prefill_binary_op = index.constant 0 : index + func.call @ggml_mul_mat_quantized_f16_prefill_wave32(%weight_format, %prefill_binary_op, %paired, %packed_input, %paired, %bounded_token_count, %input_size, %output_size, %channel_tile, %token_tile, %input, %weight, %weight, %output, %output) : (index, index, i1, i1, i1, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer) + } else { + %input_is_f16 = scalar.constant true : i1 + template.apply<@ggml.mul_mat_q6_k_packed.prefill_wave32.body>(%input_is_f16, %bounded_token_count, %input_size, %output_size, %channel_tile, %token_tile, %input, %weight, %output) : (i1, index, index, index, index, index, buffer, buffer, buffer) + } + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_swiglu_f32_f32_decode.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_swiglu_f32_f32_decode.loom new file mode 100644 index 000000000000..d0cc802ddcf8 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_swiglu_f32_f32_decode.loom @@ -0,0 +1,68 @@ +// Dense GGML decode gate/up matmul fused with an ordered binary op. +// +// The shared dual decode core computes the two adjacent gate/up row pairs. +// This file supplies the SwiGLU-specific publish hook and launcher. +template.decl @ggml.mul_mat_f32_f32_decode.dual_dispatch(%op: index, %lhs_weight_format: index, %rhs_weight_format: index, %token_capacity: index, %token_count: index, %input_size: index, %output_size: index, %input: buffer, %lhs_weight: buffer, %rhs_weight: buffer, %output: buffer) + +template.decl @ggml.mul_mat_f32_f32_decode.publish_dual_pair2(%op: index, %publish_pair: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %row1_valid: i1, %lhs0: f32, %lhs1: f32, %rhs0: f32, %rhs1: f32, %output: buffer) + +template.decl @ggml.binary_f32.apply(%op: index, %lhs: f32, %rhs: f32) -> (f32) + +amdgpu.target @ggml_mul_mat_swiglu_decode_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 1)] + +config.decl @ggml.mul_mat_swiglu.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 32)] + +config.decl @ggml.mul_mat_swiglu.output_size : %value: index where [range(%value, 1, 262144)] + +config.decl @ggml.mul_mat_swiglu.gate_weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.mul_mat_swiglu.up_weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.mul_mat_swiglu.op : %value: index where [range(%value, 0, 8)] + +template.def<@ggml.mul_mat_f32_f32_decode.publish_dual_pair2> device priority(20) @ggml_mul_mat_swiglu_decode_publish_dual_pair2(%op: index, %publish_pair: i1, %token0: index, %channel: index, %token_count0: index, %output_size0: index, %row1_valid: i1, %gate0: f32, %gate1: f32, %up0: f32, %up1: f32, %output: buffer) { + %token_count = index.assume %token_count0 [range(%token_count0, 1, 2048)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 1, 1048576)] : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %token, %launch_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %output_noalias = buffer.assume.noalias %output : buffer + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%output_size]xf32> + scf.if %publish_pair { + %value0 = template.apply<@ggml.binary_f32.apply>(%op, %gate0, %up0) : (index, f32, f32) -> (f32) + view.store %value0, %output_view[%token, %channel] : f32, view<[%launch_token_count]x[%output_size]xf32> + } + scf.if %row1_valid { + scf.if %publish_pair { + %row1 = index.add %channel, %c1 : index + %value1 = template.apply<@ggml.binary_f32.apply>(%op, %gate1, %up1) : (index, f32, f32) -> (f32) + view.store %value1, %output_view[%token, %row1] : f32, view<[%launch_token_count]x[%output_size]xf32> + } + } + template.return +} + +kernel.def target(@ggml_mul_mat_swiglu_decode_gfx11_wave64) @ggml_mul_mat_swiglu_f32_f32_decode_wave64(%token_count: index) { + %token_capacity = config.get @ggml.workload.token_capacity : index + %output_capacity = config.get @ggml.mul_mat_swiglu.output_size : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %padded_output_size = index.add %output_capacity, %c1 : index + %output_pairs = index.div %padded_output_size, %c2 : index + %padded_token_count = index.add %token_capacity, %c63 : index + %launch_tokens = index.div %padded_token_count, %c64 : index + kernel.launch.config workgroups(%output_pairs, %launch_tokens, %c1) workgroup_size(%c64, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) where [range(%token_count, 1, 1)] { + %input_size = config.get @ggml.mul_mat_swiglu.input_size : index + %output_size = config.get @ggml.mul_mat_swiglu.output_size : index + %gate_weight_format = config.get @ggml.mul_mat_swiglu.gate_weight_format : index + %up_weight_format = config.get @ggml.mul_mat_swiglu.up_weight_format : index + %op = config.get @ggml.mul_mat_swiglu.op : index + %launch_token_capacity = config.get @ggml.workload.token_capacity : index + template.apply<@ggml.mul_mat_f32_f32_decode.dual_dispatch>(%op, %gate_weight_format, %up_weight_format, %launch_token_capacity, %token_count, %input_size, %output_size, %input, %gate_weight, %up_weight, %output) : (index, index, index, index, index, index, index, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_swiglu_f32_f32_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_swiglu_f32_f32_wmma.loom new file mode 100644 index 000000000000..42ce7df3297a --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_swiglu_f32_f32_wmma.loom @@ -0,0 +1,1491 @@ +// Dense GGML gate/up matmul fused with an ordered binary op for contiguous 2D F32 activations. +template.decl @ggml.mul_mat_swiglu.apply_vector4(%op: index, %lhs: vector<4xf32>, %rhs: vector<4xf32>) -> (vector<4xf32>) + +func.decl @ggml_q4k_scale_min_from_header(%scale0: i32, %scale1: i32, %scale2: i32, %q4_group: index) -> (i32, i32) + +func.decl @ggml_dot_u8_s8_vector8_f32(%weight: vector<2xi32>, %activation: vector<2xi32>) -> (f32) + +func.decl @ggml_dequant_f32_vector4(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight_format: index, %weight: buffer, %row_byte_base: offset, %input_size: index, %quant_block: index, %quant_group: index, %packet: index, %k: index) -> (vector<4xf32>) + +func.decl @ggml_q4k_native_row64_f32_vector4(%weight: buffer, %row: index, %input_size: index, %q4_block: index, %q4_group: index, %packet: index) -> (vector<4xf32>) + +template.decl @ggml.unary_f32.apply_vector4(%op: index, %values: vector<4xf32>) -> (vector<4xf32>) + +template.decl @ggml.mul_mat_swiglu.body(%op: index, %gate_weight_format: index, %up_weight_format: index, %token_count: index, %input_size0: index, %output_size0: index, %channel_tile: index, %token_tile: index, %input: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) + +template.decl @ggml.mul_mat_swiglu.launch(%token_capacity: index, %output_size: index) -> (index, index, index, index) + +amdgpu.target @ggml_mul_mat_swiglu_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.mul_mat_swiglu.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 32)] + +config.decl @ggml.mul_mat_swiglu.output_size : %value: index where [range(%value, 1, 262144)] + +config.decl @ggml.mul_mat_swiglu.gate_weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.mul_mat_swiglu.up_weight_format : %value: index where [range(%value, 4, 81)] + +config.decl @ggml.mul_mat_swiglu.op : %value: index where [range(%value, 0, 8)] + +config.decl @ggml.mul_mat_swiglu.f16_output_layout : %value: index where [range(%value, 0, 1)] + +config.decl @ggml.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +func.decl @ggml_dequant_weight_row_bytes(%weight_format: index, %hidden_size: index) -> (offset) + +func.decl @ggml_dequant_weight_tile_bytes(%weight_format: index) -> (offset) + +func.decl @ggml_iq4nl_table_i8() -> (vector<16xi8>) + +func.decl @ggml_iq3s_grid_f32_halves() -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + +func.decl @ggml_dequant_f16_vector4(%iq4nl_table: vector<16xi8>, %iq3s_grid_low0: vector<32xf32>, %iq3s_grid_low1: vector<32xf32>, %iq3s_grid_low2: vector<32xf32>, %iq3s_grid_low3: vector<32xf32>, %iq3s_grid_low4: vector<32xf32>, %iq3s_grid_low5: vector<32xf32>, %iq3s_grid_low6: vector<32xf32>, %iq3s_grid_low7: vector<32xf32>, %iq3s_grid_low8: vector<32xf32>, %iq3s_grid_low9: vector<32xf32>, %iq3s_grid_low10: vector<32xf32>, %iq3s_grid_low11: vector<32xf32>, %iq3s_grid_low12: vector<32xf32>, %iq3s_grid_low13: vector<32xf32>, %iq3s_grid_low14: vector<32xf32>, %iq3s_grid_low15: vector<32xf32>, %iq3s_grid_high0: vector<32xf32>, %iq3s_grid_high1: vector<32xf32>, %iq3s_grid_high2: vector<32xf32>, %iq3s_grid_high3: vector<32xf32>, %iq3s_grid_high4: vector<32xf32>, %iq3s_grid_high5: vector<32xf32>, %iq3s_grid_high6: vector<32xf32>, %iq3s_grid_high7: vector<32xf32>, %iq3s_grid_high8: vector<32xf32>, %iq3s_grid_high9: vector<32xf32>, %iq3s_grid_high10: vector<32xf32>, %iq3s_grid_high11: vector<32xf32>, %iq3s_grid_high12: vector<32xf32>, %iq3s_grid_high13: vector<32xf32>, %iq3s_grid_high14: vector<32xf32>, %iq3s_grid_high15: vector<32xf32>, %weight_format: index, %weight: buffer, %row_byte_base: offset, %input_size: index, %quant_block: index, %quant_group: index, %packet: index, %k: index) -> (vector<4xf16>) + +template.def<@ggml.mul_mat_swiglu.apply_vector4> device @ggml_mul_mat_swiglu_apply_vector4(%op: index, %lhs: vector<4xf32>, %rhs: vector<4xf32>) -> (vector<4xf32>) { + %c0_5_f32x4 = vector.constant 0.5 : vector<4xf32> + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %c1_f32x4 = vector.constant 1.0 : vector<4xf32> + %c2_f32x4 = vector.constant 2.0 : vector<4xf32> + %cn2_f32x4 = vector.constant -2.0 : vector<4xf32> + %gelu_coef = vector.constant 0.044715 : vector<4xf32> + %sqrt_2_over_pi = vector.constant 0.79788456080286544 : vector<4xf32> + + %op_add = index.constant 0 : index + %sum = vector.addf %lhs, %rhs : vector<4xf32> + %is_add = index.cmp eq, %op, %op_add : index + %add_selected = scf.select %is_add, %sum, %lhs : vector<4xf32> + + %op_sub = index.constant 1 : index + %difference = vector.subf %lhs, %rhs : vector<4xf32> + %is_sub = index.cmp eq, %op, %op_sub : index + %sub_selected = scf.select %is_sub, %difference, %add_selected : vector<4xf32> + + %op_mul = index.constant 2 : index + %product = vector.mulf %lhs, %rhs : vector<4xf32> + %is_mul = index.cmp eq, %op, %op_mul : index + %mul_selected = scf.select %is_mul, %product, %sub_selected : vector<4xf32> + + %op_div = index.constant 3 : index + %quotient = vector.divf %lhs, %rhs : vector<4xf32> + %is_div = index.cmp eq, %op, %op_div : index + %div_selected = scf.select %is_div, %quotient, %mul_selected : vector<4xf32> + + %op_swiglu = index.constant 4 : index + %negative_lhs = vector.subf %c0_f32x4, %lhs : vector<4xf32> + %lhs_exp = vector.expf %negative_lhs : vector<4xf32> + %denominator = vector.addf %c1_f32x4, %lhs_exp : vector<4xf32> + %sigmoid = vector.divf %c1_f32x4, %denominator : vector<4xf32> + %silu = vector.mulf %lhs, %sigmoid : vector<4xf32> + %swiglu = vector.mulf %silu, %rhs : vector<4xf32> + %is_swiglu = index.cmp eq, %op, %op_swiglu : index + %swiglu_selected = scf.select %is_swiglu, %swiglu, %div_selected : vector<4xf32> + + %op_geglu = index.constant 5 : index + %gelu_x2 = vector.mulf %lhs, %lhs : vector<4xf32> + %gelu_poly0 = vector.mulf %gelu_coef, %gelu_x2 : vector<4xf32> + %gelu_poly = vector.addf %c1_f32x4, %gelu_poly0 : vector<4xf32> + %gelu_inner0 = vector.mulf %lhs, %gelu_poly : vector<4xf32> + %gelu_inner = vector.mulf %sqrt_2_over_pi, %gelu_inner0 : vector<4xf32> + %gelu_tanh_scaled = vector.mulf %cn2_f32x4, %gelu_inner : vector<4xf32> + %gelu_tanh_exp = vector.expf %gelu_tanh_scaled : vector<4xf32> + %gelu_tanh_denominator = vector.addf %c1_f32x4, %gelu_tanh_exp : vector<4xf32> + %gelu_tanh_ratio = vector.divf %c2_f32x4, %gelu_tanh_denominator : vector<4xf32> + %gelu_tanh = vector.subf %gelu_tanh_ratio, %c1_f32x4 : vector<4xf32> + %gelu_one_plus = vector.addf %c1_f32x4, %gelu_tanh : vector<4xf32> + %gelu_half_x = vector.mulf %c0_5_f32x4, %lhs : vector<4xf32> + %gelu = vector.mulf %gelu_half_x, %gelu_one_plus : vector<4xf32> + %geglu = vector.mulf %gelu, %rhs : vector<4xf32> + %is_geglu = index.cmp eq, %op, %op_geglu : index + %geglu_selected = scf.select %is_geglu, %geglu, %swiglu_selected : vector<4xf32> + + %op_reglu = index.constant 6 : index + %unary_relu = index.constant 2 : index + %relu_lhs = template.apply<@ggml.unary_f32.apply_vector4>(%unary_relu, %lhs) : (index, vector<4xf32>) -> (vector<4xf32>) + %reglu = vector.mulf %relu_lhs, %rhs : vector<4xf32> + %is_reglu = index.cmp eq, %op, %op_reglu : index + %reglu_selected = scf.select %is_reglu, %reglu, %geglu_selected : vector<4xf32> + + %op_geglu_erf = index.constant 7 : index + %unary_gelu_erf = index.constant 20 : index + %gelu_erf_lhs = template.apply<@ggml.unary_f32.apply_vector4>(%unary_gelu_erf, %lhs) : (index, vector<4xf32>) -> (vector<4xf32>) + %geglu_erf = vector.mulf %gelu_erf_lhs, %rhs : vector<4xf32> + %is_geglu_erf = index.cmp eq, %op, %op_geglu_erf : index + %geglu_erf_selected = scf.select %is_geglu_erf, %geglu_erf, %reglu_selected : vector<4xf32> + + %op_geglu_quick = index.constant 8 : index + %unary_gelu_quick = index.constant 14 : index + %gelu_quick_lhs = template.apply<@ggml.unary_f32.apply_vector4>(%unary_gelu_quick, %lhs) : (index, vector<4xf32>) -> (vector<4xf32>) + %geglu_quick = vector.mulf %gelu_quick_lhs, %rhs : vector<4xf32> + %is_geglu_quick = index.cmp eq, %op, %op_geglu_quick : index + %result = scf.select %is_geglu_quick, %geglu_quick, %geglu_erf_selected : vector<4xf32> + template.return %result : vector<4xf32> +} + +template.def<@ggml.mul_mat_swiglu.body> device @ggml_mul_mat_swiglu_f32_f32_wmma_body(%op: index, %gate_weight_format: index, %up_weight_format: index, %token_count: index, %input_size0: index, %output_size0: index, %channel_tile: index, %token_tile: index, %input: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 2, 2048)] : index + %bounded_input_size = index.assume %input_size0 [range(%input_size0, 256, 32768), mul(%input_size0, 32)] : index + %bounded_output_size = index.assume %output_size0 [range(%output_size0, 1, 262144)] : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 1)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c40 = index.constant 40 : index + %c64 = index.constant 64 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %weight_stage_bytes = index.constant 5120 : offset + %activation_stage_bytes = index.constant 2560 : offset + %wave_result_stage_bytes = index.constant 1024 : offset + %result_stage_bytes = index.constant 2048 : offset + %c0_f16x4 = vector.constant 0.0 : vector<4xf16> + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %zero_accumulator = vector.constant 0.0 : vector<4xf32> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %padded_input_size = index.add %bounded_input_size, %c255 : index + %quant_block_count = index.div %padded_input_size, %c256 : index + %gate_weight_row_bytes = func.call @ggml_dequant_weight_row_bytes(%gate_weight_format, %bounded_input_size) : (index, index) -> (offset) + %up_weight_row_bytes = func.call @ggml_dequant_weight_row_bytes(%up_weight_format, %bounded_input_size) : (index, index) -> (offset) + %input_noalias, %gate_weight_noalias, %up_weight_noalias, %output_noalias = buffer.assume.noalias %input, %gate_weight, %up_weight, %output : buffer, buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_input_size]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_output_size]xf32> + %gate_weight_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %up_weight_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %activation_stage = buffer.alloca align(16) %activation_stage_bytes : buffer + %gate_result_stage = buffer.alloca align(16) %result_stage_bytes : buffer + %up_result_stage = buffer.alloca align(16) %result_stage_bytes : buffer + %gate_weight_stage_view = buffer.view %gate_weight_stage[%c0_offset] : buffer -> view<64x40xf16> + %up_weight_stage_view = buffer.view %up_weight_stage[%c0_offset] : buffer -> view<64x40xf16> + %activation_stage_physical_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x40xf16> + %activation_fragment_layout = encoding.layout.strided [1, %c40] : encoding + %activation_fragment_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x32xf16, %activation_fragment_layout> + %wave_result_stage_offset = index.scale %subgroup, %wave_result_stage_bytes : index, offset -> offset + %result_fragment_layout = encoding.layout.strided [1, %c16] : encoding + %gate_result_fragment_view = buffer.view %gate_result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf32, %result_fragment_layout> + %up_result_fragment_view = buffer.view %up_result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf32, %result_fragment_layout> + %gate_result_physical_view = buffer.view %gate_result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf32> + %up_result_physical_view = buffer.view %up_result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf32> + %channel_tile_base = index.mul %channel_tile, %c64 : index + %token_tile_base = index.mul %token_tile, %c32 : index + %load_packet = index.rem %workitem, %c8 : index + %load_k = index.mul %load_packet, %c4 : index + %load_row0 = index.div %workitem, %c8 : index + %load_row = index.assume %load_row0 [range(%load_row0, 0, 15)] : index + %subgroup_channel_add = index.mul %subgroup, %c32 : index + %subgroup_channel1 = index.add %subgroup_channel_add, %c16 : index + %init00 = vector.fragment %zero_accumulator shape [%m, %n] : vector<4xf32> + %init01 = vector.fragment %zero_accumulator shape [%m, %n] : vector<4xf32> + %init10 = vector.fragment %zero_accumulator shape [%m, %n] : vector<4xf32> + %init11 = vector.fragment %zero_accumulator shape [%m, %n] : vector<4xf32> + %iq4nl_table = func.call @ggml_iq4nl_table_i8() : () -> (vector<16xi8>) + %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15 = func.call @ggml_iq3s_grid_f32_halves() : () -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + %gate_result00, %gate_result01, %gate_result10, %gate_result11, %up_result00, %up_result01, %up_result10, %up_result11 = scf.for %quant_block = [%c0 to %quant_block_count step %c1](%gate_block_acc00 = %init00 : vector<4xf32>, %gate_block_acc01 = %init01 : vector<4xf32>, %gate_block_acc10 = %init10 : vector<4xf32>, %gate_block_acc11 = %init11 : vector<4xf32>, %up_block_acc00 = %init00 : vector<4xf32>, %up_block_acc01 = %init01 : vector<4xf32>, %up_block_acc10 = %init10 : vector<4xf32>, %up_block_acc11 = %init11 : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + %gate_block_result00, %gate_block_result01, %gate_block_result10, %gate_block_result11, %up_block_result00, %up_block_result01, %up_block_result10, %up_block_result11 = scf.for %quant_group = [%c0 to %c8 step %c1](%gate_acc00 = %gate_block_acc00 : vector<4xf32>, %gate_acc01 = %gate_block_acc01 : vector<4xf32>, %gate_acc10 = %gate_block_acc10 : vector<4xf32>, %gate_acc11 = %gate_block_acc11 : vector<4xf32>, %up_acc00 = %up_block_acc00 : vector<4xf32>, %up_acc01 = %up_block_acc01 : vector<4xf32>, %up_acc10 = %up_block_acc10 : vector<4xf32>, %up_acc11 = %up_block_acc11 : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + %block_k_base = index.mul %quant_block, %c256 : index + %group_k_add = index.mul %quant_group, %c32 : index + %k_origin = index.add %block_k_base, %group_k_add : index + scf.for %row_offset = [%c0 to %c64 step %c16] unroll { + %local_row0 = index.add %load_row, %row_offset : index + %local_row = index.assume %local_row0 [range(%local_row0, 0, 63)] : index + %channel = index.add %channel_tile_base, %local_row : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %weight_k = index.add %k_origin, %load_k : index + %valid_k = index.cmp ult, %weight_k, %bounded_input_size : index + %valid_weight = scalar.andi %valid_channel, %valid_k : i1 + %gate_weight_values = scf.if %valid_weight -> (vector<4xf16>) { + %row_byte_base = index.scale %channel, %gate_weight_row_bytes : index, offset -> offset + %decoded = func.call @ggml_dequant_f16_vector4(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %gate_weight_format, %gate_weight_noalias, %row_byte_base, %bounded_input_size, %quant_block, %quant_group, %load_packet, %weight_k) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, buffer, offset, index, index, index, index, index) -> (vector<4xf16>) + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + %up_weight_values = scf.if %valid_weight -> (vector<4xf16>) { + %row_byte_base = index.scale %channel, %up_weight_row_bytes : index, offset -> offset + %decoded = func.call @ggml_dequant_f16_vector4(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %up_weight_format, %up_weight_noalias, %row_byte_base, %bounded_input_size, %quant_block, %quant_group, %load_packet, %weight_k) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, buffer, offset, index, index, index, index, index) -> (vector<4xf16>) + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %gate_weight_values, %gate_weight_stage_view[%local_row, %load_k] : vector<4xf16>, view<64x40xf16> + vector.store %up_weight_values, %up_weight_stage_view[%local_row, %load_k] : vector<4xf16>, view<64x40xf16> + %is_activation_row = index.cmp ult, %local_row, %c32 : index + scf.if %is_activation_row { + %activation_row = index.assume %local_row [range(%local_row, 0, 31)] : index + %token = index.add %token_tile_base, %activation_row : index + %valid_token = index.cmp ult, %token, %bounded_token_count : index + %input_k = index.add %k_origin, %load_k : index + %valid_input = index.cmp ult, %input_k, %bounded_input_size : index + %valid_activation = scalar.andi %valid_token, %valid_input : i1 + %activation_values = scf.if %valid_activation -> (vector<4xf16>) { + %bounded_token, %input_token_count = index.assume %token, %bounded_token_count [lt(%token, %bounded_token_count)] : index, index + %mask = vector.mask.range [%input_k to %bounded_input_size step %c1] : index -> vector<4xi1> + %loaded = vector.load.mask %input_view[%bounded_token, %input_k], %mask, %c0_f32x4 : view<[%bounded_token_count]x[%bounded_input_size]xf32>, vector<4xi1>, vector<4xf32> + %converted = vector.fptrunc %loaded : vector<4xf32> to vector<4xf16> + scf.yield %converted : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %activation_values, %activation_stage_physical_view[%activation_row, %load_k] : vector<4xf16>, view<32x40xf16> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %gate_next00, %gate_next01, %gate_next10, %gate_next11, %up_next00, %up_next01, %up_next10, %up_next11 = scf.for %k_half = [%c0 to %c32 step %c16](%gate_half_acc00 = %gate_acc00 : vector<4xf32>, %gate_half_acc01 = %gate_acc01 : vector<4xf32>, %gate_half_acc10 = %gate_acc10 : vector<4xf32>, %gate_half_acc11 = %gate_acc11 : vector<4xf32>, %up_half_acc00 = %up_acc00 : vector<4xf32>, %up_half_acc01 = %up_acc01 : vector<4xf32>, %up_half_acc10 = %up_acc10 : vector<4xf32>, %up_half_acc11 = %up_acc11 : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) unroll { + %gate_lhs0 = vector.fragment.load %gate_weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %gate_lhs1 = vector.fragment.load %gate_weight_stage_view[%subgroup_channel1, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %up_lhs0 = vector.fragment.load %up_weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %up_lhs1 = vector.fragment.load %up_weight_stage_view[%subgroup_channel1, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %rhs0 = vector.fragment.load %activation_fragment_view[%k_half, %c0] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %rhs1 = vector.fragment.load %activation_fragment_view[%k_half, %c16] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %gate_half_next00 = vector.mma %gate_lhs0, %rhs0, %gate_half_acc00 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %gate_half_next01 = vector.mma %gate_lhs0, %rhs1, %gate_half_acc01 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %gate_half_next10 = vector.mma %gate_lhs1, %rhs0, %gate_half_acc10 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %gate_half_next11 = vector.mma %gate_lhs1, %rhs1, %gate_half_acc11 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %up_half_next00 = vector.mma %up_lhs0, %rhs0, %up_half_acc00 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %up_half_next01 = vector.mma %up_lhs0, %rhs1, %up_half_acc01 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %up_half_next10 = vector.mma %up_lhs1, %rhs0, %up_half_acc10 : vector<16xf16>, vector<16xf16>, vector<4xf32> + %up_half_next11 = vector.mma %up_lhs1, %rhs1, %up_half_acc11 : vector<16xf16>, vector<16xf16>, vector<4xf32> + scf.yield %gate_half_next00, %gate_half_next01, %gate_half_next10, %gate_half_next11, %up_half_next00, %up_half_next01, %up_half_next10, %up_half_next11 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %gate_next00, %gate_next01, %gate_next10, %gate_next11, %up_next00, %up_next01, %up_next10, %up_next11 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + scf.yield %gate_block_result00, %gate_block_result01, %gate_block_result10, %gate_block_result11, %up_block_result00, %up_block_result01, %up_block_result10, %up_block_result11 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + %publish_token0 = index.div %lane, %c4 : index + %publish_token = index.assume %publish_token0 [range(%publish_token0, 0, 15)] : index + %publish_packet0 = index.rem %lane, %c4 : index + %publish_packet = index.assume %publish_packet0 [range(%publish_packet0, 0, 3)] : index + %publish_channel_add = index.mul %publish_packet, %c4 : index + %token0 = index.add %token_tile_base, %publish_token : index + %publish_token1 = index.add %publish_token, %c16 : index + %token1 = index.add %token_tile_base, %publish_token1 : index + %subgroup_channel_base = index.add %channel_tile_base, %subgroup_channel_add : index + %channel0 = index.add %subgroup_channel_base, %publish_channel_add : index + %channel1_base = index.add %subgroup_channel_base, %c16 : index + %channel1 = index.add %channel1_base, %publish_channel_add : index + %valid_token0 = index.cmp ult, %token0, %bounded_token_count : index + %valid_token1 = index.cmp ult, %token1, %bounded_token_count : index + %valid_channel0 = index.cmp ult, %channel0, %bounded_output_size : index + %valid_channel1 = index.cmp ult, %channel1, %bounded_output_size : index + %writes00 = scalar.andi %valid_token0, %valid_channel0 : i1 + %writes01 = scalar.andi %valid_token1, %valid_channel0 : i1 + %writes10 = scalar.andi %valid_token0, %valid_channel1 : i1 + %writes11 = scalar.andi %valid_token1, %valid_channel1 : i1 + vector.fragment.store %gate_result00, %gate_result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + vector.fragment.store %up_result00, %up_result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes00 { + %bounded_token, %output_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %gate_values = vector.load %gate_result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + %up_values = vector.load %up_result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + %activated = template.apply<@ggml.mul_mat_swiglu.apply_vector4>(%op, %gate_values, %up_values) : (index, vector<4xf32>, vector<4xf32>) -> (vector<4xf32>) + %mask = vector.mask.range [%channel0 to %bounded_output_size step %c1] : index -> vector<4xi1> + vector.store.mask %activated, %output_view[%bounded_token, %channel0], %mask : vector<4xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32>, vector<4xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %gate_result01, %gate_result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + vector.fragment.store %up_result01, %up_result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes01 { + %bounded_token, %output_token_count = index.assume %token1, %bounded_token_count [lt(%token1, %bounded_token_count)] : index, index + %gate_values = vector.load %gate_result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + %up_values = vector.load %up_result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + %activated = template.apply<@ggml.mul_mat_swiglu.apply_vector4>(%op, %gate_values, %up_values) : (index, vector<4xf32>, vector<4xf32>) -> (vector<4xf32>) + %mask = vector.mask.range [%channel0 to %bounded_output_size step %c1] : index -> vector<4xi1> + vector.store.mask %activated, %output_view[%bounded_token, %channel0], %mask : vector<4xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32>, vector<4xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %gate_result10, %gate_result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + vector.fragment.store %up_result10, %up_result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes10 { + %bounded_token, %output_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %gate_values = vector.load %gate_result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + %up_values = vector.load %up_result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + %activated = template.apply<@ggml.mul_mat_swiglu.apply_vector4>(%op, %gate_values, %up_values) : (index, vector<4xf32>, vector<4xf32>) -> (vector<4xf32>) + %mask = vector.mask.range [%channel1 to %bounded_output_size step %c1] : index -> vector<4xi1> + vector.store.mask %activated, %output_view[%bounded_token, %channel1], %mask : vector<4xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32>, vector<4xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %gate_result11, %gate_result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + vector.fragment.store %up_result11, %up_result_fragment_view[%c0, %c0] shape [%m, %n] : vector<4xf32>, view<16x16xf32, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes11 { + %bounded_token, %output_token_count = index.assume %token1, %bounded_token_count [lt(%token1, %bounded_token_count)] : index, index + %gate_values = vector.load %gate_result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + %up_values = vector.load %up_result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf32> -> vector<4xf32> + %activated = template.apply<@ggml.mul_mat_swiglu.apply_vector4>(%op, %gate_values, %up_values) : (index, vector<4xf32>, vector<4xf32>) -> (vector<4xf32>) + %mask = vector.mask.range [%channel1 to %bounded_output_size step %c1] : index -> vector<4xi1> + vector.store.mask %activated, %output_view[%bounded_token, %channel1], %mask : vector<4xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32>, vector<4xi1> + } + template.return +} + +template.def<@ggml.mul_mat_swiglu.launch> @ggml_mul_mat_swiglu_f32_f32_wmma_launch(%token_capacity: index, %output_size: index) -> (index, index, index, index) { + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %padded_token_count = index.add %token_capacity, %c31 : index + %token_tiles = index.div %padded_token_count, %c32 : index + template.return %output_tiles, %token_tiles, %c1, %c128 : index, index, index, index +} + +kernel.def target(@ggml_mul_mat_swiglu_gfx11_wave64) @ggml_mul_mat_swiglu_f32_f32_wmma(%token_count: index) { + %output_size = config.get @ggml.mul_mat_swiglu.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %padded_token_count = index.add %token_capacity, %c31 : index + %token_tiles = index.div %padded_token_count, %c32 : index + kernel.launch.config workgroups(%output_tiles, %token_tiles, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) where [range(%token_count, 2, 2048)] { + %input_size = config.get @ggml.mul_mat_swiglu.input_size : index + %output_size = config.get @ggml.mul_mat_swiglu.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %gate_weight_format = config.get @ggml.mul_mat_swiglu.gate_weight_format : index + %up_weight_format = config.get @ggml.mul_mat_swiglu.up_weight_format : index + %op = config.get @ggml.mul_mat_swiglu.op : index + %bounded_token_count = index.assume %token_count [range(%token_count, 2, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + template.apply<@ggml.mul_mat_swiglu.body>(%op, %gate_weight_format, %up_weight_format, %bounded_token_count, %input_size, %output_size, %channel_tile, %token_tile, %input, %gate_weight, %up_weight, %output) : (index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer) + kernel.return +} + +amdgpu.target @ggml_mul_mat_swiglu_gfx11_wave32 {subgroup_size = 32} + +// Low-token specializations keep FP32 operands and accumulation and avoid +// padding a handful of tokens to a full matrix tile. +func.def inline @ggml_mul_mat_swiglu_lowtoken_dot(%op: index, %gate_format: index, %up_format: index, %token_count: index, %input_size: index, %output_size: index, %channel_tile: index, %input: buffer, %gate: buffer, %up: buffer, %output: buffer) { + %iq4nl_table = func.call @ggml_iq4nl_table_i8() : () -> (vector<16xi8>) + %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15 = func.call @ggml_iq3s_grid_f32_halves() : () -> (vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c256 = index.constant 256 : index + %base = index.constant 0 : offset + %zero = vector.constant 0.0 : vector<4xf32> + %zero_scalar = scalar.constant 0.0 : f32 + %tokens = index.assume %token_count [range(%token_count, 1, 5)] : index + %k = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %n = index.assume %output_size [range(%output_size, 1, 262144)] : index + %subgroup = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %c16 = index.constant 16 : index + %c64 = index.constant 64 : index + %row_base = index.mul %channel_tile, %c8 : index + %wave_row = index.mul %subgroup, %c2 : index + %cohort = index.div %lane, %c16 : index + %local_row = index.add %wave_row, %cohort : index + %row = index.add %row_base, %local_row : index + %lane16 = index.rem %lane, %c16 : index + %valid_row = index.cmp ult, %row, %n : index + %lane_group = index.div %lane16, %c8 : index + %packet = index.rem %lane16, %c8 : index + %lane_k = index.mul %lane16, %c4 : index + %blocks = index.div %k, %c256 : index + %steps = index.div %k, %c64 : index + %gate_bytes = func.call @ggml_dequant_weight_tile_bytes(%gate_format) : (index) -> (offset) + %up_bytes = func.call @ggml_dequant_weight_tile_bytes(%up_format) : (index) -> (offset) + %gate_row_bytes = index.scale %blocks, %gate_bytes : index, offset -> offset + %up_row_bytes = index.scale %blocks, %up_bytes : index, offset -> offset + %gate_base = index.scale %row, %gate_row_bytes : index, offset -> offset + %up_base = index.scale %row, %up_row_bytes : index, offset -> offset + %a_na, %g_na, %u_na, %o_na = buffer.assume.noalias %input, %gate, %up, %output : buffer, buffer, buffer, buffer + %av = buffer.view %a_na[%base] : buffer -> view<[%tokens]x[%k]xf32> + %ov = buffer.view %o_na[%base] : buffer -> view<[%tokens]x[%n]xf32> + scf.if %valid_row { + %total_g0, %total_u0, %total_g1, %total_u1, %total_g2, %total_u2, %total_g3, %total_u3, %total_g4, %total_u4 = scf.for %step = [%c0 to %steps step %c1](%acc_g0 = %zero : vector<4xf32>, %acc_u0 = %zero : vector<4xf32>, %acc_g1 = %zero : vector<4xf32>, %acc_u1 = %zero : vector<4xf32>, %acc_g2 = %zero : vector<4xf32>, %acc_u2 = %zero : vector<4xf32>, %acc_g3 = %zero : vector<4xf32>, %acc_u3 = %zero : vector<4xf32>, %acc_g4 = %zero : vector<4xf32>, %acc_u4 = %zero : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + %block = index.div %step, %c4 : index + %quarter = index.rem %step, %c4 : index + %quarter_group = index.mul %quarter, %c2 : index + %group = index.add %quarter_group, %lane_group : index + %block_k = index.mul %step, %c64 : index + %kk = index.add %block_k, %lane_k : index + %g_packed_format = index.constant 44 : index + %g_packed = index.cmp eq, %gate_format, %g_packed_format : index + %gh = scf.if %g_packed -> (vector<4xf32>) { + %values = func.call @ggml_q4k_native_row64_f32_vector4(%g_na, %row, %k, %block, %group, %packet) : (buffer, index, index, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + %values = func.call @ggml_dequant_f32_vector4(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %gate_format, %g_na, %gate_base, %k, %block, %group, %packet, %kk) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, buffer, offset, index, index, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } + %u_packed_format = index.constant 44 : index + %u_packed = index.cmp eq, %up_format, %u_packed_format : index + %uh = scf.if %u_packed -> (vector<4xf32>) { + %values = func.call @ggml_q4k_native_row64_f32_vector4(%u_na, %row, %k, %block, %group, %packet) : (buffer, index, index, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + %values = func.call @ggml_dequant_f32_vector4(%iq4nl_table, %iq3s_grid_low0, %iq3s_grid_low1, %iq3s_grid_low2, %iq3s_grid_low3, %iq3s_grid_low4, %iq3s_grid_low5, %iq3s_grid_low6, %iq3s_grid_low7, %iq3s_grid_low8, %iq3s_grid_low9, %iq3s_grid_low10, %iq3s_grid_low11, %iq3s_grid_low12, %iq3s_grid_low13, %iq3s_grid_low14, %iq3s_grid_low15, %iq3s_grid_high0, %iq3s_grid_high1, %iq3s_grid_high2, %iq3s_grid_high3, %iq3s_grid_high4, %iq3s_grid_high5, %iq3s_grid_high6, %iq3s_grid_high7, %iq3s_grid_high8, %iq3s_grid_high9, %iq3s_grid_high10, %iq3s_grid_high11, %iq3s_grid_high12, %iq3s_grid_high13, %iq3s_grid_high14, %iq3s_grid_high15, %up_format, %u_na, %up_base, %k, %block, %group, %packet, %kk) : (vector<16xi8>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>, index, buffer, offset, index, index, index, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } + %t0 = index.constant 0 : index + %valid0 = index.cmp ult, %t0, %tokens : index + %next_g0, %next_u0 = scf.if %valid0 -> (vector<4xf32>, vector<4xf32>) { + %a = vector.load %av[%t0, %kk] : view<[%tokens]x[%k]xf32> -> vector<4xf32> + %gr = vector.fmaf %gh, %a, %acc_g0 : vector<4xf32> + %ur = vector.fmaf %uh, %a, %acc_u0 : vector<4xf32> + scf.yield %gr, %ur : vector<4xf32>, vector<4xf32> + } else { + scf.yield %acc_g0, %acc_u0 : vector<4xf32>, vector<4xf32> + } + %t1 = index.constant 1 : index + %valid1 = index.cmp ult, %t1, %tokens : index + %next_g1, %next_u1 = scf.if %valid1 -> (vector<4xf32>, vector<4xf32>) { + %a = vector.load %av[%t1, %kk] : view<[%tokens]x[%k]xf32> -> vector<4xf32> + %gr = vector.fmaf %gh, %a, %acc_g1 : vector<4xf32> + %ur = vector.fmaf %uh, %a, %acc_u1 : vector<4xf32> + scf.yield %gr, %ur : vector<4xf32>, vector<4xf32> + } else { + scf.yield %acc_g1, %acc_u1 : vector<4xf32>, vector<4xf32> + } + %t2 = index.constant 2 : index + %valid2 = index.cmp ult, %t2, %tokens : index + %next_g2, %next_u2 = scf.if %valid2 -> (vector<4xf32>, vector<4xf32>) { + %a = vector.load %av[%t2, %kk] : view<[%tokens]x[%k]xf32> -> vector<4xf32> + %gr = vector.fmaf %gh, %a, %acc_g2 : vector<4xf32> + %ur = vector.fmaf %uh, %a, %acc_u2 : vector<4xf32> + scf.yield %gr, %ur : vector<4xf32>, vector<4xf32> + } else { + scf.yield %acc_g2, %acc_u2 : vector<4xf32>, vector<4xf32> + } + %t3 = index.constant 3 : index + %valid3 = index.cmp ult, %t3, %tokens : index + %next_g3, %next_u3 = scf.if %valid3 -> (vector<4xf32>, vector<4xf32>) { + %a = vector.load %av[%t3, %kk] : view<[%tokens]x[%k]xf32> -> vector<4xf32> + %gr = vector.fmaf %gh, %a, %acc_g3 : vector<4xf32> + %ur = vector.fmaf %uh, %a, %acc_u3 : vector<4xf32> + scf.yield %gr, %ur : vector<4xf32>, vector<4xf32> + } else { + scf.yield %acc_g3, %acc_u3 : vector<4xf32>, vector<4xf32> + } + %t4 = index.constant 4 : index + %valid4 = index.cmp ult, %t4, %tokens : index + %next_g4, %next_u4 = scf.if %valid4 -> (vector<4xf32>, vector<4xf32>) { + %a = vector.load %av[%t4, %kk] : view<[%tokens]x[%k]xf32> -> vector<4xf32> + %gr = vector.fmaf %gh, %a, %acc_g4 : vector<4xf32> + %ur = vector.fmaf %uh, %a, %acc_u4 : vector<4xf32> + scf.yield %gr, %ur : vector<4xf32>, vector<4xf32> + } else { + scf.yield %acc_g4, %acc_u4 : vector<4xf32>, vector<4xf32> + } + scf.yield %next_g0, %next_u0, %next_g1, %next_u1, %next_g2, %next_u2, %next_g3, %next_u3, %next_g4, %next_u4 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> + } + %leader = index.cmp eq, %lane16, %c0 : index + %p0 = index.constant 0 : index + %publishes0 = index.cmp ult, %p0, %tokens : index + scf.if %publishes0 { + %g_lane = vector.reduce %total_g0, %zero_scalar : vector<4xf32>, f32 + %u_lane = vector.reduce %total_u0, %zero_scalar : vector<4xf32>, f32 + %g_width = scalar.constant 32 : i32 + %g_xor1 = scalar.constant 1 : i32 + %g_peer1, %g_ok1 = kernel.subgroup.shuffle %g_lane, %g_xor1, %g_width : f32, i32, i32 + %g_sum1 = scalar.addf %g_lane, %g_peer1 : f32 + %g_xor2 = scalar.constant 2 : i32 + %g_peer2, %g_ok2 = kernel.subgroup.shuffle %g_sum1, %g_xor2, %g_width : f32, i32, i32 + %g_sum2 = scalar.addf %g_sum1, %g_peer2 : f32 + %g_xor4 = scalar.constant 4 : i32 + %g_peer4, %g_ok4 = kernel.subgroup.shuffle %g_sum2, %g_xor4, %g_width : f32, i32, i32 + %g_sum4 = scalar.addf %g_sum2, %g_peer4 : f32 + %g_xor8 = scalar.constant 8 : i32 + %g_peer8, %g_ok8 = kernel.subgroup.shuffle %g_sum4, %g_xor8, %g_width : f32, i32, i32 + %g = scalar.addf %g_sum4, %g_peer8 : f32 + %u_width = scalar.constant 32 : i32 + %u_xor1 = scalar.constant 1 : i32 + %u_peer1, %u_ok1 = kernel.subgroup.shuffle %u_lane, %u_xor1, %u_width : f32, i32, i32 + %u_sum1 = scalar.addf %u_lane, %u_peer1 : f32 + %u_xor2 = scalar.constant 2 : i32 + %u_peer2, %u_ok2 = kernel.subgroup.shuffle %u_sum1, %u_xor2, %u_width : f32, i32, i32 + %u_sum2 = scalar.addf %u_sum1, %u_peer2 : f32 + %u_xor4 = scalar.constant 4 : i32 + %u_peer4, %u_ok4 = kernel.subgroup.shuffle %u_sum2, %u_xor4, %u_width : f32, i32, i32 + %u_sum4 = scalar.addf %u_sum2, %u_peer4 : f32 + %u_xor8 = scalar.constant 8 : i32 + %u_peer8, %u_ok8 = kernel.subgroup.shuffle %u_sum4, %u_xor8, %u_width : f32, i32, i32 + %u = scalar.addf %u_sum4, %u_peer8 : f32 + scf.if %leader { + %gv = vector.splat %g : vector<4xf32> + %uv = vector.splat %u : vector<4xf32> + %result = template.apply<@ggml.mul_mat_swiglu.apply_vector4>(%op, %gv, %uv) : (index, vector<4xf32>, vector<4xf32>) -> (vector<4xf32>) + %value = vector.extract %result[0] : vector<4xf32> -> f32 + view.store %value, %ov[%p0, %row] : f32, view<[%tokens]x[%n]xf32> + } + } + %p1 = index.constant 1 : index + %publishes1 = index.cmp ult, %p1, %tokens : index + scf.if %publishes1 { + %g_lane = vector.reduce %total_g1, %zero_scalar : vector<4xf32>, f32 + %u_lane = vector.reduce %total_u1, %zero_scalar : vector<4xf32>, f32 + %g_width = scalar.constant 32 : i32 + %g_xor1 = scalar.constant 1 : i32 + %g_peer1, %g_ok1 = kernel.subgroup.shuffle %g_lane, %g_xor1, %g_width : f32, i32, i32 + %g_sum1 = scalar.addf %g_lane, %g_peer1 : f32 + %g_xor2 = scalar.constant 2 : i32 + %g_peer2, %g_ok2 = kernel.subgroup.shuffle %g_sum1, %g_xor2, %g_width : f32, i32, i32 + %g_sum2 = scalar.addf %g_sum1, %g_peer2 : f32 + %g_xor4 = scalar.constant 4 : i32 + %g_peer4, %g_ok4 = kernel.subgroup.shuffle %g_sum2, %g_xor4, %g_width : f32, i32, i32 + %g_sum4 = scalar.addf %g_sum2, %g_peer4 : f32 + %g_xor8 = scalar.constant 8 : i32 + %g_peer8, %g_ok8 = kernel.subgroup.shuffle %g_sum4, %g_xor8, %g_width : f32, i32, i32 + %g = scalar.addf %g_sum4, %g_peer8 : f32 + %u_width = scalar.constant 32 : i32 + %u_xor1 = scalar.constant 1 : i32 + %u_peer1, %u_ok1 = kernel.subgroup.shuffle %u_lane, %u_xor1, %u_width : f32, i32, i32 + %u_sum1 = scalar.addf %u_lane, %u_peer1 : f32 + %u_xor2 = scalar.constant 2 : i32 + %u_peer2, %u_ok2 = kernel.subgroup.shuffle %u_sum1, %u_xor2, %u_width : f32, i32, i32 + %u_sum2 = scalar.addf %u_sum1, %u_peer2 : f32 + %u_xor4 = scalar.constant 4 : i32 + %u_peer4, %u_ok4 = kernel.subgroup.shuffle %u_sum2, %u_xor4, %u_width : f32, i32, i32 + %u_sum4 = scalar.addf %u_sum2, %u_peer4 : f32 + %u_xor8 = scalar.constant 8 : i32 + %u_peer8, %u_ok8 = kernel.subgroup.shuffle %u_sum4, %u_xor8, %u_width : f32, i32, i32 + %u = scalar.addf %u_sum4, %u_peer8 : f32 + scf.if %leader { + %gv = vector.splat %g : vector<4xf32> + %uv = vector.splat %u : vector<4xf32> + %result = template.apply<@ggml.mul_mat_swiglu.apply_vector4>(%op, %gv, %uv) : (index, vector<4xf32>, vector<4xf32>) -> (vector<4xf32>) + %value = vector.extract %result[0] : vector<4xf32> -> f32 + view.store %value, %ov[%p1, %row] : f32, view<[%tokens]x[%n]xf32> + } + } + %p2 = index.constant 2 : index + %publishes2 = index.cmp ult, %p2, %tokens : index + scf.if %publishes2 { + %g_lane = vector.reduce %total_g2, %zero_scalar : vector<4xf32>, f32 + %u_lane = vector.reduce %total_u2, %zero_scalar : vector<4xf32>, f32 + %g_width = scalar.constant 32 : i32 + %g_xor1 = scalar.constant 1 : i32 + %g_peer1, %g_ok1 = kernel.subgroup.shuffle %g_lane, %g_xor1, %g_width : f32, i32, i32 + %g_sum1 = scalar.addf %g_lane, %g_peer1 : f32 + %g_xor2 = scalar.constant 2 : i32 + %g_peer2, %g_ok2 = kernel.subgroup.shuffle %g_sum1, %g_xor2, %g_width : f32, i32, i32 + %g_sum2 = scalar.addf %g_sum1, %g_peer2 : f32 + %g_xor4 = scalar.constant 4 : i32 + %g_peer4, %g_ok4 = kernel.subgroup.shuffle %g_sum2, %g_xor4, %g_width : f32, i32, i32 + %g_sum4 = scalar.addf %g_sum2, %g_peer4 : f32 + %g_xor8 = scalar.constant 8 : i32 + %g_peer8, %g_ok8 = kernel.subgroup.shuffle %g_sum4, %g_xor8, %g_width : f32, i32, i32 + %g = scalar.addf %g_sum4, %g_peer8 : f32 + %u_width = scalar.constant 32 : i32 + %u_xor1 = scalar.constant 1 : i32 + %u_peer1, %u_ok1 = kernel.subgroup.shuffle %u_lane, %u_xor1, %u_width : f32, i32, i32 + %u_sum1 = scalar.addf %u_lane, %u_peer1 : f32 + %u_xor2 = scalar.constant 2 : i32 + %u_peer2, %u_ok2 = kernel.subgroup.shuffle %u_sum1, %u_xor2, %u_width : f32, i32, i32 + %u_sum2 = scalar.addf %u_sum1, %u_peer2 : f32 + %u_xor4 = scalar.constant 4 : i32 + %u_peer4, %u_ok4 = kernel.subgroup.shuffle %u_sum2, %u_xor4, %u_width : f32, i32, i32 + %u_sum4 = scalar.addf %u_sum2, %u_peer4 : f32 + %u_xor8 = scalar.constant 8 : i32 + %u_peer8, %u_ok8 = kernel.subgroup.shuffle %u_sum4, %u_xor8, %u_width : f32, i32, i32 + %u = scalar.addf %u_sum4, %u_peer8 : f32 + scf.if %leader { + %gv = vector.splat %g : vector<4xf32> + %uv = vector.splat %u : vector<4xf32> + %result = template.apply<@ggml.mul_mat_swiglu.apply_vector4>(%op, %gv, %uv) : (index, vector<4xf32>, vector<4xf32>) -> (vector<4xf32>) + %value = vector.extract %result[0] : vector<4xf32> -> f32 + view.store %value, %ov[%p2, %row] : f32, view<[%tokens]x[%n]xf32> + } + } + %p3 = index.constant 3 : index + %publishes3 = index.cmp ult, %p3, %tokens : index + scf.if %publishes3 { + %g_lane = vector.reduce %total_g3, %zero_scalar : vector<4xf32>, f32 + %u_lane = vector.reduce %total_u3, %zero_scalar : vector<4xf32>, f32 + %g_width = scalar.constant 32 : i32 + %g_xor1 = scalar.constant 1 : i32 + %g_peer1, %g_ok1 = kernel.subgroup.shuffle %g_lane, %g_xor1, %g_width : f32, i32, i32 + %g_sum1 = scalar.addf %g_lane, %g_peer1 : f32 + %g_xor2 = scalar.constant 2 : i32 + %g_peer2, %g_ok2 = kernel.subgroup.shuffle %g_sum1, %g_xor2, %g_width : f32, i32, i32 + %g_sum2 = scalar.addf %g_sum1, %g_peer2 : f32 + %g_xor4 = scalar.constant 4 : i32 + %g_peer4, %g_ok4 = kernel.subgroup.shuffle %g_sum2, %g_xor4, %g_width : f32, i32, i32 + %g_sum4 = scalar.addf %g_sum2, %g_peer4 : f32 + %g_xor8 = scalar.constant 8 : i32 + %g_peer8, %g_ok8 = kernel.subgroup.shuffle %g_sum4, %g_xor8, %g_width : f32, i32, i32 + %g = scalar.addf %g_sum4, %g_peer8 : f32 + %u_width = scalar.constant 32 : i32 + %u_xor1 = scalar.constant 1 : i32 + %u_peer1, %u_ok1 = kernel.subgroup.shuffle %u_lane, %u_xor1, %u_width : f32, i32, i32 + %u_sum1 = scalar.addf %u_lane, %u_peer1 : f32 + %u_xor2 = scalar.constant 2 : i32 + %u_peer2, %u_ok2 = kernel.subgroup.shuffle %u_sum1, %u_xor2, %u_width : f32, i32, i32 + %u_sum2 = scalar.addf %u_sum1, %u_peer2 : f32 + %u_xor4 = scalar.constant 4 : i32 + %u_peer4, %u_ok4 = kernel.subgroup.shuffle %u_sum2, %u_xor4, %u_width : f32, i32, i32 + %u_sum4 = scalar.addf %u_sum2, %u_peer4 : f32 + %u_xor8 = scalar.constant 8 : i32 + %u_peer8, %u_ok8 = kernel.subgroup.shuffle %u_sum4, %u_xor8, %u_width : f32, i32, i32 + %u = scalar.addf %u_sum4, %u_peer8 : f32 + scf.if %leader { + %gv = vector.splat %g : vector<4xf32> + %uv = vector.splat %u : vector<4xf32> + %result = template.apply<@ggml.mul_mat_swiglu.apply_vector4>(%op, %gv, %uv) : (index, vector<4xf32>, vector<4xf32>) -> (vector<4xf32>) + %value = vector.extract %result[0] : vector<4xf32> -> f32 + view.store %value, %ov[%p3, %row] : f32, view<[%tokens]x[%n]xf32> + } + } + %p4 = index.constant 4 : index + %publishes4 = index.cmp ult, %p4, %tokens : index + scf.if %publishes4 { + %g_lane = vector.reduce %total_g4, %zero_scalar : vector<4xf32>, f32 + %u_lane = vector.reduce %total_u4, %zero_scalar : vector<4xf32>, f32 + %g_width = scalar.constant 32 : i32 + %g_xor1 = scalar.constant 1 : i32 + %g_peer1, %g_ok1 = kernel.subgroup.shuffle %g_lane, %g_xor1, %g_width : f32, i32, i32 + %g_sum1 = scalar.addf %g_lane, %g_peer1 : f32 + %g_xor2 = scalar.constant 2 : i32 + %g_peer2, %g_ok2 = kernel.subgroup.shuffle %g_sum1, %g_xor2, %g_width : f32, i32, i32 + %g_sum2 = scalar.addf %g_sum1, %g_peer2 : f32 + %g_xor4 = scalar.constant 4 : i32 + %g_peer4, %g_ok4 = kernel.subgroup.shuffle %g_sum2, %g_xor4, %g_width : f32, i32, i32 + %g_sum4 = scalar.addf %g_sum2, %g_peer4 : f32 + %g_xor8 = scalar.constant 8 : i32 + %g_peer8, %g_ok8 = kernel.subgroup.shuffle %g_sum4, %g_xor8, %g_width : f32, i32, i32 + %g = scalar.addf %g_sum4, %g_peer8 : f32 + %u_width = scalar.constant 32 : i32 + %u_xor1 = scalar.constant 1 : i32 + %u_peer1, %u_ok1 = kernel.subgroup.shuffle %u_lane, %u_xor1, %u_width : f32, i32, i32 + %u_sum1 = scalar.addf %u_lane, %u_peer1 : f32 + %u_xor2 = scalar.constant 2 : i32 + %u_peer2, %u_ok2 = kernel.subgroup.shuffle %u_sum1, %u_xor2, %u_width : f32, i32, i32 + %u_sum2 = scalar.addf %u_sum1, %u_peer2 : f32 + %u_xor4 = scalar.constant 4 : i32 + %u_peer4, %u_ok4 = kernel.subgroup.shuffle %u_sum2, %u_xor4, %u_width : f32, i32, i32 + %u_sum4 = scalar.addf %u_sum2, %u_peer4 : f32 + %u_xor8 = scalar.constant 8 : i32 + %u_peer8, %u_ok8 = kernel.subgroup.shuffle %u_sum4, %u_xor8, %u_width : f32, i32, i32 + %u = scalar.addf %u_sum4, %u_peer8 : f32 + scf.if %leader { + %gv = vector.splat %g : vector<4xf32> + %uv = vector.splat %u : vector<4xf32> + %result = template.apply<@ggml.mul_mat_swiglu.apply_vector4>(%op, %gv, %uv) : (index, vector<4xf32>, vector<4xf32>) -> (vector<4xf32>) + %value = vector.extract %result[0] : vector<4xf32> -> f32 + view.store %value, %ov[%p4, %row] : f32, view<[%tokens]x[%n]xf32> + } + } + } + func.return +} + +kernel.def target(@ggml_mul_mat_swiglu_gfx11_wave32) @ggml_mul_mat_swiglu_f32_f32_lowtoken_dot(%token_count: index) { + %output_size = config.get @ggml.mul_mat_swiglu.output_size : index + %c1 = index.constant 1 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c128 = index.constant 128 : index + %last_row = index.add %output_size, %c7 : index + %channel_tiles = index.div %last_row, %c8 : index + kernel.launch.config workgroups(%channel_tiles, %c1, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) where [range(%token_count, 1, 5)] { + %input_size = config.get @ggml.mul_mat_swiglu.input_size : index + %output_size = config.get @ggml.mul_mat_swiglu.output_size : index + %gate_format = config.get @ggml.mul_mat_swiglu.gate_weight_format : index + %up_format = config.get @ggml.mul_mat_swiglu.up_weight_format : index + %tokens = index.assume %token_count [range(%token_count, 1, 5)] : index + %channel_tile = kernel.workgroup.id : index + %op = config.get @ggml.mul_mat_swiglu.op : index + func.call @ggml_mul_mat_swiglu_lowtoken_dot(%op, %gate_format, %up_format, %tokens, %input_size, %output_size, %channel_tile, %input, %gate_weight, %up_weight, %output) : (index, index, index, index, index, index, index, buffer, buffer, buffer, buffer) + kernel.return +} + +template.decl @ggml.binary_f32.apply(%op: index, %lhs: f32, %rhs: f32) -> (f32) + +func.def inline @ggml_mul_mat_swiglu_apply_scalar(%op: index, %lhs: f32, %rhs: f32) -> (f32) { + %op_swiglu = index.constant 4 : index + %is_swiglu = index.cmp eq, %op, %op_swiglu : index + %result = scf.if %is_swiglu -> (f32) { + %sigmoid = scalar.logisticf %lhs : f32 + %silu = scalar.mulf %lhs, %sigmoid : f32 + %value = scalar.mulf %silu, %rhs : f32 + scf.yield %value : f32 + } else { + %value = template.apply<@ggml.binary_f32.apply>(%op, %lhs, %rhs) : (index, f32, f32) -> (f32) + scf.yield %value : f32 + } + func.return %result : f32 +} + +// Q4_K weights retain their native affine values; only activations use the +// existing Q8_1_x4 producer. Four 16-lane cohorts share each wave's K256 step. +func.def inline @ggml_swiglu_q8_values(%local: i1, %input: buffer, %stage: buffer, %record_base: offset, %word0: index, %word1: index, %meta0: index, %meta1: index, %sum0: index, %sum1: index) -> (vector<2xi32>, vector<2xi32>, f32, f32, f32, f32) { + %payload_offset = index.constant 16 : offset + %payload_base = index.add %record_base, %payload_offset : offset + %a0, %a1, %d0, %d1, %s0, %s1 = scf.if %local -> (vector<2xi32>, vector<2xi32>, f32, f32, f32, f32) { + %payload = buffer.view %stage[%payload_base] : buffer -> view<32xi32> + %metadata = buffer.view %stage[%record_base] : buffer -> view<8xf16> + %a0 = vector.load %payload[%word0] : view<32xi32> -> vector<2xi32> + %a1 = vector.load %payload[%word1] : view<32xi32> -> vector<2xi32> + %packet = vector.load %metadata[%meta0] : view<8xf16> -> vector<4xf16> + %values = vector.extf %packet : vector<4xf16> to vector<4xf32> + %d0 = vector.extract %values[0] : vector<4xf32> -> f32 + %d1 = vector.extract %values[2] : vector<4xf32> -> f32 + %s0 = vector.extract %values[1] : vector<4xf32> -> f32 + %s1 = vector.extract %values[3] : vector<4xf32> -> f32 + scf.yield %a0, %a1, %d0, %d1, %s0, %s1 : vector<2xi32>, vector<2xi32>, f32, f32, f32, f32 + } else { + %payload = buffer.view %input[%payload_base] : buffer -> view<32xi32> + %metadata = buffer.view %input[%record_base] : buffer -> view<8xf16> + %a0 = vector.load %payload[%word0] : view<32xi32> -> vector<2xi32> + %a1 = vector.load %payload[%word1] : view<32xi32> -> vector<2xi32> + %d0h = view.load %metadata[%meta0] : view<8xf16> -> f16 + %d1h = view.load %metadata[%meta1] : view<8xf16> -> f16 + %d0 = scalar.extf %d0h : f16 to f32 + %d1 = scalar.extf %d1h : f16 to f32 + %s0h = view.load %metadata[%sum0] : view<8xf16> -> f16 + %s1h = view.load %metadata[%sum1] : view<8xf16> -> f16 + %s0 = scalar.extf %s0h : f16 to f32 + %s1 = scalar.extf %s1h : f16 to f32 + scf.yield %a0, %a1, %d0, %d1, %s0, %s1 : vector<2xi32>, vector<2xi32>, f32, f32, f32, f32 + } + func.return %a0, %a1, %d0, %d1, %s0, %s1 : vector<2xi32>, vector<2xi32>, f32, f32, f32, f32 +} + +func.def inline @ggml_q4_paired_prefetch_words(%gate: buffer, %up: buffer, %group_block: index, %row_offset: offset, %row_lane: index, %payload_field: index, %payload_word: index) -> (vector<4xi32>, vector<4xi32>, vector<2xi32>, vector<2xi32>) { + %c0 = index.constant 0 : index + %record_bytes = index.constant 9216 : offset + %payload_add = index.constant 1024 : offset + %group_base = index.scale %group_block, %record_bytes : index, offset -> offset + %header_base = index.add %group_base, %row_offset : offset + %payload_base = index.add %group_base, %payload_add : offset + %ghv = buffer.view %gate[%header_base] : buffer -> view<4xi32> + %uhv = buffer.view %up[%header_base] : buffer -> view<4xi32> + %gpv = buffer.view %gate[%payload_base] : buffer -> view<8x64x4xi32> + %upv = buffer.view %up[%payload_base] : buffer -> view<8x64x4xi32> + %gh = vector.load %ghv[%c0] : view<4xi32> -> vector<4xi32> + %gr = vector.load %gpv[%payload_field, %row_lane, %payload_word] : view<8x64x4xi32> -> vector<2xi32> + %uh = vector.load %uhv[%c0] : view<4xi32> -> vector<4xi32> + %ur = vector.load %upv[%payload_field, %row_lane, %payload_word] : view<8x64x4xi32> -> vector<2xi32> + func.return %gh, %uh, %gr, %ur : vector<4xi32>, vector<4xi32>, vector<2xi32>, vector<2xi32> +} + +// Two K teams amortize native weight service on large paired projections. +// Keep smaller/cache-resident and odd-block reductions on the original mapping. +func.def pure inline @ggml_swiglu_q8_reduction_partitions(%local: i1, %input_size: index, %output_size: index) -> (index) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c256 = index.constant 256 : index + %minimum_blocks = index.constant 262144 : index + %blocks = index.div %input_size, %c256 : index + %tail = index.rem %blocks, %c2 : index + %even = index.cmp eq, %tail, %c0 : index + %weight_blocks = index.mul %blocks, %output_size : index + %large = index.cmp uge, %weight_blocks, %minimum_blocks : index + %eligible = scalar.andi %local, %even : i1 + %partitioned = scalar.andi %eligible, %large : i1 + %partitions = scf.select %partitioned, %c2, %c1 : index + func.return %partitions : index +} + +func.decl @ggml_lowtoken_reduce_cohort_f32(%value: f32) -> (f32) + +func.def inline @ggml_swiglu_q8_reduce(%value: f32, %partitioned: i1) -> (f32) { + %width = scalar.constant 32 : i32 + %sum8 = func.call @ggml_lowtoken_reduce_cohort_f32(%value) : (f32) -> (f32) + %result = scf.if %partitioned -> (f32) { + %off16 = scalar.constant 16 : i32 + %peer16, %ok16 = kernel.subgroup.shuffle %sum8, %off16, %width : f32, i32, i32 + %sum = scalar.addf %sum8, %peer16 : f32 + scf.yield %sum : f32 + } else { + %zero = scalar.constant 0.0 : f32 + %sum = scalar.addf %sum8, %zero : f32 + scf.yield %sum : f32 + } + func.return %result : f32 +} + +template.decl @ggml.mul_mat_swiglu.q4_q8_lowtoken_body(%publish_q8: i1, %token_count: index, %input: buffer, %gate: buffer, %up: buffer, %output: buffer, %q8_output: buffer) + +template.decl @ggml.quantize_q8_1_x4.publish_vector4_strict(%publish_word: i1, %token_output_byte_base: offset, %channel: index, %values: vector<4xf32>, %scratch_values: view<256xf32>, %scratch_d: view<32xf32>, %output: buffer) + +kernel.def target(@ggml_mul_mat_swiglu_gfx11_wave64) @ggml_mul_mat_swiglu_q4_q8_1_x4_lowtoken_dot(%token_count: index) { + %n = config.get @ggml.mul_mat_swiglu.output_size : index + %k = config.get @ggml.mul_mat_swiglu.input_size : index + %capacity = config.get @ggml.workload.token_capacity : index + %local_q8_c1 = index.constant 1 : index + %local_q8_c31 = index.constant 31 : index + %local_q8_c32 = index.constant 32 : index + %local_q8_c128 = index.constant 128 : index + %local_q8_c144 = index.constant 144 : index + %local_q8_c32768 = index.constant 32768 : index + %local_q8_groups = index.div %k, %local_q8_c128 : index + %local_q8_records = index.mul %local_q8_groups, %capacity : index + %local_q8_bytes = index.mul %local_q8_records, %local_q8_c144 : index + %local_q8_padded_rows = index.add %n, %local_q8_c31 : index + %local_q8_workgroups = index.div %local_q8_padded_rows, %local_q8_c32 : index + %local_q8_batched = index.cmp ugt, %capacity, %local_q8_c1 : index + %local_q8_fits = index.cmp ule, %local_q8_bytes, %local_q8_c32768 : index + %local_q8_enough_work = index.cmp uge, %local_q8_workgroups, %local_q8_c128 : index + %local_q8_eligible = scalar.andi %local_q8_batched, %local_q8_fits : i1 + %local = scalar.andi %local_q8_eligible, %local_q8_enough_work : i1 + %c1 = index.constant 1 : index + %c8 = index.constant 8 : index + %c32 = index.constant 32 : index + %c128 = index.constant 128 : index + %c512 = index.constant 512 : index + %partitions = func.call pure @ggml_swiglu_q8_reduction_partitions(%local, %k, %n) : (i1, index, index) -> (index) + %local_rows = index.div %c32, %partitions : index + %rows = scf.select %local, %local_rows, %c8 : index + %padding = index.sub %rows, %c1 : index + %padded_n = index.add %n, %padding : index + %groups = index.div %padded_n, %rows : index + %threads = scf.select %local, %c512, %c128 : index + kernel.launch.config workgroups(%groups, %c1, %c1) workgroup_size(%threads, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %gate: buffer, %up: buffer, %output: buffer) { + %publish_q8 = scalar.constant false : i1 + template.apply<@ggml.mul_mat_swiglu.q4_q8_lowtoken_body>(%publish_q8, %token_count, %input, %gate, %up, %output, %output) : (i1, index, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +kernel.def target(@ggml_mul_mat_swiglu_gfx11_wave64) @ggml_mul_mat_swiglu_q4_q8_1_x4_lowtoken_dot_q8_output(%token_count: index) { + %n0 = config.get @ggml.mul_mat_swiglu.output_size : index + %n = index.assume %n0 [range(%n0, 1, 262144), mul(%n0, 128)] : index + %k0 = config.get @ggml.mul_mat_swiglu.input_size : index + %k = index.assume %k0 [range(%k0, 256, 32768), mul(%k0, 256)] : index + %capacity0 = config.get @ggml.workload.token_capacity : index + %capacity = index.assume %capacity0 [range(%capacity0, 1, 5)] : index + %local_q8_c1 = index.constant 1 : index + %local_q8_c31 = index.constant 31 : index + %local_q8_c32 = index.constant 32 : index + %local_q8_c128 = index.constant 128 : index + %local_q8_c144 = index.constant 144 : index + %local_q8_c32768 = index.constant 32768 : index + %local_q8_groups = index.div %k, %local_q8_c128 : index + %local_q8_records = index.mul %local_q8_groups, %capacity : index + %local_q8_bytes = index.mul %local_q8_records, %local_q8_c144 : index + %local_q8_padded_rows = index.add %n, %local_q8_c31 : index + %local_q8_workgroups = index.div %local_q8_padded_rows, %local_q8_c32 : index + %local_q8_batched = index.cmp ugt, %capacity, %local_q8_c1 : index + %local_q8_fits = index.cmp ule, %local_q8_bytes, %local_q8_c32768 : index + %local_q8_enough_work = index.cmp uge, %local_q8_workgroups, %local_q8_c128 : index + %local_q8_eligible = scalar.andi %local_q8_batched, %local_q8_fits : i1 + %c1 = index.constant 1 : index + %c8 = index.constant 8 : index + %c32 = index.constant 32 : index + %c128 = index.constant 128 : index + %c512 = index.constant 512 : index + %c1024 = index.constant 1024 : index + %local = scalar.andi %local_q8_eligible, %local_q8_enough_work : i1 + %partitions = func.call pure @ggml_swiglu_q8_reduction_partitions(%local, %k, %n) : (i1, index, index) -> (index) + %partitioned = index.cmp ugt, %partitions, %c1 : index + %rows = index.constant 32 : index + %padding = index.sub %rows, %c1 : index + %padded_n = index.add %n, %padding : index + %groups = index.div %padded_n, %rows : index + %threads = scf.select %partitioned, %c1024, %c512 : index + kernel.launch.config workgroups(%groups, %c1, %c1) workgroup_size(%threads, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %gate: buffer, %up: buffer, %output: buffer, %q8_output: buffer) where [range(%token_count, 1, 5)] { + %publish_q8 = scalar.constant true : i1 + template.apply<@ggml.mul_mat_swiglu.q4_q8_lowtoken_body>(%publish_q8, %token_count, %input, %gate, %up, %output, %q8_output) : (i1, index, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +template.def<@ggml.mul_mat_swiglu.q4_q8_lowtoken_body> device @ggml_mul_mat_swiglu_q4_q8_lowtoken_body(%publish_q8: i1, %token_count: index, %input: buffer, %gate: buffer, %up: buffer, %output: buffer, %q8_output: buffer) { + %op = config.get @ggml.mul_mat_swiglu.op : index + %k0 = config.get @ggml.mul_mat_swiglu.input_size : index + %n0 = config.get @ggml.mul_mat_swiglu.output_size : index + %capacity = config.get @ggml.workload.token_capacity : index + %tokens = index.assume %token_count [range(%token_count, 1, 5), le(%token_count, %capacity)] : index + %k = index.assume %k0 [range(%k0, 256, 32768), mul(%k0, 256)] : index + %n = index.assume %n0 [range(%n0, 1, 262144)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %base = index.constant 0 : offset + %block_bytes = index.constant 144 : offset + %code_offset = index.constant 16 : offset + %zero = scalar.constant 0.0 : f32 + %lane0 = kernel.subgroup.lane.id : index + %lane = index.assume %lane0 [range(%lane0, 0, 63)] : index + %wave = kernel.subgroup.id : index + %wg = kernel.workgroup.id : index + %local_q8_c1 = index.constant 1 : index + %local_q8_c31 = index.constant 31 : index + %local_q8_c32 = index.constant 32 : index + %local_q8_c128 = index.constant 128 : index + %local_q8_c144 = index.constant 144 : index + %local_q8_c32768 = index.constant 32768 : index + %local_q8_groups = index.div %k, %local_q8_c128 : index + %local_q8_records = index.mul %local_q8_groups, %capacity : index + %local_q8_bytes = index.mul %local_q8_records, %local_q8_c144 : index + %local_q8_padded_rows = index.add %n, %local_q8_c31 : index + %local_q8_workgroups = index.div %local_q8_padded_rows, %local_q8_c32 : index + %local_q8_batched = index.cmp ugt, %capacity, %local_q8_c1 : index + %local_q8_fits = index.cmp ule, %local_q8_bytes, %local_q8_c32768 : index + %local_q8_enough_work = index.cmp uge, %local_q8_workgroups, %local_q8_c128 : index + %local_q8_eligible = scalar.andi %local_q8_batched, %local_q8_fits : i1 + %use_local = scalar.andi %local_q8_eligible, %local_q8_enough_work : i1 + %partitions = func.call pure @ggml_swiglu_q8_reduction_partitions(%use_local, %k, %n) : (i1, index, index) -> (index) + %partitioned = index.cmp ugt, %partitions, %c1 : index + %publication_rows = index.constant 32 : index + %publication_row_budget = index.mul %publication_rows, %partitions : index + %local_row_budget = scf.select %publish_q8, %publication_row_budget, %c32 : index + %local_rows = index.div %local_row_budget, %partitions : index + %default_rows_per_workgroup = scf.select %use_local, %local_rows, %c8 : index + %rows_per_workgroup = scf.select %publish_q8, %publication_rows, %default_rows_per_workgroup : index + %cohort_width = index.mul %c16, %partitions : index + %wave_rows = index.div %c4, %partitions : index + %row_base = index.mul %wg, %rows_per_workgroup : index + %wave_row = index.mul %wave, %wave_rows : index + %row_half = index.div %lane, %cohort_width : index + %local_row = index.add %wave_row, %row_half : index + %row = index.add %row_base, %local_row : index + %valid_row = index.cmp ult, %row, %n : index + %cohort = index.div %lane, %cohort_width : index + %team_lane = index.rem %lane, %cohort_width : index + %partition = index.div %team_lane, %c16 : index + %block_lane = index.rem %lane, %c16 : index + %pair = index.div %block_lane, %c4 : index + %packet = index.rem %block_lane, %c4 : index + %group0 = index.mul %pair, %c2 : index + %group1 = index.add %group0, %c1 : index + %word_page = index.mul %pair, %c8 : index + %word_add = index.mul %packet, %c2 : index + %word_index = index.add %word_page, %word_add : index + %meta_lane = index.cmp eq, %packet, %c0 : index + %blocks = index.div %k, %c256 : index + %partition_blocks = index.div %blocks, %partitions : index + %partition_origin = index.mul %partition, %partition_blocks : index + %weight_row_bytes = index.scale %blocks, %block_bytes : index, offset -> offset + %weight_row_base = index.scale %row, %weight_row_bytes : index, offset -> offset + %q8_groups = index.div %k, %c128 : index + %q8_row_bytes = index.scale %q8_groups, %block_bytes : index, offset -> offset + %q8_half = index.div %pair, %c2 : index + %q8_pair = index.rem %pair, %c2 : index + %q8_group0 = index.mul %q8_pair, %c2 : index + %q8_group1 = index.add %q8_group0, %c1 : index + %q8_words_base = index.mul %q8_group0, %c8 : index + %q8_word0 = index.add %q8_words_base, %word_add : index + %q8_word1 = index.add %q8_word0, %c8 : index + %q8_meta0 = index.mul %q8_group0, %c2 : index + %q8_meta1 = index.mul %q8_group1, %c2 : index + %q8_sum0 = index.add %q8_meta0, %c1 : index + %q8_sum1 = index.add %q8_meta1, %c1 : index + %c64 = index.constant 64 : index + %record_bytes = index.constant 9216 : offset + %row_field_bytes = index.constant 16 : offset + %payload_add = index.constant 1024 : offset + %row_group = index.div %row, %c64 : index + %row_lane = index.rem %row, %c64 : index + %row_group_start = index.mul %row_group, %blocks : index + %row_group_block = index.add %row_group_start, %partition_origin : index + %row_offset = index.scale %row_lane, %row_field_bytes : index, offset -> offset + %payload_field = index.div %word_index, %c4 : index + %payload_word = index.rem %word_index, %c4 : index + %input_na, %gate_na, %up_na, %output_na = buffer.assume.noalias %input, %gate, %up, %output : buffer, buffer, buffer, buffer + // One immutable Q8 activation record stream is shared by all row cohorts. + %activation_rows = index.mul %capacity, %q8_groups : index + %activation_active_rows = index.mul %tokens, %q8_groups : index + %activation_bytes = index.scale %activation_rows, %block_bytes : index, offset -> offset + %local_bytes = scf.select %use_local, %activation_bytes, %base : offset + %activation_stage = buffer.alloca align(16) %local_bytes : buffer + %epilogue_capacity_bytes = index.constant 640 : offset + %epilogue_bytes = scf.select %publish_q8, %epilogue_capacity_bytes, %base : offset + %epilogue_stage = buffer.alloca align(16) %epilogue_bytes : buffer + %epilogue_view = buffer.view %epilogue_stage[%base] : buffer -> view<5x32xf32> + %epilogue_scratch_capacity_bytes = index.constant 1024 : offset + %epilogue_scratch_bytes = scf.select %publish_q8, %epilogue_scratch_capacity_bytes, %base : offset + %epilogue_scratch_stage = buffer.alloca align(16) %epilogue_scratch_bytes : buffer + %epilogue_row = index.rem %row, %c32 : index + scf.if %use_local { + %record_words = index.constant 36 : index + %record_packets = index.constant 9 : index + %activation_words = index.mul %activation_rows, %record_words : index + %activation_packets = index.mul %activation_rows, %record_packets : index + %active_packets = index.mul %activation_active_rows, %record_packets : index + %source_words = buffer.view %input_na[%base] : buffer -> view<[%activation_words]xi32> + %stage_words = buffer.view %activation_stage[%base] : buffer -> view<[%activation_words]xi32> + %stage_workgroup = kernel.workgroup.size : index + %stage_round_add = index.sub %stage_workgroup, %c1 : index + %stage_rounded = index.add %activation_packets, %stage_round_add : index + %stage_rounds = index.div %stage_rounded, %stage_workgroup : index + %stage_thread = kernel.workitem.id : index + scf.for %round = [%c0 to %stage_rounds step %c1] { + %round_origin = index.mul %round, %stage_workgroup : index + %packet = index.add %round_origin, %stage_thread : index + %live_packet = index.cmp ult, %packet, %active_packets : index + scf.if %live_packet { + %word = index.mul %packet, %c4 : index + %words = vector.load %source_words[%word] : view<[%activation_words]xi32> -> vector<4xi32> + vector.store %words, %stage_words[%word] : vector<4xi32>, view<[%activation_words]xi32> + scf.yield + } + scf.yield + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield + } + scf.if %valid_row { + %initial_g_headers, %initial_u_headers, %initial_g_raw, %initial_u_raw = scf.if %use_local -> (vector<4xi32>, vector<4xi32>, vector<2xi32>, vector<2xi32>) { + %gh, %uh, %gr, %ur = func.call @ggml_q4_paired_prefetch_words(%gate_na, %up_na, %row_group_block, %row_offset, %row_lane, %payload_field, %payload_word) : (buffer, buffer, index, offset, index, index, index) -> (vector<4xi32>, vector<4xi32>, vector<2xi32>, vector<2xi32>) + scf.yield %gh, %uh, %gr, %ur : vector<4xi32>, vector<4xi32>, vector<2xi32>, vector<2xi32> + } else { + %zero_header = vector.constant 0 : vector<4xi32> + %zero_payload = vector.constant 0 : vector<2xi32> + scf.yield %zero_header, %zero_header, %zero_payload, %zero_payload : vector<4xi32>, vector<4xi32>, vector<2xi32>, vector<2xi32> + } + %total_g0, %total_g1, %total_g2, %total_g3, %total_g4, %total_u0, %total_u1, %total_u2, %total_u3, %total_u4, %unused_g_headers, %unused_u_headers, %unused_g_raw, %unused_u_raw = scf.for %block_base = [%c0 to %partition_blocks step %c1](%acc_g0 = %zero : f32, %acc_g1 = %zero : f32, %acc_g2 = %zero : f32, %acc_g3 = %zero : f32, %acc_g4 = %zero : f32, %acc_u0 = %zero : f32, %acc_u1 = %zero : f32, %acc_u2 = %zero : f32, %acc_u3 = %zero : f32, %acc_u4 = %zero : f32, %current_g_headers = %initial_g_headers : vector<4xi32>, %current_u_headers = %initial_u_headers : vector<4xi32>, %current_g_raw = %initial_g_raw : vector<2xi32>, %current_u_raw = %initial_u_raw : vector<2xi32>) -> (f32, f32, f32, f32, f32, f32, f32, f32, f32, f32, vector<4xi32>, vector<4xi32>, vector<2xi32>, vector<2xi32>) { + %block = index.add %block_base, %partition_origin : index + %valid_block = index.cmp ult, %block, %blocks : index + %safe_block = scf.select %valid_block, %block, %c0 : index + %group_block = index.add %row_group_start, %safe_block : index + %group_base = index.scale %group_block, %record_bytes : index, offset -> offset + %wb = index.add %group_base, %row_offset : offset + %wc = index.add %group_base, %payload_add : offset + %g_header_view = buffer.view %gate_na[%wb] : buffer -> view<4xi32> + %g_code_view = buffer.view %gate_na[%wc] : buffer -> view<8x64x4xi32> + %g_headers = scf.if %use_local -> (vector<4xi32>) { + scf.yield %current_g_headers : vector<4xi32> + } else { + %loaded = vector.load %g_header_view[%c0] : view<4xi32> -> vector<4xi32> + scf.yield %loaded : vector<4xi32> + } + %g_dm_i32 = vector.extract %g_headers[0] : vector<4xi32> -> i32 + %g_dm_word = vector.splat %g_dm_i32 : vector<1xi32> + %g_dm = vector.bitcast %g_dm_word : vector<1xi32> to vector<2xf16> + %g_d0h = vector.extract %g_dm[0] : vector<2xf16> -> f16 + %g_d1h = vector.extract %g_dm[1] : vector<2xf16> -> f16 + %g_d0 = scalar.extf %g_d0h : f16 to f32 + %g_d1 = scalar.extf %g_d1h : f16 to f32 + %g_s0 = vector.extract %g_headers[1] : vector<4xi32> -> i32 + %g_s1 = vector.extract %g_headers[2] : vector<4xi32> -> i32 + %g_s2 = vector.extract %g_headers[3] : vector<4xi32> -> i32 + %g_scale0i, %g_min0i = func.call @ggml_q4k_scale_min_from_header(%g_s0, %g_s1, %g_s2, %group0) : (i32, i32, i32, index) -> (i32, i32) + %g_scale0f = scalar.uitofp %g_scale0i : i32 to f32 + %g_min0f = scalar.uitofp %g_min0i : i32 to f32 + %g_ws0 = scalar.mulf %g_d0, %g_scale0f : f32 + %g_wm0 = scalar.mulf %g_d1, %g_min0f : f32 + %g_scale1i, %g_min1i = func.call @ggml_q4k_scale_min_from_header(%g_s0, %g_s1, %g_s2, %group1) : (i32, i32, i32, index) -> (i32, i32) + %g_scale1f = scalar.uitofp %g_scale1i : i32 to f32 + %g_min1f = scalar.uitofp %g_min1i : i32 to f32 + %g_ws1 = scalar.mulf %g_d0, %g_scale1f : f32 + %g_wm1 = scalar.mulf %g_d1, %g_min1f : f32 + %g_raw = scf.if %use_local -> (vector<2xi32>) { + scf.yield %current_g_raw : vector<2xi32> + } else { + %loaded = vector.load %g_code_view[%payload_field, %row_lane, %payload_word] : view<8x64x4xi32> -> vector<2xi32> + scf.yield %loaded : vector<2xi32> + } + %g_mask = vector.constant 252645135 : vector<2xi32> + %g_shift = vector.constant 4 : vector<2xi32> + %g_lo = vector.andi %g_raw, %g_mask : vector<2xi32> + %g_high = vector.shrui %g_raw, %g_shift : vector<2xi32> + %g_hi = vector.andi %g_high, %g_mask : vector<2xi32> + %u_header_view = buffer.view %up_na[%wb] : buffer -> view<4xi32> + %u_code_view = buffer.view %up_na[%wc] : buffer -> view<8x64x4xi32> + %u_headers = scf.if %use_local -> (vector<4xi32>) { + scf.yield %current_u_headers : vector<4xi32> + } else { + %loaded = vector.load %u_header_view[%c0] : view<4xi32> -> vector<4xi32> + scf.yield %loaded : vector<4xi32> + } + %u_dm_i32 = vector.extract %u_headers[0] : vector<4xi32> -> i32 + %u_dm_word = vector.splat %u_dm_i32 : vector<1xi32> + %u_dm = vector.bitcast %u_dm_word : vector<1xi32> to vector<2xf16> + %u_d0h = vector.extract %u_dm[0] : vector<2xf16> -> f16 + %u_d1h = vector.extract %u_dm[1] : vector<2xf16> -> f16 + %u_d0 = scalar.extf %u_d0h : f16 to f32 + %u_d1 = scalar.extf %u_d1h : f16 to f32 + %u_s0 = vector.extract %u_headers[1] : vector<4xi32> -> i32 + %u_s1 = vector.extract %u_headers[2] : vector<4xi32> -> i32 + %u_s2 = vector.extract %u_headers[3] : vector<4xi32> -> i32 + %u_scale0i, %u_min0i = func.call @ggml_q4k_scale_min_from_header(%u_s0, %u_s1, %u_s2, %group0) : (i32, i32, i32, index) -> (i32, i32) + %u_scale0f = scalar.uitofp %u_scale0i : i32 to f32 + %u_min0f = scalar.uitofp %u_min0i : i32 to f32 + %u_ws0 = scalar.mulf %u_d0, %u_scale0f : f32 + %u_wm0 = scalar.mulf %u_d1, %u_min0f : f32 + %u_scale1i, %u_min1i = func.call @ggml_q4k_scale_min_from_header(%u_s0, %u_s1, %u_s2, %group1) : (i32, i32, i32, index) -> (i32, i32) + %u_scale1f = scalar.uitofp %u_scale1i : i32 to f32 + %u_min1f = scalar.uitofp %u_min1i : i32 to f32 + %u_ws1 = scalar.mulf %u_d0, %u_scale1f : f32 + %u_wm1 = scalar.mulf %u_d1, %u_min1f : f32 + %u_raw = scf.if %use_local -> (vector<2xi32>) { + scf.yield %current_u_raw : vector<2xi32> + } else { + %loaded = vector.load %u_code_view[%payload_field, %row_lane, %payload_word] : view<8x64x4xi32> -> vector<2xi32> + scf.yield %loaded : vector<2xi32> + } + %u_mask = vector.constant 252645135 : vector<2xi32> + %u_shift = vector.constant 4 : vector<2xi32> + %u_lo = vector.andi %u_raw, %u_mask : vector<2xi32> + %u_high = vector.shrui %u_raw, %u_shift : vector<2xi32> + %u_hi = vector.andi %u_high, %u_mask : vector<2xi32> + %q8_block0 = index.mul %safe_block, %c2 : index + %q8_block = index.add %q8_block0, %q8_half : index + %q8_block_add = index.scale %q8_block, %block_bytes : index, offset -> offset + scf.if %use_local { + scf.schedule.fence + scf.yield + } + %next_block = index.add %block_base, %c1 : index + %has_next = index.cmp ult, %next_block, %partition_blocks : index + %prefetch = scalar.andi %use_local, %has_next : i1 + %next_g_headers, %next_u_headers, %next_g_raw, %next_u_raw = scf.if %prefetch -> (vector<4xi32>, vector<4xi32>, vector<2xi32>, vector<2xi32>) { + %next_group_block = index.add %row_group_block, %next_block : index + %gh, %uh, %gr, %ur = func.call @ggml_q4_paired_prefetch_words(%gate_na, %up_na, %next_group_block, %row_offset, %row_lane, %payload_field, %payload_word) : (buffer, buffer, index, offset, index, index, index) -> (vector<4xi32>, vector<4xi32>, vector<2xi32>, vector<2xi32>) + scf.yield %gh, %uh, %gr, %ur : vector<4xi32>, vector<4xi32>, vector<2xi32>, vector<2xi32> + } else { + %zero_header = vector.constant 0 : vector<4xi32> + %zero_payload = vector.constant 0 : vector<2xi32> + scf.yield %zero_header, %zero_header, %zero_payload, %zero_payload : vector<4xi32>, vector<4xi32>, vector<2xi32>, vector<2xi32> + } + %t0 = index.constant 0 : index + %valid0 = index.cmp ult, %t0, %tokens : index + %next_g0, %next_u0 = scf.if %valid0 -> (f32, f32) { + %tb = index.scale %t0, %q8_row_bytes : index, offset -> offset + %qb = index.add %tb, %q8_block_add : offset + %a0, %a1, %ad0, %ad1, %as0_raw, %as1_raw = func.call @ggml_swiglu_q8_values(%use_local, %input_na, %activation_stage, %qb, %q8_word0, %q8_word1, %q8_meta0, %q8_meta1, %q8_sum0, %q8_sum1) : (i1, buffer, buffer, offset, index, index, index, index, index, index) -> (vector<2xi32>, vector<2xi32>, f32, f32, f32, f32) + %as0 = scf.select %meta_lane, %as0_raw, %zero : f32 + %as1 = scf.select %meta_lane, %as1_raw, %zero : f32 + %g_dot0 = func.call @ggml_dot_u8_s8_vector8_f32(%g_lo, %a0) : (vector<2xi32>, vector<2xi32>) -> (f32) + %g_dot1 = func.call @ggml_dot_u8_s8_vector8_f32(%g_hi, %a1) : (vector<2xi32>, vector<2xi32>) -> (f32) + %g_s0 = scalar.mulf %g_ws0, %ad0 : f32 + %g_s1 = scalar.mulf %g_ws1, %ad1 : f32 + %g_c0 = scalar.mulf %g_dot0, %g_s0 : f32 + %g_contribution = scalar.fmaf %g_dot1, %g_s1, %g_c0 : f32 + %g_cor0 = scalar.mulf %g_wm0, %as0 : f32 + %g_cor = scalar.fmaf %g_wm1, %as1, %g_cor0 : f32 + %g_corrected = scalar.subf %g_contribution, %g_cor : f32 + %g_valid = scf.select %valid_block, %g_corrected, %zero : f32 + %g_updated = scalar.addf %acc_g0, %g_valid : f32 + %u_dot0 = func.call @ggml_dot_u8_s8_vector8_f32(%u_lo, %a0) : (vector<2xi32>, vector<2xi32>) -> (f32) + %u_dot1 = func.call @ggml_dot_u8_s8_vector8_f32(%u_hi, %a1) : (vector<2xi32>, vector<2xi32>) -> (f32) + %u_s0 = scalar.mulf %u_ws0, %ad0 : f32 + %u_s1 = scalar.mulf %u_ws1, %ad1 : f32 + %u_c0 = scalar.mulf %u_dot0, %u_s0 : f32 + %u_contribution = scalar.fmaf %u_dot1, %u_s1, %u_c0 : f32 + %u_cor0 = scalar.mulf %u_wm0, %as0 : f32 + %u_cor = scalar.fmaf %u_wm1, %as1, %u_cor0 : f32 + %u_corrected = scalar.subf %u_contribution, %u_cor : f32 + %u_valid = scf.select %valid_block, %u_corrected, %zero : f32 + %u_updated = scalar.addf %acc_u0, %u_valid : f32 + scf.yield %g_updated, %u_updated : f32, f32 + } else { + scf.yield %acc_g0, %acc_u0 : f32, f32 + } + %t1 = index.constant 1 : index + %valid1 = index.cmp ult, %t1, %tokens : index + %next_g1, %next_u1 = scf.if %valid1 -> (f32, f32) { + %tb = index.scale %t1, %q8_row_bytes : index, offset -> offset + %qb = index.add %tb, %q8_block_add : offset + %a0, %a1, %ad0, %ad1, %as0_raw, %as1_raw = func.call @ggml_swiglu_q8_values(%use_local, %input_na, %activation_stage, %qb, %q8_word0, %q8_word1, %q8_meta0, %q8_meta1, %q8_sum0, %q8_sum1) : (i1, buffer, buffer, offset, index, index, index, index, index, index) -> (vector<2xi32>, vector<2xi32>, f32, f32, f32, f32) + %as0 = scf.select %meta_lane, %as0_raw, %zero : f32 + %as1 = scf.select %meta_lane, %as1_raw, %zero : f32 + %g_dot0 = func.call @ggml_dot_u8_s8_vector8_f32(%g_lo, %a0) : (vector<2xi32>, vector<2xi32>) -> (f32) + %g_dot1 = func.call @ggml_dot_u8_s8_vector8_f32(%g_hi, %a1) : (vector<2xi32>, vector<2xi32>) -> (f32) + %g_s0 = scalar.mulf %g_ws0, %ad0 : f32 + %g_s1 = scalar.mulf %g_ws1, %ad1 : f32 + %g_c0 = scalar.mulf %g_dot0, %g_s0 : f32 + %g_contribution = scalar.fmaf %g_dot1, %g_s1, %g_c0 : f32 + %g_cor0 = scalar.mulf %g_wm0, %as0 : f32 + %g_cor = scalar.fmaf %g_wm1, %as1, %g_cor0 : f32 + %g_corrected = scalar.subf %g_contribution, %g_cor : f32 + %g_valid = scf.select %valid_block, %g_corrected, %zero : f32 + %g_updated = scalar.addf %acc_g1, %g_valid : f32 + %u_dot0 = func.call @ggml_dot_u8_s8_vector8_f32(%u_lo, %a0) : (vector<2xi32>, vector<2xi32>) -> (f32) + %u_dot1 = func.call @ggml_dot_u8_s8_vector8_f32(%u_hi, %a1) : (vector<2xi32>, vector<2xi32>) -> (f32) + %u_s0 = scalar.mulf %u_ws0, %ad0 : f32 + %u_s1 = scalar.mulf %u_ws1, %ad1 : f32 + %u_c0 = scalar.mulf %u_dot0, %u_s0 : f32 + %u_contribution = scalar.fmaf %u_dot1, %u_s1, %u_c0 : f32 + %u_cor0 = scalar.mulf %u_wm0, %as0 : f32 + %u_cor = scalar.fmaf %u_wm1, %as1, %u_cor0 : f32 + %u_corrected = scalar.subf %u_contribution, %u_cor : f32 + %u_valid = scf.select %valid_block, %u_corrected, %zero : f32 + %u_updated = scalar.addf %acc_u1, %u_valid : f32 + scf.yield %g_updated, %u_updated : f32, f32 + } else { + scf.yield %acc_g1, %acc_u1 : f32, f32 + } + scf.if %publish_q8 { + scf.schedule.fence + } + %t2 = index.constant 2 : index + %valid2 = index.cmp ult, %t2, %tokens : index + %next_g2, %next_u2 = scf.if %valid2 -> (f32, f32) { + %tb = index.scale %t2, %q8_row_bytes : index, offset -> offset + %qb = index.add %tb, %q8_block_add : offset + %a0, %a1, %ad0, %ad1, %as0_raw, %as1_raw = func.call @ggml_swiglu_q8_values(%use_local, %input_na, %activation_stage, %qb, %q8_word0, %q8_word1, %q8_meta0, %q8_meta1, %q8_sum0, %q8_sum1) : (i1, buffer, buffer, offset, index, index, index, index, index, index) -> (vector<2xi32>, vector<2xi32>, f32, f32, f32, f32) + %as0 = scf.select %meta_lane, %as0_raw, %zero : f32 + %as1 = scf.select %meta_lane, %as1_raw, %zero : f32 + %g_dot0 = func.call @ggml_dot_u8_s8_vector8_f32(%g_lo, %a0) : (vector<2xi32>, vector<2xi32>) -> (f32) + %g_dot1 = func.call @ggml_dot_u8_s8_vector8_f32(%g_hi, %a1) : (vector<2xi32>, vector<2xi32>) -> (f32) + %g_s0 = scalar.mulf %g_ws0, %ad0 : f32 + %g_s1 = scalar.mulf %g_ws1, %ad1 : f32 + %g_c0 = scalar.mulf %g_dot0, %g_s0 : f32 + %g_contribution = scalar.fmaf %g_dot1, %g_s1, %g_c0 : f32 + %g_cor0 = scalar.mulf %g_wm0, %as0 : f32 + %g_cor = scalar.fmaf %g_wm1, %as1, %g_cor0 : f32 + %g_corrected = scalar.subf %g_contribution, %g_cor : f32 + %g_valid = scf.select %valid_block, %g_corrected, %zero : f32 + %g_updated = scalar.addf %acc_g2, %g_valid : f32 + %u_dot0 = func.call @ggml_dot_u8_s8_vector8_f32(%u_lo, %a0) : (vector<2xi32>, vector<2xi32>) -> (f32) + %u_dot1 = func.call @ggml_dot_u8_s8_vector8_f32(%u_hi, %a1) : (vector<2xi32>, vector<2xi32>) -> (f32) + %u_s0 = scalar.mulf %u_ws0, %ad0 : f32 + %u_s1 = scalar.mulf %u_ws1, %ad1 : f32 + %u_c0 = scalar.mulf %u_dot0, %u_s0 : f32 + %u_contribution = scalar.fmaf %u_dot1, %u_s1, %u_c0 : f32 + %u_cor0 = scalar.mulf %u_wm0, %as0 : f32 + %u_cor = scalar.fmaf %u_wm1, %as1, %u_cor0 : f32 + %u_corrected = scalar.subf %u_contribution, %u_cor : f32 + %u_valid = scf.select %valid_block, %u_corrected, %zero : f32 + %u_updated = scalar.addf %acc_u2, %u_valid : f32 + scf.yield %g_updated, %u_updated : f32, f32 + } else { + scf.yield %acc_g2, %acc_u2 : f32, f32 + } + %t3 = index.constant 3 : index + %valid3 = index.cmp ult, %t3, %tokens : index + %next_g3, %next_u3 = scf.if %valid3 -> (f32, f32) { + %tb = index.scale %t3, %q8_row_bytes : index, offset -> offset + %qb = index.add %tb, %q8_block_add : offset + %a0, %a1, %ad0, %ad1, %as0_raw, %as1_raw = func.call @ggml_swiglu_q8_values(%use_local, %input_na, %activation_stage, %qb, %q8_word0, %q8_word1, %q8_meta0, %q8_meta1, %q8_sum0, %q8_sum1) : (i1, buffer, buffer, offset, index, index, index, index, index, index) -> (vector<2xi32>, vector<2xi32>, f32, f32, f32, f32) + %as0 = scf.select %meta_lane, %as0_raw, %zero : f32 + %as1 = scf.select %meta_lane, %as1_raw, %zero : f32 + %g_dot0 = func.call @ggml_dot_u8_s8_vector8_f32(%g_lo, %a0) : (vector<2xi32>, vector<2xi32>) -> (f32) + %g_dot1 = func.call @ggml_dot_u8_s8_vector8_f32(%g_hi, %a1) : (vector<2xi32>, vector<2xi32>) -> (f32) + %g_s0 = scalar.mulf %g_ws0, %ad0 : f32 + %g_s1 = scalar.mulf %g_ws1, %ad1 : f32 + %g_c0 = scalar.mulf %g_dot0, %g_s0 : f32 + %g_contribution = scalar.fmaf %g_dot1, %g_s1, %g_c0 : f32 + %g_cor0 = scalar.mulf %g_wm0, %as0 : f32 + %g_cor = scalar.fmaf %g_wm1, %as1, %g_cor0 : f32 + %g_corrected = scalar.subf %g_contribution, %g_cor : f32 + %g_valid = scf.select %valid_block, %g_corrected, %zero : f32 + %g_updated = scalar.addf %acc_g3, %g_valid : f32 + %u_dot0 = func.call @ggml_dot_u8_s8_vector8_f32(%u_lo, %a0) : (vector<2xi32>, vector<2xi32>) -> (f32) + %u_dot1 = func.call @ggml_dot_u8_s8_vector8_f32(%u_hi, %a1) : (vector<2xi32>, vector<2xi32>) -> (f32) + %u_s0 = scalar.mulf %u_ws0, %ad0 : f32 + %u_s1 = scalar.mulf %u_ws1, %ad1 : f32 + %u_c0 = scalar.mulf %u_dot0, %u_s0 : f32 + %u_contribution = scalar.fmaf %u_dot1, %u_s1, %u_c0 : f32 + %u_cor0 = scalar.mulf %u_wm0, %as0 : f32 + %u_cor = scalar.fmaf %u_wm1, %as1, %u_cor0 : f32 + %u_corrected = scalar.subf %u_contribution, %u_cor : f32 + %u_valid = scf.select %valid_block, %u_corrected, %zero : f32 + %u_updated = scalar.addf %acc_u3, %u_valid : f32 + scf.yield %g_updated, %u_updated : f32, f32 + } else { + scf.yield %acc_g3, %acc_u3 : f32, f32 + } + scf.if %publish_q8 { + scf.schedule.fence + } + %t4 = index.constant 4 : index + %valid4 = index.cmp ult, %t4, %tokens : index + %next_g4, %next_u4 = scf.if %valid4 -> (f32, f32) { + %tb = index.scale %t4, %q8_row_bytes : index, offset -> offset + %qb = index.add %tb, %q8_block_add : offset + %a0, %a1, %ad0, %ad1, %as0_raw, %as1_raw = func.call @ggml_swiglu_q8_values(%use_local, %input_na, %activation_stage, %qb, %q8_word0, %q8_word1, %q8_meta0, %q8_meta1, %q8_sum0, %q8_sum1) : (i1, buffer, buffer, offset, index, index, index, index, index, index) -> (vector<2xi32>, vector<2xi32>, f32, f32, f32, f32) + %as0 = scf.select %meta_lane, %as0_raw, %zero : f32 + %as1 = scf.select %meta_lane, %as1_raw, %zero : f32 + %g_dot0 = func.call @ggml_dot_u8_s8_vector8_f32(%g_lo, %a0) : (vector<2xi32>, vector<2xi32>) -> (f32) + %g_dot1 = func.call @ggml_dot_u8_s8_vector8_f32(%g_hi, %a1) : (vector<2xi32>, vector<2xi32>) -> (f32) + %g_s0 = scalar.mulf %g_ws0, %ad0 : f32 + %g_s1 = scalar.mulf %g_ws1, %ad1 : f32 + %g_c0 = scalar.mulf %g_dot0, %g_s0 : f32 + %g_contribution = scalar.fmaf %g_dot1, %g_s1, %g_c0 : f32 + %g_cor0 = scalar.mulf %g_wm0, %as0 : f32 + %g_cor = scalar.fmaf %g_wm1, %as1, %g_cor0 : f32 + %g_corrected = scalar.subf %g_contribution, %g_cor : f32 + %g_valid = scf.select %valid_block, %g_corrected, %zero : f32 + %g_updated = scalar.addf %acc_g4, %g_valid : f32 + %u_dot0 = func.call @ggml_dot_u8_s8_vector8_f32(%u_lo, %a0) : (vector<2xi32>, vector<2xi32>) -> (f32) + %u_dot1 = func.call @ggml_dot_u8_s8_vector8_f32(%u_hi, %a1) : (vector<2xi32>, vector<2xi32>) -> (f32) + %u_s0 = scalar.mulf %u_ws0, %ad0 : f32 + %u_s1 = scalar.mulf %u_ws1, %ad1 : f32 + %u_c0 = scalar.mulf %u_dot0, %u_s0 : f32 + %u_contribution = scalar.fmaf %u_dot1, %u_s1, %u_c0 : f32 + %u_cor0 = scalar.mulf %u_wm0, %as0 : f32 + %u_cor = scalar.fmaf %u_wm1, %as1, %u_cor0 : f32 + %u_corrected = scalar.subf %u_contribution, %u_cor : f32 + %u_valid = scf.select %valid_block, %u_corrected, %zero : f32 + %u_updated = scalar.addf %acc_u4, %u_valid : f32 + scf.yield %g_updated, %u_updated : f32, f32 + } else { + scf.yield %acc_g4, %acc_u4 : f32, f32 + } + scf.yield %next_g0, %next_g1, %next_g2, %next_g3, %next_g4, %next_u0, %next_u1, %next_u2, %next_u3, %next_u4, %next_g_headers, %next_u_headers, %next_g_raw, %next_u_raw : f32, f32, f32, f32, f32, f32, f32, f32, f32, f32, vector<4xi32>, vector<4xi32>, vector<2xi32>, vector<2xi32> + } + %leader = index.cmp eq, %team_lane, %c0 : index + %out_view = buffer.view %output_na[%base] : buffer -> view<[%tokens]x[%n]xf32> + %p0 = index.constant 0 : index + %pv0 = index.cmp ult, %p0, %tokens : index + scf.if %pv0 { + %g_sum = func.call @ggml_swiglu_q8_reduce(%total_g0, %partitioned) : (f32, i1) -> (f32) + %u_sum = func.call @ggml_swiglu_q8_reduce(%total_u0, %partitioned) : (f32, i1) -> (f32) + scf.if %leader { + %value = func.call @ggml_mul_mat_swiglu_apply_scalar(%op, %g_sum, %u_sum) : (index, f32, f32) -> (f32) + view.store %value, %out_view[%p0, %row] : f32, view<[%tokens]x[%n]xf32> + scf.if %publish_q8 { + view.store %value, %epilogue_view[%p0, %epilogue_row] : f32, view<5x32xf32> + } + } + } + %p1 = index.constant 1 : index + %pv1 = index.cmp ult, %p1, %tokens : index + scf.if %pv1 { + %g_sum = func.call @ggml_swiglu_q8_reduce(%total_g1, %partitioned) : (f32, i1) -> (f32) + %u_sum = func.call @ggml_swiglu_q8_reduce(%total_u1, %partitioned) : (f32, i1) -> (f32) + scf.if %leader { + %value = func.call @ggml_mul_mat_swiglu_apply_scalar(%op, %g_sum, %u_sum) : (index, f32, f32) -> (f32) + view.store %value, %out_view[%p1, %row] : f32, view<[%tokens]x[%n]xf32> + scf.if %publish_q8 { + view.store %value, %epilogue_view[%p1, %epilogue_row] : f32, view<5x32xf32> + } + } + } + %p2 = index.constant 2 : index + %pv2 = index.cmp ult, %p2, %tokens : index + scf.if %pv2 { + %g_sum = func.call @ggml_swiglu_q8_reduce(%total_g2, %partitioned) : (f32, i1) -> (f32) + %u_sum = func.call @ggml_swiglu_q8_reduce(%total_u2, %partitioned) : (f32, i1) -> (f32) + scf.if %leader { + %value = func.call @ggml_mul_mat_swiglu_apply_scalar(%op, %g_sum, %u_sum) : (index, f32, f32) -> (f32) + view.store %value, %out_view[%p2, %row] : f32, view<[%tokens]x[%n]xf32> + scf.if %publish_q8 { + view.store %value, %epilogue_view[%p2, %epilogue_row] : f32, view<5x32xf32> + } + } + } + %p3 = index.constant 3 : index + %pv3 = index.cmp ult, %p3, %tokens : index + scf.if %pv3 { + %g_sum = func.call @ggml_swiglu_q8_reduce(%total_g3, %partitioned) : (f32, i1) -> (f32) + %u_sum = func.call @ggml_swiglu_q8_reduce(%total_u3, %partitioned) : (f32, i1) -> (f32) + scf.if %leader { + %value = func.call @ggml_mul_mat_swiglu_apply_scalar(%op, %g_sum, %u_sum) : (index, f32, f32) -> (f32) + view.store %value, %out_view[%p3, %row] : f32, view<[%tokens]x[%n]xf32> + scf.if %publish_q8 { + view.store %value, %epilogue_view[%p3, %epilogue_row] : f32, view<5x32xf32> + } + } + } + %p4 = index.constant 4 : index + %pv4 = index.cmp ult, %p4, %tokens : index + scf.if %pv4 { + %g_sum = func.call @ggml_swiglu_q8_reduce(%total_g4, %partitioned) : (f32, i1) -> (f32) + %u_sum = func.call @ggml_swiglu_q8_reduce(%total_u4, %partitioned) : (f32, i1) -> (f32) + scf.if %leader { + %value = func.call @ggml_mul_mat_swiglu_apply_scalar(%op, %g_sum, %u_sum) : (index, f32, f32) -> (f32) + view.store %value, %out_view[%p4, %row] : f32, view<[%tokens]x[%n]xf32> + scf.if %publish_q8 { + view.store %value, %epilogue_view[%p4, %epilogue_row] : f32, view<5x32xf32> + } + } + } + } + scf.if %publish_q8 { + kernel.barrier scope(workgroup) ordering(acq_rel) + %epilogue_tid = kernel.workitem.id : index + %epilogue_threads = index.mul %tokens, %c8 : index + %epilogue_active = index.cmp ult, %epilogue_tid, %epilogue_threads : index + scf.if %epilogue_active { + %epilogue_token = index.div %epilogue_tid, %c8 : index + %epilogue_word = index.rem %epilogue_tid, %c8 : index + %epilogue_col = index.mul %epilogue_word, %c4 : index + %epilogue_linear = index.madd %epilogue_token, %c32, %epilogue_col : index + %epilogue_flat = buffer.view %epilogue_stage[%base] : buffer -> view<160xf32> + %epilogue_values = vector.load %epilogue_flat[%epilogue_linear] : view<160xf32> -> vector<4xf32> + %epilogue_channel = index.madd %wg, %c32, %epilogue_col : index + %epilogue_groups = index.div %n, %c128 : index + %epilogue_token_groups = index.mul %epilogue_token, %epilogue_groups : index + %epilogue_token_bytes = index.scale %epilogue_token_groups, %block_bytes : index, offset -> offset + %epilogue_scratch = buffer.view %epilogue_scratch_stage[%base] : buffer -> view<256xf32> + %epilogue_scratch_d = buffer.view %epilogue_scratch_stage[%base] : buffer -> view<32xf32> + %epilogue_publish = scalar.constant true : i1 + template.apply<@ggml.quantize_q8_1_x4.publish_vector4_strict>(%epilogue_publish, %epilogue_token_bytes, %epilogue_channel, %epilogue_values, %epilogue_scratch, %epilogue_scratch_d, %q8_output) : (i1, offset, index, vector<4xf32>, view<256xf32>, view<32xf32>, buffer) + } + } + template.return +} + +func.decl @ggml_mul_mat_quantized_f16_prefill_wave32(%weight_format: index, %binary_op: index, %paired: i1, %packed_input: i1, %packed_output: i1, %token_count: index, %input_size: index, %output_size: index, %channel_tile: index, %token_tile: index, %input: buffer, %weight: buffer, %up_weight: buffer, %output: buffer, %f16_output: buffer) + +// Enough complete M512 tiles and workgroups to amortize the private F16 copy. +func.def pure inline @ggml_swiglu_use_packed_f16(%tokens: index, %outputs: index) -> (i1) { + %zero = index.constant 0 : index + %c64 = index.constant 64 : index + %c512 = index.constant 512 : index + %tail = index.rem %tokens, %c512 : index + %complete = index.cmp eq, %tail, %zero : index + %m_tiles = index.div %tokens, %c512 : index + %n_tiles = index.div %outputs, %c64 : index + %groups = index.mul %m_tiles, %n_tiles : index + %enough_groups = index.cmp uge, %groups, %c64 : index + %packed = scalar.andi %complete, %enough_groups : i1 + func.return %packed : i1 +} + +kernel.def target(@ggml_mul_mat_swiglu_gfx11_wave32) export("ggml_mul_mat_swiglu_q4_k_f16_wmma_prefill_wave32") @ggml_mul_mat_swiglu_q4_k_f16_wmma_prefill_wave32(%token_count: index) { + %output_size = config.get @ggml.mul_mat_swiglu.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %c512 = index.constant 512 : index + %c32 = index.constant 32 : index + %packed_input = func.call pure @ggml_swiglu_use_packed_f16(%token_capacity, %output_size) : (index, index) -> (i1) + %token_tile = scf.select %packed_input, %c512, %c256 : index + %channel_tile = scf.select %packed_input, %c32, %c64 : index + %output_tiles = index.div %output_size, %channel_tile : index + %token_tiles = index.div %token_capacity, %token_tile : index + kernel.launch.config workgroups(%token_tiles, %output_tiles, %c1) workgroup_size(%c512, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %up_weight: buffer, %output: buffer, %f16_output: buffer) { + %input_size = config.get @ggml.mul_mat_swiglu.input_size : index + %output_size = config.get @ggml.mul_mat_swiglu.output_size : index + %token_capacity = config.get @ggml.workload.token_capacity : index + %packed_input = func.call pure @ggml_swiglu_use_packed_f16(%token_capacity, %output_size) : (index, index) -> (i1) + %bounded_token_count = scf.if %packed_input -> (index) { + // The dispatch specializes and launches the same complete input extent. + %count, %capacity = index.assume %token_count, %token_capacity [range(%token_count, 512, 2048), mul(%token_count, 512), eq(%token_count, %token_capacity)] : index, index + scf.yield %count : index + } else { + %count = index.assume %token_count [range(%token_count, 256, 2048), mul(%token_count, 256), le(%token_count, %token_capacity)] : index + scf.yield %count : index + } + %bounded_output_size = index.assume %output_size [range(%output_size, 64, 262144), mul(%output_size, 64)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + %paired = scalar.constant true : i1 + %prefill_weight_format = index.constant 44 : index + %binary_op = config.get @ggml.mul_mat_swiglu.op : index + %f16_output_layout = config.get @ggml.mul_mat_swiglu.f16_output_layout : index + %packed_layout = index.constant 1 : index + %packed_output = index.cmp eq, %f16_output_layout, %packed_layout : index + func.call @ggml_mul_mat_quantized_f16_prefill_wave32(%prefill_weight_format, %binary_op, %paired, %packed_input, %packed_output, %bounded_token_count, %input_size, %bounded_output_size, %channel_tile, %token_tile, %input, %weight, %up_weight, %output, %f16_output) : (index, index, i1, i1, i1, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_swiglu_symmetric_i4_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_swiglu_symmetric_i4_wmma.loom new file mode 100644 index 000000000000..daddfb757b9b --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_swiglu_symmetric_i4_wmma.loom @@ -0,0 +1,1646 @@ +// Generic dense gate/up contraction using backend-resident symmetric-I4/K32 +// weights and activations. The fused kernel applies SwiGLU and publishes a +// Q8_1 plane consumed by a following projection in the same graph match. +// Shape and source-quant selection belong to the graph dispatcher; this module +// contains no model identity or fixed model dimensions. + +template.decl @ggml.quantize_symmetric_i4_k32.body(%plane_major: i1, %src: buffer, %qs: buffer, %ds: buffer, %sums: buffer) + +// Sixteen lanes quantize one K64 block; each packs four values while the two K32 halves share a scale and keep separate sums. +config.decl @ggml.quantize_symmetric_i4_k32.input_size : %value: index where [range(%value, 32, 65536), mul(%value, 32)] + +config.decl @ggml.quantize_symmetric_i4_k32.token_count : %value: index where [range(%value, 1, 16384)] + +kernel.def export("ggml_quantize_f32_symmetric_i4_k32") @ggml_quantize_f32_symmetric_i4_k32() { + %unit = index.constant 1 : index + %wg = index.constant 256 : index + %wg_groups_m1 = index.constant 127 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %k = config.get @ggml.quantize_symmetric_i4_k32.input_size : index + %cols = config.get @ggml.quantize_symmetric_i4_k32.token_count : index + %total = index.mul %k, %cols : index + %groups = index.div %total, %c64 : index + %rounded = index.add %groups, %wg_groups_m1 : index + %workgroups = index.div %rounded, %c128 : index + kernel.launch.config workgroups(%workgroups, %unit, %unit) workgroup_size(%wg, %unit, %unit) : index +} launch(%src: buffer, %qs: buffer, %ds: buffer, %sums: buffer) { + %plane_major = scalar.constant false : i1 + template.apply<@ggml.quantize_symmetric_i4_k32.body>(%plane_major, %src, %qs, %ds, %sums) : (i1, buffer, buffer, buffer, buffer) + kernel.return +} + +template.def<@ggml.quantize_symmetric_i4_k32.body> device @ggml_quantize_f32_symmetric_i4_k32_body(%plane_major: i1, %src: buffer, %qs: buffer, %ds: buffer, %sums: buffer) { + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %c0_f32 = scalar.constant 0.0 : f32 + %amax_eps = scalar.constant 1.0000000000000001e-30 : f32 + %inv7 = scalar.constant 0.14285714285714285 : f32 + %f7 = scalar.constant 7.0 : f32 + %xor1 = scalar.constant 1 : i32 + %xor2 = scalar.constant 2 : i32 + %xor4 = scalar.constant 4 : i32 + %xor8 = scalar.constant 8 : i32 + %shift4 = scalar.constant 4 : i32 + %shift8 = scalar.constant 8 : i32 + %shift12 = scalar.constant 12 : i32 + %shuffle_width = scalar.constant 32 : i32 + %mask15 = vector.constant 15 : vector<4xi32> + + %k = config.get @ggml.quantize_symmetric_i4_k32.input_size : index + %cols = config.get @ggml.quantize_symmetric_i4_k32.token_count : index + %k_b = index.assume %k [range(%k, 32, 65536)] : index + %cols_b = index.assume %cols [range(%cols, 1, 16384)] : index + %total = index.mul %k_b, %cols_b : index + %groups0 = index.div %total, %c32 : index + %groups = index.assume %groups0 [range(%groups0, 1, 33554432)] : index + %groups64_0 = index.div %total, %c64 : index + %groups64 = index.assume %groups64_0 [range(%groups64_0, 1, 16777216)] : index + %nchunks = index.div %k_b, %c64 : index + + %wgid0 = kernel.workgroup.id : index + %tid0 = kernel.workitem.id : index + %wgid = index.assume %wgid0 [range(%wgid0, 0, 131071)] : index + %tid = index.assume %tid0 [range(%tid0, 0, 255)] : index + %wave = index.div %tid, %c32 : index + %lane = index.rem %tid, %c32 : index + %group_in_wave = index.div %lane, %c16 : index + %lane_in_group = index.rem %lane, %c16 : index + %group_half = index.div %lane_in_group, %c8 : index + %lane_in_half = index.rem %lane_in_group, %c8 : index + %is_scale_leader = index.cmp eq, %lane_in_group, %c0 : index + %is_sum_leader = index.cmp eq, %lane_in_half, %c0 : index + %wg_group_base = index.mul %wgid, %c128 : index + %wave_group_base = index.mul %wave, %c2 : index + %local_group_base = index.add %wg_group_base, %wave_group_base : index + + %src_g = buffer.assume.memory_space %src : buffer + %qs_g = buffer.assume.memory_space %qs : buffer + %ds_g = buffer.assume.memory_space %ds : buffer + %sums_g = buffer.assume.memory_space %sums : buffer + %src_na, %qs_na, %ds_na, %sums_na = buffer.assume.noalias %src_g, %qs_g, %ds_g, %sums_g : buffer, buffer, buffer, buffer + %src_elems = index.assume %total [range(%total, 32, 1073741824)] : index + %src_view = buffer.view %src_na[%base] : buffer -> view<[%src_elems]xf32> + %qs_halfwords = index.div %total, %c4 : index + %qs_view = buffer.view %qs_na[%base] : buffer -> view<[%qs_halfwords]xi16> + %ds_view = buffer.view %ds_na[%base] : buffer -> view<[%groups]xf32> + %sum_view = buffer.view %sums_na[%base] : buffer -> view<[%groups]xi32> + + scf.for %batch = [%c0 to %c8 step %c1] { + %batch_group_off = index.mul %batch, %c16 : index + %gid0 = index.add %local_group_base, %batch_group_off : index + %gid1 = index.add %gid0, %group_in_wave : index + %gid = index.assume %gid1 [range(%gid1, 0, 16777215)] : index + %in_range = index.cmp ult, %gid, %groups64 : index + scf.if %in_range { + %ebase0 = index.mul %gid, %c64 : index + %lane_elem_off = index.mul %lane_in_group, %c4 : index + %ebase1 = index.add %ebase0, %lane_elem_off : index + %ebase = index.assume %ebase1 [range(%ebase1, 0, 1073741820)] : index + %v = vector.load %src_view[%ebase] : view<[%src_elems]xf32> -> vector<4xf32> + %av = vector.absf %v : vector<4xf32> + %lane_max = vector.reduce %av, %c0_f32 : vector<4xf32>, f32 + %max_x1_peer, %max_x1_valid = kernel.subgroup.shuffle %lane_max, %xor1, %shuffle_width : f32, i32, i32 + %max_x1 = scalar.maxnumf %lane_max, %max_x1_peer : f32 + %max_x2_peer, %max_x2_valid = kernel.subgroup.shuffle %max_x1, %xor2, %shuffle_width : f32, i32, i32 + %max_x2 = scalar.maxnumf %max_x1, %max_x2_peer : f32 + %max_x4_peer, %max_x4_valid = kernel.subgroup.shuffle %max_x2, %xor4, %shuffle_width : f32, i32, i32 + %max_x4 = scalar.maxnumf %max_x2, %max_x4_peer : f32 + %max_x8_peer, %max_x8_valid = kernel.subgroup.shuffle %max_x4, %xor8, %shuffle_width : f32, i32, i32 + %group_max = scalar.maxnumf %max_x4, %max_x8_peer : f32 + %amax = scalar.maxnumf %group_max, %amax_eps : f32 + %a_scale = scalar.mulf %amax, %inv7 : f32 + %a_rscale = scalar.divf %f7, %amax : f32 + %gid32_base = index.mul %gid, %c2 : index + %gid32_high = index.add %gid32_base, %c1 : index + %token = index.div %gid, %nchunks : index + %chunk = index.rem %gid, %nchunks : index + %plane_chunk32 = index.mul %chunk, %c2 : index + %plane_low_base = index.mul %plane_chunk32, %cols_b : index + %plane_gid32_low = index.add %plane_low_base, %token : index + %plane_chunk32_high = index.add %plane_chunk32, %c1 : index + %plane_high_base = index.mul %plane_chunk32_high, %cols_b : index + %plane_gid32_high = index.add %plane_high_base, %token : index + %scale_low_index = scf.if %plane_major -> (index) { + scf.yield %plane_gid32_low : index + } else { + scf.yield %gid32_base : index + } + %scale_high_index = scf.if %plane_major -> (index) { + scf.yield %plane_gid32_high : index + } else { + scf.yield %gid32_high : index + } + scf.if %is_scale_leader { + view.store %a_scale, %ds_view[%scale_low_index] : f32, view<[%groups]xf32> + view.store %a_scale, %ds_view[%scale_high_index] : f32, view<[%groups]xf32> + } + + %rs = vector.splat %a_rscale : vector<4xf32> + %scaled = vector.mulf %v, %rs : vector<4xf32> + %rounded = vector.roundf %scaled : vector<4xf32> + %q32 = vector.fptosi %rounded : vector<4xf32> to vector<4xi32> + %qn = vector.andi %q32, %mask15 : vector<4xi32> + %q0 = vector.extract %qn[0] : vector<4xi32> -> i32 + %q1 = vector.extract %qn[1] : vector<4xi32> -> i32 + %q2 = vector.extract %qn[2] : vector<4xi32> -> i32 + %q3 = vector.extract %qn[3] : vector<4xi32> -> i32 + %q1s = scalar.shli %q1, %shift4 : i32 + %q2s = scalar.shli %q2, %shift8 : i32 + %q3s = scalar.shli %q3, %shift12 : i32 + %q01 = scalar.ori %q0, %q1s : i32 + %q23 = scalar.ori %q2s, %q3s : i32 + %qpacked32 = scalar.ori %q01, %q23 : i32 + %qpacked = scalar.trunci %qpacked32 : i32 to i16 + %plane_chunk_base = index.mul %chunk, %cols_b : index + %plane_gid = index.add %plane_chunk_base, %token : index + %payload_gid = scf.if %plane_major -> (index) { + scf.yield %plane_gid : index + } else { + scf.yield %gid : index + } + %word_base = index.mul %payload_gid, %c16 : index + %word_index0 = index.add %word_base, %lane_in_group : index + %word_index = index.assume %word_index0 [range(%word_index0, 0, 268435455)] : index + view.store %qpacked, %qs_view[%word_index] : i16, view<[%qs_halfwords]xi16> + + %lane_sum = vector.reduce %rounded, %c0_f32 : vector<4xf32>, f32 + %sum_x1_peer, %sum_x1_valid = kernel.subgroup.shuffle %lane_sum, %xor1, %shuffle_width : f32, i32, i32 + %sum_x1 = scalar.addf %lane_sum, %sum_x1_peer : f32 + %sum_x2_peer, %sum_x2_valid = kernel.subgroup.shuffle %sum_x1, %xor2, %shuffle_width : f32, i32, i32 + %sum_x2 = scalar.addf %sum_x1, %sum_x2_peer : f32 + %sum_x4_peer, %sum_x4_valid = kernel.subgroup.shuffle %sum_x2, %xor4, %shuffle_width : f32, i32, i32 + %group_sum = scalar.addf %sum_x2, %sum_x4_peer : f32 + scf.if %is_sum_leader { + %sum_value = scalar.fptosi %group_sum : f32 to i32 + %row_sum_index = index.add %gid32_base, %group_half : index + %plane_sum_chunk = index.add %plane_chunk32, %group_half : index + %plane_sum_base = index.mul %plane_sum_chunk, %cols_b : index + %plane_sum_index = index.add %plane_sum_base, %token : index + %sum_index = scf.if %plane_major -> (index) { + scf.yield %plane_sum_index : index + } else { + scf.yield %row_sum_index : index + } + view.store %sum_value, %sum_view[%sum_index] : i32, view<[%groups]xi32> + } + } + } + template.return +} + +// --- gate/up contraction --- +template.decl @ggml.mul_mat_swiglu.symmetric_i4.m128n64_wg128.body(%apply_swiglu: i1, %publish_f32: i1, %publish_f16: i1, %publish_u4: i1, %publish_q8_plane: i1, %gate_na: buffer, %src0_na: buffer, %dst_na: buffer, %aq_na: buffer, %as_na: buffer, %asum_na: buffer, %qout_qs_na: buffer, %qout_ds_na: buffer, %qout_sums_na: buffer, %qout_scratch: buffer, %base: offset, %qact_scale_base: offset, %c0: index, %c1: index, %c2: index, %c4: index, %c8: index, %c16: index, %c32: index, %c17: index, %c34: index, %k_kc: index, %k_astride: index, %k_wstride: index, %k_m: index, %k_n: index, %k_blocks: index, %k_wave: index, %k_wm: index, %fzero: vector<8xf32>, %izero: vector<8xi32>, %k_b: index, %rows_b: index, %cols_b: index, %nchunks: index, %kblocks: index, %col_tile: index, %row_tile: index, %tid: index, %wave: index) + +kernel.def export("ggml_mul_mat_swiglu_symmetric_i4_wmma_q8_plane") @ggml_mul_mat_swiglu_symmetric_i4_wmma_q8_plane() { + %unit = index.constant 1 : index + %k_n = index.constant 64 : index + %k_m = index.constant 128 : index + %n_m1 = index.constant 63 : index + %m_m1 = index.constant 127 : index + %wg = index.constant 128 : index + %rows = config.get @ggml.mul_mat_swiglu.symmetric_i4.output_size : index + %cols = config.get @ggml.mul_mat_swiglu.symmetric_i4.token_count : index + %rows_up = index.add %rows, %n_m1 : index + %row_tiles = index.div %rows_up, %k_n : index + %cols_up = index.add %cols, %m_m1 : index + %col_tiles = index.div %cols_up, %k_m : index + kernel.launch.config workgroups(%col_tiles, %row_tiles, %unit) workgroup_size(%wg, %unit, %unit) : index +} launch(%gate_weight: buffer, %up_weight: buffer, %input: buffer, %dst: buffer, %aq: buffer, %as: buffer, %asum: buffer, %q8_output: buffer) { + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c17 = index.constant 17 : index + %c34 = index.constant 34 : index + %k_kc = index.constant 64 : index + %k_astride = index.constant 40 : index + %k_wstride = index.constant 40 : index + %k_m = index.constant 128 : index + %k_n = index.constant 64 : index + %k_blocks = index.constant 2 : index + %k_wave = index.constant 32 : index + %k_wm = index.constant 4 : index + %fzero = vector.constant 0.0 : vector<8xf32> + %izero = vector.constant 0 : vector<8xi32> + %apply_swiglu = scalar.constant true : i1 + %publish_f32 = scalar.constant false : i1 + %publish_f16 = scalar.constant false : i1 + %publish_u4 = scalar.constant false : i1 + %publish_q8_plane = scalar.constant true : i1 + %scratch_bytes = index.constant 16384 : offset + %scratch = buffer.alloca align(16) %scratch_bytes : buffer + + %k = config.get @ggml.mul_mat_swiglu.symmetric_i4.input_size : index + %rows = config.get @ggml.mul_mat_swiglu.symmetric_i4.output_size : index + %cols = config.get @ggml.mul_mat_swiglu.symmetric_i4.token_count : index + %k_b = index.assume %k [range(%k, 64, 32768)] : index + %rows_b = index.assume %rows [range(%rows, 64, 262144), mul(%rows, 64)] : index + %cols_b = index.assume %cols [range(%cols, 128, 32768), mul(%cols, 128)] : index + %nchunks = index.div %k_b, %k_kc : index + %kblocks = index.div %k_b, %c32 : index + %col_tile0 = kernel.workgroup.id : index + %row_tile0 = kernel.workgroup.id : index + %tid0 = kernel.workitem.id : index + %col_tile = index.assume %col_tile0 [range(%col_tile0, 0, 255)] : index + %row_tile = index.assume %row_tile0 [range(%row_tile0, 0, 8191)] : index + %tid = index.assume %tid0 [range(%tid0, 0, 127)] : index + %wave = index.div %tid, %k_wave : index + + %gate_g = buffer.assume.memory_space %gate_weight : buffer + %up_g = buffer.assume.memory_space %up_weight : buffer + %dst_g = buffer.assume.memory_space %dst : buffer + %aq_g = buffer.assume.memory_space %aq : buffer + %as_g = buffer.assume.memory_space %as : buffer + %asum_g = buffer.assume.memory_space %asum : buffer + %q8_output_g = buffer.assume.memory_space %q8_output : buffer + %gate_na, %up_na, %dst_na, %aq_na, %as_na, %asum_na, %q8_output_na = buffer.assume.noalias %gate_g, %up_g, %dst_g, %aq_g, %as_g, %asum_g, %q8_output_g : buffer, buffer, buffer, buffer, buffer, buffer, buffer + + template.apply<@ggml.mul_mat_swiglu.symmetric_i4.m128n64_wg128.body>(%apply_swiglu, %publish_f32, %publish_f16, %publish_u4, %publish_q8_plane, %up_na, %gate_na, %dst_na, %aq_na, %as_na, %asum_na, %q8_output_na, %dst_na, %dst_na, %scratch, %base, %base, %c0, %c1, %c2, %c4, %c8, %c16, %c32, %c17, %c34, %k_kc, %k_astride, %k_wstride, %k_m, %k_n, %k_blocks, %k_wave, %k_wm, %fzero, %izero, %k_b, %rows_b, %cols_b, %nchunks, %kblocks, %col_tile, %row_tile, %tid, %wave) : (i1, i1, i1, i1, i1, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, offset, offset, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, vector<8xf32>, vector<8xi32>, index, index, index, index, index, index, index, index, index) + kernel.return +} + +config.decl @ggml.mul_mat_swiglu.symmetric_i4.input_size : %value: index where [range(%value, 32, 32768), mul(%value, 32)] + +config.decl @ggml.mul_mat_swiglu.symmetric_i4.output_size : %value: index where [range(%value, 1, 262144)] + +config.decl @ggml.mul_mat_swiglu.symmetric_i4.token_count : %value: index where [range(%value, 1, 32768)] + +template.def<@ggml.mul_mat_swiglu.symmetric_i4.m128n64_wg128.body> device @ggml_mul_mat_swiglu_symmetric_i4_m128n64_wg128_body(%apply_swiglu: i1, %publish_f32: i1, %publish_f16: i1, %publish_u4: i1, %publish_q8_plane: i1, %gate_na: buffer, %src0_na: buffer, %dst_na: buffer, %aq_na: buffer, %as_na: buffer, %asum_na: buffer, %qout_qs_na: buffer, %qout_ds_na: buffer, %qout_sums_na: buffer, %qout_scratch: buffer, %base: offset, %qact_scale_base: offset, %c0: index, %c1: index, %c2: index, %c4: index, %c8: index, %c16: index, %c32: index, %c17: index, %c34: index, %k_kc: index, %k_astride: index, %k_wstride: index, %k_m: index, %k_n: index, %k_blocks: index, %k_wave: index, %k_wm: index, %fzero: vector<8xf32>, %izero: vector<8xi32>, %k_b: index, %rows_b: index, %cols_b: index, %nchunks: index, %kblocks: index, %col_tile: index, %row_tile: index, %tid: index, %wave: index) { + %aq_flat = buffer.view %aq_na[%base] : buffer -> view<1073741824xi8> + %as_flat = buffer.view %as_na[%qact_scale_base] : buffer -> view<33554432xf32> + %q4_k_block = index.constant 256 : index + // dst is [rows, cols] with rows contiguous, and C is [cols, rows]. + %dst_view = buffer.view %dst_na[%base] : buffer -> view<[%cols_b]x[%rows_b]xf32> + %dst_f16_view = buffer.view %dst_na[%base] : buffer -> view<[%cols_b]x[%rows_b]xf16> + + // Packed I4 operands use 32 payload bytes per logical K64 row. + // A 40-byte LDS stride retains the conflict-avoiding padding of the IU8 body. + %scratch = buffer.assume.memory_space %qout_scratch : buffer + %i4_schema = encoding.define #encoding.operand : encoding + %al_flat = buffer.view %scratch[%base] : buffer -> view<5120xi8> + %w_off = index.constant 5120 : offset + %wl_flat = buffer.view %scratch[%w_off] : buffer -> view<5120xi8> + %as_off = index.constant 10240 : offset + %asl_view = buffer.view %scratch[%as_off] : buffer -> view<256xf32> + %ws_off = index.constant 11264 : offset + %wsl_view = buffer.view %scratch[%ws_off] : buffer -> view<256xf32> + %col_base = index.mul %col_tile, %k_m : index + %row_base = index.mul %row_tile, %k_n : index + %k_asn = index.constant 256 : index + %k_lanes0 = index.constant 0 : index + %k_lanes1 = index.constant 128 : index + + // Four waves own M32xN64 from gate and M32xN64 from up. + %wcol = index.rem %wave, %k_wm : index + %wrow = index.div %wave, %k_wm : index + %k_mspan = index.constant 16 : index + %k_nspan = index.constant 64 : index + %wm_off = index.mul %wcol, %k_mspan : index + %wn_off = index.mul %wrow, %k_nspan : index + %m_out = index.add %col_base, %wm_off : index + %n_out = index.add %row_base, %wn_off : index + // RDNA3 WMMA RHS lanes map to column lane % 16, so each lane starts at its output row. + %lane = index.rem %tid, %k_wave : index + %lane_lo = index.rem %lane, %c16 : index + %lane_hi = index.div %lane, %c16 : index + %wn_lane = index.add %wn_off, %lane_lo : index + %lm0 = index.add %wm_off, %c0 : index + %gm0 = index.add %m_out, %c0 : index + %k_ma1 = index.constant 16 : index + %lm1 = index.add %wm_off, %c0 : index + %gm1 = index.add %m_out, %c0 : index + %k_ma2 = index.constant 64 : index + %lm2 = index.add %wm_off, %k_ma2 : index + %gm2 = index.add %m_out, %k_ma2 : index + %lm3 = index.add %wm_off, %k_ma2 : index + %gm3 = index.add %m_out, %k_ma2 : index + %ln0 = index.add %wn_lane, %c0 : index + %gn0 = index.add %n_out, %c0 : index + %k_nb1 = index.constant 16 : index + %ln1 = index.add %wn_lane, %k_nb1 : index + %gn1 = index.add %n_out, %k_nb1 : index + %k_nb2 = index.constant 32 : index + %ln2 = index.add %wn_lane, %k_nb2 : index + %gn2 = index.add %n_out, %k_nb2 : index + %k_nb3 = index.constant 48 : index + %ln3 = index.add %wn_lane, %k_nb3 : index + %gn3 = index.add %n_out, %k_nb3 : index + %k_nb4 = index.constant 64 : index + %ln4 = index.add %wn_lane, %k_nb4 : index + %k_nb5 = index.constant 80 : index + %ln5 = index.add %wn_lane, %k_nb5 : index + %k_nb6 = index.constant 96 : index + %ln6 = index.add %wn_lane, %k_nb6 : index + %k_nb7 = index.constant 112 : index + %ln7 = index.add %wn_lane, %k_nb7 : index + %wsb0_0 = index.add %ln0, %c0 : index + %wsb0 = index.mul %wsb0_0, %k_blocks : index + %wsb1_0 = index.add %ln1, %c0 : index + %wsb1 = index.mul %wsb1_0, %k_blocks : index + %wsb2_0 = index.add %ln2, %c0 : index + %wsb2 = index.mul %wsb2_0, %k_blocks : index + %wsb3_0 = index.add %ln3, %c0 : index + %wsb3 = index.mul %wsb3_0, %k_blocks : index + %wsb4_0 = index.add %ln4, %c0 : index + %wsb4 = index.mul %wsb4_0, %k_blocks : index + %wsb5_0 = index.add %ln5, %c0 : index + %wsb5 = index.mul %wsb5_0, %k_blocks : index + %wsb6_0 = index.add %ln6, %c0 : index + %wsb6 = index.mul %wsb6_0, %k_blocks : index + %wsb7_0 = index.add %ln7, %c0 : index + %wsb7 = index.mul %wsb7_0, %k_blocks : index + // Activation scale metadata is [tile][block][lane-half][register]. + %as_tile0 = index.div %lm0, %c16 : index + %asb0 = index.mul %as_tile0, %k_blocks : index + %as_tile1 = index.div %lm1, %c16 : index + %asb1 = index.mul %as_tile1, %k_blocks : index + %as_tile2 = index.div %lm2, %c16 : index + %asb2 = index.mul %as_tile2, %k_blocks : index + %as_tile3 = index.div %lm3, %c16 : index + %asb3 = index.mul %as_tile3, %k_blocks : index + + %f0_0, %f0_1, %f0_2, %f0_3, %f1_0, %f1_1, %f1_2, %f1_3, %f2_0, %f2_1, %f2_2, %f2_3, %f3_0, %f3_1, %f3_2, %f3_3 = scf.for %chunk = [%c0 to %nchunks step %c1](%fc0_0 = %fzero : vector<8xf32>, %fc0_1 = %fzero : vector<8xf32>, %fc0_2 = %fzero : vector<8xf32>, %fc0_3 = %fzero : vector<8xf32>, %fc1_0 = %fzero : vector<8xf32>, %fc1_1 = %fzero : vector<8xf32>, %fc1_2 = %fzero : vector<8xf32>, %fc1_3 = %fzero : vector<8xf32>, %fc2_0 = %fzero : vector<8xf32>, %fc2_1 = %fzero : vector<8xf32>, %fc2_2 = %fzero : vector<8xf32>, %fc2_3 = %fzero : vector<8xf32>, %fc3_0 = %fzero : vector<8xf32>, %fc3_1 = %fzero : vector<8xf32>, %fc3_2 = %fzero : vector<8xf32>, %fc3_3 = %fzero : vector<8xf32>) -> (vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) { + %chunk_k = index.mul %chunk, %k_kc : index + %chunk_blk = index.mul %chunk, %k_blocks : index + + // All lanes stage the unique M128 activation rows. + %q4a_row0 = index.add %col_base, %tid : index + %q4a_row = index.assume %q4a_row0 [range(%q4a_row0, 0, 32767)] : index + %q4a_row_bytes = index.div %k_b, %c2 : index + %q4a_row_base = index.mul %q4a_row, %q4a_row_bytes : index + %q4a_chunk_bytes = index.div %chunk_k, %c2 : index + %q4a_src0_0 = index.add %q4a_row_base, %q4a_chunk_bytes : index + %q4a_src0 = index.assume %q4a_src0_0 [range(%q4a_src0_0, 0, 536870896)] : index + %q4a_v0 = vector.load %aq_flat[%q4a_src0] : view<1073741824xi8> -> vector<16xi8> + %q4a_src1_0 = index.add %q4a_src0, %c16 : index + %q4a_src1 = index.assume %q4a_src1_0 [range(%q4a_src1_0, 16, 536870912)] : index + %q4a_v1 = vector.load %aq_flat[%q4a_src1] : view<1073741824xi8> -> vector<16xi8> + %q4a_dst0_0 = index.mul %tid, %k_astride : index + %q4a_dst0 = index.assume %q4a_dst0_0 [range(%q4a_dst0_0, 0, 5080)] : index + vector.store %q4a_v0, %al_flat[%q4a_dst0] : vector<16xi8>, view<5120xi8> + %q4a_dst1_0 = index.add %q4a_dst0, %c16 : index + %q4a_dst1 = index.assume %q4a_dst1_0 [range(%q4a_dst1_0, 16, 5096)] : index + vector.store %q4a_v1, %al_flat[%q4a_dst1] : vector<16xi8>, view<5120xi8> + + // The 144-byte block holds eight duplicated K64 scale slots followed by eight signed-I4 K32 payloads. + %q4_weight_lane_count = index.constant 128 : index + %q4_weight_lane_active = index.cmp ult, %tid, %q4_weight_lane_count : index + scf.if %q4_weight_lane_active { + %q4_projection_rows = index.constant 64 : index + %q4_weight_tid0 = index.rem %tid, %q4_projection_rows : index + %q4_weight_tid = index.assume %q4_weight_tid0 [range(%q4_weight_tid0, 0, 63)] : index + %q4_lds_row = index.assume %tid [range(%tid, 0, 127)] : index + %q4_is_up = index.cmp uge, %tid, %q4_projection_rows : index + %q4_block_count = index.div %k_b, %q4_k_block : index + %q4_block = index.div %chunk, %c4 : index + %q4_pair = index.rem %chunk, %c4 : index + %q4_row0 = index.add %row_base, %q4_weight_tid : index + %q4_row = index.assume %q4_row0 [range(%q4_row0, 0, 262143)] : index + %q4_row_group_size = index.constant 64 : index + %q4_field_count = index.constant 9 : index + %q4_field_bytes = index.constant 16 : index + %q4_row_group = index.div %q4_row, %q4_row_group_size : index + %q4_row_lane = index.rem %q4_row, %q4_row_group_size : index + %q4_group_block0 = index.mul %q4_row_group, %q4_block_count : index + %q4_group_block = index.add %q4_group_block0, %q4_block : index + %q4_group_field_base = index.mul %q4_group_block, %q4_field_count : index + %q4_header_lane_base = index.mul %q4_group_field_base, %q4_row_group_size : index + %q4_header_lane = index.add %q4_header_lane_base, %q4_row_lane : index + %q4_header_byte_index = index.mul %q4_header_lane, %q4_field_bytes : index + %q4_scale_byte_stride = index.constant 4 : index + %q4_scale_byte_offset = index.mul %q4_pair, %q4_scale_byte_stride : index + %q4_scale_byte_index = index.add %q4_header_byte_index, %q4_scale_byte_offset : index + %q4_scale_byte_base = index.cast %q4_scale_byte_index : index to offset + %q4_scale_pair = scf.if %q4_is_up -> (vector<2xf16>) { + %q4_up_scale_view = buffer.view %gate_na[%q4_scale_byte_base] : buffer -> view<2xf16> + %q4_up_scale_pair = vector.load %q4_up_scale_view[%c0] : view<2xf16> -> vector<2xf16> + scf.yield %q4_up_scale_pair : vector<2xf16> + } else { + %q4_gate_scale_view = buffer.view %src0_na[%q4_scale_byte_base] : buffer -> view<2xf16> + %q4_gate_scale_pair = vector.load %q4_gate_scale_view[%c0] : view<2xf16> -> vector<2xf16> + scf.yield %q4_gate_scale_pair : vector<2xf16> + } + %q4_scale_low_f16 = vector.extract %q4_scale_pair[0] : vector<2xf16> -> f16 + %q4_scale_high_f16 = vector.extract %q4_scale_pair[1] : vector<2xf16> -> f16 + %q4_d0 = scalar.extf %q4_scale_low_f16 : f16 to f32 + %q4_d1 = scalar.extf %q4_scale_high_f16 : f16 to f32 + %q4_low_group = index.mul %q4_pair, %c2 : index + %q4_high_group = index.add %q4_low_group, %c1 : index + %q4_low_field0 = index.add %q4_low_group, %c1 : index + %q4_high_field0 = index.add %q4_high_group, %c1 : index + %q4_low_field = index.add %q4_group_field_base, %q4_low_field0 : index + %q4_high_field = index.add %q4_group_field_base, %q4_high_field0 : index + %q4_low_lane_base = index.mul %q4_low_field, %q4_row_group_size : index + %q4_high_lane_base = index.mul %q4_high_field, %q4_row_group_size : index + %q4_low_lane = index.add %q4_low_lane_base, %q4_row_lane : index + %q4_high_lane = index.add %q4_high_lane_base, %q4_row_lane : index + %q4_low_byte_index = index.mul %q4_low_lane, %q4_field_bytes : index + %q4_high_byte_index = index.mul %q4_high_lane, %q4_field_bytes : index + %q4_low_byte_base = index.cast %q4_low_byte_index : index to offset + %q4_high_byte_base = index.cast %q4_high_byte_index : index to offset + %q4_low_words, %q4_high_words = scf.if %q4_is_up -> (vector<4xi32>, vector<4xi32>) { + %q4_up_low_view = buffer.view %gate_na[%q4_low_byte_base] : buffer -> view<4xi32> + %q4_up_high_view = buffer.view %gate_na[%q4_high_byte_base] : buffer -> view<4xi32> + %q4_up_low_words = vector.load %q4_up_low_view[%c0] : view<4xi32> -> vector<4xi32> + %q4_up_high_words = vector.load %q4_up_high_view[%c0] : view<4xi32> -> vector<4xi32> + scf.yield %q4_up_low_words, %q4_up_high_words : vector<4xi32>, vector<4xi32> + } else { + %q4_gate_low_view = buffer.view %src0_na[%q4_low_byte_base] : buffer -> view<4xi32> + %q4_gate_high_view = buffer.view %src0_na[%q4_high_byte_base] : buffer -> view<4xi32> + %q4_gate_low_words = vector.load %q4_gate_low_view[%c0] : view<4xi32> -> vector<4xi32> + %q4_gate_high_words = vector.load %q4_gate_high_view[%c0] : view<4xi32> -> vector<4xi32> + scf.yield %q4_gate_low_words, %q4_gate_high_words : vector<4xi32>, vector<4xi32> + } + %q4_low_packed = vector.bitcast %q4_low_words : vector<4xi32> to vector<16xi8> + %q4_high_packed = vector.bitcast %q4_high_words : vector<4xi32> to vector<16xi8> + %q4_weight_row_base = index.mul %q4_lds_row, %k_wstride : index + %q4_w0 = index.assume %q4_weight_row_base [range(%q4_weight_row_base, 0, 5080)] : index + %q4_w1_0 = index.add %q4_weight_row_base, %c16 : index + %q4_w1 = index.assume %q4_w1_0 [range(%q4_w1_0, 16, 5096)] : index + vector.store %q4_low_packed, %wl_flat[%q4_w0] : vector<16xi8>, view<5120xi8> + vector.store %q4_high_packed, %wl_flat[%q4_w1] : vector<16xi8>, view<5120xi8> + %q4_meta0_0 = index.mul %q4_lds_row, %c2 : index + %q4_meta0 = index.assume %q4_meta0_0 [range(%q4_meta0_0, 0, 254)] : index + %q4_meta1_0 = index.add %q4_meta0, %c1 : index + %q4_meta1 = index.assume %q4_meta1_0 [range(%q4_meta1_0, 1, 255)] : index + view.store %q4_d0, %wsl_view[%q4_meta0] : f32, view<256xf32> + view.store %q4_d1, %wsl_view[%q4_meta1] : f32, view<256xf32> + } + // Scale planes: 256 activation and 128 weight entries. + %sca0_slot0 = index.add %tid, %k_lanes0 : index + %sca0_slot = index.assume %sca0_slot0 [range(%sca0_slot0, 0, 127)] : index + %scin_a0 = index.cmp ult, %sca0_slot, %k_asn : index + scf.if %scin_a0 { + %sc0_col = index.div %sca0_slot, %k_blocks : index + %sc0_blk = index.rem %sca0_slot, %k_blocks : index + %sca0_c = index.add %col_base, %sc0_col : index + %sca0_0 = index.mul %sca0_c, %kblocks : index + %sca0_1 = index.add %sca0_0, %chunk_blk : index + %sca0_2 = index.add %sca0_1, %sc0_blk : index + %sca0 = index.assume %sca0_2 [range(%sca0_2, 0, 33554431)] : index + %scav0 = view.load %as_flat[%sca0] : view<33554432xf32> -> f32 + %sca0_tile = index.div %sc0_col, %c16 : index + %sca0_inner = index.rem %sc0_col, %c16 : index + %sca0_par = index.rem %sca0_inner, %c2 : index + %sca0_v = index.div %sca0_inner, %c2 : index + %sca0_dst0 = index.mul %sca0_tile, %k_blocks : index + %sca0_dst1 = index.add %sca0_dst0, %sc0_blk : index + %sca0_dst2 = index.mul %sca0_dst1, %c2 : index + %sca0_dst3 = index.add %sca0_dst2, %sca0_par : index + %sca0_dst4 = index.mul %sca0_dst3, %c8 : index + %sca0_dst5 = index.add %sca0_dst4, %sca0_v : index + %sca0_dst = index.assume %sca0_dst5 [range(%sca0_dst5, 0, 255)] : index + view.store %scav0, %asl_view[%sca0_dst] : f32, view<256xf32> + } + %sca1_slot0 = index.add %tid, %k_lanes1 : index + %sca1_slot = index.assume %sca1_slot0 [range(%sca1_slot0, 128, 255)] : index + %scin_a1 = index.cmp ult, %sca1_slot, %k_asn : index + scf.if %scin_a1 { + %sc1_col = index.div %sca1_slot, %k_blocks : index + %sc1_blk = index.rem %sca1_slot, %k_blocks : index + %sca1_c = index.add %col_base, %sc1_col : index + %sca1_0 = index.mul %sca1_c, %kblocks : index + %sca1_1 = index.add %sca1_0, %chunk_blk : index + %sca1_2 = index.add %sca1_1, %sc1_blk : index + %sca1 = index.assume %sca1_2 [range(%sca1_2, 0, 33554431)] : index + %scav1 = view.load %as_flat[%sca1] : view<33554432xf32> -> f32 + %sca1_tile = index.div %sc1_col, %c16 : index + %sca1_inner = index.rem %sc1_col, %c16 : index + %sca1_par = index.rem %sca1_inner, %c2 : index + %sca1_v = index.div %sca1_inner, %c2 : index + %sca1_dst0 = index.mul %sca1_tile, %k_blocks : index + %sca1_dst1 = index.add %sca1_dst0, %sc1_blk : index + %sca1_dst2 = index.mul %sca1_dst1, %c2 : index + %sca1_dst3 = index.add %sca1_dst2, %sca1_par : index + %sca1_dst4 = index.mul %sca1_dst3, %c8 : index + %sca1_dst5 = index.add %sca1_dst4, %sca1_v : index + %sca1_dst = index.assume %sca1_dst5 [range(%sca1_dst5, 0, 255)] : index + view.store %scav1, %asl_view[%sca1_dst] : f32, view<256xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + + // Two K32 blocks explicitly unrolled. Arithmetic order and the + // enclosing K64 staging/barrier cadence are unchanged. + %blk_k_u0 = index.mul %c0, %c16 : index + %blk_k16_u0 = index.add %blk_k_u0, %c8 : index + %iz_u0 = vector.fragment %izero shape [%c16, %c16] : vector<8xi32> + %wsi0_2_u0 = index.add %wsb0, %c0 : index + %wsi0_u0 = index.assume %wsi0_2_u0 [range(%wsi0_2_u0, 0, 255)] : index + %wsc0_u0 = view.load %wsl_view[%wsi0_u0] : view<256xf32> -> f32 + %wsv0_u0 = vector.splat %wsc0_u0 : vector<8xf32> + %wsi1_2_u0 = index.add %wsb1, %c0 : index + %wsi1_u0 = index.assume %wsi1_2_u0 [range(%wsi1_2_u0, 0, 255)] : index + %wsc1_u0 = view.load %wsl_view[%wsi1_u0] : view<256xf32> -> f32 + %wsv1_u0 = vector.splat %wsc1_u0 : vector<8xf32> + %wsi4_2_u0 = index.add %wsb4, %c0 : index + %wsi4_u0 = index.assume %wsi4_2_u0 [range(%wsi4_2_u0, 0, 255)] : index + %wsc4_u0 = view.load %wsl_view[%wsi4_u0] : view<256xf32> -> f32 + %wsv4_u0 = vector.splat %wsc4_u0 : vector<8xf32> + %wsi5_2_u0 = index.add %wsb5, %c0 : index + %wsi5_u0 = index.assume %wsi5_2_u0 [range(%wsi5_2_u0, 0, 255)] : index + %wsc5_u0 = view.load %wsl_view[%wsi5_u0] : view<256xf32> -> f32 + %wsv5_u0 = vector.splat %wsc5_u0 : vector<8xf32> + %lf0_0_lane_row_u0 = index.add %lm0, %lane_lo : index + %lf0_0_row_u0 = index.mul %lf0_0_lane_row_u0, %k_astride : index + %lf0_0_base_u0 = index.add %lf0_0_row_u0, %blk_k_u0 : index + %lf0_0_idx_u0 = index.assume %lf0_0_base_u0 [range(%lf0_0_base_u0, 0, 5112)] : index + %lf0_0_raw_u0 = vector.load %al_flat[%lf0_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %lf0_0_words_u0 = vector.bitcast %lf0_0_raw_u0 : vector<8xi8> to vector<2xi32> + %lf0_0_u0 = vector.fragment %lf0_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %lf0_1_lane_row_u0 = index.add %lm0, %lane_lo : index + %lf0_1_row_u0 = index.mul %lf0_1_lane_row_u0, %k_astride : index + %lf0_1_base_u0 = index.add %lf0_1_row_u0, %blk_k16_u0 : index + %lf0_1_idx_u0 = index.assume %lf0_1_base_u0 [range(%lf0_1_base_u0, 0, 5112)] : index + %lf0_1_raw_u0 = vector.load %al_flat[%lf0_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %lf0_1_words_u0 = vector.bitcast %lf0_1_raw_u0 : vector<8xi8> to vector<2xi32> + %lf0_1_u0 = vector.fragment %lf0_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %asi0_0_u0 = index.add %asb0, %c0 : index + %asi0_1_u0 = index.mul %asi0_0_u0, %c2 : index + %asi0_2_u0 = index.add %asi0_1_u0, %lane_hi : index + %asi0_3_u0 = index.mul %asi0_2_u0, %c8 : index + %asi0_u0 = index.assume %asi0_3_u0 [range(%asi0_3_u0, 0, 120)] : index + %asv0_u0 = vector.load %asl_view[%asi0_u0] : view<256xf32> -> vector<8xf32> + %rfs0_0_0_row_u0 = index.mul %ln0, %k_wstride : index + %rfs0_0_0_base_u0 = index.add %rfs0_0_0_row_u0, %blk_k_u0 : index + %rfs0_0_0_idx_u0 = index.assume %rfs0_0_0_base_u0 [range(%rfs0_0_0_base_u0, 0, 5112)] : index + %rfs0_0_0_raw_u0 = vector.load %wl_flat[%rfs0_0_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs0_0_0_words_u0 = vector.bitcast %rfs0_0_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs0_0_0_u0 = vector.fragment %rfs0_0_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_0_1_row_u0 = index.mul %ln0, %k_wstride : index + %rfs0_0_1_base_u0 = index.add %rfs0_0_1_row_u0, %blk_k16_u0 : index + %rfs0_0_1_idx_u0 = index.assume %rfs0_0_1_base_u0 [range(%rfs0_0_1_base_u0, 0, 5112)] : index + %rfs0_0_1_raw_u0 = vector.load %wl_flat[%rfs0_0_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs0_0_1_words_u0 = vector.bitcast %rfs0_0_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs0_0_1_u0 = vector.fragment %rfs0_0_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_0_u0 = vector.mma %lf0_0_u0, %rfs0_0_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_0_u0 = vector.mma %lf0_1_u0, %rfs0_0_1_u0, %i0_0_0_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff0_0_u0 = vector.sitofp %i1_0_0_u0 : vector<8xi32> to vector<8xf32> + %sv0_0_u0 = vector.mulf %asv0_u0, %wsv0_u0 : vector<8xf32> + %fn0_0_u0 = vector.fmaf %ff0_0_u0, %sv0_0_u0, %fc0_0 : vector<8xf32> + scf.schedule.fence + + %rfs0_1_0_row_u0 = index.mul %ln1, %k_wstride : index + %rfs0_1_0_base_u0 = index.add %rfs0_1_0_row_u0, %blk_k_u0 : index + %rfs0_1_0_idx_u0 = index.assume %rfs0_1_0_base_u0 [range(%rfs0_1_0_base_u0, 0, 5112)] : index + %rfs0_1_0_raw_u0 = vector.load %wl_flat[%rfs0_1_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs0_1_0_words_u0 = vector.bitcast %rfs0_1_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs0_1_0_u0 = vector.fragment %rfs0_1_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_1_1_row_u0 = index.mul %ln1, %k_wstride : index + %rfs0_1_1_base_u0 = index.add %rfs0_1_1_row_u0, %blk_k16_u0 : index + %rfs0_1_1_idx_u0 = index.assume %rfs0_1_1_base_u0 [range(%rfs0_1_1_base_u0, 0, 5112)] : index + %rfs0_1_1_raw_u0 = vector.load %wl_flat[%rfs0_1_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs0_1_1_words_u0 = vector.bitcast %rfs0_1_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs0_1_1_u0 = vector.fragment %rfs0_1_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_1_u0 = vector.mma %lf0_0_u0, %rfs0_1_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_1_u0 = vector.mma %lf0_1_u0, %rfs0_1_1_u0, %i0_0_1_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff0_1_u0 = vector.sitofp %i1_0_1_u0 : vector<8xi32> to vector<8xf32> + %sv0_1_u0 = vector.mulf %asv0_u0, %wsv1_u0 : vector<8xf32> + %fn0_1_u0 = vector.fmaf %ff0_1_u0, %sv0_1_u0, %fc0_1 : vector<8xf32> + scf.schedule.fence + + %rfs0_2_0_row_u0 = index.mul %ln4, %k_wstride : index + %rfs0_2_0_base_u0 = index.add %rfs0_2_0_row_u0, %blk_k_u0 : index + %rfs0_2_0_idx_u0 = index.assume %rfs0_2_0_base_u0 [range(%rfs0_2_0_base_u0, 0, 5112)] : index + %rfs0_2_0_raw_u0 = vector.load %wl_flat[%rfs0_2_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs0_2_0_words_u0 = vector.bitcast %rfs0_2_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs0_2_0_u0 = vector.fragment %rfs0_2_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_2_1_row_u0 = index.mul %ln4, %k_wstride : index + %rfs0_2_1_base_u0 = index.add %rfs0_2_1_row_u0, %blk_k16_u0 : index + %rfs0_2_1_idx_u0 = index.assume %rfs0_2_1_base_u0 [range(%rfs0_2_1_base_u0, 0, 5112)] : index + %rfs0_2_1_raw_u0 = vector.load %wl_flat[%rfs0_2_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs0_2_1_words_u0 = vector.bitcast %rfs0_2_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs0_2_1_u0 = vector.fragment %rfs0_2_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_2_u0 = vector.mma %lf0_0_u0, %rfs0_2_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_2_u0 = vector.mma %lf0_1_u0, %rfs0_2_1_u0, %i0_0_2_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff0_2_u0 = vector.sitofp %i1_0_2_u0 : vector<8xi32> to vector<8xf32> + %sv0_2_u0 = vector.mulf %asv0_u0, %wsv4_u0 : vector<8xf32> + %fn0_2_u0 = vector.fmaf %ff0_2_u0, %sv0_2_u0, %fc0_2 : vector<8xf32> + scf.schedule.fence + + %rfs0_3_0_row_u0 = index.mul %ln5, %k_wstride : index + %rfs0_3_0_base_u0 = index.add %rfs0_3_0_row_u0, %blk_k_u0 : index + %rfs0_3_0_idx_u0 = index.assume %rfs0_3_0_base_u0 [range(%rfs0_3_0_base_u0, 0, 5112)] : index + %rfs0_3_0_raw_u0 = vector.load %wl_flat[%rfs0_3_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs0_3_0_words_u0 = vector.bitcast %rfs0_3_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs0_3_0_u0 = vector.fragment %rfs0_3_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_3_1_row_u0 = index.mul %ln5, %k_wstride : index + %rfs0_3_1_base_u0 = index.add %rfs0_3_1_row_u0, %blk_k16_u0 : index + %rfs0_3_1_idx_u0 = index.assume %rfs0_3_1_base_u0 [range(%rfs0_3_1_base_u0, 0, 5112)] : index + %rfs0_3_1_raw_u0 = vector.load %wl_flat[%rfs0_3_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs0_3_1_words_u0 = vector.bitcast %rfs0_3_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs0_3_1_u0 = vector.fragment %rfs0_3_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_3_u0 = vector.mma %lf0_0_u0, %rfs0_3_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_3_u0 = vector.mma %lf0_1_u0, %rfs0_3_1_u0, %i0_0_3_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff0_3_u0 = vector.sitofp %i1_0_3_u0 : vector<8xi32> to vector<8xf32> + %sv0_3_u0 = vector.mulf %asv0_u0, %wsv5_u0 : vector<8xf32> + %fn0_3_u0 = vector.fmaf %ff0_3_u0, %sv0_3_u0, %fc0_3 : vector<8xf32> + + scf.schedule.fence + %wsi2_2_u0 = index.add %wsb2, %c0 : index + %wsi2_u0 = index.assume %wsi2_2_u0 [range(%wsi2_2_u0, 0, 255)] : index + %wsc2_u0 = view.load %wsl_view[%wsi2_u0] : view<256xf32> -> f32 + %wsv2_u0 = vector.splat %wsc2_u0 : vector<8xf32> + %wsi3_2_u0 = index.add %wsb3, %c0 : index + %wsi3_u0 = index.assume %wsi3_2_u0 [range(%wsi3_2_u0, 0, 255)] : index + %wsc3_u0 = view.load %wsl_view[%wsi3_u0] : view<256xf32> -> f32 + %wsv3_u0 = vector.splat %wsc3_u0 : vector<8xf32> + %wsi6_2_u0 = index.add %wsb6, %c0 : index + %wsi6_u0 = index.assume %wsi6_2_u0 [range(%wsi6_2_u0, 0, 255)] : index + %wsc6_u0 = view.load %wsl_view[%wsi6_u0] : view<256xf32> -> f32 + %wsv6_u0 = vector.splat %wsc6_u0 : vector<8xf32> + %wsi7_2_u0 = index.add %wsb7, %c0 : index + %wsi7_u0 = index.assume %wsi7_2_u0 [range(%wsi7_2_u0, 0, 255)] : index + %wsc7_u0 = view.load %wsl_view[%wsi7_u0] : view<256xf32> -> f32 + %wsv7_u0 = vector.splat %wsc7_u0 : vector<8xf32> + %rfs1_0_0_row_u0 = index.mul %ln2, %k_wstride : index + %rfs1_0_0_base_u0 = index.add %rfs1_0_0_row_u0, %blk_k_u0 : index + %rfs1_0_0_idx_u0 = index.assume %rfs1_0_0_base_u0 [range(%rfs1_0_0_base_u0, 0, 5112)] : index + %rfs1_0_0_raw_u0 = vector.load %wl_flat[%rfs1_0_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs1_0_0_words_u0 = vector.bitcast %rfs1_0_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs1_0_0_u0 = vector.fragment %rfs1_0_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs1_0_1_row_u0 = index.mul %ln2, %k_wstride : index + %rfs1_0_1_base_u0 = index.add %rfs1_0_1_row_u0, %blk_k16_u0 : index + %rfs1_0_1_idx_u0 = index.assume %rfs1_0_1_base_u0 [range(%rfs1_0_1_base_u0, 0, 5112)] : index + %rfs1_0_1_raw_u0 = vector.load %wl_flat[%rfs1_0_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs1_0_1_words_u0 = vector.bitcast %rfs1_0_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs1_0_1_u0 = vector.fragment %rfs1_0_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_1_0_u0 = vector.mma %lf0_0_u0, %rfs1_0_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_1_0_u0 = vector.mma %lf0_1_u0, %rfs1_0_1_u0, %i0_1_0_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff1_0_u0 = vector.sitofp %i1_1_0_u0 : vector<8xi32> to vector<8xf32> + %sv1_0_u0 = vector.mulf %asv0_u0, %wsv2_u0 : vector<8xf32> + %fn1_0_u0 = vector.fmaf %ff1_0_u0, %sv1_0_u0, %fc1_0 : vector<8xf32> + scf.schedule.fence + + %rfs1_1_0_row_u0 = index.mul %ln3, %k_wstride : index + %rfs1_1_0_base_u0 = index.add %rfs1_1_0_row_u0, %blk_k_u0 : index + %rfs1_1_0_idx_u0 = index.assume %rfs1_1_0_base_u0 [range(%rfs1_1_0_base_u0, 0, 5112)] : index + %rfs1_1_0_raw_u0 = vector.load %wl_flat[%rfs1_1_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs1_1_0_words_u0 = vector.bitcast %rfs1_1_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs1_1_0_u0 = vector.fragment %rfs1_1_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs1_1_1_row_u0 = index.mul %ln3, %k_wstride : index + %rfs1_1_1_base_u0 = index.add %rfs1_1_1_row_u0, %blk_k16_u0 : index + %rfs1_1_1_idx_u0 = index.assume %rfs1_1_1_base_u0 [range(%rfs1_1_1_base_u0, 0, 5112)] : index + %rfs1_1_1_raw_u0 = vector.load %wl_flat[%rfs1_1_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs1_1_1_words_u0 = vector.bitcast %rfs1_1_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs1_1_1_u0 = vector.fragment %rfs1_1_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_1_1_u0 = vector.mma %lf0_0_u0, %rfs1_1_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_1_1_u0 = vector.mma %lf0_1_u0, %rfs1_1_1_u0, %i0_1_1_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff1_1_u0 = vector.sitofp %i1_1_1_u0 : vector<8xi32> to vector<8xf32> + %sv1_1_u0 = vector.mulf %asv0_u0, %wsv3_u0 : vector<8xf32> + %fn1_1_u0 = vector.fmaf %ff1_1_u0, %sv1_1_u0, %fc1_1 : vector<8xf32> + scf.schedule.fence + + %rfs1_2_0_row_u0 = index.mul %ln6, %k_wstride : index + %rfs1_2_0_base_u0 = index.add %rfs1_2_0_row_u0, %blk_k_u0 : index + %rfs1_2_0_idx_u0 = index.assume %rfs1_2_0_base_u0 [range(%rfs1_2_0_base_u0, 0, 5112)] : index + %rfs1_2_0_raw_u0 = vector.load %wl_flat[%rfs1_2_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs1_2_0_words_u0 = vector.bitcast %rfs1_2_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs1_2_0_u0 = vector.fragment %rfs1_2_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs1_2_1_row_u0 = index.mul %ln6, %k_wstride : index + %rfs1_2_1_base_u0 = index.add %rfs1_2_1_row_u0, %blk_k16_u0 : index + %rfs1_2_1_idx_u0 = index.assume %rfs1_2_1_base_u0 [range(%rfs1_2_1_base_u0, 0, 5112)] : index + %rfs1_2_1_raw_u0 = vector.load %wl_flat[%rfs1_2_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs1_2_1_words_u0 = vector.bitcast %rfs1_2_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs1_2_1_u0 = vector.fragment %rfs1_2_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_1_2_u0 = vector.mma %lf0_0_u0, %rfs1_2_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_1_2_u0 = vector.mma %lf0_1_u0, %rfs1_2_1_u0, %i0_1_2_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff1_2_u0 = vector.sitofp %i1_1_2_u0 : vector<8xi32> to vector<8xf32> + %sv1_2_u0 = vector.mulf %asv0_u0, %wsv6_u0 : vector<8xf32> + %fn1_2_u0 = vector.fmaf %ff1_2_u0, %sv1_2_u0, %fc1_2 : vector<8xf32> + scf.schedule.fence + + %rfs1_3_0_row_u0 = index.mul %ln7, %k_wstride : index + %rfs1_3_0_base_u0 = index.add %rfs1_3_0_row_u0, %blk_k_u0 : index + %rfs1_3_0_idx_u0 = index.assume %rfs1_3_0_base_u0 [range(%rfs1_3_0_base_u0, 0, 5112)] : index + %rfs1_3_0_raw_u0 = vector.load %wl_flat[%rfs1_3_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs1_3_0_words_u0 = vector.bitcast %rfs1_3_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs1_3_0_u0 = vector.fragment %rfs1_3_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs1_3_1_row_u0 = index.mul %ln7, %k_wstride : index + %rfs1_3_1_base_u0 = index.add %rfs1_3_1_row_u0, %blk_k16_u0 : index + %rfs1_3_1_idx_u0 = index.assume %rfs1_3_1_base_u0 [range(%rfs1_3_1_base_u0, 0, 5112)] : index + %rfs1_3_1_raw_u0 = vector.load %wl_flat[%rfs1_3_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs1_3_1_words_u0 = vector.bitcast %rfs1_3_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs1_3_1_u0 = vector.fragment %rfs1_3_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_1_3_u0 = vector.mma %lf0_0_u0, %rfs1_3_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_1_3_u0 = vector.mma %lf0_1_u0, %rfs1_3_1_u0, %i0_1_3_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff1_3_u0 = vector.sitofp %i1_1_3_u0 : vector<8xi32> to vector<8xf32> + %sv1_3_u0 = vector.mulf %asv0_u0, %wsv7_u0 : vector<8xf32> + %fn1_3_u0 = vector.fmaf %ff1_3_u0, %sv1_3_u0, %fc1_3 : vector<8xf32> + + scf.schedule.fence + %lf2_0_lane_row_u0 = index.add %lm2, %lane_lo : index + %lf2_0_row_u0 = index.mul %lf2_0_lane_row_u0, %k_astride : index + %lf2_0_base_u0 = index.add %lf2_0_row_u0, %blk_k_u0 : index + %lf2_0_idx_u0 = index.assume %lf2_0_base_u0 [range(%lf2_0_base_u0, 0, 5112)] : index + %lf2_0_raw_u0 = vector.load %al_flat[%lf2_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %lf2_0_words_u0 = vector.bitcast %lf2_0_raw_u0 : vector<8xi8> to vector<2xi32> + %lf2_0_u0 = vector.fragment %lf2_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %lf2_1_lane_row_u0 = index.add %lm2, %lane_lo : index + %lf2_1_row_u0 = index.mul %lf2_1_lane_row_u0, %k_astride : index + %lf2_1_base_u0 = index.add %lf2_1_row_u0, %blk_k16_u0 : index + %lf2_1_idx_u0 = index.assume %lf2_1_base_u0 [range(%lf2_1_base_u0, 0, 5112)] : index + %lf2_1_raw_u0 = vector.load %al_flat[%lf2_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %lf2_1_words_u0 = vector.bitcast %lf2_1_raw_u0 : vector<8xi8> to vector<2xi32> + %lf2_1_u0 = vector.fragment %lf2_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %asi2_0_u0 = index.add %asb2, %c0 : index + %asi2_1_u0 = index.mul %asi2_0_u0, %c2 : index + %asi2_2_u0 = index.add %asi2_1_u0, %lane_hi : index + %asi2_3_u0 = index.mul %asi2_2_u0, %c8 : index + %asi2_u0 = index.assume %asi2_3_u0 [range(%asi2_3_u0, 128, 248)] : index + %asv2_u0 = vector.load %asl_view[%asi2_u0] : view<256xf32> -> vector<8xf32> + %rfs2_0_0_row_u0 = index.mul %ln0, %k_wstride : index + %rfs2_0_0_base_u0 = index.add %rfs2_0_0_row_u0, %blk_k_u0 : index + %rfs2_0_0_idx_u0 = index.assume %rfs2_0_0_base_u0 [range(%rfs2_0_0_base_u0, 0, 5112)] : index + %rfs2_0_0_raw_u0 = vector.load %wl_flat[%rfs2_0_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs2_0_0_words_u0 = vector.bitcast %rfs2_0_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs2_0_0_u0 = vector.fragment %rfs2_0_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs2_0_1_row_u0 = index.mul %ln0, %k_wstride : index + %rfs2_0_1_base_u0 = index.add %rfs2_0_1_row_u0, %blk_k16_u0 : index + %rfs2_0_1_idx_u0 = index.assume %rfs2_0_1_base_u0 [range(%rfs2_0_1_base_u0, 0, 5112)] : index + %rfs2_0_1_raw_u0 = vector.load %wl_flat[%rfs2_0_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs2_0_1_words_u0 = vector.bitcast %rfs2_0_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs2_0_1_u0 = vector.fragment %rfs2_0_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_2_0_u0 = vector.mma %lf2_0_u0, %rfs2_0_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_2_0_u0 = vector.mma %lf2_1_u0, %rfs2_0_1_u0, %i0_2_0_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff2_0_u0 = vector.sitofp %i1_2_0_u0 : vector<8xi32> to vector<8xf32> + %sv2_0_u0 = vector.mulf %asv2_u0, %wsv0_u0 : vector<8xf32> + %fn2_0_u0 = vector.fmaf %ff2_0_u0, %sv2_0_u0, %fc2_0 : vector<8xf32> + scf.schedule.fence + + %rfs2_1_0_row_u0 = index.mul %ln1, %k_wstride : index + %rfs2_1_0_base_u0 = index.add %rfs2_1_0_row_u0, %blk_k_u0 : index + %rfs2_1_0_idx_u0 = index.assume %rfs2_1_0_base_u0 [range(%rfs2_1_0_base_u0, 0, 5112)] : index + %rfs2_1_0_raw_u0 = vector.load %wl_flat[%rfs2_1_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs2_1_0_words_u0 = vector.bitcast %rfs2_1_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs2_1_0_u0 = vector.fragment %rfs2_1_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs2_1_1_row_u0 = index.mul %ln1, %k_wstride : index + %rfs2_1_1_base_u0 = index.add %rfs2_1_1_row_u0, %blk_k16_u0 : index + %rfs2_1_1_idx_u0 = index.assume %rfs2_1_1_base_u0 [range(%rfs2_1_1_base_u0, 0, 5112)] : index + %rfs2_1_1_raw_u0 = vector.load %wl_flat[%rfs2_1_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs2_1_1_words_u0 = vector.bitcast %rfs2_1_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs2_1_1_u0 = vector.fragment %rfs2_1_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_2_1_u0 = vector.mma %lf2_0_u0, %rfs2_1_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_2_1_u0 = vector.mma %lf2_1_u0, %rfs2_1_1_u0, %i0_2_1_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff2_1_u0 = vector.sitofp %i1_2_1_u0 : vector<8xi32> to vector<8xf32> + %sv2_1_u0 = vector.mulf %asv2_u0, %wsv1_u0 : vector<8xf32> + %fn2_1_u0 = vector.fmaf %ff2_1_u0, %sv2_1_u0, %fc2_1 : vector<8xf32> + scf.schedule.fence + + %rfs2_2_0_row_u0 = index.mul %ln4, %k_wstride : index + %rfs2_2_0_base_u0 = index.add %rfs2_2_0_row_u0, %blk_k_u0 : index + %rfs2_2_0_idx_u0 = index.assume %rfs2_2_0_base_u0 [range(%rfs2_2_0_base_u0, 0, 5112)] : index + %rfs2_2_0_raw_u0 = vector.load %wl_flat[%rfs2_2_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs2_2_0_words_u0 = vector.bitcast %rfs2_2_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs2_2_0_u0 = vector.fragment %rfs2_2_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs2_2_1_row_u0 = index.mul %ln4, %k_wstride : index + %rfs2_2_1_base_u0 = index.add %rfs2_2_1_row_u0, %blk_k16_u0 : index + %rfs2_2_1_idx_u0 = index.assume %rfs2_2_1_base_u0 [range(%rfs2_2_1_base_u0, 0, 5112)] : index + %rfs2_2_1_raw_u0 = vector.load %wl_flat[%rfs2_2_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs2_2_1_words_u0 = vector.bitcast %rfs2_2_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs2_2_1_u0 = vector.fragment %rfs2_2_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_2_2_u0 = vector.mma %lf2_0_u0, %rfs2_2_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_2_2_u0 = vector.mma %lf2_1_u0, %rfs2_2_1_u0, %i0_2_2_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff2_2_u0 = vector.sitofp %i1_2_2_u0 : vector<8xi32> to vector<8xf32> + %sv2_2_u0 = vector.mulf %asv2_u0, %wsv4_u0 : vector<8xf32> + %fn2_2_u0 = vector.fmaf %ff2_2_u0, %sv2_2_u0, %fc2_2 : vector<8xf32> + scf.schedule.fence + + %rfs2_3_0_row_u0 = index.mul %ln5, %k_wstride : index + %rfs2_3_0_base_u0 = index.add %rfs2_3_0_row_u0, %blk_k_u0 : index + %rfs2_3_0_idx_u0 = index.assume %rfs2_3_0_base_u0 [range(%rfs2_3_0_base_u0, 0, 5112)] : index + %rfs2_3_0_raw_u0 = vector.load %wl_flat[%rfs2_3_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs2_3_0_words_u0 = vector.bitcast %rfs2_3_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs2_3_0_u0 = vector.fragment %rfs2_3_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs2_3_1_row_u0 = index.mul %ln5, %k_wstride : index + %rfs2_3_1_base_u0 = index.add %rfs2_3_1_row_u0, %blk_k16_u0 : index + %rfs2_3_1_idx_u0 = index.assume %rfs2_3_1_base_u0 [range(%rfs2_3_1_base_u0, 0, 5112)] : index + %rfs2_3_1_raw_u0 = vector.load %wl_flat[%rfs2_3_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs2_3_1_words_u0 = vector.bitcast %rfs2_3_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs2_3_1_u0 = vector.fragment %rfs2_3_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_2_3_u0 = vector.mma %lf2_0_u0, %rfs2_3_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_2_3_u0 = vector.mma %lf2_1_u0, %rfs2_3_1_u0, %i0_2_3_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff2_3_u0 = vector.sitofp %i1_2_3_u0 : vector<8xi32> to vector<8xf32> + %sv2_3_u0 = vector.mulf %asv2_u0, %wsv5_u0 : vector<8xf32> + %fn2_3_u0 = vector.fmaf %ff2_3_u0, %sv2_3_u0, %fc2_3 : vector<8xf32> + + scf.schedule.fence + %rfs3_0_0_row_u0 = index.mul %ln2, %k_wstride : index + %rfs3_0_0_base_u0 = index.add %rfs3_0_0_row_u0, %blk_k_u0 : index + %rfs3_0_0_idx_u0 = index.assume %rfs3_0_0_base_u0 [range(%rfs3_0_0_base_u0, 0, 5112)] : index + %rfs3_0_0_raw_u0 = vector.load %wl_flat[%rfs3_0_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs3_0_0_words_u0 = vector.bitcast %rfs3_0_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs3_0_0_u0 = vector.fragment %rfs3_0_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs3_0_1_row_u0 = index.mul %ln2, %k_wstride : index + %rfs3_0_1_base_u0 = index.add %rfs3_0_1_row_u0, %blk_k16_u0 : index + %rfs3_0_1_idx_u0 = index.assume %rfs3_0_1_base_u0 [range(%rfs3_0_1_base_u0, 0, 5112)] : index + %rfs3_0_1_raw_u0 = vector.load %wl_flat[%rfs3_0_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs3_0_1_words_u0 = vector.bitcast %rfs3_0_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs3_0_1_u0 = vector.fragment %rfs3_0_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_3_0_u0 = vector.mma %lf2_0_u0, %rfs3_0_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_3_0_u0 = vector.mma %lf2_1_u0, %rfs3_0_1_u0, %i0_3_0_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff3_0_u0 = vector.sitofp %i1_3_0_u0 : vector<8xi32> to vector<8xf32> + %sv3_0_u0 = vector.mulf %asv2_u0, %wsv2_u0 : vector<8xf32> + %fn3_0_u0 = vector.fmaf %ff3_0_u0, %sv3_0_u0, %fc3_0 : vector<8xf32> + scf.schedule.fence + + %rfs3_1_0_row_u0 = index.mul %ln3, %k_wstride : index + %rfs3_1_0_base_u0 = index.add %rfs3_1_0_row_u0, %blk_k_u0 : index + %rfs3_1_0_idx_u0 = index.assume %rfs3_1_0_base_u0 [range(%rfs3_1_0_base_u0, 0, 5112)] : index + %rfs3_1_0_raw_u0 = vector.load %wl_flat[%rfs3_1_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs3_1_0_words_u0 = vector.bitcast %rfs3_1_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs3_1_0_u0 = vector.fragment %rfs3_1_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs3_1_1_row_u0 = index.mul %ln3, %k_wstride : index + %rfs3_1_1_base_u0 = index.add %rfs3_1_1_row_u0, %blk_k16_u0 : index + %rfs3_1_1_idx_u0 = index.assume %rfs3_1_1_base_u0 [range(%rfs3_1_1_base_u0, 0, 5112)] : index + %rfs3_1_1_raw_u0 = vector.load %wl_flat[%rfs3_1_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs3_1_1_words_u0 = vector.bitcast %rfs3_1_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs3_1_1_u0 = vector.fragment %rfs3_1_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_3_1_u0 = vector.mma %lf2_0_u0, %rfs3_1_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_3_1_u0 = vector.mma %lf2_1_u0, %rfs3_1_1_u0, %i0_3_1_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff3_1_u0 = vector.sitofp %i1_3_1_u0 : vector<8xi32> to vector<8xf32> + %sv3_1_u0 = vector.mulf %asv2_u0, %wsv3_u0 : vector<8xf32> + %fn3_1_u0 = vector.fmaf %ff3_1_u0, %sv3_1_u0, %fc3_1 : vector<8xf32> + scf.schedule.fence + + %rfs3_2_0_row_u0 = index.mul %ln6, %k_wstride : index + %rfs3_2_0_base_u0 = index.add %rfs3_2_0_row_u0, %blk_k_u0 : index + %rfs3_2_0_idx_u0 = index.assume %rfs3_2_0_base_u0 [range(%rfs3_2_0_base_u0, 0, 5112)] : index + %rfs3_2_0_raw_u0 = vector.load %wl_flat[%rfs3_2_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs3_2_0_words_u0 = vector.bitcast %rfs3_2_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs3_2_0_u0 = vector.fragment %rfs3_2_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs3_2_1_row_u0 = index.mul %ln6, %k_wstride : index + %rfs3_2_1_base_u0 = index.add %rfs3_2_1_row_u0, %blk_k16_u0 : index + %rfs3_2_1_idx_u0 = index.assume %rfs3_2_1_base_u0 [range(%rfs3_2_1_base_u0, 0, 5112)] : index + %rfs3_2_1_raw_u0 = vector.load %wl_flat[%rfs3_2_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs3_2_1_words_u0 = vector.bitcast %rfs3_2_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs3_2_1_u0 = vector.fragment %rfs3_2_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_3_2_u0 = vector.mma %lf2_0_u0, %rfs3_2_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_3_2_u0 = vector.mma %lf2_1_u0, %rfs3_2_1_u0, %i0_3_2_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff3_2_u0 = vector.sitofp %i1_3_2_u0 : vector<8xi32> to vector<8xf32> + %sv3_2_u0 = vector.mulf %asv2_u0, %wsv6_u0 : vector<8xf32> + %fn3_2_u0 = vector.fmaf %ff3_2_u0, %sv3_2_u0, %fc3_2 : vector<8xf32> + scf.schedule.fence + + %rfs3_3_0_row_u0 = index.mul %ln7, %k_wstride : index + %rfs3_3_0_base_u0 = index.add %rfs3_3_0_row_u0, %blk_k_u0 : index + %rfs3_3_0_idx_u0 = index.assume %rfs3_3_0_base_u0 [range(%rfs3_3_0_base_u0, 0, 5112)] : index + %rfs3_3_0_raw_u0 = vector.load %wl_flat[%rfs3_3_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs3_3_0_words_u0 = vector.bitcast %rfs3_3_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs3_3_0_u0 = vector.fragment %rfs3_3_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs3_3_1_row_u0 = index.mul %ln7, %k_wstride : index + %rfs3_3_1_base_u0 = index.add %rfs3_3_1_row_u0, %blk_k16_u0 : index + %rfs3_3_1_idx_u0 = index.assume %rfs3_3_1_base_u0 [range(%rfs3_3_1_base_u0, 0, 5112)] : index + %rfs3_3_1_raw_u0 = vector.load %wl_flat[%rfs3_3_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs3_3_1_words_u0 = vector.bitcast %rfs3_3_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs3_3_1_u0 = vector.fragment %rfs3_3_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_3_3_u0 = vector.mma %lf2_0_u0, %rfs3_3_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_3_3_u0 = vector.mma %lf2_1_u0, %rfs3_3_1_u0, %i0_3_3_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff3_3_u0 = vector.sitofp %i1_3_3_u0 : vector<8xi32> to vector<8xf32> + %sv3_3_u0 = vector.mulf %asv2_u0, %wsv7_u0 : vector<8xf32> + %fn3_3_u0 = vector.fmaf %ff3_3_u0, %sv3_3_u0, %fc3_3 : vector<8xf32> + + scf.schedule.fence + %blk_k_u1 = index.mul %c1, %c16 : index + %blk_k16_u1 = index.add %blk_k_u1, %c8 : index + %wsi0_2_u1 = index.add %wsb0, %c1 : index + %wsi0_u1 = index.assume %wsi0_2_u1 [range(%wsi0_2_u1, 1, 255)] : index + %wsc0_u1 = view.load %wsl_view[%wsi0_u1] : view<256xf32> -> f32 + %wsv0_u1 = vector.splat %wsc0_u1 : vector<8xf32> + %wsi1_2_u1 = index.add %wsb1, %c1 : index + %wsi1_u1 = index.assume %wsi1_2_u1 [range(%wsi1_2_u1, 1, 255)] : index + %wsc1_u1 = view.load %wsl_view[%wsi1_u1] : view<256xf32> -> f32 + %wsv1_u1 = vector.splat %wsc1_u1 : vector<8xf32> + %wsi4_2_u1 = index.add %wsb4, %c1 : index + %wsi4_u1 = index.assume %wsi4_2_u1 [range(%wsi4_2_u1, 1, 255)] : index + %wsc4_u1 = view.load %wsl_view[%wsi4_u1] : view<256xf32> -> f32 + %wsv4_u1 = vector.splat %wsc4_u1 : vector<8xf32> + %wsi5_2_u1 = index.add %wsb5, %c1 : index + %wsi5_u1 = index.assume %wsi5_2_u1 [range(%wsi5_2_u1, 1, 255)] : index + %wsc5_u1 = view.load %wsl_view[%wsi5_u1] : view<256xf32> -> f32 + %wsv5_u1 = vector.splat %wsc5_u1 : vector<8xf32> + %lf0_0_lane_row_u1 = index.add %lm0, %lane_lo : index + %lf0_0_row_u1 = index.mul %lf0_0_lane_row_u1, %k_astride : index + %lf0_0_base_u1 = index.add %lf0_0_row_u1, %blk_k_u1 : index + %lf0_0_idx_u1 = index.assume %lf0_0_base_u1 [range(%lf0_0_base_u1, 0, 5112)] : index + %lf0_0_raw_u1 = vector.load %al_flat[%lf0_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %lf0_0_words_u1 = vector.bitcast %lf0_0_raw_u1 : vector<8xi8> to vector<2xi32> + %lf0_0_u1 = vector.fragment %lf0_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %lf0_1_lane_row_u1 = index.add %lm0, %lane_lo : index + %lf0_1_row_u1 = index.mul %lf0_1_lane_row_u1, %k_astride : index + %lf0_1_base_u1 = index.add %lf0_1_row_u1, %blk_k16_u1 : index + %lf0_1_idx_u1 = index.assume %lf0_1_base_u1 [range(%lf0_1_base_u1, 0, 5112)] : index + %lf0_1_raw_u1 = vector.load %al_flat[%lf0_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %lf0_1_words_u1 = vector.bitcast %lf0_1_raw_u1 : vector<8xi8> to vector<2xi32> + %lf0_1_u1 = vector.fragment %lf0_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_0_0_row_u1 = index.mul %ln0, %k_wstride : index + %rfs0_0_0_base_u1 = index.add %rfs0_0_0_row_u1, %blk_k_u1 : index + %rfs0_0_0_idx_u1 = index.assume %rfs0_0_0_base_u1 [range(%rfs0_0_0_base_u1, 0, 5112)] : index + %rfs0_0_0_raw_u1 = vector.load %wl_flat[%rfs0_0_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs0_0_0_words_u1 = vector.bitcast %rfs0_0_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs0_0_0_u1 = vector.fragment %rfs0_0_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_0_1_row_u1 = index.mul %ln0, %k_wstride : index + %rfs0_0_1_base_u1 = index.add %rfs0_0_1_row_u1, %blk_k16_u1 : index + %rfs0_0_1_idx_u1 = index.assume %rfs0_0_1_base_u1 [range(%rfs0_0_1_base_u1, 0, 5112)] : index + %rfs0_0_1_raw_u1 = vector.load %wl_flat[%rfs0_0_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs0_0_1_words_u1 = vector.bitcast %rfs0_0_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs0_0_1_u1 = vector.fragment %rfs0_0_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_0_u1 = vector.mma %lf0_0_u1, %rfs0_0_0_u1, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_0_u1 = vector.mma %lf0_1_u1, %rfs0_0_1_u1, %i0_0_0_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + + %ff0_0_u1 = vector.sitofp %i1_0_0_u1 : vector<8xi32> to vector<8xf32> + %sv0_0_u1 = vector.mulf %asv0_u0, %wsv0_u1 : vector<8xf32> + %fn0_0 = vector.fmaf %ff0_0_u1, %sv0_0_u1, %fn0_0_u0 : vector<8xf32> + scf.schedule.fence + %rfs0_1_0_row_u1 = index.mul %ln1, %k_wstride : index + %rfs0_1_0_base_u1 = index.add %rfs0_1_0_row_u1, %blk_k_u1 : index + %rfs0_1_0_idx_u1 = index.assume %rfs0_1_0_base_u1 [range(%rfs0_1_0_base_u1, 0, 5112)] : index + %rfs0_1_0_raw_u1 = vector.load %wl_flat[%rfs0_1_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs0_1_0_words_u1 = vector.bitcast %rfs0_1_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs0_1_0_u1 = vector.fragment %rfs0_1_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_1_1_row_u1 = index.mul %ln1, %k_wstride : index + %rfs0_1_1_base_u1 = index.add %rfs0_1_1_row_u1, %blk_k16_u1 : index + %rfs0_1_1_idx_u1 = index.assume %rfs0_1_1_base_u1 [range(%rfs0_1_1_base_u1, 0, 5112)] : index + %rfs0_1_1_raw_u1 = vector.load %wl_flat[%rfs0_1_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs0_1_1_words_u1 = vector.bitcast %rfs0_1_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs0_1_1_u1 = vector.fragment %rfs0_1_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_1_u1 = vector.mma %lf0_0_u1, %rfs0_1_0_u1, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_1_u1 = vector.mma %lf0_1_u1, %rfs0_1_1_u1, %i0_0_1_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + + %ff0_1_u1 = vector.sitofp %i1_0_1_u1 : vector<8xi32> to vector<8xf32> + %sv0_1_u1 = vector.mulf %asv0_u0, %wsv1_u1 : vector<8xf32> + %fn0_1 = vector.fmaf %ff0_1_u1, %sv0_1_u1, %fn0_1_u0 : vector<8xf32> + scf.schedule.fence + %rfs0_2_0_row_u1 = index.mul %ln4, %k_wstride : index + %rfs0_2_0_base_u1 = index.add %rfs0_2_0_row_u1, %blk_k_u1 : index + %rfs0_2_0_idx_u1 = index.assume %rfs0_2_0_base_u1 [range(%rfs0_2_0_base_u1, 0, 5112)] : index + %rfs0_2_0_raw_u1 = vector.load %wl_flat[%rfs0_2_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs0_2_0_words_u1 = vector.bitcast %rfs0_2_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs0_2_0_u1 = vector.fragment %rfs0_2_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_2_1_row_u1 = index.mul %ln4, %k_wstride : index + %rfs0_2_1_base_u1 = index.add %rfs0_2_1_row_u1, %blk_k16_u1 : index + %rfs0_2_1_idx_u1 = index.assume %rfs0_2_1_base_u1 [range(%rfs0_2_1_base_u1, 0, 5112)] : index + %rfs0_2_1_raw_u1 = vector.load %wl_flat[%rfs0_2_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs0_2_1_words_u1 = vector.bitcast %rfs0_2_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs0_2_1_u1 = vector.fragment %rfs0_2_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_2_u1 = vector.mma %lf0_0_u1, %rfs0_2_0_u1, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_2_u1 = vector.mma %lf0_1_u1, %rfs0_2_1_u1, %i0_0_2_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + + %ff0_2_u1 = vector.sitofp %i1_0_2_u1 : vector<8xi32> to vector<8xf32> + %sv0_2_u1 = vector.mulf %asv0_u0, %wsv4_u1 : vector<8xf32> + %fn0_2 = vector.fmaf %ff0_2_u1, %sv0_2_u1, %fn0_2_u0 : vector<8xf32> + scf.schedule.fence + %rfs0_3_0_row_u1 = index.mul %ln5, %k_wstride : index + %rfs0_3_0_base_u1 = index.add %rfs0_3_0_row_u1, %blk_k_u1 : index + %rfs0_3_0_idx_u1 = index.assume %rfs0_3_0_base_u1 [range(%rfs0_3_0_base_u1, 0, 5112)] : index + %rfs0_3_0_raw_u1 = vector.load %wl_flat[%rfs0_3_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs0_3_0_words_u1 = vector.bitcast %rfs0_3_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs0_3_0_u1 = vector.fragment %rfs0_3_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_3_1_row_u1 = index.mul %ln5, %k_wstride : index + %rfs0_3_1_base_u1 = index.add %rfs0_3_1_row_u1, %blk_k16_u1 : index + %rfs0_3_1_idx_u1 = index.assume %rfs0_3_1_base_u1 [range(%rfs0_3_1_base_u1, 0, 5112)] : index + %rfs0_3_1_raw_u1 = vector.load %wl_flat[%rfs0_3_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs0_3_1_words_u1 = vector.bitcast %rfs0_3_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs0_3_1_u1 = vector.fragment %rfs0_3_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_3_u1 = vector.mma %lf0_0_u1, %rfs0_3_0_u1, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_3_u1 = vector.mma %lf0_1_u1, %rfs0_3_1_u1, %i0_0_3_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + + %ff0_3_u1 = vector.sitofp %i1_0_3_u1 : vector<8xi32> to vector<8xf32> + %sv0_3_u1 = vector.mulf %asv0_u0, %wsv5_u1 : vector<8xf32> + %fn0_3 = vector.fmaf %ff0_3_u1, %sv0_3_u1, %fn0_3_u0 : vector<8xf32> + scf.schedule.fence + %wsi2_2_u1 = index.add %wsb2, %c1 : index + %wsi2_u1 = index.assume %wsi2_2_u1 [range(%wsi2_2_u1, 1, 255)] : index + %wsc2_u1 = view.load %wsl_view[%wsi2_u1] : view<256xf32> -> f32 + %wsv2_u1 = vector.splat %wsc2_u1 : vector<8xf32> + %wsi3_2_u1 = index.add %wsb3, %c1 : index + %wsi3_u1 = index.assume %wsi3_2_u1 [range(%wsi3_2_u1, 1, 255)] : index + %wsc3_u1 = view.load %wsl_view[%wsi3_u1] : view<256xf32> -> f32 + %wsv3_u1 = vector.splat %wsc3_u1 : vector<8xf32> + %wsi6_2_u1 = index.add %wsb6, %c1 : index + %wsi6_u1 = index.assume %wsi6_2_u1 [range(%wsi6_2_u1, 1, 255)] : index + %wsc6_u1 = view.load %wsl_view[%wsi6_u1] : view<256xf32> -> f32 + %wsv6_u1 = vector.splat %wsc6_u1 : vector<8xf32> + %wsi7_2_u1 = index.add %wsb7, %c1 : index + %wsi7_u1 = index.assume %wsi7_2_u1 [range(%wsi7_2_u1, 1, 255)] : index + %wsc7_u1 = view.load %wsl_view[%wsi7_u1] : view<256xf32> -> f32 + %wsv7_u1 = vector.splat %wsc7_u1 : vector<8xf32> + %rfs1_0_0_row_u1 = index.mul %ln2, %k_wstride : index + %rfs1_0_0_base_u1 = index.add %rfs1_0_0_row_u1, %blk_k_u1 : index + %rfs1_0_0_idx_u1 = index.assume %rfs1_0_0_base_u1 [range(%rfs1_0_0_base_u1, 0, 5112)] : index + %rfs1_0_0_raw_u1 = vector.load %wl_flat[%rfs1_0_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs1_0_0_words_u1 = vector.bitcast %rfs1_0_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs1_0_0_u1 = vector.fragment %rfs1_0_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs1_0_1_row_u1 = index.mul %ln2, %k_wstride : index + %rfs1_0_1_base_u1 = index.add %rfs1_0_1_row_u1, %blk_k16_u1 : index + %rfs1_0_1_idx_u1 = index.assume %rfs1_0_1_base_u1 [range(%rfs1_0_1_base_u1, 0, 5112)] : index + %rfs1_0_1_raw_u1 = vector.load %wl_flat[%rfs1_0_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs1_0_1_words_u1 = vector.bitcast %rfs1_0_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs1_0_1_u1 = vector.fragment %rfs1_0_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_1_0_u1 = vector.mma %lf0_0_u1, %rfs1_0_0_u1, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_1_0_u1 = vector.mma %lf0_1_u1, %rfs1_0_1_u1, %i0_1_0_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + + %ff1_0_u1 = vector.sitofp %i1_1_0_u1 : vector<8xi32> to vector<8xf32> + %sv1_0_u1 = vector.mulf %asv0_u0, %wsv2_u1 : vector<8xf32> + %fn1_0 = vector.fmaf %ff1_0_u1, %sv1_0_u1, %fn1_0_u0 : vector<8xf32> + scf.schedule.fence + %rfs1_1_0_row_u1 = index.mul %ln3, %k_wstride : index + %rfs1_1_0_base_u1 = index.add %rfs1_1_0_row_u1, %blk_k_u1 : index + %rfs1_1_0_idx_u1 = index.assume %rfs1_1_0_base_u1 [range(%rfs1_1_0_base_u1, 0, 5112)] : index + %rfs1_1_0_raw_u1 = vector.load %wl_flat[%rfs1_1_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs1_1_0_words_u1 = vector.bitcast %rfs1_1_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs1_1_0_u1 = vector.fragment %rfs1_1_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs1_1_1_row_u1 = index.mul %ln3, %k_wstride : index + %rfs1_1_1_base_u1 = index.add %rfs1_1_1_row_u1, %blk_k16_u1 : index + %rfs1_1_1_idx_u1 = index.assume %rfs1_1_1_base_u1 [range(%rfs1_1_1_base_u1, 0, 5112)] : index + %rfs1_1_1_raw_u1 = vector.load %wl_flat[%rfs1_1_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs1_1_1_words_u1 = vector.bitcast %rfs1_1_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs1_1_1_u1 = vector.fragment %rfs1_1_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_1_1_u1 = vector.mma %lf0_0_u1, %rfs1_1_0_u1, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_1_1_u1 = vector.mma %lf0_1_u1, %rfs1_1_1_u1, %i0_1_1_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + + %ff1_1_u1 = vector.sitofp %i1_1_1_u1 : vector<8xi32> to vector<8xf32> + %sv1_1_u1 = vector.mulf %asv0_u0, %wsv3_u1 : vector<8xf32> + %fn1_1 = vector.fmaf %ff1_1_u1, %sv1_1_u1, %fn1_1_u0 : vector<8xf32> + scf.schedule.fence + %rfs1_2_0_row_u1 = index.mul %ln6, %k_wstride : index + %rfs1_2_0_base_u1 = index.add %rfs1_2_0_row_u1, %blk_k_u1 : index + %rfs1_2_0_idx_u1 = index.assume %rfs1_2_0_base_u1 [range(%rfs1_2_0_base_u1, 0, 5112)] : index + %rfs1_2_0_raw_u1 = vector.load %wl_flat[%rfs1_2_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs1_2_0_words_u1 = vector.bitcast %rfs1_2_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs1_2_0_u1 = vector.fragment %rfs1_2_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs1_2_1_row_u1 = index.mul %ln6, %k_wstride : index + %rfs1_2_1_base_u1 = index.add %rfs1_2_1_row_u1, %blk_k16_u1 : index + %rfs1_2_1_idx_u1 = index.assume %rfs1_2_1_base_u1 [range(%rfs1_2_1_base_u1, 0, 5112)] : index + %rfs1_2_1_raw_u1 = vector.load %wl_flat[%rfs1_2_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs1_2_1_words_u1 = vector.bitcast %rfs1_2_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs1_2_1_u1 = vector.fragment %rfs1_2_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_1_2_u1 = vector.mma %lf0_0_u1, %rfs1_2_0_u1, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_1_2_u1 = vector.mma %lf0_1_u1, %rfs1_2_1_u1, %i0_1_2_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + + %ff1_2_u1 = vector.sitofp %i1_1_2_u1 : vector<8xi32> to vector<8xf32> + %sv1_2_u1 = vector.mulf %asv0_u0, %wsv6_u1 : vector<8xf32> + %fn1_2 = vector.fmaf %ff1_2_u1, %sv1_2_u1, %fn1_2_u0 : vector<8xf32> + scf.schedule.fence + %rfs1_3_0_row_u1 = index.mul %ln7, %k_wstride : index + %rfs1_3_0_base_u1 = index.add %rfs1_3_0_row_u1, %blk_k_u1 : index + %rfs1_3_0_idx_u1 = index.assume %rfs1_3_0_base_u1 [range(%rfs1_3_0_base_u1, 0, 5112)] : index + %rfs1_3_0_raw_u1 = vector.load %wl_flat[%rfs1_3_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs1_3_0_words_u1 = vector.bitcast %rfs1_3_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs1_3_0_u1 = vector.fragment %rfs1_3_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs1_3_1_row_u1 = index.mul %ln7, %k_wstride : index + %rfs1_3_1_base_u1 = index.add %rfs1_3_1_row_u1, %blk_k16_u1 : index + %rfs1_3_1_idx_u1 = index.assume %rfs1_3_1_base_u1 [range(%rfs1_3_1_base_u1, 0, 5112)] : index + %rfs1_3_1_raw_u1 = vector.load %wl_flat[%rfs1_3_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs1_3_1_words_u1 = vector.bitcast %rfs1_3_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs1_3_1_u1 = vector.fragment %rfs1_3_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_1_3_u1 = vector.mma %lf0_0_u1, %rfs1_3_0_u1, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_1_3_u1 = vector.mma %lf0_1_u1, %rfs1_3_1_u1, %i0_1_3_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + + %ff1_3_u1 = vector.sitofp %i1_1_3_u1 : vector<8xi32> to vector<8xf32> + %sv1_3_u1 = vector.mulf %asv0_u0, %wsv7_u1 : vector<8xf32> + %fn1_3 = vector.fmaf %ff1_3_u1, %sv1_3_u1, %fn1_3_u0 : vector<8xf32> + scf.schedule.fence + %lf2_0_lane_row_u1 = index.add %lm2, %lane_lo : index + %lf2_0_row_u1 = index.mul %lf2_0_lane_row_u1, %k_astride : index + %lf2_0_base_u1 = index.add %lf2_0_row_u1, %blk_k_u1 : index + %lf2_0_idx_u1 = index.assume %lf2_0_base_u1 [range(%lf2_0_base_u1, 0, 5112)] : index + %lf2_0_raw_u1 = vector.load %al_flat[%lf2_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %lf2_0_words_u1 = vector.bitcast %lf2_0_raw_u1 : vector<8xi8> to vector<2xi32> + %lf2_0_u1 = vector.fragment %lf2_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %lf2_1_lane_row_u1 = index.add %lm2, %lane_lo : index + %lf2_1_row_u1 = index.mul %lf2_1_lane_row_u1, %k_astride : index + %lf2_1_base_u1 = index.add %lf2_1_row_u1, %blk_k16_u1 : index + %lf2_1_idx_u1 = index.assume %lf2_1_base_u1 [range(%lf2_1_base_u1, 0, 5112)] : index + %lf2_1_raw_u1 = vector.load %al_flat[%lf2_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %lf2_1_words_u1 = vector.bitcast %lf2_1_raw_u1 : vector<8xi8> to vector<2xi32> + %lf2_1_u1 = vector.fragment %lf2_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs2_0_0_row_u1 = index.mul %ln0, %k_wstride : index + %rfs2_0_0_base_u1 = index.add %rfs2_0_0_row_u1, %blk_k_u1 : index + %rfs2_0_0_idx_u1 = index.assume %rfs2_0_0_base_u1 [range(%rfs2_0_0_base_u1, 0, 5112)] : index + %rfs2_0_0_raw_u1 = vector.load %wl_flat[%rfs2_0_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs2_0_0_words_u1 = vector.bitcast %rfs2_0_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs2_0_0_u1 = vector.fragment %rfs2_0_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs2_0_1_row_u1 = index.mul %ln0, %k_wstride : index + %rfs2_0_1_base_u1 = index.add %rfs2_0_1_row_u1, %blk_k16_u1 : index + %rfs2_0_1_idx_u1 = index.assume %rfs2_0_1_base_u1 [range(%rfs2_0_1_base_u1, 0, 5112)] : index + %rfs2_0_1_raw_u1 = vector.load %wl_flat[%rfs2_0_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs2_0_1_words_u1 = vector.bitcast %rfs2_0_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs2_0_1_u1 = vector.fragment %rfs2_0_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_2_0_u1 = vector.mma %lf2_0_u1, %rfs2_0_0_u1, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_2_0_u1 = vector.mma %lf2_1_u1, %rfs2_0_1_u1, %i0_2_0_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + + %ff2_0_u1 = vector.sitofp %i1_2_0_u1 : vector<8xi32> to vector<8xf32> + %sv2_0_u1 = vector.mulf %asv2_u0, %wsv0_u1 : vector<8xf32> + %fn2_0 = vector.fmaf %ff2_0_u1, %sv2_0_u1, %fn2_0_u0 : vector<8xf32> + scf.schedule.fence + %rfs2_1_0_row_u1 = index.mul %ln1, %k_wstride : index + %rfs2_1_0_base_u1 = index.add %rfs2_1_0_row_u1, %blk_k_u1 : index + %rfs2_1_0_idx_u1 = index.assume %rfs2_1_0_base_u1 [range(%rfs2_1_0_base_u1, 0, 5112)] : index + %rfs2_1_0_raw_u1 = vector.load %wl_flat[%rfs2_1_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs2_1_0_words_u1 = vector.bitcast %rfs2_1_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs2_1_0_u1 = vector.fragment %rfs2_1_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs2_1_1_row_u1 = index.mul %ln1, %k_wstride : index + %rfs2_1_1_base_u1 = index.add %rfs2_1_1_row_u1, %blk_k16_u1 : index + %rfs2_1_1_idx_u1 = index.assume %rfs2_1_1_base_u1 [range(%rfs2_1_1_base_u1, 0, 5112)] : index + %rfs2_1_1_raw_u1 = vector.load %wl_flat[%rfs2_1_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs2_1_1_words_u1 = vector.bitcast %rfs2_1_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs2_1_1_u1 = vector.fragment %rfs2_1_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_2_1_u1 = vector.mma %lf2_0_u1, %rfs2_1_0_u1, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_2_1_u1 = vector.mma %lf2_1_u1, %rfs2_1_1_u1, %i0_2_1_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + + %ff2_1_u1 = vector.sitofp %i1_2_1_u1 : vector<8xi32> to vector<8xf32> + %sv2_1_u1 = vector.mulf %asv2_u0, %wsv1_u1 : vector<8xf32> + %fn2_1 = vector.fmaf %ff2_1_u1, %sv2_1_u1, %fn2_1_u0 : vector<8xf32> + scf.schedule.fence + %rfs2_2_0_row_u1 = index.mul %ln4, %k_wstride : index + %rfs2_2_0_base_u1 = index.add %rfs2_2_0_row_u1, %blk_k_u1 : index + %rfs2_2_0_idx_u1 = index.assume %rfs2_2_0_base_u1 [range(%rfs2_2_0_base_u1, 0, 5112)] : index + %rfs2_2_0_raw_u1 = vector.load %wl_flat[%rfs2_2_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs2_2_0_words_u1 = vector.bitcast %rfs2_2_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs2_2_0_u1 = vector.fragment %rfs2_2_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs2_2_1_row_u1 = index.mul %ln4, %k_wstride : index + %rfs2_2_1_base_u1 = index.add %rfs2_2_1_row_u1, %blk_k16_u1 : index + %rfs2_2_1_idx_u1 = index.assume %rfs2_2_1_base_u1 [range(%rfs2_2_1_base_u1, 0, 5112)] : index + %rfs2_2_1_raw_u1 = vector.load %wl_flat[%rfs2_2_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs2_2_1_words_u1 = vector.bitcast %rfs2_2_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs2_2_1_u1 = vector.fragment %rfs2_2_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_2_2_u1 = vector.mma %lf2_0_u1, %rfs2_2_0_u1, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_2_2_u1 = vector.mma %lf2_1_u1, %rfs2_2_1_u1, %i0_2_2_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + + %ff2_2_u1 = vector.sitofp %i1_2_2_u1 : vector<8xi32> to vector<8xf32> + %sv2_2_u1 = vector.mulf %asv2_u0, %wsv4_u1 : vector<8xf32> + %fn2_2 = vector.fmaf %ff2_2_u1, %sv2_2_u1, %fn2_2_u0 : vector<8xf32> + scf.schedule.fence + %rfs2_3_0_row_u1 = index.mul %ln5, %k_wstride : index + %rfs2_3_0_base_u1 = index.add %rfs2_3_0_row_u1, %blk_k_u1 : index + %rfs2_3_0_idx_u1 = index.assume %rfs2_3_0_base_u1 [range(%rfs2_3_0_base_u1, 0, 5112)] : index + %rfs2_3_0_raw_u1 = vector.load %wl_flat[%rfs2_3_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs2_3_0_words_u1 = vector.bitcast %rfs2_3_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs2_3_0_u1 = vector.fragment %rfs2_3_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs2_3_1_row_u1 = index.mul %ln5, %k_wstride : index + %rfs2_3_1_base_u1 = index.add %rfs2_3_1_row_u1, %blk_k16_u1 : index + %rfs2_3_1_idx_u1 = index.assume %rfs2_3_1_base_u1 [range(%rfs2_3_1_base_u1, 0, 5112)] : index + %rfs2_3_1_raw_u1 = vector.load %wl_flat[%rfs2_3_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs2_3_1_words_u1 = vector.bitcast %rfs2_3_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs2_3_1_u1 = vector.fragment %rfs2_3_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_2_3_u1 = vector.mma %lf2_0_u1, %rfs2_3_0_u1, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_2_3_u1 = vector.mma %lf2_1_u1, %rfs2_3_1_u1, %i0_2_3_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + + %ff2_3_u1 = vector.sitofp %i1_2_3_u1 : vector<8xi32> to vector<8xf32> + %sv2_3_u1 = vector.mulf %asv2_u0, %wsv5_u1 : vector<8xf32> + %fn2_3 = vector.fmaf %ff2_3_u1, %sv2_3_u1, %fn2_3_u0 : vector<8xf32> + scf.schedule.fence + %rfs3_0_0_row_u1 = index.mul %ln2, %k_wstride : index + %rfs3_0_0_base_u1 = index.add %rfs3_0_0_row_u1, %blk_k_u1 : index + %rfs3_0_0_idx_u1 = index.assume %rfs3_0_0_base_u1 [range(%rfs3_0_0_base_u1, 0, 5112)] : index + %rfs3_0_0_raw_u1 = vector.load %wl_flat[%rfs3_0_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs3_0_0_words_u1 = vector.bitcast %rfs3_0_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs3_0_0_u1 = vector.fragment %rfs3_0_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs3_0_1_row_u1 = index.mul %ln2, %k_wstride : index + %rfs3_0_1_base_u1 = index.add %rfs3_0_1_row_u1, %blk_k16_u1 : index + %rfs3_0_1_idx_u1 = index.assume %rfs3_0_1_base_u1 [range(%rfs3_0_1_base_u1, 0, 5112)] : index + %rfs3_0_1_raw_u1 = vector.load %wl_flat[%rfs3_0_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs3_0_1_words_u1 = vector.bitcast %rfs3_0_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs3_0_1_u1 = vector.fragment %rfs3_0_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_3_0_u1 = vector.mma %lf2_0_u1, %rfs3_0_0_u1, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_3_0_u1 = vector.mma %lf2_1_u1, %rfs3_0_1_u1, %i0_3_0_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + + %ff3_0_u1 = vector.sitofp %i1_3_0_u1 : vector<8xi32> to vector<8xf32> + %sv3_0_u1 = vector.mulf %asv2_u0, %wsv2_u1 : vector<8xf32> + %fn3_0 = vector.fmaf %ff3_0_u1, %sv3_0_u1, %fn3_0_u0 : vector<8xf32> + scf.schedule.fence + %rfs3_1_0_row_u1 = index.mul %ln3, %k_wstride : index + %rfs3_1_0_base_u1 = index.add %rfs3_1_0_row_u1, %blk_k_u1 : index + %rfs3_1_0_idx_u1 = index.assume %rfs3_1_0_base_u1 [range(%rfs3_1_0_base_u1, 0, 5112)] : index + %rfs3_1_0_raw_u1 = vector.load %wl_flat[%rfs3_1_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs3_1_0_words_u1 = vector.bitcast %rfs3_1_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs3_1_0_u1 = vector.fragment %rfs3_1_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs3_1_1_row_u1 = index.mul %ln3, %k_wstride : index + %rfs3_1_1_base_u1 = index.add %rfs3_1_1_row_u1, %blk_k16_u1 : index + %rfs3_1_1_idx_u1 = index.assume %rfs3_1_1_base_u1 [range(%rfs3_1_1_base_u1, 0, 5112)] : index + %rfs3_1_1_raw_u1 = vector.load %wl_flat[%rfs3_1_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs3_1_1_words_u1 = vector.bitcast %rfs3_1_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs3_1_1_u1 = vector.fragment %rfs3_1_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_3_1_u1 = vector.mma %lf2_0_u1, %rfs3_1_0_u1, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_3_1_u1 = vector.mma %lf2_1_u1, %rfs3_1_1_u1, %i0_3_1_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + + %ff3_1_u1 = vector.sitofp %i1_3_1_u1 : vector<8xi32> to vector<8xf32> + %sv3_1_u1 = vector.mulf %asv2_u0, %wsv3_u1 : vector<8xf32> + %fn3_1 = vector.fmaf %ff3_1_u1, %sv3_1_u1, %fn3_1_u0 : vector<8xf32> + scf.schedule.fence + %rfs3_2_0_row_u1 = index.mul %ln6, %k_wstride : index + %rfs3_2_0_base_u1 = index.add %rfs3_2_0_row_u1, %blk_k_u1 : index + %rfs3_2_0_idx_u1 = index.assume %rfs3_2_0_base_u1 [range(%rfs3_2_0_base_u1, 0, 5112)] : index + %rfs3_2_0_raw_u1 = vector.load %wl_flat[%rfs3_2_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs3_2_0_words_u1 = vector.bitcast %rfs3_2_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs3_2_0_u1 = vector.fragment %rfs3_2_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs3_2_1_row_u1 = index.mul %ln6, %k_wstride : index + %rfs3_2_1_base_u1 = index.add %rfs3_2_1_row_u1, %blk_k16_u1 : index + %rfs3_2_1_idx_u1 = index.assume %rfs3_2_1_base_u1 [range(%rfs3_2_1_base_u1, 0, 5112)] : index + %rfs3_2_1_raw_u1 = vector.load %wl_flat[%rfs3_2_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs3_2_1_words_u1 = vector.bitcast %rfs3_2_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs3_2_1_u1 = vector.fragment %rfs3_2_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_3_2_u1 = vector.mma %lf2_0_u1, %rfs3_2_0_u1, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_3_2_u1 = vector.mma %lf2_1_u1, %rfs3_2_1_u1, %i0_3_2_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + + %ff3_2_u1 = vector.sitofp %i1_3_2_u1 : vector<8xi32> to vector<8xf32> + %sv3_2_u1 = vector.mulf %asv2_u0, %wsv6_u1 : vector<8xf32> + %fn3_2 = vector.fmaf %ff3_2_u1, %sv3_2_u1, %fn3_2_u0 : vector<8xf32> + scf.schedule.fence + %rfs3_3_0_row_u1 = index.mul %ln7, %k_wstride : index + %rfs3_3_0_base_u1 = index.add %rfs3_3_0_row_u1, %blk_k_u1 : index + %rfs3_3_0_idx_u1 = index.assume %rfs3_3_0_base_u1 [range(%rfs3_3_0_base_u1, 0, 5112)] : index + %rfs3_3_0_raw_u1 = vector.load %wl_flat[%rfs3_3_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs3_3_0_words_u1 = vector.bitcast %rfs3_3_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs3_3_0_u1 = vector.fragment %rfs3_3_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs3_3_1_row_u1 = index.mul %ln7, %k_wstride : index + %rfs3_3_1_base_u1 = index.add %rfs3_3_1_row_u1, %blk_k16_u1 : index + %rfs3_3_1_idx_u1 = index.assume %rfs3_3_1_base_u1 [range(%rfs3_3_1_base_u1, 0, 5112)] : index + %rfs3_3_1_raw_u1 = vector.load %wl_flat[%rfs3_3_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs3_3_1_words_u1 = vector.bitcast %rfs3_3_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs3_3_1_u1 = vector.fragment %rfs3_3_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_3_3_u1 = vector.mma %lf2_0_u1, %rfs3_3_0_u1, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_3_3_u1 = vector.mma %lf2_1_u1, %rfs3_3_1_u1, %i0_3_3_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + + %ff3_3_u1 = vector.sitofp %i1_3_3_u1 : vector<8xi32> to vector<8xf32> + %sv3_3_u1 = vector.mulf %asv2_u0, %wsv7_u1 : vector<8xf32> + %fn3_3 = vector.fmaf %ff3_3_u1, %sv3_3_u1, %fn3_3_u0 : vector<8xf32> + scf.schedule.fence + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %fn0_0, %fn0_1, %fn0_2, %fn0_3, %fn1_0, %fn1_1, %fn1_2, %fn1_3, %fn2_0, %fn2_1, %fn2_2, %fn2_3, %fn3_0, %fn3_1, %fn3_2, %fn3_3 : vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32> + } + + %gate_m0_n0 = vector.siluf %f0_0 : vector<8xf32> + %gate_m0_n1 = vector.siluf %f0_1 : vector<8xf32> + %gate_m1_n0 = vector.siluf %f1_0 : vector<8xf32> + %gate_m1_n1 = vector.siluf %f1_1 : vector<8xf32> + %out_m0_n0 = vector.mulf %gate_m0_n0, %f0_2 : vector<8xf32> + %out_m0_n1 = vector.mulf %gate_m0_n1, %f0_3 : vector<8xf32> + %out_m1_n0 = vector.mulf %gate_m1_n0, %f1_2 : vector<8xf32> + %out_m1_n1 = vector.mulf %gate_m1_n1, %f1_3 : vector<8xf32> + %gate_m2_n0 = vector.siluf %f2_0 : vector<8xf32> + %gate_m2_n1 = vector.siluf %f2_1 : vector<8xf32> + %gate_m3_n0 = vector.siluf %f3_0 : vector<8xf32> + %gate_m3_n1 = vector.siluf %f3_1 : vector<8xf32> + %out_m2_n0 = vector.mulf %gate_m2_n0, %f2_2 : vector<8xf32> + %out_m2_n1 = vector.mulf %gate_m2_n1, %f2_3 : vector<8xf32> + %out_m3_n0 = vector.mulf %gate_m3_n0, %f3_2 : vector<8xf32> + %out_m3_n1 = vector.mulf %gate_m3_n1, %f3_3 : vector<8xf32> + + scf.if %publish_f32 { + vector.fragment.store %out_m0_n0, %dst_view[%gm0, %gn0] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out_m0_n1, %dst_view[%gm0, %gn1] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out_m1_n0, %dst_view[%gm0, %gn2] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out_m1_n1, %dst_view[%gm0, %gn3] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out_m2_n0, %dst_view[%gm2, %gn0] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out_m2_n1, %dst_view[%gm2, %gn1] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out_m3_n0, %dst_view[%gm2, %gn2] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out_m3_n1, %dst_view[%gm2, %gn3] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + } + + // Materialize the Q6_K down input through LDS for contiguous F16 stores and matching conversion order. + scf.if %publish_f16 { + %f16_tile = buffer.view %qout_scratch[%base] : buffer -> view<128x32xf32> + %f16_c4 = index.constant 4 : index + %f16_c8 = index.constant 8 : index + %f16_c32 = index.constant 32 : index + %f16_c512 = index.constant 512 : index + vector.fragment.store %out_m0_n0, %f16_tile[%lm0, %c0] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out_m0_n1, %f16_tile[%lm0, %c16] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out_m1_n0, %f16_tile[%lm1, %c0] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out_m1_n1, %f16_tile[%lm1, %c16] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.for %f16_batch = [%c0 to %f16_c8 step %c1] { + %f16_batch_base = index.mul %f16_batch, %f16_c512 : index + %f16_thread_base = index.mul %tid, %f16_c4 : index + %f16_linear0 = index.add %f16_batch_base, %f16_thread_base : index + %f16_linear = index.assume %f16_linear0 [range(%f16_linear0, 0, 4092), mul(%f16_linear0, 4)] : index + %f16_token_local = index.div %f16_linear, %f16_c32 : index + %f16_row_local = index.rem %f16_linear, %f16_c32 : index + %f16_wide = vector.load %f16_tile[%f16_token_local, %f16_row_local] : view<128x32xf32> -> vector<4xf32> + %f16_narrow = vector.fptrunc %f16_wide : vector<4xf32> to vector<4xf16> + %f16_token0 = index.add %col_base, %f16_token_local : index + %f16_row0 = index.add %row_base, %f16_row_local : index + %f16_row_end0 = index.add %f16_row0, %f16_c4 : index + %f16_token, %f16_row, %f16_row_end, %f16_cols, %f16_rows = index.assume %f16_token0, %f16_row0, %f16_row_end0, %cols_b, %rows_b [lt(%f16_token0, %cols_b), le(%f16_row_end0, %rows_b)] : index, index, index, index, index + vector.store %f16_narrow, %dst_f16_view[%f16_token, %f16_row] : vector<4xf16>, view<[%cols_b]x[%rows_b]xf16> + } + } + + // Publish affine-U4 K32 blocks while the SwiGLU tile is live to avoid an F32 round trip and separate quantizer. + scf.if %publish_u4 { + %u4_tile = buffer.view %qout_scratch[%base] : buffer -> view<128x32xf32> + + %u4_c4 = index.constant 4 : index + %u4_c8 = index.constant 8 : index + %u4_c16 = index.constant 16 : index + %u4_c32 = index.constant 32 : index + %u4_zero_f32 = scalar.constant 0.0 : f32 + %u4_zero_i32 = scalar.constant 0 : i32 + %u4_eps = scalar.constant 1.0000000000000001e-30 : f32 + %u4_f15 = scalar.constant 15.0 : f32 + %u4_vzero = vector.constant 0.0 : vector<4xf32> + %u4_v15 = vector.constant 15.0 : vector<4xf32> + %u4_bias_neg8 = vector.constant -8 : vector<4xi32> + %u4_nibble_mask = vector.constant 15 : vector<4xi32> + %u4_c8_i32 = scalar.constant 8 : i32 + %u4_xor1 = scalar.constant 1 : i32 + %u4_xor2 = scalar.constant 2 : i32 + %u4_xor4 = scalar.constant 4 : i32 + %u4_shift4 = scalar.constant 4 : i32 + %u4_shift8 = scalar.constant 8 : i32 + %u4_shift12 = scalar.constant 12 : i32 + %u4_shuffle_width = scalar.constant 32 : i32 + + %u4_total = index.mul %rows_b, %cols_b : index + %u4_groups0 = index.div %u4_total, %u4_c32 : index + %u4_groups = index.assume %u4_groups0 [range(%u4_groups0, 1, 33554432)] : index + %u4_halfwords = index.div %u4_total, %u4_c4 : index + %u4_qs = buffer.view %qout_qs_na[%base] : buffer -> view<[%u4_halfwords]xi16> + %u4_ds = buffer.view %qout_ds_na[%base] : buffer -> view<[%u4_groups]xf32> + %u4_meta_count = index.mul %u4_groups, %c2 : index + %u4_meta = buffer.view %qout_sums_na[%base] : buffer -> view<[%u4_meta_count]xi32> + + %u4_group_in_wave = index.div %lane, %u4_c8 : index + %u4_lane_in_group = index.rem %lane, %u4_c8 : index + %u4_is_leader = index.cmp eq, %u4_lane_in_group, %c0 : index + %u4_wave_group_base = index.mul %wave, %c4 : index + %u4_local_group_base = index.add %u4_wave_group_base, %u4_group_in_wave : index + %u4_row_group = index.assume %row_tile [range(%row_tile, 0, 8191)] : index + %u4_row_groups = index.div %rows_b, %u4_c32 : index + + vector.fragment.store %out_m0_n0, %u4_tile[%lm0, %c0] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out_m0_n1, %u4_tile[%lm0, %c16] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out_m1_n0, %u4_tile[%lm1, %c0] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out_m1_n1, %u4_tile[%lm1, %c16] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + + scf.for %u4_batch = [%c0 to %u4_c8 step %c1] { + %u4_batch_group = index.mul %u4_batch, %u4_c16 : index + %u4_local_group0 = index.add %u4_batch_group, %u4_local_group_base : index + %u4_token_local = index.assume %u4_local_group0 [range(%u4_local_group0, 0, 127)] : index + %u4_lane_row0 = index.mul %u4_lane_in_group, %u4_c4 : index + %u4_lane_row = index.assume %u4_lane_row0 [range(%u4_lane_row0, 0, 28)] : index + %u4_values = vector.load %u4_tile[%u4_token_local, %u4_lane_row] : view<128x32xf32> -> vector<4xf32> + + %u4_lane_max = vector.reduce %u4_values, %u4_zero_f32 : vector<4xf32>, f32 + %u4_neg_values = vector.subf %u4_vzero, %u4_values : vector<4xf32> + %u4_lane_neg_min = vector.reduce %u4_neg_values, %u4_zero_f32 : vector<4xf32>, f32 + %u4_max1_peer, %u4_max1_valid = kernel.subgroup.shuffle %u4_lane_max, %u4_xor1, %u4_shuffle_width : f32, i32, i32 + %u4_max1 = scalar.maxnumf %u4_lane_max, %u4_max1_peer : f32 + %u4_max2_peer, %u4_max2_valid = kernel.subgroup.shuffle %u4_max1, %u4_xor2, %u4_shuffle_width : f32, i32, i32 + %u4_max2 = scalar.maxnumf %u4_max1, %u4_max2_peer : f32 + %u4_max4_peer, %u4_max4_valid = kernel.subgroup.shuffle %u4_max2, %u4_xor4, %u4_shuffle_width : f32, i32, i32 + %u4_group_max = scalar.maxnumf %u4_max2, %u4_max4_peer : f32 + %u4_min1_peer, %u4_min1_valid = kernel.subgroup.shuffle %u4_lane_neg_min, %u4_xor1, %u4_shuffle_width : f32, i32, i32 + %u4_min1 = scalar.maxnumf %u4_lane_neg_min, %u4_min1_peer : f32 + %u4_min2_peer, %u4_min2_valid = kernel.subgroup.shuffle %u4_min1, %u4_xor2, %u4_shuffle_width : f32, i32, i32 + %u4_min2 = scalar.maxnumf %u4_min1, %u4_min2_peer : f32 + %u4_min4_peer, %u4_min4_valid = kernel.subgroup.shuffle %u4_min2, %u4_xor4, %u4_shuffle_width : f32, i32, i32 + %u4_group_neg_min = scalar.maxnumf %u4_min2, %u4_min4_peer : f32 + %u4_range0 = scalar.addf %u4_group_max, %u4_group_neg_min : f32 + %u4_range = scalar.maxnumf %u4_range0, %u4_eps : f32 + %u4_scale = scalar.divf %u4_range, %u4_f15 : f32 + %u4_rscale = scalar.divf %u4_f15, %u4_range : f32 + %u4_zp_raw = scalar.mulf %u4_group_neg_min, %u4_rscale : f32 + %u4_zp_round = scalar.roundf %u4_zp_raw : f32 + %u4_zp_clamped = scalar.clampf %u4_zp_round, %u4_zero_f32, %u4_f15 : f32 + %u4_zp = scalar.fptosi %u4_zp_clamped : f32 to i32 + + %u4_token = index.add %col_base, %u4_token_local : index + %u4_gid_base = index.mul %u4_token, %u4_row_groups : index + %u4_gid0 = index.add %u4_gid_base, %u4_row_group : index + %u4_gid = index.assume %u4_gid0 [range(%u4_gid0, 0, 33554431)] : index + scf.if %u4_is_leader { + view.store %u4_scale, %u4_ds[%u4_gid] : f32, view<[%u4_groups]xf32> + } + + %u4_rs = vector.splat %u4_rscale : vector<4xf32> + %u4_zpv = vector.splat %u4_zp_clamped : vector<4xf32> + %u4_scaled = vector.mulf %u4_values, %u4_rs : vector<4xf32> + %u4_shifted = vector.addf %u4_scaled, %u4_zpv : vector<4xf32> + %u4_rounded = vector.roundf %u4_shifted : vector<4xf32> + %u4_nonnegative = vector.maxnumf %u4_rounded, %u4_vzero : vector<4xf32> + %u4_to_ceiling = vector.subf %u4_v15, %u4_nonnegative : vector<4xf32> + %u4_ceiling_nonnegative = vector.maxnumf %u4_to_ceiling, %u4_vzero : vector<4xf32> + %u4_clamped = vector.subf %u4_v15, %u4_ceiling_nonnegative : vector<4xf32> + %u4_q32_unsigned = vector.fptosi %u4_clamped : vector<4xf32> to vector<4xi32> + %u4_q32 = vector.addi %u4_q32_unsigned, %u4_bias_neg8 : vector<4xi32> + %u4_qbits = vector.andi %u4_q32, %u4_nibble_mask : vector<4xi32> + %u4_q0 = vector.extract %u4_qbits[0] : vector<4xi32> -> i32 + %u4_q1 = vector.extract %u4_qbits[1] : vector<4xi32> -> i32 + %u4_q2 = vector.extract %u4_qbits[2] : vector<4xi32> -> i32 + %u4_q3 = vector.extract %u4_qbits[3] : vector<4xi32> -> i32 + %u4_q1s = scalar.shli %u4_q1, %u4_shift4 : i32 + %u4_q2s = scalar.shli %u4_q2, %u4_shift8 : i32 + %u4_q3s = scalar.shli %u4_q3, %u4_shift12 : i32 + %u4_q01 = scalar.ori %u4_q0, %u4_q1s : i32 + %u4_q23 = scalar.ori %u4_q2s, %u4_q3s : i32 + %u4_packed32 = scalar.ori %u4_q01, %u4_q23 : i32 + %u4_packed = scalar.trunci %u4_packed32 : i32 to i16 + %u4_word_base = index.mul %u4_gid, %u4_c8 : index + %u4_word_index0 = index.add %u4_word_base, %u4_lane_in_group : index + %u4_word_index = index.assume %u4_word_index0 [range(%u4_word_index0, 0, 268435455)] : index + view.store %u4_packed, %u4_qs[%u4_word_index] : i16, view<[%u4_halfwords]xi16> + + %u4_lane_sum = vector.reduce %u4_q32, %u4_zero_i32 : vector<4xi32>, i32 + %u4_sum1_peer, %u4_sum1_valid = kernel.subgroup.shuffle %u4_lane_sum, %u4_xor1, %u4_shuffle_width : i32, i32, i32 + %u4_sum1 = scalar.addi %u4_lane_sum, %u4_sum1_peer : i32 + %u4_sum2_peer, %u4_sum2_valid = kernel.subgroup.shuffle %u4_sum1, %u4_xor2, %u4_shuffle_width : i32, i32, i32 + %u4_sum2 = scalar.addi %u4_sum1, %u4_sum2_peer : i32 + %u4_sum4_peer, %u4_sum4_valid = kernel.subgroup.shuffle %u4_sum2, %u4_xor4, %u4_shuffle_width : i32, i32, i32 + %u4_group_sum = scalar.addi %u4_sum2, %u4_sum4_peer : i32 + scf.if %u4_is_leader { + %u4_zp_signed = scalar.subi %u4_zp, %u4_c8_i32 : i32 + %u4_meta0 = index.mul %u4_gid, %c2 : index + %u4_meta1 = index.add %u4_meta0, %c1 : index + view.store %u4_group_sum, %u4_meta[%u4_meta0] : i32, view<[%u4_meta_count]xi32> + view.store %u4_zp_signed, %u4_meta[%u4_meta1] : i32, view<[%u4_meta_count]xi32> + } + } + } + + // Publish the M128xN64 SwiGLU tile as two block-major Q8 K32 planes. + scf.if %publish_q8_plane { + %q8_tile = buffer.view %qout_scratch[%base] : buffer -> view<64x64xf32> + %q8_c4 = index.constant 4 : index + %q8_c8 = index.constant 8 : index + %q8_c16 = index.constant 16 : index + %q8_c32 = index.constant 32 : index + %q8_c64 = index.constant 64 : index + %q8_zero_f32 = scalar.constant 0.0 : f32 + %q8_one_f32 = scalar.constant 1.0 : f32 + %q8_f127 = scalar.constant 127.0 : f32 + %q8_xor1 = scalar.constant 1 : i32 + %q8_xor2 = scalar.constant 2 : i32 + %q8_xor4 = scalar.constant 4 : i32 + %q8_shuffle_width = scalar.constant 32 : i32 + + %q8_group_in_wave = index.div %lane, %q8_c8 : index + %q8_lane_in_block = index.rem %lane, %q8_c8 : index + %q8_is_leader = index.cmp eq, %q8_lane_in_block, %c0 : index + %q8_wave_group_base = index.mul %wave, %c4 : index + %q8_local_group_base = index.add %q8_wave_group_base, %q8_group_in_wave : index + %q8_lane_row_base = index.mul %q8_lane_in_block, %q8_c4 : index + %q8_payload_bytes = index.mul %rows_b, %cols_b : index + %q8_payload_words = index.div %q8_payload_bytes, %q8_c4 : index + %q8_payload_end = index.cast %q8_payload_bytes : index to offset + %q8_block_count = index.div %rows_b, %q8_c32 : index + %q8_metadata_count = index.mul %q8_block_count, %cols_b : index + %q8_payload = buffer.view %qout_qs_na[%base] : buffer -> view<[%q8_payload_words]xi32> + %q8_metadata = buffer.view %qout_qs_na[%q8_payload_end] : buffer -> view<[%q8_metadata_count]xi32> + + vector.fragment.store %out_m0_n0, %q8_tile[%lm0, %c0] shape [%c16, %c16] : vector<8xf32>, view<64x64xf32> + vector.fragment.store %out_m0_n1, %q8_tile[%lm0, %c16] shape [%c16, %c16] : vector<8xf32>, view<64x64xf32> + vector.fragment.store %out_m1_n0, %q8_tile[%lm0, %c32] shape [%c16, %c16] : vector<8xf32>, view<64x64xf32> + %q8_c48 = index.constant 48 : index + vector.fragment.store %out_m1_n1, %q8_tile[%lm0, %q8_c48] shape [%c16, %c16] : vector<8xf32>, view<64x64xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + + scf.for %q8_batch = [%c0 to %q8_c8 step %c1] { + %q8_batch_group = index.mul %q8_batch, %q8_c16 : index + %q8_local_group0 = index.add %q8_batch_group, %q8_local_group_base : index + %q8_local_group = index.assume %q8_local_group0 [range(%q8_local_group0, 0, 127)] : index + %q8_token_local0 = index.div %q8_local_group, %c2 : index + %q8_token_local = index.assume %q8_token_local0 [range(%q8_token_local0, 0, 63)] : index + %q8_row_block0 = index.rem %q8_local_group, %c2 : index + %q8_row_block = index.assume %q8_row_block0 [range(%q8_row_block0, 0, 1)] : index + %q8_row_block_base = index.mul %q8_row_block, %q8_c32 : index + %q8_lane_row0 = index.add %q8_row_block_base, %q8_lane_row_base : index + %q8_lane_row = index.assume %q8_lane_row0 [range(%q8_lane_row0, 0, 60), mul(%q8_lane_row0, 4)] : index + %q8_values = vector.load %q8_tile[%q8_token_local, %q8_lane_row] : view<64x64xf32> -> vector<4xf32> + %q8_absolute_values = vector.absf %q8_values : vector<4xf32> + %q8_lane_max = vector.reduce %q8_absolute_values, %q8_zero_f32 : vector<4xf32>, f32 + %q8_max1_peer, %q8_max1_valid = kernel.subgroup.shuffle %q8_lane_max, %q8_xor1, %q8_shuffle_width : f32, i32, i32 + %q8_max1 = scalar.maxnumf %q8_lane_max, %q8_max1_peer : f32 + %q8_max2_peer, %q8_max2_valid = kernel.subgroup.shuffle %q8_max1, %q8_xor2, %q8_shuffle_width : f32, i32, i32 + %q8_max2 = scalar.maxnumf %q8_max1, %q8_max2_peer : f32 + %q8_max4_peer, %q8_max4_valid = kernel.subgroup.shuffle %q8_max2, %q8_xor4, %q8_shuffle_width : f32, i32, i32 + %q8_amax = scalar.maxnumf %q8_max2, %q8_max4_peer : f32 + %q8_d = scalar.divf %q8_amax, %q8_f127 : f32 + %q8_d_nonzero = scalar.cmpf one, %q8_d, %q8_zero_f32 : f32 + %q8_d_inverse = scf.if %q8_d_nonzero -> (f32) { + %q8_inverse = scalar.divf %q8_one_f32, %q8_d : f32 + scf.yield %q8_inverse : f32 + } else { + scf.yield %q8_zero_f32 : f32 + } + %q8_d_inverse_vector = vector.splat %q8_d_inverse : vector<4xf32> + %q8_scaled_values = vector.mulf %q8_values, %q8_d_inverse_vector : vector<4xf32> + %q8_rounded_values = vector.roundf %q8_scaled_values : vector<4xf32> + %q8_quantized_values = vector.fptosi %q8_rounded_values : vector<4xf32> to vector<4xi8> + %q8_packed_word = vector.bitcast %q8_quantized_values : vector<4xi8> to vector<1xi32> + + %q8_token = index.add %col_base, %q8_token_local : index + %q8_global_block_base = index.mul %row_tile, %c2 : index + %q8_global_block0 = index.add %q8_global_block_base, %q8_row_block : index + %q8_global_block = index.assume %q8_global_block0 [range(%q8_global_block0, 0, 16383)] : index + %q8_block_token_base = index.mul %q8_global_block, %cols_b : index + %q8_block_token0 = index.add %q8_block_token_base, %q8_token : index + %q8_block_token = index.assume %q8_block_token0 [range(%q8_block_token0, 0, 268435455)] : index + %q8_word_base = index.mul %q8_block_token, %q8_c8 : index + %q8_word_index0 = index.add %q8_word_base, %q8_lane_in_block : index + %q8_word_index = index.assume %q8_word_index0 [range(%q8_word_index0, 0, 2147483647)] : index + vector.store %q8_packed_word, %q8_payload[%q8_word_index] : vector<1xi32>, view<[%q8_payload_words]xi32> + + %q8_lane_sum = vector.reduce %q8_rounded_values, %q8_zero_f32 : vector<4xf32>, f32 + %q8_sum1_peer, %q8_sum1_valid = kernel.subgroup.shuffle %q8_lane_sum, %q8_xor1, %q8_shuffle_width : f32, i32, i32 + %q8_sum1 = scalar.addf %q8_lane_sum, %q8_sum1_peer : f32 + %q8_sum2_peer, %q8_sum2_valid = kernel.subgroup.shuffle %q8_sum1, %q8_xor2, %q8_shuffle_width : f32, i32, i32 + %q8_sum2 = scalar.addf %q8_sum1, %q8_sum2_peer : f32 + %q8_sum4_peer, %q8_sum4_valid = kernel.subgroup.shuffle %q8_sum2, %q8_xor4, %q8_shuffle_width : f32, i32, i32 + %q8_quantized_sum = scalar.addf %q8_sum2, %q8_sum4_peer : f32 + scf.if %q8_is_leader { + %q8_s = scalar.mulf %q8_quantized_sum, %q8_d : f32 + %q8_d_f16 = scalar.fptrunc %q8_d : f32 to f16 + %q8_s_f16 = scalar.fptrunc %q8_s : f32 to f16 + %q8_ds_pair = vector.from_elements %q8_d_f16, %q8_s_f16 : vector<2xf16> + %q8_ds_word = vector.bitcast %q8_ds_pair : vector<2xf16> to vector<1xi32> + vector.store %q8_ds_word, %q8_metadata[%q8_block_token] : vector<1xi32>, view<[%q8_metadata_count]xi32> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + vector.fragment.store %out_m2_n0, %q8_tile[%lm0, %c0] shape [%c16, %c16] : vector<8xf32>, view<64x64xf32> + vector.fragment.store %out_m2_n1, %q8_tile[%lm0, %c16] shape [%c16, %c16] : vector<8xf32>, view<64x64xf32> + vector.fragment.store %out_m3_n0, %q8_tile[%lm0, %q8_c32] shape [%c16, %c16] : vector<8xf32>, view<64x64xf32> + vector.fragment.store %out_m3_n1, %q8_tile[%lm0, %q8_c48] shape [%c16, %c16] : vector<8xf32>, view<64x64xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.for %q8b_batch = [%c0 to %q8_c8 step %c1] { + %q8b_batch_group = index.mul %q8b_batch, %q8_c16 : index + %q8b_local_group0 = index.add %q8b_batch_group, %q8_local_group_base : index + %q8b_local_group = index.assume %q8b_local_group0 [range(%q8b_local_group0, 0, 127)] : index + %q8b_token_local0 = index.div %q8b_local_group, %c2 : index + %q8b_token_local = index.assume %q8b_token_local0 [range(%q8b_token_local0, 0, 63)] : index + %q8b_row_block0 = index.rem %q8b_local_group, %c2 : index + %q8b_row_block = index.assume %q8b_row_block0 [range(%q8b_row_block0, 0, 1)] : index + %q8b_row_block_base = index.mul %q8b_row_block, %q8_c32 : index + %q8b_lane_row0 = index.add %q8b_row_block_base, %q8_lane_row_base : index + %q8b_lane_row = index.assume %q8b_lane_row0 [range(%q8b_lane_row0, 0, 60), mul(%q8b_lane_row0, 4)] : index + %q8b_values = vector.load %q8_tile[%q8b_token_local, %q8b_lane_row] : view<64x64xf32> -> vector<4xf32> + %q8b_absolute_values = vector.absf %q8b_values : vector<4xf32> + %q8b_lane_max = vector.reduce %q8b_absolute_values, %q8_zero_f32 : vector<4xf32>, f32 + %q8b_max1_peer, %q8b_max1_valid = kernel.subgroup.shuffle %q8b_lane_max, %q8_xor1, %q8_shuffle_width : f32, i32, i32 + %q8b_max1 = scalar.maxnumf %q8b_lane_max, %q8b_max1_peer : f32 + %q8b_max2_peer, %q8b_max2_valid = kernel.subgroup.shuffle %q8b_max1, %q8_xor2, %q8_shuffle_width : f32, i32, i32 + %q8b_max2 = scalar.maxnumf %q8b_max1, %q8b_max2_peer : f32 + %q8b_max4_peer, %q8b_max4_valid = kernel.subgroup.shuffle %q8b_max2, %q8_xor4, %q8_shuffle_width : f32, i32, i32 + %q8b_amax = scalar.maxnumf %q8b_max2, %q8b_max4_peer : f32 + %q8b_d = scalar.divf %q8b_amax, %q8_f127 : f32 + %q8b_d_nonzero = scalar.cmpf one, %q8b_d, %q8_zero_f32 : f32 + %q8b_d_inverse = scf.if %q8b_d_nonzero -> (f32) { + %q8b_inverse = scalar.divf %q8_one_f32, %q8b_d : f32 + scf.yield %q8b_inverse : f32 + } else { + scf.yield %q8_zero_f32 : f32 + } + %q8b_d_inverse_vector = vector.splat %q8b_d_inverse : vector<4xf32> + %q8b_scaled_values = vector.mulf %q8b_values, %q8b_d_inverse_vector : vector<4xf32> + %q8b_rounded_values = vector.roundf %q8b_scaled_values : vector<4xf32> + %q8b_quantized_values = vector.fptosi %q8b_rounded_values : vector<4xf32> to vector<4xi8> + %q8b_packed_word = vector.bitcast %q8b_quantized_values : vector<4xi8> to vector<1xi32> + + %q8b_token_within_tile = index.add %q8b_token_local, %q8_c64 : index + %q8b_token = index.add %col_base, %q8b_token_within_tile : index + %q8b_global_block_base = index.mul %row_tile, %c2 : index + %q8b_global_block0 = index.add %q8b_global_block_base, %q8b_row_block : index + %q8b_global_block = index.assume %q8b_global_block0 [range(%q8b_global_block0, 0, 16383)] : index + %q8b_block_token_base = index.mul %q8b_global_block, %cols_b : index + %q8b_block_token0 = index.add %q8b_block_token_base, %q8b_token : index + %q8b_block_token = index.assume %q8b_block_token0 [range(%q8b_block_token0, 0, 268435455)] : index + %q8b_word_base = index.mul %q8b_block_token, %q8_c8 : index + %q8b_word_index0 = index.add %q8b_word_base, %q8_lane_in_block : index + %q8b_word_index = index.assume %q8b_word_index0 [range(%q8b_word_index0, 0, 2147483647)] : index + vector.store %q8b_packed_word, %q8_payload[%q8b_word_index] : vector<1xi32>, view<[%q8_payload_words]xi32> + + %q8b_lane_sum = vector.reduce %q8b_rounded_values, %q8_zero_f32 : vector<4xf32>, f32 + %q8b_sum1_peer, %q8b_sum1_valid = kernel.subgroup.shuffle %q8b_lane_sum, %q8_xor1, %q8_shuffle_width : f32, i32, i32 + %q8b_sum1 = scalar.addf %q8b_lane_sum, %q8b_sum1_peer : f32 + %q8b_sum2_peer, %q8b_sum2_valid = kernel.subgroup.shuffle %q8b_sum1, %q8_xor2, %q8_shuffle_width : f32, i32, i32 + %q8b_sum2 = scalar.addf %q8b_sum1, %q8b_sum2_peer : f32 + %q8b_sum4_peer, %q8b_sum4_valid = kernel.subgroup.shuffle %q8b_sum2, %q8_xor4, %q8_shuffle_width : f32, i32, i32 + %q8b_quantized_sum = scalar.addf %q8b_sum2, %q8b_sum4_peer : f32 + scf.if %q8_is_leader { + %q8b_s = scalar.mulf %q8b_quantized_sum, %q8b_d : f32 + %q8b_d_f16 = scalar.fptrunc %q8b_d : f32 to f16 + %q8b_s_f16 = scalar.fptrunc %q8b_s : f32 to f16 + %q8b_ds_pair = vector.from_elements %q8b_d_f16, %q8b_s_f16 : vector<2xf16> + %q8b_ds_word = vector.bitcast %q8b_ds_pair : vector<2xf16> to vector<1xi32> + vector.store %q8b_ds_word, %q8_metadata[%q8b_block_token] : vector<1xi32>, view<[%q8_metadata_count]xi32> + } + } + } + template.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_symmetric_i4_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_symmetric_i4_wmma.loom new file mode 100644 index 000000000000..af363bf16e04 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_symmetric_i4_wmma.loom @@ -0,0 +1,3451 @@ +template.decl @ggml.quantize_symmetric_i4_k64.body(%plane_major: i1, %src: buffer, %qs: buffer, %ds: buffer, %sums: buffer) + +config.decl @ggml.quantize_symmetric_i4_k64.input_size : %value: index where [range(%value, 32, 65536), mul(%value, 32)] + +config.decl @ggml.quantize_symmetric_i4_k64.token_count : %value: index where [range(%value, 1, 16384)] + +kernel.def export("ggml_quantize_f32_symmetric_i4_k64_plane") @ggml_quantize_f32_symmetric_i4_k64_plane() { + %unit = index.constant 1 : index + %wg = index.constant 256 : index + %wg_groups_m1 = index.constant 127 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %k = config.get @ggml.quantize_symmetric_i4_k64.input_size : index + %cols = config.get @ggml.quantize_symmetric_i4_k64.token_count : index + %total = index.mul %k, %cols : index + %groups = index.div %total, %c64 : index + %rounded = index.add %groups, %wg_groups_m1 : index + %workgroups = index.div %rounded, %c128 : index + kernel.launch.config workgroups(%workgroups, %unit, %unit) workgroup_size(%wg, %unit, %unit) : index +} launch(%src: buffer, %qs: buffer, %ds: buffer, %sums: buffer) { + %plane_major = scalar.constant true : i1 + template.apply<@ggml.quantize_symmetric_i4_k64.body>(%plane_major, %src, %qs, %ds, %sums) : (i1, buffer, buffer, buffer, buffer) + kernel.return +} + +// Exact low-column direct-dot specializations for adjacent projections. +func.def inline @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_half_dot(%weight_words: vector<2xi32>, %group: index, %token: index, %half: index, %k: index, %group_count: index, %qact: buffer, %scales: buffer) -> (i32, f32) { + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %zero_i32 = scalar.constant 0 : i32 + %zero_i32x2 = vector.constant 0 : vector<2xi32> + %qact_flat = buffer.view %qact[%base] : buffer -> view<1073741824xi8> + %scale_flat = buffer.view %scales[%base] : buffer -> view<33554432xf32> + + %token_bytes = index.div %k, %c2 : index + %token_payload_base = index.mul %token, %token_bytes : index + %group_payload_add = index.mul %group, %c16 : index + %half_payload_add = index.mul %half, %c8 : index + %payload_index0 = index.add %token_payload_base, %group_payload_add : index + %payload_index1 = index.add %payload_index0, %half_payload_add : index + %payload_index = index.assume %payload_index1 [range(%payload_index1, 0, 1073741816)] : index + %activation_bytes = vector.load %qact_flat[%payload_index] : view<1073741824xi8> -> vector<8xi8> + %activation_words = vector.bitcast %activation_bytes : vector<8xi8> to vector<2xi32> + %dot_parts = vector.dot8i4 %weight_words, %activation_words, %zero_i32x2 : vector<2xi32> + %dot = vector.reduce %dot_parts, %zero_i32 : vector<2xi32>, i32 + + %token_scale_base = index.mul %token, %group_count : index + %scale_index0 = index.add %token_scale_base, %group : index + %scale_index = index.assume %scale_index0 [range(%scale_index0, 0, 33554431)] : index + %activation_scale = view.load %scale_flat[%scale_index] : view<33554432xf32> -> f32 + func.return %dot, %activation_scale : i32, f32 +} + +kernel.def export("ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c1") @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c1() { + %rows = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c15 = index.constant 15 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %rows_up = index.add %rows, %c15 : index + %row_workgroups = index.div %rows_up, %c16 : index + %workgroups = index.mul %row_workgroups, %c2 : index + kernel.launch.config workgroups(%workgroups, %c1, %c1) workgroup_size(%c32, %c1, %c1) : index +} launch(%gate_weight: buffer, %up_weight: buffer, %gate_output: buffer, %up_output: buffer, %qact: buffer, %scales: buffer, %sums: buffer) { + %k0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.input_size : index + %rows0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %cols0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.token_count : index + %k = index.assume %k0 [range(%k0, 256, 32768), mul(%k0, 256)] : index + %rows = index.assume %rows0 [range(%rows0, 1, 262144)] : index + %cols = index.assume %cols0 [range(%cols0, 1, 1)] : index + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %zero_f32 = scalar.constant 0.0 : f32 + %xor_one = scalar.constant 1 : i32 + %shuffle_width = scalar.constant 32 : i32 + + %workgroup0 = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %projection = index.rem %workgroup0, %c2 : index + %workgroup = index.div %workgroup0, %c2 : index + %is_up = index.cmp eq, %projection, %c1 : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 31)] : index + %row_in_wave = index.div %workitem, %c2 : index + %half = index.rem %workitem, %c2 : index + %row_base = index.mul %workgroup, %c16 : index + %row = index.add %row_base, %row_in_wave : index + %valid_row = index.cmp ult, %row, %rows : index + + %block_count = index.div %k, %c256 : index + %group_count = index.div %k, %c32 : index + %row_group0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.row_group_size : index + %rows_plus = index.add %rows, %c255 : index + %row_group = index.assume %row_group0 [range(%row_group0, 32, 262144), mul(%row_group0, 32)] : index + %row_group_id = index.div %row, %row_group : index + %row_lane = index.rem %row, %row_group : index + %cohort_lane = index.div %row_lane, %c4 : index + %payload_plane_bytes = index.mul %row_group, %c16 : index + %payload_block_bytes = index.mul %payload_plane_bytes, %c8 : index + %scale_cohorts = index.div %row_group, %c4 : index + %scale_plane_bytes = index.mul %scale_cohorts, %c16 : index + %block_bytes = index.add %payload_block_bytes, %scale_plane_bytes : index + %row_group_bytes = index.mul %block_count, %block_bytes : index + %group_base = index.mul %row_group_id, %row_group_bytes : index + %scale_region_add = index.mul %block_count, %payload_block_bytes : index + %scale_region = index.add %group_base, %scale_region_add : index + + %gate_weight_global = buffer.assume.memory_space %gate_weight : buffer + %up_weight_global = buffer.assume.memory_space %up_weight : buffer + %gate_output_global = buffer.assume.memory_space %gate_output : buffer + %up_output_global = buffer.assume.memory_space %up_output : buffer + %qact_global = buffer.assume.memory_space %qact : buffer + %scales_global = buffer.assume.memory_space %scales : buffer + %weight_global = scf.select %is_up, %up_weight_global, %gate_weight_global : buffer + %output_global = scf.select %is_up, %up_output_global, %gate_output_global : buffer + %weight_na, %output_na, %qact_na, %scales_na = buffer.assume.noalias %weight_global, %output_global, %qact_global, %scales_global : buffer, buffer, buffer, buffer + + %result0 = scf.if %valid_row -> (f32) { + %acc0 = scf.for %group = [%c0 to %group_count step %c1](%iter0 = %zero_f32 : f32) -> (f32) { + %block = index.div %group, %c8 : index + %group_in_block = index.rem %group, %c8 : index + %block_payload_add = index.mul %block, %payload_block_bytes : index + %group_payload_add = index.mul %group_in_block, %payload_plane_bytes : index + %row_payload_add = index.mul %row_lane, %c16 : index + %half_payload_add = index.mul %half, %c8 : index + %payload_index0 = index.add %group_base, %block_payload_add : index + %payload_index1 = index.add %payload_index0, %group_payload_add : index + %payload_index2 = index.add %payload_index1, %row_payload_add : index + %payload_index = index.add %payload_index2, %half_payload_add : index + %payload_base = index.cast %payload_index : index to offset + %weight_view = buffer.view %weight_na[%payload_base] : buffer -> view<2xi32> + %weight_words = vector.load %weight_view[%c0] : view<2xi32> -> vector<2xi32> + + %block_scale_add = index.mul %block, %scale_plane_bytes : index + %cohort_scale_add = index.mul %cohort_lane, %c16 : index + %group_scale_add = index.mul %group_in_block, %c2 : index + %scale_index0 = index.add %scale_region, %block_scale_add : index + %scale_index1 = index.add %scale_index0, %cohort_scale_add : index + %scale_byte_index = index.add %scale_index1, %group_scale_add : index + %scale_byte_base = index.cast %scale_byte_index : index to offset + %weight_scale_view = buffer.view %weight_na[%scale_byte_base] : buffer -> view<1xf16> + %weight_scale_f16 = view.load %weight_scale_view[%c0] : view<1xf16> -> f16 + %weight_scale = scalar.extf %weight_scale_f16 : f16 to f32 + + %dot0, %as0 = func.call @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_half_dot(%weight_words, %group, %c0, %half, %k, %group_count, %qact_na, %scales_na) : (vector<2xi32>, index, index, index, index, index, buffer, buffer) -> (i32, f32) + %peer0, %valid0 = kernel.subgroup.shuffle %dot0, %xor_one, %shuffle_width : i32, i32, i32 + %full0 = scalar.addi %dot0, %peer0 : i32 + %fp0 = scalar.sitofp %full0 : i32 to f32 + %scale0 = scalar.mulf %weight_scale, %as0 : f32 + %next0 = scalar.fmaf %fp0, %scale0, %iter0 : f32 + scf.yield %next0 : f32 + } + scf.yield %acc0 : f32 + } else { + scf.yield %zero_f32 : f32 + } + + %is_writer = index.cmp eq, %half, %c0 : index + %writes = scalar.andi %valid_row, %is_writer : i1 + scf.if %writes { + %output_view = buffer.view %output_na[%base] : buffer -> view<[%cols]x[%rows]xf32> + view.store %result0, %output_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + } + kernel.return +} + +kernel.def export("ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c2") @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c2() { + %rows = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c15 = index.constant 15 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %rows_up = index.add %rows, %c15 : index + %row_workgroups = index.div %rows_up, %c16 : index + %workgroups = index.mul %row_workgroups, %c2 : index + kernel.launch.config workgroups(%workgroups, %c1, %c1) workgroup_size(%c32, %c1, %c1) : index +} launch(%gate_weight: buffer, %up_weight: buffer, %gate_output: buffer, %up_output: buffer, %qact: buffer, %scales: buffer, %sums: buffer) { + %k0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.input_size : index + %rows0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %cols0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.token_count : index + %k = index.assume %k0 [range(%k0, 256, 32768), mul(%k0, 256)] : index + %rows = index.assume %rows0 [range(%rows0, 1, 262144)] : index + %cols = index.assume %cols0 [range(%cols0, 2, 2)] : index + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %zero_f32 = scalar.constant 0.0 : f32 + %xor_one = scalar.constant 1 : i32 + %shuffle_width = scalar.constant 32 : i32 + + %workgroup0 = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %projection = index.rem %workgroup0, %c2 : index + %workgroup = index.div %workgroup0, %c2 : index + %is_up = index.cmp eq, %projection, %c1 : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 31)] : index + %row_in_wave = index.div %workitem, %c2 : index + %half = index.rem %workitem, %c2 : index + %row_base = index.mul %workgroup, %c16 : index + %row = index.add %row_base, %row_in_wave : index + %valid_row = index.cmp ult, %row, %rows : index + + %block_count = index.div %k, %c256 : index + %group_count = index.div %k, %c32 : index + %row_group0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.row_group_size : index + %rows_plus = index.add %rows, %c255 : index + %row_group = index.assume %row_group0 [range(%row_group0, 32, 262144), mul(%row_group0, 32)] : index + %row_group_id = index.div %row, %row_group : index + %row_lane = index.rem %row, %row_group : index + %cohort_lane = index.div %row_lane, %c4 : index + %payload_plane_bytes = index.mul %row_group, %c16 : index + %payload_block_bytes = index.mul %payload_plane_bytes, %c8 : index + %scale_cohorts = index.div %row_group, %c4 : index + %scale_plane_bytes = index.mul %scale_cohorts, %c16 : index + %block_bytes = index.add %payload_block_bytes, %scale_plane_bytes : index + %row_group_bytes = index.mul %block_count, %block_bytes : index + %group_base = index.mul %row_group_id, %row_group_bytes : index + %scale_region_add = index.mul %block_count, %payload_block_bytes : index + %scale_region = index.add %group_base, %scale_region_add : index + + %gate_weight_global = buffer.assume.memory_space %gate_weight : buffer + %up_weight_global = buffer.assume.memory_space %up_weight : buffer + %gate_output_global = buffer.assume.memory_space %gate_output : buffer + %up_output_global = buffer.assume.memory_space %up_output : buffer + %qact_global = buffer.assume.memory_space %qact : buffer + %scales_global = buffer.assume.memory_space %scales : buffer + %weight_global = scf.select %is_up, %up_weight_global, %gate_weight_global : buffer + %output_global = scf.select %is_up, %up_output_global, %gate_output_global : buffer + %weight_na, %output_na, %qact_na, %scales_na = buffer.assume.noalias %weight_global, %output_global, %qact_global, %scales_global : buffer, buffer, buffer, buffer + + %result0, %result1 = scf.if %valid_row -> (f32, f32) { + %acc0, %acc1 = scf.for %group = [%c0 to %group_count step %c1](%iter0 = %zero_f32 : f32, %iter1 = %zero_f32 : f32) -> (f32, f32) { + %block = index.div %group, %c8 : index + %group_in_block = index.rem %group, %c8 : index + %block_payload_add = index.mul %block, %payload_block_bytes : index + %group_payload_add = index.mul %group_in_block, %payload_plane_bytes : index + %row_payload_add = index.mul %row_lane, %c16 : index + %half_payload_add = index.mul %half, %c8 : index + %payload_index0 = index.add %group_base, %block_payload_add : index + %payload_index1 = index.add %payload_index0, %group_payload_add : index + %payload_index2 = index.add %payload_index1, %row_payload_add : index + %payload_index = index.add %payload_index2, %half_payload_add : index + %payload_base = index.cast %payload_index : index to offset + %weight_view = buffer.view %weight_na[%payload_base] : buffer -> view<2xi32> + %weight_words = vector.load %weight_view[%c0] : view<2xi32> -> vector<2xi32> + + %block_scale_add = index.mul %block, %scale_plane_bytes : index + %cohort_scale_add = index.mul %cohort_lane, %c16 : index + %group_scale_add = index.mul %group_in_block, %c2 : index + %scale_index0 = index.add %scale_region, %block_scale_add : index + %scale_index1 = index.add %scale_index0, %cohort_scale_add : index + %scale_byte_index = index.add %scale_index1, %group_scale_add : index + %scale_byte_base = index.cast %scale_byte_index : index to offset + %weight_scale_view = buffer.view %weight_na[%scale_byte_base] : buffer -> view<1xf16> + %weight_scale_f16 = view.load %weight_scale_view[%c0] : view<1xf16> -> f16 + %weight_scale = scalar.extf %weight_scale_f16 : f16 to f32 + + %dot0, %as0 = func.call @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_half_dot(%weight_words, %group, %c0, %half, %k, %group_count, %qact_na, %scales_na) : (vector<2xi32>, index, index, index, index, index, buffer, buffer) -> (i32, f32) + %dot1, %as1 = func.call @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_half_dot(%weight_words, %group, %c1, %half, %k, %group_count, %qact_na, %scales_na) : (vector<2xi32>, index, index, index, index, index, buffer, buffer) -> (i32, f32) + %peer0, %valid0 = kernel.subgroup.shuffle %dot0, %xor_one, %shuffle_width : i32, i32, i32 + %peer1, %valid1 = kernel.subgroup.shuffle %dot1, %xor_one, %shuffle_width : i32, i32, i32 + %full0 = scalar.addi %dot0, %peer0 : i32 + %full1 = scalar.addi %dot1, %peer1 : i32 + %fp0 = scalar.sitofp %full0 : i32 to f32 + %fp1 = scalar.sitofp %full1 : i32 to f32 + %scale0 = scalar.mulf %weight_scale, %as0 : f32 + %scale1 = scalar.mulf %weight_scale, %as1 : f32 + %next0 = scalar.fmaf %fp0, %scale0, %iter0 : f32 + %next1 = scalar.fmaf %fp1, %scale1, %iter1 : f32 + scf.yield %next0, %next1 : f32, f32 + } + scf.yield %acc0, %acc1 : f32, f32 + } else { + scf.yield %zero_f32, %zero_f32 : f32, f32 + } + + %is_writer = index.cmp eq, %half, %c0 : index + %writes = scalar.andi %valid_row, %is_writer : i1 + scf.if %writes { + %output_view = buffer.view %output_na[%base] : buffer -> view<[%cols]x[%rows]xf32> + view.store %result0, %output_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result1, %output_view[%c1, %row] : f32, view<[%cols]x[%rows]xf32> + } + kernel.return +} + +kernel.def export("ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c3") @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c3() { + %rows = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c15 = index.constant 15 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %rows_up = index.add %rows, %c15 : index + %row_workgroups = index.div %rows_up, %c16 : index + %workgroups = index.mul %row_workgroups, %c2 : index + kernel.launch.config workgroups(%workgroups, %c1, %c1) workgroup_size(%c32, %c1, %c1) : index +} launch(%gate_weight: buffer, %up_weight: buffer, %gate_output: buffer, %up_output: buffer, %qact: buffer, %scales: buffer, %sums: buffer) { + %k0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.input_size : index + %rows0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %cols0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.token_count : index + %k = index.assume %k0 [range(%k0, 256, 32768), mul(%k0, 256)] : index + %rows = index.assume %rows0 [range(%rows0, 1, 262144)] : index + %cols = index.assume %cols0 [range(%cols0, 3, 3)] : index + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %zero_f32 = scalar.constant 0.0 : f32 + %xor_one = scalar.constant 1 : i32 + %shuffle_width = scalar.constant 32 : i32 + + %workgroup0 = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %projection = index.rem %workgroup0, %c2 : index + %workgroup = index.div %workgroup0, %c2 : index + %is_up = index.cmp eq, %projection, %c1 : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 31)] : index + %row_in_wave = index.div %workitem, %c2 : index + %half = index.rem %workitem, %c2 : index + %row_base = index.mul %workgroup, %c16 : index + %row = index.add %row_base, %row_in_wave : index + %valid_row = index.cmp ult, %row, %rows : index + + %block_count = index.div %k, %c256 : index + %group_count = index.div %k, %c32 : index + %row_group0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.row_group_size : index + %rows_plus = index.add %rows, %c255 : index + %row_group = index.assume %row_group0 [range(%row_group0, 32, 262144), mul(%row_group0, 32)] : index + %row_group_id = index.div %row, %row_group : index + %row_lane = index.rem %row, %row_group : index + %cohort_lane = index.div %row_lane, %c4 : index + %payload_plane_bytes = index.mul %row_group, %c16 : index + %payload_block_bytes = index.mul %payload_plane_bytes, %c8 : index + %scale_cohorts = index.div %row_group, %c4 : index + %scale_plane_bytes = index.mul %scale_cohorts, %c16 : index + %block_bytes = index.add %payload_block_bytes, %scale_plane_bytes : index + %row_group_bytes = index.mul %block_count, %block_bytes : index + %group_base = index.mul %row_group_id, %row_group_bytes : index + %scale_region_add = index.mul %block_count, %payload_block_bytes : index + %scale_region = index.add %group_base, %scale_region_add : index + + %gate_weight_global = buffer.assume.memory_space %gate_weight : buffer + %up_weight_global = buffer.assume.memory_space %up_weight : buffer + %gate_output_global = buffer.assume.memory_space %gate_output : buffer + %up_output_global = buffer.assume.memory_space %up_output : buffer + %qact_global = buffer.assume.memory_space %qact : buffer + %scales_global = buffer.assume.memory_space %scales : buffer + %weight_global = scf.select %is_up, %up_weight_global, %gate_weight_global : buffer + %output_global = scf.select %is_up, %up_output_global, %gate_output_global : buffer + %weight_na, %output_na, %qact_na, %scales_na = buffer.assume.noalias %weight_global, %output_global, %qact_global, %scales_global : buffer, buffer, buffer, buffer + + %result0, %result1, %result2 = scf.if %valid_row -> (f32, f32, f32) { + %acc0, %acc1, %acc2 = scf.for %group = [%c0 to %group_count step %c1](%iter0 = %zero_f32 : f32, %iter1 = %zero_f32 : f32, %iter2 = %zero_f32 : f32) -> (f32, f32, f32) { + %block = index.div %group, %c8 : index + %group_in_block = index.rem %group, %c8 : index + %block_payload_add = index.mul %block, %payload_block_bytes : index + %group_payload_add = index.mul %group_in_block, %payload_plane_bytes : index + %row_payload_add = index.mul %row_lane, %c16 : index + %half_payload_add = index.mul %half, %c8 : index + %payload_index0 = index.add %group_base, %block_payload_add : index + %payload_index1 = index.add %payload_index0, %group_payload_add : index + %payload_index2 = index.add %payload_index1, %row_payload_add : index + %payload_index = index.add %payload_index2, %half_payload_add : index + %payload_base = index.cast %payload_index : index to offset + %weight_view = buffer.view %weight_na[%payload_base] : buffer -> view<2xi32> + %weight_words = vector.load %weight_view[%c0] : view<2xi32> -> vector<2xi32> + + %block_scale_add = index.mul %block, %scale_plane_bytes : index + %cohort_scale_add = index.mul %cohort_lane, %c16 : index + %group_scale_add = index.mul %group_in_block, %c2 : index + %scale_index0 = index.add %scale_region, %block_scale_add : index + %scale_index1 = index.add %scale_index0, %cohort_scale_add : index + %scale_byte_index = index.add %scale_index1, %group_scale_add : index + %scale_byte_base = index.cast %scale_byte_index : index to offset + %weight_scale_view = buffer.view %weight_na[%scale_byte_base] : buffer -> view<1xf16> + %weight_scale_f16 = view.load %weight_scale_view[%c0] : view<1xf16> -> f16 + %weight_scale = scalar.extf %weight_scale_f16 : f16 to f32 + + %dot0, %as0 = func.call @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_half_dot(%weight_words, %group, %c0, %half, %k, %group_count, %qact_na, %scales_na) : (vector<2xi32>, index, index, index, index, index, buffer, buffer) -> (i32, f32) + %dot1, %as1 = func.call @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_half_dot(%weight_words, %group, %c1, %half, %k, %group_count, %qact_na, %scales_na) : (vector<2xi32>, index, index, index, index, index, buffer, buffer) -> (i32, f32) + %dot2, %as2 = func.call @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_half_dot(%weight_words, %group, %c2, %half, %k, %group_count, %qact_na, %scales_na) : (vector<2xi32>, index, index, index, index, index, buffer, buffer) -> (i32, f32) + %peer0, %valid0 = kernel.subgroup.shuffle %dot0, %xor_one, %shuffle_width : i32, i32, i32 + %peer1, %valid1 = kernel.subgroup.shuffle %dot1, %xor_one, %shuffle_width : i32, i32, i32 + %peer2, %valid2 = kernel.subgroup.shuffle %dot2, %xor_one, %shuffle_width : i32, i32, i32 + %full0 = scalar.addi %dot0, %peer0 : i32 + %full1 = scalar.addi %dot1, %peer1 : i32 + %full2 = scalar.addi %dot2, %peer2 : i32 + %fp0 = scalar.sitofp %full0 : i32 to f32 + %fp1 = scalar.sitofp %full1 : i32 to f32 + %fp2 = scalar.sitofp %full2 : i32 to f32 + %scale0 = scalar.mulf %weight_scale, %as0 : f32 + %scale1 = scalar.mulf %weight_scale, %as1 : f32 + %scale2 = scalar.mulf %weight_scale, %as2 : f32 + %next0 = scalar.fmaf %fp0, %scale0, %iter0 : f32 + %next1 = scalar.fmaf %fp1, %scale1, %iter1 : f32 + %next2 = scalar.fmaf %fp2, %scale2, %iter2 : f32 + scf.yield %next0, %next1, %next2 : f32, f32, f32 + } + scf.yield %acc0, %acc1, %acc2 : f32, f32, f32 + } else { + scf.yield %zero_f32, %zero_f32, %zero_f32 : f32, f32, f32 + } + + %is_writer = index.cmp eq, %half, %c0 : index + %writes = scalar.andi %valid_row, %is_writer : i1 + scf.if %writes { + %output_view = buffer.view %output_na[%base] : buffer -> view<[%cols]x[%rows]xf32> + view.store %result0, %output_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result1, %output_view[%c1, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result2, %output_view[%c2, %row] : f32, view<[%cols]x[%rows]xf32> + } + kernel.return +} + +kernel.def export("ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c4") @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c4() { + %rows = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c15 = index.constant 15 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %rows_up = index.add %rows, %c15 : index + %row_workgroups = index.div %rows_up, %c16 : index + %workgroups = index.mul %row_workgroups, %c2 : index + kernel.launch.config workgroups(%workgroups, %c1, %c1) workgroup_size(%c32, %c1, %c1) : index +} launch(%gate_weight: buffer, %up_weight: buffer, %gate_output: buffer, %up_output: buffer, %qact: buffer, %scales: buffer, %sums: buffer) { + %k0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.input_size : index + %rows0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %cols0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.token_count : index + %k = index.assume %k0 [range(%k0, 256, 32768), mul(%k0, 256)] : index + %rows = index.assume %rows0 [range(%rows0, 1, 262144)] : index + %cols = index.assume %cols0 [range(%cols0, 4, 4)] : index + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %zero_f32 = scalar.constant 0.0 : f32 + %xor_one = scalar.constant 1 : i32 + %shuffle_width = scalar.constant 32 : i32 + + %workgroup0 = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %projection = index.rem %workgroup0, %c2 : index + %workgroup = index.div %workgroup0, %c2 : index + %is_up = index.cmp eq, %projection, %c1 : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 31)] : index + %row_in_wave = index.div %workitem, %c2 : index + %half = index.rem %workitem, %c2 : index + %row_base = index.mul %workgroup, %c16 : index + %row = index.add %row_base, %row_in_wave : index + %valid_row = index.cmp ult, %row, %rows : index + + %block_count = index.div %k, %c256 : index + %group_count = index.div %k, %c32 : index + %row_group0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.row_group_size : index + %rows_plus = index.add %rows, %c255 : index + %row_group = index.assume %row_group0 [range(%row_group0, 32, 262144), mul(%row_group0, 32)] : index + %row_group_id = index.div %row, %row_group : index + %row_lane = index.rem %row, %row_group : index + %cohort_lane = index.div %row_lane, %c4 : index + %payload_plane_bytes = index.mul %row_group, %c16 : index + %payload_block_bytes = index.mul %payload_plane_bytes, %c8 : index + %scale_cohorts = index.div %row_group, %c4 : index + %scale_plane_bytes = index.mul %scale_cohorts, %c16 : index + %block_bytes = index.add %payload_block_bytes, %scale_plane_bytes : index + %row_group_bytes = index.mul %block_count, %block_bytes : index + %group_base = index.mul %row_group_id, %row_group_bytes : index + %scale_region_add = index.mul %block_count, %payload_block_bytes : index + %scale_region = index.add %group_base, %scale_region_add : index + + %gate_weight_global = buffer.assume.memory_space %gate_weight : buffer + %up_weight_global = buffer.assume.memory_space %up_weight : buffer + %gate_output_global = buffer.assume.memory_space %gate_output : buffer + %up_output_global = buffer.assume.memory_space %up_output : buffer + %qact_global = buffer.assume.memory_space %qact : buffer + %scales_global = buffer.assume.memory_space %scales : buffer + %weight_global = scf.select %is_up, %up_weight_global, %gate_weight_global : buffer + %output_global = scf.select %is_up, %up_output_global, %gate_output_global : buffer + %weight_na, %output_na, %qact_na, %scales_na = buffer.assume.noalias %weight_global, %output_global, %qact_global, %scales_global : buffer, buffer, buffer, buffer + + %result0, %result1, %result2, %result3 = scf.if %valid_row -> (f32, f32, f32, f32) { + %acc0, %acc1, %acc2, %acc3 = scf.for %group = [%c0 to %group_count step %c1](%iter0 = %zero_f32 : f32, %iter1 = %zero_f32 : f32, %iter2 = %zero_f32 : f32, %iter3 = %zero_f32 : f32) -> (f32, f32, f32, f32) { + %block = index.div %group, %c8 : index + %group_in_block = index.rem %group, %c8 : index + %block_payload_add = index.mul %block, %payload_block_bytes : index + %group_payload_add = index.mul %group_in_block, %payload_plane_bytes : index + %row_payload_add = index.mul %row_lane, %c16 : index + %half_payload_add = index.mul %half, %c8 : index + %payload_index0 = index.add %group_base, %block_payload_add : index + %payload_index1 = index.add %payload_index0, %group_payload_add : index + %payload_index2 = index.add %payload_index1, %row_payload_add : index + %payload_index = index.add %payload_index2, %half_payload_add : index + %payload_base = index.cast %payload_index : index to offset + %weight_view = buffer.view %weight_na[%payload_base] : buffer -> view<2xi32> + %weight_words = vector.load %weight_view[%c0] : view<2xi32> -> vector<2xi32> + + %block_scale_add = index.mul %block, %scale_plane_bytes : index + %cohort_scale_add = index.mul %cohort_lane, %c16 : index + %group_scale_add = index.mul %group_in_block, %c2 : index + %scale_index0 = index.add %scale_region, %block_scale_add : index + %scale_index1 = index.add %scale_index0, %cohort_scale_add : index + %scale_byte_index = index.add %scale_index1, %group_scale_add : index + %scale_byte_base = index.cast %scale_byte_index : index to offset + %weight_scale_view = buffer.view %weight_na[%scale_byte_base] : buffer -> view<1xf16> + %weight_scale_f16 = view.load %weight_scale_view[%c0] : view<1xf16> -> f16 + %weight_scale = scalar.extf %weight_scale_f16 : f16 to f32 + + %dot0, %as0 = func.call @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_half_dot(%weight_words, %group, %c0, %half, %k, %group_count, %qact_na, %scales_na) : (vector<2xi32>, index, index, index, index, index, buffer, buffer) -> (i32, f32) + %dot1, %as1 = func.call @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_half_dot(%weight_words, %group, %c1, %half, %k, %group_count, %qact_na, %scales_na) : (vector<2xi32>, index, index, index, index, index, buffer, buffer) -> (i32, f32) + %dot2, %as2 = func.call @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_half_dot(%weight_words, %group, %c2, %half, %k, %group_count, %qact_na, %scales_na) : (vector<2xi32>, index, index, index, index, index, buffer, buffer) -> (i32, f32) + %dot3, %as3 = func.call @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_half_dot(%weight_words, %group, %c3, %half, %k, %group_count, %qact_na, %scales_na) : (vector<2xi32>, index, index, index, index, index, buffer, buffer) -> (i32, f32) + %peer0, %valid0 = kernel.subgroup.shuffle %dot0, %xor_one, %shuffle_width : i32, i32, i32 + %peer1, %valid1 = kernel.subgroup.shuffle %dot1, %xor_one, %shuffle_width : i32, i32, i32 + %peer2, %valid2 = kernel.subgroup.shuffle %dot2, %xor_one, %shuffle_width : i32, i32, i32 + %peer3, %valid3 = kernel.subgroup.shuffle %dot3, %xor_one, %shuffle_width : i32, i32, i32 + %full0 = scalar.addi %dot0, %peer0 : i32 + %full1 = scalar.addi %dot1, %peer1 : i32 + %full2 = scalar.addi %dot2, %peer2 : i32 + %full3 = scalar.addi %dot3, %peer3 : i32 + %fp0 = scalar.sitofp %full0 : i32 to f32 + %fp1 = scalar.sitofp %full1 : i32 to f32 + %fp2 = scalar.sitofp %full2 : i32 to f32 + %fp3 = scalar.sitofp %full3 : i32 to f32 + %scale0 = scalar.mulf %weight_scale, %as0 : f32 + %scale1 = scalar.mulf %weight_scale, %as1 : f32 + %scale2 = scalar.mulf %weight_scale, %as2 : f32 + %scale3 = scalar.mulf %weight_scale, %as3 : f32 + %next0 = scalar.fmaf %fp0, %scale0, %iter0 : f32 + %next1 = scalar.fmaf %fp1, %scale1, %iter1 : f32 + %next2 = scalar.fmaf %fp2, %scale2, %iter2 : f32 + %next3 = scalar.fmaf %fp3, %scale3, %iter3 : f32 + scf.yield %next0, %next1, %next2, %next3 : f32, f32, f32, f32 + } + scf.yield %acc0, %acc1, %acc2, %acc3 : f32, f32, f32, f32 + } else { + scf.yield %zero_f32, %zero_f32, %zero_f32, %zero_f32 : f32, f32, f32, f32 + } + + %is_writer = index.cmp eq, %half, %c0 : index + %writes = scalar.andi %valid_row, %is_writer : i1 + scf.if %writes { + %output_view = buffer.view %output_na[%base] : buffer -> view<[%cols]x[%rows]xf32> + view.store %result0, %output_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result1, %output_view[%c1, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result2, %output_view[%c2, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result3, %output_view[%c3, %row] : f32, view<[%cols]x[%rows]xf32> + } + kernel.return +} + +kernel.def export("ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c5") @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_c5() { + %rows = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c15 = index.constant 15 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %rows_up = index.add %rows, %c15 : index + %row_workgroups = index.div %rows_up, %c16 : index + %workgroups = index.mul %row_workgroups, %c2 : index + kernel.launch.config workgroups(%workgroups, %c1, %c1) workgroup_size(%c32, %c1, %c1) : index +} launch(%gate_weight: buffer, %up_weight: buffer, %gate_output: buffer, %up_output: buffer, %qact: buffer, %scales: buffer, %sums: buffer) { + %k0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.input_size : index + %rows0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %cols0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.token_count : index + %k = index.assume %k0 [range(%k0, 256, 32768), mul(%k0, 256)] : index + %rows = index.assume %rows0 [range(%rows0, 1, 262144)] : index + %cols = index.assume %cols0 [range(%cols0, 5, 5)] : index + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %zero_f32 = scalar.constant 0.0 : f32 + %xor_one = scalar.constant 1 : i32 + %shuffle_width = scalar.constant 32 : i32 + + %workgroup0 = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %projection = index.rem %workgroup0, %c2 : index + %workgroup = index.div %workgroup0, %c2 : index + %is_up = index.cmp eq, %projection, %c1 : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 31)] : index + %row_in_wave = index.div %workitem, %c2 : index + %half = index.rem %workitem, %c2 : index + %row_base = index.mul %workgroup, %c16 : index + %row = index.add %row_base, %row_in_wave : index + %valid_row = index.cmp ult, %row, %rows : index + + %block_count = index.div %k, %c256 : index + %group_count = index.div %k, %c32 : index + %row_group0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.row_group_size : index + %rows_plus = index.add %rows, %c255 : index + %row_group = index.assume %row_group0 [range(%row_group0, 32, 262144), mul(%row_group0, 32)] : index + %row_group_id = index.div %row, %row_group : index + %row_lane = index.rem %row, %row_group : index + %cohort_lane = index.div %row_lane, %c4 : index + %payload_plane_bytes = index.mul %row_group, %c16 : index + %payload_block_bytes = index.mul %payload_plane_bytes, %c8 : index + %scale_cohorts = index.div %row_group, %c4 : index + %scale_plane_bytes = index.mul %scale_cohorts, %c16 : index + %block_bytes = index.add %payload_block_bytes, %scale_plane_bytes : index + %row_group_bytes = index.mul %block_count, %block_bytes : index + %group_base = index.mul %row_group_id, %row_group_bytes : index + %scale_region_add = index.mul %block_count, %payload_block_bytes : index + %scale_region = index.add %group_base, %scale_region_add : index + + %gate_weight_global = buffer.assume.memory_space %gate_weight : buffer + %up_weight_global = buffer.assume.memory_space %up_weight : buffer + %gate_output_global = buffer.assume.memory_space %gate_output : buffer + %up_output_global = buffer.assume.memory_space %up_output : buffer + %qact_global = buffer.assume.memory_space %qact : buffer + %scales_global = buffer.assume.memory_space %scales : buffer + %weight_global = scf.select %is_up, %up_weight_global, %gate_weight_global : buffer + %output_global = scf.select %is_up, %up_output_global, %gate_output_global : buffer + %weight_na, %output_na, %qact_na, %scales_na = buffer.assume.noalias %weight_global, %output_global, %qact_global, %scales_global : buffer, buffer, buffer, buffer + + %result0, %result1, %result2, %result3, %result4 = scf.if %valid_row -> (f32, f32, f32, f32, f32) { + %acc0, %acc1, %acc2, %acc3, %acc4 = scf.for %group = [%c0 to %group_count step %c1](%iter0 = %zero_f32 : f32, %iter1 = %zero_f32 : f32, %iter2 = %zero_f32 : f32, %iter3 = %zero_f32 : f32, %iter4 = %zero_f32 : f32) -> (f32, f32, f32, f32, f32) { + %block = index.div %group, %c8 : index + %group_in_block = index.rem %group, %c8 : index + %block_payload_add = index.mul %block, %payload_block_bytes : index + %group_payload_add = index.mul %group_in_block, %payload_plane_bytes : index + %row_payload_add = index.mul %row_lane, %c16 : index + %half_payload_add = index.mul %half, %c8 : index + %payload_index0 = index.add %group_base, %block_payload_add : index + %payload_index1 = index.add %payload_index0, %group_payload_add : index + %payload_index2 = index.add %payload_index1, %row_payload_add : index + %payload_index = index.add %payload_index2, %half_payload_add : index + %payload_base = index.cast %payload_index : index to offset + %weight_view = buffer.view %weight_na[%payload_base] : buffer -> view<2xi32> + %weight_words = vector.load %weight_view[%c0] : view<2xi32> -> vector<2xi32> + + %block_scale_add = index.mul %block, %scale_plane_bytes : index + %cohort_scale_add = index.mul %cohort_lane, %c16 : index + %group_scale_add = index.mul %group_in_block, %c2 : index + %scale_index0 = index.add %scale_region, %block_scale_add : index + %scale_index1 = index.add %scale_index0, %cohort_scale_add : index + %scale_byte_index = index.add %scale_index1, %group_scale_add : index + %scale_byte_base = index.cast %scale_byte_index : index to offset + %weight_scale_view = buffer.view %weight_na[%scale_byte_base] : buffer -> view<1xf16> + %weight_scale_f16 = view.load %weight_scale_view[%c0] : view<1xf16> -> f16 + %weight_scale = scalar.extf %weight_scale_f16 : f16 to f32 + + %dot0, %as0 = func.call @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_half_dot(%weight_words, %group, %c0, %half, %k, %group_count, %qact_na, %scales_na) : (vector<2xi32>, index, index, index, index, index, buffer, buffer) -> (i32, f32) + %dot1, %as1 = func.call @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_half_dot(%weight_words, %group, %c1, %half, %k, %group_count, %qact_na, %scales_na) : (vector<2xi32>, index, index, index, index, index, buffer, buffer) -> (i32, f32) + %dot2, %as2 = func.call @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_half_dot(%weight_words, %group, %c2, %half, %k, %group_count, %qact_na, %scales_na) : (vector<2xi32>, index, index, index, index, index, buffer, buffer) -> (i32, f32) + %dot3, %as3 = func.call @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_half_dot(%weight_words, %group, %c3, %half, %k, %group_count, %qact_na, %scales_na) : (vector<2xi32>, index, index, index, index, index, buffer, buffer) -> (i32, f32) + %dot4, %as4 = func.call @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_direct_dot_half_dot(%weight_words, %group, %c4, %half, %k, %group_count, %qact_na, %scales_na) : (vector<2xi32>, index, index, index, index, index, buffer, buffer) -> (i32, f32) + %peer0, %valid0 = kernel.subgroup.shuffle %dot0, %xor_one, %shuffle_width : i32, i32, i32 + %peer1, %valid1 = kernel.subgroup.shuffle %dot1, %xor_one, %shuffle_width : i32, i32, i32 + %peer2, %valid2 = kernel.subgroup.shuffle %dot2, %xor_one, %shuffle_width : i32, i32, i32 + %peer3, %valid3 = kernel.subgroup.shuffle %dot3, %xor_one, %shuffle_width : i32, i32, i32 + %peer4, %valid4 = kernel.subgroup.shuffle %dot4, %xor_one, %shuffle_width : i32, i32, i32 + %full0 = scalar.addi %dot0, %peer0 : i32 + %full1 = scalar.addi %dot1, %peer1 : i32 + %full2 = scalar.addi %dot2, %peer2 : i32 + %full3 = scalar.addi %dot3, %peer3 : i32 + %full4 = scalar.addi %dot4, %peer4 : i32 + %fp0 = scalar.sitofp %full0 : i32 to f32 + %fp1 = scalar.sitofp %full1 : i32 to f32 + %fp2 = scalar.sitofp %full2 : i32 to f32 + %fp3 = scalar.sitofp %full3 : i32 to f32 + %fp4 = scalar.sitofp %full4 : i32 to f32 + %scale0 = scalar.mulf %weight_scale, %as0 : f32 + %scale1 = scalar.mulf %weight_scale, %as1 : f32 + %scale2 = scalar.mulf %weight_scale, %as2 : f32 + %scale3 = scalar.mulf %weight_scale, %as3 : f32 + %scale4 = scalar.mulf %weight_scale, %as4 : f32 + %next0 = scalar.fmaf %fp0, %scale0, %iter0 : f32 + %next1 = scalar.fmaf %fp1, %scale1, %iter1 : f32 + %next2 = scalar.fmaf %fp2, %scale2, %iter2 : f32 + %next3 = scalar.fmaf %fp3, %scale3, %iter3 : f32 + %next4 = scalar.fmaf %fp4, %scale4, %iter4 : f32 + scf.yield %next0, %next1, %next2, %next3, %next4 : f32, f32, f32, f32, f32 + } + scf.yield %acc0, %acc1, %acc2, %acc3, %acc4 : f32, f32, f32, f32, f32 + } else { + scf.yield %zero_f32, %zero_f32, %zero_f32, %zero_f32, %zero_f32 : f32, f32, f32, f32, f32 + } + + %is_writer = index.cmp eq, %half, %c0 : index + %writes = scalar.andi %valid_row, %is_writer : i1 + scf.if %writes { + %output_view = buffer.view %output_na[%base] : buffer -> view<[%cols]x[%rows]xf32> + view.store %result0, %output_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result1, %output_view[%c1, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result2, %output_view[%c2, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result3, %output_view[%c3, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result4, %output_view[%c4, %row] : f32, view<[%cols]x[%rows]xf32> + } + kernel.return +} + +template.def<@ggml.quantize_symmetric_i4_k64.body> device @ggml_quantize_symmetric_i4_k64_body(%plane_major: i1, %src: buffer, %qs: buffer, %ds: buffer, %sums: buffer) { + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %c0_f32 = scalar.constant 0.0 : f32 + %amax_eps = scalar.constant 1.0000000000000001e-30 : f32 + %inv7 = scalar.constant 0.14285714285714285 : f32 + %f7 = scalar.constant 7.0 : f32 + %xor1 = scalar.constant 1 : i32 + %xor2 = scalar.constant 2 : i32 + %xor4 = scalar.constant 4 : i32 + %xor8 = scalar.constant 8 : i32 + %shift4 = scalar.constant 4 : i32 + %shift8 = scalar.constant 8 : i32 + %shift12 = scalar.constant 12 : i32 + %shuffle_width = scalar.constant 32 : i32 + %mask15 = vector.constant 15 : vector<4xi32> + + %k = config.get @ggml.quantize_symmetric_i4_k64.input_size : index + %cols = config.get @ggml.quantize_symmetric_i4_k64.token_count : index + %k_b = index.assume %k [range(%k, 32, 65536)] : index + %cols_b = index.assume %cols [range(%cols, 1, 16384)] : index + %total = index.mul %k_b, %cols_b : index + %groups0 = index.div %total, %c32 : index + %groups = index.assume %groups0 [range(%groups0, 1, 33554432)] : index + %groups64_0 = index.div %total, %c64 : index + %groups64 = index.assume %groups64_0 [range(%groups64_0, 1, 16777216)] : index + %nchunks = index.div %k_b, %c64 : index + + %wgid0 = kernel.workgroup.id : index + %tid0 = kernel.workitem.id : index + %wgid = index.assume %wgid0 [range(%wgid0, 0, 131071)] : index + %tid = index.assume %tid0 [range(%tid0, 0, 255)] : index + %wave = index.div %tid, %c32 : index + %lane = index.rem %tid, %c32 : index + %group_in_wave = index.div %lane, %c16 : index + %lane_in_group = index.rem %lane, %c16 : index + %group_half = index.div %lane_in_group, %c8 : index + %lane_in_half = index.rem %lane_in_group, %c8 : index + %is_scale_leader = index.cmp eq, %lane_in_group, %c0 : index + %is_sum_leader = index.cmp eq, %lane_in_half, %c0 : index + %wg_group_base = index.mul %wgid, %c128 : index + %wave_group_base = index.mul %wave, %c2 : index + %local_group_base = index.add %wg_group_base, %wave_group_base : index + + %src_g = buffer.assume.memory_space %src : buffer + %qs_g = buffer.assume.memory_space %qs : buffer + %ds_g = buffer.assume.memory_space %ds : buffer + %sums_g = buffer.assume.memory_space %sums : buffer + %src_na, %qs_na, %ds_na, %sums_na = buffer.assume.noalias %src_g, %qs_g, %ds_g, %sums_g : buffer, buffer, buffer, buffer + %src_elems = index.assume %total [range(%total, 32, 1073741824)] : index + %src_view = buffer.view %src_na[%base] : buffer -> view<[%src_elems]xf32> + %qs_halfwords = index.div %total, %c4 : index + %qs_view = buffer.view %qs_na[%base] : buffer -> view<[%qs_halfwords]xi16> + %ds_view = buffer.view %ds_na[%base] : buffer -> view<[%groups]xf32> + %sum_view = buffer.view %sums_na[%base] : buffer -> view<[%groups]xi32> + + scf.for %batch = [%c0 to %c8 step %c1] { + %batch_group_off = index.mul %batch, %c16 : index + %gid0 = index.add %local_group_base, %batch_group_off : index + %gid1 = index.add %gid0, %group_in_wave : index + %gid = index.assume %gid1 [range(%gid1, 0, 16777215)] : index + %in_range = index.cmp ult, %gid, %groups64 : index + scf.if %in_range { + %ebase0 = index.mul %gid, %c64 : index + %lane_elem_off = index.mul %lane_in_group, %c4 : index + %ebase1 = index.add %ebase0, %lane_elem_off : index + %ebase = index.assume %ebase1 [range(%ebase1, 0, 1073741820)] : index + %v = vector.load %src_view[%ebase] : view<[%src_elems]xf32> -> vector<4xf32> + %av = vector.absf %v : vector<4xf32> + %lane_max = vector.reduce %av, %c0_f32 : vector<4xf32>, f32 + %max_x1_peer, %max_x1_valid = kernel.subgroup.shuffle %lane_max, %xor1, %shuffle_width : f32, i32, i32 + %max_x1 = scalar.maxnumf %lane_max, %max_x1_peer : f32 + %max_x2_peer, %max_x2_valid = kernel.subgroup.shuffle %max_x1, %xor2, %shuffle_width : f32, i32, i32 + %max_x2 = scalar.maxnumf %max_x1, %max_x2_peer : f32 + %max_x4_peer, %max_x4_valid = kernel.subgroup.shuffle %max_x2, %xor4, %shuffle_width : f32, i32, i32 + %max_x4 = scalar.maxnumf %max_x2, %max_x4_peer : f32 + %max_x8_peer, %max_x8_valid = kernel.subgroup.shuffle %max_x4, %xor8, %shuffle_width : f32, i32, i32 + %group_max = scalar.maxnumf %max_x4, %max_x8_peer : f32 + %amax = scalar.maxnumf %group_max, %amax_eps : f32 + %a_scale = scalar.mulf %amax, %inv7 : f32 + %a_rscale = scalar.divf %f7, %amax : f32 + %gid32_base = index.mul %gid, %c2 : index + %gid32_high = index.add %gid32_base, %c1 : index + %token = index.div %gid, %nchunks : index + %chunk = index.rem %gid, %nchunks : index + %plane_chunk32 = index.mul %chunk, %c2 : index + %plane_low_base = index.mul %plane_chunk32, %cols_b : index + %plane_gid32_low = index.add %plane_low_base, %token : index + %plane_chunk32_high = index.add %plane_chunk32, %c1 : index + %plane_high_base = index.mul %plane_chunk32_high, %cols_b : index + %plane_gid32_high = index.add %plane_high_base, %token : index + %scale_low_index = scf.if %plane_major -> (index) { + scf.yield %plane_gid32_low : index + } else { + scf.yield %gid32_base : index + } + %scale_high_index = scf.if %plane_major -> (index) { + scf.yield %plane_gid32_high : index + } else { + scf.yield %gid32_high : index + } + scf.if %is_scale_leader { + view.store %a_scale, %ds_view[%scale_low_index] : f32, view<[%groups]xf32> + view.store %a_scale, %ds_view[%scale_high_index] : f32, view<[%groups]xf32> + } + + %rs = vector.splat %a_rscale : vector<4xf32> + %scaled = vector.mulf %v, %rs : vector<4xf32> + %rounded = vector.roundf %scaled : vector<4xf32> + %q32 = vector.fptosi %rounded : vector<4xf32> to vector<4xi32> + %qn = vector.andi %q32, %mask15 : vector<4xi32> + %q0 = vector.extract %qn[0] : vector<4xi32> -> i32 + %q1 = vector.extract %qn[1] : vector<4xi32> -> i32 + %q2 = vector.extract %qn[2] : vector<4xi32> -> i32 + %q3 = vector.extract %qn[3] : vector<4xi32> -> i32 + %q1s = scalar.shli %q1, %shift4 : i32 + %q2s = scalar.shli %q2, %shift8 : i32 + %q3s = scalar.shli %q3, %shift12 : i32 + %q01 = scalar.ori %q0, %q1s : i32 + %q23 = scalar.ori %q2s, %q3s : i32 + %qpacked32 = scalar.ori %q01, %q23 : i32 + %qpacked = scalar.trunci %qpacked32 : i32 to i16 + %plane_chunk_base = index.mul %chunk, %cols_b : index + %plane_gid = index.add %plane_chunk_base, %token : index + %payload_gid = scf.if %plane_major -> (index) { + scf.yield %plane_gid : index + } else { + scf.yield %gid : index + } + %word_base = index.mul %payload_gid, %c16 : index + %word_index0 = index.add %word_base, %lane_in_group : index + %word_index = index.assume %word_index0 [range(%word_index0, 0, 268435455)] : index + view.store %qpacked, %qs_view[%word_index] : i16, view<[%qs_halfwords]xi16> + + %lane_sum = vector.reduce %rounded, %c0_f32 : vector<4xf32>, f32 + %sum_x1_peer, %sum_x1_valid = kernel.subgroup.shuffle %lane_sum, %xor1, %shuffle_width : f32, i32, i32 + %sum_x1 = scalar.addf %lane_sum, %sum_x1_peer : f32 + %sum_x2_peer, %sum_x2_valid = kernel.subgroup.shuffle %sum_x1, %xor2, %shuffle_width : f32, i32, i32 + %sum_x2 = scalar.addf %sum_x1, %sum_x2_peer : f32 + %sum_x4_peer, %sum_x4_valid = kernel.subgroup.shuffle %sum_x2, %xor4, %shuffle_width : f32, i32, i32 + %group_sum = scalar.addf %sum_x2, %sum_x4_peer : f32 + scf.if %is_sum_leader { + %sum_value = scalar.fptosi %group_sum : f32 to i32 + %row_sum_index = index.add %gid32_base, %group_half : index + %plane_sum_chunk = index.add %plane_chunk32, %group_half : index + %plane_sum_base = index.mul %plane_sum_chunk, %cols_b : index + %plane_sum_index = index.add %plane_sum_base, %token : index + %sum_index = scf.if %plane_major -> (index) { + scf.yield %plane_sum_index : index + } else { + scf.yield %row_sum_index : index + } + view.store %sum_value, %sum_view[%sum_index] : i32, view<[%groups]xi32> + } + } + } + template.return +} + +template.decl @ggml.mul_mat.symmetric_i4.m128n128_wg256.body(%apply_swiglu: i1, %publish_f32: i1, %publish_f16: i1, %publish_u4: i1, %gate_na: buffer, %src0_na: buffer, %dst_na: buffer, %aq_na: buffer, %as_na: buffer, %asum_na: buffer, %qout_qs_na: buffer, %qout_ds_na: buffer, %qout_sums_na: buffer, %qout_scratch: buffer, %base: offset, %qact_scale_base: offset, %c0: index, %c1: index, %c2: index, %c4: index, %c8: index, %c16: index, %c32: index, %c17: index, %c34: index, %k_kc: index, %k_astride: index, %k_wstride: index, %k_m: index, %k_n: index, %k_blocks: index, %k_wave: index, %k_wm: index, %fzero: vector<8xf32>, %izero: vector<8xi32>, %k_b: index, %rows_b: index, %cols_b: index, %nchunks: index, %kblocks: index, %col_tile: index, %row_tile: index, %tid: index, %wave: index) + +config.decl @ggml.mul_mat.symmetric_i4.input_size : %value: index where [range(%value, 32, 32768), mul(%value, 32)] + +config.decl @ggml.mul_mat.symmetric_i4.output_size : %value: index where [range(%value, 1, 262144)] + +config.decl @ggml.mul_mat.symmetric_i4.token_count : %value: index where [range(%value, 1, 32768)] + +// Shared Q8_0/IU8 WMMA contraction schedule. QAct storage layout is +// normalized by the two roots before this compile-time template is applied. +template.def<@ggml.mul_mat.symmetric_i4.m128n128_wg256.body> device @ggml_mul_mat_symmetric_i4_m128n128_wg256_body(%apply_swiglu: i1, %publish_f32: i1, %publish_f16: i1, %publish_u4: i1, %gate_na: buffer, %src0_na: buffer, %dst_na: buffer, %aq_na: buffer, %as_na: buffer, %asum_na: buffer, %qout_qs_na: buffer, %qout_ds_na: buffer, %qout_sums_na: buffer, %qout_scratch: buffer, %base: offset, %qact_scale_base: offset, %c0: index, %c1: index, %c2: index, %c4: index, %c8: index, %c16: index, %c32: index, %c17: index, %c34: index, %k_kc: index, %k_astride: index, %k_wstride: index, %k_m: index, %k_n: index, %k_blocks: index, %k_wave: index, %k_wm: index, %fzero: vector<8xf32>, %izero: vector<8xi32>, %k_b: index, %rows_b: index, %cols_b: index, %nchunks: index, %kblocks: index, %col_tile: index, %row_tile: index, %tid: index, %wave: index) { + %src0_i8 = buffer.view %src0_na[%base] : buffer -> view<1073741824xi8> + %src0_f16 = buffer.view %src0_na[%base] : buffer -> view<536870912xf16> + %aq_flat = buffer.view %aq_na[%base] : buffer -> view<1073741824xi8> + %as_flat = buffer.view %as_na[%qact_scale_base] : buffer -> view<33554432xf32> + %asum_flat = buffer.view %asum_na[%base] : buffer -> view<33554432xi32> + %c0_i32 = scalar.constant 0 : i32 + %pf_zero_v16i8 = vector.constant 0 : vector<16xi8> + %pf_payload_bytes = index.constant 32 : index + %q4_k_block = index.constant 256 : index + %q4_block_bytes = index.constant 144 : index + %q4_k48 = index.constant 48 : index + // dst is [rows, cols] with rows contiguous, and C is [cols, rows]. + %dst_view = buffer.view %dst_na[%base] : buffer -> view<[%cols_b]x[%rows_b]xf32> + %dst_f16_view = buffer.view %dst_na[%base] : buffer -> view<[%cols_b]x[%rows_b]xf16> + %gate_view = buffer.view %gate_na[%base] : buffer -> view<[%cols_b]x[%rows_b]xf32> + + // Packed I4 operands use 32 payload bytes per logical K64 row. + // A 40-byte LDS stride retains the conflict-avoiding padding of the IU8 body. + %scratch = buffer.assume.memory_space %qout_scratch : buffer + %i4_schema = encoding.define #encoding.operand : encoding + %u4_schema = encoding.define #encoding.operand : encoding + %al_flat = buffer.view %scratch[%base] : buffer -> view<5120xi8> + %w_off = index.constant 5120 : offset + %wl_flat = buffer.view %scratch[%w_off] : buffer -> view<5120xi8> + %as_off = index.constant 10240 : offset + %asl_view = buffer.view %scratch[%as_off] : buffer -> view<256xf32> + %ws_off = index.constant 11264 : offset + %wsl_view = buffer.view %scratch[%ws_off] : buffer -> view<256xf32> + %asum_off = index.constant 12288 : offset + %asuml_view = buffer.view %scratch[%asum_off] : buffer -> view<256xf32> + %wc_off = index.constant 13312 : offset + %wcl_view = buffer.view %scratch[%wc_off] : buffer -> view<256xf32> + + %col_base = index.mul %col_tile, %k_m : index + %row_base = index.mul %row_tile, %k_n : index + + // Staging shares. + %k_aper = index.constant 64 : index + %k_wper = index.constant 64 : index + %ashare = index.mul %tid, %k_aper : index + %ascol = index.div %ashare, %k_kc : index + %asoff = index.rem %ashare, %k_kc : index + %adbase0 = index.mul %ascol, %k_astride : index + %adbase = index.add %adbase0, %asoff : index + %ascol_g = index.add %col_base, %ascol : index + %asbase = index.mul %ascol_g, %k_b : index + // WG64 ILP8 staging maps the complete native-Q8 tile in explicit passes. + // Two scale passes cover 128 K32 blocks; four payload passes cover 256 halves. + // Scale staging: scalar legacy order or eight-scale WMMA-native vectors. + %k_asn = index.constant 256 : index + %k_aspan = index.constant 32 : index + %k_wsn = index.constant 128 : index + %k_lanes0 = index.constant 0 : index + %k_lanes1 = index.constant 128 : index + + // Four waves own four 32-token blocks and share one 64-row weight tile. + %wcol = index.rem %wave, %k_wm : index + %wrow = index.div %wave, %k_wm : index + %k_mspan = index.constant 32 : index + %k_nspan = index.constant 64 : index + %wm_off = index.mul %wcol, %k_mspan : index + %wn_off = index.mul %wrow, %k_nspan : index + %m_out = index.add %col_base, %wm_off : index + %n_out = index.add %row_base, %wn_off : index + // RDNA3 WMMA RHS payload lanes map to columns lane % 16. Point each + // lane at its own output row before constructing the four N fragments. + %lane = index.rem %tid, %k_wave : index + %lane_lo = index.rem %lane, %c16 : index + %lane_hi = index.div %lane, %c16 : index + %wn_lane = index.add %wn_off, %lane_lo : index + %lm0 = index.add %wm_off, %c0 : index + %gm0 = index.add %m_out, %c0 : index + %k_ma1 = index.constant 16 : index + %lm1 = index.add %wm_off, %k_ma1 : index + %gm1 = index.add %m_out, %k_ma1 : index + %ln0 = index.add %wn_lane, %c0 : index + %gn0 = index.add %n_out, %c0 : index + %k_nb1 = index.constant 16 : index + %ln1 = index.add %wn_lane, %k_nb1 : index + %gn1 = index.add %n_out, %k_nb1 : index + %k_nb2 = index.constant 32 : index + %ln2 = index.add %wn_lane, %k_nb2 : index + %gn2 = index.add %n_out, %k_nb2 : index + %k_nb3 = index.constant 48 : index + %ln3 = index.add %wn_lane, %k_nb3 : index + %gn3 = index.add %n_out, %k_nb3 : index + %wsb0_0 = index.add %ln0, %c0 : index + %wsb0 = index.mul %wsb0_0, %k_blocks : index + %wsb1_0 = index.add %ln1, %c0 : index + %wsb1 = index.mul %wsb1_0, %k_blocks : index + %wsb2_0 = index.add %ln2, %c0 : index + %wsb2 = index.mul %wsb2_0, %k_blocks : index + %wsb3_0 = index.add %ln3, %c0 : index + %wsb3 = index.mul %wsb3_0, %k_blocks : index + // Activation scale metadata is [tile][block][lane-half][register]. + %as_tile0 = index.div %lm0, %c16 : index + %asb0 = index.mul %as_tile0, %k_blocks : index + // Activation scale metadata is [tile][block][lane-half][register]. + %as_tile1 = index.div %lm1, %c16 : index + %asb1 = index.mul %as_tile1, %k_blocks : index + + %pf_init_row0 = index.add %col_base, %tid : index + %pf_init_row = index.assume %pf_init_row0 [range(%pf_init_row0, 0, 32767)] : index + %pf_init_plane_bytes = index.mul %cols_b, %pf_payload_bytes : index + %pf_init_chunk_base = index.mul %c0, %pf_init_plane_bytes : index + %pf_init_row_offset = index.mul %pf_init_row, %pf_payload_bytes : index + %pf_init_src0_0 = index.add %pf_init_chunk_base, %pf_init_row_offset : index + %pf_init_src0 = index.assume %pf_init_src0_0 [range(%pf_init_src0_0, 0, 536870896)] : index + %pf_init_active = index.cmp ult, %tid, %k_m : index + %pf_init_v0, %pf_init_v1 = scf.if %pf_init_active -> (vector<16xi8>, vector<16xi8>) { + %pf_init_l0 = vector.load %aq_flat[%pf_init_src0] : view<1073741824xi8> -> vector<16xi8> + %pf_init_src1_0 = index.add %pf_init_src0, %c16 : index + %pf_init_src1 = index.assume %pf_init_src1_0 [range(%pf_init_src1_0, 16, 536870912)] : index + %pf_init_l1 = vector.load %aq_flat[%pf_init_src1] : view<1073741824xi8> -> vector<16xi8> + scf.yield %pf_init_l0, %pf_init_l1 : vector<16xi8>, vector<16xi8> + } else { + scf.yield %pf_zero_v16i8, %pf_zero_v16i8 : vector<16xi8>, vector<16xi8> + } + %f0_0, %f0_1, %f0_2, %f0_3, %f1_0, %f1_1, %f1_2, %f1_3, %pf_unused_v0, %pf_unused_v1 = scf.for %chunk = [%c0 to %nchunks step %c1](%fc0_0 = %fzero : vector<8xf32>, %fc0_1 = %fzero : vector<8xf32>, %fc0_2 = %fzero : vector<8xf32>, %fc0_3 = %fzero : vector<8xf32>, %fc1_0 = %fzero : vector<8xf32>, %fc1_1 = %fzero : vector<8xf32>, %fc1_2 = %fzero : vector<8xf32>, %fc1_3 = %fzero : vector<8xf32>, %pf_carry_v0 = %pf_init_v0 : vector<16xi8>, %pf_carry_v1 = %pf_init_v1 : vector<16xi8>) -> (vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<16xi8>, vector<16xi8>) { + %chunk_k = index.mul %chunk, %k_kc : index + %chunk_blk = index.mul %chunk, %k_blocks : index + + %pf_store_active = index.cmp ult, %tid, %k_m : index + scf.if %pf_store_active { + %pf_dst0_0 = index.mul %tid, %k_astride : index + %pf_dst0 = index.assume %pf_dst0_0 [range(%pf_dst0_0, 0, 5080)] : index + vector.store %pf_carry_v0, %al_flat[%pf_dst0] : vector<16xi8>, view<5120xi8> + %pf_dst1_0 = index.add %pf_dst0, %c16 : index + %pf_dst1 = index.assume %pf_dst1_0 [range(%pf_dst1_0, 16, 5096)] : index + vector.store %pf_carry_v1, %al_flat[%pf_dst1] : vector<16xi8>, view<5120xi8> + } + // The transformed 144-byte block stores eight f16 scales in field 0 + // and eight signed-I4 K32 payloads in fields 1..8. Adjacent scale pairs + // are equal because this candidate quantizes both weights and activations + // with one K64 scale. + %q4_weight_lane_active = index.cmp ult, %tid, %k_n : index + scf.if %q4_weight_lane_active { + %q4_weight_tid0 = index.rem %tid, %k_n : index + %q4_weight_tid = index.assume %q4_weight_tid0 [range(%q4_weight_tid0, 0, 127)] : index + %q4_block_count = index.div %k_b, %q4_k_block : index + %q4_block = index.div %chunk, %c4 : index + %q4_pair = index.rem %chunk, %c4 : index + %q4_row0 = index.add %row_base, %q4_weight_tid : index + %q4_row = index.assume %q4_row0 [range(%q4_row0, 0, 262143)] : index + %q4_row_group_size = index.constant 64 : index + %q4_field_count = index.constant 9 : index + %q4_field_bytes = index.constant 16 : index + %q4_row_group = index.div %q4_row, %q4_row_group_size : index + %q4_row_lane = index.rem %q4_row, %q4_row_group_size : index + %q4_group_block0 = index.mul %q4_row_group, %q4_block_count : index + %q4_group_block = index.add %q4_group_block0, %q4_block : index + %q4_group_field_base = index.mul %q4_group_block, %q4_field_count : index + %q4_header_lane_base = index.mul %q4_group_field_base, %q4_row_group_size : index + %q4_header_lane = index.add %q4_header_lane_base, %q4_row_lane : index + %q4_header_byte_index = index.mul %q4_header_lane, %q4_field_bytes : index + %q4_scale_byte_stride = index.constant 4 : index + %q4_scale_byte_offset = index.mul %q4_pair, %q4_scale_byte_stride : index + %q4_scale_byte_index = index.add %q4_header_byte_index, %q4_scale_byte_offset : index + %q4_scale_byte_base = index.cast %q4_scale_byte_index : index to offset + %q4_scale_view = buffer.view %src0_na[%q4_scale_byte_base] : buffer -> view<2xf16> + %q4_scale_pair = vector.load %q4_scale_view[%c0] : view<2xf16> -> vector<2xf16> + %q4_scale_low_f16 = vector.extract %q4_scale_pair[0] : vector<2xf16> -> f16 + %q4_scale_high_f16 = vector.extract %q4_scale_pair[1] : vector<2xf16> -> f16 + %q4_d0 = scalar.extf %q4_scale_low_f16 : f16 to f32 + %q4_d1 = scalar.extf %q4_scale_high_f16 : f16 to f32 + %q4_low_group = index.mul %q4_pair, %c2 : index + %q4_high_group = index.add %q4_low_group, %c1 : index + %q4_low_field0 = index.add %q4_low_group, %c1 : index + %q4_high_field0 = index.add %q4_high_group, %c1 : index + %q4_low_field = index.add %q4_group_field_base, %q4_low_field0 : index + %q4_high_field = index.add %q4_group_field_base, %q4_high_field0 : index + %q4_low_lane_base = index.mul %q4_low_field, %q4_row_group_size : index + %q4_high_lane_base = index.mul %q4_high_field, %q4_row_group_size : index + %q4_low_lane = index.add %q4_low_lane_base, %q4_row_lane : index + %q4_high_lane = index.add %q4_high_lane_base, %q4_row_lane : index + %q4_low_byte_index = index.mul %q4_low_lane, %q4_field_bytes : index + %q4_high_byte_index = index.mul %q4_high_lane, %q4_field_bytes : index + %q4_low_byte_base = index.cast %q4_low_byte_index : index to offset + %q4_high_byte_base = index.cast %q4_high_byte_index : index to offset + %q4_low_view = buffer.view %src0_na[%q4_low_byte_base] : buffer -> view<4xi32> + %q4_high_view = buffer.view %src0_na[%q4_high_byte_base] : buffer -> view<4xi32> + %q4_low_words = vector.load %q4_low_view[%c0] : view<4xi32> -> vector<4xi32> + %q4_high_words = vector.load %q4_high_view[%c0] : view<4xi32> -> vector<4xi32> + %q4_low_packed = vector.bitcast %q4_low_words : vector<4xi32> to vector<16xi8> + %q4_high_packed = vector.bitcast %q4_high_words : vector<4xi32> to vector<16xi8> + %q4_weight_row_base = index.mul %q4_weight_tid, %k_wstride : index + %q4_w0 = index.assume %q4_weight_row_base [range(%q4_weight_row_base, 0, 5080)] : index + %q4_w1_0 = index.add %q4_weight_row_base, %c16 : index + %q4_w1 = index.assume %q4_w1_0 [range(%q4_w1_0, 16, 5096)] : index + vector.store %q4_low_packed, %wl_flat[%q4_w0] : vector<16xi8>, view<5120xi8> + vector.store %q4_high_packed, %wl_flat[%q4_w1] : vector<16xi8>, view<5120xi8> + %q4_meta0_0 = index.mul %q4_weight_tid, %c2 : index + %q4_meta0 = index.assume %q4_meta0_0 [range(%q4_meta0_0, 0, 254)] : index + %q4_meta1_0 = index.add %q4_meta0, %c1 : index + %q4_meta1 = index.assume %q4_meta1_0 [range(%q4_meta1_0, 1, 255)] : index + view.store %q4_d0, %wsl_view[%q4_meta0] : f32, view<256xf32> + view.store %q4_d1, %wsl_view[%q4_meta1] : f32, view<256xf32> + } + // Scale planes: 256 activation and 128 weight entries. + %sca0_slot0 = index.add %tid, %k_lanes0 : index + %sca0_slot = index.assume %sca0_slot0 [range(%sca0_slot0, 0, 255)] : index + %scin_a0 = index.cmp ult, %tid, %k_lanes1 : index + scf.if %scin_a0 { + %sc0_col = index.div %sca0_slot, %k_blocks : index + %sc0_blk = index.rem %sca0_slot, %k_blocks : index + %sca0_c = index.add %col_base, %sc0_col : index + %sca0_0 = index.add %chunk_blk, %sc0_blk : index + %sca0_1 = index.mul %sca0_0, %cols_b : index + %sca0_2 = index.add %sca0_1, %sca0_c : index + %sca0 = index.assume %sca0_2 [range(%sca0_2, 0, 33554431)] : index + %scav0 = view.load %as_flat[%sca0] : view<33554432xf32> -> f32 + %sca0_tile = index.div %sc0_col, %c16 : index + %sca0_inner = index.rem %sc0_col, %c16 : index + %sca0_par = index.rem %sca0_inner, %c2 : index + %sca0_v = index.div %sca0_inner, %c2 : index + %sca0_dst0 = index.mul %sca0_tile, %k_blocks : index + %sca0_dst1 = index.add %sca0_dst0, %sc0_blk : index + %sca0_dst2 = index.mul %sca0_dst1, %c2 : index + %sca0_dst3 = index.add %sca0_dst2, %sca0_par : index + %sca0_dst4 = index.mul %sca0_dst3, %c8 : index + %sca0_dst5 = index.add %sca0_dst4, %sca0_v : index + %sca0_dst = index.assume %sca0_dst5 [range(%sca0_dst5, 0, 255)] : index + view.store %scav0, %asl_view[%sca0_dst] : f32, view<256xf32> + } + %sca1_slot0 = index.add %tid, %k_lanes1 : index + %sca1_slot = index.assume %sca1_slot0 [range(%sca1_slot0, 128, 383)] : index + %scin_a1 = index.cmp ult, %sca1_slot, %k_asn : index + scf.if %scin_a1 { + %sc1_col = index.div %sca1_slot, %k_blocks : index + %sc1_blk = index.rem %sca1_slot, %k_blocks : index + %sca1_c = index.add %col_base, %sc1_col : index + %sca1_0 = index.add %chunk_blk, %sc1_blk : index + %sca1_1 = index.mul %sca1_0, %cols_b : index + %sca1_2 = index.add %sca1_1, %sca1_c : index + %sca1 = index.assume %sca1_2 [range(%sca1_2, 0, 33554431)] : index + %scav1 = view.load %as_flat[%sca1] : view<33554432xf32> -> f32 + %sca1_tile = index.div %sc1_col, %c16 : index + %sca1_inner = index.rem %sc1_col, %c16 : index + %sca1_par = index.rem %sca1_inner, %c2 : index + %sca1_v = index.div %sca1_inner, %c2 : index + %sca1_dst0 = index.mul %sca1_tile, %k_blocks : index + %sca1_dst1 = index.add %sca1_dst0, %sc1_blk : index + %sca1_dst2 = index.mul %sca1_dst1, %c2 : index + %sca1_dst3 = index.add %sca1_dst2, %sca1_par : index + %sca1_dst4 = index.mul %sca1_dst3, %c8 : index + %sca1_dst5 = index.add %sca1_dst4, %sca1_v : index + %sca1_dst = index.assume %sca1_dst5 [range(%sca1_dst5, 0, 255)] : index + view.store %scav1, %asl_view[%sca1_dst] : f32, view<256xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %pf_next0 = index.add %chunk, %c1 : index + %pf_has_next = index.cmp ult, %pf_next0, %nchunks : index + %pf_next_chunk = scf.if %pf_has_next -> (index) { + scf.yield %pf_next0 : index + } else { + scf.yield %chunk : index + } + + %pf_next_row0 = index.add %col_base, %tid : index + %pf_next_row = index.assume %pf_next_row0 [range(%pf_next_row0, 0, 32767)] : index + %pf_next_plane_bytes = index.mul %cols_b, %pf_payload_bytes : index + %pf_next_chunk_base = index.mul %pf_next_chunk, %pf_next_plane_bytes : index + %pf_next_row_offset = index.mul %pf_next_row, %pf_payload_bytes : index + %pf_next_src0_0 = index.add %pf_next_chunk_base, %pf_next_row_offset : index + %pf_next_src0 = index.assume %pf_next_src0_0 [range(%pf_next_src0_0, 0, 536870896)] : index + %pf_next_active = index.cmp ult, %tid, %k_m : index + %pf_next_v0, %pf_next_v1 = scf.if %pf_next_active -> (vector<16xi8>, vector<16xi8>) { + %pf_next_l0 = vector.load %aq_flat[%pf_next_src0] : view<1073741824xi8> -> vector<16xi8> + %pf_next_src1_0 = index.add %pf_next_src0, %c16 : index + %pf_next_src1 = index.assume %pf_next_src1_0 [range(%pf_next_src1_0, 16, 536870912)] : index + %pf_next_l1 = vector.load %aq_flat[%pf_next_src1] : view<1073741824xi8> -> vector<16xi8> + scf.yield %pf_next_l0, %pf_next_l1 : vector<16xi8>, vector<16xi8> + } else { + scf.yield %pf_zero_v16i8, %pf_zero_v16i8 : vector<16xi8>, vector<16xi8> + } + + // Chain each pair of K32 I4 products directly in I32. This keeps + // only one matrix result live before the shared K64 scale is applied. + %blk_k_u0 = index.mul %c0, %c16 : index + %blk_k16_u0 = index.add %blk_k_u0, %c8 : index + %iz_u0 = vector.fragment %izero shape [%c16, %c16] : vector<8xi32> + %wsi0_2_u0 = index.add %wsb0, %c0 : index + %wsi0_u0 = index.assume %wsi0_2_u0 [range(%wsi0_2_u0, 0, 255)] : index + %wsc0_u0 = view.load %wsl_view[%wsi0_u0] : view<256xf32> -> f32 + %wsv0_u0 = vector.splat %wsc0_u0 : vector<8xf32> + %wsi1_2_u0 = index.add %wsb1, %c0 : index + %wsi1_u0 = index.assume %wsi1_2_u0 [range(%wsi1_2_u0, 0, 255)] : index + %wsc1_u0 = view.load %wsl_view[%wsi1_u0] : view<256xf32> -> f32 + %wsv1_u0 = vector.splat %wsc1_u0 : vector<8xf32> + %wsi2_2_u0 = index.add %wsb2, %c0 : index + %wsi2_u0 = index.assume %wsi2_2_u0 [range(%wsi2_2_u0, 0, 255)] : index + %wsc2_u0 = view.load %wsl_view[%wsi2_u0] : view<256xf32> -> f32 + %wsv2_u0 = vector.splat %wsc2_u0 : vector<8xf32> + %wsi3_2_u0 = index.add %wsb3, %c0 : index + %wsi3_u0 = index.assume %wsi3_2_u0 [range(%wsi3_2_u0, 0, 255)] : index + %wsc3_u0 = view.load %wsl_view[%wsi3_u0] : view<256xf32> -> f32 + %wsv3_u0 = vector.splat %wsc3_u0 : vector<8xf32> + %blk_k_u1 = index.mul %c1, %c16 : index + %blk_k16_u1 = index.add %blk_k_u1, %c8 : index + %lf0_0_lane_row_u0 = index.add %lm0, %lane_lo : index + %lf0_0_row_u0 = index.mul %lf0_0_lane_row_u0, %k_astride : index + %lf0_0_base_u0 = index.add %lf0_0_row_u0, %blk_k_u0 : index + %lf0_0_idx_u0 = index.assume %lf0_0_base_u0 [range(%lf0_0_base_u0, 0, 5112)] : index + %lf0_0_raw_u0 = vector.load %al_flat[%lf0_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %lf0_0_words_u0 = vector.bitcast %lf0_0_raw_u0 : vector<8xi8> to vector<2xi32> + %lf0_0_u0 = vector.fragment %lf0_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %lf0_1_lane_row_u0 = index.add %lm0, %lane_lo : index + %lf0_1_row_u0 = index.mul %lf0_1_lane_row_u0, %k_astride : index + %lf0_1_base_u0 = index.add %lf0_1_row_u0, %blk_k16_u0 : index + %lf0_1_idx_u0 = index.assume %lf0_1_base_u0 [range(%lf0_1_base_u0, 0, 5112)] : index + %lf0_1_raw_u0 = vector.load %al_flat[%lf0_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %lf0_1_words_u0 = vector.bitcast %lf0_1_raw_u0 : vector<8xi8> to vector<2xi32> + %lf0_1_u0 = vector.fragment %lf0_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %asi0_0_u0 = index.add %asb0, %c0 : index + %asi0_1_u0 = index.mul %asi0_0_u0, %c2 : index + %asi0_2_u0 = index.add %asi0_1_u0, %lane_hi : index + %asi0_3_u0 = index.mul %asi0_2_u0, %c8 : index + %asi0_u0 = index.assume %asi0_3_u0 [range(%asi0_3_u0, 0, 120)] : index + %asv0_u0 = vector.load %asl_view[%asi0_u0] : view<256xf32> -> vector<8xf32> + %lf0_0_lane_row_u1 = index.add %lm0, %lane_lo : index + %lf0_0_row_u1 = index.mul %lf0_0_lane_row_u1, %k_astride : index + %lf0_0_base_u1 = index.add %lf0_0_row_u1, %blk_k_u1 : index + %lf0_0_idx_u1 = index.assume %lf0_0_base_u1 [range(%lf0_0_base_u1, 0, 5112)] : index + %lf0_0_raw_u1 = vector.load %al_flat[%lf0_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %lf0_0_words_u1 = vector.bitcast %lf0_0_raw_u1 : vector<8xi8> to vector<2xi32> + %lf0_0_u1 = vector.fragment %lf0_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %lf0_1_lane_row_u1 = index.add %lm0, %lane_lo : index + %lf0_1_row_u1 = index.mul %lf0_1_lane_row_u1, %k_astride : index + %lf0_1_base_u1 = index.add %lf0_1_row_u1, %blk_k16_u1 : index + %lf0_1_idx_u1 = index.assume %lf0_1_base_u1 [range(%lf0_1_base_u1, 0, 5112)] : index + %lf0_1_raw_u1 = vector.load %al_flat[%lf0_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %lf0_1_words_u1 = vector.bitcast %lf0_1_raw_u1 : vector<8xi8> to vector<2xi32> + %lf0_1_u1 = vector.fragment %lf0_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_0_0_row_u0 = index.mul %ln0, %k_wstride : index + %rfs0_0_0_base_u0 = index.add %rfs0_0_0_row_u0, %blk_k_u0 : index + %rfs0_0_0_idx_u0 = index.assume %rfs0_0_0_base_u0 [range(%rfs0_0_0_base_u0, 0, 5112)] : index + %rfs0_0_0_raw_u0 = vector.load %wl_flat[%rfs0_0_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs0_0_0_words_u0 = vector.bitcast %rfs0_0_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs0_0_0_u0 = vector.fragment %rfs0_0_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_0_1_row_u0 = index.mul %ln0, %k_wstride : index + %rfs0_0_1_base_u0 = index.add %rfs0_0_1_row_u0, %blk_k16_u0 : index + %rfs0_0_1_idx_u0 = index.assume %rfs0_0_1_base_u0 [range(%rfs0_0_1_base_u0, 0, 5112)] : index + %rfs0_0_1_raw_u0 = vector.load %wl_flat[%rfs0_0_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs0_0_1_words_u0 = vector.bitcast %rfs0_0_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs0_0_1_u0 = vector.fragment %rfs0_0_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_0_u0 = vector.mma %lf0_0_u0, %rfs0_0_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_0_u0 = vector.mma %lf0_1_u0, %rfs0_0_1_u0, %i0_0_0_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %rfs0_0_0_row_u1 = index.mul %ln0, %k_wstride : index + %rfs0_0_0_base_u1 = index.add %rfs0_0_0_row_u1, %blk_k_u1 : index + %rfs0_0_0_idx_u1 = index.assume %rfs0_0_0_base_u1 [range(%rfs0_0_0_base_u1, 0, 5112)] : index + %rfs0_0_0_raw_u1 = vector.load %wl_flat[%rfs0_0_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs0_0_0_words_u1 = vector.bitcast %rfs0_0_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs0_0_0_u1 = vector.fragment %rfs0_0_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_0_1_row_u1 = index.mul %ln0, %k_wstride : index + %rfs0_0_1_base_u1 = index.add %rfs0_0_1_row_u1, %blk_k16_u1 : index + %rfs0_0_1_idx_u1 = index.assume %rfs0_0_1_base_u1 [range(%rfs0_0_1_base_u1, 0, 5112)] : index + %rfs0_0_1_raw_u1 = vector.load %wl_flat[%rfs0_0_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs0_0_1_words_u1 = vector.bitcast %rfs0_0_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs0_0_1_u1 = vector.fragment %rfs0_0_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_0_u1 = vector.mma %lf0_0_u1, %rfs0_0_0_u1, %i1_0_0_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_0_u1 = vector.mma %lf0_1_u1, %rfs0_0_1_u1, %i0_0_0_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff0_0_k64 = vector.sitofp %i1_0_0_u1 : vector<8xi32> to vector<8xf32> + %sv0_0_k64 = vector.mulf %asv0_u0, %wsv0_u0 : vector<8xf32> + %fn0_0 = vector.fmaf %ff0_0_k64, %sv0_0_k64, %fc0_0 : vector<8xf32> + scf.schedule.fence + %rfs0_1_0_row_u0 = index.mul %ln1, %k_wstride : index + %rfs0_1_0_base_u0 = index.add %rfs0_1_0_row_u0, %blk_k_u0 : index + %rfs0_1_0_idx_u0 = index.assume %rfs0_1_0_base_u0 [range(%rfs0_1_0_base_u0, 0, 5112)] : index + %rfs0_1_0_raw_u0 = vector.load %wl_flat[%rfs0_1_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs0_1_0_words_u0 = vector.bitcast %rfs0_1_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs0_1_0_u0 = vector.fragment %rfs0_1_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_1_1_row_u0 = index.mul %ln1, %k_wstride : index + %rfs0_1_1_base_u0 = index.add %rfs0_1_1_row_u0, %blk_k16_u0 : index + %rfs0_1_1_idx_u0 = index.assume %rfs0_1_1_base_u0 [range(%rfs0_1_1_base_u0, 0, 5112)] : index + %rfs0_1_1_raw_u0 = vector.load %wl_flat[%rfs0_1_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs0_1_1_words_u0 = vector.bitcast %rfs0_1_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs0_1_1_u0 = vector.fragment %rfs0_1_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_1_u0 = vector.mma %lf0_0_u0, %rfs0_1_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_1_u0 = vector.mma %lf0_1_u0, %rfs0_1_1_u0, %i0_0_1_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %rfs0_1_0_row_u1 = index.mul %ln1, %k_wstride : index + %rfs0_1_0_base_u1 = index.add %rfs0_1_0_row_u1, %blk_k_u1 : index + %rfs0_1_0_idx_u1 = index.assume %rfs0_1_0_base_u1 [range(%rfs0_1_0_base_u1, 0, 5112)] : index + %rfs0_1_0_raw_u1 = vector.load %wl_flat[%rfs0_1_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs0_1_0_words_u1 = vector.bitcast %rfs0_1_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs0_1_0_u1 = vector.fragment %rfs0_1_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_1_1_row_u1 = index.mul %ln1, %k_wstride : index + %rfs0_1_1_base_u1 = index.add %rfs0_1_1_row_u1, %blk_k16_u1 : index + %rfs0_1_1_idx_u1 = index.assume %rfs0_1_1_base_u1 [range(%rfs0_1_1_base_u1, 0, 5112)] : index + %rfs0_1_1_raw_u1 = vector.load %wl_flat[%rfs0_1_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs0_1_1_words_u1 = vector.bitcast %rfs0_1_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs0_1_1_u1 = vector.fragment %rfs0_1_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_1_u1 = vector.mma %lf0_0_u1, %rfs0_1_0_u1, %i1_0_1_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_1_u1 = vector.mma %lf0_1_u1, %rfs0_1_1_u1, %i0_0_1_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff0_1_k64 = vector.sitofp %i1_0_1_u1 : vector<8xi32> to vector<8xf32> + %sv0_1_k64 = vector.mulf %asv0_u0, %wsv1_u0 : vector<8xf32> + %fn0_1 = vector.fmaf %ff0_1_k64, %sv0_1_k64, %fc0_1 : vector<8xf32> + scf.schedule.fence + %rfs0_2_0_row_u0 = index.mul %ln2, %k_wstride : index + %rfs0_2_0_base_u0 = index.add %rfs0_2_0_row_u0, %blk_k_u0 : index + %rfs0_2_0_idx_u0 = index.assume %rfs0_2_0_base_u0 [range(%rfs0_2_0_base_u0, 0, 5112)] : index + %rfs0_2_0_raw_u0 = vector.load %wl_flat[%rfs0_2_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs0_2_0_words_u0 = vector.bitcast %rfs0_2_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs0_2_0_u0 = vector.fragment %rfs0_2_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_2_1_row_u0 = index.mul %ln2, %k_wstride : index + %rfs0_2_1_base_u0 = index.add %rfs0_2_1_row_u0, %blk_k16_u0 : index + %rfs0_2_1_idx_u0 = index.assume %rfs0_2_1_base_u0 [range(%rfs0_2_1_base_u0, 0, 5112)] : index + %rfs0_2_1_raw_u0 = vector.load %wl_flat[%rfs0_2_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs0_2_1_words_u0 = vector.bitcast %rfs0_2_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs0_2_1_u0 = vector.fragment %rfs0_2_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_2_u0 = vector.mma %lf0_0_u0, %rfs0_2_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_2_u0 = vector.mma %lf0_1_u0, %rfs0_2_1_u0, %i0_0_2_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %rfs0_2_0_row_u1 = index.mul %ln2, %k_wstride : index + %rfs0_2_0_base_u1 = index.add %rfs0_2_0_row_u1, %blk_k_u1 : index + %rfs0_2_0_idx_u1 = index.assume %rfs0_2_0_base_u1 [range(%rfs0_2_0_base_u1, 0, 5112)] : index + %rfs0_2_0_raw_u1 = vector.load %wl_flat[%rfs0_2_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs0_2_0_words_u1 = vector.bitcast %rfs0_2_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs0_2_0_u1 = vector.fragment %rfs0_2_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_2_1_row_u1 = index.mul %ln2, %k_wstride : index + %rfs0_2_1_base_u1 = index.add %rfs0_2_1_row_u1, %blk_k16_u1 : index + %rfs0_2_1_idx_u1 = index.assume %rfs0_2_1_base_u1 [range(%rfs0_2_1_base_u1, 0, 5112)] : index + %rfs0_2_1_raw_u1 = vector.load %wl_flat[%rfs0_2_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs0_2_1_words_u1 = vector.bitcast %rfs0_2_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs0_2_1_u1 = vector.fragment %rfs0_2_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_2_u1 = vector.mma %lf0_0_u1, %rfs0_2_0_u1, %i1_0_2_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_2_u1 = vector.mma %lf0_1_u1, %rfs0_2_1_u1, %i0_0_2_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff0_2_k64 = vector.sitofp %i1_0_2_u1 : vector<8xi32> to vector<8xf32> + %sv0_2_k64 = vector.mulf %asv0_u0, %wsv2_u0 : vector<8xf32> + %fn0_2 = vector.fmaf %ff0_2_k64, %sv0_2_k64, %fc0_2 : vector<8xf32> + scf.schedule.fence + %rfs0_3_0_row_u0 = index.mul %ln3, %k_wstride : index + %rfs0_3_0_base_u0 = index.add %rfs0_3_0_row_u0, %blk_k_u0 : index + %rfs0_3_0_idx_u0 = index.assume %rfs0_3_0_base_u0 [range(%rfs0_3_0_base_u0, 0, 5112)] : index + %rfs0_3_0_raw_u0 = vector.load %wl_flat[%rfs0_3_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs0_3_0_words_u0 = vector.bitcast %rfs0_3_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs0_3_0_u0 = vector.fragment %rfs0_3_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_3_1_row_u0 = index.mul %ln3, %k_wstride : index + %rfs0_3_1_base_u0 = index.add %rfs0_3_1_row_u0, %blk_k16_u0 : index + %rfs0_3_1_idx_u0 = index.assume %rfs0_3_1_base_u0 [range(%rfs0_3_1_base_u0, 0, 5112)] : index + %rfs0_3_1_raw_u0 = vector.load %wl_flat[%rfs0_3_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs0_3_1_words_u0 = vector.bitcast %rfs0_3_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs0_3_1_u0 = vector.fragment %rfs0_3_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_3_u0 = vector.mma %lf0_0_u0, %rfs0_3_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_3_u0 = vector.mma %lf0_1_u0, %rfs0_3_1_u0, %i0_0_3_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %rfs0_3_0_row_u1 = index.mul %ln3, %k_wstride : index + %rfs0_3_0_base_u1 = index.add %rfs0_3_0_row_u1, %blk_k_u1 : index + %rfs0_3_0_idx_u1 = index.assume %rfs0_3_0_base_u1 [range(%rfs0_3_0_base_u1, 0, 5112)] : index + %rfs0_3_0_raw_u1 = vector.load %wl_flat[%rfs0_3_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs0_3_0_words_u1 = vector.bitcast %rfs0_3_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs0_3_0_u1 = vector.fragment %rfs0_3_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_3_1_row_u1 = index.mul %ln3, %k_wstride : index + %rfs0_3_1_base_u1 = index.add %rfs0_3_1_row_u1, %blk_k16_u1 : index + %rfs0_3_1_idx_u1 = index.assume %rfs0_3_1_base_u1 [range(%rfs0_3_1_base_u1, 0, 5112)] : index + %rfs0_3_1_raw_u1 = vector.load %wl_flat[%rfs0_3_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs0_3_1_words_u1 = vector.bitcast %rfs0_3_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs0_3_1_u1 = vector.fragment %rfs0_3_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_3_u1 = vector.mma %lf0_0_u1, %rfs0_3_0_u1, %i1_0_3_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_3_u1 = vector.mma %lf0_1_u1, %rfs0_3_1_u1, %i0_0_3_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff0_3_k64 = vector.sitofp %i1_0_3_u1 : vector<8xi32> to vector<8xf32> + %sv0_3_k64 = vector.mulf %asv0_u0, %wsv3_u0 : vector<8xf32> + %fn0_3 = vector.fmaf %ff0_3_k64, %sv0_3_k64, %fc0_3 : vector<8xf32> + scf.schedule.fence + %lf1_0_lane_row_u0 = index.add %lm1, %lane_lo : index + %lf1_0_row_u0 = index.mul %lf1_0_lane_row_u0, %k_astride : index + %lf1_0_base_u0 = index.add %lf1_0_row_u0, %blk_k_u0 : index + %lf1_0_idx_u0 = index.assume %lf1_0_base_u0 [range(%lf1_0_base_u0, 0, 5112)] : index + %lf1_0_raw_u0 = vector.load %al_flat[%lf1_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %lf1_0_words_u0 = vector.bitcast %lf1_0_raw_u0 : vector<8xi8> to vector<2xi32> + %lf1_0_u0 = vector.fragment %lf1_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %lf1_1_lane_row_u0 = index.add %lm1, %lane_lo : index + %lf1_1_row_u0 = index.mul %lf1_1_lane_row_u0, %k_astride : index + %lf1_1_base_u0 = index.add %lf1_1_row_u0, %blk_k16_u0 : index + %lf1_1_idx_u0 = index.assume %lf1_1_base_u0 [range(%lf1_1_base_u0, 0, 5112)] : index + %lf1_1_raw_u0 = vector.load %al_flat[%lf1_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %lf1_1_words_u0 = vector.bitcast %lf1_1_raw_u0 : vector<8xi8> to vector<2xi32> + %lf1_1_u0 = vector.fragment %lf1_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %asi1_0_u0 = index.add %asb1, %c0 : index + %asi1_1_u0 = index.mul %asi1_0_u0, %c2 : index + %asi1_2_u0 = index.add %asi1_1_u0, %lane_hi : index + %asi1_3_u0 = index.mul %asi1_2_u0, %c8 : index + %asi1_u0 = index.assume %asi1_3_u0 [range(%asi1_3_u0, 0, 120)] : index + %asv1_u0 = vector.load %asl_view[%asi1_u0] : view<256xf32> -> vector<8xf32> + %lf1_0_lane_row_u1 = index.add %lm1, %lane_lo : index + %lf1_0_row_u1 = index.mul %lf1_0_lane_row_u1, %k_astride : index + %lf1_0_base_u1 = index.add %lf1_0_row_u1, %blk_k_u1 : index + %lf1_0_idx_u1 = index.assume %lf1_0_base_u1 [range(%lf1_0_base_u1, 0, 5112)] : index + %lf1_0_raw_u1 = vector.load %al_flat[%lf1_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %lf1_0_words_u1 = vector.bitcast %lf1_0_raw_u1 : vector<8xi8> to vector<2xi32> + %lf1_0_u1 = vector.fragment %lf1_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %lf1_1_lane_row_u1 = index.add %lm1, %lane_lo : index + %lf1_1_row_u1 = index.mul %lf1_1_lane_row_u1, %k_astride : index + %lf1_1_base_u1 = index.add %lf1_1_row_u1, %blk_k16_u1 : index + %lf1_1_idx_u1 = index.assume %lf1_1_base_u1 [range(%lf1_1_base_u1, 0, 5112)] : index + %lf1_1_raw_u1 = vector.load %al_flat[%lf1_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %lf1_1_words_u1 = vector.bitcast %lf1_1_raw_u1 : vector<8xi8> to vector<2xi32> + %lf1_1_u1 = vector.fragment %lf1_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs1_0_0_row_u0 = index.mul %ln0, %k_wstride : index + %rfs1_0_0_base_u0 = index.add %rfs1_0_0_row_u0, %blk_k_u0 : index + %rfs1_0_0_idx_u0 = index.assume %rfs1_0_0_base_u0 [range(%rfs1_0_0_base_u0, 0, 5112)] : index + %rfs1_0_0_raw_u0 = vector.load %wl_flat[%rfs1_0_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs1_0_0_words_u0 = vector.bitcast %rfs1_0_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs1_0_0_u0 = vector.fragment %rfs1_0_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs1_0_1_row_u0 = index.mul %ln0, %k_wstride : index + %rfs1_0_1_base_u0 = index.add %rfs1_0_1_row_u0, %blk_k16_u0 : index + %rfs1_0_1_idx_u0 = index.assume %rfs1_0_1_base_u0 [range(%rfs1_0_1_base_u0, 0, 5112)] : index + %rfs1_0_1_raw_u0 = vector.load %wl_flat[%rfs1_0_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs1_0_1_words_u0 = vector.bitcast %rfs1_0_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs1_0_1_u0 = vector.fragment %rfs1_0_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_1_0_u0 = vector.mma %lf1_0_u0, %rfs1_0_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_1_0_u0 = vector.mma %lf1_1_u0, %rfs1_0_1_u0, %i0_1_0_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %rfs1_0_0_row_u1 = index.mul %ln0, %k_wstride : index + %rfs1_0_0_base_u1 = index.add %rfs1_0_0_row_u1, %blk_k_u1 : index + %rfs1_0_0_idx_u1 = index.assume %rfs1_0_0_base_u1 [range(%rfs1_0_0_base_u1, 0, 5112)] : index + %rfs1_0_0_raw_u1 = vector.load %wl_flat[%rfs1_0_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs1_0_0_words_u1 = vector.bitcast %rfs1_0_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs1_0_0_u1 = vector.fragment %rfs1_0_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs1_0_1_row_u1 = index.mul %ln0, %k_wstride : index + %rfs1_0_1_base_u1 = index.add %rfs1_0_1_row_u1, %blk_k16_u1 : index + %rfs1_0_1_idx_u1 = index.assume %rfs1_0_1_base_u1 [range(%rfs1_0_1_base_u1, 0, 5112)] : index + %rfs1_0_1_raw_u1 = vector.load %wl_flat[%rfs1_0_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs1_0_1_words_u1 = vector.bitcast %rfs1_0_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs1_0_1_u1 = vector.fragment %rfs1_0_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_1_0_u1 = vector.mma %lf1_0_u1, %rfs1_0_0_u1, %i1_1_0_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_1_0_u1 = vector.mma %lf1_1_u1, %rfs1_0_1_u1, %i0_1_0_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff1_0_k64 = vector.sitofp %i1_1_0_u1 : vector<8xi32> to vector<8xf32> + %sv1_0_k64 = vector.mulf %asv1_u0, %wsv0_u0 : vector<8xf32> + %fn1_0 = vector.fmaf %ff1_0_k64, %sv1_0_k64, %fc1_0 : vector<8xf32> + scf.schedule.fence + %rfs1_1_0_row_u0 = index.mul %ln1, %k_wstride : index + %rfs1_1_0_base_u0 = index.add %rfs1_1_0_row_u0, %blk_k_u0 : index + %rfs1_1_0_idx_u0 = index.assume %rfs1_1_0_base_u0 [range(%rfs1_1_0_base_u0, 0, 5112)] : index + %rfs1_1_0_raw_u0 = vector.load %wl_flat[%rfs1_1_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs1_1_0_words_u0 = vector.bitcast %rfs1_1_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs1_1_0_u0 = vector.fragment %rfs1_1_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs1_1_1_row_u0 = index.mul %ln1, %k_wstride : index + %rfs1_1_1_base_u0 = index.add %rfs1_1_1_row_u0, %blk_k16_u0 : index + %rfs1_1_1_idx_u0 = index.assume %rfs1_1_1_base_u0 [range(%rfs1_1_1_base_u0, 0, 5112)] : index + %rfs1_1_1_raw_u0 = vector.load %wl_flat[%rfs1_1_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs1_1_1_words_u0 = vector.bitcast %rfs1_1_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs1_1_1_u0 = vector.fragment %rfs1_1_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_1_1_u0 = vector.mma %lf1_0_u0, %rfs1_1_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_1_1_u0 = vector.mma %lf1_1_u0, %rfs1_1_1_u0, %i0_1_1_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %rfs1_1_0_row_u1 = index.mul %ln1, %k_wstride : index + %rfs1_1_0_base_u1 = index.add %rfs1_1_0_row_u1, %blk_k_u1 : index + %rfs1_1_0_idx_u1 = index.assume %rfs1_1_0_base_u1 [range(%rfs1_1_0_base_u1, 0, 5112)] : index + %rfs1_1_0_raw_u1 = vector.load %wl_flat[%rfs1_1_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs1_1_0_words_u1 = vector.bitcast %rfs1_1_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs1_1_0_u1 = vector.fragment %rfs1_1_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs1_1_1_row_u1 = index.mul %ln1, %k_wstride : index + %rfs1_1_1_base_u1 = index.add %rfs1_1_1_row_u1, %blk_k16_u1 : index + %rfs1_1_1_idx_u1 = index.assume %rfs1_1_1_base_u1 [range(%rfs1_1_1_base_u1, 0, 5112)] : index + %rfs1_1_1_raw_u1 = vector.load %wl_flat[%rfs1_1_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs1_1_1_words_u1 = vector.bitcast %rfs1_1_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs1_1_1_u1 = vector.fragment %rfs1_1_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_1_1_u1 = vector.mma %lf1_0_u1, %rfs1_1_0_u1, %i1_1_1_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_1_1_u1 = vector.mma %lf1_1_u1, %rfs1_1_1_u1, %i0_1_1_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff1_1_k64 = vector.sitofp %i1_1_1_u1 : vector<8xi32> to vector<8xf32> + %sv1_1_k64 = vector.mulf %asv1_u0, %wsv1_u0 : vector<8xf32> + %fn1_1 = vector.fmaf %ff1_1_k64, %sv1_1_k64, %fc1_1 : vector<8xf32> + scf.schedule.fence + %rfs1_2_0_row_u0 = index.mul %ln2, %k_wstride : index + %rfs1_2_0_base_u0 = index.add %rfs1_2_0_row_u0, %blk_k_u0 : index + %rfs1_2_0_idx_u0 = index.assume %rfs1_2_0_base_u0 [range(%rfs1_2_0_base_u0, 0, 5112)] : index + %rfs1_2_0_raw_u0 = vector.load %wl_flat[%rfs1_2_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs1_2_0_words_u0 = vector.bitcast %rfs1_2_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs1_2_0_u0 = vector.fragment %rfs1_2_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs1_2_1_row_u0 = index.mul %ln2, %k_wstride : index + %rfs1_2_1_base_u0 = index.add %rfs1_2_1_row_u0, %blk_k16_u0 : index + %rfs1_2_1_idx_u0 = index.assume %rfs1_2_1_base_u0 [range(%rfs1_2_1_base_u0, 0, 5112)] : index + %rfs1_2_1_raw_u0 = vector.load %wl_flat[%rfs1_2_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs1_2_1_words_u0 = vector.bitcast %rfs1_2_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs1_2_1_u0 = vector.fragment %rfs1_2_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_1_2_u0 = vector.mma %lf1_0_u0, %rfs1_2_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_1_2_u0 = vector.mma %lf1_1_u0, %rfs1_2_1_u0, %i0_1_2_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %rfs1_2_0_row_u1 = index.mul %ln2, %k_wstride : index + %rfs1_2_0_base_u1 = index.add %rfs1_2_0_row_u1, %blk_k_u1 : index + %rfs1_2_0_idx_u1 = index.assume %rfs1_2_0_base_u1 [range(%rfs1_2_0_base_u1, 0, 5112)] : index + %rfs1_2_0_raw_u1 = vector.load %wl_flat[%rfs1_2_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs1_2_0_words_u1 = vector.bitcast %rfs1_2_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs1_2_0_u1 = vector.fragment %rfs1_2_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs1_2_1_row_u1 = index.mul %ln2, %k_wstride : index + %rfs1_2_1_base_u1 = index.add %rfs1_2_1_row_u1, %blk_k16_u1 : index + %rfs1_2_1_idx_u1 = index.assume %rfs1_2_1_base_u1 [range(%rfs1_2_1_base_u1, 0, 5112)] : index + %rfs1_2_1_raw_u1 = vector.load %wl_flat[%rfs1_2_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs1_2_1_words_u1 = vector.bitcast %rfs1_2_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs1_2_1_u1 = vector.fragment %rfs1_2_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_1_2_u1 = vector.mma %lf1_0_u1, %rfs1_2_0_u1, %i1_1_2_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_1_2_u1 = vector.mma %lf1_1_u1, %rfs1_2_1_u1, %i0_1_2_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff1_2_k64 = vector.sitofp %i1_1_2_u1 : vector<8xi32> to vector<8xf32> + %sv1_2_k64 = vector.mulf %asv1_u0, %wsv2_u0 : vector<8xf32> + %fn1_2 = vector.fmaf %ff1_2_k64, %sv1_2_k64, %fc1_2 : vector<8xf32> + scf.schedule.fence + %rfs1_3_0_row_u0 = index.mul %ln3, %k_wstride : index + %rfs1_3_0_base_u0 = index.add %rfs1_3_0_row_u0, %blk_k_u0 : index + %rfs1_3_0_idx_u0 = index.assume %rfs1_3_0_base_u0 [range(%rfs1_3_0_base_u0, 0, 5112)] : index + %rfs1_3_0_raw_u0 = vector.load %wl_flat[%rfs1_3_0_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs1_3_0_words_u0 = vector.bitcast %rfs1_3_0_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs1_3_0_u0 = vector.fragment %rfs1_3_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs1_3_1_row_u0 = index.mul %ln3, %k_wstride : index + %rfs1_3_1_base_u0 = index.add %rfs1_3_1_row_u0, %blk_k16_u0 : index + %rfs1_3_1_idx_u0 = index.assume %rfs1_3_1_base_u0 [range(%rfs1_3_1_base_u0, 0, 5112)] : index + %rfs1_3_1_raw_u0 = vector.load %wl_flat[%rfs1_3_1_idx_u0] : view<5120xi8> -> vector<8xi8> + %rfs1_3_1_words_u0 = vector.bitcast %rfs1_3_1_raw_u0 : vector<8xi8> to vector<2xi32> + %rfs1_3_1_u0 = vector.fragment %rfs1_3_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_1_3_u0 = vector.mma %lf1_0_u0, %rfs1_3_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_1_3_u0 = vector.mma %lf1_1_u0, %rfs1_3_1_u0, %i0_1_3_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %rfs1_3_0_row_u1 = index.mul %ln3, %k_wstride : index + %rfs1_3_0_base_u1 = index.add %rfs1_3_0_row_u1, %blk_k_u1 : index + %rfs1_3_0_idx_u1 = index.assume %rfs1_3_0_base_u1 [range(%rfs1_3_0_base_u1, 0, 5112)] : index + %rfs1_3_0_raw_u1 = vector.load %wl_flat[%rfs1_3_0_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs1_3_0_words_u1 = vector.bitcast %rfs1_3_0_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs1_3_0_u1 = vector.fragment %rfs1_3_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs1_3_1_row_u1 = index.mul %ln3, %k_wstride : index + %rfs1_3_1_base_u1 = index.add %rfs1_3_1_row_u1, %blk_k16_u1 : index + %rfs1_3_1_idx_u1 = index.assume %rfs1_3_1_base_u1 [range(%rfs1_3_1_base_u1, 0, 5112)] : index + %rfs1_3_1_raw_u1 = vector.load %wl_flat[%rfs1_3_1_idx_u1] : view<5120xi8> -> vector<8xi8> + %rfs1_3_1_words_u1 = vector.bitcast %rfs1_3_1_raw_u1 : vector<8xi8> to vector<2xi32> + %rfs1_3_1_u1 = vector.fragment %rfs1_3_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_1_3_u1 = vector.mma %lf1_0_u1, %rfs1_3_0_u1, %i1_1_3_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_1_3_u1 = vector.mma %lf1_1_u1, %rfs1_3_1_u1, %i0_1_3_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff1_3_k64 = vector.sitofp %i1_1_3_u1 : vector<8xi32> to vector<8xf32> + %sv1_3_k64 = vector.mulf %asv1_u0, %wsv3_u0 : vector<8xf32> + %fn1_3 = vector.fmaf %ff1_3_k64, %sv1_3_k64, %fc1_3 : vector<8xf32> + scf.schedule.fence + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %fn0_0, %fn0_1, %fn0_2, %fn0_3, %fn1_0, %fn1_1, %fn1_2, %fn1_3, %pf_next_v0, %pf_next_v1 : vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<16xi8>, vector<16xi8> + } + + %out0_0 = scf.if %apply_swiglu -> (vector<8xf32>) { + %gate0_0 = vector.fragment.load %gate_view[%gm0, %gn0] shape [%c16, %c16] : view<[%cols_b]x[%rows_b]xf32> -> vector<8xf32> + %silu0_0 = vector.siluf %gate0_0 : vector<8xf32> + %result0_0 = vector.mulf %silu0_0, %f0_0 : vector<8xf32> + scf.yield %result0_0 : vector<8xf32> + } else { + scf.yield %f0_0 : vector<8xf32> + } + %out0_1 = scf.if %apply_swiglu -> (vector<8xf32>) { + %gate0_1 = vector.fragment.load %gate_view[%gm0, %gn1] shape [%c16, %c16] : view<[%cols_b]x[%rows_b]xf32> -> vector<8xf32> + %silu0_1 = vector.siluf %gate0_1 : vector<8xf32> + %result0_1 = vector.mulf %silu0_1, %f0_1 : vector<8xf32> + scf.yield %result0_1 : vector<8xf32> + } else { + scf.yield %f0_1 : vector<8xf32> + } + %out0_2 = scf.if %apply_swiglu -> (vector<8xf32>) { + %gate0_2 = vector.fragment.load %gate_view[%gm0, %gn2] shape [%c16, %c16] : view<[%cols_b]x[%rows_b]xf32> -> vector<8xf32> + %silu0_2 = vector.siluf %gate0_2 : vector<8xf32> + %result0_2 = vector.mulf %silu0_2, %f0_2 : vector<8xf32> + scf.yield %result0_2 : vector<8xf32> + } else { + scf.yield %f0_2 : vector<8xf32> + } + %out0_3 = scf.if %apply_swiglu -> (vector<8xf32>) { + %gate0_3 = vector.fragment.load %gate_view[%gm0, %gn3] shape [%c16, %c16] : view<[%cols_b]x[%rows_b]xf32> -> vector<8xf32> + %silu0_3 = vector.siluf %gate0_3 : vector<8xf32> + %result0_3 = vector.mulf %silu0_3, %f0_3 : vector<8xf32> + scf.yield %result0_3 : vector<8xf32> + } else { + scf.yield %f0_3 : vector<8xf32> + } + %out1_0 = scf.if %apply_swiglu -> (vector<8xf32>) { + %gate1_0 = vector.fragment.load %gate_view[%gm1, %gn0] shape [%c16, %c16] : view<[%cols_b]x[%rows_b]xf32> -> vector<8xf32> + %silu1_0 = vector.siluf %gate1_0 : vector<8xf32> + %result1_0 = vector.mulf %silu1_0, %f1_0 : vector<8xf32> + scf.yield %result1_0 : vector<8xf32> + } else { + scf.yield %f1_0 : vector<8xf32> + } + %out1_1 = scf.if %apply_swiglu -> (vector<8xf32>) { + %gate1_1 = vector.fragment.load %gate_view[%gm1, %gn1] shape [%c16, %c16] : view<[%cols_b]x[%rows_b]xf32> -> vector<8xf32> + %silu1_1 = vector.siluf %gate1_1 : vector<8xf32> + %result1_1 = vector.mulf %silu1_1, %f1_1 : vector<8xf32> + scf.yield %result1_1 : vector<8xf32> + } else { + scf.yield %f1_1 : vector<8xf32> + } + %out1_2 = scf.if %apply_swiglu -> (vector<8xf32>) { + %gate1_2 = vector.fragment.load %gate_view[%gm1, %gn2] shape [%c16, %c16] : view<[%cols_b]x[%rows_b]xf32> -> vector<8xf32> + %silu1_2 = vector.siluf %gate1_2 : vector<8xf32> + %result1_2 = vector.mulf %silu1_2, %f1_2 : vector<8xf32> + scf.yield %result1_2 : vector<8xf32> + } else { + scf.yield %f1_2 : vector<8xf32> + } + %out1_3 = scf.if %apply_swiglu -> (vector<8xf32>) { + %gate1_3 = vector.fragment.load %gate_view[%gm1, %gn3] shape [%c16, %c16] : view<[%cols_b]x[%rows_b]xf32> -> vector<8xf32> + %silu1_3 = vector.siluf %gate1_3 : vector<8xf32> + %result1_3 = vector.mulf %silu1_3, %f1_3 : vector<8xf32> + scf.yield %result1_3 : vector<8xf32> + } else { + scf.yield %f1_3 : vector<8xf32> + } + + scf.if %publish_f32 { + vector.fragment.store %out0_0, %dst_view[%gm0, %gn0] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out0_1, %dst_view[%gm0, %gn1] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out0_2, %dst_view[%gm0, %gn2] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out0_3, %dst_view[%gm0, %gn3] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out1_0, %dst_view[%gm1, %gn0] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out1_1, %dst_view[%gm1, %gn1] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out1_2, %dst_view[%gm1, %gn2] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out1_3, %dst_view[%gm1, %gn3] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + } + + scf.if %publish_f16 { + %f16_tile = buffer.view %qout_scratch[%base] : buffer -> view<128x32xf32> + %f16_c4 = index.constant 4 : index + %f16_c8 = index.constant 8 : index + %f16_c32 = index.constant 32 : index + %f16_c512 = index.constant 512 : index + scf.for %f16_phase = [%c0 to %c2 step %c1] { + %f16_first_half = index.cmp eq, %f16_phase, %c0 : index + scf.if %f16_first_half { + vector.fragment.store %out0_0, %f16_tile[%lm0, %c0] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out0_1, %f16_tile[%lm0, %c16] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out1_0, %f16_tile[%lm1, %c0] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out1_1, %f16_tile[%lm1, %c16] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + } else { + vector.fragment.store %out0_2, %f16_tile[%lm0, %c0] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out0_3, %f16_tile[%lm0, %c16] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out1_2, %f16_tile[%lm1, %c0] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out1_3, %f16_tile[%lm1, %c16] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.for %f16_batch = [%c0 to %f16_c8 step %c1] { + %f16_batch_base = index.mul %f16_batch, %f16_c512 : index + %f16_thread_base = index.mul %tid, %f16_c4 : index + %f16_linear0 = index.add %f16_batch_base, %f16_thread_base : index + %f16_linear = index.assume %f16_linear0 [range(%f16_linear0, 0, 4092), mul(%f16_linear0, 4)] : index + %f16_token_local = index.div %f16_linear, %f16_c32 : index + %f16_row_local = index.rem %f16_linear, %f16_c32 : index + %f16_wide = vector.load %f16_tile[%f16_token_local, %f16_row_local] : view<128x32xf32> -> vector<4xf32> + %f16_narrow = vector.fptrunc %f16_wide : vector<4xf32> to vector<4xf16> + %f16_token0 = index.add %col_base, %f16_token_local : index + %f16_phase_row = index.mul %f16_phase, %f16_c32 : index + %f16_row0 = index.add %row_base, %f16_phase_row : index + %f16_row1 = index.add %f16_row0, %f16_row_local : index + %f16_row_end0 = index.add %f16_row1, %f16_c4 : index + %f16_token, %f16_row, %f16_row_end, %f16_cols, %f16_rows = index.assume %f16_token0, %f16_row1, %f16_row_end0, %cols_b, %rows_b [lt(%f16_token0, %cols_b), le(%f16_row_end0, %rows_b)] : index, index, index, index, index + vector.store %f16_narrow, %dst_f16_view[%f16_token, %f16_row] : vector<4xf16>, view<[%cols_b]x[%rows_b]xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + } + } + + // The following down projection consumes affine-U4 K32 blocks. Publish + // that representation directly while each 128-token x 64-row SwiGLU tile + // is live, avoiding both an F32 round trip and a standalone quantizer. + scf.if %publish_u4 { + %u4_tile = buffer.view %qout_scratch[%base] : buffer -> view<128x32xf32> + + %u4_c4 = index.constant 4 : index + %u4_c8 = index.constant 8 : index + %u4_c16 = index.constant 16 : index + %u4_c32 = index.constant 32 : index + %u4_zero_f32 = scalar.constant 0.0 : f32 + %u4_zero_i32 = scalar.constant 0 : i32 + %u4_eps = scalar.constant 1.0000000000000001e-30 : f32 + %u4_f15 = scalar.constant 15.0 : f32 + %u4_vzero = vector.constant 0.0 : vector<4xf32> + %u4_v15 = vector.constant 15.0 : vector<4xf32> + %u4_bias_neg8 = vector.constant -8 : vector<4xi32> + %u4_nibble_mask = vector.constant 15 : vector<4xi32> + %u4_c8_i32 = scalar.constant 8 : i32 + %u4_xor1 = scalar.constant 1 : i32 + %u4_xor2 = scalar.constant 2 : i32 + %u4_xor4 = scalar.constant 4 : i32 + %u4_shift4 = scalar.constant 4 : i32 + %u4_shift8 = scalar.constant 8 : i32 + %u4_shift12 = scalar.constant 12 : i32 + %u4_shuffle_width = scalar.constant 32 : i32 + + %u4_total = index.mul %rows_b, %cols_b : index + %u4_groups0 = index.div %u4_total, %u4_c32 : index + %u4_groups = index.assume %u4_groups0 [range(%u4_groups0, 1, 33554432)] : index + %u4_halfwords = index.div %u4_total, %u4_c4 : index + %u4_qs = buffer.view %qout_qs_na[%base] : buffer -> view<[%u4_halfwords]xi16> + %u4_ds = buffer.view %qout_ds_na[%base] : buffer -> view<[%u4_groups]xf32> + %u4_meta_count = index.mul %u4_groups, %c2 : index + %u4_meta = buffer.view %qout_sums_na[%base] : buffer -> view<[%u4_meta_count]xi32> + + %u4_group_in_wave = index.div %lane, %u4_c8 : index + %u4_lane_in_group = index.rem %lane, %u4_c8 : index + %u4_is_leader = index.cmp eq, %u4_lane_in_group, %c0 : index + %u4_wave_group_base = index.mul %wave, %c4 : index + %u4_local_group_base = index.add %u4_wave_group_base, %u4_group_in_wave : index + %u4_row_group_base = index.mul %row_tile, %c2 : index + %u4_row_groups = index.div %rows_b, %u4_c32 : index + + scf.for %u4_phase = [%c0 to %c2 step %c1] { + %u4_phase0 = index.cmp eq, %u4_phase, %c0 : index + scf.if %u4_phase0 { + vector.fragment.store %out0_0, %u4_tile[%lm0, %c0] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out0_1, %u4_tile[%lm0, %c16] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out1_0, %u4_tile[%lm1, %c0] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out1_1, %u4_tile[%lm1, %c16] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + } else { + vector.fragment.store %out0_2, %u4_tile[%lm0, %c0] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out0_3, %u4_tile[%lm0, %c16] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out1_2, %u4_tile[%lm1, %c0] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + vector.fragment.store %out1_3, %u4_tile[%lm1, %c16] shape [%c16, %c16] : vector<8xf32>, view<128x32xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + + scf.for %u4_batch = [%c0 to %u4_c8 step %c1] { + %u4_batch_group = index.mul %u4_batch, %u4_c16 : index + %u4_local_group0 = index.add %u4_batch_group, %u4_local_group_base : index + %u4_token_local = index.assume %u4_local_group0 [range(%u4_local_group0, 0, 127)] : index + %u4_lane_row0 = index.mul %u4_lane_in_group, %u4_c4 : index + %u4_lane_row = index.assume %u4_lane_row0 [range(%u4_lane_row0, 0, 28)] : index + %u4_values = vector.load %u4_tile[%u4_token_local, %u4_lane_row] : view<128x32xf32> -> vector<4xf32> + + %u4_lane_max = vector.reduce %u4_values, %u4_zero_f32 : vector<4xf32>, f32 + %u4_neg_values = vector.subf %u4_vzero, %u4_values : vector<4xf32> + %u4_lane_neg_min = vector.reduce %u4_neg_values, %u4_zero_f32 : vector<4xf32>, f32 + %u4_max1_peer, %u4_max1_valid = kernel.subgroup.shuffle %u4_lane_max, %u4_xor1, %u4_shuffle_width : f32, i32, i32 + %u4_max1 = scalar.maxnumf %u4_lane_max, %u4_max1_peer : f32 + %u4_max2_peer, %u4_max2_valid = kernel.subgroup.shuffle %u4_max1, %u4_xor2, %u4_shuffle_width : f32, i32, i32 + %u4_max2 = scalar.maxnumf %u4_max1, %u4_max2_peer : f32 + %u4_max4_peer, %u4_max4_valid = kernel.subgroup.shuffle %u4_max2, %u4_xor4, %u4_shuffle_width : f32, i32, i32 + %u4_group_max = scalar.maxnumf %u4_max2, %u4_max4_peer : f32 + %u4_min1_peer, %u4_min1_valid = kernel.subgroup.shuffle %u4_lane_neg_min, %u4_xor1, %u4_shuffle_width : f32, i32, i32 + %u4_min1 = scalar.maxnumf %u4_lane_neg_min, %u4_min1_peer : f32 + %u4_min2_peer, %u4_min2_valid = kernel.subgroup.shuffle %u4_min1, %u4_xor2, %u4_shuffle_width : f32, i32, i32 + %u4_min2 = scalar.maxnumf %u4_min1, %u4_min2_peer : f32 + %u4_min4_peer, %u4_min4_valid = kernel.subgroup.shuffle %u4_min2, %u4_xor4, %u4_shuffle_width : f32, i32, i32 + %u4_group_neg_min = scalar.maxnumf %u4_min2, %u4_min4_peer : f32 + %u4_range0 = scalar.addf %u4_group_max, %u4_group_neg_min : f32 + %u4_range = scalar.maxnumf %u4_range0, %u4_eps : f32 + %u4_scale = scalar.divf %u4_range, %u4_f15 : f32 + %u4_rscale = scalar.divf %u4_f15, %u4_range : f32 + %u4_zp_raw = scalar.mulf %u4_group_neg_min, %u4_rscale : f32 + %u4_zp_round = scalar.roundf %u4_zp_raw : f32 + %u4_zp_clamped = scalar.clampf %u4_zp_round, %u4_zero_f32, %u4_f15 : f32 + %u4_zp = scalar.fptosi %u4_zp_clamped : f32 to i32 + + %u4_token = index.add %col_base, %u4_token_local : index + %u4_row_group = index.add %u4_row_group_base, %u4_phase : index + %u4_gid_base = index.mul %u4_token, %u4_row_groups : index + %u4_gid0 = index.add %u4_gid_base, %u4_row_group : index + %u4_gid = index.assume %u4_gid0 [range(%u4_gid0, 0, 33554431)] : index + scf.if %u4_is_leader { + view.store %u4_scale, %u4_ds[%u4_gid] : f32, view<[%u4_groups]xf32> + } + + %u4_rs = vector.splat %u4_rscale : vector<4xf32> + %u4_zpv = vector.splat %u4_zp_clamped : vector<4xf32> + %u4_scaled = vector.mulf %u4_values, %u4_rs : vector<4xf32> + %u4_shifted = vector.addf %u4_scaled, %u4_zpv : vector<4xf32> + %u4_rounded = vector.roundf %u4_shifted : vector<4xf32> + %u4_nonnegative = vector.maxnumf %u4_rounded, %u4_vzero : vector<4xf32> + %u4_to_ceiling = vector.subf %u4_v15, %u4_nonnegative : vector<4xf32> + %u4_ceiling_nonnegative = vector.maxnumf %u4_to_ceiling, %u4_vzero : vector<4xf32> + %u4_clamped = vector.subf %u4_v15, %u4_ceiling_nonnegative : vector<4xf32> + %u4_q32_unsigned = vector.fptosi %u4_clamped : vector<4xf32> to vector<4xi32> + %u4_q32 = vector.addi %u4_q32_unsigned, %u4_bias_neg8 : vector<4xi32> + %u4_qbits = vector.andi %u4_q32, %u4_nibble_mask : vector<4xi32> + %u4_q0 = vector.extract %u4_qbits[0] : vector<4xi32> -> i32 + %u4_q1 = vector.extract %u4_qbits[1] : vector<4xi32> -> i32 + %u4_q2 = vector.extract %u4_qbits[2] : vector<4xi32> -> i32 + %u4_q3 = vector.extract %u4_qbits[3] : vector<4xi32> -> i32 + %u4_q1s = scalar.shli %u4_q1, %u4_shift4 : i32 + %u4_q2s = scalar.shli %u4_q2, %u4_shift8 : i32 + %u4_q3s = scalar.shli %u4_q3, %u4_shift12 : i32 + %u4_q01 = scalar.ori %u4_q0, %u4_q1s : i32 + %u4_q23 = scalar.ori %u4_q2s, %u4_q3s : i32 + %u4_packed32 = scalar.ori %u4_q01, %u4_q23 : i32 + %u4_packed = scalar.trunci %u4_packed32 : i32 to i16 + %u4_word_base = index.mul %u4_gid, %u4_c8 : index + %u4_word_index0 = index.add %u4_word_base, %u4_lane_in_group : index + %u4_word_index = index.assume %u4_word_index0 [range(%u4_word_index0, 0, 268435455)] : index + view.store %u4_packed, %u4_qs[%u4_word_index] : i16, view<[%u4_halfwords]xi16> + + %u4_lane_sum = vector.reduce %u4_q32, %u4_zero_i32 : vector<4xi32>, i32 + %u4_sum1_peer, %u4_sum1_valid = kernel.subgroup.shuffle %u4_lane_sum, %u4_xor1, %u4_shuffle_width : i32, i32, i32 + %u4_sum1 = scalar.addi %u4_lane_sum, %u4_sum1_peer : i32 + %u4_sum2_peer, %u4_sum2_valid = kernel.subgroup.shuffle %u4_sum1, %u4_xor2, %u4_shuffle_width : i32, i32, i32 + %u4_sum2 = scalar.addi %u4_sum1, %u4_sum2_peer : i32 + %u4_sum4_peer, %u4_sum4_valid = kernel.subgroup.shuffle %u4_sum2, %u4_xor4, %u4_shuffle_width : i32, i32, i32 + %u4_group_sum = scalar.addi %u4_sum2, %u4_sum4_peer : i32 + scf.if %u4_is_leader { + %u4_zp_signed = scalar.subi %u4_zp, %u4_c8_i32 : i32 + %u4_meta0 = index.mul %u4_gid, %c2 : index + %u4_meta1 = index.add %u4_meta0, %c1 : index + view.store %u4_group_sum, %u4_meta[%u4_meta0] : i32, view<[%u4_meta_count]xi32> + view.store %u4_zp_signed, %u4_meta[%u4_meta1] : i32, view<[%u4_meta_count]xi32> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + } + } + template.return +} + +kernel.def export("ggml_mul_mat_symmetric_i4_wmma") @ggml_mul_mat_symmetric_i4_wmma() { + %unit = index.constant 1 : index + %k_n = index.constant 128 : index + %k_m = index.constant 128 : index + %n_m1 = index.constant 127 : index + %m_m1 = index.constant 127 : index + %wg = index.constant 256 : index + %rows = config.get @ggml.mul_mat.symmetric_i4.output_size : index + %cols = config.get @ggml.mul_mat.symmetric_i4.token_count : index + %rows_up = index.add %rows, %n_m1 : index + %row_tiles = index.div %rows_up, %k_n : index + %cols_up = index.add %cols, %m_m1 : index + %col_tiles = index.div %cols_up, %k_m : index + %workgroups = index.mul %col_tiles, %row_tiles : index + kernel.launch.config workgroups(%workgroups, %unit, %unit) workgroup_size(%wg, %unit, %unit) : index +} launch(%src0: buffer, %src1: buffer, %dst: buffer, %aq: buffer, %as: buffer, %asum: buffer) { + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c17 = index.constant 17 : index + %c34 = index.constant 34 : index + %k_kc = index.constant 64 : index + %k_astride = index.constant 40 : index + %k_wstride = index.constant 40 : index + %k_m = index.constant 128 : index + %k_n = index.constant 128 : index + %k_blocks = index.constant 2 : index + %k_wave = index.constant 32 : index + %k_wm = index.constant 4 : index + %fzero = vector.constant 0.0 : vector<8xf32> + %izero = vector.constant 0 : vector<8xi32> + %apply_swiglu = scalar.constant false : i1 + %publish_f32 = scalar.constant true : i1 + %publish_f16 = scalar.constant false : i1 + %publish_u4 = scalar.constant false : i1 + %work_scratch_bytes = index.constant 14336 : offset + %work_scratch = buffer.alloca align(16) %work_scratch_bytes : buffer + + %k = config.get @ggml.mul_mat.symmetric_i4.input_size : index + %rows = config.get @ggml.mul_mat.symmetric_i4.output_size : index + %cols = config.get @ggml.mul_mat.symmetric_i4.token_count : index + %k_b = index.assume %k [range(%k, 64, 32768)] : index + %rows_b = index.assume %rows [range(%rows, 1, 262144)] : index + %cols_b = index.assume %cols [range(%cols, 1, 32768)] : index + %nchunks = index.div %k_b, %k_kc : index + %kblocks = index.div %k_b, %c32 : index + + // Keep each 16-row-tile weight footprint inside the 32 MiB Infinity Cache, + // while preserving immediate reuse across all token tiles in the block. + %grid_c127 = index.constant 127 : index + %grid_block_rows = index.constant 16 : index + %grid_row_tiles_up = index.add %rows_b, %grid_c127 : index + %grid_row_tiles = index.div %grid_row_tiles_up, %k_n : index + %grid_col_tiles_up = index.add %cols_b, %grid_c127 : index + %grid_col_tiles = index.div %grid_col_tiles_up, %k_m : index + %grid_full_blocks = index.div %grid_row_tiles, %grid_block_rows : index + %grid_full_rows = index.mul %grid_full_blocks, %grid_block_rows : index + %grid_full_groups = index.mul %grid_full_rows, %grid_col_tiles : index + %grid_tail_rows = index.sub %grid_row_tiles, %grid_full_rows : index + %grid_block_span = index.mul %grid_block_rows, %grid_col_tiles : index + %grid_wgid0 = kernel.workgroup.id : index + %grid_is_full = index.cmp ult, %grid_wgid0, %grid_full_groups : index + %col_tile0, %row_tile0 = scf.if %grid_is_full -> (index, index) { + %grid_block = index.div %grid_wgid0, %grid_block_span : index + %grid_within = index.rem %grid_wgid0, %grid_block_span : index + %grid_col = index.div %grid_within, %grid_block_rows : index + %grid_row_in_block = index.rem %grid_within, %grid_block_rows : index + %grid_row_base = index.mul %grid_block, %grid_block_rows : index + %grid_row = index.add %grid_row_base, %grid_row_in_block : index + scf.yield %grid_col, %grid_row : index, index + } else { + %grid_tail_wgid = index.sub %grid_wgid0, %grid_full_groups : index + %grid_col = index.div %grid_tail_wgid, %grid_tail_rows : index + %grid_row_in_tail = index.rem %grid_tail_wgid, %grid_tail_rows : index + %grid_row = index.add %grid_full_rows, %grid_row_in_tail : index + scf.yield %grid_col, %grid_row : index, index + } + %tid0 = kernel.workitem.id : index + %col_tile = index.assume %col_tile0 [range(%col_tile0, 0, 511)] : index + %row_tile = index.assume %row_tile0 [range(%row_tile0, 0, 4095)] : index + %tid = index.assume %tid0 [range(%tid0, 0, 255)] : index + %wave = index.div %tid, %k_wave : index + + %src0_g = buffer.assume.memory_space %src0 : buffer + %dst_g = buffer.assume.memory_space %dst : buffer + %aq_g = buffer.assume.memory_space %aq : buffer + %as_g = buffer.assume.memory_space %as : buffer + %asum_g = buffer.assume.memory_space %asum : buffer + %src0_na, %dst_na, %aq_na, %as_na, %asum_na = buffer.assume.noalias %src0_g, %dst_g, %aq_g, %as_g, %asum_g : buffer, buffer, buffer, buffer, buffer + + template.apply<@ggml.mul_mat.symmetric_i4.m128n128_wg256.body>(%apply_swiglu, %publish_f32, %publish_f16, %publish_u4, %dst, %src0_na, %dst_na, %aq_na, %as_na, %asum_na, %dst_na, %dst_na, %dst_na, %work_scratch, %base, %base, %c0, %c1, %c2, %c4, %c8, %c16, %c32, %c17, %c34, %k_kc, %k_astride, %k_wstride, %k_m, %k_n, %k_blocks, %k_wave, %k_wm, %fzero, %izero, %k_b, %rows_b, %cols_b, %nchunks, %kblocks, %col_tile, %row_tile, %tid, %wave) : (i1, i1, i1, i1, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, offset, offset, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, vector<8xf32>, vector<8xi32>, index, index, index, index, index, index, index, index, index) + kernel.return +} + +check.case public @ggml_mul_mat_symmetric_i4_wmma_k5120_n6144_m512_case { + %weight = check.generate.fill value(17) : tensor<17694720xi8> + %input = check.generate.fill value(0.125) : tensor<2621440xf32> + %output = check.generate.fill value(0.0) : tensor<3145728xf32> + %qact = check.generate.fill value(17) : tensor<1310720xi8> + %scale = check.generate.fill value(0.03125) : tensor<81920xf32> + %sum = check.generate.fill value(32) : tensor<81920xi32> + kernel.launch @ggml_mul_mat_symmetric_i4_wmma(%weight, %input, %output, %qact, %scale, %sum) : (tensor<17694720xi8>, tensor<2621440xf32>, tensor<3145728xf32>, tensor<1310720xi8>, tensor<81920xf32>, tensor<81920xi32>) + check.return +} + +check.benchmark<@ggml_mul_mat_symmetric_i4_wmma_k5120_n6144_m512_case> @ggml_mul_mat_symmetric_i4_wmma_k5120_n6144_m512 + +check.case public @ggml_mul_mat_symmetric_i4_wmma_k5120_n10240_m512_case { + %weight = check.generate.fill value(17) : tensor<29491200xi8> + %input = check.generate.fill value(0.125) : tensor<2621440xf32> + %output = check.generate.fill value(0.0) : tensor<5242880xf32> + %qact = check.generate.fill value(17) : tensor<1310720xi8> + %scale = check.generate.fill value(0.03125) : tensor<81920xf32> + %sum = check.generate.fill value(32) : tensor<81920xi32> + kernel.launch @ggml_mul_mat_symmetric_i4_wmma(%weight, %input, %output, %qact, %scale, %sum) : (tensor<29491200xi8>, tensor<2621440xf32>, tensor<5242880xf32>, tensor<1310720xi8>, tensor<81920xf32>, tensor<81920xi32>) + check.return +} + +check.benchmark<@ggml_mul_mat_symmetric_i4_wmma_k5120_n10240_m512_case> @ggml_mul_mat_symmetric_i4_wmma_k5120_n10240_m512 + +template.decl @ggml.mul_mat.symmetric_i4.lowrow.m16n16.body(%scratch: buffer, %src0_na: buffer, %dst_na: buffer, %aq_na: buffer, %as_na: buffer, %asum_na: buffer, %compact_q2: i1, %split_k: i1, %partial_na: buffer, %completion_counters_na: buffer, %base: offset, %qact_scale_base: offset, %c0: index, %c1: index, %c2: index, %c4: index, %c8: index, %c16: index, %c32: index, %c17: index, %c34: index, %k_kc: index, %k_astride: index, %k_wstride: index, %k_m: index, %k_n: index, %k_blocks: index, %k_wave: index, %k_wm: index, %fzero: vector<8xf32>, %izero: vector<8xi32>, %k_b: index, %rows_b: index, %cols_b: index, %nchunks: index, %kblocks: index, %col_tile: index, %row_tile: index, %tid: index, %wave: index, %partition: index, %chunk_begin: index, %chunk_end: index) + +// Capacity-neutral symmetric-I4 weights with shared K64 scales. +// The dual-projection roots share one 128-token activation tile. +config.decl @ggml.mul_mat.symmetric_i4.lowrow.input_size : %value: index where [range(%value, 64, 32768), mul(%value, 64)] + +config.decl @ggml.mul_mat.symmetric_i4.lowrow.output_size : %value: index where [range(%value, 1, 262144)] + +config.decl @ggml.mul_mat.symmetric_i4.lowrow.token_count : %value: index where [range(%value, 1, 16)] + +config.decl @ggml.mul_mat.symmetric_i4.lowrow.row_group_size : %value: index where [range(%value, 32, 262144), mul(%value, 32)] + +template.def<@ggml.mul_mat.symmetric_i4.lowrow.m16n16.body> device @ggml_mul_mat_symmetric_i4_lowrow_m16n16_body(%scratch: buffer, %src0_na: buffer, %dst_na: buffer, %aq_na: buffer, %as_na: buffer, %asum_na: buffer, %compact_q2: i1, %split_k: i1, %partial_na: buffer, %completion_counters_na: buffer, %base: offset, %qact_scale_base: offset, %c0: index, %c1: index, %c2: index, %c4: index, %c8: index, %c16: index, %c32: index, %c17: index, %c34: index, %k_kc: index, %k_astride: index, %k_wstride: index, %k_m: index, %k_n: index, %k_blocks: index, %k_wave: index, %k_wm: index, %fzero: vector<8xf32>, %izero: vector<8xi32>, %k_b: index, %rows_b: index, %cols_b: index, %nchunks: index, %kblocks: index, %col_tile: index, %row_tile: index, %tid: index, %wave: index, %partition: index, %chunk_begin: index, %chunk_end: index) { + %src0_i8 = buffer.view %src0_na[%base] : buffer -> view<1073741824xi8> + %src0_f16 = buffer.view %src0_na[%base] : buffer -> view<536870912xf16> + %aq_flat = buffer.view %aq_na[%base] : buffer -> view<1073741824xi8> + %as_flat = buffer.view %as_na[%qact_scale_base] : buffer -> view<33554432xf32> + %asum_flat = buffer.view %asum_na[%base] : buffer -> view<67108864xi32> + %c0_i32 = scalar.constant 0 : i32 + %q4_k_block = index.constant 256 : index + %q4_block_bytes = index.constant 144 : index + %q4_k48 = index.constant 48 : index + %q4_shift4v = vector.constant 4 : vector<4xi32> + %q4_shift8v = vector.constant 8 : vector<4xi32> + %q4_shift16v = vector.constant 16 : vector<4xi32> + %q4_nibble_mask = vector.constant 252645135 : vector<4xi32> + %q4_pair_mask = vector.constant 16711935 : vector<4xi32> + %q4_half_mask = vector.constant 65535 : vector<4xi32> + %q4_f32_32 = scalar.constant 32.0 : f32 + %c15 = index.constant 15 : index + // dst is [rows, cols] with rows contiguous, and C is [cols, rows]. + %dst_view = buffer.view %dst_na[%base] : buffer -> view<[%cols_b]x[%rows_b]xf32> + + // Packed I4 operands use 32 payload bytes per logical K64 row. + // A 40-byte LDS stride retains the conflict-avoiding padding of the IU8 body. + %i4_schema = encoding.define #encoding.operand : encoding + %al_flat = buffer.view %scratch[%base] : buffer -> view<640xi8> + %as_off = index.constant 640 : offset + %asl_view = buffer.view %scratch[%as_off] : buffer -> view<32xf32> + %azp_off = index.constant 768 : offset + %azpl_view = buffer.view %scratch[%azp_off] : buffer -> view<32xf32> + %asum_off = index.constant 896 : offset + %asuml_view = buffer.view %scratch[%asum_off] : buffer -> view<32xf32> + %result_off = index.constant 1024 : offset + %result_wave_bytes = index.constant 1024 : offset + %result_wave_offset = index.scale %wave, %result_wave_bytes : index, offset -> offset + %result_offset = index.add %result_off, %result_wave_offset : offset + %result_fragment_view = buffer.view %scratch[%result_offset] : buffer -> view<16x16xf32> + %result_physical_view = buffer.view %scratch[%result_offset] : buffer -> view<16x16xf32> + %arrival_scratch_view = buffer.view %scratch[%result_off] : buffer -> view<1xi32> + + %col_base = index.mul %col_tile, %k_m : index + %row_base = index.mul %row_tile, %k_n : index + + %k_aper = index.constant 64 : index + %k_wper = index.constant 64 : index + %ashare = index.mul %tid, %k_aper : index + %ascol = index.div %ashare, %k_kc : index + %asoff = index.rem %ashare, %k_kc : index + %adbase0 = index.mul %ascol, %k_astride : index + %adbase = index.add %adbase0, %asoff : index + %ascol_g = index.add %col_base, %ascol : index + %asbase = index.mul %ascol_g, %k_b : index + // WG64 ILP8 stages 128 scales in two passes and 256 payload halves in four passes. + %k_asn = index.mul %cols_b, %k_blocks : index + %k_aspan = index.constant 32 : index + %k_wsn = index.constant 128 : index + %k_lanes0 = index.constant 0 : index + %k_lanes1 = index.constant 128 : index + + // All waves share the padded M16 token tile; each owns one disjoint N16 weight slice. + %wm_off = index.constant 0 : index + %k_nspan = index.constant 16 : index + %wn_off = index.mul %wave, %k_nspan : index + %m_out = index.add %col_base, %wm_off : index + %n_out = index.add %row_base, %wn_off : index + // RDNA3 WMMA RHS lanes map to column lane % 16, so each lane starts at its output row. + %lane = index.rem %tid, %k_wave : index + %lane_lo = index.rem %lane, %c16 : index + %lane_hi = index.div %lane, %c16 : index + %wn_lane = index.add %wn_off, %lane_lo : index + %lm0 = index.add %wm_off, %c0 : index + %gm0 = index.add %m_out, %c0 : index + %ln0 = index.add %wn_lane, %c0 : index + %gn0 = index.add %n_out, %c0 : index + %wsb0_0 = index.add %ln0, %c0 : index + %wsb0 = index.mul %wsb0_0, %k_blocks : index + // Activation scale metadata is [tile][block][lane-half][register]. + %as_tile0 = index.div %lm0, %c16 : index + %asb0 = index.mul %as_tile0, %k_blocks : index + + // Carry one prefetched K64 activation payload between loop iterations. + %q4a_lane_active = index.cmp ult, %tid, %cols_b : index + %q4a_row0 = index.add %col_base, %tid : index + %q4a_row = index.assume %q4a_row0 [range(%q4a_row0, 0, 32767)] : index + %q4a_row_bytes = index.div %k_b, %c2 : index + %q4a_row_base = index.mul %q4a_row, %q4a_row_bytes : index + %q4a_c16 = index.constant 16 : index + %q4a_zero = vector.constant 0 : vector<16xi8> + %q4a_initial_v0, %q4a_initial_v1 = scf.if %q4a_lane_active -> (vector<16xi8>, vector<16xi8>) { + %q4a_initial_chunk_k = index.mul %chunk_begin, %k_kc : index + %q4a_initial_chunk_bytes = index.div %q4a_initial_chunk_k, %c2 : index + %q4a_initial_src0_0 = index.add %q4a_row_base, %q4a_initial_chunk_bytes : index + %q4a_initial_src0 = index.assume %q4a_initial_src0_0 [range(%q4a_initial_src0_0, 0, 536870896)] : index + %q4a_loaded_initial_v0 = vector.load %aq_flat[%q4a_initial_src0] : view<1073741824xi8> -> vector<16xi8> + %q4a_initial_src1_0 = index.add %q4a_initial_src0, %q4a_c16 : index + %q4a_initial_src1 = index.assume %q4a_initial_src1_0 [range(%q4a_initial_src1_0, 16, 536870912)] : index + %q4a_loaded_initial_v1 = vector.load %aq_flat[%q4a_initial_src1] : view<1073741824xi8> -> vector<16xi8> + scf.yield %q4a_loaded_initial_v0, %q4a_loaded_initial_v1 : vector<16xi8>, vector<16xi8> + } else { + scf.yield %q4a_zero, %q4a_zero : vector<16xi8>, vector<16xi8> + } + + %f0_0, %q4a_unused_v0, %q4a_unused_v1 = scf.for %chunk = [%chunk_begin to %chunk_end step %c1](%fc0_0 = %fzero : vector<8xf32>, %q4a_current_v0 = %q4a_initial_v0 : vector<16xi8>, %q4a_current_v1 = %q4a_initial_v1 : vector<16xi8>) -> (vector<8xf32>, vector<16xi8>, vector<16xi8>) { + %chunk_k = index.mul %chunk, %k_kc : index + %chunk_blk = index.mul %chunk, %k_blocks : index + + // The first 128 lanes stage the unchanged M128 activation tile. + scf.if %q4a_lane_active { + %q4a_dst0_0 = index.mul %tid, %k_astride : index + %q4a_dst0 = index.assume %q4a_dst0_0 [range(%q4a_dst0_0, 0, 600)] : index + vector.store %q4a_current_v0, %al_flat[%q4a_dst0] : vector<16xi8>, view<640xi8> + %q4a_dst1_0 = index.add %q4a_dst0, %q4a_c16 : index + %q4a_dst1 = index.assume %q4a_dst1_0 [range(%q4a_dst1_0, 16, 616)] : index + vector.store %q4a_current_v1, %al_flat[%q4a_dst1] : vector<16xi8>, view<640xi8> + } + + // Issue the following K64 payload before weight unpacking and contraction. + // The uniform last-iteration guard prevents a speculative next-row access. + %q4a_next_chunk = index.add %chunk, %c1 : index + %q4a_has_next = index.cmp ult, %q4a_next_chunk, %chunk_end : index + %q4a_load_next = scalar.andi %q4a_has_next, %q4a_lane_active : i1 + %q4a_next_v0, %q4a_next_v1 = scf.if %q4a_load_next -> (vector<16xi8>, vector<16xi8>) { + %q4a_next_chunk_k = index.mul %q4a_next_chunk, %k_kc : index + %q4a_next_chunk_bytes = index.div %q4a_next_chunk_k, %c2 : index + %q4a_next_src0_0 = index.add %q4a_row_base, %q4a_next_chunk_bytes : index + %q4a_next_src0 = index.assume %q4a_next_src0_0 [range(%q4a_next_src0_0, 0, 536870896)] : index + %q4a_loaded_v0 = vector.load %aq_flat[%q4a_next_src0] : view<1073741824xi8> -> vector<16xi8> + %q4a_next_src1_0 = index.add %q4a_next_src0, %q4a_c16 : index + %q4a_next_src1 = index.assume %q4a_next_src1_0 [range(%q4a_next_src1_0, 16, 536870912)] : index + %q4a_loaded_v1 = vector.load %aq_flat[%q4a_next_src1] : view<1073741824xi8> -> vector<16xi8> + scf.yield %q4a_loaded_v0, %q4a_loaded_v1 : vector<16xi8>, vector<16xi8> + } else { + scf.yield %q4a_zero, %q4a_zero : vector<16xi8>, vector<16xi8> + } + scf.schedule.fence + // One half-wave owns each signed-I4 K32 group for the same + // 16 output rows. Four adjacent rows share each K32 scale. + %q4_weight_tid0 = index.rem %tid, %c16 : index + %q4_weight_tid = index.assume %q4_weight_tid0 [range(%q4_weight_tid0, 0, 15)] : index + %q4_group_half = index.div %lane, %c16 : index + %q4_block_count = index.div %k_b, %q4_k_block : index + %q4_block = index.div %chunk, %c4 : index + %q4_pair = index.rem %chunk, %c4 : index + %q4_row0 = index.add %n_out, %q4_weight_tid : index + %q4_row = index.assume %q4_row0 [range(%q4_row0, 0, 262143)] : index + %q4_row_group_size0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.row_group_size : index + %q4_row_group_size = index.assume %q4_row_group_size0 [range(%q4_row_group_size0, 32, 262144), mul(%q4_row_group_size0, 32)] : index + %q4_scale_cohort_size = index.constant 4 : index + %q4_field_bytes = index.constant 16 : index + %q2_field_bytes = index.constant 8 : index + %weight_field_bytes = scf.select %compact_q2, %q2_field_bytes, %q4_field_bytes : index + %q4_scale_plane_bytes = index.mul %q4_row_group_size, %c4 : index + %q4_payload_plane_bytes = index.mul %q4_row_group_size, %weight_field_bytes : index + %q4_full_bytes_per_row = index.constant 132 : index + %q2_compact_bytes_per_row = index.constant 68 : index + %q4_bytes_per_row = scf.select %compact_q2, %q2_compact_bytes_per_row, %q4_full_bytes_per_row : index + %q4_full_payload_bytes_per_row = index.constant 128 : index + %q2_compact_payload_bytes_per_row = index.constant 64 : index + %q4_payload_bytes_per_row = scf.select %compact_q2, %q2_compact_payload_bytes_per_row, %q4_full_payload_bytes_per_row : index + %q4_group_block_bytes = index.mul %q4_row_group_size, %q4_bytes_per_row : index + %q4_payload_block_bytes = index.mul %q4_row_group_size, %q4_payload_bytes_per_row : index + %q4_row_group = index.div %q4_row, %q4_row_group_size : index + %q4_row_lane = index.rem %q4_row, %q4_row_group_size : index + %q4_row_group_span0 = index.mul %q4_block_count, %q4_group_block_bytes : index + %q4_row_group_byte_base = index.mul %q4_row_group, %q4_row_group_span0 : index + %q4_payload_block_offset = index.mul %q4_block, %q4_payload_block_bytes : index + %q4_payload_block_base = index.add %q4_row_group_byte_base, %q4_payload_block_offset : index + %q4_all_payload_bytes = index.mul %q4_block_count, %q4_payload_block_bytes : index + %q4_scale_region = index.add %q4_row_group_byte_base, %q4_all_payload_bytes : index + %q4_scale_block_offset = index.mul %q4_block, %q4_scale_plane_bytes : index + %q4_scale_block_base = index.add %q4_scale_region, %q4_scale_block_offset : index + %q4_scale_cohort = index.div %q4_row_lane, %q4_scale_cohort_size : index + %q4_header_cohort_byte = index.mul %q4_scale_cohort, %q4_field_bytes : index + %q4_header_byte_index0 = index.add %q4_scale_block_base, %q4_header_cohort_byte : index + %q4_scale_byte_stride = index.constant 4 : index + %q4_scale_byte_offset = index.mul %q4_pair, %q4_scale_byte_stride : index + %q4_scale_byte_index = index.add %q4_header_byte_index0, %q4_scale_byte_offset : index + %q4_scale_byte_base = index.cast %q4_scale_byte_index : index to offset + %q4_scale_view = buffer.view %src0_na[%q4_scale_byte_base] : buffer -> view<2xf16> + %q4_scale_pair = vector.load %q4_scale_view[%c0] : view<2xf16> -> vector<2xf16> + %q4_scale_low_f16 = vector.extract %q4_scale_pair[0] : vector<2xf16> -> f16 + %q4_scale_high_f16 = vector.extract %q4_scale_pair[1] : vector<2xf16> -> f16 + %q4_scale_low = scalar.extf %q4_scale_low_f16 : f16 to f32 + %q4_scale_high = scalar.extf %q4_scale_high_f16 : f16 to f32 + %q4_scale_is_high_half = index.cmp eq, %q4_group_half, %c1 : index + %q4_d0 = scf.select %q4_scale_is_high_half, %q4_scale_high, %q4_scale_low : f32 + %q4_group_base = index.mul %q4_pair, %c2 : index + %q4_group = index.add %q4_group_base, %q4_group_half : index + %q4_payload_group_byte = index.mul %q4_group, %q4_payload_plane_bytes : index + %q4_payload_row_byte = index.mul %q4_row_lane, %weight_field_bytes : index + %q4_payload_base1 = index.add %q4_payload_block_base, %q4_payload_group_byte : index + %q4_byte_index = index.add %q4_payload_base1, %q4_payload_row_byte : index + %q4_byte_base = index.cast %q4_byte_index : index to offset + %q4_words = scf.if %compact_q2 -> (vector<4xi32>) { + %q2_view = buffer.view %src0_na[%q4_byte_base] : buffer -> view<2xi32> + %q2_packed = vector.load %q2_view[%c0] : view<2xi32> -> vector<2xi32> + %q2_mask16 = vector.constant 65535 : vector<2xi32> + %q2_mask8 = vector.constant 16711935 : vector<2xi32> + %q2_mask4 = vector.constant 252645135 : vector<2xi32> + %q2_mask2 = vector.constant 858993459 : vector<2xi32> + %q2_sign = vector.constant 572662306 : vector<2xi32> + %q2_shift16 = vector.constant 16 : vector<2xi32> + %q2_shift8 = vector.constant 8 : vector<2xi32> + %q2_shift4 = vector.constant 4 : vector<2xi32> + %q2_shift2 = vector.constant 2 : vector<2xi32> + %q2_shift1 = vector.constant 1 : vector<2xi32> + %q2_lo0 = vector.andi %q2_packed, %q2_mask16 : vector<2xi32> + %q2_hi0 = vector.shrui %q2_packed, %q2_shift16 : vector<2xi32> + %q2_lo1s = vector.shli %q2_lo0, %q2_shift8 : vector<2xi32> + %q2_lo1o = vector.ori %q2_lo0, %q2_lo1s : vector<2xi32> + %q2_lo1 = vector.andi %q2_lo1o, %q2_mask8 : vector<2xi32> + %q2_lo2s = vector.shli %q2_lo1, %q2_shift4 : vector<2xi32> + %q2_lo2o = vector.ori %q2_lo1, %q2_lo2s : vector<2xi32> + %q2_lo2 = vector.andi %q2_lo2o, %q2_mask4 : vector<2xi32> + %q2_lo3s = vector.shli %q2_lo2, %q2_shift2 : vector<2xi32> + %q2_lo3o = vector.ori %q2_lo2, %q2_lo3s : vector<2xi32> + %q2_lo3 = vector.andi %q2_lo3o, %q2_mask2 : vector<2xi32> + %q2_losign = vector.andi %q2_lo3, %q2_sign : vector<2xi32> + %q2_losign2 = vector.shli %q2_losign, %q2_shift1 : vector<2xi32> + %q2_losign3 = vector.shli %q2_losign, %q2_shift2 : vector<2xi32> + %q2_lo4 = vector.ori %q2_lo3, %q2_losign2 : vector<2xi32> + %q2_lo = vector.ori %q2_lo4, %q2_losign3 : vector<2xi32> + %q2_hi1s = vector.shli %q2_hi0, %q2_shift8 : vector<2xi32> + %q2_hi1o = vector.ori %q2_hi0, %q2_hi1s : vector<2xi32> + %q2_hi1 = vector.andi %q2_hi1o, %q2_mask8 : vector<2xi32> + %q2_hi2s = vector.shli %q2_hi1, %q2_shift4 : vector<2xi32> + %q2_hi2o = vector.ori %q2_hi1, %q2_hi2s : vector<2xi32> + %q2_hi2 = vector.andi %q2_hi2o, %q2_mask4 : vector<2xi32> + %q2_hi3s = vector.shli %q2_hi2, %q2_shift2 : vector<2xi32> + %q2_hi3o = vector.ori %q2_hi2, %q2_hi3s : vector<2xi32> + %q2_hi3 = vector.andi %q2_hi3o, %q2_mask2 : vector<2xi32> + %q2_hisign = vector.andi %q2_hi3, %q2_sign : vector<2xi32> + %q2_hisign2 = vector.shli %q2_hisign, %q2_shift1 : vector<2xi32> + %q2_hisign3 = vector.shli %q2_hisign, %q2_shift2 : vector<2xi32> + %q2_hi4 = vector.ori %q2_hi3, %q2_hisign2 : vector<2xi32> + %q2_hi = vector.ori %q2_hi4, %q2_hisign3 : vector<2xi32> + %q2_word0 = vector.extract %q2_lo[0] : vector<2xi32> -> i32 + %q2_word1 = vector.extract %q2_hi[0] : vector<2xi32> -> i32 + %q2_word2 = vector.extract %q2_lo[1] : vector<2xi32> -> i32 + %q2_word3 = vector.extract %q2_hi[1] : vector<2xi32> -> i32 + %q2_words = vector.from_elements %q2_word0, %q2_word1, %q2_word2, %q2_word3 : vector<4xi32> + scf.yield %q2_words : vector<4xi32> + } else { + %q4_view = buffer.view %src0_na[%q4_byte_base] : buffer -> view<4xi32> + %q4_loaded = vector.load %q4_view[%c0] : view<4xi32> -> vector<4xi32> + scf.yield %q4_loaded : vector<4xi32> + } + // Scale planes: 256 activation and 128 weight entries. + %sca0_slot0 = index.add %tid, %k_lanes0 : index + %sca0_slot = index.assume %sca0_slot0 [range(%sca0_slot0, 0, 63)] : index + %scin_a0 = index.cmp ult, %sca0_slot, %k_asn : index + scf.if %scin_a0 { + %sc0_col = index.div %sca0_slot, %k_blocks : index + %sc0_blk = index.rem %sca0_slot, %k_blocks : index + %sca0_c = index.add %col_base, %sc0_col : index + %sca0_0 = index.mul %sca0_c, %kblocks : index + %sca0_1 = index.add %sca0_0, %chunk_blk : index + %sca0_2 = index.add %sca0_1, %sc0_blk : index + %sca0 = index.assume %sca0_2 [range(%sca0_2, 0, 33554431)] : index + %scav0 = view.load %as_flat[%sca0] : view<33554432xf32> -> f32 + %sca0_tile = index.div %sc0_col, %c16 : index + %sca0_inner = index.rem %sc0_col, %c16 : index + %sca0_par = index.rem %sca0_inner, %c2 : index + %sca0_v = index.div %sca0_inner, %c2 : index + %sca0_dst0 = index.mul %sca0_tile, %k_blocks : index + %sca0_dst1 = index.add %sca0_dst0, %sc0_blk : index + %sca0_dst2 = index.mul %sca0_dst1, %c2 : index + %sca0_dst3 = index.add %sca0_dst2, %sca0_par : index + %sca0_dst4 = index.mul %sca0_dst3, %c8 : index + %sca0_dst5 = index.add %sca0_dst4, %sca0_v : index + %sca0_dst = index.assume %sca0_dst5 [range(%sca0_dst5, 0, 31)] : index + view.store %scav0, %asl_view[%sca0_dst] : f32, view<32xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + + %q4_xor16 = scalar.constant 16 : i32 + %q4_shuffle_width = scalar.constant 32 : i32 + %q4_word0_local = vector.extract %q4_words[0] : vector<4xi32> -> i32 + %q4_word1_local = vector.extract %q4_words[1] : vector<4xi32> -> i32 + %q4_word2_local = vector.extract %q4_words[2] : vector<4xi32> -> i32 + %q4_word3_local = vector.extract %q4_words[3] : vector<4xi32> -> i32 + %q4_word0_peer, %q4_word0_valid = kernel.subgroup.shuffle %q4_word0_local, %q4_xor16, %q4_shuffle_width : i32, i32, i32 + %q4_word1_peer, %q4_word1_valid = kernel.subgroup.shuffle %q4_word1_local, %q4_xor16, %q4_shuffle_width : i32, i32, i32 + %q4_word2_peer, %q4_word2_valid = kernel.subgroup.shuffle %q4_word2_local, %q4_xor16, %q4_shuffle_width : i32, i32, i32 + %q4_word3_peer, %q4_word3_valid = kernel.subgroup.shuffle %q4_word3_local, %q4_xor16, %q4_shuffle_width : i32, i32, i32 + %q4_d0_peer, %q4_d0_valid = kernel.subgroup.shuffle %q4_d0, %q4_xor16, %q4_shuffle_width : f32, i32, i32 + %q4_is_high_half = index.cmp eq, %q4_group_half, %c1 : index + %q4_word0_u0 = scf.select %q4_is_high_half, %q4_word0_peer, %q4_word0_local : i32 + %q4_word1_u0 = scf.select %q4_is_high_half, %q4_word1_peer, %q4_word1_local : i32 + %q4_word2_u0 = scf.select %q4_is_high_half, %q4_word2_peer, %q4_word2_local : i32 + %q4_word3_u0 = scf.select %q4_is_high_half, %q4_word3_peer, %q4_word3_local : i32 + %q4_d_u0 = scf.select %q4_is_high_half, %q4_d0_peer, %q4_d0 : f32 + %q4_d_u1 = scf.select %q4_is_high_half, %q4_d0, %q4_d0_peer : f32 + %q4_word0_u1 = scf.select %q4_is_high_half, %q4_word0_local, %q4_word0_peer : i32 + %q4_word1_u1 = scf.select %q4_is_high_half, %q4_word1_local, %q4_word1_peer : i32 + %q4_word2_u1 = scf.select %q4_is_high_half, %q4_word2_local, %q4_word2_peer : i32 + %q4_word3_u1 = scf.select %q4_is_high_half, %q4_word3_local, %q4_word3_peer : i32 + + // Apply each signed-I4 K32 dot product with its own weight scale. + %blk_k_u0 = index.mul %c0, %c16 : index + %blk_k16_u0 = index.add %blk_k_u0, %c8 : index + %iz_u0 = vector.fragment %izero shape [%c16, %c16] : vector<8xi32> + %wsv0_u0 = vector.splat %q4_d_u0 : vector<8xf32> + %lf0_0_lane_row_u0 = index.add %lm0, %lane_lo : index + %lf0_0_row_u0 = index.mul %lf0_0_lane_row_u0, %k_astride : index + %lf0_0_base_u0 = index.add %lf0_0_row_u0, %blk_k_u0 : index + %lf0_0_idx_u0 = index.assume %lf0_0_base_u0 [range(%lf0_0_base_u0, 0, 632)] : index + %lf0_0_raw_u0 = vector.load %al_flat[%lf0_0_idx_u0] : view<640xi8> -> vector<8xi8> + %lf0_0_words_u0 = vector.bitcast %lf0_0_raw_u0 : vector<8xi8> to vector<2xi32> + %lf0_0_u0 = vector.fragment %lf0_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %lf0_1_lane_row_u0 = index.add %lm0, %lane_lo : index + %lf0_1_row_u0 = index.mul %lf0_1_lane_row_u0, %k_astride : index + %lf0_1_base_u0 = index.add %lf0_1_row_u0, %blk_k16_u0 : index + %lf0_1_idx_u0 = index.assume %lf0_1_base_u0 [range(%lf0_1_base_u0, 0, 632)] : index + %lf0_1_raw_u0 = vector.load %al_flat[%lf0_1_idx_u0] : view<640xi8> -> vector<8xi8> + %lf0_1_words_u0 = vector.bitcast %lf0_1_raw_u0 : vector<8xi8> to vector<2xi32> + %lf0_1_u0 = vector.fragment %lf0_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %asi0_0_u0 = index.add %asb0, %c0 : index + %asi0_1_u0 = index.mul %asi0_0_u0, %c2 : index + %asi0_2_u0 = index.add %asi0_1_u0, %lane_hi : index + %asi0_3_u0 = index.mul %asi0_2_u0, %c8 : index + %asi0_u0 = index.assume %asi0_3_u0 [range(%asi0_3_u0, 0, 24)] : index + %asv0_u0 = vector.load %asl_view[%asi0_u0] : view<32xf32> -> vector<8xf32> + %rfs0_0_0_words_u0 = vector.from_elements %q4_word0_u0, %q4_word1_u0 : vector<2xi32> + %rfs0_0_0_u0 = vector.fragment %rfs0_0_0_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_0_1_words_u0 = vector.from_elements %q4_word2_u0, %q4_word3_u0 : vector<2xi32> + %rfs0_0_1_u0 = vector.fragment %rfs0_0_1_words_u0 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_0_u0 = vector.mma %lf0_0_u0, %rfs0_0_0_u0, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_0_u0 = vector.mma %lf0_1_u0, %rfs0_0_1_u0, %i0_0_0_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff0_0_u0 = vector.sitofp %i1_0_0_u0 : vector<8xi32> to vector<8xf32> + %sv0_0_u0 = vector.mulf %asv0_u0, %wsv0_u0 : vector<8xf32> + %fn0_0_u0 = vector.fmaf %ff0_0_u0, %sv0_0_u0, %fc0_0 : vector<8xf32> + scf.schedule.fence + + %blk_k_u1 = index.mul %c1, %c16 : index + %blk_k16_u1 = index.add %blk_k_u1, %c8 : index + %lf0_0_lane_row_u1 = index.add %lm0, %lane_lo : index + %lf0_0_row_u1 = index.mul %lf0_0_lane_row_u1, %k_astride : index + %lf0_0_base_u1 = index.add %lf0_0_row_u1, %blk_k_u1 : index + %lf0_0_idx_u1 = index.assume %lf0_0_base_u1 [range(%lf0_0_base_u1, 0, 632)] : index + %lf0_0_raw_u1 = vector.load %al_flat[%lf0_0_idx_u1] : view<640xi8> -> vector<8xi8> + %lf0_0_words_u1 = vector.bitcast %lf0_0_raw_u1 : vector<8xi8> to vector<2xi32> + %lf0_0_u1 = vector.fragment %lf0_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %lf0_1_lane_row_u1 = index.add %lm0, %lane_lo : index + %lf0_1_row_u1 = index.mul %lf0_1_lane_row_u1, %k_astride : index + %lf0_1_base_u1 = index.add %lf0_1_row_u1, %blk_k16_u1 : index + %lf0_1_idx_u1 = index.assume %lf0_1_base_u1 [range(%lf0_1_base_u1, 0, 632)] : index + %lf0_1_raw_u1 = vector.load %al_flat[%lf0_1_idx_u1] : view<640xi8> -> vector<8xi8> + %lf0_1_words_u1 = vector.bitcast %lf0_1_raw_u1 : vector<8xi8> to vector<2xi32> + %lf0_1_u1 = vector.fragment %lf0_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_0_0_words_u1 = vector.from_elements %q4_word0_u1, %q4_word1_u1 : vector<2xi32> + %rfs0_0_0_u1 = vector.fragment %rfs0_0_0_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %rfs0_0_1_words_u1 = vector.from_elements %q4_word2_u1, %q4_word3_u1 : vector<2xi32> + %rfs0_0_1_u1 = vector.fragment %rfs0_0_1_words_u1 shape [%c16, %c16] using {schema = %i4_schema : encoding} : vector<2xi32> + %i0_0_0_u1 = vector.mma %lf0_0_u1, %rfs0_0_0_u1, %iz_u0 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %i1_0_0_u1 = vector.mma %lf0_1_u1, %rfs0_0_1_u1, %i0_0_0_u1 : vector<2xi32>, vector<2xi32>, vector<8xi32> + %ff0_0_u1 = vector.sitofp %i1_0_0_u1 : vector<8xi32> to vector<8xf32> + %asi0_0_u1 = index.add %asb0, %c1 : index + %asi0_1_u1 = index.mul %asi0_0_u1, %c2 : index + %asi0_2_u1 = index.add %asi0_1_u1, %lane_hi : index + %asi0_3_u1 = index.mul %asi0_2_u1, %c8 : index + %asi0_u1 = index.assume %asi0_3_u1 [range(%asi0_3_u1, 0, 24)] : index + %asv0_u1 = vector.load %asl_view[%asi0_u1] : view<32xf32> -> vector<8xf32> + %wsv0_u1 = vector.splat %q4_d_u1 : vector<8xf32> + %sv0_0_u1 = vector.mulf %asv0_u1, %wsv0_u1 : vector<8xf32> + %fn0_0 = vector.fmaf %ff0_0_u1, %sv0_0_u1, %fn0_0_u0 : vector<8xf32> + scf.schedule.fence + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %fn0_0, %q4a_next_v0, %q4a_next_v1 : vector<8xf32>, vector<16xi8>, vector<16xi8> + } + + vector.fragment.store %f0_0, %result_fragment_view[%c0, %c0] shape [%c16, %c16] : vector<8xf32>, view<16x16xf32> + kernel.barrier scope(subgroup) ordering(acq_rel) + %publish_two = index.constant 2 : index + %publish_eight = index.constant 8 : index + %publish_token = index.div %lane, %publish_two : index + %publish_packet = index.rem %lane, %publish_two : index + %publish_channel_add = index.mul %publish_packet, %publish_eight : index + %token = index.add %gm0, %publish_token : index + %channel = index.add %n_out, %publish_channel_add : index + %valid_token = index.cmp ult, %token, %cols_b : index + %valid_channel = index.cmp ult, %channel, %rows_b : index + %writes = scalar.andi %valid_token, %valid_channel : i1 + scf.if %split_k { + %row_tiles_up = index.add %rows_b, %c15 : index + %row_tile_count = index.div %row_tiles_up, %c16 : index + %col_tiles_up = index.add %cols_b, %c15 : index + %col_tile_count = index.div %col_tiles_up, %c16 : index + %completion_counter_count = index.mul %row_tile_count, %col_tile_count : index + %completion_counters_aligned = buffer.assume.alignment %completion_counters_na {minimum_alignment = 16} : buffer + %completion_counter_view = buffer.view %completion_counters_aligned[%base] : buffer -> view<[%completion_counter_count]xi32> + %is_partition_zero = index.cmp eq, %partition, %c0 : index + scf.if %writes { + %bounded_token, %output_token_count = index.assume %token, %cols_b [lt(%token, %cols_b)] : index, index + %bounded_output_view = buffer.view %dst_na[%base] : buffer -> view<[%output_token_count]x[%rows_b]xf32> + %bounded_partial_view = buffer.view %partial_na[%base] : buffer -> view<[%output_token_count]x[%rows_b]xf32> + %values = vector.load %result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf32> -> vector<8xf32> + %mask = vector.mask.range [%channel to %rows_b step %c1] : index -> vector<8xi1> + scf.if %is_partition_zero { + vector.store.mask %values, %bounded_output_view[%bounded_token, %channel], %mask : vector<8xf32>, view<[%output_token_count]x[%rows_b]xf32>, vector<8xi1> + } else { + vector.store.mask %values, %bounded_partial_view[%bounded_token, %channel], %mask : vector<8xf32>, view<[%output_token_count]x[%rows_b]xf32>, vector<8xi1> + } + } + + // Both K partitions publish before the second arrival combines P0 then P1 and resets the tile counter. + kernel.barrier scope(workgroup) ordering(release) + %is_arrival_workitem = index.cmp eq, %tid, %c0 : index + %counter_group_base = index.mul %col_tile, %row_tile_count : index + %counter_index = index.add %counter_group_base, %row_tile : index + %c1_i32 = scalar.constant 1 : i32 + %cneg2_i32 = scalar.constant -2 : i32 + scf.if %is_arrival_workitem { + %old_counter = view.atomic.rmw %c1_i32, %completion_counter_view[%counter_index] {ordering = acq_rel, scope = device} : i32, view<[%completion_counter_count]xi32> -> i32 + view.store %old_counter, %arrival_scratch_view[%c0] : i32, view<1xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %old_counter = view.load %arrival_scratch_view[%c0] : view<1xi32> -> i32 + %is_last_partition = scalar.cmpi eq, %old_counter, %c1_i32 : i32 + scf.if %is_last_partition { + kernel.barrier scope(workgroup) ordering(acquire) + scf.if %writes { + %bounded_token, %output_token_count = index.assume %token, %cols_b [lt(%token, %cols_b)] : index, index + %bounded_output_view = buffer.view %dst_na[%base] : buffer -> view<[%output_token_count]x[%rows_b]xf32> + %bounded_partial_view = buffer.view %partial_na[%base] : buffer -> view<[%output_token_count]x[%rows_b]xf32> + %mask = vector.mask.range [%channel to %rows_b step %c1] : index -> vector<8xi1> + %partition_zero_values = vector.load.mask %bounded_output_view[%bounded_token, %channel], %mask, %fzero : view<[%output_token_count]x[%rows_b]xf32>, vector<8xi1>, vector<8xf32> + %partition_one_values = vector.load.mask %bounded_partial_view[%bounded_token, %channel], %mask, %fzero : view<[%output_token_count]x[%rows_b]xf32>, vector<8xi1>, vector<8xf32> + %combined_values = vector.addf %partition_zero_values, %partition_one_values : vector<8xf32> + vector.store.mask %combined_values, %bounded_output_view[%bounded_token, %channel], %mask : vector<8xf32>, view<[%output_token_count]x[%rows_b]xf32>, vector<8xi1> + } + kernel.barrier scope(workgroup) ordering(release) + scf.if %is_arrival_workitem { + view.atomic.reduce %cneg2_i32, %completion_counter_view[%counter_index] {ordering = release, scope = device} : i32, view<[%completion_counter_count]xi32> + } + } + } else { + scf.if %writes { + %bounded_token, %output_token_count = index.assume %token, %cols_b [lt(%token, %cols_b)] : index, index + %bounded_output_view = buffer.view %dst_na[%base] : buffer -> view<[%output_token_count]x[%rows_b]xf32> + %values = vector.load %result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf32> -> vector<8xf32> + %mask = vector.mask.range [%channel to %rows_b step %c1] : index -> vector<8xi1> + vector.store.mask %values, %bounded_output_view[%bounded_token, %channel], %mask : vector<8xf32>, view<[%output_token_count]x[%rows_b]xf32>, vector<8xi1> + } + } + template.return +} + +kernel.def export("ggml_mul_mat_symmetric_i4_lowrow_wmma") @ggml_mul_mat_symmetric_i4_lowrow_wmma() { + %unit = index.constant 1 : index + %c2 = index.constant 2 : index + %k_n = index.constant 16 : index + %k_m = index.constant 16 : index + %n_m1 = index.constant 15 : index + %m_m1 = index.constant 15 : index + %wg = index.constant 64 : index + %rows = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %cols = config.get @ggml.mul_mat.symmetric_i4.lowrow.token_count : index + %rows_up = index.add %rows, %n_m1 : index + %row_tiles = index.div %rows_up, %k_n : index + %row_tiles_paired_up = index.add %row_tiles, %unit : index + %row_groups = index.div %row_tiles_paired_up, %c2 : index + %cols_up = index.add %cols, %m_m1 : index + %col_tiles = index.div %cols_up, %k_m : index + kernel.launch.config workgroups(%col_tiles, %row_groups, %unit) workgroup_size(%wg, %unit, %unit) : index +} launch(%weight: buffer, %input: buffer, %output: buffer, %aq: buffer, %as: buffer, %asum: buffer) { + %base = index.constant 0 : offset + %lds_bytes = index.constant 3072 : offset + %scratch = buffer.alloca align(16) %lds_bytes : buffer + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c17 = index.constant 17 : index + %c34 = index.constant 34 : index + %k_kc = index.constant 64 : index + %k_astride = index.constant 40 : index + %k_wstride = index.constant 40 : index + %k_m = index.constant 16 : index + %k_n = index.constant 16 : index + %k_blocks = index.constant 2 : index + %k_wave = index.constant 32 : index + %k_wm = index.constant 1 : index + %fzero = vector.constant 0.0 : vector<8xf32> + %izero = vector.constant 0 : vector<8xi32> + %no_split_k = scalar.constant false : i1 + %compact_q2 = scalar.constant false : i1 + + %k = config.get @ggml.mul_mat.symmetric_i4.lowrow.input_size : index + %rows = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %cols = config.get @ggml.mul_mat.symmetric_i4.lowrow.token_count : index + %k_b = index.assume %k [range(%k, 64, 32768)] : index + %rows_b = index.assume %rows [range(%rows, 1, 262144)] : index + %cols_b = index.assume %cols [range(%cols, 1, 32768)] : index + %nchunks = index.div %k_b, %k_kc : index + %kblocks = index.div %k_b, %c32 : index + + %col_tile0 = kernel.workgroup.id : index + %row_group0 = kernel.workgroup.id : index + %tid0 = kernel.workitem.id : index + %col_tile = index.assume %col_tile0 [range(%col_tile0, 0, 511)] : index + %row_group = index.assume %row_group0 [range(%row_group0, 0, 8191)] : index + %row_tile0 = index.mul %row_group, %c2 : index + %row_tile = index.assume %row_tile0 [range(%row_tile0, 0, 16383)] : index + %tid = index.assume %tid0 [range(%tid0, 0, 63)] : index + %wave = index.div %tid, %k_wave : index + + %weight_g = buffer.assume.memory_space %weight : buffer + %output_g = buffer.assume.memory_space %output : buffer + %aq_g = buffer.assume.memory_space %aq : buffer + %as_g = buffer.assume.memory_space %as : buffer + %asum_g = buffer.assume.memory_space %asum : buffer + %weight_na, %output_na, %aq_na, %as_na, %asum_na = buffer.assume.noalias %weight_g, %output_g, %aq_g, %as_g, %asum_g : buffer, buffer, buffer, buffer, buffer + template.apply<@ggml.mul_mat.symmetric_i4.lowrow.m16n16.body>(%scratch, %weight_na, %output_na, %aq_na, %as_na, %asum_na, %compact_q2, %no_split_k, %output_na, %output_na, %base, %base, %c0, %c1, %c2, %c4, %c8, %c16, %c32, %c17, %c34, %k_kc, %k_astride, %k_wstride, %k_m, %k_n, %k_blocks, %k_wave, %k_wm, %fzero, %izero, %k_b, %rows_b, %cols_b, %nchunks, %kblocks, %col_tile, %row_tile, %tid, %wave, %c0, %c0, %nchunks) : (buffer, buffer, buffer, buffer, buffer, buffer, i1, i1, buffer, buffer, offset, offset, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, vector<8xf32>, vector<8xi32>, index, index, index, index, index, index, index, index, index, index, index, index) + kernel.return +} + +// Two adjacent projections share one launch and one quantized activation. +// Each projection keeps the same low-row contraction body and owns disjoint +// output storage; the graph matcher supplies any semantically related pair. +kernel.def export("ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma") @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma() { + %unit = index.constant 1 : index + %c2 = index.constant 2 : index + %k_n = index.constant 16 : index + %k_m = index.constant 16 : index + %n_m1 = index.constant 15 : index + %m_m1 = index.constant 15 : index + %wg = index.constant 64 : index + %rows = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %cols = config.get @ggml.mul_mat.symmetric_i4.lowrow.token_count : index + %rows_up = index.add %rows, %n_m1 : index + %row_tiles = index.div %rows_up, %k_n : index + %row_tiles_paired_up = index.add %row_tiles, %unit : index + %row_groups = index.div %row_tiles_paired_up, %c2 : index + %projection_row_groups = index.mul %row_groups, %c2 : index + %cols_up = index.add %cols, %m_m1 : index + %col_tiles = index.div %cols_up, %k_m : index + kernel.launch.config workgroups(%col_tiles, %projection_row_groups, %unit) workgroup_size(%wg, %unit, %unit) : index +} launch(%first_weight: buffer, %second_weight: buffer, %first_output: buffer, %second_output: buffer, %aq: buffer, %as: buffer, %asum: buffer) { + %base = index.constant 0 : offset + %lds_bytes = index.constant 3072 : offset + %scratch = buffer.alloca align(16) %lds_bytes : buffer + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c17 = index.constant 17 : index + %c34 = index.constant 34 : index + %k_kc = index.constant 64 : index + %k_astride = index.constant 40 : index + %k_wstride = index.constant 40 : index + %k_m = index.constant 16 : index + %k_n = index.constant 16 : index + %k_blocks = index.constant 2 : index + %k_wave = index.constant 32 : index + %k_wm = index.constant 1 : index + %fzero = vector.constant 0.0 : vector<8xf32> + %izero = vector.constant 0 : vector<8xi32> + %no_split_k = scalar.constant false : i1 + %compact_q2 = scalar.constant false : i1 + + %k = config.get @ggml.mul_mat.symmetric_i4.lowrow.input_size : index + %rows = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %cols = config.get @ggml.mul_mat.symmetric_i4.lowrow.token_count : index + %k_b = index.assume %k [range(%k, 64, 32768)] : index + %rows_b = index.assume %rows [range(%rows, 1, 262144)] : index + %cols_b = index.assume %cols [range(%cols, 1, 16)] : index + %nchunks = index.div %k_b, %k_kc : index + %kblocks = index.div %k_b, %c32 : index + + %col_tile0 = kernel.workgroup.id : index + %projection_row_group0 = kernel.workgroup.id : index + %tid0 = kernel.workitem.id : index + %col_tile = index.assume %col_tile0 [range(%col_tile0, 0, 511)] : index + %projection_row_group = index.assume %projection_row_group0 [range(%projection_row_group0, 0, 16383)] : index + %projection = index.rem %projection_row_group, %c2 : index + %row_group = index.div %projection_row_group, %c2 : index + %row_tile0 = index.mul %row_group, %c2 : index + %row_tile = index.assume %row_tile0 [range(%row_tile0, 0, 16383)] : index + %use_second = index.cmp eq, %projection, %c1 : index + %tid = index.assume %tid0 [range(%tid0, 0, 63)] : index + %wave = index.div %tid, %k_wave : index + + %first_weight_g = buffer.assume.memory_space %first_weight : buffer + %second_weight_g = buffer.assume.memory_space %second_weight : buffer + %first_output_g = buffer.assume.memory_space %first_output : buffer + %second_output_g = buffer.assume.memory_space %second_output : buffer + %aq_g = buffer.assume.memory_space %aq : buffer + %as_g = buffer.assume.memory_space %as : buffer + %asum_g = buffer.assume.memory_space %asum : buffer + scf.if %use_second { + %weight_na, %output_na, %aq_na, %as_na, %asum_na = buffer.assume.noalias %second_weight_g, %second_output_g, %aq_g, %as_g, %asum_g : buffer, buffer, buffer, buffer, buffer + template.apply<@ggml.mul_mat.symmetric_i4.lowrow.m16n16.body>(%scratch, %weight_na, %output_na, %aq_na, %as_na, %asum_na, %compact_q2, %no_split_k, %output_na, %output_na, %base, %base, %c0, %c1, %c2, %c4, %c8, %c16, %c32, %c17, %c34, %k_kc, %k_astride, %k_wstride, %k_m, %k_n, %k_blocks, %k_wave, %k_wm, %fzero, %izero, %k_b, %rows_b, %cols_b, %nchunks, %kblocks, %col_tile, %row_tile, %tid, %wave, %c0, %c0, %nchunks) : (buffer, buffer, buffer, buffer, buffer, buffer, i1, i1, buffer, buffer, offset, offset, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, vector<8xf32>, vector<8xi32>, index, index, index, index, index, index, index, index, index, index, index, index) + } else { + %weight_na, %output_na, %aq_na, %as_na, %asum_na = buffer.assume.noalias %first_weight_g, %first_output_g, %aq_g, %as_g, %asum_g : buffer, buffer, buffer, buffer, buffer + template.apply<@ggml.mul_mat.symmetric_i4.lowrow.m16n16.body>(%scratch, %weight_na, %output_na, %aq_na, %as_na, %asum_na, %compact_q2, %no_split_k, %output_na, %output_na, %base, %base, %c0, %c1, %c2, %c4, %c8, %c16, %c32, %c17, %c34, %k_kc, %k_astride, %k_wstride, %k_m, %k_n, %k_blocks, %k_wave, %k_wm, %fzero, %izero, %k_b, %rows_b, %cols_b, %nchunks, %kblocks, %col_tile, %row_tile, %tid, %wave, %c0, %c0, %nchunks) : (buffer, buffer, buffer, buffer, buffer, buffer, i1, i1, buffer, buffer, offset, offset, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, vector<8xf32>, vector<8xi32>, index, index, index, index, index, index, index, index, index, index, index, index) + } + kernel.return +} + +// Exact low-column direct-dot specializations for expanding split-K projections. +func.def inline @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_full_dot(%acc: f32, %weight_words: vector<4xi32>, %weight_scale: f32, %group: index, %token: index, %k: index, %group_count: index, %qact: buffer, %scales: buffer) -> (f32) { + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c16 = index.constant 16 : index + %zero_i32 = scalar.constant 0 : i32 + %zero_i32x4 = vector.constant 0 : vector<4xi32> + %qact_flat = buffer.view %qact[%base] : buffer -> view<1073741824xi8> + %scale_flat = buffer.view %scales[%base] : buffer -> view<33554432xf32> + + %token_bytes = index.div %k, %c2 : index + %token_payload_base = index.mul %token, %token_bytes : index + %group_payload_add = index.mul %group, %c16 : index + %payload_index0 = index.add %token_payload_base, %group_payload_add : index + %payload_index = index.assume %payload_index0 [range(%payload_index0, 0, 1073741808)] : index + %activation_bytes = vector.load %qact_flat[%payload_index] : view<1073741824xi8> -> vector<16xi8> + %activation_words = vector.bitcast %activation_bytes : vector<16xi8> to vector<4xi32> + %dot_parts = vector.dot8i4 %weight_words, %activation_words, %zero_i32x4 : vector<4xi32> + %dot_i32 = vector.reduce %dot_parts, %zero_i32 : vector<4xi32>, i32 + %dot = scalar.sitofp %dot_i32 : i32 to f32 + + %token_scale_base = index.mul %token, %group_count : index + %scale_index0 = index.add %token_scale_base, %group : index + %scale_index = index.assume %scale_index0 [range(%scale_index0, 0, 33554431)] : index + %activation_scale = view.load %scale_flat[%scale_index] : view<33554432xf32> -> f32 + %combined_scale = scalar.mulf %weight_scale, %activation_scale : f32 + %next = scalar.fmaf %dot, %combined_scale, %acc : f32 + func.return %next : f32 +} + +kernel.def export("ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c1") @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c1() { + %rows = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c127 = index.constant 127 : index + %c128 = index.constant 128 : index + %rows_up = index.add %rows, %c127 : index + %row_workgroups = index.div %rows_up, %c128 : index + %workgroups = index.mul %row_workgroups, %c2 : index + kernel.launch.config workgroups(%workgroups, %c1, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%weight: buffer, %input: buffer, %output: buffer, %qact: buffer, %scales: buffer, %sums: buffer, %partial: buffer, %completion_counters: buffer) { + %k0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.input_size : index + %rows0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %cols0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.token_count : index + %k = index.assume %k0 [range(%k0, 256, 32768), mul(%k0, 256)] : index + %rows = index.assume %rows0 [range(%rows0, 1, 262144)] : index + %cols = index.assume %cols0 [range(%cols0, 1, 1)] : index + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c128 = index.constant 128 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %zero_f32 = scalar.constant 0.0 : f32 + + %split_workgroup = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %partition = index.rem %split_workgroup, %c2 : index + %workgroup = index.div %split_workgroup, %c2 : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 127)] : index + %row_base = index.mul %workgroup, %c128 : index + %row = index.add %row_base, %workitem : index + %valid_row = index.cmp ult, %row, %rows : index + + %block_count = index.div %k, %c256 : index + %group_count = index.div %k, %c32 : index + %groups_per_partition = index.div %group_count, %c2 : index + %group_begin = index.mul %partition, %groups_per_partition : index + %group_end = index.add %group_begin, %groups_per_partition : index + %row_group0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.row_group_size : index + %rows_plus = index.add %rows, %c255 : index + %row_group = index.assume %row_group0 [range(%row_group0, 32, 262144), mul(%row_group0, 32)] : index + %row_group_id = index.div %row, %row_group : index + %row_lane = index.rem %row, %row_group : index + %cohort_lane = index.div %row_lane, %c4 : index + %payload_plane_bytes = index.mul %row_group, %c16 : index + %payload_block_bytes = index.mul %payload_plane_bytes, %c8 : index + %scale_cohorts = index.div %row_group, %c4 : index + %scale_plane_bytes = index.mul %scale_cohorts, %c16 : index + %block_bytes = index.add %payload_block_bytes, %scale_plane_bytes : index + %row_group_bytes = index.mul %block_count, %block_bytes : index + %group_base = index.mul %row_group_id, %row_group_bytes : index + %scale_region_add = index.mul %block_count, %payload_block_bytes : index + %scale_region = index.add %group_base, %scale_region_add : index + + %weight_global = buffer.assume.memory_space %weight : buffer + %output_global = buffer.assume.memory_space %output : buffer + %qact_global = buffer.assume.memory_space %qact : buffer + %scales_global = buffer.assume.memory_space %scales : buffer + %partial_global = buffer.assume.memory_space %partial : buffer + %counters_global = buffer.assume.memory_space %completion_counters : buffer + %weight_na, %output_na, %qact_na, %scales_na, %partial_na, %counters_na = buffer.assume.noalias %weight_global, %output_global, %qact_global, %scales_global, %partial_global, %counters_global : buffer, buffer, buffer, buffer, buffer, buffer + + %result0 = scf.if %valid_row -> (f32) { + %acc0 = scf.for %group = [%group_begin to %group_end step %c1](%iter0 = %zero_f32 : f32) -> (f32) { + %block = index.div %group, %c8 : index + %group_in_block = index.rem %group, %c8 : index + %block_payload_add = index.mul %block, %payload_block_bytes : index + %group_payload_add = index.mul %group_in_block, %payload_plane_bytes : index + %row_payload_add = index.mul %row_lane, %c16 : index + %payload_index0 = index.add %group_base, %block_payload_add : index + %payload_index1 = index.add %payload_index0, %group_payload_add : index + %payload_index = index.add %payload_index1, %row_payload_add : index + %payload_base = index.cast %payload_index : index to offset + %weight_view = buffer.view %weight_na[%payload_base] : buffer -> view<4xi32> + %weight_words = vector.load %weight_view[%c0] : view<4xi32> -> vector<4xi32> + + %block_scale_add = index.mul %block, %scale_plane_bytes : index + %cohort_scale_add = index.mul %cohort_lane, %c16 : index + %group_scale_add = index.mul %group_in_block, %c2 : index + %scale_index0 = index.add %scale_region, %block_scale_add : index + %scale_index1 = index.add %scale_index0, %cohort_scale_add : index + %scale_byte_index = index.add %scale_index1, %group_scale_add : index + %scale_byte_base = index.cast %scale_byte_index : index to offset + %weight_scale_view = buffer.view %weight_na[%scale_byte_base] : buffer -> view<1xf16> + %weight_scale_f16 = view.load %weight_scale_view[%c0] : view<1xf16> -> f16 + %weight_scale = scalar.extf %weight_scale_f16 : f16 to f32 + + %next0 = func.call @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_full_dot(%iter0, %weight_words, %weight_scale, %group, %c0, %k, %group_count, %qact_na, %scales_na) : (f32, vector<4xi32>, f32, index, index, index, index, buffer, buffer) -> (f32) + scf.yield %next0 : f32 + } + scf.yield %acc0 : f32 + } else { + scf.yield %zero_f32 : f32 + } + + %output_view = buffer.view %output_na[%base] : buffer -> view<[%cols]x[%rows]xf32> + %partial_view = buffer.view %partial_na[%base] : buffer -> view<[%cols]x[%rows]xf32> + %is_partition_zero = index.cmp eq, %partition, %c0 : index + scf.if %valid_row { + scf.if %is_partition_zero { + view.store %result0, %output_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + } else { + view.store %result0, %partial_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + } + } + + kernel.barrier scope(workgroup) ordering(release) + %counter_count = index.div %rows_plus, %c128 : index + %counter_view = buffer.view %counters_na[%base] : buffer -> view<[%counter_count]xi32> + %is_arrival_lane = index.cmp eq, %workitem, %c0 : index + %one_i32 = scalar.constant 1 : i32 + %zero_i32 = scalar.constant 0 : i32 + %negative_two_i32 = scalar.constant -2 : i32 + %local_old_counter = scf.if %is_arrival_lane -> (i32) { + %old_counter = view.atomic.rmw %one_i32, %counter_view[%workgroup] {ordering = acq_rel, scope = device} : i32, view<[%counter_count]xi32> -> i32 + scf.yield %old_counter : i32 + } else { + scf.yield %zero_i32 : i32 + } + %old_counter = kernel.workgroup.reduce %local_old_counter : i32 + %is_last_partition = scalar.cmpi eq, %old_counter, %one_i32 : i32 + scf.if %is_last_partition { + kernel.barrier scope(workgroup) ordering(acquire) + scf.if %valid_row { + %partition_zero0 = view.load %output_view[%c0, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_one0 = view.load %partial_view[%c0, %row] : view<[%cols]x[%rows]xf32> -> f32 + %combined0 = scalar.addf %partition_zero0, %partition_one0 : f32 + view.store %combined0, %output_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + } + kernel.barrier scope(workgroup) ordering(release) + scf.if %is_arrival_lane { + view.atomic.reduce %negative_two_i32, %counter_view[%workgroup] {ordering = release, scope = device} : i32, view<[%counter_count]xi32> + } + } + kernel.return +} + +kernel.def export("ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c2") @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c2() { + %rows = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c127 = index.constant 127 : index + %c128 = index.constant 128 : index + %rows_up = index.add %rows, %c127 : index + %row_workgroups = index.div %rows_up, %c128 : index + %workgroups = index.mul %row_workgroups, %c2 : index + kernel.launch.config workgroups(%workgroups, %c1, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%weight: buffer, %input: buffer, %output: buffer, %qact: buffer, %scales: buffer, %sums: buffer, %partial: buffer, %completion_counters: buffer) { + %k0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.input_size : index + %rows0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %cols0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.token_count : index + %k = index.assume %k0 [range(%k0, 256, 32768), mul(%k0, 256)] : index + %rows = index.assume %rows0 [range(%rows0, 1, 262144)] : index + %cols = index.assume %cols0 [range(%cols0, 2, 2)] : index + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c128 = index.constant 128 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %zero_f32 = scalar.constant 0.0 : f32 + + %split_workgroup = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %partition = index.rem %split_workgroup, %c2 : index + %workgroup = index.div %split_workgroup, %c2 : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 127)] : index + %row_base = index.mul %workgroup, %c128 : index + %row = index.add %row_base, %workitem : index + %valid_row = index.cmp ult, %row, %rows : index + + %block_count = index.div %k, %c256 : index + %group_count = index.div %k, %c32 : index + %groups_per_partition = index.div %group_count, %c2 : index + %group_begin = index.mul %partition, %groups_per_partition : index + %group_end = index.add %group_begin, %groups_per_partition : index + %row_group0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.row_group_size : index + %rows_plus = index.add %rows, %c255 : index + %row_group = index.assume %row_group0 [range(%row_group0, 32, 262144), mul(%row_group0, 32)] : index + %row_group_id = index.div %row, %row_group : index + %row_lane = index.rem %row, %row_group : index + %cohort_lane = index.div %row_lane, %c4 : index + %payload_plane_bytes = index.mul %row_group, %c16 : index + %payload_block_bytes = index.mul %payload_plane_bytes, %c8 : index + %scale_cohorts = index.div %row_group, %c4 : index + %scale_plane_bytes = index.mul %scale_cohorts, %c16 : index + %block_bytes = index.add %payload_block_bytes, %scale_plane_bytes : index + %row_group_bytes = index.mul %block_count, %block_bytes : index + %group_base = index.mul %row_group_id, %row_group_bytes : index + %scale_region_add = index.mul %block_count, %payload_block_bytes : index + %scale_region = index.add %group_base, %scale_region_add : index + + %weight_global = buffer.assume.memory_space %weight : buffer + %output_global = buffer.assume.memory_space %output : buffer + %qact_global = buffer.assume.memory_space %qact : buffer + %scales_global = buffer.assume.memory_space %scales : buffer + %partial_global = buffer.assume.memory_space %partial : buffer + %counters_global = buffer.assume.memory_space %completion_counters : buffer + %weight_na, %output_na, %qact_na, %scales_na, %partial_na, %counters_na = buffer.assume.noalias %weight_global, %output_global, %qact_global, %scales_global, %partial_global, %counters_global : buffer, buffer, buffer, buffer, buffer, buffer + + %result0, %result1 = scf.if %valid_row -> (f32, f32) { + %acc0, %acc1 = scf.for %group = [%group_begin to %group_end step %c1](%iter0 = %zero_f32 : f32, %iter1 = %zero_f32 : f32) -> (f32, f32) { + %block = index.div %group, %c8 : index + %group_in_block = index.rem %group, %c8 : index + %block_payload_add = index.mul %block, %payload_block_bytes : index + %group_payload_add = index.mul %group_in_block, %payload_plane_bytes : index + %row_payload_add = index.mul %row_lane, %c16 : index + %payload_index0 = index.add %group_base, %block_payload_add : index + %payload_index1 = index.add %payload_index0, %group_payload_add : index + %payload_index = index.add %payload_index1, %row_payload_add : index + %payload_base = index.cast %payload_index : index to offset + %weight_view = buffer.view %weight_na[%payload_base] : buffer -> view<4xi32> + %weight_words = vector.load %weight_view[%c0] : view<4xi32> -> vector<4xi32> + + %block_scale_add = index.mul %block, %scale_plane_bytes : index + %cohort_scale_add = index.mul %cohort_lane, %c16 : index + %group_scale_add = index.mul %group_in_block, %c2 : index + %scale_index0 = index.add %scale_region, %block_scale_add : index + %scale_index1 = index.add %scale_index0, %cohort_scale_add : index + %scale_byte_index = index.add %scale_index1, %group_scale_add : index + %scale_byte_base = index.cast %scale_byte_index : index to offset + %weight_scale_view = buffer.view %weight_na[%scale_byte_base] : buffer -> view<1xf16> + %weight_scale_f16 = view.load %weight_scale_view[%c0] : view<1xf16> -> f16 + %weight_scale = scalar.extf %weight_scale_f16 : f16 to f32 + + %next0 = func.call @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_full_dot(%iter0, %weight_words, %weight_scale, %group, %c0, %k, %group_count, %qact_na, %scales_na) : (f32, vector<4xi32>, f32, index, index, index, index, buffer, buffer) -> (f32) + %next1 = func.call @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_full_dot(%iter1, %weight_words, %weight_scale, %group, %c1, %k, %group_count, %qact_na, %scales_na) : (f32, vector<4xi32>, f32, index, index, index, index, buffer, buffer) -> (f32) + scf.yield %next0, %next1 : f32, f32 + } + scf.yield %acc0, %acc1 : f32, f32 + } else { + scf.yield %zero_f32, %zero_f32 : f32, f32 + } + + %output_view = buffer.view %output_na[%base] : buffer -> view<[%cols]x[%rows]xf32> + %partial_view = buffer.view %partial_na[%base] : buffer -> view<[%cols]x[%rows]xf32> + %is_partition_zero = index.cmp eq, %partition, %c0 : index + scf.if %valid_row { + scf.if %is_partition_zero { + view.store %result0, %output_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result1, %output_view[%c1, %row] : f32, view<[%cols]x[%rows]xf32> + } else { + view.store %result0, %partial_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result1, %partial_view[%c1, %row] : f32, view<[%cols]x[%rows]xf32> + } + } + + kernel.barrier scope(workgroup) ordering(release) + %counter_count = index.div %rows_plus, %c128 : index + %counter_view = buffer.view %counters_na[%base] : buffer -> view<[%counter_count]xi32> + %is_arrival_lane = index.cmp eq, %workitem, %c0 : index + %one_i32 = scalar.constant 1 : i32 + %zero_i32 = scalar.constant 0 : i32 + %negative_two_i32 = scalar.constant -2 : i32 + %local_old_counter = scf.if %is_arrival_lane -> (i32) { + %old_counter = view.atomic.rmw %one_i32, %counter_view[%workgroup] {ordering = acq_rel, scope = device} : i32, view<[%counter_count]xi32> -> i32 + scf.yield %old_counter : i32 + } else { + scf.yield %zero_i32 : i32 + } + %old_counter = kernel.workgroup.reduce %local_old_counter : i32 + %is_last_partition = scalar.cmpi eq, %old_counter, %one_i32 : i32 + scf.if %is_last_partition { + kernel.barrier scope(workgroup) ordering(acquire) + scf.if %valid_row { + %partition_zero0 = view.load %output_view[%c0, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_zero1 = view.load %output_view[%c1, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_one0 = view.load %partial_view[%c0, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_one1 = view.load %partial_view[%c1, %row] : view<[%cols]x[%rows]xf32> -> f32 + %combined0 = scalar.addf %partition_zero0, %partition_one0 : f32 + %combined1 = scalar.addf %partition_zero1, %partition_one1 : f32 + view.store %combined0, %output_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %combined1, %output_view[%c1, %row] : f32, view<[%cols]x[%rows]xf32> + } + kernel.barrier scope(workgroup) ordering(release) + scf.if %is_arrival_lane { + view.atomic.reduce %negative_two_i32, %counter_view[%workgroup] {ordering = release, scope = device} : i32, view<[%counter_count]xi32> + } + } + kernel.return +} + +kernel.def export("ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c3") @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c3() { + %rows = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c127 = index.constant 127 : index + %c128 = index.constant 128 : index + %rows_up = index.add %rows, %c127 : index + %row_workgroups = index.div %rows_up, %c128 : index + %workgroups = index.mul %row_workgroups, %c2 : index + kernel.launch.config workgroups(%workgroups, %c1, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%weight: buffer, %input: buffer, %output: buffer, %qact: buffer, %scales: buffer, %sums: buffer, %partial: buffer, %completion_counters: buffer) { + %k0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.input_size : index + %rows0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %cols0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.token_count : index + %k = index.assume %k0 [range(%k0, 256, 32768), mul(%k0, 256)] : index + %rows = index.assume %rows0 [range(%rows0, 1, 262144)] : index + %cols = index.assume %cols0 [range(%cols0, 3, 3)] : index + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c128 = index.constant 128 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %zero_f32 = scalar.constant 0.0 : f32 + + %split_workgroup = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %partition = index.rem %split_workgroup, %c2 : index + %workgroup = index.div %split_workgroup, %c2 : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 127)] : index + %row_base = index.mul %workgroup, %c128 : index + %row = index.add %row_base, %workitem : index + %valid_row = index.cmp ult, %row, %rows : index + + %block_count = index.div %k, %c256 : index + %group_count = index.div %k, %c32 : index + %groups_per_partition = index.div %group_count, %c2 : index + %group_begin = index.mul %partition, %groups_per_partition : index + %group_end = index.add %group_begin, %groups_per_partition : index + %row_group0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.row_group_size : index + %rows_plus = index.add %rows, %c255 : index + %row_group = index.assume %row_group0 [range(%row_group0, 32, 262144), mul(%row_group0, 32)] : index + %row_group_id = index.div %row, %row_group : index + %row_lane = index.rem %row, %row_group : index + %cohort_lane = index.div %row_lane, %c4 : index + %payload_plane_bytes = index.mul %row_group, %c16 : index + %payload_block_bytes = index.mul %payload_plane_bytes, %c8 : index + %scale_cohorts = index.div %row_group, %c4 : index + %scale_plane_bytes = index.mul %scale_cohorts, %c16 : index + %block_bytes = index.add %payload_block_bytes, %scale_plane_bytes : index + %row_group_bytes = index.mul %block_count, %block_bytes : index + %group_base = index.mul %row_group_id, %row_group_bytes : index + %scale_region_add = index.mul %block_count, %payload_block_bytes : index + %scale_region = index.add %group_base, %scale_region_add : index + + %weight_global = buffer.assume.memory_space %weight : buffer + %output_global = buffer.assume.memory_space %output : buffer + %qact_global = buffer.assume.memory_space %qact : buffer + %scales_global = buffer.assume.memory_space %scales : buffer + %partial_global = buffer.assume.memory_space %partial : buffer + %counters_global = buffer.assume.memory_space %completion_counters : buffer + %weight_na, %output_na, %qact_na, %scales_na, %partial_na, %counters_na = buffer.assume.noalias %weight_global, %output_global, %qact_global, %scales_global, %partial_global, %counters_global : buffer, buffer, buffer, buffer, buffer, buffer + + %result0, %result1, %result2 = scf.if %valid_row -> (f32, f32, f32) { + %acc0, %acc1, %acc2 = scf.for %group = [%group_begin to %group_end step %c1](%iter0 = %zero_f32 : f32, %iter1 = %zero_f32 : f32, %iter2 = %zero_f32 : f32) -> (f32, f32, f32) { + %block = index.div %group, %c8 : index + %group_in_block = index.rem %group, %c8 : index + %block_payload_add = index.mul %block, %payload_block_bytes : index + %group_payload_add = index.mul %group_in_block, %payload_plane_bytes : index + %row_payload_add = index.mul %row_lane, %c16 : index + %payload_index0 = index.add %group_base, %block_payload_add : index + %payload_index1 = index.add %payload_index0, %group_payload_add : index + %payload_index = index.add %payload_index1, %row_payload_add : index + %payload_base = index.cast %payload_index : index to offset + %weight_view = buffer.view %weight_na[%payload_base] : buffer -> view<4xi32> + %weight_words = vector.load %weight_view[%c0] : view<4xi32> -> vector<4xi32> + + %block_scale_add = index.mul %block, %scale_plane_bytes : index + %cohort_scale_add = index.mul %cohort_lane, %c16 : index + %group_scale_add = index.mul %group_in_block, %c2 : index + %scale_index0 = index.add %scale_region, %block_scale_add : index + %scale_index1 = index.add %scale_index0, %cohort_scale_add : index + %scale_byte_index = index.add %scale_index1, %group_scale_add : index + %scale_byte_base = index.cast %scale_byte_index : index to offset + %weight_scale_view = buffer.view %weight_na[%scale_byte_base] : buffer -> view<1xf16> + %weight_scale_f16 = view.load %weight_scale_view[%c0] : view<1xf16> -> f16 + %weight_scale = scalar.extf %weight_scale_f16 : f16 to f32 + + %next0 = func.call @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_full_dot(%iter0, %weight_words, %weight_scale, %group, %c0, %k, %group_count, %qact_na, %scales_na) : (f32, vector<4xi32>, f32, index, index, index, index, buffer, buffer) -> (f32) + %next1 = func.call @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_full_dot(%iter1, %weight_words, %weight_scale, %group, %c1, %k, %group_count, %qact_na, %scales_na) : (f32, vector<4xi32>, f32, index, index, index, index, buffer, buffer) -> (f32) + %next2 = func.call @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_full_dot(%iter2, %weight_words, %weight_scale, %group, %c2, %k, %group_count, %qact_na, %scales_na) : (f32, vector<4xi32>, f32, index, index, index, index, buffer, buffer) -> (f32) + scf.yield %next0, %next1, %next2 : f32, f32, f32 + } + scf.yield %acc0, %acc1, %acc2 : f32, f32, f32 + } else { + scf.yield %zero_f32, %zero_f32, %zero_f32 : f32, f32, f32 + } + + %output_view = buffer.view %output_na[%base] : buffer -> view<[%cols]x[%rows]xf32> + %partial_view = buffer.view %partial_na[%base] : buffer -> view<[%cols]x[%rows]xf32> + %is_partition_zero = index.cmp eq, %partition, %c0 : index + scf.if %valid_row { + scf.if %is_partition_zero { + view.store %result0, %output_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result1, %output_view[%c1, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result2, %output_view[%c2, %row] : f32, view<[%cols]x[%rows]xf32> + } else { + view.store %result0, %partial_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result1, %partial_view[%c1, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result2, %partial_view[%c2, %row] : f32, view<[%cols]x[%rows]xf32> + } + } + + kernel.barrier scope(workgroup) ordering(release) + %counter_count = index.div %rows_plus, %c128 : index + %counter_view = buffer.view %counters_na[%base] : buffer -> view<[%counter_count]xi32> + %is_arrival_lane = index.cmp eq, %workitem, %c0 : index + %one_i32 = scalar.constant 1 : i32 + %zero_i32 = scalar.constant 0 : i32 + %negative_two_i32 = scalar.constant -2 : i32 + %local_old_counter = scf.if %is_arrival_lane -> (i32) { + %old_counter = view.atomic.rmw %one_i32, %counter_view[%workgroup] {ordering = acq_rel, scope = device} : i32, view<[%counter_count]xi32> -> i32 + scf.yield %old_counter : i32 + } else { + scf.yield %zero_i32 : i32 + } + %old_counter = kernel.workgroup.reduce %local_old_counter : i32 + %is_last_partition = scalar.cmpi eq, %old_counter, %one_i32 : i32 + scf.if %is_last_partition { + kernel.barrier scope(workgroup) ordering(acquire) + scf.if %valid_row { + %partition_zero0 = view.load %output_view[%c0, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_zero1 = view.load %output_view[%c1, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_zero2 = view.load %output_view[%c2, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_one0 = view.load %partial_view[%c0, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_one1 = view.load %partial_view[%c1, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_one2 = view.load %partial_view[%c2, %row] : view<[%cols]x[%rows]xf32> -> f32 + %combined0 = scalar.addf %partition_zero0, %partition_one0 : f32 + %combined1 = scalar.addf %partition_zero1, %partition_one1 : f32 + %combined2 = scalar.addf %partition_zero2, %partition_one2 : f32 + view.store %combined0, %output_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %combined1, %output_view[%c1, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %combined2, %output_view[%c2, %row] : f32, view<[%cols]x[%rows]xf32> + } + kernel.barrier scope(workgroup) ordering(release) + scf.if %is_arrival_lane { + view.atomic.reduce %negative_two_i32, %counter_view[%workgroup] {ordering = release, scope = device} : i32, view<[%counter_count]xi32> + } + } + kernel.return +} + +kernel.def export("ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c4") @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c4() { + %rows = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c127 = index.constant 127 : index + %c128 = index.constant 128 : index + %rows_up = index.add %rows, %c127 : index + %row_workgroups = index.div %rows_up, %c128 : index + %workgroups = index.mul %row_workgroups, %c2 : index + kernel.launch.config workgroups(%workgroups, %c1, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%weight: buffer, %input: buffer, %output: buffer, %qact: buffer, %scales: buffer, %sums: buffer, %partial: buffer, %completion_counters: buffer) { + %k0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.input_size : index + %rows0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %cols0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.token_count : index + %k = index.assume %k0 [range(%k0, 256, 32768), mul(%k0, 256)] : index + %rows = index.assume %rows0 [range(%rows0, 1, 262144)] : index + %cols = index.assume %cols0 [range(%cols0, 4, 4)] : index + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c128 = index.constant 128 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %zero_f32 = scalar.constant 0.0 : f32 + + %split_workgroup = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %partition = index.rem %split_workgroup, %c2 : index + %workgroup = index.div %split_workgroup, %c2 : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 127)] : index + %row_base = index.mul %workgroup, %c128 : index + %row = index.add %row_base, %workitem : index + %valid_row = index.cmp ult, %row, %rows : index + + %block_count = index.div %k, %c256 : index + %group_count = index.div %k, %c32 : index + %groups_per_partition = index.div %group_count, %c2 : index + %group_begin = index.mul %partition, %groups_per_partition : index + %group_end = index.add %group_begin, %groups_per_partition : index + %row_group0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.row_group_size : index + %rows_plus = index.add %rows, %c255 : index + %row_group = index.assume %row_group0 [range(%row_group0, 32, 262144), mul(%row_group0, 32)] : index + %row_group_id = index.div %row, %row_group : index + %row_lane = index.rem %row, %row_group : index + %cohort_lane = index.div %row_lane, %c4 : index + %payload_plane_bytes = index.mul %row_group, %c16 : index + %payload_block_bytes = index.mul %payload_plane_bytes, %c8 : index + %scale_cohorts = index.div %row_group, %c4 : index + %scale_plane_bytes = index.mul %scale_cohorts, %c16 : index + %block_bytes = index.add %payload_block_bytes, %scale_plane_bytes : index + %row_group_bytes = index.mul %block_count, %block_bytes : index + %group_base = index.mul %row_group_id, %row_group_bytes : index + %scale_region_add = index.mul %block_count, %payload_block_bytes : index + %scale_region = index.add %group_base, %scale_region_add : index + + %weight_global = buffer.assume.memory_space %weight : buffer + %output_global = buffer.assume.memory_space %output : buffer + %qact_global = buffer.assume.memory_space %qact : buffer + %scales_global = buffer.assume.memory_space %scales : buffer + %partial_global = buffer.assume.memory_space %partial : buffer + %counters_global = buffer.assume.memory_space %completion_counters : buffer + %weight_na, %output_na, %qact_na, %scales_na, %partial_na, %counters_na = buffer.assume.noalias %weight_global, %output_global, %qact_global, %scales_global, %partial_global, %counters_global : buffer, buffer, buffer, buffer, buffer, buffer + + %result0, %result1, %result2, %result3 = scf.if %valid_row -> (f32, f32, f32, f32) { + %acc0, %acc1, %acc2, %acc3 = scf.for %group = [%group_begin to %group_end step %c1](%iter0 = %zero_f32 : f32, %iter1 = %zero_f32 : f32, %iter2 = %zero_f32 : f32, %iter3 = %zero_f32 : f32) -> (f32, f32, f32, f32) { + %block = index.div %group, %c8 : index + %group_in_block = index.rem %group, %c8 : index + %block_payload_add = index.mul %block, %payload_block_bytes : index + %group_payload_add = index.mul %group_in_block, %payload_plane_bytes : index + %row_payload_add = index.mul %row_lane, %c16 : index + %payload_index0 = index.add %group_base, %block_payload_add : index + %payload_index1 = index.add %payload_index0, %group_payload_add : index + %payload_index = index.add %payload_index1, %row_payload_add : index + %payload_base = index.cast %payload_index : index to offset + %weight_view = buffer.view %weight_na[%payload_base] : buffer -> view<4xi32> + %weight_words = vector.load %weight_view[%c0] : view<4xi32> -> vector<4xi32> + + %block_scale_add = index.mul %block, %scale_plane_bytes : index + %cohort_scale_add = index.mul %cohort_lane, %c16 : index + %group_scale_add = index.mul %group_in_block, %c2 : index + %scale_index0 = index.add %scale_region, %block_scale_add : index + %scale_index1 = index.add %scale_index0, %cohort_scale_add : index + %scale_byte_index = index.add %scale_index1, %group_scale_add : index + %scale_byte_base = index.cast %scale_byte_index : index to offset + %weight_scale_view = buffer.view %weight_na[%scale_byte_base] : buffer -> view<1xf16> + %weight_scale_f16 = view.load %weight_scale_view[%c0] : view<1xf16> -> f16 + %weight_scale = scalar.extf %weight_scale_f16 : f16 to f32 + + %next0 = func.call @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_full_dot(%iter0, %weight_words, %weight_scale, %group, %c0, %k, %group_count, %qact_na, %scales_na) : (f32, vector<4xi32>, f32, index, index, index, index, buffer, buffer) -> (f32) + %next1 = func.call @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_full_dot(%iter1, %weight_words, %weight_scale, %group, %c1, %k, %group_count, %qact_na, %scales_na) : (f32, vector<4xi32>, f32, index, index, index, index, buffer, buffer) -> (f32) + %next2 = func.call @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_full_dot(%iter2, %weight_words, %weight_scale, %group, %c2, %k, %group_count, %qact_na, %scales_na) : (f32, vector<4xi32>, f32, index, index, index, index, buffer, buffer) -> (f32) + %next3 = func.call @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_full_dot(%iter3, %weight_words, %weight_scale, %group, %c3, %k, %group_count, %qact_na, %scales_na) : (f32, vector<4xi32>, f32, index, index, index, index, buffer, buffer) -> (f32) + scf.yield %next0, %next1, %next2, %next3 : f32, f32, f32, f32 + } + scf.yield %acc0, %acc1, %acc2, %acc3 : f32, f32, f32, f32 + } else { + scf.yield %zero_f32, %zero_f32, %zero_f32, %zero_f32 : f32, f32, f32, f32 + } + + %output_view = buffer.view %output_na[%base] : buffer -> view<[%cols]x[%rows]xf32> + %partial_view = buffer.view %partial_na[%base] : buffer -> view<[%cols]x[%rows]xf32> + %is_partition_zero = index.cmp eq, %partition, %c0 : index + scf.if %valid_row { + scf.if %is_partition_zero { + view.store %result0, %output_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result1, %output_view[%c1, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result2, %output_view[%c2, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result3, %output_view[%c3, %row] : f32, view<[%cols]x[%rows]xf32> + } else { + view.store %result0, %partial_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result1, %partial_view[%c1, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result2, %partial_view[%c2, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result3, %partial_view[%c3, %row] : f32, view<[%cols]x[%rows]xf32> + } + } + + kernel.barrier scope(workgroup) ordering(release) + %counter_count = index.div %rows_plus, %c128 : index + %counter_view = buffer.view %counters_na[%base] : buffer -> view<[%counter_count]xi32> + %is_arrival_lane = index.cmp eq, %workitem, %c0 : index + %one_i32 = scalar.constant 1 : i32 + %zero_i32 = scalar.constant 0 : i32 + %negative_two_i32 = scalar.constant -2 : i32 + %local_old_counter = scf.if %is_arrival_lane -> (i32) { + %old_counter = view.atomic.rmw %one_i32, %counter_view[%workgroup] {ordering = acq_rel, scope = device} : i32, view<[%counter_count]xi32> -> i32 + scf.yield %old_counter : i32 + } else { + scf.yield %zero_i32 : i32 + } + %old_counter = kernel.workgroup.reduce %local_old_counter : i32 + %is_last_partition = scalar.cmpi eq, %old_counter, %one_i32 : i32 + scf.if %is_last_partition { + kernel.barrier scope(workgroup) ordering(acquire) + scf.if %valid_row { + %partition_zero0 = view.load %output_view[%c0, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_zero1 = view.load %output_view[%c1, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_zero2 = view.load %output_view[%c2, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_zero3 = view.load %output_view[%c3, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_one0 = view.load %partial_view[%c0, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_one1 = view.load %partial_view[%c1, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_one2 = view.load %partial_view[%c2, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_one3 = view.load %partial_view[%c3, %row] : view<[%cols]x[%rows]xf32> -> f32 + %combined0 = scalar.addf %partition_zero0, %partition_one0 : f32 + %combined1 = scalar.addf %partition_zero1, %partition_one1 : f32 + %combined2 = scalar.addf %partition_zero2, %partition_one2 : f32 + %combined3 = scalar.addf %partition_zero3, %partition_one3 : f32 + view.store %combined0, %output_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %combined1, %output_view[%c1, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %combined2, %output_view[%c2, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %combined3, %output_view[%c3, %row] : f32, view<[%cols]x[%rows]xf32> + } + kernel.barrier scope(workgroup) ordering(release) + scf.if %is_arrival_lane { + view.atomic.reduce %negative_two_i32, %counter_view[%workgroup] {ordering = release, scope = device} : i32, view<[%counter_count]xi32> + } + } + kernel.return +} + +kernel.def export("ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c5") @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_c5() { + %rows = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c127 = index.constant 127 : index + %c128 = index.constant 128 : index + %rows_up = index.add %rows, %c127 : index + %row_workgroups = index.div %rows_up, %c128 : index + %workgroups = index.mul %row_workgroups, %c2 : index + kernel.launch.config workgroups(%workgroups, %c1, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%weight: buffer, %input: buffer, %output: buffer, %qact: buffer, %scales: buffer, %sums: buffer, %partial: buffer, %completion_counters: buffer) { + %k0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.input_size : index + %rows0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %cols0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.token_count : index + %k = index.assume %k0 [range(%k0, 256, 32768), mul(%k0, 256)] : index + %rows = index.assume %rows0 [range(%rows0, 1, 262144)] : index + %cols = index.assume %cols0 [range(%cols0, 5, 5)] : index + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c128 = index.constant 128 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %zero_f32 = scalar.constant 0.0 : f32 + + %split_workgroup = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %partition = index.rem %split_workgroup, %c2 : index + %workgroup = index.div %split_workgroup, %c2 : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 127)] : index + %row_base = index.mul %workgroup, %c128 : index + %row = index.add %row_base, %workitem : index + %valid_row = index.cmp ult, %row, %rows : index + + %block_count = index.div %k, %c256 : index + %group_count = index.div %k, %c32 : index + %groups_per_partition = index.div %group_count, %c2 : index + %group_begin = index.mul %partition, %groups_per_partition : index + %group_end = index.add %group_begin, %groups_per_partition : index + %row_group0 = config.get @ggml.mul_mat.symmetric_i4.lowrow.row_group_size : index + %rows_plus = index.add %rows, %c255 : index + %row_group = index.assume %row_group0 [range(%row_group0, 32, 262144), mul(%row_group0, 32)] : index + %row_group_id = index.div %row, %row_group : index + %row_lane = index.rem %row, %row_group : index + %cohort_lane = index.div %row_lane, %c4 : index + %payload_plane_bytes = index.mul %row_group, %c16 : index + %payload_block_bytes = index.mul %payload_plane_bytes, %c8 : index + %scale_cohorts = index.div %row_group, %c4 : index + %scale_plane_bytes = index.mul %scale_cohorts, %c16 : index + %block_bytes = index.add %payload_block_bytes, %scale_plane_bytes : index + %row_group_bytes = index.mul %block_count, %block_bytes : index + %group_base = index.mul %row_group_id, %row_group_bytes : index + %scale_region_add = index.mul %block_count, %payload_block_bytes : index + %scale_region = index.add %group_base, %scale_region_add : index + + %weight_global = buffer.assume.memory_space %weight : buffer + %output_global = buffer.assume.memory_space %output : buffer + %qact_global = buffer.assume.memory_space %qact : buffer + %scales_global = buffer.assume.memory_space %scales : buffer + %partial_global = buffer.assume.memory_space %partial : buffer + %counters_global = buffer.assume.memory_space %completion_counters : buffer + %weight_na, %output_na, %qact_na, %scales_na, %partial_na, %counters_na = buffer.assume.noalias %weight_global, %output_global, %qact_global, %scales_global, %partial_global, %counters_global : buffer, buffer, buffer, buffer, buffer, buffer + + %result0, %result1, %result2, %result3, %result4 = scf.if %valid_row -> (f32, f32, f32, f32, f32) { + %acc0, %acc1, %acc2, %acc3, %acc4 = scf.for %group = [%group_begin to %group_end step %c1](%iter0 = %zero_f32 : f32, %iter1 = %zero_f32 : f32, %iter2 = %zero_f32 : f32, %iter3 = %zero_f32 : f32, %iter4 = %zero_f32 : f32) -> (f32, f32, f32, f32, f32) { + %block = index.div %group, %c8 : index + %group_in_block = index.rem %group, %c8 : index + %block_payload_add = index.mul %block, %payload_block_bytes : index + %group_payload_add = index.mul %group_in_block, %payload_plane_bytes : index + %row_payload_add = index.mul %row_lane, %c16 : index + %payload_index0 = index.add %group_base, %block_payload_add : index + %payload_index1 = index.add %payload_index0, %group_payload_add : index + %payload_index = index.add %payload_index1, %row_payload_add : index + %payload_base = index.cast %payload_index : index to offset + %weight_view = buffer.view %weight_na[%payload_base] : buffer -> view<4xi32> + %weight_words = vector.load %weight_view[%c0] : view<4xi32> -> vector<4xi32> + + %block_scale_add = index.mul %block, %scale_plane_bytes : index + %cohort_scale_add = index.mul %cohort_lane, %c16 : index + %group_scale_add = index.mul %group_in_block, %c2 : index + %scale_index0 = index.add %scale_region, %block_scale_add : index + %scale_index1 = index.add %scale_index0, %cohort_scale_add : index + %scale_byte_index = index.add %scale_index1, %group_scale_add : index + %scale_byte_base = index.cast %scale_byte_index : index to offset + %weight_scale_view = buffer.view %weight_na[%scale_byte_base] : buffer -> view<1xf16> + %weight_scale_f16 = view.load %weight_scale_view[%c0] : view<1xf16> -> f16 + %weight_scale = scalar.extf %weight_scale_f16 : f16 to f32 + + %next0 = func.call @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_full_dot(%iter0, %weight_words, %weight_scale, %group, %c0, %k, %group_count, %qact_na, %scales_na) : (f32, vector<4xi32>, f32, index, index, index, index, buffer, buffer) -> (f32) + %next1 = func.call @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_full_dot(%iter1, %weight_words, %weight_scale, %group, %c1, %k, %group_count, %qact_na, %scales_na) : (f32, vector<4xi32>, f32, index, index, index, index, buffer, buffer) -> (f32) + %next2 = func.call @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_full_dot(%iter2, %weight_words, %weight_scale, %group, %c2, %k, %group_count, %qact_na, %scales_na) : (f32, vector<4xi32>, f32, index, index, index, index, buffer, buffer) -> (f32) + %next3 = func.call @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_full_dot(%iter3, %weight_words, %weight_scale, %group, %c3, %k, %group_count, %qact_na, %scales_na) : (f32, vector<4xi32>, f32, index, index, index, index, buffer, buffer) -> (f32) + %next4 = func.call @ggml_mul_mat_symmetric_i4_lowrow_split_k2_direct_dot_full_dot(%iter4, %weight_words, %weight_scale, %group, %c4, %k, %group_count, %qact_na, %scales_na) : (f32, vector<4xi32>, f32, index, index, index, index, buffer, buffer) -> (f32) + scf.yield %next0, %next1, %next2, %next3, %next4 : f32, f32, f32, f32, f32 + } + scf.yield %acc0, %acc1, %acc2, %acc3, %acc4 : f32, f32, f32, f32, f32 + } else { + scf.yield %zero_f32, %zero_f32, %zero_f32, %zero_f32, %zero_f32 : f32, f32, f32, f32, f32 + } + + %output_view = buffer.view %output_na[%base] : buffer -> view<[%cols]x[%rows]xf32> + %partial_view = buffer.view %partial_na[%base] : buffer -> view<[%cols]x[%rows]xf32> + %is_partition_zero = index.cmp eq, %partition, %c0 : index + scf.if %valid_row { + scf.if %is_partition_zero { + view.store %result0, %output_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result1, %output_view[%c1, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result2, %output_view[%c2, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result3, %output_view[%c3, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result4, %output_view[%c4, %row] : f32, view<[%cols]x[%rows]xf32> + } else { + view.store %result0, %partial_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result1, %partial_view[%c1, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result2, %partial_view[%c2, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result3, %partial_view[%c3, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %result4, %partial_view[%c4, %row] : f32, view<[%cols]x[%rows]xf32> + } + } + + kernel.barrier scope(workgroup) ordering(release) + %counter_count = index.div %rows_plus, %c128 : index + %counter_view = buffer.view %counters_na[%base] : buffer -> view<[%counter_count]xi32> + %is_arrival_lane = index.cmp eq, %workitem, %c0 : index + %one_i32 = scalar.constant 1 : i32 + %zero_i32 = scalar.constant 0 : i32 + %negative_two_i32 = scalar.constant -2 : i32 + %local_old_counter = scf.if %is_arrival_lane -> (i32) { + %old_counter = view.atomic.rmw %one_i32, %counter_view[%workgroup] {ordering = acq_rel, scope = device} : i32, view<[%counter_count]xi32> -> i32 + scf.yield %old_counter : i32 + } else { + scf.yield %zero_i32 : i32 + } + %old_counter = kernel.workgroup.reduce %local_old_counter : i32 + %is_last_partition = scalar.cmpi eq, %old_counter, %one_i32 : i32 + scf.if %is_last_partition { + kernel.barrier scope(workgroup) ordering(acquire) + scf.if %valid_row { + %partition_zero0 = view.load %output_view[%c0, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_zero1 = view.load %output_view[%c1, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_zero2 = view.load %output_view[%c2, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_zero3 = view.load %output_view[%c3, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_zero4 = view.load %output_view[%c4, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_one0 = view.load %partial_view[%c0, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_one1 = view.load %partial_view[%c1, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_one2 = view.load %partial_view[%c2, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_one3 = view.load %partial_view[%c3, %row] : view<[%cols]x[%rows]xf32> -> f32 + %partition_one4 = view.load %partial_view[%c4, %row] : view<[%cols]x[%rows]xf32> -> f32 + %combined0 = scalar.addf %partition_zero0, %partition_one0 : f32 + %combined1 = scalar.addf %partition_zero1, %partition_one1 : f32 + %combined2 = scalar.addf %partition_zero2, %partition_one2 : f32 + %combined3 = scalar.addf %partition_zero3, %partition_one3 : f32 + %combined4 = scalar.addf %partition_zero4, %partition_one4 : f32 + view.store %combined0, %output_view[%c0, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %combined1, %output_view[%c1, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %combined2, %output_view[%c2, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %combined3, %output_view[%c3, %row] : f32, view<[%cols]x[%rows]xf32> + view.store %combined4, %output_view[%c4, %row] : f32, view<[%cols]x[%rows]xf32> + } + kernel.barrier scope(workgroup) ordering(release) + scf.if %is_arrival_lane { + view.atomic.reduce %negative_two_i32, %counter_view[%workgroup] {ordering = release, scope = device} : i32, view<[%counter_count]xi32> + } + } + kernel.return +} + +// Two K partitions increase independent low-row work while retaining the +// canonical signed-I4 contraction and fixed-order F32 publication path. +kernel.def export("ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma") @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma() { + %unit = index.constant 1 : index + %c2 = index.constant 2 : index + %k_n = index.constant 16 : index + %k_m = index.constant 16 : index + %n_m1 = index.constant 15 : index + %m_m1 = index.constant 15 : index + %wg = index.constant 64 : index + %rows = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %cols = config.get @ggml.mul_mat.symmetric_i4.lowrow.token_count : index + %rows_up = index.add %rows, %n_m1 : index + %row_tiles = index.div %rows_up, %k_n : index + %row_groups_up = index.add %row_tiles, %unit : index + %row_groups = index.div %row_groups_up, %c2 : index + %split_row_groups = index.mul %row_groups, %c2 : index + %cols_up = index.add %cols, %m_m1 : index + %col_tiles = index.div %cols_up, %k_m : index + kernel.launch.config workgroups(%col_tiles, %split_row_groups, %unit) workgroup_size(%wg, %unit, %unit) : index +} launch(%weight: buffer, %input: buffer, %output: buffer, %aq: buffer, %as: buffer, %asum: buffer, %partial: buffer, %completion_counters: buffer) { + %base = index.constant 0 : offset + %lds_bytes = index.constant 3072 : offset + %scratch = buffer.alloca align(16) %lds_bytes : buffer + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c17 = index.constant 17 : index + %c34 = index.constant 34 : index + %k_kc = index.constant 64 : index + %k_astride = index.constant 40 : index + %k_wstride = index.constant 40 : index + %k_m = index.constant 16 : index + %k_n = index.constant 16 : index + %k_blocks = index.constant 2 : index + %k_wave = index.constant 32 : index + %k_wm = index.constant 1 : index + %fzero = vector.constant 0.0 : vector<8xf32> + %izero = vector.constant 0 : vector<8xi32> + %split_k = scalar.constant true : i1 + %compact_q2 = scalar.constant false : i1 + + %k = config.get @ggml.mul_mat.symmetric_i4.lowrow.input_size : index + %rows = config.get @ggml.mul_mat.symmetric_i4.lowrow.output_size : index + %cols = config.get @ggml.mul_mat.symmetric_i4.lowrow.token_count : index + %k_b = index.assume %k [range(%k, 64, 32768)] : index + %rows_b = index.assume %rows [range(%rows, 1, 262144)] : index + %cols_b = index.assume %cols [range(%cols, 1, 32768)] : index + %nchunks = index.div %k_b, %k_kc : index + %kblocks = index.div %k_b, %c32 : index + %partition_chunks = index.div %nchunks, %c2 : index + + %col_tile0 = kernel.workgroup.id : index + %split_row_group0 = kernel.workgroup.id : index + %tid0 = kernel.workitem.id : index + %col_tile = index.assume %col_tile0 [range(%col_tile0, 0, 511)] : index + %split_row_group = index.assume %split_row_group0 [range(%split_row_group0, 0, 32767)] : index + %row_group = index.div %split_row_group, %c2 : index + %partition = index.rem %split_row_group, %c2 : index + %row_tile0 = index.mul %row_group, %c2 : index + %row_tile = index.assume %row_tile0 [range(%row_tile0, 0, 16383)] : index + %chunk_begin = index.mul %partition, %partition_chunks : index + %chunk_end = index.add %chunk_begin, %partition_chunks : index + %tid = index.assume %tid0 [range(%tid0, 0, 63)] : index + %wave = index.div %tid, %k_wave : index + + %weight_g = buffer.assume.memory_space %weight : buffer + %output_g = buffer.assume.memory_space %output : buffer + %aq_g = buffer.assume.memory_space %aq : buffer + %as_g = buffer.assume.memory_space %as : buffer + %asum_g = buffer.assume.memory_space %asum : buffer + %partial_g = buffer.assume.memory_space %partial : buffer + %completion_counters_g = buffer.assume.memory_space %completion_counters : buffer + %weight_na, %output_na, %aq_na, %as_na, %asum_na, %partial_na, %completion_counters_na = buffer.assume.noalias %weight_g, %output_g, %aq_g, %as_g, %asum_g, %partial_g, %completion_counters_g : buffer, buffer, buffer, buffer, buffer, buffer, buffer + template.apply<@ggml.mul_mat.symmetric_i4.lowrow.m16n16.body>(%scratch, %weight_na, %output_na, %aq_na, %as_na, %asum_na, %compact_q2, %split_k, %partial_na, %completion_counters_na, %base, %base, %c0, %c1, %c2, %c4, %c8, %c16, %c32, %c17, %c34, %k_kc, %k_astride, %k_wstride, %k_m, %k_n, %k_blocks, %k_wave, %k_wm, %fzero, %izero, %k_b, %rows_b, %cols_b, %nchunks, %kblocks, %col_tile, %row_tile, %tid, %wave, %partition, %chunk_begin, %chunk_end) : (buffer, buffer, buffer, buffer, buffer, buffer, i1, i1, buffer, buffer, offset, offset, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, vector<8xf32>, vector<8xi32>, index, index, index, index, index, index, index, index, index, index, index, index) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_symmetric_i8_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_symmetric_i8_wmma.loom new file mode 100644 index 000000000000..e755fa01056a --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/mul_mat_symmetric_i8_wmma.loom @@ -0,0 +1,762 @@ +config.decl @ggml.quantize_symmetric_i8_k256.token_count : %value: index where [range(%value, 256, 2048), mul(%value, 256)] + +config.decl @ggml.quantize_symmetric_i8_k256.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +func.def inline @ggml_symmetric_i8_wmmai8_stage_values(%token256: i1, %scratch: buffer, %base: offset, %adbase: index, %ad1: index, %values0: vector<32xi8>, %values1: vector<32xi8>) { + scf.if %token256 { + %al_flat = buffer.view %scratch[%base] : buffer -> view<20480xi8> + %ad0_bounded = index.assume %adbase [range(%adbase, 0, 20400)] : index + %ad1_bounded = index.assume %ad1 [range(%ad1, 32, 20432)] : index + vector.store %values0, %al_flat[%ad0_bounded] : vector<32xi8>, view<20480xi8> + vector.store %values1, %al_flat[%ad1_bounded] : vector<32xi8>, view<20480xi8> + } else { + %al_flat = buffer.view %scratch[%base] : buffer -> view<10240xi8> + %ad0_bounded = index.assume %adbase [range(%adbase, 0, 10160)] : index + %ad1_bounded = index.assume %ad1 [range(%ad1, 32, 10192)] : index + vector.store %values0, %al_flat[%ad0_bounded] : vector<32xi8>, view<10240xi8> + vector.store %values1, %al_flat[%ad1_bounded] : vector<32xi8>, view<10240xi8> + } + func.return +} + +func.def inline @ggml_symmetric_i8_wmmai8_stage_metadata(%token256: i1, %scratch: buffer, %as_off: offset, %asum_off: offset, %dst0: index, %dst1: index, %d0: f32, %d1: f32, %sum0: f32, %sum1: f32) { + scf.if %token256 { + %scale_view = buffer.view %scratch[%as_off] : buffer -> view<512xf32> + %sum_view = buffer.view %scratch[%asum_off] : buffer -> view<512xf32> + %dst0_bounded = index.assume %dst0 [range(%dst0, 0, 495)] : index + %dst1_bounded = index.assume %dst1 [range(%dst1, 16, 511)] : index + view.store %d0, %scale_view[%dst0_bounded] : f32, view<512xf32> + view.store %d1, %scale_view[%dst1_bounded] : f32, view<512xf32> + view.store %sum0, %sum_view[%dst0_bounded] : f32, view<512xf32> + view.store %sum1, %sum_view[%dst1_bounded] : f32, view<512xf32> + } else { + %scale_view = buffer.view %scratch[%as_off] : buffer -> view<256xf32> + %sum_view = buffer.view %scratch[%asum_off] : buffer -> view<256xf32> + %dst0_bounded = index.assume %dst0 [range(%dst0, 0, 239)] : index + %dst1_bounded = index.assume %dst1 [range(%dst1, 16, 255)] : index + view.store %d0, %scale_view[%dst0_bounded] : f32, view<256xf32> + view.store %d1, %scale_view[%dst1_bounded] : f32, view<256xf32> + view.store %sum0, %sum_view[%dst0_bounded] : f32, view<256xf32> + view.store %sum1, %sum_view[%dst1_bounded] : f32, view<256xf32> + } + func.return +} + +func.def inline @ggml_symmetric_i8_wmmai8_load_lhs(%token256: i1, %scratch: buffer, %base: offset, %row: index, %column: index, %c16: index) -> (vector<16xi8>) { + %lhs_layout = encoding.layout.strided [80, 1] : encoding + %result = scf.if %token256 -> (vector<16xi8>) { + %al_view = buffer.view %scratch[%base] : buffer -> view<256x64xi8, %lhs_layout> + %values = vector.fragment.load %al_view[%row, %column] shape [%c16, %c16] : view<256x64xi8, %lhs_layout> -> vector<16xi8> + scf.yield %values : vector<16xi8> + } else { + %al_view = buffer.view %scratch[%base] : buffer -> view<128x64xi8, %lhs_layout> + %values = vector.fragment.load %al_view[%row, %column] shape [%c16, %c16] : view<128x64xi8, %lhs_layout> -> vector<16xi8> + scf.yield %values : vector<16xi8> + } + func.return %result : vector<16xi8> +} + +func.def inline @ggml_symmetric_i8_wmmai8_load_metadata(%token256: i1, %scratch: buffer, %as_off: offset, %asum_off: offset, %index: index) -> (vector<8xf32>, vector<8xf32>) { + %scale, %sum = scf.if %token256 -> (vector<8xf32>, vector<8xf32>) { + %scale_view = buffer.view %scratch[%as_off] : buffer -> view<512xf32> + %sum_view = buffer.view %scratch[%asum_off] : buffer -> view<512xf32> + %bounded_index = index.assume %index [range(%index, 0, 504)] : index + %scale_values = vector.load %scale_view[%bounded_index] : view<512xf32> -> vector<8xf32> + %sum_values = vector.load %sum_view[%bounded_index] : view<512xf32> -> vector<8xf32> + scf.yield %scale_values, %sum_values : vector<8xf32>, vector<8xf32> + } else { + %scale_view = buffer.view %scratch[%as_off] : buffer -> view<256xf32> + %sum_view = buffer.view %scratch[%asum_off] : buffer -> view<256xf32> + %bounded_index = index.assume %index [range(%index, 0, 120)] : index + %scale_values = vector.load %scale_view[%bounded_index] : view<256xf32> -> vector<8xf32> + %sum_values = vector.load %sum_view[%bounded_index] : view<256xf32> -> vector<8xf32> + scf.yield %scale_values, %sum_values : vector<8xf32>, vector<8xf32> + } + func.return %scale, %sum : vector<8xf32>, vector<8xf32> +} + +amdgpu.target @ggml_quantize_symmetric_i8_k256_gfx11_wave32 {subgroup_size = 32} + +// Quantize one K256 activation block with two least-squares scale refinements. +kernel.def target(@ggml_quantize_symmetric_i8_k256_gfx11_wave32) export("ggml_quantize_f32_symmetric_i8_k256") @ggml_quantize_f32_symmetric_i8_k256(%token_count: index, %input_size_arg: index) { + %token_capacity = config.get @ggml.quantize_symmetric_i8_k256.token_count : index + %input_size = config.get @ggml.quantize_symmetric_i8_k256.input_size : index + %c256 = index.constant 256 : index + %c32 = index.constant 32 : index + %c1 = index.constant 1 : index + %block_count = index.div %input_size, %c256 : index + %group_count = index.mul %token_capacity, %block_count : index + kernel.launch.config workgroups(%group_count, %c1, %c1) workgroup_size(%c32, %c1, %c1) : index +} launch(%token_count: index, %input_size_arg: index, %input: buffer, %output: buffer) { + %token_capacity0 = config.get @ggml.quantize_symmetric_i8_k256.token_count : index + %input_size0 = config.get @ggml.quantize_symmetric_i8_k256.input_size : index + %token_capacity, %input_size = index.assume %token_capacity0, %input_size0 [range(%token_capacity0, 1, 2048), range(%input_size0, 256, 32768), mul(%input_size0, 256)] : index, index + %bounded_token_count, %bounded_input_size = index.assume %token_count, %input_size_arg [eq(%token_count, %token_capacity), eq(%input_size_arg, %input_size)] : index, index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c32 = index.constant 32 : index + %c256 = index.constant 256 : index + %c127 = scalar.constant 127.0 : f32 + %cn127 = scalar.constant -127.0 : f32 + %cn1 = scalar.constant -1.0 : f32 + %c1f = scalar.constant 1.0 : f32 + %c0f = scalar.constant 0.0 : f32 + %c0h = scalar.constant 0.0 : f16 + %v127 = vector.splat %c127 : vector<8xf32> + %vn127 = vector.splat %cn127 : vector<8xf32> + %vn1 = vector.splat %cn1 : vector<8xf32> + %base = index.constant 0 : offset + %lane0 = kernel.workitem.id : index + %lane = index.assume %lane0 [range(%lane0, 0, 31)] : index + %group0 = kernel.workgroup.id : index + %block_count = index.div %bounded_input_size, %c256 : index + %group_count = index.mul %bounded_token_count, %block_count : index + %group = index.assume %group0 [range(%group0, 0, 262143), lt(%group0, %group_count)] : index + %block = index.div %group, %bounded_token_count : index + %token = index.rem %group, %bounded_token_count : index + %token_base = index.mul %token, %bounded_input_size : index + %block_add = index.mul %block, %c256 : index + %lane_add = index.mul %lane, %c8 : index + %input_index0 = index.add %token_base, %block_add : index + %input_index = index.add %input_index0, %lane_add : index + %element_count = index.mul %bounded_token_count, %bounded_input_size : index + %input_noalias, %output_noalias = buffer.assume.noalias %input, %output : buffer, buffer + %input_view = buffer.view %input_noalias[%base] : buffer -> view<[%element_count]xf32> + %values = vector.load %input_view[%input_index] : view<[%element_count]xf32> -> vector<8xf32> + %absolute = vector.absf %values : vector<8xf32> + %lane_max = vector.reduce %absolute, %c0f : vector<8xf32>, f32 + %amax = kernel.subgroup.reduce %lane_max : f32 + %nonzero = scalar.cmpf one, %amax, %c0f : f32 + %scale0 = scf.if %nonzero -> (f32) { + %value = scalar.divf %amax, %c127 : f32 + scf.yield %value : f32 + } else { + scf.yield %c0f : f32 + } + %inv0 = scf.if %nonzero -> (f32) { + %value = scalar.divf %c1f, %scale0 : f32 + scf.yield %value : f32 + } else { + scf.yield %c0f : f32 + } + %inv0v = vector.splat %inv0 : vector<8xf32> + %scaled0 = vector.mulf %values, %inv0v : vector<8xf32> + %rounded0 = vector.roundf %scaled0 : vector<8xf32> + %lower0 = vector.maxnumf %rounded0, %vn127 : vector<8xf32> + %neg0 = vector.mulf %lower0, %vn1 : vector<8xf32> + %negclamp0 = vector.maxnumf %neg0, %vn127 : vector<8xf32> + %clamped0 = vector.mulf %negclamp0, %vn1 : vector<8xf32> + %numerator_terms0 = vector.mulf %values, %clamped0 : vector<8xf32> + %denominator_terms0 = vector.mulf %clamped0, %clamped0 : vector<8xf32> + %lane_numerator0 = vector.reduce %numerator_terms0, %c0f : vector<8xf32>, f32 + %lane_denominator0 = vector.reduce %denominator_terms0, %c0f : vector<8xf32>, f32 + %numerator0 = kernel.subgroup.reduce %lane_numerator0 : f32 + %denominator0 = kernel.subgroup.reduce %lane_denominator0 : f32 + %has_denominator0 = scalar.cmpf one, %denominator0, %c0f : f32 + %scale1 = scf.if %has_denominator0 -> (f32) { + %value = scalar.divf %numerator0, %denominator0 : f32 + scf.yield %value : f32 + } else { + scf.yield %scale0 : f32 + } + %inv1 = scf.if %nonzero -> (f32) { + %value = scalar.divf %c1f, %scale1 : f32 + scf.yield %value : f32 + } else { + scf.yield %c0f : f32 + } + %inv1v = vector.splat %inv1 : vector<8xf32> + %scaled1 = vector.mulf %values, %inv1v : vector<8xf32> + %rounded1 = vector.roundf %scaled1 : vector<8xf32> + %lower1 = vector.maxnumf %rounded1, %vn127 : vector<8xf32> + %neg1 = vector.mulf %lower1, %vn1 : vector<8xf32> + %negclamp1 = vector.maxnumf %neg1, %vn127 : vector<8xf32> + %clamped1 = vector.mulf %negclamp1, %vn1 : vector<8xf32> + %numerator_terms1 = vector.mulf %values, %clamped1 : vector<8xf32> + %denominator_terms1 = vector.mulf %clamped1, %clamped1 : vector<8xf32> + %lane_numerator1 = vector.reduce %numerator_terms1, %c0f : vector<8xf32>, f32 + %lane_denominator1 = vector.reduce %denominator_terms1, %c0f : vector<8xf32>, f32 + %numerator1 = kernel.subgroup.reduce %lane_numerator1 : f32 + %denominator1 = kernel.subgroup.reduce %lane_denominator1 : f32 + %has_denominator1 = scalar.cmpf one, %denominator1, %c0f : f32 + %scale2 = scf.if %has_denominator1 -> (f32) { + %value = scalar.divf %numerator1, %denominator1 : f32 + scf.yield %value : f32 + } else { + scf.yield %scale1 : f32 + } + %encoded_scale = scalar.fptrunc %scale2 : f32 to f16 + %decoded_scale = scalar.extf %encoded_scale : f16 to f32 + %inverse = scf.if %nonzero -> (f32) { + %value = scalar.divf %c1f, %decoded_scale : f32 + scf.yield %value : f32 + } else { + scf.yield %c0f : f32 + } + %inverse_v = vector.splat %inverse : vector<8xf32> + %scaled = vector.mulf %values, %inverse_v : vector<8xf32> + %rounded = vector.roundf %scaled : vector<8xf32> + %lower = vector.maxnumf %rounded, %vn127 : vector<8xf32> + %neg = vector.mulf %lower, %vn1 : vector<8xf32> + %negclamp = vector.maxnumf %neg, %vn127 : vector<8xf32> + %clamped = vector.mulf %negclamp, %vn1 : vector<8xf32> + %quantized = vector.fptosi %clamped : vector<8xf32> to vector<8xi8> + %chunk = index.div %lane, %c4 : index + %word = index.rem %lane, %c4 : index + %block_chunk0 = index.mul %block, %c8 : index + %block_chunk = index.add %block_chunk0, %chunk : index + %payload_row0 = index.mul %block_chunk, %bounded_token_count : index + %payload_row = index.add %payload_row0, %token : index + %payload_base = index.mul %payload_row, %c32 : index + %word_add = index.mul %word, %c8 : index + %payload_index = index.add %payload_base, %word_add : index + %output_payload = buffer.view %output_noalias[%base] : buffer -> view<[%element_count]xi8> + vector.store %quantized, %output_payload[%payload_index] : vector<8xi8>, view<[%element_count]xi8> + %metadata_offset = index.cast %element_count : index to offset + %metadata = buffer.view %output_noalias[%metadata_offset] : buffer -> view<[%group_count]xi32> + %is_leader = index.cmp eq, %lane, %c0 : index + scf.if %is_leader { + %metadata_pair = vector.from_elements %encoded_scale, %c0h : vector<2xf16> + %metadata_word = vector.bitcast %metadata_pair : vector<2xf16> to vector<1xi32> + vector.store %metadata_word, %metadata[%group] : vector<1xi32>, view<[%group_count]xi32> + } + kernel.return +} + +kernel.def target(@ggml_symmetric_i8_gfx11_wave32) export("ggml_mul_mat_q5_k_symmetric_i8_wmma") @ggml_mul_mat_q5_k_symmetric_i8_wmma(%token_count: index) { + %unit = index.constant 1 : index + %tile_rows = index.constant 64 : index + %tile_tokens = index.constant 256 : index + %workgroup_size = index.constant 256 : index + %output_size = config.get @ggml.mul_mat.symmetric_i8.output_size : index + %token_capacity = config.get @ggml.mul_mat.symmetric_i8.token_count : index + %row_tiles = index.div %output_size, %tile_rows : index + %token_tiles = index.div %token_capacity, %tile_tokens : index + kernel.launch.config workgroups(%token_tiles, %row_tiles, %unit) workgroup_size(%workgroup_size, %unit, %unit) : index +} launch(%token_count: index, %q8_input: buffer, %weight: buffer, %output: buffer) { + %base = index.constant 0 : offset + %lds_bytes = index.constant 30720 : offset + %w_off = index.constant 20480 : offset + %as_off = index.constant 25600 : offset + %ws_off = index.constant 27648 : offset + %asum_off = index.constant 28160 : offset + %wc_off = index.constant 30208 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c17 = index.constant 17 : index + %c34 = index.constant 34 : index + %k_kc = index.constant 64 : index + %k_astride = index.constant 80 : index + %k_wstride = index.constant 80 : index + %k_m = index.constant 256 : index + %k_n = index.constant 64 : index + %k_blocks = index.constant 2 : index + %k_wave = index.constant 32 : index + %k_wm = index.constant 8 : index + %fzero = vector.constant 0.0 : vector<8xf32> + %izero = vector.constant 0 : vector<8xi32> + %is_iq4xs = scalar.constant false : i1 + %token256 = scalar.constant true : i1 + %q8_plane = scalar.constant true : i1 + %apply_swiglu = scalar.constant false : i1 + + %input_size0 = config.get @ggml.mul_mat.symmetric_i8.input_size : index + %output_size0 = config.get @ggml.mul_mat.symmetric_i8.output_size : index + %token_capacity = config.get @ggml.mul_mat.symmetric_i8.token_count : index + %input_size = index.assume %input_size0 [range(%input_size0, 256, 32768), mul(%input_size0, 256)] : index + %output_size = index.assume %output_size0 [range(%output_size0, 64, 262144), mul(%output_size0, 64)] : index + %bounded_token_count, %configured_token_capacity = index.assume %token_count, %token_capacity [range(%token_count, 256, 2048), mul(%token_count, 256), eq(%token_count, %token_capacity)] : index, index + %nchunks = index.div %input_size, %k_kc : index + %kblocks = index.div %input_size, %c32 : index + + %col_tile0 = kernel.workgroup.id : index + %row_tile0 = kernel.workgroup.id : index + %tid0 = kernel.workitem.id : index + %col_tile = index.assume %col_tile0 [range(%col_tile0, 0, 7)] : index + %row_tile = index.assume %row_tile0 [range(%row_tile0, 0, 4095)] : index + %tid = index.assume %tid0 [range(%tid0, 0, 255)] : index + %wave = index.div %tid, %k_wave : index + + %weight_g = buffer.assume.memory_space %weight : buffer + %output_g = buffer.assume.memory_space %output : buffer + %q8_g = buffer.assume.memory_space %q8_input : buffer + %weight_noalias, %output_noalias, %q8_noalias = buffer.assume.noalias %weight_g, %output_g, %q8_g : buffer, buffer, buffer + + template.apply<@ggml.mul_mat.symmetric_i8.wmmai8.body>(%is_iq4xs, %token256, %q8_plane, %apply_swiglu, %output, %weight_noalias, %output_noalias, %q8_noalias, %base, %lds_bytes, %w_off, %as_off, %ws_off, %asum_off, %wc_off, %c0, %c1, %c2, %c4, %c8, %c16, %c32, %c17, %c34, %k_kc, %k_astride, %k_wstride, %k_m, %k_n, %k_blocks, %k_wave, %k_wm, %fzero, %izero, %input_size, %output_size, %bounded_token_count, %nchunks, %kblocks, %col_tile, %row_tile, %tid, %wave) : (i1, i1, i1, i1, buffer, buffer, buffer, buffer, offset, offset, offset, offset, offset, offset, offset, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, index, vector<8xf32>, vector<8xi32>, index, index, index, index, index, index, index, index, index) + kernel.return +} + +template.decl @ggml.mul_mat.symmetric_i8.wmmai8.body(%is_iq4xs: i1, %token256: i1, %q8_plane: i1, %apply_swiglu: i1, %gate: buffer, %src0_na: buffer, %dst_na: buffer, %q8_na: buffer, %base: offset, %lds_bytes: offset, %w_off: offset, %as_off: offset, %ws_off: offset, %asum_off: offset, %wc_off: offset, %c0: index, %c1: index, %c2: index, %c4: index, %c8: index, %c16: index, %c32: index, %c17: index, %c34: index, %k_kc: index, %k_astride: index, %k_wstride: index, %k_m: index, %k_n: index, %k_blocks: index, %k_wave: index, %k_wm: index, %fzero: vector<8xf32>, %izero: vector<8xi32>, %k_b: index, %rows_b: index, %cols_b: index, %nchunks: index, %kblocks: index, %col_tile: index, %row_tile: index, %tid: index, %wave: index) + +amdgpu.target @ggml_symmetric_i8_gfx11_wave32 {subgroup_size = 32} + +config.def @ggml.mul_mat.symmetric_i8.input_size = 5120 : index + +config.def @ggml.mul_mat.symmetric_i8.output_size = 17408 : index + +config.def @ggml.mul_mat.symmetric_i8.token_count = 512 : index + +config.def @ggml.mul_mat.symmetric_i8.activation_cadence = 256 : index + +config.def @ggml.mul_mat.symmetric_i8.chunks_per_activation = 4 : index + +config.def @ggml.mul_mat.symmetric_i8.activation_blocks_per_weight = 1 : index + +// Raw Q5_K x packed-Q8_1_x4 IU8 WMMA using the existing Loom activation layout. +template.def<@ggml.mul_mat.symmetric_i8.wmmai8.body> device @ggml_mul_mat_symmetric_i8_wmmai8_body(%is_iq4xs: i1, %token256: i1, %q8_plane: i1, %apply_swiglu: i1, %gate: buffer, %src0_na: buffer, %dst_na: buffer, %q8_na: buffer, %base: offset, %lds_bytes: offset, %w_off: offset, %as_off: offset, %ws_off: offset, %asum_off: offset, %wc_off: offset, %c0: index, %c1: index, %c2: index, %c4: index, %c8: index, %c16: index, %c32: index, %c17: index, %c34: index, %k_kc: index, %k_astride: index, %k_wstride: index, %k_m: index, %k_n: index, %k_blocks: index, %k_wave: index, %k_wm: index, %fzero: vector<8xf32>, %izero: vector<8xi32>, %k_b: index, %rows_b: index, %cols_b: index, %nchunks: index, %kblocks: index, %col_tile: index, %row_tile: index, %tid: index, %wave: index) { + %q4_k_block = index.constant 256 : index + %q5_block_bytes = index.constant 176 : offset + %q5_high_offset = index.constant 16 : offset + %q5_code_offset = index.constant 48 : offset + %q8_group_bytes = index.constant 144 : offset + %q8_payload_offset = index.constant 16 : offset + %q8_elements_per_group = index.constant 128 : index + // dst is [rows, cols] with rows contiguous, and C is [cols, rows]. + %dst_view = buffer.view %dst_na[%base] : buffer -> view<[%cols_b]x[%rows_b]xf32> + %gate_global = buffer.assume.memory_space %gate : buffer + %gate_view = buffer.view %gate_global[%base] : buffer -> view<[%cols_b]x[%rows_b]xf32> + + // Four or eight waves share one 64-row weight tile across the token tile. + %scratch = buffer.alloca align(16) %lds_bytes : buffer + %i8_schema = encoding.define #encoding.operand : encoding + // A compile-time format selects unsigned Q5_K or signed IQ4_XS operands before target lowering. + %signed_rhs_schema = encoding.define #encoding.operand : encoding + %unsigned_rhs_schema = encoding.define #encoding.operand : encoding + %u8_schema = scf.select %is_iq4xs, %signed_rhs_schema, %unsigned_rhs_schema : encoding + %rhs_layout = encoding.layout.strided [1, 80] : encoding + %wl_view = buffer.view %scratch[%w_off] : buffer -> view<64x64xi8, %rhs_layout> + %wl_flat = buffer.view %scratch[%w_off] : buffer -> view<5120xi8> + %wsl_view = buffer.view %scratch[%ws_off] : buffer -> view<128xf32> + %wcl_view = buffer.view %scratch[%wc_off] : buffer -> view<128xf32> + + %col_base = index.mul %col_tile, %k_m : index + %row_base = index.mul %row_tile, %k_n : index + + // One thread stages both K32 blocks for a token from the 144-byte Q8_1_x4 group. + %ascol_g = index.add %col_base, %tid : index + %adbase = index.mul %tid, %k_astride : index + %q8_groups_per_row = index.div %k_b, %q8_elements_per_group : index + %q8_row_bytes = index.scale %q8_groups_per_row, %q8_group_bytes : index, offset -> offset + %q8_row_byte_base = index.scale %ascol_g, %q8_row_bytes : index, offset -> offset + %q8_plane_payload_elements = index.mul %k_b, %cols_b : index + %q8_plane_payload_words = index.div %q8_plane_payload_elements, %c4 : index + %q8_plane_metadata_offset = index.cast %q8_plane_payload_elements : index to offset + %q8_plane_metadata_count = index.mul %kblocks, %cols_b : index + %q8_plane_payload = buffer.view %q8_na[%base] : buffer -> view<[%q8_plane_payload_words]xi32> + %q8_plane_metadata = buffer.view %q8_na[%q8_plane_metadata_offset] : buffer -> view<[%q8_plane_metadata_count]xi32> + + // The waves own 32-token blocks and share one 64-row weight tile. + %wcol = index.rem %wave, %k_wm : index + %wrow = index.div %wave, %k_wm : index + %k_mspan = index.constant 32 : index + %k_nspan = index.constant 64 : index + %wm_off = index.mul %wcol, %k_mspan : index + %wn_off = index.mul %wrow, %k_nspan : index + %m_out = index.add %col_base, %wm_off : index + %n_out = index.add %row_base, %wn_off : index + %lm0 = index.add %wm_off, %c0 : index + %gm0 = index.add %m_out, %c0 : index + %k_ma1 = index.constant 16 : index + %lm1 = index.add %wm_off, %k_ma1 : index + %gm1 = index.add %m_out, %k_ma1 : index + %ln0 = index.add %wn_off, %c0 : index + %gn0 = index.add %n_out, %c0 : index + %k_nb1 = index.constant 16 : index + %ln1 = index.add %wn_off, %k_nb1 : index + %gn1 = index.add %n_out, %k_nb1 : index + %k_nb2 = index.constant 32 : index + %ln2 = index.add %wn_off, %k_nb2 : index + %gn2 = index.add %n_out, %k_nb2 : index + %k_nb3 = index.constant 48 : index + %ln3 = index.add %wn_off, %k_nb3 : index + %gn3 = index.add %n_out, %k_nb3 : index + %lane = index.rem %tid, %k_wave : index + %lane_lo = index.rem %lane, %c16 : index + %lane_hi = index.div %lane, %c16 : index + %wsb0_0 = index.add %ln0, %lane_lo : index + %wsb0 = index.mul %wsb0_0, %k_blocks : index + %wsb1_0 = index.add %ln1, %lane_lo : index + %wsb1 = index.mul %wsb1_0, %k_blocks : index + %wsb2_0 = index.add %ln2, %lane_lo : index + %wsb2 = index.mul %wsb2_0, %k_blocks : index + %wsb3_0 = index.add %ln3, %lane_lo : index + %wsb3 = index.mul %wsb3_0, %k_blocks : index + // Activation scale metadata is [tile][block][lane-half][register]. + %as_tile0 = index.div %lm0, %c16 : index + %asb0 = index.mul %as_tile0, %k_blocks : index + %as_tile1 = index.div %lm1, %c16 : index + %asb1 = index.mul %as_tile1, %k_blocks : index + + %i8_c64 = index.constant 64 : index + %i8_c258 = index.constant 258 : index + %i8_c256 = index.constant 256 : index + %i8_act_cadence = config.get @ggml.mul_mat.symmetric_i8.activation_cadence : index + %i8_chunks_per_act = config.get @ggml.mul_mat.symmetric_i8.chunks_per_activation : index + %i8_act_blocks_per_weight = config.get @ggml.mul_mat.symmetric_i8.activation_blocks_per_weight : index + %i8_c48_outer = index.constant 48 : index + %i8_record_bytes = index.constant 16512 : index + %i8_scale_bytes = index.constant 128 : index + %i8_field_bytes = index.constant 4096 : index + %i8_payload_elements = index.mul %k_b, %cols_b : index + %i8_payload_words = index.div %i8_payload_elements, %c4 : index + %i8_meta_offset = index.cast %i8_payload_elements : index to offset + %i8_act_block_count = index.div %k_b, %i8_act_cadence : index + %i8_weight_block_count = index.div %k_b, %i8_c256 : index + %i8_meta_count = index.mul %i8_act_block_count, %cols_b : index + %i8_q8_payload = buffer.view %q8_na[%base] : buffer -> view<[%i8_payload_words]xi32> + %i8_q8_meta = buffer.view %q8_na[%i8_meta_offset] : buffer -> view<[%i8_meta_count]xi32> + %i8_weight_records = index.mul %rows_b, %i8_weight_block_count : index + %i8_weight_bytes = index.mul %i8_weight_records, %i8_c258 : index + %i8_weight_view = buffer.view %src0_na[%base] : buffer -> view<[%i8_weight_bytes]xi8> + %f0_0, %f0_1, %f0_2, %f0_3, %f1_0, %f1_1, %f1_2, %f1_3 = scf.for %i8_act_block = [%c0 to %i8_act_block_count step %c1](%fc0_0 = %fzero : vector<8xf32>, %fc0_1 = %fzero : vector<8xf32>, %fc0_2 = %fzero : vector<8xf32>, %fc0_3 = %fzero : vector<8xf32>, %fc1_0 = %fzero : vector<8xf32>, %fc1_1 = %fzero : vector<8xf32>, %fc1_2 = %fzero : vector<8xf32>, %fc1_3 = %fzero : vector<8xf32>) -> (vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>) { + %i8_source_row_base = index.constant 0 : index + %i8_weight_block = index.div %i8_act_block, %i8_act_blocks_per_weight : index + %i8_act_in_weight = index.rem %i8_act_block, %i8_act_blocks_per_weight : index + %i8_weight_inner_base = index.mul %i8_act_in_weight, %i8_chunks_per_act : index + %i8_tile_block0 = index.mul %row_tile, %i8_weight_block_count : index + %i8_tile_block = index.add %i8_tile_block0, %i8_weight_block : index + %i8_record_base = index.mul %i8_tile_block, %i8_record_bytes : index + %i8_i0_0, %i8_i0_1, %i8_i0_2, %i8_i0_3, %i8_i1_0, %i8_i1_1, %i8_i1_2, %i8_i1_3 = scf.for %inner = [%c0 to %i8_chunks_per_act step %c1](%i8_ic0_0 = %izero : vector<8xi32>, %i8_ic0_1 = %izero : vector<8xi32>, %i8_ic0_2 = %izero : vector<8xi32>, %i8_ic0_3 = %izero : vector<8xi32>, %i8_ic1_0 = %izero : vector<8xi32>, %i8_ic1_1 = %izero : vector<8xi32>, %i8_ic1_2 = %izero : vector<8xi32>, %i8_ic1_3 = %izero : vector<8xi32>) -> (vector<8xi32>, vector<8xi32>, vector<8xi32>, vector<8xi32>, vector<8xi32>, vector<8xi32>, vector<8xi32>, vector<8xi32>) { + %i8_chunk0_loop = index.mul %i8_act_block, %i8_chunks_per_act : index + %i8_chunk_loop = index.add %i8_chunk0_loop, %inner : index + %i8_q8_block0_loop = index.mul %i8_chunk_loop, %c2 : index + %i8_q8_block1_loop = index.add %i8_q8_block0_loop, %c1 : index + %i8_q8_token_base0_loop = index.mul %i8_q8_block0_loop, %cols_b : index + %i8_q8_token_base1_loop = index.mul %i8_q8_block1_loop, %cols_b : index + %i8_q8_token0_loop = index.add %i8_q8_token_base0_loop, %ascol_g : index + %i8_q8_token1_loop = index.add %i8_q8_token_base1_loop, %ascol_g : index + %i8_q8_word00_loop = index.mul %i8_q8_token0_loop, %c8 : index + %i8_q8_word10_loop = index.mul %i8_q8_token1_loop, %c8 : index + %i8_q8_word0_loop = index.assume %i8_q8_word00_loop [range(%i8_q8_word00_loop, 0, 268435448), mul(%i8_q8_word00_loop, 8)] : index + %i8_q8_word1_loop = index.assume %i8_q8_word10_loop [range(%i8_q8_word10_loop, 8, 268435448), mul(%i8_q8_word10_loop, 8)] : index + %i8_q8_words0_loop = vector.load %i8_q8_payload[%i8_q8_word0_loop] : view<[%i8_payload_words]xi32> -> vector<8xi32> + %i8_q8_words1_loop = vector.load %i8_q8_payload[%i8_q8_word1_loop] : view<[%i8_payload_words]xi32> -> vector<8xi32> + %i8_q8_values0_loop = vector.bitcast %i8_q8_words0_loop : vector<8xi32> to vector<32xi8> + %i8_q8_values1_loop = vector.bitcast %i8_q8_words1_loop : vector<8xi32> to vector<32xi8> + %i8_ad1_loop = index.add %adbase, %c32 : index + func.call @ggml_symmetric_i8_wmmai8_stage_values(%token256, %scratch, %base, %adbase, %i8_ad1_loop, %i8_q8_values0_loop, %i8_q8_values1_loop) : (i1, buffer, offset, index, index, vector<32xi8>, vector<32xi8>) + %i8_meta_block_base_loop = index.mul %i8_act_block, %cols_b : index + %i8_meta_index_loop = index.add %i8_meta_block_base_loop, %ascol_g : index + %i8_meta_word_loop = view.load %i8_q8_meta[%i8_meta_index_loop] : view<[%i8_meta_count]xi32> -> i32 + %i8_meta_vector_loop = vector.from_elements %i8_meta_word_loop : vector<1xi32> + %i8_meta_halves_loop = vector.bitcast %i8_meta_vector_loop : vector<1xi32> to vector<2xf16> + %i8_qscale_f16_loop = vector.extract %i8_meta_halves_loop[0] : vector<2xf16> -> f16 + %i8_qscale_loop = scalar.extf %i8_qscale_f16_loop : f16 to f32 + %i8_qtoken_tile_loop = index.div %tid, %c16 : index + %i8_qtoken_inner_loop = index.rem %tid, %c16 : index + %i8_qtoken_parity_loop = index.rem %i8_qtoken_inner_loop, %c2 : index + %i8_qtoken_vector_loop = index.div %i8_qtoken_inner_loop, %c2 : index + %i8_qscale_tile0_loop = index.mul %i8_qtoken_tile_loop, %c4 : index + %i8_qscale_tile1_loop = index.add %i8_qscale_tile0_loop, %i8_qtoken_parity_loop : index + %i8_qscale_tile2_loop = index.mul %i8_qscale_tile1_loop, %c8 : index + %i8_qscale_dst0_loop = index.add %i8_qscale_tile2_loop, %i8_qtoken_vector_loop : index + %i8_qscale_dst1_loop = index.add %i8_qscale_dst0_loop, %c16 : index + func.call @ggml_symmetric_i8_wmmai8_stage_metadata(%token256, %scratch, %as_off, %asum_off, %i8_qscale_dst0_loop, %i8_qscale_dst1_loop, %i8_qscale_loop, %i8_qscale_loop, %i8_qscale_loop, %i8_qscale_loop) : (i1, buffer, offset, offset, index, index, f32, f32, f32, f32) + %i8_word0 = index.rem %tid, %c4 : index + %i8_word = index.assume %i8_word0 [range(%i8_word0, 0, 3)] : index + %i8_byte_position0 = index.mul %i8_word, %c16 : index + %i8_byte_position = index.assume %i8_byte_position0 [range(%i8_byte_position0, 0, 48), mul(%i8_byte_position0, 16)] : index + %i8_local_row0 = index.div %tid, %c4 : index + %i8_local_row = index.assume %i8_local_row0 [range(%i8_local_row0, 0, 63)] : index + %i8_weight_inner = index.add %i8_weight_inner_base, %inner : index + %i8_inner_bytes = index.mul %i8_weight_inner, %i8_field_bytes : index + %i8_row_byte0 = index.mul %i8_local_row, %i8_c64 : index + %i8_row_byte = index.add %i8_row_byte0, %i8_byte_position : index + %i8_payload_base = index.add %i8_record_base, %i8_scale_bytes : index + %i8_chunk_base = index.add %i8_payload_base, %i8_inner_bytes : index + %i8_address = index.add %i8_chunk_base, %i8_row_byte : index + %i8_values = vector.load %i8_weight_view[%i8_address] : view<[%i8_weight_bytes]xi8> -> vector<16xi8> + %i8_weight_row_base0 = index.mul %i8_local_row, %k_wstride : index + %i8_weight_row_base = index.add %i8_weight_row_base0, %i8_byte_position : index + %i8_weight_destination = index.assume %i8_weight_row_base [range(%i8_weight_row_base, 0, 5104), mul(%i8_weight_row_base, 16)] : index + vector.store %i8_values, %wl_flat[%i8_weight_destination] : vector<16xi8>, view<5120xi8> + kernel.barrier scope(workgroup) ordering(acq_rel) + %i8_lhs0_0_raw = func.call @ggml_symmetric_i8_wmmai8_load_lhs(%token256, %scratch, %base, %lm0, %c0, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %i8_lhs0_0_words = vector.bitcast %i8_lhs0_0_raw : vector<16xi8> to vector<4xi32> + %i8_lhs0_0 = vector.fragment %i8_lhs0_0_words shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %i8_lhs0_1_raw = func.call @ggml_symmetric_i8_wmmai8_load_lhs(%token256, %scratch, %base, %lm0, %c16, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %i8_lhs0_1_words = vector.bitcast %i8_lhs0_1_raw : vector<16xi8> to vector<4xi32> + %i8_lhs0_1 = vector.fragment %i8_lhs0_1_words shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %i8_lhs0_2_raw = func.call @ggml_symmetric_i8_wmmai8_load_lhs(%token256, %scratch, %base, %lm0, %c32, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %i8_lhs0_2_words = vector.bitcast %i8_lhs0_2_raw : vector<16xi8> to vector<4xi32> + %i8_lhs0_2 = vector.fragment %i8_lhs0_2_words shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %i8_lhs0_3_raw = func.call @ggml_symmetric_i8_wmmai8_load_lhs(%token256, %scratch, %base, %lm0, %i8_c48_outer, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %i8_lhs0_3_words = vector.bitcast %i8_lhs0_3_raw : vector<16xi8> to vector<4xi32> + %i8_lhs0_3 = vector.fragment %i8_lhs0_3_words shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %i8_lhs1_0_raw = func.call @ggml_symmetric_i8_wmmai8_load_lhs(%token256, %scratch, %base, %lm1, %c0, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %i8_lhs1_0_words = vector.bitcast %i8_lhs1_0_raw : vector<16xi8> to vector<4xi32> + %i8_lhs1_0 = vector.fragment %i8_lhs1_0_words shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %i8_lhs1_1_raw = func.call @ggml_symmetric_i8_wmmai8_load_lhs(%token256, %scratch, %base, %lm1, %c16, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %i8_lhs1_1_words = vector.bitcast %i8_lhs1_1_raw : vector<16xi8> to vector<4xi32> + %i8_lhs1_1 = vector.fragment %i8_lhs1_1_words shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %i8_lhs1_2_raw = func.call @ggml_symmetric_i8_wmmai8_load_lhs(%token256, %scratch, %base, %lm1, %c32, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %i8_lhs1_2_words = vector.bitcast %i8_lhs1_2_raw : vector<16xi8> to vector<4xi32> + %i8_lhs1_2 = vector.fragment %i8_lhs1_2_words shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %i8_lhs1_3_raw = func.call @ggml_symmetric_i8_wmmai8_load_lhs(%token256, %scratch, %base, %lm1, %i8_c48_outer, %c16) : (i1, buffer, offset, index, index, index) -> (vector<16xi8>) + %i8_lhs1_3_words = vector.bitcast %i8_lhs1_3_raw : vector<16xi8> to vector<4xi32> + %i8_lhs1_3 = vector.fragment %i8_lhs1_3_words shape [%c16, %c16] using {schema = %i8_schema : encoding} : vector<4xi32> + %i8_rhs0_0_raw = vector.fragment.load %wl_view[%c0, %ln0] shape [%c16, %c16] : view<64x64xi8, %rhs_layout> -> vector<16xi8> + %i8_rhs0_0_words = vector.bitcast %i8_rhs0_0_raw : vector<16xi8> to vector<4xi32> + %i8_rhs0_0 = vector.fragment %i8_rhs0_0_words shape [%c16, %c16] using {schema = %signed_rhs_schema : encoding} : vector<4xi32> + %i8_rhs0_1_raw = vector.fragment.load %wl_view[%c16, %ln0] shape [%c16, %c16] : view<64x64xi8, %rhs_layout> -> vector<16xi8> + %i8_rhs0_1_words = vector.bitcast %i8_rhs0_1_raw : vector<16xi8> to vector<4xi32> + %i8_rhs0_1 = vector.fragment %i8_rhs0_1_words shape [%c16, %c16] using {schema = %signed_rhs_schema : encoding} : vector<4xi32> + %i8_rhs0_2_raw = vector.fragment.load %wl_view[%c32, %ln0] shape [%c16, %c16] : view<64x64xi8, %rhs_layout> -> vector<16xi8> + %i8_rhs0_2_words = vector.bitcast %i8_rhs0_2_raw : vector<16xi8> to vector<4xi32> + %i8_rhs0_2 = vector.fragment %i8_rhs0_2_words shape [%c16, %c16] using {schema = %signed_rhs_schema : encoding} : vector<4xi32> + %i8_rhs0_3_raw = vector.fragment.load %wl_view[%i8_c48_outer, %ln0] shape [%c16, %c16] : view<64x64xi8, %rhs_layout> -> vector<16xi8> + %i8_rhs0_3_words = vector.bitcast %i8_rhs0_3_raw : vector<16xi8> to vector<4xi32> + %i8_rhs0_3 = vector.fragment %i8_rhs0_3_words shape [%c16, %c16] using {schema = %signed_rhs_schema : encoding} : vector<4xi32> + %i8_rhs1_0_raw = vector.fragment.load %wl_view[%c0, %ln1] shape [%c16, %c16] : view<64x64xi8, %rhs_layout> -> vector<16xi8> + %i8_rhs1_0_words = vector.bitcast %i8_rhs1_0_raw : vector<16xi8> to vector<4xi32> + %i8_rhs1_0 = vector.fragment %i8_rhs1_0_words shape [%c16, %c16] using {schema = %signed_rhs_schema : encoding} : vector<4xi32> + %i8_rhs1_1_raw = vector.fragment.load %wl_view[%c16, %ln1] shape [%c16, %c16] : view<64x64xi8, %rhs_layout> -> vector<16xi8> + %i8_rhs1_1_words = vector.bitcast %i8_rhs1_1_raw : vector<16xi8> to vector<4xi32> + %i8_rhs1_1 = vector.fragment %i8_rhs1_1_words shape [%c16, %c16] using {schema = %signed_rhs_schema : encoding} : vector<4xi32> + %i8_rhs1_2_raw = vector.fragment.load %wl_view[%c32, %ln1] shape [%c16, %c16] : view<64x64xi8, %rhs_layout> -> vector<16xi8> + %i8_rhs1_2_words = vector.bitcast %i8_rhs1_2_raw : vector<16xi8> to vector<4xi32> + %i8_rhs1_2 = vector.fragment %i8_rhs1_2_words shape [%c16, %c16] using {schema = %signed_rhs_schema : encoding} : vector<4xi32> + %i8_rhs1_3_raw = vector.fragment.load %wl_view[%i8_c48_outer, %ln1] shape [%c16, %c16] : view<64x64xi8, %rhs_layout> -> vector<16xi8> + %i8_rhs1_3_words = vector.bitcast %i8_rhs1_3_raw : vector<16xi8> to vector<4xi32> + %i8_rhs1_3 = vector.fragment %i8_rhs1_3_words shape [%c16, %c16] using {schema = %signed_rhs_schema : encoding} : vector<4xi32> + %i8_rhs2_0_raw = vector.fragment.load %wl_view[%c0, %ln2] shape [%c16, %c16] : view<64x64xi8, %rhs_layout> -> vector<16xi8> + %i8_rhs2_0_words = vector.bitcast %i8_rhs2_0_raw : vector<16xi8> to vector<4xi32> + %i8_rhs2_0 = vector.fragment %i8_rhs2_0_words shape [%c16, %c16] using {schema = %signed_rhs_schema : encoding} : vector<4xi32> + %i8_rhs2_1_raw = vector.fragment.load %wl_view[%c16, %ln2] shape [%c16, %c16] : view<64x64xi8, %rhs_layout> -> vector<16xi8> + %i8_rhs2_1_words = vector.bitcast %i8_rhs2_1_raw : vector<16xi8> to vector<4xi32> + %i8_rhs2_1 = vector.fragment %i8_rhs2_1_words shape [%c16, %c16] using {schema = %signed_rhs_schema : encoding} : vector<4xi32> + %i8_rhs2_2_raw = vector.fragment.load %wl_view[%c32, %ln2] shape [%c16, %c16] : view<64x64xi8, %rhs_layout> -> vector<16xi8> + %i8_rhs2_2_words = vector.bitcast %i8_rhs2_2_raw : vector<16xi8> to vector<4xi32> + %i8_rhs2_2 = vector.fragment %i8_rhs2_2_words shape [%c16, %c16] using {schema = %signed_rhs_schema : encoding} : vector<4xi32> + %i8_rhs2_3_raw = vector.fragment.load %wl_view[%i8_c48_outer, %ln2] shape [%c16, %c16] : view<64x64xi8, %rhs_layout> -> vector<16xi8> + %i8_rhs2_3_words = vector.bitcast %i8_rhs2_3_raw : vector<16xi8> to vector<4xi32> + %i8_rhs2_3 = vector.fragment %i8_rhs2_3_words shape [%c16, %c16] using {schema = %signed_rhs_schema : encoding} : vector<4xi32> + %i8_rhs3_0_raw = vector.fragment.load %wl_view[%c0, %ln3] shape [%c16, %c16] : view<64x64xi8, %rhs_layout> -> vector<16xi8> + %i8_rhs3_0_words = vector.bitcast %i8_rhs3_0_raw : vector<16xi8> to vector<4xi32> + %i8_rhs3_0 = vector.fragment %i8_rhs3_0_words shape [%c16, %c16] using {schema = %signed_rhs_schema : encoding} : vector<4xi32> + %i8_rhs3_1_raw = vector.fragment.load %wl_view[%c16, %ln3] shape [%c16, %c16] : view<64x64xi8, %rhs_layout> -> vector<16xi8> + %i8_rhs3_1_words = vector.bitcast %i8_rhs3_1_raw : vector<16xi8> to vector<4xi32> + %i8_rhs3_1 = vector.fragment %i8_rhs3_1_words shape [%c16, %c16] using {schema = %signed_rhs_schema : encoding} : vector<4xi32> + %i8_rhs3_2_raw = vector.fragment.load %wl_view[%c32, %ln3] shape [%c16, %c16] : view<64x64xi8, %rhs_layout> -> vector<16xi8> + %i8_rhs3_2_words = vector.bitcast %i8_rhs3_2_raw : vector<16xi8> to vector<4xi32> + %i8_rhs3_2 = vector.fragment %i8_rhs3_2_words shape [%c16, %c16] using {schema = %signed_rhs_schema : encoding} : vector<4xi32> + %i8_rhs3_3_raw = vector.fragment.load %wl_view[%i8_c48_outer, %ln3] shape [%c16, %c16] : view<64x64xi8, %rhs_layout> -> vector<16xi8> + %i8_rhs3_3_words = vector.bitcast %i8_rhs3_3_raw : vector<16xi8> to vector<4xi32> + %i8_rhs3_3 = vector.fragment %i8_rhs3_3_words shape [%c16, %c16] using {schema = %signed_rhs_schema : encoding} : vector<4xi32> + %i8_init0_0 = vector.fragment %i8_ic0_0 shape [%c16, %c16] : vector<8xi32> + %i8_mma0_0_0 = vector.mma %i8_lhs0_0, %i8_rhs0_0, %i8_init0_0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma0_0_1 = vector.mma %i8_lhs0_1, %i8_rhs0_1, %i8_mma0_0_0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma0_0_2 = vector.mma %i8_lhs0_2, %i8_rhs0_2, %i8_mma0_0_1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma0_0_3 = vector.mma %i8_lhs0_3, %i8_rhs0_3, %i8_mma0_0_2 : vector<4xi32>, vector<4xi32>, vector<8xi32> + scf.schedule.fence + %i8_init0_1 = vector.fragment %i8_ic0_1 shape [%c16, %c16] : vector<8xi32> + %i8_mma0_1_0 = vector.mma %i8_lhs0_0, %i8_rhs1_0, %i8_init0_1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma0_1_1 = vector.mma %i8_lhs0_1, %i8_rhs1_1, %i8_mma0_1_0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma0_1_2 = vector.mma %i8_lhs0_2, %i8_rhs1_2, %i8_mma0_1_1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma0_1_3 = vector.mma %i8_lhs0_3, %i8_rhs1_3, %i8_mma0_1_2 : vector<4xi32>, vector<4xi32>, vector<8xi32> + scf.schedule.fence + %i8_init0_2 = vector.fragment %i8_ic0_2 shape [%c16, %c16] : vector<8xi32> + %i8_mma0_2_0 = vector.mma %i8_lhs0_0, %i8_rhs2_0, %i8_init0_2 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma0_2_1 = vector.mma %i8_lhs0_1, %i8_rhs2_1, %i8_mma0_2_0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma0_2_2 = vector.mma %i8_lhs0_2, %i8_rhs2_2, %i8_mma0_2_1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma0_2_3 = vector.mma %i8_lhs0_3, %i8_rhs2_3, %i8_mma0_2_2 : vector<4xi32>, vector<4xi32>, vector<8xi32> + scf.schedule.fence + %i8_init0_3 = vector.fragment %i8_ic0_3 shape [%c16, %c16] : vector<8xi32> + %i8_mma0_3_0 = vector.mma %i8_lhs0_0, %i8_rhs3_0, %i8_init0_3 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma0_3_1 = vector.mma %i8_lhs0_1, %i8_rhs3_1, %i8_mma0_3_0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma0_3_2 = vector.mma %i8_lhs0_2, %i8_rhs3_2, %i8_mma0_3_1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma0_3_3 = vector.mma %i8_lhs0_3, %i8_rhs3_3, %i8_mma0_3_2 : vector<4xi32>, vector<4xi32>, vector<8xi32> + scf.schedule.fence + %i8_init1_0 = vector.fragment %i8_ic1_0 shape [%c16, %c16] : vector<8xi32> + %i8_mma1_0_0 = vector.mma %i8_lhs1_0, %i8_rhs0_0, %i8_init1_0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma1_0_1 = vector.mma %i8_lhs1_1, %i8_rhs0_1, %i8_mma1_0_0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma1_0_2 = vector.mma %i8_lhs1_2, %i8_rhs0_2, %i8_mma1_0_1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma1_0_3 = vector.mma %i8_lhs1_3, %i8_rhs0_3, %i8_mma1_0_2 : vector<4xi32>, vector<4xi32>, vector<8xi32> + scf.schedule.fence + %i8_init1_1 = vector.fragment %i8_ic1_1 shape [%c16, %c16] : vector<8xi32> + %i8_mma1_1_0 = vector.mma %i8_lhs1_0, %i8_rhs1_0, %i8_init1_1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma1_1_1 = vector.mma %i8_lhs1_1, %i8_rhs1_1, %i8_mma1_1_0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma1_1_2 = vector.mma %i8_lhs1_2, %i8_rhs1_2, %i8_mma1_1_1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma1_1_3 = vector.mma %i8_lhs1_3, %i8_rhs1_3, %i8_mma1_1_2 : vector<4xi32>, vector<4xi32>, vector<8xi32> + scf.schedule.fence + %i8_init1_2 = vector.fragment %i8_ic1_2 shape [%c16, %c16] : vector<8xi32> + %i8_mma1_2_0 = vector.mma %i8_lhs1_0, %i8_rhs2_0, %i8_init1_2 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma1_2_1 = vector.mma %i8_lhs1_1, %i8_rhs2_1, %i8_mma1_2_0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma1_2_2 = vector.mma %i8_lhs1_2, %i8_rhs2_2, %i8_mma1_2_1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma1_2_3 = vector.mma %i8_lhs1_3, %i8_rhs2_3, %i8_mma1_2_2 : vector<4xi32>, vector<4xi32>, vector<8xi32> + scf.schedule.fence + %i8_init1_3 = vector.fragment %i8_ic1_3 shape [%c16, %c16] : vector<8xi32> + %i8_mma1_3_0 = vector.mma %i8_lhs1_0, %i8_rhs3_0, %i8_init1_3 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma1_3_1 = vector.mma %i8_lhs1_1, %i8_rhs3_1, %i8_mma1_3_0 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma1_3_2 = vector.mma %i8_lhs1_2, %i8_rhs3_2, %i8_mma1_3_1 : vector<4xi32>, vector<4xi32>, vector<8xi32> + %i8_mma1_3_3 = vector.mma %i8_lhs1_3, %i8_rhs3_3, %i8_mma1_3_2 : vector<4xi32>, vector<4xi32>, vector<8xi32> + scf.schedule.fence + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %i8_mma0_0_3, %i8_mma0_1_3, %i8_mma0_2_3, %i8_mma0_3_3, %i8_mma1_0_3, %i8_mma1_1_3, %i8_mma1_2_3, %i8_mma1_3_3 : vector<8xi32>, vector<8xi32>, vector<8xi32>, vector<8xi32>, vector<8xi32>, vector<8xi32>, vector<8xi32>, vector<8xi32> + } + %i8_asi0_0 = index.add %asb0, %c0 : index + %i8_asi0_1 = index.mul %i8_asi0_0, %c2 : index + %i8_asi0_2 = index.add %i8_asi0_1, %lane_hi : index + %i8_asi0_3 = index.mul %i8_asi0_2, %c8 : index + %i8_asv0, %i8_unused_sum0 = func.call @ggml_symmetric_i8_wmmai8_load_metadata(%token256, %scratch, %as_off, %asum_off, %i8_asi0_3) : (i1, buffer, offset, offset, index) -> (vector<8xf32>, vector<8xf32>) + %i8_asi1_0 = index.add %asb1, %c0 : index + %i8_asi1_1 = index.mul %i8_asi1_0, %c2 : index + %i8_asi1_2 = index.add %i8_asi1_1, %lane_hi : index + %i8_asi1_3 = index.mul %i8_asi1_2, %c8 : index + %i8_asv1, %i8_unused_sum1 = func.call @ggml_symmetric_i8_wmmai8_load_metadata(%token256, %scratch, %as_off, %asum_off, %i8_asi1_3) : (i1, buffer, offset, offset, index) -> (vector<8xf32>, vector<8xf32>) + %i8_wsi0_0 = index.add %ln0, %lane_lo : index + %i8_wsi0 = index.assume %i8_wsi0_0 [range(%i8_wsi0_0, 0, 63)] : index + %i8_scale_source_row0 = index.add %i8_source_row_base, %i8_wsi0 : index + %i8_scale_byte0_0 = index.mul %i8_scale_source_row0, %c2 : index + %i8_scale_byte0 = index.add %i8_record_base, %i8_scale_byte0_0 : index + %i8_scale_bytes0 = vector.load %i8_weight_view[%i8_scale_byte0] : view<[%i8_weight_bytes]xi8> -> vector<2xi8> + %i8_scale_halfv0 = vector.bitcast %i8_scale_bytes0 : vector<2xi8> to vector<1xf16> + %i8_scale_half0 = vector.extract %i8_scale_halfv0[0] : vector<1xf16> -> f16 + %i8_wscale0_scalar = scalar.extf %i8_scale_half0 : f16 to f32 + %i8_wscale0 = vector.splat %i8_wscale0_scalar : vector<8xf32> + %i8_wsi1_0 = index.add %ln1, %lane_lo : index + %i8_wsi1 = index.assume %i8_wsi1_0 [range(%i8_wsi1_0, 0, 63)] : index + %i8_scale_source_row1 = index.add %i8_source_row_base, %i8_wsi1 : index + %i8_scale_byte1_0 = index.mul %i8_scale_source_row1, %c2 : index + %i8_scale_byte1 = index.add %i8_record_base, %i8_scale_byte1_0 : index + %i8_scale_bytes1 = vector.load %i8_weight_view[%i8_scale_byte1] : view<[%i8_weight_bytes]xi8> -> vector<2xi8> + %i8_scale_halfv1 = vector.bitcast %i8_scale_bytes1 : vector<2xi8> to vector<1xf16> + %i8_scale_half1 = vector.extract %i8_scale_halfv1[0] : vector<1xf16> -> f16 + %i8_wscale1_scalar = scalar.extf %i8_scale_half1 : f16 to f32 + %i8_wscale1 = vector.splat %i8_wscale1_scalar : vector<8xf32> + %i8_wsi2_0 = index.add %ln2, %lane_lo : index + %i8_wsi2 = index.assume %i8_wsi2_0 [range(%i8_wsi2_0, 0, 63)] : index + %i8_scale_source_row2 = index.add %i8_source_row_base, %i8_wsi2 : index + %i8_scale_byte2_0 = index.mul %i8_scale_source_row2, %c2 : index + %i8_scale_byte2 = index.add %i8_record_base, %i8_scale_byte2_0 : index + %i8_scale_bytes2 = vector.load %i8_weight_view[%i8_scale_byte2] : view<[%i8_weight_bytes]xi8> -> vector<2xi8> + %i8_scale_halfv2 = vector.bitcast %i8_scale_bytes2 : vector<2xi8> to vector<1xf16> + %i8_scale_half2 = vector.extract %i8_scale_halfv2[0] : vector<1xf16> -> f16 + %i8_wscale2_scalar = scalar.extf %i8_scale_half2 : f16 to f32 + %i8_wscale2 = vector.splat %i8_wscale2_scalar : vector<8xf32> + %i8_wsi3_0 = index.add %ln3, %lane_lo : index + %i8_wsi3 = index.assume %i8_wsi3_0 [range(%i8_wsi3_0, 0, 63)] : index + %i8_scale_source_row3 = index.add %i8_source_row_base, %i8_wsi3 : index + %i8_scale_byte3_0 = index.mul %i8_scale_source_row3, %c2 : index + %i8_scale_byte3 = index.add %i8_record_base, %i8_scale_byte3_0 : index + %i8_scale_bytes3 = vector.load %i8_weight_view[%i8_scale_byte3] : view<[%i8_weight_bytes]xi8> -> vector<2xi8> + %i8_scale_halfv3 = vector.bitcast %i8_scale_bytes3 : vector<2xi8> to vector<1xf16> + %i8_scale_half3 = vector.extract %i8_scale_halfv3[0] : vector<1xf16> -> f16 + %i8_wscale3_scalar = scalar.extf %i8_scale_half3 : f16 to f32 + %i8_wscale3 = vector.splat %i8_wscale3_scalar : vector<8xf32> + %i8_fp0_0 = vector.sitofp %i8_i0_0 : vector<8xi32> to vector<8xf32> + %i8_scale0_0 = vector.mulf %i8_asv0, %i8_wscale0 : vector<8xf32> + %i8_next0_0 = vector.fmaf %i8_fp0_0, %i8_scale0_0, %fc0_0 : vector<8xf32> + scf.schedule.fence + %i8_fp0_1 = vector.sitofp %i8_i0_1 : vector<8xi32> to vector<8xf32> + %i8_scale0_1 = vector.mulf %i8_asv0, %i8_wscale1 : vector<8xf32> + %i8_next0_1 = vector.fmaf %i8_fp0_1, %i8_scale0_1, %fc0_1 : vector<8xf32> + scf.schedule.fence + %i8_fp0_2 = vector.sitofp %i8_i0_2 : vector<8xi32> to vector<8xf32> + %i8_scale0_2 = vector.mulf %i8_asv0, %i8_wscale2 : vector<8xf32> + %i8_next0_2 = vector.fmaf %i8_fp0_2, %i8_scale0_2, %fc0_2 : vector<8xf32> + scf.schedule.fence + %i8_fp0_3 = vector.sitofp %i8_i0_3 : vector<8xi32> to vector<8xf32> + %i8_scale0_3 = vector.mulf %i8_asv0, %i8_wscale3 : vector<8xf32> + %i8_next0_3 = vector.fmaf %i8_fp0_3, %i8_scale0_3, %fc0_3 : vector<8xf32> + scf.schedule.fence + %i8_fp1_0 = vector.sitofp %i8_i1_0 : vector<8xi32> to vector<8xf32> + %i8_scale1_0 = vector.mulf %i8_asv1, %i8_wscale0 : vector<8xf32> + %i8_next1_0 = vector.fmaf %i8_fp1_0, %i8_scale1_0, %fc1_0 : vector<8xf32> + scf.schedule.fence + %i8_fp1_1 = vector.sitofp %i8_i1_1 : vector<8xi32> to vector<8xf32> + %i8_scale1_1 = vector.mulf %i8_asv1, %i8_wscale1 : vector<8xf32> + %i8_next1_1 = vector.fmaf %i8_fp1_1, %i8_scale1_1, %fc1_1 : vector<8xf32> + scf.schedule.fence + %i8_fp1_2 = vector.sitofp %i8_i1_2 : vector<8xi32> to vector<8xf32> + %i8_scale1_2 = vector.mulf %i8_asv1, %i8_wscale2 : vector<8xf32> + %i8_next1_2 = vector.fmaf %i8_fp1_2, %i8_scale1_2, %fc1_2 : vector<8xf32> + scf.schedule.fence + %i8_fp1_3 = vector.sitofp %i8_i1_3 : vector<8xi32> to vector<8xf32> + %i8_scale1_3 = vector.mulf %i8_asv1, %i8_wscale3 : vector<8xf32> + %i8_next1_3 = vector.fmaf %i8_fp1_3, %i8_scale1_3, %fc1_3 : vector<8xf32> + scf.schedule.fence + scf.yield %i8_next0_0, %i8_next0_1, %i8_next0_2, %i8_next0_3, %i8_next1_0, %i8_next1_1, %i8_next1_2, %i8_next1_3 : vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32>, vector<8xf32> + } + %out0_0 = scf.if %apply_swiglu -> (vector<8xf32>) { + %gate0_0 = vector.fragment.load %gate_view[%gm0, %gn0] shape [%c16, %c16] : view<[%cols_b]x[%rows_b]xf32> -> vector<8xf32> + %silu0_0 = vector.siluf %gate0_0 : vector<8xf32> + %result0_0 = vector.mulf %silu0_0, %f0_0 : vector<8xf32> + scf.yield %result0_0 : vector<8xf32> + } else { + scf.yield %f0_0 : vector<8xf32> + } + %out0_1 = scf.if %apply_swiglu -> (vector<8xf32>) { + %gate0_1 = vector.fragment.load %gate_view[%gm0, %gn1] shape [%c16, %c16] : view<[%cols_b]x[%rows_b]xf32> -> vector<8xf32> + %silu0_1 = vector.siluf %gate0_1 : vector<8xf32> + %result0_1 = vector.mulf %silu0_1, %f0_1 : vector<8xf32> + scf.yield %result0_1 : vector<8xf32> + } else { + scf.yield %f0_1 : vector<8xf32> + } + %out0_2 = scf.if %apply_swiglu -> (vector<8xf32>) { + %gate0_2 = vector.fragment.load %gate_view[%gm0, %gn2] shape [%c16, %c16] : view<[%cols_b]x[%rows_b]xf32> -> vector<8xf32> + %silu0_2 = vector.siluf %gate0_2 : vector<8xf32> + %result0_2 = vector.mulf %silu0_2, %f0_2 : vector<8xf32> + scf.yield %result0_2 : vector<8xf32> + } else { + scf.yield %f0_2 : vector<8xf32> + } + %out0_3 = scf.if %apply_swiglu -> (vector<8xf32>) { + %gate0_3 = vector.fragment.load %gate_view[%gm0, %gn3] shape [%c16, %c16] : view<[%cols_b]x[%rows_b]xf32> -> vector<8xf32> + %silu0_3 = vector.siluf %gate0_3 : vector<8xf32> + %result0_3 = vector.mulf %silu0_3, %f0_3 : vector<8xf32> + scf.yield %result0_3 : vector<8xf32> + } else { + scf.yield %f0_3 : vector<8xf32> + } + %out1_0 = scf.if %apply_swiglu -> (vector<8xf32>) { + %gate1_0 = vector.fragment.load %gate_view[%gm1, %gn0] shape [%c16, %c16] : view<[%cols_b]x[%rows_b]xf32> -> vector<8xf32> + %silu1_0 = vector.siluf %gate1_0 : vector<8xf32> + %result1_0 = vector.mulf %silu1_0, %f1_0 : vector<8xf32> + scf.yield %result1_0 : vector<8xf32> + } else { + scf.yield %f1_0 : vector<8xf32> + } + %out1_1 = scf.if %apply_swiglu -> (vector<8xf32>) { + %gate1_1 = vector.fragment.load %gate_view[%gm1, %gn1] shape [%c16, %c16] : view<[%cols_b]x[%rows_b]xf32> -> vector<8xf32> + %silu1_1 = vector.siluf %gate1_1 : vector<8xf32> + %result1_1 = vector.mulf %silu1_1, %f1_1 : vector<8xf32> + scf.yield %result1_1 : vector<8xf32> + } else { + scf.yield %f1_1 : vector<8xf32> + } + %out1_2 = scf.if %apply_swiglu -> (vector<8xf32>) { + %gate1_2 = vector.fragment.load %gate_view[%gm1, %gn2] shape [%c16, %c16] : view<[%cols_b]x[%rows_b]xf32> -> vector<8xf32> + %silu1_2 = vector.siluf %gate1_2 : vector<8xf32> + %result1_2 = vector.mulf %silu1_2, %f1_2 : vector<8xf32> + scf.yield %result1_2 : vector<8xf32> + } else { + scf.yield %f1_2 : vector<8xf32> + } + %out1_3 = scf.if %apply_swiglu -> (vector<8xf32>) { + %gate1_3 = vector.fragment.load %gate_view[%gm1, %gn3] shape [%c16, %c16] : view<[%cols_b]x[%rows_b]xf32> -> vector<8xf32> + %silu1_3 = vector.siluf %gate1_3 : vector<8xf32> + %result1_3 = vector.mulf %silu1_3, %f1_3 : vector<8xf32> + scf.yield %result1_3 : vector<8xf32> + } else { + scf.yield %f1_3 : vector<8xf32> + } + + vector.fragment.store %out0_0, %dst_view[%gm0, %gn0] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out0_1, %dst_view[%gm0, %gn1] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out0_2, %dst_view[%gm0, %gn2] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out0_3, %dst_view[%gm0, %gn3] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out1_0, %dst_view[%gm1, %gn0] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out1_1, %dst_view[%gm1, %gn1] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out1_2, %dst_view[%gm1, %gn2] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + vector.fragment.store %out1_3, %dst_view[%gm1, %gn3] shape [%c16, %c16] : vector<8xf32>, view<[%cols_b]x[%rows_b]xf32> + template.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/res_scale_pair_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/res_scale_pair_f32.loom new file mode 100644 index 000000000000..da5b0f616085 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/res_scale_pair_f32.loom @@ -0,0 +1,139 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// One side of ZAYA's residual scale in one dispatch: +// out = (a + b) * s [+ y] +// a, y and out are [row_size, rows]; b and s are [row_size] and broadcast over rows. The bias b +// (config has_bias) and the addend y (config has_addend: the other side, already computed) are +// optional. Every element uses the same add, mul, add order as the ADD / MUL nodes it replaces, +// so the result is bitwise the same; only the dispatches (and their gaps) go. + +amdgpu.target @ggml_res_scale_pair_f32_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.res_scale_pair_f32.has_bias : %value: index where [range(%value, 0, 1)] + +config.decl @ggml.res_scale_pair_f32.has_addend : %value: index where [range(%value, 0, 1)] + +// Grid: x over the channels of a row (256 per workgroup), y over the rows, so no thread divides +// by the runtime row size (AMDGPU Loom lowers only constant divisors). +kernel.def target(@ggml_res_scale_pair_f32_gfx11_wave64) export("ggml_res_scale_pair_f32") @ggml_res_scale_pair_f32(%element_count: index, %row_size: index, %row_count: index) { + %one = index.constant 1 : index + %twofiftysix = index.constant 256 : index + %rounding = index.constant 255 : index + %rounded = index.add %row_size, %rounding : index + %column_groups = index.div %rounded, %twofiftysix : index + kernel.launch.config workgroups(%column_groups, %row_count, %one) workgroup_size(%twofiftysix, %one, %one) : index +} launch(%element_count: index, %row_size: index, %row_count: index, %a: buffer, %b: buffer, %s: buffer, %y: buffer, %output: buffer) { + %count = index.assume %element_count [range(%element_count, 1, 134217728)] : index + %row = index.assume %row_size [range(%row_size, 1, 134217728)] : index + %has_bias = config.get @ggml.res_scale_pair_f32.has_bias : index + %has_addend = config.get @ggml.res_scale_pair_f32.has_addend : index + %column_group = kernel.workgroup.id : index + %row_index = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %twofiftysix = index.constant 256 : index + %zero = index.constant 0 : index + %column_base = index.mul %column_group, %twofiftysix : index + %column0 = index.add %column_base, %workitem : index + %column = index.assume %column0 [range(%column0, 0, 134217983)] : index + %row_offset = index.mul %row_index, %row : index + %linear = index.add %row_offset, %column : index + %in_bounds = index.cmp ult, %column, %row : index + %zero_offset = index.constant 0 : offset + %a_view = buffer.view %a[%zero_offset] : buffer -> view<[%count]xf32> + %b_view = buffer.view %b[%zero_offset] : buffer -> view<[%row]xf32> + %s_view = buffer.view %s[%zero_offset] : buffer -> view<[%row]xf32> + %y_view = buffer.view %y[%zero_offset] : buffer -> view<[%count]xf32> + %output_view = buffer.view %output[%zero_offset] : buffer -> view<[%count]xf32> + %in_rows = index.cmp ult, %linear, %count : index + scf.if %in_bounds { + scf.if %in_rows { + %a_value = view.load %a_view[%linear] : view<[%count]xf32> -> f32 + %s_value = view.load %s_view[%column] : view<[%row]xf32> -> f32 + %use_bias = index.cmp ne, %has_bias, %zero : index + %use_addend = index.cmp ne, %has_addend, %zero : index + %biased = scf.if %use_bias -> (f32) { + %b_value = view.load %b_view[%column] : view<[%row]xf32> -> f32 + %sum = scalar.addf %a_value, %b_value : f32 + scf.yield %sum : f32 + } else { + scf.yield %a_value : f32 + } + %scaled = scalar.mulf %biased, %s_value : f32 + %result = scf.if %use_addend -> (f32) { + %y_value = view.load %y_view[%linear] : view<[%count]xf32> -> f32 + %total = scalar.addf %scaled, %y_value : f32 + scf.yield %total : f32 + } else { + scf.yield %scaled : f32 + } + view.store %result, %output_view[%linear] : f32, view<[%count]xf32> + } + } + kernel.return +} + +// Cases. Rows of 8 channels, 2 rows; period(8) makes an iota a per-channel pattern, so each +// expectation checks the per-channel broadcast of b and s and the per-element a and y. Run with +// --config=ggml.res_scale_pair_f32.has_bias=1 --config=ggml.res_scale_pair_f32.has_addend=1 +// for the first two, and both =0 for the third. + +// (a + b) * s + y with a = c + 1, b = c, s = 2, y = c for channel c: 5c + 2. +check.case public @ggml_res_scale_pair_f32_bias_addend_case { + %count = check.literal value(16) : index + %row = check.literal value(8) : index + %rows = check.literal value(2) : index + %a = check.generate.iota offset(1.0) step(1.0) period(8) : tensor<16xf32> + %b = check.generate.iota offset(0.0) step(1.0) : tensor<8xf32> + %s = check.generate.fill value(2.0) : tensor<8xf32> + %y = check.generate.iota offset(0.0) step(1.0) period(8) : tensor<16xf32> + %output = check.generate.fill value(-7.0) : tensor<16xf32> + %expected = check.generate.iota offset(2.0) step(5.0) period(8) : tensor<16xf32> + kernel.launch @ggml_res_scale_pair_f32[%count, %row, %rows](%count, %row, %rows, %a, %b, %s, %y, %output) : [index, index, index](index, index, index, tensor<16xf32>, tensor<8xf32>, tensor<8xf32>, tensor<16xf32>, tensor<16xf32>) + check.expect.equal actual(%output) expected(%expected) : tensor<16xf32> + check.return +} + +// Per-channel scale: a = 1.5, b = 0.5, s = c + 1, y = 3: 2c + 5. +check.case public @ggml_res_scale_pair_f32_channel_scale_case { + %count = check.literal value(16) : index + %row = check.literal value(8) : index + %rows = check.literal value(2) : index + %a = check.generate.fill value(1.5) : tensor<16xf32> + %b = check.generate.fill value(0.5) : tensor<8xf32> + %s = check.generate.iota offset(1.0) step(1.0) : tensor<8xf32> + %y = check.generate.fill value(3.0) : tensor<16xf32> + %output = check.generate.fill value(-7.0) : tensor<16xf32> + %expected = check.generate.iota offset(5.0) step(2.0) period(8) : tensor<16xf32> + kernel.launch @ggml_res_scale_pair_f32[%count, %row, %rows](%count, %row, %rows, %a, %b, %s, %y, %output) : [index, index, index](index, index, index, tensor<16xf32>, tensor<8xf32>, tensor<8xf32>, tensor<16xf32>, tensor<16xf32>) + check.expect.equal actual(%output) expected(%expected) : tensor<16xf32> + check.return +} + +// Without bias and addend, b and y must not be read: a = c + 1, s = 3, b = y = 100: 3c + 3. +check.case public @ggml_res_scale_pair_f32_plain_case { + %count = check.literal value(16) : index + %row = check.literal value(8) : index + %rows = check.literal value(2) : index + %a = check.generate.iota offset(1.0) step(1.0) period(8) : tensor<16xf32> + %b = check.generate.fill value(100.0) : tensor<8xf32> + %s = check.generate.fill value(3.0) : tensor<8xf32> + %y = check.generate.fill value(100.0) : tensor<16xf32> + %output = check.generate.fill value(-7.0) : tensor<16xf32> + %expected = check.generate.iota offset(3.0) step(3.0) period(8) : tensor<16xf32> + kernel.launch @ggml_res_scale_pair_f32[%count, %row, %rows](%count, %row, %rows, %a, %b, %s, %y, %output) : [index, index, index](index, index, index, tensor<16xf32>, tensor<8xf32>, tensor<8xf32>, tensor<16xf32>, tensor<16xf32>) + check.expect.equal actual(%output) expected(%expected) : tensor<16xf32> + check.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/rmsnorm_binary_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/rmsnorm_binary_f32.loom new file mode 100644 index 000000000000..17dc05f2c92f --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/rmsnorm_binary_f32.loom @@ -0,0 +1,891 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +template.decl @ggml.binary_f32.apply(%arg0: index, %arg1: f32, %arg2: f32) -> (f32) + +template.decl @ggml.rmsnorm_f32.apply(%arg0: f32, %arg1: f32) -> (f32) + +template.decl @ggml.rmsnorm_f32.row_scale(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: f32, %arg5: buffer) -> (f32, index, index, i1) + +amdgpu.target @ggml_rmsnorm_binary_gfx11_wave32 {subgroup_size = 32} + +config.decl @ggml.rmsnorm_binary_f32.hidden_size : %value: index where [range(%value, 128, 32768), mul(%value, 128)] + +config.decl @ggml.rmsnorm_binary_f32.rms_epsilon : f32 + +config.decl @ggml.rmsnorm_binary_f32.op : %value: index where [range(%value, 0, 3)] + +config.decl @ggml.add_rmsnorm_binary_symmetric_i4.hidden_size : %value: index where [range(%value, 128, 32768), mul(%value, 128)] + +config.decl @ggml.add_rmsnorm_binary_symmetric_i4.rms_epsilon : f32 + +config.decl @ggml.rmsnorm_binary_symmetric_i4.hidden_size : %value: index where [range(%value, 128, 32768), mul(%value, 128)] + +config.decl @ggml.rmsnorm_binary_symmetric_i4.rms_epsilon : f32 + + + +config.decl @ggml.rmsnorm_gate_silu_mul_symmetric_i4.hidden_size : %value: index where [range(%value, 128, 32768), mul(%value, 128)] + +config.decl @ggml.rmsnorm_gate_silu_mul_symmetric_i4.rms_epsilon : f32 + +kernel.def target(@ggml_rmsnorm_binary_gfx11_wave32) export("ggml_rmsnorm_binary_f32") @ggml_rmsnorm_binary_f32(%token_count: index) { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%token_count, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %rhs: buffer, %output: buffer) where [range(%token_count, 1, 1048576)] { + %hidden_size0 = config.get @ggml.rmsnorm_binary_f32.hidden_size : index + %epsilon = config.get @ggml.rmsnorm_binary_f32.rms_epsilon : f32 + %op = config.get @ggml.rmsnorm_binary_f32.op : index + %hidden_size = index.assume %hidden_size0 [range(%hidden_size0, 128, 32768), mul(%hidden_size0, 128)] : index + %token0 = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %input_noalias, %rhs_noalias, %output_noalias = buffer.assume.noalias %input, %rhs, %output : buffer, buffer, buffer + %row_scale, %token, %launch_token_count, %valid_token = template.apply<@ggml.rmsnorm_f32.row_scale>(%token_count, %token0, %hidden_size, %hidden_size, %epsilon, %input_noalias) : (index, index, index, index, f32, buffer) -> (f32, index, index, i1) + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%hidden_size]xf32> + %rhs_view = buffer.view %rhs_noalias[%c0_offset] : buffer -> view<[%hidden_size]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%hidden_size]xf32> + scf.if %valid_token { + scf.for %channel = [%workitem to %hidden_size step %c256] { + %value = view.load %input_view[%token, %channel] : view<[%launch_token_count]x[%hidden_size]xf32> -> f32 + %rhs_value = view.load %rhs_view[%channel] : view<[%hidden_size]xf32> -> f32 + %normalized = template.apply<@ggml.rmsnorm_f32.apply>(%value, %row_scale) : (f32, f32) -> (f32) + %result = template.apply<@ggml.binary_f32.apply>(%op, %normalized, %rhs_value) : (index, f32, f32) -> (f32) + view.store %result, %output_view[%token, %channel] : f32, view<[%launch_token_count]x[%hidden_size]xf32> + } + } + kernel.return +} + +// Retain two original lane streams per thread and publish both F32 and K16 F16. +kernel.def target(@ggml_rmsnorm_binary_gfx11_wave32) export("ggml_rmsnorm_binary_f32_k16") @ggml_rmsnorm_binary_f32_k16(%token_count: index) { + %c1 = index.constant 1 : index + %c128 = index.constant 128 : index + kernel.launch.config workgroups(%token_count, %c1, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %rhs: buffer, %output: buffer, %f16_output: buffer) where [range(%token_count, 1, 1048576)] { + %hidden0 = config.get @ggml.rmsnorm_binary_f32.hidden_size : index + %hidden = index.assume %hidden0 [range(%hidden0, 1024, 8192), mul(%hidden0, 1024)] : index + %epsilon = config.get @ggml.rmsnorm_binary_f32.rms_epsilon : f32 + %op = config.get @ggml.rmsnorm_binary_f32.op : index + %tid0 = kernel.workitem.id : index + %tid = index.assume %tid0 [range(%tid0, 0, 127)] : index + %token0 = kernel.workgroup.id : index + %token = index.assume %token0 [lt(%token0, %token_count)] : index + %lane = kernel.subgroup.lane.id : index + %wave = kernel.subgroup.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c256 = index.constant 256 : index + %base = index.constant 0 : offset + %zero = scalar.constant 0.0 : f32 + %cache_zero = vector.constant 0.0 : vector<32xf32> + %input_na, %rhs_na, %output_na = buffer.assume.noalias %input, %rhs, %output : buffer, buffer, buffer + %input_view = buffer.view %input_na[%base] : buffer -> view<[%token_count]x[%hidden]xf32> + %rhs_view = buffer.view %rhs_na[%base] : buffer -> view<[%hidden]xf32> + %output_view = buffer.view %output_na[%base] : buffer -> view<[%token_count]x[%hidden]xf32> + %packed_columns = index.div %hidden, %c16 : index + %packed_view = buffer.view %f16_output[%base] : buffer -> view<[%packed_columns]x[%token_count]x16xf16> + %first_channel = index.mul %tid, %c2 : index + %stripe_count = index.div %hidden, %c256 : index + %sum0, %sum1, %cache0, %cache1, %cache2, %cache3 = scf.for %stripe = [%c0 to %stripe_count step %c1](%running0 = %zero : f32, %running1 = %zero : f32, %values0 = %cache_zero : vector<32xf32>, %values1 = %cache_zero : vector<32xf32>, %values2 = %cache_zero : vector<32xf32>, %values3 = %cache_zero : vector<32xf32>) -> (f32, f32, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32>) unroll { + %channel0 = index.madd %stripe, %c256, %first_channel : index + %channel = index.assume %channel0 [range(%channel0, 0, 8190), mul(%channel0, 2), lt(%channel0, %hidden)] : index + %input2 = vector.load %input_view[%token, %channel] : view<[%token_count]x[%hidden]xf32> -> vector<2xf32> + %cache_phase = index.div %stripe, %c16 : index + %cache_stripe = index.rem %stripe, %c16 : index + %cache_base = index.mul %cache_stripe, %c2 : index + %value0 = vector.extract %input2[%c0] : vector<2xf32> -> f32 + %square0 = scalar.mulf %value0, %value0 : f32 + %next_sum0 = scalar.addf %running0, %square0 : f32 + %cache_index0 = index.add %cache_base, %c0 : index + %value1 = vector.extract %input2[%c1] : vector<2xf32> -> f32 + %square1 = scalar.mulf %value1, %value1 : f32 + %next_sum1 = scalar.addf %running1, %square1 : f32 + %cache_index1 = index.add %cache_base, %c1 : index + %phase0 = index.cmp eq, %cache_phase, %c0 : index + %next_values0 = scf.if %phase0 -> (vector<32xf32>) { + %insert0_0 = vector.insert %value0 into %values0[%cache_index0] : f32, vector<32xf32> + %insert0_1 = vector.insert %value1 into %insert0_0[%cache_index1] : f32, vector<32xf32> + scf.yield %insert0_1 : vector<32xf32> + } else { + scf.yield %values0 : vector<32xf32> + } + %phase1 = index.cmp eq, %cache_phase, %c1 : index + %next_values1 = scf.if %phase1 -> (vector<32xf32>) { + %insert1_0 = vector.insert %value0 into %values1[%cache_index0] : f32, vector<32xf32> + %insert1_1 = vector.insert %value1 into %insert1_0[%cache_index1] : f32, vector<32xf32> + scf.yield %insert1_1 : vector<32xf32> + } else { + scf.yield %values1 : vector<32xf32> + } + %phase2 = index.cmp eq, %cache_phase, %c2 : index + %next_values2 = scf.if %phase2 -> (vector<32xf32>) { + %insert2_0 = vector.insert %value0 into %values2[%cache_index0] : f32, vector<32xf32> + %insert2_1 = vector.insert %value1 into %insert2_0[%cache_index1] : f32, vector<32xf32> + scf.yield %insert2_1 : vector<32xf32> + } else { + scf.yield %values2 : vector<32xf32> + } + %phase3 = index.cmp eq, %cache_phase, %c3 : index + %next_values3 = scf.if %phase3 -> (vector<32xf32>) { + %insert3_0 = vector.insert %value0 into %values3[%cache_index0] : f32, vector<32xf32> + %insert3_1 = vector.insert %value1 into %insert3_0[%cache_index1] : f32, vector<32xf32> + scf.yield %insert3_1 : vector<32xf32> + } else { + scf.yield %values3 : vector<32xf32> + } + %fence_pos = index.rem %stripe, %c4 : index + %end_packet = index.cmp eq, %fence_pos, %c3 : index + scf.if %end_packet { + scf.schedule.fence + } + scf.yield %next_sum0, %next_sum1, %next_values0, %next_values1, %next_values2, %next_values3 : f32, f32, vector<32xf32>, vector<32xf32>, vector<32xf32>, vector<32xf32> + } + // Two original lane sums form a pair; sixteen threads form one original wave. + %pair_sum = scalar.addf %sum0, %sum1 : f32 + %shuffle_width = scalar.constant 32 : i32 + %xor1 = scalar.constant 1 : i32 + %xor2 = scalar.constant 2 : i32 + %xor4 = scalar.constant 4 : i32 + %xor8 = scalar.constant 8 : i32 + %peer1, %valid1 = kernel.subgroup.shuffle %pair_sum, %xor1, %shuffle_width : f32, i32, i32 + %sum4 = scalar.addf %pair_sum, %peer1 : f32 + %peer2, %valid2 = kernel.subgroup.shuffle %sum4, %xor2, %shuffle_width : f32, i32, i32 + %sum8 = scalar.addf %sum4, %peer2 : f32 + %peer4, %valid4 = kernel.subgroup.shuffle %sum8, %xor4, %shuffle_width : f32, i32, i32 + %sum16 = scalar.addf %sum8, %peer4 : f32 + %peer8, %valid8 = kernel.subgroup.shuffle %sum16, %xor8, %shuffle_width : f32, i32, i32 + %original_wave_sum = scalar.addf %sum16, %peer8 : f32 + %original_wave = index.div %tid, %c16 : index + %cohort_lane = index.rem %tid, %c16 : index + %cohort_leader = index.cmp eq, %cohort_lane, %c0 : index + %scratch_bytes = index.constant 1024 : offset + %scratch = buffer.alloca align(16) %scratch_bytes : buffer + %scratch_view = buffer.view %scratch[%base] : buffer -> view<256xf32> + scf.if %cohort_leader { + view.store %original_wave_sum, %scratch_view[%original_wave] : f32, view<256xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %first_wave = index.cmp eq, %wave, %c0 : index + %live_partial = index.cmp ult, %lane, %c8 : index + %loads_partial = scalar.andi %first_wave, %live_partial : i1 + %partial = scf.if %loads_partial -> (f32) { + %loaded = view.load %scratch_view[%lane] : view<256xf32> -> f32 + scf.yield %loaded : f32 + } else { + scf.yield %zero : f32 + } + %row_sum = kernel.subgroup.reduce %partial : f32 + %leader = index.cmp eq, %tid, %c0 : index + scf.if %leader { + %hidden_i32 = index.cast %hidden : index to i32 + %hidden_f32 = scalar.sitofp %hidden_i32 : i32 to f32 + %mean = scalar.divf %row_sum, %hidden_f32 : f32 + %biased_mean = scalar.addf %mean, %epsilon : f32 + %scale = scalar.rsqrtf %biased_mean : f32 + view.store %scale, %scratch_view[%c0] : f32, view<256xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %row_scale = view.load %scratch_view[%c0] : view<256xf32> -> f32 + %pair_count = index.div %stripe_count, %c2 : index + %parity = index.rem %tid, %c2 : index + %odd = index.cmp ne, %parity, %c0 : index + %aligned_tid = index.sub %tid, %parity : index + %aligned_channel = index.mul %aligned_tid, %c2 : index + %stripe_lane = index.mul %parity, %c256 : index + %output_origin = index.add %aligned_channel, %stripe_lane : index + %c512 = index.constant 512 : index + scf.for %pair = [%c0 to %pair_count step %c1] unroll { + %channel0 = index.madd %pair, %c512, %output_origin : index + %channel = index.assume %channel0 [range(%channel0, 0, 8188), mul(%channel0, 4), lt(%channel0, %hidden)] : index + %stripe = index.mul %pair, %c2 : index + %phase = index.div %stripe, %c16 : index + %phase0 = index.cmp eq, %phase, %c0 : index + %phase1 = index.cmp eq, %phase, %c1 : index + %phase2 = index.cmp eq, %phase, %c2 : index + %values23 = scf.select %phase2, %cache2, %cache3 : vector<32xf32> + %values123 = scf.select %phase1, %cache1, %values23 : vector<32xf32> + %values = scf.select %phase0, %cache0, %values123 : vector<32xf32> + %cache_stripe = index.rem %stripe, %c16 : index + %cache_base = index.mul %cache_stripe, %c2 : index + %cache_index0 = index.add %cache_base, %c0 : index + %cached0 = vector.extract %values[%cache_index0] : vector<32xf32> -> f32 + %cache_index1 = index.add %cache_base, %c1 : index + %cached1 = vector.extract %values[%cache_index1] : vector<32xf32> -> f32 + %cache_index2 = index.add %cache_base, %c2 : index + %cached2 = vector.extract %values[%cache_index2] : vector<32xf32> -> f32 + %cache_index3 = index.add %cache_base, %c3 : index + %cached3 = vector.extract %values[%cache_index3] : vector<32xf32> -> f32 + %send0 = scf.select %odd, %cached0, %cached2 : f32 + %send1 = scf.select %odd, %cached1, %cached3 : f32 + %received0, %receive_valid0 = kernel.subgroup.shuffle %send0, %xor1, %shuffle_width : f32, i32, i32 + %received1, %receive_valid1 = kernel.subgroup.shuffle %send1, %xor1, %shuffle_width : f32, i32, i32 + %value0 = scf.select %odd, %received0, %cached0 : f32 + %value1 = scf.select %odd, %received1, %cached1 : f32 + %value2 = scf.select %odd, %cached2, %received0 : f32 + %value3 = scf.select %odd, %cached3, %received1 : f32 + %weights = vector.load %rhs_view[%channel] : view<[%hidden]xf32> -> vector<4xf32> + %result_zero = vector.constant 0.0 : vector<4xf32> + %weight0 = vector.extract %weights[%c0] : vector<4xf32> -> f32 + %norm0 = template.apply<@ggml.rmsnorm_f32.apply>(%value0, %row_scale) : (f32, f32) -> (f32) + %out0 = template.apply<@ggml.binary_f32.apply>(%op, %norm0, %weight0) : (index, f32, f32) -> (f32) + %result0 = vector.insert %out0 into %result_zero[%c0] : f32, vector<4xf32> + %weight1 = vector.extract %weights[%c1] : vector<4xf32> -> f32 + %norm1 = template.apply<@ggml.rmsnorm_f32.apply>(%value1, %row_scale) : (f32, f32) -> (f32) + %out1 = template.apply<@ggml.binary_f32.apply>(%op, %norm1, %weight1) : (index, f32, f32) -> (f32) + %result1 = vector.insert %out1 into %result0[%c1] : f32, vector<4xf32> + %weight2 = vector.extract %weights[%c2] : vector<4xf32> -> f32 + %norm2 = template.apply<@ggml.rmsnorm_f32.apply>(%value2, %row_scale) : (f32, f32) -> (f32) + %out2 = template.apply<@ggml.binary_f32.apply>(%op, %norm2, %weight2) : (index, f32, f32) -> (f32) + %result2 = vector.insert %out2 into %result1[%c2] : f32, vector<4xf32> + %weight3 = vector.extract %weights[%c3] : vector<4xf32> -> f32 + %norm3 = template.apply<@ggml.rmsnorm_f32.apply>(%value3, %row_scale) : (f32, f32) -> (f32) + %out3 = template.apply<@ggml.binary_f32.apply>(%op, %norm3, %weight3) : (index, f32, f32) -> (f32) + %result3 = vector.insert %out3 into %result2[%c3] : f32, vector<4xf32> + vector.store %result3, %output_view[%token, %channel] : vector<4xf32>, view<[%token_count]x[%hidden]xf32> + %packed_column = index.div %channel, %c16 : index + %packed_lane = index.rem %channel, %c16 : index + %half = vector.fptrunc %result3 : vector<4xf32> to vector<4xf16> + vector.store %half, %packed_view[%packed_column, %token, %packed_lane] : vector<4xf16>, view<[%packed_columns]x[%token_count]x16xf16> + %fence_pos = index.rem %pair, %c4 : index + %end_packet = index.cmp eq, %fence_pos, %c3 : index + scf.if %end_packet { + scf.schedule.fence + } + } + kernel.return +} + +kernel.def target(@ggml_rmsnorm_binary_gfx11_wave32) export("ggml_rmsnorm_binary_symmetric_i4_k32") @ggml_rmsnorm_binary_symmetric_i4_k32(%token_count: index) { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%token_count, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer, %i4_qs: buffer, %i4_ds: buffer, %i4_sums: buffer) where [range(%token_count, 1, 16)] { + %hidden_size0 = config.get @ggml.rmsnorm_binary_symmetric_i4.hidden_size : index + %epsilon = config.get @ggml.rmsnorm_binary_symmetric_i4.rms_epsilon : f32 + %hidden_size = index.assume %hidden_size0 [range(%hidden_size0, 128, 32768), mul(%hidden_size0, 128)] : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 16)] : index + %token0 = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %scratch_bytes = index.constant 1024 : offset + %c0_offset = index.constant 0 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %amax_epsilon = scalar.constant 1.0000000000000001e-30 : f32 + %one_seventh = scalar.constant 0.14285714285714285 : f32 + %seven = scalar.constant 7.0 : f32 + %xor1 = scalar.constant 1 : i32 + %xor2 = scalar.constant 2 : i32 + %xor4 = scalar.constant 4 : i32 + %xor8 = scalar.constant 8 : i32 + %shuffle_width = scalar.constant 32 : i32 + %shift4 = scalar.constant 4 : i32 + %shift8 = scalar.constant 8 : i32 + %shift12 = scalar.constant 12 : i32 + %nibble_mask = vector.constant 15 : vector<4xi32> + %valid_token = index.cmp ult, %token0, %bounded_token_count : index + %safe_token0 = scf.select %valid_token, %token0, %c0 : index + %token = index.assume %safe_token0 [lt(%safe_token0, %bounded_token_count)] : index + + %input_noalias, %weight_noalias, %output_noalias, %i4_qs_noalias, %i4_ds_noalias, %i4_sums_noalias = buffer.assume.noalias %input, %weight, %output, %i4_qs, %i4_ds, %i4_sums : buffer, buffer, buffer, buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%hidden_size]xf32> + %weight_view = buffer.view %weight_noalias[%c0_offset] : buffer -> view<[%hidden_size]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%hidden_size]xf32> + %element_count = index.mul %bounded_token_count, %hidden_size : index + %group32_count = index.div %element_count, %c32 : index + %group64_count_per_row = index.div %hidden_size, %c64 : index + %qs_halfword_count = index.div %element_count, %c4 : index + %i4_qs_view = buffer.view %i4_qs_noalias[%c0_offset] : buffer -> view<[%qs_halfword_count]xi16> + %i4_ds_view = buffer.view %i4_ds_noalias[%c0_offset] : buffer -> view<[%group32_count]xf32> + %i4_sums_view = buffer.view %i4_sums_noalias[%c0_offset] : buffer -> view<[%group32_count]xi32> + + %thread_sum = scf.for %channel = [%workitem to %hidden_size step %c256](%running_sum = %c0_f32 : f32) -> (f32) { + %input_value = view.load %input_view[%token, %channel] : view<[%bounded_token_count]x[%hidden_size]xf32> -> f32 + %square = scalar.mulf %input_value, %input_value : f32 + %next_sum = scalar.addf %running_sum, %square : f32 + scf.yield %next_sum : f32 + } + %subgroup_sum = kernel.subgroup.reduce %thread_sum : f32 + %scratch = buffer.alloca align(16) %scratch_bytes : buffer + %scratch_view = buffer.view %scratch[%c0_offset] : buffer -> view<256xf32> + %is_subgroup_leader = index.cmp eq, %lane, %c0 : index + scf.if %is_subgroup_leader { + view.store %subgroup_sum, %scratch_view[%subgroup] : f32, view<256xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %is_reduction_subgroup = index.cmp eq, %subgroup, %c0 : index + %is_reduction_lane = index.cmp ult, %lane, %c8 : index + %loads_subgroup_sum = scalar.andi %is_reduction_subgroup, %is_reduction_lane : i1 + %subgroup_partial = scf.if %loads_subgroup_sum -> (f32) { + %value = view.load %scratch_view[%lane] : view<256xf32> -> f32 + scf.yield %value : f32 + } else { + scf.yield %c0_f32 : f32 + } + %row_sum = kernel.subgroup.reduce %subgroup_partial : f32 + %writes_scale = scalar.andi %is_reduction_subgroup, %is_subgroup_leader : i1 + scf.if %writes_scale { + %hidden_size_i32 = index.cast %hidden_size : index to i32 + %hidden_size_f32 = scalar.sitofp %hidden_size_i32 : i32 to f32 + %mean = scalar.divf %row_sum, %hidden_size_f32 : f32 + %biased_mean = scalar.addf %mean, %epsilon : f32 + %row_scale = scalar.rsqrtf %biased_mean : f32 + view.store %row_scale, %scratch_view[%c0] : f32, view<256xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %row_scale = view.load %scratch_view[%c0] : view<256xf32> -> f32 + + %wave = index.div %workitem, %c32 : index + %group_in_wave = index.div %lane, %c16 : index + %lane_in_group = index.rem %lane, %c16 : index + %group_half = index.div %lane_in_group, %c8 : index + %lane_in_half = index.rem %lane_in_group, %c8 : index + %is_scale_leader = index.cmp eq, %lane_in_group, %c0 : index + %is_sum_leader = index.cmp eq, %lane_in_half, %c0 : index + %wave_group_base = index.mul %wave, %c2 : index + %local_group = index.add %wave_group_base, %group_in_wave : index + %token_group_base = index.mul %token, %group64_count_per_row : index + scf.if %valid_token { + scf.for %group_base = [%c0 to %group64_count_per_row step %c16] { + %group_in_row = index.add %group_base, %local_group : index + %in_range = index.cmp ult, %group_in_row, %group64_count_per_row : index + scf.if %in_range { + %channel_base = index.mul %group_in_row, %c64 : index + %lane_channel = index.mul %lane_in_group, %c4 : index + %channel = index.add %channel_base, %lane_channel : index + %input_values = vector.load %input_view[%token, %channel] : view<[%bounded_token_count]x[%hidden_size]xf32> -> vector<4xf32> + %learned_weights = vector.load %weight_view[%channel] : view<[%hidden_size]xf32> -> vector<4xf32> + %scale_vector = vector.splat %row_scale : vector<4xf32> + %normalized = vector.mulf %input_values, %scale_vector : vector<4xf32> + %result = vector.mulf %normalized, %learned_weights : vector<4xf32> + vector.store %result, %output_view[%token, %channel] : vector<4xf32>, view<[%bounded_token_count]x[%hidden_size]xf32> + + %absolute_values = vector.absf %result : vector<4xf32> + %lane_max = vector.reduce %absolute_values, %c0_f32 : vector<4xf32>, f32 + %max_x1_peer, %max_x1_valid = kernel.subgroup.shuffle %lane_max, %xor1, %shuffle_width : f32, i32, i32 + %max_x1 = scalar.maxnumf %lane_max, %max_x1_peer : f32 + %max_x2_peer, %max_x2_valid = kernel.subgroup.shuffle %max_x1, %xor2, %shuffle_width : f32, i32, i32 + %max_x2 = scalar.maxnumf %max_x1, %max_x2_peer : f32 + %max_x4_peer, %max_x4_valid = kernel.subgroup.shuffle %max_x2, %xor4, %shuffle_width : f32, i32, i32 + %max_x4 = scalar.maxnumf %max_x2, %max_x4_peer : f32 + %max_x8_peer, %max_x8_valid = kernel.subgroup.shuffle %max_x4, %xor8, %shuffle_width : f32, i32, i32 + %group_max = scalar.maxnumf %max_x4, %max_x8_peer : f32 + %amax = scalar.maxnumf %group_max, %amax_epsilon : f32 + %activation_scale = scalar.mulf %amax, %one_seventh : f32 + %activation_rscale = scalar.divf %seven, %amax : f32 + %group64 = index.add %token_group_base, %group_in_row : index + %group32_base = index.mul %group64, %c2 : index + %group32_high = index.add %group32_base, %c1 : index + scf.if %is_scale_leader { + view.store %activation_scale, %i4_ds_view[%group32_base] : f32, view<[%group32_count]xf32> + view.store %activation_scale, %i4_ds_view[%group32_high] : f32, view<[%group32_count]xf32> + } + + %rscale_vector = vector.splat %activation_rscale : vector<4xf32> + %scaled = vector.mulf %result, %rscale_vector : vector<4xf32> + %rounded = vector.roundf %scaled : vector<4xf32> + %quantized_i32 = vector.fptosi %rounded : vector<4xf32> to vector<4xi32> + %quantized_nibbles = vector.andi %quantized_i32, %nibble_mask : vector<4xi32> + %q0 = vector.extract %quantized_nibbles[0] : vector<4xi32> -> i32 + %q1 = vector.extract %quantized_nibbles[1] : vector<4xi32> -> i32 + %q2 = vector.extract %quantized_nibbles[2] : vector<4xi32> -> i32 + %q3 = vector.extract %quantized_nibbles[3] : vector<4xi32> -> i32 + %q1_shifted = scalar.shli %q1, %shift4 : i32 + %q2_shifted = scalar.shli %q2, %shift8 : i32 + %q3_shifted = scalar.shli %q3, %shift12 : i32 + %q01 = scalar.ori %q0, %q1_shifted : i32 + %q23 = scalar.ori %q2_shifted, %q3_shifted : i32 + %packed_i32 = scalar.ori %q01, %q23 : i32 + %packed_i16 = scalar.trunci %packed_i32 : i32 to i16 + %word_base = index.mul %group64, %c16 : index + %word_index = index.add %word_base, %lane_in_group : index + view.store %packed_i16, %i4_qs_view[%word_index] : i16, view<[%qs_halfword_count]xi16> + + %lane_sum = vector.reduce %rounded, %c0_f32 : vector<4xf32>, f32 + %sum_x1_peer, %sum_x1_valid = kernel.subgroup.shuffle %lane_sum, %xor1, %shuffle_width : f32, i32, i32 + %sum_x1 = scalar.addf %lane_sum, %sum_x1_peer : f32 + %sum_x2_peer, %sum_x2_valid = kernel.subgroup.shuffle %sum_x1, %xor2, %shuffle_width : f32, i32, i32 + %sum_x2 = scalar.addf %sum_x1, %sum_x2_peer : f32 + %sum_x4_peer, %sum_x4_valid = kernel.subgroup.shuffle %sum_x2, %xor4, %shuffle_width : f32, i32, i32 + %group_sum = scalar.addf %sum_x2, %sum_x4_peer : f32 + scf.if %is_sum_leader { + %sum_value = scalar.fptosi %group_sum : f32 to i32 + %sum_index = index.add %group32_base, %group_half : index + view.store %sum_value, %i4_sums_view[%sum_index] : i32, view<[%group32_count]xi32> + } + } + } + } + kernel.return +} + +// Publishes residual ADD, RMSNorm, learned scale, and a signed-I4/K64 alternate. +kernel.def target(@ggml_rmsnorm_binary_gfx11_wave32) export("ggml_add_rmsnorm_binary_symmetric_i4_k32") @ggml_add_rmsnorm_binary_symmetric_i4_k32(%token_count: index) { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%token_count, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %lhs: buffer, %rhs: buffer, %residual_output: buffer, %weight: buffer, %output: buffer, %i4_qs: buffer, %i4_ds: buffer, %i4_sums: buffer) where [range(%token_count, 1, 16)] { + %hidden_size0 = config.get @ggml.add_rmsnorm_binary_symmetric_i4.hidden_size : index + %epsilon = config.get @ggml.add_rmsnorm_binary_symmetric_i4.rms_epsilon : f32 + %hidden_size = index.assume %hidden_size0 [range(%hidden_size0, 128, 32768), mul(%hidden_size0, 128)] : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 16)] : index + %token0 = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %scratch_bytes = index.constant 1024 : offset + %c0_offset = index.constant 0 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %amax_epsilon = scalar.constant 1.0000000000000001e-30 : f32 + %one_seventh = scalar.constant 0.14285714285714285 : f32 + %seven = scalar.constant 7.0 : f32 + %xor1 = scalar.constant 1 : i32 + %xor2 = scalar.constant 2 : i32 + %xor4 = scalar.constant 4 : i32 + %xor8 = scalar.constant 8 : i32 + %shuffle_width = scalar.constant 32 : i32 + %shift4 = scalar.constant 4 : i32 + %shift8 = scalar.constant 8 : i32 + %shift12 = scalar.constant 12 : i32 + %nibble_mask = vector.constant 15 : vector<4xi32> + %valid_token = index.cmp ult, %token0, %bounded_token_count : index + %safe_token0 = scf.select %valid_token, %token0, %c0 : index + %token = index.assume %safe_token0 [lt(%safe_token0, %bounded_token_count)] : index + + %lhs_noalias, %rhs_noalias, %residual_noalias, %weight_noalias, %output_noalias, %i4_qs_noalias, %i4_ds_noalias, %i4_sums_noalias = buffer.assume.noalias %lhs, %rhs, %residual_output, %weight, %output, %i4_qs, %i4_ds, %i4_sums : buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer + %lhs_view = buffer.view %lhs_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%hidden_size]xf32> + %rhs_view = buffer.view %rhs_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%hidden_size]xf32> + %residual_view = buffer.view %residual_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%hidden_size]xf32> + %weight_view = buffer.view %weight_noalias[%c0_offset] : buffer -> view<[%hidden_size]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%hidden_size]xf32> + %element_count = index.mul %bounded_token_count, %hidden_size : index + %group32_count = index.div %element_count, %c32 : index + %group64_count_per_row = index.div %hidden_size, %c64 : index + %qs_halfword_count = index.div %element_count, %c4 : index + %i4_qs_view = buffer.view %i4_qs_noalias[%c0_offset] : buffer -> view<[%qs_halfword_count]xi16> + %i4_ds_view = buffer.view %i4_ds_noalias[%c0_offset] : buffer -> view<[%group32_count]xf32> + %i4_sums_view = buffer.view %i4_sums_noalias[%c0_offset] : buffer -> view<[%group32_count]xi32> + + %thread_sum = scf.for %channel = [%workitem to %hidden_size step %c256](%running_sum = %c0_f32 : f32) -> (f32) { + %lhs_value = view.load %lhs_view[%token, %channel] : view<[%bounded_token_count]x[%hidden_size]xf32> -> f32 + %rhs_value = view.load %rhs_view[%token, %channel] : view<[%bounded_token_count]x[%hidden_size]xf32> -> f32 + %residual = scalar.addf %lhs_value, %rhs_value : f32 + view.store %residual, %residual_view[%token, %channel] : f32, view<[%bounded_token_count]x[%hidden_size]xf32> + %square = scalar.mulf %residual, %residual : f32 + %next_sum = scalar.addf %running_sum, %square : f32 + scf.yield %next_sum : f32 + } + %subgroup_sum = kernel.subgroup.reduce %thread_sum : f32 + %scratch = buffer.alloca align(16) %scratch_bytes : buffer + %scratch_view = buffer.view %scratch[%c0_offset] : buffer -> view<256xf32> + %is_subgroup_leader = index.cmp eq, %lane, %c0 : index + scf.if %is_subgroup_leader { + view.store %subgroup_sum, %scratch_view[%subgroup] : f32, view<256xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %is_reduction_subgroup = index.cmp eq, %subgroup, %c0 : index + %is_reduction_lane = index.cmp ult, %lane, %c8 : index + %loads_subgroup_sum = scalar.andi %is_reduction_subgroup, %is_reduction_lane : i1 + %subgroup_partial = scf.if %loads_subgroup_sum -> (f32) { + %value = view.load %scratch_view[%lane] : view<256xf32> -> f32 + scf.yield %value : f32 + } else { + scf.yield %c0_f32 : f32 + } + %row_sum = kernel.subgroup.reduce %subgroup_partial : f32 + %writes_scale = scalar.andi %is_reduction_subgroup, %is_subgroup_leader : i1 + scf.if %writes_scale { + %hidden_size_i32 = index.cast %hidden_size : index to i32 + %hidden_size_f32 = scalar.sitofp %hidden_size_i32 : i32 to f32 + %mean = scalar.divf %row_sum, %hidden_size_f32 : f32 + %biased_mean = scalar.addf %mean, %epsilon : f32 + %row_scale = scalar.rsqrtf %biased_mean : f32 + view.store %row_scale, %scratch_view[%c0] : f32, view<256xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %row_scale = view.load %scratch_view[%c0] : view<256xf32> -> f32 + kernel.barrier scope(workgroup) ordering(acq_rel) + + %wave = index.div %workitem, %c32 : index + %group_in_wave = index.div %lane, %c16 : index + %lane_in_group = index.rem %lane, %c16 : index + %group_half = index.div %lane_in_group, %c8 : index + %lane_in_half = index.rem %lane_in_group, %c8 : index + %is_scale_leader = index.cmp eq, %lane_in_group, %c0 : index + %is_sum_leader = index.cmp eq, %lane_in_half, %c0 : index + %wave_group_base = index.mul %wave, %c2 : index + %local_group = index.add %wave_group_base, %group_in_wave : index + %token_group_base = index.mul %token, %group64_count_per_row : index + scf.if %valid_token { + scf.for %group_base = [%c0 to %group64_count_per_row step %c16] { + %group_in_row = index.add %group_base, %local_group : index + %in_range = index.cmp ult, %group_in_row, %group64_count_per_row : index + scf.if %in_range { + %channel_base = index.mul %group_in_row, %c64 : index + %lane_channel = index.mul %lane_in_group, %c4 : index + %channel = index.add %channel_base, %lane_channel : index + %residual_values = vector.load %residual_view[%token, %channel] : view<[%bounded_token_count]x[%hidden_size]xf32> -> vector<4xf32> + %learned_weights = vector.load %weight_view[%channel] : view<[%hidden_size]xf32> -> vector<4xf32> + %scale_vector = vector.splat %row_scale : vector<4xf32> + %normalized = vector.mulf %residual_values, %scale_vector : vector<4xf32> + %result = vector.mulf %normalized, %learned_weights : vector<4xf32> + vector.store %result, %output_view[%token, %channel] : vector<4xf32>, view<[%bounded_token_count]x[%hidden_size]xf32> + + %absolute_values = vector.absf %result : vector<4xf32> + %lane_max = vector.reduce %absolute_values, %c0_f32 : vector<4xf32>, f32 + %max_x1_peer, %max_x1_valid = kernel.subgroup.shuffle %lane_max, %xor1, %shuffle_width : f32, i32, i32 + %max_x1 = scalar.maxnumf %lane_max, %max_x1_peer : f32 + %max_x2_peer, %max_x2_valid = kernel.subgroup.shuffle %max_x1, %xor2, %shuffle_width : f32, i32, i32 + %max_x2 = scalar.maxnumf %max_x1, %max_x2_peer : f32 + %max_x4_peer, %max_x4_valid = kernel.subgroup.shuffle %max_x2, %xor4, %shuffle_width : f32, i32, i32 + %max_x4 = scalar.maxnumf %max_x2, %max_x4_peer : f32 + %max_x8_peer, %max_x8_valid = kernel.subgroup.shuffle %max_x4, %xor8, %shuffle_width : f32, i32, i32 + %group_max = scalar.maxnumf %max_x4, %max_x8_peer : f32 + %amax = scalar.maxnumf %group_max, %amax_epsilon : f32 + %activation_scale = scalar.mulf %amax, %one_seventh : f32 + %activation_rscale = scalar.divf %seven, %amax : f32 + %group64 = index.add %token_group_base, %group_in_row : index + %group32_base = index.mul %group64, %c2 : index + %group32_high = index.add %group32_base, %c1 : index + scf.if %is_scale_leader { + view.store %activation_scale, %i4_ds_view[%group32_base] : f32, view<[%group32_count]xf32> + view.store %activation_scale, %i4_ds_view[%group32_high] : f32, view<[%group32_count]xf32> + } + + %rscale_vector = vector.splat %activation_rscale : vector<4xf32> + %scaled = vector.mulf %result, %rscale_vector : vector<4xf32> + %rounded = vector.roundf %scaled : vector<4xf32> + %quantized_i32 = vector.fptosi %rounded : vector<4xf32> to vector<4xi32> + %quantized_nibbles = vector.andi %quantized_i32, %nibble_mask : vector<4xi32> + %q0 = vector.extract %quantized_nibbles[0] : vector<4xi32> -> i32 + %q1 = vector.extract %quantized_nibbles[1] : vector<4xi32> -> i32 + %q2 = vector.extract %quantized_nibbles[2] : vector<4xi32> -> i32 + %q3 = vector.extract %quantized_nibbles[3] : vector<4xi32> -> i32 + %q1_shifted = scalar.shli %q1, %shift4 : i32 + %q2_shifted = scalar.shli %q2, %shift8 : i32 + %q3_shifted = scalar.shli %q3, %shift12 : i32 + %q01 = scalar.ori %q0, %q1_shifted : i32 + %q23 = scalar.ori %q2_shifted, %q3_shifted : i32 + %packed_i32 = scalar.ori %q01, %q23 : i32 + %packed_i16 = scalar.trunci %packed_i32 : i32 to i16 + %word_base = index.mul %group64, %c16 : index + %word_index = index.add %word_base, %lane_in_group : index + view.store %packed_i16, %i4_qs_view[%word_index] : i16, view<[%qs_halfword_count]xi16> + + %lane_sum = vector.reduce %rounded, %c0_f32 : vector<4xf32>, f32 + %sum_x1_peer, %sum_x1_valid = kernel.subgroup.shuffle %lane_sum, %xor1, %shuffle_width : f32, i32, i32 + %sum_x1 = scalar.addf %lane_sum, %sum_x1_peer : f32 + %sum_x2_peer, %sum_x2_valid = kernel.subgroup.shuffle %sum_x1, %xor2, %shuffle_width : f32, i32, i32 + %sum_x2 = scalar.addf %sum_x1, %sum_x2_peer : f32 + %sum_x4_peer, %sum_x4_valid = kernel.subgroup.shuffle %sum_x2, %xor4, %shuffle_width : f32, i32, i32 + %group_sum = scalar.addf %sum_x2, %sum_x4_peer : f32 + scf.if %is_sum_leader { + %sum_value = scalar.fptosi %group_sum : f32 to i32 + %sum_index = index.add %group32_base, %group_half : index + view.store %sum_value, %i4_sums_view[%sum_index] : i32, view<[%group32_count]xi32> + } + } + } + } + kernel.return +} + +// Publishes RMSNorm, learned scale, SiLU gate product, and a signed-I4/K64 alternate. +kernel.def target(@ggml_rmsnorm_binary_gfx11_wave32) export("ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32") @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32(%token_count: index) { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%token_count, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %raw_gate: buffer, %output: buffer, %i4_qs: buffer, %i4_ds: buffer, %i4_sums: buffer) where [range(%token_count, 1, 1048576)] { + %hidden_size0 = config.get @ggml.rmsnorm_gate_silu_mul_symmetric_i4.hidden_size : index + %epsilon = config.get @ggml.rmsnorm_gate_silu_mul_symmetric_i4.rms_epsilon : f32 + %hidden_size = index.assume %hidden_size0 [range(%hidden_size0, 128, 32768), mul(%hidden_size0, 128)] : index + %token0 = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c0_offset = index.constant 0 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %amax_epsilon = scalar.constant 1.0000000000000001e-30 : f32 + %one_seventh = scalar.constant 0.14285714285714285 : f32 + %seven = scalar.constant 7.0 : f32 + %xor1 = scalar.constant 1 : i32 + %xor2 = scalar.constant 2 : i32 + %xor4 = scalar.constant 4 : i32 + %xor8 = scalar.constant 8 : i32 + %shuffle_width = scalar.constant 32 : i32 + %shift4 = scalar.constant 4 : i32 + %shift8 = scalar.constant 8 : i32 + %shift12 = scalar.constant 12 : i32 + %nibble_mask = vector.constant 15 : vector<4xi32> + + %input_noalias, %weight_noalias, %raw_gate_noalias, %output_noalias, %i4_qs_noalias, %i4_ds_noalias, %i4_sums_noalias = buffer.assume.noalias %input, %weight, %raw_gate, %output, %i4_qs, %i4_ds, %i4_sums : buffer, buffer, buffer, buffer, buffer, buffer, buffer + %row_scale, %token, %bounded_token_count, %valid_token = template.apply<@ggml.rmsnorm_f32.row_scale>(%token_count, %token0, %hidden_size, %hidden_size, %epsilon, %input_noalias) : (index, index, index, index, f32, buffer) -> (f32, index, index, i1) + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%hidden_size]xf32> + %weight_view = buffer.view %weight_noalias[%c0_offset] : buffer -> view<[%hidden_size]xf32> + %raw_gate_view = buffer.view %raw_gate_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%hidden_size]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%hidden_size]xf32> + %element_count = index.mul %bounded_token_count, %hidden_size : index + %group32_count = index.div %element_count, %c32 : index + %group64_count_per_row = index.div %hidden_size, %c64 : index + %qs_halfword_count = index.div %element_count, %c4 : index + %i4_qs_view = buffer.view %i4_qs_noalias[%c0_offset] : buffer -> view<[%qs_halfword_count]xi16> + %i4_ds_view = buffer.view %i4_ds_noalias[%c0_offset] : buffer -> view<[%group32_count]xf32> + %i4_sums_view = buffer.view %i4_sums_noalias[%c0_offset] : buffer -> view<[%group32_count]xi32> + + %wave = index.div %workitem, %c32 : index + %group_in_wave = index.div %lane, %c16 : index + %lane_in_group = index.rem %lane, %c16 : index + %group_half = index.div %lane_in_group, %c8 : index + %lane_in_half = index.rem %lane_in_group, %c8 : index + %is_scale_leader = index.cmp eq, %lane_in_group, %c0 : index + %is_sum_leader = index.cmp eq, %lane_in_half, %c0 : index + %wave_group_base = index.mul %wave, %c2 : index + %local_group = index.add %wave_group_base, %group_in_wave : index + %token_group_base = index.mul %token, %group64_count_per_row : index + scf.if %valid_token { + scf.for %group_base = [%c0 to %group64_count_per_row step %c16] { + %group_in_row = index.add %group_base, %local_group : index + %in_range = index.cmp ult, %group_in_row, %group64_count_per_row : index + scf.if %in_range { + %channel_base = index.mul %group_in_row, %c64 : index + %lane_channel = index.mul %lane_in_group, %c4 : index + %channel = index.add %channel_base, %lane_channel : index + %values = vector.load %input_view[%token, %channel] : view<[%bounded_token_count]x[%hidden_size]xf32> -> vector<4xf32> + %learned_weights = vector.load %weight_view[%channel] : view<[%hidden_size]xf32> -> vector<4xf32> + %raw_gate_values = vector.load %raw_gate_view[%token, %channel] : view<[%bounded_token_count]x[%hidden_size]xf32> -> vector<4xf32> + %scale_vector = vector.splat %row_scale : vector<4xf32> + %normalized = vector.mulf %values, %scale_vector : vector<4xf32> + %side = vector.mulf %normalized, %learned_weights : vector<4xf32> + %activated = vector.siluf %raw_gate_values : vector<4xf32> + %result = vector.mulf %side, %activated : vector<4xf32> + vector.store %result, %output_view[%token, %channel] : vector<4xf32>, view<[%bounded_token_count]x[%hidden_size]xf32> + + %absolute_values = vector.absf %result : vector<4xf32> + %lane_max = vector.reduce %absolute_values, %c0_f32 : vector<4xf32>, f32 + %max_x1_peer, %max_x1_valid = kernel.subgroup.shuffle %lane_max, %xor1, %shuffle_width : f32, i32, i32 + %max_x1 = scalar.maxnumf %lane_max, %max_x1_peer : f32 + %max_x2_peer, %max_x2_valid = kernel.subgroup.shuffle %max_x1, %xor2, %shuffle_width : f32, i32, i32 + %max_x2 = scalar.maxnumf %max_x1, %max_x2_peer : f32 + %max_x4_peer, %max_x4_valid = kernel.subgroup.shuffle %max_x2, %xor4, %shuffle_width : f32, i32, i32 + %max_x4 = scalar.maxnumf %max_x2, %max_x4_peer : f32 + %max_x8_peer, %max_x8_valid = kernel.subgroup.shuffle %max_x4, %xor8, %shuffle_width : f32, i32, i32 + %group_max = scalar.maxnumf %max_x4, %max_x8_peer : f32 + %amax = scalar.maxnumf %group_max, %amax_epsilon : f32 + %activation_scale = scalar.mulf %amax, %one_seventh : f32 + %activation_rscale = scalar.divf %seven, %amax : f32 + %group64 = index.add %token_group_base, %group_in_row : index + %group32_base = index.mul %group64, %c2 : index + %group32_high = index.add %group32_base, %c1 : index + scf.if %is_scale_leader { + view.store %activation_scale, %i4_ds_view[%group32_base] : f32, view<[%group32_count]xf32> + view.store %activation_scale, %i4_ds_view[%group32_high] : f32, view<[%group32_count]xf32> + } + + %rscale_vector = vector.splat %activation_rscale : vector<4xf32> + %scaled = vector.mulf %result, %rscale_vector : vector<4xf32> + %rounded = vector.roundf %scaled : vector<4xf32> + %quantized_i32 = vector.fptosi %rounded : vector<4xf32> to vector<4xi32> + %quantized_nibbles = vector.andi %quantized_i32, %nibble_mask : vector<4xi32> + %q0 = vector.extract %quantized_nibbles[0] : vector<4xi32> -> i32 + %q1 = vector.extract %quantized_nibbles[1] : vector<4xi32> -> i32 + %q2 = vector.extract %quantized_nibbles[2] : vector<4xi32> -> i32 + %q3 = vector.extract %quantized_nibbles[3] : vector<4xi32> -> i32 + %q1_shifted = scalar.shli %q1, %shift4 : i32 + %q2_shifted = scalar.shli %q2, %shift8 : i32 + %q3_shifted = scalar.shli %q3, %shift12 : i32 + %q01 = scalar.ori %q0, %q1_shifted : i32 + %q23 = scalar.ori %q2_shifted, %q3_shifted : i32 + %packed_i32 = scalar.ori %q01, %q23 : i32 + %packed_i16 = scalar.trunci %packed_i32 : i32 to i16 + %word_base = index.mul %group64, %c16 : index + %word_index = index.add %word_base, %lane_in_group : index + view.store %packed_i16, %i4_qs_view[%word_index] : i16, view<[%qs_halfword_count]xi16> + + %lane_sum = vector.reduce %rounded, %c0_f32 : vector<4xf32>, f32 + %sum_x1_peer, %sum_x1_valid = kernel.subgroup.shuffle %lane_sum, %xor1, %shuffle_width : f32, i32, i32 + %sum_x1 = scalar.addf %lane_sum, %sum_x1_peer : f32 + %sum_x2_peer, %sum_x2_valid = kernel.subgroup.shuffle %sum_x1, %xor2, %shuffle_width : f32, i32, i32 + %sum_x2 = scalar.addf %sum_x1, %sum_x2_peer : f32 + %sum_x4_peer, %sum_x4_valid = kernel.subgroup.shuffle %sum_x2, %xor4, %shuffle_width : f32, i32, i32 + %group_sum = scalar.addf %sum_x2, %sum_x4_peer : f32 + scf.if %is_sum_leader { + %sum_value = scalar.fptosi %group_sum : f32 to i32 + %sum_index = index.add %group32_base, %group_half : index + view.store %sum_value, %i4_sums_view[%sum_index] : i32, view<[%group32_count]xi32> + } + } + } + } + kernel.return +} + +template.decl @ggml.rmsnorm_f32.subgroup_row_scale(%token_count: index, %token: index, %hidden_size: index, %epsilon: f32, %input: buffer) -> (f32) + +template.decl @ggml.unary_f32.apply_vector4(%op: index, %values: vector<4xf32>) -> (vector<4xf32>) + +config.decl @ggml.rmsnorm_gate_f32.hidden_size : %value: index where [range(%value, 128, 1024), mul(%value, 128)] +config.decl @ggml.rmsnorm_gate_f32.rms_epsilon : f32 +config.decl @ggml.rmsnorm_gate_f32.gate_op : %value: index where [range(%value, 0, 23)] +config.decl @ggml.rmsnorm_gate_f32.f16_output_row_width : %value: index where [range(%value, 0, 32768), mul(%value, 128)] +template.decl @ggml.quantize_q8_1_x4.publish_vector4_strict(%publish_word: i1, %token_output_byte_base: offset, %channel: index, %values: vector<4xf32>, %scratch_values: view<256xf32>, %scratch_d: view<32xf32>, %output: buffer) + +template.decl @ggml.rmsnorm_gate_f32.publish_body(%publish_q8: i1, %token_count: index, %input: buffer, %weight: buffer, %raw_gate: buffer, %output: buffer, %alternate_output: buffer) + +kernel.def target(@ggml_rmsnorm_binary_gfx11_wave32) export("ggml_rmsnorm_gate_f32_f16") @ggml_rmsnorm_gate_f32_f16(%token_count: index) { + %one = index.constant 1 : index + %eight = index.constant 8 : index + %seven = index.constant 7 : index + %threads = index.constant 256 : index + %round = index.add %token_count, %seven : index + %groups = index.div %round, %eight : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%threads, %one, %one) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %raw_gate: buffer, %output: buffer, %f16_output: buffer) { + %publish_q8 = scalar.constant false : i1 + template.apply<@ggml.rmsnorm_gate_f32.publish_body>(%publish_q8, %token_count, %input, %weight, %raw_gate, %output, %f16_output) : (i1, index, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +kernel.def target(@ggml_rmsnorm_binary_gfx11_wave32) export("ggml_rmsnorm_gate_f32_q8_1_x4") @ggml_rmsnorm_gate_f32_q8_1_x4(%token_count: index) { + %one = index.constant 1 : index + %eight = index.constant 8 : index + %seven = index.constant 7 : index + %threads = index.constant 256 : index + %round = index.add %token_count, %seven : index + %groups = index.div %round, %eight : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%threads, %one, %one) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %raw_gate: buffer, %output: buffer, %q8_output: buffer) { + %publish_q8 = scalar.constant true : i1 + template.apply<@ggml.rmsnorm_gate_f32.publish_body>(%publish_q8, %token_count, %input, %weight, %raw_gate, %output, %q8_output) : (i1, index, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +template.def<@ggml.rmsnorm_gate_f32.publish_body> device @ggml_rmsnorm_gate_f32_publish_body(%publish_q8: i1, %token_count: index, %input: buffer, %weight: buffer, %raw_gate: buffer, %output: buffer, %alternate_output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 1048576)] : index + %hidden0 = config.get @ggml.rmsnorm_gate_f32.hidden_size : index + %hidden = index.assume %hidden0 [range(%hidden0, 128, 1024), mul(%hidden0, 128)] : index + %epsilon = config.get @ggml.rmsnorm_gate_f32.rms_epsilon : f32 + %gate_op = config.get @ggml.rmsnorm_gate_f32.gate_op : index + %row_width = config.get @ggml.rmsnorm_gate_f32.f16_output_row_width : index + %z = index.constant 0 : index + %zero_offset = index.constant 0 : offset + %four = index.constant 4 : index + %eight = index.constant 8 : index + %c128 = index.constant 128 : index + %wg = kernel.workgroup.id : index + %wave = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %row0 = index.madd %wg, %eight, %wave : index + %valid = index.cmp ult, %row0, %bounded_token_count : index + %safe = scf.select %valid, %row0, %z : index + %row = index.assume %safe [lt(%safe, %bounded_token_count)] : index + %channel0 = index.mul %lane, %four : index + %input_n, %weight_n, %raw_gate_n, %output_n, %alternate_n = buffer.assume.noalias %input, %weight, %raw_gate, %output, %alternate_output : buffer, buffer, buffer, buffer, buffer + %iv = buffer.view %input_n[%zero_offset] : buffer -> view<[%bounded_token_count]x[%hidden]xf32> + %wv = buffer.view %weight_n[%zero_offset] : buffer -> view<[%hidden]xf32> + %gv = buffer.view %raw_gate_n[%zero_offset] : buffer -> view<[%bounded_token_count]x[%hidden]xf32> + %ov = buffer.view %output_n[%zero_offset] : buffer -> view<[%bounded_token_count]x[%hidden]xf32> + %scale = template.apply<@ggml.rmsnorm_f32.subgroup_row_scale>(%bounded_token_count, %row, %hidden, %epsilon, %input_n) : (index, index, index, f32, buffer) -> (f32) + %scale_v = vector.splat %scale : vector<4xf32> + %last_channel = index.sub %hidden, %four : index + scf.if %valid { + scf.for %channel = [%channel0 to %hidden step %c128] { + %bounded_channel = index.assume %channel [range(%channel, 0, 1020), mul(%channel, 4), le(%channel, %last_channel)] : index + %x = vector.load %iv[%row, %bounded_channel] : view<[%bounded_token_count]x[%hidden]xf32> -> vector<4xf32> + %w = vector.load %wv[%bounded_channel] : view<[%hidden]xf32> -> vector<4xf32> + %g = vector.load %gv[%row, %bounded_channel] : view<[%bounded_token_count]x[%hidden]xf32> -> vector<4xf32> + %norm = vector.mulf %x, %scale_v : vector<4xf32> + %side = vector.mulf %norm, %w : vector<4xf32> + %activated = template.apply<@ggml.unary_f32.apply_vector4>(%gate_op, %g) : (index, vector<4xf32>) -> (vector<4xf32>) + %out = vector.mulf %side, %activated : vector<4xf32> + vector.store %out, %ov[%row, %bounded_channel] : vector<4xf32>, view<[%bounded_token_count]x[%hidden]xf32> + scf.if %publish_q8 { + %scratch_bytes = index.constant 1152 : offset + %scratch_d_offset = index.constant 1024 : offset + %scratch = buffer.alloca align(16) %scratch_bytes : buffer + %scratch_values = buffer.view %scratch[%zero_offset] : buffer -> view<256xf32> + %scratch_d = buffer.view %scratch[%scratch_d_offset] : buffer -> view<32xf32> + %flat_channel = index.madd %row, %hidden, %bounded_channel : index + template.apply<@ggml.quantize_q8_1_x4.publish_vector4_strict>(%valid, %zero_offset, %flat_channel, %out, %scratch_values, %scratch_d, %alternate_n) : (i1, offset, index, vector<4xf32>, view<256xf32>, view<32xf32>, buffer) + } else { + %hv = buffer.view %alternate_n[%zero_offset] : buffer -> view<[%bounded_token_count]x[%hidden]xf16> + %out_half = vector.fptrunc %out : vector<4xf32> to vector<4xf16> + %packed = index.cmp ugt, %row_width, %z : index + scf.if %packed { + %c16 = index.constant 16 : index + %element_count = index.mul %bounded_token_count, %hidden : index + %packed_rows = index.div %element_count, %row_width : index + %packed_columns = index.div %row_width, %c16 : index + %packed_view = buffer.view %alternate_n[%zero_offset] : buffer -> view<[%packed_columns]x[%packed_rows]x16xf16> + %flat = index.madd %row, %hidden, %bounded_channel : index + %packed_row = index.div %flat, %row_width : index + %packed_channel = index.rem %flat, %row_width : index + %packed_column = index.div %packed_channel, %c16 : index + %packed_lane = index.rem %packed_channel, %c16 : index + vector.store %out_half, %packed_view[%packed_column, %packed_row, %packed_lane] : vector<4xf16>, view<[%packed_columns]x[%packed_rows]x16xf16> + } else { + vector.store %out_half, %hv[%row, %bounded_channel] : vector<4xf16>, view<[%bounded_token_count]x[%hidden]xf16> + } + } + } + } + template.return +} + +check.case public @ggml_rmsnorm_binary_f32_mul_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(2.0) : tensor<128xf32> + %rhs = check.generate.iota offset(-1.0) step(0.015625) : tensor<128xf32> + %output = check.generate.fill value(0.0) : tensor<128xf32> + %expected = check.generate.iota offset(-1.0) step(0.015625) : tensor<128xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<128xf32>, tensor<128xf32>, tensor<128xf32>) + check.expect.close actual(%output) expected(%expected) atol(9.9999999999999995e-07) rtol(9.9999999999999995e-07) nan(same) : tensor<128xf32> + check.return +} + +check.case public @ggml_rmsnorm_binary_f32_add_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(2.0) : tensor<128xf32> + %rhs = check.generate.iota offset(0.0) step(0.015625) : tensor<128xf32> + %output = check.generate.fill value(0.0) : tensor<128xf32> + %expected = check.generate.iota offset(1.0) step(0.015625) : tensor<128xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<128xf32>, tensor<128xf32>, tensor<128xf32>) + check.expect.close actual(%output) expected(%expected) atol(9.9999999999999995e-07) rtol(9.9999999999999995e-07) nan(same) : tensor<128xf32> + check.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/rmsnorm_binary_q8_1_x4.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/rmsnorm_binary_q8_1_x4.loom new file mode 100644 index 000000000000..a4b44f79bda1 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/rmsnorm_binary_q8_1_x4.loom @@ -0,0 +1,121 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +template.decl @ggml.binary_f32.apply(%arg0: index, %arg1: f32, %arg2: f32) -> (f32) + +template.decl @ggml.quantize_q8_1_x4.publish_vector4(%arg0: i1, %arg1: offset, %arg2: index, %arg3: vector<4xf32>, %arg4: view<256xf32>, %arg5: view<32xf32>, %arg6: buffer) + +template.decl @ggml.rmsnorm_f32.row_scale(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: f32, %arg5: buffer) -> (f32, index, index, i1) + +amdgpu.target @ggml_rmsnorm_binary_q8_1_x4_gfx11_wave32 {subgroup_size = 32} + +config.decl @ggml.rmsnorm_binary_q8_1_x4.hidden_size : %value: index where [range(%value, 128, 32768), mul(%value, 128)] + +config.decl @ggml.rmsnorm_binary_q8_1_x4.rms_epsilon : f32 + +config.decl @ggml.rmsnorm_binary_q8_1_x4.op : %value: index where [range(%value, 0, 3)] + +kernel.decl @ggml_rmsnorm_binary_f32(%token_count$23: index) launch(%token_count$24: index, %input: buffer, %rhs: buffer, %output: buffer) + +kernel.decl @ggml_quantize_q8_1_x4_f32(%token_count$28: index, %input_size$29: index) launch(%token_count$30: index, %input_size$31: index, %input: buffer, %output: buffer) + +kernel.def target(@ggml_rmsnorm_binary_q8_1_x4_gfx11_wave32) export("ggml_rmsnorm_binary_q8_1_x4") @ggml_rmsnorm_binary_q8_1_x4(%token_count: index) { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + %hidden = config.get @ggml.rmsnorm_binary_q8_1_x4.hidden_size : index + %tile_width = index.constant 1024 : index + %tile_round = index.constant 1023 : index + %decode_limit = index.constant 6 : index + %rounded_hidden = index.add %hidden, %tile_round : index + %stripe_count = index.div %rounded_hidden, %tile_width : index + %is_decode = index.cmp ult, %token_count, %decode_limit : index + %groups_per_row = scf.select %is_decode, %stripe_count, %c1 : index + %groups = index.mul %token_count, %groups_per_row : index + kernel.launch.config workgroups(%groups, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %rhs: buffer, %output: buffer, %q8_output: buffer) where [range(%token_count, 1, 2048)] { + %hidden_size0 = config.get @ggml.rmsnorm_binary_q8_1_x4.hidden_size : index + %epsilon = config.get @ggml.rmsnorm_binary_q8_1_x4.rms_epsilon : f32 + %op = config.get @ggml.rmsnorm_binary_q8_1_x4.op : index + %hidden_size = index.assume %hidden_size0 [range(%hidden_size0, 128, 32768), mul(%hidden_size0, 128)] : index + %dispatch_group = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %c1024 = index.constant 1024 : index + %tile_round = index.constant 1023 : index + %decode_limit = index.constant 6 : index + %rounded_hidden = index.add %hidden_size, %tile_round : index + %stripe_count = index.div %rounded_hidden, %c1024 : index + %is_decode = index.cmp ult, %token_count, %decode_limit : index + %groups_per_row = scf.select %is_decode, %stripe_count, %c1 : index + %token0 = index.div %dispatch_group, %groups_per_row : index + %stripe_id = index.rem %dispatch_group, %groups_per_row : index + %stripe_start = index.mul %stripe_id, %c1024 : index + %stripe_step = index.mul %groups_per_row, %c1024 : index + %group_bytes = index.constant 144 : offset + %scratch_d_byte_add = index.constant 1024 : offset + %scratch_bytes = index.constant 1152 : offset + %c0_offset = index.constant 0 : offset + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %input_noalias, %rhs_noalias, %output_noalias, %q8_output_noalias = buffer.assume.noalias %input, %rhs, %output, %q8_output : buffer, buffer, buffer, buffer + %row_scale, %token, %launch_token_count, %valid_token = template.apply<@ggml.rmsnorm_f32.row_scale>(%token_count, %token0, %hidden_size, %hidden_size, %epsilon, %input_noalias) : (index, index, index, index, f32, buffer) -> (f32, index, index, i1) + %row_scale_vector = vector.splat %row_scale : vector<4xf32> + %physical_group_count = index.div %hidden_size, %c128 : index + %row_bytes = index.scale %physical_group_count, %group_bytes : index, offset -> offset + %token_output_byte_base = index.scale %token, %row_bytes : index, offset -> offset + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%hidden_size]xf32> + %rhs_view = buffer.view %rhs_noalias[%c0_offset] : buffer -> view<[%hidden_size]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%hidden_size]xf32> + %scratch = buffer.alloca align(16) %scratch_bytes : buffer + %scratch_values = buffer.view %scratch[%c0_offset] : buffer -> view<256xf32> + %scratch_d = buffer.view %scratch[%scratch_d_byte_add] : buffer -> view<32xf32> + scf.for %stripe_base = [%stripe_start to %hidden_size step %stripe_step] { + %word_element_add = index.mul %workitem, %c4 : index + %channel = index.add %stripe_base, %word_element_add : index + %valid_word0 = index.cmp ult, %channel, %hidden_size : index + %valid_word = scalar.andi %valid_word0, %valid_token : i1 + %mask = vector.mask.range [%channel to %hidden_size step %c1] : index -> vector<4xi1> + %input_values = vector.load.mask %input_view[%token, %channel], %mask, %c0_f32x4 : view<[%launch_token_count]x[%hidden_size]xf32>, vector<4xi1>, vector<4xf32> + %rhs_values = vector.load.mask %rhs_view[%channel], %mask, %c0_f32x4 : view<[%hidden_size]xf32>, vector<4xi1>, vector<4xf32> + %normalized = vector.mulf %input_values, %row_scale_vector : vector<4xf32> + %result0 = vector.extract %normalized[0] : vector<4xf32> -> f32 + %result1 = vector.extract %normalized[1] : vector<4xf32> -> f32 + %result2 = vector.extract %normalized[2] : vector<4xf32> -> f32 + %result3 = vector.extract %normalized[3] : vector<4xf32> -> f32 + %rhs0 = vector.extract %rhs_values[0] : vector<4xf32> -> f32 + %rhs1 = vector.extract %rhs_values[1] : vector<4xf32> -> f32 + %rhs2 = vector.extract %rhs_values[2] : vector<4xf32> -> f32 + %rhs3 = vector.extract %rhs_values[3] : vector<4xf32> -> f32 + %binary0 = template.apply<@ggml.binary_f32.apply>(%op, %result0, %rhs0) : (index, f32, f32) -> (f32) + %binary1 = template.apply<@ggml.binary_f32.apply>(%op, %result1, %rhs1) : (index, f32, f32) -> (f32) + %binary2 = template.apply<@ggml.binary_f32.apply>(%op, %result2, %rhs2) : (index, f32, f32) -> (f32) + %binary3 = template.apply<@ggml.binary_f32.apply>(%op, %result3, %rhs3) : (index, f32, f32) -> (f32) + %packed0 = vector.insert %binary0 into %c0_f32x4[0] : f32, vector<4xf32> + %packed1 = vector.insert %binary1 into %packed0[1] : f32, vector<4xf32> + %packed2 = vector.insert %binary2 into %packed1[2] : f32, vector<4xf32> + %packed = vector.insert %binary3 into %packed2[3] : f32, vector<4xf32> + vector.store.mask %packed, %output_view[%token, %channel], %mask : vector<4xf32>, view<[%launch_token_count]x[%hidden_size]xf32>, vector<4xi1> + template.apply<@ggml.quantize_q8_1_x4.publish_vector4>(%valid_word, %token_output_byte_base, %channel, %packed, %scratch_values, %scratch_d, %q8_output_noalias) : (i1, offset, index, vector<4xf32>, view<256xf32>, view<32xf32>, buffer) + } + kernel.return +} + +check.case public @ggml_rmsnorm_binary_q8_1_x4_mul_case { + %token_count = check.literal value(1) : index + %hidden_size = check.literal value(128) : index + %input = check.generate.fill value(2.0) : tensor<128xf32> + %rhs = check.generate.iota offset(-1.0) step(0.015625) : tensor<128xf32> + %f32 = check.generate.fill value(0.0) : tensor<128xf32> + %expected = check.generate.fill value(0) : tensor<144xi8> + %actual = check.generate.fill value(1) : tensor<144xi8> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %f32) : [index](index, tensor<128xf32>, tensor<128xf32>, tensor<128xf32>) + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %hidden_size](%token_count, %hidden_size, %f32, %expected) : [index, index](index, index, tensor<128xf32>, tensor<144xi8>) + %actual_f32 = check.generate.fill value(0.0) : tensor<128xf32> + kernel.launch @ggml_rmsnorm_binary_q8_1_x4[%token_count](%token_count, %input, %rhs, %actual_f32, %actual) : [index](index, tensor<128xf32>, tensor<128xf32>, tensor<128xf32>, tensor<144xi8>) + check.expect.equal actual(%actual_f32) expected(%f32) : tensor<128xf32> + check.expect.equal actual(%actual) expected(%expected) : tensor<144xi8> + check.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/rmsnorm_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/rmsnorm_f32.loom new file mode 100644 index 000000000000..38803798b7c8 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/rmsnorm_f32.loom @@ -0,0 +1,284 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +template.decl @ggml.rmsnorm_f32.apply(%arg0: f32, %arg1: f32) -> (f32) + +template.decl @ggml.rmsnorm_f32.row_scale(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: f32, %arg5: buffer) -> (f32, index, index, i1) + +amdgpu.target @ggml_rmsnorm_gfx11_wave32 {subgroup_size = 32} + +config.decl @ggml.rmsnorm_f32.hidden_size : %value: index where [range(%value, 64, 32768), mul(%value, 64)] + +config.decl @ggml.rmsnorm_f32.input_stride : %value: index where [range(%value, 64, 1048576)] + +config.decl @ggml.rmsnorm_f32.rms_epsilon : f32 + +kernel.def target(@ggml_rmsnorm_gfx11_wave32) export("ggml_rmsnorm_f32") @ggml_rmsnorm_f32(%token_count: index) { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%token_count, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %hidden_size0 = config.get @ggml.rmsnorm_f32.hidden_size : index + %input_stride0 = config.get @ggml.rmsnorm_f32.input_stride : index + %epsilon = config.get @ggml.rmsnorm_f32.rms_epsilon : f32 + %hidden_size1 = index.assume %hidden_size0 [range(%hidden_size0, 64, 32768), mul(%hidden_size0, 64)] : index + %input_stride1 = index.assume %input_stride0 [range(%input_stride0, 64, 1048576)] : index + %hidden_size, %input_stride = index.assume %hidden_size1, %input_stride1 [le(%hidden_size1, %input_stride1)] : index, index + %token0 = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %input_noalias, %output_noalias = buffer.assume.noalias %input, %output : buffer, buffer + %row_scale, %token, %launch_token_count, %valid_token = template.apply<@ggml.rmsnorm_f32.row_scale>(%token_count, %token0, %hidden_size, %input_stride, %epsilon, %input_noalias) : (index, index, index, index, f32, buffer) -> (f32, index, index, i1) + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%input_stride]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%hidden_size]xf32> + scf.if %valid_token { + scf.for %channel = [%workitem to %hidden_size step %c256] { + %value = view.load %input_view[%token, %channel] : view<[%launch_token_count]x[%input_stride]xf32> -> f32 + %normalized = template.apply<@ggml.rmsnorm_f32.apply>(%value, %row_scale) : (f32, f32) -> (f32) + view.store %normalized, %output_view[%token, %channel] : f32, view<[%launch_token_count]x[%hidden_size]xf32> + } + } + kernel.return +} + +check.case public @ggml_rmsnorm_f32_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(2.0) : tensor<128xf32> + %output = check.generate.fill value(0.0) : tensor<128xf32> + %expected = check.generate.fill value(1.0) : tensor<128xf32> + kernel.launch @ggml_rmsnorm_f32[%token_count](%token_count, %input, %output) : [index](index, tensor<128xf32>, tensor<128xf32>) + check.expect.close actual(%output) expected(%expected) atol(9.9999999999999995e-07) rtol(9.9999999999999995e-07) nan(same) : tensor<128xf32> + check.return +} +// One workgroup preserves each strided RMSNorm row before scale and RoPE. +config.decl @ggml.rmsnorm_mul_rope.hidden_size : %value: index where [range(%value, 2, 512)] + +config.decl @ggml.rmsnorm_mul_rope.ne1 : %value: index where [range(%value, 1, 1048576)] + +config.decl @ggml.rmsnorm_mul_rope.ne2 : %value: index where [range(%value, 1, 1048576)] + +config.decl @ggml.rmsnorm_mul_rope.ne3 : %value: index where [range(%value, 1, 1024)] + +config.decl @ggml.rmsnorm_mul_rope.input_stride1 : %value: index where [range(%value, 0, 1073741823)] + +config.decl @ggml.rmsnorm_mul_rope.input_stride2 : %value: index where [range(%value, 0, 1073741823)] + +config.decl @ggml.rmsnorm_mul_rope.input_stride3 : %value: index where [range(%value, 0, 1073741823)] + +config.decl @ggml.rmsnorm_mul_rope.output_stride1 : %value: index where [range(%value, 0, 268435456)] + +config.decl @ggml.rmsnorm_mul_rope.output_stride2 : %value: index where [range(%value, 0, 1073741823)] + +config.decl @ggml.rmsnorm_mul_rope.output_stride3 : %value: index where [range(%value, 0, 1073741823)] + +config.decl @ggml.rmsnorm_mul_rope.n_dims : %value: index where [range(%value, 2, 512)] + +config.decl @ggml.rmsnorm_mul_rope.section0 : %value: index where [range(%value, 0, 65536)] + +config.decl @ggml.rmsnorm_mul_rope.section1 : %value: index where [range(%value, 0, 65536)] + +config.decl @ggml.rmsnorm_mul_rope.section2 : %value: index where [range(%value, 0, 65536)] + +config.decl @ggml.rmsnorm_mul_rope.section3 : %value: index where [range(%value, 0, 65536)] + +config.decl @ggml.rmsnorm_mul_rope.mode : %value: index where [range(%value, 8, 40)] + +config.decl @ggml.rmsnorm_mul_rope.workgroup_size : %value: index where [range(%value, 32, 1024), mul(%value, 32)] + +config.decl @ggml.rmsnorm_mul_rope.epsilon : f32 + +config.decl @ggml.rmsnorm_mul_rope.freq_base : f32 + +config.decl @ggml.rmsnorm_mul_rope.freq_scale : f32 + +config.decl @ggml.rmsnorm_mul_rope.attn_factor : f32 + +kernel.def target(@ggml_rmsnorm_gfx11_wave32) export("ggml_rmsnorm_mul_rope_f32") @ggml_rmsnorm_mul_rope_f32() { + %unit = index.constant 1 : index + %ne1 = config.get @ggml.rmsnorm_mul_rope.ne1 : index + %ne2 = config.get @ggml.rmsnorm_mul_rope.ne2 : index + %ne3 = config.get @ggml.rmsnorm_mul_rope.ne3 : index + %wg = config.get @ggml.rmsnorm_mul_rope.workgroup_size : index + %n12 = index.mul %ne1, %ne2 : index + %rows = index.mul %n12, %ne3 : index + kernel.launch.config workgroups(%rows, %unit, %unit) workgroup_size(%wg, %unit, %unit) : index +} launch(%src0: buffer, %weight: buffer, %pos: buffer, %dst: buffer) { + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c40 = index.constant 40 : index + %c0_f32 = scalar.constant 0.0 : f32 + %cneg2_f32 = scalar.constant -2.0 : f32 + %eps = config.get @ggml.rmsnorm_mul_rope.epsilon : f32 + %freq_base = config.get @ggml.rmsnorm_mul_rope.freq_base : f32 + %freq_scale = config.get @ggml.rmsnorm_mul_rope.freq_scale : f32 + %attn_factor = config.get @ggml.rmsnorm_mul_rope.attn_factor : f32 + + %ncols = config.get @ggml.rmsnorm_mul_rope.hidden_size : index + %ne1 = config.get @ggml.rmsnorm_mul_rope.ne1 : index + %ne2 = config.get @ggml.rmsnorm_mul_rope.ne2 : index + %s1 = config.get @ggml.rmsnorm_mul_rope.input_stride1 : index + %s2 = config.get @ggml.rmsnorm_mul_rope.input_stride2 : index + %s3 = config.get @ggml.rmsnorm_mul_rope.input_stride3 : index + %d1 = config.get @ggml.rmsnorm_mul_rope.output_stride1 : index + %d2 = config.get @ggml.rmsnorm_mul_rope.output_stride2 : index + %d3 = config.get @ggml.rmsnorm_mul_rope.output_stride3 : index + %n_dims = config.get @ggml.rmsnorm_mul_rope.n_dims : index + %sec0 = config.get @ggml.rmsnorm_mul_rope.section0 : index + %sec1 = config.get @ggml.rmsnorm_mul_rope.section1 : index + %sec2 = config.get @ggml.rmsnorm_mul_rope.section2 : index + %sec3 = config.get @ggml.rmsnorm_mul_rope.section3 : index + %mode = config.get @ggml.rmsnorm_mul_rope.mode : index + %wg = config.get @ggml.rmsnorm_mul_rope.workgroup_size : index + + %ncols_b = index.assume %ncols [range(%ncols, 2, 512)] : index + %ne1_b = index.assume %ne1 [range(%ne1, 1, 1048576)] : index + %ne2_b = index.assume %ne2 [range(%ne2, 1, 1048576)] : index + %n_dims_b = index.assume %n_dims [range(%n_dims, 2, 512)] : index + + %row0 = kernel.workgroup.id : index + %row = index.assume %row0 [range(%row0, 0, 1073741823)] : index + %lane0 = kernel.workitem.id : index + %lane = index.assume %lane0 [range(%lane0, 0, 1023)] : index + + %src_global = buffer.assume.memory_space %src0 : buffer + %weight_global = buffer.assume.memory_space %weight : buffer + %pos_global = buffer.assume.memory_space %pos : buffer + %dst_global = buffer.assume.memory_space %dst : buffer + %src_na, %weight_na, %pos_na, %dst_na = buffer.assume.noalias %src_global, %weight_global, %pos_global, %dst_global : buffer, buffer, buffer, buffer + %src_view = buffer.view %src_na[%base] : buffer -> view<1073741824xf32> + %weight_view = buffer.view %weight_na[%base] : buffer -> view<65536xf32> + %pos_view = buffer.view %pos_na[%base] : buffer -> view<1073741824xi32> + %dst_view = buffer.view %dst_na[%base] : buffer -> view<1073741824xf32> + + %n12 = index.mul %ne1_b, %ne2_b : index + %i3 = index.div %row, %n12 : index + %r12 = index.rem %row, %n12 : index + %i2 = index.div %r12, %ne1_b : index + %i1 = index.rem %r12, %ne1_b : index + %o1 = index.mul %i1, %s1 : index + %o2 = index.mul %i2, %s2 : index + %o3 = index.mul %i3, %s3 : index + %o12 = index.add %o1, %o2 : index + %src_base0 = index.add %o12, %o3 : index + %src_base = index.assume %src_base0 [range(%src_base0, 0, 1073741823)] : index + + %sum = scf.for %col = [%lane to %ncols_b step %wg](%acc = %c0_f32 : f32) -> (f32) { + %si0 = index.add %src_base, %col : index + %si = index.assume %si0 [range(%si0, 0, 1073741823)] : index + %value = view.load %src_view[%si] : view<1073741824xf32> -> f32 + %square = scalar.mulf %value, %value : f32 + %next = scalar.addf %acc, %square : f32 + scf.yield %next : f32 + } + %row_sum = kernel.workgroup.reduce %sum : f32 + %ncols_i32 = index.cast %ncols_b : index to i32 + %ncols_f32 = scalar.sitofp %ncols_i32 : i32 to f32 + %mean = scalar.divf %row_sum, %ncols_f32 : f32 + %biased = scalar.addf %mean, %eps : f32 + %rms_scale = scalar.rsqrtf %biased : f32 + + %half = index.div %ncols_b, %c2 : index + %active = index.cmp ult, %lane, %half : index + scf.if %active { + %ih = index.add %lane, %c0 : index + %i0 = index.mul %ih, %c2 : index + %half_dims = index.div %n_dims_b, %c2 : index + %rotate = index.cmp ult, %i0, %n_dims_b : index + + %sd01 = index.add %sec0, %sec1 : index + %sd012 = index.add %sd01, %sec2 : index + %sect_dims = index.add %sd012, %sec3 : index + %sector = index.rem %ih, %sect_dims : index + %m3 = index.rem %sector, %c3 : index + %s0lim = index.mul %sec0, %c3 : index + %s1lim = index.mul %sec1, %c3 : index + %s2lim = index.mul %sec2, %c3 : index + %is0 = index.cmp eq, %m3, %c0 : index + %is1 = index.cmp eq, %m3, %c1 : index + %is2 = index.cmp eq, %m3, %c2 : index + %ok0 = index.cmp ult, %sector, %s0lim : index + %ok1 = index.cmp ult, %sector, %s1lim : index + %ok2 = index.cmp ult, %sector, %s2lim : index + %im_p0 = scf.select %ok0, %c0, %c3 : index + %im_a = scf.select %is0, %im_p0, %c3 : index + %im_p2 = scf.select %ok2, %c2, %c3 : index + %im_b = scf.select %is2, %im_p2, %im_a : index + %im_p1 = scf.select %ok1, %c1, %c3 : index + %plane_im = scf.select %is1, %im_p1, %im_b : index + + %sec_w = index.add %sec0, %sec1 : index + %sec_w2 = index.add %sec_w, %sec2 : index + %lt0 = index.cmp ult, %sector, %sec0 : index + %ltw = index.cmp ult, %sector, %sec_w : index + %ltw2 = index.cmp ult, %sector, %sec_w2 : index + %mr_a = scf.select %ltw2, %c2, %c3 : index + %mr_b = scf.select %ltw, %c1, %mr_a : index + %plane_mr = scf.select %lt0, %c0, %mr_b : index + %is_im = index.cmp eq, %mode, %c40 : index + %plane = scf.select %is_im, %plane_im, %plane_mr : index + + %plane_off = index.mul %plane, %ne2_b : index + %pos_i0 = index.add %i2, %plane_off : index + %pos_i = index.assume %pos_i0 [range(%pos_i0, 0, 1073741823)] : index + %pos_i32 = view.load %pos_view[%pos_i] : view<1073741824xi32> -> i32 + %pos_f = scalar.sitofp %pos_i32 : i32 to f32 + %n_dims_i32 = index.cast %n_dims_b : index to i32 + %n_dims_f32 = scalar.sitofp %n_dims_i32 : i32 to f32 + %ih_i32 = index.cast %ih : index to i32 + %ih_f = scalar.sitofp %ih_i32 : i32 to f32 + %scaled_ih = scalar.mulf %cneg2_f32, %ih_f : f32 + %theta_exponent = scalar.divf %scaled_ih, %n_dims_f32 : f32 + %freq = scalar.powf %freq_base, %theta_exponent : f32 + %theta_base = scalar.mulf %pos_f, %freq : f32 + %theta = scalar.mulf %theta_base, %freq_scale : f32 + %cos_theta_raw = scalar.cosf %theta : f32 + %s0v = scalar.sinf %theta : f32 + %cos_theta = scalar.mulf %cos_theta_raw, %attn_factor : f32 + %sin_theta = scalar.mulf %s0v, %attn_factor : f32 + + %off_a = scf.select %rotate, %ih, %i0 : index + %rot_b = index.add %ih, %half_dims : index + %pass_b = index.add %i0, %c1 : index + %off_b = scf.select %rotate, %rot_b, %pass_b : index + %sa0 = index.add %src_base, %off_a : index + %sa = index.assume %sa0 [range(%sa0, 0, 1073741823)] : index + %sb0 = index.add %src_base, %off_b : index + %sb = index.assume %sb0 [range(%sb0, 0, 1073741823)] : index + + %raw0 = view.load %src_view[%sa] : view<1073741824xf32> -> f32 + %raw1 = view.load %src_view[%sb] : view<1073741824xf32> -> f32 + %rms0 = scalar.mulf %raw0, %rms_scale : f32 + %rms1 = scalar.mulf %raw1, %rms_scale : f32 + %w0 = view.load %weight_view[%off_a] : view<65536xf32> -> f32 + %w1 = view.load %weight_view[%off_b] : view<65536xf32> -> f32 + %x0 = scalar.mulf %rms0, %w0 : f32 + %x1 = scalar.mulf %rms1, %w1 : f32 + + %t0 = scalar.mulf %x0, %cos_theta : f32 + %t1 = scalar.mulf %x1, %sin_theta : f32 + %rot0 = scalar.subf %t0, %t1 : f32 + %t2 = scalar.mulf %x0, %sin_theta : f32 + %t3 = scalar.mulf %x1, %cos_theta : f32 + %rot1 = scalar.addf %t2, %t3 : f32 + %out0 = scf.select %rotate, %rot0, %x0 : f32 + %out1 = scf.select %rotate, %rot1, %x1 : f32 + + %db1 = index.mul %i1, %d1 : index + %db2 = index.mul %i2, %d2 : index + %db3 = index.mul %i3, %d3 : index + %db01 = index.add %db1, %db2 : index + %dst_row = index.add %db01, %db3 : index + %da0 = index.add %dst_row, %off_a : index + %da = index.assume %da0 [range(%da0, 0, 1073741823)] : index + %db0 = index.add %dst_row, %off_b : index + %db = index.assume %db0 [range(%db0, 0, 1073741823)] : index + view.store %out0, %dst_view[%da] : f32, view<1073741824xf32> + view.store %out1, %dst_view[%db] : f32, view<1073741824xf32> + } + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/rope_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/rope_f32.loom new file mode 100644 index 000000000000..78d54b38e7d1 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/rope_f32.loom @@ -0,0 +1,150 @@ +// Generic GGML ROPE for F32 head-major rows. +template.decl @ggml.rope_f32.body(%token_count: index, %positions: buffer, %input: buffer, %theta: buffer, %freq_factors: buffer, %output: buffer) + +amdgpu.target @ggml_rope_f32_gfx11_wave32 {subgroup_size = 32} + +config.decl @ggml.rope_f32.head_size : %value: index where [range(%value, 4, 1024), mul(%value, 4)] + +config.decl @ggml.rope_f32.n_dims : %value: index where [range(%value, 4, 1024), mul(%value, 4)] + +config.decl @ggml.rope_f32.head_count : %value: index where [range(%value, 1, 64)] + +config.decl @ggml.rope_f32.token_capacity : %value: index where [range(%value, 1, 2048)] + +config.decl @ggml.rope_f32.input_stride1 : %value: index where [range(%value, 4, 1073741824)] + +config.decl @ggml.rope_f32.input_stride2 : %value: index where [range(%value, 4, 1073741824)] + +config.decl @ggml.rope_f32.input_span : %value: index where [range(%value, 1, 1073741824)] + +config.decl @ggml.rope_f32.mscale : f32 + +config.decl @ggml.rope_f32.mode : %value: index where [range(%value, 0, 2)] + +func.decl @ggml_rope_f32_pair_packet(%position: f32, %theta: vector<2xf32>, %freq_factors: vector<2xf32>, %x_values: vector<2xf32>, %y_values: vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) + +template.def<@ggml.rope_f32.body> device @ggml_rope_f32_body(%token_count: index, %positions: buffer, %input: buffer, %theta: buffer, %freq_factors: buffer, %output: buffer) { + %token_capacity = config.get @ggml.rope_f32.token_capacity : index + %head_count0 = config.get @ggml.rope_f32.head_count : index + %head_size0 = config.get @ggml.rope_f32.head_size : index + %n_dims0 = config.get @ggml.rope_f32.n_dims : index + %input_stride1_0 = config.get @ggml.rope_f32.input_stride1 : index + %input_stride2_0 = config.get @ggml.rope_f32.input_stride2 : index + %input_span0 = config.get @ggml.rope_f32.input_span : index + %mscale = config.get @ggml.rope_f32.mscale : f32 + %mode = config.get @ggml.rope_f32.mode : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %bounded_input_span = index.assume %input_span0 [range(%input_span0, 1, 1073741824)] : index + %head_count, %head_size, %n_dims = index.assume %head_count0, %head_size0, %n_dims0 [range(%head_count0, 1, 64), range(%head_size0, 4, 1024), range(%n_dims0, 4, 1024), mul(%head_size0, 4), mul(%n_dims0, 4), le(%n_dims0, %head_size0)] : index, index, index + %input_stride1, %input_stride2 = index.assume %input_stride1_0, %input_stride2_0 [range(%input_stride1_0, 4, 1073741824), range(%input_stride2_0, 4, 1073741824)] : index, index + %head = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %channel0 = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %mode_neox = index.constant 2 : index + %c0_offset = index.constant 0 : offset + %half_n_dims = index.div %n_dims, %c2 : index + %pair_packet_count = index.div %head_size, %c4 : index + %rotary_packet_count = index.div %n_dims, %c4 : index + %valid_token = index.cmp ult, %token0, %bounded_token_count : index + %active_channel = index.cmp ult, %channel0, %pair_packet_count : index + %rotary_channel = index.cmp ult, %channel0, %rotary_packet_count : index + %tail_channel = index.cmp uge, %channel0, %rotary_packet_count : index + %copy_channel = scalar.andi %active_channel, %tail_channel : i1 + %publish = scalar.andi %valid_token, %active_channel : i1 + %publish_rotary0 = scalar.andi %valid_token, %rotary_channel : i1 + %publish_rotary = scalar.andi %publish_rotary0, %active_channel : i1 + %publish_copy = scalar.andi %valid_token, %copy_channel : i1 + %token = scf.select %valid_token, %token0, %c0 : index + %is_neox = index.cmp eq, %mode, %mode_neox : index + %is_normal = index.cmp eq, %mode, %c0 : index + %positions_noalias, %input_noalias, %theta_noalias, %freq_factors_noalias, %output_noalias = buffer.assume.noalias %positions, %input, %theta, %freq_factors, %output : buffer, buffer, buffer, buffer, buffer + scf.if %publish_rotary { + %channel = index.assume %channel0 [lt(%channel0, %pair_packet_count)] : index + %pair_channel = index.mul %channel, %c2 : index + %normal_channel = index.mul %channel, %c4 : index + %positions_view = buffer.view %positions_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]xi32> + %theta_view = buffer.view %theta_noalias[%c0_offset] : buffer -> view<[%half_n_dims]xf32> + %freq_factors_view = buffer.view %freq_factors_noalias[%c0_offset] : buffer -> view<[%half_n_dims]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%head_count]x[%head_size]xf32> + %input_token_offset = index.mul %token, %input_stride2 : index + %input_head_offset = index.mul %head, %input_stride1 : index + %input_base = index.add %input_token_offset, %input_head_offset : index + %position_i32 = view.load %positions_view[%token] : view<[%bounded_token_count]xi32> -> i32 + %position = scalar.sitofp %position_i32 : i32 to f32 + %theta_packet = vector.load %theta_view[%pair_channel] : view<[%half_n_dims]xf32> -> vector<2xf32> + %freq_factors_packet = vector.load %freq_factors_view[%pair_channel] : view<[%half_n_dims]xf32> -> vector<2xf32> + scf.if %is_neox { + %mscale_vector = vector.splat %mscale : vector<2xf32> + %paired_channel = index.add %pair_channel, %half_n_dims : index + %low_index0 = index.add %input_base, %pair_channel : index + %high_index0 = index.add %input_base, %paired_channel : index + %low_end0 = index.add %low_index0, %c2 : index + %high_end0 = index.add %high_index0, %c2 : index + %low_index, %low_end = index.assume %low_index0, %low_end0 [le(%low_end0, %bounded_input_span)] : index, index + %high_index, %high_end = index.assume %high_index0, %high_end0 [le(%high_end0, %bounded_input_span)] : index, index + %low_input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%low_end]xf32> + %high_input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%high_end]xf32> + %low_values = vector.load %low_input_view[%low_index] : view<[%low_end]xf32> -> vector<2xf32> + %high_values = vector.load %high_input_view[%high_index] : view<[%high_end]xf32> -> vector<2xf32> + %rotated_low, %rotated_high = func.call @ggml_rope_f32_pair_packet(%position, %theta_packet, %freq_factors_packet, %low_values, %high_values) : (f32, vector<2xf32>, vector<2xf32>, vector<2xf32>, vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) + %scaled_low = vector.mulf %rotated_low, %mscale_vector : vector<2xf32> + %scaled_high = vector.mulf %rotated_high, %mscale_vector : vector<2xf32> + vector.store %scaled_low, %output_view[%token, %head, %pair_channel] : vector<2xf32>, view<[%bounded_token_count]x[%head_count]x[%head_size]xf32> + vector.store %scaled_high, %output_view[%token, %head, %paired_channel] : vector<2xf32>, view<[%bounded_token_count]x[%head_count]x[%head_size]xf32> + } + scf.if %is_normal { + %mscale_vector = vector.splat %mscale : vector<2xf32> + %normal_index0 = index.add %input_base, %normal_channel : index + %normal_end0 = index.add %normal_index0, %c4 : index + %normal_index, %normal_end = index.assume %normal_index0, %normal_end0 [le(%normal_end0, %bounded_input_span)] : index, index + %normal_input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%normal_end]xf32> + %values = vector.load %normal_input_view[%normal_index] : view<[%normal_end]xf32> -> vector<4xf32> + %x0 = vector.extract %values[0] : vector<4xf32> -> f32 + %y0 = vector.extract %values[1] : vector<4xf32> -> f32 + %x1 = vector.extract %values[2] : vector<4xf32> -> f32 + %y1 = vector.extract %values[3] : vector<4xf32> -> f32 + %x_values = vector.from_elements %x0, %x1 : vector<2xf32> + %y_values = vector.from_elements %y0, %y1 : vector<2xf32> + %rotated_x, %rotated_y = func.call @ggml_rope_f32_pair_packet(%position, %theta_packet, %freq_factors_packet, %x_values, %y_values) : (f32, vector<2xf32>, vector<2xf32>, vector<2xf32>, vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) + %scaled_x = vector.mulf %rotated_x, %mscale_vector : vector<2xf32> + %scaled_y = vector.mulf %rotated_y, %mscale_vector : vector<2xf32> + %scaled_x0 = vector.extract %scaled_x[0] : vector<2xf32> -> f32 + %scaled_x1 = vector.extract %scaled_x[1] : vector<2xf32> -> f32 + %scaled_y0 = vector.extract %scaled_y[0] : vector<2xf32> -> f32 + %scaled_y1 = vector.extract %scaled_y[1] : vector<2xf32> -> f32 + %rotated = vector.from_elements %scaled_x0, %scaled_y0, %scaled_x1, %scaled_y1 : vector<4xf32> + vector.store %rotated, %output_view[%token, %head, %normal_channel] : vector<4xf32>, view<[%bounded_token_count]x[%head_count]x[%head_size]xf32> + } + } + scf.if %publish_copy { + %channel = index.assume %channel0 [lt(%channel0, %pair_packet_count)] : index + %normal_channel = index.mul %channel, %c4 : index + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%head_count]x[%head_size]xf32> + %input_token_offset = index.mul %token, %input_stride2 : index + %input_head_offset = index.mul %head, %input_stride1 : index + %input_base = index.add %input_token_offset, %input_head_offset : index + %normal_index0 = index.add %input_base, %normal_channel : index + %normal_end0 = index.add %normal_index0, %c4 : index + %normal_index, %normal_end = index.assume %normal_index0, %normal_end0 [le(%normal_end0, %bounded_input_span)] : index, index + %normal_input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%normal_end]xf32> + %values = vector.load %normal_input_view[%normal_index] : view<[%normal_end]xf32> -> vector<4xf32> + vector.store %values, %output_view[%token, %head, %normal_channel] : vector<4xf32>, view<[%bounded_token_count]x[%head_count]x[%head_size]xf32> + } + template.return +} + +kernel.def target(@ggml_rope_f32_gfx11_wave32) @ggml_rope_f32(%token_count: index) { + %token_capacity = config.get @ggml.rope_f32.token_capacity : index + %head_count = config.get @ggml.rope_f32.head_count : index + %head_size = config.get @ggml.rope_f32.head_size : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %pair_packet_count = index.div %head_size, %c4 : index + kernel.launch.config workgroups(%head_count, %token_capacity, %c1) workgroup_size(%pair_packet_count, %c1, %c1) : index +} launch(%token_count: index, %positions: buffer, %input: buffer, %theta: buffer, %freq_factors: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + template.apply<@ggml.rope_f32.body>(%token_count, %positions, %input, %theta, %freq_factors, %output) : (index, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/rope_set_rows_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/rope_set_rows_f32.loom new file mode 100644 index 000000000000..fb649c0f5361 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/rope_set_rows_f32.loom @@ -0,0 +1,209 @@ +// Fused F32 ROPE followed by indexed cache publication. +template.decl @ggml.rope_set_rows_f32.body(%token_count: index, %cache_row_count: index, %positions: buffer, %indices: buffer, %input: buffer, %theta: buffer, %freq_factors: buffer, %cache: buffer) + +template.decl @ggml.rope_set_rows_f32.store(%output_format: index, %publish: i1, %cache_row: index, %head: index, %channel: index, %cache_row_count: index, %head_count: index, %head_size: index, %values: vector<2xf32>, %cache: buffer) + +template.decl @ggml.rope_set_rows_f32.store4(%output_format: index, %publish: i1, %cache_row: index, %head: index, %channel: index, %cache_row_count: index, %head_count: index, %head_size: index, %values: vector<4xf32>, %cache: buffer) + +amdgpu.target @ggml_rope_set_rows_f32_gfx11_wave32 {subgroup_size = 32} + +config.decl @ggml.rope_set_rows_f32.head_size : %value: index where [range(%value, 4, 1024), mul(%value, 4)] + +config.decl @ggml.rope_set_rows_f32.n_dims : %value: index where [range(%value, 4, 1024), mul(%value, 4)] + +config.decl @ggml.rope_set_rows_f32.head_count : %value: index where [range(%value, 1, 64)] + +config.decl @ggml.rope_set_rows_f32.token_capacity : %value: index where [range(%value, 1, 2048)] + +config.decl @ggml.rope_set_rows_f32.input_stride1 : %value: index where [range(%value, 4, 1073741824)] + +config.decl @ggml.rope_set_rows_f32.input_stride2 : %value: index where [range(%value, 4, 1073741824)] + +config.decl @ggml.rope_set_rows_f32.input_span : %value: index where [range(%value, 1, 1073741824)] + +config.decl @ggml.rope_set_rows_f32.mscale : f32 + +config.decl @ggml.rope_set_rows_f32.output_format : %value: index where [range(%value, 16, 32)] + +config.decl @ggml.rope_set_rows_f32.mode : %value: index where [range(%value, 0, 2)] + +func.decl @ggml_rope_f32_pair_packet(%position: f32, %theta: vector<2xf32>, %freq_factors: vector<2xf32>, %x_values: vector<2xf32>, %y_values: vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) + +template.def<@ggml.rope_set_rows_f32.store> device @ggml_rope_set_rows_f32_store(%output_format: index, %publish: i1, %cache_row: index, %head: index, %channel: index, %cache_row_count: index, %head_count: index, %head_size: index, %values: vector<2xf32>, %cache: buffer) { + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c0_offset = index.constant 0 : offset + %is_f16 = index.cmp eq, %output_format, %c16 : index + %is_f32 = index.cmp eq, %output_format, %c32 : index + %publish_f16 = scalar.andi %publish, %is_f16 : i1 + %publish_f32 = scalar.andi %publish, %is_f32 : i1 + scf.if %publish_f16 { + %cache_view = buffer.view %cache[%c0_offset] : buffer -> view<[%cache_row_count]x[%head_count]x[%head_size]xf16> + %truncated = vector.fptrunc %values : vector<2xf32> to vector<2xf16> + vector.store %truncated, %cache_view[%cache_row, %head, %channel] : vector<2xf16>, view<[%cache_row_count]x[%head_count]x[%head_size]xf16> + } + scf.if %publish_f32 { + %cache_view = buffer.view %cache[%c0_offset] : buffer -> view<[%cache_row_count]x[%head_count]x[%head_size]xf32> + vector.store %values, %cache_view[%cache_row, %head, %channel] : vector<2xf32>, view<[%cache_row_count]x[%head_count]x[%head_size]xf32> + } + template.return +} + +template.def<@ggml.rope_set_rows_f32.store4> device @ggml_rope_set_rows_f32_store4(%output_format: index, %publish: i1, %cache_row: index, %head: index, %channel: index, %cache_row_count: index, %head_count: index, %head_size: index, %values: vector<4xf32>, %cache: buffer) { + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c0_offset = index.constant 0 : offset + %is_f16 = index.cmp eq, %output_format, %c16 : index + %is_f32 = index.cmp eq, %output_format, %c32 : index + %publish_f16 = scalar.andi %publish, %is_f16 : i1 + %publish_f32 = scalar.andi %publish, %is_f32 : i1 + scf.if %publish_f16 { + %cache_view = buffer.view %cache[%c0_offset] : buffer -> view<[%cache_row_count]x[%head_count]x[%head_size]xf16> + %truncated = vector.fptrunc %values : vector<4xf32> to vector<4xf16> + vector.store %truncated, %cache_view[%cache_row, %head, %channel] : vector<4xf16>, view<[%cache_row_count]x[%head_count]x[%head_size]xf16> + } + scf.if %publish_f32 { + %cache_view = buffer.view %cache[%c0_offset] : buffer -> view<[%cache_row_count]x[%head_count]x[%head_size]xf32> + vector.store %values, %cache_view[%cache_row, %head, %channel] : vector<4xf32>, view<[%cache_row_count]x[%head_count]x[%head_size]xf32> + } + template.return +} + +template.def<@ggml.rope_set_rows_f32.body> device @ggml_rope_set_rows_f32_body(%token_count: index, %cache_row_count: index, %positions: buffer, %indices: buffer, %input: buffer, %theta: buffer, %freq_factors: buffer, %cache: buffer) { + %token_capacity = config.get @ggml.rope_set_rows_f32.token_capacity : index + %head_count0 = config.get @ggml.rope_set_rows_f32.head_count : index + %head_size0 = config.get @ggml.rope_set_rows_f32.head_size : index + %n_dims0 = config.get @ggml.rope_set_rows_f32.n_dims : index + %input_stride1_0 = config.get @ggml.rope_set_rows_f32.input_stride1 : index + %input_stride2_0 = config.get @ggml.rope_set_rows_f32.input_stride2 : index + %input_span0 = config.get @ggml.rope_set_rows_f32.input_span : index + %mscale = config.get @ggml.rope_set_rows_f32.mscale : f32 + %output_format = config.get @ggml.rope_set_rows_f32.output_format : index + %mode = config.get @ggml.rope_set_rows_f32.mode : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %bounded_cache_row_count = index.assume %cache_row_count [range(%cache_row_count, 1, 1048576)] : index + %bounded_input_span = index.assume %input_span0 [range(%input_span0, 1, 1073741824)] : index + %head_count, %head_size, %n_dims = index.assume %head_count0, %head_size0, %n_dims0 [range(%head_count0, 1, 64), range(%head_size0, 4, 1024), range(%n_dims0, 4, 1024), mul(%head_size0, 4), mul(%n_dims0, 4), le(%n_dims0, %head_size0)] : index, index, index + %input_stride1, %input_stride2 = index.assume %input_stride1_0, %input_stride2_0 [range(%input_stride1_0, 4, 1073741824), range(%input_stride2_0, 4, 1073741824)] : index, index + %head = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %channel0 = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %mode_neox = index.constant 2 : index + %c0_i64 = scalar.constant 0 : i64 + %c1048575_i64 = scalar.constant 1048575 : i64 + %c0_offset = index.constant 0 : offset + %half_n_dims = index.div %n_dims, %c2 : index + %pair_packet_count = index.div %head_size, %c4 : index + %rotary_packet_count = index.div %n_dims, %c4 : index + %valid_token = index.cmp ult, %token0, %bounded_token_count : index + %active_channel = index.cmp ult, %channel0, %pair_packet_count : index + %rotary_channel = index.cmp ult, %channel0, %rotary_packet_count : index + %tail_channel = index.cmp uge, %channel0, %rotary_packet_count : index + %copy_channel = scalar.andi %active_channel, %tail_channel : i1 + %publish0 = scalar.andi %valid_token, %active_channel : i1 + %token = scf.select %valid_token, %token0, %c0 : index + %is_neox = index.cmp eq, %mode, %mode_neox : index + %is_normal = index.cmp eq, %mode, %c0 : index + %positions_noalias, %indices_noalias, %input_noalias, %theta_noalias, %freq_factors_noalias, %cache_noalias = buffer.assume.noalias %positions, %indices, %input, %theta, %freq_factors, %cache : buffer, buffer, buffer, buffer, buffer, buffer + %positions_view = buffer.view %positions_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]xi32> + %indices_view = buffer.view %indices_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]xi64> + %theta_view = buffer.view %theta_noalias[%c0_offset] : buffer -> view<[%half_n_dims]xf32> + %freq_factors_view = buffer.view %freq_factors_noalias[%c0_offset] : buffer -> view<[%half_n_dims]xf32> + %index_raw = view.load %indices_view[%token] : view<[%bounded_token_count]xi64> -> i64 + %index_nonnegative = scalar.cmpi sge, %index_raw, %c0_i64 : i64 + %index_in_cast_range = scalar.cmpi sle, %index_raw, %c1048575_i64 : i64 + %valid_index = scalar.andi %index_nonnegative, %index_in_cast_range : i1 + %safe_index0_i64 = scf.select %valid_index, %index_raw, %c0_i64 : i64 + %safe_index_i64 = scalar.assume %safe_index0_i64 [range(%safe_index0_i64, 0, 1048575)] : i64 + %cache_row0 = index.cast %safe_index_i64 : i64 to index + %valid_row = index.cmp ult, %cache_row0, %bounded_cache_row_count : index + %publish1 = scalar.andi %publish0, %valid_index : i1 + %publish = scalar.andi %publish1, %valid_row : i1 + %cache_row = scf.select %valid_row, %cache_row0, %c0 : index + %publish_rotary = scalar.andi %publish, %rotary_channel : i1 + %publish_copy = scalar.andi %publish, %copy_channel : i1 + scf.if %publish_rotary { + %channel = index.assume %channel0 [lt(%channel0, %pair_packet_count)] : index + %pair_channel = index.mul %channel, %c2 : index + %normal_channel = index.mul %channel, %c4 : index + %input_token_offset = index.mul %token, %input_stride2 : index + %input_head_offset = index.mul %head, %input_stride1 : index + %input_base = index.add %input_token_offset, %input_head_offset : index + %position_i32 = view.load %positions_view[%token] : view<[%bounded_token_count]xi32> -> i32 + %position = scalar.sitofp %position_i32 : i32 to f32 + %theta_packet = vector.load %theta_view[%pair_channel] : view<[%half_n_dims]xf32> -> vector<2xf32> + %freq_factors_packet = vector.load %freq_factors_view[%pair_channel] : view<[%half_n_dims]xf32> -> vector<2xf32> + scf.if %is_neox { + %mscale_vector = vector.splat %mscale : vector<2xf32> + %paired_channel = index.add %pair_channel, %half_n_dims : index + %low_index0 = index.add %input_base, %pair_channel : index + %high_index0 = index.add %input_base, %paired_channel : index + %low_end0 = index.add %low_index0, %c2 : index + %high_end0 = index.add %high_index0, %c2 : index + %low_index, %low_end = index.assume %low_index0, %low_end0 [le(%low_end0, %bounded_input_span)] : index, index + %high_index, %high_end = index.assume %high_index0, %high_end0 [le(%high_end0, %bounded_input_span)] : index, index + %low_input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%low_end]xf32> + %high_input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%high_end]xf32> + %low_values = vector.load %low_input_view[%low_index] : view<[%low_end]xf32> -> vector<2xf32> + %high_values = vector.load %high_input_view[%high_index] : view<[%high_end]xf32> -> vector<2xf32> + %rotated_low, %rotated_high = func.call @ggml_rope_f32_pair_packet(%position, %theta_packet, %freq_factors_packet, %low_values, %high_values) : (f32, vector<2xf32>, vector<2xf32>, vector<2xf32>, vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) + %scaled_low = vector.mulf %rotated_low, %mscale_vector : vector<2xf32> + %scaled_high = vector.mulf %rotated_high, %mscale_vector : vector<2xf32> + template.apply<@ggml.rope_set_rows_f32.store>(%output_format, %publish, %cache_row, %head, %pair_channel, %bounded_cache_row_count, %head_count, %head_size, %scaled_low, %cache_noalias) : (index, i1, index, index, index, index, index, index, vector<2xf32>, buffer) + template.apply<@ggml.rope_set_rows_f32.store>(%output_format, %publish, %cache_row, %head, %paired_channel, %bounded_cache_row_count, %head_count, %head_size, %scaled_high, %cache_noalias) : (index, i1, index, index, index, index, index, index, vector<2xf32>, buffer) + } + scf.if %is_normal { + %mscale_vector = vector.splat %mscale : vector<2xf32> + %normal_index0 = index.add %input_base, %normal_channel : index + %normal_end0 = index.add %normal_index0, %c4 : index + %normal_index, %normal_end = index.assume %normal_index0, %normal_end0 [le(%normal_end0, %bounded_input_span)] : index, index + %normal_input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%normal_end]xf32> + %values = vector.load %normal_input_view[%normal_index] : view<[%normal_end]xf32> -> vector<4xf32> + %x0 = vector.extract %values[0] : vector<4xf32> -> f32 + %y0 = vector.extract %values[1] : vector<4xf32> -> f32 + %x1 = vector.extract %values[2] : vector<4xf32> -> f32 + %y1 = vector.extract %values[3] : vector<4xf32> -> f32 + %x_values = vector.from_elements %x0, %x1 : vector<2xf32> + %y_values = vector.from_elements %y0, %y1 : vector<2xf32> + %rotated_x, %rotated_y = func.call @ggml_rope_f32_pair_packet(%position, %theta_packet, %freq_factors_packet, %x_values, %y_values) : (f32, vector<2xf32>, vector<2xf32>, vector<2xf32>, vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) + %scaled_x = vector.mulf %rotated_x, %mscale_vector : vector<2xf32> + %scaled_y = vector.mulf %rotated_y, %mscale_vector : vector<2xf32> + %scaled_x0 = vector.extract %scaled_x[0] : vector<2xf32> -> f32 + %scaled_x1 = vector.extract %scaled_x[1] : vector<2xf32> -> f32 + %scaled_y0 = vector.extract %scaled_y[0] : vector<2xf32> -> f32 + %scaled_y1 = vector.extract %scaled_y[1] : vector<2xf32> -> f32 + %rotated = vector.from_elements %scaled_x0, %scaled_y0, %scaled_x1, %scaled_y1 : vector<4xf32> + template.apply<@ggml.rope_set_rows_f32.store4>(%output_format, %publish, %cache_row, %head, %normal_channel, %bounded_cache_row_count, %head_count, %head_size, %rotated, %cache_noalias) : (index, i1, index, index, index, index, index, index, vector<4xf32>, buffer) + } + } + scf.if %publish_copy { + %channel = index.assume %channel0 [lt(%channel0, %pair_packet_count)] : index + %normal_channel = index.mul %channel, %c4 : index + %input_token_offset = index.mul %token, %input_stride2 : index + %input_head_offset = index.mul %head, %input_stride1 : index + %input_base = index.add %input_token_offset, %input_head_offset : index + %normal_index0 = index.add %input_base, %normal_channel : index + %normal_end0 = index.add %normal_index0, %c4 : index + %normal_index, %normal_end = index.assume %normal_index0, %normal_end0 [le(%normal_end0, %bounded_input_span)] : index, index + %normal_input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%normal_end]xf32> + %values = vector.load %normal_input_view[%normal_index] : view<[%normal_end]xf32> -> vector<4xf32> + template.apply<@ggml.rope_set_rows_f32.store4>(%output_format, %publish, %cache_row, %head, %normal_channel, %bounded_cache_row_count, %head_count, %head_size, %values, %cache_noalias) : (index, i1, index, index, index, index, index, index, vector<4xf32>, buffer) + } + template.return +} + +kernel.def target(@ggml_rope_set_rows_f32_gfx11_wave32) @ggml_rope_set_rows_f32(%token_count: index, %cache_row_count: index) { + %token_capacity = config.get @ggml.rope_set_rows_f32.token_capacity : index + %head_count = config.get @ggml.rope_set_rows_f32.head_count : index + %head_size = config.get @ggml.rope_set_rows_f32.head_size : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %pair_packet_count = index.div %head_size, %c4 : index + kernel.launch.config workgroups(%head_count, %token_capacity, %c1) workgroup_size(%pair_packet_count, %c1, %c1) : index +} launch(%token_count: index, %cache_row_count: index, %positions: buffer, %indices: buffer, %input: buffer, %theta: buffer, %freq_factors: buffer, %cache: buffer) where [range(%token_count, 1, 2048)] { + template.apply<@ggml.rope_set_rows_f32.body>(%token_count, %cache_row_count, %positions, %indices, %input, %theta, %freq_factors, %cache) : (index, index, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/scale_bias_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/scale_bias_f32.loom new file mode 100644 index 000000000000..30abcd59c564 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/scale_bias_f32.loom @@ -0,0 +1,50 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +amdgpu.target @ggml_scale_bias_f32_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.scale.scale : f32 + +config.decl @ggml.scale.bias : f32 + +kernel.def target(@ggml_scale_bias_f32_gfx11_wave64) export("ggml_scale_bias_f32") @ggml_scale_bias_f32(%element_count: index) { + %one = index.constant 1 : index + %twofiftysix = index.constant 256 : index + %rounding = index.constant 255 : index + %rounded = index.add %element_count, %rounding : index + %workgroup_count = index.div %rounded, %twofiftysix : index + kernel.launch.config workgroups(%workgroup_count, %one, %one) workgroup_size(%twofiftysix, %one, %one) : index +} launch(%element_count: index, %input: buffer, %output: buffer) { + %count = index.assume %element_count [range(%element_count, 1, 134217728)] : index + %workgroup = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %twofiftysix = index.constant 256 : index + %base = index.mul %workgroup, %twofiftysix : index + %linear0 = index.add %base, %workitem : index + %linear = index.assume %linear0 [range(%linear0, 0, 134217983)] : index + %in_bounds = index.cmp ult, %linear, %count : index + %scale = config.get @ggml.scale.scale : f32 + %bias = config.get @ggml.scale.bias : f32 + %zero_offset = index.constant 0 : offset + %input_view = buffer.view %input[%zero_offset] : buffer -> view<[%count]xf32> + %output_view = buffer.view %output[%zero_offset] : buffer -> view<[%count]xf32> + scf.if %in_bounds { + %input_value = view.load %input_view[%linear] : view<[%count]xf32> -> f32 + %scaled = scalar.mulf %input_value, %scale : f32 + %result = scalar.addf %scaled, %bias : f32 + view.store %result, %output_view[%linear] : f32, view<[%count]xf32> + } + kernel.return +} + +check.case public @ggml_scale_bias_f32_small_case { + %four = check.literal value(4) : index + %input = check.generate.iota offset(-2.0) step(1.0) : tensor<4xf32> + %output = check.generate.fill value(0.0) : tensor<4xf32> + %expected = check.generate.iota offset(-3.0) step(2.0) : tensor<4xf32> + kernel.launch @ggml_scale_bias_f32[%four](%four, %input, %output) : [index](index, tensor<4xf32>, tensor<4xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<4xf32> + check.return +} + +check.benchmark<@ggml_scale_bias_f32_small_case> @ggml_scale_bias_f32_small diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/scale_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/scale_f32.loom new file mode 100644 index 000000000000..6cef7d497296 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/scale_f32.loom @@ -0,0 +1,39 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +amdgpu.target @ggml_scale_f32_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.scale_f32.scale : f32 + +config.decl @ggml.scale_f32.bias : f32 + +kernel.def target(@ggml_scale_f32_gfx11_wave64) export("ggml_scale_f32") @ggml_scale_f32(%element_count: index) { + %one = index.constant 1 : index + %twofiftysix = index.constant 256 : index + %rounding = index.constant 255 : index + %rounded = index.add %element_count, %rounding : index + %workgroup_count = index.div %rounded, %twofiftysix : index + kernel.launch.config workgroups(%workgroup_count, %one, %one) workgroup_size(%twofiftysix, %one, %one) : index +} launch(%element_count: index, %input: buffer, %output: buffer) { + %count = index.assume %element_count [range(%element_count, 1, 134217728)] : index + %scale = config.get @ggml.scale_f32.scale : f32 + %bias = config.get @ggml.scale_f32.bias : f32 + %workgroup = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %twofiftysix = index.constant 256 : index + %base0 = index.mul %workgroup, %twofiftysix : index + %linear0 = index.add %base0, %workitem : index + %linear = index.assume %linear0 [range(%linear0, 0, 134217983)] : index + %in_bounds = index.cmp ult, %linear, %count : index + %zero_offset = index.constant 0 : offset + %input_noalias, %output_noalias = buffer.assume.noalias %input, %output : buffer, buffer + %input_view = buffer.view %input_noalias[%zero_offset] : buffer -> view<[%count]xf32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%count]xf32> + scf.if %in_bounds { + %value = view.load %input_view[%linear] : view<[%count]xf32> -> f32 + %scaled = scalar.mulf %value, %scale : f32 + %result = scalar.addf %scaled, %bias : f32 + view.store %result, %output_view[%linear] : f32, view<[%count]xf32> + } + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/set_rows.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/set_rows.loom new file mode 100644 index 000000000000..4d7a5e518e00 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/set_rows.loom @@ -0,0 +1,184 @@ +// Generic GGML SET_ROWS for contiguous 2D row caches. +template.decl @ggml.set_rows.body(%token_count: index, %cache_row_count: index, %hidden_size: index, %rows: buffer, %indices: buffer, %cache: buffer) + +template.decl @ggml.set_rows.launch(%hidden_capacity: index, %token_capacity: index) -> (index, index, index, index) + +template.decl @ggml.set_rows.load_f32_vector4(%input_format: index, %token: index, %channel: index, %token_count: index, %input_stride: index, %rows: buffer) -> (vector<4xf32>) + +template.decl @ggml.set_rows.store_f32_vector4(%output_format: index, %publish: i1, %cache_row: index, %channel: index, %cache_row_count: index, %hidden_size: index, %values: vector<4xf32>, %cache: buffer) + +amdgpu.target @ggml_set_rows_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.set_rows.token_capacity : %value: index where [range(%value, 1, 2048)] + +config.decl @ggml.set_rows.hidden_capacity : %value: index where [range(%value, 4, 32768), mul(%value, 4)] + +config.decl @ggml.set_rows.input_format : %value: index where [range(%value, 16, 32)] + +config.decl @ggml.set_rows.output_format : %value: index where [range(%value, 16, 32)] + +config.decl @ggml.set_rows.input_stride : %value: index where [range(%value, 4, 1048576)] + +template.def<@ggml.set_rows.load_f32_vector4> device @ggml_set_rows_load_f32_vector4(%input_format: index, %token: index, %channel: index, %token_count: index, %input_stride: index, %rows: buffer) -> (vector<4xf32>) { + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c0_offset = index.constant 0 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %zero = vector.splat %c0_f32 : vector<4xf32> + %is_f16 = index.cmp eq, %input_format, %c16 : index + %is_f32 = index.cmp eq, %input_format, %c32 : index + %f16_values = scf.if %is_f16 -> (vector<4xf32>) { + %rows_view = buffer.view %rows[%c0_offset] : buffer -> view<[%token_count]x[%input_stride]xf16> + %loaded = vector.load %rows_view[%token, %channel] : view<[%token_count]x[%input_stride]xf16> -> vector<4xf16> + %widened = vector.extf %loaded : vector<4xf16> to vector<4xf32> + scf.yield %widened : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %f32_values = scf.if %is_f32 -> (vector<4xf32>) { + %rows_view = buffer.view %rows[%c0_offset] : buffer -> view<[%token_count]x[%input_stride]xf32> + %loaded = vector.load %rows_view[%token, %channel] : view<[%token_count]x[%input_stride]xf32> -> vector<4xf32> + scf.yield %loaded : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } + %selected = scf.select %is_f32, %f32_values, %f16_values : vector<4xf32> + template.return %selected : vector<4xf32> +} + +template.def<@ggml.set_rows.store_f32_vector4> device @ggml_set_rows_store_f32_vector4(%output_format: index, %publish: i1, %cache_row: index, %channel: index, %cache_row_count: index, %hidden_size: index, %values: vector<4xf32>, %cache: buffer) { + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c0_offset = index.constant 0 : offset + %is_f16 = index.cmp eq, %output_format, %c16 : index + %is_f32 = index.cmp eq, %output_format, %c32 : index + %publish_f16 = scalar.andi %publish, %is_f16 : i1 + %publish_f32 = scalar.andi %publish, %is_f32 : i1 + scf.if %publish_f16 { + %cache_view = buffer.view %cache[%c0_offset] : buffer -> view<[%cache_row_count]x[%hidden_size]xf16> + %truncated = vector.fptrunc %values : vector<4xf32> to vector<4xf16> + vector.store %truncated, %cache_view[%cache_row, %channel] : vector<4xf16>, view<[%cache_row_count]x[%hidden_size]xf16> + } + scf.if %publish_f32 { + %cache_view = buffer.view %cache[%c0_offset] : buffer -> view<[%cache_row_count]x[%hidden_size]xf32> + vector.store %values, %cache_view[%cache_row, %channel] : vector<4xf32>, view<[%cache_row_count]x[%hidden_size]xf32> + } + template.return +} + +template.def<@ggml.set_rows.body> device @ggml_set_rows_body(%token_count: index, %cache_row_count: index, %hidden_size: index, %rows: buffer, %indices: buffer, %cache: buffer) { + %token_capacity = config.get @ggml.set_rows.token_capacity : index + %hidden_capacity = config.get @ggml.set_rows.hidden_capacity : index + %input_format = config.get @ggml.set_rows.input_format : index + %output_format = config.get @ggml.set_rows.output_format : index + %input_stride = config.get @ggml.set_rows.input_stride : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %bounded_cache_row_count = index.assume %cache_row_count [range(%cache_row_count, 1, 1048576)] : index + %bounded_hidden_size = index.assume %hidden_size [range(%hidden_size, 4, 32768), mul(%hidden_size, 4), le(%hidden_size, %hidden_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %c256 = index.constant 256 : index + %c0_i64 = scalar.constant 0 : i64 + %c1048575_i64 = scalar.constant 1048575 : i64 + %c0_offset = index.constant 0 : offset + %packet_tile_base = index.mul %channel_tile, %c256 : index + %packet0 = index.add %packet_tile_base, %workitem0 : index + %packet_count = index.div %bounded_hidden_size, %c4 : index + %valid_token = index.cmp ult, %token0, %bounded_token_count : index + %valid_packet = index.cmp ult, %packet0, %packet_count : index + %publish0 = scalar.andi %valid_token, %valid_packet : i1 + %token = scf.select %valid_token, %token0, %c0 : index + %packet = scf.select %valid_packet, %packet0, %c0 : index + %channel = index.mul %packet, %c4 : index + %rows_noalias, %indices_noalias, %cache_noalias = buffer.assume.noalias %rows, %indices, %cache : buffer, buffer, buffer + %indices_view = buffer.view %indices_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]xi64> + %index_raw = view.load %indices_view[%token] : view<[%bounded_token_count]xi64> -> i64 + %index_nonnegative = scalar.cmpi sge, %index_raw, %c0_i64 : i64 + %index_in_cast_range = scalar.cmpi sle, %index_raw, %c1048575_i64 : i64 + %valid_index = scalar.andi %index_nonnegative, %index_in_cast_range : i1 + %safe_index0_i64 = scf.select %valid_index, %index_raw, %c0_i64 : i64 + %safe_index_i64 = scalar.assume %safe_index0_i64 [range(%safe_index0_i64, 0, 1048575)] : i64 + %cache_row0 = index.cast %safe_index_i64 : i64 to index + %valid_row = index.cmp ult, %cache_row0, %bounded_cache_row_count : index + %publish1 = scalar.andi %publish0, %valid_index : i1 + %publish = scalar.andi %publish1, %valid_row : i1 + %cache_row = scf.select %valid_row, %cache_row0, %c0 : index + %values = template.apply<@ggml.set_rows.load_f32_vector4>(%input_format, %token, %channel, %bounded_token_count, %input_stride, %rows_noalias) : (index, index, index, index, index, buffer) -> (vector<4xf32>) + template.apply<@ggml.set_rows.store_f32_vector4>(%output_format, %publish, %cache_row, %channel, %bounded_cache_row_count, %bounded_hidden_size, %values, %cache_noalias) : (index, i1, index, index, index, index, vector<4xf32>, buffer) + template.return +} + +template.def<@ggml.set_rows.launch> @ggml_set_rows_launch(%hidden_capacity: index, %token_capacity: index) -> (index, index, index, index) { + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %packet_count = index.div %hidden_capacity, %c4 : index + %padded_packet_count = index.add %packet_count, %c255 : index + %packet_tiles = index.div %padded_packet_count, %c256 : index + template.return %packet_tiles, %token_capacity, %c1, %c256 : index, index, index, index +} + +kernel.def target(@ggml_set_rows_gfx11_wave64) @ggml_set_rows(%token_count: index, %cache_row_count: index, %hidden_size: index) { + %token_capacity = config.get @ggml.set_rows.token_capacity : index + %hidden_capacity = config.get @ggml.set_rows.hidden_capacity : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c255 = index.constant 255 : index + %workgroup_size = index.constant 256 : index + %packet_count = index.div %hidden_capacity, %c4 : index + %padded_packet_count = index.add %packet_count, %c255 : index + %packet_tiles = index.div %padded_packet_count, %workgroup_size : index + kernel.launch.config workgroups(%packet_tiles, %token_capacity, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %cache_row_count: index, %hidden_size: index, %rows: buffer, %indices: buffer, %cache: buffer) where [range(%token_count, 1, 2048)] { + template.apply<@ggml.set_rows.body>(%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : (index, index, index, buffer, buffer, buffer) + kernel.return +} + +// Scalar SET_ROWS for element-wise KV-cache scatter. The non-FA attention path flattens +// the K/V rows, so hidden_size == 1 and each value is written to its own cache row. +config.decl @ggml.set_rows_scatter.token_capacity : %value: index where [range(%value, 1, 1048576)] + +kernel.def target(@ggml_set_rows_gfx11_wave64) @ggml_set_rows_scatter(%token_count: index, %cache_row_count: index) { + %token_capacity = config.get @ggml.set_rows_scatter.token_capacity : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %c1 = index.constant 1 : index + %padded = index.add %token_capacity, %c255 : index + %workgroups = index.div %padded, %c256 : index + kernel.launch.config workgroups(%workgroups, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %cache_row_count: index, %rows: buffer, %indices: buffer, %cache: buffer) where [range(%token_count, 1, 1048576)] { + %token_capacity = config.get @ggml.set_rows_scatter.token_capacity : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 1048576), le(%token_count, %token_capacity)] : index + %bounded_cache_row_count = index.assume %cache_row_count [range(%cache_row_count, 1, 134217728)] : index + %workitem = kernel.workitem.id : index + %workgroup = kernel.workgroup.id : index + %c256 = index.constant 256 : index + %c0 = index.constant 0 : index + %c0_offset = index.constant 0 : offset + %i0 = index.mul %workgroup, %c256 : index + %i = index.add %i0, %workitem : index + %valid = index.cmp ult, %i, %bounded_token_count : index + %safe_i = scf.select %valid, %i, %c0 : index + %rows_view = buffer.view %rows[%c0_offset] : buffer -> view<[%bounded_token_count]xf32> + %indices_view = buffer.view %indices[%c0_offset] : buffer -> view<[%bounded_token_count]xi64> + %cache_view = buffer.view %cache[%c0_offset] : buffer -> view<[%bounded_cache_row_count]xf16> + %value = view.load %rows_view[%safe_i] : view<[%bounded_token_count]xf32> -> f32 + %idx_i64 = view.load %indices_view[%safe_i] : view<[%bounded_token_count]xi64> -> i64 + %c0_i64 = scalar.constant 0 : i64 + %idx_nonneg = scalar.cmpi sge, %idx_i64, %c0_i64 : i64 + %safe_idx0_i64 = scf.select %idx_nonneg, %idx_i64, %c0_i64 : i64 + %safe_idx_i64 = scalar.assume %safe_idx0_i64 [range(%safe_idx0_i64, 0, 134217727)] : i64 + %idx = index.cast %safe_idx_i64 : i64 to index + %idx_in_range = index.cmp ult, %idx, %bounded_cache_row_count : index + %publish0 = scalar.andi %valid, %idx_nonneg : i1 + %publish = scalar.andi %publish0, %idx_in_range : i1 + %f16 = scalar.fptrunc %value : f32 to f16 + scf.if %publish { + view.store %f16, %cache_view[%idx] : f16, view<[%bounded_cache_row_count]xf16> + } + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/small_rows_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/small_rows_f32.loom new file mode 100644 index 000000000000..881c3b9095d4 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/small_rows_f32.loom @@ -0,0 +1,1105 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// Row ops for short rows (MoE routers, head groups): SOFT_MAX (no mask, scale 1), SUM_ROWS, +// ARGSORT and GET_ROWS; and NORM (LayerNorm without affine, as encoders such as ModernBERT use), +// one workgroup per row. They run where a model's router or head-group reduction would +// otherwise leave the GPU for a few dozen floats per token. One workitem per row (softmax, +// sum) or per element (argsort rank, gather); rows are short, so no cross-lane reduction. + +amdgpu.target @ggml_small_rows_f32_gfx11_wave64 {subgroup_size = 64} + +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_softmax_rows_f32") @ggml_softmax_rows_f32(%column_count: index, %row_count: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %rounded = index.add %row_count, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%column_count: index, %row_count: index, %input: buffer, %output: buffer) { + %cols = index.assume %column_count [range(%column_count, 1, 4096)] : index + %rows = index.assume %row_count [range(%row_count, 1, 16777216)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %c0_f32 = scalar.constant 0.0 : f32 + %lowest = scalar.constant -3.40282347e+38 : f32 + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %row0 = index.add %base, %lane : index + %valid = index.cmp ult, %row0, %rows : index + %row = scf.select %valid, %row0, %c0 : index + %input_noalias, %output_noalias = buffer.assume.noalias %input, %output : buffer, buffer + %input_view = buffer.view %input_noalias[%zero_offset] : buffer -> view<[%rows]x[%cols]xf32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%rows]x[%cols]xf32> + %maximum = scf.for %column = [%c0 to %cols step %c1](%running = %lowest : f32) -> (f32) { + %value = view.load %input_view[%row, %column] : view<[%rows]x[%cols]xf32> -> f32 + %next = scalar.maxnumf %running, %value : f32 + scf.yield %next : f32 + } + %total = scf.for %column = [%c0 to %cols step %c1](%running = %c0_f32 : f32) -> (f32) { + %value = view.load %input_view[%row, %column] : view<[%rows]x[%cols]xf32> -> f32 + %shifted = scalar.subf %value, %maximum : f32 + %exponential = scalar.expf %shifted : f32 + %next = scalar.addf %running, %exponential : f32 + scf.yield %next : f32 + } + scf.if %valid { + %unused = scf.for %column = [%c0 to %cols step %c1](%carry = %c0_f32 : f32) -> (f32) { + %value = view.load %input_view[%row, %column] : view<[%rows]x[%cols]xf32> -> f32 + %shifted = scalar.subf %value, %maximum : f32 + %exponential = scalar.expf %shifted : f32 + %probability = scalar.divf %exponential, %total : f32 + view.store %probability, %output_view[%row, %column] : f32, view<[%rows]x[%cols]xf32> + scf.yield %carry : f32 + } + } + kernel.return +} + +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_sum_rows_f32") @ggml_sum_rows_f32(%column_count: index, %row_count: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %rounded = index.add %row_count, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%column_count: index, %row_count: index, %input: buffer, %output: buffer) { + %cols = index.assume %column_count [range(%column_count, 1, 4096)] : index + %rows = index.assume %row_count [range(%row_count, 1, 16777216)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %c0_f32 = scalar.constant 0.0 : f32 + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %row0 = index.add %base, %lane : index + %valid = index.cmp ult, %row0, %rows : index + %row = scf.select %valid, %row0, %c0 : index + %input_noalias, %output_noalias = buffer.assume.noalias %input, %output : buffer, buffer + %input_view = buffer.view %input_noalias[%zero_offset] : buffer -> view<[%rows]x[%cols]xf32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%rows]xf32> + %total = scf.for %column = [%c0 to %cols step %c1](%running = %c0_f32 : f32) -> (f32) { + %value = view.load %input_view[%row, %column] : view<[%rows]x[%cols]xf32> -> f32 + %next = scalar.addf %running, %value : f32 + scf.yield %next : f32 + } + scf.if %valid { + view.store %total, %output_view[%row] : f32, view<[%rows]xf32> + } + kernel.return +} + +// ARGSORT: element i of a row goes to position rank(i), where rank counts the elements that +// sort before it (larger ones for descending order, smaller for ascending; ties by index). +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_argsort_rows_f32") @ggml_argsort_rows_f32(%column_count: index, %row_count: index, %descending: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %elements = index.mul %column_count, %row_count : index + %rounded = index.add %elements, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%column_count: index, %row_count: index, %descending: index, %input: buffer, %output: buffer) { + %cols = index.assume %column_count [range(%column_count, 1, 1024)] : index + %rows = index.assume %row_count [range(%row_count, 1, 16777216)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %linear0 = index.add %base, %lane : index + %elements = index.mul %cols, %rows : index + %valid = index.cmp ult, %linear0, %elements : index + %linear = scf.select %valid, %linear0, %c0 : index + %row = index.div %linear, %cols : index + %element = index.rem %linear, %cols : index + %want_descending = index.cmp ne, %descending, %c0 : index + %input_noalias, %output_noalias = buffer.assume.noalias %input, %output : buffer, buffer + %input_view = buffer.view %input_noalias[%zero_offset] : buffer -> view<[%rows]x[%cols]xf32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%rows]x[%cols]xi32> + %mine = view.load %input_view[%row, %element] : view<[%rows]x[%cols]xf32> -> f32 + %rank = scf.for %other = [%c0 to %cols step %c1](%count = %c0 : index) -> (index) { + %value = view.load %input_view[%row, %other] : view<[%rows]x[%cols]xf32> -> f32 + %greater = scalar.cmpf ogt, %value, %mine : f32 + %less = scalar.cmpf olt, %value, %mine : f32 + %equal = scalar.cmpf oeq, %value, %mine : f32 + %earlier = index.cmp ult, %other, %element : index + %tie_before = scalar.andi %equal, %earlier : i1 + %strict = scf.select %want_descending, %greater, %less : i1 + %before = scalar.ori %strict, %tie_before : i1 + %step = scf.select %before, %c1, %c0 : index + %next = index.add %count, %step : index + scf.yield %next : index + } + scf.if %valid { + %element_i32 = index.cast %element : index to i32 + view.store %element_i32, %output_view[%row, %rank] : i32, view<[%rows]x[%cols]xi32> + } + kernel.return +} + +// GET_ROWS within batches, for narrow rows: output[b][r][c] = input[b][ids[b][r]][c]. The ids may be a +// view with a row stride (the first k columns of an ARGSORT, as top-k selection makes). +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_get_rows_small_f32") @ggml_get_rows_small_f32(%width: index, %id_count: index, %batch_count: index, %source_rows: index, %id_stride: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %per_batch = index.mul %width, %id_count : index + %elements = index.mul %per_batch, %batch_count : index + %rounded = index.add %elements, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%width: index, %id_count: index, %batch_count: index, %source_rows: index, %id_stride: index, %input: buffer, %ids: buffer, %output: buffer) { + %w = index.assume %width [range(%width, 1, 65536)] : index + %r = index.assume %id_count [range(%id_count, 1, 65536)] : index + %b = index.assume %batch_count [range(%batch_count, 1, 65536)] : index + %s = index.assume %source_rows [range(%source_rows, 1, 65536)] : index + %t = index.assume %id_stride [range(%id_stride, 1, 1048576)] : index + %c0 = index.constant 0 : index + %c64 = index.constant 64 : index + %c0_i32 = scalar.constant 0 : i32 + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %linear0 = index.add %base, %lane : index + %per_batch = index.mul %w, %r : index + %elements = index.mul %per_batch, %b : index + %valid = index.cmp ult, %linear0, %elements : index + %linear = scf.select %valid, %linear0, %c0 : index + %batch = index.div %linear, %per_batch : index + %within = index.rem %linear, %per_batch : index + %slot = index.div %within, %w : index + %column = index.rem %within, %w : index + %source_total = index.mul %s, %b : index + %id_total = index.mul %r, %b : index + %output_total = index.mul %id_total, %w : index + %input_noalias, %ids_noalias, %output_noalias = buffer.assume.noalias %input, %ids, %output : buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%zero_offset] : buffer -> view<[%source_total]x[%w]xf32> + %ids_view = buffer.view %ids_noalias[%zero_offset] : buffer -> view<[%b]x[%t]xi32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%output_total]xf32> + %id_raw = view.load %ids_view[%batch, %slot] : view<[%b]x[%t]xi32> -> i32 + %id_nonnegative = scalar.cmpi sge, %id_raw, %c0_i32 : i32 + %id_safe_i32 = scf.select %id_nonnegative, %id_raw, %c0_i32 : i32 + %id0 = index.cast %id_safe_i32 : i32 to index + %in_range = index.cmp ult, %id0, %s : index + %id = scf.select %in_range, %id0, %c0 : index + %source_base = index.mul %batch, %s : index + %source_row = index.add %source_base, %id : index + %value = view.load %input_view[%source_row, %column] : view<[%source_total]x[%w]xf32> -> f32 + scf.if %valid { + view.store %value, %output_view[%linear] : f32, view<[%output_total]xf32> + } + kernel.return +} + +// CONT of a strided (permuted, transposed or sliced) F32 view into a packed output: one workitem +// per output element, reading the source through its four element strides. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_copy_strided_f32") @ggml_copy_strided_f32(%ne0: index, %ne1: index, %ne2: index, %ne3: index, %s0: index, %s1: index, %s2: index, %s3: index, %source_extent: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %e01 = index.mul %ne0, %ne1 : index + %e012 = index.mul %e01, %ne2 : index + %elements = index.mul %e012, %ne3 : index + %rounded = index.add %elements, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%ne0: index, %ne1: index, %ne2: index, %ne3: index, %s0: index, %s1: index, %s2: index, %s3: index, %source_extent: index, %input: buffer, %output: buffer) { + %n0 = index.assume %ne0 [range(%ne0, 1, 16777216)] : index + %n1 = index.assume %ne1 [range(%ne1, 1, 16777216)] : index + %n2 = index.assume %ne2 [range(%ne2, 1, 16777216)] : index + %n3 = index.assume %ne3 [range(%ne3, 1, 16777216)] : index + %extent = index.assume %source_extent [range(%source_extent, 1, 268435456)] : index + %c0 = index.constant 0 : index + %c64 = index.constant 64 : index + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %linear0 = index.add %base, %lane : index + %e01 = index.mul %n0, %n1 : index + %e012 = index.mul %e01, %n2 : index + %elements = index.mul %e012, %n3 : index + %valid = index.cmp ult, %linear0, %elements : index + %linear = scf.select %valid, %linear0, %c0 : index + %i0 = index.rem %linear, %n0 : index + %q0 = index.div %linear, %n0 : index + %i1 = index.rem %q0, %n1 : index + %q1 = index.div %q0, %n1 : index + %i2 = index.rem %q1, %n2 : index + %i3 = index.div %q1, %n2 : index + %o0 = index.mul %i0, %s0 : index + %o1 = index.mul %i1, %s1 : index + %o2 = index.mul %i2, %s2 : index + %o3 = index.mul %i3, %s3 : index + %o01 = index.add %o0, %o1 : index + %o23 = index.add %o2, %o3 : index + %source0 = index.add %o01, %o23 : index + %in_range = index.cmp ult, %source0, %extent : index + %source = scf.select %in_range, %source0, %c0 : index + %input_noalias, %output_noalias = buffer.assume.noalias %input, %output : buffer, buffer + %input_view = buffer.view %input_noalias[%zero_offset] : buffer -> view<[%extent]xf32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%elements]xf32> + %value = view.load %input_view[%source] : view<[%extent]xf32> -> f32 + scf.if %valid { + view.store %value, %output_view[%linear] : f32, view<[%elements]xf32> + } + kernel.return +} + +// NORM: y = (x - mean) / sqrt(var + eps) per row, var of the centered values (ggml's order). +// One 256-lane workgroup per row: lanes stride the columns, two workgroup reductions. +config.decl @ggml.norm_rows_f32.epsilon : f32 + +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_norm_rows_f32") @ggml_norm_rows_f32(%column_count: index, %row_count: index) { + %one = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%row_count, %one, %one) workgroup_size(%c256, %one, %one) : index +} launch(%column_count: index, %row_count: index, %input: buffer, %output: buffer) { + %cols = index.assume %column_count [range(%column_count, 1, 65536)] : index + %rows = index.assume %row_count [range(%row_count, 1, 16777216)] : index + %eps = config.get @ggml.norm_rows_f32.epsilon : f32 + %c256 = index.constant 256 : index + %c0_f32 = scalar.constant 0.0 : f32 + %zero_offset = index.constant 0 : offset + %row0 = kernel.workgroup.id : index + %row = index.assume %row0 [range(%row0, 0, 16777215)] : index + %lane = kernel.workitem.id : index + %input_noalias, %output_noalias = buffer.assume.noalias %input, %output : buffer, buffer + %input_view = buffer.view %input_noalias[%zero_offset] : buffer -> view<[%rows]x[%cols]xf32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%rows]x[%cols]xf32> + %sum = scf.for %column = [%lane to %cols step %c256](%acc = %c0_f32 : f32) -> (f32) { + %value = view.load %input_view[%row, %column] : view<[%rows]x[%cols]xf32> -> f32 + %next = scalar.addf %acc, %value : f32 + scf.yield %next : f32 + } + %row_sum = kernel.workgroup.reduce %sum : f32 + %cols_i32 = index.cast %cols : index to i32 + %cols_f32 = scalar.sitofp %cols_i32 : i32 to f32 + %mean = scalar.divf %row_sum, %cols_f32 : f32 + %squares = scf.for %column = [%lane to %cols step %c256](%acc = %c0_f32 : f32) -> (f32) { + %value = view.load %input_view[%row, %column] : view<[%rows]x[%cols]xf32> -> f32 + %centered = scalar.subf %value, %mean : f32 + %square = scalar.mulf %centered, %centered : f32 + %next = scalar.addf %acc, %square : f32 + scf.yield %next : f32 + } + %row_squares = kernel.workgroup.reduce %squares : f32 + %variance = scalar.divf %row_squares, %cols_f32 : f32 + %biased = scalar.addf %variance, %eps : f32 + %scale = scalar.rsqrtf %biased : f32 + %unused = scf.for %column = [%lane to %cols step %c256](%carry = %c0_f32 : f32) -> (f32) { + %value = view.load %input_view[%row, %column] : view<[%rows]x[%cols]xf32> -> f32 + %centered = scalar.subf %value, %mean : f32 + %normalized = scalar.mulf %centered, %scale : f32 + view.store %normalized, %output_view[%row, %column] : f32, view<[%rows]x[%cols]xf32> + scf.yield %carry : f32 + } + kernel.return +} + +// ADD / SUB / MUL / DIV with either input strided or broadcast (stride 0 on a broadcast dim) into a +// packed output: the cases the packed binary kernels do not take. One workitem per output element. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_binary_strided_f32") @ggml_binary_strided_f32(%ne0: index, %ne1: index, %ne2: index, %ne3: index, %a0: index, %a1: index, %a2: index, %a3: index, %b0: index, %b1: index, %b2: index, %b3: index, %a_extent: index, %b_extent: index, %op: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %e01 = index.mul %ne0, %ne1 : index + %e012 = index.mul %e01, %ne2 : index + %elements = index.mul %e012, %ne3 : index + %rounded = index.add %elements, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%ne0: index, %ne1: index, %ne2: index, %ne3: index, %a0: index, %a1: index, %a2: index, %a3: index, %b0: index, %b1: index, %b2: index, %b3: index, %a_extent: index, %b_extent: index, %op: index, %lhs: buffer, %rhs: buffer, %output: buffer) { + %n0 = index.assume %ne0 [range(%ne0, 1, 16777216)] : index + %n1 = index.assume %ne1 [range(%ne1, 1, 16777216)] : index + %n2 = index.assume %ne2 [range(%ne2, 1, 16777216)] : index + %n3 = index.assume %ne3 [range(%ne3, 1, 16777216)] : index + %ae = index.assume %a_extent [range(%a_extent, 1, 268435456)] : index + %be = index.assume %b_extent [range(%b_extent, 1, 268435456)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c64 = index.constant 64 : index + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %linear0 = index.add %base, %lane : index + %e01 = index.mul %n0, %n1 : index + %e012 = index.mul %e01, %n2 : index + %elements = index.mul %e012, %n3 : index + %valid = index.cmp ult, %linear0, %elements : index + %linear = scf.select %valid, %linear0, %c0 : index + %i0 = index.rem %linear, %n0 : index + %q0 = index.div %linear, %n0 : index + %i1 = index.rem %q0, %n1 : index + %q1 = index.div %q0, %n1 : index + %i2 = index.rem %q1, %n2 : index + %i3 = index.div %q1, %n2 : index + %x0 = index.mul %i0, %a0 : index + %x1 = index.mul %i1, %a1 : index + %x2 = index.mul %i2, %a2 : index + %x3 = index.mul %i3, %a3 : index + %x01 = index.add %x0, %x1 : index + %x23 = index.add %x2, %x3 : index + %xa0 = index.add %x01, %x23 : index + %xin = index.cmp ult, %xa0, %ae : index + %xa = scf.select %xin, %xa0, %c0 : index + %y0 = index.mul %i0, %b0 : index + %y1 = index.mul %i1, %b1 : index + %y2 = index.mul %i2, %b2 : index + %y3 = index.mul %i3, %b3 : index + %y01 = index.add %y0, %y1 : index + %y23 = index.add %y2, %y3 : index + %yb0 = index.add %y01, %y23 : index + %yin = index.cmp ult, %yb0, %be : index + %yb = scf.select %yin, %yb0, %c0 : index + %lhs_na, %rhs_na, %out_na = buffer.assume.noalias %lhs, %rhs, %output : buffer, buffer, buffer + %lhs_view = buffer.view %lhs_na[%zero_offset] : buffer -> view<[%ae]xf32> + %rhs_view = buffer.view %rhs_na[%zero_offset] : buffer -> view<[%be]xf32> + %out_view = buffer.view %out_na[%zero_offset] : buffer -> view<[%elements]xf32> + %a = view.load %lhs_view[%xa] : view<[%ae]xf32> -> f32 + %b = view.load %rhs_view[%yb] : view<[%be]xf32> -> f32 + %sum = scalar.addf %a, %b : f32 + %difference = scalar.subf %a, %b : f32 + %product = scalar.mulf %a, %b : f32 + %quotient = scalar.divf %a, %b : f32 + %is_add = index.cmp eq, %op, %c0 : index + %is_sub = index.cmp eq, %op, %c1 : index + %is_mul = index.cmp eq, %op, %c2 : index + %r0 = scf.select %is_mul, %product, %quotient : f32 + %r1 = scf.select %is_sub, %difference, %r0 : f32 + %result = scf.select %is_add, %sum, %r1 : f32 + scf.if %valid { + view.store %result, %out_view[%linear] : f32, view<[%elements]xf32> + } + kernel.return +} + +// CLAMP of a packed F32 tensor, standalone (the MoE router's clamp is fused elsewhere). +config.decl @ggml.clamp_f32.min : f32 + +config.decl @ggml.clamp_f32.max : f32 + +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_clamp_f32") @ggml_clamp_f32(%element_count: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %rounded = index.add %element_count, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%element_count: index, %input: buffer, %output: buffer) { + %n = index.assume %element_count [range(%element_count, 1, 268435456)] : index + %lo = config.get @ggml.clamp_f32.min : f32 + %hi = config.get @ggml.clamp_f32.max : f32 + %c0 = index.constant 0 : index + %c64 = index.constant 64 : index + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %i0 = index.add %base, %lane : index + %valid = index.cmp ult, %i0, %n : index + %i = scf.select %valid, %i0, %c0 : index + %input_noalias, %output_noalias = buffer.assume.noalias %input, %output : buffer, buffer + %input_view = buffer.view %input_noalias[%zero_offset] : buffer -> view<[%n]xf32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%n]xf32> + %value = view.load %input_view[%i] : view<[%n]xf32> -> f32 + %above = scalar.maxnumf %value, %lo : f32 + %clamped = scalar.minnumf %above, %hi : f32 + scf.if %valid { + view.store %clamped, %output_view[%i] : f32, view<[%n]xf32> + } + kernel.return +} + +// CLAMP in place (ggml_clamp's output is a view of its input): one buffer, read then write. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_clamp_inplace_f32") @ggml_clamp_inplace_f32(%element_count: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %rounded = index.add %element_count, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%element_count: index, %data: buffer) { + %n = index.assume %element_count [range(%element_count, 1, 268435456)] : index + %lo = config.get @ggml.clamp_f32.min : f32 + %hi = config.get @ggml.clamp_f32.max : f32 + %c0 = index.constant 0 : index + %c64 = index.constant 64 : index + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %i0 = index.add %base, %lane : index + %valid = index.cmp ult, %i0, %n : index + %i = scf.select %valid, %i0, %c0 : index + %view = buffer.view %data[%zero_offset] : buffer -> view<[%n]xf32> + %value = view.load %view[%i] : view<[%n]xf32> -> f32 + %above = scalar.maxnumf %value, %lo : f32 + %clamped = scalar.minnumf %above, %hi : f32 + scf.if %valid { + view.store %clamped, %view[%i] : f32, view<[%n]xf32> + } + kernel.return +} + +// CPY of an F32 tensor (any element strides) into a packed F16 destination, such as an attention +// mask converted for flash attention. One workitem per element. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_copy_strided_f32_f16") @ggml_copy_strided_f32_f16(%ne0: index, %ne1: index, %ne2: index, %ne3: index, %s0: index, %s1: index, %s2: index, %s3: index, %source_extent: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %e01 = index.mul %ne0, %ne1 : index + %e012 = index.mul %e01, %ne2 : index + %elements = index.mul %e012, %ne3 : index + %rounded = index.add %elements, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%ne0: index, %ne1: index, %ne2: index, %ne3: index, %s0: index, %s1: index, %s2: index, %s3: index, %source_extent: index, %input: buffer, %output: buffer) { + %n0 = index.assume %ne0 [range(%ne0, 1, 16777216)] : index + %n1 = index.assume %ne1 [range(%ne1, 1, 16777216)] : index + %n2 = index.assume %ne2 [range(%ne2, 1, 16777216)] : index + %n3 = index.assume %ne3 [range(%ne3, 1, 16777216)] : index + %extent = index.assume %source_extent [range(%source_extent, 1, 268435456)] : index + %c0 = index.constant 0 : index + %c64 = index.constant 64 : index + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %linear0 = index.add %base, %lane : index + %e01 = index.mul %n0, %n1 : index + %e012 = index.mul %e01, %n2 : index + %elements = index.mul %e012, %n3 : index + %valid = index.cmp ult, %linear0, %elements : index + %linear = scf.select %valid, %linear0, %c0 : index + %i0 = index.rem %linear, %n0 : index + %q0 = index.div %linear, %n0 : index + %i1 = index.rem %q0, %n1 : index + %q1 = index.div %q0, %n1 : index + %i2 = index.rem %q1, %n2 : index + %i3 = index.div %q1, %n2 : index + %o0 = index.mul %i0, %s0 : index + %o1 = index.mul %i1, %s1 : index + %o2 = index.mul %i2, %s2 : index + %o3 = index.mul %i3, %s3 : index + %o01 = index.add %o0, %o1 : index + %o23 = index.add %o2, %o3 : index + %source0 = index.add %o01, %o23 : index + %in_range = index.cmp ult, %source0, %extent : index + %source = scf.select %in_range, %source0, %c0 : index + %input_noalias, %output_noalias = buffer.assume.noalias %input, %output : buffer, buffer + %input_view = buffer.view %input_noalias[%zero_offset] : buffer -> view<[%extent]xf32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%elements]xf16> + %value = view.load %input_view[%source] : view<[%extent]xf32> -> f32 + %half = scalar.fptrunc %value : f32 to f16 + scf.if %valid { + view.store %half, %output_view[%linear] : f16, view<[%elements]xf16> + } + kernel.return +} + +// FLASH_ATTN_EXT for encoder-sized inputs with any Q/K/V/mask strides (F32 query, F16 key, value +// and mask, no ALiBi or softcap): the layouts the llama-shaped flash-attention kernels do not +// take, such as one contiguous block per head. One workgroup per (query token, head), one lane +// per value dimension, online softmax over the keys (running max and sum), so no lane has to +// wait for another; each lane computes the query-key dot products itself. +config.decl @ggml.attention_strided.scale : f32 + +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_attention_strided_f32_f16") @ggml_attention_strided_f32_f16(%qk_size: index, %v_size: index, %q_count: index, %kv_count: index, %head_count: index, %kv_head_count: index, %q_s1: index, %q_s2: index, %k_s1: index, %k_s2: index, %v_s1: index, %v_s2: index, %m_s1: index, %q_extent: index, %k_extent: index, %v_extent: index, %m_extent: index) { + %one = index.constant 1 : index + %groups = index.mul %q_count, %head_count : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%v_size, %one, %one) : index +} launch(%qk_size: index, %v_size: index, %q_count: index, %kv_count: index, %head_count: index, %kv_head_count: index, %q_s1: index, %q_s2: index, %k_s1: index, %k_s2: index, %v_s1: index, %v_s2: index, %m_s1: index, %q_extent: index, %k_extent: index, %v_extent: index, %m_extent: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %output: buffer) { + %d = index.assume %qk_size [range(%qk_size, 1, 1024)] : index + %dv = index.assume %v_size [range(%v_size, 1, 1024)] : index + %nq = index.assume %q_count [range(%q_count, 1, 65536)] : index + %nkv = index.assume %kv_count [range(%kv_count, 1, 65536)] : index + %nh = index.assume %head_count [range(%head_count, 1, 1024)] : index + %nhkv = index.assume %kv_head_count [range(%kv_head_count, 1, 1024)] : index + %qe = index.assume %q_extent [range(%q_extent, 1, 268435456)] : index + %ke = index.assume %k_extent [range(%k_extent, 1, 268435456)] : index + %ve = index.assume %v_extent [range(%v_extent, 1, 268435456)] : index + %me = index.assume %m_extent [range(%m_extent, 1, 268435456)] : index + %scale = config.get @ggml.attention_strided.scale : f32 + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %zero_offset = index.constant 0 : offset + %f0 = scalar.constant 0.0 : f32 + %lowest = scalar.constant -3.40282347e+38 : f32 + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %qi = index.div %group, %nh : index + %h = index.rem %group, %nh : index + %ratio = index.div %nh, %nhkv : index + %hkv = index.div %h, %ratio : index + %q_na, %k_na, %v_na, %m_na, %o_na = buffer.assume.noalias %query, %key, %value, %mask, %output : buffer, buffer, buffer, buffer, buffer + %q_view = buffer.view %q_na[%zero_offset] : buffer -> view<[%qe]xf32> + %k_view = buffer.view %k_na[%zero_offset] : buffer -> view<[%ke]xf16> + %v_view = buffer.view %v_na[%zero_offset] : buffer -> view<[%ve]xf16> + %m_view = buffer.view %m_na[%zero_offset] : buffer -> view<[%me]xf16> + %out_count0 = index.mul %dv, %nh : index + %out_count = index.mul %out_count0, %nq : index + %o_view = buffer.view %o_na[%zero_offset] : buffer -> view<[%out_count]xf32> + %qa = index.mul %qi, %q_s1 : index + %qb = index.mul %h, %q_s2 : index + %q_base = index.add %qa, %qb : index + %kb = index.mul %hkv, %k_s2 : index + %vb = index.mul %hkv, %v_s2 : index + %m_base = index.mul %qi, %m_s1 : index + %m_final, %l_final, %acc_final = scf.for %j = [%c0 to %nkv step %c1](%m = %lowest : f32, %l = %f0 : f32, %acc = %f0 : f32) -> (f32, f32, f32) { + %kj = index.mul %j, %k_s1 : index + %k_base = index.add %kj, %kb : index + %dot = scf.for %t = [%c0 to %d step %c1](%sum = %f0 : f32) -> (f32) { + %qt0 = index.add %q_base, %t : index + %qt = index.assume %qt0 [range(%qt0, 0, 268435455)] : index + %kt0 = index.add %k_base, %t : index + %kt = index.assume %kt0 [range(%kt0, 0, 268435455)] : index + %qv = view.load %q_view[%qt] : view<[%qe]xf32> -> f32 + %kv16 = view.load %k_view[%kt] : view<[%ke]xf16> -> f16 + %kv = scalar.extf %kv16 : f16 to f32 + %next = scalar.fmaf %qv, %kv, %sum : f32 + scf.yield %next : f32 + } + %mi0 = index.add %m_base, %j : index + %mi = index.assume %mi0 [range(%mi0, 0, 268435455)] : index + %mask16 = view.load %m_view[%mi] : view<[%me]xf16> -> f16 + %maskv = scalar.extf %mask16 : f16 to f32 + %scaled = scalar.mulf %dot, %scale : f32 + %score = scalar.addf %scaled, %maskv : f32 + %m_new = scalar.maxnumf %m, %score : f32 + %d_old = scalar.subf %m, %m_new : f32 + %corr = scalar.expf %d_old : f32 + %d_new = scalar.subf %score, %m_new : f32 + %p = scalar.expf %d_new : f32 + %l_scaled = scalar.mulf %l, %corr : f32 + %l_new = scalar.addf %l_scaled, %p : f32 + %vj = index.mul %j, %v_s1 : index + %v_row = index.add %vj, %vb : index + %vi0 = index.add %v_row, %lane : index + %vi = index.assume %vi0 [range(%vi0, 0, 268435455)] : index + %v16 = view.load %v_view[%vi] : view<[%ve]xf16> -> f16 + %vv = scalar.extf %v16 : f16 to f32 + %acc_scaled = scalar.mulf %acc, %corr : f32 + %acc_new = scalar.fmaf %p, %vv, %acc_scaled : f32 + scf.yield %m_new, %l_new, %acc_new : f32, f32, f32 + } + %result = scalar.divf %acc_final, %l_final : f32 + %oh = index.mul %h, %dv : index + %oq0 = index.mul %qi, %nh : index + %oq = index.mul %oq0, %dv : index + %o0 = index.add %oq, %oh : index + %oi0 = index.add %o0, %lane : index + %oi = index.assume %oi0 [range(%oi0, 0, 268435455)] : index + view.store %result, %o_view[%oi] : f32, view<[%out_count]xf32> + kernel.return +} + +// MUL_MAT of a small F16 weight [K, N] (any K, such as a classification head's) with F32 columns +// [K, T] (any element strides), into a packed [N, T]: one workitem per output, a loop over K. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_mul_mat_small_f16_f32") @ggml_mul_mat_small_f16_f32(%k_size: index, %n_size: index, %t_count: index, %w_s1: index, %x_s1: index, %w_extent: index, %x_extent: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %outputs = index.mul %n_size, %t_count : index + %rounded = index.add %outputs, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%k_size: index, %n_size: index, %t_count: index, %w_s1: index, %x_s1: index, %w_extent: index, %x_extent: index, %weight: buffer, %input: buffer, %output: buffer) { + %k = index.assume %k_size [range(%k_size, 1, 1048576)] : index + %n = index.assume %n_size [range(%n_size, 1, 1048576)] : index + %t = index.assume %t_count [range(%t_count, 1, 1048576)] : index + %we = index.assume %w_extent [range(%w_extent, 1, 268435456)] : index + %xe = index.assume %x_extent [range(%x_extent, 1, 268435456)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %f0 = scalar.constant 0.0 : f32 + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %o0 = index.add %base, %lane : index + %outputs = index.mul %n, %t : index + %valid = index.cmp ult, %o0, %outputs : index + %o = scf.select %valid, %o0, %c0 : index + %col = index.rem %o, %n : index + %tok = index.div %o, %n : index + %w_na, %x_na, %out_na = buffer.assume.noalias %weight, %input, %output : buffer, buffer, buffer + %w_view = buffer.view %w_na[%zero_offset] : buffer -> view<[%we]xf16> + %x_view = buffer.view %x_na[%zero_offset] : buffer -> view<[%xe]xf32> + %out_view = buffer.view %out_na[%zero_offset] : buffer -> view<[%outputs]xf32> + %w_base = index.mul %col, %w_s1 : index + %x_base = index.mul %tok, %x_s1 : index + %dot = scf.for %i = [%c0 to %k step %c1](%sum = %f0 : f32) -> (f32) { + %wi0 = index.add %w_base, %i : index + %wi = index.assume %wi0 [range(%wi0, 0, 268435455)] : index + %xi0 = index.add %x_base, %i : index + %xi = index.assume %xi0 [range(%xi0, 0, 268435455)] : index + %w16 = view.load %w_view[%wi] : view<[%we]xf16> -> f16 + %w = scalar.extf %w16 : f16 to f32 + %x = view.load %x_view[%xi] : view<[%xe]xf32> -> f32 + %next = scalar.fmaf %w, %x, %sum : f32 + scf.yield %next : f32 + } + scf.if %valid { + view.store %dot, %out_view[%o] : f32, view<[%outputs]xf32> + } + kernel.return +} + +// MUL_MAT of a small F32 weight [K, N] (any K, such as a classification head's) with F32 columns +// [K, T] (any element strides), into a packed [N, T]: one workitem per output, a loop over K. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_mul_mat_small_f32_f32") @ggml_mul_mat_small_f32_f32(%k_size: index, %n_size: index, %t_count: index, %w_s1: index, %x_s1: index, %w_extent: index, %x_extent: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %outputs = index.mul %n_size, %t_count : index + %rounded = index.add %outputs, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%k_size: index, %n_size: index, %t_count: index, %w_s1: index, %x_s1: index, %w_extent: index, %x_extent: index, %weight: buffer, %input: buffer, %output: buffer) { + %k = index.assume %k_size [range(%k_size, 1, 1048576)] : index + %n = index.assume %n_size [range(%n_size, 1, 1048576)] : index + %t = index.assume %t_count [range(%t_count, 1, 1048576)] : index + %we = index.assume %w_extent [range(%w_extent, 1, 268435456)] : index + %xe = index.assume %x_extent [range(%x_extent, 1, 268435456)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %f0 = scalar.constant 0.0 : f32 + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %o0 = index.add %base, %lane : index + %outputs = index.mul %n, %t : index + %valid = index.cmp ult, %o0, %outputs : index + %o = scf.select %valid, %o0, %c0 : index + %col = index.rem %o, %n : index + %tok = index.div %o, %n : index + %w_na, %x_na, %out_na = buffer.assume.noalias %weight, %input, %output : buffer, buffer, buffer + %w_view = buffer.view %w_na[%zero_offset] : buffer -> view<[%we]xf32> + %x_view = buffer.view %x_na[%zero_offset] : buffer -> view<[%xe]xf32> + %out_view = buffer.view %out_na[%zero_offset] : buffer -> view<[%outputs]xf32> + %w_base = index.mul %col, %w_s1 : index + %x_base = index.mul %tok, %x_s1 : index + %dot = scf.for %i = [%c0 to %k step %c1](%sum = %f0 : f32) -> (f32) { + %wi0 = index.add %w_base, %i : index + %wi = index.assume %wi0 [range(%wi0, 0, 268435455)] : index + %xi0 = index.add %x_base, %i : index + %xi = index.assume %xi0 [range(%xi0, 0, 268435455)] : index + %w = view.load %w_view[%wi] : view<[%we]xf32> -> f32 + %x = view.load %x_view[%xi] : view<[%xe]xf32> -> f32 + %next = scalar.fmaf %w, %x, %sum : f32 + scf.yield %next : f32 + } + scf.if %valid { + view.store %dot, %out_view[%o] : f32, view<[%outputs]xf32> + } + kernel.return +} + +// MUL_MAT of a Q8_0 weight [K, N] with K any multiple of 32 (the tiled kernels need 256), with F32 +// columns [K, T] (any element strides), into a packed [N, T]: one workitem per output, a loop over +// the 34-byte blocks (f16 scale, 32 int8 codes). +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_mul_mat_small_q8_0_f32") @ggml_mul_mat_small_q8_0_f32(%k_size: index, %n_size: index, %t_count: index, %w_s1: index, %x_s1: index, %w_extent: index, %x_extent: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %outputs = index.mul %n_size, %t_count : index + %rounded = index.add %outputs, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%k_size: index, %n_size: index, %t_count: index, %w_s1: index, %x_s1: index, %w_extent: index, %x_extent: index, %weight: buffer, %input: buffer, %output: buffer) { + %k = index.assume %k_size [range(%k_size, 32, 1048576)] : index + %n = index.assume %n_size [range(%n_size, 1, 1048576)] : index + %t = index.assume %t_count [range(%t_count, 1, 1048576)] : index + %xe = index.assume %x_extent [range(%x_extent, 1, 268435456)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %f0 = scalar.constant 0.0 : f32 + %zero_offset = index.constant 0 : offset + %block_bytes = index.constant 34 : offset + %code_offset = index.constant 2 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %o0 = index.add %base, %lane : index + %outputs = index.mul %n, %t : index + %valid = index.cmp ult, %o0, %outputs : index + %o = scf.select %valid, %o0, %c0 : index + %col = index.rem %o, %n : index + %tok = index.div %o, %n : index + %x_na, %out_na = buffer.assume.noalias %input, %output : buffer, buffer + %x_view = buffer.view %x_na[%zero_offset] : buffer -> view<[%xe]xf32> + %out_view = buffer.view %out_na[%zero_offset] : buffer -> view<[%outputs]xf32> + %blocks = index.div %k, %c32 : index + %row_block = index.mul %col, %w_s1 : index + %x_base = index.mul %tok, %x_s1 : index + %dot = scf.for %b = [%c0 to %blocks step %c1](%sum = %f0 : f32) -> (f32) { + %gb = index.add %row_block, %b : index + %block_base = index.scale %gb, %block_bytes : index, offset -> offset + %code_base = index.add %block_base, %code_offset : offset + %d_view = buffer.view %weight[%block_base] : buffer -> view<1xf16> + %code_view = buffer.view %weight[%code_base] : buffer -> view<32xi8> + %d16 = view.load %d_view[%c0] : view<1xf16> -> f16 + %d = scalar.extf %d16 : f16 to f32 + %xb = index.mul %b, %c32 : index + %xrow = index.add %x_base, %xb : index + %partial = scf.for %i = [%c0 to %c32 step %c1](%acc = %f0 : f32) -> (f32) { + %code = view.load %code_view[%i] : view<32xi8> -> i8 + %q = scalar.sitofp %code : i8 to f32 + %xi0 = index.add %xrow, %i : index + %xi = index.assume %xi0 [range(%xi0, 0, 268435455)] : index + %x = view.load %x_view[%xi] : view<[%xe]xf32> -> f32 + %next = scalar.fmaf %q, %x, %acc : f32 + scf.yield %next : f32 + } + %next_sum = scalar.fmaf %d, %partial, %sum : f32 + scf.yield %next_sum : f32 + } + scf.if %valid { + view.store %dot, %out_view[%o] : f32, view<[%outputs]xf32> + } + kernel.return +} + +// Q8_0 MUL_MAT, one 64-lane workgroup per output: the lanes split the K blocks (reads along each +// weight row stay together), then one workgroup reduction. For K = 2624 (ModernBERT's FFN down +// projection) the per-output kernel above reads 64 different rows per wave. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_mul_mat_rows_q8_0_f32") @ggml_mul_mat_rows_q8_0_f32(%k_size: index, %n_size: index, %t_count: index, %w_s1: index, %x_s1: index, %w_extent: index, %x_extent: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %outputs = index.mul %n_size, %t_count : index + kernel.launch.config workgroups(%outputs, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%k_size: index, %n_size: index, %t_count: index, %w_s1: index, %x_s1: index, %w_extent: index, %x_extent: index, %weight: buffer, %input: buffer, %output: buffer) { + %k = index.assume %k_size [range(%k_size, 32, 1048576)] : index + %n = index.assume %n_size [range(%n_size, 1, 1048576)] : index + %t = index.assume %t_count [range(%t_count, 1, 1048576)] : index + %xe = index.assume %x_extent [range(%x_extent, 1, 268435456)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %f0 = scalar.constant 0.0 : f32 + %zero_offset = index.constant 0 : offset + %block_bytes = index.constant 34 : offset + %code_offset = index.constant 2 : offset + %o = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %col = index.rem %o, %n : index + %tok = index.div %o, %n : index + %x_na, %out_na = buffer.assume.noalias %input, %output : buffer, buffer + %x_view = buffer.view %x_na[%zero_offset] : buffer -> view<[%xe]xf32> + %outputs = index.mul %n, %t : index + %out_view = buffer.view %out_na[%zero_offset] : buffer -> view<[%outputs]xf32> + %blocks = index.div %k, %c32 : index + %row_block = index.mul %col, %w_s1 : index + %x_base = index.mul %tok, %x_s1 : index + %mine = scf.for %b = [%lane to %blocks step %c64](%sum = %f0 : f32) -> (f32) { + %gb = index.add %row_block, %b : index + %block_base = index.scale %gb, %block_bytes : index, offset -> offset + %code_base = index.add %block_base, %code_offset : offset + %d_view = buffer.view %weight[%block_base] : buffer -> view<1xf16> + %code_view = buffer.view %weight[%code_base] : buffer -> view<32xi8> + %d16 = view.load %d_view[%c0] : view<1xf16> -> f16 + %d = scalar.extf %d16 : f16 to f32 + %xb = index.mul %b, %c32 : index + %xrow = index.add %x_base, %xb : index + %partial = scf.for %i = [%c0 to %c32 step %c1](%acc = %f0 : f32) -> (f32) { + %code = view.load %code_view[%i] : view<32xi8> -> i8 + %q = scalar.sitofp %code : i8 to f32 + %xi0 = index.add %xrow, %i : index + %xi = index.assume %xi0 [range(%xi0, 0, 268435455)] : index + %x = view.load %x_view[%xi] : view<[%xe]xf32> -> f32 + %next = scalar.fmaf %q, %x, %acc : f32 + scf.yield %next : f32 + } + %next_sum = scalar.fmaf %d, %partial, %sum : f32 + scf.yield %next_sum : f32 + } + %total = kernel.workgroup.reduce %mine : f32 + %first = index.cmp eq, %lane, %c0 : index + scf.if %first { + view.store %total, %out_view[%o] : f32, view<[%outputs]xf32> + } + kernel.return +} + +// The same attention with each score computed once: one 64-lane workgroup per (query token, head) +// stages the query row in workgroup memory, each lane scores its keys (j = lane, lane + 64, ...) +// into a workgroup score row, workgroup max and sum reductions give the softmax, and the lanes +// then take the value dimensions (reads along each value row). Up to 2048 keys, head sizes <= 256. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_attention_rows_f32_f16") @ggml_attention_rows_f32_f16(%qk_size: index, %v_size: index, %q_count: index, %kv_count: index, %head_count: index, %kv_head_count: index, %q_s1: index, %q_s2: index, %k_s1: index, %k_s2: index, %v_s1: index, %v_s2: index, %m_s1: index, %q_extent: index, %k_extent: index, %v_extent: index, %m_extent: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %groups = index.mul %q_count, %head_count : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%qk_size: index, %v_size: index, %q_count: index, %kv_count: index, %head_count: index, %kv_head_count: index, %q_s1: index, %q_s2: index, %k_s1: index, %k_s2: index, %v_s1: index, %v_s2: index, %m_s1: index, %q_extent: index, %k_extent: index, %v_extent: index, %m_extent: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %output: buffer) { + %d = index.assume %qk_size [range(%qk_size, 1, 256)] : index + %dv = index.assume %v_size [range(%v_size, 1, 256)] : index + %nq = index.assume %q_count [range(%q_count, 1, 65536)] : index + %nkv = index.assume %kv_count [range(%kv_count, 1, 2048)] : index + %nh = index.assume %head_count [range(%head_count, 1, 1024)] : index + %nhkv = index.assume %kv_head_count [range(%kv_head_count, 1, 1024)] : index + %qe = index.assume %q_extent [range(%q_extent, 1, 268435456)] : index + %ke = index.assume %k_extent [range(%k_extent, 1, 268435456)] : index + %ve = index.assume %v_extent [range(%v_extent, 1, 268435456)] : index + %me = index.assume %m_extent [range(%m_extent, 1, 268435456)] : index + %scale = config.get @ggml.attention_strided.scale : f32 + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %zero = index.constant 0 : offset + %f0 = scalar.constant 0.0 : f32 + %lowest = scalar.constant -3.40282347e+38 : f32 + %q_bytes = index.constant 1024 : offset + %s_bytes = index.constant 8192 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %qi = index.div %group, %nh : index + %h = index.rem %group, %nh : index + %ratio = index.div %nh, %nhkv : index + %hkv = index.div %h, %ratio : index + %q_na, %k_na, %v_na, %m_na, %o_na = buffer.assume.noalias %query, %key, %value, %mask, %output : buffer, buffer, buffer, buffer, buffer + %q_view = buffer.view %q_na[%zero] : buffer -> view<[%qe]xf32> + %k_view = buffer.view %k_na[%zero] : buffer -> view<[%ke]xf16> + %v_view = buffer.view %v_na[%zero] : buffer -> view<[%ve]xf16> + %m_view = buffer.view %m_na[%zero] : buffer -> view<[%me]xf16> + %out_count0 = index.mul %dv, %nh : index + %out_count = index.mul %out_count0, %nq : index + %o_view = buffer.view %o_na[%zero] : buffer -> view<[%out_count]xf32> + %q_shared = buffer.alloca align(16) %q_bytes : buffer + %q_row = buffer.view %q_shared[%zero] : buffer -> view<256xf32> + %s_shared = buffer.alloca align(16) %s_bytes : buffer + %s_row = buffer.view %s_shared[%zero] : buffer -> view<2048xf32> + %qa = index.mul %qi, %q_s1 : index + %qb = index.mul %h, %q_s2 : index + %q_base = index.add %qa, %qb : index + %u0 = scf.for %t = [%lane to %d step %c64](%carry0 = %f0 : f32) -> (f32) { + %qt0 = index.add %q_base, %t : index + %qt = index.assume %qt0 [range(%qt0, 0, 268435455)] : index + %qv = view.load %q_view[%qt] : view<[%qe]xf32> -> f32 + view.store %qv, %q_row[%t] : f32, view<256xf32> + scf.yield %carry0 : f32 + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %kb = index.mul %hkv, %k_s2 : index + %m_base = index.mul %qi, %m_s1 : index + %lane_max = scf.for %j = [%lane to %nkv step %c64](%mx = %lowest : f32) -> (f32) { + %kj = index.mul %j, %k_s1 : index + %k_base = index.add %kj, %kb : index + %dot = scf.for %t = [%c0 to %d step %c1](%sum = %f0 : f32) -> (f32) { + %kt0 = index.add %k_base, %t : index + %kt = index.assume %kt0 [range(%kt0, 0, 268435455)] : index + %qv = view.load %q_row[%t] : view<256xf32> -> f32 + %kv16 = view.load %k_view[%kt] : view<[%ke]xf16> -> f16 + %kv = scalar.extf %kv16 : f16 to f32 + %next = scalar.fmaf %qv, %kv, %sum : f32 + scf.yield %next : f32 + } + %mi0 = index.add %m_base, %j : index + %mi = index.assume %mi0 [range(%mi0, 0, 268435455)] : index + %mask16 = view.load %m_view[%mi] : view<[%me]xf16> -> f16 + %maskv = scalar.extf %mask16 : f16 to f32 + %scaled = scalar.mulf %dot, %scale : f32 + %score = scalar.addf %scaled, %maskv : f32 + view.store %score, %s_row[%j] : f32, view<2048xf32> + %next_max = scalar.maxnumf %mx, %score : f32 + scf.yield %next_max : f32 + } + %row_max = kernel.workgroup.reduce %lane_max : f32 + %lane_sum = scf.for %j = [%lane to %nkv step %c64](%acc = %f0 : f32) -> (f32) { + %score = view.load %s_row[%j] : view<2048xf32> -> f32 + %shifted = scalar.subf %score, %row_max : f32 + %p = scalar.expf %shifted : f32 + view.store %p, %s_row[%j] : f32, view<2048xf32> + %next = scalar.addf %acc, %p : f32 + scf.yield %next : f32 + } + %row_sum = kernel.workgroup.reduce %lane_sum : f32 + kernel.barrier scope(workgroup) ordering(acq_rel) + %vb = index.mul %hkv, %v_s2 : index + %oh = index.mul %h, %dv : index + %oq0 = index.mul %qi, %nh : index + %oq = index.mul %oq0, %dv : index + %o_base = index.add %oq, %oh : index + %u1 = scf.for %c = [%lane to %dv step %c64](%carry1 = %f0 : f32) -> (f32) { + %acc_v = scf.for %j = [%c0 to %nkv step %c1](%acc = %f0 : f32) -> (f32) { + %p = view.load %s_row[%j] : view<2048xf32> -> f32 + %vj = index.mul %j, %v_s1 : index + %v_row = index.add %vj, %vb : index + %vi0 = index.add %v_row, %c : index + %vi = index.assume %vi0 [range(%vi0, 0, 268435455)] : index + %v16 = view.load %v_view[%vi] : view<[%ve]xf16> -> f16 + %vv = scalar.extf %v16 : f16 to f32 + %next = scalar.fmaf %p, %vv, %acc : f32 + scf.yield %next : f32 + } + %result = scalar.divf %acc_v, %row_sum : f32 + %oi0 = index.add %o_base, %c : index + %oi = index.assume %oi0 [range(%oi0, 0, 268435455)] : index + view.store %result, %o_view[%oi] : f32, view<[%out_count]xf32> + scf.yield %carry1 : f32 + } + kernel.return +} + +// Rotate-half RoPE as graph compilers lower it (x * cos + concat(-x[half:], x[:half]) * sin, eight +// nodes), in one pass: x packed [d, T, H], cos and sin [d, T] (row strides given), broadcast over +// heads. One workitem per element. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_rope_rotate_half_f32") @ggml_rope_rotate_half_f32(%ne0: index, %ne1: index, %ne2: index, %cos_s1: index, %sin_s1: index, %cos_extent: index, %sin_extent: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %e01 = index.mul %ne0, %ne1 : index + %elements = index.mul %e01, %ne2 : index + %rounded = index.add %elements, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%ne0: index, %ne1: index, %ne2: index, %cos_s1: index, %sin_s1: index, %cos_extent: index, %sin_extent: index, %input: buffer, %cos: buffer, %sin: buffer, %output: buffer) { + %d = index.assume %ne0 [range(%ne0, 2, 4096)] : index + %t = index.assume %ne1 [range(%ne1, 1, 1048576)] : index + %h = index.assume %ne2 [range(%ne2, 1, 4096)] : index + %ce = index.assume %cos_extent [range(%cos_extent, 1, 268435456)] : index + %se = index.assume %sin_extent [range(%sin_extent, 1, 268435456)] : index + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c64 = index.constant 64 : index + %zero = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %lin0 = index.add %base, %lane : index + %e01 = index.mul %d, %t : index + %elements = index.mul %e01, %h : index + %valid = index.cmp ult, %lin0, %elements : index + %lin = scf.select %valid, %lin0, %c0 : index + %i0 = index.rem %lin, %d : index + %q0 = index.div %lin, %d : index + %ti = index.rem %q0, %t : index + %half = index.div %d, %c2 : index + %low = index.cmp ult, %i0, %half : index + %up = index.add %lin, %half : index + %down0 = index.sub %lin, %half : index + %down = scf.select %low, %lin, %down0 : index + %partner0 = scf.select %low, %up, %down : index + %partner_in = index.cmp ult, %partner0, %elements : index + %partner = scf.select %partner_in, %partner0, %c0 : index + %x_na, %c_na, %s_na, %o_na = buffer.assume.noalias %input, %cos, %sin, %output : buffer, buffer, buffer, buffer + %x_view = buffer.view %x_na[%zero] : buffer -> view<[%elements]xf32> + %c_view = buffer.view %c_na[%zero] : buffer -> view<[%ce]xf32> + %s_view = buffer.view %s_na[%zero] : buffer -> view<[%se]xf32> + %o_view = buffer.view %o_na[%zero] : buffer -> view<[%elements]xf32> + %x = view.load %x_view[%lin] : view<[%elements]xf32> -> f32 + %other = view.load %x_view[%partner] : view<[%elements]xf32> -> f32 + %neg_other = scalar.negf %other : f32 + %rot = scf.select %low, %neg_other, %other : f32 + %cr = index.mul %ti, %cos_s1 : index + %ci0 = index.add %cr, %i0 : index + %ci_in = index.cmp ult, %ci0, %ce : index + %ci = scf.select %ci_in, %ci0, %c0 : index + %sr = index.mul %ti, %sin_s1 : index + %si0 = index.add %sr, %i0 : index + %si_in = index.cmp ult, %si0, %se : index + %si = scf.select %si_in, %si0, %c0 : index + %cv = view.load %c_view[%ci] : view<[%ce]xf32> -> f32 + %sv = view.load %s_view[%si] : view<[%se]xf32> -> f32 + %xc = scalar.mulf %x, %cv : f32 + %result = scalar.fmaf %rot, %sv, %xc : f32 + scf.if %valid { + view.store %result, %o_view[%lin] : f32, view<[%elements]xf32> + } + kernel.return +} + +// GEGLU lowered to three nodes (CONT of the gate half, GELU, MUL by the up half): gelu(a) * b with +// a and b strided F32 views [n, T] (rows a_s1 / b_s1 apart), into a packed [n, T]. GELU in ggml's +// tanh form, 0.5 x (1 + tanh(sqrt(2/pi) x (1 + 0.044715 x^2))), tanh(z) = 1 - 2 / (exp(2z) + 1). +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_geglu_strided_f32") @ggml_geglu_strided_f32(%n_size: index, %t_count: index, %a_s1: index, %b_s1: index, %a_extent: index, %b_extent: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %elements = index.mul %n_size, %t_count : index + %rounded = index.add %elements, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%n_size: index, %t_count: index, %a_s1: index, %b_s1: index, %a_extent: index, %b_extent: index, %gate: buffer, %up: buffer, %output: buffer) { + %n = index.assume %n_size [range(%n_size, 1, 1048576)] : index + %t = index.assume %t_count [range(%t_count, 1, 1048576)] : index + %ae = index.assume %a_extent [range(%a_extent, 1, 268435456)] : index + %be = index.assume %b_extent [range(%b_extent, 1, 268435456)] : index + %c0 = index.constant 0 : index + %c64 = index.constant 64 : index + %zero = index.constant 0 : offset + %half = scalar.constant 0.5 : f32 + %one_f = scalar.constant 1.0 : f32 + %two = scalar.constant 2.0 : f32 + %k0 = scalar.constant 0.7978845608 : f32 + %k1 = scalar.constant 0.044715 : f32 + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %lin0 = index.add %base, %lane : index + %elements = index.mul %n, %t : index + %valid = index.cmp ult, %lin0, %elements : index + %lin = scf.select %valid, %lin0, %c0 : index + %i = index.rem %lin, %n : index + %row = index.div %lin, %n : index + %ar = index.mul %row, %a_s1 : index + %ai0 = index.add %ar, %i : index + %ai_in = index.cmp ult, %ai0, %ae : index + %ai = scf.select %ai_in, %ai0, %c0 : index + %br = index.mul %row, %b_s1 : index + %bi0 = index.add %br, %i : index + %bi_in = index.cmp ult, %bi0, %be : index + %bi = scf.select %bi_in, %bi0, %c0 : index + %a_na, %b_na, %o_na = buffer.assume.noalias %gate, %up, %output : buffer, buffer, buffer + %a_view = buffer.view %a_na[%zero] : buffer -> view<[%ae]xf32> + %b_view = buffer.view %b_na[%zero] : buffer -> view<[%be]xf32> + %o_view = buffer.view %o_na[%zero] : buffer -> view<[%elements]xf32> + %x = view.load %a_view[%ai] : view<[%ae]xf32> -> f32 + %u = view.load %b_view[%bi] : view<[%be]xf32> -> f32 + %x2 = scalar.mulf %x, %x : f32 + %poly = scalar.fmaf %k1, %x2, %one_f : f32 + %inner0 = scalar.mulf %x, %poly : f32 + %inner = scalar.mulf %k0, %inner0 : f32 + %twice = scalar.mulf %inner, %two : f32 + %e = scalar.expf %twice : f32 + %e1 = scalar.addf %e, %one_f : f32 + %frac = scalar.divf %two, %e1 : f32 + %tanh = scalar.subf %one_f, %frac : f32 + %onep = scalar.addf %one_f, %tanh : f32 + %hx = scalar.mulf %half, %x : f32 + %g = scalar.mulf %hx, %onep : f32 + %result = scalar.mulf %g, %u : f32 + scf.if %valid { + view.store %result, %o_view[%lin] : f32, view<[%elements]xf32> + } + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/softplus_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/softplus_f32.loom new file mode 100644 index 000000000000..a15baea5ffca --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/softplus_f32.loom @@ -0,0 +1,94 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// Element-wise softplus for contiguous F32 values, as ggml's CPU backend defines it: +// y = x > 20 ? x : log(1 + exp(x)) +// The generic ggml_unary_f32 kernel does not cover softplus. Qwen3.5/3.8 gated delta-net layers +// apply it to their gate; one-sequence decode fuses it into the delta-net kernels, but batches +// with several sequences (llama-server --parallel) leave it standalone. + +amdgpu.target @ggml_softplus_f32_gfx11_wave64 {subgroup_size = 64} + +kernel.def target(@ggml_softplus_f32_gfx11_wave64) export("ggml_softplus_f32") @ggml_softplus_f32(%element_count: index) { + %one = index.constant 1 : index + %c256 = index.constant 256 : index + %c255 = index.constant 255 : index + %rounded = index.add %element_count, %c255 : index + %groups = index.div %rounded, %c256 : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%c256, %one, %one) : index +} launch(%element_count: index, %input: buffer, %output: buffer) { + %count = index.assume %element_count [range(%element_count, 1, 134217728)] : index + %wg = kernel.workgroup.id : index + %item = kernel.workitem.id : index + %c256 = index.constant 256 : index + %base = index.mul %wg, %c256 : index + %linear0 = index.add %base, %item : index + %linear = index.assume %linear0 [range(%linear0, 0, 134217983)] : index + %in_bounds = index.cmp ult, %linear, %count : index + %zero_offset = index.constant 0 : offset + %one_f = scalar.constant 1.0 : f32 + %threshold = scalar.constant 20.0 : f32 + %input_view = buffer.view %input[%zero_offset] : buffer -> view<[%count]xf32> + %output_view = buffer.view %output[%zero_offset] : buffer -> view<[%count]xf32> + scf.if %in_bounds { + %x = view.load %input_view[%linear] : view<[%count]xf32> -> f32 + %large = scalar.cmpf ogt, %x, %threshold : f32 + %e = scalar.expf %x : f32 + %e1 = scalar.addf %one_f, %e : f32 + %l = scalar.logf %e1 : f32 + %y = scf.select %large, %x, %l : f32 + view.store %y, %output_view[%linear] : f32, view<[%count]xf32> + } + kernel.return +} + +// Reference for the case: the overflow-safe form max(x, 0) + log(1 + exp(-|x|)). +kernel.def target(@ggml_softplus_f32_gfx11_wave64) export("ggml_softplus_reference_f32") @ggml_softplus_reference_f32(%element_count: index) { + %one = index.constant 1 : index + kernel.launch.config workgroups(%element_count, %one, %one) workgroup_size(%one, %one, %one) : index +} launch(%element_count: index, %input: buffer, %output: buffer) { + %count = index.assume %element_count [range(%element_count, 1, 134217728)] : index + %i0 = kernel.workgroup.id : index + %i, %bound = index.assume %i0, %count [lt(%i0, %count)] : index, index + %zero_offset = index.constant 0 : offset + %zero = scalar.constant 0.0 : f32 + %one_f = scalar.constant 1.0 : f32 + %iv = buffer.view %input[%zero_offset] : buffer -> view<[%bound]xf32> + %ov = buffer.view %output[%zero_offset] : buffer -> view<[%bound]xf32> + %x = view.load %iv[%i] : view<[%bound]xf32> -> f32 + %pos = scalar.cmpf ogt, %x, %zero : f32 + %relu = scf.select %pos, %x, %zero : f32 + %neg = scalar.negf %x : f32 + %abs = scf.select %pos, %x, %neg : f32 + %nabs = scalar.negf %abs : f32 + %e = scalar.expf %nabs : f32 + %e1 = scalar.addf %one_f, %e : f32 + %l = scalar.logf %e1 : f32 + %y = scalar.addf %relu, %l : f32 + view.store %y, %ov[%i] : f32, view<[%bound]xf32> + kernel.return +} + +// iota -24..40 (step 0.5): both branches of the x > 20 cut. +check.case public @ggml_softplus_f32_case { + %count = check.literal value(128) : index + %input = check.generate.iota offset(-24.0) step(0.5) : tensor<128xf32> + %output = check.generate.fill value(-7.0) : tensor<128xf32> + %expected = check.generate.fill value(7.0) : tensor<128xf32> + kernel.launch @ggml_softplus_reference_f32[%count](%count, %input, %expected) : [index](index, tensor<128xf32>, tensor<128xf32>) + kernel.launch @ggml_softplus_f32[%count](%count, %input, %output) : [index](index, tensor<128xf32>, tensor<128xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-6) rtol(1.0e-5) nan(same) : tensor<128xf32> + check.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/ssm_conv_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/ssm_conv_f32.loom new file mode 100644 index 000000000000..384a78804cf6 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/ssm_conv_f32.loom @@ -0,0 +1,1191 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +config.decl @llm.ssm_conv.snapshot.d_conv : %value: index where [range(%value, 4, 4)] + +config.decl @llm.ssm_conv.snapshot.d_inner : %value: index where [range(%value, 8192, 10240), mul(%value, 32)] + +config.decl @llm.ssm_conv.snapshot.n_t : %value: index where [range(%value, 512, 512)] + +config.decl @llm.ssm_conv.snapshot.n_s : %value: index where [range(%value, 1, 1)] + +config.decl @llm.ssm_conv.snapshot.state_row_stride : %value: index where [range(%value, 1, 1)] + +config.decl @llm.ssm_conv.snapshot.state_channel_stride : %value: index where [range(%value, 3, 3)] + +config.decl @llm.ssm_conv.snapshot.x_row_stride : %value: index where [range(%value, 1, 16777216)] + +config.decl @llm.ssm_conv.snapshot.dst_row_stride : %value: index where [range(%value, 1, 16777216)] + +config.decl @llm.ssm_conv.snapshot.cache_row_stride : %value: index where [range(%value, 1, 1)] + +config.decl @llm.ssm_conv.snapshot.cache_channel_stride : %value: index where [range(%value, 3, 3)] + +config.decl @llm.ssm_conv.snapshot.workgroup_size : %value: index where [range(%value, 32, 1024), mul(%value, 32)] + +kernel.def export("llm_ssm_conv_snapshot_window_tail_f32") @llm_ssm_conv_snapshot_window_tail_f32() { + %unit = index.constant 1 : index + %cneg1 = index.constant -1 : index + %x_snapshot_rows = index.constant 61 : index + %d_conv = config.get @llm.ssm_conv.snapshot.d_conv : index + %d_inner = config.get @llm.ssm_conv.snapshot.d_inner : index + %n_s = config.get @llm.ssm_conv.snapshot.n_s : index + %wg = config.get @llm.ssm_conv.snapshot.workgroup_size : index + %carry = index.add %d_conv, %cneg1 : index + %snapshot_rows = index.add %carry, %x_snapshot_rows : index + %rows = index.mul %snapshot_rows, %n_s : index + %total = index.mul %rows, %d_inner : index + %rounding = index.add %wg, %cneg1 : index + %rounded = index.add %total, %rounding : index + %groups = index.div %rounded, %wg : index + kernel.launch.config workgroups(%groups, %unit, %unit) workgroup_size(%wg, %unit, %unit) : index +} launch(%state: buffer, %x: buffer, %dst: buffer, %cache: buffer) { + %base = index.constant 0 : offset + %cneg1 = index.constant -1 : index + %x_snapshot_rows = index.constant 61 : index + + %d_conv = config.get @llm.ssm_conv.snapshot.d_conv : index + %d_inner = config.get @llm.ssm_conv.snapshot.d_inner : index + %n_t = config.get @llm.ssm_conv.snapshot.n_t : index + %n_s = config.get @llm.ssm_conv.snapshot.n_s : index + %srs = config.get @llm.ssm_conv.snapshot.state_row_stride : index + %scs = config.get @llm.ssm_conv.snapshot.state_channel_stride : index + %xrs = config.get @llm.ssm_conv.snapshot.x_row_stride : index + %drs = config.get @llm.ssm_conv.snapshot.dst_row_stride : index + %crs = config.get @llm.ssm_conv.snapshot.cache_row_stride : index + %ccs = config.get @llm.ssm_conv.snapshot.cache_channel_stride : index + %wg = config.get @llm.ssm_conv.snapshot.workgroup_size : index + + %di = index.assume %d_inner [range(%d_inner, 8192, 10240), mul(%d_inner, 32)] : index + %nt = index.assume %n_t [range(%n_t, 512, 512)] : index + %ns = index.assume %n_s [range(%n_s, 1, 1)] : index + %srsb = index.assume %srs [range(%srs, 1, 1)] : index + %scsb = index.assume %scs [range(%scs, 3, 3)] : index + %xrsb = index.assume %xrs [range(%xrs, 1, 16777216)] : index + %drsb = index.assume %drs [range(%drs, 1, 16777216)] : index + %crsb = index.assume %crs [range(%crs, 1, 1)] : index + %ccsb = index.assume %ccs [range(%ccs, 3, 3)] : index + %carry0 = index.add %d_conv, %cneg1 : index + %carry = index.assume %carry0 [range(%carry0, 3, 3)] : index + %snapshot_rows0 = index.add %carry, %x_snapshot_rows : index + %snapshot_rows = index.assume %snapshot_rows0 [range(%snapshot_rows0, 64, 64)] : index + + %rows0 = index.mul %snapshot_rows, %ns : index + %total0 = index.mul %rows0, %di : index + %total = index.assume %total0 [range(%total0, 1, 1073741823)] : index + + %group0 = kernel.workgroup.id : index + %lane0 = kernel.workitem.id : index + %group = index.assume %group0 [range(%group0, 0, 16777215)] : index + %lane = index.assume %lane0 [range(%lane0, 0, 1023)] : index + %off0 = index.mul %group, %wg : index + %idx0 = index.add %off0, %lane : index + %idx = index.assume %idx0 [range(%idx0, 0, 1073741823)] : index + %in_bounds = index.cmp ult, %idx, %total : index + + %state_g = buffer.assume.memory_space %state : buffer + %x_g = buffer.assume.memory_space %x : buffer + %dst_g = buffer.assume.memory_space %dst : buffer + %c_g = buffer.assume.memory_space %cache : buffer + %state_na, %x_na, %dst_na, %c_na = buffer.assume.noalias %state_g, %x_g, %dst_g, %c_g : buffer, buffer, buffer, buffer + %state_view = buffer.view %state_na[%base] : buffer -> view<1073741824xf32> + %x_view = buffer.view %x_na[%base] : buffer -> view<1073741824xf32> + %dst_view = buffer.view %dst_na[%base] : buffer -> view<1073741824xf32> + %c_view = buffer.view %c_na[%base] : buffer -> view<1073741824xf32> + + scf.if %in_bounds { + // Consecutive lanes take consecutive channels, so both sides are contiguous + // bursts across the wave. + %chan = index.rem %idx, %di : index + %flat_row = index.div %idx, %di : index + %snapshot_row = index.rem %flat_row, %snapshot_rows : index + %seq = index.div %flat_row, %snapshot_rows : index + + // All 64 snapshot rows are written to their natural CONCAT row. + %d_win0 = index.add %nt, %carry : index + %d_span = index.mul %d_win0, %drsb : index + %d_seq = index.mul %seq, %d_span : index + %oi0 = index.mul %snapshot_row, %drsb : index + %oi1 = index.add %oi0, %d_seq : index + %oi2 = index.add %oi1, %chan : index + %oi = index.assume %oi2 [range(%oi2, 0, 1073741823)] : index + + %is_state = index.cmp ult, %snapshot_row, %carry : index + %snapshot = scf.if %is_state -> (f32) { + // Rows 0..2 snapshot the old recurrent state. In the same workitem, + // update the corresponding persistent-cache row from x rows 509..511. + %state_span = index.mul %di, %scsb : index + %state_seq = index.mul %seq, %state_span : index + %si0 = index.mul %snapshot_row, %srsb : index + %si1 = index.add %si0, %state_seq : index + %si_chan = index.mul %chan, %scsb : index + %si2 = index.add %si1, %si_chan : index + %si = index.assume %si2 [range(%si2, 0, 1073741823)] : index + + %x_tail0 = index.sub %nt, %carry : index + %x_tail = index.add %x_tail0, %snapshot_row : index + %x_span = index.mul %nt, %xrsb : index + %x_seq = index.mul %seq, %x_span : index + %xi0 = index.mul %x_tail, %xrsb : index + %xi1 = index.add %xi0, %x_seq : index + %xi2 = index.add %xi1, %chan : index + %xi = index.assume %xi2 [range(%xi2, 0, 1073741823)] : index + + %cache_span = index.mul %di, %ccsb : index + %cache_seq = index.mul %seq, %cache_span : index + %ci0 = index.mul %snapshot_row, %crsb : index + %ci1 = index.add %ci0, %cache_seq : index + %ci_chan = index.mul %chan, %ccsb : index + %ci2 = index.add %ci1, %ci_chan : index + %ci = index.assume %ci2 [range(%ci2, 0, 1073741823)] : index + + %old_state = view.load %state_view[%si] : view<1073741824xf32> -> f32 + %new_state = view.load %x_view[%xi] : view<1073741824xf32> -> f32 + view.store %new_state, %c_view[%ci] : f32, view<1073741824xf32> + scf.yield %old_state : f32 + } else { + // Rows 3..63 snapshot x rows 0..60 before their arena span is reused. + %x_row = index.sub %snapshot_row, %carry : index + %x_span = index.mul %nt, %xrsb : index + %x_seq = index.mul %seq, %x_span : index + %xi0 = index.mul %x_row, %xrsb : index + %xi1 = index.add %xi0, %x_seq : index + %xi2 = index.add %xi1, %chan : index + %xi = index.assume %xi2 [range(%xi2, 0, 1073741823)] : index + %x_value = view.load %x_view[%xi] : view<1073741824xf32> -> f32 + scf.yield %x_value : f32 + } + view.store %snapshot, %dst_view[%oi] : f32, view<1073741824xf32> + } + kernel.return +} + +template.decl @llm.ssm_conv.rollback_store_carried(%slot: index, %cache_count: index, %n_t: index, %cache_sequence_base: index, %channel: index, %old0: f32, %old1: f32, %old2: f32, %cache: buffer) + +template.decl @llm.ssm_conv.rollback_store_x(%slot: index, %cache_count: index, %n_t: index, %cache_sequence_base: index, %token: index, %channel: index, %value: f32, %cache: buffer) + +template.decl @llm.ssm_conv.state_materialized_body(%d_inner0: index, %n_t0: index, %n_s0: index, %cache_count0: index, %state: buffer, %x: buffer, %filt: buffer, %dst: buffer, %cache0: buffer, %cache1: buffer, %cache2: buffer, %cache3: buffer, %cache4: buffer) + +// Token-one channels-first SSM assigns one channel to each workitem for the four-tap FMA, SiLU, and cache shift. +// A separate export lets Loom specialize decode independently from prefill. +config.decl @llm.ssm_conv.decode.d_inner : %value: index where [range(%value, 32, 65536), mul(%value, 32)] + +config.decl @llm.ssm_conv.decode.workgroup_size : %value: index where [range(%value, 256, 256)] + +config.decl @llm.ssm_conv.rollback.d_inner : %value: index where [range(%value, 32, 65536), mul(%value, 32)] + +config.decl @llm.ssm_conv.rollback.n_t : %value: index where [range(%value, 1, 512)] + +config.decl @llm.ssm_conv.rollback.n_s : %value: index where [range(%value, 1, 3)] + +config.decl @llm.ssm_conv.rollback.cache_count : %value: index where [range(%value, 1, 5)] + +config.decl @llm.ssm_conv.rollback.workgroup_size : %value: index where [range(%value, 256, 256)] + +template.def<@llm.ssm_conv.rollback_store_carried> device @llm_ssm_conv_rollback_store_carried(%slot: index, %cache_count: index, %n_t: index, %cache_sequence_base: index, %channel: index, %old0: f32, %old1: f32, %old2: f32, %cache: buffer) { + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %active = index.cmp ult, %slot, %cache_count : index + scf.if %active { + %has_full_history = index.cmp ult, %slot, %n_t : index + %start = scf.if %has_full_history -> (index) { + %candidate = index.sub %n_t, %slot : index + scf.yield %candidate : index + } else { + scf.yield %c0 : index + } + %cache_g = buffer.assume.memory_space %cache : buffer + %cache_v = buffer.view %cache_g[%base] : buffer -> view<589824xf32> + scf.for %row = [%c0 to %c3 step %c1] { + %source_row = index.add %start, %row : index + %is_carried = index.cmp ult, %source_row, %c3 : index + scf.if %is_carried { + %is_zero = index.cmp eq, %source_row, %c0 : index + %is_one = index.cmp eq, %source_row, %c1 : index + %old12 = scf.select %is_one, %old1, %old2 : f32 + %value = scf.select %is_zero, %old0, %old12 : f32 + %channel_base0 = index.mul %channel, %c3 : index + %channel_base = index.assume %channel_base0 [range(%channel_base0, 0, 196605), mul(%channel_base0, 3)] : index + %cache_channel_index0 = index.add %channel_base, %row : index + %cache_channel_index = index.assume %cache_channel_index0 [range(%cache_channel_index0, 0, 196607)] : index + %cache_index0 = index.add %cache_sequence_base, %cache_channel_index : index + %cache_index = index.assume %cache_index0 [range(%cache_index0, 0, 589823)] : index + view.store %value, %cache_v[%cache_index] : f32, view<589824xf32> + } + } + } + template.return +} + +template.def<@llm.ssm_conv.rollback_store_x> device @llm_ssm_conv_rollback_store_x(%slot: index, %cache_count: index, %n_t: index, %cache_sequence_base: index, %token: index, %channel: index, %value: f32, %cache: buffer) { + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c3 = index.constant 3 : index + %active = index.cmp ult, %slot, %cache_count : index + %has_full_history = index.cmp ult, %slot, %n_t : index + %start = scf.if %has_full_history -> (index) { + %candidate = index.sub %n_t, %slot : index + scf.yield %candidate : index + } else { + scf.yield %c0 : index + } + %source_row = index.add %token, %c3 : index + %end = index.add %start, %c3 : index + %at_or_after_start = index.cmp uge, %source_row, %start : index + %before_end = index.cmp ult, %source_row, %end : index + %within_window = scalar.andi %at_or_after_start, %before_end : i1 + %writes = scalar.andi %active, %within_window : i1 + scf.if %writes { + %row0 = index.sub %source_row, %start : index + %row = index.assume %row0 [range(%row0, 0, 2)] : index + %channel_base0 = index.mul %channel, %c3 : index + %channel_base = index.assume %channel_base0 [range(%channel_base0, 0, 196605), mul(%channel_base0, 3)] : index + %cache_channel_index0 = index.add %channel_base, %row : index + %cache_channel_index = index.assume %cache_channel_index0 [range(%cache_channel_index0, 0, 196607)] : index + %cache_index0 = index.add %cache_sequence_base, %cache_channel_index : index + %cache_index = index.assume %cache_index0 [range(%cache_index0, 0, 589823)] : index + %cache_g = buffer.assume.memory_space %cache : buffer + %cache_v = buffer.view %cache_g[%base] : buffer -> view<589824xf32> + view.store %value, %cache_v[%cache_index] : f32, view<589824xf32> + } + template.return +} + +template.def<@llm.ssm_conv.state_materialized_body> device @llm_ssm_conv_state_materialized_body(%d_inner0: index, %n_t0: index, %n_s0: index, %cache_count0: index, %state: buffer, %x: buffer, %filt: buffer, %dst: buffer, %cache0: buffer, %cache1: buffer, %cache2: buffer, %cache3: buffer, %cache4: buffer) { + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %d_inner = index.assume %d_inner0 [range(%d_inner0, 32, 65536), mul(%d_inner0, 32)] : index + %n_t = index.assume %n_t0 [range(%n_t0, 1, 512)] : index + %n_s = index.assume %n_s0 [range(%n_s0, 1, 3)] : index + %cache_count = index.assume %cache_count0 [range(%cache_count0, 1, 5)] : index + %group0 = kernel.workgroup.id : index + %sequence0 = kernel.workgroup.id : index + %lane0 = kernel.workitem.id : index + %group = index.assume %group0 [range(%group0, 0, 255)] : index + %sequence, %launch_n_s = index.assume %sequence0, %n_s [lt(%sequence0, %n_s)] : index, index + %lane = index.assume %lane0 [range(%lane0, 0, 255)] : index + %c256 = index.constant 256 : index + %channel_base = index.mul %group, %c256 : index + %channel0 = index.add %channel_base, %lane : index + %channel = index.assume %channel0 [range(%channel0, 0, 65535)] : index + %in_bounds = index.cmp ult, %channel, %d_inner : index + + %state_g = buffer.assume.memory_space %state : buffer + %x_g = buffer.assume.memory_space %x : buffer + %filt_g = buffer.assume.memory_space %filt : buffer + %dst_g = buffer.assume.memory_space %dst : buffer + %state_na, %x_na, %filt_na = buffer.assume.noalias %state_g, %x_g, %filt_g : buffer, buffer, buffer + %state_v = buffer.view %state_na[%base] : buffer -> view<589824xf32> + %x_v = buffer.view %x_na[%base] : buffer -> view<100663296xf32> + %filt_v = buffer.view %filt_na[%base] : buffer -> view<262144xf32> + %dst_v = buffer.view %dst_g[%base] : buffer -> view<100663296xf32> + + %state_sequence_span0 = index.mul %d_inner, %c3 : index + %state_sequence_span = index.assume %state_sequence_span0 [range(%state_sequence_span0, 96, 196608), mul(%state_sequence_span0, 96)] : index + %state_sequence_base0 = index.mul %sequence, %state_sequence_span : index + %state_sequence_base = index.assume %state_sequence_base0 [range(%state_sequence_base0, 0, 393216)] : index + %x_sequence_span0 = index.mul %n_t, %d_inner : index + %x_sequence_span = index.assume %x_sequence_span0 [range(%x_sequence_span0, 32, 33554432), mul(%x_sequence_span0, 32)] : index + %x_sequence_base0 = index.mul %sequence, %x_sequence_span : index + %x_sequence_base = index.assume %x_sequence_base0 [range(%x_sequence_base0, 0, 67108864)] : index + + scf.if %in_bounds { + %state_channel0 = index.mul %channel, %c3 : index + %state_channel = index.assume %state_channel0 [range(%state_channel0, 0, 196605), mul(%state_channel0, 3)] : index + %state_row0 = index.add %state_sequence_base, %state_channel : index + %state_row = index.assume %state_row0 [range(%state_row0, 0, 589821)] : index + %state_row1_0 = index.add %state_row, %c1 : index + %state_row1 = index.assume %state_row1_0 [range(%state_row1_0, 1, 589822)] : index + %state_row2_0 = index.add %state_row, %c2 : index + %state_row2 = index.assume %state_row2_0 [range(%state_row2_0, 2, 589823)] : index + + %old0 = view.load %state_v[%state_row] : view<589824xf32> -> f32 + %old1 = view.load %state_v[%state_row1] : view<589824xf32> -> f32 + %old2 = view.load %state_v[%state_row2] : view<589824xf32> -> f32 + + %filter_row0 = index.mul %channel, %c4 : index + %filter_row = index.assume %filter_row0 [range(%filter_row0, 0, 262140), mul(%filter_row0, 4)] : index + %weights = vector.load %filt_v[%filter_row] : view<262144xf32> -> vector<4xf32> + %w0 = vector.extract %weights[%c0] : vector<4xf32> -> f32 + %w1 = vector.extract %weights[%c1] : vector<4xf32> -> f32 + %w2 = vector.extract %weights[%c2] : vector<4xf32> -> f32 + %w3 = vector.extract %weights[%c3] : vector<4xf32> -> f32 + + template.apply<@llm.ssm_conv.rollback_store_carried>(%c0, %cache_count, %n_t, %state_sequence_base, %channel, %old0, %old1, %old2, %cache0) : (index, index, index, index, index, f32, f32, f32, buffer) + template.apply<@llm.ssm_conv.rollback_store_carried>(%c1, %cache_count, %n_t, %state_sequence_base, %channel, %old0, %old1, %old2, %cache1) : (index, index, index, index, index, f32, f32, f32, buffer) + template.apply<@llm.ssm_conv.rollback_store_carried>(%c2, %cache_count, %n_t, %state_sequence_base, %channel, %old0, %old1, %old2, %cache2) : (index, index, index, index, index, f32, f32, f32, buffer) + template.apply<@llm.ssm_conv.rollback_store_carried>(%c3, %cache_count, %n_t, %state_sequence_base, %channel, %old0, %old1, %old2, %cache3) : (index, index, index, index, index, f32, f32, f32, buffer) + template.apply<@llm.ssm_conv.rollback_store_carried>(%c4, %cache_count, %n_t, %state_sequence_base, %channel, %old0, %old1, %old2, %cache4) : (index, index, index, index, index, f32, f32, f32, buffer) + + // Replace the full three-value history at each rolled loop boundary. + // Keep only three steps and the final partial block unrolled. + %full_trip_count = index.div %n_t, %c3 : index + %full_token_count = index.mul %full_trip_count, %c3 : index + %trip0, %trip1, %trip2 = scf.for %block = [%c0 to %full_token_count step %c3](%block0 = %old0 : f32, %block1 = %old1 : f32, %block2 = %old2 : f32) -> (f32, f32, f32) { + %step0, %step1, %step2 = scf.for %phase = [%c0 to %c3 step %c1](%prev0 = %block0 : f32, %prev1 = %block1 : f32, %prev2 = %block2 : f32) -> (f32, f32, f32) unroll { + %token = index.add %block, %phase : index + + %row_base0 = index.mul %token, %d_inner : index + %row_base = index.assume %row_base0 [range(%row_base0, 0, 33488896), mul(%row_base0, 32)] : index + %row_index0 = index.add %row_base, %channel : index + %row_index = index.assume %row_index0 [range(%row_index0, 0, 33554431)] : index + %index0 = index.add %x_sequence_base, %row_index : index + %index = index.assume %index0 [range(%index0, 0, 100663295)] : index + %new_x = view.load %x_v[%index] : view<100663296xf32> -> f32 + %p0 = scalar.mulf %prev0, %w0 : f32 + %p1 = scalar.fmaf %prev1, %w1, %p0 : f32 + %p2 = scalar.fmaf %prev2, %w2, %p1 : f32 + %acc = scalar.fmaf %new_x, %w3, %p2 : f32 + %activated = scalar.siluf %acc : f32 + + template.apply<@llm.ssm_conv.rollback_store_x>(%c0, %cache_count, %n_t, %state_sequence_base, %token, %channel, %new_x, %cache0) : (index, index, index, index, index, index, f32, buffer) + template.apply<@llm.ssm_conv.rollback_store_x>(%c1, %cache_count, %n_t, %state_sequence_base, %token, %channel, %new_x, %cache1) : (index, index, index, index, index, index, f32, buffer) + template.apply<@llm.ssm_conv.rollback_store_x>(%c2, %cache_count, %n_t, %state_sequence_base, %token, %channel, %new_x, %cache2) : (index, index, index, index, index, index, f32, buffer) + template.apply<@llm.ssm_conv.rollback_store_x>(%c3, %cache_count, %n_t, %state_sequence_base, %token, %channel, %new_x, %cache3) : (index, index, index, index, index, index, f32, buffer) + template.apply<@llm.ssm_conv.rollback_store_x>(%c4, %cache_count, %n_t, %state_sequence_base, %token, %channel, %new_x, %cache4) : (index, index, index, index, index, index, f32, buffer) + view.store %activated, %dst_v[%index] : f32, view<100663296xf32> + scf.yield %prev1, %prev2, %new_x : f32, f32, f32 + } + scf.yield %step0, %step1, %step2 : f32, f32, f32 + } + %final0, %final1, %final2 = scf.for %token = [%full_token_count to %n_t step %c1](%prev0 = %trip0 : f32, %prev1 = %trip1 : f32, %prev2 = %trip2 : f32) -> (f32, f32, f32) unroll { + %row_base0 = index.mul %token, %d_inner : index + %row_base = index.assume %row_base0 [range(%row_base0, 0, 33488896), mul(%row_base0, 32)] : index + %row_index0 = index.add %row_base, %channel : index + %row_index = index.assume %row_index0 [range(%row_index0, 0, 33554431)] : index + %index0 = index.add %x_sequence_base, %row_index : index + %index = index.assume %index0 [range(%index0, 0, 100663295)] : index + %new_x = view.load %x_v[%index] : view<100663296xf32> -> f32 + %p0 = scalar.mulf %prev0, %w0 : f32 + %p1 = scalar.fmaf %prev1, %w1, %p0 : f32 + %p2 = scalar.fmaf %prev2, %w2, %p1 : f32 + %acc = scalar.fmaf %new_x, %w3, %p2 : f32 + %activated = scalar.siluf %acc : f32 + + template.apply<@llm.ssm_conv.rollback_store_x>(%c0, %cache_count, %n_t, %state_sequence_base, %token, %channel, %new_x, %cache0) : (index, index, index, index, index, index, f32, buffer) + template.apply<@llm.ssm_conv.rollback_store_x>(%c1, %cache_count, %n_t, %state_sequence_base, %token, %channel, %new_x, %cache1) : (index, index, index, index, index, index, f32, buffer) + template.apply<@llm.ssm_conv.rollback_store_x>(%c2, %cache_count, %n_t, %state_sequence_base, %token, %channel, %new_x, %cache2) : (index, index, index, index, index, index, f32, buffer) + template.apply<@llm.ssm_conv.rollback_store_x>(%c3, %cache_count, %n_t, %state_sequence_base, %token, %channel, %new_x, %cache3) : (index, index, index, index, index, index, f32, buffer) + template.apply<@llm.ssm_conv.rollback_store_x>(%c4, %cache_count, %n_t, %state_sequence_base, %token, %channel, %new_x, %cache4) : (index, index, index, index, index, index, f32, buffer) + view.store %activated, %dst_v[%index] : f32, view<100663296xf32> + scf.yield %prev1, %prev2, %new_x : f32, f32, f32 + } + } + template.return +} + +kernel.def export("llm_ssm_conv_dconv4_silu_decode_f32") @llm_ssm_conv_dconv4_silu_decode_f32() { + %c1 = index.constant 1 : index + %cneg1 = index.constant -1 : index + %d_inner = config.get @llm.ssm_conv.decode.d_inner : index + %wg = config.get @llm.ssm_conv.decode.workgroup_size : index + %rounding = index.add %wg, %cneg1 : index + %rounded = index.add %d_inner, %rounding : index + %groups = index.div %rounded, %wg : index + kernel.launch.config workgroups(%groups, %c1, %c1) workgroup_size(%wg, %c1, %c1) : index +} launch(%state: buffer, %x: buffer, %filt: buffer, %dst: buffer, %cache: buffer) { + %c1 = index.constant 1 : index + %d_inner = config.get @llm.ssm_conv.decode.d_inner : index + template.apply<@llm.ssm_conv.state_materialized_body>(%d_inner, %c1, %c1, %c1, %state, %x, %filt, %dst, %cache, %cache, %cache, %cache, %cache) : (index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +kernel.def export("llm_ssm_conv_dconv4_silu_rollback_f32") @llm_ssm_conv_dconv4_silu_rollback_f32() { + %c1 = index.constant 1 : index + %cneg1 = index.constant -1 : index + %d_inner = config.get @llm.ssm_conv.rollback.d_inner : index + %n_s = config.get @llm.ssm_conv.rollback.n_s : index + %wg = config.get @llm.ssm_conv.rollback.workgroup_size : index + %rounding = index.add %wg, %cneg1 : index + %rounded = index.add %d_inner, %rounding : index + %groups = index.div %rounded, %wg : index + kernel.launch.config workgroups(%groups, %n_s, %c1) workgroup_size(%wg, %c1, %c1) : index +} launch(%state: buffer, %x: buffer, %filt: buffer, %dst: buffer, %cache0: buffer, %cache1: buffer, %cache2: buffer, %cache3: buffer, %cache4: buffer) { + %d_inner = config.get @llm.ssm_conv.rollback.d_inner : index + %n_t = config.get @llm.ssm_conv.rollback.n_t : index + %n_s = config.get @llm.ssm_conv.rollback.n_s : index + %cache_count = config.get @llm.ssm_conv.rollback.cache_count : index + template.apply<@llm.ssm_conv.state_materialized_body>(%d_inner, %n_t, %n_s, %cache_count, %state, %x, %filt, %dst, %cache0, %cache1, %cache2, %cache3, %cache4) : (index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +config.decl @llm.ssm_conv.prefill.d_conv : %value: index where [range(%value, 4, 4)] + +config.decl @llm.ssm_conv.prefill.d_inner : %value: index where [range(%value, 8192, 10240), mul(%value, 32)] + +config.decl @llm.ssm_conv.prefill.n_t : %value: index where [range(%value, 512, 512)] + +config.decl @llm.ssm_conv.prefill.n_s : %value: index where [range(%value, 1, 1)] + +config.decl @llm.ssm_conv.prefill.state_row_stride : %value: index where [range(%value, 8192, 10240), mul(%value, 32)] + +config.decl @llm.ssm_conv.prefill.x_row_stride : %value: index where [range(%value, 8192, 10240), mul(%value, 32)] + +func.def inline @llm_ssm_conv_prefill_channel_packet(%channels: index) -> (index) { + %zero = index.constant 0 : index + %one = index.constant 1 : index + %two = index.constant 2 : index + %alignment = index.constant 64 : index + %remainder = index.rem %channels, %alignment : index + %aligned = index.cmp eq, %remainder, %zero : index + %width = scf.select %aligned, %two, %one : index + func.return %width : index +} + +kernel.def export("llm_ssm_conv_dconv4_silu_prefill_512_wg1024") @llm_ssm_conv_dconv4_silu_prefill_512_wg1024() { + %c1 = index.constant 1 : index + %channel_lanes = index.constant 32 : index + %threads = index.constant 1024 : index + %d_inner = config.get @llm.ssm_conv.prefill.d_inner : index + %n_s = config.get @llm.ssm_conv.prefill.n_s : index + %packet_width = func.call pure @llm_ssm_conv_prefill_channel_packet(%d_inner) : (index) -> (index) + %channels = index.mul %packet_width, %channel_lanes : index + %groups_x = index.div %d_inner, %channels : index + kernel.launch.config workgroups(%groups_x, %n_s, %c1) workgroup_size(%threads, %c1, %c1) : index +} launch(%state: buffer, %x: buffer, %filt: buffer, %dst: buffer) { + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c5 = index.constant 5 : index + %c6 = index.constant 6 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c9 = index.constant 9 : index + %c10 = index.constant 10 : index + %c11 = index.constant 11 : index + %c12 = index.constant 12 : index + %c13 = index.constant 13 : index + %c14 = index.constant 14 : index + %c15 = index.constant 15 : index + %c16 = index.constant 16 : index + %c17 = index.constant 17 : index + %c18 = index.constant 18 : index + %c19 = index.constant 19 : index + %c20 = index.constant 20 : index + %c21 = index.constant 21 : index + %c22 = index.constant 22 : index + %c23 = index.constant 23 : index + %c24 = index.constant 24 : index + %c25 = index.constant 25 : index + %c26 = index.constant 26 : index + %c27 = index.constant 27 : index + %c28 = index.constant 28 : index + %c29 = index.constant 29 : index + %c30 = index.constant 30 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c61 = index.constant 61 : index + %packet_channels = config.get @llm.ssm_conv.prefill.d_inner : index + %packet_width = func.call pure @llm_ssm_conv_prefill_channel_packet(%packet_channels) : (index) -> (index) + %channel_limit = index.constant 10240 : index + %filter_limit = index.constant 163840 : index + %flat_limit = index.constant 1073741824 : index + %channel_last = index.sub %channel_limit, %packet_width : index + %packet_bytes = index.mul %packet_width, %c16 : index + %filter_last = index.sub %filter_limit, %packet_bytes : index + %flat_last = index.sub %flat_limit, %packet_width : index + %c0_f32 = vector.constant 0.0 : vector<[%packet_width]xf32> + + %d_inner0 = config.get @llm.ssm_conv.prefill.d_inner : index + %n_t0 = config.get @llm.ssm_conv.prefill.n_t : index + %n_s0 = config.get @llm.ssm_conv.prefill.n_s : index + %state_row_stride0 = config.get @llm.ssm_conv.prefill.state_row_stride : index + %x_row_stride0 = config.get @llm.ssm_conv.prefill.x_row_stride : index + %d_inner = index.assume %d_inner0 [range(%d_inner0, 8192, 10240), mul(%d_inner0, 32)] : index + %n_t = index.assume %n_t0 [range(%n_t0, 512, 512)] : index + %n_s = index.assume %n_s0 [range(%n_s0, 1, 1)] : index + %state_row_stride = index.assume %state_row_stride0 [range(%state_row_stride0, 8192, 10240), mul(%state_row_stride0, 32)] : index + %x_row_stride = index.assume %x_row_stride0 [range(%x_row_stride0, 8192, 10240), mul(%x_row_stride0, 32)] : index + + %group_x0 = kernel.workgroup.id : index + %sequence0 = kernel.workgroup.id : index + %local0 = kernel.workitem.id : index + %group_x = index.assume %group_x0 [range(%group_x0, 0, 319)] : index + %sequence = index.assume %sequence0 [range(%sequence0, 0, 0)] : index + %local = index.assume %local0 [range(%local0, 0, 1023)] : index + %wave0 = index.div %local, %c32 : index + %wave = index.assume %wave0 [range(%wave0, 0, 31)] : index + %lane0 = index.rem %local, %c32 : index + %lane = index.assume %lane0 [range(%lane0, 0, 31)] : index + %first_wave = index.cmp eq, %wave, %c0 : index + + %group_channels = index.mul %packet_width, %c32 : index + %channel_base0 = index.mul %group_x, %group_channels : index + %channel_base = index.assume %channel_base0 [range(%channel_base0, 0, 10208), mul(%channel_base0, 32)] : index + %lane_channel = index.mul %lane, %packet_width : index + %channel0 = index.add %channel_base, %lane_channel : index + %channel = index.assume %channel0 [range(%channel0, 0, %channel_last), mul(%channel0, %packet_width)] : index + %wave_token_base0 = index.mul %wave, %c16 : index + %wave_token_base = index.assume %wave_token_base0 [range(%wave_token_base0, 0, 496), mul(%wave_token_base0, 16)] : index + %prior_wave = scf.if %first_wave -> (index) { + scf.yield %c0 : index + } else { + %prior_wave0 = index.sub %wave, %c1 : index + %prior_wave1 = index.assume %prior_wave0 [range(%prior_wave0, 0, 30)] : index + scf.yield %prior_wave1 : index + } + %prior_wave_base0 = index.mul %prior_wave, %c16 : index + %prior_wave_base = index.assume %prior_wave_base0 [range(%prior_wave_base0, 0, 480), mul(%prior_wave_base0, 16)] : index + + %state_g = buffer.assume.memory_space %state : buffer + %x_g = buffer.assume.memory_space %x : buffer + %filt_g = buffer.assume.memory_space %filt : buffer + %dst_g = buffer.assume.memory_space %dst : buffer + // dst deliberately does not participate: its physical span may equal x. + %state_na, %x_na, %filt_na = buffer.assume.noalias %state_g, %x_g, %filt_g : buffer, buffer, buffer + %state_v = buffer.view %state_na[%base] : buffer -> view<1073741824xf32> + %x_v = buffer.view %x_na[%base] : buffer -> view<1073741824xf32> + %dst_v = buffer.view %dst_g[%base] : buffer -> view<1073741824xf32> + + %x_sequence_span0 = index.mul %n_t, %x_row_stride : index + %x_sequence_span = index.assume %x_sequence_span0 [range(%x_sequence_span0, 4194304, 5242880), mul(%x_sequence_span0, 16384)] : index + %x_sequence_base0 = index.mul %sequence, %x_sequence_span : index + %x_sequence_base = index.assume %x_sequence_base0 [range(%x_sequence_base0, 0, %flat_last)] : index + + %state_sequence_rows0 = index.add %n_t, %c3 : index + %state_sequence_rows = index.assume %state_sequence_rows0 [range(%state_sequence_rows0, 515, 515)] : index + %state_sequence_span0 = index.mul %state_sequence_rows, %state_row_stride : index + %state_sequence_span = index.assume %state_sequence_span0 [range(%state_sequence_span0, 4218880, 5273600), mul(%state_sequence_span0, 16480)] : index + %state_sequence_base0 = index.mul %sequence, %state_sequence_span : index + %state_sequence_base = index.assume %state_sequence_base0 [range(%state_sequence_base0, 0, %flat_last)] : index + + // Snapshot slot k maps to x row wave_base-3+k; wave zero gets its first three taps from carried state. + // Other waves preload all 19 slots, and rows 0..60 use the immutable CONCAT window. + %xr0, %xr1, %xr2 = scf.if %first_wave -> (vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>) { + scf.yield %c0_f32, %c0_f32, %c0_f32 : vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32> + } else { + %x_row_00 = index.add %prior_wave_base, %c13 : index + %x_row_0 = index.assume %x_row_00 [range(%x_row_00, 13, 508)] : index + %x_row_10 = index.add %prior_wave_base, %c14 : index + %x_row_1 = index.assume %x_row_10 [range(%x_row_10, 14, 509)] : index + %x_row_20 = index.add %prior_wave_base, %c15 : index + %x_row_2 = index.assume %x_row_20 [range(%x_row_20, 15, 510)] : index + %early_is_snapshotted = index.cmp ult, %wave, %c4 : index + %early_value_0, %early_value_1, %early_value_2 = scf.if %early_is_snapshotted -> (vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>) { + %snapshot_x_row_early_0 = index.assume %x_row_0 [range(%x_row_0, 0, 60)] : index + %snapshot_row_early_00 = index.add %snapshot_x_row_early_0, %c3 : index + %snapshot_row_early_0 = index.assume %snapshot_row_early_00 [range(%snapshot_row_early_00, 3, 63)] : index + %snapshot_row_early_0_offset0 = index.mul %snapshot_row_early_0, %state_row_stride : index + %snapshot_row_early_0_offset1 = index.add %state_sequence_base, %snapshot_row_early_0_offset0 : index + %snapshot_index_early_00 = index.add %snapshot_row_early_0_offset1, %channel : index + %snapshot_index_early_0 = index.assume %snapshot_index_early_00 [range(%snapshot_index_early_00, 0, %flat_last)] : index + %early_snapshot_value_0 = vector.load %state_v[%snapshot_index_early_0] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %snapshot_x_row_early_1 = index.assume %x_row_1 [range(%x_row_1, 0, 60)] : index + %snapshot_row_early_10 = index.add %snapshot_x_row_early_1, %c3 : index + %snapshot_row_early_1 = index.assume %snapshot_row_early_10 [range(%snapshot_row_early_10, 3, 63)] : index + %snapshot_row_early_1_offset0 = index.mul %snapshot_row_early_1, %state_row_stride : index + %snapshot_row_early_1_offset1 = index.add %state_sequence_base, %snapshot_row_early_1_offset0 : index + %snapshot_index_early_10 = index.add %snapshot_row_early_1_offset1, %channel : index + %snapshot_index_early_1 = index.assume %snapshot_index_early_10 [range(%snapshot_index_early_10, 0, %flat_last)] : index + %early_snapshot_value_1 = vector.load %state_v[%snapshot_index_early_1] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %snapshot_x_row_early_2 = index.assume %x_row_2 [range(%x_row_2, 0, 60)] : index + %snapshot_row_early_20 = index.add %snapshot_x_row_early_2, %c3 : index + %snapshot_row_early_2 = index.assume %snapshot_row_early_20 [range(%snapshot_row_early_20, 3, 63)] : index + %snapshot_row_early_2_offset0 = index.mul %snapshot_row_early_2, %state_row_stride : index + %snapshot_row_early_2_offset1 = index.add %state_sequence_base, %snapshot_row_early_2_offset0 : index + %snapshot_index_early_20 = index.add %snapshot_row_early_2_offset1, %channel : index + %snapshot_index_early_2 = index.assume %snapshot_index_early_20 [range(%snapshot_index_early_20, 0, %flat_last)] : index + %early_snapshot_value_2 = vector.load %state_v[%snapshot_index_early_2] : view<1073741824xf32> -> vector<[%packet_width]xf32> + scf.yield %early_snapshot_value_0, %early_snapshot_value_1, %early_snapshot_value_2 : vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32> + } else { + %x_row_early_00 = index.mul %x_row_0, %x_row_stride : index + %x_row_early_01 = index.add %x_sequence_base, %x_row_early_00 : index + %x_index_early_00 = index.add %x_row_early_01, %channel : index + %x_index_early_0 = index.assume %x_index_early_00 [range(%x_index_early_00, 0, %flat_last)] : index + %early_raw_value_0 = vector.load %x_v[%x_index_early_0] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %x_row_early_10 = index.mul %x_row_1, %x_row_stride : index + %x_row_early_11 = index.add %x_sequence_base, %x_row_early_10 : index + %x_index_early_10 = index.add %x_row_early_11, %channel : index + %x_index_early_1 = index.assume %x_index_early_10 [range(%x_index_early_10, 0, %flat_last)] : index + %early_raw_value_1 = vector.load %x_v[%x_index_early_1] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %x_row_early_20 = index.mul %x_row_2, %x_row_stride : index + %x_row_early_21 = index.add %x_sequence_base, %x_row_early_20 : index + %x_index_early_20 = index.add %x_row_early_21, %channel : index + %x_index_early_2 = index.assume %x_index_early_20 [range(%x_index_early_20, 0, %flat_last)] : index + %early_raw_value_2 = vector.load %x_v[%x_index_early_2] : view<1073741824xf32> -> vector<[%packet_width]xf32> + scf.yield %early_raw_value_0, %early_raw_value_1, %early_raw_value_2 : vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32> + } + scf.yield %early_value_0, %early_value_1, %early_value_2 : vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32> + } + %x_row_common_30 = index.add %wave_token_base, %c0 : index + %x_row_common_3 = index.assume %x_row_common_30 [range(%x_row_common_30, 0, 496)] : index + %x_row_common_40 = index.add %wave_token_base, %c1 : index + %x_row_common_4 = index.assume %x_row_common_40 [range(%x_row_common_40, 1, 497)] : index + %x_row_common_50 = index.add %wave_token_base, %c2 : index + %x_row_common_5 = index.assume %x_row_common_50 [range(%x_row_common_50, 2, 498)] : index + %x_row_common_60 = index.add %wave_token_base, %c3 : index + %x_row_common_6 = index.assume %x_row_common_60 [range(%x_row_common_60, 3, 499)] : index + %x_row_common_70 = index.add %wave_token_base, %c4 : index + %x_row_common_7 = index.assume %x_row_common_70 [range(%x_row_common_70, 4, 500)] : index + %x_row_common_80 = index.add %wave_token_base, %c5 : index + %x_row_common_8 = index.assume %x_row_common_80 [range(%x_row_common_80, 5, 501)] : index + %x_row_common_90 = index.add %wave_token_base, %c6 : index + %x_row_common_9 = index.assume %x_row_common_90 [range(%x_row_common_90, 6, 502)] : index + %x_row_common_100 = index.add %wave_token_base, %c7 : index + %x_row_common_10 = index.assume %x_row_common_100 [range(%x_row_common_100, 7, 503)] : index + %x_row_common_110 = index.add %wave_token_base, %c8 : index + %x_row_common_11 = index.assume %x_row_common_110 [range(%x_row_common_110, 8, 504)] : index + %x_row_common_120 = index.add %wave_token_base, %c9 : index + %x_row_common_12 = index.assume %x_row_common_120 [range(%x_row_common_120, 9, 505)] : index + %x_row_common_130 = index.add %wave_token_base, %c10 : index + %x_row_common_13 = index.assume %x_row_common_130 [range(%x_row_common_130, 10, 506)] : index + %x_row_common_140 = index.add %wave_token_base, %c11 : index + %x_row_common_14 = index.assume %x_row_common_140 [range(%x_row_common_140, 11, 507)] : index + %x_row_common_150 = index.add %wave_token_base, %c12 : index + %x_row_common_15 = index.assume %x_row_common_150 [range(%x_row_common_150, 12, 508)] : index + %common_low_is_snapshotted = index.cmp ult, %wave, %c4 : index + %xr3, %xr4, %xr5, %xr6, %xr7, %xr8, %xr9, %xr10, %xr11, %xr12, %xr13, %xr14, %xr15 = scf.if %common_low_is_snapshotted -> (vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>) { + %snapshot_x_row_common_low_3 = index.assume %x_row_common_3 [range(%x_row_common_3, 0, 60)] : index + %snapshot_row_common_low_30 = index.add %snapshot_x_row_common_low_3, %c3 : index + %snapshot_row_common_low_3 = index.assume %snapshot_row_common_low_30 [range(%snapshot_row_common_low_30, 3, 63)] : index + %snapshot_row_common_low_3_offset0 = index.mul %snapshot_row_common_low_3, %state_row_stride : index + %snapshot_row_common_low_3_offset1 = index.add %state_sequence_base, %snapshot_row_common_low_3_offset0 : index + %snapshot_index_common_low_30 = index.add %snapshot_row_common_low_3_offset1, %channel : index + %snapshot_index_common_low_3 = index.assume %snapshot_index_common_low_30 [range(%snapshot_index_common_low_30, 0, %flat_last)] : index + %common_low_snapshot_value_3 = vector.load %state_v[%snapshot_index_common_low_3] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %snapshot_x_row_common_low_4 = index.assume %x_row_common_4 [range(%x_row_common_4, 0, 60)] : index + %snapshot_row_common_low_40 = index.add %snapshot_x_row_common_low_4, %c3 : index + %snapshot_row_common_low_4 = index.assume %snapshot_row_common_low_40 [range(%snapshot_row_common_low_40, 3, 63)] : index + %snapshot_row_common_low_4_offset0 = index.mul %snapshot_row_common_low_4, %state_row_stride : index + %snapshot_row_common_low_4_offset1 = index.add %state_sequence_base, %snapshot_row_common_low_4_offset0 : index + %snapshot_index_common_low_40 = index.add %snapshot_row_common_low_4_offset1, %channel : index + %snapshot_index_common_low_4 = index.assume %snapshot_index_common_low_40 [range(%snapshot_index_common_low_40, 0, %flat_last)] : index + %common_low_snapshot_value_4 = vector.load %state_v[%snapshot_index_common_low_4] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %snapshot_x_row_common_low_5 = index.assume %x_row_common_5 [range(%x_row_common_5, 0, 60)] : index + %snapshot_row_common_low_50 = index.add %snapshot_x_row_common_low_5, %c3 : index + %snapshot_row_common_low_5 = index.assume %snapshot_row_common_low_50 [range(%snapshot_row_common_low_50, 3, 63)] : index + %snapshot_row_common_low_5_offset0 = index.mul %snapshot_row_common_low_5, %state_row_stride : index + %snapshot_row_common_low_5_offset1 = index.add %state_sequence_base, %snapshot_row_common_low_5_offset0 : index + %snapshot_index_common_low_50 = index.add %snapshot_row_common_low_5_offset1, %channel : index + %snapshot_index_common_low_5 = index.assume %snapshot_index_common_low_50 [range(%snapshot_index_common_low_50, 0, %flat_last)] : index + %common_low_snapshot_value_5 = vector.load %state_v[%snapshot_index_common_low_5] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %snapshot_x_row_common_low_6 = index.assume %x_row_common_6 [range(%x_row_common_6, 0, 60)] : index + %snapshot_row_common_low_60 = index.add %snapshot_x_row_common_low_6, %c3 : index + %snapshot_row_common_low_6 = index.assume %snapshot_row_common_low_60 [range(%snapshot_row_common_low_60, 3, 63)] : index + %snapshot_row_common_low_6_offset0 = index.mul %snapshot_row_common_low_6, %state_row_stride : index + %snapshot_row_common_low_6_offset1 = index.add %state_sequence_base, %snapshot_row_common_low_6_offset0 : index + %snapshot_index_common_low_60 = index.add %snapshot_row_common_low_6_offset1, %channel : index + %snapshot_index_common_low_6 = index.assume %snapshot_index_common_low_60 [range(%snapshot_index_common_low_60, 0, %flat_last)] : index + %common_low_snapshot_value_6 = vector.load %state_v[%snapshot_index_common_low_6] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %snapshot_x_row_common_low_7 = index.assume %x_row_common_7 [range(%x_row_common_7, 0, 60)] : index + %snapshot_row_common_low_70 = index.add %snapshot_x_row_common_low_7, %c3 : index + %snapshot_row_common_low_7 = index.assume %snapshot_row_common_low_70 [range(%snapshot_row_common_low_70, 3, 63)] : index + %snapshot_row_common_low_7_offset0 = index.mul %snapshot_row_common_low_7, %state_row_stride : index + %snapshot_row_common_low_7_offset1 = index.add %state_sequence_base, %snapshot_row_common_low_7_offset0 : index + %snapshot_index_common_low_70 = index.add %snapshot_row_common_low_7_offset1, %channel : index + %snapshot_index_common_low_7 = index.assume %snapshot_index_common_low_70 [range(%snapshot_index_common_low_70, 0, %flat_last)] : index + %common_low_snapshot_value_7 = vector.load %state_v[%snapshot_index_common_low_7] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %snapshot_x_row_common_low_8 = index.assume %x_row_common_8 [range(%x_row_common_8, 0, 60)] : index + %snapshot_row_common_low_80 = index.add %snapshot_x_row_common_low_8, %c3 : index + %snapshot_row_common_low_8 = index.assume %snapshot_row_common_low_80 [range(%snapshot_row_common_low_80, 3, 63)] : index + %snapshot_row_common_low_8_offset0 = index.mul %snapshot_row_common_low_8, %state_row_stride : index + %snapshot_row_common_low_8_offset1 = index.add %state_sequence_base, %snapshot_row_common_low_8_offset0 : index + %snapshot_index_common_low_80 = index.add %snapshot_row_common_low_8_offset1, %channel : index + %snapshot_index_common_low_8 = index.assume %snapshot_index_common_low_80 [range(%snapshot_index_common_low_80, 0, %flat_last)] : index + %common_low_snapshot_value_8 = vector.load %state_v[%snapshot_index_common_low_8] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %snapshot_x_row_common_low_9 = index.assume %x_row_common_9 [range(%x_row_common_9, 0, 60)] : index + %snapshot_row_common_low_90 = index.add %snapshot_x_row_common_low_9, %c3 : index + %snapshot_row_common_low_9 = index.assume %snapshot_row_common_low_90 [range(%snapshot_row_common_low_90, 3, 63)] : index + %snapshot_row_common_low_9_offset0 = index.mul %snapshot_row_common_low_9, %state_row_stride : index + %snapshot_row_common_low_9_offset1 = index.add %state_sequence_base, %snapshot_row_common_low_9_offset0 : index + %snapshot_index_common_low_90 = index.add %snapshot_row_common_low_9_offset1, %channel : index + %snapshot_index_common_low_9 = index.assume %snapshot_index_common_low_90 [range(%snapshot_index_common_low_90, 0, %flat_last)] : index + %common_low_snapshot_value_9 = vector.load %state_v[%snapshot_index_common_low_9] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %snapshot_x_row_common_low_10 = index.assume %x_row_common_10 [range(%x_row_common_10, 0, 60)] : index + %snapshot_row_common_low_100 = index.add %snapshot_x_row_common_low_10, %c3 : index + %snapshot_row_common_low_10 = index.assume %snapshot_row_common_low_100 [range(%snapshot_row_common_low_100, 3, 63)] : index + %snapshot_row_common_low_10_offset0 = index.mul %snapshot_row_common_low_10, %state_row_stride : index + %snapshot_row_common_low_10_offset1 = index.add %state_sequence_base, %snapshot_row_common_low_10_offset0 : index + %snapshot_index_common_low_100 = index.add %snapshot_row_common_low_10_offset1, %channel : index + %snapshot_index_common_low_10 = index.assume %snapshot_index_common_low_100 [range(%snapshot_index_common_low_100, 0, %flat_last)] : index + %common_low_snapshot_value_10 = vector.load %state_v[%snapshot_index_common_low_10] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %snapshot_x_row_common_low_11 = index.assume %x_row_common_11 [range(%x_row_common_11, 0, 60)] : index + %snapshot_row_common_low_110 = index.add %snapshot_x_row_common_low_11, %c3 : index + %snapshot_row_common_low_11 = index.assume %snapshot_row_common_low_110 [range(%snapshot_row_common_low_110, 3, 63)] : index + %snapshot_row_common_low_11_offset0 = index.mul %snapshot_row_common_low_11, %state_row_stride : index + %snapshot_row_common_low_11_offset1 = index.add %state_sequence_base, %snapshot_row_common_low_11_offset0 : index + %snapshot_index_common_low_110 = index.add %snapshot_row_common_low_11_offset1, %channel : index + %snapshot_index_common_low_11 = index.assume %snapshot_index_common_low_110 [range(%snapshot_index_common_low_110, 0, %flat_last)] : index + %common_low_snapshot_value_11 = vector.load %state_v[%snapshot_index_common_low_11] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %snapshot_x_row_common_low_12 = index.assume %x_row_common_12 [range(%x_row_common_12, 0, 60)] : index + %snapshot_row_common_low_120 = index.add %snapshot_x_row_common_low_12, %c3 : index + %snapshot_row_common_low_12 = index.assume %snapshot_row_common_low_120 [range(%snapshot_row_common_low_120, 3, 63)] : index + %snapshot_row_common_low_12_offset0 = index.mul %snapshot_row_common_low_12, %state_row_stride : index + %snapshot_row_common_low_12_offset1 = index.add %state_sequence_base, %snapshot_row_common_low_12_offset0 : index + %snapshot_index_common_low_120 = index.add %snapshot_row_common_low_12_offset1, %channel : index + %snapshot_index_common_low_12 = index.assume %snapshot_index_common_low_120 [range(%snapshot_index_common_low_120, 0, %flat_last)] : index + %common_low_snapshot_value_12 = vector.load %state_v[%snapshot_index_common_low_12] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %snapshot_x_row_common_low_13 = index.assume %x_row_common_13 [range(%x_row_common_13, 0, 60)] : index + %snapshot_row_common_low_130 = index.add %snapshot_x_row_common_low_13, %c3 : index + %snapshot_row_common_low_13 = index.assume %snapshot_row_common_low_130 [range(%snapshot_row_common_low_130, 3, 63)] : index + %snapshot_row_common_low_13_offset0 = index.mul %snapshot_row_common_low_13, %state_row_stride : index + %snapshot_row_common_low_13_offset1 = index.add %state_sequence_base, %snapshot_row_common_low_13_offset0 : index + %snapshot_index_common_low_130 = index.add %snapshot_row_common_low_13_offset1, %channel : index + %snapshot_index_common_low_13 = index.assume %snapshot_index_common_low_130 [range(%snapshot_index_common_low_130, 0, %flat_last)] : index + %common_low_snapshot_value_13 = vector.load %state_v[%snapshot_index_common_low_13] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %snapshot_x_row_common_low_14 = index.assume %x_row_common_14 [range(%x_row_common_14, 0, 60)] : index + %snapshot_row_common_low_140 = index.add %snapshot_x_row_common_low_14, %c3 : index + %snapshot_row_common_low_14 = index.assume %snapshot_row_common_low_140 [range(%snapshot_row_common_low_140, 3, 63)] : index + %snapshot_row_common_low_14_offset0 = index.mul %snapshot_row_common_low_14, %state_row_stride : index + %snapshot_row_common_low_14_offset1 = index.add %state_sequence_base, %snapshot_row_common_low_14_offset0 : index + %snapshot_index_common_low_140 = index.add %snapshot_row_common_low_14_offset1, %channel : index + %snapshot_index_common_low_14 = index.assume %snapshot_index_common_low_140 [range(%snapshot_index_common_low_140, 0, %flat_last)] : index + %common_low_snapshot_value_14 = vector.load %state_v[%snapshot_index_common_low_14] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %snapshot_x_row_common_low_15 = index.assume %x_row_common_15 [range(%x_row_common_15, 0, 60)] : index + %snapshot_row_common_low_150 = index.add %snapshot_x_row_common_low_15, %c3 : index + %snapshot_row_common_low_15 = index.assume %snapshot_row_common_low_150 [range(%snapshot_row_common_low_150, 3, 63)] : index + %snapshot_row_common_low_15_offset0 = index.mul %snapshot_row_common_low_15, %state_row_stride : index + %snapshot_row_common_low_15_offset1 = index.add %state_sequence_base, %snapshot_row_common_low_15_offset0 : index + %snapshot_index_common_low_150 = index.add %snapshot_row_common_low_15_offset1, %channel : index + %snapshot_index_common_low_15 = index.assume %snapshot_index_common_low_150 [range(%snapshot_index_common_low_150, 0, %flat_last)] : index + %common_low_snapshot_value_15 = vector.load %state_v[%snapshot_index_common_low_15] : view<1073741824xf32> -> vector<[%packet_width]xf32> + scf.yield %common_low_snapshot_value_3, %common_low_snapshot_value_4, %common_low_snapshot_value_5, %common_low_snapshot_value_6, %common_low_snapshot_value_7, %common_low_snapshot_value_8, %common_low_snapshot_value_9, %common_low_snapshot_value_10, %common_low_snapshot_value_11, %common_low_snapshot_value_12, %common_low_snapshot_value_13, %common_low_snapshot_value_14, %common_low_snapshot_value_15 : vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32> + } else { + %x_row_common_low_30 = index.mul %x_row_common_3, %x_row_stride : index + %x_row_common_low_31 = index.add %x_sequence_base, %x_row_common_low_30 : index + %x_index_common_low_30 = index.add %x_row_common_low_31, %channel : index + %x_index_common_low_3 = index.assume %x_index_common_low_30 [range(%x_index_common_low_30, 0, %flat_last)] : index + %common_low_raw_value_3 = vector.load %x_v[%x_index_common_low_3] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %x_row_common_low_40 = index.mul %x_row_common_4, %x_row_stride : index + %x_row_common_low_41 = index.add %x_sequence_base, %x_row_common_low_40 : index + %x_index_common_low_40 = index.add %x_row_common_low_41, %channel : index + %x_index_common_low_4 = index.assume %x_index_common_low_40 [range(%x_index_common_low_40, 0, %flat_last)] : index + %common_low_raw_value_4 = vector.load %x_v[%x_index_common_low_4] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %x_row_common_low_50 = index.mul %x_row_common_5, %x_row_stride : index + %x_row_common_low_51 = index.add %x_sequence_base, %x_row_common_low_50 : index + %x_index_common_low_50 = index.add %x_row_common_low_51, %channel : index + %x_index_common_low_5 = index.assume %x_index_common_low_50 [range(%x_index_common_low_50, 0, %flat_last)] : index + %common_low_raw_value_5 = vector.load %x_v[%x_index_common_low_5] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %x_row_common_low_60 = index.mul %x_row_common_6, %x_row_stride : index + %x_row_common_low_61 = index.add %x_sequence_base, %x_row_common_low_60 : index + %x_index_common_low_60 = index.add %x_row_common_low_61, %channel : index + %x_index_common_low_6 = index.assume %x_index_common_low_60 [range(%x_index_common_low_60, 0, %flat_last)] : index + %common_low_raw_value_6 = vector.load %x_v[%x_index_common_low_6] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %x_row_common_low_70 = index.mul %x_row_common_7, %x_row_stride : index + %x_row_common_low_71 = index.add %x_sequence_base, %x_row_common_low_70 : index + %x_index_common_low_70 = index.add %x_row_common_low_71, %channel : index + %x_index_common_low_7 = index.assume %x_index_common_low_70 [range(%x_index_common_low_70, 0, %flat_last)] : index + %common_low_raw_value_7 = vector.load %x_v[%x_index_common_low_7] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %x_row_common_low_80 = index.mul %x_row_common_8, %x_row_stride : index + %x_row_common_low_81 = index.add %x_sequence_base, %x_row_common_low_80 : index + %x_index_common_low_80 = index.add %x_row_common_low_81, %channel : index + %x_index_common_low_8 = index.assume %x_index_common_low_80 [range(%x_index_common_low_80, 0, %flat_last)] : index + %common_low_raw_value_8 = vector.load %x_v[%x_index_common_low_8] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %x_row_common_low_90 = index.mul %x_row_common_9, %x_row_stride : index + %x_row_common_low_91 = index.add %x_sequence_base, %x_row_common_low_90 : index + %x_index_common_low_90 = index.add %x_row_common_low_91, %channel : index + %x_index_common_low_9 = index.assume %x_index_common_low_90 [range(%x_index_common_low_90, 0, %flat_last)] : index + %common_low_raw_value_9 = vector.load %x_v[%x_index_common_low_9] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %x_row_common_low_100 = index.mul %x_row_common_10, %x_row_stride : index + %x_row_common_low_101 = index.add %x_sequence_base, %x_row_common_low_100 : index + %x_index_common_low_100 = index.add %x_row_common_low_101, %channel : index + %x_index_common_low_10 = index.assume %x_index_common_low_100 [range(%x_index_common_low_100, 0, %flat_last)] : index + %common_low_raw_value_10 = vector.load %x_v[%x_index_common_low_10] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %x_row_common_low_110 = index.mul %x_row_common_11, %x_row_stride : index + %x_row_common_low_111 = index.add %x_sequence_base, %x_row_common_low_110 : index + %x_index_common_low_110 = index.add %x_row_common_low_111, %channel : index + %x_index_common_low_11 = index.assume %x_index_common_low_110 [range(%x_index_common_low_110, 0, %flat_last)] : index + %common_low_raw_value_11 = vector.load %x_v[%x_index_common_low_11] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %x_row_common_low_120 = index.mul %x_row_common_12, %x_row_stride : index + %x_row_common_low_121 = index.add %x_sequence_base, %x_row_common_low_120 : index + %x_index_common_low_120 = index.add %x_row_common_low_121, %channel : index + %x_index_common_low_12 = index.assume %x_index_common_low_120 [range(%x_index_common_low_120, 0, %flat_last)] : index + %common_low_raw_value_12 = vector.load %x_v[%x_index_common_low_12] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %x_row_common_low_130 = index.mul %x_row_common_13, %x_row_stride : index + %x_row_common_low_131 = index.add %x_sequence_base, %x_row_common_low_130 : index + %x_index_common_low_130 = index.add %x_row_common_low_131, %channel : index + %x_index_common_low_13 = index.assume %x_index_common_low_130 [range(%x_index_common_low_130, 0, %flat_last)] : index + %common_low_raw_value_13 = vector.load %x_v[%x_index_common_low_13] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %x_row_common_low_140 = index.mul %x_row_common_14, %x_row_stride : index + %x_row_common_low_141 = index.add %x_sequence_base, %x_row_common_low_140 : index + %x_index_common_low_140 = index.add %x_row_common_low_141, %channel : index + %x_index_common_low_14 = index.assume %x_index_common_low_140 [range(%x_index_common_low_140, 0, %flat_last)] : index + %common_low_raw_value_14 = vector.load %x_v[%x_index_common_low_14] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %x_row_common_low_150 = index.mul %x_row_common_15, %x_row_stride : index + %x_row_common_low_151 = index.add %x_sequence_base, %x_row_common_low_150 : index + %x_index_common_low_150 = index.add %x_row_common_low_151, %channel : index + %x_index_common_low_15 = index.assume %x_index_common_low_150 [range(%x_index_common_low_150, 0, %flat_last)] : index + %common_low_raw_value_15 = vector.load %x_v[%x_index_common_low_15] : view<1073741824xf32> -> vector<[%packet_width]xf32> + scf.yield %common_low_raw_value_3, %common_low_raw_value_4, %common_low_raw_value_5, %common_low_raw_value_6, %common_low_raw_value_7, %common_low_raw_value_8, %common_low_raw_value_9, %common_low_raw_value_10, %common_low_raw_value_11, %common_low_raw_value_12, %common_low_raw_value_13, %common_low_raw_value_14, %common_low_raw_value_15 : vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32> + } + %x_row_common_160 = index.add %wave_token_base, %c13 : index + %x_row_common_16 = index.assume %x_row_common_160 [range(%x_row_common_160, 13, 509)] : index + %x_row_common_170 = index.add %wave_token_base, %c14 : index + %x_row_common_17 = index.assume %x_row_common_170 [range(%x_row_common_170, 14, 510)] : index + %x_row_common_180 = index.add %wave_token_base, %c15 : index + %x_row_common_18 = index.assume %x_row_common_180 [range(%x_row_common_180, 15, 511)] : index + %common_high_is_snapshotted = index.cmp ult, %wave, %c3 : index + %xr16, %xr17, %xr18 = scf.if %common_high_is_snapshotted -> (vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32>) { + %snapshot_x_row_common_high_16 = index.assume %x_row_common_16 [range(%x_row_common_16, 0, 60)] : index + %snapshot_row_common_high_160 = index.add %snapshot_x_row_common_high_16, %c3 : index + %snapshot_row_common_high_16 = index.assume %snapshot_row_common_high_160 [range(%snapshot_row_common_high_160, 3, 63)] : index + %snapshot_row_common_high_16_offset0 = index.mul %snapshot_row_common_high_16, %state_row_stride : index + %snapshot_row_common_high_16_offset1 = index.add %state_sequence_base, %snapshot_row_common_high_16_offset0 : index + %snapshot_index_common_high_160 = index.add %snapshot_row_common_high_16_offset1, %channel : index + %snapshot_index_common_high_16 = index.assume %snapshot_index_common_high_160 [range(%snapshot_index_common_high_160, 0, %flat_last)] : index + %common_high_snapshot_value_16 = vector.load %state_v[%snapshot_index_common_high_16] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %snapshot_x_row_common_high_17 = index.assume %x_row_common_17 [range(%x_row_common_17, 0, 60)] : index + %snapshot_row_common_high_170 = index.add %snapshot_x_row_common_high_17, %c3 : index + %snapshot_row_common_high_17 = index.assume %snapshot_row_common_high_170 [range(%snapshot_row_common_high_170, 3, 63)] : index + %snapshot_row_common_high_17_offset0 = index.mul %snapshot_row_common_high_17, %state_row_stride : index + %snapshot_row_common_high_17_offset1 = index.add %state_sequence_base, %snapshot_row_common_high_17_offset0 : index + %snapshot_index_common_high_170 = index.add %snapshot_row_common_high_17_offset1, %channel : index + %snapshot_index_common_high_17 = index.assume %snapshot_index_common_high_170 [range(%snapshot_index_common_high_170, 0, %flat_last)] : index + %common_high_snapshot_value_17 = vector.load %state_v[%snapshot_index_common_high_17] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %snapshot_x_row_common_high_18 = index.assume %x_row_common_18 [range(%x_row_common_18, 0, 60)] : index + %snapshot_row_common_high_180 = index.add %snapshot_x_row_common_high_18, %c3 : index + %snapshot_row_common_high_18 = index.assume %snapshot_row_common_high_180 [range(%snapshot_row_common_high_180, 3, 63)] : index + %snapshot_row_common_high_18_offset0 = index.mul %snapshot_row_common_high_18, %state_row_stride : index + %snapshot_row_common_high_18_offset1 = index.add %state_sequence_base, %snapshot_row_common_high_18_offset0 : index + %snapshot_index_common_high_180 = index.add %snapshot_row_common_high_18_offset1, %channel : index + %snapshot_index_common_high_18 = index.assume %snapshot_index_common_high_180 [range(%snapshot_index_common_high_180, 0, %flat_last)] : index + %common_high_snapshot_value_18 = vector.load %state_v[%snapshot_index_common_high_18] : view<1073741824xf32> -> vector<[%packet_width]xf32> + scf.yield %common_high_snapshot_value_16, %common_high_snapshot_value_17, %common_high_snapshot_value_18 : vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32> + } else { + %x_row_common_high_160 = index.mul %x_row_common_16, %x_row_stride : index + %x_row_common_high_161 = index.add %x_sequence_base, %x_row_common_high_160 : index + %x_index_common_high_160 = index.add %x_row_common_high_161, %channel : index + %x_index_common_high_16 = index.assume %x_index_common_high_160 [range(%x_index_common_high_160, 0, %flat_last)] : index + %common_high_raw_value_16 = vector.load %x_v[%x_index_common_high_16] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %x_row_common_high_170 = index.mul %x_row_common_17, %x_row_stride : index + %x_row_common_high_171 = index.add %x_sequence_base, %x_row_common_high_170 : index + %x_index_common_high_170 = index.add %x_row_common_high_171, %channel : index + %x_index_common_high_17 = index.assume %x_index_common_high_170 [range(%x_index_common_high_170, 0, %flat_last)] : index + %common_high_raw_value_17 = vector.load %x_v[%x_index_common_high_17] : view<1073741824xf32> -> vector<[%packet_width]xf32> + %x_row_common_high_180 = index.mul %x_row_common_18, %x_row_stride : index + %x_row_common_high_181 = index.add %x_sequence_base, %x_row_common_high_180 : index + %x_index_common_high_180 = index.add %x_row_common_high_181, %channel : index + %x_index_common_high_18 = index.assume %x_index_common_high_180 [range(%x_index_common_high_180, 0, %flat_last)] : index + %common_high_raw_value_18 = vector.load %x_v[%x_index_common_high_18] : view<1073741824xf32> -> vector<[%packet_width]xf32> + scf.yield %common_high_raw_value_16, %common_high_raw_value_17, %common_high_raw_value_18 : vector<[%packet_width]xf32>, vector<[%packet_width]xf32>, vector<[%packet_width]xf32> + } + // Filter and state are disjoint from x/dst. Each wave loads its lane's four + // filter values once and reuses them for 16 outputs. + %filter_byte0 = index.mul %channel, %c16 : index + %filter_byte = index.assume %filter_byte0 [range(%filter_byte0, 0, %filter_last), mul(%filter_byte0, %packet_bytes)] : index + %filter_offset = index.cast %filter_byte : index to offset + %filter_v = buffer.view %filt_na[%filter_offset] : buffer -> view<[%packet_width]x4xf32> + %filter_values = vector.load %filter_v[%c0, %c0] : view<[%packet_width]x4xf32> -> vector<[%packet_width]x4xf32> + %filter_transposed = vector.transpose<[1, 0]> %filter_values : vector<[%packet_width]x4xf32> -> vector<4x[%packet_width]xf32> + %w0 = vector.extract %filter_transposed[0] : vector<4x[%packet_width]xf32> -> vector<[%packet_width]xf32> + %w1 = vector.extract %filter_transposed[1] : vector<4x[%packet_width]xf32> -> vector<[%packet_width]xf32> + %w2 = vector.extract %filter_transposed[2] : vector<4x[%packet_width]xf32> -> vector<[%packet_width]xf32> + %w3 = vector.extract %filter_transposed[3] : vector<4x[%packet_width]xf32> -> vector<[%packet_width]xf32> + + %state0 = scf.if %first_wave -> (vector<[%packet_width]xf32>) { + %state_row_00 = index.mul %c0, %state_row_stride : index + %state_row_01 = index.add %state_sequence_base, %state_row_00 : index + %state_index_00 = index.add %state_row_01, %channel : index + %state_index_0 = index.assume %state_index_00 [range(%state_index_00, 0, %flat_last)] : index + %loaded = vector.load %state_v[%state_index_0] : view<1073741824xf32> -> vector<[%packet_width]xf32> + scf.yield %loaded : vector<[%packet_width]xf32> + } else { + scf.yield %c0_f32 : vector<[%packet_width]xf32> + } + + %state1 = scf.if %first_wave -> (vector<[%packet_width]xf32>) { + %state_row_10 = index.mul %c1, %state_row_stride : index + %state_row_11 = index.add %state_sequence_base, %state_row_10 : index + %state_index_10 = index.add %state_row_11, %channel : index + %state_index_1 = index.assume %state_index_10 [range(%state_index_10, 0, %flat_last)] : index + %loaded = vector.load %state_v[%state_index_1] : view<1073741824xf32> -> vector<[%packet_width]xf32> + scf.yield %loaded : vector<[%packet_width]xf32> + } else { + scf.yield %c0_f32 : vector<[%packet_width]xf32> + } + + %state2 = scf.if %first_wave -> (vector<[%packet_width]xf32>) { + %state_row_20 = index.mul %c2, %state_row_stride : index + %state_row_21 = index.add %state_sequence_base, %state_row_20 : index + %state_index_20 = index.add %state_row_21, %channel : index + %state_index_2 = index.assume %state_index_20 [range(%state_index_20, 0, %flat_last)] : index + %loaded = vector.load %state_v[%state_index_2] : view<1073741824xf32> -> vector<[%packet_width]xf32> + scf.yield %loaded : vector<[%packet_width]xf32> + } else { + scf.yield %c0_f32 : vector<[%packet_width]xf32> + } + + // Consume every X preload before the barrier because gfx11 s_barrier does not drain VMEM. + // Only the 16 activated values remain live across the rendezvous. + %v0_0 = scf.select %first_wave, %state0, %xr0 : vector<[%packet_width]xf32> + %v0_1 = scf.select %first_wave, %state1, %xr1 : vector<[%packet_width]xf32> + %v0_2 = scf.select %first_wave, %state2, %xr2 : vector<[%packet_width]xf32> + %p0_0 = vector.mulf %v0_0, %w0 : vector<[%packet_width]xf32> + %p0_1 = vector.fmaf %v0_1, %w1, %p0_0 : vector<[%packet_width]xf32> + %p0_2 = vector.fmaf %v0_2, %w2, %p0_1 : vector<[%packet_width]xf32> + %acc0 = vector.fmaf %xr3, %w3, %p0_2 : vector<[%packet_width]xf32> + %activated0 = vector.siluf %acc0 : vector<[%packet_width]xf32> + + %v1_0 = scf.select %first_wave, %state1, %xr1 : vector<[%packet_width]xf32> + %v1_1 = scf.select %first_wave, %state2, %xr2 : vector<[%packet_width]xf32> + %p1_0 = vector.mulf %v1_0, %w0 : vector<[%packet_width]xf32> + %p1_1 = vector.fmaf %v1_1, %w1, %p1_0 : vector<[%packet_width]xf32> + %p1_2 = vector.fmaf %xr3, %w2, %p1_1 : vector<[%packet_width]xf32> + %acc1 = vector.fmaf %xr4, %w3, %p1_2 : vector<[%packet_width]xf32> + %activated1 = vector.siluf %acc1 : vector<[%packet_width]xf32> + + %v2_0 = scf.select %first_wave, %state2, %xr2 : vector<[%packet_width]xf32> + %p2_0 = vector.mulf %v2_0, %w0 : vector<[%packet_width]xf32> + %p2_1 = vector.fmaf %xr3, %w1, %p2_0 : vector<[%packet_width]xf32> + %p2_2 = vector.fmaf %xr4, %w2, %p2_1 : vector<[%packet_width]xf32> + %acc2 = vector.fmaf %xr5, %w3, %p2_2 : vector<[%packet_width]xf32> + %activated2 = vector.siluf %acc2 : vector<[%packet_width]xf32> + + %p3_0 = vector.mulf %xr3, %w0 : vector<[%packet_width]xf32> + %p3_1 = vector.fmaf %xr4, %w1, %p3_0 : vector<[%packet_width]xf32> + %p3_2 = vector.fmaf %xr5, %w2, %p3_1 : vector<[%packet_width]xf32> + %acc3 = vector.fmaf %xr6, %w3, %p3_2 : vector<[%packet_width]xf32> + %activated3 = vector.siluf %acc3 : vector<[%packet_width]xf32> + + %p4_0 = vector.mulf %xr4, %w0 : vector<[%packet_width]xf32> + %p4_1 = vector.fmaf %xr5, %w1, %p4_0 : vector<[%packet_width]xf32> + %p4_2 = vector.fmaf %xr6, %w2, %p4_1 : vector<[%packet_width]xf32> + %acc4 = vector.fmaf %xr7, %w3, %p4_2 : vector<[%packet_width]xf32> + %activated4 = vector.siluf %acc4 : vector<[%packet_width]xf32> + + %p5_0 = vector.mulf %xr5, %w0 : vector<[%packet_width]xf32> + %p5_1 = vector.fmaf %xr6, %w1, %p5_0 : vector<[%packet_width]xf32> + %p5_2 = vector.fmaf %xr7, %w2, %p5_1 : vector<[%packet_width]xf32> + %acc5 = vector.fmaf %xr8, %w3, %p5_2 : vector<[%packet_width]xf32> + %activated5 = vector.siluf %acc5 : vector<[%packet_width]xf32> + + %p6_0 = vector.mulf %xr6, %w0 : vector<[%packet_width]xf32> + %p6_1 = vector.fmaf %xr7, %w1, %p6_0 : vector<[%packet_width]xf32> + %p6_2 = vector.fmaf %xr8, %w2, %p6_1 : vector<[%packet_width]xf32> + %acc6 = vector.fmaf %xr9, %w3, %p6_2 : vector<[%packet_width]xf32> + %activated6 = vector.siluf %acc6 : vector<[%packet_width]xf32> + + %p7_0 = vector.mulf %xr7, %w0 : vector<[%packet_width]xf32> + %p7_1 = vector.fmaf %xr8, %w1, %p7_0 : vector<[%packet_width]xf32> + %p7_2 = vector.fmaf %xr9, %w2, %p7_1 : vector<[%packet_width]xf32> + %acc7 = vector.fmaf %xr10, %w3, %p7_2 : vector<[%packet_width]xf32> + %activated7 = vector.siluf %acc7 : vector<[%packet_width]xf32> + + %p8_0 = vector.mulf %xr8, %w0 : vector<[%packet_width]xf32> + %p8_1 = vector.fmaf %xr9, %w1, %p8_0 : vector<[%packet_width]xf32> + %p8_2 = vector.fmaf %xr10, %w2, %p8_1 : vector<[%packet_width]xf32> + %acc8 = vector.fmaf %xr11, %w3, %p8_2 : vector<[%packet_width]xf32> + %activated8 = vector.siluf %acc8 : vector<[%packet_width]xf32> + + %p9_0 = vector.mulf %xr9, %w0 : vector<[%packet_width]xf32> + %p9_1 = vector.fmaf %xr10, %w1, %p9_0 : vector<[%packet_width]xf32> + %p9_2 = vector.fmaf %xr11, %w2, %p9_1 : vector<[%packet_width]xf32> + %acc9 = vector.fmaf %xr12, %w3, %p9_2 : vector<[%packet_width]xf32> + %activated9 = vector.siluf %acc9 : vector<[%packet_width]xf32> + + %p10_0 = vector.mulf %xr10, %w0 : vector<[%packet_width]xf32> + %p10_1 = vector.fmaf %xr11, %w1, %p10_0 : vector<[%packet_width]xf32> + %p10_2 = vector.fmaf %xr12, %w2, %p10_1 : vector<[%packet_width]xf32> + %acc10 = vector.fmaf %xr13, %w3, %p10_2 : vector<[%packet_width]xf32> + %activated10 = vector.siluf %acc10 : vector<[%packet_width]xf32> + + %p11_0 = vector.mulf %xr11, %w0 : vector<[%packet_width]xf32> + %p11_1 = vector.fmaf %xr12, %w1, %p11_0 : vector<[%packet_width]xf32> + %p11_2 = vector.fmaf %xr13, %w2, %p11_1 : vector<[%packet_width]xf32> + %acc11 = vector.fmaf %xr14, %w3, %p11_2 : vector<[%packet_width]xf32> + %activated11 = vector.siluf %acc11 : vector<[%packet_width]xf32> + + %p12_0 = vector.mulf %xr12, %w0 : vector<[%packet_width]xf32> + %p12_1 = vector.fmaf %xr13, %w1, %p12_0 : vector<[%packet_width]xf32> + %p12_2 = vector.fmaf %xr14, %w2, %p12_1 : vector<[%packet_width]xf32> + %acc12 = vector.fmaf %xr15, %w3, %p12_2 : vector<[%packet_width]xf32> + %activated12 = vector.siluf %acc12 : vector<[%packet_width]xf32> + + %p13_0 = vector.mulf %xr13, %w0 : vector<[%packet_width]xf32> + %p13_1 = vector.fmaf %xr14, %w1, %p13_0 : vector<[%packet_width]xf32> + %p13_2 = vector.fmaf %xr15, %w2, %p13_1 : vector<[%packet_width]xf32> + %acc13 = vector.fmaf %xr16, %w3, %p13_2 : vector<[%packet_width]xf32> + %activated13 = vector.siluf %acc13 : vector<[%packet_width]xf32> + + %p14_0 = vector.mulf %xr14, %w0 : vector<[%packet_width]xf32> + %p14_1 = vector.fmaf %xr15, %w1, %p14_0 : vector<[%packet_width]xf32> + %p14_2 = vector.fmaf %xr16, %w2, %p14_1 : vector<[%packet_width]xf32> + %acc14 = vector.fmaf %xr17, %w3, %p14_2 : vector<[%packet_width]xf32> + %activated14 = vector.siluf %acc14 : vector<[%packet_width]xf32> + + %p15_0 = vector.mulf %xr15, %w0 : vector<[%packet_width]xf32> + %p15_1 = vector.fmaf %xr16, %w1, %p15_0 : vector<[%packet_width]xf32> + %p15_2 = vector.fmaf %xr17, %w2, %p15_1 : vector<[%packet_width]xf32> + %acc15 = vector.fmaf %xr18, %w3, %p15_2 : vector<[%packet_width]xf32> + %activated15 = vector.siluf %acc15 : vector<[%packet_width]xf32> + + // Uniform load/compute-before-store boundary. There is no LDS allocation + // and no destination store above this point. + kernel.barrier scope(workgroup) ordering(acq_rel) + + %dst_sequence_span0 = index.mul %n_t, %d_inner : index + %dst_sequence_span = index.assume %dst_sequence_span0 [range(%dst_sequence_span0, 4194304, 5242880), mul(%dst_sequence_span0, 16384)] : index + %dst_sequence_base0 = index.mul %sequence, %dst_sequence_span : index + %dst_sequence_base = index.assume %dst_sequence_base0 [range(%dst_sequence_base0, 0, %flat_last)] : index + + %token_0_raw = index.add %wave_token_base, %c0 : index + %token_0 = index.assume %token_0_raw [range(%token_0_raw, 0, 496)] : index + %dst_row_0_raw = index.mul %token_0, %d_inner : index + %dst_row_0 = index.add %dst_sequence_base, %dst_row_0_raw : index + %dst_index_0_raw = index.add %dst_row_0, %channel : index + %dst_index_0 = index.assume %dst_index_0_raw [range(%dst_index_0_raw, 0, %flat_last)] : index + vector.store %activated0, %dst_v[%dst_index_0] : vector<[%packet_width]xf32>, view<1073741824xf32> + %token_1_raw = index.add %wave_token_base, %c1 : index + %token_1 = index.assume %token_1_raw [range(%token_1_raw, 1, 497)] : index + %dst_row_1_raw = index.mul %token_1, %d_inner : index + %dst_row_1 = index.add %dst_sequence_base, %dst_row_1_raw : index + %dst_index_1_raw = index.add %dst_row_1, %channel : index + %dst_index_1 = index.assume %dst_index_1_raw [range(%dst_index_1_raw, 0, %flat_last)] : index + vector.store %activated1, %dst_v[%dst_index_1] : vector<[%packet_width]xf32>, view<1073741824xf32> + %token_2_raw = index.add %wave_token_base, %c2 : index + %token_2 = index.assume %token_2_raw [range(%token_2_raw, 2, 498)] : index + %dst_row_2_raw = index.mul %token_2, %d_inner : index + %dst_row_2 = index.add %dst_sequence_base, %dst_row_2_raw : index + %dst_index_2_raw = index.add %dst_row_2, %channel : index + %dst_index_2 = index.assume %dst_index_2_raw [range(%dst_index_2_raw, 0, %flat_last)] : index + vector.store %activated2, %dst_v[%dst_index_2] : vector<[%packet_width]xf32>, view<1073741824xf32> + %token_3_raw = index.add %wave_token_base, %c3 : index + %token_3 = index.assume %token_3_raw [range(%token_3_raw, 3, 499)] : index + %dst_row_3_raw = index.mul %token_3, %d_inner : index + %dst_row_3 = index.add %dst_sequence_base, %dst_row_3_raw : index + %dst_index_3_raw = index.add %dst_row_3, %channel : index + %dst_index_3 = index.assume %dst_index_3_raw [range(%dst_index_3_raw, 0, %flat_last)] : index + vector.store %activated3, %dst_v[%dst_index_3] : vector<[%packet_width]xf32>, view<1073741824xf32> + %token_4_raw = index.add %wave_token_base, %c4 : index + %token_4 = index.assume %token_4_raw [range(%token_4_raw, 4, 500)] : index + %dst_row_4_raw = index.mul %token_4, %d_inner : index + %dst_row_4 = index.add %dst_sequence_base, %dst_row_4_raw : index + %dst_index_4_raw = index.add %dst_row_4, %channel : index + %dst_index_4 = index.assume %dst_index_4_raw [range(%dst_index_4_raw, 0, %flat_last)] : index + vector.store %activated4, %dst_v[%dst_index_4] : vector<[%packet_width]xf32>, view<1073741824xf32> + %token_5_raw = index.add %wave_token_base, %c5 : index + %token_5 = index.assume %token_5_raw [range(%token_5_raw, 5, 501)] : index + %dst_row_5_raw = index.mul %token_5, %d_inner : index + %dst_row_5 = index.add %dst_sequence_base, %dst_row_5_raw : index + %dst_index_5_raw = index.add %dst_row_5, %channel : index + %dst_index_5 = index.assume %dst_index_5_raw [range(%dst_index_5_raw, 0, %flat_last)] : index + vector.store %activated5, %dst_v[%dst_index_5] : vector<[%packet_width]xf32>, view<1073741824xf32> + %token_6_raw = index.add %wave_token_base, %c6 : index + %token_6 = index.assume %token_6_raw [range(%token_6_raw, 6, 502)] : index + %dst_row_6_raw = index.mul %token_6, %d_inner : index + %dst_row_6 = index.add %dst_sequence_base, %dst_row_6_raw : index + %dst_index_6_raw = index.add %dst_row_6, %channel : index + %dst_index_6 = index.assume %dst_index_6_raw [range(%dst_index_6_raw, 0, %flat_last)] : index + vector.store %activated6, %dst_v[%dst_index_6] : vector<[%packet_width]xf32>, view<1073741824xf32> + %token_7_raw = index.add %wave_token_base, %c7 : index + %token_7 = index.assume %token_7_raw [range(%token_7_raw, 7, 503)] : index + %dst_row_7_raw = index.mul %token_7, %d_inner : index + %dst_row_7 = index.add %dst_sequence_base, %dst_row_7_raw : index + %dst_index_7_raw = index.add %dst_row_7, %channel : index + %dst_index_7 = index.assume %dst_index_7_raw [range(%dst_index_7_raw, 0, %flat_last)] : index + vector.store %activated7, %dst_v[%dst_index_7] : vector<[%packet_width]xf32>, view<1073741824xf32> + %token_8_raw = index.add %wave_token_base, %c8 : index + %token_8 = index.assume %token_8_raw [range(%token_8_raw, 8, 504)] : index + %dst_row_8_raw = index.mul %token_8, %d_inner : index + %dst_row_8 = index.add %dst_sequence_base, %dst_row_8_raw : index + %dst_index_8_raw = index.add %dst_row_8, %channel : index + %dst_index_8 = index.assume %dst_index_8_raw [range(%dst_index_8_raw, 0, %flat_last)] : index + vector.store %activated8, %dst_v[%dst_index_8] : vector<[%packet_width]xf32>, view<1073741824xf32> + %token_9_raw = index.add %wave_token_base, %c9 : index + %token_9 = index.assume %token_9_raw [range(%token_9_raw, 9, 505)] : index + %dst_row_9_raw = index.mul %token_9, %d_inner : index + %dst_row_9 = index.add %dst_sequence_base, %dst_row_9_raw : index + %dst_index_9_raw = index.add %dst_row_9, %channel : index + %dst_index_9 = index.assume %dst_index_9_raw [range(%dst_index_9_raw, 0, %flat_last)] : index + vector.store %activated9, %dst_v[%dst_index_9] : vector<[%packet_width]xf32>, view<1073741824xf32> + %token_10_raw = index.add %wave_token_base, %c10 : index + %token_10 = index.assume %token_10_raw [range(%token_10_raw, 10, 506)] : index + %dst_row_10_raw = index.mul %token_10, %d_inner : index + %dst_row_10 = index.add %dst_sequence_base, %dst_row_10_raw : index + %dst_index_10_raw = index.add %dst_row_10, %channel : index + %dst_index_10 = index.assume %dst_index_10_raw [range(%dst_index_10_raw, 0, %flat_last)] : index + vector.store %activated10, %dst_v[%dst_index_10] : vector<[%packet_width]xf32>, view<1073741824xf32> + %token_11_raw = index.add %wave_token_base, %c11 : index + %token_11 = index.assume %token_11_raw [range(%token_11_raw, 11, 507)] : index + %dst_row_11_raw = index.mul %token_11, %d_inner : index + %dst_row_11 = index.add %dst_sequence_base, %dst_row_11_raw : index + %dst_index_11_raw = index.add %dst_row_11, %channel : index + %dst_index_11 = index.assume %dst_index_11_raw [range(%dst_index_11_raw, 0, %flat_last)] : index + vector.store %activated11, %dst_v[%dst_index_11] : vector<[%packet_width]xf32>, view<1073741824xf32> + %token_12_raw = index.add %wave_token_base, %c12 : index + %token_12 = index.assume %token_12_raw [range(%token_12_raw, 12, 508)] : index + %dst_row_12_raw = index.mul %token_12, %d_inner : index + %dst_row_12 = index.add %dst_sequence_base, %dst_row_12_raw : index + %dst_index_12_raw = index.add %dst_row_12, %channel : index + %dst_index_12 = index.assume %dst_index_12_raw [range(%dst_index_12_raw, 0, %flat_last)] : index + vector.store %activated12, %dst_v[%dst_index_12] : vector<[%packet_width]xf32>, view<1073741824xf32> + %token_13_raw = index.add %wave_token_base, %c13 : index + %token_13 = index.assume %token_13_raw [range(%token_13_raw, 13, 509)] : index + %dst_row_13_raw = index.mul %token_13, %d_inner : index + %dst_row_13 = index.add %dst_sequence_base, %dst_row_13_raw : index + %dst_index_13_raw = index.add %dst_row_13, %channel : index + %dst_index_13 = index.assume %dst_index_13_raw [range(%dst_index_13_raw, 0, %flat_last)] : index + vector.store %activated13, %dst_v[%dst_index_13] : vector<[%packet_width]xf32>, view<1073741824xf32> + %token_14_raw = index.add %wave_token_base, %c14 : index + %token_14 = index.assume %token_14_raw [range(%token_14_raw, 14, 510)] : index + %dst_row_14_raw = index.mul %token_14, %d_inner : index + %dst_row_14 = index.add %dst_sequence_base, %dst_row_14_raw : index + %dst_index_14_raw = index.add %dst_row_14, %channel : index + %dst_index_14 = index.assume %dst_index_14_raw [range(%dst_index_14_raw, 0, %flat_last)] : index + vector.store %activated14, %dst_v[%dst_index_14] : vector<[%packet_width]xf32>, view<1073741824xf32> + %token_15_raw = index.add %wave_token_base, %c15 : index + %token_15 = index.assume %token_15_raw [range(%token_15_raw, 15, 511)] : index + %dst_row_15_raw = index.mul %token_15, %d_inner : index + %dst_row_15 = index.add %dst_sequence_base, %dst_row_15_raw : index + %dst_index_15_raw = index.add %dst_row_15, %channel : index + %dst_index_15 = index.assume %dst_index_15_raw [range(%dst_index_15_raw, 0, %flat_last)] : index + vector.store %activated15, %dst_v[%dst_index_15] : vector<[%packet_width]xf32>, view<1073741824xf32> + kernel.return +} + +// Regression: distinct initial history must survive every rolled boundary. +// With history [1, 2, 3], five input tokens equal to 4, and filter +// [1, 2, 3, 4], the convolution is [30, 36, 39, 40, 40]. +// SiLU rounds each of these positive values back to itself in F32. +kernel.def @ssm_conv_rolling_history_check() { + %one = index.constant 1 : index + %wg = index.constant 256 : index + kernel.launch.config workgroups(%one, %one, %one) workgroup_size(%wg, %one, %one) : index +} launch(%state: buffer, %input: buffer, %filter: buffer, %output: buffer, %cache: buffer) { + %channels = index.constant 32 : index + %tokens = index.constant 5 : index + %one = index.constant 1 : index + template.apply<@llm.ssm_conv.state_materialized_body>(%channels, %tokens, %one, %one, %state, %input, %filter, %output, %cache, %cache, %cache, %cache, %cache) : (index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +check.case public @ssm_conv_rolling_history { + %state = check.generate.iota offset(1.0) step(1.0) period(3) : tensor<96xf32> + %input = check.generate.fill value(4.0) : tensor<160xf32> + %filter = check.generate.iota offset(1.0) step(1.0) period(4) : tensor<128xf32> + %separate = check.generate.fill value(-999.0) : tensor<160xf32> + %inplace = check.generate.fill value(-999.0) : tensor<160xf32> + %cache = check.generate.fill value(-999.0) : tensor<96xf32> + %expected_cache = check.generate.fill value(4.0) : tensor<96xf32> + kernel.launch @ssm_conv_rolling_history_check[](%state, %input, %filter, %separate, %cache) : [](tensor<96xf32>, tensor<160xf32>, tensor<128xf32>, tensor<160xf32>, tensor<96xf32>) + kernel.launch @ssm_conv_rolling_history_check[](%state, %input, %filter, %inplace, %state) : [](tensor<96xf32>, tensor<160xf32>, tensor<128xf32>, tensor<160xf32>, tensor<96xf32>) + check.expect.equal actual(%cache) expected(%expected_cache) : tensor<96xf32> + check.expect.equal actual(%state) expected(%expected_cache) : tensor<96xf32> + check.expect.equal actual(%inplace) expected(%separate) : tensor<160xf32> + %row0 = check.tensor.view %separate offset(0) : tensor<160xf32> -> tensor<32xf32> + %expected0 = check.generate.fill value(30.0) : tensor<32xf32> + check.expect.equal actual(%row0) expected(%expected0) : tensor<32xf32> + %row1 = check.tensor.view %separate offset(128) : tensor<160xf32> -> tensor<32xf32> + %expected1 = check.generate.fill value(36.0) : tensor<32xf32> + check.expect.equal actual(%row1) expected(%expected1) : tensor<32xf32> + %row2 = check.tensor.view %separate offset(256) : tensor<160xf32> -> tensor<32xf32> + %expected2 = check.generate.fill value(39.0) : tensor<32xf32> + check.expect.equal actual(%row2) expected(%expected2) : tensor<32xf32> + %row3 = check.tensor.view %separate offset(384) : tensor<160xf32> -> tensor<32xf32> + %expected3 = check.generate.fill value(40.0) : tensor<32xf32> + check.expect.equal actual(%row3) expected(%expected3) : tensor<32xf32> + %row4 = check.tensor.view %separate offset(512) : tensor<160xf32> -> tensor<32xf32> + %expected4 = check.generate.fill value(40.0) : tensor<32xf32> + check.expect.equal actual(%row4) expected(%expected4) : tensor<32xf32> + check.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/ssm_conv_generic_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/ssm_conv_generic_f32.loom new file mode 100644 index 000000000000..27e86f14b073 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/ssm_conv_generic_f32.loom @@ -0,0 +1,276 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +template.decl @ggml.unary_f32.apply(%op: index, %value: f32) -> (f32) + +template.decl @ggml.binary_f32.apply(%op: index, %lhs: f32, %rhs: f32) -> (f32) + +config.decl @llm.ssm_conv.generic.d_conv : %value: index where [range(%value, 1, 16)] + +config.decl @llm.ssm_conv.generic.d_inner : %value: index where [range(%value, 32, 65536), mul(%value, 32)] + +config.decl @llm.ssm_conv.generic.n_t : %value: index where [range(%value, 1, 512)] + +config.decl @llm.ssm_conv.generic.n_s : %value: index where [range(%value, 1, 4)] + +config.decl @llm.ssm_conv.generic.unary_op : %value: index where [range(%value, 0, 23)] + +config.decl @llm.ssm_conv.generic.binary_op : %value: index where [range(%value, 0, 8)] + +config.decl @llm.ssm_conv.generic.binary_lhs : %value: index where [range(%value, 0, 1)] + +config.decl @llm.ssm_conv.generic.workgroup_size : %value: index where [range(%value, 256, 256)] + +kernel.def export("llm_ssm_conv_f32") @llm_ssm_conv_f32() { + %c1 = index.constant 1 : index + %cneg1 = index.constant -1 : index + %d_inner = config.get @llm.ssm_conv.generic.d_inner : index + %n_t = config.get @llm.ssm_conv.generic.n_t : index + %n_s = config.get @llm.ssm_conv.generic.n_s : index + %wg = config.get @llm.ssm_conv.generic.workgroup_size : index + %tokens = index.mul %d_inner, %n_t : index + %total = index.mul %tokens, %n_s : index + %rounding = index.add %wg, %cneg1 : index + %rounded = index.add %total, %rounding : index + %groups = index.div %rounded, %wg : index + kernel.launch.config workgroups(%groups, %c1, %c1) workgroup_size(%wg, %c1, %c1) : index +} launch(%window: buffer, %filter: buffer, %output: buffer) { + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c0_f32 = scalar.constant 0.0 : f32 + + %d_conv0 = config.get @llm.ssm_conv.generic.d_conv : index + %d_inner0 = config.get @llm.ssm_conv.generic.d_inner : index + %n_t0 = config.get @llm.ssm_conv.generic.n_t : index + %n_s0 = config.get @llm.ssm_conv.generic.n_s : index + %unary_op = config.get @llm.ssm_conv.generic.unary_op : index + %wg = config.get @llm.ssm_conv.generic.workgroup_size : index + %d_conv = index.assume %d_conv0 [range(%d_conv0, 1, 16)] : index + %d_inner = index.assume %d_inner0 [range(%d_inner0, 32, 65536), mul(%d_inner0, 32)] : index + %n_t = index.assume %n_t0 [range(%n_t0, 1, 512)] : index + %n_s = index.assume %n_s0 [range(%n_s0, 1, 4)] : index + + %token_span0 = index.mul %d_inner, %n_t : index + %token_span = index.assume %token_span0 [range(%token_span0, 32, 33554432), mul(%token_span0, 32)] : index + %total0 = index.mul %token_span, %n_s : index + %total = index.assume %total0 [range(%total0, 32, 134217728), mul(%total0, 32)] : index + + %group0 = kernel.workgroup.id : index + %lane0 = kernel.workitem.id : index + %group = index.assume %group0 [range(%group0, 0, 524287)] : index + %lane = index.assume %lane0 [range(%lane0, 0, 255)] : index + %base_idx0 = index.mul %group, %wg : index + %linear0 = index.add %base_idx0, %lane : index + %linear = index.assume %linear0 [range(%linear0, 0, 134217983)] : index + %in_bounds = index.cmp ult, %linear, %total : index + + %window_g = buffer.assume.memory_space %window : buffer + %filter_g = buffer.assume.memory_space %filter : buffer + %output_g = buffer.assume.memory_space %output : buffer + %window_na, %filter_na, %output_na = buffer.assume.noalias %window_g, %filter_g, %output_g : buffer, buffer, buffer + %window_v = buffer.view %window_na[%base] : buffer -> view<1073741824xf32> + %filter_v = buffer.view %filter_na[%base] : buffer -> view<1048576xf32> + %output_v = buffer.view %output_na[%base] : buffer -> view<134217728xf32> + + scf.if %in_bounds { + %channel = index.rem %linear, %d_inner : index + %linear_div_inner = index.div %linear, %d_inner : index + %token = index.rem %linear_div_inner, %n_t : index + %sequence = index.div %linear_div_inner, %n_t : index + + %window_rows0 = index.add %n_t, %d_conv : index + %window_rows = index.sub %window_rows0, %c1 : index + %window_sequence_span0 = index.mul %window_rows, %d_inner : index + %window_sequence_span = index.assume %window_sequence_span0 [range(%window_sequence_span0, 32, 34603008), mul(%window_sequence_span0, 32)] : index + %window_sequence_base = index.mul %sequence, %window_sequence_span : index + %window_channel_base0 = index.mul %channel, %window_rows : index + %window_channel_base = index.assume %window_channel_base0 [range(%window_channel_base0, 0, 34603008)] : index + + %filter_channel_base0 = index.mul %channel, %d_conv : index + %filter_channel_base = index.assume %filter_channel_base0 [range(%filter_channel_base0, 0, 1048560)] : index + + %acc = scf.for %tap = [%c0 to %d_conv step %c1](%running = %c0_f32 : f32) -> (f32) { + %window_row = index.add %token, %tap : index + %window_index0 = index.add %window_sequence_base, %window_channel_base : index + %window_index1 = index.add %window_index0, %window_row : index + %window_index = index.assume %window_index1 [range(%window_index1, 0, 1073741823)] : index + %filter_index0 = index.add %filter_channel_base, %tap : index + %filter_index = index.assume %filter_index0 [range(%filter_index0, 0, 1048575)] : index + %window_value = view.load %window_v[%window_index] : view<1073741824xf32> -> f32 + %filter_value = view.load %filter_v[%filter_index] : view<1048576xf32> -> f32 + %product = scalar.mulf %window_value, %filter_value : f32 + %next = scalar.addf %running, %product : f32 + scf.yield %next : f32 + } + %activated = template.apply<@ggml.unary_f32.apply>(%unary_op, %acc) : (index, f32) -> (f32) + view.store %activated, %output_v[%linear] : f32, view<134217728xf32> + } + kernel.return +} + +kernel.def export("llm_ssm_conv_binary_f32") @llm_ssm_conv_binary_f32() { + %c1 = index.constant 1 : index + %cneg1 = index.constant -1 : index + %d_inner = config.get @llm.ssm_conv.generic.d_inner : index + %n_t = config.get @llm.ssm_conv.generic.n_t : index + %n_s = config.get @llm.ssm_conv.generic.n_s : index + %wg = config.get @llm.ssm_conv.generic.workgroup_size : index + %tokens = index.mul %d_inner, %n_t : index + %total = index.mul %tokens, %n_s : index + %rounding = index.add %wg, %cneg1 : index + %rounded = index.add %total, %rounding : index + %groups = index.div %rounded, %wg : index + kernel.launch.config workgroups(%groups, %c1, %c1) workgroup_size(%wg, %c1, %c1) : index +} launch(%window: buffer, %filter: buffer, %operand: buffer, %output: buffer) { + %base = index.constant 0 : offset + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c0_f32 = scalar.constant 0.0 : f32 + + %d_conv0 = config.get @llm.ssm_conv.generic.d_conv : index + %d_inner0 = config.get @llm.ssm_conv.generic.d_inner : index + %n_t0 = config.get @llm.ssm_conv.generic.n_t : index + %n_s0 = config.get @llm.ssm_conv.generic.n_s : index + %unary_op = config.get @llm.ssm_conv.generic.unary_op : index + %binary_op = config.get @llm.ssm_conv.generic.binary_op : index + %binary_lhs = config.get @llm.ssm_conv.generic.binary_lhs : index + %wg = config.get @llm.ssm_conv.generic.workgroup_size : index + %d_conv = index.assume %d_conv0 [range(%d_conv0, 1, 16)] : index + %d_inner = index.assume %d_inner0 [range(%d_inner0, 32, 65536), mul(%d_inner0, 32)] : index + %n_t = index.assume %n_t0 [range(%n_t0, 1, 512)] : index + %n_s = index.assume %n_s0 [range(%n_s0, 1, 4)] : index + + %token_span0 = index.mul %d_inner, %n_t : index + %token_span = index.assume %token_span0 [range(%token_span0, 32, 33554432), mul(%token_span0, 32)] : index + %total0 = index.mul %token_span, %n_s : index + %total = index.assume %total0 [range(%total0, 32, 134217728), mul(%total0, 32)] : index + + %group0 = kernel.workgroup.id : index + %lane0 = kernel.workitem.id : index + %group = index.assume %group0 [range(%group0, 0, 524287)] : index + %lane = index.assume %lane0 [range(%lane0, 0, 255)] : index + %base_idx0 = index.mul %group, %wg : index + %linear0 = index.add %base_idx0, %lane : index + %linear = index.assume %linear0 [range(%linear0, 0, 134217983)] : index + %in_bounds = index.cmp ult, %linear, %total : index + + %window_g = buffer.assume.memory_space %window : buffer + %filter_g = buffer.assume.memory_space %filter : buffer + %operand_g = buffer.assume.memory_space %operand : buffer + %output_g = buffer.assume.memory_space %output : buffer + %window_na, %filter_na, %operand_na, %output_na = buffer.assume.noalias %window_g, %filter_g, %operand_g, %output_g : buffer, buffer, buffer, buffer + %window_v = buffer.view %window_na[%base] : buffer -> view<1073741824xf32> + %filter_v = buffer.view %filter_na[%base] : buffer -> view<1048576xf32> + %operand_v = buffer.view %operand_na[%base] : buffer -> view<134217728xf32> + %output_v = buffer.view %output_na[%base] : buffer -> view<134217728xf32> + + scf.if %in_bounds { + %channel = index.rem %linear, %d_inner : index + %linear_div_inner = index.div %linear, %d_inner : index + %token = index.rem %linear_div_inner, %n_t : index + %sequence = index.div %linear_div_inner, %n_t : index + + %window_rows0 = index.add %n_t, %d_conv : index + %window_rows = index.sub %window_rows0, %c1 : index + %window_sequence_span0 = index.mul %window_rows, %d_inner : index + %window_sequence_span = index.assume %window_sequence_span0 [range(%window_sequence_span0, 32, 34603008), mul(%window_sequence_span0, 32)] : index + %window_sequence_base = index.mul %sequence, %window_sequence_span : index + %window_channel_base0 = index.mul %channel, %window_rows : index + %window_channel_base = index.assume %window_channel_base0 [range(%window_channel_base0, 0, 34603008)] : index + + %filter_channel_base0 = index.mul %channel, %d_conv : index + %filter_channel_base = index.assume %filter_channel_base0 [range(%filter_channel_base0, 0, 1048560)] : index + + %acc = scf.for %tap = [%c0 to %d_conv step %c1](%running = %c0_f32 : f32) -> (f32) { + %window_row = index.add %token, %tap : index + %window_index0 = index.add %window_sequence_base, %window_channel_base : index + %window_index1 = index.add %window_index0, %window_row : index + %window_index = index.assume %window_index1 [range(%window_index1, 0, 1073741823)] : index + %filter_index0 = index.add %filter_channel_base, %tap : index + %filter_index = index.assume %filter_index0 [range(%filter_index0, 0, 1048575)] : index + %window_value = view.load %window_v[%window_index] : view<1073741824xf32> -> f32 + %filter_value = view.load %filter_v[%filter_index] : view<1048576xf32> -> f32 + %product = scalar.mulf %window_value, %filter_value : f32 + %next = scalar.addf %running, %product : f32 + scf.yield %next : f32 + } + %activated = template.apply<@ggml.unary_f32.apply>(%unary_op, %acc) : (index, f32) -> (f32) + %operand_value = view.load %operand_v[%linear] : view<134217728xf32> -> f32 + %is_lhs = index.cmp eq, %binary_lhs, %c1 : index + %lhs = scf.select %is_lhs, %activated, %operand_value : f32 + %rhs = scf.select %is_lhs, %operand_value, %activated : f32 + %result = template.apply<@ggml.binary_f32.apply>(%binary_op, %lhs, %rhs) : (index, f32, f32) -> (f32) + view.store %result, %output_v[%linear] : f32, view<134217728xf32> + } + kernel.return +} + +amdgpu.target @llm_ssm_conv_prefill_finish_wave32 {subgroup_size = 32} + +kernel.def target(@llm_ssm_conv_prefill_finish_wave32) @llm_ssm_conv_dconv4_silu_prefill_finish_f32() { + %width = config.get @llm.ssm_conv.generic.d_inner : index + %one = index.constant 1 : index + %threads = index.constant 256 : index + %round = index.constant 255 : index + %rounded = index.add %width, %round : index + %groups = index.div %rounded, %threads : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%threads, %one, %one) : index +} launch(%state: buffer, %filter: buffer, %edges: buffer, %output: buffer, %cache: buffer) { + %width = config.get @llm.ssm_conv.generic.d_inner : index + %zero = index.constant 0 : index + %one = index.constant 1 : index + %two = index.constant 2 : index + %three = index.constant 3 : index + %four = index.constant 4 : index + %five = index.constant 5 : index + %threads = index.constant 256 : index + %base = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %channel0 = index.madd %group, %threads, %lane : index + %valid = index.cmp ult, %channel0, %width : index + scf.if %valid { + %channel, %bounded_width = index.assume %channel0, %width [lt(%channel0, %width)] : index, index + %state_view = buffer.view %state[%base] : buffer -> view<[%bounded_width]x3xf32> + %filter_view = buffer.view %filter[%base] : buffer -> view<[%bounded_width]x4xf32> + %edge_view = buffer.view %edges[%base] : buffer -> view<[%bounded_width]x6xf32> + %output_view = buffer.view %output[%base] : buffer -> view<512x[%bounded_width]xf32> + %cache_view = buffer.view %cache[%base] : buffer -> view<[%bounded_width]x3xf32> + %s0 = view.load %state_view[%channel, %zero] : view<[%bounded_width]x3xf32> -> f32 + %s1 = view.load %state_view[%channel, %one] : view<[%bounded_width]x3xf32> -> f32 + %s2 = view.load %state_view[%channel, %two] : view<[%bounded_width]x3xf32> -> f32 + %x0 = view.load %edge_view[%channel, %zero] : view<[%bounded_width]x6xf32> -> f32 + %x1 = view.load %edge_view[%channel, %one] : view<[%bounded_width]x6xf32> -> f32 + %x2 = view.load %edge_view[%channel, %two] : view<[%bounded_width]x6xf32> -> f32 + %w0 = view.load %filter_view[%channel, %zero] : view<[%bounded_width]x4xf32> -> f32 + %w1 = view.load %filter_view[%channel, %one] : view<[%bounded_width]x4xf32> -> f32 + %w2 = view.load %filter_view[%channel, %two] : view<[%bounded_width]x4xf32> -> f32 + %w3 = view.load %filter_view[%channel, %three] : view<[%bounded_width]x4xf32> -> f32 + %a00 = scalar.mulf %s0, %w0 : f32 + %a01 = scalar.fmaf %s1, %w1, %a00 : f32 + %a02 = scalar.fmaf %s2, %w2, %a01 : f32 + %a03 = scalar.fmaf %x0, %w3, %a02 : f32 + %y0 = scalar.siluf %a03 : f32 + %a10 = scalar.mulf %s1, %w0 : f32 + %a11 = scalar.fmaf %s2, %w1, %a10 : f32 + %a12 = scalar.fmaf %x0, %w2, %a11 : f32 + %a13 = scalar.fmaf %x1, %w3, %a12 : f32 + %y1 = scalar.siluf %a13 : f32 + %a20 = scalar.mulf %s2, %w0 : f32 + %a21 = scalar.fmaf %x0, %w1, %a20 : f32 + %a22 = scalar.fmaf %x1, %w2, %a21 : f32 + %a23 = scalar.fmaf %x2, %w3, %a22 : f32 + %y2 = scalar.siluf %a23 : f32 + view.store %y0, %output_view[%zero, %channel] : f32, view<512x[%bounded_width]xf32> + view.store %y1, %output_view[%one, %channel] : f32, view<512x[%bounded_width]xf32> + view.store %y2, %output_view[%two, %channel] : f32, view<512x[%bounded_width]xf32> + %tail0 = view.load %edge_view[%channel, %three] : view<[%bounded_width]x6xf32> -> f32 + %tail1 = view.load %edge_view[%channel, %four] : view<[%bounded_width]x6xf32> -> f32 + %tail2 = view.load %edge_view[%channel, %five] : view<[%bounded_width]x6xf32> -> f32 + view.store %tail0, %cache_view[%channel, %zero] : f32, view<[%bounded_width]x3xf32> + view.store %tail1, %cache_view[%channel, %one] : f32, view<[%bounded_width]x3xf32> + view.store %tail2, %cache_view[%channel, %two] : f32, view<[%bounded_width]x3xf32> + } + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/swiglu_oai_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/swiglu_oai_f32.loom new file mode 100644 index 000000000000..2f2250c47f3f --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/swiglu_oai_f32.loom @@ -0,0 +1,141 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// GGML_GLU_OP_SWIGLU_OAI on F32 (gpt-oss's clamped SwiGLU), gate and up as two contiguous tensors of the same shape: +// x = min(gate, limit), y = clamp(up, -limit, limit), output = x / (1 + exp(-alpha x)) * (y + 1) +// (ggml-cpu ggml_compute_forward_swiglu_oai_f32). Element-wise over the flat tensors, four values per workitem. + +amdgpu.target @ggml_swiglu_oai_gfx11_wave32 {subgroup_size = 32} + +config.decl @ggml.swiglu_oai.alpha : f32 + +config.decl @ggml.swiglu_oai.limit : f32 + +func.def inline @ggml_swiglu_oai_value(%gate: f32, %up: f32, %alpha: f32, %limit: f32) -> (f32) { + %one = scalar.constant 1.0 : f32 + %gate_over = scalar.cmpf ogt, %gate, %limit : f32 + %x = scf.select %gate_over, %limit, %gate : f32 + %neg_limit = scalar.negf %limit : f32 + %up_over = scalar.cmpf ogt, %up, %limit : f32 + %up_under = scalar.cmpf olt, %up, %neg_limit : f32 + %y0 = scf.select %up_over, %limit, %up : f32 + %y = scf.select %up_under, %neg_limit, %y0 : f32 + %neg_x = scalar.negf %x : f32 + %ax = scalar.mulf %alpha, %neg_x : f32 + %e = scalar.expf %ax : f32 + %den = scalar.addf %one, %e : f32 + %glu = scalar.divf %x, %den : f32 + %y1 = scalar.addf %y, %one : f32 + %out = scalar.mulf %glu, %y1 : f32 + func.return %out : f32 +} + +kernel.def target(@ggml_swiglu_oai_gfx11_wave32) export("ggml_swiglu_oai_f32") @ggml_swiglu_oai_f32(%element_count: index) { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + %c1023 = index.constant 1023 : index + %c1024 = index.constant 1024 : index + %padded = index.add %element_count, %c1023 : index + %groups = index.div %padded, %c1024 : index + kernel.launch.config workgroups(%groups, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%element_count: index, %gate: buffer, %up: buffer, %output: buffer) where [range(%element_count, 1, 134217728)] { + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c1024 = index.constant 1024 : index + %c0_offset = index.constant 0 : offset + %zero4 = vector.constant 0.0 : vector<4xf32> + %alpha = config.get @ggml.swiglu_oai.alpha : f32 + %limit = config.get @ggml.swiglu_oai.limit : f32 + %group = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %group_base = index.mul %group, %c1024 : index + %item_offset = index.mul %workitem, %c4 : index + %first = index.add %group_base, %item_offset : index + %in_range = index.cmp ult, %first, %element_count : index + %gate_noalias, %up_noalias, %output_noalias = buffer.assume.noalias %gate, %up, %output : buffer, buffer, buffer + %gate_view = buffer.view %gate_noalias[%c0_offset] : buffer -> view<[%element_count]xf32> + %up_view = buffer.view %up_noalias[%c0_offset] : buffer -> view<[%element_count]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%element_count]xf32> + scf.if %in_range { + %mask = vector.mask.range [%first to %element_count step %c1] : index -> vector<4xi1> + %g = vector.load.mask %gate_view[%first], %mask, %zero4 : view<[%element_count]xf32>, vector<4xi1>, vector<4xf32> + %u = vector.load.mask %up_view[%first], %mask, %zero4 : view<[%element_count]xf32>, vector<4xi1>, vector<4xf32> + %g0 = vector.extract %g[0] : vector<4xf32> -> f32 + %g1 = vector.extract %g[1] : vector<4xf32> -> f32 + %g2 = vector.extract %g[2] : vector<4xf32> -> f32 + %g3 = vector.extract %g[3] : vector<4xf32> -> f32 + %u0 = vector.extract %u[0] : vector<4xf32> -> f32 + %u1 = vector.extract %u[1] : vector<4xf32> -> f32 + %u2 = vector.extract %u[2] : vector<4xf32> -> f32 + %u3 = vector.extract %u[3] : vector<4xf32> -> f32 + %o0 = func.call @ggml_swiglu_oai_value(%g0, %u0, %alpha, %limit) : (f32, f32, f32, f32) -> (f32) + %o1 = func.call @ggml_swiglu_oai_value(%g1, %u1, %alpha, %limit) : (f32, f32, f32, f32) -> (f32) + %o2 = func.call @ggml_swiglu_oai_value(%g2, %u2, %alpha, %limit) : (f32, f32, f32, f32) -> (f32) + %o3 = func.call @ggml_swiglu_oai_value(%g3, %u3, %alpha, %limit) : (f32, f32, f32, f32) -> (f32) + %o = vector.from_elements %o0, %o1, %o2, %o3 : vector<4xf32> + vector.store.mask %o, %output_view[%first], %mask : vector<4xf32>, view<[%element_count]xf32>, vector<4xi1> + } + kernel.return +} + +// Reference for the case: one workitem per value, written out without the shared helper. +kernel.def target(@ggml_swiglu_oai_gfx11_wave32) export("ggml_swiglu_oai_reference_f32") @ggml_swiglu_oai_reference_f32(%element_count: index) { + %one = index.constant 1 : index + kernel.launch.config workgroups(%element_count, %one, %one) workgroup_size(%one, %one, %one) : index +} launch(%element_count: index, %gate: buffer, %up: buffer, %output: buffer) where [range(%element_count, 1, 134217728)] { + %c0_offset = index.constant 0 : offset + %one = scalar.constant 1.0 : f32 + %alpha = config.get @ggml.swiglu_oai.alpha : f32 + %limit = config.get @ggml.swiglu_oai.limit : f32 + %i0 = kernel.workgroup.id : index + %i, %bound = index.assume %i0, %element_count [lt(%i0, %element_count)] : index, index + %gv = buffer.view %gate[%c0_offset] : buffer -> view<[%bound]xf32> + %uv = buffer.view %up[%c0_offset] : buffer -> view<[%bound]xf32> + %ov = buffer.view %output[%c0_offset] : buffer -> view<[%bound]xf32> + %g = view.load %gv[%i] : view<[%bound]xf32> -> f32 + %u = view.load %uv[%i] : view<[%bound]xf32> -> f32 + %g_low = scalar.cmpf ole, %g, %limit : f32 + %x = scf.select %g_low, %g, %limit : f32 + %neg_limit = scalar.negf %limit : f32 + %u_low = scalar.cmpf ole, %u, %limit : f32 + %y0 = scf.select %u_low, %u, %limit : f32 + %y_high = scalar.cmpf oge, %y0, %neg_limit : f32 + %y = scf.select %y_high, %y0, %neg_limit : f32 + %nx = scalar.negf %x : f32 + %ax = scalar.mulf %alpha, %nx : f32 + %e = scalar.expf %ax : f32 + %den = scalar.addf %one, %e : f32 + %glu = scalar.divf %x, %den : f32 + %y1 = scalar.addf %y, %one : f32 + %out = scalar.mulf %glu, %y1 : f32 + view.store %out, %ov[%i] : f32, view<[%bound]xf32> + kernel.return +} + +// 2050 values (the last workgroup partial, the last vector partial), inputs past +-limit on both sides. +// Run with --config=ggml.swiglu_oai.alpha=1.702 --config=ggml.swiglu_oai.limit=7.0 (gpt-oss). +check.case public @ggml_swiglu_oai_f32_case { + %gate_seed = check.param.seed base(7300000000000067001) count(1) : i64 + %up_seed = check.param.seed base(7300000000000067002) count(1) : i64 + %count = check.literal value(2050) : index + %gate = check.generate.random.uniform seed(%gate_seed) range(-12.0 to 12.0) : tensor<2050xf32> + %up = check.generate.random.uniform seed(%up_seed) range(-12.0 to 12.0) : tensor<2050xf32> + %output = check.generate.fill value(-7.0) : tensor<2050xf32> + %expected = check.generate.fill value(7.0) : tensor<2050xf32> + kernel.launch @ggml_swiglu_oai_reference_f32[%count](%count, %gate, %up, %expected) : [index](index, tensor<2050xf32>, tensor<2050xf32>, tensor<2050xf32>) + kernel.launch @ggml_swiglu_oai_f32[%count](%count, %gate, %up, %output) : [index](index, tensor<2050xf32>, tensor<2050xf32>, tensor<2050xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-6) rtol(1.0e-6) nan(same) : tensor<2050xf32> + check.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/unary_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/unary_f32.loom new file mode 100644 index 000000000000..b858a1650c06 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/unary_f32.loom @@ -0,0 +1,37 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +template.decl @ggml.unary_f32.apply(%arg0: index, %arg1: f32) -> (f32) + +amdgpu.target @ggml_unary_f32_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.unary_f32.op : %value: index where [range(%value, 0, 23)] + +kernel.def target(@ggml_unary_f32_gfx11_wave64) export("ggml_unary_f32") @ggml_unary_f32(%element_count: index) { + %one = index.constant 1 : index + %twofiftysix = index.constant 256 : index + %rounding = index.constant 255 : index + %rounded = index.add %element_count, %rounding : index + %workgroup_count = index.div %rounded, %twofiftysix : index + kernel.launch.config workgroups(%workgroup_count, %one, %one) workgroup_size(%twofiftysix, %one, %one) : index +} launch(%element_count: index, %input: buffer, %output: buffer) { + %count = index.assume %element_count [range(%element_count, 1, 134217728)] : index + %op = config.get @ggml.unary_f32.op : index + %workgroup = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %twofiftysix = index.constant 256 : index + %base0 = index.mul %workgroup, %twofiftysix : index + %linear0 = index.add %base0, %workitem : index + %linear = index.assume %linear0 [range(%linear0, 0, 134217983)] : index + %in_bounds = index.cmp ult, %linear, %count : index + %zero_offset = index.constant 0 : offset + %input_noalias, %output_noalias = buffer.assume.noalias %input, %output : buffer, buffer + %input_view = buffer.view %input_noalias[%zero_offset] : buffer -> view<[%count]xf32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%count]xf32> + scf.if %in_bounds { + %value = view.load %input_view[%linear] : view<[%count]xf32> -> f32 + %result = template.apply<@ggml.unary_f32.apply>(%op, %value) : (index, f32) -> (f32) + view.store %result, %output_view[%linear] : f32, view<[%count]xf32> + } + kernel.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/zaya_cca_conv_decode_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/zaya_cca_conv_decode_f32.loom new file mode 100644 index 000000000000..f0eb32334449 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/zaya_cca_conv_decode_f32.loom @@ -0,0 +1,272 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// ZAYA's CCA convolution for one token of one sequence (decode), in one dispatch. llama.cpp builds +// it from ~35 graph nodes (13 dispatches on HRX): concat Q and K, append them to the two-step conv +// state, a depthwise 2-tap conv plus bias, then per group of 128 channels a 2-tap grouped conv +// (two F16 matmuls over the depthwise outputs at steps t-1 and t), their sum, and a bias. +// +// Per channel c (x = Q or K projection at c, (s0, s1) the conv state, (w0, w1) and b the depthwise +// weights and bias): +// d0[c] = s0 * w0 + s1 * w1 + b d1[c] = s1 * w0 + x * w1 + b +// Per output o in group g = o / 128 (W is [128 ic, channels, 2 taps], ic fastest): +// out[o] = sum_ic (W[ic, o, 0] * d0[128 g + ic] + W[ic, o, 1] * d1[128 g + ic]) + grp_bias[o] +// and the new conv state of channel o is (s1, x). +// +// One wave64 per output: lanes cover the 128 input channels of the group two at a time, so the F16 +// weight reads are contiguous, and each lane recomputes the two depthwise values it needs (a few +// loads) instead of staging them through memory. Group size 128 and the channel counts are +// compile-time config, so every divisor is a constant. + +amdgpu.target @ggml_zaya_cca_conv_decode_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.zaya_cca_conv.channels : %value: index where [range(%value, 128, 8192), mul(%value, 128)] + +config.decl @ggml.zaya_cca_conv.q_size : %value: index where [range(%value, 1, 8191)] + +// The depthwise conv of channel %c at both steps, and the channel's new input x. +func.def inline @ggml_zaya_cca_depthwise(%c: index, %channels: index, %q_size: index, %qraw: buffer, %kraw: buffer, %state: buffer, %dw: buffer, %dw_bias: buffer) -> (f32, f32, f32, f32) { + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %zero_offset = index.constant 0 : offset + %state_size = index.mul %channels, %c2 : index + %i0_raw = index.mul %c, %c2 : index + %i1_raw = index.add %i0_raw, %c1 : index + %i0, %state_bound0 = index.assume %i0_raw, %state_size [lt(%i0_raw, %state_size)] : index, index + %i1, %state_bound1 = index.assume %i1_raw, %state_size [lt(%i1_raw, %state_size)] : index, index + %state_view0 = buffer.view %state[%zero_offset] : buffer -> view<[%state_bound0]xf32> + %state_view1 = buffer.view %state[%zero_offset] : buffer -> view<[%state_bound1]xf32> + %dw_view0 = buffer.view %dw[%zero_offset] : buffer -> view<[%state_bound0]xf32> + %dw_view1 = buffer.view %dw[%zero_offset] : buffer -> view<[%state_bound1]xf32> + %s0 = view.load %state_view0[%i0] : view<[%state_bound0]xf32> -> f32 + %s1 = view.load %state_view1[%i1] : view<[%state_bound1]xf32> -> f32 + %w0 = view.load %dw_view0[%i0] : view<[%state_bound0]xf32> -> f32 + %w1 = view.load %dw_view1[%i1] : view<[%state_bound1]xf32> -> f32 + %cb, %channel_bound = index.assume %c, %channels [lt(%c, %channels)] : index, index + %bias_view = buffer.view %dw_bias[%zero_offset] : buffer -> view<[%channel_bound]xf32> + %b = view.load %bias_view[%cb] : view<[%channel_bound]xf32> -> f32 + %is_q = index.cmp ult, %c, %q_size : index + %x = scf.if %is_q -> (f32) { + %qi, %q_bound = index.assume %c, %q_size [lt(%c, %q_size)] : index, index + %q_view = buffer.view %qraw[%zero_offset] : buffer -> view<[%q_bound]xf32> + %qv = view.load %q_view[%qi] : view<[%q_bound]xf32> -> f32 + scf.yield %qv : f32 + } else { + %k_size = index.sub %channels, %q_size : index + %ki_raw = index.sub %c, %q_size : index + %ki, %k_bound = index.assume %ki_raw, %k_size [lt(%ki_raw, %k_size)] : index, index + %k_view = buffer.view %kraw[%zero_offset] : buffer -> view<[%k_bound]xf32> + %kv = view.load %k_view[%ki] : view<[%k_bound]xf32> -> f32 + scf.yield %kv : f32 + } + %p00 = scalar.mulf %s0, %w0 : f32 + %p01 = scalar.mulf %s1, %w1 : f32 + %sum0 = scalar.addf %p00, %p01 : f32 + %d0 = scalar.addf %sum0, %b : f32 + %p10 = scalar.mulf %s1, %w0 : f32 + %p11 = scalar.mulf %x, %w1 : f32 + %sum1 = scalar.addf %p10, %p11 : f32 + %d1 = scalar.addf %sum1, %b : f32 + func.return %d0, %d1, %x, %s1 : f32, f32, f32, f32 +} + +// W[ic, o, tap] as f32. +func.def inline @ggml_zaya_cca_weight(%ic: index, %o: index, %tap: index, %channels: index, %weight: buffer) -> (f32) { + %c128 = index.constant 128 : index + %zero_offset = index.constant 0 : offset + %tap_size = index.mul %channels, %c128 : index + %c2 = index.constant 2 : index + %weight_size = index.mul %tap_size, %c2 : index + %row = index.mul %o, %c128 : index + %in_tap = index.add %row, %ic : index + %tap_base = index.mul %tap, %tap_size : index + %wi_raw = index.add %tap_base, %in_tap : index + %wi, %weight_bound = index.assume %wi_raw, %weight_size [lt(%wi_raw, %weight_size)] : index, index + %weight_view = buffer.view %weight[%zero_offset] : buffer -> view<[%weight_bound]xf16> + %wh = view.load %weight_view[%wi] : view<[%weight_bound]xf16> -> f16 + %wf = scalar.extf %wh : f16 to f32 + func.return %wf : f32 +} + +kernel.def target(@ggml_zaya_cca_conv_decode_gfx11_wave64) export("ggml_zaya_cca_conv_decode_f32") @ggml_zaya_cca_conv_decode_f32() { + %channels = config.get @ggml.zaya_cca_conv.channels : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c256 = index.constant 256 : index + %workgroups = index.div %channels, %c4 : index + kernel.launch.config workgroups(%workgroups, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%qraw: buffer, %kraw: buffer, %state: buffer, %dw: buffer, %dw_bias: buffer, %weight: buffer, %grp_bias: buffer, %output: buffer, %new_state: buffer) { + %channels = config.get @ggml.zaya_cca_conv.channels : index + %q_size = config.get @ggml.zaya_cca_conv.q_size : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c128 = index.constant 128 : index + %c0_f32 = scalar.constant 0.0 : f32 + %zero_offset = index.constant 0 : offset + %workgroup = kernel.workgroup.id : index + %wave0 = kernel.subgroup.id : index + %wave = index.assume %wave0 [range(%wave0, 0, 3)] : index + %lane0 = kernel.subgroup.lane.id : index + %lane = index.assume %lane0 [range(%lane0, 0, 63)] : index + %o_base = index.mul %workgroup, %c4 : index + %o0 = index.add %o_base, %wave : index + %valid = index.cmp ult, %o0, %channels : index + scf.if %valid { + %o, %output_bound = index.assume %o0, %channels [lt(%o0, %channels)] : index, index + %group = index.div %o, %c128 : index + %group_base = index.mul %group, %c128 : index + %ic_base = index.mul %lane, %c2 : index + %acc = scf.for %j = [%c0 to %c2 step %c1](%partial = %c0_f32 : f32) -> (f32) { + %ic = index.add %ic_base, %j : index + %cin = index.add %group_base, %ic : index + %d0, %d1, %x_in, %s1_in = func.call @ggml_zaya_cca_depthwise(%cin, %channels, %q_size, %qraw, %kraw, %state, %dw, %dw_bias) : (index, index, index, buffer, buffer, buffer, buffer, buffer) -> (f32, f32, f32, f32) + %w0 = func.call @ggml_zaya_cca_weight(%ic, %o, %c0, %channels, %weight) : (index, index, index, index, buffer) -> (f32) + %w1 = func.call @ggml_zaya_cca_weight(%ic, %o, %c1, %channels, %weight) : (index, index, index, index, buffer) -> (f32) + %t0 = scalar.mulf %w0, %d0 : f32 + %t1 = scalar.mulf %w1, %d1 : f32 + %both = scalar.addf %t0, %t1 : f32 + %next = scalar.addf %partial, %both : f32 + scf.yield %next : f32 + } + %sum = kernel.subgroup.reduce %acc : f32 + %is_lane_zero = index.cmp eq, %lane, %c0 : index + scf.if %is_lane_zero { + %grp_view = buffer.view %grp_bias[%zero_offset] : buffer -> view<[%output_bound]xf32> + %gb = view.load %grp_view[%o] : view<[%output_bound]xf32> -> f32 + %result = scalar.addf %sum, %gb : f32 + %out_view = buffer.view %output[%zero_offset] : buffer -> view<[%output_bound]xf32> + view.store %result, %out_view[%o] : f32, view<[%output_bound]xf32> + %e0, %e1, %x_own, %s1_own = func.call @ggml_zaya_cca_depthwise(%o, %channels, %q_size, %qraw, %kraw, %state, %dw, %dw_bias) : (index, index, index, buffer, buffer, buffer, buffer, buffer) -> (f32, f32, f32, f32) + %state_size = index.mul %channels, %c2 : index + %n0_raw = index.mul %o, %c2 : index + %n1_raw = index.add %n0_raw, %c1 : index + %n0, %new_bound0 = index.assume %n0_raw, %state_size [lt(%n0_raw, %state_size)] : index, index + %n1, %new_bound1 = index.assume %n1_raw, %state_size [lt(%n1_raw, %state_size)] : index, index + %new_view0 = buffer.view %new_state[%zero_offset] : buffer -> view<[%new_bound0]xf32> + %new_view1 = buffer.view %new_state[%zero_offset] : buffer -> view<[%new_bound1]xf32> + view.store %s1_own, %new_view0[%n0] : f32, view<[%new_bound0]xf32> + view.store %x_own, %new_view1[%n1] : f32, view<[%new_bound1]xf32> + } + } + kernel.return +} + +// Reference for the cases: one thread per output, a plain loop over the group's input channels. +kernel.def target(@ggml_zaya_cca_conv_decode_gfx11_wave64) export("ggml_zaya_cca_conv_decode_reference_f32") @ggml_zaya_cca_conv_decode_reference_f32() { + %channels = config.get @ggml.zaya_cca_conv.channels : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %workgroups = index.div %channels, %c64 : index + kernel.launch.config workgroups(%workgroups, %c1, %c1) workgroup_size(%c64, %c1, %c1) : index +} launch(%qraw: buffer, %kraw: buffer, %state: buffer, %dw: buffer, %dw_bias: buffer, %weight: buffer, %grp_bias: buffer, %output: buffer, %new_state: buffer) { + %channels = config.get @ggml.zaya_cca_conv.channels : index + %q_size = config.get @ggml.zaya_cca_conv.q_size : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %c0_f32 = scalar.constant 0.0 : f32 + %zero_offset = index.constant 0 : offset + %workgroup = kernel.workgroup.id : index + %item = kernel.workitem.id : index + %o_base = index.mul %workgroup, %c64 : index + %o0 = index.add %o_base, %item : index + %valid = index.cmp ult, %o0, %channels : index + scf.if %valid { + %o, %output_bound = index.assume %o0, %channels [lt(%o0, %channels)] : index, index + %group = index.div %o, %c128 : index + %group_base = index.mul %group, %c128 : index + %acc = scf.for %ic = [%c0 to %c128 step %c1](%partial = %c0_f32 : f32) -> (f32) { + %cin = index.add %group_base, %ic : index + %d0, %d1, %x_in, %s1_in = func.call @ggml_zaya_cca_depthwise(%cin, %channels, %q_size, %qraw, %kraw, %state, %dw, %dw_bias) : (index, index, index, buffer, buffer, buffer, buffer, buffer) -> (f32, f32, f32, f32) + %w0 = func.call @ggml_zaya_cca_weight(%ic, %o, %c0, %channels, %weight) : (index, index, index, index, buffer) -> (f32) + %w1 = func.call @ggml_zaya_cca_weight(%ic, %o, %c1, %channels, %weight) : (index, index, index, index, buffer) -> (f32) + %t0 = scalar.mulf %w0, %d0 : f32 + %t1 = scalar.mulf %w1, %d1 : f32 + %p0 = scalar.addf %partial, %t0 : f32 + %p1 = scalar.addf %p0, %t1 : f32 + scf.yield %p1 : f32 + } + %grp_view = buffer.view %grp_bias[%zero_offset] : buffer -> view<[%output_bound]xf32> + %gb = view.load %grp_view[%o] : view<[%output_bound]xf32> -> f32 + %result = scalar.addf %acc, %gb : f32 + %out_view = buffer.view %output[%zero_offset] : buffer -> view<[%output_bound]xf32> + view.store %result, %out_view[%o] : f32, view<[%output_bound]xf32> + %e0, %e1, %x_own, %s1_own = func.call @ggml_zaya_cca_depthwise(%o, %channels, %q_size, %qraw, %kraw, %state, %dw, %dw_bias) : (index, index, index, buffer, buffer, buffer, buffer, buffer) -> (f32, f32, f32, f32) + %state_size = index.mul %channels, %c2 : index + %n0_raw = index.mul %o, %c2 : index + %n1_raw = index.add %n0_raw, %c1 : index + %n0, %new_bound0 = index.assume %n0_raw, %state_size [lt(%n0_raw, %state_size)] : index, index + %n1, %new_bound1 = index.assume %n1_raw, %state_size [lt(%n1_raw, %state_size)] : index, index + %new_view0 = buffer.view %new_state[%zero_offset] : buffer -> view<[%new_bound0]xf32> + %new_view1 = buffer.view %new_state[%zero_offset] : buffer -> view<[%new_bound1]xf32> + view.store %s1_own, %new_view0[%n0] : f32, view<[%new_bound0]xf32> + view.store %x_own, %new_view1[%n1] : f32, view<[%new_bound1]xf32> + } + kernel.return +} + +// Cases: channels = 256 (two groups), q_size = 160 (the Q/K split falls inside group 1). +// Run with --config=ggml.zaya_cca_conv.channels=256 --config=ggml.zaya_cca_conv.q_size=160. + +// Exact: state 1, (w0, w1) = (0, 1), b = 1, x = 2, W = 0.25: d0 = 2, d1 = 3, so every output is +// 128 * 0.25 * 5 = 160 plus grp_bias[o] = o; the new state is (s1, x) = (1, 2). +check.case public @ggml_zaya_cca_conv_decode_exact_case { + %qraw = check.generate.fill value(2.0) : tensor<160xf32> + %kraw = check.generate.fill value(2.0) : tensor<96xf32> + %state = check.generate.fill value(1.0) : tensor<512xf32> + %dw = check.generate.iota offset(0.0) step(1.0) period(2) : tensor<512xf32> + %dw_bias = check.generate.fill value(1.0) : tensor<256xf32> + %weight = check.generate.fill value(0.25) : tensor<65536xf16> + %grp_bias = check.generate.iota offset(0.0) step(1.0) : tensor<256xf32> + %output = check.generate.fill value(-7.0) : tensor<256xf32> + %new_state = check.generate.fill value(-7.0) : tensor<512xf32> + %expected = check.generate.iota offset(160.0) step(1.0) : tensor<256xf32> + %expected_state = check.generate.iota offset(1.0) step(1.0) period(2) : tensor<512xf32> + kernel.launch @ggml_zaya_cca_conv_decode_f32[](%qraw, %kraw, %state, %dw, %dw_bias, %weight, %grp_bias, %output, %new_state) : [](tensor<160xf32>, tensor<96xf32>, tensor<512xf32>, tensor<512xf32>, tensor<256xf32>, tensor<65536xf16>, tensor<256xf32>, tensor<256xf32>, tensor<512xf32>) + check.expect.equal actual(%output) expected(%expected) : tensor<256xf32> + check.expect.equal actual(%new_state) expected(%expected_state) : tensor<512xf32> + check.return +} + +// Differential: independent random inputs for every binding against the reference kernel. +check.case public @ggml_zaya_cca_conv_decode_random_case { + %qraw_seed = check.param.seed base(7300000000011000033) count(1) : i64 + %kraw_seed = check.param.seed base(7300000000012000036) count(1) : i64 + %state_seed = check.param.seed base(7300000000013000039) count(1) : i64 + %dw_seed = check.param.seed base(7300000000014000042) count(1) : i64 + %dw_bias_seed = check.param.seed base(7300000000015000045) count(1) : i64 + %weight_seed = check.param.seed base(7300000000016000048) count(1) : i64 + %grp_bias_seed = check.param.seed base(7300000000017000051) count(1) : i64 + %qraw = check.generate.random.uniform seed(%qraw_seed) range(-1.0 to 1.0) : tensor<160xf32> + %kraw = check.generate.random.uniform seed(%kraw_seed) range(-1.0 to 1.0) : tensor<96xf32> + %state = check.generate.random.uniform seed(%state_seed) range(-1.0 to 1.0) : tensor<512xf32> + %dw = check.generate.random.uniform seed(%dw_seed) range(-1.0 to 1.0) : tensor<512xf32> + %dw_bias = check.generate.random.uniform seed(%dw_bias_seed) range(-0.5 to 0.5) : tensor<256xf32> + %weight = check.generate.random.uniform seed(%weight_seed) range(-0.25 to 0.25) : tensor<65536xf16> + %grp_bias = check.generate.random.uniform seed(%grp_bias_seed) range(-0.5 to 0.5) : tensor<256xf32> + %output = check.generate.fill value(-7.0) : tensor<256xf32> + %new_state = check.generate.fill value(-7.0) : tensor<512xf32> + %expected = check.generate.fill value(7.0) : tensor<256xf32> + %expected_state = check.generate.fill value(7.0) : tensor<512xf32> + kernel.launch @ggml_zaya_cca_conv_decode_reference_f32[](%qraw, %kraw, %state, %dw, %dw_bias, %weight, %grp_bias, %expected, %expected_state) : [](tensor<160xf32>, tensor<96xf32>, tensor<512xf32>, tensor<512xf32>, tensor<256xf32>, tensor<65536xf16>, tensor<256xf32>, tensor<256xf32>, tensor<512xf32>) + kernel.launch @ggml_zaya_cca_conv_decode_f32[](%qraw, %kraw, %state, %dw, %dw_bias, %weight, %grp_bias, %output, %new_state) : [](tensor<160xf32>, tensor<96xf32>, tensor<512xf32>, tensor<512xf32>, tensor<256xf32>, tensor<65536xf16>, tensor<256xf32>, tensor<256xf32>, tensor<512xf32>) + check.expect.close actual(%output) expected(%expected) atol(1.0e-4) rtol(1.0e-4) nan(same) : tensor<256xf32> + check.expect.equal actual(%new_state) expected(%expected_state) : tensor<512xf32> + check.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/zaya_cca_qk_norm_decode_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/zaya_cca_qk_norm_decode_f32.loom new file mode 100644 index 000000000000..3d4b00041489 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/zaya_cca_qk_norm_decode_f32.loom @@ -0,0 +1,237 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// ZAYA's CCA query/key mixing and normalization for one token (decode), in one dispatch. llama.cpp +// builds it from ~30 graph nodes (about 16 dispatches on HRX) between the CCA convolution and RoPE. +// With Q = Qraw, K = Kraw (the plain projections), C = the convolution output (Q part, then K part), +// head_dim D, n_head query heads and n_head_kv key heads (gqa = n_head / n_head_kv): +// query head h: v = C_q[h, d] + (Q[h, d] + K[h / gqa, d]) * 0.5 +// key head k: m = (sum_j Q[k * gqa + j, d]) * (1 / gqa) +// v = C_k[k, d] + (m + K[k, d]) * 0.5 +// then per head RMSNorm without weight (eps from the graph), and each key head times its scale. +// One workgroup per head, one thread per channel of the head. + +amdgpu.target @ggml_zaya_cca_qk_norm_gfx11_wave64 {subgroup_size = 64} + +config.decl @ggml.zaya_cca_qk_norm.head_dim : %value: index where [range(%value, 64, 256), mul(%value, 64)] + +config.decl @ggml.zaya_cca_qk_norm.n_head : %value: index where [range(%value, 1, 128)] + +config.decl @ggml.zaya_cca_qk_norm.n_head_kv : %value: index where [range(%value, 1, 128)] + +config.decl @ggml.zaya_cca_qk_norm.gqa : %value: index where [range(%value, 1, 128)] + +config.decl @ggml.zaya_cca_qk_norm.rms_epsilon : f32 + +// The pre-normalization value of channel %d of head %h. +func.def inline @ggml_zaya_cca_qk_value(%h: index, %d: index, %head_dim: index, %n_head: index, %n_head_kv: index, %gqa: index, %conv: buffer, %qraw: buffer, %kraw: buffer) -> (f32) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %zero_offset = index.constant 0 : offset + %half = scalar.constant 0.5 : f32 + %c0_f32 = scalar.constant 0.0 : f32 + %q_size = index.mul %n_head, %head_dim : index + %k_size = index.mul %n_head_kv, %head_dim : index + %conv_size = index.add %q_size, %k_size : index + %is_q = index.cmp ult, %h, %n_head : index + %value = scf.if %is_q -> (f32) { + %qi_row = index.mul %h, %head_dim : index + %qi_raw = index.add %qi_row, %d : index + %kv = index.div %h, %gqa : index + %ki_row = index.mul %kv, %head_dim : index + %ki_raw = index.add %ki_row, %d : index + %qi, %q_bound = index.assume %qi_raw, %q_size [lt(%qi_raw, %q_size)] : index, index + %ki, %k_bound = index.assume %ki_raw, %k_size [lt(%ki_raw, %k_size)] : index, index + %ci, %c_bound = index.assume %qi_raw, %conv_size [lt(%qi_raw, %conv_size)] : index, index + %q_view = buffer.view %qraw[%zero_offset] : buffer -> view<[%q_bound]xf32> + %k_view = buffer.view %kraw[%zero_offset] : buffer -> view<[%k_bound]xf32> + %c_view = buffer.view %conv[%zero_offset] : buffer -> view<[%c_bound]xf32> + %q = view.load %q_view[%qi] : view<[%q_bound]xf32> -> f32 + %k = view.load %k_view[%ki] : view<[%k_bound]xf32> -> f32 + %c = view.load %c_view[%ci] : view<[%c_bound]xf32> -> f32 + %qk = scalar.addf %q, %k : f32 + %mean = scalar.mulf %qk, %half : f32 + %v = scalar.addf %c, %mean : f32 + scf.yield %v : f32 + } else { + %kv = index.sub %h, %n_head : index + %group_base = index.mul %kv, %gqa : index + %sum = scf.for %j = [%c0 to %gqa step %c1](%partial = %c0_f32 : f32) -> (f32) { + %qh = index.add %group_base, %j : index + %qj_row = index.mul %qh, %head_dim : index + %qj_raw = index.add %qj_row, %d : index + %qj, %qj_bound = index.assume %qj_raw, %q_size [lt(%qj_raw, %q_size)] : index, index + %qj_view = buffer.view %qraw[%zero_offset] : buffer -> view<[%qj_bound]xf32> + %qv = view.load %qj_view[%qj] : view<[%qj_bound]xf32> -> f32 + %next = scalar.addf %partial, %qv : f32 + scf.yield %next : f32 + } + %gqa_i32 = index.cast %gqa : index to i32 + %gqa_f32 = scalar.sitofp %gqa_i32 : i32 to f32 + %one = scalar.constant 1.0 : f32 + %inv_gqa = scalar.divf %one, %gqa_f32 : f32 + %qmean = scalar.mulf %sum, %inv_gqa : f32 + %ki_row = index.mul %kv, %head_dim : index + %ki_raw = index.add %ki_row, %d : index + %ki, %k_bound = index.assume %ki_raw, %k_size [lt(%ki_raw, %k_size)] : index, index + %k_view = buffer.view %kraw[%zero_offset] : buffer -> view<[%k_bound]xf32> + %k = view.load %k_view[%ki] : view<[%k_bound]xf32> -> f32 + %ci_raw = index.add %q_size, %ki_raw : index + %ci, %c_bound = index.assume %ci_raw, %conv_size [lt(%ci_raw, %conv_size)] : index, index + %c_view = buffer.view %conv[%zero_offset] : buffer -> view<[%c_bound]xf32> + %c = view.load %c_view[%ci] : view<[%c_bound]xf32> -> f32 + %qk = scalar.addf %qmean, %k : f32 + %mean = scalar.mulf %qk, %half : f32 + %v = scalar.addf %c, %mean : f32 + scf.yield %v : f32 + } + func.return %value : f32 +} + +// Store channel %d of head %h, normalized by %scale (and the key scale for key heads). +func.def inline @ggml_zaya_cca_qk_store(%h: index, %d: index, %value: f32, %scale: f32, %head_dim: index, %n_head: index, %n_head_kv: index, %k_scale: buffer, %q_out: buffer, %k_out: buffer) { + %zero_offset = index.constant 0 : offset + %q_size = index.mul %n_head, %head_dim : index + %k_size = index.mul %n_head_kv, %head_dim : index + %normalized = scalar.mulf %value, %scale : f32 + %is_q = index.cmp ult, %h, %n_head : index + scf.if %is_q { + %qi_row = index.mul %h, %head_dim : index + %qi_raw = index.add %qi_row, %d : index + %qi, %q_bound = index.assume %qi_raw, %q_size [lt(%qi_raw, %q_size)] : index, index + %q_view = buffer.view %q_out[%zero_offset] : buffer -> view<[%q_bound]xf32> + view.store %normalized, %q_view[%qi] : f32, view<[%q_bound]xf32> + } else { + %kv_raw = index.sub %h, %n_head : index + %kv, %kv_bound = index.assume %kv_raw, %n_head_kv [lt(%kv_raw, %n_head_kv)] : index, index + %scale_view = buffer.view %k_scale[%zero_offset] : buffer -> view<[%kv_bound]xf32> + %ks = view.load %scale_view[%kv] : view<[%kv_bound]xf32> -> f32 + %scaled = scalar.mulf %normalized, %ks : f32 + %ki_row = index.mul %kv_raw, %head_dim : index + %ki_raw = index.add %ki_row, %d : index + %ki, %k_bound = index.assume %ki_raw, %k_size [lt(%ki_raw, %k_size)] : index, index + %k_view = buffer.view %k_out[%zero_offset] : buffer -> view<[%k_bound]xf32> + view.store %scaled, %k_view[%ki] : f32, view<[%k_bound]xf32> + } + func.return +} + +kernel.def target(@ggml_zaya_cca_qk_norm_gfx11_wave64) export("ggml_zaya_cca_qk_norm_decode_f32") @ggml_zaya_cca_qk_norm_decode_f32() { + %head_dim = config.get @ggml.zaya_cca_qk_norm.head_dim : index + %n_head = config.get @ggml.zaya_cca_qk_norm.n_head : index + %n_head_kv = config.get @ggml.zaya_cca_qk_norm.n_head_kv : index + %c1 = index.constant 1 : index + %heads = index.add %n_head, %n_head_kv : index + kernel.launch.config workgroups(%heads, %c1, %c1) workgroup_size(%head_dim, %c1, %c1) : index +} launch(%conv: buffer, %qraw: buffer, %kraw: buffer, %k_scale: buffer, %q_out: buffer, %k_out: buffer) { + %head_dim = config.get @ggml.zaya_cca_qk_norm.head_dim : index + %n_head = config.get @ggml.zaya_cca_qk_norm.n_head : index + %n_head_kv = config.get @ggml.zaya_cca_qk_norm.n_head_kv : index + %gqa = config.get @ggml.zaya_cca_qk_norm.gqa : index + %epsilon = config.get @ggml.zaya_cca_qk_norm.rms_epsilon : f32 + %h = kernel.workgroup.id : index + %d0 = kernel.workitem.id : index + %d, %d_bound = index.assume %d0, %head_dim [lt(%d0, %head_dim)] : index, index + %value = func.call @ggml_zaya_cca_qk_value(%h, %d, %head_dim, %n_head, %n_head_kv, %gqa, %conv, %qraw, %kraw) : (index, index, index, index, index, index, buffer, buffer, buffer) -> (f32) + %square = scalar.mulf %value, %value : f32 + %sum = kernel.workgroup.reduce %square : f32 + %head_dim_i32 = index.cast %head_dim : index to i32 + %head_dim_f32 = scalar.sitofp %head_dim_i32 : i32 to f32 + %mean = scalar.divf %sum, %head_dim_f32 : f32 + %biased = scalar.addf %mean, %epsilon : f32 + %scale = scalar.rsqrtf %biased : f32 + func.call @ggml_zaya_cca_qk_store(%h, %d, %value, %scale, %head_dim, %n_head, %n_head_kv, %k_scale, %q_out, %k_out) : (index, index, f32, f32, index, index, index, buffer, buffer, buffer) -> () + kernel.return +} + +// Reference for the cases: one thread per head, plain loops. +kernel.def target(@ggml_zaya_cca_qk_norm_gfx11_wave64) export("ggml_zaya_cca_qk_norm_decode_reference_f32") @ggml_zaya_cca_qk_norm_decode_reference_f32() { + %n_head = config.get @ggml.zaya_cca_qk_norm.n_head : index + %n_head_kv = config.get @ggml.zaya_cca_qk_norm.n_head_kv : index + %c1 = index.constant 1 : index + %heads = index.add %n_head, %n_head_kv : index + kernel.launch.config workgroups(%heads, %c1, %c1) workgroup_size(%c1, %c1, %c1) : index +} launch(%conv: buffer, %qraw: buffer, %kraw: buffer, %k_scale: buffer, %q_out: buffer, %k_out: buffer) { + %head_dim = config.get @ggml.zaya_cca_qk_norm.head_dim : index + %n_head = config.get @ggml.zaya_cca_qk_norm.n_head : index + %n_head_kv = config.get @ggml.zaya_cca_qk_norm.n_head_kv : index + %gqa = config.get @ggml.zaya_cca_qk_norm.gqa : index + %epsilon = config.get @ggml.zaya_cca_qk_norm.rms_epsilon : f32 + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c0_f32 = scalar.constant 0.0 : f32 + %h = kernel.workgroup.id : index + %sum = scf.for %d = [%c0 to %head_dim step %c1](%partial = %c0_f32 : f32) -> (f32) { + %v = func.call @ggml_zaya_cca_qk_value(%h, %d, %head_dim, %n_head, %n_head_kv, %gqa, %conv, %qraw, %kraw) : (index, index, index, index, index, index, buffer, buffer, buffer) -> (f32) + %sq = scalar.mulf %v, %v : f32 + %next = scalar.addf %partial, %sq : f32 + scf.yield %next : f32 + } + %head_dim_i32 = index.cast %head_dim : index to i32 + %head_dim_f32 = scalar.sitofp %head_dim_i32 : i32 to f32 + %mean = scalar.divf %sum, %head_dim_f32 : f32 + %biased = scalar.addf %mean, %epsilon : f32 + %root = scalar.sqrtf %biased : f32 + %one = scalar.constant 1.0 : f32 + %scale = scalar.divf %one, %root : f32 + scf.for %d = [%c0 to %head_dim step %c1] { + %v = func.call @ggml_zaya_cca_qk_value(%h, %d, %head_dim, %n_head, %n_head_kv, %gqa, %conv, %qraw, %kraw) : (index, index, index, index, index, index, buffer, buffer, buffer) -> (f32) + func.call @ggml_zaya_cca_qk_store(%h, %d, %v, %scale, %head_dim, %n_head, %n_head_kv, %k_scale, %q_out, %k_out) : (index, index, f32, f32, index, index, index, buffer, buffer, buffer) -> () + } + kernel.return +} + +// Cases: head_dim 64, 4 query heads, 2 key heads (gqa 2). Run with +// --config=ggml.zaya_cca_qk_norm.head_dim=64 --config=ggml.zaya_cca_qk_norm.n_head=4 +// --config=ggml.zaya_cca_qk_norm.n_head_kv=2 --config=ggml.zaya_cca_qk_norm.gqa=2 +// --config=ggml.zaya_cca_qk_norm.rms_epsilon=1e-24 + +// Constant heads normalize to 1: C = 1, Q = 2, K = 4 give v = 4 for query heads and +// 1 + (2 + 4) / 2 = 4 for key heads, so queries are 1 and keys are their scale (3). +check.case public @ggml_zaya_cca_qk_norm_decode_exact_case { + %conv = check.generate.fill value(1.0) : tensor<384xf32> + %qraw = check.generate.fill value(2.0) : tensor<256xf32> + %kraw = check.generate.fill value(4.0) : tensor<128xf32> + %k_scale = check.generate.fill value(3.0) : tensor<2xf32> + %q_out = check.generate.fill value(-7.0) : tensor<256xf32> + %k_out = check.generate.fill value(-7.0) : tensor<128xf32> + %q_expected = check.generate.fill value(1.0) : tensor<256xf32> + %k_expected = check.generate.fill value(3.0) : tensor<128xf32> + kernel.launch @ggml_zaya_cca_qk_norm_decode_f32[](%conv, %qraw, %kraw, %k_scale, %q_out, %k_out) : [](tensor<384xf32>, tensor<256xf32>, tensor<128xf32>, tensor<2xf32>, tensor<256xf32>, tensor<128xf32>) + check.expect.close actual(%q_out) expected(%q_expected) atol(1.0e-5) rtol(1.0e-5) nan(same) : tensor<256xf32> + check.expect.close actual(%k_out) expected(%k_expected) atol(1.0e-5) rtol(1.0e-5) nan(same) : tensor<128xf32> + check.return +} + +// Differential: independent random inputs against the reference kernel. +check.case public @ggml_zaya_cca_qk_norm_decode_random_case { + %conv_seed = check.param.seed base(7300000000000021011) count(1) : i64 + %qraw_seed = check.param.seed base(7300000000000021012) count(1) : i64 + %kraw_seed = check.param.seed base(7300000000000021013) count(1) : i64 + %scale_seed = check.param.seed base(7300000000000021014) count(1) : i64 + %conv = check.generate.random.uniform seed(%conv_seed) range(-1.0 to 1.0) : tensor<384xf32> + %qraw = check.generate.random.uniform seed(%qraw_seed) range(-1.0 to 1.0) : tensor<256xf32> + %kraw = check.generate.random.uniform seed(%kraw_seed) range(-1.0 to 1.0) : tensor<128xf32> + %k_scale = check.generate.random.uniform seed(%scale_seed) range(0.5 to 2.0) : tensor<2xf32> + %q_out = check.generate.fill value(-7.0) : tensor<256xf32> + %k_out = check.generate.fill value(-7.0) : tensor<128xf32> + %q_expected = check.generate.fill value(7.0) : tensor<256xf32> + %k_expected = check.generate.fill value(7.0) : tensor<128xf32> + kernel.launch @ggml_zaya_cca_qk_norm_decode_reference_f32[](%conv, %qraw, %kraw, %k_scale, %q_expected, %k_expected) : [](tensor<384xf32>, tensor<256xf32>, tensor<128xf32>, tensor<2xf32>, tensor<256xf32>, tensor<128xf32>) + kernel.launch @ggml_zaya_cca_qk_norm_decode_f32[](%conv, %qraw, %kraw, %k_scale, %q_out, %k_out) : [](tensor<384xf32>, tensor<256xf32>, tensor<128xf32>, tensor<2xf32>, tensor<256xf32>, tensor<128xf32>) + check.expect.close actual(%q_out) expected(%q_expected) atol(1.0e-5) rtol(1.0e-4) nan(same) : tensor<256xf32> + check.expect.close actual(%k_out) expected(%k_expected) atol(1.0e-5) rtol(1.0e-4) nan(same) : tensor<128xf32> + check.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/attention_metadata.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/attention_metadata.loom new file mode 100644 index 000000000000..a75a3fa0059d --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/attention_metadata.loom @@ -0,0 +1,184 @@ +// Copyright 2026 The IREE Authors +// +// Licensed under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception + +// Derives positions, separate K/V cache indices, and the dense causal-mask bit +// pattern on device from the owned runtime's compact context-base control word. +amdgpu.target @qwen_attention_metadata_gfx11_wave64 {subgroup_size = 64} + +kernel.def target(@qwen_attention_metadata_gfx11_wave64) export("qwen_attention_metadata") @qwen_attention_metadata(%token_count: index, %context_capacity: index) { + %bounded_context_capacity = index.assume %context_capacity [range(%context_capacity, 1, 32768)] : index + %c1 = index.constant 1 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %padded_context_capacity = index.add %bounded_context_capacity, %c255 : index + %key_workgroup_count = index.div %padded_context_capacity, %c256 : index + kernel.launch.config workgroups(%key_workgroup_count, %token_count, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %context_capacity: index, %control: buffer, %positions: buffer, %key_cache_indices: buffer, %value_cache_indices: buffer, %attention_mask: buffer) where [range(%token_count, 1, 2048)] { + %bounded_context_capacity = index.assume %context_capacity [range(%context_capacity, 1, 32768)] : index + %query0 = kernel.workgroup.id : index + %query, %launch_token_count = index.assume %query0, %token_count [lt(%query0, %token_count)] : index, index + %key_workgroup = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + %c0_i16 = scalar.constant 0 : i16 + // 0xfc00 is the F16 negative-infinity bit pattern. + %negative_infinity_i16 = scalar.constant -1024 : i16 + %c0_offset = index.constant 0 : offset + %control_noalias, %positions_noalias, %key_cache_indices_noalias, %value_cache_indices_noalias, %attention_mask_noalias = buffer.assume.noalias %control, %positions, %key_cache_indices, %value_cache_indices, %attention_mask : buffer, buffer, buffer, buffer, buffer + %control_view = buffer.view %control_noalias[%c0_offset] : buffer -> view<1xi32> + %context_base_raw = view.load %control_view[%c0] : view<1xi32> -> i32 + %context_base_i32 = scalar.assume %context_base_raw [range(%context_base_raw, 0, 32767)] : i32 + %context_base0 = index.cast %context_base_i32 : i32 to index + %context_base = index.assume %context_base0 [range(%context_base0, 0, 32767)] : index + %visible_count0 = index.add %context_base, %token_count : index + %visible_count, %launch_context_capacity = index.assume %visible_count0, %bounded_context_capacity [le(%visible_count0, %bounded_context_capacity)] : index, index + %key_workgroup_base = index.mul %key_workgroup, %c256 : index + %key0 = index.add %key_workgroup_base, %workitem : index + %valid_key = index.cmp ult, %key0, %launch_context_capacity : index + %safe_key0 = scf.select %valid_key, %key0, %c0 : index + %key = index.assume %safe_key0 [lt(%safe_key0, %launch_context_capacity)] : index + %mask_element_count = index.mul %launch_token_count, %launch_context_capacity : index + %mask_row_base = index.mul %query, %launch_context_capacity : index + %mask_element0 = index.add %mask_row_base, %key : index + %mask_element = index.assume %mask_element0 [lt(%mask_element0, %mask_element_count)] : index + %positions_view = buffer.view %positions_noalias[%c0_offset] : buffer -> view<[%launch_token_count]xi32> + %key_cache_indices_view = buffer.view %key_cache_indices_noalias[%c0_offset] : buffer -> view<[%launch_token_count]xi64> + %value_cache_indices_view = buffer.view %value_cache_indices_noalias[%c0_offset] : buffer -> view<[%launch_token_count]xi64> + %attention_mask_view = buffer.view %attention_mask_noalias[%c0_offset] : buffer -> view<[%mask_element_count]xi16> + %query_position0 = index.add %context_base, %query : index + %query_position = index.assume %query_position0 [lt(%query_position0, %visible_count)] : index + %causal_end = index.add %query_position, %c1 : index + %before_causal_end = index.cmp ult, %key, %causal_end : index + %before_visible_count = index.cmp ult, %key, %visible_count : index + %is_visible = scalar.andi %before_causal_end, %before_visible_count : i1 + %mask_bits = scf.select %is_visible, %c0_i16, %negative_infinity_i16 : i16 + scf.if %valid_key { + view.store %mask_bits, %attention_mask_view[%mask_element] : i16, view<[%mask_element_count]xi16> + } + %is_first_key_workgroup = index.cmp eq, %key_workgroup, %c0 : index + %is_first_workitem = index.cmp eq, %workitem, %c0 : index + %publishes_metadata = scalar.andi %is_first_key_workgroup, %is_first_workitem : i1 + scf.if %publishes_metadata { + %absolute_position_i32 = index.cast %query_position : index to i32 + %cache_index_i64 = index.cast %query_position : index to i64 + // This fixed no-ring experiment maps each logical position directly to + // the same physical row in both cache planes. + view.store %absolute_position_i32, %positions_view[%query] : i32, view<[%launch_token_count]xi32> + view.store %cache_index_i64, %key_cache_indices_view[%query] : i64, view<[%launch_token_count]xi64> + view.store %cache_index_i64, %value_cache_indices_view[%query] : i64, view<[%launch_token_count]xi64> + } + kernel.return +} + +// Publishes the single position and cache row needed by one decode issue. +// Decode attention bounds itself with the same request control word and does +// not need a materialized causal mask: every prior row is visible. +kernel.def target(@qwen_attention_metadata_gfx11_wave64) export("qwen_decode_attention_metadata") @qwen_decode_attention_metadata() { + %c1 = index.constant 1 : index + kernel.launch.config workgroups(%c1, %c1, %c1) workgroup_size(%c1, %c1, %c1) : index +} launch(%control: buffer, %positions: buffer, %key_cache_indices: buffer, %value_cache_indices: buffer) { + %c0 = index.constant 0 : index + %c0_offset = index.constant 0 : offset + %control_noalias, %positions_noalias, %key_cache_indices_noalias, %value_cache_indices_noalias = buffer.assume.noalias %control, %positions, %key_cache_indices, %value_cache_indices : buffer, buffer, buffer, buffer + %control_view = buffer.view %control_noalias[%c0_offset] : buffer -> view<1xi32> + %positions_view = buffer.view %positions_noalias[%c0_offset] : buffer -> view<1xi32> + %key_cache_indices_view = buffer.view %key_cache_indices_noalias[%c0_offset] : buffer -> view<1xi64> + %value_cache_indices_view = buffer.view %value_cache_indices_noalias[%c0_offset] : buffer -> view<1xi64> + %context_base_raw = view.load %control_view[%c0] : view<1xi32> -> i32 + %context_base = scalar.assume %context_base_raw [range(%context_base_raw, 0, 32767)] : i32 + %cache_index = scalar.extsi %context_base : i32 to i64 + view.store %context_base, %positions_view[%c0] : i32, view<1xi32> + view.store %cache_index, %key_cache_indices_view[%c0] : i64, view<1xi64> + view.store %cache_index, %value_cache_indices_view[%c0] : i64, view<1xi64> + kernel.return +} + +// Position zero exposes the causal edge directly as `[0, -inf]`. A second +// invocation proves that a nonzero context base advances all three metadata +// streams and makes every prior cache row visible. +check.case public @qwen_attention_metadata_causal_and_nonzero_base_case { + %c1 = check.literal value(1) : index + %c2 = check.literal value(2) : index + %five = check.literal value(5) : index + %zero_control = check.generate.fill value(0) : tensor<1xi32> + %zero_positions = check.generate.fill value(-1) : tensor<1xi32> + %zero_key_indices = check.generate.fill value(-1) : tensor<1xi64> + %zero_value_indices = check.generate.fill value(-1) : tensor<1xi64> + %zero_mask = check.generate.fill value(1) : tensor<2xi16> + %expected_zero = check.generate.fill value(0) : tensor<1xi32> + %expected_zero_indices = check.generate.fill value(0) : tensor<1xi64> + %expected_zero_mask = check.generate.iota offset(0) step(-1024) : tensor<2xi16> + kernel.launch @qwen_attention_metadata[%c1, %c2](%c1, %c2, %zero_control, %zero_positions, %zero_key_indices, %zero_value_indices, %zero_mask) : [index, index](index, index, tensor<1xi32>, tensor<1xi32>, tensor<1xi64>, tensor<1xi64>, tensor<2xi16>) + check.expect.equal actual(%zero_positions) expected(%expected_zero) : tensor<1xi32> + check.expect.equal actual(%zero_key_indices) expected(%expected_zero_indices) : tensor<1xi64> + check.expect.equal actual(%zero_value_indices) expected(%expected_zero_indices) : tensor<1xi64> + check.expect.equal actual(%zero_mask) expected(%expected_zero_mask) : tensor<2xi16> + %nonzero_control = check.generate.fill value(4) : tensor<1xi32> + %nonzero_positions = check.generate.fill value(-1) : tensor<1xi32> + %nonzero_key_indices = check.generate.fill value(-1) : tensor<1xi64> + %nonzero_value_indices = check.generate.fill value(-1) : tensor<1xi64> + %nonzero_mask = check.generate.fill value(1) : tensor<5xi16> + %expected_nonzero = check.generate.fill value(4) : tensor<1xi32> + %expected_nonzero_indices = check.generate.fill value(4) : tensor<1xi64> + %expected_nonzero_mask = check.generate.fill value(0) : tensor<5xi16> + kernel.launch @qwen_attention_metadata[%c1, %five](%c1, %five, %nonzero_control, %nonzero_positions, %nonzero_key_indices, %nonzero_value_indices, %nonzero_mask) : [index, index](index, index, tensor<1xi32>, tensor<1xi32>, tensor<1xi64>, tensor<1xi64>, tensor<5xi16>) + check.expect.equal actual(%nonzero_positions) expected(%expected_nonzero) : tensor<1xi32> + check.expect.equal actual(%nonzero_key_indices) expected(%expected_nonzero_indices) : tensor<1xi64> + check.expect.equal actual(%nonzero_value_indices) expected(%expected_nonzero_indices) : tensor<1xi64> + check.expect.equal actual(%nonzero_mask) expected(%expected_nonzero_mask) : tensor<5xi16> + check.return +} + +check.case public @qwen_attention_metadata_benchmark_case { + %token_count = check.param.choice values([32, 128, 512]) name("token_count") : index + %control = check.generate.fill value(0) : tensor<1xi32> + %positions = check.generate.fill value(-1) : tensor<[%token_count]xi32> + %key_cache_indices = check.generate.fill value(-1) : tensor<[%token_count]xi64> + %value_cache_indices = check.generate.fill value(-1) : tensor<[%token_count]xi64> + %attention_mask = check.generate.fill value(1) : tensor<[%token_count]x[%token_count]xi16> + %expected_positions = check.generate.iota offset(0) step(1) : tensor<[%token_count]xi32> + %expected_cache_indices = check.generate.iota offset(0) step(1) : tensor<[%token_count]xi64> + kernel.launch @qwen_attention_metadata[%token_count, %token_count](%token_count, %token_count, %control, %positions, %key_cache_indices, %value_cache_indices, %attention_mask) : [index, index](index, index, tensor<1xi32>, tensor<[%token_count]xi32>, tensor<[%token_count]xi64>, tensor<[%token_count]xi64>, tensor<[%token_count]x[%token_count]xi16>) + check.expect.equal actual(%positions) expected(%expected_positions) : tensor<[%token_count]xi32> + check.expect.equal actual(%key_cache_indices) expected(%expected_cache_indices) : tensor<[%token_count]xi64> + check.expect.equal actual(%value_cache_indices) expected(%expected_cache_indices) : tensor<[%token_count]xi64> + check.return +} + +check.case public @qwen_decode_attention_metadata_sequence_case { + %control0 = check.generate.fill value(0) : tensor<1xi32> + %positions0 = check.generate.fill value(-1) : tensor<1xi32> + %key_indices0 = check.generate.fill value(-1) : tensor<1xi64> + %value_indices0 = check.generate.fill value(-1) : tensor<1xi64> + %expected_position0 = check.generate.fill value(0) : tensor<1xi32> + %expected_indices0 = check.generate.fill value(0) : tensor<1xi64> + kernel.launch @qwen_decode_attention_metadata(%control0, %positions0, %key_indices0, %value_indices0) : (tensor<1xi32>, tensor<1xi32>, tensor<1xi64>, tensor<1xi64>) + check.expect.equal actual(%positions0) expected(%expected_position0) : tensor<1xi32> + check.expect.equal actual(%key_indices0) expected(%expected_indices0) : tensor<1xi64> + check.expect.equal actual(%value_indices0) expected(%expected_indices0) : tensor<1xi64> + + %control575 = check.generate.fill value(575) : tensor<1xi32> + %positions575 = check.generate.fill value(-1) : tensor<1xi32> + %key_indices575 = check.generate.fill value(-1) : tensor<1xi64> + %value_indices575 = check.generate.fill value(-1) : tensor<1xi64> + %expected_position575 = check.generate.fill value(575) : tensor<1xi32> + %expected_indices575 = check.generate.fill value(575) : tensor<1xi64> + kernel.launch @qwen_decode_attention_metadata(%control575, %positions575, %key_indices575, %value_indices575) : (tensor<1xi32>, tensor<1xi32>, tensor<1xi64>, tensor<1xi64>) + check.expect.equal actual(%positions575) expected(%expected_position575) : tensor<1xi32> + check.expect.equal actual(%key_indices575) expected(%expected_indices575) : tensor<1xi64> + check.expect.equal actual(%value_indices575) expected(%expected_indices575) : tensor<1xi64> + check.return +} + +check.benchmark<@qwen_attention_metadata_causal_and_nonzero_base_case> @qwen_attention_metadata_causal_and_nonzero_base + +check.benchmark<@qwen_attention_metadata_benchmark_case> @qwen_attention_metadata_prefill_32 {token_count = 32} + +check.benchmark<@qwen_attention_metadata_benchmark_case> @qwen_attention_metadata_prefill_128 {token_count = 128} + +check.benchmark<@qwen_attention_metadata_benchmark_case> @qwen_attention_metadata_prefill_512 {token_count = 512} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/attention_metadata_bringup_workaround.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/attention_metadata_bringup_workaround.loom new file mode 100644 index 000000000000..6c06d4baa682 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/attention_metadata_bringup_workaround.loom @@ -0,0 +1,170 @@ +// Copyright 2026 The IREE Authors +// +// Licensed under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception + +// Temporary, non-sanctioned one-function kernel for Qwen bring-up. This is not +// a metadata-kernel framework or a second Loom authoring path. The owned +// runtime carries one compact context-base control word and needs positions, +// separate K/V cache indices, and the dense causal-mask bit pattern derived on +// device before attention begins. Delete this fork when the Qwen kernel corpus +// provides the canonical producer. +amdgpu.target @qwen_attention_metadata_gfx11_wave64 {subgroup_size = 64} + +kernel.def target(@qwen_attention_metadata_gfx11_wave64) export("qwen_attention_metadata_bringup_workaround") @qwen_attention_metadata_bringup_workaround(%token_count: index, %context_capacity: index) { + %one = index.constant 1 : index + %twofiftyfive = index.constant 255 : index + %twofiftysix = index.constant 256 : index + %padded_context_capacity = index.add %context_capacity, %twofiftyfive : index + %key_workgroup_count = index.div %padded_context_capacity, %twofiftysix : index + kernel.launch.config workgroups(%key_workgroup_count, %token_count, %one) workgroup_size(%twofiftysix, %one, %one) : index +} launch(%token_count: index, %context_capacity: index, %control: buffer, %positions: buffer, %key_cache_indices: buffer, %value_cache_indices: buffer, %attention_mask: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_context_capacity = index.assume %context_capacity [range(%context_capacity, 1, 32768)] : index + %query0 = kernel.workgroup.id : index + %query, %launch_token_count = index.assume %query0, %bounded_token_count [lt(%query0, %bounded_token_count)] : index, index + %key_workgroup = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %zero = index.constant 0 : index + %one = index.constant 1 : index + %twofiftysix = index.constant 256 : index + %zero_i16 = scalar.constant 0 : i16 + // 0xfc00 is the F16 negative-infinity bit pattern. + %negative_infinity_i16 = scalar.constant -1024 : i16 + %zero_offset = index.constant 0 : offset + %control_noalias, %positions_noalias, %key_cache_indices_noalias, %value_cache_indices_noalias, %attention_mask_noalias = buffer.assume.noalias %control, %positions, %key_cache_indices, %value_cache_indices, %attention_mask : buffer, buffer, buffer, buffer, buffer + %control_view = buffer.view %control_noalias[%zero_offset] : buffer -> view<1xi32> + %context_base_raw = view.load %control_view[%zero] : view<1xi32> -> i32 + %context_base_i32 = scalar.assume %context_base_raw [range(%context_base_raw, 0, 32767)] : i32 + %context_base0 = index.cast %context_base_i32 : i32 to index + %context_base = index.assume %context_base0 [range(%context_base0, 0, 32767)] : index + %visible_count0 = index.add %context_base, %bounded_token_count : index + %visible_count, %launch_context_capacity = index.assume %visible_count0, %bounded_context_capacity [le(%visible_count0, %bounded_context_capacity)] : index, index + %key_workgroup_base = index.mul %key_workgroup, %twofiftysix : index + %key0 = index.add %key_workgroup_base, %workitem : index + %valid_key = index.cmp ult, %key0, %launch_context_capacity : index + %safe_key0 = scf.select %valid_key, %key0, %zero : index + %key = index.assume %safe_key0 [lt(%safe_key0, %launch_context_capacity)] : index + %mask_element_count = index.mul %launch_token_count, %launch_context_capacity : index + %mask_row_base = index.mul %query, %launch_context_capacity : index + %mask_element0 = index.add %mask_row_base, %key : index + %mask_element = index.assume %mask_element0 [lt(%mask_element0, %mask_element_count)] : index + %positions_view = buffer.view %positions_noalias[%zero_offset] : buffer -> view<[%launch_token_count]xi32> + %key_cache_indices_view = buffer.view %key_cache_indices_noalias[%zero_offset] : buffer -> view<[%launch_token_count]xi64> + %value_cache_indices_view = buffer.view %value_cache_indices_noalias[%zero_offset] : buffer -> view<[%launch_token_count]xi64> + %attention_mask_view = buffer.view %attention_mask_noalias[%zero_offset] : buffer -> view<[%mask_element_count]xi16> + %query_position0 = index.add %context_base, %query : index + %query_position = index.assume %query_position0 [lt(%query_position0, %visible_count)] : index + %causal_end = index.add %query_position, %one : index + %before_causal_end = index.cmp ult, %key, %causal_end : index + %before_visible_count = index.cmp ult, %key, %visible_count : index + %is_visible = scalar.andi %before_causal_end, %before_visible_count : i1 + %mask_bits = scf.select %is_visible, %zero_i16, %negative_infinity_i16 : i16 + scf.if %valid_key { + view.store %mask_bits, %attention_mask_view[%mask_element] : i16, view<[%mask_element_count]xi16> + } + %is_first_key_workgroup = index.cmp eq, %key_workgroup, %zero : index + %is_first_workitem = index.cmp eq, %workitem, %zero : index + %publishes_metadata = scalar.andi %is_first_key_workgroup, %is_first_workitem : i1 + scf.if %publishes_metadata { + %absolute_position_i32 = index.cast %query_position : index to i32 + %cache_index_i64 = index.cast %query_position : index to i64 + // This fixed no-ring experiment maps each logical position directly to + // the same physical row in both cache planes. + view.store %absolute_position_i32, %positions_view[%query] : i32, view<[%launch_token_count]xi32> + view.store %cache_index_i64, %key_cache_indices_view[%query] : i64, view<[%launch_token_count]xi64> + view.store %cache_index_i64, %value_cache_indices_view[%query] : i64, view<[%launch_token_count]xi64> + } + kernel.return +} + +// Position zero exposes the causal edge directly as `[0, -inf]`. A second +// invocation proves that a nonzero context base advances all three metadata +// streams and makes every prior cache row visible. +check.case public @qwen_attention_metadata_causal_and_nonzero_base_case { + %one = check.literal value(1) : index + %two = check.literal value(2) : index + %five = check.literal value(5) : index + %zero_control = check.generate.fill value(0) : tensor<1xi32> + %zero_positions = check.generate.fill value(-1) : tensor<1xi32> + %zero_key_indices = check.generate.fill value(-1) : tensor<1xi64> + %zero_value_indices = check.generate.fill value(-1) : tensor<1xi64> + %zero_mask = check.generate.fill value(1) : tensor<2xi16> + %expected_zero = check.generate.fill value(0) : tensor<1xi32> + %expected_zero_indices = check.generate.fill value(0) : tensor<1xi64> + %expected_zero_mask = check.generate.iota offset(0) step(-1024) : tensor<2xi16> + kernel.launch @qwen_attention_metadata_bringup_workaround[%one, %two](%one, %two, %zero_control, %zero_positions, %zero_key_indices, %zero_value_indices, %zero_mask) : [index, index](index, index, tensor<1xi32>, tensor<1xi32>, tensor<1xi64>, tensor<1xi64>, tensor<2xi16>) + check.expect.equal actual(%zero_positions) expected(%expected_zero) : tensor<1xi32> + check.expect.equal actual(%zero_key_indices) expected(%expected_zero_indices) : tensor<1xi64> + check.expect.equal actual(%zero_value_indices) expected(%expected_zero_indices) : tensor<1xi64> + check.expect.equal actual(%zero_mask) expected(%expected_zero_mask) : tensor<2xi16> + %nonzero_control = check.generate.fill value(4) : tensor<1xi32> + %nonzero_positions = check.generate.fill value(-1) : tensor<1xi32> + %nonzero_key_indices = check.generate.fill value(-1) : tensor<1xi64> + %nonzero_value_indices = check.generate.fill value(-1) : tensor<1xi64> + %nonzero_mask = check.generate.fill value(1) : tensor<5xi16> + %expected_nonzero = check.generate.fill value(4) : tensor<1xi32> + %expected_nonzero_indices = check.generate.fill value(4) : tensor<1xi64> + %expected_nonzero_mask = check.generate.fill value(0) : tensor<5xi16> + kernel.launch @qwen_attention_metadata_bringup_workaround[%one, %five](%one, %five, %nonzero_control, %nonzero_positions, %nonzero_key_indices, %nonzero_value_indices, %nonzero_mask) : [index, index](index, index, tensor<1xi32>, tensor<1xi32>, tensor<1xi64>, tensor<1xi64>, tensor<5xi16>) + check.expect.equal actual(%nonzero_positions) expected(%expected_nonzero) : tensor<1xi32> + check.expect.equal actual(%nonzero_key_indices) expected(%expected_nonzero_indices) : tensor<1xi64> + check.expect.equal actual(%nonzero_value_indices) expected(%expected_nonzero_indices) : tensor<1xi64> + check.expect.equal actual(%nonzero_mask) expected(%expected_nonzero_mask) : tensor<5xi16> + check.return +} + +check.case public @qwen_attention_metadata_benchmark_case { + %token_count = check.param.choice values([32, 128, 512]) name("token_count") : index + %control = check.generate.fill value(0) : tensor<1xi32> + %positions = check.generate.fill value(-1) : tensor<[%token_count]xi32> + %key_cache_indices = check.generate.fill value(-1) : tensor<[%token_count]xi64> + %value_cache_indices = check.generate.fill value(-1) : tensor<[%token_count]xi64> + %attention_mask = check.generate.fill value(1) : tensor<[%token_count]x[%token_count]xi16> + %expected_positions = check.generate.iota offset(0) step(1) : tensor<[%token_count]xi32> + %expected_cache_indices = check.generate.iota offset(0) step(1) : tensor<[%token_count]xi64> + kernel.launch @qwen_attention_metadata_bringup_workaround[%token_count, %token_count](%token_count, %token_count, %control, %positions, %key_cache_indices, %value_cache_indices, %attention_mask) : [index, index](index, index, tensor<1xi32>, tensor<[%token_count]xi32>, tensor<[%token_count]xi64>, tensor<[%token_count]xi64>, tensor<[%token_count]x[%token_count]xi16>) + check.expect.equal actual(%positions) expected(%expected_positions) : tensor<[%token_count]xi32> + check.expect.equal actual(%key_cache_indices) expected(%expected_cache_indices) : tensor<[%token_count]xi64> + check.expect.equal actual(%value_cache_indices) expected(%expected_cache_indices) : tensor<[%token_count]xi64> + check.return +} + +// Locked decode geometry: one query row addressing the explicit 768-row KV +// bucket recovered from the live llama.cpp graph. +check.case public @qwen_attention_metadata_decode_768_case { + %token_count = check.literal value(1) : index + %context_capacity = check.literal value(768) : index + %control = check.generate.fill value(767) : tensor<1xi32> + %positions = check.generate.fill value(-1) : tensor<1xi32> + %key_cache_indices = check.generate.fill value(-1) : tensor<1xi64> + %value_cache_indices = check.generate.fill value(-1) : tensor<1xi64> + %attention_mask = check.generate.fill value(1) : tensor<768xi16> + kernel.launch @qwen_attention_metadata_bringup_workaround[%token_count, %context_capacity](%token_count, %context_capacity, %control, %positions, %key_cache_indices, %value_cache_indices, %attention_mask) : [index, index](index, index, tensor<1xi32>, tensor<1xi32>, tensor<1xi64>, tensor<1xi64>, tensor<768xi16>) + check.return +} + +check.case public @qwen_attention_metadata_prefill_512_case { + %token_count = check.literal value(512) : index + %context_capacity = check.literal value(512) : index + %control = check.generate.fill value(0) : tensor<1xi32> + %positions = check.generate.fill value(-1) : tensor<512xi32> + %key_cache_indices = check.generate.fill value(-1) : tensor<512xi64> + %value_cache_indices = check.generate.fill value(-1) : tensor<512xi64> + %attention_mask = check.generate.fill value(1) : tensor<512x512xi16> + kernel.launch @qwen_attention_metadata_bringup_workaround[%token_count, %context_capacity](%token_count, %context_capacity, %control, %positions, %key_cache_indices, %value_cache_indices, %attention_mask) : [index, index](index, index, tensor<1xi32>, tensor<512xi32>, tensor<512xi64>, tensor<512xi64>, tensor<512x512xi16>) + check.return +} + +check.benchmark<@qwen_attention_metadata_causal_and_nonzero_base_case> @qwen_attention_metadata_causal_and_nonzero_base + +check.benchmark<@qwen_attention_metadata_decode_768_case> @qwen_attention_metadata_decode_768 + +check.benchmark<@qwen_attention_metadata_prefill_512_case> @qwen_attention_metadata_model_prefill_512 + +check.benchmark<@qwen_attention_metadata_benchmark_case> @qwen_attention_metadata_prefill_32 {token_count = 32} + +check.benchmark<@qwen_attention_metadata_benchmark_case> @qwen_attention_metadata_prefill_128 {token_count = 128} + +check.benchmark<@qwen_attention_metadata_benchmark_case> @qwen_attention_metadata_prefill_512 {token_count = 512} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/attention_state_initialize.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/attention_state_initialize.loom new file mode 100644 index 000000000000..9754207c7210 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/attention_state_initialize.loom @@ -0,0 +1,112 @@ +// Copyright 2026 The IREE Authors +// +// Licensed under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception + +amdgpu.target @qwen_attention_state_gfx11_wave64 {subgroup_size = 64} + +// Captures the position before the metadata producer overwrites the positions +// buffer with the canonical sequence used by attention. +kernel.def target(@qwen_attention_state_gfx11_wave64) export("qwen_attention_context_base_capture") @qwen_attention_context_base_capture() { + %c1 = index.constant 1 : index + kernel.launch.config workgroups(%c1, %c1, %c1) workgroup_size(%c1, %c1, %c1) : index +} launch(%positions: buffer, %control: buffer) { + %c0 = index.constant 0 : index + %c0_offset = index.constant 0 : offset + %positions_noalias, %control_noalias = buffer.assume.noalias %positions, %control : buffer, buffer + %positions_view = buffer.view %positions_noalias[%c0_offset] : buffer -> view<1xi32> + %control_view = buffer.view %control_noalias[%c0_offset] : buffer -> view<1xi32> + %context_base = view.load %positions_view[%c0] : view<1xi32> -> i32 + view.store %context_base, %control_view[%c0] : i32, view<1xi32> + kernel.return +} + +// Decode has one query row, so the input position is already the absolute +// position consumed by attention. Publish its cache indices and causal mask +// while initializing every self-resetting completion counter in the replay. +kernel.def target(@qwen_attention_state_gfx11_wave64) export("qwen_attention_decode_state_initialize") @qwen_attention_decode_state_initialize(%context_capacity: index, %completion_counter_count: index) { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%c1, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%context_capacity: index, %completion_counter_count: index, %positions: buffer, %key_cache_indices: buffer, %value_cache_indices: buffer, %attention_mask: buffer, %completion_counters: buffer) { + %bounded_context_capacity = index.assume %context_capacity [range(%context_capacity, 1, 32768)] : index + %counter_count = index.assume %completion_counter_count [range(%completion_counter_count, 1, 64)] : index + %workitem0 = kernel.workitem.id : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 255)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + %c0_i16 = scalar.constant 0 : i16 + %c0_i32 = scalar.constant 0 : i32 + // 0xfc00 is the F16 negative-infinity bit pattern. + %negative_infinity_i16 = scalar.constant -1024 : i16 + %c0_offset = index.constant 0 : offset + %positions_noalias, %key_indices_noalias, %value_indices_noalias, %mask_noalias, %counters_noalias = buffer.assume.noalias %positions, %key_cache_indices, %value_cache_indices, %attention_mask, %completion_counters : buffer, buffer, buffer, buffer, buffer + %positions_view = buffer.view %positions_noalias[%c0_offset] : buffer -> view<1xi32> + %key_indices_view = buffer.view %key_indices_noalias[%c0_offset] : buffer -> view<1xi64> + %value_indices_view = buffer.view %value_indices_noalias[%c0_offset] : buffer -> view<1xi64> + %mask_view = buffer.view %mask_noalias[%c0_offset] : buffer -> view<[%bounded_context_capacity]xi16> + %counters_view = buffer.view %counters_noalias[%c0_offset] : buffer -> view<[%counter_count]xi32> + %context_base_raw = view.load %positions_view[%c0] : view<1xi32> -> i32 + %context_base_i32 = scalar.assume %context_base_raw [range(%context_base_raw, 0, 32767)] : i32 + %context_base0 = index.cast %context_base_i32 : i32 to index + %context_base = index.assume %context_base0 [range(%context_base0, 0, 32767)] : index + %visible_count0 = index.add %context_base, %c1 : index + %visible_count, %launch_context_capacity = index.assume %visible_count0, %bounded_context_capacity [le(%visible_count0, %bounded_context_capacity)] : index, index + scf.for %key_base = [%c0 to %launch_context_capacity step %c256] { + %key0 = index.add %key_base, %workitem : index + %valid_key = index.cmp ult, %key0, %launch_context_capacity : index + %safe_key0 = scf.select %valid_key, %key0, %c0 : index + %key = index.assume %safe_key0 [lt(%safe_key0, %launch_context_capacity)] : index + %is_visible = index.cmp ult, %key, %visible_count : index + %mask_bits = scf.select %is_visible, %c0_i16, %negative_infinity_i16 : i16 + scf.if %valid_key { + view.store %mask_bits, %mask_view[%key] : i16, view<[%bounded_context_capacity]xi16> + } + } + %is_first = index.cmp eq, %workitem, %c0 : index + scf.if %is_first { + %cache_index = index.cast %context_base : index to i64 + view.store %cache_index, %key_indices_view[%c0] : i64, view<1xi64> + view.store %cache_index, %value_indices_view[%c0] : i64, view<1xi64> + } + %is_counter = index.cmp ult, %workitem, %counter_count : index + scf.if %is_counter { + %counter = index.assume %workitem [lt(%workitem, %counter_count)] : index + view.store %c0_i32, %counters_view[%counter] : i32, view<[%counter_count]xi32> + } + kernel.return +} + +check.case public @qwen_attention_context_base_capture_case { + %positions = check.generate.fill value(7) : tensor<1xi32> + %control = check.generate.fill value(-1) : tensor<1xi32> + %expected = check.generate.fill value(7) : tensor<1xi32> + kernel.launch @qwen_attention_context_base_capture(%positions, %control) : (tensor<1xi32>, tensor<1xi32>) + check.expect.equal actual(%control) expected(%expected) : tensor<1xi32> + check.return +} + +check.case public @qwen_attention_decode_state_initialize_case { + %context_capacity = check.literal value(768) : index + %counter_count = check.literal value(56) : index + %positions = check.generate.fill value(767) : tensor<1xi32> + %key_cache_indices = check.generate.fill value(-1) : tensor<1xi64> + %value_cache_indices = check.generate.fill value(-1) : tensor<1xi64> + %attention_mask = check.generate.fill value(1) : tensor<768xi16> + %counters = check.generate.fill value(-1) : tensor<56xi32> + %expected_cache_indices = check.generate.fill value(767) : tensor<1xi64> + %expected_mask = check.generate.fill value(0) : tensor<768xi16> + %expected_counters = check.generate.fill value(0) : tensor<56xi32> + kernel.launch @qwen_attention_decode_state_initialize[%context_capacity, %counter_count](%context_capacity, %counter_count, %positions, %key_cache_indices, %value_cache_indices, %attention_mask, %counters) : [index, index](index, index, tensor<1xi32>, tensor<1xi64>, tensor<1xi64>, tensor<768xi16>, tensor<56xi32>) + check.expect.equal actual(%key_cache_indices) expected(%expected_cache_indices) : tensor<1xi64> + check.expect.equal actual(%value_cache_indices) expected(%expected_cache_indices) : tensor<1xi64> + check.expect.equal actual(%attention_mask) expected(%expected_mask) : tensor<768xi16> + check.expect.equal actual(%counters) expected(%expected_counters) : tensor<56xi32> + check.return +} + +check.benchmark<@qwen_attention_context_base_capture_case> @qwen_attention_context_base_capture_benchmark + +check.benchmark<@qwen_attention_decode_state_initialize_case> @qwen_attention_decode_state_initialize_benchmark diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/manifest.json b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/manifest.json new file mode 100644 index 000000000000..da6f8462d037 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/manifest.json @@ -0,0 +1,195 @@ +{ + "schema": "ggml-hrx-kernel-corpus-v1", + "upstream_revision": "local", + "files": [ + { + "path": "token_embedding_bringup_workaround.loom" + }, + { + "path": "attention_state_initialize.loom" + }, + { + "path": "attention_metadata_bringup_workaround.loom" + } + ], + "exports": [ + { + "name": "qwen_attention_context_base_capture", + "symbol": "qwen_attention_context_base_capture", + "source": "attention_state_initialize.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "positions", + "control" + ], + "binding_access": [ + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "attention_state_initialize.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [], + "family": "qwen" + }, + { + "name": "qwen_attention_decode_state_initialize", + "symbol": "qwen_attention_decode_state_initialize", + "source": "attention_state_initialize.loom", + "workload_parameters": [ + { + "name": "context_capacity", + "type": "index" + }, + { + "name": "completion_counter_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "context_capacity", + "type": "index" + }, + { + "name": "completion_counter_count", + "type": "index" + } + ], + "bindings": [ + "positions", + "key_cache_indices", + "value_cache_indices", + "attention_mask", + "completion_counters" + ], + "binding_access": [ + "read", + "write", + "write", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "attention_state_initialize.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [], + "family": "qwen" + }, + { + "name": "qwen_attention_metadata_bringup_workaround", + "symbol": "qwen_attention_metadata_bringup_workaround", + "source": "attention_metadata_bringup_workaround.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "context_capacity", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "context_capacity", + "type": "index" + } + ], + "bindings": [ + "control", + "positions", + "key_cache_indices", + "value_cache_indices", + "attention_mask" + ], + "binding_access": [ + "read", + "read_write", + "read_write", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "attention_metadata_bringup_workaround.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [], + "family": "qwen" + }, + { + "name": "qwen_token_embedding_q4k_bringup_workaround", + "symbol": "qwen_token_embedding_q4k_bringup_workaround", + "source": "token_embedding_bringup_workaround.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "vocabulary_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "vocabulary_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ], + "bindings": [ + "token_ids", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "token_embedding_bringup_workaround.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [], + "family": "qwen" + } + ], + "link_modules": [], + "plan_cases": [] +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/token_embedding_bringup_workaround.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/token_embedding_bringup_workaround.loom new file mode 100644 index 000000000000..2a4db9725a33 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/token_embedding_bringup_workaround.loom @@ -0,0 +1,224 @@ +// Copyright 2026 The IREE Authors +// +// Licensed under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception + +// Temporary, non-sanctioned one-function kernel for Qwen bring-up. This is not +// an embedding-kernel framework, a generator, or a second Loom authoring path. +// It gathers token rows directly from the model's unmodified GGUF Q4_K payload +// and decodes them to the owned F32 hidden-state layout. Delete this file when +// the Qwen kernel corpus provides the canonical token-embedding producer. +// +// The owned model row contract is: +// Q4_K: [vocabulary row][hidden_size / 256 blocks][144 bytes] +// -> [hidden_size x f32] +amdgpu.target @qwen_token_embedding_gfx11_wave64 {subgroup_size = 64} + +kernel.def target(@qwen_token_embedding_gfx11_wave64) export("qwen_token_embedding_q4k_bringup_workaround") @qwen_token_embedding_q4k_bringup_workaround(%token_count: index, %vocabulary_count: index, %hidden_size: index) { + %one = index.constant 1 : index + %onethousandtwentyfour = index.constant 1024 : index + %workgroups_per_token = index.div %hidden_size, %onethousandtwentyfour : index + %workgroup_size = index.constant 256 : index + kernel.launch.config workgroups(%workgroups_per_token, %token_count, %one) workgroup_size(%workgroup_size, %one, %one) : index +} launch(%token_count: index, %vocabulary_count: index, %hidden_size: index, %token_ids: buffer, %weight: buffer, %output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_vocabulary_count = index.assume %vocabulary_count [range(%vocabulary_count, 1, 262144)] : index + %bounded_hidden_size = index.assume %hidden_size [range(%hidden_size, 2048, 3072), mul(%hidden_size, 1024)] : index + %workgroup = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %zero = index.constant 0 : index + %two = index.constant 2 : index + %three = index.constant 3 : index + %four = index.constant 4 : index + %eight = index.constant 8 : index + %sixtyfour = index.constant 64 : index + %twofiftysix = index.constant 256 : index + %packets_per_token0 = index.div %bounded_hidden_size, %four : index + %packets_per_token = index.assume %packets_per_token0 [range(%packets_per_token0, 512, 768)] : index + %q4_block_count0 = index.div %bounded_hidden_size, %twofiftysix : index + %q4_block_count = index.assume %q4_block_count0 [range(%q4_block_count0, 8, 12)] : index + %zero_offset = index.constant 0 : offset + %block_bytes = index.constant 144 : offset + %scale_offset = index.constant 4 : offset + %code_offset = index.constant 16 : offset + %row_bytes = index.scale %q4_block_count, %block_bytes : index, offset -> offset + %two_i32 = scalar.constant 2 : i32 + %four_i32 = scalar.constant 4 : i32 + %fifteen = vector.constant 15 : vector<1xi32> + %fortyeight = vector.constant 48 : vector<1xi32> + %q4_mask = vector.constant 252645135 : vector<1xi32> + %token, %launch_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %row_packet0 = index.madd %workgroup, %twofiftysix, %workitem : index + %row_packet, %launch_packets_per_token = index.assume %row_packet0, %packets_per_token [range(%row_packet0, 0, 767), lt(%row_packet0, %packets_per_token)] : index, index + %q4_block0 = index.div %row_packet, %sixtyfour : index + %q4_block, %weight_q4_block_count = index.assume %q4_block0, %q4_block_count [range(%q4_block0, 0, 11), lt(%q4_block0, %q4_block_count)] : index, index + %block_packet0 = index.rem %row_packet, %sixtyfour : index + %block_packet = index.assume %block_packet0 [range(%block_packet0, 0, 63)] : index + %q4_group0 = index.div %block_packet, %eight : index + %q4_group = index.assume %q4_group0 [range(%q4_group0, 0, 7)] : index + %load_packet0 = index.rem %block_packet, %eight : index + %load_packet = index.assume %load_packet0 [range(%load_packet0, 0, 7)] : index + %token_ids_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %token_ids, %weight, %output : buffer, buffer, buffer + %token_ids_view = buffer.view %token_ids_noalias[%zero_offset] : buffer -> view<[%launch_token_count]xi32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%launch_token_count]x[%bounded_hidden_size]xf32> + // Token IDs are validated at the request boundary. This trusted consumer + // carries that fact into row addressing instead of adding a fallback row. + %token_id_raw = view.load %token_ids_view[%token] : view<[%launch_token_count]xi32> -> i32 + %token_id_i32 = scalar.assume %token_id_raw [range(%token_id_raw, 0, 262143)] : i32 + %token_id0 = index.cast %token_id_i32 : i32 to index + %token_id, %weight_vocabulary_count = index.assume %token_id0, %bounded_vocabulary_count [lt(%token_id0, %bounded_vocabulary_count)] : index, index + %row_byte_base = index.scale %token_id, %row_bytes : index, offset -> offset + %block_byte_add = index.scale %q4_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %scale_byte_base = index.add %block_byte_base, %scale_offset : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %dm_view = buffer.view %weight_noalias[%block_byte_base] : buffer -> view<2xf16> + %scale_view = buffer.view %weight_noalias[%scale_byte_base] : buffer -> view<3xi32> + %code_view = buffer.view %weight_noalias[%code_byte_base] : buffer -> view<32xi32> + %dm = vector.load %dm_view[%zero] : view<2xf16> -> vector<2xf16> + %scales = vector.load %scale_view[%zero] : view<3xi32> -> vector<3xi32> + %d_f16 = vector.extract %dm[0] : vector<2xf16> -> f16 + %dmin_f16 = vector.extract %dm[1] : vector<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %q_page0 = index.div %q4_group, %two : index + %q_page = index.mul %q_page0, %eight : index + %q_word_index0 = index.add %q_page, %load_packet : index + %q_word_index = index.assume %q_word_index0 [range(%q_word_index0, 0, 31)] : index + %is_low = index.cmp ult, %q4_group, %four : index + %scale_lane = index.rem %q4_group, %four : index + %scale_shift_index = index.mul %scale_lane, %eight : index + %scale_shift_i32 = index.cast %scale_shift_index : index to i32 + %scale_shift = vector.splat %scale_shift_i32 : vector<1xi32> + %scale0_i32 = vector.extract %scales[0] : vector<3xi32> -> i32 + %scale1_i32 = vector.extract %scales[1] : vector<3xi32> -> i32 + %scale2_i32 = vector.extract %scales[2] : vector<3xi32> -> i32 + %scale0 = vector.splat %scale0_i32 : vector<1xi32> + %scale1 = vector.splat %scale1_i32 : vector<1xi32> + %scale2 = vector.splat %scale2_i32 : vector<1xi32> + %high_shift_i32 = scalar.addi %scale_shift_i32, %two_i32 : i32 + %minimum_shift_i32 = scalar.addi %scale_shift_i32, %four_i32 : i32 + %selected_scale_source = scf.select %is_low, %scale0, %scale2 : vector<1xi32> + %selected_minimum_source = scf.select %is_low, %scale1, %scale2 : vector<1xi32> + %selected_scale_high_shift_i32 = scf.select %is_low, %scale_shift_i32, %high_shift_i32 : i32 + %selected_minimum_low_shift_i32 = scf.select %is_low, %scale_shift_i32, %minimum_shift_i32 : i32 + %selected_scale_high_shift = vector.splat %selected_scale_high_shift_i32 : vector<1xi32> + %selected_minimum_low_shift = vector.splat %selected_minimum_low_shift_i32 : vector<1xi32> + %scale_low0 = vector.shrui %selected_scale_source, %scale_shift : vector<1xi32> + %scale_low = vector.andi %scale_low0, %fifteen : vector<1xi32> + %scale_high0 = vector.shrui %scale0, %selected_scale_high_shift : vector<1xi32> + %scale_high = vector.andi %scale_high0, %fortyeight : vector<1xi32> + %scale = vector.ori %scale_low, %scale_high : vector<1xi32> + %minimum_low0 = vector.shrui %selected_minimum_source, %selected_minimum_low_shift : vector<1xi32> + %minimum_low = vector.andi %minimum_low0, %fifteen : vector<1xi32> + %minimum_high0 = vector.shrui %scale1, %selected_scale_high_shift : vector<1xi32> + %minimum_high = vector.andi %minimum_high0, %fortyeight : vector<1xi32> + %minimum = vector.ori %minimum_low, %minimum_high : vector<1xi32> + %scale_f32 = vector.uitofp %scale : vector<1xi32> to vector<1xf32> + %minimum_f32 = vector.uitofp %minimum : vector<1xi32> to vector<1xf32> + %d_vector1 = vector.splat %d : vector<1xf32> + %dmin_vector1 = vector.splat %dmin : vector<1xf32> + %d_scale_vector1 = vector.mulf %d_vector1, %scale_f32 : vector<1xf32> + %minimum_scale_vector1 = vector.mulf %dmin_vector1, %minimum_f32 : vector<1xf32> + %d_scale = vector.extract %d_scale_vector1[0] : vector<1xf32> -> f32 + %minimum_scale = vector.extract %minimum_scale_vector1[0] : vector<1xf32> -> f32 + %q_word = vector.load %code_view[%q_word_index] : view<32xi32> -> vector<1xi32> + %q_half = index.rem %q4_group, %two : index + %q_shift_index = index.mul %q_half, %four : index + %q_shift_i32 = index.cast %q_shift_index : index to i32 + %q_shift = vector.splat %q_shift_i32 : vector<1xi32> + %shifted_q = vector.shrui %q_word, %q_shift : vector<1xi32> + %masked_q = vector.andi %shifted_q, %q4_mask : vector<1xi32> + %q_i8 = vector.bitcast %masked_q : vector<1xi32> to vector<4xi8> + %q_f32 = vector.uitofp %q_i8 : vector<4xi8> to vector<4xf32> + %negative_minimum_scale = scalar.negf %minimum_scale : f32 + %q0 = vector.extract %q_f32[0] : vector<4xf32> -> f32 + %q1 = vector.extract %q_f32[1] : vector<4xf32> -> f32 + %q2 = vector.extract %q_f32[2] : vector<4xf32> -> f32 + %q3 = vector.extract %q_f32[3] : vector<4xf32> -> f32 + %value0 = scalar.fmaf %q0, %d_scale, %negative_minimum_scale : f32 + %value1 = scalar.fmaf %q1, %d_scale, %negative_minimum_scale : f32 + %value2 = scalar.fmaf %q2, %d_scale, %negative_minimum_scale : f32 + %value3 = scalar.fmaf %q3, %d_scale, %negative_minimum_scale : f32 + %values = vector.from_elements %value0, %value1, %value2, %value3 : vector<4xf32> + %output_channel0 = index.mul %row_packet, %four : index + %output_channel_end = index.sub %bounded_hidden_size, %three : index + %output_channel = index.assume %output_channel0 [range(%output_channel0, 0, 3068), lt(%output_channel0, %output_channel_end)] : index + vector.store %values, %output_view[%token, %output_channel] : vector<4xf32>, view<[%launch_token_count]x[%bounded_hidden_size]xf32> + kernel.return +} + +// Uniform 0x55 Q4_K bytes encode d=dmin=85.3125, scale=minimum=21, +// and q=5 in every group, producing 85.3125 * 21 * (5 - 1) = 7166.25. +check.case public @qwen_token_embedding_q4k_decode_case { + %token_count = check.literal value(1) : index + %vocabulary_count = check.literal value(1) : index + %hidden_size = check.literal value(2048) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(85) : tensor<1x8x144xi8> + %output = check.generate.fill value(0.0) : tensor<1x2048xf32> + %expected = check.generate.fill value(7166.25) : tensor<1x2048xf32> + kernel.launch @qwen_token_embedding_q4k_bringup_workaround[%token_count, %vocabulary_count, %hidden_size](%token_count, %vocabulary_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<1x8x144xi8>, tensor<1x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<1x2048xf32> + check.return +} + +// Reverse-order IDs exercise the first, interior, and final legal row while +// remaining suitable for access-sanitized execution. +check.case public @qwen_token_embedding_q4k_row_access_case { + %token_count = check.literal value(3) : index + %vocabulary_count = check.literal value(3) : index + %hidden_size = check.literal value(2048) : index + %token_ids = check.generate.iota offset(2) step(-1) : tensor<3xi32> + %weight = check.generate.fill value(0) : tensor<3x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<3x2048xf32> + %expected = check.generate.fill value(0.0) : tensor<3x2048xf32> + kernel.launch @qwen_token_embedding_q4k_bringup_workaround[%token_count, %vocabulary_count, %hidden_size](%token_count, %vocabulary_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<3xi32>, tensor<3x8x144xi8>, tensor<3x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<3x2048xf32> + check.return +} + +check.case public @qwen_token_embedding_q4k_benchmark_case { + %token_count = check.param.choice values([1, 32, 128, 512]) name("token_count") : index + %vocabulary_count = check.literal value(151936) : index + %hidden_size = check.literal value(2048) : index + %token_ids = check.generate.iota offset(151424) step(1) period(512) : tensor<[%token_count]xi32> + %weight = check.generate.fill value(0) : tensor<151936x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + kernel.launch @qwen_token_embedding_q4k_bringup_workaround[%token_count, %vocabulary_count, %hidden_size](%token_count, %vocabulary_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<[%token_count]xi32>, tensor<151936x8x144xi8>, tensor<[%token_count]x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x2048xf32> + check.return +} + +// Exercises the final legal checkpoint row at the exact Qwen3.5 geometry. +check.case public @qwen35_token_embedding_q4k_row_access_case { + %token_count = check.literal value(1) : index + %vocabulary_count = check.literal value(248320) : index + %hidden_size = check.literal value(3072) : index + %token_ids = check.generate.fill value(248319) : tensor<1xi32> + %weight = check.generate.fill value(0) : tensor<248320x12x144xi8> + %output = check.generate.fill value(1.0) : tensor<1x3072xf32> + %expected = check.generate.fill value(0.0) : tensor<1x3072xf32> + kernel.launch @qwen_token_embedding_q4k_bringup_workaround[%token_count, %vocabulary_count, %hidden_size](%token_count, %vocabulary_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<248320x12x144xi8>, tensor<1x3072xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<1x3072xf32> + check.return +} + +check.benchmark<@qwen_token_embedding_q4k_decode_case> @qwen_token_embedding_q4k_decode + +// Production decode specialization. Unlike the tiny differential case above, +// this carries the locked Q4_K_M vocabulary geometry into LoomC. +check.benchmark<@qwen_token_embedding_q4k_benchmark_case> @qwen_token_embedding_q4k_model_decode {token_count = 1} + +check.benchmark<@qwen_token_embedding_q4k_row_access_case> @qwen_token_embedding_q4k_row_access + +check.benchmark<@qwen35_token_embedding_q4k_row_access_case> @qwen35_token_embedding_q4k_decode + +check.benchmark<@qwen_token_embedding_q4k_benchmark_case> @qwen_token_embedding_q4k_prefill_32 {token_count = 32} + +check.benchmark<@qwen_token_embedding_q4k_benchmark_case> @qwen_token_embedding_q4k_prefill_128 {token_count = 128} + +check.benchmark<@qwen_token_embedding_q4k_benchmark_case> @qwen_token_embedding_q4k_prefill_512 {token_count = 512} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/token_embedding_q4k.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/token_embedding_q4k.loom new file mode 100644 index 000000000000..1bce11c86515 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen/token_embedding_q4k.loom @@ -0,0 +1,191 @@ +// Copyright 2026 The IREE Authors +// +// Licensed under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception + +// Gathers token rows directly from the model's unmodified GGUF Q4_K payload +// and decodes them to the owned F32 hidden-state layout. The fixed model row +// contract is: +// Q4_K: [vocabulary row][8 blocks][144 bytes] -> [2048xf32] +amdgpu.target @qwen_token_embedding_gfx11_wave64 {subgroup_size = 64} + +kernel.def target(@qwen_token_embedding_gfx11_wave64) export("qwen_token_embedding_q4k") @qwen_token_embedding_q4k(%token_count: index, %vocabulary_count: index) { + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %workgroup_count = index.mul %token_count, %c2 : index + %workgroup_size = index.constant 256 : index + kernel.launch.config workgroups(%workgroup_count, %c1, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %vocabulary_count: index, %token_ids: buffer, %weight: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %bounded_vocabulary_count = index.assume %vocabulary_count [range(%vocabulary_count, 1, 262144)] : index + %workgroup = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %packets_per_token = index.constant 512 : index + %c0_offset = index.constant 0 : offset + %block_bytes = index.constant 144 : offset + %scale_offset = index.constant 4 : offset + %code_offset = index.constant 16 : offset + %row_bytes = index.constant 1152 : offset + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15 = vector.constant 15 : vector<1xi32> + %c48 = vector.constant 48 : vector<1xi32> + %q4_mask = vector.constant 252645135 : vector<1xi32> + %packet0 = index.madd %workgroup, %c256, %workitem : index + %packet_count = index.mul %token_count, %packets_per_token : index + %packet, %launch_packet_count = index.assume %packet0, %packet_count [lt(%packet0, %packet_count)] : index, index + %token0 = index.div %packet, %packets_per_token : index + %token, %launch_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %row_packet0 = index.rem %packet, %packets_per_token : index + %row_packet = index.assume %row_packet0 [range(%row_packet0, 0, 511)] : index + %q4_block0 = index.div %row_packet, %c64 : index + %q4_block = index.assume %q4_block0 [range(%q4_block0, 0, 7)] : index + %block_packet0 = index.rem %row_packet, %c64 : index + %block_packet = index.assume %block_packet0 [range(%block_packet0, 0, 63)] : index + %q4_group0 = index.div %block_packet, %c8 : index + %q4_group = index.assume %q4_group0 [range(%q4_group0, 0, 7)] : index + %load_packet0 = index.rem %block_packet, %c8 : index + %load_packet = index.assume %load_packet0 [range(%load_packet0, 0, 7)] : index + %token_ids_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %token_ids, %weight, %output : buffer, buffer, buffer + %token_ids_view = buffer.view %token_ids_noalias[%c0_offset] : buffer -> view<[%launch_token_count]xi32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x2048xf32> + // Token IDs are validated at the request boundary. This trusted consumer + // carries that fact into row addressing instead of adding a fallback row. + %token_id_raw = view.load %token_ids_view[%token] : view<[%launch_token_count]xi32> -> i32 + %token_id_i32 = scalar.assume %token_id_raw [range(%token_id_raw, 0, 262143)] : i32 + %token_id0 = index.cast %token_id_i32 : i32 to index + %token_id, %weight_vocabulary_count = index.assume %token_id0, %bounded_vocabulary_count [lt(%token_id0, %bounded_vocabulary_count)] : index, index + %row_byte_base = index.scale %token_id, %row_bytes : index, offset -> offset + %block_byte_add = index.scale %q4_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %scale_byte_base = index.add %block_byte_base, %scale_offset : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %dm_view = buffer.view %weight_noalias[%block_byte_base] : buffer -> view<2xf16> + %scale_view = buffer.view %weight_noalias[%scale_byte_base] : buffer -> view<3xi32> + %code_view = buffer.view %weight_noalias[%code_byte_base] : buffer -> view<32xi32> + %dm = vector.load %dm_view[%c0] : view<2xf16> -> vector<2xf16> + %scales = vector.load %scale_view[%c0] : view<3xi32> -> vector<3xi32> + %d_f16 = vector.extract %dm[0] : vector<2xf16> -> f16 + %dmin_f16 = vector.extract %dm[1] : vector<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %q_page0 = index.div %q4_group, %c2 : index + %q_page = index.mul %q_page0, %c8 : index + %q_word_index0 = index.add %q_page, %load_packet : index + %q_word_index = index.assume %q_word_index0 [range(%q_word_index0, 0, 31)] : index + %is_low = index.cmp ult, %q4_group, %c4 : index + %scale_lane = index.rem %q4_group, %c4 : index + %scale_shift_index = index.mul %scale_lane, %c8 : index + %scale_shift_i32 = index.cast %scale_shift_index : index to i32 + %scale_shift = vector.splat %scale_shift_i32 : vector<1xi32> + %scale0_i32 = vector.extract %scales[0] : vector<3xi32> -> i32 + %scale1_i32 = vector.extract %scales[1] : vector<3xi32> -> i32 + %scale2_i32 = vector.extract %scales[2] : vector<3xi32> -> i32 + %scale0 = vector.splat %scale0_i32 : vector<1xi32> + %scale1 = vector.splat %scale1_i32 : vector<1xi32> + %scale2 = vector.splat %scale2_i32 : vector<1xi32> + %high_shift_i32 = scalar.addi %scale_shift_i32, %c2_i32 : i32 + %minimum_shift_i32 = scalar.addi %scale_shift_i32, %c4_i32 : i32 + %selected_scale_source = scf.select %is_low, %scale0, %scale2 : vector<1xi32> + %selected_minimum_source = scf.select %is_low, %scale1, %scale2 : vector<1xi32> + %selected_scale_high_shift_i32 = scf.select %is_low, %scale_shift_i32, %high_shift_i32 : i32 + %selected_minimum_low_shift_i32 = scf.select %is_low, %scale_shift_i32, %minimum_shift_i32 : i32 + %selected_scale_high_shift = vector.splat %selected_scale_high_shift_i32 : vector<1xi32> + %selected_minimum_low_shift = vector.splat %selected_minimum_low_shift_i32 : vector<1xi32> + %scale_low0 = vector.shrui %selected_scale_source, %scale_shift : vector<1xi32> + %scale_low = vector.andi %scale_low0, %c15 : vector<1xi32> + %scale_high0 = vector.shrui %scale0, %selected_scale_high_shift : vector<1xi32> + %scale_high = vector.andi %scale_high0, %c48 : vector<1xi32> + %scale = vector.ori %scale_low, %scale_high : vector<1xi32> + %minimum_low0 = vector.shrui %selected_minimum_source, %selected_minimum_low_shift : vector<1xi32> + %minimum_low = vector.andi %minimum_low0, %c15 : vector<1xi32> + %minimum_high0 = vector.shrui %scale1, %selected_scale_high_shift : vector<1xi32> + %minimum_high = vector.andi %minimum_high0, %c48 : vector<1xi32> + %minimum = vector.ori %minimum_low, %minimum_high : vector<1xi32> + %scale_f32 = vector.uitofp %scale : vector<1xi32> to vector<1xf32> + %minimum_f32 = vector.uitofp %minimum : vector<1xi32> to vector<1xf32> + %d_vector1 = vector.splat %d : vector<1xf32> + %dmin_vector1 = vector.splat %dmin : vector<1xf32> + %d_scale_vector1 = vector.mulf %d_vector1, %scale_f32 : vector<1xf32> + %minimum_scale_vector1 = vector.mulf %dmin_vector1, %minimum_f32 : vector<1xf32> + %d_scale = vector.extract %d_scale_vector1[0] : vector<1xf32> -> f32 + %minimum_scale = vector.extract %minimum_scale_vector1[0] : vector<1xf32> -> f32 + %q_word = vector.load %code_view[%q_word_index] : view<32xi32> -> vector<1xi32> + %q_half = index.rem %q4_group, %c2 : index + %q_shift_index = index.mul %q_half, %c4 : index + %q_shift_i32 = index.cast %q_shift_index : index to i32 + %q_shift = vector.splat %q_shift_i32 : vector<1xi32> + %shifted_q = vector.shrui %q_word, %q_shift : vector<1xi32> + %masked_q = vector.andi %shifted_q, %q4_mask : vector<1xi32> + %q_i8 = vector.bitcast %masked_q : vector<1xi32> to vector<4xi8> + %q_f32 = vector.uitofp %q_i8 : vector<4xi8> to vector<4xf32> + %negative_minimum_scale = scalar.negf %minimum_scale : f32 + %q0 = vector.extract %q_f32[0] : vector<4xf32> -> f32 + %q1 = vector.extract %q_f32[1] : vector<4xf32> -> f32 + %q2 = vector.extract %q_f32[2] : vector<4xf32> -> f32 + %q3 = vector.extract %q_f32[3] : vector<4xf32> -> f32 + %value0 = scalar.fmaf %q0, %d_scale, %negative_minimum_scale : f32 + %value1 = scalar.fmaf %q1, %d_scale, %negative_minimum_scale : f32 + %value2 = scalar.fmaf %q2, %d_scale, %negative_minimum_scale : f32 + %value3 = scalar.fmaf %q3, %d_scale, %negative_minimum_scale : f32 + %values = vector.from_elements %value0, %value1, %value2, %value3 : vector<4xf32> + %output_channel = index.mul %row_packet, %c4 : index + vector.store %values, %output_view[%token, %output_channel] : vector<4xf32>, view<[%launch_token_count]x2048xf32> + kernel.return +} + +// Uniform 0x55 Q4_K bytes encode d=dmin=85.3125, scale=minimum=21, +// and q=5 in every group, producing 85.3125 * 21 * (5 - 1) = 7166.25. +check.case public @qwen_token_embedding_q4k_decode_case { + %token_count = check.literal value(1) : index + %vocabulary_count = check.literal value(1) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(85) : tensor<1x8x144xi8> + %output = check.generate.fill value(0.0) : tensor<1x2048xf32> + %expected = check.generate.fill value(7166.25) : tensor<1x2048xf32> + kernel.launch @qwen_token_embedding_q4k[%token_count, %vocabulary_count](%token_count, %vocabulary_count, %token_ids, %weight, %output) : [index, index](index, index, tensor<1xi32>, tensor<1x8x144xi8>, tensor<1x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<1x2048xf32> + check.return +} + +// Reverse-order IDs exercise the first, interior, and final legal row while +// remaining suitable for access-sanitized execution. +check.case public @qwen_token_embedding_q4k_row_access_case { + %token_count = check.literal value(3) : index + %vocabulary_count = check.literal value(3) : index + %token_ids = check.generate.iota offset(2) step(-1) : tensor<3xi32> + %weight = check.generate.fill value(0) : tensor<3x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<3x2048xf32> + %expected = check.generate.fill value(0.0) : tensor<3x2048xf32> + kernel.launch @qwen_token_embedding_q4k[%token_count, %vocabulary_count](%token_count, %vocabulary_count, %token_ids, %weight, %output) : [index, index](index, index, tensor<3xi32>, tensor<3x8x144xi8>, tensor<3x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<3x2048xf32> + check.return +} + +check.case public @qwen_token_embedding_q4k_benchmark_case { + %token_count = check.param.choice values([32, 128, 512]) name("token_count") : index + %vocabulary_count = check.literal value(151936) : index + %token_ids = check.generate.iota offset(151424) step(1) period(512) : tensor<[%token_count]xi32> + %weight = check.generate.fill value(0) : tensor<151936x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + kernel.launch @qwen_token_embedding_q4k[%token_count, %vocabulary_count](%token_count, %vocabulary_count, %token_ids, %weight, %output) : [index, index](index, index, tensor<[%token_count]xi32>, tensor<151936x8x144xi8>, tensor<[%token_count]x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x2048xf32> + check.return +} + +check.benchmark<@qwen_token_embedding_q4k_decode_case> @qwen_token_embedding_q4k_decode + +check.benchmark<@qwen_token_embedding_q4k_row_access_case> @qwen_token_embedding_q4k_row_access + +check.benchmark<@qwen_token_embedding_q4k_benchmark_case> @qwen_token_embedding_q4k_prefill_32 {token_count = 32} + +check.benchmark<@qwen_token_embedding_q4k_benchmark_case> @qwen_token_embedding_q4k_prefill_128 {token_count = 128} + +check.benchmark<@qwen_token_embedding_q4k_benchmark_case> @qwen_token_embedding_q4k_prefill_512 {token_count = 512} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/ggml/linear_q6k_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/ggml/linear_q6k_f32.loom new file mode 100644 index 000000000000..c5908bb01fcc --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/ggml/linear_q6k_f32.loom @@ -0,0 +1,425 @@ +// Contracts GGML Q6_K weight rows directly with F32 activation rows using +// the decode schedule selected by llama.cpp's Vulkan backend. +// +// One 64-workitem workgroup computes two adjacent output rows. Four cohorts +// of 16 lanes each process four Q6_K blocks in parallel, while each lane +// decodes four values from the 0, 32, 64, and 96 element quarters of its +// block. Per-group Q6_K scales are exchanged through two 256-byte LDS frames, +// one for each output row. The frames preserve the producer/consumer shape of +// the Vulkan oracle without repacking the persistent GGUF weights. +// +// This provider is intentionally independent of the Q8_1 integer-dot path. +// On gfx11, direct F32 loads and software Q6_K decode can beat activation +// quantization for decode-sized batches. Keeping both algorithms in the same +// corpus lets JIT selection depend on the specialized shape and target. +template.decl @ggml.linear_q6k_f32.body(%publish_output: i1, %token_count: index, %token0: index, %pair: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %output: buffer) + +amdgpu.target @ggml_q6k_gfx11_wave64 {subgroup_size = 64} + +amdgpu.target @ggml_q6k_gfx11_wave32 {subgroup_size = 32} + +config.decl @ggml.linear_q6k_f32.token_capacity : %value: index where [range(%value, 1, 2048)] + +config.decl @ggml.linear_q6k_f32.output_capacity : %value: index where [range(%value, 1, 262144)] + +// Q8_1 reference declarations used only by the differential case. +kernel.decl @ggml_quantize_q8_1_x4_f32(%token_count$11: index, %input_size$12: index) launch(%token_count$13: index, %input_size$14: index, %input: buffer, %output: buffer) + +kernel.decl @ggml_linear_q6k_q8_1_x4(%token_count$17: index, %input_size$18: index, %output_size$19: index) launch(%token_count$20: index, %input_size$21: index, %output_size$22: index, %q8_input: buffer, %weight: buffer, %output: buffer) + +// Computes an ordered four-element F32 dot product. The scalar recurrence +// mirrors the GLSL oracle and gives target scheduling four independent dot +// chains per decoded Q6_K block. +func.def inline @ggml_q6k_dot4_f32(%lhs: vector<4xf32>, %rhs: vector<4xf32>) -> (f32) { + %c0 = scalar.constant 0.0 : f32 + %lhs0 = vector.extract %lhs[0] : vector<4xf32> -> f32 + %lhs1 = vector.extract %lhs[1] : vector<4xf32> -> f32 + %lhs2 = vector.extract %lhs[2] : vector<4xf32> -> f32 + %lhs3 = vector.extract %lhs[3] : vector<4xf32> -> f32 + %rhs0 = vector.extract %rhs[0] : vector<4xf32> -> f32 + %rhs1 = vector.extract %rhs[1] : vector<4xf32> -> f32 + %rhs2 = vector.extract %rhs[2] : vector<4xf32> -> f32 + %rhs3 = vector.extract %rhs[3] : vector<4xf32> -> f32 + %sum0 = scalar.fmaf %lhs0, %rhs0, %c0 : f32 + %sum1 = scalar.fmaf %lhs1, %rhs1, %sum0 : f32 + %sum2 = scalar.fmaf %lhs2, %rhs2, %sum1 : f32 + %sum3 = scalar.fmaf %lhs3, %rhs3, %sum2 : f32 + func.return %sum3 : f32 +} + +// Loads one row contraction's activation packet. The two output rows read the +// same packet, but each load remains adjacent to its consumer. Hoisting the +// packet across the intervening LDS barrier halves issued activation traffic +// while extending sixteen F32 live values across synchronization; on gfx1151 +// that longer live range is slower than the repeated cache-hot read. +func.def inline @ggml_q6k_load_f32_block(%token_count: index, %input_size: index, %token: index, %block: index, %lane: index, %input: buffer) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) { + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c96 = index.constant 96 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %bounded_lane = index.assume %lane [range(%lane, 0, 63)] : index + %block_count = index.div %input_size, %c256 : index + %valid_block = index.cmp ult, %block, %block_count : index + %safe_block0 = scf.select %valid_block, %block, %c0 : index + %safe_block, %launch_block_count = index.assume %safe_block0, %block_count [lt(%safe_block0, %block_count)] : index, index + %itid = index.rem %bounded_lane, %c16 : index + %vector_half = index.div %itid, %c8 : index + %vector_index = index.rem %itid, %c8 : index + %input_view = buffer.view %input[%c0_offset] : buffer -> view<[%token_count]x[%launch_block_count]x256xf32> + %vector_half_base = index.mul %vector_half, %c128 : index + %vector_offset = index.mul %vector_index, %c4 : index + %input_index0 = index.add %vector_half_base, %vector_offset : index + %input_index1 = index.add %input_index0, %c32 : index + %input_index2 = index.add %input_index0, %c64 : index + %input_index3 = index.add %input_index0, %c96 : index + %input0 = vector.load %input_view[%token, %safe_block, %input_index0] : view<[%token_count]x[%launch_block_count]x256xf32> -> vector<4xf32> + %input1 = vector.load %input_view[%token, %safe_block, %input_index1] : view<[%token_count]x[%launch_block_count]x256xf32> -> vector<4xf32> + %input2 = vector.load %input_view[%token, %safe_block, %input_index2] : view<[%token_count]x[%launch_block_count]x256xf32> -> vector<4xf32> + %input3 = vector.load %input_view[%token, %safe_block, %input_index3] : view<[%token_count]x[%launch_block_count]x256xf32> -> vector<4xf32> + func.return %input0, %input1, %input2, %input3 : vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32> +} + +// Loads the signed group scale owned by one lane of a Q6_K block. +func.def inline @ggml_q6k_load_f32_scale(%input_size: index, %row: index, %block: index, %lane: index, %weight: buffer) -> (f32) { + %c0 = index.constant 0 : index + %c16 = index.constant 16 : index + %c192 = index.constant 192 : offset + %c210_bytes = index.constant 210 : offset + %c256 = index.constant 256 : index + %bounded_lane = index.assume %lane [range(%lane, 0, 63)] : index + %block_count = index.div %input_size, %c256 : index + %valid_block = index.cmp ult, %block, %block_count : index + %safe_block = scf.if %valid_block -> (index) { + scf.yield %block : index + } else { + scf.yield %c0 : index + } + %itid = index.rem %bounded_lane, %c16 : index + %weight_row_bytes = index.scale %block_count, %c210_bytes : index, offset -> offset + %weight_row_byte_base = index.scale %row, %weight_row_bytes : index, offset -> offset + %block_byte_add = index.scale %safe_block, %c210_bytes : index, offset -> offset + %weight_block_byte_base = index.add %weight_row_byte_base, %block_byte_add : offset + %scale_byte_base = index.add %weight_block_byte_base, %c192 : offset + %scale_view = buffer.view %weight[%scale_byte_base] : buffer -> view<16xi8> + %scale_i8 = view.load %scale_view[%itid] : view<16xi8> -> i8 + %scale = scalar.sitofp %scale_i8 : i8 to f32 + func.return %scale : f32 +} + +// Stages one row's group scales and synchronizes before contraction. +func.def inline @ggml_q6k_stage_f32_scales(%input_size: index, %row: index, %block: index, %frame_count: index, %frame: index, %lane: index, %weight: buffer, %scale_stage: buffer) { + %c0_offset = index.constant 0 : offset + %bounded_lane = index.assume %lane [range(%lane, 0, 63)] : index + %scale_stage_view = buffer.view %scale_stage[%c0_offset] : buffer -> view<[%frame_count]x64xf32> + %scale = func.call @ggml_q6k_load_f32_scale(%input_size, %row, %block, %bounded_lane, %weight) : (index, index, index, index, buffer) -> (f32) + view.store %scale, %scale_stage_view[%frame, %bounded_lane] : f32, view<[%frame_count]x64xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + func.return +} + +// Computes one lane contribution for one output row and one four-block +// cohort after its group scales have been staged. +func.def inline @ggml_q6k_f32_block_row(%input_size: index, %row: index, %block: index, %frame_count: index, %frame: index, %lane: index, %weight: buffer, %scale_stage: buffer, %input0: vector<4xf32>, %input1: vector<4xf32>, %input2: vector<4xf32>, %input3: vector<4xf32>) -> (f32) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c6 = index.constant 6 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c128 = index.constant 128 : offset + %c208 = index.constant 208 : offset + %c210_bytes = index.constant 210 : offset + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %c4_i32v = vector.constant 4 : vector<1xi32> + %c2_i32v = vector.constant 2 : vector<1xi32> + %nibble_mask = vector.constant 252645135 : vector<1xi32> + %high0_mask = vector.constant 50529027 : vector<1xi32> + %high2_mask = vector.constant 202116108 : vector<1xi32> + %high4_mask = vector.constant 808464432 : vector<1xi32> + %high6_mask = vector.constant -1061109568 : vector<1xi32> + %c32_f32v = vector.constant 32.0 : vector<4xf32> + %c0_f32 = scalar.constant 0.0 : f32 + %bounded_lane = index.assume %lane [range(%lane, 0, 63)] : index + %block_count = index.div %input_size, %c256 : index + %valid_block = index.cmp ult, %block, %block_count : index + %safe_block = scf.if %valid_block -> (index) { + scf.yield %block : index + } else { + scf.yield %c0 : index + } + %itid = index.rem %bounded_lane, %c16 : index + %cohort = index.div %bounded_lane, %c16 : index + %vector_half = index.div %itid, %c8 : index + %vector_index = index.rem %itid, %c8 : index + %vector_quarter = index.div %vector_index, %c4 : index + %weight_row_bytes = index.scale %block_count, %c210_bytes : index, offset -> offset + %weight_row_byte_base = index.scale %row, %weight_row_bytes : index, offset -> offset + %block_byte_add = index.scale %safe_block, %c210_bytes : index, offset -> offset + %weight_block_byte_base = index.add %weight_row_byte_base, %block_byte_add : offset + %qh_byte_base = index.add %weight_block_byte_base, %c128 : offset + %d_byte_base = index.add %weight_block_byte_base, %c208 : offset + %ql_view = buffer.view %weight[%weight_block_byte_base] : buffer -> view<32xi32> + %qh_view = buffer.view %weight[%qh_byte_base] : buffer -> view<16xi32> + %d_view = buffer.view %weight[%d_byte_base] : buffer -> view<1xf16> + %scale_stage_view = buffer.view %scale_stage[%c0_offset] : buffer -> view<[%frame_count]x64xf32> + %scale_cohort_base = index.mul %cohort, %c16 : index + %scale_half_base = index.mul %vector_half, %c8 : index + %scale_index0 = index.add %scale_half_base, %vector_quarter : index + %scale_index1 = index.add %scale_index0, %c2 : index + %scale_index2 = index.add %scale_index0, %c4 : index + %scale_index3 = index.add %scale_index0, %c6 : index + %stage_scale_index0 = index.add %scale_cohort_base, %scale_index0 : index + %stage_scale_index1 = index.add %scale_cohort_base, %scale_index1 : index + %stage_scale_index2 = index.add %scale_cohort_base, %scale_index2 : index + %stage_scale_index3 = index.add %scale_cohort_base, %scale_index3 : index + %scale0 = view.load %scale_stage_view[%frame, %stage_scale_index0] : view<[%frame_count]x64xf32> -> f32 + %scale1 = view.load %scale_stage_view[%frame, %stage_scale_index1] : view<[%frame_count]x64xf32> -> f32 + %scale2 = view.load %scale_stage_view[%frame, %stage_scale_index2] : view<[%frame_count]x64xf32> -> f32 + %scale3 = view.load %scale_stage_view[%frame, %stage_scale_index3] : view<[%frame_count]x64xf32> -> f32 + %ql_half_word_base = index.mul %vector_half, %c16 : index + %ql_word_index00 = index.add %ql_half_word_base, %vector_index : index + %ql_word_index10 = index.add %ql_word_index00, %c8 : index + %qh_half_word_base = index.mul %vector_half, %c8 : index + %qh_word_index0 = index.add %qh_half_word_base, %vector_index : index + %ql_word_index0, %ql_word_index1, %qh_word_index = index.assume %ql_word_index00, %ql_word_index10, %qh_word_index0 [range(%ql_word_index00, 0, 23), range(%ql_word_index10, 8, 31), range(%qh_word_index0, 0, 15)] : index, index, index + %ql_word0 = vector.load %ql_view[%ql_word_index0] : view<32xi32> -> vector<1xi32> + %ql_word1 = vector.load %ql_view[%ql_word_index1] : view<32xi32> -> vector<1xi32> + %qh_word = vector.load %qh_view[%qh_word_index] : view<16xi32> -> vector<1xi32> + %ql0 = vector.andi %ql_word0, %nibble_mask : vector<1xi32> + %ql1 = vector.andi %ql_word1, %nibble_mask : vector<1xi32> + %ql_word0_high = vector.shrui %ql_word0, %c4_i32v : vector<1xi32> + %ql_word1_high = vector.shrui %ql_word1, %c4_i32v : vector<1xi32> + %ql2 = vector.andi %ql_word0_high, %nibble_mask : vector<1xi32> + %ql3 = vector.andi %ql_word1_high, %nibble_mask : vector<1xi32> + %qh0_low = vector.andi %qh_word, %high0_mask : vector<1xi32> + %qh1_low = vector.andi %qh_word, %high2_mask : vector<1xi32> + %qh2 = vector.andi %qh_word, %high4_mask : vector<1xi32> + %qh3_high = vector.andi %qh_word, %high6_mask : vector<1xi32> + %qh0 = vector.shli %qh0_low, %c4_i32v : vector<1xi32> + %qh1 = vector.shli %qh1_low, %c2_i32v : vector<1xi32> + %qh3 = vector.shrui %qh3_high, %c2_i32v : vector<1xi32> + %code0 = vector.ori %ql0, %qh0 : vector<1xi32> + %code1 = vector.ori %ql1, %qh1 : vector<1xi32> + %code2 = vector.ori %ql2, %qh2 : vector<1xi32> + %code3 = vector.ori %ql3, %qh3 : vector<1xi32> + %code0_i8 = vector.bitcast %code0 : vector<1xi32> to vector<4xi8> + %code1_i8 = vector.bitcast %code1 : vector<1xi32> to vector<4xi8> + %code2_i8 = vector.bitcast %code2 : vector<1xi32> to vector<4xi8> + %code3_i8 = vector.bitcast %code3 : vector<1xi32> to vector<4xi8> + %code0_f32 = vector.uitofp %code0_i8 : vector<4xi8> to vector<4xf32> + %code1_f32 = vector.uitofp %code1_i8 : vector<4xi8> to vector<4xf32> + %code2_f32 = vector.uitofp %code2_i8 : vector<4xi8> to vector<4xf32> + %code3_f32 = vector.uitofp %code3_i8 : vector<4xi8> to vector<4xf32> + %q0 = vector.subf %code0_f32, %c32_f32v : vector<4xf32> + %q1 = vector.subf %code1_f32, %c32_f32v : vector<4xf32> + %q2 = vector.subf %code2_f32, %c32_f32v : vector<4xf32> + %q3 = vector.subf %code3_f32, %c32_f32v : vector<4xf32> + %dot0 = func.call @ggml_q6k_dot4_f32(%input0, %q0) : (vector<4xf32>, vector<4xf32>) -> (f32) + %dot1 = func.call @ggml_q6k_dot4_f32(%input1, %q1) : (vector<4xf32>, vector<4xf32>) -> (f32) + %dot2 = func.call @ggml_q6k_dot4_f32(%input2, %q2) : (vector<4xf32>, vector<4xf32>) -> (f32) + %dot3 = func.call @ggml_q6k_dot4_f32(%input3, %q3) : (vector<4xf32>, vector<4xf32>) -> (f32) + %scaled3 = scalar.mulf %dot3, %scale3 : f32 + %scaled2 = scalar.fmaf %dot2, %scale2, %scaled3 : f32 + %scaled1 = scalar.fmaf %dot1, %scale1, %scaled2 : f32 + %scaled0 = scalar.fmaf %dot0, %scale0, %scaled1 : f32 + %d_f16 = view.load %d_view[0] : view<1xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %contribution0 = scalar.mulf %scaled0, %d : f32 + %contribution = scf.if %valid_block -> (f32) { + scf.yield %contribution0 : f32 + } else { + scf.yield %c0_f32 : f32 + } + func.return %contribution : f32 +} + +// Shared direct-F32 matrix-vector schedule. Target entry points below vary +// subgroup width without duplicating the Q6_K decoding or memory schedule. +template.def<@ggml.linear_q6k_f32.body> device @ggml_linear_q6k_f32_body(%publish_output: i1, %token_count: index, %token0: index, %pair: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 262144)] : index + %lane0 = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c256 = index.constant 256 : index + %c0_f32 = scalar.constant 0.0 : f32 + %c0_offset = index.constant 0 : offset + %scale_stage_bytes = index.constant 512 : offset + %token, %launch_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %lane = index.assume %lane0 [range(%lane0, 0, 63)] : index + %row00 = index.mul %pair, %c2 : index + %row0, %launch_output_size = index.assume %row00, %bounded_output_size [lt(%row00, %bounded_output_size)] : index, index + %row1 = index.add %row0, %c1 : index + %row1_valid = index.cmp ult, %row1, %launch_output_size : index + %block_count = index.div %bounded_input_size, %c256 : index + %cohort = index.div %lane, %c16 : index + scf.if %publish_output { + %input_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input, %weight, %output : buffer, buffer, buffer + %scale_stage = buffer.alloca align(16) %scale_stage_bytes : buffer + %acc0, %acc1 = scf.for %block_base = [%c0 to %block_count step %c4](%row_acc0 = %c0_f32 : f32, %row_acc1 = %c0_f32 : f32) -> (f32, f32) { + %block = index.add %block_base, %cohort : index + func.call @ggml_q6k_stage_f32_scales(%bounded_input_size, %row0, %block, %c2, %c0, %lane, %weight_noalias, %scale_stage) : (index, index, index, index, index, index, buffer, buffer) + %row0_input0, %row0_input1, %row0_input2, %row0_input3 = func.call @ggml_q6k_load_f32_block(%launch_token_count, %bounded_input_size, %token, %block, %lane, %input_noalias) : (index, index, index, index, index, buffer) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) + %contribution0 = func.call @ggml_q6k_f32_block_row(%bounded_input_size, %row0, %block, %c2, %c0, %lane, %weight_noalias, %scale_stage, %row0_input0, %row0_input1, %row0_input2, %row0_input3) : (index, index, index, index, index, index, buffer, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) -> (f32) + %contribution1 = scf.if %row1_valid -> (f32) { + func.call @ggml_q6k_stage_f32_scales(%bounded_input_size, %row1, %block, %c2, %c1, %lane, %weight_noalias, %scale_stage) : (index, index, index, index, index, index, buffer, buffer) + %row1_input0, %row1_input1, %row1_input2, %row1_input3 = func.call @ggml_q6k_load_f32_block(%launch_token_count, %bounded_input_size, %token, %block, %lane, %input_noalias) : (index, index, index, index, index, buffer) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) + %row1_contribution = func.call @ggml_q6k_f32_block_row(%bounded_input_size, %row1, %block, %c2, %c1, %lane, %weight_noalias, %scale_stage, %row1_input0, %row1_input1, %row1_input2, %row1_input3) : (index, index, index, index, index, index, buffer, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) -> (f32) + scf.yield %row1_contribution : f32 + } else { + scf.yield %c0_f32 : f32 + } + %next0 = scalar.addf %row_acc0, %contribution0 : f32 + %next1 = scalar.addf %row_acc1, %contribution1 : f32 + scf.yield %next0, %next1 : f32, f32 + } + %sum0 = kernel.workgroup.reduce %acc0 : f32 + %sum1 = kernel.workgroup.reduce %acc1 : f32 + %is_lane_zero = index.cmp eq, %lane, %c0 : index + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%launch_output_size]xf32> + scf.if %is_lane_zero { + view.store %sum0, %output_view[%token, %row0] : f32, view<[%launch_token_count]x[%launch_output_size]xf32> + } + scf.if %row1_valid { + scf.if %is_lane_zero { + view.store %sum1, %output_view[%token, %row1] : f32, view<[%launch_token_count]x[%launch_output_size]xf32> + } + } + } + template.return +} + +// Wave64 provider matching the Vulkan oracle's subgroup schedule. +kernel.def target(@ggml_q6k_gfx11_wave64) @ggml_linear_q6k_f32_wave64(%token_count: index, %input_size: index, %output_size: index) { + %token_capacity = config.get @ggml.linear_q6k_f32.token_capacity : index + %output_capacity = config.get @ggml.linear_q6k_f32.output_capacity : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c64 = index.constant 64 : index + %padded_output_size = index.add %output_capacity, %c1 : index + %output_pairs = index.div %padded_output_size, %c2 : index + kernel.launch.config workgroups(%output_pairs, %token_capacity, %c1) workgroup_size(%c64, %c1, %c1) : index +} launch(%token_count: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %output: buffer) { + %token_capacity = config.get @ggml.linear_q6k_f32.token_capacity : index + %output_capacity = config.get @ggml.linear_q6k_f32.output_capacity : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 262144), le(%output_size, %output_capacity)] : index + %c2 = index.constant 2 : index + %pair = kernel.workgroup.id : index + %token = kernel.workgroup.id : index + %row0 = index.mul %pair, %c2 : index + %valid_token = index.cmp ult, %token, %bounded_token_count : index + %valid_row = index.cmp ult, %row0, %bounded_output_size : index + %publish_output = scalar.andi %valid_token, %valid_row : i1 + %c0 = index.constant 0 : index + %safe_token = scf.select %valid_token, %token, %c0 : index + %safe_pair = scf.select %valid_row, %pair, %c0 : index + template.apply<@ggml.linear_q6k_f32.body>(%publish_output, %bounded_token_count, %safe_token, %safe_pair, %input_size, %bounded_output_size, %input, %weight, %output) : (i1, index, index, index, index, index, buffer, buffer, buffer) + kernel.return +} + +// Wave32 provider testing whether two subgroups hide direct-F32 decode latency +// better on targets where wave64 is not the default execution mode. +kernel.def target(@ggml_q6k_gfx11_wave32) @ggml_linear_q6k_f32_wave32(%token_count: index, %input_size: index, %output_size: index) { + %token_capacity = config.get @ggml.linear_q6k_f32.token_capacity : index + %output_capacity = config.get @ggml.linear_q6k_f32.output_capacity : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c64 = index.constant 64 : index + %padded_output_size = index.add %output_capacity, %c1 : index + %output_pairs = index.div %padded_output_size, %c2 : index + kernel.launch.config workgroups(%output_pairs, %token_capacity, %c1) workgroup_size(%c64, %c1, %c1) : index +} launch(%token_count: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %output: buffer) { + %token_capacity = config.get @ggml.linear_q6k_f32.token_capacity : index + %output_capacity = config.get @ggml.linear_q6k_f32.output_capacity : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 262144), le(%output_size, %output_capacity)] : index + %c2 = index.constant 2 : index + %pair = kernel.workgroup.id : index + %token = kernel.workgroup.id : index + %row0 = index.mul %pair, %c2 : index + %valid_token = index.cmp ult, %token, %bounded_token_count : index + %valid_row = index.cmp ult, %row0, %bounded_output_size : index + %publish_output = scalar.andi %valid_token, %valid_row : i1 + %c0 = index.constant 0 : index + %safe_token = scf.select %valid_token, %token, %c0 : index + %safe_pair = scf.select %valid_row, %pair, %c0 : index + template.apply<@ggml.linear_q6k_f32.body>(%publish_output, %bounded_token_count, %safe_token, %safe_pair, %input_size, %bounded_output_size, %input, %weight, %output) : (i1, index, index, index, index, index, buffer, buffer, buffer) + kernel.return +} + +// Exact-representable activations make the direct-F32 and Q8_1 contraction +// paths comparable while nonzero packed bytes exercise every Q6_K field. +check.case public @ggml_linear_q6k_f32_differential_case { + %token_count = check.literal value(2) : index + %input_size = check.literal value(2048) : index + %output_size = check.literal value(9) : index + %input = check.generate.fill value(0.00390625) : tensor<2x2048xf32> + %q8_input = check.generate.fill value(0) : tensor<2x2304xi8> + %weight = check.generate.fill value(-86) : tensor<9x8x210xi8> + %expected = check.generate.fill value(0.0) : tensor<2x9xf32> + %actual_wave64 = check.generate.fill value(1.0) : tensor<2x9xf32> + %actual_wave32 = check.generate.fill value(1.0) : tensor<2x9xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<2x2048xf32>, tensor<2x2304xi8>) + kernel.launch @ggml_linear_q6k_q8_1_x4[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %q8_input, %weight, %expected) : [index, index, index](index, index, index, tensor<2x2304xi8>, tensor<9x8x210xi8>, tensor<2x9xf32>) + kernel.launch @ggml_linear_q6k_f32_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %actual_wave64) : [index, index, index](index, index, index, tensor<2x2048xf32>, tensor<9x8x210xi8>, tensor<2x9xf32>) + kernel.launch @ggml_linear_q6k_f32_wave32[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %actual_wave32) : [index, index, index](index, index, index, tensor<2x2048xf32>, tensor<9x8x210xi8>, tensor<2x9xf32>) + check.expect.close actual(%actual_wave64) expected(%expected) atol(0.25) rtol(0.01) nan(same) : tensor<2x9xf32> + check.expect.close actual(%actual_wave32) expected(%expected) atol(0.25) rtol(0.01) nan(same) : tensor<2x9xf32> + check.return +} + +check.case public @ggml_linear_q6k_f32_wave64_dense_v_benchmark_case { + %token_count = check.param.choice values([1, 8, 32, 128, 512]) name("token_count") : index + %input_size = check.literal value(2048) : index + %output_size = check.literal value(512) : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %weight = check.generate.fill value(0) : tensor<512x8x210xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x512xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x512xf32> + kernel.launch @ggml_linear_q6k_f32_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<[%token_count]x2048xf32>, tensor<512x8x210xi8>, tensor<[%token_count]x512xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x512xf32> + check.return +} + +check.case public @ggml_linear_q6k_f32_wave32_dense_v_benchmark_case { + %token_count = check.param.choice values([1, 8, 32, 128, 512]) name("token_count") : index + %input_size = check.literal value(2048) : index + %output_size = check.literal value(512) : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %weight = check.generate.fill value(0) : tensor<512x8x210xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x512xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x512xf32> + kernel.launch @ggml_linear_q6k_f32_wave32[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<[%token_count]x2048xf32>, tensor<512x8x210xi8>, tensor<[%token_count]x512xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x512xf32> + check.return +} + +check.benchmark<@ggml_linear_q6k_f32_differential_case> @ggml_linear_q6k_f32_differential + +check.benchmark<@ggml_linear_q6k_f32_wave64_dense_v_benchmark_case> @ggml_linear_q6k_f32_wave64_dense_v_decode {token_count = 1} + +check.benchmark<@ggml_linear_q6k_f32_wave64_dense_v_benchmark_case> @ggml_linear_q6k_f32_wave64_dense_v_prefill_32 {token_count = 32} + +check.benchmark<@ggml_linear_q6k_f32_wave64_dense_v_benchmark_case> @ggml_linear_q6k_f32_wave64_dense_v_prefill_128 {token_count = 128} + +check.benchmark<@ggml_linear_q6k_f32_wave64_dense_v_benchmark_case> @ggml_linear_q6k_f32_wave64_dense_v_prefill_512 {token_count = 512} + +check.benchmark<@ggml_linear_q6k_f32_wave32_dense_v_benchmark_case> @ggml_linear_q6k_f32_wave32_dense_v_decode {token_count = 1} + +check.benchmark<@ggml_linear_q6k_f32_wave32_dense_v_benchmark_case> @ggml_linear_q6k_f32_wave32_dense_v_prefill_32 {token_count = 32} + +check.benchmark<@ggml_linear_q6k_f32_wave32_dense_v_benchmark_case> @ggml_linear_q6k_f32_wave32_dense_v_prefill_128 {token_count = 128} + +check.benchmark<@ggml_linear_q6k_f32_wave32_dense_v_benchmark_case> @ggml_linear_q6k_f32_wave32_dense_v_prefill_512 {token_count = 512} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/ggml/linear_q6k_q8_1_x4.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/ggml/linear_q6k_q8_1_x4.loom new file mode 100644 index 000000000000..2cc4cd97a3d8 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/ggml/linear_q6k_q8_1_x4.loom @@ -0,0 +1,510 @@ +// Contracts GGML Q6_K weight rows directly with the Q8_1 x4 activation +// layout used by llama.cpp's Vulkan matrix path. One 210-byte Q6_K block +// represents 256 signed six-bit values: +// +// i8 ql[128]; // low four bits +// i8 qh[64]; // high two bits +// i8 scales[16]; // signed scale for each 16-value group +// f16 d; // block-wide scale +// +// A wave owns one output value. Each lane contracts eight values from every +// Q6_K block as two native signed dot4 packets, and the wave reduces those +// partials. The kernel is the raw-layout correctness and decode baseline shared +// by dense and routed projections; prefill schedules can reuse the inline +// packed-row primitive while staging weights across multiple activation rows. +// The shared packer is linked beside this module. +template.decl @ggml.linear_q6k_q8_1_x4.body(%publish_output: i1, %token_count: index, %token0: index, %input_size: index, %output_size: index, %q8_input: buffer, %weight: buffer) -> (f32, index, index, i1, index, index) + +amdgpu.target @ggml_q6k_q8_gfx1151_wave64 {subgroup_size = 64} + +config.decl @ggml.linear_q6k_q8_1_x4.token_capacity : %value: index where [range(%value, 1, 2048)] + +config.decl @ggml.linear_q6k_q8_1_x4.output_capacity : %value: index where [range(%value, 1, 262144)] + +kernel.decl @ggml_quantize_q8_1_x4_f32(%token_count$15: index, %input_size$16: index) launch(%token_count$17: index, %input_size$18: index, %input: buffer, %output: buffer) + +func.decl @ggml_q8_1_x4_word(%q8_input: buffer, %row_byte_base: offset, %q8_block: index, %word_in_block: index) -> (vector<4xi8>, f32) + +// Sign-extends four packed six-bit values without scalar lane extraction. +// Each byte enters with bits [5:0] populated and leaves as signed i8. +func.def inline @ggml_q6k_sign_extend_dot4(%code: vector<1xi32>) -> (vector<4xi8>) { + %c1_i32v = vector.constant 1 : vector<1xi32> + %c2_i32v = vector.constant 2 : vector<1xi32> + %low5_mask = vector.constant 522133279 : vector<1xi32> + %bit5_mask = vector.constant 538976288 : vector<1xi32> + %sign_mask = vector.constant -522133280 : vector<1xi32> + %low5 = vector.andi %code, %low5_mask : vector<1xi32> + %bit5 = vector.andi %code, %bit5_mask : vector<1xi32> + %bit6 = vector.shli %bit5, %c1_i32v : vector<1xi32> + %bit7 = vector.shli %bit5, %c2_i32v : vector<1xi32> + %high01 = vector.ori %bit5, %bit6 : vector<1xi32> + %high = vector.ori %high01, %bit7 : vector<1xi32> + %sign = vector.xori %high, %sign_mask : vector<1xi32> + %signed_i32 = vector.ori %low5, %sign : vector<1xi32> + %signed = vector.bitcast %signed_i32 : vector<1xi32> to vector<4xi8> + func.return %signed : vector<4xi8> +} + +// Decodes four adjacent values from one 32-value Q6_K group for FP16 matrix +// staging. The group and packet coordinates match the contiguous K dimension +// consumed by WMMA tiles, unlike the lane/part mapping used by the Q8_1 dot +// contraction below. +func.def inline @ggml_q6k_f16_vector4(%weight: buffer, %weight_row_byte_base: offset, %q6_block: index, %q6_group: index, %packet: index) -> (vector<4xf16>) { + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %block_bytes = index.constant 210 : offset + %qh_byte_add = index.constant 128 : offset + %scale_byte_add = index.constant 192 : offset + %d_byte_add = index.constant 208 : offset + %c4_i32v = vector.constant 4 : vector<1xi32> + %nibble_mask = vector.constant 252645135 : vector<1xi32> + %high_mask = vector.constant 50529027 : vector<1xi32> + %c32_f32v = vector.constant 32.0 : vector<4xf32> + %bounded_group = index.assume %q6_group [range(%q6_group, 0, 7)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q6_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %weight_row_byte_base, %block_byte_add : offset + %qh_byte_base = index.add %block_byte_base, %qh_byte_add : offset + %scale_byte_base = index.add %block_byte_base, %scale_byte_add : offset + %d_byte_base = index.add %block_byte_base, %d_byte_add : offset + %ql_view = buffer.view %weight[%block_byte_base] : buffer -> view<32xi32> + %qh_view = buffer.view %weight[%qh_byte_base] : buffer -> view<16xi32> + %scale_view = buffer.view %weight[%scale_byte_base] : buffer -> view<16xi8> + %d_view = buffer.view %weight[%d_byte_base] : buffer -> view<1xf16> + %group_in_half = index.rem %bounded_group, %c4 : index + %half = index.div %bounded_group, %c4 : index + %ql_side = index.rem %group_in_half, %c2 : index + %ql_half_word_base = index.mul %half, %c16 : index + %ql_side_word_add = index.mul %ql_side, %c8 : index + %ql_word_base = index.add %ql_half_word_base, %ql_side_word_add : index + %ql_word_index = index.add %ql_word_base, %bounded_packet : index + %qh_half_word_base = index.mul %half, %c8 : index + %qh_word_index = index.add %qh_half_word_base, %bounded_packet : index + %nibble = index.div %group_in_half, %c2 : index + %nibble_shift_index = index.mul %nibble, %c4 : index + %nibble_shift_i32 = index.cast %nibble_shift_index : index to i32 + %nibble_shift = vector.splat %nibble_shift_i32 : vector<1xi32> + %qh_shift_index = index.mul %group_in_half, %c2 : index + %qh_shift_i32 = index.cast %qh_shift_index : index to i32 + %qh_shift = vector.splat %qh_shift_i32 : vector<1xi32> + %scale_packet_half = index.div %bounded_packet, %c4 : index + %scale_group_base = index.mul %bounded_group, %c2 : index + %scale_index = index.add %scale_group_base, %scale_packet_half : index + %ql_word = vector.load %ql_view[%ql_word_index] : view<32xi32> -> vector<1xi32> + %qh_word = vector.load %qh_view[%qh_word_index] : view<16xi32> -> vector<1xi32> + %ql_shifted = vector.shrui %ql_word, %nibble_shift : vector<1xi32> + %ql = vector.andi %ql_shifted, %nibble_mask : vector<1xi32> + %qh_shifted = vector.shrui %qh_word, %qh_shift : vector<1xi32> + %qh_low = vector.andi %qh_shifted, %high_mask : vector<1xi32> + %qh = vector.shli %qh_low, %c4_i32v : vector<1xi32> + %code = vector.ori %ql, %qh : vector<1xi32> + %code_i8 = vector.bitcast %code : vector<1xi32> to vector<4xi8> + %code_f32 = vector.uitofp %code_i8 : vector<4xi8> to vector<4xf32> + %centered = vector.subf %code_f32, %c32_f32v : vector<4xf32> + %scale_i8 = view.load %scale_view[%scale_index] : view<16xi8> -> i8 + %d_f16 = view.load %d_view[0] : view<1xf16> -> f16 + %scale = scalar.sitofp %scale_i8 : i8 to f32 + %d = scalar.extf %d_f16 : f16 to f32 + %combined_scale = scalar.mulf %scale, %d : f32 + %combined_scale_vector = vector.splat %combined_scale : vector<4xf32> + %values_f32 = vector.mulf %centered, %combined_scale_vector : vector<4xf32> + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +// Contracts both four-value packets assigned to one lane in a Q6_K block. +// QL, QH, and the block scale are shared by the two packed parts, and their +// contributions are summed before entering the row-wide recurrence. +func.def inline @ggml_q6k_q8_1_x4_block_lane(%weight: buffer, %weight_row_byte_base: offset, %q8_input: buffer, %q8_row_byte_base: offset, %q6_block: index, %lane: index) -> (f32) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %block_bytes = index.constant 210 : offset + %qh_byte_add = index.constant 128 : offset + %scale_byte_add = index.constant 192 : offset + %d_byte_add = index.constant 208 : offset + %c4_i32v = vector.constant 4 : vector<1xi32> + %c0_i32v = vector.constant 0 : vector<1xi32> + %nibble_mask = vector.constant 252645135 : vector<1xi32> + %high_mask = vector.constant 808464432 : vector<1xi32> + %c0_f32 = scalar.constant 0.0 : f32 + %bounded_lane = index.assume %lane [range(%lane, 0, 31)] : index + %block_byte_add = index.scale %q6_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %weight_row_byte_base, %block_byte_add : offset + %qh_byte_base = index.add %block_byte_base, %qh_byte_add : offset + %scale_byte_base = index.add %block_byte_base, %scale_byte_add : offset + %d_byte_base = index.add %block_byte_base, %d_byte_add : offset + %ql_view = buffer.view %weight[%block_byte_base] : buffer -> view<32xi32> + %qh_view = buffer.view %weight[%qh_byte_base] : buffer -> view<16xi32> + %scale_view = buffer.view %weight[%scale_byte_base] : buffer -> view<16xi8> + %d_view = buffer.view %weight[%d_byte_base] : buffer -> view<1xf16> + %lane_mod8 = index.rem %bounded_lane, %c8 : index + %lane_mod16 = index.rem %bounded_lane, %c16 : index + %lane_div16 = index.div %bounded_lane, %c16 : index + %lane_div8_in_16 = index.div %lane_mod16, %c8 : index + %lane_div4_in_16 = index.div %lane_mod16, %c4 : index + %qh_high_base = index.mul %lane_div16, %c8 : index + %qh_index0 = index.add %qh_high_base, %lane_mod8 : index + %qh_index = index.assume %qh_index0 [range(%qh_index0, 0, 15)] : index + %ql_word = vector.load %ql_view[%bounded_lane] : view<32xi32> -> vector<1xi32> + %qh_word = vector.load %qh_view[%qh_index] : view<16xi32> -> vector<1xi32> + %qh_base_shift_index = index.mul %lane_div8_in_16, %c2 : index + %qh_base_shift_i32 = index.cast %qh_base_shift_index : index to i32 + %q8_block_base = index.mul %q6_block, %c8 : index + %q8_high_add = index.mul %lane_div16, %c4 : index + %q8_quadrant = index.add %q8_high_add, %lane_div8_in_16 : index + %scale_high_base = index.mul %lane_div16, %c8 : index + %scale_lane0 = index.add %scale_high_base, %lane_div4_in_16 : index + %d_f16 = view.load %d_view[0] : view<1xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %sum = scf.for %part = [%c0 to %c2 step %c1](%accumulator = %c0_f32 : f32) -> (f32) unroll { + %bounded_part = index.assume %part [range(%part, 0, 1)] : index + %part_shift_index = index.mul %bounded_part, %c4 : index + %part_shift_i32 = index.cast %part_shift_index : index to i32 + %part_shift = vector.splat %part_shift_i32 : vector<1xi32> + %ql_shifted = vector.shrui %ql_word, %part_shift : vector<1xi32> + %ql = vector.andi %ql_shifted, %nibble_mask : vector<1xi32> + %qh_shift_i32 = scalar.addi %qh_base_shift_i32, %part_shift_i32 : i32 + %qh_shift = vector.splat %qh_shift_i32 : vector<1xi32> + %qh_shifted = vector.shrui %qh_word, %qh_shift : vector<1xi32> + %qh_positioned = vector.shli %qh_shifted, %c4_i32v : vector<1xi32> + %qh = vector.andi %qh_positioned, %high_mask : vector<1xi32> + %code = vector.ori %ql, %qh : vector<1xi32> + %signed_weight = func.call @ggml_q6k_sign_extend_dot4(%code) : (vector<1xi32>) -> (vector<4xi8>) + %q8_part_add = index.mul %bounded_part, %c2 : index + %q8_block_part = index.add %q8_block_base, %q8_part_add : index + %q8_block = index.add %q8_block_part, %q8_quadrant : index + %q8_values, %q8_d = func.call @ggml_q8_1_x4_word(%q8_input, %q8_row_byte_base, %q8_block, %lane_mod8) : (buffer, offset, index, index) -> (vector<4xi8>, f32) + %scale_lane = index.add %scale_lane0, %part_shift_index : index + %scale_i8 = view.load %scale_view[%scale_lane] : view<16xi8> -> i8 + %scale = scalar.sitofp %scale_i8 : i8 to f32 + %dot = vector.dot4i %signed_weight, %q8_values, %c0_i32v : vector<4xi8>, vector<4xi8>, vector<1xi32> + %dot_i32 = vector.extract %dot[0] : vector<1xi32> -> i32 + %dot_f32 = scalar.sitofp %dot_i32 : i32 to f32 + %scaled0 = scalar.mulf %dot_f32, %scale : f32 + %scaled1 = scalar.mulf %scaled0, %d : f32 + %contribution = scalar.mulf %scaled1, %q8_d : f32 + %next = scalar.addf %accumulator, %contribution : f32 + scf.yield %next : f32 + } + func.return %sum : f32 +} + +// Computes one lane's partial for a complete Q6_K row. Callers choose how +// rows and activations are tiled, then reduce the returned value by subgroup. +func.def inline @ggml_q6k_q8_1_x4_row_lane(%input_size: index, %weight: buffer, %weight_row_byte_base: offset, %q8_input: buffer, %q8_row_byte_base: offset, %lane: index) -> (f32) { + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + %c0_f32 = scalar.constant 0.0 : f32 + %block_count = index.div %bounded_input_size, %c256 : index + %result = scf.for %block = [%c0 to %block_count step %c1](%block_acc = %c0_f32 : f32) -> (f32) { + %contribution = func.call @ggml_q6k_q8_1_x4_block_lane(%weight, %weight_row_byte_base, %q8_input, %q8_row_byte_base, %block, %lane) : (buffer, offset, buffer, offset, index, index) -> (f32) + %next = scalar.addf %block_acc, %contribution : f32 + scf.yield %next : f32 + } + func.return %result : f32 +} + +// Maps four Q6_K blocks across the four 16-lane partitions of a wave64. Each +// physical lane contracts the pair of virtual wave32 lanes that own the same +// packed positions in the low and high halves of one block. The subgroup +// reduction therefore combines four complete blocks per loop iteration while +// preserving the canonical wave32 unpacking and dot-product primitive. +func.def inline @ggml_q6k_q8_1_x4_row_lane_wave64_block4(%input_size: index, %weight: buffer, %weight_row_byte_base: offset, %q8_input: buffer, %q8_row_byte_base: offset, %lane: index) -> (f32) { + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %bounded_lane = index.assume %lane [range(%lane, 0, 63)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c256 = index.constant 256 : index + %c0_f32 = scalar.constant 0.0 : f32 + %block_count = index.div %bounded_input_size, %c256 : index + %padded_block_count = index.add %block_count, %c3 : index + %iteration_count = index.div %padded_block_count, %c4 : index + %block_in_iteration = index.div %bounded_lane, %c16 : index + %virtual_lane0 = index.rem %bounded_lane, %c16 : index + %virtual_lane1 = index.add %virtual_lane0, %c16 : index + %result = scf.for %iteration = [%c0 to %iteration_count step %c1](%row_acc = %c0_f32 : f32) -> (f32) { + %block_base = index.mul %iteration, %c4 : index + %block0 = index.add %block_base, %block_in_iteration : index + %valid_block = index.cmp ult, %block0, %block_count : index + %block_acc = scf.if %valid_block -> (f32) { + %low = func.call @ggml_q6k_q8_1_x4_block_lane(%weight, %weight_row_byte_base, %q8_input, %q8_row_byte_base, %block0, %virtual_lane0) : (buffer, offset, buffer, offset, index, index) -> (f32) + %high = func.call @ggml_q6k_q8_1_x4_block_lane(%weight, %weight_row_byte_base, %q8_input, %q8_row_byte_base, %block0, %virtual_lane1) : (buffer, offset, buffer, offset, index, index) -> (f32) + %sum = scalar.addf %low, %high : f32 + scf.yield %sum : f32 + } else { + scf.yield %c0_f32 : f32 + } + %next = scalar.addf %row_acc, %block_acc : f32 + scf.yield %next : f32 + } + func.return %result : f32 +} + +// Contracts one output channel for the calling Q6_K projection kernel. The +// caller owns publication so the same canonical contraction can feed either a +// dense logits tensor or an endpoint reduction. +template.def<@ggml.linear_q6k_q8_1_x4.body> device @ggml_linear_q6k_q8_1_x4_body(%publish_output: i1, %token_count: index, %token0: index, %input_size: index, %output_size: index, %q8_input: buffer, %weight: buffer) -> (f32, index, index, i1, index, index) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 262144)] : index + %channel_tile = kernel.workgroup.id : index + %subgroup0 = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c8 = index.constant 8 : index + %c128 = index.constant 128 : index + %c210_bytes = index.constant 210 : offset + %c256 = index.constant 256 : index + %c144_bytes = index.constant 144 : offset + %token, %launch_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %channel_base = index.mul %channel_tile, %c8 : index + %channel = index.add %channel_base, %subgroup : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %q6_block_count = index.div %bounded_input_size, %c256 : index + %weight_row_bytes = index.scale %q6_block_count, %c210_bytes : index, offset -> offset + %weight_row_byte_base = index.scale %channel, %weight_row_bytes : index, offset -> offset + %q8_group_count = index.div %bounded_input_size, %c128 : index + %q8_row_bytes = index.scale %q8_group_count, %c144_bytes : index, offset -> offset + %q8_row_byte_base = index.scale %token, %q8_row_bytes : index, offset -> offset + %lane_acc = scf.if %publish_output -> (f32) { + %channel_acc = scf.if %valid_channel -> (f32) { + %value = func.call @ggml_q6k_q8_1_x4_row_lane(%bounded_input_size, %weight, %weight_row_byte_base, %q8_input, %q8_row_byte_base, %lane) : (index, buffer, offset, buffer, offset, index) -> (f32) + scf.yield %value : f32 + } else { + %c0_f32 = scalar.constant 0.0 : f32 + scf.yield %c0_f32 : f32 + } + scf.yield %channel_acc : f32 + } else { + %c0_f32 = scalar.constant 0.0 : f32 + scf.yield %c0_f32 : f32 + } + %dot = kernel.subgroup.reduce %lane_acc : f32 + template.return %dot, %token, %channel, %valid_channel, %launch_token_count, %bounded_output_size : f32, index, index, i1, index, index +} + +// Dense raw-layout baseline. A 256-thread workgroup carries eight independent +// wave32 output rows so decode launches enough waves without duplicating Q6_K +// decoding within a row. +kernel.def export("ggml_linear_q6k_q8_1_x4") @ggml_linear_q6k_q8_1_x4(%token_count: index, %input_size: index, %output_size: index) { + %token_capacity = config.get @ggml.linear_q6k_q8_1_x4.token_capacity : index + %output_capacity = config.get @ggml.linear_q6k_q8_1_x4.output_capacity : index + %c8 = index.constant 8 : index + %c7 = index.constant 7 : index + %c1 = index.constant 1 : index + %workgroup_size = index.constant 256 : index + %padded_output_size = index.add %output_capacity, %c7 : index + %output_tiles = index.div %padded_output_size, %c8 : index + kernel.launch.config workgroups(%output_tiles, %token_capacity, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %input_size: index, %output_size: index, %q8_input: buffer, %weight: buffer, %output: buffer) { + %q8_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %q8_input, %weight, %output : buffer, buffer, buffer + %token_capacity = config.get @ggml.linear_q6k_q8_1_x4.token_capacity : index + %output_capacity = config.get @ggml.linear_q6k_q8_1_x4.output_capacity : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 262144), le(%output_size, %output_capacity)] : index + %c0 = index.constant 0 : index + %token0 = kernel.workgroup.id : index + %valid_token = index.cmp ult, %token0, %bounded_token_count : index + %safe_token = scf.select %valid_token, %token0, %c0 : index + %dot, %token, %channel, %valid_channel, %launch_token_count, %output_bound = template.apply<@ggml.linear_q6k_q8_1_x4.body>(%valid_token, %bounded_token_count, %safe_token, %input_size, %bounded_output_size, %q8_noalias, %weight_noalias) : (i1, index, index, index, index, buffer, buffer) -> (f32, index, index, i1, index, index) + %c0_i32 = scalar.constant 0 : i32 + %c0_offset = index.constant 0 : offset + %lane = kernel.subgroup.lane.id : index + %lane_i32 = index.cast %lane : index to i32 + %is_lane_zero = scalar.cmpi eq, %lane_i32, %c0_i32 : i32 + scf.if %valid_token { + scf.if %valid_channel { + scf.if %is_lane_zero { + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%output_bound]xf32> + view.store %dot, %output_view[%token, %channel] : f32, view<[%launch_token_count]x[%output_bound]xf32> + } + } + } + kernel.return +} + +// Wave64 vocabulary schedule matching llama.cpp Vulkan's one-row workgroup and +// four-block concurrency on AMD targets with native wave64 execution. +kernel.def target(@ggml_q6k_q8_gfx1151_wave64) export("ggml_linear_q6k_q8_1_x4") @ggml_linear_q6k_q8_1_x4_wave64_block4(%token_count: index, %input_size: index, %output_size: index) { + %token_capacity = config.get @ggml.linear_q6k_q8_1_x4.token_capacity : index + %output_capacity = config.get @ggml.linear_q6k_q8_1_x4.output_capacity : index + %c1 = index.constant 1 : index + %subgroup_size = target.subgroup.size : index + kernel.launch.config workgroups(%output_capacity, %token_capacity, %c1) workgroup_size(%subgroup_size, %c1, %c1) : index +} launch(%token_count: index, %input_size: index, %output_size: index, %q8_input: buffer, %weight: buffer, %output: buffer) { + %q8_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %q8_input, %weight, %output : buffer, buffer, buffer + %token_capacity = config.get @ggml.linear_q6k_q8_1_x4.token_capacity : index + %output_capacity = config.get @ggml.linear_q6k_q8_1_x4.output_capacity : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 262144), le(%output_size, %output_capacity)] : index + %token = kernel.workgroup.id : index + %channel = kernel.workgroup.id : index + %lane = kernel.subgroup.lane.id : index + %valid_token = index.cmp ult, %token, %bounded_token_count : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %publish_output = scalar.andi %valid_token, %valid_channel : i1 + %c0 = index.constant 0 : index + %c128 = index.constant 128 : index + %c144_bytes = index.constant 144 : offset + %c210_bytes = index.constant 210 : offset + %c256 = index.constant 256 : index + %safe_token = scf.select %valid_token, %token, %c0 : index + %safe_channel = scf.select %valid_channel, %channel, %c0 : index + %q6_block_count = index.div %bounded_input_size, %c256 : index + %weight_row_bytes = index.scale %q6_block_count, %c210_bytes : index, offset -> offset + %weight_row_byte_base = index.scale %safe_channel, %weight_row_bytes : index, offset -> offset + %q8_group_count = index.div %bounded_input_size, %c128 : index + %q8_row_bytes = index.scale %q8_group_count, %c144_bytes : index, offset -> offset + %q8_row_byte_base = index.scale %safe_token, %q8_row_bytes : index, offset -> offset + %lane_acc = scf.if %publish_output -> (f32) { + %value = func.call @ggml_q6k_q8_1_x4_row_lane_wave64_block4(%bounded_input_size, %weight_noalias, %weight_row_byte_base, %q8_noalias, %q8_row_byte_base, %lane) : (index, buffer, offset, buffer, offset, index) -> (f32) + scf.yield %value : f32 + } else { + %c0_f32 = scalar.constant 0.0 : f32 + scf.yield %c0_f32 : f32 + } + %dot = kernel.subgroup.reduce %lane_acc : f32 + %is_lane_zero = index.cmp eq, %lane, %c0 : index + scf.if %publish_output { + scf.if %is_lane_zero { + %c0_offset = index.constant 0 : offset + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_output_size]xf32> + view.store %dot, %output_view[%token, %channel] : f32, view<[%bounded_token_count]x[%bounded_output_size]xf32> + } + } + kernel.return +} + +// Uniform packed bytes exercise both low/high Q6 fields and signed per-group +// scales. The output width crosses the eight-wave tile boundary, while K=2048 +// covers the exact mixed-format dense V contraction depth. +check.case public @ggml_linear_q6k_q8_1_x4_nonzero_tail_case { + %token_count = check.literal value(2) : index + %input_size = check.literal value(2048) : index + %output_size = check.literal value(9) : index + %input = check.generate.fill value(0.00390625) : tensor<2x2048xf32> + %q8_input = check.generate.fill value(0) : tensor<2x2304xi8> + %weight = check.generate.fill value(-86) : tensor<9x8x210xi8> + %output = check.generate.fill value(0.0) : tensor<2x9xf32> + %expected = check.generate.fill value(358.1715087890625) : tensor<2x9xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<2x2048xf32>, tensor<2x2304xi8>) + kernel.launch @ggml_linear_q6k_q8_1_x4[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %q8_input, %weight, %output) : [index, index, index](index, index, index, tensor<2x2304xi8>, tensor<9x8x210xi8>, tensor<2x9xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.25) rtol(9.9999999999999995e-07) nan(same) : tensor<2x9xf32> + check.return +} + +// The gfx1151 schedule has an exact target requirement, so its nonzero packed +// data case remains independently selectable from the generic gfx11 case. +check.case public @ggml_linear_q6k_q8_1_x4_wave64_block4_nonzero_tail_case { + %token_count = check.literal value(2) : index + %input_size = check.literal value(2048) : index + %output_size = check.literal value(9) : index + %input = check.generate.fill value(0.00390625) : tensor<2x2048xf32> + %q8_input = check.generate.fill value(0) : tensor<2x2304xi8> + %weight = check.generate.fill value(-86) : tensor<9x8x210xi8> + %output = check.generate.fill value(0.0) : tensor<2x9xf32> + %expected = check.generate.fill value(358.1715087890625) : tensor<2x9xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<2x2048xf32>, tensor<2x2304xi8>) + kernel.launch @ggml_linear_q6k_q8_1_x4_wave64_block4[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %q8_input, %weight, %output) : [index, index, index](index, index, index, tensor<2x2304xi8>, tensor<9x8x210xi8>, tensor<2x9xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.25) rtol(9.9999999999999995e-07) nan(same) : tensor<2x9xf32> + check.return +} + +check.case public @ggml_linear_q6k_q8_1_x4_benchmark_case { + %token_count = check.param.choice values([1, 17, 32, 63, 128, 129, 512]) name("token_count") : index + %input_size = check.literal value(768) : index + %output_size = check.literal value(2048) : index + %q8_input = check.generate.fill value(0) : tensor<[%token_count]x864xi8> + %weight = check.generate.fill value(0) : tensor<2048x3x210xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + kernel.launch @ggml_linear_q6k_q8_1_x4[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %q8_input, %weight, %output) : [index, index, index](index, index, index, tensor<[%token_count]x864xi8>, tensor<2048x3x210xi8>, tensor<[%token_count]x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x2048xf32> + check.return +} + +// Production vocabulary shape used to compare the eight-row wave32 baseline +// with the one-row wave64 block schedule without model-runtime overhead. +check.case public @ggml_linear_q6k_q8_1_x4_vocabulary_wave32_benchmark_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(2048) : index + %output_size = check.literal value(151936) : index + %q8_input = check.generate.fill value(1) : tensor<1x2304xi8> + %weight = check.generate.fill value(0) : tensor<151936x8x210xi8> + %output = check.generate.fill value(1.0) : tensor<1x151936xf32> + %expected = check.generate.fill value(0.0) : tensor<1x151936xf32> + kernel.launch @ggml_linear_q6k_q8_1_x4[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %q8_input, %weight, %output) : [index, index, index](index, index, index, tensor<1x2304xi8>, tensor<151936x8x210xi8>, tensor<1x151936xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<1x151936xf32> + check.return +} + +check.case public @ggml_linear_q6k_q8_1_x4_vocabulary_wave64_block4_benchmark_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(2048) : index + %output_size = check.literal value(151936) : index + %q8_input = check.generate.fill value(1) : tensor<1x2304xi8> + %weight = check.generate.fill value(0) : tensor<151936x8x210xi8> + %output = check.generate.fill value(1.0) : tensor<1x151936xf32> + %expected = check.generate.fill value(0.0) : tensor<1x151936xf32> + kernel.launch @ggml_linear_q6k_q8_1_x4_wave64_block4[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %q8_input, %weight, %output) : [index, index, index](index, index, index, tensor<1x2304xi8>, tensor<151936x8x210xi8>, tensor<1x151936xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<1x151936xf32> + check.return +} + +// Exercises the complete activation-pack and dense Q6_K projection boundary at +// the K=2048, M=512 shape used by mixed-format attention V weights. The token +// buckets cover decode, awkward prefill tails, llama.cpp's 512-token +// microbatch, and larger schedules available to callers without host routing. +check.case public @ggml_linear_q6k_q8_1_x4_dense_v_benchmark_case { + %token_count = check.param.choice values([1, 17, 32, 63, 128, 129, 512, 1024, 2048]) name("token_count") : index + %input_size = check.literal value(2048) : index + %output_size = check.literal value(512) : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %q8_input = check.generate.fill value(1) : tensor<[%token_count]x2304xi8> + %weight = check.generate.fill value(0) : tensor<512x8x210xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x512xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x512xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<[%token_count]x2048xf32>, tensor<[%token_count]x2304xi8>) + kernel.launch @ggml_linear_q6k_q8_1_x4[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %q8_input, %weight, %output) : [index, index, index](index, index, index, tensor<[%token_count]x2304xi8>, tensor<512x8x210xi8>, tensor<[%token_count]x512xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x512xf32> + check.return +} + +check.benchmark<@ggml_linear_q6k_q8_1_x4_nonzero_tail_case> @ggml_linear_q6k_q8_1_x4_small + +check.benchmark<@ggml_linear_q6k_q8_1_x4_benchmark_case> @ggml_linear_q6k_q8_1_x4_decode {token_count = 1} + +check.benchmark<@ggml_linear_q6k_q8_1_x4_vocabulary_wave32_benchmark_case> @ggml_linear_q6k_q8_1_x4_vocabulary_wave32 + +check.benchmark<@ggml_linear_q6k_q8_1_x4_vocabulary_wave64_block4_benchmark_case> @ggml_linear_q6k_q8_1_x4_vocabulary_wave64_block4 + +check.benchmark<@ggml_linear_q6k_q8_1_x4_benchmark_case> @ggml_linear_q6k_q8_1_x4_prefill_32 {token_count = 32} + +check.benchmark<@ggml_linear_q6k_q8_1_x4_benchmark_case> @ggml_linear_q6k_q8_1_x4_prefill_128 {token_count = 128} + +check.benchmark<@ggml_linear_q6k_q8_1_x4_benchmark_case> @ggml_linear_q6k_q8_1_x4_prefill_512 {token_count = 512} + +check.benchmark<@ggml_linear_q6k_q8_1_x4_dense_v_benchmark_case> @ggml_linear_q6k_q8_1_x4_dense_v_decode {token_count = 1} + +check.benchmark<@ggml_linear_q6k_q8_1_x4_dense_v_benchmark_case> @ggml_linear_q6k_q8_1_x4_dense_v_prefill_32 {token_count = 32} + +check.benchmark<@ggml_linear_q6k_q8_1_x4_dense_v_benchmark_case> @ggml_linear_q6k_q8_1_x4_dense_v_prefill_128 {token_count = 128} + +check.benchmark<@ggml_linear_q6k_q8_1_x4_dense_v_benchmark_case> @ggml_linear_q6k_q8_1_x4_dense_v_prefill_512 {token_count = 512} + +check.benchmark<@ggml_linear_q6k_q8_1_x4_dense_v_benchmark_case> @ggml_linear_q6k_q8_1_x4_dense_v_prefill_1024 {token_count = 1024} + +check.benchmark<@ggml_linear_q6k_q8_1_x4_dense_v_benchmark_case> @ggml_linear_q6k_q8_1_x4_dense_v_prefill_2048 {token_count = 2048} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/ggml/quantize_q8_1_x4.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/ggml/quantize_q8_1_x4.loom new file mode 100644 index 000000000000..95c96b299276 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/ggml/quantize_q8_1_x4.loom @@ -0,0 +1,276 @@ +// Defines GGML's block_q8_1_x4 physical-layout accessors and packs contiguous +// F32 activations into that layout. Four logical 32-element Q8_1 blocks share +// one 144-byte physical group: +// +// struct block_q8_1_x4 { +// f16 ds[4][2]; // per-block (scale, quantized_sum * scale) +// i32 qs[4][8]; // four signed i8 values per packed word +// }; +template.decl @ggml.quantize_q8_1_x4.group_body(%publish_output: i1, %group_count0: index, %group0: index, %input: buffer, %output: buffer) + +config.decl @ggml.quantize_q8_1_x4.group_capacity : %value: index where [range(%value, 1, 524288)] + +// Loads one logical 32-element Q8_1 block from the four-way physical packing. +// `row_byte_base` addresses the first x4 group for one activation row. +func.def inline @ggml_q8_1_x4_block(%q8_input: buffer, %row_byte_base: offset, %q8_block: index) -> (vector<32xi8>, f32, f32) { + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %group_bytes = index.constant 144 : offset + %payload_byte_add = index.constant 16 : offset + %group = index.div %q8_block, %c4 : index + %inner0 = index.rem %q8_block, %c4 : index + %inner = index.assume %inner0 [range(%inner0, 0, 3)] : index + %group_byte_add = index.scale %group, %group_bytes : index, offset -> offset + %group_byte_base = index.add %row_byte_base, %group_byte_add : offset + %payload_byte_base = index.add %group_byte_base, %payload_byte_add : offset + %ds_view = buffer.view %q8_input[%group_byte_base] : buffer -> view<8xf16> + %payload_view = buffer.view %q8_input[%payload_byte_base] : buffer -> view<32xi32> + %d_index = index.mul %inner, %c2 : index + %s_index = index.add %d_index, %c1 : index + %word_index = index.mul %inner, %c8 : index + %d_f16 = view.load %ds_view[%d_index] : view<8xf16> -> f16 + %s_f16 = view.load %ds_view[%s_index] : view<8xf16> -> f16 + %words = vector.load %payload_view[%word_index] : view<32xi32> -> vector<8xi32> + %values = vector.bitcast %words : vector<8xi32> to vector<32xi8> + %d = scalar.extf %d_f16 : f16 to f32 + %s = scalar.extf %s_f16 : f16 to f32 + func.return %values, %d, %s : vector<32xi8>, f32, f32 +} + +// Loads one packed four-value word and its block scale. Dot-product kernels +// use this narrower form so a lane does not load the other seven words owned +// by its subgroup peers. +func.def inline @ggml_q8_1_x4_word(%q8_input: buffer, %row_byte_base: offset, %q8_block: index, %word_in_block: index) -> (vector<4xi8>, f32) { + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %group_bytes = index.constant 144 : offset + %payload_byte_add = index.constant 16 : offset + %group = index.div %q8_block, %c4 : index + %inner0 = index.rem %q8_block, %c4 : index + %inner = index.assume %inner0 [range(%inner0, 0, 3)] : index + %word0 = index.assume %word_in_block [range(%word_in_block, 0, 7)] : index + %group_byte_add = index.scale %group, %group_bytes : index, offset -> offset + %group_byte_base = index.add %row_byte_base, %group_byte_add : offset + %payload_byte_base = index.add %group_byte_base, %payload_byte_add : offset + %ds_view = buffer.view %q8_input[%group_byte_base] : buffer -> view<8xf16> + %payload_view = buffer.view %q8_input[%payload_byte_base] : buffer -> view<32xi32> + %d_index = index.mul %inner, %c2 : index + %inner_word_base = index.mul %inner, %c8 : index + %word_index = index.add %inner_word_base, %word0 : index + %d_f16 = view.load %ds_view[%d_index] : view<8xf16> -> f16 + %packed = vector.load %payload_view[%word_index] : view<32xi32> -> vector<1xi32> + %values = vector.bitcast %packed : vector<1xi32> to vector<4xi8> + %d = scalar.extf %d_f16 : f16 to f32 + func.return %values, %d : vector<4xi8>, f32 +} + +// Packs one explicit physical group. Callers provide the complete physical +// group domain and an in-range ordinal, and guarantee that one complete +// 32-lane wave executes after all 128 source values are visible. The uniform +// publication predicate suppresses only destination writes so every lane still +// reaches the workgroup barriers. +template.def<@ggml.quantize_q8_1_x4.group_body> device @ggml_quantize_q8_1_x4_group_body(%publish_output: i1, %group_count0: index, %group0: index, %input: buffer, %output: buffer) { + %group_count, %group = index.assume %group_count0, %group0 [range(%group_count0, 1, 524288), lt(%group0, %group_count0)] : index, index + %lane = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c32 = index.constant 32 : index + %c128 = index.constant 128 : index + %group_bytes = index.constant 144 : offset + %payload_byte_add = index.constant 16 : offset + %scratch_d_byte_add = index.constant 128 : offset + %scratch_bytes = index.constant 144 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %c1_f32 = scalar.constant 1.0 : f32 + %c127 = scalar.constant 127.0 : f32 + %c0_offset = index.constant 0 : offset + %launched_element_count = index.mul %group_count, %c128 : index + %block_in_group0 = index.div %lane, %c8 : index + %block_in_group = index.assume %block_in_group0 [range(%block_in_group0, 0, 3)] : index + %word_in_block0 = index.rem %lane, %c8 : index + %word_in_block = index.assume %word_in_block0 [range(%word_in_block0, 0, 7)] : index + %group_element_base = index.mul %group, %c128 : index + %block_element_add = index.mul %block_in_group, %c32 : index + %block_word_add = index.mul %block_in_group, %c8 : index + %word_element_add = index.mul %word_in_block, %c4 : index + %input_block_base = index.add %group_element_base, %block_element_add : index + %input_index = index.add %input_block_base, %word_element_add : index + %input_noalias, %output_noalias = buffer.assume.noalias %input, %output : buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%launched_element_count]xf32> + %input_values = vector.load %input_view[%input_index] : view<[%launched_element_count]xf32> -> vector<4xf32> + %absolute_values = vector.absf %input_values : vector<4xf32> + %thread_max = vector.reduce %absolute_values, %c0_f32 : vector<4xf32>, f32 + %scratch = buffer.alloca align(16) %scratch_bytes : buffer + %scratch_values = buffer.view %scratch[%c0_offset] : buffer -> view<32xf32> + %scratch_d = buffer.view %scratch[%scratch_d_byte_add] : buffer -> view<4xf32> + view.store %thread_max, %scratch_values[%lane] : f32, view<32xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + %is_cohort_leader = index.cmp eq, %word_in_block, %c0 : index + scf.if %is_cohort_leader { + %cohort_base = index.mul %block_in_group, %c8 : index + %cohort_maxima = vector.load %scratch_values[%cohort_base] : view<32xf32> -> vector<8xf32> + %amax = vector.reduce %cohort_maxima, %c0_f32 : vector<8xf32>, f32 + %d = scalar.divf %amax, %c127 : f32 + view.store %d, %scratch_d[%block_in_group] : f32, view<4xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %d = view.load %scratch_d[%block_in_group] : view<4xf32> -> f32 + %d_nonzero = scalar.cmpf one, %d, %c0_f32 : f32 + %d_inverse = scf.if %d_nonzero -> (f32) { + %inverse = scalar.divf %c1_f32, %d : f32 + scf.yield %inverse : f32 + } else { + scf.yield %c0_f32 : f32 + } + %d_inverse_vector = vector.splat %d_inverse : vector<4xf32> + %scaled_values = vector.mulf %input_values, %d_inverse_vector : vector<4xf32> + %rounded_values = vector.roundf %scaled_values : vector<4xf32> + %quantized_values = vector.fptosi %rounded_values : vector<4xf32> to vector<4xi8> + %packed_word = vector.bitcast %quantized_values : vector<4xi8> to vector<1xi32> + %group_byte_offset = index.scale %group, %group_bytes : index, offset -> offset + %payload_byte_offset = index.add %group_byte_offset, %payload_byte_add : offset + %group_ds = buffer.view %output_noalias[%group_byte_offset] : buffer -> view<8xf16> + %group_qs = buffer.view %output_noalias[%payload_byte_offset] : buffer -> view<32xi32> + %packed_word_index0 = index.add %block_word_add, %word_in_block : index + %packed_word_index = index.assume %packed_word_index0 [range(%packed_word_index0, 0, 31)] : index + scf.if %publish_output { + vector.store %packed_word, %group_qs[%packed_word_index] : vector<1xi32>, view<32xi32> + } + %thread_sum = vector.reduce %rounded_values, %c0_f32 : vector<4xf32>, f32 + view.store %thread_sum, %scratch_values[%lane] : f32, view<32xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + %publishes_metadata = scalar.andi %is_cohort_leader, %publish_output : i1 + scf.if %publishes_metadata { + %cohort_base = index.mul %block_in_group, %c8 : index + %cohort_sums = vector.load %scratch_values[%cohort_base] : view<32xf32> -> vector<8xf32> + %quantized_sum = vector.reduce %cohort_sums, %c0_f32 : vector<8xf32>, f32 + %s = scalar.mulf %quantized_sum, %d : f32 + %d_f16 = scalar.fptrunc %d : f32 to f16 + %s_f16 = scalar.fptrunc %s : f32 to f16 + %ds_index = index.mul %block_in_group, %c2 : index + view.store %d_f16, %group_ds[%ds_index] : f16, view<8xf16> + %s_index = index.add %ds_index, %c1 : index + view.store %s_f16, %group_ds[%s_index] : f16, view<8xf16> + } + template.return +} + +// A 32-lane workgroup owns one physical group. Each eight-lane cohort loads +// four values, computes one Q8_1 block, and writes disjoint metadata and packed +// words. The LDS reductions make the independent eight-lane cohorts explicit. +kernel.def @ggml_quantize_q8_1_x4_f32(%token_count: index, %input_size: index) { + %group_capacity = config.get @ggml.quantize_q8_1_x4.group_capacity : index + %unit = index.constant 1 : index + %workgroup_size = index.constant 32 : index + kernel.launch.config workgroups(%group_capacity, %unit, %unit) workgroup_size(%workgroup_size, %unit, %unit) : index +} launch(%token_count: index, %input_size: index, %input: buffer, %output: buffer) { + %group_capacity = config.get @ggml.quantize_q8_1_x4.group_capacity : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_input_size = index.assume %input_size [range(%input_size, 128, 32768), mul(%input_size, 128)] : index + %elements_per_group = index.constant 128 : index + %c0 = index.constant 0 : index + %element_count = index.mul %bounded_token_count, %bounded_input_size : index + %group_count0 = index.div %element_count, %elements_per_group : index + %group_count = index.assume %group_count0 [range(%group_count0, 1, 524288), le(%group_count0, %group_capacity)] : index + %group0 = kernel.workgroup.id : index + %valid_group = index.cmp ult, %group0, %group_count : index + %safe_group = scf.select %valid_group, %group0, %c0 : index + %group = index.assume %safe_group [lt(%safe_group, %group_count)] : index + template.apply<@ggml.quantize_q8_1_x4.group_body>(%valid_group, %group_count, %group, %input, %output) : (i1, index, index, buffer, buffer) + kernel.return +} + +// This check-only inspector keeps the production packer ABI exact while +// exposing packed words and heterogeneous metadata as comparable tensors. +kernel.def @ggml_q8_1_x4_inspect_one_group() { + %unit = index.constant 1 : index + %workgroup_size = index.constant 32 : index + kernel.launch.config workgroups(%unit, %unit, %unit) workgroup_size(%workgroup_size, %unit, %unit) : index +} launch(%packed: buffer, %words: buffer, %d_values: buffer, %s_values: buffer) { + %lane = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %payload_offset = index.constant 16 : offset + %c0_offset = index.constant 0 : offset + %packed_ds = buffer.view %packed[%c0_offset] : buffer -> view<8xf16> + %packed_words = buffer.view %packed[%payload_offset] : buffer -> view<32xi32> + %word_output = buffer.view %words[%c0_offset] : buffer -> view<32xi32> + %d_output = buffer.view %d_values[%c0_offset] : buffer -> view<4xf32> + %s_output = buffer.view %s_values[%c0_offset] : buffer -> view<4xf32> + %word = view.load %packed_words[%lane] : view<32xi32> -> i32 + view.store %word, %word_output[%lane] : i32, view<32xi32> + %is_metadata_lane = index.cmp ult, %lane, %c4 : index + scf.if %is_metadata_lane { + %ds_index = index.mul %lane, %c2 : index + %s_index = index.add %ds_index, %c1 : index + %d_f16 = view.load %packed_ds[%ds_index] : view<8xf16> -> f16 + %s_f16 = view.load %packed_ds[%s_index] : view<8xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %s = scalar.extf %s_f16 : f16 to f32 + view.store %d, %d_output[%lane] : f32, view<4xf32> + view.store %s, %s_output[%lane] : f32, view<4xf32> + } + kernel.return +} + +check.case public @ggml_quantize_q8_1_x4_f32_nonzero_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(128) : index + %positive_input = check.generate.fill value(1.0) : tensor<128xf32> + %negative_input = check.generate.fill value(-1.0) : tensor<128xf32> + %positive_packed = check.generate.fill value(0) : tensor<144xi8> + %negative_packed = check.generate.fill value(0) : tensor<144xi8> + %positive_words = check.generate.fill value(0) : tensor<32xi32> + %negative_words = check.generate.fill value(0) : tensor<32xi32> + %positive_d = check.generate.fill value(0.0) : tensor<4xf32> + %negative_d = check.generate.fill value(0.0) : tensor<4xf32> + %positive_s = check.generate.fill value(0.0) : tensor<4xf32> + %negative_s = check.generate.fill value(0.0) : tensor<4xf32> + %expected_positive_words = check.generate.fill value(2139062143) : tensor<32xi32> + %expected_negative_words = check.generate.fill value(-2122219135) : tensor<32xi32> + %expected_d = check.generate.fill value(0.00787353515625) : tensor<4xf32> + %expected_positive_s = check.generate.fill value(32.0) : tensor<4xf32> + %expected_negative_s = check.generate.fill value(-32.0) : tensor<4xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %positive_input, %positive_packed) : [index, index](index, index, tensor<128xf32>, tensor<144xi8>) + kernel.launch @ggml_q8_1_x4_inspect_one_group(%positive_packed, %positive_words, %positive_d, %positive_s) : (tensor<144xi8>, tensor<32xi32>, tensor<4xf32>, tensor<4xf32>) + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %negative_input, %negative_packed) : [index, index](index, index, tensor<128xf32>, tensor<144xi8>) + kernel.launch @ggml_q8_1_x4_inspect_one_group(%negative_packed, %negative_words, %negative_d, %negative_s) : (tensor<144xi8>, tensor<32xi32>, tensor<4xf32>, tensor<4xf32>) + check.expect.equal actual(%positive_words) expected(%expected_positive_words) : tensor<32xi32> + check.expect.equal actual(%negative_words) expected(%expected_negative_words) : tensor<32xi32> + check.expect.close actual(%positive_d) expected(%expected_d) atol(0.0) rtol(0.0) nan(same) : tensor<4xf32> + check.expect.close actual(%negative_d) expected(%expected_d) atol(0.0) rtol(0.0) nan(same) : tensor<4xf32> + check.expect.close actual(%positive_s) expected(%expected_positive_s) atol(0.0) rtol(0.0) nan(same) : tensor<4xf32> + check.expect.close actual(%negative_s) expected(%expected_negative_s) atol(0.0) rtol(0.0) nan(same) : tensor<4xf32> + check.return +} + +check.case public @ggml_quantize_q8_1_x4_f32_benchmark_case { + %token_count = check.param.choice values([1, 4, 8, 17, 32, 63, 128, 129, 512]) name("token_count") : index + %input_size = check.literal value(2048) : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %packed = check.generate.fill value(1) : tensor<[%token_count]x2304xi8> + %expected = check.generate.fill value(0) : tensor<[%token_count]x2304xi8> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %packed) : [index, index](index, index, tensor<[%token_count]x2048xf32>, tensor<[%token_count]x2304xi8>) + check.expect.equal actual(%packed) expected(%expected) : tensor<[%token_count]x2304xi8> + check.return +} + +check.benchmark<@ggml_quantize_q8_1_x4_f32_nonzero_case> @ggml_quantize_q8_1_x4_f32_small + +check.benchmark<@ggml_quantize_q8_1_x4_f32_benchmark_case> @ggml_quantize_q8_1_x4_f32_decode {token_count = 1} + +check.benchmark<@ggml_quantize_q8_1_x4_f32_benchmark_case> @ggml_quantize_q8_1_x4_f32_small_batch_8 {token_count = 8} + +check.benchmark<@ggml_quantize_q8_1_x4_f32_benchmark_case> @ggml_quantize_q8_1_x4_f32_prefill_32 {token_count = 32} + +check.benchmark<@ggml_quantize_q8_1_x4_f32_benchmark_case> @ggml_quantize_q8_1_x4_f32_prefill_128 {token_count = 128} + +check.benchmark<@ggml_quantize_q8_1_x4_f32_benchmark_case> @ggml_quantize_q8_1_x4_f32_prefill_512 {token_count = 512} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/manifest.json b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/manifest.json new file mode 100644 index 000000000000..e627b8c0cacf --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/manifest.json @@ -0,0 +1,4828 @@ +{ + "schema": "ggml-hrx-kernel-corpus-v1", + "upstream_revision": "local", + "files": [ + { + "path": "ggml/linear_q6k_f32.loom" + }, + { + "path": "ggml/linear_q6k_q8_1_x4.loom" + }, + { + "path": "ggml/quantize_q8_1_x4.loom" + }, + { + "path": "qwen3_moe/attention_postprocess_f32_f16.loom" + }, + { + "path": "qwen3_moe/attention_prepare_quantized.loom" + }, + { + "path": "qwen3_moe/attention_qkv_postprocess_fused.loom" + }, + { + "path": "qwen3_moe/attention_qkv_quantized.loom" + }, + { + "path": "qwen3_moe/attention_qkv_same_format_prefill.loom" + }, + { + "path": "qwen3_moe/batched_decode_expert_dispatch.loom" + }, + { + "path": "qwen3_moe/batched_decode_gate_up_q4k.loom" + }, + { + "path": "qwen3_moe/dense_linear_quantized_f16_wmma.loom" + }, + { + "path": "qwen3_moe/expert_table_partition_fused.loom" + }, + { + "path": "qwen3_moe/flash_attention_decode_f32_f16_wmma.loom" + }, + { + "path": "qwen3_moe/flash_attention_decode_q128_f32_f16_wmma.loom" + }, + { + "path": "qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom" + }, + { + "path": "qwen3_moe/flash_attention_decode_split_next_q8_test.loom" + }, + { + "path": "qwen3_moe/flash_attention_f32_f16_wmma.loom" + }, + { + "path": "qwen3_moe/model_config.loom" + }, + { + "path": "qwen3_moe/routed_down_q4k.loom" + }, + { + "path": "qwen3_moe/routed_down_q6k.loom" + }, + { + "path": "qwen3_moe/routed_down_next_q8.loom" + }, + { + "path": "qwen3_moe/routed_down_quantized_f16_wmma.loom" + }, + { + "path": "qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_f32.loom" + }, + { + "path": "qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_q8_1_x4.loom" + }, + { + "path": "qwen3_moe/routed_gate_up_swiglu_q4k.loom" + }, + { + "path": "qwen3_moe/routed_linear_q4k_f16_wmma.loom" + }, + { + "path": "qwen3_moe/router_projection_f32.loom" + }, + { + "path": "qwen3_moe/router_projection_top8_fused_f32.loom" + }, + { + "path": "qwen3_moe/router_top8_f32.loom" + }, + { + "path": "../qwen/token_embedding_q4k.loom", + "upstream_path": "experimental/qwen/kernels/token_embedding_q4k.loom" + }, + { + "path": "../qwen/attention_metadata.loom", + "upstream_path": "experimental/qwen/kernels/attention_metadata.loom" + } + ], + "exports": [ + { + "name": "ggml_linear_q6k_f32_wave32", + "symbol": "ggml_linear_q6k_f32_wave32", + "source": "ggml/linear_q6k_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "linear_q6k_f32_linked", + "primary_sources": [ + "ggml/linear_q6k_f32.loom" + ], + "library_sources": [ + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_linear_q6k_f32_wave64", + "symbol": "ggml_linear_q6k_f32_wave64", + "source": "ggml/linear_q6k_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "linear_q6k_f32_linked", + "primary_sources": [ + "ggml/linear_q6k_f32.loom" + ], + "library_sources": [ + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_linear_q6k_q8_1_x4", + "symbol": "ggml_linear_q6k_q8_1_x4", + "source": "ggml/linear_q6k_q8_1_x4.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "linear_q6k_q8_1_x4_linked", + "primary_sources": [ + "ggml/linear_q6k_q8_1_x4.loom" + ], + "library_sources": [ + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_linear_q6k_q8_1_x4", + "symbol": "ggml_linear_q6k_q8_1_x4_wave64_block4", + "source": "ggml/linear_q6k_q8_1_x4.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "gfx1151", + "compile_recipe": { + "mode": "archive", + "link_module": "linear_q6k_q8_1_x4_linked", + "primary_sources": [ + "ggml/linear_q6k_q8_1_x4.loom" + ], + "library_sources": [ + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "ggml_q8_1_x4_inspect_one_group", + "symbol": "ggml_q8_1_x4_inspect_one_group", + "source": "ggml/quantize_q8_1_x4.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "packed", + "words", + "d_values", + "s_values" + ], + "binding_access": [ + "read", + "write", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ggml/quantize_q8_1_x4.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_quantize_q8_1_x4_f32", + "symbol": "ggml_quantize_q8_1_x4_f32", + "source": "ggml/quantize_q8_1_x4.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ggml/quantize_q8_1_x4.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "qwen3_moe_attention_key_q4_aggregate_prefill_512", + "symbol": "qwen3_moe_attention_key_q4_aggregate_prefill_512", + "source": "qwen3_moe/attention_qkv_same_format_prefill.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "combined_weight", + "combined_output" + ], + "binding_access": [ + "read", + "read", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "attention_qkv_same_format_prefill_linked", + "primary_sources": [ + "qwen3_moe/attention_qkv_same_format_prefill.loom" + ], + "library_sources": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_attention_postprocess_f32_f16", + "symbol": "qwen3_moe_attention_postprocess_f32_f16", + "source": "qwen3_moe/attention_postprocess_f32_f16.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + } + ], + "bindings": [ + "positions", + "key_cache_indices", + "value_cache_indices", + "query_input", + "key_input", + "value_input", + "query_norm_weight", + "key_norm_weight", + "inverse_frequencies", + "query_output", + "key_cache", + "value_cache" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read_write", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "attention_postprocess_f32_f16_linked", + "primary_sources": [ + "qwen3_moe/attention_postprocess_f32_f16.loom" + ], + "library_sources": [ + "qwen3_moe/model_config.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/model_config.loom" + ] + }, + { + "name": "qwen3_moe_attention_qkv_postprocess_fused_decode", + "symbol": "qwen3_moe_attention_qkv_postprocess_fused_decode", + "source": "qwen3_moe/attention_qkv_postprocess_fused.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "query_weight", + "key_weight", + "value_weight", + "positions", + "key_cache_indices", + "value_cache_indices", + "query_output_raw", + "key_output_raw", + "value_output_raw", + "query_norm_weight", + "key_norm_weight", + "inverse_frequencies", + "query_output", + "key_cache", + "value_cache", + "completion_counters" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read_write", + "read_write", + "read_write", + "read", + "read", + "read", + "read_write", + "read_write", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "attention_qkv_postprocess_fused_linked", + "primary_sources": [ + "qwen3_moe/attention_qkv_postprocess_fused.loom" + ], + "library_sources": [ + "qwen3_moe/attention_postprocess_f32_f16.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/attention_qkv_quantized.loom", + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/attention_postprocess_f32_f16.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/attention_qkv_quantized.loom", + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_attention_qkv_postprocess_fused_decode_q4", + "symbol": "qwen3_moe_attention_qkv_postprocess_fused_decode_q4", + "source": "qwen3_moe/attention_qkv_postprocess_fused.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "query_weight", + "key_weight", + "value_weight", + "positions", + "key_cache_indices", + "value_cache_indices", + "query_output_raw", + "key_output_raw", + "value_output_raw", + "query_norm_weight", + "key_norm_weight", + "inverse_frequencies", + "query_output", + "key_cache", + "value_cache", + "completion_counters" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read_write", + "read_write", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "attention_qkv_postprocess_fused_linked", + "primary_sources": [ + "qwen3_moe/attention_qkv_postprocess_fused.loom" + ], + "library_sources": [ + "qwen3_moe/attention_postprocess_f32_f16.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/attention_qkv_quantized.loom", + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/attention_postprocess_f32_f16.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/attention_qkv_quantized.loom", + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_attention_qkv_postprocess_fused_decode_q6", + "symbol": "qwen3_moe_attention_qkv_postprocess_fused_decode_q6", + "source": "qwen3_moe/attention_qkv_postprocess_fused.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "query_weight", + "key_weight", + "value_weight", + "positions", + "key_cache_indices", + "value_cache_indices", + "query_output_raw", + "key_output_raw", + "value_output_raw", + "query_norm_weight", + "key_norm_weight", + "inverse_frequencies", + "query_output", + "key_cache", + "value_cache", + "completion_counters" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read", + "read_write", + "read_write", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "attention_qkv_postprocess_fused_linked", + "primary_sources": [ + "qwen3_moe/attention_qkv_postprocess_fused.loom" + ], + "library_sources": [ + "qwen3_moe/attention_postprocess_f32_f16.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/attention_qkv_quantized.loom", + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/attention_postprocess_f32_f16.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/attention_qkv_quantized.loom", + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_attention_qkv_q4_prefill_512", + "symbol": "qwen3_moe_attention_qkv_q4_prefill_512", + "source": "qwen3_moe/attention_qkv_same_format_prefill.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "combined_weight", + "combined_output" + ], + "binding_access": [ + "read", + "read", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "attention_qkv_same_format_prefill_linked", + "primary_sources": [ + "qwen3_moe/attention_qkv_same_format_prefill.loom" + ], + "library_sources": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_attention_qkv_quantized", + "symbol": "qwen3_moe_attention_qkv_quantized", + "source": "qwen3_moe/attention_qkv_quantized.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "query_weight", + "key_weight", + "value_weight", + "query_output", + "key_output", + "value_output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "attention_qkv_quantized_linked", + "primary_sources": [ + "qwen3_moe/attention_qkv_quantized.loom" + ], + "library_sources": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_attention_query_q4_aggregate_prefill_512", + "symbol": "qwen3_moe_attention_query_q4_aggregate_prefill_512", + "source": "qwen3_moe/attention_qkv_same_format_prefill.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "combined_weight", + "combined_output" + ], + "binding_access": [ + "read", + "read", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "attention_qkv_same_format_prefill_linked", + "primary_sources": [ + "qwen3_moe/attention_qkv_same_format_prefill.loom" + ], + "library_sources": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_attention_rmsnorm_quantize_q8_1_x4", + "symbol": "qwen3_moe_attention_rmsnorm_quantize_q8_1_x4", + "source": "qwen3_moe/attention_prepare_quantized.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "attention_prepare_quantized_linked", + "primary_sources": [ + "qwen3_moe/attention_prepare_quantized.loom" + ], + "library_sources": [ + "qwen3_moe/model_config.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/model_config.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_attention_value_q4_aggregate_prefill_512", + "symbol": "qwen3_moe_attention_value_q4_aggregate_prefill_512", + "source": "qwen3_moe/attention_qkv_same_format_prefill.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "combined_weight", + "combined_output" + ], + "binding_access": [ + "read", + "read", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "attention_qkv_same_format_prefill_linked", + "primary_sources": [ + "qwen3_moe/attention_qkv_same_format_prefill.loom" + ], + "library_sources": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_batched_decode_gate_up_q4k_rows2", + "symbol": "qwen3_moe_batched_decode_gate_up_q4k_rows2", + "source": "qwen3_moe/batched_decode_gate_up_q4k.loom", + "workload_parameters": [ + { + "name": "descriptor_count", + "type": "index" + }, + { + "name": "queue_ordinal", + "type": "index" + }, + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "descriptor_count", + "type": "index" + }, + { + "name": "queue_ordinal", + "type": "index" + }, + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "bindings": [ + "queue_descriptors", + "assignment_ordinals", + "q8_input", + "gate_weight", + "up_weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "batched_decode_gate_up_q4k_linked", + "primary_sources": [ + "qwen3_moe/batched_decode_gate_up_q4k.loom" + ], + "library_sources": [ + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/batched_decode_expert_dispatch.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom" + ] + }, + "compile_dependencies": [ + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/batched_decode_expert_dispatch.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom" + ] + }, + { + "name": "qwen3_moe_build_batched_decode_expert_dispatch", + "symbol": "qwen3_moe_build_batched_decode_expert_dispatch", + "source": "qwen3_moe/batched_decode_expert_dispatch.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ], + "bindings": [ + "route_ids", + "assignment_ordinals", + "queue_counts", + "queue_descriptors" + ], + "binding_access": [ + "read", + "read", + "read", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "batched_decode_expert_dispatch_linked", + "primary_sources": [ + "qwen3_moe/batched_decode_expert_dispatch.loom" + ], + "library_sources": [ + "qwen3_moe/model_config.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/model_config.loom" + ] + }, + { + "name": "qwen3_moe_build_batched_decode_expert_dispatch_reference", + "symbol": "qwen3_moe_build_batched_decode_expert_dispatch_reference", + "source": "qwen3_moe/batched_decode_expert_dispatch.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ], + "bindings": [ + "route_ids", + "assignment_ordinals", + "queue_counts", + "queue_descriptors" + ], + "binding_access": [ + "read", + "read", + "read", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "batched_decode_expert_dispatch_linked", + "primary_sources": [ + "qwen3_moe/batched_decode_expert_dispatch.loom" + ], + "library_sources": [ + "qwen3_moe/model_config.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/model_config.loom" + ] + }, + { + "name": "qwen3_moe_build_expert_partition_table", + "symbol": "qwen3_moe_build_expert_partition_table", + "source": "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ], + "bindings": [ + "expert_table", + "partition_table" + ], + "binding_access": [ + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_gate_up_swiglu_q4k_linked", + "primary_sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom" + ], + "library_sources": [ + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_build_expert_table", + "symbol": "qwen3_moe_build_expert_table", + "source": "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ], + "bindings": [ + "route_ids", + "expert_table" + ], + "binding_access": [ + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_gate_up_swiglu_q4k_linked", + "primary_sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom" + ], + "library_sources": [ + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_build_expert_table_partition_prefill_512", + "symbol": "qwen3_moe_build_expert_table_partition_prefill_512", + "source": "qwen3_moe/expert_table_partition_fused.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + } + ], + "bindings": [ + "route_ids", + "expert_table", + "partition_table", + "completion_counter" + ], + "binding_access": [ + "read", + "write", + "write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "expert_table_partition_fused_linked", + "primary_sources": [ + "qwen3_moe/expert_table_partition_fused.loom" + ], + "library_sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_dense_linear_q4k_f16_wmma", + "symbol": "qwen3_moe_dense_linear_q4k_f16_wmma", + "source": "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "dense_linear_quantized_f16_wmma_linked", + "primary_sources": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom" + ], + "library_sources": [ + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_dense_linear_q4k_f16_wmma_parameterized", + "symbol": "qwen3_moe_dense_linear_q4k_f16_wmma_parameterized", + "source": "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + }, + { + "name": "output_accumulation", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + }, + { + "name": "output_accumulation", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "dense_linear_quantized_f16_wmma_linked", + "primary_sources": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom" + ], + "library_sources": [ + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_dense_linear_q4k_q8_1_x4", + "symbol": "qwen3_moe_dense_linear_q4k_q8_1_x4", + "source": "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "dense_linear_quantized_f16_wmma_linked", + "primary_sources": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom" + ], + "library_sources": [ + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8", + "symbol": "qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8", + "source": "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "weight", + "output", + "norm_weight", + "normalized_output", + "completion_counter", + "next_q8_output" + ], + "binding_access": [ + "read", + "read", + "read_write", + "read", + "read_write", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "dense_linear_quantized_f16_wmma_linked", + "primary_sources": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom" + ], + "library_sources": [ + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_dense_linear_q6k_f16_wmma", + "symbol": "qwen3_moe_dense_linear_q6k_f16_wmma", + "source": "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "dense_linear_quantized_f16_wmma_linked", + "primary_sources": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom" + ], + "library_sources": [ + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_dense_linear_q6k_f16_wmma_parameterized", + "symbol": "qwen3_moe_dense_linear_q6k_f16_wmma_parameterized", + "source": "qwen3_moe/dense_linear_quantized_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + }, + { + "name": "output_accumulation", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + }, + { + "name": "output_accumulation", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "dense_linear_quantized_f16_wmma_linked", + "primary_sources": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom" + ], + "library_sources": [ + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_flash_attention_decode_f32_f16_wmma", + "symbol": "qwen3_moe_flash_attention_decode_f32_f16_wmma", + "source": "qwen3_moe/flash_attention_decode_f32_f16_wmma.loom", + "workload_parameters": [ + { + "name": "key_value_token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "key_value_token_count", + "type": "index" + } + ], + "bindings": [ + "query", + "key", + "value", + "mask", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "qwen3_moe/flash_attention_decode_f32_f16_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "qwen3_moe_flash_attention_decode_q128_fused_f32_f16_wmma", + "symbol": "qwen3_moe_flash_attention_decode_q128_fused_f32_f16_wmma", + "source": "qwen3_moe/flash_attention_decode_q128_f32_f16_wmma.loom", + "workload_parameters": [ + { + "name": "key_value_token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "key_value_token_count", + "type": "index" + } + ], + "bindings": [ + "query", + "key", + "value", + "mask", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "qwen3_moe/flash_attention_decode_q128_f32_f16_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "qwen3_moe_flash_attention_decode_split_f32_f16_wmma", + "symbol": "qwen3_moe_flash_attention_decode_split_f32_f16_wmma", + "source": "qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom", + "workload_parameters": [ + { + "name": "key_value_token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "key_value_token_count", + "type": "index" + } + ], + "bindings": [ + "query", + "key", + "value", + "mask", + "partial_max", + "partial_sum", + "partial_output", + "completion_counter", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write", + "read_write", + "read_write", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "qwen3_moe_flash_attention_decode_split_f32_f16_wmma_next_q8", + "symbol": "qwen3_moe_flash_attention_decode_split_f32_f16_wmma_next_q8", + "source": "qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom", + "workload_parameters": [ + { + "name": "key_value_token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "key_value_token_count", + "type": "index" + } + ], + "bindings": [ + "query", + "key", + "value", + "mask", + "partial_max", + "partial_sum", + "partial_output", + "completion_counter", + "output", + "next_q8_output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write", + "read_write", + "read_write", + "read_write", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "qwen3_moe_flash_attention_decode_split_mask_513_of_576", + "symbol": "qwen3_moe_flash_attention_decode_split_mask_513_of_576", + "source": "qwen3_moe/flash_attention_decode_split_next_q8_test.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "mask" + ], + "binding_access": [ + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "flash_attention_decode_split_f32_f16_wmma_next_q8_linked", + "primary_sources": [ + "qwen3_moe/flash_attention_decode_split_next_q8_test.loom" + ], + "library_sources": [ + "qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_flash_attention_decode_split_pack_completed_q8_test", + "symbol": "qwen3_moe_flash_attention_decode_split_pack_completed_q8_test", + "source": "qwen3_moe/flash_attention_decode_split_next_q8_test.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "flash_attention_decode_split_f32_f16_wmma_next_q8_linked", + "primary_sources": [ + "qwen3_moe/flash_attention_decode_split_next_q8_test.loom" + ], + "library_sources": [ + "qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_flash_attention_decode_split_produce_partials_f32_f16_wmma", + "symbol": "qwen3_moe_flash_attention_decode_split_produce_partials_f32_f16_wmma", + "source": "qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom", + "workload_parameters": [ + { + "name": "key_value_token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "key_value_token_count", + "type": "index" + } + ], + "bindings": [ + "query", + "key", + "value", + "mask", + "partial_max", + "partial_sum", + "partial_output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "qwen3_moe_flash_attention_decode_split_quantize_reference_4096", + "symbol": "qwen3_moe_flash_attention_decode_split_quantize_reference_4096", + "source": "qwen3_moe/flash_attention_decode_split_next_q8_test.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "flash_attention_decode_split_f32_f16_wmma_next_q8_linked", + "primary_sources": [ + "qwen3_moe/flash_attention_decode_split_next_q8_test.loom" + ], + "library_sources": [ + "qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_flash_attention_decode_split_reduce_f32", + "symbol": "qwen3_moe_flash_attention_decode_split_reduce_f32", + "source": "qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom", + "workload_parameters": [ + { + "name": "key_value_token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "key_value_token_count", + "type": "index" + } + ], + "bindings": [ + "partial_max", + "partial_sum", + "partial_output", + "output" + ], + "binding_access": [ + "read_write", + "read_write", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "qwen3_moe_flash_attention_f32_f16_wmma", + "symbol": "qwen3_moe_flash_attention_f32_f16_wmma", + "source": "qwen3_moe/flash_attention_f32_f16_wmma.loom", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ], + "bindings": [ + "query", + "key", + "value", + "mask", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "qwen3_moe/flash_attention_f32_f16_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "qwen3_moe_flash_attention_test_extract_row", + "symbol": "qwen3_moe_flash_attention_test_extract_row", + "source": "qwen3_moe/flash_attention_f32_f16_wmma.loom", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "context_count", + "type": "index" + }, + { + "name": "source_row", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "context_count", + "type": "index" + }, + { + "name": "source_row", + "type": "index" + } + ], + "bindings": [ + "source_query", + "source_mask", + "target_query", + "target_mask" + ], + "binding_access": [ + "read", + "read", + "read", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "qwen3_moe/flash_attention_f32_f16_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "qwen3_moe_flash_attention_test_make_causal_mask", + "symbol": "qwen3_moe_flash_attention_test_make_causal_mask", + "source": "qwen3_moe/flash_attention_f32_f16_wmma.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "mask" + ], + "binding_access": [ + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "qwen3_moe/flash_attention_f32_f16_wmma.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "qwen3_moe_rmsnorm_f32", + "symbol": "qwen3_moe_rmsnorm_f32", + "source": "qwen3_moe/attention_prepare_quantized.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "attention_prepare_quantized_linked", + "primary_sources": [ + "qwen3_moe/attention_prepare_quantized.loom" + ], + "library_sources": [ + "qwen3_moe/model_config.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/model_config.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_rmsnorm_f32_quantize_q8_1_x4", + "symbol": "qwen3_moe_rmsnorm_f32_quantize_q8_1_x4", + "source": "qwen3_moe/attention_prepare_quantized.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "normalized_output", + "q8_output" + ], + "binding_access": [ + "read", + "read", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "attention_prepare_quantized_linked", + "primary_sources": [ + "qwen3_moe/attention_prepare_quantized.loom" + ], + "library_sources": [ + "qwen3_moe/model_config.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/model_config.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_rmsnorm_quantize_q8_1_x4_wave64_production_check", + "symbol": "qwen3_moe_rmsnorm_quantize_q8_1_x4_wave64_production_check", + "source": "qwen3_moe/routed_down_q6k.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "input", + "norm_weight", + "q8_output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_down_q6k_linked", + "primary_sources": [ + "qwen3_moe/routed_down_q6k.loom" + ], + "library_sources": [ + "ggml/linear_q6k_f32.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom" + ] + }, + "compile_dependencies": [ + "ggml/linear_q6k_f32.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom" + ] + }, + { + "name": "qwen3_moe_routed_down_q4k_f16_wmma_grouped", + "symbol": "qwen3_moe_routed_down_q4k_f16_wmma_grouped", + "source": "qwen3_moe/routed_down_quantized_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "expert_table", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_down_quantized_f16_wmma_linked", + "primary_sources": [ + "qwen3_moe/routed_down_quantized_f16_wmma.loom" + ], + "library_sources": [ + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/routed_down_q6k.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/routed_down_q6k.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_routed_down_q4k_q8_1_x4", + "symbol": "qwen3_moe_routed_down_q4k_q8_1_x4", + "source": "qwen3_moe/routed_down_q4k.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "route_ids", + "route_weights", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_down_q4k_linked", + "primary_sources": [ + "qwen3_moe/routed_down_q4k.loom" + ], + "library_sources": [ + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_routed_down_q4k_q8_1_x4_next_q8", + "symbol": "qwen3_moe_routed_down_q4k_q8_1_x4_next_q8", + "source": "qwen3_moe/routed_down_q4k.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "route_ids", + "route_weights", + "weight", + "output", + "norm_weight", + "completion_counter", + "next_q8_output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write", + "read", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_down_q4k_linked", + "primary_sources": [ + "qwen3_moe/routed_down_q4k.loom" + ], + "library_sources": [ + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_routed_down_q6k_f16_wmma_grouped", + "symbol": "qwen3_moe_routed_down_q6k_f16_wmma_grouped", + "source": "qwen3_moe/routed_down_quantized_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "expert_table", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_down_quantized_f16_wmma_linked", + "primary_sources": [ + "qwen3_moe/routed_down_quantized_f16_wmma.loom" + ], + "library_sources": [ + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/routed_down_q6k.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/routed_down_q6k.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_routed_down_q6k_f32_wave64", + "symbol": "qwen3_moe_routed_down_q6k_f32_wave64", + "source": "qwen3_moe/routed_down_q6k.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "bindings": [ + "input", + "route_ids", + "route_weights", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_down_q6k_linked", + "primary_sources": [ + "qwen3_moe/routed_down_q6k.loom" + ], + "library_sources": [ + "ggml/linear_q6k_f32.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom" + ] + }, + "compile_dependencies": [ + "ggml/linear_q6k_f32.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom" + ] + }, + { + "name": "qwen3_moe_routed_down_q6k_f32_wave64_next_q8", + "symbol": "qwen3_moe_routed_down_q6k_f32_wave64_next_q8", + "source": "qwen3_moe/routed_down_q6k.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "bindings": [ + "input", + "route_ids", + "route_weights", + "weight", + "output", + "norm_weight", + "completion_counter", + "next_q8_output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write", + "read", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_down_q6k_linked", + "primary_sources": [ + "qwen3_moe/routed_down_q6k.loom" + ], + "library_sources": [ + "ggml/linear_q6k_f32.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom" + ] + }, + "compile_dependencies": [ + "ggml/linear_q6k_f32.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom" + ] + }, + { + "name": "qwen3_moe_routed_down_q6k_q8_1_x4", + "symbol": "qwen3_moe_routed_down_q6k_q8_1_x4", + "source": "qwen3_moe/routed_down_q6k.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "route_ids", + "route_weights", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_down_q6k_linked", + "primary_sources": [ + "qwen3_moe/routed_down_q6k.loom" + ], + "library_sources": [ + "ggml/linear_q6k_f32.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom" + ] + }, + "compile_dependencies": [ + "ggml/linear_q6k_f32.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom" + ] + }, + { + "name": "qwen3_moe_routed_down_q6k_q8_1_x4_next_q8", + "symbol": "qwen3_moe_routed_down_q6k_q8_1_x4_next_q8", + "source": "qwen3_moe/routed_down_q6k.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "route_ids", + "route_weights", + "weight", + "output", + "norm_weight", + "completion_counter", + "next_q8_output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write", + "read", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_down_q6k_linked", + "primary_sources": [ + "qwen3_moe/routed_down_q6k.loom" + ], + "library_sources": [ + "ggml/linear_q6k_f32.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom" + ] + }, + "compile_dependencies": [ + "ggml/linear_q6k_f32.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom" + ] + }, + { + "name": "qwen3_moe_routed_down_weighted_reduce_f16_f32", + "symbol": "qwen3_moe_routed_down_weighted_reduce_f16_f32", + "source": "qwen3_moe/routed_down_quantized_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "route_weights", + "routed_output", + "residual_input", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_down_quantized_f16_wmma_linked", + "primary_sources": [ + "qwen3_moe/routed_down_quantized_f16_wmma.loom" + ], + "library_sources": [ + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/routed_down_q6k.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/routed_down_q6k.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32", + "symbol": "qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32", + "source": "qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "route_weights", + "routed_output", + "hidden_state", + "next_norm_weight", + "next_projection_input" + ], + "binding_access": [ + "read", + "read", + "read_write", + "read", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_down_weighted_reduce_next_rmsnorm_f32_linked", + "primary_sources": [ + "qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_f32.loom" + ], + "library_sources": [ + "qwen3_moe/routed_down_quantized_f16_wmma.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/routed_down_q6k.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/routed_down_quantized_f16_wmma.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/routed_down_q6k.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_q8_1_x4", + "symbol": "qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_q8_1_x4", + "source": "qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_q8_1_x4.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "route_weights", + "routed_output", + "hidden_state", + "next_norm_weight", + "next_projection_input" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_down_weighted_reduce_next_rmsnorm_q8_1_x4_linked", + "primary_sources": [ + "qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_q8_1_x4.loom" + ], + "library_sources": [ + "qwen3_moe/routed_down_quantized_f16_wmma.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/routed_down_q6k.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/routed_down_quantized_f16_wmma.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/routed_down_q6k.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma", + "symbol": "qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma", + "source": "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "expert_table", + "partition_table", + "gate_weight", + "up_weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_gate_up_swiglu_q4k_f16_wmma_linked", + "primary_sources": [ + "qwen3_moe/routed_linear_q4k_f16_wmma.loom" + ], + "library_sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_routed_gate_up_swiglu_q4k_q8", + "symbol": "qwen3_moe_routed_gate_up_swiglu_q4k_q8", + "source": "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "route_ids", + "gate_weight", + "up_weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_gate_up_swiglu_q4k_linked", + "primary_sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom" + ], + "library_sources": [ + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8", + "symbol": "qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8", + "source": "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "route_stride", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "route_ids", + "gate_weight", + "up_weight", + "output", + "completion_counters", + "next_q8_output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_gate_up_swiglu_q4k_linked", + "primary_sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom" + ], + "library_sources": [ + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped", + "symbol": "qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped", + "source": "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_count", + "type": "index" + }, + { + "name": "expert_count", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ], + "bindings": [ + "q8_input", + "expert_table", + "gate_weight", + "up_weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_gate_up_swiglu_q4k_linked", + "primary_sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom" + ], + "library_sources": [ + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_routed_linear_q4k_f16_wmma", + "symbol": "qwen3_moe_routed_linear_q4k_f16_wmma", + "source": "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "expert_table", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_gate_up_swiglu_q4k_f16_wmma_linked", + "primary_sources": [ + "qwen3_moe/routed_linear_q4k_f16_wmma.loom" + ], + "library_sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_routed_swiglu_f16", + "symbol": "qwen3_moe_routed_swiglu_f16", + "source": "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "gate", + "up", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_gate_up_swiglu_q4k_f16_wmma_linked", + "primary_sources": [ + "qwen3_moe/routed_linear_q4k_f16_wmma.loom" + ], + "library_sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_routed_swiglu_f32", + "symbol": "qwen3_moe_routed_swiglu_f32", + "source": "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "gate", + "up", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "routed_gate_up_swiglu_q4k_f16_wmma_linked", + "primary_sources": [ + "qwen3_moe/routed_linear_q4k_f16_wmma.loom" + ], + "library_sources": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "qwen3_moe_router_projection_f32_four_row_wave32", + "symbol": "qwen3_moe_router_projection_f32_four_row_wave32", + "source": "qwen3_moe/router_projection_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "router_projection_f32_linked", + "primary_sources": [ + "qwen3_moe/router_projection_f32.loom" + ], + "library_sources": [ + "qwen3_moe/model_config.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/model_config.loom" + ] + }, + { + "name": "qwen3_moe_router_projection_f32_one_row_wave64", + "symbol": "qwen3_moe_router_projection_f32_one_row_wave64", + "source": "qwen3_moe/router_projection_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "router_projection_f32_linked", + "primary_sources": [ + "qwen3_moe/router_projection_f32.loom" + ], + "library_sources": [ + "qwen3_moe/model_config.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/model_config.loom" + ] + }, + { + "name": "qwen3_moe_router_projection_f32_reference", + "symbol": "qwen3_moe_router_projection_f32_reference", + "source": "qwen3_moe/router_projection_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "router_projection_f32_linked", + "primary_sources": [ + "qwen3_moe/router_projection_f32.loom" + ], + "library_sources": [ + "qwen3_moe/model_config.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/model_config.loom" + ] + }, + { + "name": "qwen3_moe_router_projection_top8_fused_decode_f32", + "symbol": "qwen3_moe_router_projection_top8_fused_decode_f32", + "source": "qwen3_moe/router_projection_top8_fused_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + } + ], + "bindings": [ + "input", + "weight", + "logits", + "completion_counter", + "route_ids", + "route_weights" + ], + "binding_access": [ + "read", + "read", + "read_write", + "read_write", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "router_projection_top8_fused_f32_linked", + "primary_sources": [ + "qwen3_moe/router_projection_top8_fused_f32.loom" + ], + "library_sources": [ + "qwen3_moe/model_config.loom", + "qwen3_moe/router_projection_f32.loom", + "qwen3_moe/router_top8_f32.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/model_config.loom", + "qwen3_moe/router_projection_f32.loom", + "qwen3_moe/router_top8_f32.loom" + ] + }, + { + "name": "qwen3_moe_router_top8_f32", + "symbol": "qwen3_moe_router_top8_f32", + "source": "qwen3_moe/router_top8_f32.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "route_id_stride", + "type": "index" + } + ], + "bindings": [ + "logits", + "route_ids", + "route_weights" + ], + "binding_access": [ + "read", + "write", + "write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "router_top8_f32_linked", + "primary_sources": [ + "qwen3_moe/router_top8_f32.loom" + ], + "library_sources": [ + "qwen3_moe/model_config.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/model_config.loom" + ] + }, + { + "name": "qwen3_moe_router_top8_wide_stride_reference", + "symbol": "qwen3_moe_router_top8_wide_stride_reference", + "source": "qwen3_moe/router_top8_f32.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "route_ids" + ], + "binding_access": [ + "read" + ], + "target_selector": "", + "compile_recipe": { + "mode": "archive", + "link_module": "router_top8_f32_linked", + "primary_sources": [ + "qwen3_moe/router_top8_f32.loom" + ], + "library_sources": [ + "qwen3_moe/model_config.loom" + ] + }, + "compile_dependencies": [ + "qwen3_moe/model_config.loom" + ] + }, + { + "name": "qwen_attention_metadata", + "symbol": "qwen_attention_metadata", + "source": "../qwen/attention_metadata.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "context_capacity", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "context_capacity", + "type": "index" + } + ], + "bindings": [ + "control", + "positions", + "key_cache_indices", + "value_cache_indices", + "attention_mask" + ], + "binding_access": [ + "read", + "read_write", + "read_write", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "../qwen/attention_metadata.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "qwen_decode_attention_metadata", + "symbol": "qwen_decode_attention_metadata", + "source": "../qwen/attention_metadata.loom", + "workload_parameters": [], + "launch_parameters": [], + "bindings": [ + "control", + "positions", + "key_cache_indices", + "value_cache_indices" + ], + "binding_access": [ + "read", + "read_write", + "read_write", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "../qwen/attention_metadata.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "qwen_token_embedding_q4k", + "symbol": "qwen_token_embedding_q4k", + "source": "../qwen/token_embedding_q4k.loom", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "vocabulary_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "vocabulary_count", + "type": "index" + } + ], + "bindings": [ + "token_ids", + "weight", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "../qwen/token_embedding_q4k.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + } + ], + "link_modules": [ + { + "name": "router_projection_f32_linked", + "srcs": [ + "qwen3_moe/router_projection_f32.loom" + ], + "libraries": [ + "qwen3_moe/model_config.loom" + ] + }, + { + "name": "router_top8_f32_linked", + "srcs": [ + "qwen3_moe/router_top8_f32.loom" + ], + "libraries": [ + "qwen3_moe/model_config.loom" + ] + }, + { + "name": "router_projection_top8_fused_f32_linked", + "srcs": [ + "qwen3_moe/router_projection_top8_fused_f32.loom" + ], + "libraries": [ + "qwen3_moe/model_config.loom", + "qwen3_moe/router_projection_f32.loom", + "qwen3_moe/router_top8_f32.loom" + ] + }, + { + "name": "attention_postprocess_f32_f16_linked", + "srcs": [ + "qwen3_moe/attention_postprocess_f32_f16.loom" + ], + "libraries": [ + "qwen3_moe/model_config.loom" + ] + }, + { + "name": "attention_prepare_quantized_linked", + "srcs": [ + "qwen3_moe/attention_prepare_quantized.loom" + ], + "libraries": [ + "qwen3_moe/model_config.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "attention_qkv_quantized_linked", + "srcs": [ + "qwen3_moe/attention_qkv_quantized.loom" + ], + "libraries": [ + ":dense_linear_quantized_f16_wmma_linked" + ] + }, + { + "name": "attention_qkv_postprocess_fused_linked", + "srcs": [ + "qwen3_moe/attention_qkv_postprocess_fused.loom" + ], + "libraries": [ + ":attention_postprocess_f32_f16_linked", + ":attention_qkv_quantized_linked" + ] + }, + { + "name": "attention_qkv_same_format_prefill_linked", + "srcs": [ + "qwen3_moe/attention_qkv_same_format_prefill.loom" + ], + "libraries": [ + ":dense_linear_quantized_f16_wmma_linked" + ] + }, + { + "name": "linear_q6k_q8_1_x4_linked", + "srcs": [ + "ggml/linear_q6k_q8_1_x4.loom" + ], + "libraries": [ + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "linear_q6k_f32_linked", + "srcs": [ + "ggml/linear_q6k_f32.loom" + ], + "libraries": [ + ":linear_q6k_q8_1_x4_linked" + ] + }, + { + "name": "routed_down_quantized_f16_wmma_linked", + "srcs": [ + "qwen3_moe/routed_down_quantized_f16_wmma.loom" + ], + "libraries": [ + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/routed_down_q6k.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "routed_down_weighted_reduce_next_rmsnorm_f32_linked", + "srcs": [ + "qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_f32.loom" + ], + "libraries": [ + "qwen3_moe/routed_down_quantized_f16_wmma.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/routed_down_q6k.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "routed_down_weighted_reduce_next_rmsnorm_q8_1_x4_linked", + "srcs": [ + "qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_q8_1_x4.loom" + ], + "libraries": [ + "qwen3_moe/routed_down_quantized_f16_wmma.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_down_q4k.loom", + "qwen3_moe/routed_down_q6k.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "routed_down_q4k_linked", + "srcs": [ + "qwen3_moe/routed_down_q4k.loom" + ], + "libraries": [ + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "routed_down_q6k_linked", + "srcs": [ + "qwen3_moe/routed_down_q6k.loom" + ], + "libraries": [ + "ggml/linear_q6k_f32.loom", + "ggml/linear_q6k_q8_1_x4.loom", + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_down_next_q8.loom" + ] + }, + { + "name": "routed_gate_up_swiglu_q4k_linked", + "srcs": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom" + ], + "libraries": [ + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "batched_decode_expert_dispatch_linked", + "srcs": [ + "qwen3_moe/batched_decode_expert_dispatch.loom" + ], + "libraries": [ + "qwen3_moe/model_config.loom" + ] + }, + { + "name": "batched_decode_gate_up_q4k_linked", + "srcs": [ + "qwen3_moe/batched_decode_gate_up_q4k.loom" + ], + "libraries": [ + "ggml/quantize_q8_1_x4.loom", + "qwen3_moe/batched_decode_expert_dispatch.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom" + ] + }, + { + "name": "expert_table_partition_fused_linked", + "srcs": [ + "qwen3_moe/expert_table_partition_fused.loom" + ], + "libraries": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "routed_gate_up_swiglu_q4k_f16_wmma_linked", + "srcs": [ + "qwen3_moe/routed_linear_q4k_f16_wmma.loom" + ], + "libraries": [ + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "dense_linear_quantized_f16_wmma_linked", + "srcs": [ + "qwen3_moe/dense_linear_quantized_f16_wmma.loom" + ], + "libraries": [ + "ggml/linear_q6k_q8_1_x4.loom", + "qwen3_moe/attention_prepare_quantized.loom", + "qwen3_moe/model_config.loom", + "qwen3_moe/routed_linear_q4k_f16_wmma.loom", + "qwen3_moe/routed_gate_up_swiglu_q4k.loom", + "ggml/quantize_q8_1_x4.loom" + ] + }, + { + "name": "flash_attention_decode_split_f32_f16_wmma_next_q8_linked", + "srcs": [ + "qwen3_moe/flash_attention_decode_split_next_q8_test.loom" + ], + "libraries": [ + "qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom", + "ggml/quantize_q8_1_x4.loom" + ] + } + ], + "plan_cases": [ + { + "name": "router_projection_f32_plan_test", + "args": [ + "$(location :router_projection_f32_linked)", + "--benchmark=@qwen3_moe_router_projection_f32_decode", + "--config=qwen3_moe.model.hidden_size=2048", + "--config=qwen3_moe.router.expert_count=128", + "--config=qwen3_moe.workload.token_capacity=1", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "router_projection_f32_linked" + }, + { + "name": "router_top8_f32_plan_test", + "args": [ + "$(location :router_top8_f32_linked)", + "--benchmark=@qwen3_moe_router_top8_f32_decode", + "--config=qwen3_moe.router.expert_count=128", + "--config=qwen3_moe.router.route_count=8", + "--config=qwen3_moe.workload.token_capacity=1", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "router_top8_f32_linked" + }, + { + "name": "router_projection_top8_fused_f32_plan_test", + "args": [ + "$(location :router_projection_top8_fused_f32_linked)", + "--benchmark=@qwen3_moe_router_projection_top8_fused_decode", + "--config=qwen3_moe.model.hidden_size=2048", + "--config=qwen3_moe.router.expert_count=128", + "--config=qwen3_moe.router.route_count=8", + "--config=qwen3_moe.workload.token_capacity=1", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "router_projection_top8_fused_f32_linked" + }, + { + "name": "attention_postprocess_f32_f16_plan_test", + "args": [ + "$(location :attention_postprocess_f32_f16_linked)", + "--benchmark=@qwen3_moe_attention_postprocess_decode", + "--config=qwen3_moe.attention.head_size=128", + "--config=qwen3_moe.attention.key_value_size=512", + "--config=qwen3_moe.attention.query_size=4096", + "--config=qwen3_moe.model.rms_epsilon=0.000001", + "--config=qwen3_moe.workload.token_capacity=1", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "attention_postprocess_f32_f16_linked" + }, + { + "name": "attention_prepare_quantized_plan_test", + "args": [ + "$(location :attention_prepare_quantized_linked)", + "--config=qwen3_moe.model.hidden_size=2048", + "--config=qwen3_moe.model.rms_epsilon=0.000001", + "--config=ggml.quantize_q8_1_x4.group_capacity=8192", + "--config=qwen3_moe.workload.token_capacity=512", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "attention_prepare_quantized_linked" + }, + { + "name": "attention_qkv_quantized_plan_test", + "args": [ + "$(location :attention_qkv_quantized_linked)", + "--benchmark=@qwen3_moe_attention_qkv_full_q6_decode", + "--config=qwen3_moe.attention.key_value_size=512", + "--config=qwen3_moe.attention.query_size=4096", + "--config=qwen3_moe.attention.value_uses_q6=1", + "--config=qwen3_moe.dense_quantized.input_size=2048", + "--config=qwen3_moe.dense_quantized.output_accumulation=0", + "--config=qwen3_moe.dense_quantized.output_size=512", + "--config=qwen3_moe.model.hidden_size=2048", + "--config=qwen3_moe.model.rms_epsilon=0.000001", + "--config=qwen3_moe.workload.token_capacity=32", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "attention_qkv_quantized_linked" + }, + { + "name": "attention_qkv_postprocess_fused_plan_test", + "args": [ + "$(location :attention_qkv_postprocess_fused_linked)", + "--benchmark=@qwen3_moe_attention_qkv_postprocess_fused_boundary_decode", + "--config=qwen3_moe.attention.head_size=128", + "--config=qwen3_moe.attention.key_value_size=512", + "--config=qwen3_moe.attention.query_size=4096", + "--config=qwen3_moe.attention.value_uses_q6=1", + "--config=qwen3_moe.dense_quantized.input_size=2048", + "--config=qwen3_moe.dense_quantized.output_accumulation=0", + "--config=qwen3_moe.dense_quantized.output_size=512", + "--config=qwen3_moe.model.hidden_size=2048", + "--config=qwen3_moe.model.rms_epsilon=0.000001", + "--config=ggml.quantize_q8_1_x4.group_capacity=16", + "--config=qwen3_moe.workload.token_capacity=1", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "attention_qkv_postprocess_fused_linked" + }, + { + "name": "attention_qkv_same_format_prefill_plan_test", + "args": [ + "$(location :attention_qkv_same_format_prefill_linked)", + "--benchmark=@qwen3_moe_attention_qkv_q4_prefill_512_fused", + "--config=qwen3_moe.attention.key_value_size=512", + "--config=qwen3_moe.attention.query_size=4096", + "--config=qwen3_moe.dense_quantized.output_accumulation=0", + "--config=qwen3_moe.model.hidden_size=2048", + "--config=qwen3_moe.workload.token_capacity=512", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "attention_qkv_same_format_prefill_linked" + }, + { + "name": "quantize_q8_1_x4_plan_test", + "args": [ + "$(location ggml/quantize_q8_1_x4.loom)", + "--config=ggml.quantize_q8_1_x4.group_capacity=8192", + "--dry-run", + "--output-format=jsonl" + ], + "source": "ggml/quantize_q8_1_x4.loom" + }, + { + "name": "routed_gate_up_swiglu_q4k_plan_test", + "args": [ + "$(location :routed_gate_up_swiglu_q4k_linked)", + "--benchmark=@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_prefill_512", + "--config=qwen3_moe.routed_gate_up.expert_count=128", + "--config=qwen3_moe.routed_gate_up.input_size=2048", + "--config=qwen3_moe.routed_gate_up.output_size=768", + "--config=qwen3_moe.routed_gate_up.route_count=8", + "--config=qwen3_moe.workload.token_capacity=512", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "routed_gate_up_swiglu_q4k_linked" + }, + { + "name": "linear_q6k_f32_plan_test", + "args": [ + "$(location :linear_q6k_f32_linked)", + "--benchmark=@ggml_linear_q6k_f32_wave64_dense_v_decode", + "--config=ggml.linear_q6k_f32.output_capacity=512", + "--config=ggml.linear_q6k_f32.token_capacity=2048", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "linear_q6k_f32_linked" + }, + { + "name": "linear_q6k_q8_1_x4_plan_test", + "args": [ + "$(location :linear_q6k_q8_1_x4_linked)", + "--config=ggml.linear_q6k_q8_1_x4.output_capacity=151936", + "--config=ggml.linear_q6k_q8_1_x4.token_capacity=2048", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "linear_q6k_q8_1_x4_linked" + }, + { + "name": "routed_down_weighted_reduce_next_rmsnorm_f32_plan_test", + "args": [ + "$(location :routed_down_weighted_reduce_next_rmsnorm_f32_linked)", + "--benchmark=@qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32_fused_prefill_512", + "--config=qwen3_moe.model.hidden_size=2048", + "--config=qwen3_moe.model.rms_epsilon=0.000001", + "--config=qwen3_moe.routed_down.output_size=2048", + "--config=qwen3_moe.routed_down.route_count=8", + "--config=qwen3_moe.workload.token_capacity=512", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "routed_down_weighted_reduce_next_rmsnorm_f32_linked" + }, + { + "name": "routed_down_weighted_reduce_next_rmsnorm_q8_1_x4_plan_test", + "args": [ + "$(location :routed_down_weighted_reduce_next_rmsnorm_q8_1_x4_linked)", + "--benchmark=@qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_q8_1_x4_fused_prefill_14", + "--config=qwen3_moe.model.hidden_size=2048", + "--config=qwen3_moe.model.rms_epsilon=0.000001", + "--config=qwen3_moe.routed_down.output_size=2048", + "--config=qwen3_moe.routed_down.route_count=8", + "--config=qwen3_moe.workload.token_capacity=512", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "routed_down_weighted_reduce_next_rmsnorm_q8_1_x4_linked" + }, + { + "name": "routed_down_quantized_f16_wmma_plan_test", + "args": [ + "$(location :routed_down_quantized_f16_wmma_linked)", + "--config=qwen3_moe.routed_down.expert_count=128", + "--config=qwen3_moe.routed_down.input_size=768", + "--config=qwen3_moe.routed_down.output_size=2048", + "--config=qwen3_moe.routed_down.route_count=8", + "--config=qwen3_moe.workload.token_capacity=512", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "routed_down_quantized_f16_wmma_linked" + }, + { + "name": "routed_down_q4k_plan_test", + "args": [ + "$(location :routed_down_q4k_linked)", + "--benchmark=@qwen3_moe_routed_down_q4k_q8_1_x4_prefill_512", + "--config=qwen3_moe.routed_down.expert_count=128", + "--config=qwen3_moe.routed_down.input_size=768", + "--config=qwen3_moe.routed_down.output_size=2048", + "--config=qwen3_moe.routed_down.route_count=8", + "--config=qwen3_moe.workload.token_capacity=512", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "routed_down_q4k_linked" + }, + { + "name": "routed_down_q6k_plan_test", + "args": [ + "$(location :routed_down_q6k_linked)", + "--benchmark=@qwen3_moe_routed_down_q6k_f32_wave64_decode", + "--config=qwen3_moe.routed_down.expert_count=128", + "--config=qwen3_moe.routed_down.input_size=768", + "--config=qwen3_moe.routed_down.output_size=2048", + "--config=qwen3_moe.routed_down.route_count=8", + "--config=qwen3_moe.workload.token_capacity=1", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "routed_down_q6k_linked" + }, + { + "name": "batched_decode_gate_up_q4k_plan_test", + "args": [ + "$(location :batched_decode_gate_up_q4k_linked)", + "--benchmark=@qwen3_moe_batched_decode_gate_up_q4k_rows2_benchmark", + "--config=qwen3_moe.batched_decode.rows2_descriptor_capacity=64", + "--config=qwen3_moe.batched_decode.schedule0_row_limit=1", + "--config=qwen3_moe.batched_decode.schedule1_row_limit=2", + "--config=qwen3_moe.batched_decode.schedule2_row_limit=4", + "--config=qwen3_moe.router.expert_count=128", + "--config=qwen3_moe.router.route_count=8", + "--config=qwen3_moe.routed_gate_up.expert_count=128", + "--config=qwen3_moe.routed_gate_up.input_size=2048", + "--config=qwen3_moe.routed_gate_up.output_size=768", + "--config=qwen3_moe.routed_gate_up.route_count=8", + "--config=qwen3_moe.workload.token_capacity=16", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "batched_decode_gate_up_q4k_linked" + }, + { + "name": "batched_decode_expert_dispatch_plan_test", + "args": [ + "$(location :batched_decode_expert_dispatch_linked)", + "--benchmark=@qwen3_moe_batched_decode_expert_dispatch_diverse", + "--config=qwen3_moe.batched_decode.schedule0_row_limit=1", + "--config=qwen3_moe.batched_decode.schedule1_row_limit=2", + "--config=qwen3_moe.batched_decode.schedule2_row_limit=4", + "--config=qwen3_moe.router.expert_count=128", + "--config=qwen3_moe.router.route_count=8", + "--config=qwen3_moe.workload.token_capacity=16", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "batched_decode_expert_dispatch_linked" + }, + { + "name": "batched_decode_expert_dispatch_configurable_plan_test", + "args": [ + "$(location :batched_decode_expert_dispatch_linked)", + "--benchmark=@qwen3_moe_batched_decode_expert_dispatch_configurable", + "--config=qwen3_moe.batched_decode.schedule0_row_limit=1", + "--config=qwen3_moe.batched_decode.schedule1_row_limit=3", + "--config=qwen3_moe.batched_decode.schedule2_row_limit=6", + "--config=qwen3_moe.router.expert_count=32", + "--config=qwen3_moe.router.route_count=4", + "--config=qwen3_moe.workload.token_capacity=8", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "batched_decode_expert_dispatch_linked" + }, + { + "name": "expert_table_partition_fused_plan_test", + "args": [ + "$(location :expert_table_partition_fused_linked)", + "--benchmark=@qwen3_moe_expert_table_partition_fused_prefill_512", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "expert_table_partition_fused_linked" + }, + { + "name": "dense_linear_q6k_f16_wmma_plan_test", + "args": [ + "$(location :dense_linear_quantized_f16_wmma_linked)", + "--benchmark=@qwen3_moe_dense_linear_q6k_f16_wmma_v_prefill_128", + "--config=qwen3_moe.dense_quantized.input_size=2048", + "--config=qwen3_moe.dense_quantized.output_accumulation=0", + "--config=qwen3_moe.dense_quantized.output_size=512", + "--config=qwen3_moe.workload.token_capacity=128", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "dense_linear_quantized_f16_wmma_linked" + }, + { + "name": "dense_linear_q4k_f16_wmma_o_plan_test", + "args": [ + "$(location :dense_linear_quantized_f16_wmma_linked)", + "--benchmark=@qwen3_moe_dense_linear_q4k_f16_wmma_o_prefill_512", + "--config=qwen3_moe.dense_quantized.input_size=4096", + "--config=qwen3_moe.dense_quantized.output_accumulation=1", + "--config=qwen3_moe.dense_quantized.output_size=2048", + "--config=qwen3_moe.workload.token_capacity=512", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "dense_linear_quantized_f16_wmma_linked" + }, + { + "name": "dense_linear_q4k_q8_1_x4_plan_test", + "args": [ + "$(location :dense_linear_quantized_f16_wmma_linked)", + "--benchmark=@qwen3_moe_dense_linear_q4k_q8_1_x4_k_decode", + "--config=qwen3_moe.dense_quantized.input_size=2048", + "--config=qwen3_moe.dense_quantized.output_accumulation=0", + "--config=qwen3_moe.dense_quantized.output_size=512", + "--config=qwen3_moe.workload.token_capacity=1", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "dense_linear_quantized_f16_wmma_linked" + }, + { + "name": "dense_linear_q4k_q8_1_x4_next_q8_plan_test", + "args": [ + "$(location :dense_linear_quantized_f16_wmma_linked)", + "--benchmark=@qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8_decode", + "--config=qwen3_moe.dense_quantized.input_size=4096", + "--config=qwen3_moe.dense_quantized.output_accumulation=1", + "--config=qwen3_moe.dense_quantized.output_size=2048", + "--config=qwen3_moe.model.hidden_size=2048", + "--config=qwen3_moe.model.rms_epsilon=0.000001", + "--config=qwen3_moe.workload.token_capacity=1", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "dense_linear_quantized_f16_wmma_linked" + }, + { + "name": "routed_gate_up_swiglu_q4k_f16_wmma_plan_test", + "args": [ + "$(location :routed_gate_up_swiglu_q4k_f16_wmma_linked)", + "--config=qwen3_moe.routed_gate_up.expert_count=128", + "--config=qwen3_moe.routed_gate_up.input_size=2048", + "--config=qwen3_moe.routed_gate_up.output_size=768", + "--config=qwen3_moe.routed_gate_up.route_count=8", + "--config=qwen3_moe.workload.token_capacity=2048", + "--dry-run", + "--output-format=jsonl" + ], + "link_module": "routed_gate_up_swiglu_q4k_f16_wmma_linked" + }, + { + "name": "flash_attention_decode_f32_f16_wmma_plan_test", + "args": [ + "$(location qwen3_moe/flash_attention_decode_f32_f16_wmma.loom)", + "--config=qwen3_moe.attention.decode.output_partition_count=1", + "--config=qwen3_moe.attention.key_value_head_count=4", + "--config=qwen3_moe.attention.query_head_count=32", + "--dry-run", + "--output-format=jsonl" + ], + "source": "qwen3_moe/flash_attention_decode_f32_f16_wmma.loom" + }, + { + "name": "flash_attention_decode_q128_f32_f16_wmma_plan_test", + "args": [ + "$(location qwen3_moe/flash_attention_decode_q128_f32_f16_wmma.loom)", + "--config=qwen3_moe.attention.key_value_head_count=4", + "--config=qwen3_moe.attention.query_head_count=32", + "--dry-run", + "--output-format=jsonl" + ], + "source": "qwen3_moe/flash_attention_decode_q128_f32_f16_wmma.loom" + }, + { + "name": "flash_attention_decode_split_f32_f16_wmma_plan_test", + "args": [ + "$(location qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom)", + "--config=qwen3_moe.attention.key_value_head_count=4", + "--config=qwen3_moe.attention.key_value_token_capacity=2048", + "--config=qwen3_moe.attention.query_head_count=32", + "--dry-run", + "--output-format=jsonl" + ], + "source": "qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom" + }, + { + "name": "flash_attention_f32_f16_wmma_plan_test", + "args": [ + "$(location qwen3_moe/flash_attention_f32_f16_wmma.loom)", + "--config=qwen3_moe.attention.key_value_head_count=4", + "--config=qwen3_moe.attention.query_head_count=32", + "--config=qwen3_moe.workload.token_capacity=2048", + "--dry-run", + "--output-format=jsonl" + ], + "source": "qwen3_moe/flash_attention_f32_f16_wmma.loom" + }, + { + "name": "owned_token_embedding_decode_plan_test", + "args": [ + "$(location ../qwen/token_embedding_q4k.loom)", + "--benchmark=@qwen_token_embedding_q4k_decode", + "--dry-run", + "--output-format=jsonl", + "--sample-compilation=per_sample" + ], + "source": "../qwen/token_embedding_q4k.loom" + }, + { + "name": "owned_token_embedding_prefill_plan_test", + "args": [ + "$(location ../qwen/token_embedding_q4k.loom)", + "--benchmark=@qwen_token_embedding_q4k_prefill_512", + "--dry-run", + "--output-format=jsonl", + "--sample-compilation=per_sample" + ], + "source": "../qwen/token_embedding_q4k.loom" + }, + { + "name": "owned_attention_metadata_prefill_plan_test", + "args": [ + "$(location ../qwen/attention_metadata.loom)", + "--benchmark=@qwen_attention_metadata_prefill_512", + "--dry-run", + "--output-format=jsonl", + "--sample-compilation=per_sample" + ], + "source": "../qwen/attention_metadata.loom" + } + ] +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/attention_postprocess_f32_f16.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/attention_postprocess_f32_f16.loom new file mode 100644 index 000000000000..ed381348a7a4 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/attention_postprocess_f32_f16.loom @@ -0,0 +1,304 @@ +// Qwen grouped-query postprocessing with direct F16 cache publication. +// +// Raw projection rows are physically `[token][head][channel]`: the reshape +// operations in the reference graph only expose that existing layout. Query +// and key rows receive their independent per-head RMSNorm and NEOX rotary +// transform. Queries remain F32 for attention, while keys and values are +// converted directly into their indexed F16 cache rows. No reshape, +// transpose, normalized-row, or rotated-key allocation is materialized. +// +// Logical positions and cache rows are separate runtime inputs. This preserves +// continuous batching, where a token's rotary position need not equal the +// physical cache row selected by the allocator. K and V indices also remain +// distinct so the kernel does not impose an ordering contract on the cache. +// The stage scheduler owns those indices and guarantees they select valid cache +// rows; violating that trusted contract is an error rather than a skipped write. +template.decl @qwen3_moe.attention.postprocess.head_body(%publish_output: i1, %token_count: index, %cache_row_count: index, %head_domain0: index, %token0: index, %positions: buffer, %key_cache_indices: buffer, %value_cache_indices: buffer, %query_input: buffer, %key_input: buffer, %value_input: buffer, %query_norm_weight: buffer, %key_norm_weight: buffer, %inverse_frequencies: buffer, %query_output: buffer, %key_cache: buffer, %value_cache: buffer) + +amdgpu.target @qwen3_moe_attention_postprocess_gfx11_wave32 {subgroup_size = 32} + +config.decl @qwen3_moe.model.rms_epsilon : f32 + +config.decl @qwen3_moe.attention.head_size : %value: index where [range(%value, 4, 1024), mul(%value, 4)] + +config.decl @qwen3_moe.attention.query_size : %value: index where [range(%value, 1, 262144)] + +config.decl @qwen3_moe.attention.key_value_size : %value: index where [range(%value, 1, 262144)] + +config.decl @qwen3_moe.attention.rope_mscale : f32 + +// Applies one RMSNorm and two adjacent NEOX pairs. NEOX pairs corresponding +// channels from the low and high halves, while keeping each half contiguous. +// A two-pair packet therefore gives both halves naturally coalesced loads and +// stores without changing the model's pairing rule. +func.def inline @qwen3_moe_rmsnorm_neox_packet(%row_sum: f32, %head_size: f32, %epsilon: f32, %position: f32, %rope_mscale: f32, %inverse_frequencies: vector<2xf32>, %low_values: vector<2xf32>, %high_values: vector<2xf32>, %low_weights: vector<2xf32>, %high_weights: vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) { + %mean = scalar.divf %row_sum, %head_size : f32 + %biased_mean = scalar.addf %mean, %epsilon : f32 + %scale = scalar.rsqrtf %biased_mean : f32 + %scale_vector = vector.splat %scale : vector<2xf32> + %position_vector = vector.splat %position : vector<2xf32> + %inverse_two_pi = scalar.constant 0.15915494309189535 : f32 + %inverse_two_pi_vector = vector.splat %inverse_two_pi : vector<2xf32> + %normalized_low = vector.mulf %low_values, %scale_vector : vector<2xf32> + %normalized_high = vector.mulf %high_values, %scale_vector : vector<2xf32> + %scaled_low = vector.mulf %normalized_low, %low_weights : vector<2xf32> + %scaled_high = vector.mulf %normalized_high, %high_weights : vector<2xf32> + %angles = vector.mulf %position_vector, %inverse_frequencies : vector<2xf32> + %turns = vector.mulf %angles, %inverse_two_pi_vector : vector<2xf32> + %cosines = vector.costurnsf %turns : vector<2xf32> + %sines = vector.sinturnsf %turns : vector<2xf32> + %low_cosines = vector.mulf %scaled_low, %cosines : vector<2xf32> + %high_sines = vector.mulf %scaled_high, %sines : vector<2xf32> + %low_sines = vector.mulf %scaled_low, %sines : vector<2xf32> + %high_cosines = vector.mulf %scaled_high, %cosines : vector<2xf32> + %rotated_low0 = vector.subf %low_cosines, %high_sines : vector<2xf32> + %rotated_high0 = vector.addf %low_sines, %high_cosines : vector<2xf32> + %rope_mscale_vector = vector.splat %rope_mscale : vector<2xf32> + %rotated_low = vector.mulf %rotated_low0, %rope_mscale_vector : vector<2xf32> + %rotated_high = vector.mulf %rotated_high0, %rope_mscale_vector : vector<2xf32> + func.return %rotated_low, %rotated_high : vector<2xf32>, vector<2xf32> +} + +// Processes one head domain after its raw projection row is visible. Callers +// may use a larger workgroup than the packet count; inactive workitems +// contribute zero to normalization and perform no loads or stores. +template.def<@qwen3_moe.attention.postprocess.head_body> device @qwen3_moe_attention_postprocess_head_body(%publish_output: i1, %token_count: index, %cache_row_count: index, %head_domain0: index, %token0: index, %positions: buffer, %key_cache_indices: buffer, %value_cache_indices: buffer, %query_input: buffer, %key_input: buffer, %value_input: buffer, %query_norm_weight: buffer, %key_norm_weight: buffer, %inverse_frequencies: buffer, %query_output: buffer, %key_cache: buffer, %value_cache: buffer) { + %query_size0 = config.get @qwen3_moe.attention.query_size : index + %key_value_size0 = config.get @qwen3_moe.attention.key_value_size : index + %head_size0 = config.get @qwen3_moe.attention.head_size : index + %rope_mscale = config.get @qwen3_moe.attention.rope_mscale : f32 + %epsilon = config.get @qwen3_moe.model.rms_epsilon : f32 + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_cache_row_count = index.assume %cache_row_count [range(%cache_row_count, 1, 1048576)] : index + %query_size, %key_value_size, %head_size = index.assume %query_size0, %key_value_size0, %head_size0 [range(%query_size0, 1, 262144), range(%key_value_size0, 1, 262144), range(%head_size0, 4, 1024), mul(%head_size0, 4), mul(%query_size0, %head_size0), mul(%key_value_size0, %head_size0)] : index, index, index + %channel0 = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c0_f32 = scalar.constant 0.0 : f32 + %c0_offset = index.constant 0 : offset + %token, %launch_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %half_head_size = index.div %head_size, %c2 : index + %pair_packet_count = index.div %head_size, %c4 : index + %query_head_count = index.div %query_size, %head_size : index + %key_value_head_count = index.div %key_value_size, %head_size : index + %key_value_domain_count = index.mul %key_value_head_count, %c2 : index + %head_domain_count = index.add %query_head_count, %key_value_domain_count : index + %head_domain = index.assume %head_domain0 [lt(%head_domain0, %head_domain_count)] : index + %key_domain_end = index.add %query_head_count, %key_value_head_count : index + %is_query = index.cmp ult, %head_domain, %query_head_count : index + %is_query_or_key = index.cmp ult, %head_domain, %key_domain_end : index + %key_value_head = index.rem %head_domain, %key_value_head_count : index + %active_channel = index.cmp ult, %channel0, %pair_packet_count : index + %head_size_i32 = index.cast %head_size : index to i32 + %head_size_f32 = scalar.sitofp %head_size_i32 : i32 to f32 + %positions_noalias, %key_cache_indices_noalias, %value_cache_indices_noalias, %query_input_noalias, %key_input_noalias, %value_input_noalias, %query_norm_weight_noalias, %key_norm_weight_noalias, %inverse_frequencies_noalias, %query_output_noalias, %key_cache_noalias, %value_cache_noalias = buffer.assume.noalias %positions, %key_cache_indices, %value_cache_indices, %query_input, %key_input, %value_input, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %query_output, %key_cache, %value_cache : buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer + %positions_view = buffer.view %positions_noalias[%c0_offset] : buffer -> view<[%launch_token_count]xi32> + %key_cache_indices_view = buffer.view %key_cache_indices_noalias[%c0_offset] : buffer -> view<[%launch_token_count]xi64> + %value_cache_indices_view = buffer.view %value_cache_indices_noalias[%c0_offset] : buffer -> view<[%launch_token_count]xi64> + %query_input_view = buffer.view %query_input_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%query_head_count]x[%head_size]xf32> + %key_input_view = buffer.view %key_input_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%key_value_head_count]x[%head_size]xf32> + %value_input_view = buffer.view %value_input_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%key_value_head_count]x[%head_size]xf32> + %query_norm_weight_view = buffer.view %query_norm_weight_noalias[%c0_offset] : buffer -> view<[%head_size]xf32> + %key_norm_weight_view = buffer.view %key_norm_weight_noalias[%c0_offset] : buffer -> view<[%head_size]xf32> + %inverse_frequencies_view = buffer.view %inverse_frequencies_noalias[%c0_offset] : buffer -> view<[%half_head_size]xf32> + %query_output_view = buffer.view %query_output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%query_head_count]x[%head_size]xf32> + %key_cache_words_view = buffer.view %key_cache_noalias[%c0_offset] : buffer -> view<[%bounded_cache_row_count]x[%key_value_head_count]x[%half_head_size]xi32> + %value_cache_words_view = buffer.view %value_cache_noalias[%c0_offset] : buffer -> view<[%bounded_cache_row_count]x[%key_value_head_count]x[%half_head_size]xi32> + %row_sum = scf.if %is_query_or_key -> (f32) { + %partial_sum = scf.if %publish_output -> (f32) { + %channel_sum = scf.if %active_channel -> (f32) { + %channel = index.assume %channel0 [lt(%channel0, %pair_packet_count)] : index + %reduction_channel = index.mul %channel, %c4 : index + %reduction_values = scf.if %is_query -> (vector<4xf32>) { + %query_values = vector.load %query_input_view[%token, %head_domain, %reduction_channel] : view<[%launch_token_count]x[%query_head_count]x[%head_size]xf32> -> vector<4xf32> + scf.yield %query_values : vector<4xf32> + } else { + %key_values = vector.load %key_input_view[%token, %key_value_head, %reduction_channel] : view<[%launch_token_count]x[%key_value_head_count]x[%head_size]xf32> -> vector<4xf32> + scf.yield %key_values : vector<4xf32> + } + %squares = vector.mulf %reduction_values, %reduction_values : vector<4xf32> + %sum = vector.reduce %squares, %c0_f32 : vector<4xf32>, f32 + scf.yield %sum : f32 + } else { + scf.yield %c0_f32 : f32 + } + scf.yield %channel_sum : f32 + } else { + scf.yield %c0_f32 : f32 + } + %sum = kernel.workgroup.reduce %partial_sum : f32 + scf.yield %sum : f32 + } else { + scf.yield %c0_f32 : f32 + } + scf.if %publish_output { + scf.if %active_channel { + %channel = index.assume %channel0 [lt(%channel0, %pair_packet_count)] : index + %pair_channel = index.mul %channel, %c2 : index + %paired_channel = index.add %pair_channel, %half_head_size : index + %paired_word = index.add %channel, %pair_packet_count : index + %reduction_channel = index.mul %channel, %c4 : index + scf.if %is_query { + %low_values = vector.load %query_input_view[%token, %head_domain, %pair_channel] : view<[%launch_token_count]x[%query_head_count]x[%head_size]xf32> -> vector<2xf32> + %high_values = vector.load %query_input_view[%token, %head_domain, %paired_channel] : view<[%launch_token_count]x[%query_head_count]x[%head_size]xf32> -> vector<2xf32> + %low_weights = vector.load %query_norm_weight_view[%pair_channel] : view<[%head_size]xf32> -> vector<2xf32> + %high_weights = vector.load %query_norm_weight_view[%paired_channel] : view<[%head_size]xf32> -> vector<2xf32> + %inverse_frequencies_packet = vector.load %inverse_frequencies_view[%pair_channel] : view<[%half_head_size]xf32> -> vector<2xf32> + %position_i32 = view.load %positions_view[%token] : view<[%launch_token_count]xi32> -> i32 + %position = scalar.sitofp %position_i32 : i32 to f32 + %rotated_low, %rotated_high = func.call @qwen3_moe_rmsnorm_neox_packet(%row_sum, %head_size_f32, %epsilon, %position, %rope_mscale, %inverse_frequencies_packet, %low_values, %high_values, %low_weights, %high_weights) : (f32, f32, f32, f32, f32, vector<2xf32>, vector<2xf32>, vector<2xf32>, vector<2xf32>, vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) + vector.store %rotated_low, %query_output_view[%token, %head_domain, %pair_channel] : vector<2xf32>, view<[%launch_token_count]x[%query_head_count]x[%head_size]xf32> + vector.store %rotated_high, %query_output_view[%token, %head_domain, %paired_channel] : vector<2xf32>, view<[%launch_token_count]x[%query_head_count]x[%head_size]xf32> + } else { + scf.if %is_query_or_key { + %low_values = vector.load %key_input_view[%token, %key_value_head, %pair_channel] : view<[%launch_token_count]x[%key_value_head_count]x[%head_size]xf32> -> vector<2xf32> + %high_values = vector.load %key_input_view[%token, %key_value_head, %paired_channel] : view<[%launch_token_count]x[%key_value_head_count]x[%head_size]xf32> -> vector<2xf32> + %low_weights = vector.load %key_norm_weight_view[%pair_channel] : view<[%head_size]xf32> -> vector<2xf32> + %high_weights = vector.load %key_norm_weight_view[%paired_channel] : view<[%head_size]xf32> -> vector<2xf32> + %inverse_frequencies_packet = vector.load %inverse_frequencies_view[%pair_channel] : view<[%half_head_size]xf32> -> vector<2xf32> + %position_i32 = view.load %positions_view[%token] : view<[%launch_token_count]xi32> -> i32 + %position = scalar.sitofp %position_i32 : i32 to f32 + %cache_index_raw = view.load %key_cache_indices_view[%token] : view<[%launch_token_count]xi64> -> i64 + %cache_index_i64 = scalar.assume %cache_index_raw [range(%cache_index_raw, 0, 1048575)] : i64 + %cache_index0 = index.cast %cache_index_i64 : i64 to index + %cache_index = index.assume %cache_index0 [lt(%cache_index0, %bounded_cache_row_count)] : index + %rotated_low, %rotated_high = func.call @qwen3_moe_rmsnorm_neox_packet(%row_sum, %head_size_f32, %epsilon, %position, %rope_mscale, %inverse_frequencies_packet, %low_values, %high_values, %low_weights, %high_weights) : (f32, f32, f32, f32, f32, vector<2xf32>, vector<2xf32>, vector<2xf32>, vector<2xf32>, vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) + %half_low = vector.fptrunc %rotated_low : vector<2xf32> to vector<2xf16> + %half_high = vector.fptrunc %rotated_high : vector<2xf32> to vector<2xf16> + %packed_low = vector.bitcast %half_low : vector<2xf16> to vector<1xi32> + %packed_high = vector.bitcast %half_high : vector<2xf16> to vector<1xi32> + vector.store %packed_low, %key_cache_words_view[%cache_index, %key_value_head, %channel] : vector<1xi32>, view<[%bounded_cache_row_count]x[%key_value_head_count]x[%half_head_size]xi32> + vector.store %packed_high, %key_cache_words_view[%cache_index, %key_value_head, %paired_word] : vector<1xi32>, view<[%bounded_cache_row_count]x[%key_value_head_count]x[%half_head_size]xi32> + } else { + %cache_index_raw = view.load %value_cache_indices_view[%token] : view<[%launch_token_count]xi64> -> i64 + %cache_index_i64 = scalar.assume %cache_index_raw [range(%cache_index_raw, 0, 1048575)] : i64 + %cache_index0 = index.cast %cache_index_i64 : i64 to index + %cache_index = index.assume %cache_index0 [lt(%cache_index0, %bounded_cache_row_count)] : index + %values = vector.load %value_input_view[%token, %key_value_head, %reduction_channel] : view<[%launch_token_count]x[%key_value_head_count]x[%head_size]xf32> -> vector<4xf32> + %half_values = vector.fptrunc %values : vector<4xf32> to vector<4xf16> + %packed_values = vector.bitcast %half_values : vector<4xf16> to vector<2xi32> + vector.store %packed_values, %value_cache_words_view[%cache_index, %key_value_head, %pair_channel] : vector<2xi32>, view<[%bounded_cache_row_count]x[%key_value_head_count]x[%half_head_size]xi32> + } + } + } + } + template.return +} + +kernel.def target(@qwen3_moe_attention_postprocess_gfx11_wave32) @qwen3_moe_attention_postprocess_f32_f16(%token_count: index, %cache_row_count: index) { + %query_size = config.get @qwen3_moe.attention.query_size : index + %key_value_size = config.get @qwen3_moe.attention.key_value_size : index + %head_size = config.get @qwen3_moe.attention.head_size : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %pair_packet_count = index.div %head_size, %c4 : index + %query_head_count = index.div %query_size, %head_size : index + %key_value_head_count = index.div %key_value_size, %head_size : index + %key_value_domain_count = index.mul %key_value_head_count, %c2 : index + %head_domain_count = index.add %query_head_count, %key_value_domain_count : index + kernel.launch.config workgroups(%head_domain_count, %token_count, %c1) workgroup_size(%pair_packet_count, %c1, %c1) : index +} launch(%token_count: index, %cache_row_count: index, %positions: buffer, %key_cache_indices: buffer, %value_cache_indices: buffer, %query_input: buffer, %key_input: buffer, %value_input: buffer, %query_norm_weight: buffer, %key_norm_weight: buffer, %inverse_frequencies: buffer, %query_output: buffer, %key_cache: buffer, %value_cache: buffer) where [range(%token_count, 1, 2048)] { + %head_domain = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %c0 = index.constant 0 : index + %valid_token = index.cmp ult, %token0, %token_count : index + %safe_token = scf.select %valid_token, %token0, %c0 : index + template.apply<@qwen3_moe.attention.postprocess.head_body>(%valid_token, %token_count, %cache_row_count, %head_domain, %safe_token, %positions, %key_cache_indices, %value_cache_indices, %query_input, %key_input, %value_input, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %query_output, %key_cache, %value_cache) : (i1, index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +// At position zero NEOX is the identity. Uniform rows make the RMSNorm result +// analytically one to within epsilon, while the distinct V value proves that +// all three output domains reach their intended bindings. +check.case public @qwen3_moe_attention_postprocess_identity_case { + %token_count = check.literal value(2) : index + %cache_row_count = check.literal value(2) : index + %positions = check.generate.fill value(0) : tensor<2xi32> + %key_cache_indices = check.generate.iota offset(0) step(1) : tensor<2xi64> + %value_cache_indices = check.generate.iota offset(0) step(1) : tensor<2xi64> + %query_input = check.generate.fill value(1.0) : tensor<2x1x4xf32> + %key_input = check.generate.fill value(1.0) : tensor<2x1x4xf32> + %value_input = check.generate.fill value(2.0) : tensor<2x1x4xf32> + %query_norm_weight = check.generate.fill value(1.0) : tensor<4xf32> + %key_norm_weight = check.generate.fill value(1.0) : tensor<4xf32> + %inverse_frequencies = check.generate.fill value(1.0) : tensor<2xf32> + %query_output = check.generate.fill value(0.0) : tensor<2x1x4xf32> + %key_cache = check.generate.fill value(0.0) : tensor<2x1x4xf16> + %value_cache = check.generate.fill value(0.0) : tensor<2x1x4xf16> + %expected_query = check.generate.fill value(1.0) : tensor<2x1x4xf32> + %expected_key = check.generate.fill value(1.0) : tensor<2x1x4xf16> + %expected_value = check.generate.fill value(2.0) : tensor<2x1x4xf16> + kernel.launch @qwen3_moe_attention_postprocess_f32_f16[%token_count, %cache_row_count](%token_count, %cache_row_count, %positions, %key_cache_indices, %value_cache_indices, %query_input, %key_input, %value_input, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %query_output, %key_cache, %value_cache) : [index, index](index, index, tensor<2xi32>, tensor<2xi64>, tensor<2xi64>, tensor<2x1x4xf32>, tensor<2x1x4xf32>, tensor<2x1x4xf32>, tensor<4xf32>, tensor<4xf32>, tensor<2xf32>, tensor<2x1x4xf32>, tensor<2x1x4xf16>, tensor<2x1x4xf16>) + check.expect.close actual(%query_output) expected(%expected_query) atol(1.0000000000000001e-05) rtol(1.0000000000000001e-05) nan(same) : tensor<2x1x4xf32> + check.expect.close actual(%key_cache) expected(%expected_key) atol(0.001) rtol(0.001) nan(same) : tensor<2x1x4xf16> + check.expect.close actual(%value_cache) expected(%expected_value) atol(0.0) rtol(0.0) nan(same) : tensor<2x1x4xf16> + check.return +} + +// Distinct angles make the two NEOX pairs rotate `[1, 3]` and `[2, 4]` +// into `[sqrt(5), sqrt(5)]` and `[sqrt(10), sqrt(10)]`. Their flattened +// `[A, B, A, B]` result catches adjacent-channel pairing and packet-order +// mistakes without embedding a second implementation in the test. +check.case public @qwen3_moe_attention_postprocess_neox_case { + %token_count = check.literal value(2) : index + %cache_row_count = check.literal value(2) : index + %positions = check.generate.fill value(1) : tensor<2xi32> + %key_cache_indices = check.generate.iota offset(0) step(1) : tensor<2xi64> + %value_cache_indices = check.generate.iota offset(0) step(1) : tensor<2xi64> + %query_input = check.generate.fill value(1.0) : tensor<2x1x4xf32> + %key_input = check.generate.fill value(1.0) : tensor<2x1x4xf32> + %value_input = check.generate.fill value(0.0) : tensor<2x1x4xf32> + %query_norm_weight = check.generate.iota offset(1.0) step(1.0) : tensor<4xf32> + %key_norm_weight = check.generate.iota offset(1.0) step(1.0) : tensor<4xf32> + %inverse_frequencies = check.generate.iota offset(-0.46364760900080609) step(0.14189705460416391) : tensor<2xf32> + %query_output = check.generate.fill value(0.0) : tensor<2x1x4xf32> + %key_cache = check.generate.fill value(0.0) : tensor<2x1x4xf16> + %value_cache = check.generate.fill value(0.0) : tensor<2x1x4xf16> + %expected_query = check.generate.iota offset(2.2360668182373047) step(0.92620921134948736) period(2) : tensor<2x1x4xf32> + %expected_key = check.generate.iota offset(2.236328125) step(0.92578125) period(2) : tensor<2x1x4xf16> + %expected_value = check.generate.fill value(0.0) : tensor<2x1x4xf16> + kernel.launch @qwen3_moe_attention_postprocess_f32_f16[%token_count, %cache_row_count](%token_count, %cache_row_count, %positions, %key_cache_indices, %value_cache_indices, %query_input, %key_input, %value_input, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %query_output, %key_cache, %value_cache) : [index, index](index, index, tensor<2xi32>, tensor<2xi64>, tensor<2xi64>, tensor<2x1x4xf32>, tensor<2x1x4xf32>, tensor<2x1x4xf32>, tensor<4xf32>, tensor<4xf32>, tensor<2xf32>, tensor<2x1x4xf32>, tensor<2x1x4xf16>, tensor<2x1x4xf16>) + check.expect.close actual(%query_output) expected(%expected_query) atol(0.0001) rtol(0.0001) nan(same) : tensor<2x1x4xf32> + check.expect.close actual(%key_cache) expected(%expected_key) atol(0.002) rtol(0.002) nan(same) : tensor<2x1x4xf16> + check.expect.close actual(%value_cache) expected(%expected_value) atol(0.0) rtol(0.0) nan(same) : tensor<2x1x4xf16> + check.return +} + +check.case public @qwen3_moe_attention_postprocess_benchmark_case { + %token_count = check.param.choice values([1, 32, 128, 512]) name("token_count") : index + %cache_row_count = check.literal value(4096) : index + %positions = check.generate.iota offset(0) step(1) : tensor<[%token_count]xi32> + %key_cache_indices = check.generate.iota offset(0) step(1) : tensor<[%token_count]xi64> + %value_cache_indices = check.generate.iota offset(0) step(1) : tensor<[%token_count]xi64> + %query_input = check.generate.fill value(0.0) : tensor<[%token_count]x32x128xf32> + %key_input = check.generate.fill value(0.0) : tensor<[%token_count]x4x128xf32> + %value_input = check.generate.fill value(0.0) : tensor<[%token_count]x4x128xf32> + %query_norm_weight = check.generate.fill value(1.0) : tensor<128xf32> + %key_norm_weight = check.generate.fill value(1.0) : tensor<128xf32> + %inverse_frequencies = check.generate.fill value(1.0) : tensor<64xf32> + %query_output = check.generate.fill value(1.0) : tensor<[%token_count]x32x128xf32> + %key_cache = check.generate.fill value(0.0) : tensor<4096x4x128xf16> + %value_cache = check.generate.fill value(0.0) : tensor<4096x4x128xf16> + %expected_query = check.generate.fill value(0.0) : tensor<[%token_count]x32x128xf32> + %expected_key_cache = check.generate.fill value(0.0) : tensor<4096x4x128xf16> + %expected_value_cache = check.generate.fill value(0.0) : tensor<4096x4x128xf16> + kernel.launch @qwen3_moe_attention_postprocess_f32_f16[%token_count, %cache_row_count](%token_count, %cache_row_count, %positions, %key_cache_indices, %value_cache_indices, %query_input, %key_input, %value_input, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %query_output, %key_cache, %value_cache) : [index, index](index, index, tensor<[%token_count]xi32>, tensor<[%token_count]xi64>, tensor<[%token_count]xi64>, tensor<[%token_count]x32x128xf32>, tensor<[%token_count]x4x128xf32>, tensor<[%token_count]x4x128xf32>, tensor<128xf32>, tensor<128xf32>, tensor<64xf32>, tensor<[%token_count]x32x128xf32>, tensor<4096x4x128xf16>, tensor<4096x4x128xf16>) + check.expect.close actual(%query_output) expected(%expected_query) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x32x128xf32> + check.expect.close actual(%key_cache) expected(%expected_key_cache) atol(0.0) rtol(0.0) nan(same) : tensor<4096x4x128xf16> + check.expect.close actual(%value_cache) expected(%expected_value_cache) atol(0.0) rtol(0.0) nan(same) : tensor<4096x4x128xf16> + check.return +} + +check.benchmark<@qwen3_moe_attention_postprocess_identity_case> @qwen3_moe_attention_postprocess_identity + +check.benchmark<@qwen3_moe_attention_postprocess_neox_case> @qwen3_moe_attention_postprocess_neox + +check.benchmark<@qwen3_moe_attention_postprocess_benchmark_case> @qwen3_moe_attention_postprocess_decode {token_count = 1} + +check.benchmark<@qwen3_moe_attention_postprocess_benchmark_case> @qwen3_moe_attention_postprocess_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_attention_postprocess_benchmark_case> @qwen3_moe_attention_postprocess_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_attention_postprocess_benchmark_case> @qwen3_moe_attention_postprocess_prefill_512 {token_count = 512} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/attention_prepare_quantized.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/attention_prepare_quantized.loom new file mode 100644 index 000000000000..84a966e0f15a --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/attention_prepare_quantized.loom @@ -0,0 +1,434 @@ +// Qwen attention preparation for raw GGUF quantized projections. +// +// Decode contractions consume a shared Q8_1 x4 activation row. Producing that +// row directly from the residual stream removes the materialized F32 attention +// RMSNorm tensor and avoids rereading it solely for quantization. The fused +// producer preserves GGML's physical Q8_1 x4 contract: +// +// struct block_q8_1_x4 { +// f16 ds[4][2]; +// i32 qs[4][8]; +// }; +// +// One 256-workitem workgroup owns one token. It first reduces the complete +// hidden row to one reciprocal RMS scale, then visits 1024-element stripes. +// Every workitem packs one four-value word per stripe; adjacent eight-lane +// cohorts reduce the maximum and quantized sum for one logical Q8_1 block. +// Scratch is fixed at 1152 bytes regardless of hidden size. +template.decl @qwen3_moe.rmsnorm_quantize_q8_1_x4.body(%publish_normalized: i1, %reduction_subgroup_count0: index, %token_count: index, %token0: index, %input: buffer, %weight: buffer, %normalized_output: buffer, %q8_output: buffer) + +amdgpu.target @qwen3_moe_attention_prepare_gfx11_wave32 {subgroup_size = 32} + +config.decl @qwen3_moe.model.hidden_size : %value: index where [range(%value, 128, 32768), mul(%value, 128)] + +config.decl @qwen3_moe.model.rms_epsilon : f32 + +// Shared Q8_1 x4 packer used by the standalone differential path. +kernel.decl @ggml_quantize_q8_1_x4_f32(%token_count$10: index, %input_size$11: index) launch(%token_count$12: index, %input_size$13: index, %input: buffer, %output: buffer) + +// Materializes the ordinary Qwen RMSNorm boundary. This remains useful for +// prefill schedules that reuse one normalized row across many output tiles; +// decode uses the fused Q8_1 producer below. +kernel.def target(@qwen3_moe_attention_prepare_gfx11_wave32) @qwen3_moe_rmsnorm_f32(%token_count: index) { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%token_count, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %hidden_size0 = config.get @qwen3_moe.model.hidden_size : index + %epsilon = config.get @qwen3_moe.model.rms_epsilon : f32 + %hidden_size = index.assume %hidden_size0 [range(%hidden_size0, 128, 32768), mul(%hidden_size0, 128)] : index + %token0 = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c8 = index.constant 8 : index + %c256 = index.constant 256 : index + %scratch_bytes = index.constant 1024 : offset + %c0_offset = index.constant 0 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %valid_token = index.cmp ult, %token0, %token_count : index + %safe_token0 = scf.select %valid_token, %token0, %c0 : index + %token, %launch_token_count = index.assume %safe_token0, %token_count [lt(%safe_token0, %token_count)] : index, index + %hidden_size_i32 = index.cast %hidden_size : index to i32 + %hidden_size_f32 = scalar.sitofp %hidden_size_i32 : i32 to f32 + %input_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input, %weight, %output : buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%hidden_size]xf32> + %weight_view = buffer.view %weight_noalias[%c0_offset] : buffer -> view<[%hidden_size]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%hidden_size]xf32> + %thread_sum = scf.for %channel = [%workitem to %hidden_size step %c256](%running_sum = %c0_f32 : f32) -> (f32) { + %value = view.load %input_view[%token, %channel] : view<[%launch_token_count]x[%hidden_size]xf32> -> f32 + %square = scalar.mulf %value, %value : f32 + %next_sum = scalar.addf %running_sum, %square : f32 + scf.yield %next_sum : f32 + } + %subgroup_sum = kernel.subgroup.reduce %thread_sum : f32 + %scratch = buffer.alloca align(16) %scratch_bytes : buffer + %scratch_view = buffer.view %scratch[%c0_offset] : buffer -> view<256xf32> + %is_subgroup_leader = index.cmp eq, %lane, %c0 : index + scf.if %is_subgroup_leader { + view.store %subgroup_sum, %scratch_view[%subgroup] : f32, view<256xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %is_reduction_subgroup = index.cmp eq, %subgroup, %c0 : index + %is_reduction_lane = index.cmp ult, %lane, %c8 : index + %loads_subgroup_sum = scalar.andi %is_reduction_subgroup, %is_reduction_lane : i1 + %subgroup_partial = scf.if %loads_subgroup_sum -> (f32) { + %value = view.load %scratch_view[%lane] : view<256xf32> -> f32 + scf.yield %value : f32 + } else { + scf.yield %c0_f32 : f32 + } + %row_sum = kernel.subgroup.reduce %subgroup_partial : f32 + %writes_scale = scalar.andi %is_reduction_subgroup, %is_subgroup_leader : i1 + scf.if %writes_scale { + %mean = scalar.divf %row_sum, %hidden_size_f32 : f32 + %biased_mean = scalar.addf %mean, %epsilon : f32 + %scale = scalar.rsqrtf %biased_mean : f32 + view.store %scale, %scratch_view[%c0] : f32, view<256xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %scale = view.load %scratch_view[%c0] : view<256xf32> -> f32 + scf.if %valid_token { + scf.for %channel = [%workitem to %hidden_size step %c256] { + %value = view.load %input_view[%token, %channel] : view<[%launch_token_count]x[%hidden_size]xf32> -> f32 + %learned_weight = view.load %weight_view[%channel] : view<[%hidden_size]xf32> -> f32 + %normalized = scalar.mulf %value, %scale : f32 + %result = scalar.mulf %normalized, %learned_weight : f32 + view.store %result, %output_view[%token, %channel] : f32, view<[%launch_token_count]x[%hidden_size]xf32> + } + } + kernel.return +} + +// Shared RMSNorm and GGML Q8_1 x4 row producer. The caller owns launch geometry +// and supplies the logical token whose complete hidden row this workgroup owns. +// The attention export discards the ordinary F32 row, while feed-forward +// publishes it for the router without rereading and requantizing the normalized +// values. +template.def<@qwen3_moe.rmsnorm_quantize_q8_1_x4.body> device @qwen3_moe_rmsnorm_quantize_q8_1_x4_body(%publish_normalized: i1, %reduction_subgroup_count0: index, %token_count: index, %token0: index, %input: buffer, %weight: buffer, %normalized_output: buffer, %q8_output: buffer) { + %hidden_size0 = config.get @qwen3_moe.model.hidden_size : index + %epsilon = config.get @qwen3_moe.model.rms_epsilon : f32 + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %reduction_subgroup_count = index.assume %reduction_subgroup_count0 [range(%reduction_subgroup_count0, 1, 8)] : index + %hidden_size = index.assume %hidden_size0 [range(%hidden_size0, 128, 32768), mul(%hidden_size0, 128)] : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %c1024 = index.constant 1024 : index + %group_bytes = index.constant 144 : offset + %payload_byte_add = index.constant 16 : offset + %scratch_d_byte_add = index.constant 1024 : offset + %scratch_bytes = index.constant 1152 : offset + %c0_offset = index.constant 0 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %c1_f32 = scalar.constant 1.0 : f32 + %c127 = scalar.constant 127.0 : f32 + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %valid_token = index.cmp ult, %token0, %bounded_token_count : index + %safe_token0 = scf.select %valid_token, %token0, %c0 : index + %token, %launch_token_count = index.assume %safe_token0, %bounded_token_count [lt(%safe_token0, %bounded_token_count)] : index, index + %hidden_size_i32 = index.cast %hidden_size : index to i32 + %hidden_size_f32 = scalar.sitofp %hidden_size_i32 : i32 to f32 + %physical_group_count = index.div %hidden_size, %c128 : index + %row_bytes = index.scale %physical_group_count, %group_bytes : index, offset -> offset + %token_output_byte_base = index.scale %token, %row_bytes : index, offset -> offset + %input_noalias, %weight_noalias, %q8_output_noalias = buffer.assume.noalias %input, %weight, %q8_output : buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%hidden_size]xf32> + %weight_view = buffer.view %weight_noalias[%c0_offset] : buffer -> view<[%hidden_size]xf32> + %normalized_output_view = buffer.view %normalized_output[%c0_offset] : buffer -> view<[%launch_token_count]x[%hidden_size]xf32> + %scratch = buffer.alloca align(16) %scratch_bytes : buffer + %scratch_values = buffer.view %scratch[%c0_offset] : buffer -> view<256xf32> + %scratch_d = buffer.view %scratch[%scratch_d_byte_add] : buffer -> view<32xf32> + // Reduce the complete row before any block-local quantization. + %thread_sum = scf.for %channel = [%workitem to %hidden_size step %c256](%running_sum = %c0_f32 : f32) -> (f32) { + %value = view.load %input_view[%token, %channel] : view<[%launch_token_count]x[%hidden_size]xf32> -> f32 + %square = scalar.mulf %value, %value : f32 + %next_sum = scalar.addf %running_sum, %square : f32 + scf.yield %next_sum : f32 + } + %subgroup_sum = kernel.subgroup.reduce %thread_sum : f32 + %is_subgroup_leader = index.cmp eq, %lane, %c0 : index + scf.if %is_subgroup_leader { + view.store %subgroup_sum, %scratch_values[%subgroup] : f32, view<256xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %is_reduction_subgroup = index.cmp eq, %subgroup, %c0 : index + %is_reduction_lane = index.cmp ult, %lane, %reduction_subgroup_count : index + %loads_subgroup_sum = scalar.andi %is_reduction_subgroup, %is_reduction_lane : i1 + %subgroup_partial = scf.if %loads_subgroup_sum -> (f32) { + %value = view.load %scratch_values[%lane] : view<256xf32> -> f32 + scf.yield %value : f32 + } else { + scf.yield %c0_f32 : f32 + } + %row_sum = kernel.subgroup.reduce %subgroup_partial : f32 + %writes_scale = scalar.andi %is_reduction_subgroup, %is_subgroup_leader : i1 + scf.if %writes_scale { + %mean = scalar.divf %row_sum, %hidden_size_f32 : f32 + %biased_mean = scalar.addf %mean, %epsilon : f32 + %scale = scalar.rsqrtf %biased_mean : f32 + view.store %scale, %scratch_values[%c0] : f32, view<256xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %row_scale = view.load %scratch_values[%c0] : view<256xf32> -> f32 + %row_scale_vector = vector.splat %row_scale : vector<4xf32> + %publishes_normalized = scalar.andi %publish_normalized, %valid_token : i1 + // Reuse one scratch frame for each 1024-element stripe. Hidden sizes need + // only be divisible by 128; inactive workitems in the final stripe carry + // zeros and never publish. + scf.for %stripe_base = [%c0 to %hidden_size step %c1024] { + %word_element_add = index.mul %workitem, %c4 : index + %channel = index.add %stripe_base, %word_element_add : index + %valid_word = index.cmp ult, %channel, %hidden_size : index + %mask = vector.mask.range [%channel to %hidden_size step %c1] : index -> vector<4xi1> + %input_values = vector.load.mask %input_view[%token, %channel], %mask, %c0_f32x4 : view<[%launch_token_count]x[%hidden_size]xf32>, vector<4xi1>, vector<4xf32> + %learned_weights = vector.load.mask %weight_view[%channel], %mask, %c0_f32x4 : view<[%hidden_size]xf32>, vector<4xi1>, vector<4xf32> + %normalized0 = vector.mulf %input_values, %row_scale_vector : vector<4xf32> + %normalized = vector.mulf %normalized0, %learned_weights : vector<4xf32> + scf.if %publishes_normalized { + vector.store.mask %normalized, %normalized_output_view[%token, %channel], %mask : vector<4xf32>, view<[%launch_token_count]x[%hidden_size]xf32>, vector<4xi1> + } + %absolute_values = vector.absf %normalized : vector<4xf32> + %thread_max = vector.reduce %absolute_values, %c0_f32 : vector<4xf32>, f32 + view.store %thread_max, %scratch_values[%workitem] : f32, view<256xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + %word_in_block = index.rem %workitem, %c8 : index + %block_in_stripe = index.div %workitem, %c8 : index + %is_block_leader = index.cmp eq, %word_in_block, %c0 : index + %writes_block_d = scalar.andi %valid_word, %is_block_leader : i1 + scf.if %writes_block_d { + %cohort_base = index.mul %block_in_stripe, %c8 : index + %cohort_maxima = vector.load %scratch_values[%cohort_base] : view<256xf32> -> vector<8xf32> + %amax = vector.reduce %cohort_maxima, %c0_f32 : vector<8xf32>, f32 + %d = scalar.divf %amax, %c127 : f32 + view.store %d, %scratch_d[%block_in_stripe] : f32, view<32xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %d = scf.if %valid_word -> (f32) { + %block_d = view.load %scratch_d[%block_in_stripe] : view<32xf32> -> f32 + scf.yield %block_d : f32 + } else { + scf.yield %c0_f32 : f32 + } + %d_nonzero = scalar.cmpf one, %d, %c0_f32 : f32 + %d_inverse = scf.if %d_nonzero -> (f32) { + %inverse = scalar.divf %c1_f32, %d : f32 + scf.yield %inverse : f32 + } else { + scf.yield %c0_f32 : f32 + } + %d_inverse_vector = vector.splat %d_inverse : vector<4xf32> + %scaled_values = vector.mulf %normalized, %d_inverse_vector : vector<4xf32> + %rounded_values = vector.roundf %scaled_values : vector<4xf32> + %quantized_values = vector.fptosi %rounded_values : vector<4xf32> to vector<4xi8> + %packed_word = vector.bitcast %quantized_values : vector<4xi8> to vector<1xi32> + %publishes_q8_word = scalar.andi %valid_word, %valid_token : i1 + scf.if %publishes_q8_word { + %q8_block = index.div %channel, %c32 : index + %physical_group = index.div %q8_block, %c4 : index + %block_in_group = index.rem %q8_block, %c4 : index + %group_byte_add = index.scale %physical_group, %group_bytes : index, offset -> offset + %group_byte_offset = index.add %token_output_byte_base, %group_byte_add : offset + %payload_byte_offset = index.add %group_byte_offset, %payload_byte_add : offset + %group_ds = buffer.view %q8_output_noalias[%group_byte_offset] : buffer -> view<8xf16> + %group_qs = buffer.view %q8_output_noalias[%payload_byte_offset] : buffer -> view<32xi32> + %block_word_base = index.mul %block_in_group, %c8 : index + %packed_word_index0 = index.add %block_word_base, %word_in_block : index + %packed_word_index = index.assume %packed_word_index0 [range(%packed_word_index0, 0, 31)] : index + vector.store %packed_word, %group_qs[%packed_word_index] : vector<1xi32>, view<32xi32> + } + %thread_quantized_sum = vector.reduce %rounded_values, %c0_f32 : vector<4xf32>, f32 + view.store %thread_quantized_sum, %scratch_values[%workitem] : f32, view<256xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + %publishes_block_ds = scalar.andi %writes_block_d, %valid_token : i1 + scf.if %publishes_block_ds { + %cohort_base = index.mul %block_in_stripe, %c8 : index + %cohort_sums = vector.load %scratch_values[%cohort_base] : view<256xf32> -> vector<8xf32> + %quantized_sum = vector.reduce %cohort_sums, %c0_f32 : vector<8xf32>, f32 + %s = scalar.mulf %quantized_sum, %d : f32 + %q8_block = index.div %channel, %c32 : index + %physical_group = index.div %q8_block, %c4 : index + %block_in_group = index.rem %q8_block, %c4 : index + %group_byte_add = index.scale %physical_group, %group_bytes : index, offset -> offset + %group_byte_offset = index.add %token_output_byte_base, %group_byte_add : offset + %group_ds = buffer.view %q8_output_noalias[%group_byte_offset] : buffer -> view<8xf16> + %d_f16 = scalar.fptrunc %d : f32 to f16 + %s_f16 = scalar.fptrunc %s : f32 to f16 + %ds_index = index.mul %block_in_group, %c2 : index + %s_index = index.add %ds_index, %c1 : index + view.store %d_f16, %group_ds[%ds_index] : f16, view<8xf16> + view.store %s_f16, %group_ds[%s_index] : f16, view<8xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + } + template.return +} + +// Fuses attention RMSNorm directly into the Q8_1 x4 row consumed by decode +// projections. The output has hidden_size / 128 physical groups of 144 bytes. +kernel.def target(@qwen3_moe_attention_prepare_gfx11_wave32) @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4(%token_count: index) { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%token_count, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %publish_normalized = scalar.constant false : i1 + %reduction_subgroup_count = index.constant 8 : index + %token = kernel.workgroup.id : index + %input_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input, %weight, %output : buffer, buffer, buffer + template.apply<@qwen3_moe.rmsnorm_quantize_q8_1_x4.body>(%publish_normalized, %reduction_subgroup_count, %token_count, %token, %input_noalias, %weight_noalias, %output_noalias, %output_noalias) : (i1, index, index, index, buffer, buffer, buffer, buffer) + kernel.return +} + +// Publishes both the ordinary F32 RMSNorm row required by routing and the Q8_1 +// x4 row consumed by direct decode contractions. +kernel.def target(@qwen3_moe_attention_prepare_gfx11_wave32) @qwen3_moe_rmsnorm_f32_quantize_q8_1_x4(%token_count: index) { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%token_count, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %normalized_output: buffer, %q8_output: buffer) where [range(%token_count, 1, 2048)] { + %publish_normalized = scalar.constant true : i1 + %reduction_subgroup_count = index.constant 8 : index + %token = kernel.workgroup.id : index + %input_noalias, %weight_noalias, %normalized_output_noalias, %q8_output_noalias = buffer.assume.noalias %input, %weight, %normalized_output, %q8_output : buffer, buffer, buffer, buffer + template.apply<@qwen3_moe.rmsnorm_quantize_q8_1_x4.body>(%publish_normalized, %reduction_subgroup_count, %token_count, %token, %input_noalias, %weight_noalias, %normalized_output_noalias, %q8_output_noalias) : (i1, index, index, index, buffer, buffer, buffer, buffer) + kernel.return +} + +// A uniform residual makes RMSNorm equal to the learned weight up to the +// configured epsilon. The nonuniform weights cross two physical Q8_1 groups, +// exercise signed values, and compare every packed metadata and payload byte +// against the ordinary RMSNorm-plus-packer path. +check.case public @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4_differential_case { + %token_count = check.literal value(1) : index + %hidden_size = check.literal value(256) : index + %input = check.generate.fill value(2.0) : tensor<256xf32> + %weight = check.generate.iota offset(-1.0) step(0.0078125) : tensor<256xf32> + %normalized = check.generate.fill value(0.0) : tensor<256xf32> + %expected = check.generate.fill value(0) : tensor<288xi8> + %actual = check.generate.fill value(1) : tensor<288xi8> + %dual_normalized = check.generate.fill value(1.0) : tensor<256xf32> + %dual_q8 = check.generate.fill value(1) : tensor<288xi8> + kernel.launch @qwen3_moe_rmsnorm_f32[%token_count](%token_count, %input, %weight, %normalized) : [index](index, tensor<256xf32>, tensor<256xf32>, tensor<256xf32>) + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %hidden_size](%token_count, %hidden_size, %normalized, %expected) : [index, index](index, index, tensor<256xf32>, tensor<288xi8>) + kernel.launch @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4[%token_count](%token_count, %input, %weight, %actual) : [index](index, tensor<256xf32>, tensor<256xf32>, tensor<288xi8>) + kernel.launch @qwen3_moe_rmsnorm_f32_quantize_q8_1_x4[%token_count](%token_count, %input, %weight, %dual_normalized, %dual_q8) : [index](index, tensor<256xf32>, tensor<256xf32>, tensor<256xf32>, tensor<288xi8>) + check.expect.equal actual(%actual) expected(%expected) : tensor<288xi8> + check.expect.close actual(%dual_normalized) expected(%normalized) atol(0.0) rtol(0.0) nan(same) : tensor<256xf32> + check.expect.equal actual(%dual_q8) expected(%expected) : tensor<288xi8> + check.return +} + +// Fourteen distinct production-width rows cross two complete 1024-element +// scratch stripes. Comparing every packed byte with the unfused path locks +// both stripe reuse and row-byte addressing across independent workgroups. +check.case public @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4_multistripe_case { + %token_count = check.literal value(14) : index + %hidden_size = check.literal value(2048) : index + %input_seed = check.param.seed base(5858425849397790032) count(1) : i64 + %input = check.generate.random.uniform seed(%input_seed) range(-1.0 to 1.0) : tensor<14x2048xf32> + %weight = check.generate.iota offset(-1.0) step(0.0009765625) : tensor<2048xf32> + %normalized = check.generate.fill value(0.0) : tensor<14x2048xf32> + %expected = check.generate.fill value(0) : tensor<14x2304xi8> + %actual = check.generate.fill value(1) : tensor<14x2304xi8> + %dual_normalized = check.generate.fill value(1.0) : tensor<14x2048xf32> + %dual_q8 = check.generate.fill value(1) : tensor<14x2304xi8> + kernel.launch @qwen3_moe_rmsnorm_f32[%token_count](%token_count, %input, %weight, %normalized) : [index](index, tensor<14x2048xf32>, tensor<2048xf32>, tensor<14x2048xf32>) + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %hidden_size](%token_count, %hidden_size, %normalized, %expected) : [index, index](index, index, tensor<14x2048xf32>, tensor<14x2304xi8>) + kernel.launch @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4[%token_count](%token_count, %input, %weight, %actual) : [index](index, tensor<14x2048xf32>, tensor<2048xf32>, tensor<14x2304xi8>) + kernel.launch @qwen3_moe_rmsnorm_f32_quantize_q8_1_x4[%token_count](%token_count, %input, %weight, %dual_normalized, %dual_q8) : [index](index, tensor<14x2048xf32>, tensor<2048xf32>, tensor<14x2048xf32>, tensor<14x2304xi8>) + check.expect.equal actual(%actual) expected(%expected) : tensor<14x2304xi8> + check.expect.close actual(%dual_normalized) expected(%normalized) atol(0.0) rtol(0.0) nan(same) : tensor<14x2048xf32> + check.expect.equal actual(%dual_q8) expected(%expected) : tensor<14x2304xi8> + check.return +} + +check.case public @qwen3_moe_attention_rmsnorm_f32_benchmark_case { + %token_count = check.param.choice values([1, 8, 32, 128, 512]) name("token_count") : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %weight = check.generate.fill value(1.0) : tensor<2048xf32> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + kernel.launch @qwen3_moe_rmsnorm_f32[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<[%token_count]x2048xf32>, tensor<2048xf32>, tensor<[%token_count]x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x2048xf32> + check.return +} + +check.case public @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4_benchmark_case { + %token_count = check.param.choice values([1, 8, 32, 128, 512]) name("token_count") : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %weight = check.generate.fill value(1.0) : tensor<2048xf32> + %output = check.generate.fill value(1) : tensor<[%token_count]x2304xi8> + %expected = check.generate.fill value(0) : tensor<[%token_count]x2304xi8> + kernel.launch @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<[%token_count]x2048xf32>, tensor<2048xf32>, tensor<[%token_count]x2304xi8>) + check.expect.equal actual(%output) expected(%expected) : tensor<[%token_count]x2304xi8> + check.return +} + +check.case public @qwen3_moe_rmsnorm_f32_quantize_q8_1_x4_benchmark_case { + %token_count = check.param.choice values([1, 8, 32, 128, 512]) name("token_count") : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %weight = check.generate.fill value(1.0) : tensor<2048xf32> + %normalized_output = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %q8_output = check.generate.fill value(1) : tensor<[%token_count]x2304xi8> + %expected_normalized = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %expected_q8 = check.generate.fill value(0) : tensor<[%token_count]x2304xi8> + kernel.launch @qwen3_moe_rmsnorm_f32_quantize_q8_1_x4[%token_count](%token_count, %input, %weight, %normalized_output, %q8_output) : [index](index, tensor<[%token_count]x2048xf32>, tensor<2048xf32>, tensor<[%token_count]x2048xf32>, tensor<[%token_count]x2304xi8>) + check.expect.close actual(%normalized_output) expected(%expected_normalized) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x2048xf32> + check.expect.equal actual(%q8_output) expected(%expected_q8) : tensor<[%token_count]x2304xi8> + check.return +} + +check.case public @qwen3_moe_attention_rmsnorm_then_quantize_q8_1_x4_benchmark_case { + %token_count = check.param.choice values([1, 8, 32, 128, 512]) name("token_count") : index + %hidden_size = check.literal value(2048) : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %weight = check.generate.fill value(1.0) : tensor<2048xf32> + %normalized = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %output = check.generate.fill value(1) : tensor<[%token_count]x2304xi8> + %expected = check.generate.fill value(0) : tensor<[%token_count]x2304xi8> + kernel.launch @qwen3_moe_rmsnorm_f32[%token_count](%token_count, %input, %weight, %normalized) : [index](index, tensor<[%token_count]x2048xf32>, tensor<2048xf32>, tensor<[%token_count]x2048xf32>) + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %hidden_size](%token_count, %hidden_size, %normalized, %output) : [index, index](index, index, tensor<[%token_count]x2048xf32>, tensor<[%token_count]x2304xi8>) + check.expect.equal actual(%output) expected(%expected) : tensor<[%token_count]x2304xi8> + check.return +} + +check.benchmark<@qwen3_moe_attention_rmsnorm_quantize_q8_1_x4_differential_case> @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4_differential + +check.benchmark<@qwen3_moe_attention_rmsnorm_f32_benchmark_case> @qwen3_moe_attention_rmsnorm_f32_decode {token_count = 1} + +check.benchmark<@qwen3_moe_attention_rmsnorm_f32_benchmark_case> @qwen3_moe_attention_rmsnorm_f32_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_attention_rmsnorm_f32_benchmark_case> @qwen3_moe_attention_rmsnorm_f32_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_attention_rmsnorm_f32_benchmark_case> @qwen3_moe_attention_rmsnorm_f32_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_attention_rmsnorm_quantize_q8_1_x4_benchmark_case> @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4_decode {token_count = 1} + +check.benchmark<@qwen3_moe_rmsnorm_f32_quantize_q8_1_x4_benchmark_case> @qwen3_moe_rmsnorm_f32_quantize_q8_1_x4_decode {token_count = 1} + +check.benchmark<@qwen3_moe_attention_rmsnorm_quantize_q8_1_x4_benchmark_case> @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_attention_rmsnorm_quantize_q8_1_x4_benchmark_case> @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_attention_rmsnorm_quantize_q8_1_x4_benchmark_case> @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_attention_rmsnorm_then_quantize_q8_1_x4_benchmark_case> @qwen3_moe_attention_rmsnorm_then_quantize_q8_1_x4_decode {token_count = 1} + +check.benchmark<@qwen3_moe_attention_rmsnorm_then_quantize_q8_1_x4_benchmark_case> @qwen3_moe_attention_rmsnorm_then_quantize_q8_1_x4_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_attention_rmsnorm_then_quantize_q8_1_x4_benchmark_case> @qwen3_moe_attention_rmsnorm_then_quantize_q8_1_x4_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_attention_rmsnorm_then_quantize_q8_1_x4_benchmark_case> @qwen3_moe_attention_rmsnorm_then_quantize_q8_1_x4_prefill_512 {token_count = 512} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/attention_qkv_postprocess_fused.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/attention_qkv_postprocess_fused.loom new file mode 100644 index 000000000000..ffb2415b5dcd --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/attention_qkv_postprocess_fused.loom @@ -0,0 +1,282 @@ +// Decode-only Qwen Q/K/V projection with last-arrival head postprocessing. +// +// The accepted quantized projection publishes eight adjacent raw F32 rows per +// 256-workitem workgroup. Head width is fixed at 128 for this model, so sixteen +// workgroups publish each Q, K, or V head. One device-scope completion counter +// per head forms a release sequence over those stores. The last arrival +// acquires the complete raw row and invokes the canonical per-head +// normalization, NEOX rotary, and cache-publication body. +// +// Raw Q/K/V rows remain explicit because they are the inter-workgroup +// communication surface. The fusion removes the command-buffer boundary and +// lets the postprocess consume freshly published rows; it does not impose a +// hidden cache-index or position relationship. Counters return to zero only +// after semantic outputs are visible, preserving reusable command buffers. +template.decl @qwen3_moe.attention.postprocess.head_body(%arg0: i1, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: buffer, %arg6: buffer, %arg7: buffer, %arg8: buffer, %arg9: buffer, %arg10: buffer, %arg11: buffer, %arg12: buffer, %arg13: buffer, %arg14: buffer, %arg15: buffer, %arg16: buffer) + +template.decl @qwen3_moe.attention.qkv_postprocess_fused_decode.body(%value_uses_q6_index: index, %token_count: index, %cache_row_count: index, %q8_input: buffer, %query_weight: buffer, %key_weight: buffer, %value_weight: buffer, %positions: buffer, %key_cache_indices: buffer, %value_cache_indices: buffer, %query_output_raw: buffer, %key_output_raw: buffer, %value_output_raw: buffer, %query_norm_weight: buffer, %key_norm_weight: buffer, %inverse_frequencies: buffer, %query_output: buffer, %key_cache: buffer, %value_cache: buffer, %completion_counters: buffer) + +template.decl @qwen3_moe.attention.qkv_postprocess_fused_decode.launch() -> (index, index, index) + +template.decl @qwen3_moe.attention.qkv_quantized.body(%arg0: index, %arg1: i1, %arg2: index, %arg3: index, %arg4: buffer, %arg5: buffer, %arg6: buffer, %arg7: buffer, %arg8: buffer, %arg9: buffer, %arg10: buffer) + +amdgpu.target @qwen3_moe_attention_qkv_postprocess_gfx11_wave32 {subgroup_size = 32} + +config.decl @qwen3_moe.attention.head_size : %value: index where [range(%value, 4, 1024), mul(%value, 4)] + +config.decl @qwen3_moe.attention.query_size : %value: index where [range(%value, 1, 262144)] + +config.decl @qwen3_moe.attention.key_value_size : %value: index where [range(%value, 1, 262144)] + +config.decl @qwen3_moe.attention.value_uses_q6 : %value: index where [range(%value, 0, 1)] + +// Reference entry points used by differential and benchmark cases. +kernel.decl @ggml_quantize_q8_1_x4_f32(%token_count$55: index, %input_size$56: index) launch(%token_count$57: index, %input_size$58: index, %input: buffer, %output: buffer) + +kernel.decl @qwen3_moe_attention_qkv_quantized(%token_count$61: index) launch(%token_count$62: index, %q8_input: buffer, %query_weight: buffer, %key_weight: buffer, %value_weight: buffer, %query_output: buffer, %key_output: buffer, %value_output: buffer) + +kernel.decl @qwen3_moe_attention_postprocess_f32_f16(%token_count$70: index, %cache_row_count$71: index) launch(%token_count$72: index, %cache_row_count$73: index, %positions: buffer, %key_cache_indices: buffer, %value_cache_indices: buffer, %query_input: buffer, %key_input: buffer, %value_input: buffer, %query_norm_weight: buffer, %key_norm_weight: buffer, %inverse_frequencies: buffer, %query_output: buffer, %key_cache: buffer, %value_cache: buffer) + +// Storage-specific exports pass a constant value format into this shared body, +// while the compatibility export continues to source it from model config. +template.def<@qwen3_moe.attention.qkv_postprocess_fused_decode.body> device @qwen3_moe_attention_qkv_postprocess_fused_decode_body(%value_uses_q6_index: index, %token_count: index, %cache_row_count: index, %q8_input: buffer, %query_weight: buffer, %key_weight: buffer, %value_weight: buffer, %positions: buffer, %key_cache_indices: buffer, %value_cache_indices: buffer, %query_output_raw: buffer, %key_output_raw: buffer, %value_output_raw: buffer, %query_norm_weight: buffer, %key_norm_weight: buffer, %inverse_frequencies: buffer, %query_output: buffer, %key_cache: buffer, %value_cache: buffer, %completion_counters: buffer) { + %publish_projection_output = scalar.constant true : i1 + %projection_token = kernel.workgroup.id : index + template.apply<@qwen3_moe.attention.qkv_quantized.body>(%value_uses_q6_index, %publish_projection_output, %token_count, %projection_token, %q8_input, %query_weight, %key_weight, %value_weight, %query_output_raw, %key_output_raw, %value_output_raw) : (index, i1, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + %query_size0 = config.get @qwen3_moe.attention.query_size : index + %key_value_size0 = config.get @qwen3_moe.attention.key_value_size : index + %head_size0 = config.get @qwen3_moe.attention.head_size : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 1)] : index + %query_size, %key_value_size, %head_size = index.assume %query_size0, %key_value_size0, %head_size0 [range(%query_size0, 128, 262144), range(%key_value_size0, 128, 262144), range(%head_size0, 128, 128), mul(%query_size0, %head_size0), mul(%key_value_size0, %head_size0)] : index, index, index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c8 = index.constant 8 : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c0_offset = index.constant 0 : offset + %counter_scratch_bytes = index.constant 4 : offset + %channel_tile = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %token = index.assume %c0 [lt(%c0, %bounded_token_count)] : index + %head_tile_count = index.div %head_size, %c8 : index + %head_domain0 = index.div %channel_tile, %head_tile_count : index + %query_head_count = index.div %query_size, %head_size : index + %key_value_head_count = index.div %key_value_size, %head_size : index + %key_value_domain_count = index.mul %key_value_head_count, %c2 : index + %head_domain_count = index.add %query_head_count, %key_value_domain_count : index + %head_domain, %completion_counter_count = index.assume %head_domain0, %head_domain_count [lt(%head_domain0, %head_domain_count)] : index, index + %is_arrival_workitem = index.cmp eq, %workitem, %c0 : index + %completion_counters_aligned = buffer.assume.alignment %completion_counters {minimum_alignment = 16} : buffer + %completion_counters_view = buffer.view %completion_counters_aligned[%c0_offset] : buffer -> view<[%completion_counter_count]xi32> + %counter_scratch = buffer.alloca align(4) %counter_scratch_bytes : buffer + %counter_scratch_view = buffer.view %counter_scratch[%c0_offset] : buffer -> view<1xi32> + // Publish every producer's projection stores before the leader advances one + // workgroup arrival. The last arrival then acquires the complete head. + kernel.barrier scope(workgroup) ordering(release) + scf.if %is_arrival_workitem { + %old_counter = view.atomic.rmw %c1_i32, %completion_counters_view[%head_domain] {ordering = acq_rel, scope = device} : i32, view<[%completion_counter_count]xi32> -> i32 + view.store %old_counter, %counter_scratch_view[%c0] : i32, view<1xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %old_counter = view.load %counter_scratch_view[%c0] : view<1xi32> -> i32 + %head_tile_count_i32 = index.cast %head_tile_count : index to i32 + %last_head_tile_i32 = scalar.subi %head_tile_count_i32, %c1_i32 : i32 + %negative_head_tile_count_i32 = scalar.subi %c0_i32, %head_tile_count_i32 : i32 + %is_last_head_tile = scalar.cmpi eq, %old_counter, %last_head_tile_i32 : i32 + scf.if %is_last_head_tile { + kernel.barrier scope(workgroup) ordering(acquire) + %publish_postprocess_output = scalar.constant true : i1 + template.apply<@qwen3_moe.attention.postprocess.head_body>(%publish_postprocess_output, %bounded_token_count, %cache_row_count, %head_domain, %token, %positions, %key_cache_indices, %value_cache_indices, %query_output_raw, %key_output_raw, %value_output_raw, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %query_output, %key_cache, %value_cache) : (i1, index, index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + // The counter becomes reusable only after all semantic output stores. + kernel.barrier scope(workgroup) ordering(release) + scf.if %is_arrival_workitem { + view.atomic.reduce %negative_head_tile_count_i32, %completion_counters_view[%head_domain] {ordering = release, scope = device} : i32, view<[%completion_counter_count]xi32> + } + } + template.return +} + +template.def<@qwen3_moe.attention.qkv_postprocess_fused_decode.launch> @qwen3_moe_attention_qkv_postprocess_fused_decode_launch() -> (index, index, index) { + %query_size = config.get @qwen3_moe.attention.query_size : index + %key_value_size = config.get @qwen3_moe.attention.key_value_size : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c256 = index.constant 256 : index + %key_value_output_size = index.mul %key_value_size, %c2 : index + %output_size = index.add %query_size, %key_value_output_size : index + %padded_output_size = index.add %output_size, %c7 : index + %output_tiles = index.div %padded_output_size, %c8 : index + template.return %output_tiles, %c1, %c256 : index, index, index +} + +kernel.def target(@qwen3_moe_attention_qkv_postprocess_gfx11_wave32) @qwen3_moe_attention_qkv_postprocess_fused_decode(%token_count: index, %cache_row_count: index) { + %query_size = config.get @qwen3_moe.attention.query_size : index + %c1 = index.constant 1 : index + %key_value_size = config.get @qwen3_moe.attention.key_value_size : index + %c2 = index.constant 2 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c256 = index.constant 256 : index + %key_value_output_size = index.mul %key_value_size, %c2 : index + %output_size = index.add %query_size, %key_value_output_size : index + %padded_output_size = index.add %output_size, %c7 : index + %output_tiles = index.div %padded_output_size, %c8 : index + kernel.launch.config workgroups(%output_tiles, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %cache_row_count: index, %q8_input: buffer, %query_weight: buffer, %key_weight: buffer, %value_weight: buffer, %positions: buffer, %key_cache_indices: buffer, %value_cache_indices: buffer, %query_output_raw: buffer, %key_output_raw: buffer, %value_output_raw: buffer, %query_norm_weight: buffer, %key_norm_weight: buffer, %inverse_frequencies: buffer, %query_output: buffer, %key_cache: buffer, %value_cache: buffer, %completion_counters: buffer) { + %value_uses_q6_index = config.get @qwen3_moe.attention.value_uses_q6 : index + template.apply<@qwen3_moe.attention.qkv_postprocess_fused_decode.body>(%value_uses_q6_index, %token_count, %cache_row_count, %q8_input, %query_weight, %key_weight, %value_weight, %positions, %key_cache_indices, %value_cache_indices, %query_output_raw, %key_output_raw, %value_output_raw, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %query_output, %key_cache, %value_cache, %completion_counters) : (index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +kernel.def target(@qwen3_moe_attention_qkv_postprocess_gfx11_wave32) @qwen3_moe_attention_qkv_postprocess_fused_decode_q4(%token_count: index, %cache_row_count: index) { + %query_size = config.get @qwen3_moe.attention.query_size : index + %c1 = index.constant 1 : index + %key_value_size = config.get @qwen3_moe.attention.key_value_size : index + %c2 = index.constant 2 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c256 = index.constant 256 : index + %key_value_output_size = index.mul %key_value_size, %c2 : index + %output_size = index.add %query_size, %key_value_output_size : index + %padded_output_size = index.add %output_size, %c7 : index + %output_tiles = index.div %padded_output_size, %c8 : index + kernel.launch.config workgroups(%output_tiles, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %cache_row_count: index, %q8_input: buffer, %query_weight: buffer, %key_weight: buffer, %value_weight: buffer, %positions: buffer, %key_cache_indices: buffer, %value_cache_indices: buffer, %query_output_raw: buffer, %key_output_raw: buffer, %value_output_raw: buffer, %query_norm_weight: buffer, %key_norm_weight: buffer, %inverse_frequencies: buffer, %query_output: buffer, %key_cache: buffer, %value_cache: buffer, %completion_counters: buffer) { + %q4 = index.constant 0 : index + template.apply<@qwen3_moe.attention.qkv_postprocess_fused_decode.body>(%q4, %token_count, %cache_row_count, %q8_input, %query_weight, %key_weight, %value_weight, %positions, %key_cache_indices, %value_cache_indices, %query_output_raw, %key_output_raw, %value_output_raw, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %query_output, %key_cache, %value_cache, %completion_counters) : (index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +kernel.def target(@qwen3_moe_attention_qkv_postprocess_gfx11_wave32) @qwen3_moe_attention_qkv_postprocess_fused_decode_q6(%token_count: index, %cache_row_count: index) { + %query_size = config.get @qwen3_moe.attention.query_size : index + %c1 = index.constant 1 : index + %key_value_size = config.get @qwen3_moe.attention.key_value_size : index + %c2 = index.constant 2 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c256 = index.constant 256 : index + %key_value_output_size = index.mul %key_value_size, %c2 : index + %output_size = index.add %query_size, %key_value_output_size : index + %padded_output_size = index.add %output_size, %c7 : index + %output_tiles = index.div %padded_output_size, %c8 : index + kernel.launch.config workgroups(%output_tiles, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %cache_row_count: index, %q8_input: buffer, %query_weight: buffer, %key_weight: buffer, %value_weight: buffer, %positions: buffer, %key_cache_indices: buffer, %value_cache_indices: buffer, %query_output_raw: buffer, %key_output_raw: buffer, %value_output_raw: buffer, %query_norm_weight: buffer, %key_norm_weight: buffer, %inverse_frequencies: buffer, %query_output: buffer, %key_cache: buffer, %value_cache: buffer, %completion_counters: buffer) { + %q6 = index.constant 1 : index + template.apply<@qwen3_moe.attention.qkv_postprocess_fused_decode.body>(%q6, %token_count, %cache_row_count, %q8_input, %query_weight, %key_weight, %value_weight, %positions, %key_cache_indices, %value_cache_indices, %query_output_raw, %key_output_raw, %value_output_raw, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %query_output, %key_cache, %value_cache, %completion_counters) : (index, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +// Uses the production head geometry, nonzero rotary position, and distinct K/V +// cache rows. Two fused calls reuse the same counters after comparison with +// the ordinary storage-selected composition. The value buffer has production +// Q6_K capacity and is intentionally oversized when the Q4_K config is tested. +check.case public @qwen3_moe_attention_qkv_postprocess_fused_differential_case { + %token_count = check.literal value(1) : index + %cache_row_count = check.literal value(4) : index + %hidden_size = check.literal value(2048) : index + %input = check.generate.fill value(0.00390625) : tensor<1x2048xf32> + %q8_input = check.generate.fill value(0) : tensor<2304xi8> + %query_weight = check.generate.fill value(34) : tensor<4096x8x144xi8> + %key_weight = check.generate.fill value(35) : tensor<512x8x144xi8> + %value_weight = check.generate.fill value(-86) : tensor<860160xi8> + %positions = check.generate.fill value(7) : tensor<1xi32> + %key_cache_indices = check.generate.fill value(1) : tensor<1xi64> + %value_cache_indices = check.generate.fill value(2) : tensor<1xi64> + %query_norm_weight = check.generate.iota offset(0.5) step(0.00390625) : tensor<128xf32> + %key_norm_weight = check.generate.iota offset(0.75) step(0.001953125) : tensor<128xf32> + %inverse_frequencies = check.generate.fill value(0.03125) : tensor<64xf32> + %expected_query_raw = check.generate.fill value(1.0) : tensor<1x4096xf32> + %expected_key_raw = check.generate.fill value(1.0) : tensor<1x512xf32> + %expected_value_raw = check.generate.fill value(1.0) : tensor<1x512xf32> + %expected_query = check.generate.fill value(0.0) : tensor<1x32x128xf32> + %expected_key_cache = check.generate.fill value(-1.0) : tensor<4x4x128xf16> + %expected_value_cache = check.generate.fill value(-1.0) : tensor<4x4x128xf16> + %actual_query_raw0 = check.generate.fill value(1.0) : tensor<1x4096xf32> + %actual_key_raw0 = check.generate.fill value(1.0) : tensor<1x512xf32> + %actual_value_raw0 = check.generate.fill value(1.0) : tensor<1x512xf32> + %actual_query0 = check.generate.fill value(0.0) : tensor<1x32x128xf32> + %actual_key_cache0 = check.generate.fill value(-1.0) : tensor<4x4x128xf16> + %actual_value_cache0 = check.generate.fill value(-1.0) : tensor<4x4x128xf16> + %actual_query_raw1 = check.generate.fill value(1.0) : tensor<1x4096xf32> + %actual_key_raw1 = check.generate.fill value(1.0) : tensor<1x512xf32> + %actual_value_raw1 = check.generate.fill value(1.0) : tensor<1x512xf32> + %actual_query1 = check.generate.fill value(0.0) : tensor<1x32x128xf32> + %actual_key_cache1 = check.generate.fill value(-1.0) : tensor<4x4x128xf16> + %actual_value_cache1 = check.generate.fill value(-1.0) : tensor<4x4x128xf16> + %completion_counters = check.generate.fill value(0) : tensor<40xi32> + %expected_counters = check.generate.fill value(0) : tensor<40xi32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %hidden_size](%token_count, %hidden_size, %input, %q8_input) : [index, index](index, index, tensor<1x2048xf32>, tensor<2304xi8>) + kernel.launch @qwen3_moe_attention_qkv_quantized[%token_count](%token_count, %q8_input, %query_weight, %key_weight, %value_weight, %expected_query_raw, %expected_key_raw, %expected_value_raw) : [index](index, tensor<2304xi8>, tensor<4096x8x144xi8>, tensor<512x8x144xi8>, tensor<860160xi8>, tensor<1x4096xf32>, tensor<1x512xf32>, tensor<1x512xf32>) + kernel.launch @qwen3_moe_attention_postprocess_f32_f16[%token_count, %cache_row_count](%token_count, %cache_row_count, %positions, %key_cache_indices, %value_cache_indices, %expected_query_raw, %expected_key_raw, %expected_value_raw, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %expected_query, %expected_key_cache, %expected_value_cache) : [index, index](index, index, tensor<1xi32>, tensor<1xi64>, tensor<1xi64>, tensor<1x4096xf32>, tensor<1x512xf32>, tensor<1x512xf32>, tensor<128xf32>, tensor<128xf32>, tensor<64xf32>, tensor<1x32x128xf32>, tensor<4x4x128xf16>, tensor<4x4x128xf16>) + kernel.launch @qwen3_moe_attention_qkv_postprocess_fused_decode[%token_count, %cache_row_count](%token_count, %cache_row_count, %q8_input, %query_weight, %key_weight, %value_weight, %positions, %key_cache_indices, %value_cache_indices, %actual_query_raw0, %actual_key_raw0, %actual_value_raw0, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %actual_query0, %actual_key_cache0, %actual_value_cache0, %completion_counters) : [index, index](index, index, tensor<2304xi8>, tensor<4096x8x144xi8>, tensor<512x8x144xi8>, tensor<860160xi8>, tensor<1xi32>, tensor<1xi64>, tensor<1xi64>, tensor<1x4096xf32>, tensor<1x512xf32>, tensor<1x512xf32>, tensor<128xf32>, tensor<128xf32>, tensor<64xf32>, tensor<1x32x128xf32>, tensor<4x4x128xf16>, tensor<4x4x128xf16>, tensor<40xi32>) + kernel.launch @qwen3_moe_attention_qkv_postprocess_fused_decode[%token_count, %cache_row_count](%token_count, %cache_row_count, %q8_input, %query_weight, %key_weight, %value_weight, %positions, %key_cache_indices, %value_cache_indices, %actual_query_raw1, %actual_key_raw1, %actual_value_raw1, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %actual_query1, %actual_key_cache1, %actual_value_cache1, %completion_counters) : [index, index](index, index, tensor<2304xi8>, tensor<4096x8x144xi8>, tensor<512x8x144xi8>, tensor<860160xi8>, tensor<1xi32>, tensor<1xi64>, tensor<1xi64>, tensor<1x4096xf32>, tensor<1x512xf32>, tensor<1x512xf32>, tensor<128xf32>, tensor<128xf32>, tensor<64xf32>, tensor<1x32x128xf32>, tensor<4x4x128xf16>, tensor<4x4x128xf16>, tensor<40xi32>) + check.expect.close actual(%actual_query_raw0) expected(%expected_query_raw) atol(0.0) rtol(0.0) nan(same) : tensor<1x4096xf32> + check.expect.close actual(%actual_key_raw0) expected(%expected_key_raw) atol(0.0) rtol(0.0) nan(same) : tensor<1x512xf32> + check.expect.close actual(%actual_value_raw0) expected(%expected_value_raw) atol(0.0) rtol(0.0) nan(same) : tensor<1x512xf32> + check.expect.close actual(%actual_query0) expected(%expected_query) atol(0.0001) rtol(0.0001) nan(same) : tensor<1x32x128xf32> + check.expect.close actual(%actual_key_cache0) expected(%expected_key_cache) atol(0.002) rtol(0.002) nan(same) : tensor<4x4x128xf16> + check.expect.close actual(%actual_value_cache0) expected(%expected_value_cache) atol(0.0) rtol(0.0) nan(same) : tensor<4x4x128xf16> + check.expect.close actual(%actual_query_raw1) expected(%expected_query_raw) atol(0.0) rtol(0.0) nan(same) : tensor<1x4096xf32> + check.expect.close actual(%actual_key_raw1) expected(%expected_key_raw) atol(0.0) rtol(0.0) nan(same) : tensor<1x512xf32> + check.expect.close actual(%actual_value_raw1) expected(%expected_value_raw) atol(0.0) rtol(0.0) nan(same) : tensor<1x512xf32> + check.expect.close actual(%actual_query1) expected(%expected_query) atol(0.0001) rtol(0.0001) nan(same) : tensor<1x32x128xf32> + check.expect.close actual(%actual_key_cache1) expected(%expected_key_cache) atol(0.002) rtol(0.002) nan(same) : tensor<4x4x128xf16> + check.expect.close actual(%actual_value_cache1) expected(%expected_value_cache) atol(0.0) rtol(0.0) nan(same) : tensor<4x4x128xf16> + check.expect.equal actual(%completion_counters) expected(%expected_counters) : tensor<40xi32> + check.return +} + +check.case public @qwen3_moe_attention_qkv_postprocess_composed_benchmark_case { + %token_count = check.literal value(1) : index + %cache_row_count = check.literal value(1024) : index + %q8_input = check.generate.fill value(0) : tensor<2304xi8> + %query_weight = check.generate.fill value(0) : tensor<4096x8x144xi8> + %key_weight = check.generate.fill value(0) : tensor<512x8x144xi8> + %value_weight = check.generate.fill value(0) : tensor<512x8x210xi8> + %positions = check.generate.fill value(513) : tensor<1xi32> + %key_cache_indices = check.generate.fill value(513) : tensor<1xi64> + %value_cache_indices = check.generate.fill value(513) : tensor<1xi64> + %query_output_raw = check.generate.fill value(1.0) : tensor<1x4096xf32> + %key_output_raw = check.generate.fill value(1.0) : tensor<1x512xf32> + %value_output_raw = check.generate.fill value(1.0) : tensor<1x512xf32> + %query_norm_weight = check.generate.fill value(1.0) : tensor<128xf32> + %key_norm_weight = check.generate.fill value(1.0) : tensor<128xf32> + %inverse_frequencies = check.generate.fill value(1.0) : tensor<64xf32> + %query_output = check.generate.fill value(0.0) : tensor<1x32x128xf32> + %key_cache = check.generate.fill value(0.0) : tensor<1024x4x128xf16> + %value_cache = check.generate.fill value(0.0) : tensor<1024x4x128xf16> + kernel.launch @qwen3_moe_attention_qkv_quantized[%token_count](%token_count, %q8_input, %query_weight, %key_weight, %value_weight, %query_output_raw, %key_output_raw, %value_output_raw) : [index](index, tensor<2304xi8>, tensor<4096x8x144xi8>, tensor<512x8x144xi8>, tensor<512x8x210xi8>, tensor<1x4096xf32>, tensor<1x512xf32>, tensor<1x512xf32>) + kernel.launch @qwen3_moe_attention_postprocess_f32_f16[%token_count, %cache_row_count](%token_count, %cache_row_count, %positions, %key_cache_indices, %value_cache_indices, %query_output_raw, %key_output_raw, %value_output_raw, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %query_output, %key_cache, %value_cache) : [index, index](index, index, tensor<1xi32>, tensor<1xi64>, tensor<1xi64>, tensor<1x4096xf32>, tensor<1x512xf32>, tensor<1x512xf32>, tensor<128xf32>, tensor<128xf32>, tensor<64xf32>, tensor<1x32x128xf32>, tensor<1024x4x128xf16>, tensor<1024x4x128xf16>) + check.return +} + +check.case public @qwen3_moe_attention_qkv_postprocess_fused_benchmark_case { + %token_count = check.literal value(1) : index + %cache_row_count = check.literal value(1024) : index + %q8_input = check.generate.fill value(0) : tensor<2304xi8> + %query_weight = check.generate.fill value(0) : tensor<4096x8x144xi8> + %key_weight = check.generate.fill value(0) : tensor<512x8x144xi8> + %value_weight = check.generate.fill value(0) : tensor<512x8x210xi8> + %positions = check.generate.fill value(513) : tensor<1xi32> + %key_cache_indices = check.generate.fill value(513) : tensor<1xi64> + %value_cache_indices = check.generate.fill value(513) : tensor<1xi64> + %query_output_raw = check.generate.fill value(1.0) : tensor<1x4096xf32> + %key_output_raw = check.generate.fill value(1.0) : tensor<1x512xf32> + %value_output_raw = check.generate.fill value(1.0) : tensor<1x512xf32> + %query_norm_weight = check.generate.fill value(1.0) : tensor<128xf32> + %key_norm_weight = check.generate.fill value(1.0) : tensor<128xf32> + %inverse_frequencies = check.generate.fill value(1.0) : tensor<64xf32> + %query_output = check.generate.fill value(0.0) : tensor<1x32x128xf32> + %key_cache = check.generate.fill value(0.0) : tensor<1024x4x128xf16> + %value_cache = check.generate.fill value(0.0) : tensor<1024x4x128xf16> + %completion_counters = check.generate.fill value(0) : tensor<40xi32> + kernel.launch @qwen3_moe_attention_qkv_postprocess_fused_decode[%token_count, %cache_row_count](%token_count, %cache_row_count, %q8_input, %query_weight, %key_weight, %value_weight, %positions, %key_cache_indices, %value_cache_indices, %query_output_raw, %key_output_raw, %value_output_raw, %query_norm_weight, %key_norm_weight, %inverse_frequencies, %query_output, %key_cache, %value_cache, %completion_counters) : [index, index](index, index, tensor<2304xi8>, tensor<4096x8x144xi8>, tensor<512x8x144xi8>, tensor<512x8x210xi8>, tensor<1xi32>, tensor<1xi64>, tensor<1xi64>, tensor<1x4096xf32>, tensor<1x512xf32>, tensor<1x512xf32>, tensor<128xf32>, tensor<128xf32>, tensor<64xf32>, tensor<1x32x128xf32>, tensor<1024x4x128xf16>, tensor<1024x4x128xf16>, tensor<40xi32>) + check.return +} + +check.benchmark<@qwen3_moe_attention_qkv_postprocess_composed_benchmark_case> @qwen3_moe_attention_qkv_postprocess_composed_decode + +check.benchmark<@qwen3_moe_attention_qkv_postprocess_fused_benchmark_case> @qwen3_moe_attention_qkv_postprocess_fused_boundary_decode diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/attention_qkv_quantized.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/attention_qkv_quantized.loom new file mode 100644 index 000000000000..3a5d0ddc5f79 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/attention_qkv_quantized.loom @@ -0,0 +1,338 @@ +// Co-scheduled Qwen Q/K/V projections from one GGML Q8_1 x4 activation row. +// +// Query and key weights use their native GGUF Q4_K rows. Value weights are +// selected at JIT time between Q4_K and Q6_K because the checkpoint uses both +// layer contracts. One wave owns one output row, while a 256-workitem +// workgroup carries eight rows drawn from the concatenated Q/K/V output +// domain. The three logical outputs remain distinct bindings and layouts. +// +// This kernel intentionally starts after activation packing and ends at raw +// F32 projections. The adjacent preparation producer owns RMSNorm and Q8_1 +// packing; the following attention postprocess owns per-head normalization, +// RoPE, and K/V cache publication. +template.decl @qwen3_moe.attention.qkv_quantized.body(%value_uses_q6_index: index, %publish_output: i1, %token_count: index, %token0: index, %q8_input: buffer, %query_weight: buffer, %key_weight: buffer, %value_weight: buffer, %query_output: buffer, %key_output: buffer, %value_output: buffer) + +amdgpu.target @qwen3_moe_attention_qkv_gfx11_wave32 {subgroup_size = 32} + +config.decl @qwen3_moe.model.hidden_size : %value: index where [range(%value, 128, 32768), mul(%value, 128)] + +config.decl @qwen3_moe.attention.query_size : %value: index where [range(%value, 1, 262144)] + +config.decl @qwen3_moe.attention.key_value_size : %value: index where [range(%value, 1, 262144)] + +config.decl @qwen3_moe.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +// Selects Q4_K (0) or Q6_K (1) storage for the value projection. +config.decl @qwen3_moe.attention.value_uses_q6 : %value: index where [range(%value, 0, 1)] + +// Paired-nibble Q4_K row contraction provider linked from the dense quantized +// library. +func.decl @qwen3_moe_q4k_q8_1_x4_paired_row_lane(%input_size: index, %weight: buffer, %weight_row_byte_base: offset, %q8_input: buffer, %q8_row_byte_base: offset, %lane: index) -> (f32) + +// Q6_K row contraction provider linked from the GGML quantized library. +func.decl @ggml_q6k_q8_1_x4_row_lane(%input_size: index, %weight: buffer, %weight_row_byte_base: offset, %q8_input: buffer, %q8_row_byte_base: offset, %lane: index) -> (f32) + +// Fused RMSNorm and Q8_1 x4 producer linked from attention preparation. +kernel.decl @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4(%token_count$30: index) launch(%token_count$31: index, %input: buffer, %weight: buffer, %q8_output: buffer) + +// Reference entry points used only by differential cases. +kernel.decl @ggml_quantize_q8_1_x4_f32(%token_count$35: index, %input_size$36: index) launch(%token_count$37: index, %input_size$38: index, %input: buffer, %output: buffer) + +kernel.decl @qwen3_moe_dense_linear_q4k_q8_1_x4(%token_count$41: index) launch(%token_count$42: index, %q8_input: buffer, %weight: buffer, %output: buffer) + +kernel.decl @ggml_linear_q6k_q8_1_x4(%token_count$46: index, %input_size$47: index, %output_size$48: index) launch(%token_count$49: index, %input_size$50: index, %output_size$51: index, %q8_input: buffer, %weight: buffer, %output: buffer) + +// Device body shared by the ordinary projection and completion-fused +// postprocess exports. Each caller owns its boundary after the raw row stores. +template.def<@qwen3_moe.attention.qkv_quantized.body> device @qwen3_moe_attention_qkv_quantized_body(%value_uses_q6_index: index, %publish_output: i1, %token_count: index, %token0: index, %q8_input: buffer, %query_weight: buffer, %key_weight: buffer, %value_weight: buffer, %query_output: buffer, %key_output: buffer, %value_output: buffer) { + %hidden_size0 = config.get @qwen3_moe.model.hidden_size : index + %query_size0 = config.get @qwen3_moe.attention.query_size : index + %key_value_size0 = config.get @qwen3_moe.attention.key_value_size : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %hidden_size = index.assume %hidden_size0 [range(%hidden_size0, 128, 32768), mul(%hidden_size0, 128)] : index + %query_size, %key_value_size = index.assume %query_size0, %key_value_size0 [range(%query_size0, 1, 262144), range(%key_value_size0, 1, 262144), mul(%query_size0, %key_value_size0)] : index, index + %channel_tile = kernel.workgroup.id : index + %subgroup0 = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c8 = index.constant 8 : index + %c128 = index.constant 128 : index + %c144_bytes = index.constant 144 : offset + %c210_bytes = index.constant 210 : offset + %c256 = index.constant 256 : index + %c0_i32 = scalar.constant 0 : i32 + %c0_f32 = scalar.constant 0.0 : f32 + %c0_offset = index.constant 0 : offset + %token, %launch_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %channel_base = index.mul %channel_tile, %c8 : index + %global_channel = index.add %channel_base, %subgroup : index + %key_value_end = index.add %query_size, %key_value_size : index + %key_value_output_size = index.mul %key_value_size, %c2 : index + %total_output_size = index.add %query_size, %key_value_output_size : index + %valid_channel = index.cmp ult, %global_channel, %total_output_size : index + %is_query = index.cmp ult, %global_channel, %query_size : index + %is_key = index.cmp ult, %global_channel, %key_value_end : index + %key_value_channel = index.rem %global_channel, %key_value_size : index + %value_uses_q6 = index.cmp eq, %value_uses_q6_index, %c1 : index + %lane_i32 = index.cast %lane : index to i32 + %is_lane_zero = scalar.cmpi eq, %lane_i32, %c0_i32 : i32 + %quant_block_count = index.div %hidden_size, %c256 : index + %q4_row_bytes = index.scale %quant_block_count, %c144_bytes : index, offset -> offset + %q6_row_bytes = index.scale %quant_block_count, %c210_bytes : index, offset -> offset + %q8_group_count = index.div %hidden_size, %c128 : index + %q8_row_bytes = index.scale %q8_group_count, %c144_bytes : index, offset -> offset + %q8_row_byte_base = index.scale %token, %q8_row_bytes : index, offset -> offset + %q8_noalias, %query_weight_noalias, %key_weight_noalias, %value_weight_noalias, %query_output_noalias, %key_output_noalias, %value_output_noalias = buffer.assume.noalias %q8_input, %query_weight, %key_weight, %value_weight, %query_output, %key_output, %value_output : buffer, buffer, buffer, buffer, buffer, buffer, buffer + %lane_sum = scf.if %publish_output -> (f32) { + %channel_sum = scf.if %valid_channel -> (f32) { + %projection_sum = scf.if %is_query -> (f32) { + %row_byte_base = index.scale %global_channel, %q4_row_bytes : index, offset -> offset + %sum = func.call @qwen3_moe_q4k_q8_1_x4_paired_row_lane(%hidden_size, %query_weight_noalias, %row_byte_base, %q8_noalias, %q8_row_byte_base, %lane) : (index, buffer, offset, buffer, offset, index) -> (f32) + scf.yield %sum : f32 + } else { + %key_or_value_sum = scf.if %is_key -> (f32) { + %row_byte_base = index.scale %key_value_channel, %q4_row_bytes : index, offset -> offset + %sum = func.call @qwen3_moe_q4k_q8_1_x4_paired_row_lane(%hidden_size, %key_weight_noalias, %row_byte_base, %q8_noalias, %q8_row_byte_base, %lane) : (index, buffer, offset, buffer, offset, index) -> (f32) + scf.yield %sum : f32 + } else { + %value_sum = scf.if %value_uses_q6 -> (f32) { + %row_byte_base = index.scale %key_value_channel, %q6_row_bytes : index, offset -> offset + %sum = func.call @ggml_q6k_q8_1_x4_row_lane(%hidden_size, %value_weight_noalias, %row_byte_base, %q8_noalias, %q8_row_byte_base, %lane) : (index, buffer, offset, buffer, offset, index) -> (f32) + scf.yield %sum : f32 + } else { + %row_byte_base = index.scale %key_value_channel, %q4_row_bytes : index, offset -> offset + %sum = func.call @qwen3_moe_q4k_q8_1_x4_paired_row_lane(%hidden_size, %value_weight_noalias, %row_byte_base, %q8_noalias, %q8_row_byte_base, %lane) : (index, buffer, offset, buffer, offset, index) -> (f32) + scf.yield %sum : f32 + } + scf.yield %value_sum : f32 + } + scf.yield %key_or_value_sum : f32 + } + scf.yield %projection_sum : f32 + } else { + scf.yield %c0_f32 : f32 + } + scf.yield %channel_sum : f32 + } else { + scf.yield %c0_f32 : f32 + } + %dot = kernel.subgroup.reduce %lane_sum : f32 + %query_output_view = buffer.view %query_output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%query_size]xf32> + %key_output_view = buffer.view %key_output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%key_value_size]xf32> + %value_output_view = buffer.view %value_output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%key_value_size]xf32> + scf.if %publish_output { + scf.if %valid_channel { + scf.if %is_lane_zero { + scf.if %is_query { + view.store %dot, %query_output_view[%token, %global_channel] : f32, view<[%launch_token_count]x[%query_size]xf32> + } else { + scf.if %is_key { + view.store %dot, %key_output_view[%token, %key_value_channel] : f32, view<[%launch_token_count]x[%key_value_size]xf32> + } else { + view.store %dot, %value_output_view[%token, %key_value_channel] : f32, view<[%launch_token_count]x[%key_value_size]xf32> + } + } + } + } + } + template.return +} + +kernel.def target(@qwen3_moe_attention_qkv_gfx11_wave32) @qwen3_moe_attention_qkv_quantized(%token_count: index) { + %query_size = config.get @qwen3_moe.attention.query_size : index + %key_value_size = config.get @qwen3_moe.attention.key_value_size : index + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c256 = index.constant 256 : index + %key_value_output_size = index.mul %key_value_size, %c2 : index + %output_size = index.add %query_size, %key_value_output_size : index + %padded_output_size = index.add %output_size, %c7 : index + %output_tiles = index.div %padded_output_size, %c8 : index + kernel.launch.config workgroups(%output_tiles, %token_capacity, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %q8_input: buffer, %query_weight: buffer, %key_weight: buffer, %value_weight: buffer, %query_output: buffer, %key_output: buffer, %value_output: buffer) { + %value_uses_q6_index = config.get @qwen3_moe.attention.value_uses_q6 : index + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %token0 = kernel.workgroup.id : index + %c0 = index.constant 0 : index + %valid_token = index.cmp ult, %token0, %bounded_token_count : index + %safe_token = scf.select %valid_token, %token0, %c0 : index + template.apply<@qwen3_moe.attention.qkv_quantized.body>(%value_uses_q6_index, %valid_token, %bounded_token_count, %safe_token, %q8_input, %query_weight, %key_weight, %value_weight, %query_output, %key_output, %value_output) : (index, i1, index, index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +// Distinct token rows and Q/K/V byte patterns make row ownership and +// binding-domain mistakes observable. Equal synthetic output widths let the +// two Q4_K references share the config-specialized direct provider while still +// crossing every domain edge. +check.case public @qwen3_moe_attention_qkv_q6_differential_case { + %token_count = check.literal value(14) : index + %hidden_size = check.literal value(512) : index + %output_size = check.literal value(65) : index + %input_seed = check.param.seed base(5858425849414112822) count(1) : i64 + %input = check.generate.random.uniform seed(%input_seed) range(-1.0 to 1.0) : tensor<14x512xf32> + %q8_input = check.generate.fill value(0) : tensor<14x576xi8> + %query_weight = check.generate.fill value(34) : tensor<65x2x144xi8> + %key_weight = check.generate.fill value(35) : tensor<65x2x144xi8> + %value_weight = check.generate.fill value(-86) : tensor<65x2x210xi8> + %expected_query = check.generate.fill value(0.0) : tensor<14x65xf32> + %expected_key = check.generate.fill value(0.0) : tensor<14x65xf32> + %expected_value = check.generate.fill value(0.0) : tensor<14x65xf32> + %actual_query = check.generate.fill value(1.0) : tensor<14x65xf32> + %actual_key = check.generate.fill value(1.0) : tensor<14x65xf32> + %actual_value = check.generate.fill value(1.0) : tensor<14x65xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %hidden_size](%token_count, %hidden_size, %input, %q8_input) : [index, index](index, index, tensor<14x512xf32>, tensor<14x576xi8>) + kernel.launch @qwen3_moe_dense_linear_q4k_q8_1_x4[%token_count](%token_count, %q8_input, %query_weight, %expected_query) : [index](index, tensor<14x576xi8>, tensor<65x2x144xi8>, tensor<14x65xf32>) + kernel.launch @qwen3_moe_dense_linear_q4k_q8_1_x4[%token_count](%token_count, %q8_input, %key_weight, %expected_key) : [index](index, tensor<14x576xi8>, tensor<65x2x144xi8>, tensor<14x65xf32>) + kernel.launch @ggml_linear_q6k_q8_1_x4[%token_count, %hidden_size, %output_size](%token_count, %hidden_size, %output_size, %q8_input, %value_weight, %expected_value) : [index, index, index](index, index, index, tensor<14x576xi8>, tensor<65x2x210xi8>, tensor<14x65xf32>) + kernel.launch @qwen3_moe_attention_qkv_quantized[%token_count](%token_count, %q8_input, %query_weight, %key_weight, %value_weight, %actual_query, %actual_key, %actual_value) : [index](index, tensor<14x576xi8>, tensor<65x2x144xi8>, tensor<65x2x144xi8>, tensor<65x2x210xi8>, tensor<14x65xf32>, tensor<14x65xf32>, tensor<14x65xf32>) + check.expect.close actual(%actual_query) expected(%expected_query) atol(0.25) rtol(0.01) nan(same) : tensor<14x65xf32> + check.expect.close actual(%actual_key) expected(%expected_key) atol(0.25) rtol(0.01) nan(same) : tensor<14x65xf32> + check.expect.close actual(%actual_value) expected(%expected_value) atol(0.25) rtol(0.01) nan(same) : tensor<14x65xf32> + check.return +} + +check.case public @qwen3_moe_attention_qkv_q4_differential_case { + %token_count = check.literal value(14) : index + %hidden_size = check.literal value(512) : index + %input_seed = check.param.seed base(5858425849414112820) count(1) : i64 + %input = check.generate.random.uniform seed(%input_seed) range(-1.0 to 1.0) : tensor<14x512xf32> + %q8_input = check.generate.fill value(0) : tensor<14x576xi8> + %query_weight = check.generate.fill value(34) : tensor<65x2x144xi8> + %key_weight = check.generate.fill value(35) : tensor<65x2x144xi8> + %value_weight = check.generate.fill value(36) : tensor<65x2x144xi8> + %expected_query = check.generate.fill value(0.0) : tensor<14x65xf32> + %expected_key = check.generate.fill value(0.0) : tensor<14x65xf32> + %expected_value = check.generate.fill value(0.0) : tensor<14x65xf32> + %actual_query = check.generate.fill value(1.0) : tensor<14x65xf32> + %actual_key = check.generate.fill value(1.0) : tensor<14x65xf32> + %actual_value = check.generate.fill value(1.0) : tensor<14x65xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %hidden_size](%token_count, %hidden_size, %input, %q8_input) : [index, index](index, index, tensor<14x512xf32>, tensor<14x576xi8>) + kernel.launch @qwen3_moe_dense_linear_q4k_q8_1_x4[%token_count](%token_count, %q8_input, %query_weight, %expected_query) : [index](index, tensor<14x576xi8>, tensor<65x2x144xi8>, tensor<14x65xf32>) + kernel.launch @qwen3_moe_dense_linear_q4k_q8_1_x4[%token_count](%token_count, %q8_input, %key_weight, %expected_key) : [index](index, tensor<14x576xi8>, tensor<65x2x144xi8>, tensor<14x65xf32>) + kernel.launch @qwen3_moe_dense_linear_q4k_q8_1_x4[%token_count](%token_count, %q8_input, %value_weight, %expected_value) : [index](index, tensor<14x576xi8>, tensor<65x2x144xi8>, tensor<14x65xf32>) + kernel.launch @qwen3_moe_attention_qkv_quantized[%token_count](%token_count, %q8_input, %query_weight, %key_weight, %value_weight, %actual_query, %actual_key, %actual_value) : [index](index, tensor<14x576xi8>, tensor<65x2x144xi8>, tensor<65x2x144xi8>, tensor<65x2x144xi8>, tensor<14x65xf32>, tensor<14x65xf32>, tensor<14x65xf32>) + check.expect.close actual(%actual_query) expected(%expected_query) atol(0.25) rtol(0.01) nan(same) : tensor<14x65xf32> + check.expect.close actual(%actual_key) expected(%expected_key) atol(0.25) rtol(0.01) nan(same) : tensor<14x65xf32> + check.expect.close actual(%actual_value) expected(%expected_value) atol(0.25) rtol(0.01) nan(same) : tensor<14x65xf32> + check.return +} + +check.case public @qwen3_moe_attention_qkv_q6_benchmark_case { + %token_count = check.param.choice values([1, 8, 32, 128, 512]) name("token_count") : index + %q8_input = check.generate.fill value(0) : tensor<[%token_count]x2304xi8> + %query_weight = check.generate.fill value(0) : tensor<4096x8x144xi8> + %key_weight = check.generate.fill value(0) : tensor<512x8x144xi8> + %value_weight = check.generate.fill value(0) : tensor<512x8x210xi8> + %query_output = check.generate.fill value(1.0) : tensor<[%token_count]x4096xf32> + %key_output = check.generate.fill value(1.0) : tensor<[%token_count]x512xf32> + %value_output = check.generate.fill value(1.0) : tensor<[%token_count]x512xf32> + %expected_query = check.generate.fill value(0.0) : tensor<[%token_count]x4096xf32> + %expected_key = check.generate.fill value(0.0) : tensor<[%token_count]x512xf32> + %expected_value = check.generate.fill value(0.0) : tensor<[%token_count]x512xf32> + kernel.launch @qwen3_moe_attention_qkv_quantized[%token_count](%token_count, %q8_input, %query_weight, %key_weight, %value_weight, %query_output, %key_output, %value_output) : [index](index, tensor<[%token_count]x2304xi8>, tensor<4096x8x144xi8>, tensor<512x8x144xi8>, tensor<512x8x210xi8>, tensor<[%token_count]x4096xf32>, tensor<[%token_count]x512xf32>, tensor<[%token_count]x512xf32>) + check.expect.close actual(%query_output) expected(%expected_query) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x4096xf32> + check.expect.close actual(%key_output) expected(%expected_key) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x512xf32> + check.expect.close actual(%value_output) expected(%expected_value) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x512xf32> + check.return +} + +check.case public @qwen3_moe_attention_qkv_q4_benchmark_case { + %token_count = check.param.choice values([1, 8, 32, 128, 512]) name("token_count") : index + %q8_input = check.generate.fill value(0) : tensor<[%token_count]x2304xi8> + %query_weight = check.generate.fill value(0) : tensor<4096x8x144xi8> + %key_weight = check.generate.fill value(0) : tensor<512x8x144xi8> + %value_weight = check.generate.fill value(0) : tensor<512x8x144xi8> + %query_output = check.generate.fill value(1.0) : tensor<[%token_count]x4096xf32> + %key_output = check.generate.fill value(1.0) : tensor<[%token_count]x512xf32> + %value_output = check.generate.fill value(1.0) : tensor<[%token_count]x512xf32> + %expected_query = check.generate.fill value(0.0) : tensor<[%token_count]x4096xf32> + %expected_key = check.generate.fill value(0.0) : tensor<[%token_count]x512xf32> + %expected_value = check.generate.fill value(0.0) : tensor<[%token_count]x512xf32> + kernel.launch @qwen3_moe_attention_qkv_quantized[%token_count](%token_count, %q8_input, %query_weight, %key_weight, %value_weight, %query_output, %key_output, %value_output) : [index](index, tensor<[%token_count]x2304xi8>, tensor<4096x8x144xi8>, tensor<512x8x144xi8>, tensor<512x8x144xi8>, tensor<[%token_count]x4096xf32>, tensor<[%token_count]x512xf32>, tensor<[%token_count]x512xf32>) + check.expect.close actual(%query_output) expected(%expected_query) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x4096xf32> + check.expect.close actual(%key_output) expected(%expected_key) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x512xf32> + check.expect.close actual(%value_output) expected(%expected_value) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x512xf32> + check.return +} + +// Production decode and small-prefill boundary. The two calls become two +// serialized dispatches in one reusable command buffer: the first publishes +// the only Q8_1 activation row and the second consumes it for all projections. +check.case public @qwen3_moe_attention_qkv_full_q6_benchmark_case { + %token_count = check.param.choice values([1, 8, 32]) name("token_count") : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %norm_weight = check.generate.fill value(1.0) : tensor<2048xf32> + %q8_input = check.generate.fill value(1) : tensor<[%token_count]x2304xi8> + %query_weight = check.generate.fill value(0) : tensor<4096x8x144xi8> + %key_weight = check.generate.fill value(0) : tensor<512x8x144xi8> + %value_weight = check.generate.fill value(0) : tensor<512x8x210xi8> + %query_output = check.generate.fill value(1.0) : tensor<[%token_count]x4096xf32> + %key_output = check.generate.fill value(1.0) : tensor<[%token_count]x512xf32> + %value_output = check.generate.fill value(1.0) : tensor<[%token_count]x512xf32> + %expected_query = check.generate.fill value(0.0) : tensor<[%token_count]x4096xf32> + %expected_key = check.generate.fill value(0.0) : tensor<[%token_count]x512xf32> + %expected_value = check.generate.fill value(0.0) : tensor<[%token_count]x512xf32> + kernel.launch @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4[%token_count](%token_count, %input, %norm_weight, %q8_input) : [index](index, tensor<[%token_count]x2048xf32>, tensor<2048xf32>, tensor<[%token_count]x2304xi8>) + kernel.launch @qwen3_moe_attention_qkv_quantized[%token_count](%token_count, %q8_input, %query_weight, %key_weight, %value_weight, %query_output, %key_output, %value_output) : [index](index, tensor<[%token_count]x2304xi8>, tensor<4096x8x144xi8>, tensor<512x8x144xi8>, tensor<512x8x210xi8>, tensor<[%token_count]x4096xf32>, tensor<[%token_count]x512xf32>, tensor<[%token_count]x512xf32>) + check.expect.close actual(%query_output) expected(%expected_query) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x4096xf32> + check.expect.close actual(%key_output) expected(%expected_key) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x512xf32> + check.expect.close actual(%value_output) expected(%expected_value) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x512xf32> + check.return +} + +check.case public @qwen3_moe_attention_qkv_full_q4_benchmark_case { + %token_count = check.param.choice values([1, 8, 32]) name("token_count") : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %norm_weight = check.generate.fill value(1.0) : tensor<2048xf32> + %q8_input = check.generate.fill value(1) : tensor<[%token_count]x2304xi8> + %query_weight = check.generate.fill value(0) : tensor<4096x8x144xi8> + %key_weight = check.generate.fill value(0) : tensor<512x8x144xi8> + %value_weight = check.generate.fill value(0) : tensor<512x8x144xi8> + %query_output = check.generate.fill value(1.0) : tensor<[%token_count]x4096xf32> + %key_output = check.generate.fill value(1.0) : tensor<[%token_count]x512xf32> + %value_output = check.generate.fill value(1.0) : tensor<[%token_count]x512xf32> + %expected_query = check.generate.fill value(0.0) : tensor<[%token_count]x4096xf32> + %expected_key = check.generate.fill value(0.0) : tensor<[%token_count]x512xf32> + %expected_value = check.generate.fill value(0.0) : tensor<[%token_count]x512xf32> + kernel.launch @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4[%token_count](%token_count, %input, %norm_weight, %q8_input) : [index](index, tensor<[%token_count]x2048xf32>, tensor<2048xf32>, tensor<[%token_count]x2304xi8>) + kernel.launch @qwen3_moe_attention_qkv_quantized[%token_count](%token_count, %q8_input, %query_weight, %key_weight, %value_weight, %query_output, %key_output, %value_output) : [index](index, tensor<[%token_count]x2304xi8>, tensor<4096x8x144xi8>, tensor<512x8x144xi8>, tensor<512x8x144xi8>, tensor<[%token_count]x4096xf32>, tensor<[%token_count]x512xf32>, tensor<[%token_count]x512xf32>) + check.expect.close actual(%query_output) expected(%expected_query) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x4096xf32> + check.expect.close actual(%key_output) expected(%expected_key) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x512xf32> + check.expect.close actual(%value_output) expected(%expected_value) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x512xf32> + check.return +} + +check.benchmark<@qwen3_moe_attention_qkv_q6_differential_case> @qwen3_moe_attention_qkv_q6_differential + +check.benchmark<@qwen3_moe_attention_qkv_q4_differential_case> @qwen3_moe_attention_qkv_q4_differential + +check.benchmark<@qwen3_moe_attention_qkv_q6_benchmark_case> @qwen3_moe_attention_qkv_q6_decode {token_count = 1} + +check.benchmark<@qwen3_moe_attention_qkv_q6_benchmark_case> @qwen3_moe_attention_qkv_q6_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_attention_qkv_q6_benchmark_case> @qwen3_moe_attention_qkv_q6_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_attention_qkv_q6_benchmark_case> @qwen3_moe_attention_qkv_q6_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_attention_qkv_q4_benchmark_case> @qwen3_moe_attention_qkv_q4_decode {token_count = 1} + +check.benchmark<@qwen3_moe_attention_qkv_q4_benchmark_case> @qwen3_moe_attention_qkv_q4_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_attention_qkv_q4_benchmark_case> @qwen3_moe_attention_qkv_q4_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_attention_qkv_q4_benchmark_case> @qwen3_moe_attention_qkv_q4_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_attention_qkv_full_q6_benchmark_case> @qwen3_moe_attention_qkv_full_q6_decode {token_count = 1} + +check.benchmark<@qwen3_moe_attention_qkv_full_q6_benchmark_case> @qwen3_moe_attention_qkv_full_q6_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_attention_qkv_full_q4_benchmark_case> @qwen3_moe_attention_qkv_full_q4_decode {token_count = 1} + +check.benchmark<@qwen3_moe_attention_qkv_full_q4_benchmark_case> @qwen3_moe_attention_qkv_full_q4_prefill_32 {token_count = 32} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/attention_qkv_same_format_prefill.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/attention_qkv_same_format_prefill.loom new file mode 100644 index 000000000000..02d266bdb1bf --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/attention_qkv_same_format_prefill.loom @@ -0,0 +1,193 @@ +// Same-format Q4_K attention projections over one contiguous Q/K/V tensor. +// +// The fixed parameter layout places Q, K, and Q4_K V rows consecutively. This +// bounded candidate pairs those [5120][2048] weights with a row-interleaved +// [token][5120] raw output so one uniform 80xN grid can reuse the canonical +// dense WMMA body without buffer-valued control flow, completion counters, or +// additional workgroup storage. +template.decl @qwen3_moe.dense_quantized.body(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: index, %arg5: index, %arg6: index, %arg7: buffer, %arg8: buffer, %arg9: buffer) + +amdgpu.target @qwen3_moe_attention_qkv_same_format_prefill_gfx11_wave64 {subgroup_size = 64} + +config.decl @qwen3_moe.model.hidden_size : %value: index where [range(%value, 128, 32768), mul(%value, 128)] + +config.decl @qwen3_moe.attention.query_size : %value: index where [range(%value, 1, 262144)] + +config.decl @qwen3_moe.attention.key_value_size : %value: index where [range(%value, 1, 262144)] + +config.decl @qwen3_moe.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +// Runs every Q4_K Q/K/V tile in one launch over the aggregate physical rows. +kernel.def target(@qwen3_moe_attention_qkv_same_format_prefill_gfx11_wave64) @qwen3_moe_attention_qkv_q4_prefill_512(%token_count: index) { + %query_size = config.get @qwen3_moe.attention.query_size : index + %key_value_size = config.get @qwen3_moe.attention.key_value_size : index + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %key_value_output_size = index.mul %key_value_size, %c2 : index + %combined_output_size = index.add %query_size, %key_value_output_size : index + %padded_output_size = index.add %combined_output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %padded_token_count = index.add %token_capacity, %c31 : index + %token_tiles = index.div %padded_token_count, %c32 : index + kernel.launch.config workgroups(%output_tiles, %token_tiles, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %combined_weight: buffer, %combined_output: buffer) { + %hidden_size = config.get @qwen3_moe.model.hidden_size : index + %query_size = config.get @qwen3_moe.attention.query_size : index + %key_value_size = config.get @qwen3_moe.attention.key_value_size : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + %c2 = index.constant 2 : index + %q4 = index.constant 4 : index + %overwrite = index.constant 0 : index + %key_value_output_size = index.mul %key_value_size, %c2 : index + %combined_output_size = index.add %query_size, %key_value_output_size : index + template.apply<@qwen3_moe.dense_quantized.body>(%q4, %token_count, %hidden_size, %combined_output_size, %overwrite, %channel_tile, %token_tile, %input, %combined_weight, %combined_output) : (index, index, index, index, index, index, index, buffer, buffer, buffer) + kernel.return +} + +// The three segment exports form the exact aggregate-layout baseline. Their +// workgroups and memory accesses match the combined launch; only the launch +// records differ. +kernel.def target(@qwen3_moe_attention_qkv_same_format_prefill_gfx11_wave64) @qwen3_moe_attention_query_q4_aggregate_prefill_512(%token_count: index) { + %query_size = config.get @qwen3_moe.attention.query_size : index + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %query_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %padded_token_count = index.add %token_capacity, %c31 : index + %token_tiles = index.div %padded_token_count, %c32 : index + kernel.launch.config workgroups(%output_tiles, %token_tiles, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %combined_weight: buffer, %combined_output: buffer) { + %hidden_size = config.get @qwen3_moe.model.hidden_size : index + %query_size = config.get @qwen3_moe.attention.query_size : index + %key_value_size = config.get @qwen3_moe.attention.key_value_size : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + %c2 = index.constant 2 : index + %q4 = index.constant 4 : index + %overwrite = index.constant 0 : index + %key_value_output_size = index.mul %key_value_size, %c2 : index + %combined_output_size = index.add %query_size, %key_value_output_size : index + template.apply<@qwen3_moe.dense_quantized.body>(%q4, %token_count, %hidden_size, %combined_output_size, %overwrite, %channel_tile, %token_tile, %input, %combined_weight, %combined_output) : (index, index, index, index, index, index, index, buffer, buffer, buffer) + kernel.return +} + +kernel.def target(@qwen3_moe_attention_qkv_same_format_prefill_gfx11_wave64) @qwen3_moe_attention_key_q4_aggregate_prefill_512(%token_count: index) { + %key_value_size = config.get @qwen3_moe.attention.key_value_size : index + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %key_value_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %padded_token_count = index.add %token_capacity, %c31 : index + %token_tiles = index.div %padded_token_count, %c32 : index + kernel.launch.config workgroups(%output_tiles, %token_tiles, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %combined_weight: buffer, %combined_output: buffer) { + %hidden_size = config.get @qwen3_moe.model.hidden_size : index + %query_size = config.get @qwen3_moe.attention.query_size : index + %key_value_size = config.get @qwen3_moe.attention.key_value_size : index + %local_channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + %c2 = index.constant 2 : index + %q4 = index.constant 4 : index + %c64 = index.constant 64 : index + %overwrite = index.constant 0 : index + %query_channel_tile_count = index.div %query_size, %c64 : index + %channel_tile = index.add %query_channel_tile_count, %local_channel_tile : index + %key_value_output_size = index.mul %key_value_size, %c2 : index + %combined_output_size = index.add %query_size, %key_value_output_size : index + template.apply<@qwen3_moe.dense_quantized.body>(%q4, %token_count, %hidden_size, %combined_output_size, %overwrite, %channel_tile, %token_tile, %input, %combined_weight, %combined_output) : (index, index, index, index, index, index, index, buffer, buffer, buffer) + kernel.return +} + +kernel.def target(@qwen3_moe_attention_qkv_same_format_prefill_gfx11_wave64) @qwen3_moe_attention_value_q4_aggregate_prefill_512(%token_count: index) { + %key_value_size = config.get @qwen3_moe.attention.key_value_size : index + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %key_value_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %padded_token_count = index.add %token_capacity, %c31 : index + %token_tiles = index.div %padded_token_count, %c32 : index + kernel.launch.config workgroups(%output_tiles, %token_tiles, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %combined_weight: buffer, %combined_output: buffer) { + %hidden_size = config.get @qwen3_moe.model.hidden_size : index + %query_size = config.get @qwen3_moe.attention.query_size : index + %key_value_size = config.get @qwen3_moe.attention.key_value_size : index + %local_channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + %c2 = index.constant 2 : index + %q4 = index.constant 4 : index + %c64 = index.constant 64 : index + %overwrite = index.constant 0 : index + %query_channel_tile_count = index.div %query_size, %c64 : index + %key_value_channel_tile_count = index.div %key_value_size, %c64 : index + %key_channel_tile_end = index.add %query_channel_tile_count, %key_value_channel_tile_count : index + %channel_tile = index.add %key_channel_tile_end, %local_channel_tile : index + %key_value_output_size = index.mul %key_value_size, %c2 : index + %combined_output_size = index.add %query_size, %key_value_output_size : index + template.apply<@qwen3_moe.dense_quantized.body>(%q4, %token_count, %hidden_size, %combined_output_size, %overwrite, %channel_tile, %token_tile, %input, %combined_weight, %combined_output) : (index, index, index, index, index, index, index, buffer, buffer, buffer) + kernel.return +} + +// The production shape crosses every token and channel tile. The aggregate +// output starts with different sentinels so missing or overlapping segment +// ownership is observable even though the packed bytes are uniform. +check.case public @qwen3_moe_attention_qkv_q4_prefill_512_differential_case { + %token_count = check.literal value(512) : index + %input = check.generate.fill value(0.00390625) : tensor<512x2048xf32> + %combined_weight = check.generate.fill value(34) : tensor<5120x8x144xi8> + %expected = check.generate.fill value(-1.0) : tensor<512x5120xf32> + %actual = check.generate.fill value(1.0) : tensor<512x5120xf32> + kernel.launch @qwen3_moe_attention_query_q4_aggregate_prefill_512[%token_count](%token_count, %input, %combined_weight, %expected) : [index](index, tensor<512x2048xf32>, tensor<5120x8x144xi8>, tensor<512x5120xf32>) + kernel.launch @qwen3_moe_attention_key_q4_aggregate_prefill_512[%token_count](%token_count, %input, %combined_weight, %expected) : [index](index, tensor<512x2048xf32>, tensor<5120x8x144xi8>, tensor<512x5120xf32>) + kernel.launch @qwen3_moe_attention_value_q4_aggregate_prefill_512[%token_count](%token_count, %input, %combined_weight, %expected) : [index](index, tensor<512x2048xf32>, tensor<5120x8x144xi8>, tensor<512x5120xf32>) + kernel.launch @qwen3_moe_attention_qkv_q4_prefill_512[%token_count](%token_count, %input, %combined_weight, %actual) : [index](index, tensor<512x2048xf32>, tensor<5120x8x144xi8>, tensor<512x5120xf32>) + check.expect.close actual(%actual) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<512x5120xf32> + check.return +} + +check.case public @qwen3_moe_attention_qkv_q4_prefill_512_composed_benchmark_case { + %token_count = check.literal value(512) : index + %input = check.generate.fill value(0.0) : tensor<512x2048xf32> + %combined_weight = check.generate.fill value(0) : tensor<5120x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<512x5120xf32> + kernel.launch @qwen3_moe_attention_query_q4_aggregate_prefill_512[%token_count](%token_count, %input, %combined_weight, %output) : [index](index, tensor<512x2048xf32>, tensor<5120x8x144xi8>, tensor<512x5120xf32>) + kernel.launch @qwen3_moe_attention_key_q4_aggregate_prefill_512[%token_count](%token_count, %input, %combined_weight, %output) : [index](index, tensor<512x2048xf32>, tensor<5120x8x144xi8>, tensor<512x5120xf32>) + kernel.launch @qwen3_moe_attention_value_q4_aggregate_prefill_512[%token_count](%token_count, %input, %combined_weight, %output) : [index](index, tensor<512x2048xf32>, tensor<5120x8x144xi8>, tensor<512x5120xf32>) + check.return +} + +check.case public @qwen3_moe_attention_qkv_q4_prefill_512_fused_benchmark_case { + %token_count = check.literal value(512) : index + %input = check.generate.fill value(0.0) : tensor<512x2048xf32> + %combined_weight = check.generate.fill value(0) : tensor<5120x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<512x5120xf32> + kernel.launch @qwen3_moe_attention_qkv_q4_prefill_512[%token_count](%token_count, %input, %combined_weight, %output) : [index](index, tensor<512x2048xf32>, tensor<5120x8x144xi8>, tensor<512x5120xf32>) + check.return +} + +check.benchmark<@qwen3_moe_attention_qkv_q4_prefill_512_differential_case> @qwen3_moe_attention_qkv_q4_prefill_512_differential + +check.benchmark<@qwen3_moe_attention_qkv_q4_prefill_512_composed_benchmark_case> @qwen3_moe_attention_qkv_q4_prefill_512_composed + +check.benchmark<@qwen3_moe_attention_qkv_q4_prefill_512_fused_benchmark_case> @qwen3_moe_attention_qkv_q4_prefill_512_fused diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/batched_decode_expert_dispatch.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/batched_decode_expert_dispatch.loom new file mode 100644 index 000000000000..b45ac4ea517d --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/batched_decode_expert_dispatch.loom @@ -0,0 +1,498 @@ +// Builds the GPU-resident expert work queues used by batched decode. +// +// One lane owns each configured expert, counts its assignments, and uses +// workgroup scans to assign deterministic compact offsets. Three configured +// row-count limits classify active experts into four target-selected schedule +// queues. The resulting counts are suitable for device-side indirect launch; +// route identities never cross the host boundary. +// +// assignment_ordinals is expert-contiguous. Each ordinal is the original +// flattened [token][route] position, so consumers can recover token and route +// with division and remainder by the configured route count and scatter +// results directly into the established compact routed-output layout. +// +// queue_descriptors is physically [4][expert_count][4xi32]. Each naturally +// aligned descriptor contains the expert ordinal, assignment base, row count, +// and a reserved zero field. The b128 representation keeps the ABI independent +// of model geometry while giving consumers one native-width descriptor load. +config.decl @qwen3_moe.router.expert_count : %value: index where [range(%value, 32, 512), mul(%value, 32)] + +config.decl @qwen3_moe.router.route_count : %value: index where [range(%value, 1, 32)] + +config.decl @qwen3_moe.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +config.decl @qwen3_moe.batched_decode.schedule0_row_limit : %value: index where [range(%value, 1, 2048)] + +config.decl @qwen3_moe.batched_decode.schedule1_row_limit : %value: index where [range(%value, 1, 2048)] + +config.decl @qwen3_moe.batched_decode.schedule2_row_limit : %value: index where [range(%value, 1, 2048)] + +amdgpu.target @qwen3_moe_batched_decode_dispatch_gfx11_wave32 {subgroup_size = 32} + +// Decodes the stable batched-decode expert descriptor representation. +func.def inline @qwen3_moe_unpack_batched_decode_expert_descriptor(%descriptor: vector<4xi32>) -> (index, index, index) { + %configured_expert_count = config.get @qwen3_moe.router.expert_count : index + %configured_route_count = config.get @qwen3_moe.router.route_count : index + %configured_token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %assignment_capacity = index.mul %configured_token_capacity, %configured_route_count : index + %expert_i32 = vector.extract %descriptor[0] : vector<4xi32> -> i32 + %assignment_base_i32 = vector.extract %descriptor[1] : vector<4xi32> -> i32 + %row_count_i32 = vector.extract %descriptor[2] : vector<4xi32> -> i32 + %expert0 = index.cast %expert_i32 : i32 to index + %expert, %descriptor_expert_count = index.assume %expert0, %configured_expert_count [range(%expert0, 0, 511), lt(%expert0, %configured_expert_count)] : index, index + %assignment_base0 = index.cast %assignment_base_i32 : i32 to index + %assignment_base, %descriptor_assignment_capacity = index.assume %assignment_base0, %assignment_capacity [range(%assignment_base0, 0, 65535), lt(%assignment_base0, %assignment_capacity)] : index, index + %row_count0 = index.cast %row_count_i32 : i32 to index + %row_count, %descriptor_token_capacity = index.assume %row_count0, %configured_token_capacity [range(%row_count0, 1, 2048), le(%row_count0, %configured_token_capacity)] : index, index + func.return %expert, %assignment_base, %row_count : index, index, index +} + +kernel.def target(@qwen3_moe_batched_decode_dispatch_gfx11_wave32) @qwen3_moe_build_batched_decode_expert_dispatch(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index) { + %configured_expert_count = config.get @qwen3_moe.router.expert_count : index + %c1 = index.constant 1 : index + kernel.launch.config workgroups(%c1, %c1, %c1) workgroup_size(%configured_expert_count, %c1, %c1) : index +} launch(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index, %route_ids: buffer, %assignment_ordinals: buffer, %queue_counts: buffer, %queue_descriptors: buffer) { + %configured_token_capacity0 = config.get @qwen3_moe.workload.token_capacity : index + %bounded_token_count, %configured_token_capacity = index.assume %token_count, %configured_token_capacity0 [range(%token_count, 1, 2048), le(%token_count, %configured_token_capacity0)] : index, index + %configured_route_count0 = config.get @qwen3_moe.router.route_count : index + %bounded_route_count, %configured_route_count = index.assume %route_count, %configured_route_count0 [range(%route_count, 1, 32), eq(%route_count, %configured_route_count0)] : index, index + %bounded_route_stride, %route_row_width = index.assume %route_stride, %bounded_route_count [range(%route_stride, 1, 512), le(%bounded_route_count, %route_stride)] : index, index + %configured_expert_count0 = config.get @qwen3_moe.router.expert_count : index + %bounded_expert_count, %configured_expert_count = index.assume %expert_count, %configured_expert_count0 [range(%expert_count, 32, 512), eq(%expert_count, %configured_expert_count0)] : index, index + %schedule0_row_limit0 = config.get @qwen3_moe.batched_decode.schedule0_row_limit : index + %schedule1_row_limit0 = config.get @qwen3_moe.batched_decode.schedule1_row_limit : index + %schedule2_row_limit0 = config.get @qwen3_moe.batched_decode.schedule2_row_limit : index + %schedule0_row_limit, %schedule1_row_limit, %schedule2_row_limit = index.assume %schedule0_row_limit0, %schedule1_row_limit0, %schedule2_row_limit0 [lt(%schedule0_row_limit0, %schedule1_row_limit0), lt(%schedule1_row_limit0, %schedule2_row_limit0)] : index, index, index + %lane0 = kernel.workitem.id : index + %expert, %launch_expert_count = index.assume %lane0, %bounded_expert_count [lt(%lane0, %bounded_expert_count)] : index, index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c0_offset = index.constant 0 : offset + %assignment_count = index.mul %bounded_token_count, %configured_route_count : index + %route_ids_noalias, %assignment_ordinals_noalias, %queue_counts_noalias, %queue_descriptors_noalias = buffer.assume.noalias %route_ids, %assignment_ordinals, %queue_counts, %queue_descriptors : buffer, buffer, buffer, buffer + %route_view = buffer.view %route_ids_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_route_stride]xi32> + %assignment_view = buffer.view %assignment_ordinals_noalias[%c0_offset] : buffer -> view<[%assignment_count]xi32> + %queue_count_view = buffer.view %queue_counts_noalias[%c0_offset] : buffer -> view<4xi32> + %queue_descriptor_view = buffer.view %queue_descriptors_noalias[%c0_offset] : buffer -> view<4x[%configured_expert_count]x4xi32> + %expert_i32 = index.cast %expert : index to i32 + + // Every lane walks the same tiny, cache-resident route plane. The serial + // comparisons are substantially cheaper than a host-visible sort or atomics + // and make one lane the sole owner of each expert's count. + %expert_assignment_count = scf.for %assignment = [%c0 to %assignment_count step %c1](%matched_count = %c0_i32 : i32) -> (i32) { + %token0 = index.div %assignment, %configured_route_count : index + %route0 = index.rem %assignment, %configured_route_count : index + %token, %route_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %route, %route_row_stride = index.assume %route0, %bounded_route_stride [lt(%route0, %bounded_route_stride)] : index, index + %route_expert_i32 = view.load %route_view[%token, %route] : view<[%bounded_token_count]x[%bounded_route_stride]xi32> -> i32 + %matches = scalar.cmpi eq, %route_expert_i32, %expert_i32 : i32 + %match_increment = scf.select %matches, %c1_i32, %c0_i32 : i32 + %next_matched_count = scalar.addi %matched_count, %match_increment : i32 + scf.yield %next_matched_count : i32 + } + + %expert_assignment_base_i32 = kernel.workgroup.scan %expert_assignment_count {direction = forward, mode = exclusive} : i32 + + // A second pass publishes stable original assignment ordinals into the + // expert-contiguous permutation. This metadata remains cache-resident for + // practical decode batches. + %published_count = scf.for %assignment = [%c0 to %assignment_count step %c1](%matched_count = %c0_i32 : i32) -> (i32) { + %token0 = index.div %assignment, %configured_route_count : index + %route0 = index.rem %assignment, %configured_route_count : index + %token, %route_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %route, %route_row_stride = index.assume %route0, %bounded_route_stride [lt(%route0, %bounded_route_stride)] : index, index + %route_expert_i32 = view.load %route_view[%token, %route] : view<[%bounded_token_count]x[%bounded_route_stride]xi32> -> i32 + %matches = scalar.cmpi eq, %route_expert_i32, %expert_i32 : i32 + scf.if %matches { + %compact_ordinal_i32 = scalar.addi %expert_assignment_base_i32, %matched_count : i32 + %compact_ordinal0 = index.cast %compact_ordinal_i32 : i32 to index + %bounded_compact_ordinal, %bounded_assignment_count = index.assume %compact_ordinal0, %assignment_count [range(%compact_ordinal0, 0, 65535), lt(%compact_ordinal0, %assignment_count)] : index, index + %assignment_i32 = index.cast %assignment : index to i32 + view.store %assignment_i32, %assignment_view[%bounded_compact_ordinal] : i32, view<[%assignment_count]xi32> + } + %match_increment = scf.select %matches, %c1_i32, %c0_i32 : i32 + %next_matched_count = scalar.addi %matched_count, %match_increment : i32 + scf.yield %next_matched_count : i32 + } + + // Top-k route IDs are unique within a token, so one expert owns at most the + // configured token capacity. Target-selected row limits partition that + // range without changing this producer or its ABI. + %published_count0 = index.cast %published_count : i32 to index + %published_row_count, %published_token_capacity = index.assume %published_count0, %configured_token_capacity [range(%published_count0, 0, 2048), le(%published_count0, %configured_token_capacity)] : index, index + %has_assignments = index.cmp ugt, %published_row_count, %c0 : index + %at_most_schedule0 = index.cmp ule, %published_row_count, %schedule0_row_limit : index + %above_schedule0 = index.cmp ugt, %published_row_count, %schedule0_row_limit : index + %at_most_schedule1 = index.cmp ule, %published_row_count, %schedule1_row_limit : index + %above_schedule1 = index.cmp ugt, %published_row_count, %schedule1_row_limit : index + %at_most_schedule2 = index.cmp ule, %published_row_count, %schedule2_row_limit : index + %is_schedule0 = scalar.andi %has_assignments, %at_most_schedule0 : i1 + %is_schedule1 = scalar.andi %above_schedule0, %at_most_schedule1 : i1 + %is_schedule2 = scalar.andi %above_schedule1, %at_most_schedule2 : i1 + %is_schedule3 = index.cmp ugt, %published_row_count, %schedule2_row_limit : index + %schedule0_i32 = scf.select %is_schedule0, %c1_i32, %c0_i32 : i32 + %schedule1_i32 = scf.select %is_schedule1, %c1_i32, %c0_i32 : i32 + %schedule2_i32 = scf.select %is_schedule2, %c1_i32, %c0_i32 : i32 + %schedule3_i32 = scf.select %is_schedule3, %c1_i32, %c0_i32 : i32 + + %schedule0_ordinal_i32 = kernel.workgroup.scan %schedule0_i32 {direction = forward, mode = exclusive} : i32 + %schedule1_ordinal_i32 = kernel.workgroup.scan %schedule1_i32 {direction = forward, mode = exclusive} : i32 + %schedule2_ordinal_i32 = kernel.workgroup.scan %schedule2_i32 {direction = forward, mode = exclusive} : i32 + %schedule3_ordinal_i32 = kernel.workgroup.scan %schedule3_i32 {direction = forward, mode = exclusive} : i32 + %schedule0_count = kernel.workgroup.reduce %schedule0_i32 : i32 + %schedule1_count = kernel.workgroup.reduce %schedule1_i32 : i32 + %schedule2_count = kernel.workgroup.reduce %schedule2_i32 : i32 + %schedule3_count = kernel.workgroup.reduce %schedule3_i32 : i32 + %descriptor = vector.from_elements %expert_i32, %expert_assignment_base_i32, %published_count, %c0_i32 : vector<4xi32> + + // Inactive lanes use ordinal zero so every physical origin is in bounds; + // the vector masks suppress their stores. This keeps the descriptor write a + // native b128 operation without placing it inside divergent control flow. + %safe_schedule0_ordinal_i32 = scf.select %is_schedule0, %schedule0_ordinal_i32, %c0_i32 : i32 + %schedule0_ordinal0 = index.cast %safe_schedule0_ordinal_i32 : i32 to index + %schedule0_ordinal, %schedule0_queue_capacity = index.assume %schedule0_ordinal0, %configured_expert_count [range(%schedule0_ordinal0, 0, 511), lt(%schedule0_ordinal0, %configured_expert_count)] : index, index + %schedule0_mask = vector.splat %is_schedule0 : vector<4xi1> + vector.store.mask %descriptor, %queue_descriptor_view[0, %schedule0_ordinal, 0], %schedule0_mask : vector<4xi32>, view<4x[%configured_expert_count]x4xi32>, vector<4xi1> + %safe_schedule1_ordinal_i32 = scf.select %is_schedule1, %schedule1_ordinal_i32, %c0_i32 : i32 + %schedule1_ordinal0 = index.cast %safe_schedule1_ordinal_i32 : i32 to index + %schedule1_ordinal, %schedule1_queue_capacity = index.assume %schedule1_ordinal0, %configured_expert_count [range(%schedule1_ordinal0, 0, 511), lt(%schedule1_ordinal0, %configured_expert_count)] : index, index + %schedule1_mask = vector.splat %is_schedule1 : vector<4xi1> + vector.store.mask %descriptor, %queue_descriptor_view[1, %schedule1_ordinal, 0], %schedule1_mask : vector<4xi32>, view<4x[%configured_expert_count]x4xi32>, vector<4xi1> + %safe_schedule2_ordinal_i32 = scf.select %is_schedule2, %schedule2_ordinal_i32, %c0_i32 : i32 + %schedule2_ordinal0 = index.cast %safe_schedule2_ordinal_i32 : i32 to index + %schedule2_ordinal, %schedule2_queue_capacity = index.assume %schedule2_ordinal0, %configured_expert_count [range(%schedule2_ordinal0, 0, 511), lt(%schedule2_ordinal0, %configured_expert_count)] : index, index + %schedule2_mask = vector.splat %is_schedule2 : vector<4xi1> + vector.store.mask %descriptor, %queue_descriptor_view[2, %schedule2_ordinal, 0], %schedule2_mask : vector<4xi32>, view<4x[%configured_expert_count]x4xi32>, vector<4xi1> + %safe_schedule3_ordinal_i32 = scf.select %is_schedule3, %schedule3_ordinal_i32, %c0_i32 : i32 + %schedule3_ordinal0 = index.cast %safe_schedule3_ordinal_i32 : i32 to index + %schedule3_ordinal, %schedule3_queue_capacity = index.assume %schedule3_ordinal0, %configured_expert_count [range(%schedule3_ordinal0, 0, 511), lt(%schedule3_ordinal0, %configured_expert_count)] : index, index + %schedule3_mask = vector.splat %is_schedule3 : vector<4xi1> + vector.store.mask %descriptor, %queue_descriptor_view[3, %schedule3_ordinal, 0], %schedule3_mask : vector<4xi32>, view<4x[%configured_expert_count]x4xi32>, vector<4xi1> + + %is_lane_zero = index.cmp eq, %expert, %c0 : index + scf.if %is_lane_zero { + view.store %schedule0_count, %queue_count_view[0] : i32, view<4xi32> + view.store %schedule1_count, %queue_count_view[1] : i32, view<4xi32> + view.store %schedule2_count, %queue_count_view[2] : i32, view<4xi32> + view.store %schedule3_count, %queue_count_view[3] : i32, view<4xi32> + } + kernel.return +} + +// Serial specification oracle for exact differential tests. Its ownership and +// control flow deliberately differ from the parallel scan implementation: +// one workitem visits experts in order and carries all compact offsets and +// queue tails explicitly. +kernel.def target(@qwen3_moe_batched_decode_dispatch_gfx11_wave32) @qwen3_moe_build_batched_decode_expert_dispatch_reference(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index) { + %c1 = index.constant 1 : index + kernel.launch.config workgroups(%c1, %c1, %c1) workgroup_size(%c1, %c1, %c1) : index +} launch(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index, %route_ids: buffer, %assignment_ordinals: buffer, %queue_counts: buffer, %queue_descriptors: buffer) { + %configured_token_capacity0 = config.get @qwen3_moe.workload.token_capacity : index + %bounded_token_count, %configured_token_capacity = index.assume %token_count, %configured_token_capacity0 [range(%token_count, 1, 2048), le(%token_count, %configured_token_capacity0)] : index, index + %configured_route_count0 = config.get @qwen3_moe.router.route_count : index + %bounded_route_count, %configured_route_count = index.assume %route_count, %configured_route_count0 [range(%route_count, 1, 32), eq(%route_count, %configured_route_count0)] : index, index + %bounded_route_stride, %route_row_width = index.assume %route_stride, %bounded_route_count [range(%route_stride, 1, 512), le(%bounded_route_count, %route_stride)] : index, index + %configured_expert_count0 = config.get @qwen3_moe.router.expert_count : index + %bounded_expert_count, %configured_expert_count = index.assume %expert_count, %configured_expert_count0 [range(%expert_count, 32, 512), eq(%expert_count, %configured_expert_count0)] : index, index + %schedule0_row_limit0 = config.get @qwen3_moe.batched_decode.schedule0_row_limit : index + %schedule1_row_limit0 = config.get @qwen3_moe.batched_decode.schedule1_row_limit : index + %schedule2_row_limit0 = config.get @qwen3_moe.batched_decode.schedule2_row_limit : index + %schedule0_row_limit, %schedule1_row_limit, %schedule2_row_limit = index.assume %schedule0_row_limit0, %schedule1_row_limit0, %schedule2_row_limit0 [lt(%schedule0_row_limit0, %schedule1_row_limit0), lt(%schedule1_row_limit0, %schedule2_row_limit0)] : index, index, index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c0_offset = index.constant 0 : offset + %assignment_count = index.mul %bounded_token_count, %configured_route_count : index + %route_ids_noalias, %assignment_ordinals_noalias, %queue_counts_noalias, %queue_descriptors_noalias = buffer.assume.noalias %route_ids, %assignment_ordinals, %queue_counts, %queue_descriptors : buffer, buffer, buffer, buffer + %route_view = buffer.view %route_ids_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_route_stride]xi32> + %assignment_view = buffer.view %assignment_ordinals_noalias[%c0_offset] : buffer -> view<[%assignment_count]xi32> + %queue_count_view = buffer.view %queue_counts_noalias[%c0_offset] : buffer -> view<4xi32> + %queue_descriptor_view = buffer.view %queue_descriptors_noalias[%c0_offset] : buffer -> view<4x[%configured_expert_count]x4xi32> + + %final_assignment_base, %final_schedule0_count, %final_schedule1_count, %final_schedule2_count, %final_schedule3_count = scf.for %expert = [%c0 to %bounded_expert_count step %c1](%assignment_base = %c0_i32 : i32, %schedule0_count = %c0_i32 : i32, %schedule1_count = %c0_i32 : i32, %schedule2_count = %c0_i32 : i32, %schedule3_count = %c0_i32 : i32) -> (i32, i32, i32, i32, i32) { + %expert_i32 = index.cast %expert : index to i32 + %expert_assignment_count = scf.for %assignment = [%c0 to %assignment_count step %c1](%matched_count = %c0_i32 : i32) -> (i32) { + %token0 = index.div %assignment, %configured_route_count : index + %route0 = index.rem %assignment, %configured_route_count : index + %token, %route_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %route, %route_row_stride = index.assume %route0, %bounded_route_stride [lt(%route0, %bounded_route_stride)] : index, index + %route_expert_i32 = view.load %route_view[%token, %route] : view<[%bounded_token_count]x[%bounded_route_stride]xi32> -> i32 + %matches = scalar.cmpi eq, %route_expert_i32, %expert_i32 : i32 + scf.if %matches { + %compact_ordinal_i32 = scalar.addi %assignment_base, %matched_count : i32 + %compact_ordinal0 = index.cast %compact_ordinal_i32 : i32 to index + %bounded_compact_ordinal, %bounded_assignment_count = index.assume %compact_ordinal0, %assignment_count [range(%compact_ordinal0, 0, 65535), lt(%compact_ordinal0, %assignment_count)] : index, index + %assignment_i32 = index.cast %assignment : index to i32 + view.store %assignment_i32, %assignment_view[%bounded_compact_ordinal] : i32, view<[%assignment_count]xi32> + } + %match_increment = scf.select %matches, %c1_i32, %c0_i32 : i32 + %next_matched_count = scalar.addi %matched_count, %match_increment : i32 + scf.yield %next_matched_count : i32 + } + + %expert_assignment_count0 = index.cast %expert_assignment_count : i32 to index + %expert_row_count, %expert_token_capacity = index.assume %expert_assignment_count0, %configured_token_capacity [range(%expert_assignment_count0, 0, 2048), le(%expert_assignment_count0, %configured_token_capacity)] : index, index + %has_assignments = index.cmp ugt, %expert_row_count, %c0 : index + %at_most_schedule0 = index.cmp ule, %expert_row_count, %schedule0_row_limit : index + %above_schedule0 = index.cmp ugt, %expert_row_count, %schedule0_row_limit : index + %at_most_schedule1 = index.cmp ule, %expert_row_count, %schedule1_row_limit : index + %above_schedule1 = index.cmp ugt, %expert_row_count, %schedule1_row_limit : index + %at_most_schedule2 = index.cmp ule, %expert_row_count, %schedule2_row_limit : index + %is_schedule0 = scalar.andi %has_assignments, %at_most_schedule0 : i1 + %is_schedule1 = scalar.andi %above_schedule0, %at_most_schedule1 : i1 + %is_schedule2 = scalar.andi %above_schedule1, %at_most_schedule2 : i1 + %is_schedule3 = index.cmp ugt, %expert_row_count, %schedule2_row_limit : index + %descriptor = vector.from_elements %expert_i32, %assignment_base, %expert_assignment_count, %c0_i32 : vector<4xi32> + + // The scalar oracle retains its loop-carried queue tails. Masked b128 + // stores express conditional publication without divergent branch regions, + // while safe ordinal zero keeps inactive physical origins in bounds. + %safe_schedule0_count = scf.select %is_schedule0, %schedule0_count, %c0_i32 : i32 + %schedule0_ordinal0 = index.cast %safe_schedule0_count : i32 to index + %schedule0_ordinal, %schedule0_capacity = index.assume %schedule0_ordinal0, %configured_expert_count [range(%schedule0_ordinal0, 0, 511), lt(%schedule0_ordinal0, %configured_expert_count)] : index, index + %schedule0_mask = vector.splat %is_schedule0 : vector<4xi1> + vector.store.mask %descriptor, %queue_descriptor_view[0, %schedule0_ordinal, 0], %schedule0_mask : vector<4xi32>, view<4x[%configured_expert_count]x4xi32>, vector<4xi1> + %safe_schedule1_count = scf.select %is_schedule1, %schedule1_count, %c0_i32 : i32 + %schedule1_ordinal0 = index.cast %safe_schedule1_count : i32 to index + %schedule1_ordinal, %schedule1_capacity = index.assume %schedule1_ordinal0, %configured_expert_count [range(%schedule1_ordinal0, 0, 511), lt(%schedule1_ordinal0, %configured_expert_count)] : index, index + %schedule1_mask = vector.splat %is_schedule1 : vector<4xi1> + vector.store.mask %descriptor, %queue_descriptor_view[1, %schedule1_ordinal, 0], %schedule1_mask : vector<4xi32>, view<4x[%configured_expert_count]x4xi32>, vector<4xi1> + %safe_schedule2_count = scf.select %is_schedule2, %schedule2_count, %c0_i32 : i32 + %schedule2_ordinal0 = index.cast %safe_schedule2_count : i32 to index + %schedule2_ordinal, %schedule2_capacity = index.assume %schedule2_ordinal0, %configured_expert_count [range(%schedule2_ordinal0, 0, 511), lt(%schedule2_ordinal0, %configured_expert_count)] : index, index + %schedule2_mask = vector.splat %is_schedule2 : vector<4xi1> + vector.store.mask %descriptor, %queue_descriptor_view[2, %schedule2_ordinal, 0], %schedule2_mask : vector<4xi32>, view<4x[%configured_expert_count]x4xi32>, vector<4xi1> + %safe_schedule3_count = scf.select %is_schedule3, %schedule3_count, %c0_i32 : i32 + %schedule3_ordinal0 = index.cast %safe_schedule3_count : i32 to index + %schedule3_ordinal, %schedule3_capacity = index.assume %schedule3_ordinal0, %configured_expert_count [range(%schedule3_ordinal0, 0, 511), lt(%schedule3_ordinal0, %configured_expert_count)] : index, index + %schedule3_mask = vector.splat %is_schedule3 : vector<4xi1> + vector.store.mask %descriptor, %queue_descriptor_view[3, %schedule3_ordinal, 0], %schedule3_mask : vector<4xi32>, view<4x[%configured_expert_count]x4xi32>, vector<4xi1> + + %schedule0_increment = scf.select %is_schedule0, %c1_i32, %c0_i32 : i32 + %schedule1_increment = scf.select %is_schedule1, %c1_i32, %c0_i32 : i32 + %schedule2_increment = scf.select %is_schedule2, %c1_i32, %c0_i32 : i32 + %schedule3_increment = scf.select %is_schedule3, %c1_i32, %c0_i32 : i32 + %next_assignment_base = scalar.addi %assignment_base, %expert_assignment_count : i32 + %next_schedule0_count = scalar.addi %schedule0_count, %schedule0_increment : i32 + %next_schedule1_count = scalar.addi %schedule1_count, %schedule1_increment : i32 + %next_schedule2_count = scalar.addi %schedule2_count, %schedule2_increment : i32 + %next_schedule3_count = scalar.addi %schedule3_count, %schedule3_increment : i32 + scf.yield %next_assignment_base, %next_schedule0_count, %next_schedule1_count, %next_schedule2_count, %next_schedule3_count : i32, i32, i32, i32, i32 + } + + view.store %final_schedule0_count, %queue_count_view[0] : i32, view<4xi32> + view.store %final_schedule1_count, %queue_count_view[1] : i32, view<4xi32> + view.store %final_schedule2_count, %queue_count_view[2] : i32, view<4xi32> + view.store %final_schedule3_count, %queue_count_view[3] : i32, view<4xi32> + kernel.return +} + +check.case public @qwen3_moe_batched_decode_expert_dispatch_singleton_case { + %token_count = check.literal value(16) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.iota offset(0) step(1) period(128) : tensor<16x8xi32> + %actual_assignments = check.generate.fill value(-1) : tensor<128xi32> + %actual_counts = check.generate.fill value(-1) : tensor<4xi32> + %actual_descriptors = check.generate.fill value(-1) : tensor<4x128x4xi32> + %expected_assignments = check.generate.fill value(-1) : tensor<128xi32> + %expected_counts = check.generate.fill value(-1) : tensor<4xi32> + %expected_descriptors = check.generate.fill value(-1) : tensor<4x128x4xi32> + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %actual_assignments, %actual_counts, %actual_descriptors) : [index, index, index, index](index, index, index, index, tensor<16x8xi32>, tensor<128xi32>, tensor<4xi32>, tensor<4x128x4xi32>) + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch_reference[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expected_assignments, %expected_counts, %expected_descriptors) : [index, index, index, index](index, index, index, index, tensor<16x8xi32>, tensor<128xi32>, tensor<4xi32>, tensor<4x128x4xi32>) + check.expect.equal actual(%actual_assignments) expected(%expected_assignments) : tensor<128xi32> + check.expect.equal actual(%actual_counts) expected(%expected_counts) : tensor<4xi32> + check.expect.equal actual(%actual_descriptors) expected(%expected_descriptors) : tensor<4x128x4xi32> + check.return +} + +check.case public @qwen3_moe_batched_decode_expert_dispatch_pair_case { + %token_count = check.literal value(16) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.iota offset(0) step(1) period(64) : tensor<16x8xi32> + %actual_assignments = check.generate.fill value(-1) : tensor<128xi32> + %actual_counts = check.generate.fill value(-1) : tensor<4xi32> + %actual_descriptors = check.generate.fill value(-1) : tensor<4x128x4xi32> + %expected_assignments = check.generate.fill value(-1) : tensor<128xi32> + %expected_counts = check.generate.fill value(-1) : tensor<4xi32> + %expected_descriptors = check.generate.fill value(-1) : tensor<4x128x4xi32> + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %actual_assignments, %actual_counts, %actual_descriptors) : [index, index, index, index](index, index, index, index, tensor<16x8xi32>, tensor<128xi32>, tensor<4xi32>, tensor<4x128x4xi32>) + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch_reference[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expected_assignments, %expected_counts, %expected_descriptors) : [index, index, index, index](index, index, index, index, tensor<16x8xi32>, tensor<128xi32>, tensor<4xi32>, tensor<4x128x4xi32>) + check.expect.equal actual(%actual_assignments) expected(%expected_assignments) : tensor<128xi32> + check.expect.equal actual(%actual_counts) expected(%expected_counts) : tensor<4xi32> + check.expect.equal actual(%actual_descriptors) expected(%expected_descriptors) : tensor<4x128x4xi32> + check.return +} + +check.case public @qwen3_moe_batched_decode_expert_dispatch_three_row_case { + %token_count = check.literal value(16) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.iota offset(0) step(1) period(43) : tensor<16x8xi32> + %actual_assignments = check.generate.fill value(-1) : tensor<128xi32> + %actual_counts = check.generate.fill value(-1) : tensor<4xi32> + %actual_descriptors = check.generate.fill value(-1) : tensor<4x128x4xi32> + %expected_assignments = check.generate.fill value(-1) : tensor<128xi32> + %expected_counts = check.generate.fill value(-1) : tensor<4xi32> + %expected_descriptors = check.generate.fill value(-1) : tensor<4x128x4xi32> + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %actual_assignments, %actual_counts, %actual_descriptors) : [index, index, index, index](index, index, index, index, tensor<16x8xi32>, tensor<128xi32>, tensor<4xi32>, tensor<4x128x4xi32>) + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch_reference[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expected_assignments, %expected_counts, %expected_descriptors) : [index, index, index, index](index, index, index, index, tensor<16x8xi32>, tensor<128xi32>, tensor<4xi32>, tensor<4x128x4xi32>) + check.expect.equal actual(%actual_assignments) expected(%expected_assignments) : tensor<128xi32> + check.expect.equal actual(%actual_counts) expected(%expected_counts) : tensor<4xi32> + check.expect.equal actual(%actual_descriptors) expected(%expected_descriptors) : tensor<4x128x4xi32> + check.return +} + +check.case public @qwen3_moe_batched_decode_expert_dispatch_four_row_case { + %token_count = check.literal value(16) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.iota offset(0) step(1) period(32) : tensor<16x8xi32> + %actual_assignments = check.generate.fill value(-1) : tensor<128xi32> + %actual_counts = check.generate.fill value(-1) : tensor<4xi32> + %actual_descriptors = check.generate.fill value(-1) : tensor<4x128x4xi32> + %expected_assignments = check.generate.fill value(-1) : tensor<128xi32> + %expected_counts = check.generate.fill value(-1) : tensor<4xi32> + %expected_descriptors = check.generate.fill value(-1) : tensor<4x128x4xi32> + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %actual_assignments, %actual_counts, %actual_descriptors) : [index, index, index, index](index, index, index, index, tensor<16x8xi32>, tensor<128xi32>, tensor<4xi32>, tensor<4x128x4xi32>) + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch_reference[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expected_assignments, %expected_counts, %expected_descriptors) : [index, index, index, index](index, index, index, index, tensor<16x8xi32>, tensor<128xi32>, tensor<4xi32>, tensor<4x128x4xi32>) + check.expect.equal actual(%actual_assignments) expected(%expected_assignments) : tensor<128xi32> + check.expect.equal actual(%actual_counts) expected(%expected_counts) : tensor<4xi32> + check.expect.equal actual(%actual_descriptors) expected(%expected_descriptors) : tensor<4x128x4xi32> + check.return +} + +check.case public @qwen3_moe_batched_decode_expert_dispatch_five_to_six_row_case { + %token_count = check.literal value(16) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.iota offset(0) step(1) period(25) : tensor<16x8xi32> + %actual_assignments = check.generate.fill value(-1) : tensor<128xi32> + %actual_counts = check.generate.fill value(-1) : tensor<4xi32> + %actual_descriptors = check.generate.fill value(-1) : tensor<4x128x4xi32> + %expected_assignments = check.generate.fill value(-1) : tensor<128xi32> + %expected_counts = check.generate.fill value(-1) : tensor<4xi32> + %expected_descriptors = check.generate.fill value(-1) : tensor<4x128x4xi32> + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %actual_assignments, %actual_counts, %actual_descriptors) : [index, index, index, index](index, index, index, index, tensor<16x8xi32>, tensor<128xi32>, tensor<4xi32>, tensor<4x128x4xi32>) + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch_reference[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expected_assignments, %expected_counts, %expected_descriptors) : [index, index, index, index](index, index, index, index, tensor<16x8xi32>, tensor<128xi32>, tensor<4xi32>, tensor<4x128x4xi32>) + check.expect.equal actual(%actual_assignments) expected(%expected_assignments) : tensor<128xi32> + check.expect.equal actual(%actual_counts) expected(%expected_counts) : tensor<4xi32> + check.expect.equal actual(%actual_descriptors) expected(%expected_descriptors) : tensor<4x128x4xi32> + check.return +} + +check.case public @qwen3_moe_batched_decode_expert_dispatch_fourteen_to_fifteen_row_case { + %token_count = check.literal value(16) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.iota offset(0) step(1) period(9) : tensor<16x8xi32> + %actual_assignments = check.generate.fill value(-1) : tensor<128xi32> + %actual_counts = check.generate.fill value(-1) : tensor<4xi32> + %actual_descriptors = check.generate.fill value(-1) : tensor<4x128x4xi32> + %expected_assignments = check.generate.fill value(-1) : tensor<128xi32> + %expected_counts = check.generate.fill value(-1) : tensor<4xi32> + %expected_descriptors = check.generate.fill value(-1) : tensor<4x128x4xi32> + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %actual_assignments, %actual_counts, %actual_descriptors) : [index, index, index, index](index, index, index, index, tensor<16x8xi32>, tensor<128xi32>, tensor<4xi32>, tensor<4x128x4xi32>) + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch_reference[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expected_assignments, %expected_counts, %expected_descriptors) : [index, index, index, index](index, index, index, index, tensor<16x8xi32>, tensor<128xi32>, tensor<4xi32>, tensor<4x128x4xi32>) + check.expect.equal actual(%actual_assignments) expected(%expected_assignments) : tensor<128xi32> + check.expect.equal actual(%actual_counts) expected(%expected_counts) : tensor<4xi32> + check.expect.equal actual(%actual_descriptors) expected(%expected_descriptors) : tensor<4x128x4xi32> + check.return +} + +check.case public @qwen3_moe_batched_decode_expert_dispatch_coherent_case { + %token_count = check.literal value(16) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.iota offset(0) step(1) period(8) : tensor<16x8xi32> + %actual_assignments = check.generate.fill value(-1) : tensor<128xi32> + %actual_counts = check.generate.fill value(-1) : tensor<4xi32> + %actual_descriptors = check.generate.fill value(-1) : tensor<4x128x4xi32> + %expected_assignments = check.generate.fill value(-1) : tensor<128xi32> + %expected_counts = check.generate.fill value(-1) : tensor<4xi32> + %expected_descriptors = check.generate.fill value(-1) : tensor<4x128x4xi32> + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %actual_assignments, %actual_counts, %actual_descriptors) : [index, index, index, index](index, index, index, index, tensor<16x8xi32>, tensor<128xi32>, tensor<4xi32>, tensor<4x128x4xi32>) + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch_reference[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expected_assignments, %expected_counts, %expected_descriptors) : [index, index, index, index](index, index, index, index, tensor<16x8xi32>, tensor<128xi32>, tensor<4xi32>, tensor<4x128x4xi32>) + check.expect.equal actual(%actual_assignments) expected(%expected_assignments) : tensor<128xi32> + check.expect.equal actual(%actual_counts) expected(%expected_counts) : tensor<4xi32> + check.expect.equal actual(%actual_descriptors) expected(%expected_descriptors) : tensor<4x128x4xi32> + check.return +} + +// A second model geometry proves that expert, route, batch, descriptor, and +// schedule boundaries are configuration rather than Qwen3-30B constants. +check.case public @qwen3_moe_batched_decode_expert_dispatch_configurable_geometry_case { + %token_count = check.literal value(8) : index + %route_count = check.literal value(4) : index + %route_stride = check.literal value(4) : index + %expert_count = check.literal value(32) : index + %route_ids = check.generate.iota offset(0) step(1) period(13) : tensor<8x4xi32> + %actual_assignments = check.generate.fill value(-1) : tensor<32xi32> + %actual_counts = check.generate.fill value(-1) : tensor<4xi32> + %actual_descriptors = check.generate.fill value(-1) : tensor<4x32x4xi32> + %expected_assignments = check.generate.fill value(-1) : tensor<32xi32> + %expected_counts = check.generate.fill value(-1) : tensor<4xi32> + %expected_descriptors = check.generate.fill value(-1) : tensor<4x32x4xi32> + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %actual_assignments, %actual_counts, %actual_descriptors) : [index, index, index, index](index, index, index, index, tensor<8x4xi32>, tensor<32xi32>, tensor<4xi32>, tensor<4x32x4xi32>) + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch_reference[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expected_assignments, %expected_counts, %expected_descriptors) : [index, index, index, index](index, index, index, index, tensor<8x4xi32>, tensor<32xi32>, tensor<4xi32>, tensor<4x32x4xi32>) + check.expect.equal actual(%actual_assignments) expected(%expected_assignments) : tensor<32xi32> + check.expect.equal actual(%actual_counts) expected(%expected_counts) : tensor<4xi32> + check.expect.equal actual(%actual_descriptors) expected(%expected_descriptors) : tensor<4x32x4xi32> + check.return +} + +check.case public @qwen3_moe_batched_decode_expert_dispatch_diverse_benchmark_case { + %token_count = check.literal value(16) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.iota offset(0) step(1) period(128) : tensor<16x8xi32> + %assignment_ordinals = check.generate.fill value(-1) : tensor<128xi32> + %queue_counts = check.generate.fill value(-1) : tensor<4xi32> + %queue_descriptors = check.generate.fill value(-1) : tensor<4x128x4xi32> + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %assignment_ordinals, %queue_counts, %queue_descriptors) : [index, index, index, index](index, index, index, index, tensor<16x8xi32>, tensor<128xi32>, tensor<4xi32>, tensor<4x128x4xi32>) + check.return +} + +check.case public @qwen3_moe_batched_decode_expert_dispatch_coherent_benchmark_case { + %token_count = check.literal value(16) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.iota offset(0) step(1) period(8) : tensor<16x8xi32> + %assignment_ordinals = check.generate.fill value(-1) : tensor<128xi32> + %queue_counts = check.generate.fill value(-1) : tensor<4xi32> + %queue_descriptors = check.generate.fill value(-1) : tensor<4x128x4xi32> + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %assignment_ordinals, %queue_counts, %queue_descriptors) : [index, index, index, index](index, index, index, index, tensor<16x8xi32>, tensor<128xi32>, tensor<4xi32>, tensor<4x128x4xi32>) + check.return +} + +check.case public @qwen3_moe_batched_decode_expert_dispatch_configurable_benchmark_case { + %token_count = check.literal value(8) : index + %route_count = check.literal value(4) : index + %route_stride = check.literal value(4) : index + %expert_count = check.literal value(32) : index + %route_ids = check.generate.iota offset(0) step(1) period(13) : tensor<8x4xi32> + %assignment_ordinals = check.generate.fill value(-1) : tensor<32xi32> + %queue_counts = check.generate.fill value(-1) : tensor<4xi32> + %queue_descriptors = check.generate.fill value(-1) : tensor<4x32x4xi32> + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %assignment_ordinals, %queue_counts, %queue_descriptors) : [index, index, index, index](index, index, index, index, tensor<8x4xi32>, tensor<32xi32>, tensor<4xi32>, tensor<4x32x4xi32>) + check.return +} + +check.benchmark<@qwen3_moe_batched_decode_expert_dispatch_diverse_benchmark_case> @qwen3_moe_batched_decode_expert_dispatch_diverse + +check.benchmark<@qwen3_moe_batched_decode_expert_dispatch_coherent_benchmark_case> @qwen3_moe_batched_decode_expert_dispatch_coherent + +check.benchmark<@qwen3_moe_batched_decode_expert_dispatch_configurable_benchmark_case> @qwen3_moe_batched_decode_expert_dispatch_configurable diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/batched_decode_gate_up_q4k.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/batched_decode_gate_up_q4k.loom new file mode 100644 index 000000000000..660a06551399 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/batched_decode_gate_up_q4k.loom @@ -0,0 +1,315 @@ +// Expert-grouped batched-decode gate/up consumers for raw GGUF Q4_K weights. +// +// The assignment producer owns routing and publishes expert-contiguous +// assignment ordinals plus aligned descriptors. This file owns row-count +// schedules that consume those descriptors without inspecting route IDs. The +// first provider is exact M_e=2: each wave reuses one expert/channel weight +// row across two independently quantized activation rows and scatters the +// fused SwiGLU results back to their original [token][route] assignments. +config.decl @qwen3_moe.router.expert_count : %value: index where [range(%value, 32, 512), mul(%value, 32)] + +config.decl @qwen3_moe.router.route_count : %value: index where [range(%value, 1, 32)] + +config.decl @qwen3_moe.routed_gate_up.input_size : %value: index where [range(%value, 512, 32768), mul(%value, 512)] + +config.decl @qwen3_moe.routed_gate_up.expert_count : %value: index where [range(%value, 1, 512)] + +config.decl @qwen3_moe.routed_gate_up.route_count : %value: index where [range(%value, 1, 8)] + +config.decl @qwen3_moe.routed_gate_up.output_size : %value: index where [range(%value, 1, 4096)] + +config.decl @qwen3_moe.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +// Static fallback capacity for the exact-two-row queue. Device-side indirect +// launch replaces this upper bound with the producer's exact queue count. +config.decl @qwen3_moe.batched_decode.rows2_descriptor_capacity : %value: index where [range(%value, 1, 512)] + +kernel.decl @ggml_quantize_q8_1_x4_f32(%token_count$8: index, %input_size$9: index) launch(%token_count$10: index, %input_size$11: index, %input: buffer, %output: buffer) + +func.decl @qwen3_moe_q4k_chunk_pair_global(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group_pair: index, %q4_half: index, %header_words: vector<4xi32>) -> (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) + +func.decl @qwen3_moe_q4k_q8_1_dot(%q4_values: vector<16xi8>, %d_scale: f32, %dmin_scale: f32, %q8_values: vector<16xi8>, %q8_d: f32, %q8_s: f32) -> (f32) + +func.decl @qwen3_moe_unpack_batched_decode_expert_descriptor(%descriptor: vector<4xi32>) -> (index, index, index) + +kernel.decl @qwen3_moe_build_batched_decode_expert_dispatch(%token_count$37: index, %route_count$38: index, %route_stride$39: index, %expert_count$40: index) launch(%token_count$41: index, %route_count$42: index, %route_stride$43: index, %expert_count$44: index, %route_ids: buffer, %assignment_ordinals: buffer, %queue_counts: buffer, %queue_descriptors: buffer) + +kernel.decl @qwen3_moe_routed_gate_up_swiglu_q4k_q8(%token_count$49: index, %route_count$50: index, %route_stride$51: index, %expert_count$52: index, %output_size$53: index) launch(%token_count$54: index, %route_count$55: index, %route_stride$56: index, %expert_count$57: index, %output_size$58: index, %q8_input: buffer, %route_ids: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) + +// Contracts one expert/channel Q4_K row against two independently quantized +// activation rows. Splitting gate and up into separate contractions keeps each +// recurrence small while preserving the important invariant: every weight +// byte loaded by this helper contributes to both rows. +func.def inline @qwen3_moe_batched_decode_q4k_rows2_lane(%input_size: index, %weight: buffer, %q8_input: buffer, %weight_row_byte_base: offset, %q8_row0_byte_base: offset, %q8_row1_byte_base: offset, %lane: index) -> (f32, f32) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %c1023 = index.constant 1023 : index + %c1024 = index.constant 1024 : index + %q4_block_bytes = index.constant 144 : offset + %q8_group_bytes = index.constant 144 : offset + %q8_payload_byte_add = index.constant 16 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %q4_block_count = index.div %input_size, %c256 : index + %q8_group_count = index.div %input_size, %c128 : index + %padded_input_size = index.add %input_size, %c1023 : index + %iteration_count = index.div %padded_input_size, %c1024 : index + %q4_group_pair0 = index.div %lane, %c2 : index + %q4_group_pair1 = index.rem %q4_group_pair0, %c4 : index + %q4_group_pair = index.assume %q4_group_pair1 [range(%q4_group_pair1, 0, 3)] : index + %q4_half0 = index.rem %lane, %c2 : index + %q4_half = index.assume %q4_half0 [range(%q4_half0, 0, 1)] : index + %lane_q4_block = index.div %lane, %c8 : index + %q8_group_in_block0 = index.div %q4_group_pair, %c2 : index + %q8_group_in_block = index.assume %q8_group_in_block0 [range(%q8_group_in_block0, 0, 1)] : index + %pair_in_q8_group0 = index.rem %q4_group_pair, %c2 : index + %pair_in_q8_group = index.assume %pair_in_q8_group0 [range(%pair_in_q8_group0, 0, 1)] : index + %q8_low_inner_block0 = index.mul %pair_in_q8_group, %c2 : index + %q8_low_inner_block = index.assume %q8_low_inner_block0 [range(%q8_low_inner_block0, 0, 2)] : index + %q8_high_inner_block0 = index.add %q8_low_inner_block, %c1 : index + %q8_high_inner_block = index.assume %q8_high_inner_block0 [range(%q8_high_inner_block0, 1, 3)] : index + %q8_half_word_add = index.mul %q4_half, %c4 : index + %q8_low_inner_word_base = index.mul %q8_low_inner_block, %c8 : index + %q8_low_word_index0 = index.add %q8_low_inner_word_base, %q8_half_word_add : index + %q8_low_word_index = index.assume %q8_low_word_index0 [range(%q8_low_word_index0, 0, 20)] : index + %q8_high_inner_word_base = index.mul %q8_high_inner_block, %c8 : index + %q8_high_word_index0 = index.add %q8_high_inner_word_base, %q8_half_word_add : index + %q8_high_word_index = index.assume %q8_high_word_index0 [range(%q8_high_word_index0, 8, 28)] : index + %q8_low_ds_index0 = index.mul %q8_low_inner_block, %c2 : index + %q8_low_ds_index = index.assume %q8_low_ds_index0 [range(%q8_low_ds_index0, 0, 4)] : index + + %row0_acc, %row1_acc = scf.for %iteration = [%c0 to %iteration_count step %c1](%row0_iter = %c0_f32 : f32, %row1_iter = %c0_f32 : f32) -> (f32, f32) unroll { + %iteration_q4_block = index.mul %iteration, %c4 : index + %q4_block0 = index.add %iteration_q4_block, %lane_q4_block : index + %valid_q4_block = index.cmp ult, %q4_block0, %q4_block_count : index + %row0_contribution, %row1_contribution = scf.if %valid_q4_block -> (f32, f32) { + %q4_block, %bounded_q4_block_count = index.assume %q4_block0, %q4_block_count [lt(%q4_block0, %q4_block_count)] : index, index + %q8_block_group_base = index.mul %q4_block, %c2 : index + %q8_group0 = index.add %q8_block_group_base, %q8_group_in_block : index + %q8_group, %bounded_q8_group_count = index.assume %q8_group0, %q8_group_count [lt(%q8_group0, %q8_group_count)] : index, index + %q8_group_byte_add = index.scale %q8_group, %q8_group_bytes : index, offset -> offset + %q8_group0_byte_base = index.add %q8_row0_byte_base, %q8_group_byte_add : offset + %q8_group1_byte_base = index.add %q8_row1_byte_base, %q8_group_byte_add : offset + %q8_payload0_byte_base = index.add %q8_group0_byte_base, %q8_payload_byte_add : offset + %q8_payload1_byte_base = index.add %q8_group1_byte_base, %q8_payload_byte_add : offset + %q8_ds0_view = buffer.view %q8_input[%q8_group0_byte_base] : buffer -> view<8xf16> + %q8_ds1_view = buffer.view %q8_input[%q8_group1_byte_base] : buffer -> view<8xf16> + %q8_words0_view = buffer.view %q8_input[%q8_payload0_byte_base] : buffer -> view<32xi32> + %q8_words1_view = buffer.view %q8_input[%q8_payload1_byte_base] : buffer -> view<32xi32> + %q8_ds0 = vector.load %q8_ds0_view[%q8_low_ds_index] : view<8xf16> -> vector<4xf16> + %q8_ds1 = vector.load %q8_ds1_view[%q8_low_ds_index] : view<8xf16> -> vector<4xf16> + %q8_low_d0_f16 = vector.extract %q8_ds0[0] : vector<4xf16> -> f16 + %q8_low_s0_f16 = vector.extract %q8_ds0[1] : vector<4xf16> -> f16 + %q8_high_d0_f16 = vector.extract %q8_ds0[2] : vector<4xf16> -> f16 + %q8_high_s0_f16 = vector.extract %q8_ds0[3] : vector<4xf16> -> f16 + %q8_low_d1_f16 = vector.extract %q8_ds1[0] : vector<4xf16> -> f16 + %q8_low_s1_f16 = vector.extract %q8_ds1[1] : vector<4xf16> -> f16 + %q8_high_d1_f16 = vector.extract %q8_ds1[2] : vector<4xf16> -> f16 + %q8_high_s1_f16 = vector.extract %q8_ds1[3] : vector<4xf16> -> f16 + %q8_low_d0 = scalar.extf %q8_low_d0_f16 : f16 to f32 + %q8_low_s0 = scalar.extf %q8_low_s0_f16 : f16 to f32 + %q8_high_d0 = scalar.extf %q8_high_d0_f16 : f16 to f32 + %q8_high_s0 = scalar.extf %q8_high_s0_f16 : f16 to f32 + %q8_low_d1 = scalar.extf %q8_low_d1_f16 : f16 to f32 + %q8_low_s1 = scalar.extf %q8_low_s1_f16 : f16 to f32 + %q8_high_d1 = scalar.extf %q8_high_d1_f16 : f16 to f32 + %q8_high_s1 = scalar.extf %q8_high_s1_f16 : f16 to f32 + %q8_low_words0 = vector.load %q8_words0_view[%q8_low_word_index] : view<32xi32> -> vector<4xi32> + %q8_high_words0 = vector.load %q8_words0_view[%q8_high_word_index] : view<32xi32> -> vector<4xi32> + %q8_low_words1 = vector.load %q8_words1_view[%q8_low_word_index] : view<32xi32> -> vector<4xi32> + %q8_high_words1 = vector.load %q8_words1_view[%q8_high_word_index] : view<32xi32> -> vector<4xi32> + %q8_low_values0 = vector.bitcast %q8_low_words0 : vector<4xi32> to vector<16xi8> + %q8_high_values0 = vector.bitcast %q8_high_words0 : vector<4xi32> to vector<16xi8> + %q8_low_values1 = vector.bitcast %q8_low_words1 : vector<4xi32> to vector<16xi8> + %q8_high_values1 = vector.bitcast %q8_high_words1 : vector<4xi32> to vector<16xi8> + %q4_block_byte_add = index.scale %q4_block, %q4_block_bytes : index, offset -> offset + %q4_block_byte_base = index.add %weight_row_byte_base, %q4_block_byte_add : offset + %header_view = buffer.view %weight[%q4_block_byte_base] : buffer -> view<4xi32> + %header_words = vector.load %header_view[0] : view<4xi32> -> vector<4xi32> + %q4_low, %low_d_scale, %low_dmin_scale, %q4_high, %high_d_scale, %high_dmin_scale = func.call @qwen3_moe_q4k_chunk_pair_global(%weight, %weight_row_byte_base, %q4_block, %q4_group_pair, %q4_half, %header_words) : (buffer, offset, index, index, index, vector<4xi32>) -> (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) + %row0_low = func.call @qwen3_moe_q4k_q8_1_dot(%q4_low, %low_d_scale, %low_dmin_scale, %q8_low_values0, %q8_low_d0, %q8_low_s0) : (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) -> (f32) + %row0_high = func.call @qwen3_moe_q4k_q8_1_dot(%q4_high, %high_d_scale, %high_dmin_scale, %q8_high_values0, %q8_high_d0, %q8_high_s0) : (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) -> (f32) + %row1_low = func.call @qwen3_moe_q4k_q8_1_dot(%q4_low, %low_d_scale, %low_dmin_scale, %q8_low_values1, %q8_low_d1, %q8_low_s1) : (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) -> (f32) + %row1_high = func.call @qwen3_moe_q4k_q8_1_dot(%q4_high, %high_d_scale, %high_dmin_scale, %q8_high_values1, %q8_high_d1, %q8_high_s1) : (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) -> (f32) + %row0_pair = scalar.addf %row0_low, %row0_high : f32 + %row1_pair = scalar.addf %row1_low, %row1_high : f32 + scf.yield %row0_pair, %row1_pair : f32, f32 + } else { + scf.yield %c0_f32, %c0_f32 : f32, f32 + } + %row0_next = scalar.addf %row0_iter, %row0_contribution : f32 + %row1_next = scalar.addf %row1_iter, %row1_contribution : f32 + scf.yield %row0_next, %row1_next : f32, f32 + } + func.return %row0_acc, %row1_acc : f32, f32 +} + +amdgpu.target @qwen3_moe_batched_decode_gate_up_gfx11_wave32 {subgroup_size = 32} + +// Processes the exact-two-row descriptor queue. The static launch covers a +// config-derived safe descriptor capacity and guards inactive workgroups. A +// composed device launcher instead supplies the exact descriptor count as +// both launch geometry and an ABI fact. +kernel.def target(@qwen3_moe_batched_decode_gate_up_gfx11_wave32) @qwen3_moe_batched_decode_gate_up_q4k_rows2(%descriptor_count: index, %queue_ordinal: index, %token_count: index, %route_count: index, %expert_count: index, %output_size: index) { + %configured_output_size = config.get @qwen3_moe.routed_gate_up.output_size : index + %descriptor_capacity = config.get @qwen3_moe.batched_decode.rows2_descriptor_capacity : index + %c1 = index.constant 1 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %workgroup_size = index.constant 128 : index + %padded_output_size = index.add %configured_output_size, %c3 : index + %channel_workgroup_count = index.div %padded_output_size, %c4 : index + kernel.launch.config workgroups(%channel_workgroup_count, %descriptor_capacity, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%descriptor_count: index, %queue_ordinal: index, %token_count: index, %route_count: index, %expert_count: index, %output_size: index, %queue_descriptors: buffer, %assignment_ordinals: buffer, %q8_input: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) { + %configured_token_capacity0 = config.get @qwen3_moe.workload.token_capacity : index + %bounded_token_count, %configured_token_capacity = index.assume %token_count, %configured_token_capacity0 [range(%token_count, 1, 2048), eq(%token_count, %configured_token_capacity0)] : index, index + %configured_route_count0 = config.get @qwen3_moe.router.route_count : index + %bounded_route_count, %configured_route_count = index.assume %route_count, %configured_route_count0 [range(%route_count, 1, 32), eq(%route_count, %configured_route_count0)] : index, index + %configured_expert_count0 = config.get @qwen3_moe.router.expert_count : index + %bounded_expert_count, %configured_expert_count = index.assume %expert_count, %configured_expert_count0 [range(%expert_count, 32, 512), eq(%expert_count, %configured_expert_count0)] : index, index + %configured_weight_expert_count0 = config.get @qwen3_moe.routed_gate_up.expert_count : index + %configured_weight_expert_count, %model_expert_count = index.assume %configured_weight_expert_count0, %configured_expert_count [eq(%configured_weight_expert_count0, %configured_expert_count)] : index, index + %configured_output_size0 = config.get @qwen3_moe.routed_gate_up.output_size : index + %bounded_output_size, %configured_output_size = index.assume %output_size, %configured_output_size0 [range(%output_size, 1, 4096), eq(%output_size, %configured_output_size0)] : index, index + %bounded_queue_ordinal = index.assume %queue_ordinal [range(%queue_ordinal, 0, 3)] : index + %input_size = config.get @qwen3_moe.routed_gate_up.input_size : index + %descriptor_ordinal0 = kernel.workgroup.id : index + %channel_workgroup = kernel.workgroup.id : index + %subgroup = kernel.subgroup.id : index + %lane0 = kernel.subgroup.lane.id : index + %lane = index.assume %lane0 [range(%lane0, 0, 31)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %c1_byte = index.constant 1 : offset + %q4_block_bytes = index.constant 144 : index + %c0_i32 = scalar.constant 0 : i32 + %c0_offset = index.constant 0 : offset + %assignment_count = index.mul %bounded_token_count, %configured_route_count : index + %assignment_capacity = index.mul %configured_token_capacity, %configured_route_count : index + %assignment_limited_descriptor_capacity0 = index.div %assignment_capacity, %c2 : index + %descriptor_capacity0 = index.min %configured_expert_count, %assignment_limited_descriptor_capacity0 : index + %configured_descriptor_capacity0 = config.get @qwen3_moe.batched_decode.rows2_descriptor_capacity : index + %descriptor_capacity, %configured_descriptor_capacity = index.assume %descriptor_capacity0, %configured_descriptor_capacity0 [range(%descriptor_capacity0, 1, 512), eq(%descriptor_capacity0, %configured_descriptor_capacity0)] : index, index + %bounded_descriptor_count, %queue_descriptor_capacity = index.assume %descriptor_count, %configured_descriptor_capacity [range(%descriptor_count, 1, 512), le(%descriptor_count, %configured_descriptor_capacity)] : index, index + %valid_descriptor = index.cmp ult, %descriptor_ordinal0, %bounded_descriptor_count : index + %safe_descriptor_ordinal0 = scf.select %valid_descriptor, %descriptor_ordinal0, %c0 : index + %descriptor_ordinal, %launch_descriptor_count = index.assume %safe_descriptor_ordinal0, %bounded_descriptor_count [lt(%safe_descriptor_ordinal0, %bounded_descriptor_count)] : index, index + %q4_block_count = index.div %input_size, %c256 : index + %weight_row_bytes = index.mul %q4_block_count, %q4_block_bytes : index + %weight_expert_bytes = index.mul %bounded_output_size, %weight_row_bytes : index + %q8_group_count = index.div %input_size, %c128 : index + %q8_bytes_per_token = index.mul %q8_group_count, %q4_block_bytes : index + %queue_descriptors_noalias, %assignment_ordinals_noalias, %q8_noalias, %gate_noalias, %up_noalias, %output_noalias = buffer.assume.noalias %queue_descriptors, %assignment_ordinals, %q8_input, %gate_weight, %up_weight, %output : buffer, buffer, buffer, buffer, buffer, buffer + %queue_descriptor_view = buffer.view %queue_descriptors_noalias[%c0_offset] : buffer -> view<4x[%configured_expert_count]x4xi32> + %assignment_view = buffer.view %assignment_ordinals_noalias[%c0_offset] : buffer -> view<[%assignment_count]xi32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%assignment_count]x[%bounded_output_size]xf32> + %descriptor = vector.load %queue_descriptor_view[%bounded_queue_ordinal, %descriptor_ordinal, 0] : view<4x[%configured_expert_count]x4xi32> -> vector<4xi32> + %expert, %assignment_base0, %row_count0 = func.call @qwen3_moe_unpack_batched_decode_expert_descriptor(%descriptor) : (vector<4xi32>) -> (index, index, index) + %assignment_base, %consumer_assignment_count = index.assume %assignment_base0, %assignment_count [lt(%assignment_base0, %assignment_count)] : index, index + %row_count = index.assume %row_count0 [range(%row_count0, 2, 2)] : index + %assignment1_ordinal0 = index.add %assignment_base, %c1 : index + %assignment1_ordinal, %pair_assignment_count = index.assume %assignment1_ordinal0, %assignment_count [lt(%assignment1_ordinal0, %assignment_count)] : index, index + %assignment0_i32 = view.load %assignment_view[%assignment_base] : view<[%assignment_count]xi32> -> i32 + %assignment1_i32 = view.load %assignment_view[%assignment1_ordinal] : view<[%assignment_count]xi32> -> i32 + %assignment0_index0 = index.cast %assignment0_i32 : i32 to index + %assignment1_index0 = index.cast %assignment1_i32 : i32 to index + %assignment0, %assignment1, %routed_assignment_count = index.assume %assignment0_index0, %assignment1_index0, %assignment_count [range(%assignment0_index0, 0, 65535), range(%assignment1_index0, 0, 65535), lt(%assignment0_index0, %assignment_count), lt(%assignment1_index0, %assignment_count)] : index, index, index + %token0_index0 = index.div %assignment0, %configured_route_count : index + %token1_index0 = index.div %assignment1, %configured_route_count : index + %token0, %token1, %q8_token_count = index.assume %token0_index0, %token1_index0, %bounded_token_count [lt(%token0_index0, %bounded_token_count), lt(%token1_index0, %bounded_token_count)] : index, index, index + %channel_base = index.mul %channel_workgroup, %c4 : index + %channel0 = index.add %channel_base, %subgroup : index + %valid_channel = index.cmp ult, %channel0, %bounded_output_size : index + %safe_channel0 = scf.select %valid_channel, %channel0, %c0 : index + %channel, %weight_output_size = index.assume %safe_channel0, %bounded_output_size [lt(%safe_channel0, %bounded_output_size)] : index, index + %expert_byte_base = index.mul %expert, %weight_expert_bytes : index + %channel_byte_add = index.mul %channel, %weight_row_bytes : index + %row_byte_index = index.add %expert_byte_base, %channel_byte_add : index + %row_byte_base = index.scale %row_byte_index, %c1_byte : index, offset -> offset + %q8_token0_byte_index = index.mul %token0, %q8_bytes_per_token : index + %q8_token1_byte_index = index.mul %token1, %q8_bytes_per_token : index + %q8_token0_byte_base = index.scale %q8_token0_byte_index, %c1_byte : index, offset -> offset + %q8_token1_byte_base = index.scale %q8_token1_byte_index, %c1_byte : index, offset -> offset + %gate0_acc, %gate1_acc = func.call @qwen3_moe_batched_decode_q4k_rows2_lane(%input_size, %gate_noalias, %q8_noalias, %row_byte_base, %q8_token0_byte_base, %q8_token1_byte_base, %lane) : (index, buffer, buffer, offset, offset, offset, index) -> (f32, f32) + %up0_acc, %up1_acc = func.call @qwen3_moe_batched_decode_q4k_rows2_lane(%input_size, %up_noalias, %q8_noalias, %row_byte_base, %q8_token0_byte_base, %q8_token1_byte_base, %lane) : (index, buffer, buffer, offset, offset, offset, index) -> (f32, f32) + %gate0_dot = kernel.subgroup.reduce %gate0_acc : f32 + %up0_dot = kernel.subgroup.reduce %up0_acc : f32 + %gate1_dot = kernel.subgroup.reduce %gate1_acc : f32 + %up1_dot = kernel.subgroup.reduce %up1_acc : f32 + %lane_i32 = index.cast %lane : index to i32 + %is_lane_zero = scalar.cmpi eq, %lane_i32, %c0_i32 : i32 + %valid_channel_descriptor = scalar.andi %valid_channel, %valid_descriptor : i1 + %writes_output = scalar.andi %valid_channel_descriptor, %is_lane_zero : i1 + scf.if %writes_output { + %gate0_silu = scalar.siluf %gate0_dot : f32 + %gate1_silu = scalar.siluf %gate1_dot : f32 + %result0 = scalar.mulf %gate0_silu, %up0_dot : f32 + %result1 = scalar.mulf %gate1_silu, %up1_dot : f32 + view.store %result0, %output_view[%assignment0, %channel] : f32, view<[%assignment_count]x[%bounded_output_size]xf32> + view.store %result1, %output_view[%assignment1, %channel] : f32, view<[%assignment_count]x[%bounded_output_size]xf32> + } + kernel.return +} + +// Two identical top-8 route rows produce eight exact-two-row descriptors. +// The direct route-centric provider remains the numerical oracle. +check.case public @qwen3_moe_batched_decode_gate_up_q4k_rows2_differential_case { + %descriptor_count = check.literal value(8) : index + %queue_ordinal = check.literal value(1) : index + %token_count = check.literal value(2) : index + %input_size = check.literal value(512) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(32) : index + %output_size = check.literal value(32) : index + %input = check.generate.iota offset(-0.5) step(0.001953125) period(1024) : tensor<2x512xf32> + %q8_input = check.generate.fill value(0) : tensor<2x576xi8> + %route_ids = check.generate.iota offset(0) step(1) period(8) : tensor<2x8xi32> + %assignment_ordinals = check.generate.fill value(-1) : tensor<16xi32> + %queue_counts = check.generate.fill value(-1) : tensor<4xi32> + %queue_descriptors = check.generate.fill value(-1) : tensor<4x32x4xi32> + %gate_weight = check.generate.iota offset(-72) step(1) period(144) : tensor<32x32x2x144xi8> + %up_weight = check.generate.iota offset(-71) step(1) period(144) : tensor<32x32x2x144xi8> + %expected = check.generate.fill value(0.0) : tensor<2x8x32xf32> + %actual = check.generate.fill value(1.0) : tensor<2x8x32xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<2x512xf32>, tensor<2x576xi8>) + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %assignment_ordinals, %queue_counts, %queue_descriptors) : [index, index, index, index](index, index, index, index, tensor<2x8xi32>, tensor<16xi32>, tensor<4xi32>, tensor<4x32x4xi32>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8[%token_count, %route_count, %route_stride, %expert_count, %output_size](%token_count, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %expected) : [index, index, index, index, index](index, index, index, index, index, tensor<2x576xi8>, tensor<2x8xi32>, tensor<32x32x2x144xi8>, tensor<32x32x2x144xi8>, tensor<2x8x32xf32>) + kernel.launch @qwen3_moe_batched_decode_gate_up_q4k_rows2[%descriptor_count, %queue_ordinal, %token_count, %route_count, %expert_count, %output_size](%descriptor_count, %queue_ordinal, %token_count, %route_count, %expert_count, %output_size, %queue_descriptors, %assignment_ordinals, %q8_input, %gate_weight, %up_weight, %actual) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<4x32x4xi32>, tensor<16xi32>, tensor<2x576xi8>, tensor<32x32x2x144xi8>, tensor<32x32x2x144xi8>, tensor<2x8x32xf32>) + check.expect.close actual(%actual) expected(%expected) atol(0.25) rtol(9.9999999999999995e-07) nan(same) : tensor<2x8x32xf32> + check.return +} + +// Production dimensions with a uniform two-row-per-expert distribution. +// Timing includes the assignment producer so the row schedule cannot hide its +// routing preparation cost. +check.case public @qwen3_moe_batched_decode_gate_up_q4k_rows2_benchmark_case { + %descriptor_count = check.literal value(64) : index + %queue_ordinal = check.literal value(1) : index + %token_count = check.literal value(16) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(768) : index + %q8_input = check.generate.fill value(0) : tensor<16x2304xi8> + %route_ids = check.generate.iota offset(0) step(1) period(64) : tensor<16x8xi32> + %assignment_ordinals = check.generate.fill value(-1) : tensor<128xi32> + %queue_counts = check.generate.fill value(-1) : tensor<4xi32> + %queue_descriptors = check.generate.fill value(-1) : tensor<4x128x4xi32> + %gate_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %up_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<16x8x768xf32> + kernel.launch @qwen3_moe_build_batched_decode_expert_dispatch[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %assignment_ordinals, %queue_counts, %queue_descriptors) : [index, index, index, index](index, index, index, index, tensor<16x8xi32>, tensor<128xi32>, tensor<4xi32>, tensor<4x128x4xi32>) + kernel.launch @qwen3_moe_batched_decode_gate_up_q4k_rows2[%descriptor_count, %queue_ordinal, %token_count, %route_count, %expert_count, %output_size](%descriptor_count, %queue_ordinal, %token_count, %route_count, %expert_count, %output_size, %queue_descriptors, %assignment_ordinals, %q8_input, %gate_weight, %up_weight, %output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<4x128x4xi32>, tensor<128xi32>, tensor<16x2304xi8>, tensor<128x768x8x144xi8>, tensor<128x768x8x144xi8>, tensor<16x8x768xf32>) + check.return +} + +check.benchmark<@qwen3_moe_batched_decode_gate_up_q4k_rows2_benchmark_case> @qwen3_moe_batched_decode_gate_up_q4k_rows2_benchmark diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/dense_linear_quantized_f16_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/dense_linear_quantized_f16_wmma.loom new file mode 100644 index 000000000000..7d3b3dd8e44d --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/dense_linear_quantized_f16_wmma.loom @@ -0,0 +1,957 @@ +// Dense raw-quantized projection for the Qwen attention path. +// +// Each two-wave workgroup computes 64 output channels for 32 contiguous +// tokens. Raw GGUF Q4_K or Q6_K rows are decoded directly into a padded FP16 +// LDS tile, while the same workgroup converts the corresponding F32 activation +// tile to FP16. Four wave64 WMMA accumulators cover the 32x32 result owned by +// each wave. A wave-private LDS slice transposes each accumulator for coalesced +// F32 publication. +// +// Both weight contracts use the unmodified GGUF layout: +// Q4_K: [output channel][input size / 256][144 bytes] +// Q6_K: [output channel][input size / 256][210 bytes] +// The entry point selects the format before the shared device template is +// instantiated, so inactive decode logic is absent from the emitted kernel. No +// persistent repacking or expanded-weight allocation is required. +template.decl @qwen3_moe.dense.q4k_q8_1_x4.body(%publish_output: i1, %token_count: index, %token0: index, %q8_input: buffer, %weight: buffer, %output: buffer) + +template.decl @qwen3_moe.dense_quantized.body(%weight_format: index, %token_count: index, %input_size0: index, %output_size0: index, %output_accumulation: index, %channel_tile: index, %token_tile: index, %input: buffer, %weight: buffer, %output: buffer) + +template.decl @qwen3_moe.dense_quantized.launch(%token_capacity: index, %output_size: index) -> (index, index, index, index) + +template.decl @qwen3_moe.rmsnorm_quantize_q8_1_x4.body(%arg0: i1, %arg1: index, %arg2: index, %arg3: index, %arg4: buffer, %arg5: buffer, %arg6: buffer, %arg7: buffer) + +amdgpu.target @qwen3_moe_dense_gfx11_wave64 {subgroup_size = 64} + +amdgpu.target @qwen3_moe_dense_gfx11_wave32 {subgroup_size = 32} + +target.decl @qwen3_moe_attention_prepare_gfx11_wave32 + +config.decl @qwen3_moe.dense_quantized.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +config.decl @qwen3_moe.dense_quantized.output_size : %value: index where [range(%value, 1, 262144)] + +// Selects whether publication overwrites the output (0) or accumulates the +// projection into a caller-provided residual (1). +config.decl @qwen3_moe.dense_quantized.output_accumulation : %value: index where [range(%value, 0, 1)] + +config.decl @qwen3_moe.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +// These declarations are used only by the differential correctness case. The +// dense production kernel has no routing input or dependency. +kernel.decl @qwen3_moe_build_expert_table(%token_count$34: index, %route_count$35: index, %route_stride$36: index, %expert_count$37: index) launch(%token_count$38: index, %route_count$39: index, %route_stride$40: index, %expert_count$41: index, %route_ids: buffer, %expert_table: buffer) + +kernel.decl @qwen3_moe_routed_linear_q4k_f16_wmma(%token_count$44: index) launch(%token_count$45: index, %input: buffer, %expert_table: buffer, %weight: buffer, %output: buffer) + +// Shared raw-layout Q4_K row contraction primitive. Each lane consumes both +// nibbles of one packed-code load. +func.decl @qwen3_moe_q4k_q8_1_x4_paired_row_lane(%input_size: index, %weight: buffer, %weight_row_byte_base: offset, %q8_input: buffer, %q8_row_byte_base: offset, %lane: index) -> (f32) + +// Shared raw Q6_K decoder linked from the GGML physical-layout module. +func.decl @ggml_q6k_f16_vector4(%weight: buffer, %weight_row_byte_base: offset, %q6_block: index, %q6_group: index, %packet: index) -> (vector<4xf16>) + +// These declarations are used only by the Q6_K differential correctness case. +kernel.decl @ggml_quantize_q8_1_x4_f32(%token_count$63: index, %input_size$64: index) launch(%token_count$65: index, %input_size$66: index, %input: buffer, %output: buffer) + +kernel.decl @ggml_linear_q6k_q8_1_x4(%token_count$69: index, %input_size$70: index, %output_size$71: index) launch(%token_count$72: index, %input_size$73: index, %output_size$74: index, %q8_input: buffer, %weight: buffer, %output: buffer) + +kernel.decl @qwen3_moe_rmsnorm_f32_quantize_q8_1_x4(%token_count$78: index) launch(%token_count$79: index, %input: buffer, %rms_weight: buffer, %normalized_f32_output: buffer, %q8_1_x4_output: buffer) + +// Decodes the four adjacent Q4_K values owned by one load packet. One packed +// header load supplies both half scales and all twelve six-bit group scales; +// one packed code load supplies four adjacent unsigned nibbles. The final +// affine conversion is formed in F32 before truncation to the FP16 WMMA input. +func.def inline @qwen3_moe_dense_q4k_wmma_vector4(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group: index, %packet: index) -> (vector<4xf16>) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %block_bytes = index.constant 144 : offset + %scale_offset = index.constant 4 : offset + %code_offset = index.constant 16 : offset + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15 = vector.constant 15 : vector<1xi32> + %c48 = vector.constant 48 : vector<1xi32> + %q4_mask = vector.constant 252645135 : vector<1xi32> + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q4_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %scale_byte_base = index.add %block_byte_base, %scale_offset : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %dm_view = buffer.view %weight[%block_byte_base] : buffer -> view<2xf16> + %scale_view = buffer.view %weight[%scale_byte_base] : buffer -> view<3xi32> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<32xi32> + %dm = vector.load %dm_view[%c0] : view<2xf16> -> vector<2xf16> + %scales = vector.load %scale_view[%c0] : view<3xi32> -> vector<3xi32> + %d_f16 = vector.extract %dm[0] : vector<2xf16> -> f16 + %dmin_f16 = vector.extract %dm[1] : vector<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %q_page0 = index.div %bounded_group, %c2 : index + %q_page = index.mul %q_page0, %c8 : index + %q_word_index0 = index.add %q_page, %bounded_packet : index + %q_word_index = index.assume %q_word_index0 [range(%q_word_index0, 0, 31)] : index + %is_low = index.cmp ult, %bounded_group, %c4 : index + %scale_lane = index.rem %bounded_group, %c4 : index + %scale_shift_index = index.mul %scale_lane, %c8 : index + %scale_shift_i32 = index.cast %scale_shift_index : index to i32 + %scale_shift = vector.splat %scale_shift_i32 : vector<1xi32> + %scale0_i32 = vector.extract %scales[0] : vector<3xi32> -> i32 + %scale1_i32 = vector.extract %scales[1] : vector<3xi32> -> i32 + %scale2_i32 = vector.extract %scales[2] : vector<3xi32> -> i32 + %scale0 = vector.splat %scale0_i32 : vector<1xi32> + %scale1 = vector.splat %scale1_i32 : vector<1xi32> + %scale2 = vector.splat %scale2_i32 : vector<1xi32> + %high_shift_i32 = scalar.addi %scale_shift_i32, %c2_i32 : i32 + %minimum_shift_i32 = scalar.addi %scale_shift_i32, %c4_i32 : i32 + %selected_scale_source = scf.select %is_low, %scale0, %scale2 : vector<1xi32> + %selected_minimum_source = scf.select %is_low, %scale1, %scale2 : vector<1xi32> + %selected_scale_high_shift_i32 = scf.select %is_low, %scale_shift_i32, %high_shift_i32 : i32 + %selected_minimum_low_shift_i32 = scf.select %is_low, %scale_shift_i32, %minimum_shift_i32 : i32 + %selected_scale_high_shift = vector.splat %selected_scale_high_shift_i32 : vector<1xi32> + %selected_minimum_low_shift = vector.splat %selected_minimum_low_shift_i32 : vector<1xi32> + %scale_low0 = vector.shrui %selected_scale_source, %scale_shift : vector<1xi32> + %scale_low = vector.andi %scale_low0, %c15 : vector<1xi32> + %scale_high0 = vector.shrui %scale0, %selected_scale_high_shift : vector<1xi32> + %scale_high = vector.andi %scale_high0, %c48 : vector<1xi32> + %scale = vector.ori %scale_low, %scale_high : vector<1xi32> + %minimum_low0 = vector.shrui %selected_minimum_source, %selected_minimum_low_shift : vector<1xi32> + %minimum_low = vector.andi %minimum_low0, %c15 : vector<1xi32> + %minimum_high0 = vector.shrui %scale1, %selected_scale_high_shift : vector<1xi32> + %minimum_high = vector.andi %minimum_high0, %c48 : vector<1xi32> + %minimum = vector.ori %minimum_low, %minimum_high : vector<1xi32> + %scale_f32 = vector.uitofp %scale : vector<1xi32> to vector<1xf32> + %minimum_f32 = vector.uitofp %minimum : vector<1xi32> to vector<1xf32> + %d_vector1 = vector.splat %d : vector<1xf32> + %dmin_vector1 = vector.splat %dmin : vector<1xf32> + %d_scale_vector1 = vector.mulf %d_vector1, %scale_f32 : vector<1xf32> + %minimum_scale_vector1 = vector.mulf %dmin_vector1, %minimum_f32 : vector<1xf32> + %d_scale = vector.extract %d_scale_vector1[0] : vector<1xf32> -> f32 + %minimum_scale = vector.extract %minimum_scale_vector1[0] : vector<1xf32> -> f32 + %q_word = vector.load %code_view[%q_word_index] : view<32xi32> -> vector<1xi32> + %q_half = index.rem %bounded_group, %c2 : index + %q_shift_index = index.mul %q_half, %c4 : index + %q_shift_i32 = index.cast %q_shift_index : index to i32 + %q_shift = vector.splat %q_shift_i32 : vector<1xi32> + %shifted_q = vector.shrui %q_word, %q_shift : vector<1xi32> + %masked_q = vector.andi %shifted_q, %q4_mask : vector<1xi32> + %q_i8 = vector.bitcast %masked_q : vector<1xi32> to vector<4xi8> + %q_f32 = vector.uitofp %q_i8 : vector<4xi8> to vector<4xf32> + %negative_minimum_scale = scalar.negf %minimum_scale : f32 + %q0 = vector.extract %q_f32[0] : vector<4xf32> -> f32 + %q1 = vector.extract %q_f32[1] : vector<4xf32> -> f32 + %q2 = vector.extract %q_f32[2] : vector<4xf32> -> f32 + %q3 = vector.extract %q_f32[3] : vector<4xf32> -> f32 + %value0 = scalar.fmaf %q0, %d_scale, %negative_minimum_scale : f32 + %value1 = scalar.fmaf %q1, %d_scale, %negative_minimum_scale : f32 + %value2 = scalar.fmaf %q2, %d_scale, %negative_minimum_scale : f32 + %value3 = scalar.fmaf %q3, %d_scale, %negative_minimum_scale : f32 + %half0 = scalar.fptrunc %value0 : f32 to f16 + %half1 = scalar.fptrunc %value1 : f32 to f16 + %half2 = scalar.fptrunc %value2 : f32 to f16 + %half3 = scalar.fptrunc %value3 : f32 to f16 + %result = vector.from_elements %half0, %half1, %half2, %half3 : vector<4xf16> + func.return %result : vector<4xf16> +} + +// Low-fixed-cost Q4_K body shared by the ordinary and completion-fused decode +// exports. Each wave owns one output channel and contracts a Q8_1 activation +// row directly against the original GGUF weight row. +template.def<@qwen3_moe.dense.q4k_q8_1_x4.body> device @qwen3_moe_dense_linear_q4k_q8_1_x4_body(%publish_output: i1, %token_count: index, %token0: index, %q8_input: buffer, %weight: buffer, %output: buffer) { + %input_size = config.get @qwen3_moe.dense_quantized.input_size : index + %output_size = config.get @qwen3_moe.dense_quantized.output_size : index + %output_accumulation = config.get @qwen3_moe.dense_quantized.output_accumulation : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 262144)] : index + %channel_tile = kernel.workgroup.id : index + %subgroup0 = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %c1 = index.constant 1 : index + %c8 = index.constant 8 : index + %c128 = index.constant 128 : index + %c144_bytes = index.constant 144 : offset + %c256 = index.constant 256 : index + %c0_i32 = scalar.constant 0 : i32 + %c0_f32 = scalar.constant 0.0 : f32 + %c0_offset = index.constant 0 : offset + %token, %launch_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %channel_base = index.mul %channel_tile, %c8 : index + %channel = index.add %channel_base, %subgroup : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %lane_i32 = index.cast %lane : index to i32 + %is_lane_zero = scalar.cmpi eq, %lane_i32, %c0_i32 : i32 + %q4_block_count = index.div %bounded_input_size, %c256 : index + %weight_row_bytes = index.scale %q4_block_count, %c144_bytes : index, offset -> offset + %weight_row_byte_base = index.scale %channel, %weight_row_bytes : index, offset -> offset + %q8_group_count = index.div %bounded_input_size, %c128 : index + %q8_row_bytes = index.scale %q8_group_count, %c144_bytes : index, offset -> offset + %q8_row_byte_base = index.scale %token, %q8_row_bytes : index, offset -> offset + %q8_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %q8_input, %weight, %output : buffer, buffer, buffer + %lane_sum = scf.if %publish_output -> (f32) { + %channel_sum = scf.if %valid_channel -> (f32) { + %sum = func.call @qwen3_moe_q4k_q8_1_x4_paired_row_lane(%bounded_input_size, %weight_noalias, %weight_row_byte_base, %q8_noalias, %q8_row_byte_base, %lane) : (index, buffer, offset, buffer, offset, index) -> (f32) + scf.yield %sum : f32 + } else { + scf.yield %c0_f32 : f32 + } + scf.yield %channel_sum : f32 + } else { + scf.yield %c0_f32 : f32 + } + %dot = kernel.subgroup.reduce %lane_sum : f32 + %accumulates_output = index.cmp eq, %output_accumulation, %c1 : index + scf.if %publish_output { + scf.if %valid_channel { + scf.if %is_lane_zero { + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%bounded_output_size]xf32> + %result = scf.if %accumulates_output -> (f32) { + %residual = view.load %output_view[%token, %channel] : view<[%launch_token_count]x[%bounded_output_size]xf32> -> f32 + %sum = scalar.addf %residual, %dot : f32 + scf.yield %sum : f32 + } else { + scf.yield %dot : f32 + } + view.store %result, %output_view[%token, %channel] : f32, view<[%launch_token_count]x[%bounded_output_size]xf32> + } + } + } + template.return +} + +// Direct decode entry point. This complements the grouped WMMA schedule below: +// it gives up cross-token weight reuse in exchange for doing no padded matrix +// work when fewer than 32 tokens are available. +kernel.def target(@qwen3_moe_dense_gfx11_wave32) @qwen3_moe_dense_linear_q4k_q8_1_x4(%token_count: index) { + %output_size = config.get @qwen3_moe.dense_quantized.output_size : index + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %c1 = index.constant 1 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c256 = index.constant 256 : index + %padded_output_size = index.add %output_size, %c7 : index + %output_tiles = index.div %padded_output_size, %c8 : index + kernel.launch.config workgroups(%output_tiles, %token_capacity, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %q8_input: buffer, %weight: buffer, %output: buffer) { + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %token0 = kernel.workgroup.id : index + %c0 = index.constant 0 : index + %valid_token = index.cmp ult, %token0, %bounded_token_count : index + %safe_token = scf.select %valid_token, %token0, %c0 : index + template.apply<@qwen3_moe.dense.q4k_q8_1_x4.body>(%valid_token, %bounded_token_count, %safe_token, %q8_input, %weight, %output) : (i1, index, index, buffer, buffer, buffer) + kernel.return +} + +// Decode-only route that publishes the normalized F32 and Q8_1 x4 rows +// consumed by the following feed-forward boundary. +kernel.def target(@qwen3_moe_attention_prepare_gfx11_wave32) @qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8(%token_count: index) { + %output_size = config.get @qwen3_moe.dense_quantized.output_size : index + %c1 = index.constant 1 : index + %c8 = index.constant 8 : index + %c256 = index.constant 256 : index + %output_tiles = index.div %output_size, %c8 : index + kernel.launch.config workgroups(%output_tiles, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %q8_input: buffer, %weight: buffer, %output: buffer, %norm_weight: buffer, %normalized_output: buffer, %completion_counter: buffer, %next_q8_output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 1)] : index + %publish_output = scalar.constant true : i1 + %projection_token = kernel.workgroup.id : index + template.apply<@qwen3_moe.dense.q4k_q8_1_x4.body>(%publish_output, %bounded_token_count, %projection_token, %q8_input, %weight, %output) : (i1, index, index, buffer, buffer, buffer) + %output_size0 = config.get @qwen3_moe.dense_quantized.output_size : index + %output_size = index.assume %output_size0 [range(%output_size0, 128, 32768), mul(%output_size0, 128)] : index + %token0 = kernel.workgroup.id : index + %token = index.assume %token0 [lt(%token0, %bounded_token_count)] : index + %c0 = index.constant 0 : index + %c8 = index.constant 8 : index + %workitem = kernel.workitem.id : index + %output_tile_count = index.div %output_size, %c8 : index + %is_arrival_workitem = index.cmp eq, %workitem, %c0 : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c0_offset = index.constant 0 : offset + %counter_scratch_bytes = index.constant 4 : offset + %output_noalias, %norm_weight_noalias, %normalized_output_noalias, %completion_counter_noalias, %next_q8_output_noalias = buffer.assume.noalias %output, %norm_weight, %normalized_output, %completion_counter, %next_q8_output : buffer, buffer, buffer, buffer, buffer + %completion_counter_aligned = buffer.assume.alignment %completion_counter_noalias {minimum_alignment = 16} : buffer + %completion_counter_view = buffer.view %completion_counter_aligned[%c0_offset] : buffer -> view<1xi32> + %counter_scratch = buffer.alloca align(4) %counter_scratch_bytes : buffer + %counter_scratch_view = buffer.view %counter_scratch[%c0_offset] : buffer -> view<1xi32> + // Publish every producer's residual stores before the leader advances one + // workgroup arrival. The last arrival then acquires the complete row. + kernel.barrier scope(workgroup) ordering(release) + scf.if %is_arrival_workitem { + %old_counter = view.atomic.rmw %c1_i32, %completion_counter_view[%c0] {ordering = acq_rel, scope = device} : i32, view<1xi32> -> i32 + view.store %old_counter, %counter_scratch_view[%c0] : i32, view<1xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %old_counter = view.load %counter_scratch_view[%c0] : view<1xi32> -> i32 + %output_tile_count_i32 = index.cast %output_tile_count : index to i32 + %last_output_tile_i32 = scalar.subi %output_tile_count_i32, %c1_i32 : i32 + %negative_output_tile_count_i32 = scalar.subi %c0_i32, %output_tile_count_i32 : i32 + %is_last_output_tile = scalar.cmpi eq, %old_counter, %last_output_tile_i32 : i32 + scf.if %is_last_output_tile { + kernel.barrier scope(workgroup) ordering(acquire) + %publish_normalized = scalar.constant true : i1 + template.apply<@qwen3_moe.rmsnorm_quantize_q8_1_x4.body>(%publish_normalized, %c8, %bounded_token_count, %token, %output_noalias, %norm_weight_noalias, %normalized_output_noalias, %next_q8_output_noalias) : (i1, index, index, index, buffer, buffer, buffer, buffer) + // Reset only after every normalized F32 and Q8 store completes. + kernel.barrier scope(workgroup) ordering(release) + scf.if %is_arrival_workitem { + view.atomic.reduce %negative_output_tile_count_i32, %completion_counter_view[%c0] {ordering = release, scope = device} : i32, view<1xi32> + } + } + kernel.return +} + +// Shared matrix-tile schedule. Callers own the launch grid and pass one +// validated tile coordinate plus a literal storage kind, allowing the linker +// and JIT to erase the inactive packed decoder. +template.def<@qwen3_moe.dense_quantized.body> device @qwen3_moe_dense_linear_quantized_f16_wmma_body(%weight_format: index, %token_count: index, %input_size0: index, %output_size0: index, %output_accumulation: index, %channel_tile: index, %token_tile: index, %input: buffer, %weight: buffer, %output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_input_size = index.assume %input_size0 [range(%input_size0, 256, 32768), mul(%input_size0, 256)] : index + %bounded_output_size = index.assume %output_size0 [range(%output_size0, 1, 262144)] : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 1)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c6 = index.constant 6 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c40 = index.constant 40 : index + %c64 = index.constant 64 : index + %c210_bytes = index.constant 210 : offset + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %q4_block_bytes = index.constant 144 : offset + %weight_stage_bytes = index.constant 5120 : offset + %activation_stage_bytes = index.constant 2560 : offset + %wave_result_stage_bytes = index.constant 512 : offset + %result_stage_bytes = index.constant 1024 : offset + %c0_f16x4 = vector.constant 0.0 : vector<4xf16> + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %zero_accumulator = vector.constant 0.0 : vector<8xf16> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %is_q6 = index.cmp eq, %weight_format, %c6 : index + %weight_block_bytes = scf.select %is_q6, %c210_bytes, %q4_block_bytes : offset + %quant_block_count = index.div %bounded_input_size, %c256 : index + %weight_row_bytes = index.scale %quant_block_count, %weight_block_bytes : index, offset -> offset + %input_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input, %weight, %output : buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_input_size]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_output_size]xf32> + %weight_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %activation_stage = buffer.alloca align(16) %activation_stage_bytes : buffer + %result_stage = buffer.alloca align(16) %result_stage_bytes : buffer + %weight_stage_view = buffer.view %weight_stage[%c0_offset] : buffer -> view<64x40xf16> + %activation_stage_physical_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x40xf16> + %activation_fragment_layout = encoding.layout.strided [1, %c40] : encoding + %activation_fragment_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x32xf16, %activation_fragment_layout> + %wave_result_stage_offset = index.scale %subgroup, %wave_result_stage_bytes : index, offset -> offset + %result_fragment_layout = encoding.layout.strided [1, %c16] : encoding + %result_fragment_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf16, %result_fragment_layout> + %result_physical_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf16> + %channel_tile_base = index.mul %channel_tile, %c64 : index + %token_tile_base = index.mul %token_tile, %c32 : index + %load_packet = index.rem %workitem, %c8 : index + %load_k = index.mul %load_packet, %c4 : index + %load_row0 = index.div %workitem, %c8 : index + %load_row = index.assume %load_row0 [range(%load_row0, 0, 15)] : index + %subgroup_channel_add = index.mul %subgroup, %c32 : index + %subgroup_channel1 = index.add %subgroup_channel_add, %c16 : index + %init00 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf16> + %init01 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf16> + %init10 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf16> + %init11 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf16> + %result00, %result01, %result10, %result11 = scf.for %quant_block = [%c0 to %quant_block_count step %c1](%block_acc00 = %init00 : vector<8xf16>, %block_acc01 = %init01 : vector<8xf16>, %block_acc10 = %init10 : vector<8xf16>, %block_acc11 = %init11 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16>) { + %block_result00, %block_result01, %block_result10, %block_result11 = scf.for %quant_group = [%c0 to %c8 step %c1](%acc00 = %block_acc00 : vector<8xf16>, %acc01 = %block_acc01 : vector<8xf16>, %acc10 = %block_acc10 : vector<8xf16>, %acc11 = %block_acc11 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16>) { + %block_k_base = index.mul %quant_block, %c256 : index + %group_k_add = index.mul %quant_group, %c32 : index + %k_origin = index.add %block_k_base, %group_k_add : index + scf.for %row_offset = [%c0 to %c64 step %c16] unroll { + %local_row0 = index.add %load_row, %row_offset : index + %local_row = index.assume %local_row0 [range(%local_row0, 0, 63)] : index + %channel = index.add %channel_tile_base, %local_row : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %weight_values = scf.if %valid_channel -> (vector<4xf16>) { + %row_byte_base = index.scale %channel, %weight_row_bytes : index, offset -> offset + %decoded = scf.if %is_q6 -> (vector<4xf16>) { + %q6_values = func.call @ggml_q6k_f16_vector4(%weight_noalias, %row_byte_base, %quant_block, %quant_group, %load_packet) : (buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %q6_values : vector<4xf16> + } else { + %q4_values = func.call @qwen3_moe_dense_q4k_wmma_vector4(%weight_noalias, %row_byte_base, %quant_block, %quant_group, %load_packet) : (buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %q4_values : vector<4xf16> + } + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %weight_values, %weight_stage_view[%local_row, %load_k] : vector<4xf16>, view<64x40xf16> + %is_activation_row = index.cmp ult, %local_row, %c32 : index + scf.if %is_activation_row { + %activation_row = index.assume %local_row [range(%local_row, 0, 31)] : index + %token = index.add %token_tile_base, %activation_row : index + %valid_token = index.cmp ult, %token, %bounded_token_count : index + %activation_values = scf.if %valid_token -> (vector<4xf16>) { + %bounded_token, %input_token_count = index.assume %token, %bounded_token_count [lt(%token, %bounded_token_count)] : index, index + %input_k = index.add %k_origin, %load_k : index + %loaded = vector.load %input_view[%bounded_token, %input_k] : view<[%bounded_token_count]x[%bounded_input_size]xf32> -> vector<4xf32> + %converted = vector.fptrunc %loaded : vector<4xf32> to vector<4xf16> + scf.yield %converted : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %activation_values, %activation_stage_physical_view[%activation_row, %load_k] : vector<4xf16>, view<32x40xf16> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %next00, %next01, %next10, %next11 = scf.for %k_half = [%c0 to %c32 step %c16](%half_acc00 = %acc00 : vector<8xf16>, %half_acc01 = %acc01 : vector<8xf16>, %half_acc10 = %acc10 : vector<8xf16>, %half_acc11 = %acc11 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16>) unroll { + %lhs0 = vector.fragment.load %weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %lhs1 = vector.fragment.load %weight_stage_view[%subgroup_channel1, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %rhs0 = vector.fragment.load %activation_fragment_view[%k_half, %c0] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %rhs1 = vector.fragment.load %activation_fragment_view[%k_half, %c16] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %half_next00 = vector.mma %lhs0, %rhs0, %half_acc00 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %half_next01 = vector.mma %lhs0, %rhs1, %half_acc01 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %half_next10 = vector.mma %lhs1, %rhs0, %half_acc10 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %half_next11 = vector.mma %lhs1, %rhs1, %half_acc11 : vector<16xf16>, vector<16xf16>, vector<8xf16> + scf.yield %half_next00, %half_next01, %half_next10, %half_next11 : vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %next00, %next01, %next10, %next11 : vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16> + } + scf.yield %block_result00, %block_result01, %block_result10, %block_result11 : vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16> + } + // WMMA owns [channel][token] fragments. Each wave transposes one fragment + // through its private LDS slice so every lane publishes four contiguous F32 + // channels for one logical token. + %publish_token0 = index.div %lane, %c4 : index + %publish_token = index.assume %publish_token0 [range(%publish_token0, 0, 15)] : index + %publish_packet0 = index.rem %lane, %c4 : index + %publish_packet = index.assume %publish_packet0 [range(%publish_packet0, 0, 3)] : index + %publish_channel_add = index.mul %publish_packet, %c4 : index + %token0 = index.add %token_tile_base, %publish_token : index + %publish_token1 = index.add %publish_token, %c16 : index + %token1 = index.add %token_tile_base, %publish_token1 : index + %subgroup_channel_base = index.add %channel_tile_base, %subgroup_channel_add : index + %channel0 = index.add %subgroup_channel_base, %publish_channel_add : index + %channel1_base = index.add %subgroup_channel_base, %c16 : index + %channel1 = index.add %channel1_base, %publish_channel_add : index + %valid_token0 = index.cmp ult, %token0, %bounded_token_count : index + %valid_token1 = index.cmp ult, %token1, %bounded_token_count : index + %valid_channel0 = index.cmp ult, %channel0, %bounded_output_size : index + %valid_channel1 = index.cmp ult, %channel1, %bounded_output_size : index + %writes00 = scalar.andi %valid_token0, %valid_channel0 : i1 + %writes01 = scalar.andi %valid_token1, %valid_channel0 : i1 + %writes10 = scalar.andi %valid_token0, %valid_channel1 : i1 + %writes11 = scalar.andi %valid_token1, %valid_channel1 : i1 + %accumulates_output = index.cmp eq, %output_accumulation, %c1 : index + vector.fragment.store %result00, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<8xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes00 { + %bounded_token, %output_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %values = vector.load %result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf16> -> vector<4xf16> + %wide = vector.extf %values : vector<4xf16> to vector<4xf32> + %mask = vector.mask.range [%channel0 to %bounded_output_size step %c1] : index -> vector<4xi1> + %published = scf.if %accumulates_output -> (vector<4xf32>) { + %residual = vector.load.mask %output_view[%bounded_token, %channel0], %mask, %c0_f32x4 : view<[%bounded_token_count]x[%bounded_output_size]xf32>, vector<4xi1>, vector<4xf32> + %sum = vector.addf %residual, %wide : vector<4xf32> + scf.yield %sum : vector<4xf32> + } else { + scf.yield %wide : vector<4xf32> + } + vector.store.mask %published, %output_view[%bounded_token, %channel0], %mask : vector<4xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32>, vector<4xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result01, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<8xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes01 { + %bounded_token, %output_token_count = index.assume %token1, %bounded_token_count [lt(%token1, %bounded_token_count)] : index, index + %values = vector.load %result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf16> -> vector<4xf16> + %wide = vector.extf %values : vector<4xf16> to vector<4xf32> + %mask = vector.mask.range [%channel0 to %bounded_output_size step %c1] : index -> vector<4xi1> + %published = scf.if %accumulates_output -> (vector<4xf32>) { + %residual = vector.load.mask %output_view[%bounded_token, %channel0], %mask, %c0_f32x4 : view<[%bounded_token_count]x[%bounded_output_size]xf32>, vector<4xi1>, vector<4xf32> + %sum = vector.addf %residual, %wide : vector<4xf32> + scf.yield %sum : vector<4xf32> + } else { + scf.yield %wide : vector<4xf32> + } + vector.store.mask %published, %output_view[%bounded_token, %channel0], %mask : vector<4xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32>, vector<4xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result10, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<8xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes10 { + %bounded_token, %output_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %values = vector.load %result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf16> -> vector<4xf16> + %wide = vector.extf %values : vector<4xf16> to vector<4xf32> + %mask = vector.mask.range [%channel1 to %bounded_output_size step %c1] : index -> vector<4xi1> + %published = scf.if %accumulates_output -> (vector<4xf32>) { + %residual = vector.load.mask %output_view[%bounded_token, %channel1], %mask, %c0_f32x4 : view<[%bounded_token_count]x[%bounded_output_size]xf32>, vector<4xi1>, vector<4xf32> + %sum = vector.addf %residual, %wide : vector<4xf32> + scf.yield %sum : vector<4xf32> + } else { + scf.yield %wide : vector<4xf32> + } + vector.store.mask %published, %output_view[%bounded_token, %channel1], %mask : vector<4xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32>, vector<4xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result11, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<8xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes11 { + %bounded_token, %output_token_count = index.assume %token1, %bounded_token_count [lt(%token1, %bounded_token_count)] : index, index + %values = vector.load %result_physical_view[%publish_token, %publish_channel_add] : view<16x16xf16> -> vector<4xf16> + %wide = vector.extf %values : vector<4xf16> to vector<4xf32> + %mask = vector.mask.range [%channel1 to %bounded_output_size step %c1] : index -> vector<4xi1> + %published = scf.if %accumulates_output -> (vector<4xf32>) { + %residual = vector.load.mask %output_view[%bounded_token, %channel1], %mask, %c0_f32x4 : view<[%bounded_token_count]x[%bounded_output_size]xf32>, vector<4xi1>, vector<4xf32> + %sum = vector.addf %residual, %wide : vector<4xf32> + scf.yield %sum : vector<4xf32> + } else { + scf.yield %wide : vector<4xf32> + } + vector.store.mask %published, %output_view[%bounded_token, %channel1], %mask : vector<4xf32>, view<[%bounded_token_count]x[%bounded_output_size]xf32>, vector<4xi1> + } + template.return +} + +// Shared launch geometry for configured standalone kernels and parameterized +// command-program kernels. +template.def<@qwen3_moe.dense_quantized.launch> @qwen3_moe_dense_linear_quantized_f16_wmma_launch(%token_capacity: index, %output_size: index) -> (index, index, index, index) { + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %padded_token_count = index.add %token_capacity, %c31 : index + %token_tiles = index.div %padded_token_count, %c32 : index + template.return %output_tiles, %token_tiles, %c1, %c128 : index, index, index, index +} + +// Configured entry points retain the compact standalone kernel ABI used by +// library-style callers. Shape configs specialize the native code once while +// token count remains the only per-dispatch scalar. +kernel.def target(@qwen3_moe_dense_gfx11_wave64) @qwen3_moe_dense_linear_q4k_f16_wmma(%token_count: index) { + %output_size = config.get @qwen3_moe.dense_quantized.output_size : index + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %padded_token_count = index.add %token_capacity, %c31 : index + %token_tiles = index.div %padded_token_count, %c32 : index + kernel.launch.config workgroups(%output_tiles, %token_tiles, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) { + %input_size = config.get @qwen3_moe.dense_quantized.input_size : index + %output_size = config.get @qwen3_moe.dense_quantized.output_size : index + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %output_accumulation = config.get @qwen3_moe.dense_quantized.output_accumulation : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + %q4 = index.constant 4 : index + template.apply<@qwen3_moe.dense_quantized.body>(%q4, %bounded_token_count, %input_size, %output_size, %output_accumulation, %channel_tile, %token_tile, %input, %weight, %output) : (index, index, index, index, index, index, index, buffer, buffer, buffer) + kernel.return +} + +kernel.def target(@qwen3_moe_dense_gfx11_wave64) @qwen3_moe_dense_linear_q6k_f16_wmma(%token_count: index) { + %output_size = config.get @qwen3_moe.dense_quantized.output_size : index + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %padded_token_count = index.add %token_capacity, %c31 : index + %token_tiles = index.div %padded_token_count, %c32 : index + kernel.launch.config workgroups(%output_tiles, %token_tiles, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) { + %input_size = config.get @qwen3_moe.dense_quantized.input_size : index + %output_size = config.get @qwen3_moe.dense_quantized.output_size : index + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %output_accumulation = config.get @qwen3_moe.dense_quantized.output_accumulation : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + %q6 = index.constant 6 : index + template.apply<@qwen3_moe.dense_quantized.body>(%q6, %bounded_token_count, %input_size, %output_size, %output_accumulation, %channel_tile, %token_tile, %input, %weight, %output) : (index, index, index, index, index, index, index, buffer, buffer, buffer) + kernel.return +} + +// Parameterized entry points carry launch-varying shapes through command IR. +// Command planning materializes each exact fact environment in a private unit, +// so these scalar values do not survive in the native device ABI. +kernel.def target(@qwen3_moe_dense_gfx11_wave64) @qwen3_moe_dense_linear_q4k_f16_wmma_parameterized(%token_count: index, %input_size: index, %output_size: index, %output_accumulation: index) { + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %padded_token_count = index.add %token_count, %c31 : index + %token_tiles = index.div %padded_token_count, %c32 : index + kernel.launch.config workgroups(%output_tiles, %token_tiles, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input_size: index, %output_size: index, %output_accumulation: index, %input: buffer, %weight: buffer, %output: buffer) where [range(%token_count, 1, 2048), range(%input_size, 256, 32768), mul(%input_size, 256), range(%output_size, 1, 262144), range(%output_accumulation, 0, 1)] { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + %q4 = index.constant 4 : index + template.apply<@qwen3_moe.dense_quantized.body>(%q4, %bounded_token_count, %input_size, %output_size, %output_accumulation, %channel_tile, %token_tile, %input, %weight, %output) : (index, index, index, index, index, index, index, buffer, buffer, buffer) + kernel.return +} + +kernel.def target(@qwen3_moe_dense_gfx11_wave64) @qwen3_moe_dense_linear_q6k_f16_wmma_parameterized(%token_count: index, %input_size: index, %output_size: index, %output_accumulation: index) { + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %padded_token_count = index.add %token_count, %c31 : index + %token_tiles = index.div %padded_token_count, %c32 : index + kernel.launch.config workgroups(%output_tiles, %token_tiles, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input_size: index, %output_size: index, %output_accumulation: index, %input: buffer, %weight: buffer, %output: buffer) where [range(%token_count, 1, 2048), range(%input_size, 256, 32768), mul(%input_size, 256), range(%output_size, 1, 262144), range(%output_accumulation, 0, 1)] { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %channel_tile = kernel.workgroup.id : index + %token_tile = kernel.workgroup.id : index + %q6 = index.constant 6 : index + template.apply<@qwen3_moe.dense_quantized.body>(%q6, %bounded_token_count, %input_size, %output_size, %output_accumulation, %channel_tile, %token_tile, %input, %weight, %output) : (index, index, index, index, index, index, index, buffer, buffer, buffer) + kernel.return +} + +// The exact-representable activation and nonzero packed weight bytes compare +// dense addressing against the already-certified routed WMMA provider. The +// shape crosses both the 32-token and 64-channel tile boundaries. +check.case public @qwen3_moe_dense_linear_q4k_f16_wmma_differential_case { + %token_count = check.literal value(33) : index + %route_count = check.literal value(1) : index + %route_stride = check.literal value(1) : index + %expert_count = check.literal value(1) : index + %input = check.generate.fill value(0.00390625) : tensor<33x512xf32> + %route_ids = check.generate.fill value(0) : tensor<33x1xi32> + %expert_table = check.generate.fill value(-1) : tensor<34xi32> + %weight = check.generate.fill value(34) : tensor<1x65x2x144xi8> + %expected = check.generate.fill value(0.0) : tensor<33x65xf32> + %actual = check.generate.fill value(1.0) : tensor<33x65xf32> + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<33x1xi32>, tensor<34xi32>) + kernel.launch @qwen3_moe_routed_linear_q4k_f16_wmma[%token_count](%token_count, %input, %expert_table, %weight, %expected) : [index](index, tensor<33x512xf32>, tensor<34xi32>, tensor<1x65x2x144xi8>, tensor<33x65xf32>) + kernel.launch @qwen3_moe_dense_linear_q4k_f16_wmma[%token_count](%token_count, %input, %weight, %actual) : [index](index, tensor<33x512xf32>, tensor<1x65x2x144xi8>, tensor<33x65xf32>) + check.expect.close actual(%actual) expected(%expected) atol(0.01) rtol(0.01) nan(same) : tensor<33x65xf32> + check.return +} + +// Q6_K uses the same exact-representable activation across both the Q8_1 dot +// reference and FP16 WMMA provider. The comparison crosses both matrix tile +// boundaries and validates the format-specialized packed decoder. +check.case public @qwen3_moe_dense_linear_q6k_f16_wmma_differential_case { + %token_count = check.literal value(33) : index + %input_size = check.literal value(512) : index + %output_size = check.literal value(65) : index + %input = check.generate.fill value(0.00390625) : tensor<33x512xf32> + %q8_input = check.generate.fill value(0) : tensor<33x576xi8> + %weight = check.generate.fill value(-86) : tensor<65x2x210xi8> + %expected = check.generate.fill value(0.0) : tensor<33x65xf32> + %actual = check.generate.fill value(1.0) : tensor<33x65xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<33x512xf32>, tensor<33x576xi8>) + kernel.launch @ggml_linear_q6k_q8_1_x4[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %q8_input, %weight, %expected) : [index, index, index](index, index, index, tensor<33x576xi8>, tensor<65x2x210xi8>, tensor<33x65xf32>) + kernel.launch @qwen3_moe_dense_linear_q6k_f16_wmma[%token_count](%token_count, %input, %weight, %actual) : [index](index, tensor<33x512xf32>, tensor<65x2x210xi8>, tensor<33x65xf32>) + check.expect.close actual(%actual) expected(%expected) atol(0.25) rtol(0.01) nan(same) : tensor<33x65xf32> + check.return +} + +// The O-projection dimensions double the Q/K input depth. This nonzero +// production-size differential keeps that loop regime covered independently +// of the benchmark parameter sweep. +check.case public @qwen3_moe_dense_linear_q4k_f16_wmma_o_differential_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(4096) : index + %route_count = check.literal value(1) : index + %route_stride = check.literal value(1) : index + %expert_count = check.literal value(1) : index + %input = check.generate.fill value(0.00390625) : tensor<1x4096xf32> + %q8_input = check.generate.fill value(0) : tensor<1x4608xi8> + %route_ids = check.generate.fill value(0) : tensor<1x1xi32> + %expert_table = check.generate.fill value(-1) : tensor<2xi32> + %weight = check.generate.fill value(34) : tensor<1x2048x16x144xi8> + %expected = check.generate.fill value(0.0) : tensor<1x2048xf32> + %actual = check.generate.fill value(0.0) : tensor<1x2048xf32> + %direct = check.generate.fill value(0.0) : tensor<1x2048xf32> + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<1x1xi32>, tensor<2xi32>) + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<1x4096xf32>, tensor<1x4608xi8>) + kernel.launch @qwen3_moe_routed_linear_q4k_f16_wmma[%token_count](%token_count, %input, %expert_table, %weight, %expected) : [index](index, tensor<1x4096xf32>, tensor<2xi32>, tensor<1x2048x16x144xi8>, tensor<1x2048xf32>) + kernel.launch @qwen3_moe_dense_linear_q4k_f16_wmma[%token_count](%token_count, %input, %weight, %actual) : [index](index, tensor<1x4096xf32>, tensor<1x2048x16x144xi8>, tensor<1x2048xf32>) + kernel.launch @qwen3_moe_dense_linear_q4k_q8_1_x4[%token_count](%token_count, %q8_input, %weight, %direct) : [index](index, tensor<1x4608xi8>, tensor<1x2048x16x144xi8>, tensor<1x2048xf32>) + check.expect.close actual(%actual) expected(%expected) atol(0.01) rtol(0.01) nan(same) : tensor<1x2048xf32> + check.expect.close actual(%direct) expected(%expected) atol(0.25) rtol(0.01) nan(same) : tensor<1x2048xf32> + check.return +} + +// Crosses the direct schedule's eight-wave output tile while comparing its +// Q8_1 dot path against the independently structured FP16 WMMA provider. +check.case public @qwen3_moe_dense_linear_q4k_q8_1_x4_differential_case { + %token_count = check.literal value(2) : index + %input_size = check.literal value(512) : index + %input = check.generate.fill value(0.00390625) : tensor<2x512xf32> + %q8_input = check.generate.fill value(0) : tensor<2x576xi8> + %weight = check.generate.iota offset(-72) step(1) period(144) : tensor<65x2x144xi8> + %expected = check.generate.fill value(0.0) : tensor<2x65xf32> + %actual = check.generate.fill value(0.0) : tensor<2x65xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<2x512xf32>, tensor<2x576xi8>) + kernel.launch @qwen3_moe_dense_linear_q4k_f16_wmma[%token_count](%token_count, %input, %weight, %expected) : [index](index, tensor<2x512xf32>, tensor<65x2x144xi8>, tensor<2x65xf32>) + kernel.launch @qwen3_moe_dense_linear_q4k_q8_1_x4[%token_count](%token_count, %q8_input, %weight, %actual) : [index](index, tensor<2x576xi8>, tensor<65x2x144xi8>, tensor<2x65xf32>) + check.expect.close actual(%actual) expected(%expected) atol(0.25) rtol(0.01) nan(same) : tensor<2x65xf32> + check.return +} + +// A zero projection leaves the nonzero caller-provided residual unchanged. +// Overwrite publication would instead produce zero, making the accumulation +// contract observable without coupling this case to another projection kernel. +check.case public @qwen3_moe_dense_linear_q4k_q8_1_x4_accumulation_case { + %token_count = check.literal value(2) : index + %input_size = check.literal value(512) : index + %input = check.generate.fill value(0.0) : tensor<2x512xf32> + %q8_input = check.generate.fill value(1) : tensor<2x576xi8> + %weight = check.generate.fill value(34) : tensor<65x2x144xi8> + %actual = check.generate.fill value(1.25) : tensor<2x65xf32> + %expected = check.generate.fill value(1.25) : tensor<2x65xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<2x512xf32>, tensor<2x576xi8>) + kernel.launch @qwen3_moe_dense_linear_q4k_q8_1_x4[%token_count](%token_count, %q8_input, %weight, %actual) : [index](index, tensor<2x576xi8>, tensor<65x2x144xi8>, tensor<2x65xf32>) + check.expect.close actual(%actual) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<2x65xf32> + check.return +} + +// The production output geometry requires all 256 tiles to publish the +// accumulated hidden row before the final workgroup can normalize and pack it. +// Reuse each oversized attention-input Q8 allocation in place, invoke the fused +// kernel twice through one completion word, and compare every semantic output +// with the ordinary two-dispatch composition. +check.case public @qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8_differential_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(4096) : index + %input = check.generate.iota offset(-0.5) step(0.000244140625) period(4096) : tensor<1x4096xf32> + %expected_q8 = check.generate.fill value(0) : tensor<4608xi8> + %actual_q8_0 = check.generate.fill value(0) : tensor<4608xi8> + %actual_q8_1 = check.generate.fill value(0) : tensor<4608xi8> + %weight = check.generate.fill value(34) : tensor<2048x16x144xi8> + %norm_weight = check.generate.iota offset(-1.0) step(0.0009765625) : tensor<2048xf32> + %expected_output = check.generate.iota offset(-0.25) step(0.000244140625) period(2048) : tensor<1x2048xf32> + %actual_output_0 = check.generate.iota offset(-0.25) step(0.000244140625) period(2048) : tensor<1x2048xf32> + %actual_output_1 = check.generate.iota offset(-0.25) step(0.000244140625) period(2048) : tensor<1x2048xf32> + %expected_normalized = check.generate.fill value(1.0) : tensor<1x2048xf32> + %actual_normalized_0 = check.generate.fill value(1.0) : tensor<1x2048xf32> + %actual_normalized_1 = check.generate.fill value(1.0) : tensor<1x2048xf32> + %completion_counter = check.generate.fill value(0) : tensor<1xi32> + %expected_counter = check.generate.fill value(0) : tensor<1xi32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %expected_q8) : [index, index](index, index, tensor<1x4096xf32>, tensor<4608xi8>) + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %actual_q8_0) : [index, index](index, index, tensor<1x4096xf32>, tensor<4608xi8>) + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %actual_q8_1) : [index, index](index, index, tensor<1x4096xf32>, tensor<4608xi8>) + kernel.launch @qwen3_moe_dense_linear_q4k_q8_1_x4[%token_count](%token_count, %expected_q8, %weight, %expected_output) : [index](index, tensor<4608xi8>, tensor<2048x16x144xi8>, tensor<1x2048xf32>) + kernel.launch @qwen3_moe_rmsnorm_f32_quantize_q8_1_x4[%token_count](%token_count, %expected_output, %norm_weight, %expected_normalized, %expected_q8) : [index](index, tensor<1x2048xf32>, tensor<2048xf32>, tensor<1x2048xf32>, tensor<4608xi8>) + kernel.launch @qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8[%token_count](%token_count, %actual_q8_0, %weight, %actual_output_0, %norm_weight, %actual_normalized_0, %completion_counter, %actual_q8_0) : [index](index, tensor<4608xi8>, tensor<2048x16x144xi8>, tensor<1x2048xf32>, tensor<2048xf32>, tensor<1x2048xf32>, tensor<1xi32>, tensor<4608xi8>) + kernel.launch @qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8[%token_count](%token_count, %actual_q8_1, %weight, %actual_output_1, %norm_weight, %actual_normalized_1, %completion_counter, %actual_q8_1) : [index](index, tensor<4608xi8>, tensor<2048x16x144xi8>, tensor<1x2048xf32>, tensor<2048xf32>, tensor<1x2048xf32>, tensor<1xi32>, tensor<4608xi8>) + check.expect.close actual(%actual_output_0) expected(%expected_output) atol(0.0) rtol(0.0) nan(same) : tensor<1x2048xf32> + check.expect.close actual(%actual_output_1) expected(%expected_output) atol(0.0) rtol(0.0) nan(same) : tensor<1x2048xf32> + check.expect.close actual(%actual_normalized_0) expected(%expected_normalized) atol(0.0) rtol(0.0) nan(same) : tensor<1x2048xf32> + check.expect.close actual(%actual_normalized_1) expected(%expected_normalized) atol(0.0) rtol(0.0) nan(same) : tensor<1x2048xf32> + check.expect.equal actual(%actual_q8_0) expected(%expected_q8) : tensor<4608xi8> + check.expect.equal actual(%actual_q8_1) expected(%expected_q8) : tensor<4608xi8> + check.expect.equal actual(%completion_counter) expected(%expected_counter) : tensor<1xi32> + check.return +} + +check.case public @qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8_benchmark_case { + %token_count = check.literal value(1) : index + %q8_input_and_output = check.generate.fill value(0) : tensor<4608xi8> + %weight = check.generate.fill value(0) : tensor<2048x16x144xi8> + %output = check.generate.fill value(1.0) : tensor<1x2048xf32> + %norm_weight = check.generate.fill value(1.0) : tensor<2048xf32> + %normalized_output = check.generate.fill value(0.0) : tensor<1x2048xf32> + %completion_counter = check.generate.fill value(0) : tensor<1xi32> + kernel.launch @qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8[%token_count](%token_count, %q8_input_and_output, %weight, %output, %norm_weight, %normalized_output, %completion_counter, %q8_input_and_output) : [index](index, tensor<4608xi8>, tensor<2048x16x144xi8>, tensor<1x2048xf32>, tensor<2048xf32>, tensor<1x2048xf32>, tensor<1xi32>, tensor<4608xi8>) + check.return +} + +check.case public @qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8_composed_benchmark_case { + %token_count = check.literal value(1) : index + %q8_input_and_output = check.generate.fill value(0) : tensor<4608xi8> + %weight = check.generate.fill value(0) : tensor<2048x16x144xi8> + %output = check.generate.fill value(1.0) : tensor<1x2048xf32> + %norm_weight = check.generate.fill value(1.0) : tensor<2048xf32> + %normalized_output = check.generate.fill value(0.0) : tensor<1x2048xf32> + kernel.launch @qwen3_moe_dense_linear_q4k_q8_1_x4[%token_count](%token_count, %q8_input_and_output, %weight, %output) : [index](index, tensor<4608xi8>, tensor<2048x16x144xi8>, tensor<1x2048xf32>) + kernel.launch @qwen3_moe_rmsnorm_f32_quantize_q8_1_x4[%token_count](%token_count, %output, %norm_weight, %normalized_output, %q8_input_and_output) : [index](index, tensor<1x2048xf32>, tensor<2048xf32>, tensor<1x2048xf32>, tensor<4608xi8>) + check.return +} + +check.case public @qwen3_moe_dense_linear_q4k_f16_wmma_q_projection_benchmark_case { + %token_count = check.param.choice values([1, 17, 32, 63, 128, 129, 512, 1024, 2048]) name("token_count") : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %weight = check.generate.fill value(0) : tensor<4096x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x4096xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x4096xf32> + kernel.launch @qwen3_moe_dense_linear_q4k_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<[%token_count]x2048xf32>, tensor<4096x8x144xi8>, tensor<[%token_count]x4096xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x4096xf32> + check.return +} + +check.case public @qwen3_moe_dense_linear_q4k_q8_1_x4_q_projection_benchmark_case { + %token_count = check.param.choice values([1, 17, 32, 63, 128, 129, 512, 1024, 2048]) name("token_count") : index + %input_size = check.literal value(2048) : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %q8_input = check.generate.fill value(1) : tensor<[%token_count]x2304xi8> + %weight = check.generate.fill value(0) : tensor<4096x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x4096xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x4096xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<[%token_count]x2048xf32>, tensor<[%token_count]x2304xi8>) + kernel.launch @qwen3_moe_dense_linear_q4k_q8_1_x4[%token_count](%token_count, %q8_input, %weight, %output) : [index](index, tensor<[%token_count]x2304xi8>, tensor<4096x8x144xi8>, tensor<[%token_count]x4096xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x4096xf32> + check.return +} + +check.case public @qwen3_moe_dense_linear_q4k_f16_wmma_k_projection_benchmark_case { + %token_count = check.param.choice values([1, 17, 32, 63, 128, 129, 512, 1024, 2048]) name("token_count") : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %weight = check.generate.fill value(0) : tensor<512x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x512xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x512xf32> + kernel.launch @qwen3_moe_dense_linear_q4k_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<[%token_count]x2048xf32>, tensor<512x8x144xi8>, tensor<[%token_count]x512xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x512xf32> + check.return +} + +check.case public @qwen3_moe_dense_linear_q4k_q8_1_x4_k_projection_benchmark_case { + %token_count = check.param.choice values([1, 17, 32, 63, 128, 129, 512, 1024, 2048]) name("token_count") : index + %input_size = check.literal value(2048) : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %q8_input = check.generate.fill value(1) : tensor<[%token_count]x2304xi8> + %weight = check.generate.fill value(0) : tensor<512x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x512xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x512xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<[%token_count]x2048xf32>, tensor<[%token_count]x2304xi8>) + kernel.launch @qwen3_moe_dense_linear_q4k_q8_1_x4[%token_count](%token_count, %q8_input, %weight, %output) : [index](index, tensor<[%token_count]x2304xi8>, tensor<512x8x144xi8>, tensor<[%token_count]x512xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x512xf32> + check.return +} + +check.case public @qwen3_moe_dense_linear_q6k_f16_wmma_v_projection_benchmark_case { + %token_count = check.param.choice values([1, 17, 32, 63, 128, 129, 512, 1024, 2048]) name("token_count") : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %weight = check.generate.fill value(0) : tensor<512x8x210xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x512xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x512xf32> + kernel.launch @qwen3_moe_dense_linear_q6k_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<[%token_count]x2048xf32>, tensor<512x8x210xi8>, tensor<[%token_count]x512xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x512xf32> + check.return +} + +check.case public @qwen3_moe_dense_linear_q4k_f16_wmma_o_projection_benchmark_case { + %token_count = check.param.choice values([1, 17, 32, 63, 128, 129, 512, 1024, 2048]) name("token_count") : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x4096xf32> + %weight = check.generate.fill value(0) : tensor<2048x16x144xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %expected = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + kernel.launch @qwen3_moe_dense_linear_q4k_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<[%token_count]x4096xf32>, tensor<2048x16x144xi8>, tensor<[%token_count]x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x2048xf32> + check.return +} + +check.case public @qwen3_moe_dense_linear_q4k_q8_1_x4_o_projection_benchmark_case { + %token_count = check.param.choice values([1, 17, 32, 63, 128, 129, 512, 1024, 2048]) name("token_count") : index + %input_size = check.literal value(4096) : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x4096xf32> + %q8_input = check.generate.fill value(1) : tensor<[%token_count]x4608xi8> + %weight = check.generate.fill value(0) : tensor<2048x16x144xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %expected = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<[%token_count]x4096xf32>, tensor<[%token_count]x4608xi8>) + kernel.launch @qwen3_moe_dense_linear_q4k_q8_1_x4[%token_count](%token_count, %q8_input, %weight, %output) : [index](index, tensor<[%token_count]x4608xi8>, tensor<2048x16x144xi8>, tensor<[%token_count]x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x2048xf32> + check.return +} + +check.benchmark<@qwen3_moe_dense_linear_q4k_f16_wmma_differential_case> @qwen3_moe_dense_linear_q4k_f16_wmma_differential + +check.benchmark<@qwen3_moe_dense_linear_q6k_f16_wmma_differential_case> @qwen3_moe_dense_linear_q6k_f16_wmma_differential + +check.benchmark<@qwen3_moe_dense_linear_q4k_f16_wmma_o_differential_case> @qwen3_moe_dense_linear_q4k_f16_wmma_o_differential + +check.benchmark<@qwen3_moe_dense_linear_q4k_q8_1_x4_differential_case> @qwen3_moe_dense_linear_q4k_q8_1_x4_differential + +check.benchmark<@qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8_benchmark_case> @qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8_decode + +check.benchmark<@qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8_composed_benchmark_case> @qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8_composed_decode + +check.benchmark<@qwen3_moe_dense_linear_q4k_f16_wmma_q_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_f16_wmma_q_decode {token_count = 1} + +check.benchmark<@qwen3_moe_dense_linear_q4k_f16_wmma_q_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_f16_wmma_q_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_dense_linear_q4k_f16_wmma_q_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_f16_wmma_q_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_dense_linear_q4k_f16_wmma_q_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_f16_wmma_q_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_dense_linear_q4k_q8_1_x4_q_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_q8_1_x4_q_decode {token_count = 1} + +check.benchmark<@qwen3_moe_dense_linear_q4k_q8_1_x4_q_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_q8_1_x4_q_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_dense_linear_q4k_q8_1_x4_q_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_q8_1_x4_q_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_dense_linear_q4k_q8_1_x4_q_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_q8_1_x4_q_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_dense_linear_q4k_f16_wmma_k_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_f16_wmma_k_decode {token_count = 1} + +check.benchmark<@qwen3_moe_dense_linear_q4k_f16_wmma_k_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_f16_wmma_k_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_dense_linear_q4k_f16_wmma_k_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_f16_wmma_k_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_dense_linear_q4k_f16_wmma_k_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_f16_wmma_k_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_dense_linear_q4k_q8_1_x4_k_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_q8_1_x4_k_decode {token_count = 1} + +check.benchmark<@qwen3_moe_dense_linear_q4k_q8_1_x4_k_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_q8_1_x4_k_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_dense_linear_q4k_q8_1_x4_k_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_q8_1_x4_k_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_dense_linear_q4k_q8_1_x4_k_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_q8_1_x4_k_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_dense_linear_q6k_f16_wmma_v_projection_benchmark_case> @qwen3_moe_dense_linear_q6k_f16_wmma_v_decode {token_count = 1} + +check.benchmark<@qwen3_moe_dense_linear_q6k_f16_wmma_v_projection_benchmark_case> @qwen3_moe_dense_linear_q6k_f16_wmma_v_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_dense_linear_q6k_f16_wmma_v_projection_benchmark_case> @qwen3_moe_dense_linear_q6k_f16_wmma_v_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_dense_linear_q6k_f16_wmma_v_projection_benchmark_case> @qwen3_moe_dense_linear_q6k_f16_wmma_v_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_dense_linear_q4k_f16_wmma_o_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_f16_wmma_o_decode {token_count = 1} + +check.benchmark<@qwen3_moe_dense_linear_q4k_f16_wmma_o_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_f16_wmma_o_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_dense_linear_q4k_f16_wmma_o_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_f16_wmma_o_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_dense_linear_q4k_f16_wmma_o_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_f16_wmma_o_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_dense_linear_q4k_q8_1_x4_o_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_q8_1_x4_o_decode {token_count = 1} + +check.benchmark<@qwen3_moe_dense_linear_q4k_q8_1_x4_o_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_q8_1_x4_o_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_dense_linear_q4k_q8_1_x4_o_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_q8_1_x4_o_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_dense_linear_q4k_q8_1_x4_o_projection_benchmark_case> @qwen3_moe_dense_linear_q4k_q8_1_x4_o_prefill_512 {token_count = 512} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/expert_table_partition_fused.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/expert_table_partition_fused.loom new file mode 100644 index 000000000000..ed315544ce6b --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/expert_table_partition_fused.loom @@ -0,0 +1,229 @@ +// Fuses the Prefill-512 expert assignment table and partition descriptor +// construction used by grouped routed projections. +// +// Each of 128 workgroups retains ownership of one expert row. After publishing +// its assignment count, lane zero releases a device-scope completion arrival. +// The last workgroup acquires all counts, compacts them into deterministic +// 32-row descriptors, and resets the counter so reusable command buffers can +// issue the kernel again against the same storage. +amdgpu.target @qwen3_moe_expert_table_partition_gfx11_wave32 {subgroup_size = 32} + +kernel.decl @qwen3_moe_build_expert_table(%token_count$0: index, %route_count$1: index, %route_stride$2: index, %expert_count$3: index) launch(%token_count$4: index, %route_count$5: index, %route_stride$6: index, %expert_count$7: index, %route_ids: buffer, %expert_table: buffer) + +kernel.decl @qwen3_moe_build_expert_partition_table(%token_count$10: index, %route_count$11: index, %expert_count$12: index) launch(%token_count$13: index, %route_count$14: index, %expert_count$15: index, %expert_table: buffer, %partition_table: buffer) + +kernel.def target(@qwen3_moe_expert_table_partition_gfx11_wave32) @qwen3_moe_build_expert_table_partition_prefill_512(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index) { + %c1 = index.constant 1 : index + %launch_expert_count = index.constant 128 : index + %workgroup_size = index.constant 256 : index + kernel.launch.config workgroups(%launch_expert_count, %c1, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index, %route_ids: buffer, %expert_table: buffer, %partition_table: buffer, %completion_counter: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 512, 512)] : index + %bounded_route_count = index.assume %route_count [range(%route_count, 8, 8)] : index + %bounded_route_stride = index.assume %route_stride [range(%route_stride, 8, 8)] : index + %bounded_expert_count = index.assume %expert_count [range(%expert_count, 128, 128)] : index + %expert0 = kernel.workgroup.id : index + %expert, %launch_expert_count = index.assume %expert0, %bounded_expert_count [lt(%expert0, %bounded_expert_count)] : index, index + %lane0 = kernel.workitem.id : index + %lane = index.assume %lane0 [range(%lane0, 0, 255)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %workgroup_size = index.constant 256 : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c13_i32 = scalar.constant 13 : i32 + %c0_offset = index.constant 0 : offset + %c4_bytes = index.constant 4 : offset + %assignment_count = index.mul %bounded_token_count, %bounded_route_count : index + %assignment_table_byte_base = index.scale %launch_expert_count, %c4_bytes : index, offset -> offset + %route_ids_noalias, %expert_table_noalias, %partition_table_noalias, %completion_counter_noalias = buffer.assume.noalias %route_ids, %expert_table, %partition_table, %completion_counter : buffer, buffer, buffer, buffer + %completion_counter_aligned = buffer.assume.alignment %completion_counter_noalias {minimum_alignment = 16} : buffer + %route_view = buffer.view %route_ids_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_route_stride]xi32> + %count_view = buffer.view %expert_table_noalias[%c0_offset] : buffer -> view<[%launch_expert_count]xi32> + %assignment_view = buffer.view %expert_table_noalias[%assignment_table_byte_base] : buffer -> view<[%launch_expert_count]x[%bounded_token_count]xi32> + %completion_counter_view = buffer.view %completion_counter_aligned[%c0_offset] : buffer -> view<1xi32> + + %expert_route_count = scf.for %block_base = [%c0 to %assignment_count step %workgroup_size](%matched_base = %c0_i32 : i32) -> (i32) { + %assignment = index.add %block_base, %lane : index + %in_range = index.cmp ult, %assignment, %assignment_count : index + %route_expert_i32 = scf.if %in_range -> (i32) { + %token0 = index.div %assignment, %bounded_route_count : index + %route0 = index.rem %assignment, %bounded_route_count : index + %token, %route_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %route, %route_row_stride = index.assume %route0, %bounded_route_stride [lt(%route0, %bounded_route_stride)] : index, index + %loaded = view.load %route_view[%token, %route] : view<[%bounded_token_count]x[%bounded_route_stride]xi32> -> i32 + scf.yield %loaded : i32 + } else { + %cn1_i32 = scalar.constant -1 : i32 + scf.yield %cn1_i32 : i32 + } + %route_expert0 = index.cast %route_expert_i32 : i32 to index + %route_expert = index.assume %route_expert0 [range(%route_expert0, -1, 127)] : index + %matches = index.cmp eq, %route_expert, %expert : index + %match_i32 = scf.if %matches -> (i32) { + scf.yield %c1_i32 : i32 + } else { + scf.yield %c0_i32 : i32 + } + %block_prefix = kernel.workgroup.scan %match_i32 {direction = forward, mode = exclusive} : i32 + %block_match_count_reduced = kernel.workgroup.reduce %match_i32 : i32 + %block_match_count = kernel.subgroup.broadcast.first %block_match_count_reduced : i32 + scf.if %matches { + %match_ordinal_i32 = scalar.addi %matched_base, %block_prefix : i32 + %match_ordinal0 = index.cast %match_ordinal_i32 : i32 to index + %match_ordinal = index.assume %match_ordinal0 [range(%match_ordinal0, 0, 511)] : index + %bounded_match_ordinal, %table_token_count = index.assume %match_ordinal, %bounded_token_count [lt(%match_ordinal, %bounded_token_count)] : index, index + %assignment_i32 = index.cast %assignment : index to i32 + view.store %assignment_i32, %assignment_view[%expert, %bounded_match_ordinal] : i32, view<[%launch_expert_count]x[%bounded_token_count]xi32> + } + %next_matched_base = scalar.addi %matched_base, %block_match_count : i32 + scf.yield %next_matched_base : i32 + } + + %is_lane_zero = index.cmp eq, %lane, %c0 : index + scf.if %is_lane_zero { + view.store %expert_route_count, %count_view[%expert] : i32, view<[%launch_expert_count]xi32> + } + + // Every count store precedes its workgroup's release. The last arrival + // acquires all preceding releases before any lane loads the complete table. + %local_old_counter = scf.if %is_lane_zero -> (i32) { + %old_counter = view.atomic.rmw %c1_i32, %completion_counter_view[%c0] {ordering = acq_rel, scope = device} : i32, view<1xi32> -> i32 + scf.yield %old_counter : i32 + } else { + scf.yield %c0_i32 : i32 + } + %old_counter = kernel.workgroup.reduce %local_old_counter : i32 + %expert_count_i32 = index.cast %launch_expert_count : index to i32 + %last_expert_i32 = scalar.subi %expert_count_i32, %c1_i32 : i32 + %negative_expert_count_i32 = scalar.subi %c0_i32, %expert_count_i32 : i32 + %is_last_expert = scalar.cmpi eq, %old_counter, %last_expert_i32 : i32 + scf.if %is_last_expert { + %rounded_assignment_count = index.add %assignment_count, %c31 : index + %assignment_partition_count = index.div %rounded_assignment_count, %c32 : index + %maximum_partition_count = index.add %assignment_partition_count, %launch_expert_count : index + %partition_count_view = buffer.view %partition_table_noalias[%c0_offset] : buffer -> view<1xi32> + %partition_descriptor_view = buffer.view %partition_table_noalias[%c4_bytes] : buffer -> view<[%maximum_partition_count]xi32> + %has_expert = index.cmp ult, %lane, %launch_expert_count : index + %expert_assignment_count_i32 = scf.if %has_expert -> (i32) { + %partition_expert, %table_expert_count = index.assume %lane, %launch_expert_count [lt(%lane, %launch_expert_count)] : index, index + %loaded = view.load %count_view[%partition_expert] : view<[%launch_expert_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %c0_i32 : i32 + } + %expert_assignment_count0 = index.cast %expert_assignment_count_i32 : i32 to index + %expert_assignment_count = index.assume %expert_assignment_count0 [range(%expert_assignment_count0, 0, 512)] : index + %rounded_expert_assignment_count = index.add %expert_assignment_count, %c31 : index + %expert_partition_count = index.div %rounded_expert_assignment_count, %c32 : index + %expert_partition_count_i32 = index.cast %expert_partition_count : index to i32 + %expert_partition_base_i32 = kernel.workgroup.scan %expert_partition_count_i32 {direction = forward, mode = exclusive} : i32 + %partition_count_i32 = kernel.workgroup.reduce %expert_partition_count_i32 : i32 + %partition_count = index.cast %partition_count_i32 : i32 to index + %bounded_partition_count, %table_partition_capacity = index.assume %partition_count, %maximum_partition_count [lt(%partition_count, %maximum_partition_count)] : index, index + %expert_partition_base0 = index.cast %expert_partition_base_i32 : i32 to index + %expert_partition_base = index.assume %expert_partition_base0 [range(%expert_partition_base0, 0, 255)] : index + scf.if %has_expert { + %partition_expert, %table_expert_count = index.assume %lane, %launch_expert_count [lt(%lane, %launch_expert_count)] : index, index + %expert_i32 = index.cast %partition_expert : index to i32 + scf.for %partition = [%c0 to %expert_partition_count step %c1] { + %descriptor_ordinal0 = index.add %expert_partition_base, %partition : index + %descriptor_ordinal, %descriptor_count = index.assume %descriptor_ordinal0, %bounded_partition_count [lt(%descriptor_ordinal0, %bounded_partition_count)] : index, index + %table_descriptor_ordinal, %table_descriptor_capacity = index.assume %descriptor_ordinal, %maximum_partition_count [lt(%descriptor_ordinal, %maximum_partition_count)] : index, index + %partition_remainder = index.rem %expert_assignment_count, %c32 : index + %has_partial_tail = index.cmp ne, %partition_remainder, %c0 : index + %partition_row_count = scf.if %has_partial_tail -> (index) { + %next_partition = index.add %partition, %c1 : index + %is_tail_partition = index.cmp eq, %next_partition, %expert_partition_count : index + %tail_row_count = scf.if %is_tail_partition -> (index) { + scf.yield %partition_remainder : index + } else { + scf.yield %c32 : index + } + scf.yield %tail_row_count : index + } else { + scf.yield %c32 : index + } + %partition_i32 = index.cast %partition : index to i32 + %partition_row_count_i32 = index.cast %partition_row_count : index to i32 + %packed_partition = scalar.shli %partition_i32, %c7_i32 : i32 + %partition_row_count_minus_one = scalar.subi %partition_row_count_i32, %c1_i32 : i32 + %packed_row_count = scalar.shli %partition_row_count_minus_one, %c13_i32 : i32 + %packed_expert_partition = scalar.ori %expert_i32, %packed_partition : i32 + %packed_descriptor = scalar.ori %packed_expert_partition, %packed_row_count : i32 + view.store %packed_descriptor, %partition_descriptor_view[%table_descriptor_ordinal] : i32, view<[%maximum_partition_count]xi32> + } + } + scf.if %is_lane_zero { + view.store %partition_count_i32, %partition_count_view[%c0] : i32, view<1xi32> + } + // The counter cannot become reusable until every descriptor store is + // complete and visible to the following dispatch. + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.if %is_lane_zero { + view.atomic.reduce %negative_expert_count_i32, %completion_counter_view[%c0] {ordering = release, scope = device} : i32, view<1xi32> + } + } + kernel.return +} + +// Compare both fused outputs against the production composition, then invoke +// the fused route twice against one counter to make reset correctness visible. +check.case public @qwen3_moe_expert_table_partition_fused_differential_case { + %token_count = check.literal value(512) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.iota offset(0) step(1) period(128) : tensor<512x8xi32> + %expected_expert_table = check.generate.fill value(-1) : tensor<65664xi32> + %expected_partition_table = check.generate.fill value(-1) : tensor<257xi32> + %actual_expert_table0 = check.generate.fill value(-1) : tensor<65664xi32> + %actual_partition_table0 = check.generate.fill value(-1) : tensor<257xi32> + %actual_expert_table1 = check.generate.fill value(-1) : tensor<65664xi32> + %actual_partition_table1 = check.generate.fill value(-1) : tensor<257xi32> + %completion_counter = check.generate.fill value(0) : tensor<1xi32> + %expected_counter = check.generate.fill value(0) : tensor<1xi32> + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expected_expert_table) : [index, index, index, index](index, index, index, index, tensor<512x8xi32>, tensor<65664xi32>) + kernel.launch @qwen3_moe_build_expert_partition_table[%token_count, %route_count, %expert_count](%token_count, %route_count, %expert_count, %expected_expert_table, %expected_partition_table) : [index, index, index](index, index, index, tensor<65664xi32>, tensor<257xi32>) + kernel.launch @qwen3_moe_build_expert_table_partition_prefill_512[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %actual_expert_table0, %actual_partition_table0, %completion_counter) : [index, index, index, index](index, index, index, index, tensor<512x8xi32>, tensor<65664xi32>, tensor<257xi32>, tensor<1xi32>) + kernel.launch @qwen3_moe_build_expert_table_partition_prefill_512[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %actual_expert_table1, %actual_partition_table1, %completion_counter) : [index, index, index, index](index, index, index, index, tensor<512x8xi32>, tensor<65664xi32>, tensor<257xi32>, tensor<1xi32>) + check.expect.equal actual(%actual_expert_table0) expected(%expected_expert_table) : tensor<65664xi32> + check.expect.equal actual(%actual_partition_table0) expected(%expected_partition_table) : tensor<257xi32> + check.expect.equal actual(%actual_expert_table1) expected(%expected_expert_table) : tensor<65664xi32> + check.expect.equal actual(%actual_partition_table1) expected(%expected_partition_table) : tensor<257xi32> + check.expect.equal actual(%completion_counter) expected(%expected_counter) : tensor<1xi32> + check.return +} + +check.case public @qwen3_moe_expert_table_partition_fused_benchmark_case { + %token_count = check.literal value(512) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.iota offset(0) step(1) period(128) : tensor<512x8xi32> + %expert_table = check.generate.fill value(-1) : tensor<65664xi32> + %partition_table = check.generate.fill value(-1) : tensor<257xi32> + %completion_counter = check.generate.fill value(0) : tensor<1xi32> + kernel.launch @qwen3_moe_build_expert_table_partition_prefill_512[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table, %partition_table, %completion_counter) : [index, index, index, index](index, index, index, index, tensor<512x8xi32>, tensor<65664xi32>, tensor<257xi32>, tensor<1xi32>) + check.return +} + +check.case public @qwen3_moe_expert_table_partition_composed_benchmark_case { + %token_count = check.literal value(512) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.iota offset(0) step(1) period(128) : tensor<512x8xi32> + %expert_table = check.generate.fill value(-1) : tensor<65664xi32> + %partition_table = check.generate.fill value(-1) : tensor<257xi32> + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<512x8xi32>, tensor<65664xi32>) + kernel.launch @qwen3_moe_build_expert_partition_table[%token_count, %route_count, %expert_count](%token_count, %route_count, %expert_count, %expert_table, %partition_table) : [index, index, index](index, index, index, tensor<65664xi32>, tensor<257xi32>) + check.return +} + +check.benchmark<@qwen3_moe_expert_table_partition_fused_benchmark_case> @qwen3_moe_expert_table_partition_fused_prefill_512 + +check.benchmark<@qwen3_moe_expert_table_partition_composed_benchmark_case> @qwen3_moe_expert_table_partition_composed_prefill_512 diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/flash_attention_decode_f32_f16_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/flash_attention_decode_f32_f16_wmma.loom new file mode 100644 index 000000000000..fc523520717a --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/flash_attention_decode_f32_f16_wmma.loom @@ -0,0 +1,690 @@ +// Qwen3 MoE grouped-query decode FlashAttention. +// +// One or two four-wave workgroups compute up to 16 query heads that share one +// KV head. Each workgroup walks the exact KV length in 64-row blocks while +// carrying online-softmax max, sum, and output state in registers. With two +// output partitions, each workgroup owns one disjoint 64-channel output half; +// this duplicates QK work but exposes more parallelism without publishing +// partial softmax state. The ownership changes mirror the cooperative-matrix +// schedule used by llama.cpp's Vulkan CM1 kernel: +// +// 1. All workitems stage a scaled 16x128 F16 GQA-head tile. +// 2. Each wave computes one 16x16 QK score slice. +// 3. Scores cross LDS so each wave can normalize four complete query heads. +// 4. F16 probabilities cross LDS for four P*V WMMA steps. +// 5. Each active lane retains one four-channel F16 packet for each of its +// four query heads across subsequent KV blocks. +// +// K and V remain in the row-major llama.cpp cache layout +// [KV token][KV head][128]. Their aligned F16 fragments load directly from +// global memory; there is no expanded or repacked persistent allocation. QK +// and the online-softmax statistics remain F32, while the P*V accumulation and +// carried output match the Vulkan oracle's F16 policy. Unlike the split-K +// fallback, this kernel never publishes partial tensors or completion atomics. +amdgpu.target @qwen3_moe_attention_decode_gfx11_wave64 {subgroup_size = 64} + +config.decl @qwen3_moe.attention.query_head_count : %value: index where [range(%value, 1, 64)] + +config.decl @qwen3_moe.attention.key_value_head_count : %value: index where [range(%value, 1, 64)] + +// Number of independent 64-channel output partitions per KV head. One avoids +// redundant QK work; two exposes more workgroups for short decode contexts. +config.decl @qwen3_moe.attention.decode.output_partition_count : %value: index where [range(%value, 1, 2)] + +kernel.def target(@qwen3_moe_attention_decode_gfx11_wave64) @qwen3_moe_flash_attention_decode_f32_f16_wmma(%key_value_token_count: index) { + %key_value_head_count = config.get @qwen3_moe.attention.key_value_head_count : index + %output_partition_count = config.get @qwen3_moe.attention.decode.output_partition_count : index + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + %workgroup_count = index.mul %key_value_head_count, %output_partition_count : index + kernel.launch.config workgroups(%workgroup_count, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %output: buffer) { + %bounded_key_value_token_count = index.assume %key_value_token_count [range(%key_value_token_count, 1, 32768)] : index + %query_head_count = config.get @qwen3_moe.attention.query_head_count : index + %key_value_head_count = config.get @qwen3_moe.attention.key_value_head_count : index + %output_partition_count = config.get @qwen3_moe.attention.decode.output_partition_count : index + %workgroup_x0 = kernel.workgroup.id : index + %workgroup_count0 = index.mul %key_value_head_count, %output_partition_count : index + %workgroup_x, %workgroup_count = index.assume %workgroup_x0, %workgroup_count0 [lt(%workgroup_x0, %workgroup_count0)] : index, index + %key_value_head0 = index.div %workgroup_x, %output_partition_count : index + %key_value_head = index.assume %key_value_head0 [range(%key_value_head0, 0, 63)] : index + %output_partition = index.rem %workgroup_x, %output_partition_count : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 3)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c15 = index.constant 15 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %query_stage_bytes = index.constant 4352 : offset + %score_stage_bytes = index.constant 6144 : offset + %probability_stage_bytes = index.constant 3072 : offset + %product_stage_bytes = index.constant 2048 : offset + %tail_key_value_stage_capacity = index.constant 8192 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %head_size_f32 = scalar.constant 128.0 : f32 + %attention_scale = scalar.rsqrtf %head_size_f32 : f32 + %c0_f16 = scalar.constant 0.0 : f16 + %c1_f32 = scalar.constant 1.0 : f32 + %output_zero0 = vector.constant 0.0 : vector<4xf16> + %output_zero1 = vector.constant 0.0 : vector<4xf16> + %output_zero2 = vector.constant 0.0 : vector<4xf16> + %output_zero3 = vector.constant 0.0 : vector<4xf16> + %c0_f16x8 = vector.constant 0.0 : vector<8xf16> + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %negative_f32x4 = vector.constant -1e+30 : vector<4xf32> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %query_heads_per_key_value_head0 = index.div %query_head_count, %key_value_head_count : index + %query_heads_per_key_value_head = index.assume %query_heads_per_key_value_head0 [range(%query_heads_per_key_value_head0, 1, 16)] : index + %query_head_base = index.mul %key_value_head, %query_heads_per_key_value_head : index + %key_value_width = index.mul %key_value_head_count, %c128 : index + %key_value_head_base = index.mul %key_value_head, %c128 : index + %full_key_value_block_count = index.div %bounded_key_value_token_count, %c64 : index + %full_key_value_token_count0 = index.mul %full_key_value_block_count, %c64 : index + %full_key_value_token_count = index.assume %full_key_value_token_count0 [range(%full_key_value_token_count0, 0, 32768), mul(%full_key_value_token_count0, 64)] : index + %tail_key_value_token_count = index.sub %bounded_key_value_token_count, %full_key_value_token_count : index + %has_key_value_tail = index.cmp ne, %tail_key_value_token_count, %c0 : index + %has_single_key_value_tail = index.cmp eq, %tail_key_value_token_count, %c1 : index + %tail_key_value_stage_bytes = scf.select %has_key_value_tail, %tail_key_value_stage_capacity, %c0_offset : offset + %subgroup_score_column = index.mul %subgroup, %c16 : index + %subgroup_query_row = index.mul %subgroup, %c4 : index + %query_row0 = index.add %subgroup_query_row, %c0 : index + %query_row1 = index.add %subgroup_query_row, %c1 : index + %query_row2 = index.add %subgroup_query_row, %c2 : index + %query_row3 = index.add %subgroup_query_row, %c3 : index + %query_head0 = index.add %query_head_base, %query_row0 : index + %query_head1 = index.add %query_head_base, %query_row1 : index + %query_head2 = index.add %query_head_base, %query_row2 : index + %query_head3 = index.add %query_head_base, %query_row3 : index + %query_valid0 = index.cmp ult, %query_head0, %query_head_count : index + %query_valid1 = index.cmp ult, %query_head1, %query_head_count : index + %query_valid2 = index.cmp ult, %query_head2, %query_head_count : index + %query_valid3 = index.cmp ult, %query_head3, %query_head_count : index + %query_valid = vector.from_elements %query_valid0, %query_valid1, %query_valid2, %query_valid3 : vector<4xi1> + %subgroup_product_channel = index.mul %subgroup, %c16 : index + %lane_output_tile = index.div %lane, %c16 : index + %output_tile_count = index.div %c2, %output_partition_count : index + %output_tile_end = index.add %output_partition, %output_tile_count : index + %lane_product_channel0 = index.rem %lane, %c16 : index + %lane_product_channel = index.mul %lane_product_channel0, %c4 : index + %lane_output_channel = index.mul %lane, %c4 : index + %lane_output_at_or_after_partition = index.cmp uge, %lane_output_tile, %output_partition : index + %lane_output_before_end = index.cmp ult, %lane_output_tile, %output_tile_end : index + %lane_has_output = scalar.andi %lane_output_at_or_after_partition, %lane_output_before_end : i1 + // The padded LDS rows mirror the Vulkan oracle's ownership changes. Q uses + // eight spare F16 columns after its 128 channels. Score and probability + // transpose to key-major rows with eight spare columns after 16 queries. + // These strides avoid the bank pattern produced by dense transposed rows. + %query_transposed_layout = encoding.layout.strided [1, 136] : encoding + %probability_transposed_layout = encoding.layout.strided [1, 24] : encoding + %query_noalias, %key_noalias, %value_noalias, %mask_noalias, %output_noalias = buffer.assume.noalias %query, %key, %value, %mask, %output : buffer, buffer, buffer, buffer, buffer + %query_aligned = buffer.assume.alignment %query_noalias {minimum_alignment = 16} : buffer + %key_aligned = buffer.assume.alignment %key_noalias {minimum_alignment = 16} : buffer + %value_aligned = buffer.assume.alignment %value_noalias {minimum_alignment = 16} : buffer + %mask_aligned = buffer.assume.alignment %mask_noalias {minimum_alignment = 16} : buffer + %output_aligned = buffer.assume.alignment %output_noalias {minimum_alignment = 16} : buffer + %query_view = buffer.view %query_aligned[%c0_offset] : buffer -> view<[%query_head_count]x128xf32> + %key_view = buffer.view %key_aligned[%c0_offset] : buffer -> view<[%bounded_key_value_token_count]x[%key_value_width]xf16> + %value_view = buffer.view %value_aligned[%c0_offset] : buffer -> view<[%bounded_key_value_token_count]x[%key_value_width]xf16> + %mask_view = buffer.view %mask_aligned[%c0_offset] : buffer -> view<[%bounded_key_value_token_count]xf16> + %output_view = buffer.view %output_aligned[%c0_offset] : buffer -> view<[%query_head_count]x128xf32> + %query_stage = buffer.alloca align(16) %query_stage_bytes : buffer + %score_stage = buffer.alloca align(16) %score_stage_bytes : buffer + %probability_stage = buffer.alloca align(16) %probability_stage_bytes : buffer + %product_stage = buffer.alloca align(16) %product_stage_bytes : buffer + %tail_key_value_stage = buffer.alloca align(16) %tail_key_value_stage_bytes : buffer + %query_stage_view = buffer.view %query_stage[%c0_offset] : buffer -> view<16x136xf16> + %query_transposed_view = buffer.view %query_stage[%c0_offset] : buffer -> view<128x16xf16, %query_transposed_layout> + %score_stage_view = buffer.view %score_stage[%c0_offset] : buffer -> view<64x24xf32> + %probability_stage_view = buffer.view %probability_stage[%c0_offset] : buffer -> view<16x64xf16, %probability_transposed_layout> + %product_stage_view = buffer.view %product_stage[%c0_offset] : buffer -> view<16x64xf16> + %tail_key_value_stage_view = buffer.view %tail_key_value_stage[%c0_offset] : buffer -> view<32x128xf16> + // Scale and truncate Q exactly once. The Vulkan reference does this before + // entering its KV loop, making QK a native F16 WMMA while retaining F32 + // accumulation. + scf.for %load_iteration = [%c0 to %c8 step %c1] unroll { + %linear = index.madd %load_iteration, %c256, %workitem : index + %local_query_row = index.div %linear, %c128 : index + %query_channel = index.rem %linear, %c128 : index + %local_query_head = index.add %query_head_base, %local_query_row : index + %query_valid_load = index.cmp ult, %local_query_head, %query_head_count : index + %query_value = scf.if %query_valid_load -> (f16) { + %loaded = view.load %query_view[%local_query_head, %query_channel] : view<[%query_head_count]x128xf32> -> f32 + %scaled = scalar.mulf %loaded, %attention_scale : f32 + %truncated = scalar.fptrunc %scaled : f32 to f16 + scf.yield %truncated : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %query_value, %query_stage_view[%local_query_row, %query_channel] : f16, view<16x136xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %full_max, %full_sum, %full_output0, %full_output1, %full_output2, %full_output3 = scf.for %key_origin = [%c0 to %full_key_value_token_count step %c64](%current_max = %negative_f32x4 : vector<4xf32>, %current_sum = %c0_f32x4 : vector<4xf32>, %current_output0 = %output_zero0 : vector<4xf16>, %current_output1 = %output_zero1 : vector<4xf16>, %current_output2 = %output_zero2 : vector<4xf16>, %current_output3 = %output_zero3 : vector<4xf16>) -> (vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + // Four independent wave-level WMMAs produce a 16x64 score tile. + %score_key_origin0 = index.add %key_origin, %subgroup_score_column : index + %last_full_key_tile_start = index.sub %bounded_key_value_token_count, %c15 : index + %score_key_origin = index.assume %score_key_origin0 [lt(%score_key_origin0, %last_full_key_tile_start)] : index + %score_init_values = vector.constant 0.0 : vector<4xf32> + %score_init = vector.fragment %score_init_values shape [%m, %n] : vector<4xf32> + %score_fragment = scf.for %head_tile = [%c0 to %c128 step %c16](%score_accumulator = %score_init : vector<4xf32>) -> (vector<4xf32>) unroll { + %key_channel = index.add %key_value_head_base, %head_tile : index + %key_fragment = vector.fragment.load %key_view[%score_key_origin, %key_channel] shape [%m, %k] : view<[%bounded_key_value_token_count]x[%key_value_width]xf16> -> vector<16xf16> + %query_fragment = vector.fragment.load %query_transposed_view[%head_tile, %c0] shape [%k, %n] : view<128x16xf16, %query_transposed_layout> -> vector<16xf16> + %next_score_accumulator = vector.mma %key_fragment, %query_fragment, %score_accumulator : vector<16xf16>, vector<16xf16>, vector<4xf32> + scf.yield %next_score_accumulator : vector<4xf32> + } + vector.fragment.store %score_fragment, %score_stage_view[%subgroup_score_column, %c0] shape [%m, %n] : vector<4xf32>, view<64x24xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + // LDS transposes ownership from one 16-column score slice per wave to + // four complete query rows per wave. Every lane contributes one key + // column to each of those rows. + %key_token0 = index.add %key_origin, %lane : index + %key_token = index.assume %key_token0 [lt(%key_token0, %bounded_key_value_token_count)] : index + %raw_score0 = view.load %score_stage_view[%lane, %query_row0] : view<64x24xf32> -> f32 + %raw_score1 = view.load %score_stage_view[%lane, %query_row1] : view<64x24xf32> -> f32 + %raw_score2 = view.load %score_stage_view[%lane, %query_row2] : view<64x24xf32> -> f32 + %raw_score3 = view.load %score_stage_view[%lane, %query_row3] : view<64x24xf32> -> f32 + %mask_f16 = view.load %mask_view[%key_token] : view<[%bounded_key_value_token_count]xf16> -> f16 + %mask_f32 = scalar.extf %mask_f16 : f16 to f32 + %masked_score0 = scf.if %query_valid0 -> (f32) { + %score = scalar.addf %raw_score0, %mask_f32 : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score1 = scf.if %query_valid1 -> (f32) { + %score = scalar.addf %raw_score1, %mask_f32 : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score2 = scf.if %query_valid2 -> (f32) { + %score = scalar.addf %raw_score2, %mask_f32 : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score3 = scf.if %query_valid3 -> (f32) { + %score = scalar.addf %raw_score3, %mask_f32 : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_scores = vector.from_elements %masked_score0, %masked_score1, %masked_score2, %masked_score3 : vector<4xf32> + %block_max = kernel.subgroup.reduce %masked_scores : vector<4xf32> + %next_max = vector.maxnumf %current_max, %block_max : vector<4xf32> + %score_delta = vector.subf %masked_scores, %next_max : vector<4xf32> + %raw_probability = vector.expf %score_delta : vector<4xf32> + %probability = vector.select %query_valid, %raw_probability, %c0_f32x4 : vector<4xf32> + %block_sum = kernel.subgroup.reduce %probability : vector<4xf32> + %old_delta = vector.subf %current_max, %next_max : vector<4xf32> + %old_scale = vector.expf %old_delta : vector<4xf32> + %scaled_current_sum = vector.mulf %current_sum, %old_scale : vector<4xf32> + %next_sum = vector.addf %scaled_current_sum, %block_sum : vector<4xf32> + %probability_f16 = vector.fptrunc %probability : vector<4xf32> to vector<4xf16> + %probability0 = vector.extract %probability_f16[0] : vector<4xf16> -> f16 + %probability1 = vector.extract %probability_f16[1] : vector<4xf16> -> f16 + %probability2 = vector.extract %probability_f16[2] : vector<4xf16> -> f16 + %probability3 = vector.extract %probability_f16[3] : vector<4xf16> -> f16 + view.store %probability0, %probability_stage_view[%query_row0, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability1, %probability_stage_view[%query_row1, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability2, %probability_stage_view[%query_row2, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability3, %probability_stage_view[%query_row3, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + kernel.barrier scope(workgroup) ordering(acq_rel) + // Match the Vulkan CM1 ownership schedule by computing two sequential + // 64-channel output tiles. This keeps one P*V accumulator live per wave + // and reuses a 2 KiB exchange tile instead of retaining both halves. + %old_scale_f16 = vector.fptrunc %old_scale : vector<4xf32> to vector<4xf16> + %old_scale0_scalar = vector.extract %old_scale_f16[0] : vector<4xf16> -> f16 + %old_scale1_scalar = vector.extract %old_scale_f16[1] : vector<4xf16> -> f16 + %old_scale2_scalar = vector.extract %old_scale_f16[2] : vector<4xf16> -> f16 + %old_scale3_scalar = vector.extract %old_scale_f16[3] : vector<4xf16> -> f16 + %old_scale0 = vector.splat %old_scale0_scalar : vector<4xf16> + %old_scale1 = vector.splat %old_scale1_scalar : vector<4xf16> + %old_scale2 = vector.splat %old_scale2_scalar : vector<4xf16> + %old_scale3 = vector.splat %old_scale3_scalar : vector<4xf16> + %scaled_current_output0 = vector.mulf %current_output0, %old_scale0 : vector<4xf16> + %scaled_current_output1 = vector.mulf %current_output1, %old_scale1 : vector<4xf16> + %scaled_current_output2 = vector.mulf %current_output2, %old_scale2 : vector<4xf16> + %scaled_current_output3 = vector.mulf %current_output3, %old_scale3 : vector<4xf16> + %next_output0, %next_output1, %next_output2, %next_output3 = scf.for %output_tile = [%c0 to %c2 step %c1](%tile_output0 = %scaled_current_output0 : vector<4xf16>, %tile_output1 = %scaled_current_output1 : vector<4xf16>, %tile_output2 = %scaled_current_output2 : vector<4xf16>, %tile_output3 = %scaled_current_output3 : vector<4xf16>) -> (vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) unroll { + %tile_at_or_after_partition = index.cmp uge, %output_tile, %output_partition : index + %tile_before_partition_end = index.cmp ult, %output_tile, %output_tile_end : index + %tile_selected = scalar.andi %tile_at_or_after_partition, %tile_before_partition_end : i1 + %updated_output0, %updated_output1, %updated_output2, %updated_output3 = scf.if %tile_selected -> (vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + %output_tile_channel = index.mul %output_tile, %c64 : index + %value_channel0 = index.add %key_value_head_base, %output_tile_channel : index + %value_channel = index.add %value_channel0, %subgroup_product_channel : index + %product_init = vector.fragment %c0_f16x8 shape [%m, %n] : vector<8xf16> + %product_fragment = scf.for %key_tile = [%c0 to %c64 step %c16](%product_accumulator = %product_init : vector<8xf16>) -> (vector<8xf16>) unroll { + %value_token0 = index.add %key_origin, %key_tile : index + %value_token = index.assume %value_token0 [lt(%value_token0, %last_full_key_tile_start)] : index + %probability_fragment = vector.fragment.load %probability_stage_view[%c0, %key_tile] shape [%m, %k] : view<16x64xf16, %probability_transposed_layout> -> vector<16xf16> + %value_fragment = vector.fragment.load %value_view[%value_token, %value_channel] shape [%k, %n] : view<[%bounded_key_value_token_count]x[%key_value_width]xf16> -> vector<16xf16> + %next_product_accumulator = vector.mma %probability_fragment, %value_fragment, %product_accumulator : vector<16xf16>, vector<16xf16>, vector<8xf16> + scf.yield %next_product_accumulator : vector<8xf16> + } + vector.fragment.store %product_fragment, %product_stage_view[%c0, %subgroup_product_channel] shape [%m, %n] : vector<8xf16>, view<16x64xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + %owns_output_tile = index.cmp eq, %lane_output_tile, %output_tile : index + %next_tile_output0, %next_tile_output1, %next_tile_output2, %next_tile_output3 = scf.if %owns_output_tile -> (vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + %block_output0 = vector.load %product_stage_view[%query_row0, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output1 = vector.load %product_stage_view[%query_row1, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output2 = vector.load %product_stage_view[%query_row2, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output3 = vector.load %product_stage_view[%query_row3, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %next_tile_output0 = vector.addf %tile_output0, %block_output0 : vector<4xf16> + %next_tile_output1 = vector.addf %tile_output1, %block_output1 : vector<4xf16> + %next_tile_output2 = vector.addf %tile_output2, %block_output2 : vector<4xf16> + %next_tile_output3 = vector.addf %tile_output3, %block_output3 : vector<4xf16> + scf.yield %next_tile_output0, %next_tile_output1, %next_tile_output2, %next_tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } else { + scf.yield %tile_output0, %tile_output1, %tile_output2, %tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + // Complete every read before another output tile overwrites LDS. + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %next_tile_output0, %next_tile_output1, %next_tile_output2, %next_tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } else { + scf.yield %tile_output0, %tile_output1, %tile_output2, %tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + scf.yield %updated_output0, %updated_output1, %updated_output2, %updated_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + scf.yield %next_max, %next_sum, %next_output0, %next_output1, %next_output2, %next_output3 : vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + // A single trailing KV row is cheaper as a native wave reduction and direct + // V update. Skip the general WMMA tail loop for that exact JIT-specialized + // shape; larger tails continue through the masked 32-row schedule below. + %wmma_tail_start = scf.select %has_single_key_value_tail, %bounded_key_value_token_count, %full_key_value_token_count : index + %tail_score_wave = index.cmp ult, %subgroup, %c2 : index + %wmma_tail_max, %wmma_tail_sum, %wmma_tail_output0, %wmma_tail_output1, %wmma_tail_output2, %wmma_tail_output3 = scf.for %tail_key_origin = [%wmma_tail_start to %bounded_key_value_token_count step %c32](%current_max = %full_max : vector<4xf32>, %current_sum = %full_sum : vector<4xf32>, %current_output0 = %full_output0 : vector<4xf16>, %current_output1 = %full_output1 : vector<4xf16>, %current_output2 = %full_output2 : vector<4xf16>, %current_output3 = %full_output3 : vector<4xf16>) -> (vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + %tail_remaining = index.sub %bounded_key_value_token_count, %tail_key_origin : index + %tail_key_count = index.min %tail_remaining, %c32 : index + // Cooperatively stage one K tile, explicitly zeroing the padded rows. + scf.for %load_iteration = [%c0 to %c16 step %c1] unroll { + %linear = index.madd %load_iteration, %c256, %workitem : index + %tail_key_row = index.div %linear, %c128 : index + %tail_key_channel = index.rem %linear, %c128 : index + %tail_key_valid = index.cmp ult, %tail_key_row, %tail_key_count : index + %tail_key_value = scf.if %tail_key_valid -> (f16) { + %tail_key_token0 = index.add %tail_key_origin, %tail_key_row : index + %tail_key_token = index.assume %tail_key_token0 [lt(%tail_key_token0, %bounded_key_value_token_count)] : index + %global_key_channel = index.add %key_value_head_base, %tail_key_channel : index + %loaded = view.load %key_view[%tail_key_token, %global_key_channel] : view<[%bounded_key_value_token_count]x[%key_value_width]xf16> -> f16 + scf.yield %loaded : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %tail_key_value, %tail_key_value_stage_view[%tail_key_row, %tail_key_channel] : f16, view<32x128xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + // Waves zero and one compute the two 16-column QK fragments in this tile. + scf.if %tail_score_wave { + %tail_score_subgroup = index.assume %subgroup [range(%subgroup, 0, 1)] : index + %tail_score_column = index.mul %tail_score_subgroup, %c16 : index + %tail_score_init_values = vector.constant 0.0 : vector<4xf32> + %tail_score_init = vector.fragment %tail_score_init_values shape [%m, %n] : vector<4xf32> + %tail_score_fragment = scf.for %head_tile = [%c0 to %c128 step %c16](%score_accumulator = %tail_score_init : vector<4xf32>) -> (vector<4xf32>) unroll { + %key_fragment = vector.fragment.load %tail_key_value_stage_view[%tail_score_column, %head_tile] shape [%m, %k] : view<32x128xf16> -> vector<16xf16> + %query_fragment = vector.fragment.load %query_transposed_view[%head_tile, %c0] shape [%k, %n] : view<128x16xf16, %query_transposed_layout> -> vector<16xf16> + %next_score_accumulator = vector.mma %key_fragment, %query_fragment, %score_accumulator : vector<16xf16>, vector<16xf16>, vector<4xf32> + scf.yield %next_score_accumulator : vector<4xf32> + } + vector.fragment.store %tail_score_fragment, %score_stage_view[%tail_score_column, %c0] shape [%m, %n] : vector<4xf32>, view<64x24xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + // LDS changes ownership from the QK wave to four query-row waves. Lanes + // beyond the logical tail never read the score or mask buffers. + %tail_lane_valid = index.cmp ult, %lane, %tail_key_count : index + %tail_valid0 = scalar.andi %tail_lane_valid, %query_valid0 : i1 + %tail_valid1 = scalar.andi %tail_lane_valid, %query_valid1 : i1 + %tail_valid2 = scalar.andi %tail_lane_valid, %query_valid2 : i1 + %tail_valid3 = scalar.andi %tail_lane_valid, %query_valid3 : i1 + %tail_valid = vector.from_elements %tail_valid0, %tail_valid1, %tail_valid2, %tail_valid3 : vector<4xi1> + %tail_key_token0 = index.add %tail_key_origin, %lane : index + %tail_mask_f32 = scf.if %tail_lane_valid -> (f32) { + %tail_key_token = index.assume %tail_key_token0 [lt(%tail_key_token0, %bounded_key_value_token_count)] : index + %mask_f16 = view.load %mask_view[%tail_key_token] : view<[%bounded_key_value_token_count]xf16> -> f16 + %mask_f32 = scalar.extf %mask_f16 : f16 to f32 + scf.yield %mask_f32 : f32 + } else { + scf.yield %c0_f32 : f32 + } + %masked_score0 = scf.if %tail_valid0 -> (f32) { + %raw_score = view.load %score_stage_view[%lane, %query_row0] : view<64x24xf32> -> f32 + %score = scalar.addf %raw_score, %tail_mask_f32 : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score1 = scf.if %tail_valid1 -> (f32) { + %raw_score = view.load %score_stage_view[%lane, %query_row1] : view<64x24xf32> -> f32 + %score = scalar.addf %raw_score, %tail_mask_f32 : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score2 = scf.if %tail_valid2 -> (f32) { + %raw_score = view.load %score_stage_view[%lane, %query_row2] : view<64x24xf32> -> f32 + %score = scalar.addf %raw_score, %tail_mask_f32 : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score3 = scf.if %tail_valid3 -> (f32) { + %raw_score = view.load %score_stage_view[%lane, %query_row3] : view<64x24xf32> -> f32 + %score = scalar.addf %raw_score, %tail_mask_f32 : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_scores = vector.from_elements %masked_score0, %masked_score1, %masked_score2, %masked_score3 : vector<4xf32> + %block_max = kernel.subgroup.reduce %masked_scores : vector<4xf32> + %next_max = vector.maxnumf %current_max, %block_max : vector<4xf32> + %score_delta = vector.subf %masked_scores, %next_max : vector<4xf32> + %raw_probability = vector.expf %score_delta : vector<4xf32> + %probability = vector.select %tail_valid, %raw_probability, %c0_f32x4 : vector<4xf32> + %block_sum = kernel.subgroup.reduce %probability : vector<4xf32> + %old_delta = vector.subf %current_max, %next_max : vector<4xf32> + %old_scale = vector.expf %old_delta : vector<4xf32> + %scaled_current_sum = vector.mulf %current_sum, %old_scale : vector<4xf32> + %next_sum = vector.addf %scaled_current_sum, %block_sum : vector<4xf32> + %probability_f16 = vector.fptrunc %probability : vector<4xf32> to vector<4xf16> + %probability0 = vector.extract %probability_f16[0] : vector<4xf16> -> f16 + %probability1 = vector.extract %probability_f16[1] : vector<4xf16> -> f16 + %probability2 = vector.extract %probability_f16[2] : vector<4xf16> -> f16 + %probability3 = vector.extract %probability_f16[3] : vector<4xf16> -> f16 + view.store %probability0, %probability_stage_view[%query_row0, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability1, %probability_stage_view[%query_row1, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability2, %probability_stage_view[%query_row2, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability3, %probability_stage_view[%query_row3, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + kernel.barrier scope(workgroup) ordering(acq_rel) + // Reuse product scratch for V after every probability is resident in its + // disjoint LDS tile. + scf.for %load_iteration = [%c0 to %c16 step %c1] unroll { + %linear = index.madd %load_iteration, %c256, %workitem : index + %tail_value_row = index.div %linear, %c128 : index + %tail_value_channel = index.rem %linear, %c128 : index + %tail_value_valid = index.cmp ult, %tail_value_row, %tail_key_count : index + %tail_value = scf.if %tail_value_valid -> (f16) { + %tail_value_token0 = index.add %tail_key_origin, %tail_value_row : index + %tail_value_token = index.assume %tail_value_token0 [lt(%tail_value_token0, %bounded_key_value_token_count)] : index + %global_value_channel = index.add %key_value_head_base, %tail_value_channel : index + %loaded = view.load %value_view[%tail_value_token, %global_value_channel] : view<[%bounded_key_value_token_count]x[%key_value_width]xf16> -> f16 + scf.yield %loaded : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %tail_value, %tail_key_value_stage_view[%tail_value_row, %tail_value_channel] : f16, view<32x128xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + // Keep tail V staging separate from the 2 KiB product exchange, then use + // the same two 64-channel phases as the aligned path. + %old_scale_f16 = vector.fptrunc %old_scale : vector<4xf32> to vector<4xf16> + %old_scale0_scalar = vector.extract %old_scale_f16[0] : vector<4xf16> -> f16 + %old_scale1_scalar = vector.extract %old_scale_f16[1] : vector<4xf16> -> f16 + %old_scale2_scalar = vector.extract %old_scale_f16[2] : vector<4xf16> -> f16 + %old_scale3_scalar = vector.extract %old_scale_f16[3] : vector<4xf16> -> f16 + %old_scale0 = vector.splat %old_scale0_scalar : vector<4xf16> + %old_scale1 = vector.splat %old_scale1_scalar : vector<4xf16> + %old_scale2 = vector.splat %old_scale2_scalar : vector<4xf16> + %old_scale3 = vector.splat %old_scale3_scalar : vector<4xf16> + %scaled_current_output0 = vector.mulf %current_output0, %old_scale0 : vector<4xf16> + %scaled_current_output1 = vector.mulf %current_output1, %old_scale1 : vector<4xf16> + %scaled_current_output2 = vector.mulf %current_output2, %old_scale2 : vector<4xf16> + %scaled_current_output3 = vector.mulf %current_output3, %old_scale3 : vector<4xf16> + %next_output0, %next_output1, %next_output2, %next_output3 = scf.for %output_tile = [%c0 to %c2 step %c1](%tile_output0 = %scaled_current_output0 : vector<4xf16>, %tile_output1 = %scaled_current_output1 : vector<4xf16>, %tile_output2 = %scaled_current_output2 : vector<4xf16>, %tile_output3 = %scaled_current_output3 : vector<4xf16>) -> (vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) unroll { + %tile_at_or_after_partition = index.cmp uge, %output_tile, %output_partition : index + %tile_before_partition_end = index.cmp ult, %output_tile, %output_tile_end : index + %tile_selected = scalar.andi %tile_at_or_after_partition, %tile_before_partition_end : i1 + %updated_output0, %updated_output1, %updated_output2, %updated_output3 = scf.if %tile_selected -> (vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + %output_tile_channel = index.mul %output_tile, %c64 : index + %value_channel = index.add %output_tile_channel, %subgroup_product_channel : index + %tail_product_init = vector.fragment %c0_f16x8 shape [%m, %n] : vector<8xf16> + %tail_product_fragment = scf.for %key_tile = [%c0 to %c32 step %c16](%product_accumulator = %tail_product_init : vector<8xf16>) -> (vector<8xf16>) unroll { + %probability_fragment = vector.fragment.load %probability_stage_view[%c0, %key_tile] shape [%m, %k] : view<16x64xf16, %probability_transposed_layout> -> vector<16xf16> + %value_fragment = vector.fragment.load %tail_key_value_stage_view[%key_tile, %value_channel] shape [%k, %n] : view<32x128xf16> -> vector<16xf16> + %next_product_accumulator = vector.mma %probability_fragment, %value_fragment, %product_accumulator : vector<16xf16>, vector<16xf16>, vector<8xf16> + scf.yield %next_product_accumulator : vector<8xf16> + } + vector.fragment.store %tail_product_fragment, %product_stage_view[%c0, %subgroup_product_channel] shape [%m, %n] : vector<8xf16>, view<16x64xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + %owns_output_tile = index.cmp eq, %lane_output_tile, %output_tile : index + %next_tile_output0, %next_tile_output1, %next_tile_output2, %next_tile_output3 = scf.if %owns_output_tile -> (vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + %block_output0 = vector.load %product_stage_view[%query_row0, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output1 = vector.load %product_stage_view[%query_row1, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output2 = vector.load %product_stage_view[%query_row2, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output3 = vector.load %product_stage_view[%query_row3, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %next_tile_output0 = vector.addf %tile_output0, %block_output0 : vector<4xf16> + %next_tile_output1 = vector.addf %tile_output1, %block_output1 : vector<4xf16> + %next_tile_output2 = vector.addf %tile_output2, %block_output2 : vector<4xf16> + %next_tile_output3 = vector.addf %tile_output3, %block_output3 : vector<4xf16> + scf.yield %next_tile_output0, %next_tile_output1, %next_tile_output2, %next_tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } else { + scf.yield %tile_output0, %tile_output1, %tile_output2, %tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %next_tile_output0, %next_tile_output1, %next_tile_output2, %next_tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } else { + scf.yield %tile_output0, %tile_output1, %tile_output2, %tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + scf.yield %updated_output0, %updated_output1, %updated_output2, %updated_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + scf.yield %next_max, %next_sum, %next_output0, %next_output1, %next_output2, %next_output3 : vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + %final_max, %final_sum, %final_output0, %final_output1, %final_output2, %final_output3 = scf.if %has_single_key_value_tail -> (vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + %single_tail_token = index.assume %full_key_value_token_count [lt(%full_key_value_token_count, %bounded_key_value_token_count)] : index + %single_tail_channel = index.mul %lane, %c2 : index + %single_tail_key_channel = index.add %key_value_head_base, %single_tail_channel : index + %key_pair_f16 = vector.load %key_view[%single_tail_token, %single_tail_key_channel] : view<[%bounded_key_value_token_count]x[%key_value_width]xf16> -> vector<2xf16> + %query_pair0_f16 = vector.load %query_stage_view[%query_row0, %single_tail_channel] : view<16x136xf16> -> vector<2xf16> + %query_pair1_f16 = vector.load %query_stage_view[%query_row1, %single_tail_channel] : view<16x136xf16> -> vector<2xf16> + %query_pair2_f16 = vector.load %query_stage_view[%query_row2, %single_tail_channel] : view<16x136xf16> -> vector<2xf16> + %query_pair3_f16 = vector.load %query_stage_view[%query_row3, %single_tail_channel] : view<16x136xf16> -> vector<2xf16> + %key_pair = vector.extf %key_pair_f16 : vector<2xf16> to vector<2xf32> + %query_pair0 = vector.extf %query_pair0_f16 : vector<2xf16> to vector<2xf32> + %query_pair1 = vector.extf %query_pair1_f16 : vector<2xf16> to vector<2xf32> + %query_pair2 = vector.extf %query_pair2_f16 : vector<2xf16> to vector<2xf32> + %query_pair3 = vector.extf %query_pair3_f16 : vector<2xf16> to vector<2xf32> + %product_pair0 = vector.mulf %query_pair0, %key_pair : vector<2xf32> + %product_pair1 = vector.mulf %query_pair1, %key_pair : vector<2xf32> + %product_pair2 = vector.mulf %query_pair2, %key_pair : vector<2xf32> + %product_pair3 = vector.mulf %query_pair3, %key_pair : vector<2xf32> + %partial_score0 = vector.reduce %product_pair0, %c0_f32 : vector<2xf32>, f32 + %partial_score1 = vector.reduce %product_pair1, %c0_f32 : vector<2xf32>, f32 + %partial_score2 = vector.reduce %product_pair2, %c0_f32 : vector<2xf32>, f32 + %partial_score3 = vector.reduce %product_pair3, %c0_f32 : vector<2xf32>, f32 + %partial_scores = vector.from_elements %partial_score0, %partial_score1, %partial_score2, %partial_score3 : vector<4xf32> + %reduced_scores = kernel.subgroup.reduce %partial_scores : vector<4xf32> + %mask_f16 = view.load %mask_view[%single_tail_token] : view<[%bounded_key_value_token_count]xf16> -> f16 + %mask_f32 = scalar.extf %mask_f16 : f16 to f32 + %mask_vector = vector.splat %mask_f32 : vector<4xf32> + %raw_scores = vector.addf %reduced_scores, %mask_vector : vector<4xf32> + %masked_scores = vector.select %query_valid, %raw_scores, %negative_f32x4 : vector<4xf32> + %next_max = vector.maxnumf %wmma_tail_max, %masked_scores : vector<4xf32> + %score_delta = vector.subf %masked_scores, %next_max : vector<4xf32> + %raw_probability = vector.expf %score_delta : vector<4xf32> + %probability = vector.select %query_valid, %raw_probability, %c0_f32x4 : vector<4xf32> + %old_delta = vector.subf %wmma_tail_max, %next_max : vector<4xf32> + %old_scale = vector.expf %old_delta : vector<4xf32> + %scaled_current_sum = vector.mulf %wmma_tail_sum, %old_scale : vector<4xf32> + %next_sum = vector.addf %scaled_current_sum, %probability : vector<4xf32> + %next_output0, %next_output1, %next_output2, %next_output3 = scf.if %lane_has_output -> (vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + %output_packet_channel = index.assume %lane_output_channel [range(%lane_output_channel, 0, 124), mul(%lane_output_channel, 4)] : index + %value_channel = index.add %key_value_head_base, %output_packet_channel : index + %value_packet = vector.load %value_view[%single_tail_token, %value_channel] : view<[%bounded_key_value_token_count]x[%key_value_width]xf16> -> vector<4xf16> + %old_scale_f16 = vector.fptrunc %old_scale : vector<4xf32> to vector<4xf16> + %probability_f16 = vector.fptrunc %probability : vector<4xf32> to vector<4xf16> + %old_scale0_scalar = vector.extract %old_scale_f16[0] : vector<4xf16> -> f16 + %old_scale1_scalar = vector.extract %old_scale_f16[1] : vector<4xf16> -> f16 + %old_scale2_scalar = vector.extract %old_scale_f16[2] : vector<4xf16> -> f16 + %old_scale3_scalar = vector.extract %old_scale_f16[3] : vector<4xf16> -> f16 + %probability0_scalar = vector.extract %probability_f16[0] : vector<4xf16> -> f16 + %probability1_scalar = vector.extract %probability_f16[1] : vector<4xf16> -> f16 + %probability2_scalar = vector.extract %probability_f16[2] : vector<4xf16> -> f16 + %probability3_scalar = vector.extract %probability_f16[3] : vector<4xf16> -> f16 + %old_scale0 = vector.splat %old_scale0_scalar : vector<4xf16> + %old_scale1 = vector.splat %old_scale1_scalar : vector<4xf16> + %old_scale2 = vector.splat %old_scale2_scalar : vector<4xf16> + %old_scale3 = vector.splat %old_scale3_scalar : vector<4xf16> + %probability0 = vector.splat %probability0_scalar : vector<4xf16> + %probability1 = vector.splat %probability1_scalar : vector<4xf16> + %probability2 = vector.splat %probability2_scalar : vector<4xf16> + %probability3 = vector.splat %probability3_scalar : vector<4xf16> + %scaled_current_output0 = vector.mulf %wmma_tail_output0, %old_scale0 : vector<4xf16> + %scaled_current_output1 = vector.mulf %wmma_tail_output1, %old_scale1 : vector<4xf16> + %scaled_current_output2 = vector.mulf %wmma_tail_output2, %old_scale2 : vector<4xf16> + %scaled_current_output3 = vector.mulf %wmma_tail_output3, %old_scale3 : vector<4xf16> + %tail_output0 = vector.mulf %value_packet, %probability0 : vector<4xf16> + %tail_output1 = vector.mulf %value_packet, %probability1 : vector<4xf16> + %tail_output2 = vector.mulf %value_packet, %probability2 : vector<4xf16> + %tail_output3 = vector.mulf %value_packet, %probability3 : vector<4xf16> + %updated_output0 = vector.addf %scaled_current_output0, %tail_output0 : vector<4xf16> + %updated_output1 = vector.addf %scaled_current_output1, %tail_output1 : vector<4xf16> + %updated_output2 = vector.addf %scaled_current_output2, %tail_output2 : vector<4xf16> + %updated_output3 = vector.addf %scaled_current_output3, %tail_output3 : vector<4xf16> + scf.yield %updated_output0, %updated_output1, %updated_output2, %updated_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } else { + scf.yield %wmma_tail_output0, %wmma_tail_output1, %wmma_tail_output2, %wmma_tail_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + scf.yield %next_max, %next_sum, %next_output0, %next_output1, %next_output2, %next_output3 : vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } else { + scf.yield %wmma_tail_max, %wmma_tail_sum, %wmma_tail_output0, %wmma_tail_output1, %wmma_tail_output2, %wmma_tail_output3 : vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + // Normalize and publish the lane-owned packets. Query-tail rows never + // participate in the output store. + scf.if %lane_has_output { + %output_packet_channel = index.assume %lane_output_channel [range(%lane_output_channel, 0, 124), mul(%lane_output_channel, 4)] : index + %sum0_scalar = vector.extract %final_sum[0] : vector<4xf32> -> f32 + %sum1_scalar = vector.extract %final_sum[1] : vector<4xf32> -> f32 + %sum2_scalar = vector.extract %final_sum[2] : vector<4xf32> -> f32 + %sum3_scalar = vector.extract %final_sum[3] : vector<4xf32> -> f32 + %inverse_sum0_f32 = scalar.divf %c1_f32, %sum0_scalar : f32 + %inverse_sum1_f32 = scalar.divf %c1_f32, %sum1_scalar : f32 + %inverse_sum2_f32 = scalar.divf %c1_f32, %sum2_scalar : f32 + %inverse_sum3_f32 = scalar.divf %c1_f32, %sum3_scalar : f32 + %inverse_sum0_f16 = scalar.fptrunc %inverse_sum0_f32 : f32 to f16 + %inverse_sum1_f16 = scalar.fptrunc %inverse_sum1_f32 : f32 to f16 + %inverse_sum2_f16 = scalar.fptrunc %inverse_sum2_f32 : f32 to f16 + %inverse_sum3_f16 = scalar.fptrunc %inverse_sum3_f32 : f32 to f16 + %inverse_sum0 = vector.splat %inverse_sum0_f16 : vector<4xf16> + %inverse_sum1 = vector.splat %inverse_sum1_f16 : vector<4xf16> + %inverse_sum2 = vector.splat %inverse_sum2_f16 : vector<4xf16> + %inverse_sum3 = vector.splat %inverse_sum3_f16 : vector<4xf16> + %normalized0_f16 = vector.mulf %final_output0, %inverse_sum0 : vector<4xf16> + %normalized1_f16 = vector.mulf %final_output1, %inverse_sum1 : vector<4xf16> + %normalized2_f16 = vector.mulf %final_output2, %inverse_sum2 : vector<4xf16> + %normalized3_f16 = vector.mulf %final_output3, %inverse_sum3 : vector<4xf16> + %normalized0 = vector.extf %normalized0_f16 : vector<4xf16> to vector<4xf32> + %normalized1 = vector.extf %normalized1_f16 : vector<4xf16> to vector<4xf32> + %normalized2 = vector.extf %normalized2_f16 : vector<4xf16> to vector<4xf32> + %normalized3 = vector.extf %normalized3_f16 : vector<4xf16> to vector<4xf32> + scf.if %query_valid0 { + vector.store %normalized0, %output_view[%query_head0, %output_packet_channel] : vector<4xf32>, view<[%query_head_count]x128xf32> + } + scf.if %query_valid1 { + vector.store %normalized1, %output_view[%query_head1, %output_packet_channel] : vector<4xf32>, view<[%query_head_count]x128xf32> + } + scf.if %query_valid2 { + vector.store %normalized2, %output_view[%query_head2, %output_packet_channel] : vector<4xf32>, view<[%query_head_count]x128xf32> + } + scf.if %query_valid3 { + vector.store %normalized3, %output_view[%query_head3, %output_packet_channel] : vector<4xf32>, view<[%query_head_count]x128xf32> + } + } + kernel.return +} + +// The mask selects the first KV row exactly. QK, F16 probability conversion, +// and P*V all execute, while the expected result remains an auditable iota. +check.case public @qwen3_moe_flash_attention_decode_f32_f16_wmma_selected_row_case { + %key_value_token_count = check.literal value(64) : index + %query = check.generate.fill value(1.0) : tensor<1x128xf32> + %key = check.generate.fill value(1.0) : tensor<64x1x128xf16> + %value = check.generate.iota offset(0.0) step(0.125) : tensor<64x1x128xf16> + %mask = check.generate.iota offset(0.0) step(-10000.0) : tensor<64xf16> + %output = check.generate.fill value(-1.0) : tensor<1x128xf32> + %expected = check.generate.iota offset(0.0) step(0.125) : tensor<1x128xf32> + kernel.launch @qwen3_moe_flash_attention_decode_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %output) : [index](index, tensor<1x128xf32>, tensor<64x1x128xf16>, tensor<64x1x128xf16>, tensor<64xf16>, tensor<1x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<1x128xf32> + check.return +} + +// Thirty-two query heads share four KV heads in the production GQA ratio. +// Equal scores and constant values make every output exactly two while all +// head-group addressing and output ownership execute. +check.case public @qwen3_moe_flash_attention_decode_f32_f16_wmma_gqa_case { + %key_value_token_count = check.literal value(128) : index + %query = check.generate.fill value(1.0) : tensor<32x128xf32> + %key = check.generate.fill value(1.0) : tensor<128x4x128xf16> + %value = check.generate.fill value(2.0) : tensor<128x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<128xf16> + %output = check.generate.fill value(-1.0) : tensor<32x128xf32> + %expected = check.generate.fill value(2.0) : tensor<32x128xf32> + kernel.launch @qwen3_moe_flash_attention_decode_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %output) : [index](index, tensor<32x128xf32>, tensor<128x4x128xf16>, tensor<128x4x128xf16>, tensor<128xf16>, tensor<32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<32x128xf32> + check.return +} + +// Sixty-five KV rows force one masked cleanup tile after a full WMMA block. +// The mask selects only that final row, whose iota values begin at 1024. +check.case public @qwen3_moe_flash_attention_decode_f32_f16_wmma_tail_case { + %key_value_token_count = check.literal value(65) : index + %query = check.generate.fill value(1.0) : tensor<1x128xf32> + %key = check.generate.fill value(1.0) : tensor<65x1x128xf16> + %value = check.generate.iota offset(0.0) step(0.125) : tensor<65x1x128xf16> + %mask = check.generate.iota offset(-64000.0) step(1000.0) : tensor<65xf16> + %output = check.generate.fill value(-1.0) : tensor<1x128xf32> + %expected = check.generate.iota offset(1024.0) step(0.125) : tensor<1x128xf32> + kernel.launch @qwen3_moe_flash_attention_decode_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %output) : [index](index, tensor<1x128xf32>, tensor<65x1x128xf16>, tensor<65x1x128xf16>, tensor<65xf16>, tensor<1x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<1x128xf32> + check.return +} + +check.case public @qwen3_moe_flash_attention_decode_f32_f16_wmma_benchmark_case { + %key_value_token_count = check.param.choice values([64, 65, 128, 256, 512, 768, 1024, 1280, 2048]) name("key_value_token_count") : index + %query = check.generate.fill value(0.0) : tensor<32x128xf32> + %key = check.generate.fill value(0.0) : tensor<[%key_value_token_count]x4x128xf16> + %value = check.generate.fill value(0.0) : tensor<[%key_value_token_count]x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<[%key_value_token_count]xf16> + %output = check.generate.fill value(1.0) : tensor<32x128xf32> + %expected = check.generate.fill value(0.0) : tensor<32x128xf32> + kernel.launch @qwen3_moe_flash_attention_decode_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %output) : [index](index, tensor<32x128xf32>, tensor<[%key_value_token_count]x4x128xf16>, tensor<[%key_value_token_count]x4x128xf16>, tensor<[%key_value_token_count]xf16>, tensor<32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<32x128xf32> + check.return +} + +check.benchmark<@qwen3_moe_flash_attention_decode_f32_f16_wmma_selected_row_case> @qwen3_moe_flash_attention_decode_f32_f16_wmma_selected_row + +check.benchmark<@qwen3_moe_flash_attention_decode_f32_f16_wmma_gqa_case> @qwen3_moe_flash_attention_decode_f32_f16_wmma_gqa + +check.benchmark<@qwen3_moe_flash_attention_decode_f32_f16_wmma_tail_case> @qwen3_moe_flash_attention_decode_f32_f16_wmma_tail + +check.benchmark<@qwen3_moe_flash_attention_decode_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_f32_f16_wmma_decode_64 {key_value_token_count = 64} + +check.benchmark<@qwen3_moe_flash_attention_decode_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_f32_f16_wmma_decode_65 {key_value_token_count = 65} + +check.benchmark<@qwen3_moe_flash_attention_decode_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_f32_f16_wmma_decode_128 {key_value_token_count = 128} + +check.benchmark<@qwen3_moe_flash_attention_decode_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_f32_f16_wmma_decode_256 {key_value_token_count = 256} + +check.benchmark<@qwen3_moe_flash_attention_decode_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_f32_f16_wmma_decode_512 {key_value_token_count = 512} + +check.benchmark<@qwen3_moe_flash_attention_decode_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_f32_f16_wmma_decode_768 {key_value_token_count = 768} + +check.benchmark<@qwen3_moe_flash_attention_decode_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_f32_f16_wmma_decode_1024 {key_value_token_count = 1024} + +check.benchmark<@qwen3_moe_flash_attention_decode_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_f32_f16_wmma_decode_1280 {key_value_token_count = 1280} + +check.benchmark<@qwen3_moe_flash_attention_decode_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_f32_f16_wmma_decode_2048 {key_value_token_count = 2048} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/flash_attention_decode_q128_f32_f16_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/flash_attention_decode_q128_f32_f16_wmma.loom new file mode 100644 index 000000000000..5a0c8037f856 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/flash_attention_decode_q128_f32_f16_wmma.loom @@ -0,0 +1,266 @@ +// Exact-q128 grouped-query decode FlashAttention. +// +// One eight-wave workgroup owns each KV head. The waves cooperatively cover +// all 128 KV rows for QK, normalize the complete score matrix in LDS, and then +// each own one 16-channel P*V tile. This preserves 32 active waves for the +// production 32Q/4KV shape while removing duplicate QK work, global partial +// tensors, completion atomics, and a separate resolve dispatch. +amdgpu.target @qwen3_moe_attention_decode_q128_gfx11_wave64 {subgroup_size = 64} + +config.decl @qwen3_moe.attention.query_head_count : %value: index where [range(%value, 1, 64)] + +config.decl @qwen3_moe.attention.key_value_head_count : %value: index where [range(%value, 1, 64)] + +kernel.def target(@qwen3_moe_attention_decode_q128_gfx11_wave64) @qwen3_moe_flash_attention_decode_q128_fused_f32_f16_wmma(%key_value_token_count: index) { + %key_value_head_count = config.get @qwen3_moe.attention.key_value_head_count : index + %c1 = index.constant 1 : index + %c512 = index.constant 512 : index + kernel.launch.config workgroups(%key_value_head_count, %c1, %c1) workgroup_size(%c512, %c1, %c1) : index +} launch(%key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %output: buffer) { + %query_head_count = config.get @qwen3_moe.attention.query_head_count : index + %key_value_head_count = config.get @qwen3_moe.attention.key_value_head_count : index + %workgroup_x0 = kernel.workgroup.id : index + %key_value_head = index.assume %workgroup_x0 [range(%workgroup_x0, 0, 63)] : index + %workitem0 = kernel.workitem.id : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 511)] : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %lane0 = kernel.subgroup.lane.id : index + %lane = index.assume %lane0 [range(%lane0, 0, 63)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %c512 = index.constant 512 : index + %c0_offset = index.constant 0 : offset + %query_stage_bytes = index.constant 4352 : offset + %score_stage_bytes = index.constant 12288 : offset + %probability_stage_bytes = index.constant 6144 : offset + %product_stage_bytes = index.constant 8192 : offset + %c0_f16 = scalar.constant 0.0 : f16 + %c0_f32 = scalar.constant 0.0 : f32 + %c1_f32 = scalar.constant 1.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %c0_f32x2 = vector.constant 0.0 : vector<2xf32> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %head_size_f32 = scalar.constant 128.0 : f32 + %attention_scale = scalar.rsqrtf %head_size_f32 : f32 + %query_heads_per_key_value_head0 = index.div %query_head_count, %key_value_head_count : index + %query_heads_per_key_value_head = index.assume %query_heads_per_key_value_head0 [range(%query_heads_per_key_value_head0, 1, 16)] : index + %query_head_base = index.mul %key_value_head, %query_heads_per_key_value_head : index + %key_value_width = index.mul %key_value_head_count, %c128 : index + %key_value_head_base = index.mul %key_value_head, %c128 : index + %score_key_origin = index.mul %subgroup, %c16 : index + %query_row_base = index.mul %subgroup, %c2 : index + %query_row0 = index.add %query_row_base, %c0 : index + %query_row1 = index.add %query_row_base, %c1 : index + %query_head0 = index.add %query_head_base, %query_row0 : index + %query_head1 = index.add %query_head_base, %query_row1 : index + %query_row_valid0 = index.cmp ult, %query_row0, %query_heads_per_key_value_head : index + %query_row_valid1 = index.cmp ult, %query_row1, %query_heads_per_key_value_head : index + %query_head_in_range0 = index.cmp ult, %query_head0, %query_head_count : index + %query_head_in_range1 = index.cmp ult, %query_head1, %query_head_count : index + %query_valid0 = scalar.andi %query_row_valid0, %query_head_in_range0 : i1 + %query_valid1 = scalar.andi %query_row_valid1, %query_head_in_range1 : i1 + %query_valid0x2 = vector.from_elements %query_valid0, %query_valid0 : vector<2xi1> + %query_valid1x2 = vector.from_elements %query_valid1, %query_valid1 : vector<2xi1> + %key_token0 = index.add %lane, %c0 : index + %key_token1 = index.add %lane, %c64 : index + %output_tile_channel = index.mul %subgroup, %c16 : index + %query_transposed_layout = encoding.layout.strided [1, 136] : encoding + %probability_transposed_layout = encoding.layout.strided [1, 24] : encoding + %query_noalias, %key_noalias, %value_noalias, %mask_noalias, %output_noalias = buffer.assume.noalias %query, %key, %value, %mask, %output : buffer, buffer, buffer, buffer, buffer + %query_aligned = buffer.assume.alignment %query_noalias {minimum_alignment = 16} : buffer + %key_aligned = buffer.assume.alignment %key_noalias {minimum_alignment = 16} : buffer + %value_aligned = buffer.assume.alignment %value_noalias {minimum_alignment = 16} : buffer + %mask_aligned = buffer.assume.alignment %mask_noalias {minimum_alignment = 16} : buffer + %output_aligned = buffer.assume.alignment %output_noalias {minimum_alignment = 16} : buffer + %query_view = buffer.view %query_aligned[%c0_offset] : buffer -> view<[%query_head_count]x128xf32> + %key_view = buffer.view %key_aligned[%c0_offset] : buffer -> view<128x[%key_value_width]xf16> + %value_view = buffer.view %value_aligned[%c0_offset] : buffer -> view<128x[%key_value_width]xf16> + %mask_view = buffer.view %mask_aligned[%c0_offset] : buffer -> view<128xf16> + %output_view = buffer.view %output_aligned[%c0_offset] : buffer -> view<[%query_head_count]x128xf32> + %query_stage = buffer.alloca align(16) %query_stage_bytes : buffer + %score_stage = buffer.alloca align(16) %score_stage_bytes : buffer + %probability_stage = buffer.alloca align(16) %probability_stage_bytes : buffer + %product_stage = buffer.alloca align(16) %product_stage_bytes : buffer + %query_stage_view = buffer.view %query_stage[%c0_offset] : buffer -> view<16x136xf16> + %query_transposed_view = buffer.view %query_stage[%c0_offset] : buffer -> view<128x16xf16, %query_transposed_layout> + %score_stage_view = buffer.view %score_stage[%c0_offset] : buffer -> view<128x24xf32> + %probability_stage_view = buffer.view %probability_stage[%c0_offset] : buffer -> view<16x128xf16, %probability_transposed_layout> + %product_stage_view = buffer.view %product_stage[%c0_offset] : buffer -> view<16x128xf32> + // Scale and transpose the 16-row Q tile once for all eight score waves. + scf.for %load_iteration = [%c0 to %c4 step %c1] unroll { + %linear = index.madd %load_iteration, %c512, %workitem : index + %local_query_row = index.div %linear, %c128 : index + %query_channel = index.rem %linear, %c128 : index + %local_query_head = index.add %query_head_base, %local_query_row : index + %local_query_row_valid = index.cmp ult, %local_query_row, %query_heads_per_key_value_head : index + %local_query_head_in_range = index.cmp ult, %local_query_head, %query_head_count : index + %local_query_valid = scalar.andi %local_query_row_valid, %local_query_head_in_range : i1 + %query_value = scf.if %local_query_valid -> (f16) { + %loaded = view.load %query_view[%local_query_head, %query_channel] : view<[%query_head_count]x128xf32> -> f32 + %scaled = scalar.mulf %loaded, %attention_scale : f32 + %truncated = scalar.fptrunc %scaled : f32 to f16 + scf.yield %truncated : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %query_value, %query_stage_view[%local_query_row, %query_channel] : f16, view<16x136xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + // Each wave produces one 16x16 score slice, covering all 128 KV rows. + %score_init_values = vector.constant 0.0 : vector<4xf32> + %score_init = vector.fragment %score_init_values shape [%m, %n] : vector<4xf32> + %score_fragment = scf.for %head_tile = [%c0 to %c128 step %c16](%score_accumulator = %score_init : vector<4xf32>) -> (vector<4xf32>) unroll { + %key_channel = index.add %key_value_head_base, %head_tile : index + %key_fragment = vector.fragment.load %key_view[%score_key_origin, %key_channel] shape [%m, %k] : view<128x[%key_value_width]xf16> -> vector<16xf16> + %query_fragment = vector.fragment.load %query_transposed_view[%head_tile, %c0] shape [%k, %n] : view<128x16xf16, %query_transposed_layout> -> vector<16xf16> + %next_score_accumulator = vector.mma %key_fragment, %query_fragment, %score_accumulator : vector<16xf16>, vector<16xf16>, vector<4xf32> + scf.yield %next_score_accumulator : vector<4xf32> + } + vector.fragment.store %score_fragment, %score_stage_view[%score_key_origin, %c0] shape [%m, %n] : vector<4xf32>, view<128x24xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + // Each wave normalizes two query rows. Every lane owns two KV columns, so a + // vector reduction first combines its pair and the subgroup reduction then + // spans the complete 128-token row. + %mask0_f16 = view.load %mask_view[%key_token0] : view<128xf16> -> f16 + %mask1_f16 = view.load %mask_view[%key_token1] : view<128xf16> -> f16 + %mask0 = scalar.extf %mask0_f16 : f16 to f32 + %mask1 = scalar.extf %mask1_f16 : f16 to f32 + %mask_pair = vector.from_elements %mask0, %mask1 : vector<2xf32> + %raw_score00 = view.load %score_stage_view[%key_token0, %query_row0] : view<128x24xf32> -> f32 + %raw_score01 = view.load %score_stage_view[%key_token1, %query_row0] : view<128x24xf32> -> f32 + %raw_score10 = view.load %score_stage_view[%key_token0, %query_row1] : view<128x24xf32> -> f32 + %raw_score11 = view.load %score_stage_view[%key_token1, %query_row1] : view<128x24xf32> -> f32 + %raw_scores0 = vector.from_elements %raw_score00, %raw_score01 : vector<2xf32> + %raw_scores1 = vector.from_elements %raw_score10, %raw_score11 : vector<2xf32> + %added_scores0 = vector.addf %raw_scores0, %mask_pair : vector<2xf32> + %added_scores1 = vector.addf %raw_scores1, %mask_pair : vector<2xf32> + %negative_pair = vector.constant -1e+30 : vector<2xf32> + %masked_scores0 = vector.select %query_valid0x2, %added_scores0, %negative_pair : vector<2xf32> + %masked_scores1 = vector.select %query_valid1x2, %added_scores1, %negative_pair : vector<2xf32> + %lane_max0 = vector.reduce %masked_scores0, %negative_large : vector<2xf32>, f32 + %lane_max1 = vector.reduce %masked_scores1, %negative_large : vector<2xf32>, f32 + %lane_maxima = vector.from_elements %lane_max0, %lane_max1 : vector<2xf32> + %row_maxima = kernel.subgroup.reduce %lane_maxima : vector<2xf32> + %row_maximum0 = vector.extract %row_maxima[0] : vector<2xf32> -> f32 + %row_maximum1 = vector.extract %row_maxima[1] : vector<2xf32> -> f32 + %row_maximum0x2 = vector.splat %row_maximum0 : vector<2xf32> + %row_maximum1x2 = vector.splat %row_maximum1 : vector<2xf32> + %score_delta0 = vector.subf %masked_scores0, %row_maximum0x2 : vector<2xf32> + %score_delta1 = vector.subf %masked_scores1, %row_maximum1x2 : vector<2xf32> + %raw_probability0 = vector.expf %score_delta0 : vector<2xf32> + %raw_probability1 = vector.expf %score_delta1 : vector<2xf32> + %probability0 = vector.select %query_valid0x2, %raw_probability0, %c0_f32x2 : vector<2xf32> + %probability1 = vector.select %query_valid1x2, %raw_probability1, %c0_f32x2 : vector<2xf32> + %lane_sum0 = vector.reduce %probability0, %c0_f32 : vector<2xf32>, f32 + %lane_sum1 = vector.reduce %probability1, %c0_f32 : vector<2xf32>, f32 + %lane_sums = vector.from_elements %lane_sum0, %lane_sum1 : vector<2xf32> + %row_sums = kernel.subgroup.reduce %lane_sums : vector<2xf32> + %row_sum0 = vector.extract %row_sums[0] : vector<2xf32> -> f32 + %row_sum1 = vector.extract %row_sums[1] : vector<2xf32> -> f32 + %reciprocal_sum0 = scf.if %query_valid0 -> (f32) { + %reciprocal = scalar.divf %c1_f32, %row_sum0 : f32 + scf.yield %reciprocal : f32 + } else { + scf.yield %c0_f32 : f32 + } + %reciprocal_sum1 = scf.if %query_valid1 -> (f32) { + %reciprocal = scalar.divf %c1_f32, %row_sum1 : f32 + scf.yield %reciprocal : f32 + } else { + scf.yield %c0_f32 : f32 + } + %reciprocal_sum0x2 = vector.splat %reciprocal_sum0 : vector<2xf32> + %reciprocal_sum1x2 = vector.splat %reciprocal_sum1 : vector<2xf32> + %normalized_probability0 = vector.mulf %probability0, %reciprocal_sum0x2 : vector<2xf32> + %normalized_probability1 = vector.mulf %probability1, %reciprocal_sum1x2 : vector<2xf32> + %probability0_f16 = vector.fptrunc %normalized_probability0 : vector<2xf32> to vector<2xf16> + %probability1_f16 = vector.fptrunc %normalized_probability1 : vector<2xf32> to vector<2xf16> + %probability00 = vector.extract %probability0_f16[0] : vector<2xf16> -> f16 + %probability01 = vector.extract %probability0_f16[1] : vector<2xf16> -> f16 + %probability10 = vector.extract %probability1_f16[0] : vector<2xf16> -> f16 + %probability11 = vector.extract %probability1_f16[1] : vector<2xf16> -> f16 + view.store %probability00, %probability_stage_view[%query_row0, %key_token0] : f16, view<16x128xf16, %probability_transposed_layout> + view.store %probability01, %probability_stage_view[%query_row0, %key_token1] : f16, view<16x128xf16, %probability_transposed_layout> + view.store %probability10, %probability_stage_view[%query_row1, %key_token0] : f16, view<16x128xf16, %probability_transposed_layout> + view.store %probability11, %probability_stage_view[%query_row1, %key_token1] : f16, view<16x128xf16, %probability_transposed_layout> + kernel.barrier scope(workgroup) ordering(acq_rel) + // Each wave owns one disjoint 16-channel output tile. + %product_init_values = vector.constant 0.0 : vector<4xf32> + %product_init = vector.fragment %product_init_values shape [%m, %n] : vector<4xf32> + %value_channel = index.add %key_value_head_base, %output_tile_channel : index + %product_fragment = scf.for %key_tile = [%c0 to %c128 step %c16](%product_accumulator = %product_init : vector<4xf32>) -> (vector<4xf32>) unroll { + %probability_fragment = vector.fragment.load %probability_stage_view[%c0, %key_tile] shape [%m, %k] : view<16x128xf16, %probability_transposed_layout> -> vector<16xf16> + %value_fragment = vector.fragment.load %value_view[%key_tile, %value_channel] shape [%k, %n] : view<128x[%key_value_width]xf16> -> vector<16xf16> + %next_product_accumulator = vector.mma %probability_fragment, %value_fragment, %product_accumulator : vector<16xf16>, vector<16xf16>, vector<4xf32> + scf.yield %next_product_accumulator : vector<4xf32> + } + vector.fragment.store %product_fragment, %product_stage_view[%c0, %output_tile_channel] shape [%m, %n] : vector<4xf32>, view<16x128xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + // One vector packet per workitem drains the 16x128 result tile while + // suppressing padded GQA rows. + %output_row = index.div %workitem, %c32 : index + %output_packet = index.rem %workitem, %c32 : index + %output_channel = index.mul %output_packet, %c4 : index + %output_head = index.add %query_head_base, %output_row : index + %output_row_valid = index.cmp ult, %output_row, %query_heads_per_key_value_head : index + %output_head_in_range = index.cmp ult, %output_head, %query_head_count : index + %output_valid = scalar.andi %output_row_valid, %output_head_in_range : i1 + scf.if %output_valid { + %output_packet_value = vector.load %product_stage_view[%output_row, %output_channel] : view<16x128xf32> -> vector<4xf32> + vector.store %output_packet_value, %output_view[%output_head, %output_channel] : vector<4xf32>, view<[%query_head_count]x128xf32> + } + kernel.return +} + +// A sharp mask selects the first KV row and makes the expected result an +// auditable iota while QK, global softmax, F16 probabilities, and P*V execute. +check.case public @qwen3_moe_flash_attention_decode_q128_fused_selected_row_case { + %key_value_token_count = check.literal value(128) : index + %query = check.generate.fill value(1.0) : tensor<1x128xf32> + %key = check.generate.fill value(1.0) : tensor<128x1x128xf16> + %value = check.generate.iota offset(0.0) step(0.125) : tensor<128x1x128xf16> + %mask = check.generate.iota offset(0.0) step(-10000.0) : tensor<128xf16> + %output = check.generate.fill value(-1.0) : tensor<1x128xf32> + %expected = check.generate.iota offset(0.0) step(0.125) : tensor<1x128xf32> + kernel.launch @qwen3_moe_flash_attention_decode_q128_fused_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %output) : [index](index, tensor<1x128xf32>, tensor<128x1x128xf16>, tensor<128x1x128xf16>, tensor<128xf16>, tensor<1x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<1x128xf32> + check.return +} + +// The production 32Q/4KV GQA shape verifies all head groups and output tiles. +check.case public @qwen3_moe_flash_attention_decode_q128_fused_gqa_case { + %key_value_token_count = check.literal value(128) : index + %query = check.generate.fill value(1.0) : tensor<32x128xf32> + %key = check.generate.fill value(1.0) : tensor<128x4x128xf16> + %value = check.generate.fill value(2.0) : tensor<128x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<128xf16> + %output = check.generate.fill value(-1.0) : tensor<32x128xf32> + %expected = check.generate.fill value(2.0) : tensor<32x128xf32> + kernel.launch @qwen3_moe_flash_attention_decode_q128_fused_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %output) : [index](index, tensor<32x128xf32>, tensor<128x4x128xf16>, tensor<128x4x128xf16>, tensor<128xf16>, tensor<32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<32x128xf32> + check.return +} + +check.case public @qwen3_moe_flash_attention_decode_q128_fused_benchmark_case { + %key_value_token_count = check.literal value(128) : index + %query = check.generate.fill value(0.0) : tensor<32x128xf32> + %key = check.generate.fill value(0.0) : tensor<128x4x128xf16> + %value = check.generate.fill value(0.0) : tensor<128x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<128xf16> + %output = check.generate.fill value(1.0) : tensor<32x128xf32> + %expected = check.generate.fill value(0.0) : tensor<32x128xf32> + kernel.launch @qwen3_moe_flash_attention_decode_q128_fused_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %output) : [index](index, tensor<32x128xf32>, tensor<128x4x128xf16>, tensor<128x4x128xf16>, tensor<128xf16>, tensor<32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<32x128xf32> + check.return +} + +check.benchmark<@qwen3_moe_flash_attention_decode_q128_fused_benchmark_case> @qwen3_moe_flash_attention_decode_q128_fused diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom new file mode 100644 index 000000000000..00ccc77cbcd2 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/flash_attention_decode_split_f32_f16_wmma.loom @@ -0,0 +1,1235 @@ +// Qwen3 MoE grouped-query decode FlashAttention. +// +// Each workgroup processes one 64-token KV block for all GQA query heads that +// share a KV head. The workgroups publish online-softmax state, then the last +// arrival folds every block and resets the per-KV-head completion counter for +// the next invocation. Packing GQA heads removes redundant K/V traffic while +// split-K preserves enough parallelism for decode without another dispatch. +template.decl @qwen3_moe.attention.decode_split.pack_completed_q8(%key_value_head: index, %output: buffer, %q8_output: buffer) + +template.decl @qwen3_moe.attention.decode_split.produce_partials(%key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer) + +template.decl @qwen3_moe.attention.decode_split.produce_partials.active(%key_value_token_count: index, %partial_block_capacity0: index, %launched_block_count0: index, %query: buffer, %key: buffer, %value: buffer, %lane_mask: f32, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer) + +template.decl @qwen3_moe.attention.decode_split.reduce_completed.cooperative(%partial_block_capacity0: index, %active_block_count0: index, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %output: buffer) + +template.decl @qwen3_moe.attention.decode_split.reduce_completed.direct(%partial_block_capacity0: index, %active_block_count0: index, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %output: buffer) + +template.decl @qwen3_moe.attention.decode_split.reduce_fused(%key_value_token_capacity: index, %partial_block_capacity0: index, %producer_block_count0: index, %publish_q8: i1, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %completion_counter: buffer, %output: buffer, %q8_output: buffer) + +amdgpu.target @qwen3_moe_decode_split_gfx11_wave64 {subgroup_size = 64} + +config.decl @qwen3_moe.attention.query_head_count : %value: index where [range(%value, 1, 64)] + +config.decl @qwen3_moe.attention.key_value_head_count : %value: index where [range(%value, 1, 64)] + +// Maximum K/V storage capacity available to the compiled kernel. +config.decl @qwen3_moe.attention.key_value_token_capacity : %value: index where [range(%value, 64, 32768)] + +// Computes one active online-softmax partial for a 64-row KV block. Both the +// fused short-context export and the two-dispatch long-context export reach +// this body through the block-classifying producer below, keeping their +// different binding contracts honest without duplicating the attention math. +// Partial storage retains its capacity-specialized block-axis stride while the +// launched block count bounds issue-time work. Keeping those values distinct +// preserves constant address arithmetic across changing visible prefixes. +template.def<@qwen3_moe.attention.decode_split.produce_partials.active> device @qwen3_moe_flash_attention_decode_split_produce_active_partials_body_f32_f16_wmma(%key_value_token_count: index, %partial_block_capacity0: index, %launched_block_count0: index, %query: buffer, %key: buffer, %value: buffer, %lane_mask: f32, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer) { + %bounded_key_value_token_count = index.assume %key_value_token_count [range(%key_value_token_count, 1, 32768)] : index + %partial_block_capacity, %launched_block_count = index.assume %partial_block_capacity0, %launched_block_count0 [range(%partial_block_capacity0, 1, 512), range(%launched_block_count0, 1, 512), le(%launched_block_count0, %partial_block_capacity0)] : index, index + %query_head_count = config.get @qwen3_moe.attention.query_head_count : index + %key_value_head_count = config.get @qwen3_moe.attention.key_value_head_count : index + %workgroup_x0 = kernel.workgroup.id : index + %workgroup_y0 = kernel.workgroup.id : index + %workgroup_x_in_launch = index.assume %workgroup_x0 [range(%workgroup_x0, 0, 511)] : index + %workgroup_y = index.assume %workgroup_y0 [range(%workgroup_y0, 0, 63)] : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 3)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %query_stage_bytes = index.constant 4352 : offset + %score_stage_bytes = index.constant 6144 : offset + %probability_stage_bytes = index.constant 3072 : offset + %product_stage_bytes = index.constant 4096 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %head_size_f32 = scalar.constant 128.0 : f32 + %attention_scale = scalar.rsqrtf %head_size_f32 : f32 + %c0_f16 = scalar.constant 0.0 : f16 + %c0_f16x8 = vector.constant 0.0 : vector<8xf16> + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %workgroup_x, %launch_partial_block_capacity, %launch_launched_block_count = index.assume %workgroup_x_in_launch, %partial_block_capacity, %launched_block_count [lt(%workgroup_x_in_launch, %launched_block_count)] : index, index, index + %tail_key_value_token_count = index.rem %bounded_key_value_token_count, %c64 : index + %has_no_tail = index.cmp eq, %tail_key_value_token_count, %c0 : index + %active_padded_key_value_token_count = index.add %bounded_key_value_token_count, %c63 : index + %active_key_value_block_count = index.div %active_padded_key_value_token_count, %c64 : index + %last_block_ordinal = index.sub %active_key_value_block_count, %c1 : index + %block_ordinal = index.add %workgroup_x, %c0 : index + %is_not_last_block = index.cmp ne, %block_ordinal, %last_block_ordinal : index + %is_full_block = scalar.ori %has_no_tail, %is_not_last_block : i1 + %key_value_head = index.add %workgroup_y, %c0 : index + %query_heads_per_key_value_head0 = index.div %query_head_count, %key_value_head_count : index + %query_heads_per_key_value_head = index.assume %query_heads_per_key_value_head0 [range(%query_heads_per_key_value_head0, 1, 16)] : index + %query_head_base = index.mul %key_value_head, %query_heads_per_key_value_head : index + %key_value_width = index.mul %key_value_head_count, %c128 : index + %key_value_head_base = index.mul %key_value_head, %c128 : index + %key_origin = index.mul %block_ordinal, %c64 : index + %subgroup_score_column = index.mul %subgroup, %c16 : index + %subgroup_query_row = index.mul %subgroup, %c4 : index + %query_row0 = index.add %subgroup_query_row, %c0 : index + %query_row1 = index.add %subgroup_query_row, %c1 : index + %query_row2 = index.add %subgroup_query_row, %c2 : index + %query_row3 = index.add %subgroup_query_row, %c3 : index + %query_head0 = index.add %query_head_base, %query_row0 : index + %query_head1 = index.add %query_head_base, %query_row1 : index + %query_head2 = index.add %query_head_base, %query_row2 : index + %query_head3 = index.add %query_head_base, %query_row3 : index + %query_head_valid0 = index.cmp ult, %query_head0, %query_head_count : index + %query_head_valid1 = index.cmp ult, %query_head1, %query_head_count : index + %query_head_valid2 = index.cmp ult, %query_head2, %query_head_count : index + %query_head_valid3 = index.cmp ult, %query_head3, %query_head_count : index + %query_valid = vector.from_elements %query_head_valid0, %query_head_valid1, %query_head_valid2, %query_head_valid3 : vector<4xi1> + %subgroup_output_channel = index.mul %subgroup, %c32 : index + %subgroup_output_channel1 = index.add %subgroup_output_channel, %c16 : index + %lane_output_channel = index.mul %lane, %c4 : index + %lane_has_output = index.cmp ult, %lane, %c32 : index + %lane_is_zero = index.cmp eq, %lane, %c0 : index + %query_transposed_layout = encoding.layout.strided [1, 136] : encoding + %probability_transposed_layout = encoding.layout.strided [1, 24] : encoding + %query_noalias, %key_noalias, %value_noalias, %partial_max_noalias, %partial_sum_noalias, %partial_output_noalias = buffer.assume.noalias %query, %key, %value, %partial_max, %partial_sum, %partial_output : buffer, buffer, buffer, buffer, buffer, buffer + %query_aligned = buffer.assume.alignment %query_noalias {minimum_alignment = 16} : buffer + %key_aligned = buffer.assume.alignment %key_noalias {minimum_alignment = 16} : buffer + %value_aligned = buffer.assume.alignment %value_noalias {minimum_alignment = 16} : buffer + %partial_max_aligned = buffer.assume.alignment %partial_max_noalias {minimum_alignment = 16} : buffer + %partial_sum_aligned = buffer.assume.alignment %partial_sum_noalias {minimum_alignment = 16} : buffer + %partial_output_aligned = buffer.assume.alignment %partial_output_noalias {minimum_alignment = 16} : buffer + %query_view = buffer.view %query_aligned[%c0_offset] : buffer -> view<[%query_head_count]x128xf32> + %key_view = buffer.view %key_aligned[%c0_offset] : buffer -> view<[%bounded_key_value_token_count]x[%key_value_width]xf16> + %value_view = buffer.view %value_aligned[%c0_offset] : buffer -> view<[%bounded_key_value_token_count]x[%key_value_width]xf16> + %partial_max_view = buffer.view %partial_max_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + %partial_sum_view = buffer.view %partial_sum_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + %partial_output_view = buffer.view %partial_output_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%launch_partial_block_capacity]x16x128xf16> + %query_stage = buffer.alloca align(16) %query_stage_bytes : buffer + %score_stage = buffer.alloca align(16) %score_stage_bytes : buffer + %probability_stage = buffer.alloca align(16) %probability_stage_bytes : buffer + %product_stage = buffer.alloca align(16) %product_stage_bytes : buffer + %query_stage_view = buffer.view %query_stage[%c0_offset] : buffer -> view<16x136xf16> + %query_transposed_view = buffer.view %query_stage[%c0_offset] : buffer -> view<128x16xf16, %query_transposed_layout> + %score_stage_view = buffer.view %score_stage[%c0_offset] : buffer -> view<64x24xf32> + %probability_stage_view = buffer.view %probability_stage[%c0_offset] : buffer -> view<16x64xf16, %probability_transposed_layout> + %product_stage_view = buffer.view %product_stage[%c0_offset] : buffer -> view<16x128xf16> + %tail_key_stage_view = buffer.view %product_stage[%c0_offset] : buffer -> view<64x16xf16> + %tail_value_stage_view = buffer.view %score_stage[%c0_offset] : buffer -> view<16x128xf16> + scf.for %load_iteration = [%c0 to %c8 step %c1] unroll { + %linear = index.madd %load_iteration, %c256, %workitem : index + %local_query_row = index.div %linear, %c128 : index + %query_channel = index.rem %linear, %c128 : index + %local_query_head = index.add %query_head_base, %local_query_row : index + %local_query_head_valid = index.cmp ult, %local_query_head, %query_head_count : index + %query_value = scf.if %local_query_head_valid -> (f16) { + %loaded = view.load %query_view[%local_query_head, %query_channel] : view<[%query_head_count]x128xf32> -> f32 + %scaled = scalar.mulf %loaded, %attention_scale : f32 + %truncated = scalar.fptrunc %scaled : f32 to f16 + scf.yield %truncated : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %query_value, %query_stage_view[%local_query_row, %query_channel] : f16, view<16x136xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %score_init_values = vector.constant 0.0 : vector<4xf32> + %score_init = vector.fragment %score_init_values shape [%m, %n] : vector<4xf32> + %score_fragment = scf.if %is_full_block -> (vector<4xf32>) { + %full_score_fragment = scf.for %head_tile = [%c0 to %c128 step %c16](%score_accumulator = %score_init : vector<4xf32>) -> (vector<4xf32>) unroll { + // Bound this view by the proven tile end so the full vector footprint is + // visible without treating physical tail padding as logical storage. + %score_key_origin0 = index.add %key_origin, %subgroup_score_column : index + %score_key_end0 = index.add %score_key_origin0, %c16 : index + %score_key_origin, %score_key_end = index.assume %score_key_origin0, %score_key_end0 [le(%score_key_end0, %bounded_key_value_token_count)] : index, index + %full_key_view = buffer.view %key_aligned[%c0_offset] : buffer -> view<[%score_key_end]x[%key_value_width]xf16> + %key_channel = index.add %key_value_head_base, %head_tile : index + %key_fragment = vector.fragment.load %full_key_view[%score_key_origin, %key_channel] shape [%m, %k] : view<[%score_key_end]x[%key_value_width]xf16> -> vector<16xf16> + %query_fragment = vector.fragment.load %query_transposed_view[%head_tile, %c0] shape [%k, %n] : view<128x16xf16, %query_transposed_layout> -> vector<16xf16> + %next_score_accumulator = vector.mma %key_fragment, %query_fragment, %score_accumulator : vector<16xf16>, vector<16xf16>, vector<4xf32> + scf.yield %next_score_accumulator : vector<4xf32> + } + scf.yield %full_score_fragment : vector<4xf32> + } else { + // Stage one 64x16 K panel at a time. Every physical load is guarded, so + // callers need no initialized padding beyond the logical KV length. + %tail_score_fragment = scf.for %head_tile = [%c0 to %c128 step %c16](%score_accumulator = %score_init : vector<4xf32>) -> (vector<4xf32>) unroll { + scf.for %load_iteration = [%c0 to %c4 step %c1] unroll { + %linear = index.madd %load_iteration, %c256, %workitem : index + %key_row = index.div %linear, %c16 : index + %head_channel = index.rem %linear, %c16 : index + %key_token = index.add %key_origin, %key_row : index + %key_valid = index.cmp ult, %key_token, %bounded_key_value_token_count : index + %key_value = scf.if %key_valid -> (f16) { + %head_channel0 = index.add %head_tile, %head_channel : index + %key_channel = index.add %key_value_head_base, %head_channel0 : index + %loaded = view.load %key_view[%key_token, %key_channel] : view<[%bounded_key_value_token_count]x[%key_value_width]xf16> -> f16 + scf.yield %loaded : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %key_value, %tail_key_stage_view[%key_row, %head_channel] : f16, view<64x16xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %key_fragment = vector.fragment.load %tail_key_stage_view[%subgroup_score_column, %c0] shape [%m, %k] : view<64x16xf16> -> vector<16xf16> + %query_fragment = vector.fragment.load %query_transposed_view[%head_tile, %c0] shape [%k, %n] : view<128x16xf16, %query_transposed_layout> -> vector<16xf16> + %next_score_accumulator = vector.mma %key_fragment, %query_fragment, %score_accumulator : vector<16xf16>, vector<16xf16>, vector<4xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %next_score_accumulator : vector<4xf32> + } + scf.yield %tail_score_fragment : vector<4xf32> + } + vector.fragment.store %score_fragment, %score_stage_view[%subgroup_score_column, %c0] shape [%m, %n] : vector<4xf32>, view<64x24xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + %key_token = index.add %key_origin, %lane : index + %raw_score0 = view.load %score_stage_view[%lane, %query_row0] : view<64x24xf32> -> f32 + %raw_score1 = view.load %score_stage_view[%lane, %query_row1] : view<64x24xf32> -> f32 + %raw_score2 = view.load %score_stage_view[%lane, %query_row2] : view<64x24xf32> -> f32 + %raw_score3 = view.load %score_stage_view[%lane, %query_row3] : view<64x24xf32> -> f32 + %key_valid = index.cmp ult, %key_token, %bounded_key_value_token_count : index + %score_valid0 = scalar.andi %query_head_valid0, %key_valid : i1 + %score_valid1 = scalar.andi %query_head_valid1, %key_valid : i1 + %score_valid2 = scalar.andi %query_head_valid2, %key_valid : i1 + %score_valid3 = scalar.andi %query_head_valid3, %key_valid : i1 + %masked_score0 = scf.if %score_valid0 -> (f32) { + %score = scalar.addf %raw_score0, %lane_mask : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score1 = scf.if %score_valid1 -> (f32) { + %score = scalar.addf %raw_score1, %lane_mask : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score2 = scf.if %score_valid2 -> (f32) { + %score = scalar.addf %raw_score2, %lane_mask : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score3 = scf.if %score_valid3 -> (f32) { + %score = scalar.addf %raw_score3, %lane_mask : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_scores = vector.from_elements %masked_score0, %masked_score1, %masked_score2, %masked_score3 : vector<4xf32> + %block_max = kernel.subgroup.reduce %masked_scores : vector<4xf32> + %score_delta = vector.subf %masked_scores, %block_max : vector<4xf32> + %raw_probability = vector.expf %score_delta : vector<4xf32> + %probability = vector.select %query_valid, %raw_probability, %c0_f32x4 : vector<4xf32> + %block_sum = kernel.subgroup.reduce %probability : vector<4xf32> + %probability_f16 = vector.fptrunc %probability : vector<4xf32> to vector<4xf16> + %probability0 = vector.extract %probability_f16[0] : vector<4xf16> -> f16 + %probability1 = vector.extract %probability_f16[1] : vector<4xf16> -> f16 + %probability2 = vector.extract %probability_f16[2] : vector<4xf16> -> f16 + %probability3 = vector.extract %probability_f16[3] : vector<4xf16> -> f16 + view.store %probability0, %probability_stage_view[%query_row0, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability1, %probability_stage_view[%query_row1, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability2, %probability_stage_view[%query_row2, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability3, %probability_stage_view[%query_row3, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + kernel.barrier scope(workgroup) ordering(acq_rel) + // Keep both P*V halves live so the scheduler can interleave their independent + // matrix chains. F16 accumulation and exchange preserve the Vulkan CM1 + // contract while halving the product LDS footprint. + %product_init0 = vector.fragment %c0_f16x8 shape [%m, %n] : vector<8xf16> + %product_init1 = vector.fragment %c0_f16x8 shape [%m, %n] : vector<8xf16> + %value_channel0 = index.add %key_value_head_base, %subgroup_output_channel : index + %value_channel1 = index.add %key_value_head_base, %subgroup_output_channel1 : index + %product_fragment0, %product_fragment1 = scf.if %is_full_block -> (vector<8xf16>, vector<8xf16>) { + %full_product_fragment0, %full_product_fragment1 = scf.for %key_tile = [%c0 to %c64 step %c16](%product_accumulator0 = %product_init0 : vector<8xf16>, %product_accumulator1 = %product_init1 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>) unroll { + // Retain the same logical-extent contract for the full P*V tile. + %value_token0 = index.add %key_origin, %key_tile : index + %value_token_end0 = index.add %value_token0, %c16 : index + %value_token, %value_token_end = index.assume %value_token0, %value_token_end0 [le(%value_token_end0, %bounded_key_value_token_count)] : index, index + %full_value_view = buffer.view %value_aligned[%c0_offset] : buffer -> view<[%value_token_end]x[%key_value_width]xf16> + %probability_fragment = vector.fragment.load %probability_stage_view[%c0, %key_tile] shape [%m, %k] : view<16x64xf16, %probability_transposed_layout> -> vector<16xf16> + %value_fragment0 = vector.fragment.load %full_value_view[%value_token, %value_channel0] shape [%k, %n] : view<[%value_token_end]x[%key_value_width]xf16> -> vector<16xf16> + %value_fragment1 = vector.fragment.load %full_value_view[%value_token, %value_channel1] shape [%k, %n] : view<[%value_token_end]x[%key_value_width]xf16> -> vector<16xf16> + %next_product_accumulator0 = vector.mma %probability_fragment, %value_fragment0, %product_accumulator0 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %next_product_accumulator1 = vector.mma %probability_fragment, %value_fragment1, %product_accumulator1 : vector<16xf16>, vector<16xf16>, vector<8xf16> + scf.yield %next_product_accumulator0, %next_product_accumulator1 : vector<8xf16>, vector<8xf16> + } + scf.yield %full_product_fragment0, %full_product_fragment1 : vector<8xf16>, vector<8xf16> + } else { + // Stage one 16x128 V panel at a time and zero every lane beyond the + // logical KV length before the panel participates in P*V. + %tail_product_fragment0, %tail_product_fragment1 = scf.for %key_tile = [%c0 to %c64 step %c16](%product_accumulator0 = %product_init0 : vector<8xf16>, %product_accumulator1 = %product_init1 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>) unroll { + scf.for %load_iteration = [%c0 to %c8 step %c1] unroll { + %linear = index.madd %load_iteration, %c256, %workitem : index + %key_row = index.div %linear, %c128 : index + %value_channel = index.rem %linear, %c128 : index + %value_token0 = index.add %key_origin, %key_tile : index + %value_token = index.add %value_token0, %key_row : index + %value_valid = index.cmp ult, %value_token, %bounded_key_value_token_count : index + %value_element = scf.if %value_valid -> (f16) { + %global_value_channel = index.add %key_value_head_base, %value_channel : index + %loaded = view.load %value_view[%value_token, %global_value_channel] : view<[%bounded_key_value_token_count]x[%key_value_width]xf16> -> f16 + scf.yield %loaded : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %value_element, %tail_value_stage_view[%key_row, %value_channel] : f16, view<16x128xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %probability_fragment = vector.fragment.load %probability_stage_view[%c0, %key_tile] shape [%m, %k] : view<16x64xf16, %probability_transposed_layout> -> vector<16xf16> + %value_fragment0 = vector.fragment.load %tail_value_stage_view[%c0, %subgroup_output_channel] shape [%k, %n] : view<16x128xf16> -> vector<16xf16> + %value_fragment1 = vector.fragment.load %tail_value_stage_view[%c0, %subgroup_output_channel1] shape [%k, %n] : view<16x128xf16> -> vector<16xf16> + %next_product_accumulator0 = vector.mma %probability_fragment, %value_fragment0, %product_accumulator0 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %next_product_accumulator1 = vector.mma %probability_fragment, %value_fragment1, %product_accumulator1 : vector<16xf16>, vector<16xf16>, vector<8xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %next_product_accumulator0, %next_product_accumulator1 : vector<8xf16>, vector<8xf16> + } + scf.yield %tail_product_fragment0, %tail_product_fragment1 : vector<8xf16>, vector<8xf16> + } + vector.fragment.store %product_fragment0, %product_stage_view[%c0, %subgroup_output_channel] shape [%m, %n] : vector<8xf16>, view<16x128xf16> + vector.fragment.store %product_fragment1, %product_stage_view[%c0, %subgroup_output_channel1] shape [%m, %n] : vector<8xf16>, view<16x128xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.if %lane_has_output { + %block_output0 = vector.load %product_stage_view[%query_row0, %lane_output_channel] : view<16x128xf16> -> vector<4xf16> + %block_output1 = vector.load %product_stage_view[%query_row1, %lane_output_channel] : view<16x128xf16> -> vector<4xf16> + %block_output2 = vector.load %product_stage_view[%query_row2, %lane_output_channel] : view<16x128xf16> -> vector<4xf16> + %block_output3 = vector.load %product_stage_view[%query_row3, %lane_output_channel] : view<16x128xf16> -> vector<4xf16> + scf.if %query_head_valid0 { + vector.store %block_output0, %partial_output_view[%key_value_head, %block_ordinal, %query_row0, %lane_output_channel] : vector<4xf16>, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16x128xf16> + } + scf.if %query_head_valid1 { + vector.store %block_output1, %partial_output_view[%key_value_head, %block_ordinal, %query_row1, %lane_output_channel] : vector<4xf16>, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16x128xf16> + } + scf.if %query_head_valid2 { + vector.store %block_output2, %partial_output_view[%key_value_head, %block_ordinal, %query_row2, %lane_output_channel] : vector<4xf16>, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16x128xf16> + } + scf.if %query_head_valid3 { + vector.store %block_output3, %partial_output_view[%key_value_head, %block_ordinal, %query_row3, %lane_output_channel] : vector<4xf16>, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16x128xf16> + } + } + scf.if %lane_is_zero { + scf.if %query_head_valid0 { + %maximum = vector.extract %block_max[0] : vector<4xf32> -> f32 + %sum = vector.extract %block_sum[0] : vector<4xf32> -> f32 + view.store %maximum, %partial_max_view[%key_value_head, %block_ordinal, %query_row0] : f32, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + view.store %sum, %partial_sum_view[%key_value_head, %block_ordinal, %query_row0] : f32, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + } + scf.if %query_head_valid1 { + %maximum = vector.extract %block_max[1] : vector<4xf32> -> f32 + %sum = vector.extract %block_sum[1] : vector<4xf32> -> f32 + view.store %maximum, %partial_max_view[%key_value_head, %block_ordinal, %query_row1] : f32, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + view.store %sum, %partial_sum_view[%key_value_head, %block_ordinal, %query_row1] : f32, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + } + scf.if %query_head_valid2 { + %maximum = vector.extract %block_max[2] : vector<4xf32> -> f32 + %sum = vector.extract %block_sum[2] : vector<4xf32> -> f32 + view.store %maximum, %partial_max_view[%key_value_head, %block_ordinal, %query_row2] : f32, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + view.store %sum, %partial_sum_view[%key_value_head, %block_ordinal, %query_row2] : f32, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + } + scf.if %query_head_valid3 { + %maximum = vector.extract %block_max[3] : vector<4xf32> -> f32 + %sum = vector.extract %block_sum[3] : vector<4xf32> -> f32 + view.store %maximum, %partial_max_view[%key_value_head, %block_ordinal, %query_row3] : f32, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + view.store %sum, %partial_sum_view[%key_value_head, %block_ordinal, %query_row3] : f32, view<[%key_value_head_count]x[%launch_partial_block_capacity]x16xf32> + } + } + template.return +} + +// Fully masked splits publish the online-softmax identity without running QK +// or P*V. Valid additive masks contain finite F16 values or negative infinity, +// and every finite F16 value is greater than -1e30. Comparing after extension +// therefore distinguishes active rows from masked rows exactly. +template.def<@qwen3_moe.attention.decode_split.produce_partials> device @qwen3_moe_flash_attention_decode_split_produce_partials_body_f32_f16_wmma(%key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer) { + %bounded_key_value_token_count = index.assume %key_value_token_count [range(%key_value_token_count, 1, 32768)] : index + %query_head_count = config.get @qwen3_moe.attention.query_head_count : index + %key_value_head_count = config.get @qwen3_moe.attention.key_value_head_count : index + %workgroup_x0 = kernel.workgroup.id : index + %workgroup_y0 = kernel.workgroup.id : index + %workgroup_x_in_launch = index.assume %workgroup_x0 [range(%workgroup_x0, 0, 511)] : index + %workgroup_y = index.assume %workgroup_y0 [range(%workgroup_y0, 0, 63)] : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 3)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c0_offset = index.constant 0 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %c0_f16x4 = vector.constant 0.0 : vector<4xf16> + %key_value_token_capacity = config.get @qwen3_moe.attention.key_value_token_capacity : index + %c63 = index.constant 63 : index + %padded_key_value_token_capacity = index.add %key_value_token_capacity, %c63 : index + %key_value_block_count0 = index.div %padded_key_value_token_capacity, %c64 : index + %key_value_block_count = index.assume %key_value_block_count0 [range(%key_value_block_count0, 1, 512)] : index + %workgroup_x, %launch_key_value_block_count = index.assume %workgroup_x_in_launch, %key_value_block_count [lt(%workgroup_x_in_launch, %key_value_block_count)] : index, index + %block_ordinal = index.add %workgroup_x, %c0 : index + %key_origin = index.mul %block_ordinal, %c64 : index + %key_token = index.add %key_origin, %lane : index + %key_valid = index.cmp ult, %key_token, %bounded_key_value_token_count : index + %mask_aligned = buffer.assume.alignment %mask {minimum_alignment = 16} : buffer + %mask_view = buffer.view %mask_aligned[%c0_offset] : buffer -> view<[%bounded_key_value_token_count]xf16> + %lane_mask = scf.if %key_valid -> (f32) { + %mask_f16 = view.load %mask_view[%key_token] : view<[%bounded_key_value_token_count]xf16> -> f16 + %mask_f32 = scalar.extf %mask_f16 : f16 to f32 + scf.yield %mask_f32 : f32 + } else { + scf.yield %negative_large : f32 + } + %block_mask_maximum = kernel.workgroup.reduce %lane_mask : f32 + %block_has_attention = scalar.cmpf ogt, %block_mask_maximum, %negative_large : f32 + scf.if %block_has_attention { + template.apply<@qwen3_moe.attention.decode_split.produce_partials.active>(%bounded_key_value_token_count, %launch_key_value_block_count, %launch_key_value_block_count, %query, %key, %value, %lane_mask, %partial_max, %partial_sum, %partial_output) : (index, index, index, buffer, buffer, buffer, f32, buffer, buffer, buffer) + } else { + %query_heads_per_key_value_head0 = index.div %query_head_count, %key_value_head_count : index + %query_heads_per_key_value_head = index.assume %query_heads_per_key_value_head0 [range(%query_heads_per_key_value_head0, 1, 16)] : index + %key_value_head = index.add %workgroup_y, %c0 : index + %query_head_base = index.mul %key_value_head, %query_heads_per_key_value_head : index + %subgroup_query_row = index.mul %subgroup, %c4 : index + %query_row0 = index.add %subgroup_query_row, %c0 : index + %query_row1 = index.add %subgroup_query_row, %c1 : index + %query_row2 = index.add %subgroup_query_row, %c2 : index + %query_row3 = index.add %subgroup_query_row, %c3 : index + %query_head0 = index.add %query_head_base, %query_row0 : index + %query_head1 = index.add %query_head_base, %query_row1 : index + %query_head2 = index.add %query_head_base, %query_row2 : index + %query_head3 = index.add %query_head_base, %query_row3 : index + %query_head_valid0 = index.cmp ult, %query_head0, %query_head_count : index + %query_head_valid1 = index.cmp ult, %query_head1, %query_head_count : index + %query_head_valid2 = index.cmp ult, %query_head2, %query_head_count : index + %query_head_valid3 = index.cmp ult, %query_head3, %query_head_count : index + %lane_output_channel = index.mul %lane, %c4 : index + %lane_has_output = index.cmp ult, %lane, %c32 : index + %lane_is_zero = index.cmp eq, %lane, %c0 : index + %partial_max_noalias, %partial_sum_noalias, %partial_output_noalias = buffer.assume.noalias %partial_max, %partial_sum, %partial_output : buffer, buffer, buffer + %partial_max_aligned = buffer.assume.alignment %partial_max_noalias {minimum_alignment = 16} : buffer + %partial_sum_aligned = buffer.assume.alignment %partial_sum_noalias {minimum_alignment = 16} : buffer + %partial_output_aligned = buffer.assume.alignment %partial_output_noalias {minimum_alignment = 16} : buffer + %partial_max_view = buffer.view %partial_max_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + %partial_sum_view = buffer.view %partial_sum_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + %partial_output_view = buffer.view %partial_output_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%key_value_block_count]x16x128xf16> + scf.if %lane_has_output { + scf.if %query_head_valid0 { + vector.store %c0_f16x4, %partial_output_view[%key_value_head, %block_ordinal, %query_row0, %lane_output_channel] : vector<4xf16>, view<[%key_value_head_count]x[%key_value_block_count]x16x128xf16> + } + scf.if %query_head_valid1 { + vector.store %c0_f16x4, %partial_output_view[%key_value_head, %block_ordinal, %query_row1, %lane_output_channel] : vector<4xf16>, view<[%key_value_head_count]x[%key_value_block_count]x16x128xf16> + } + scf.if %query_head_valid2 { + vector.store %c0_f16x4, %partial_output_view[%key_value_head, %block_ordinal, %query_row2, %lane_output_channel] : vector<4xf16>, view<[%key_value_head_count]x[%key_value_block_count]x16x128xf16> + } + scf.if %query_head_valid3 { + vector.store %c0_f16x4, %partial_output_view[%key_value_head, %block_ordinal, %query_row3, %lane_output_channel] : vector<4xf16>, view<[%key_value_head_count]x[%key_value_block_count]x16x128xf16> + } + } + scf.if %lane_is_zero { + scf.if %query_head_valid0 { + view.store %negative_large, %partial_max_view[%key_value_head, %block_ordinal, %query_row0] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + view.store %c0_f32, %partial_sum_view[%key_value_head, %block_ordinal, %query_row0] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + } + scf.if %query_head_valid1 { + view.store %negative_large, %partial_max_view[%key_value_head, %block_ordinal, %query_row1] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + view.store %c0_f32, %partial_sum_view[%key_value_head, %block_ordinal, %query_row1] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + } + scf.if %query_head_valid2 { + view.store %negative_large, %partial_max_view[%key_value_head, %block_ordinal, %query_row2] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + view.store %c0_f32, %partial_sum_view[%key_value_head, %block_ordinal, %query_row2] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + } + scf.if %query_head_valid3 { + view.store %negative_large, %partial_max_view[%key_value_head, %block_ordinal, %query_row3] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + view.store %c0_f32, %partial_sum_view[%key_value_head, %block_ordinal, %query_row3] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + } + } + } + template.return +} + +// Up to four split-K blocks are cheapest to fold directly. Each workitem owns +// one output element, so the reducer needs no LDS or subgroup synchronization. +template.def<@qwen3_moe.attention.decode_split.reduce_completed.direct> device @qwen3_moe_flash_attention_decode_split_reduce_completed_direct_f32(%partial_block_capacity0: index, %active_block_count0: index, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %output: buffer) where [range(%partial_block_capacity0, 1, 4)] { + %partial_block_capacity, %active_block_count = index.assume %partial_block_capacity0, %active_block_count0 [range(%partial_block_capacity0, 1, 4), range(%active_block_count0, 1, 4), le(%active_block_count0, %partial_block_capacity0)] : index, index + %query_head_count = config.get @qwen3_moe.attention.query_head_count : index + %key_value_head_count = config.get @qwen3_moe.attention.key_value_head_count : index + %workgroup_y0 = kernel.workgroup.id : index + %key_value_head = index.assume %workgroup_y0 [range(%workgroup_y0, 0, 63)] : index + %workitem = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %query_heads_per_key_value_head0 = index.div %query_head_count, %key_value_head_count : index + %query_heads_per_key_value_head = index.assume %query_heads_per_key_value_head0 [range(%query_heads_per_key_value_head0, 1, 16)] : index + %query_head_base = index.mul %key_value_head, %query_heads_per_key_value_head : index + %partial_max_aligned = buffer.assume.alignment %partial_max {minimum_alignment = 16} : buffer + %partial_sum_aligned = buffer.assume.alignment %partial_sum {minimum_alignment = 16} : buffer + %partial_output_aligned = buffer.assume.alignment %partial_output {minimum_alignment = 16} : buffer + %output_aligned = buffer.assume.alignment %output {minimum_alignment = 16} : buffer + %partial_max_view = buffer.view %partial_max_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> + %partial_sum_view = buffer.view %partial_sum_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> + %partial_output_view = buffer.view %partial_output_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%partial_block_capacity]x16x128xf16> + %output_view = buffer.view %output_aligned[%c0_offset] : buffer -> view<[%query_head_count]x128xf32> + %partial_element_count = index.mul %query_heads_per_key_value_head, %c128 : index + scf.for %linear = [%workitem to %partial_element_count step %c256] { + %query_row = index.div %linear, %c128 : index + %output_channel = index.rem %linear, %c128 : index + %query_head = index.add %query_head_base, %query_row : index + %query_head_valid = index.cmp ult, %query_head, %query_head_count : index + scf.if %query_head_valid { + %maximum = scf.for %block = [%c0 to %active_block_count step %c1](%running_maximum = %negative_large : f32) -> (f32) unroll schedule(interleaved) { + %block_maximum = view.load %partial_max_view[%key_value_head, %block, %query_row] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + %next_maximum = scalar.maxnumf %running_maximum, %block_maximum : f32 + scf.yield %next_maximum : f32 + } + %sum, %unnormalized_output = scf.for %block = [%c0 to %active_block_count step %c1](%running_sum = %c0_f32 : f32, %running_output = %c0_f32 : f32) -> (f32, f32) unroll schedule(interleaved) { + %block_maximum = view.load %partial_max_view[%key_value_head, %block, %query_row] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + %delta = scalar.subf %block_maximum, %maximum : f32 + %scale = scalar.expf %delta : f32 + %block_sum = view.load %partial_sum_view[%key_value_head, %block, %query_row] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + %scaled_sum = scalar.mulf %block_sum, %scale : f32 + %next_sum = scalar.addf %running_sum, %scaled_sum : f32 + %block_output_f16 = view.load %partial_output_view[%key_value_head, %block, %query_row, %output_channel] : view<[%key_value_head_count]x[%partial_block_capacity]x16x128xf16> -> f16 + %block_output = scalar.extf %block_output_f16 : f16 to f32 + %scaled_output = scalar.mulf %block_output, %scale : f32 + %next_output = scalar.addf %running_output, %scaled_output : f32 + scf.yield %next_sum, %next_output : f32, f32 + } + %normalized_output = scalar.divf %unnormalized_output, %sum : f32 + view.store %normalized_output, %output_view[%query_head, %output_channel] : f32, view<[%query_head_count]x128xf32> + } + } + template.return +} + +// Longer bounded contexts amortize a cooperative reducer. Each wave folds two +// query rows, stages the per-block scales in LDS, and writes two output channels +// per lane. +template.def<@qwen3_moe.attention.decode_split.reduce_completed.cooperative> device @qwen3_moe_flash_attention_decode_split_reduce_completed_cooperative_f32(%partial_block_capacity0: index, %active_block_count0: index, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %output: buffer) { + %partial_block_capacity, %active_block_count = index.assume %partial_block_capacity0, %active_block_count0 [range(%partial_block_capacity0, 1, 32), range(%active_block_count0, 1, 32), le(%active_block_count0, %partial_block_capacity0)] : index, index + %query_head_count = config.get @qwen3_moe.attention.query_head_count : index + %key_value_head_count = config.get @qwen3_moe.attention.key_value_head_count : index + %workgroup_y0 = kernel.workgroup.id : index + %key_value_head = index.assume %workgroup_y0 [range(%workgroup_y0, 0, 63)] : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 3)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c0_offset = index.constant 0 : offset + %scale_stage_bytes = index.constant 1024 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %c0_f32x2 = vector.constant 0.0 : vector<2xf32> + %query_heads_per_key_value_head0 = index.div %query_head_count, %key_value_head_count : index + %query_heads_per_key_value_head = index.assume %query_heads_per_key_value_head0 [range(%query_heads_per_key_value_head0, 1, 16)] : index + %query_head_base = index.mul %key_value_head, %query_heads_per_key_value_head : index + %lane_output_channel0 = index.mul %lane, %c2 : index + %lane_output_channel = index.assume %lane_output_channel0 [range(%lane_output_channel0, 0, 126)] : index + %lane_has_block = index.cmp ult, %lane, %active_block_count : index + %partial_max_aligned = buffer.assume.alignment %partial_max {minimum_alignment = 16} : buffer + %partial_sum_aligned = buffer.assume.alignment %partial_sum {minimum_alignment = 16} : buffer + %partial_output_aligned = buffer.assume.alignment %partial_output {minimum_alignment = 16} : buffer + %output_aligned = buffer.assume.alignment %output {minimum_alignment = 16} : buffer + %partial_max_view = buffer.view %partial_max_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> + %partial_sum_view = buffer.view %partial_sum_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> + %partial_output_view = buffer.view %partial_output_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%partial_block_capacity]x16x128xf16> + %output_view = buffer.view %output_aligned[%c0_offset] : buffer -> view<[%query_head_count]x128xf32> + %scale_stage = buffer.alloca align(16) %scale_stage_bytes : buffer + %scale_stage_view = buffer.view %scale_stage[%c0_offset] : buffer -> view<4x32x2xf32> + %padded_reduction_row_count = index.add %query_heads_per_key_value_head, %c7 : index + %reduction_phase_count = index.div %padded_reduction_row_count, %c8 : index + scf.for %phase = [%c0 to %reduction_phase_count step %c1] unroll { + %phase_row_base = index.mul %phase, %c8 : index + %subgroup_row_base = index.mul %subgroup, %c2 : index + %query_row0 = index.add %phase_row_base, %subgroup_row_base : index + %query_row1 = index.add %query_row0, %c1 : index + %query_head0 = index.add %query_head_base, %query_row0 : index + %query_head1 = index.add %query_head0, %c1 : index + // A KV head owns only its own query rows; rows past that belong to the next + // KV head's workgroup, and reducing them here overwrites its output. + %query_row_owned0 = index.cmp ult, %query_row0, %query_heads_per_key_value_head : index + %query_row_owned1 = index.cmp ult, %query_row1, %query_heads_per_key_value_head : index + %query_head_in_range0 = index.cmp ult, %query_head0, %query_head_count : index + %query_head_in_range1 = index.cmp ult, %query_head1, %query_head_count : index + %query_head_valid0 = scalar.andi %query_row_owned0, %query_head_in_range0 : i1 + %query_head_valid1 = scalar.andi %query_row_owned1, %query_head_in_range1 : i1 + %safe_query_row0 = scf.select %query_head_valid0, %query_row0, %c0 : index + %safe_query_row1 = scf.select %query_head_valid1, %query_row1, %c0 : index + %reducer_active0 = scalar.andi %lane_has_block, %query_head_valid0 : i1 + %reducer_active1 = scalar.andi %lane_has_block, %query_head_valid1 : i1 + %local_maximum0 = scf.if %reducer_active0 -> (f32) { + %block_maximum = view.load %partial_max_view[%key_value_head, %lane, %query_row0] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + scf.yield %block_maximum : f32 + } else { + scf.yield %negative_large : f32 + } + %local_maximum1 = scf.if %reducer_active1 -> (f32) { + %block_maximum = view.load %partial_max_view[%key_value_head, %lane, %query_row1] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + scf.yield %block_maximum : f32 + } else { + scf.yield %negative_large : f32 + } + %local_maximums = vector.from_elements %local_maximum0, %local_maximum1 : vector<2xf32> + %maximums = kernel.subgroup.reduce %local_maximums : vector<2xf32> + %maximum0 = vector.extract %maximums[0] : vector<2xf32> -> f32 + %maximum1 = vector.extract %maximums[1] : vector<2xf32> -> f32 + %local_sum0, %scale0 = scf.if %reducer_active0 -> (f32, f32) { + %block_maximum = view.load %partial_max_view[%key_value_head, %lane, %query_row0] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + %delta = scalar.subf %block_maximum, %maximum0 : f32 + %scale = scalar.expf %delta : f32 + %block_sum = view.load %partial_sum_view[%key_value_head, %lane, %query_row0] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + %scaled_sum = scalar.mulf %block_sum, %scale : f32 + scf.yield %scaled_sum, %scale : f32, f32 + } else { + scf.yield %c0_f32, %c0_f32 : f32, f32 + } + %local_sum1, %scale1 = scf.if %reducer_active1 -> (f32, f32) { + %block_maximum = view.load %partial_max_view[%key_value_head, %lane, %query_row1] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + %delta = scalar.subf %block_maximum, %maximum1 : f32 + %scale = scalar.expf %delta : f32 + %block_sum = view.load %partial_sum_view[%key_value_head, %lane, %query_row1] : view<[%key_value_head_count]x[%partial_block_capacity]x16xf32> -> f32 + %scaled_sum = scalar.mulf %block_sum, %scale : f32 + scf.yield %scaled_sum, %scale : f32, f32 + } else { + scf.yield %c0_f32, %c0_f32 : f32, f32 + } + scf.if %lane_has_block { + %bounded_block_lane = index.assume %lane [range(%lane, 0, 31)] : index + %scales = vector.from_elements %scale0, %scale1 : vector<2xf32> + vector.store %scales, %scale_stage_view[%subgroup, %bounded_block_lane, %c0] : vector<2xf32>, view<4x32x2xf32> + } + %local_sums = vector.from_elements %local_sum0, %local_sum1 : vector<2xf32> + %sums = kernel.subgroup.reduce %local_sums : vector<2xf32> + %sum0 = vector.extract %sums[0] : vector<2xf32> -> f32 + %sum1 = vector.extract %sums[1] : vector<2xf32> -> f32 + kernel.barrier scope(workgroup) ordering(acq_rel) + %unnormalized_output0, %unnormalized_output1 = scf.for %block = [%c0 to %active_block_count step %c1](%running_output0 = %c0_f32x2 : vector<2xf32>, %running_output1 = %c0_f32x2 : vector<2xf32>) -> (vector<2xf32>, vector<2xf32>) unroll(%c4) schedule(interleaved) { + %scales = vector.load %scale_stage_view[%subgroup, %block, %c0] : view<4x32x2xf32> -> vector<2xf32> + %scale0 = vector.extract %scales[0] : vector<2xf32> -> f32 + %scale1 = vector.extract %scales[1] : vector<2xf32> -> f32 + %block_output0_f16 = vector.load %partial_output_view[%key_value_head, %block, %safe_query_row0, %lane_output_channel] : view<[%key_value_head_count]x[%partial_block_capacity]x16x128xf16> -> vector<2xf16> + %block_output0 = vector.extf %block_output0_f16 : vector<2xf16> to vector<2xf32> + %scale_vector0 = vector.splat %scale0 : vector<2xf32> + %scaled_output0 = vector.mulf %block_output0, %scale_vector0 : vector<2xf32> + %next_output0 = vector.addf %running_output0, %scaled_output0 : vector<2xf32> + %block_output1_f16 = vector.load %partial_output_view[%key_value_head, %block, %safe_query_row1, %lane_output_channel] : view<[%key_value_head_count]x[%partial_block_capacity]x16x128xf16> -> vector<2xf16> + %block_output1 = vector.extf %block_output1_f16 : vector<2xf16> to vector<2xf32> + %scale_vector1 = vector.splat %scale1 : vector<2xf32> + %scaled_output1 = vector.mulf %block_output1, %scale_vector1 : vector<2xf32> + %next_output1 = vector.addf %running_output1, %scaled_output1 : vector<2xf32> + scf.yield %next_output0, %next_output1 : vector<2xf32>, vector<2xf32> + } + %sum_vector0 = vector.splat %sum0 : vector<2xf32> + %sum_vector1 = vector.splat %sum1 : vector<2xf32> + %normalized_output0 = vector.divf %unnormalized_output0, %sum_vector0 : vector<2xf32> + %normalized_output1 = vector.divf %unnormalized_output1, %sum_vector1 : vector<2xf32> + scf.if %query_head_valid0 { + vector.store %normalized_output0, %output_view[%query_head0, %lane_output_channel] : vector<2xf32>, view<[%query_head_count]x128xf32> + } + scf.if %query_head_valid1 { + vector.store %normalized_output1, %output_view[%query_head1, %lane_output_channel] : vector<2xf32>, view<[%query_head_count]x128xf32> + } + } + template.return +} + +// Packs the query heads owned by one completed KV head into GGML's Q8_1 x4 +// layout. Each 32-workitem cohort owns one contiguous 128-element query row; +// the production 8:1 GQA ratio therefore fills all 256 workitems. Smaller or +// larger valid ratios use the same phase loop without making barriers +// conditional. +template.def<@qwen3_moe.attention.decode_split.pack_completed_q8> device @qwen3_moe_flash_attention_decode_pack_completed_key_value_head_q8_1_x4(%key_value_head: index, %output: buffer, %q8_output: buffer) { + %query_head_count = config.get @qwen3_moe.attention.query_head_count : index + %key_value_head_count = config.get @qwen3_moe.attention.key_value_head_count : index + %query_heads_per_key_value_head0 = index.div %query_head_count, %key_value_head_count : index + %query_heads_per_key_value_head = index.assume %query_heads_per_key_value_head0 [range(%query_heads_per_key_value_head0, 1, 16)] : index + %workitem0 = kernel.workitem.id : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 255)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c32 = index.constant 32 : index + %group_bytes = index.constant 144 : offset + %payload_byte_add = index.constant 16 : offset + %scratch_d_byte_add = index.constant 1024 : offset + %scratch_bytes = index.constant 1152 : offset + %c0_offset = index.constant 0 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %c1_f32 = scalar.constant 1.0 : f32 + %c127 = scalar.constant 127.0 : f32 + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %output_noalias, %q8_output_noalias = buffer.assume.noalias %output, %q8_output : buffer, buffer + %output_aligned = buffer.assume.alignment %output_noalias {minimum_alignment = 16} : buffer + %q8_output_aligned = buffer.assume.alignment %q8_output_noalias {minimum_alignment = 16} : buffer + %output_view = buffer.view %output_aligned[%c0_offset] : buffer -> view<[%query_head_count]x128xf32> + %scratch = buffer.alloca align(16) %scratch_bytes : buffer + %scratch_values = buffer.view %scratch[%c0_offset] : buffer -> view<256xf32> + %scratch_d = buffer.view %scratch[%scratch_d_byte_add] : buffer -> view<32xf32> + %group_in_phase = index.div %workitem, %c32 : index + %lane_in_group0 = index.rem %workitem, %c32 : index + %lane_in_group = index.assume %lane_in_group0 [range(%lane_in_group0, 0, 31)] : index + %block_in_group0 = index.div %lane_in_group, %c8 : index + %block_in_group = index.assume %block_in_group0 [range(%block_in_group0, 0, 3)] : index + %word_in_block0 = index.rem %lane_in_group, %c8 : index + %word_in_block = index.assume %word_in_block0 [range(%word_in_block0, 0, 7)] : index + %padded_phase_count = index.add %query_heads_per_key_value_head, %c7 : index + %phase_count = index.div %padded_phase_count, %c8 : index + scf.for %phase = [%c0 to %phase_count step %c1] unroll { + %phase_group_base = index.mul %phase, %c8 : index + %group_in_key_value_head = index.add %phase_group_base, %group_in_phase : index + %valid_group = index.cmp ult, %group_in_key_value_head, %query_heads_per_key_value_head : index + %key_value_query_head_base = index.mul %key_value_head, %query_heads_per_key_value_head : index + %query_head0 = index.add %key_value_query_head_base, %group_in_key_value_head : index + %safe_query_head = scf.select %valid_group, %query_head0, %c0 : index + %block_element_add = index.mul %block_in_group, %c32 : index + %word_element_add = index.mul %word_in_block, %c4 : index + %input_channel0 = index.add %block_element_add, %word_element_add : index + %input_channel = index.assume %input_channel0 [range(%input_channel0, 0, 124)] : index + %input_values = scf.if %valid_group -> (vector<4xf32>) { + %values = vector.load %output_view[%safe_query_head, %input_channel] : view<[%query_head_count]x128xf32> -> vector<4xf32> + scf.yield %values : vector<4xf32> + } else { + scf.yield %c0_f32x4 : vector<4xf32> + } + %absolute_values = vector.absf %input_values : vector<4xf32> + %thread_max = vector.reduce %absolute_values, %c0_f32 : vector<4xf32>, f32 + view.store %thread_max, %scratch_values[%workitem] : f32, view<256xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + %is_cohort_leader = index.cmp eq, %word_in_block, %c0 : index + %d_index0 = index.mul %group_in_phase, %c4 : index + %d_index1 = index.add %d_index0, %block_in_group : index + %d_index = index.assume %d_index1 [range(%d_index1, 0, 31)] : index + scf.if %is_cohort_leader { + %group_thread_base = index.mul %group_in_phase, %c32 : index + %block_thread_add = index.mul %block_in_group, %c8 : index + %cohort_base = index.add %group_thread_base, %block_thread_add : index + %cohort_maxima = vector.load %scratch_values[%cohort_base] : view<256xf32> -> vector<8xf32> + %amax = vector.reduce %cohort_maxima, %c0_f32 : vector<8xf32>, f32 + %d = scalar.divf %amax, %c127 : f32 + view.store %d, %scratch_d[%d_index] : f32, view<32xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %d = view.load %scratch_d[%d_index] : view<32xf32> -> f32 + %d_nonzero = scalar.cmpf one, %d, %c0_f32 : f32 + %d_inverse = scf.if %d_nonzero -> (f32) { + %inverse = scalar.divf %c1_f32, %d : f32 + scf.yield %inverse : f32 + } else { + scf.yield %c0_f32 : f32 + } + %d_inverse_vector = vector.splat %d_inverse : vector<4xf32> + %scaled_values = vector.mulf %input_values, %d_inverse_vector : vector<4xf32> + %rounded_values = vector.roundf %scaled_values : vector<4xf32> + %quantized_values = vector.fptosi %rounded_values : vector<4xf32> to vector<4xi8> + %packed_word = vector.bitcast %quantized_values : vector<4xi8> to vector<1xi32> + %group_byte_offset = index.scale %safe_query_head, %group_bytes : index, offset -> offset + %payload_byte_offset = index.add %group_byte_offset, %payload_byte_add : offset + %group_ds = buffer.view %q8_output_aligned[%group_byte_offset] : buffer -> view<8xf16> + %group_qs = buffer.view %q8_output_aligned[%payload_byte_offset] : buffer -> view<32xi32> + %block_word_add = index.mul %block_in_group, %c8 : index + %packed_word_index0 = index.add %block_word_add, %word_in_block : index + %packed_word_index = index.assume %packed_word_index0 [range(%packed_word_index0, 0, 31)] : index + scf.if %valid_group { + vector.store %packed_word, %group_qs[%packed_word_index] : vector<1xi32>, view<32xi32> + } + %thread_sum = vector.reduce %rounded_values, %c0_f32 : vector<4xf32>, f32 + view.store %thread_sum, %scratch_values[%workitem] : f32, view<256xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.if %is_cohort_leader { + %group_thread_base = index.mul %group_in_phase, %c32 : index + %block_thread_add = index.mul %block_in_group, %c8 : index + %cohort_base = index.add %group_thread_base, %block_thread_add : index + %cohort_sums = vector.load %scratch_values[%cohort_base] : view<256xf32> -> vector<8xf32> + %quantized_sum = vector.reduce %cohort_sums, %c0_f32 : vector<8xf32>, f32 + %s = scalar.mulf %quantized_sum, %d : f32 + %d_f16 = scalar.fptrunc %d : f32 to f16 + %s_f16 = scalar.fptrunc %s : f32 to f16 + %ds_index = index.mul %block_in_group, %c2 : index + %s_index = index.add %ds_index, %c1 : index + scf.if %valid_group { + view.store %d_f16, %group_ds[%ds_index] : f16, view<8xf16> + view.store %s_f16, %group_ds[%s_index] : f16, view<8xf16> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + } + template.return +} + +// Completes bounded decode contexts inside the last arriving producer +// workgroup. Four KV-head workgroups reduce their own query heads while all +// other producers retire, erasing a second dispatch and its execution barrier. +// Capacity selects the algorithm and partial layout; producer count owns the +// issue-time completion threshold and active reduction prefix. +template.def<@qwen3_moe.attention.decode_split.reduce_fused> device priority(20) @qwen3_moe_flash_attention_decode_split_reduce_fused_direct_f32(%key_value_token_capacity: index, %partial_block_capacity0: index, %producer_block_count0: index, %publish_q8: i1, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %completion_counter: buffer, %output: buffer, %q8_output: buffer) where [range(%key_value_token_capacity, 64, 256)] { + %partial_block_capacity, %producer_block_count = index.assume %partial_block_capacity0, %producer_block_count0 [range(%partial_block_capacity0, 1, 4), range(%producer_block_count0, 1, 4), le(%producer_block_count0, %partial_block_capacity0)] : index, index + %key_value_head_count = config.get @qwen3_moe.attention.key_value_head_count : index + %workgroup_y0 = kernel.workgroup.id : index + %key_value_head = index.assume %workgroup_y0 [range(%workgroup_y0, 0, 63)] : index + %workitem = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %counter_scratch_bytes = index.constant 4 : offset + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %partial_max_noalias, %partial_sum_noalias, %partial_output_noalias, %completion_counter_noalias, %output_noalias = buffer.assume.noalias %partial_max, %partial_sum, %partial_output, %completion_counter, %output : buffer, buffer, buffer, buffer, buffer + %completion_counter_aligned = buffer.assume.alignment %completion_counter_noalias {minimum_alignment = 16} : buffer + %completion_counter_view = buffer.view %completion_counter_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]xi32> + %counter_scratch = buffer.alloca align(4) %counter_scratch_bytes : buffer + %counter_scratch_view = buffer.view %counter_scratch[%c0_offset] : buffer -> view<1xi32> + %workitem_is_zero = index.cmp eq, %workitem, %c0 : index + // Publish every producer's partial stores before the leader advances one + // workgroup arrival. The last arrival then acquires every partial. + kernel.barrier scope(workgroup) ordering(release) + scf.if %workitem_is_zero { + %old_counter = view.atomic.rmw %c1_i32, %completion_counter_view[%key_value_head] {ordering = acq_rel, scope = device} : i32, view<[%key_value_head_count]xi32> -> i32 + view.store %old_counter, %counter_scratch_view[%c0] : i32, view<1xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %old_counter = view.load %counter_scratch_view[%c0] : view<1xi32> -> i32 + %key_value_block_count_i32 = index.cast %producer_block_count : index to i32 + %last_block_ordinal_i32 = scalar.subi %key_value_block_count_i32, %c1_i32 : i32 + %negative_key_value_block_count_i32 = scalar.subi %c0_i32, %key_value_block_count_i32 : i32 + %is_last_partition = scalar.cmpi eq, %old_counter, %last_block_ordinal_i32 : i32 + scf.if %is_last_partition { + kernel.barrier scope(workgroup) ordering(acquire) + template.apply<@qwen3_moe.attention.decode_split.reduce_completed.direct>(%partial_block_capacity, %producer_block_count, %partial_max_noalias, %partial_sum_noalias, %partial_output_noalias, %output_noalias) : (index, index, buffer, buffer, buffer, buffer) + scf.if %publish_q8 { + kernel.barrier scope(workgroup) ordering(acq_rel) + kernel.barrier scope(workgroup) ordering(acq_rel) + template.apply<@qwen3_moe.attention.decode_split.pack_completed_q8>(%key_value_head, %output_noalias, %q8_output) : (index, buffer, buffer) + } + // Do not expose the reset until every final F32 and Q8 store completes. + kernel.barrier scope(workgroup) ordering(release) + scf.if %workitem_is_zero { + view.atomic.reduce %negative_key_value_block_count_i32, %completion_counter_view[%key_value_head] {ordering = release, scope = device} : i32, view<[%key_value_head_count]xi32> + } + } + template.return +} + +// Contexts without a proven short bound use the cooperative completion path. +template.def<@qwen3_moe.attention.decode_split.reduce_fused> device priority(10) @qwen3_moe_flash_attention_decode_split_reduce_fused_cooperative_f32(%key_value_token_capacity: index, %partial_block_capacity0: index, %producer_block_count0: index, %publish_q8: i1, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %completion_counter: buffer, %output: buffer, %q8_output: buffer) where [range(%key_value_token_capacity, 257, 2048)] { + %partial_block_capacity, %producer_block_count = index.assume %partial_block_capacity0, %producer_block_count0 [range(%partial_block_capacity0, 1, 32), range(%producer_block_count0, 1, 32), le(%producer_block_count0, %partial_block_capacity0)] : index, index + %key_value_head_count = config.get @qwen3_moe.attention.key_value_head_count : index + %workgroup_y0 = kernel.workgroup.id : index + %key_value_head = index.assume %workgroup_y0 [range(%workgroup_y0, 0, 63)] : index + %workitem = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %counter_scratch_bytes = index.constant 4 : offset + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %partial_max_noalias, %partial_sum_noalias, %partial_output_noalias, %completion_counter_noalias, %output_noalias = buffer.assume.noalias %partial_max, %partial_sum, %partial_output, %completion_counter, %output : buffer, buffer, buffer, buffer, buffer + %completion_counter_aligned = buffer.assume.alignment %completion_counter_noalias {minimum_alignment = 16} : buffer + %completion_counter_view = buffer.view %completion_counter_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]xi32> + %counter_scratch = buffer.alloca align(4) %counter_scratch_bytes : buffer + %counter_scratch_view = buffer.view %counter_scratch[%c0_offset] : buffer -> view<1xi32> + %workitem_is_zero = index.cmp eq, %workitem, %c0 : index + // Publish every producer's partial stores before the leader advances one + // workgroup arrival. The last arrival then acquires every partial. + kernel.barrier scope(workgroup) ordering(release) + scf.if %workitem_is_zero { + %old_counter = view.atomic.rmw %c1_i32, %completion_counter_view[%key_value_head] {ordering = acq_rel, scope = device} : i32, view<[%key_value_head_count]xi32> -> i32 + view.store %old_counter, %counter_scratch_view[%c0] : i32, view<1xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %old_counter = view.load %counter_scratch_view[%c0] : view<1xi32> -> i32 + %key_value_block_count_i32 = index.cast %producer_block_count : index to i32 + %last_block_ordinal_i32 = scalar.subi %key_value_block_count_i32, %c1_i32 : i32 + %negative_key_value_block_count_i32 = scalar.subi %c0_i32, %key_value_block_count_i32 : i32 + %is_last_partition = scalar.cmpi eq, %old_counter, %last_block_ordinal_i32 : i32 + scf.if %is_last_partition { + kernel.barrier scope(workgroup) ordering(acquire) + template.apply<@qwen3_moe.attention.decode_split.reduce_completed.cooperative>(%partial_block_capacity, %producer_block_count, %partial_max_noalias, %partial_sum_noalias, %partial_output_noalias, %output_noalias) : (index, index, buffer, buffer, buffer, buffer) + scf.if %publish_q8 { + kernel.barrier scope(workgroup) ordering(acq_rel) + kernel.barrier scope(workgroup) ordering(acq_rel) + template.apply<@qwen3_moe.attention.decode_split.pack_completed_q8>(%key_value_head, %output_noalias, %q8_output) : (index, buffer, buffer) + } + // Do not expose the reset until every final F32 and Q8 store completes. + kernel.barrier scope(workgroup) ordering(release) + scf.if %workitem_is_zero { + view.atomic.reduce %negative_key_value_block_count_i32, %completion_counter_view[%key_value_head] {ordering = release, scope = device} : i32, view<[%key_value_head_count]xi32> + } + } + template.return +} + +// Short-context export: produce and reduce in one dispatch. +kernel.def target(@qwen3_moe_decode_split_gfx11_wave64) @qwen3_moe_flash_attention_decode_split_f32_f16_wmma(%key_value_token_count: index) { + %key_value_head_count = config.get @qwen3_moe.attention.key_value_head_count : index + %key_value_token_capacity = config.get @qwen3_moe.attention.key_value_token_capacity : index + %c1 = index.constant 1 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %padded_key_value_token_capacity = index.add %key_value_token_capacity, %c63 : index + %key_value_block_count = index.div %padded_key_value_token_capacity, %c64 : index + kernel.launch.config workgroups(%key_value_block_count, %key_value_head_count, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %completion_counter: buffer, %output: buffer) { + %key_value_token_capacity = config.get @qwen3_moe.attention.key_value_token_capacity : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %key_value_token_count_in_range = index.assume %key_value_token_count [range(%key_value_token_count, 1, 2048)] : index + %bounded_key_value_token_count, %launch_key_value_token_capacity = index.assume %key_value_token_count_in_range, %key_value_token_capacity [le(%key_value_token_count_in_range, %key_value_token_capacity)] : index, index + %padded_key_value_token_capacity = index.add %launch_key_value_token_capacity, %c63 : index + %producer_block_count = index.div %padded_key_value_token_capacity, %c64 : index + %publish_q8 = scalar.constant false : i1 + template.apply<@qwen3_moe.attention.decode_split.produce_partials>(%bounded_key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output) : (index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + template.apply<@qwen3_moe.attention.decode_split.reduce_fused>(%launch_key_value_token_capacity, %producer_block_count, %producer_block_count, %publish_q8, %partial_max, %partial_sum, %partial_output, %completion_counter, %output, %partial_output) : (index, index, index, i1, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +// Short-context export that publishes the F32 attention result and the Q8_1 +// representation consumed by the following output projection. +kernel.def target(@qwen3_moe_decode_split_gfx11_wave64) @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_next_q8(%key_value_token_count: index) { + %key_value_head_count = config.get @qwen3_moe.attention.key_value_head_count : index + %key_value_token_capacity = config.get @qwen3_moe.attention.key_value_token_capacity : index + %c1 = index.constant 1 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %padded_key_value_token_capacity = index.add %key_value_token_capacity, %c63 : index + %key_value_block_count = index.div %padded_key_value_token_capacity, %c64 : index + kernel.launch.config workgroups(%key_value_block_count, %key_value_head_count, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %completion_counter: buffer, %output: buffer, %next_q8_output: buffer) { + %key_value_token_capacity = config.get @qwen3_moe.attention.key_value_token_capacity : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %key_value_token_count_in_range = index.assume %key_value_token_count [range(%key_value_token_count, 1, 2048)] : index + %bounded_key_value_token_count, %launch_key_value_token_capacity = index.assume %key_value_token_count_in_range, %key_value_token_capacity [le(%key_value_token_count_in_range, %key_value_token_capacity)] : index, index + %padded_key_value_token_capacity = index.add %launch_key_value_token_capacity, %c63 : index + %producer_block_count = index.div %padded_key_value_token_capacity, %c64 : index + %publish_q8 = scalar.constant true : i1 + template.apply<@qwen3_moe.attention.decode_split.produce_partials>(%bounded_key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output) : (index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + template.apply<@qwen3_moe.attention.decode_split.reduce_fused>(%launch_key_value_token_capacity, %producer_block_count, %producer_block_count, %publish_q8, %partial_max, %partial_sum, %partial_output, %completion_counter, %output, %next_q8_output) : (index, index, index, i1, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +// Long-context producer export: publish partials for a following parallel +// reducer without carrying short-context synchronization bindings. +kernel.def target(@qwen3_moe_decode_split_gfx11_wave64) @qwen3_moe_flash_attention_decode_split_produce_partials_f32_f16_wmma(%key_value_token_count: index) { + %key_value_head_count = config.get @qwen3_moe.attention.key_value_head_count : index + %key_value_token_capacity = config.get @qwen3_moe.attention.key_value_token_capacity : index + %c1 = index.constant 1 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %padded_key_value_token_capacity = index.add %key_value_token_capacity, %c63 : index + %key_value_block_count = index.div %padded_key_value_token_capacity, %c64 : index + kernel.launch.config workgroups(%key_value_block_count, %key_value_head_count, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer) { + %key_value_token_capacity = config.get @qwen3_moe.attention.key_value_token_capacity : index + %key_value_token_count_in_range = index.assume %key_value_token_count [range(%key_value_token_count, 1, 32768)] : index + %bounded_key_value_token_count = index.assume %key_value_token_count_in_range [le(%key_value_token_count_in_range, %key_value_token_capacity)] : index + template.apply<@qwen3_moe.attention.decode_split.produce_partials>(%bounded_key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output) : (index, buffer, buffer, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +// Long contexts have enough split-K partials that assigning the reduction to +// one last-arriving producer serializes useful work. This reducer launches one +// two-wave workgroup per query head after an execution barrier from the +// producer. The first wave computes normalization once and both waves consume +// it, avoiding the duplicate work of independent 64-channel output slices. +kernel.def target(@qwen3_moe_decode_split_gfx11_wave64) @qwen3_moe_flash_attention_decode_split_reduce_f32(%key_value_token_count: index) { + %query_head_count = config.get @qwen3_moe.attention.query_head_count : index + %c1 = index.constant 1 : index + %c128 = index.constant 128 : index + kernel.launch.config workgroups(%c1, %query_head_count, %c1) workgroup_size(%c128, %c1, %c1) : index +} launch(%key_value_token_count: index, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %output: buffer) { + %bounded_key_value_token_count = index.assume %key_value_token_count [range(%key_value_token_count, 1, 32768)] : index + %query_head_count = config.get @qwen3_moe.attention.query_head_count : index + %key_value_head_count = config.get @qwen3_moe.attention.key_value_head_count : index + %query_head0 = kernel.workgroup.id : index + %query_head = index.assume %query_head0 [range(%query_head0, 0, 63)] : index + %workitem = kernel.workitem.id : index + %subgroup = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c64 = index.constant 64 : index + %c0_offset = index.constant 0 : offset + %reduction_stage_bytes = index.constant 8 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %key_value_token_capacity = config.get @qwen3_moe.attention.key_value_token_capacity : index + %c63 = index.constant 63 : index + %padded_key_value_token_capacity = index.add %key_value_token_capacity, %c63 : index + %key_value_block_count0 = index.div %padded_key_value_token_capacity, %c64 : index + %key_value_block_count = index.assume %key_value_block_count0 [range(%key_value_block_count0, 1, 512)] : index + %query_heads_per_key_value_head0 = index.div %query_head_count, %key_value_head_count : index + %query_heads_per_key_value_head = index.assume %query_heads_per_key_value_head0 [range(%query_heads_per_key_value_head0, 1, 16)] : index + %key_value_head = index.div %query_head, %query_heads_per_key_value_head : index + %query_row = index.rem %query_head, %query_heads_per_key_value_head : index + %is_first_subgroup = index.cmp eq, %subgroup, %c0 : index + %workitem_is_zero = index.cmp eq, %workitem, %c0 : index + %output_channel = index.add %workitem, %c0 : index + %partial_max_noalias, %partial_sum_noalias, %partial_output_noalias, %output_noalias = buffer.assume.noalias %partial_max, %partial_sum, %partial_output, %output : buffer, buffer, buffer, buffer + %partial_max_aligned = buffer.assume.alignment %partial_max_noalias {minimum_alignment = 16} : buffer + %partial_sum_aligned = buffer.assume.alignment %partial_sum_noalias {minimum_alignment = 16} : buffer + %partial_output_aligned = buffer.assume.alignment %partial_output_noalias {minimum_alignment = 16} : buffer + %output_aligned = buffer.assume.alignment %output_noalias {minimum_alignment = 16} : buffer + %partial_max_view = buffer.view %partial_max_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + %partial_sum_view = buffer.view %partial_sum_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + %partial_output_view = buffer.view %partial_output_aligned[%c0_offset] : buffer -> view<[%key_value_head_count]x[%key_value_block_count]x16x128xf16> + %output_view = buffer.view %output_aligned[%c0_offset] : buffer -> view<[%query_head_count]x128xf32> + %reduction_stage = buffer.alloca align(8) %reduction_stage_bytes : buffer + %reduction_stage_view = buffer.view %reduction_stage[%c0_offset] : buffer -> view<2xf32> + %lane_maximum = scf.if %is_first_subgroup -> (f32) { + %maximum = scf.for %block = [%lane to %key_value_block_count step %c64](%running_maximum = %negative_large : f32) -> (f32) { + %block_maximum = view.load %partial_max_view[%key_value_head, %block, %query_row] : view<[%key_value_head_count]x[%key_value_block_count]x16xf32> -> f32 + %next_maximum = scalar.maxnumf %running_maximum, %block_maximum : f32 + scf.yield %next_maximum : f32 + } + scf.yield %maximum : f32 + } else { + scf.yield %negative_large : f32 + } + %subgroup_maximum = kernel.subgroup.reduce %lane_maximum : f32 + scf.if %workitem_is_zero { + view.store %subgroup_maximum, %reduction_stage_view[%c0] : f32, view<2xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %maximum = view.load %reduction_stage_view[%c0] : view<2xf32> -> f32 + %lane_sum = scf.if %is_first_subgroup -> (f32) { + %sum = scf.for %block = [%lane to %key_value_block_count step %c64](%running_sum = %c0_f32 : f32) -> (f32) { + %block_maximum = view.load %partial_max_view[%key_value_head, %block, %query_row] : view<[%key_value_head_count]x[%key_value_block_count]x16xf32> -> f32 + %delta = scalar.subf %block_maximum, %maximum : f32 + %scale = scalar.expf %delta : f32 + view.store %scale, %partial_max_view[%key_value_head, %block, %query_row] : f32, view<[%key_value_head_count]x[%key_value_block_count]x16xf32> + %block_sum = view.load %partial_sum_view[%key_value_head, %block, %query_row] : view<[%key_value_head_count]x[%key_value_block_count]x16xf32> -> f32 + %scaled_sum = scalar.mulf %block_sum, %scale : f32 + %next_sum = scalar.addf %running_sum, %scaled_sum : f32 + scf.yield %next_sum : f32 + } + scf.yield %sum : f32 + } else { + scf.yield %c0_f32 : f32 + } + %subgroup_sum = kernel.subgroup.reduce %lane_sum : f32 + scf.if %workitem_is_zero { + view.store %subgroup_sum, %reduction_stage_view[%c1] : f32, view<2xf32> + } + // Wave 0 rewrote partial_max in global memory; every wave reads it below. + kernel.barrier scope(workgroup) ordering(acq_rel) + kernel.barrier scope(workgroup) ordering(acq_rel) + %sum = view.load %reduction_stage_view[%c1] : view<2xf32> -> f32 + %unnormalized_output = scf.for %block = [%c0 to %key_value_block_count step %c1](%running_output = %c0_f32 : f32) -> (f32) unroll(%c4) schedule(interleaved) { + %scale = view.load %partial_max_view[%key_value_head, %block, %query_row] : view<[%key_value_head_count]x[%key_value_block_count]x16xf32> -> f32 + %block_output_f16 = view.load %partial_output_view[%key_value_head, %block, %query_row, %output_channel] : view<[%key_value_head_count]x[%key_value_block_count]x16x128xf16> -> f16 + %block_output = scalar.extf %block_output_f16 : f16 to f32 + %scaled_output = scalar.mulf %block_output, %scale : f32 + %next_output = scalar.addf %running_output, %scaled_output : f32 + scf.yield %next_output : f32 + } + %normalized_output = scalar.divf %unnormalized_output, %sum : f32 + view.store %normalized_output, %output_view[%query_head, %output_channel] : f32, view<[%query_head_count]x128xf32> + kernel.return +} + +check.case public @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_case { + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(1.0) : tensor<32x128xf32> + %key = check.generate.fill value(1.0) : tensor<256x4x128xf16> + %value = check.generate.fill value(2.0) : tensor<256x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<256xf16> + %partial_max = check.generate.fill value(-1.0) : tensor<4x4x16xf32> + %partial_sum = check.generate.fill value(-1.0) : tensor<4x4x16xf32> + %partial_output = check.generate.fill value(-1.0) : tensor<4x4x16x128xf16> + %completion_counter = check.generate.fill value(0) : tensor<4xi32> + %output = check.generate.fill value(-1.0) : tensor<32x128xf32> + %expected = check.generate.fill value(2.0) : tensor<32x128xf32> + kernel.launch @qwen3_moe_flash_attention_decode_split_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output, %completion_counter, %output) : [index](index, tensor<32x128xf32>, tensor<256x4x128xf16>, tensor<256x4x128xf16>, tensor<256xf16>, tensor<4x4x16xf32>, tensor<4x4x16xf32>, tensor<4x4x16x128xf16>, tensor<4xi32>, tensor<32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<32x128xf32> + check.return +} + +// Only KV row zero participates. The finite F32 iota step overflows to F16 +// negative infinity at every later row, leaving blocks one through three +// entirely masked. Each empty split must publish the online-softmax identity +// instead of evaluating -inf - -inf and contaminating the final reduction. +check.case public @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_masked_blocks_case { + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(1.0) : tensor<32x128xf32> + %key = check.generate.fill value(1.0) : tensor<256x4x128xf16> + %value = check.generate.fill value(2.0) : tensor<256x4x128xf16> + %mask = check.generate.iota offset(0.0) step(-1e+30) : tensor<256xf16> + %partial_max = check.generate.fill value(-1.0) : tensor<4x4x16xf32> + %partial_sum = check.generate.fill value(-1.0) : tensor<4x4x16xf32> + %partial_output = check.generate.fill value(-1.0) : tensor<4x4x16x128xf16> + %completion_counter = check.generate.fill value(0) : tensor<4xi32> + %output = check.generate.fill value(-1.0) : tensor<32x128xf32> + %expected = check.generate.fill value(2.0) : tensor<32x128xf32> + kernel.launch @qwen3_moe_flash_attention_decode_split_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output, %completion_counter, %output) : [index](index, tensor<32x128xf32>, tensor<256x4x128xf16>, tensor<256x4x128xf16>, tensor<256xf16>, tensor<4x4x16xf32>, tensor<4x4x16xf32>, tensor<4x4x16x128xf16>, tensor<4xi32>, tensor<32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<32x128xf32> + check.return +} + +// The production Decode-513 shape combines eight full blocks with a one-row +// tail and selects the cooperative fused reducer. Constant nonzero V keeps the +// expected result exact while all four KV heads, 32 GQA heads, completion +// counters, partial tensors, and tail guards participate. A second invocation +// reuses the partial and counter storage with a different V tensor, making the +// completion-counter reset observable. +check.case public @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_decode_513_case { + %key_value_token_count = check.literal value(513) : index + %query = check.generate.fill value(1.0) : tensor<32x128xf32> + %key = check.generate.fill value(1.0) : tensor<513x4x128xf16> + %value0 = check.generate.fill value(2.0) : tensor<513x4x128xf16> + %value1 = check.generate.fill value(3.0) : tensor<513x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<513xf16> + %partial_max = check.generate.fill value(-1.0) : tensor<4x9x16xf32> + %partial_sum = check.generate.fill value(-1.0) : tensor<4x9x16xf32> + %partial_output = check.generate.fill value(-1.0) : tensor<4x9x16x128xf16> + %completion_counter = check.generate.fill value(0) : tensor<4xi32> + %output0 = check.generate.fill value(-1.0) : tensor<32x128xf32> + %output1 = check.generate.fill value(-1.0) : tensor<32x128xf32> + %expected0 = check.generate.fill value(2.0) : tensor<32x128xf32> + %expected1 = check.generate.fill value(3.0) : tensor<32x128xf32> + kernel.launch @qwen3_moe_flash_attention_decode_split_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value0, %mask, %partial_max, %partial_sum, %partial_output, %completion_counter, %output0) : [index](index, tensor<32x128xf32>, tensor<513x4x128xf16>, tensor<513x4x128xf16>, tensor<513xf16>, tensor<4x9x16xf32>, tensor<4x9x16xf32>, tensor<4x9x16x128xf16>, tensor<4xi32>, tensor<32x128xf32>) + kernel.launch @qwen3_moe_flash_attention_decode_split_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value1, %mask, %partial_max, %partial_sum, %partial_output, %completion_counter, %output1) : [index](index, tensor<32x128xf32>, tensor<513x4x128xf16>, tensor<513x4x128xf16>, tensor<513xf16>, tensor<4x9x16xf32>, tensor<4x9x16xf32>, tensor<4x9x16x128xf16>, tensor<4xi32>, tensor<32x128xf32>) + check.expect.close actual(%output0) expected(%expected0) atol(0.001) rtol(0.001) nan(same) : tensor<32x128xf32> + check.expect.close actual(%output1) expected(%expected1) atol(0.001) rtol(0.001) nan(same) : tensor<32x128xf32> + check.return +} + +// Sixty-five KV rows force a second split containing one valid row. The mask +// selects that final row, whose iota values begin at 1024, so an omitted or +// uninitialized tail cannot accidentally satisfy the check. +check.case public @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_tail_case { + %key_value_token_count = check.literal value(65) : index + %query = check.generate.fill value(1.0) : tensor<1x128xf32> + %key = check.generate.fill value(1.0) : tensor<65x1x128xf16> + %value = check.generate.iota offset(0.0) step(0.125) : tensor<65x1x128xf16> + %mask = check.generate.iota offset(-64000.0) step(1000.0) : tensor<65xf16> + %partial_max = check.generate.fill value(-1.0) : tensor<1x2x16xf32> + %partial_sum = check.generate.fill value(-1.0) : tensor<1x2x16xf32> + %partial_output = check.generate.fill value(-1.0) : tensor<1x2x16x128xf16> + %completion_counter = check.generate.fill value(0) : tensor<1xi32> + %output = check.generate.fill value(-1.0) : tensor<1x128xf32> + %expected = check.generate.iota offset(1024.0) step(0.125) : tensor<1x128xf32> + kernel.launch @qwen3_moe_flash_attention_decode_split_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output, %completion_counter, %output) : [index](index, tensor<1x128xf32>, tensor<65x1x128xf16>, tensor<65x1x128xf16>, tensor<65xf16>, tensor<1x2x16xf32>, tensor<1x2x16xf32>, tensor<1x2x16x128xf16>, tensor<1xi32>, tensor<1x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<1x128xf32> + check.return +} + +// Thirty-two split-K blocks exercise the separate parallel reducer and the +// execution barrier between its producer and consumer dispatches. +check.case public @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_long_case { + %key_value_token_count = check.literal value(2048) : index + %query = check.generate.fill value(0.0) : tensor<32x128xf32> + %key = check.generate.fill value(0.0) : tensor<2048x4x128xf16> + %value = check.generate.fill value(0.0) : tensor<2048x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<2048xf16> + %partial_max = check.generate.fill value(-1.0) : tensor<4x32x16xf32> + %partial_sum = check.generate.fill value(-1.0) : tensor<4x32x16xf32> + %partial_output = check.generate.fill value(-1.0) : tensor<4x32x16x128xf16> + %output = check.generate.fill value(1.0) : tensor<32x128xf32> + %expected = check.generate.fill value(0.0) : tensor<32x128xf32> + kernel.launch @qwen3_moe_flash_attention_decode_split_produce_partials_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output) : [index](index, tensor<32x128xf32>, tensor<2048x4x128xf16>, tensor<2048x4x128xf16>, tensor<2048xf16>, tensor<4x32x16xf32>, tensor<4x32x16xf32>, tensor<4x32x16x128xf16>) + kernel.launch @qwen3_moe_flash_attention_decode_split_reduce_f32[%key_value_token_count](%key_value_token_count, %partial_max, %partial_sum, %partial_output, %output) : [index](index, tensor<4x32x16xf32>, tensor<4x32x16xf32>, tensor<4x32x16x128xf16>, tensor<32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<32x128xf32> + check.return +} + +check.case public @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_benchmark_case { + %key_value_token_count = check.param.choice values([64, 65, 128, 256, 512, 513, 768, 1024, 1280, 2048]) name("key_value_token_count") : index + %query = check.generate.fill value(0.0) : tensor<32x128xf32> + %key = check.generate.fill value(0.0) : tensor<[%key_value_token_count]x4x128xf16> + %value = check.generate.fill value(0.0) : tensor<[%key_value_token_count]x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<[%key_value_token_count]xf16> + // Reserve the bounded scratch capacity once; each specialization addresses + // only ceildiv(key_value_token_count, 64) blocks. + %partial_max = check.generate.fill value(-1.0) : tensor<4x512x16xf32> + %partial_sum = check.generate.fill value(-1.0) : tensor<4x512x16xf32> + %partial_output = check.generate.fill value(-1.0) : tensor<4x512x16x128xf16> + %completion_counter = check.generate.fill value(0) : tensor<4xi32> + %output = check.generate.fill value(1.0) : tensor<32x128xf32> + %expected = check.generate.fill value(0.0) : tensor<32x128xf32> + kernel.launch @qwen3_moe_flash_attention_decode_split_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output, %completion_counter, %output) : [index](index, tensor<32x128xf32>, tensor<[%key_value_token_count]x4x128xf16>, tensor<[%key_value_token_count]x4x128xf16>, tensor<[%key_value_token_count]xf16>, tensor<4x512x16xf32>, tensor<4x512x16xf32>, tensor<4x512x16x128xf16>, tensor<4xi32>, tensor<32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<32x128xf32> + check.return +} + +// Long-context execution is one reusable producer/reducer command buffer. The +// harness records an explicit dispatch execution barrier between these calls +// and profiles their complete device-side span as one semantic operation. +check.case public @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_long_benchmark_case { + %key_value_token_count = check.param.choice values([2048, 32768]) name("key_value_token_count") : index + %query = check.generate.fill value(0.0) : tensor<32x128xf32> + %key = check.generate.fill value(0.0) : tensor<[%key_value_token_count]x4x128xf16> + %value = check.generate.fill value(0.0) : tensor<[%key_value_token_count]x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<[%key_value_token_count]xf16> + %partial_max = check.generate.fill value(-1.0) : tensor<4x512x16xf32> + %partial_sum = check.generate.fill value(-1.0) : tensor<4x512x16xf32> + %partial_output = check.generate.fill value(-1.0) : tensor<4x512x16x128xf16> + %output = check.generate.fill value(1.0) : tensor<32x128xf32> + %expected = check.generate.fill value(0.0) : tensor<32x128xf32> + kernel.launch @qwen3_moe_flash_attention_decode_split_produce_partials_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value, %mask, %partial_max, %partial_sum, %partial_output) : [index](index, tensor<32x128xf32>, tensor<[%key_value_token_count]x4x128xf16>, tensor<[%key_value_token_count]x4x128xf16>, tensor<[%key_value_token_count]xf16>, tensor<4x512x16xf32>, tensor<4x512x16xf32>, tensor<4x512x16x128xf16>) + kernel.launch @qwen3_moe_flash_attention_decode_split_reduce_f32[%key_value_token_count](%key_value_token_count, %partial_max, %partial_sum, %partial_output, %output) : [index](index, tensor<4x512x16xf32>, tensor<4x512x16xf32>, tensor<4x512x16x128xf16>, tensor<32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<32x128xf32> + check.return +} + +check.benchmark<@qwen3_moe_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_decode_64 {key_value_token_count = 64} + +check.benchmark<@qwen3_moe_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_decode_65 {key_value_token_count = 65} + +check.benchmark<@qwen3_moe_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_decode_128 {key_value_token_count = 128} + +check.benchmark<@qwen3_moe_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_decode_256 {key_value_token_count = 256} + +check.benchmark<@qwen3_moe_flash_attention_decode_split_f32_f16_wmma_masked_blocks_case> @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_decode_256_masked_blocks + +check.benchmark<@qwen3_moe_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_decode_512 {key_value_token_count = 512} + +check.benchmark<@qwen3_moe_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_decode_513 {key_value_token_count = 513} + +check.benchmark<@qwen3_moe_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_decode_768 {key_value_token_count = 768} + +check.benchmark<@qwen3_moe_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_decode_1024 {key_value_token_count = 1024} + +check.benchmark<@qwen3_moe_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_decode_1280 {key_value_token_count = 1280} + +check.benchmark<@qwen3_moe_flash_attention_decode_split_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_decode_2048_fused {key_value_token_count = 2048} + +check.benchmark<@qwen3_moe_flash_attention_decode_split_f32_f16_wmma_long_benchmark_case> @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_decode_2048 {key_value_token_count = 2048} + +check.benchmark<@qwen3_moe_flash_attention_decode_split_f32_f16_wmma_long_benchmark_case> @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_decode_32768 {key_value_token_count = 32768} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/flash_attention_decode_split_next_q8_test.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/flash_attention_decode_split_next_q8_test.loom new file mode 100644 index 000000000000..58feb9fbba73 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/flash_attention_decode_split_next_q8_test.loom @@ -0,0 +1,117 @@ +// Test-only differentials for fused decode-attention Q8 publication. The +// production kernels remain the actual providers; this module contributes only +// reference packing, masked metadata setup, and comparison cases. +template.decl @ggml.quantize_q8_1_x4.group_body(%arg0: i1, %arg1: index, %arg2: index, %arg3: buffer, %arg4: buffer) + +template.decl @qwen3_moe.attention.decode_split.pack_completed_q8(%arg0: index, %arg1: buffer, %arg2: buffer) + +target.decl @qwen3_moe_decode_split_gfx11_wave64 + +kernel.decl target(@qwen3_moe_decode_split_gfx11_wave64) @qwen3_moe_flash_attention_decode_split_f32_f16_wmma(%key_value_token_count$8: index) launch(%key_value_token_count$9: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %completion_counter: buffer, %output: buffer) + +kernel.decl target(@qwen3_moe_decode_split_gfx11_wave64) @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_next_q8(%key_value_token_count$19: index) launch(%key_value_token_count$20: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %partial_max: buffer, %partial_sum: buffer, %partial_output: buffer, %completion_counter: buffer, %output: buffer, %next_q8_output: buffer) + +// Publishes the canonical reference Q8_1 x4 layout for one 4096-element row. +kernel.def @qwen3_moe_flash_attention_decode_split_quantize_reference_4096() { + %c1 = index.constant 1 : index + %c32 = index.constant 32 : index + kernel.launch.config workgroups(%c32, %c1, %c1) workgroup_size(%c32, %c1, %c1) : index +} launch(%input: buffer, %output: buffer) { + %publish_output = scalar.constant true : i1 + %group_count = index.constant 32 : index + %group = kernel.workgroup.id : index + template.apply<@ggml.quantize_q8_1_x4.group_body>(%publish_output, %group_count, %group, %input, %output) : (i1, index, index, buffer, buffer) + kernel.return +} + +// Isolates the new publication phase for access-sanitized coverage without +// inflating the complete fused attention kernel beyond a short branch's range. +kernel.def target(@qwen3_moe_decode_split_gfx11_wave64) @qwen3_moe_flash_attention_decode_split_pack_completed_q8_test() { + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%c1, %c4, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%input: buffer, %output: buffer) { + %key_value_head = kernel.workgroup.id : index + template.apply<@qwen3_moe.attention.decode_split.pack_completed_q8>(%key_value_head, %input, %output) : (index, buffer, buffer) + kernel.return +} + +// Models the request-owned Decode-513 mask inside its 576-row capacity class. +kernel.def @qwen3_moe_flash_attention_decode_split_mask_513_of_576() { + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + kernel.launch.config workgroups(%c1, %c1, %c1) workgroup_size(%c64, %c1, %c1) : index +} launch(%mask: buffer) { + %workitem = kernel.workitem.id : index + %c63 = index.constant 63 : index + %c513 = index.constant 513 : index + %c0 = index.constant 0 : index + %c0_offset = index.constant 0 : offset + %negative_large = scalar.constant -1e+30 : f32 + %negative_infinity = scalar.fptrunc %negative_large : f32 to f16 + %valid = index.cmp ult, %workitem, %c63 : index + %row0 = index.add %c513, %workitem : index + %row = scf.select %valid, %row0, %c0 : index + %mask_view = buffer.view %mask[%c0_offset] : buffer -> view<576xf16> + scf.if %valid { + view.store %negative_infinity, %mask_view[%row] : f16, view<576xf16> + } + kernel.return +} + +// The fused producer must match the ordinary attention export followed by the +// canonical GGML packer for every F32 value and every packed byte. Two value +// tensors reuse both paths' partials and completion counters, making counter +// reset and stale-publication failures observable. +check.case public @qwen3_moe_flash_attention_decode_split_next_q8_capacity_576_differential_case { + %key_value_token_count = check.literal value(576) : index + %query_seed = check.param.seed base(5858425849413783864) count(1) : i64 + %key_seed = check.param.seed base(5858425849313120568) count(1) : i64 + %value_seed0 = check.param.seed base(5858425849497669936) count(1) : i64 + %value_seed1 = check.param.seed base(5858425849497669937) count(1) : i64 + %query = check.generate.random.uniform seed(%query_seed) range(-1.0 to 1.0) : tensor<32x128xf32> + %key = check.generate.random.uniform seed(%key_seed) range(-1.0 to 1.0) : tensor<576x4x128xf16> + %value0 = check.generate.random.uniform seed(%value_seed0) range(-1.0 to 1.0) : tensor<576x4x128xf16> + %value1 = check.generate.random.uniform seed(%value_seed1) range(-1.0 to 1.0) : tensor<576x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<576xf16> + %reference_partial_max = check.generate.fill value(-1.0) : tensor<4x9x16xf32> + %reference_partial_sum = check.generate.fill value(-1.0) : tensor<4x9x16xf32> + %reference_partial_output = check.generate.fill value(-1.0) : tensor<4x9x16x128xf16> + %reference_counter = check.generate.fill value(0) : tensor<4xi32> + %actual_partial_max = check.generate.fill value(-1.0) : tensor<4x9x16xf32> + %actual_partial_sum = check.generate.fill value(-1.0) : tensor<4x9x16xf32> + %actual_partial_output = check.generate.fill value(-1.0) : tensor<4x9x16x128xf16> + %actual_counter = check.generate.fill value(0) : tensor<4xi32> + %reference_output0 = check.generate.fill value(-1.0) : tensor<32x128xf32> + %reference_output1 = check.generate.fill value(-1.0) : tensor<32x128xf32> + %actual_output0 = check.generate.fill value(-2.0) : tensor<32x128xf32> + %actual_output1 = check.generate.fill value(-2.0) : tensor<32x128xf32> + %reference_q8_0 = check.generate.fill value(0) : tensor<4608xi8> + %reference_q8_1 = check.generate.fill value(0) : tensor<4608xi8> + %actual_q8_0 = check.generate.fill value(1) : tensor<4608xi8> + %actual_q8_1 = check.generate.fill value(1) : tensor<4608xi8> + kernel.launch @qwen3_moe_flash_attention_decode_split_mask_513_of_576(%mask) : (tensor<576xf16>) + kernel.launch @qwen3_moe_flash_attention_decode_split_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value0, %mask, %reference_partial_max, %reference_partial_sum, %reference_partial_output, %reference_counter, %reference_output0) : [index](index, tensor<32x128xf32>, tensor<576x4x128xf16>, tensor<576x4x128xf16>, tensor<576xf16>, tensor<4x9x16xf32>, tensor<4x9x16xf32>, tensor<4x9x16x128xf16>, tensor<4xi32>, tensor<32x128xf32>) + kernel.launch @qwen3_moe_flash_attention_decode_split_quantize_reference_4096(%reference_output0, %reference_q8_0) : (tensor<32x128xf32>, tensor<4608xi8>) + kernel.launch @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_next_q8[%key_value_token_count](%key_value_token_count, %query, %key, %value0, %mask, %actual_partial_max, %actual_partial_sum, %actual_partial_output, %actual_counter, %actual_output0, %actual_q8_0) : [index](index, tensor<32x128xf32>, tensor<576x4x128xf16>, tensor<576x4x128xf16>, tensor<576xf16>, tensor<4x9x16xf32>, tensor<4x9x16xf32>, tensor<4x9x16x128xf16>, tensor<4xi32>, tensor<32x128xf32>, tensor<4608xi8>) + kernel.launch @qwen3_moe_flash_attention_decode_split_f32_f16_wmma[%key_value_token_count](%key_value_token_count, %query, %key, %value1, %mask, %reference_partial_max, %reference_partial_sum, %reference_partial_output, %reference_counter, %reference_output1) : [index](index, tensor<32x128xf32>, tensor<576x4x128xf16>, tensor<576x4x128xf16>, tensor<576xf16>, tensor<4x9x16xf32>, tensor<4x9x16xf32>, tensor<4x9x16x128xf16>, tensor<4xi32>, tensor<32x128xf32>) + kernel.launch @qwen3_moe_flash_attention_decode_split_quantize_reference_4096(%reference_output1, %reference_q8_1) : (tensor<32x128xf32>, tensor<4608xi8>) + kernel.launch @qwen3_moe_flash_attention_decode_split_f32_f16_wmma_next_q8[%key_value_token_count](%key_value_token_count, %query, %key, %value1, %mask, %actual_partial_max, %actual_partial_sum, %actual_partial_output, %actual_counter, %actual_output1, %actual_q8_1) : [index](index, tensor<32x128xf32>, tensor<576x4x128xf16>, tensor<576x4x128xf16>, tensor<576xf16>, tensor<4x9x16xf32>, tensor<4x9x16xf32>, tensor<4x9x16x128xf16>, tensor<4xi32>, tensor<32x128xf32>, tensor<4608xi8>) + check.expect.equal actual(%actual_output0) expected(%reference_output0) : tensor<32x128xf32> + check.expect.equal actual(%actual_q8_0) expected(%reference_q8_0) : tensor<4608xi8> + check.expect.equal actual(%actual_output1) expected(%reference_output1) : tensor<32x128xf32> + check.expect.equal actual(%actual_q8_1) expected(%reference_q8_1) : tensor<4608xi8> + check.return +} + +check.case public @qwen3_moe_flash_attention_decode_split_pack_completed_q8_differential_case { + %input_seed = check.param.seed base(5858425849396089144) count(1) : i64 + %input = check.generate.random.uniform seed(%input_seed) range(-1.0 to 1.0) : tensor<32x128xf32> + %reference = check.generate.fill value(0) : tensor<4608xi8> + %actual = check.generate.fill value(1) : tensor<4608xi8> + kernel.launch @qwen3_moe_flash_attention_decode_split_quantize_reference_4096(%input, %reference) : (tensor<32x128xf32>, tensor<4608xi8>) + kernel.launch @qwen3_moe_flash_attention_decode_split_pack_completed_q8_test(%input, %actual) : (tensor<32x128xf32>, tensor<4608xi8>) + check.expect.equal actual(%actual) expected(%reference) : tensor<4608xi8> + check.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/flash_attention_f32_f16_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/flash_attention_f32_f16_wmma.loom new file mode 100644 index 000000000000..cec99bdd30bd --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/flash_attention_f32_f16_wmma.loom @@ -0,0 +1,870 @@ +// Qwen3 MoE grouped-query FlashAttention for the F32-query/F16-cache path. +// +// One four-wave workgroup computes 16 query rows for one query head against +// 64 KV rows at a time. The ownership changes mirror the cooperative-matrix +// schedule used by llama.cpp's Vulkan CM1 kernel: +// +// 1. All workitems stage a scaled 16x128 F16 query tile. +// 2. Each wave computes one 16x16 QK score slice. +// 3. Scores cross LDS so each wave can normalize four complete query rows. +// 4. F16 probabilities cross LDS for four P*V WMMA steps. +// 5. Each active lane retains one four-channel F16 packet for each of its +// four query rows across subsequent 64-row KV blocks. +// +// K and V remain in the row-major llama.cpp cache layout +// [KV token][KV head][128]. Their aligned F16 fragments load directly from +// global memory; there is no expanded or repacked persistent allocation. QK +// and the online-softmax statistics remain F32, while the P*V accumulation and +// carried output match the Vulkan oracle's F16 policy. +amdgpu.target @qwen3_moe_attention_gfx11_wave64 {subgroup_size = 64} + +config.decl @qwen3_moe.attention.query_head_count : %value: index where [range(%value, 1, 64)] + +config.decl @qwen3_moe.attention.key_value_head_count : %value: index where [range(%value, 1, 64)] + +// Bounds the context copied by the test-only row-extraction kernel below. +config.decl @qwen3_moe.attention.test.context_capacity : %value: index where [range(%value, 1, 32768)] + +kernel.def target(@qwen3_moe_attention_gfx11_wave64) @qwen3_moe_flash_attention_f32_f16_wmma(%query_token_count: index, %key_value_token_count: index) { + %query_head_count = config.get @qwen3_moe.attention.query_head_count : index + %c1 = index.constant 1 : index + %c15 = index.constant 15 : index + %c16 = index.constant 16 : index + %c256 = index.constant 256 : index + %padded_query_token_count = index.add %query_token_count, %c15 : index + %query_tile_count = index.div %padded_query_token_count, %c16 : index + kernel.launch.config workgroups(%query_tile_count, %query_head_count, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%query_token_count: index, %key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %output: buffer) where [range(%query_token_count, 1, 2048)] { + %bounded_key_value_token_count = index.assume %key_value_token_count [range(%key_value_token_count, 1, 32768)] : index + %query_head_count = config.get @qwen3_moe.attention.query_head_count : index + %key_value_head_count = config.get @qwen3_moe.attention.key_value_head_count : index + %query_tile = kernel.workgroup.id : index + %query_head0 = kernel.workgroup.id : index + %query_head, %launch_query_head_count = index.assume %query_head0, %query_head_count [lt(%query_head0, %query_head_count)] : index, index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 3)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c15 = index.constant 15 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %query_stage_bytes = index.constant 4352 : offset + %score_stage_bytes = index.constant 6144 : offset + %probability_stage_bytes = index.constant 3072 : offset + %product_stage_bytes = index.constant 2048 : offset + %tail_key_value_stage_capacity = index.constant 8192 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %head_size_f32 = scalar.constant 128.0 : f32 + %attention_scale = scalar.rsqrtf %head_size_f32 : f32 + %c0_f16 = scalar.constant 0.0 : f16 + %c1_f32 = scalar.constant 1.0 : f32 + %output_zero0 = vector.constant 0.0 : vector<4xf16> + %output_zero1 = vector.constant 0.0 : vector<4xf16> + %output_zero2 = vector.constant 0.0 : vector<4xf16> + %output_zero3 = vector.constant 0.0 : vector<4xf16> + %c0_f16x8 = vector.constant 0.0 : vector<8xf16> + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %negative_f32x4 = vector.constant -1e+30 : vector<4xf32> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %query_heads_per_key_value_head = index.div %query_head_count, %key_value_head_count : index + %key_value_head = index.div %query_head, %query_heads_per_key_value_head : index + %key_value_width = index.mul %key_value_head_count, %c128 : index + %key_value_head_base = index.mul %key_value_head, %c128 : index + %query_origin = index.mul %query_tile, %c16 : index + %full_key_value_block_count = index.div %bounded_key_value_token_count, %c64 : index + %full_key_value_token_count0 = index.mul %full_key_value_block_count, %c64 : index + %full_key_value_token_count = index.assume %full_key_value_token_count0 [range(%full_key_value_token_count0, 0, 32768), mul(%full_key_value_token_count0, 64)] : index + %tail_key_value_token_count = index.sub %bounded_key_value_token_count, %full_key_value_token_count : index + %has_key_value_tail = index.cmp ne, %tail_key_value_token_count, %c0 : index + %tail_key_value_stage_bytes = scf.select %has_key_value_tail, %tail_key_value_stage_capacity, %c0_offset : offset + %subgroup_score_column = index.mul %subgroup, %c16 : index + %subgroup_query_row = index.mul %subgroup, %c4 : index + %query_row0 = index.add %subgroup_query_row, %c0 : index + %query_row1 = index.add %subgroup_query_row, %c1 : index + %query_row2 = index.add %subgroup_query_row, %c2 : index + %query_row3 = index.add %subgroup_query_row, %c3 : index + %query_token0 = index.add %query_origin, %query_row0 : index + %query_token1 = index.add %query_origin, %query_row1 : index + %query_token2 = index.add %query_origin, %query_row2 : index + %query_token3 = index.add %query_origin, %query_row3 : index + %query_valid0 = index.cmp ult, %query_token0, %query_token_count : index + %query_valid1 = index.cmp ult, %query_token1, %query_token_count : index + %query_valid2 = index.cmp ult, %query_token2, %query_token_count : index + %query_valid3 = index.cmp ult, %query_token3, %query_token_count : index + %query_valid = vector.from_elements %query_valid0, %query_valid1, %query_valid2, %query_valid3 : vector<4xi1> + %subgroup_product_channel = index.mul %subgroup, %c16 : index + %lane_output_tile = index.div %lane, %c16 : index + %lane_product_channel0 = index.rem %lane, %c16 : index + %lane_product_channel = index.mul %lane_product_channel0, %c4 : index + %lane_output_channel = index.mul %lane, %c4 : index + %lane_has_output = index.cmp ult, %lane, %c32 : index + // The padded LDS rows mirror the Vulkan oracle's ownership changes. Q uses + // eight spare F16 columns after its 128 channels. Score and probability + // transpose to key-major rows with eight spare columns after 16 queries. + // These strides avoid the bank pattern produced by dense transposed rows. + %query_transposed_layout = encoding.layout.strided [1, 136] : encoding + %probability_transposed_layout = encoding.layout.strided [1, 24] : encoding + %query_noalias, %key_noalias, %value_noalias, %mask_noalias, %output_noalias = buffer.assume.noalias %query, %key, %value, %mask, %output : buffer, buffer, buffer, buffer, buffer + %query_aligned = buffer.assume.alignment %query_noalias {minimum_alignment = 16} : buffer + %key_aligned = buffer.assume.alignment %key_noalias {minimum_alignment = 16} : buffer + %value_aligned = buffer.assume.alignment %value_noalias {minimum_alignment = 16} : buffer + %mask_aligned = buffer.assume.alignment %mask_noalias {minimum_alignment = 16} : buffer + %output_aligned = buffer.assume.alignment %output_noalias {minimum_alignment = 16} : buffer + %query_view = buffer.view %query_aligned[%c0_offset] : buffer -> view<[%query_token_count]x[%query_head_count]x128xf32> + %key_view = buffer.view %key_aligned[%c0_offset] : buffer -> view<[%bounded_key_value_token_count]x[%key_value_width]xf16> + %value_view = buffer.view %value_aligned[%c0_offset] : buffer -> view<[%bounded_key_value_token_count]x[%key_value_width]xf16> + %mask_view = buffer.view %mask_aligned[%c0_offset] : buffer -> view<[%query_token_count]x[%bounded_key_value_token_count]xf16> + %output_view = buffer.view %output_aligned[%c0_offset] : buffer -> view<[%query_token_count]x[%query_head_count]x128xf32> + %query_stage = buffer.alloca align(16) %query_stage_bytes : buffer + %score_stage = buffer.alloca align(16) %score_stage_bytes : buffer + %probability_stage = buffer.alloca align(16) %probability_stage_bytes : buffer + %product_stage = buffer.alloca align(16) %product_stage_bytes : buffer + %tail_key_value_stage = buffer.alloca align(16) %tail_key_value_stage_bytes : buffer + %query_stage_view = buffer.view %query_stage[%c0_offset] : buffer -> view<16x136xf16> + %query_transposed_view = buffer.view %query_stage[%c0_offset] : buffer -> view<128x16xf16, %query_transposed_layout> + %score_stage_view = buffer.view %score_stage[%c0_offset] : buffer -> view<64x24xf32> + %mask_summary_view = buffer.view %score_stage[%c0_offset] : buffer -> view<4xf32> + %probability_stage_view = buffer.view %probability_stage[%c0_offset] : buffer -> view<16x64xf16, %probability_transposed_layout> + %product_stage_view = buffer.view %product_stage[%c0_offset] : buffer -> view<16x64xf16> + %tail_key_value_stage_view = buffer.view %tail_key_value_stage[%c0_offset] : buffer -> view<32x128xf16> + // Scale and truncate Q exactly once. The Vulkan reference does this before + // entering its KV loop, making QK a native F16 WMMA while retaining F32 + // accumulation. + scf.for %load_iteration = [%c0 to %c8 step %c1] unroll { + %linear = index.madd %load_iteration, %c256, %workitem : index + %local_query_row = index.div %linear, %c128 : index + %query_channel = index.rem %linear, %c128 : index + %query_token = index.add %query_origin, %local_query_row : index + %query_valid_load = index.cmp ult, %query_token, %query_token_count : index + %query_value = scf.if %query_valid_load -> (f16) { + %loaded = view.load %query_view[%query_token, %query_head, %query_channel] : view<[%query_token_count]x[%query_head_count]x128xf32> -> f32 + %scaled = scalar.mulf %loaded, %attention_scale : f32 + %truncated = scalar.fptrunc %scaled : f32 to f16 + scf.yield %truncated : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %query_value, %query_stage_view[%local_query_row, %query_channel] : f16, view<16x136xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %full_max, %full_sum, %full_output0, %full_output1, %full_output2, %full_output3 = scf.for %key_origin = [%c0 to %full_key_value_token_count step %c64](%current_max = %negative_f32x4 : vector<4xf32>, %current_sum = %c0_f32x4 : vector<4xf32>, %current_output0 = %output_zero0 : vector<4xf16>, %current_output1 = %output_zero1 : vector<4xf16>, %current_output2 = %output_zero2 : vector<4xf16>, %current_output3 = %output_zero3 : vector<4xf16>) -> (vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + // Cache the complete 16x64 mask tile before QK. Causal masks leave future + // KV blocks entirely at negative infinity; reducing the cached tile lets + // the workgroup skip QK, softmax, and P*V for those blocks. This mirrors + // the Vulkan oracle without requiring its auxiliary compact mask buffer. + %key_token0 = index.add %key_origin, %lane : index + %key_token = index.assume %key_token0 [lt(%key_token0, %bounded_key_value_token_count)] : index + %mask0 = scf.if %query_valid0 -> (f16) { + %value = view.load %mask_view[%query_token0, %key_token] : view<[%query_token_count]x[%bounded_key_value_token_count]xf16> -> f16 + scf.yield %value : f16 + } else { + scf.yield %c0_f16 : f16 + } + %mask1 = scf.if %query_valid1 -> (f16) { + %value = view.load %mask_view[%query_token1, %key_token] : view<[%query_token_count]x[%bounded_key_value_token_count]xf16> -> f16 + scf.yield %value : f16 + } else { + scf.yield %c0_f16 : f16 + } + %mask2 = scf.if %query_valid2 -> (f16) { + %value = view.load %mask_view[%query_token2, %key_token] : view<[%query_token_count]x[%bounded_key_value_token_count]xf16> -> f16 + scf.yield %value : f16 + } else { + scf.yield %c0_f16 : f16 + } + %mask3 = scf.if %query_valid3 -> (f16) { + %value = view.load %mask_view[%query_token3, %key_token] : view<[%query_token_count]x[%bounded_key_value_token_count]xf16> -> f16 + scf.yield %value : f16 + } else { + scf.yield %c0_f16 : f16 + } + %mask_f16 = vector.from_elements %mask0, %mask1, %mask2, %mask3 : vector<4xf16> + %mask_summary_f32 = vector.extf %mask_f16 : vector<4xf16> to vector<4xf32> + %effective_mask = vector.select %query_valid, %mask_summary_f32, %negative_f32x4 : vector<4xf32> + %lane_mask_maximum = vector.reduce %effective_mask, %negative_large : vector<4xf32>, f32 + %subgroup_mask_maximum = kernel.subgroup.reduce %lane_mask_maximum : f32 + %is_subgroup_leader = index.cmp eq, %lane, %c0 : index + scf.if %is_subgroup_leader { + view.store %subgroup_mask_maximum, %mask_summary_view[%subgroup] : f32, view<4xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %subgroup_mask_maxima = vector.load %mask_summary_view[%c0] : view<4xf32> -> vector<4xf32> + %workgroup_mask_maximum = vector.reduce %subgroup_mask_maxima, %negative_large : vector<4xf32>, f32 + %block_has_attention = scalar.cmpf ogt, %workgroup_mask_maximum, %negative_large : f32 + %next_block_max, %next_block_sum, %next_block_output0, %next_block_output1, %next_block_output2, %next_block_output3 = scf.if %block_has_attention -> (vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + // Four independent wave-level WMMAs produce a 16x64 score tile. + %score_key_origin0 = index.add %key_origin, %subgroup_score_column : index + %last_full_key_tile_start = index.sub %bounded_key_value_token_count, %c15 : index + %score_key_origin = index.assume %score_key_origin0 [lt(%score_key_origin0, %last_full_key_tile_start)] : index + %score_init_values = vector.constant 0.0 : vector<4xf32> + %score_init = vector.fragment %score_init_values shape [%m, %n] : vector<4xf32> + %score_fragment = scf.for %head_tile = [%c0 to %c128 step %c16](%score_accumulator = %score_init : vector<4xf32>) -> (vector<4xf32>) unroll schedule(recurrence) { + %key_channel = index.add %key_value_head_base, %head_tile : index + %key_fragment = vector.fragment.load %key_view[%score_key_origin, %key_channel] shape [%m, %k] : view<[%bounded_key_value_token_count]x[%key_value_width]xf16> -> vector<16xf16> + %query_fragment = vector.fragment.load %query_transposed_view[%head_tile, %c0] shape [%k, %n] : view<128x16xf16, %query_transposed_layout> -> vector<16xf16> + %next_score_accumulator = vector.mma %key_fragment, %query_fragment, %score_accumulator : vector<16xf16>, vector<16xf16>, vector<4xf32> + scf.yield %next_score_accumulator : vector<4xf32> + } + vector.fragment.store %score_fragment, %score_stage_view[%subgroup_score_column, %c0] shape [%m, %n] : vector<4xf32>, view<64x24xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + // LDS transposes ownership from one 16-column score slice per wave to + // four complete query rows per wave. Every lane contributes one key + // column to each of those rows. + %raw_score0 = view.load %score_stage_view[%lane, %query_row0] : view<64x24xf32> -> f32 + %raw_score1 = view.load %score_stage_view[%lane, %query_row1] : view<64x24xf32> -> f32 + %raw_score2 = view.load %score_stage_view[%lane, %query_row2] : view<64x24xf32> -> f32 + %raw_score3 = view.load %score_stage_view[%lane, %query_row3] : view<64x24xf32> -> f32 + %mask_f32 = vector.extf %mask_f16 : vector<4xf16> to vector<4xf32> + %raw_scores = vector.from_elements %raw_score0, %raw_score1, %raw_score2, %raw_score3 : vector<4xf32> + %masked_scores0 = vector.addf %raw_scores, %mask_f32 : vector<4xf32> + %masked_scores = vector.select %query_valid, %masked_scores0, %negative_f32x4 : vector<4xf32> + %block_max = kernel.subgroup.reduce %masked_scores : vector<4xf32> + %next_max = vector.maxnumf %current_max, %block_max : vector<4xf32> + %score_delta = vector.subf %masked_scores, %next_max : vector<4xf32> + %raw_probability = vector.expf %score_delta : vector<4xf32> + %probability = vector.select %query_valid, %raw_probability, %c0_f32x4 : vector<4xf32> + %block_sum = kernel.subgroup.reduce %probability : vector<4xf32> + %old_delta = vector.subf %current_max, %next_max : vector<4xf32> + %old_scale = vector.expf %old_delta : vector<4xf32> + %scaled_current_sum = vector.mulf %current_sum, %old_scale : vector<4xf32> + %next_sum = vector.addf %scaled_current_sum, %block_sum : vector<4xf32> + %probability_f16 = vector.fptrunc %probability : vector<4xf32> to vector<4xf16> + %probability0 = vector.extract %probability_f16[0] : vector<4xf16> -> f16 + %probability1 = vector.extract %probability_f16[1] : vector<4xf16> -> f16 + %probability2 = vector.extract %probability_f16[2] : vector<4xf16> -> f16 + %probability3 = vector.extract %probability_f16[3] : vector<4xf16> -> f16 + view.store %probability0, %probability_stage_view[%query_row0, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability1, %probability_stage_view[%query_row1, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability2, %probability_stage_view[%query_row2, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability3, %probability_stage_view[%query_row3, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + kernel.barrier scope(workgroup) ordering(acq_rel) + // Match the Vulkan CM1 ownership schedule by computing two sequential + // 64-channel output tiles. This keeps one P*V accumulator live per wave + // and reuses a 2 KiB exchange tile instead of retaining both halves. + %old_scale_f16 = vector.fptrunc %old_scale : vector<4xf32> to vector<4xf16> + %old_scale0_scalar = vector.extract %old_scale_f16[0] : vector<4xf16> -> f16 + %old_scale1_scalar = vector.extract %old_scale_f16[1] : vector<4xf16> -> f16 + %old_scale2_scalar = vector.extract %old_scale_f16[2] : vector<4xf16> -> f16 + %old_scale3_scalar = vector.extract %old_scale_f16[3] : vector<4xf16> -> f16 + %old_scale0 = vector.splat %old_scale0_scalar : vector<4xf16> + %old_scale1 = vector.splat %old_scale1_scalar : vector<4xf16> + %old_scale2 = vector.splat %old_scale2_scalar : vector<4xf16> + %old_scale3 = vector.splat %old_scale3_scalar : vector<4xf16> + %scaled_current_output0 = vector.mulf %current_output0, %old_scale0 : vector<4xf16> + %scaled_current_output1 = vector.mulf %current_output1, %old_scale1 : vector<4xf16> + %scaled_current_output2 = vector.mulf %current_output2, %old_scale2 : vector<4xf16> + %scaled_current_output3 = vector.mulf %current_output3, %old_scale3 : vector<4xf16> + %next_output0, %next_output1, %next_output2, %next_output3 = scf.for %output_tile = [%c0 to %c2 step %c1](%tile_output0 = %scaled_current_output0 : vector<4xf16>, %tile_output1 = %scaled_current_output1 : vector<4xf16>, %tile_output2 = %scaled_current_output2 : vector<4xf16>, %tile_output3 = %scaled_current_output3 : vector<4xf16>) -> (vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) unroll { + %output_tile_channel = index.mul %output_tile, %c64 : index + %value_channel0 = index.add %key_value_head_base, %output_tile_channel : index + %value_channel = index.add %value_channel0, %subgroup_product_channel : index + %product_init = vector.fragment %c0_f16x8 shape [%m, %n] : vector<8xf16> + %product_fragment = scf.for %key_tile = [%c0 to %c64 step %c16](%product_accumulator = %product_init : vector<8xf16>) -> (vector<8xf16>) unroll schedule(recurrence) { + %value_token0 = index.add %key_origin, %key_tile : index + %value_token = index.assume %value_token0 [lt(%value_token0, %last_full_key_tile_start)] : index + %probability_fragment = vector.fragment.load %probability_stage_view[%c0, %key_tile] shape [%m, %k] : view<16x64xf16, %probability_transposed_layout> -> vector<16xf16> + %value_fragment = vector.fragment.load %value_view[%value_token, %value_channel] shape [%k, %n] : view<[%bounded_key_value_token_count]x[%key_value_width]xf16> -> vector<16xf16> + %next_product_accumulator = vector.mma %probability_fragment, %value_fragment, %product_accumulator : vector<16xf16>, vector<16xf16>, vector<8xf16> + scf.yield %next_product_accumulator : vector<8xf16> + } + vector.fragment.store %product_fragment, %product_stage_view[%c0, %subgroup_product_channel] shape [%m, %n] : vector<8xf16>, view<16x64xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + %owns_output_tile = index.cmp eq, %lane_output_tile, %output_tile : index + %updated_output0, %updated_output1, %updated_output2, %updated_output3 = scf.if %owns_output_tile -> (vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + %block_output0 = vector.load %product_stage_view[%query_row0, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output1 = vector.load %product_stage_view[%query_row1, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output2 = vector.load %product_stage_view[%query_row2, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output3 = vector.load %product_stage_view[%query_row3, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %updated_tile_output0 = vector.addf %tile_output0, %block_output0 : vector<4xf16> + %updated_tile_output1 = vector.addf %tile_output1, %block_output1 : vector<4xf16> + %updated_tile_output2 = vector.addf %tile_output2, %block_output2 : vector<4xf16> + %updated_tile_output3 = vector.addf %tile_output3, %block_output3 : vector<4xf16> + scf.yield %updated_tile_output0, %updated_tile_output1, %updated_tile_output2, %updated_tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } else { + scf.yield %tile_output0, %tile_output1, %tile_output2, %tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + // Complete every read before the next output tile overwrites LDS. + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %updated_output0, %updated_output1, %updated_output2, %updated_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + scf.yield %next_max, %next_sum, %next_output0, %next_output1, %next_output2, %next_output3 : vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } else { + scf.yield %current_max, %current_sum, %current_output0, %current_output1, %current_output2, %current_output3 : vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + scf.yield %next_block_max, %next_block_sum, %next_block_output0, %next_block_output1, %next_block_output2, %next_block_output3 : vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + // A masked 32-row tile handles the final 1-63 KV rows with the same WMMA + // fragments as the aligned path. K and V are staged separately into product + // scratch, so every padded element is initialized and no physical padding is + // required of the caller. Each tile rounds probabilities to F16 before P*V. + %tail_score_wave = index.cmp ult, %subgroup, %c2 : index + %final_max, %final_sum, %final_output0, %final_output1, %final_output2, %final_output3 = scf.for %tail_key_origin = [%full_key_value_token_count to %bounded_key_value_token_count step %c32](%current_max = %full_max : vector<4xf32>, %current_sum = %full_sum : vector<4xf32>, %current_output0 = %full_output0 : vector<4xf16>, %current_output1 = %full_output1 : vector<4xf16>, %current_output2 = %full_output2 : vector<4xf16>, %current_output3 = %full_output3 : vector<4xf16>) -> (vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + %tail_remaining = index.sub %bounded_key_value_token_count, %tail_key_origin : index + %tail_key_count = index.min %tail_remaining, %c32 : index + // Cooperatively stage one K tile, explicitly zeroing the padded rows. + scf.for %load_iteration = [%c0 to %c16 step %c1] unroll { + %linear = index.madd %load_iteration, %c256, %workitem : index + %tail_key_row = index.div %linear, %c128 : index + %tail_key_channel = index.rem %linear, %c128 : index + %tail_key_valid = index.cmp ult, %tail_key_row, %tail_key_count : index + %tail_key_value = scf.if %tail_key_valid -> (f16) { + %tail_key_token0 = index.add %tail_key_origin, %tail_key_row : index + %tail_key_token = index.assume %tail_key_token0 [lt(%tail_key_token0, %bounded_key_value_token_count)] : index + %global_key_channel = index.add %key_value_head_base, %tail_key_channel : index + %loaded = view.load %key_view[%tail_key_token, %global_key_channel] : view<[%bounded_key_value_token_count]x[%key_value_width]xf16> -> f16 + scf.yield %loaded : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %tail_key_value, %tail_key_value_stage_view[%tail_key_row, %tail_key_channel] : f16, view<32x128xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + // Waves zero and one compute the two 16-column QK fragments in this tile. + scf.if %tail_score_wave { + %tail_score_subgroup = index.assume %subgroup [range(%subgroup, 0, 1)] : index + %tail_score_column = index.mul %tail_score_subgroup, %c16 : index + %tail_score_init_values = vector.constant 0.0 : vector<4xf32> + %tail_score_init = vector.fragment %tail_score_init_values shape [%m, %n] : vector<4xf32> + %tail_score_fragment = scf.for %head_tile = [%c0 to %c128 step %c16](%score_accumulator = %tail_score_init : vector<4xf32>) -> (vector<4xf32>) unroll { + %key_fragment = vector.fragment.load %tail_key_value_stage_view[%tail_score_column, %head_tile] shape [%m, %k] : view<32x128xf16> -> vector<16xf16> + %query_fragment = vector.fragment.load %query_transposed_view[%head_tile, %c0] shape [%k, %n] : view<128x16xf16, %query_transposed_layout> -> vector<16xf16> + %next_score_accumulator = vector.mma %key_fragment, %query_fragment, %score_accumulator : vector<16xf16>, vector<16xf16>, vector<4xf32> + scf.yield %next_score_accumulator : vector<4xf32> + } + vector.fragment.store %tail_score_fragment, %score_stage_view[%tail_score_column, %c0] shape [%m, %n] : vector<4xf32>, view<64x24xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + // LDS changes ownership from the QK wave to four query-row waves. Lanes + // beyond the logical tail never read the score or mask buffers. + %tail_lane_valid = index.cmp ult, %lane, %tail_key_count : index + %tail_valid0 = scalar.andi %tail_lane_valid, %query_valid0 : i1 + %tail_valid1 = scalar.andi %tail_lane_valid, %query_valid1 : i1 + %tail_valid2 = scalar.andi %tail_lane_valid, %query_valid2 : i1 + %tail_valid3 = scalar.andi %tail_lane_valid, %query_valid3 : i1 + %tail_valid = vector.from_elements %tail_valid0, %tail_valid1, %tail_valid2, %tail_valid3 : vector<4xi1> + %tail_key_token0 = index.add %tail_key_origin, %lane : index + %masked_score0 = scf.if %tail_valid0 -> (f32) { + %tail_key_token = index.assume %tail_key_token0 [lt(%tail_key_token0, %bounded_key_value_token_count)] : index + %raw_score = view.load %score_stage_view[%lane, %query_row0] : view<64x24xf32> -> f32 + %mask_f16 = view.load %mask_view[%query_token0, %tail_key_token] : view<[%query_token_count]x[%bounded_key_value_token_count]xf16> -> f16 + %mask_f32 = scalar.extf %mask_f16 : f16 to f32 + %score = scalar.addf %raw_score, %mask_f32 : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score1 = scf.if %tail_valid1 -> (f32) { + %tail_key_token = index.assume %tail_key_token0 [lt(%tail_key_token0, %bounded_key_value_token_count)] : index + %raw_score = view.load %score_stage_view[%lane, %query_row1] : view<64x24xf32> -> f32 + %mask_f16 = view.load %mask_view[%query_token1, %tail_key_token] : view<[%query_token_count]x[%bounded_key_value_token_count]xf16> -> f16 + %mask_f32 = scalar.extf %mask_f16 : f16 to f32 + %score = scalar.addf %raw_score, %mask_f32 : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score2 = scf.if %tail_valid2 -> (f32) { + %tail_key_token = index.assume %tail_key_token0 [lt(%tail_key_token0, %bounded_key_value_token_count)] : index + %raw_score = view.load %score_stage_view[%lane, %query_row2] : view<64x24xf32> -> f32 + %mask_f16 = view.load %mask_view[%query_token2, %tail_key_token] : view<[%query_token_count]x[%bounded_key_value_token_count]xf16> -> f16 + %mask_f32 = scalar.extf %mask_f16 : f16 to f32 + %score = scalar.addf %raw_score, %mask_f32 : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_score3 = scf.if %tail_valid3 -> (f32) { + %tail_key_token = index.assume %tail_key_token0 [lt(%tail_key_token0, %bounded_key_value_token_count)] : index + %raw_score = view.load %score_stage_view[%lane, %query_row3] : view<64x24xf32> -> f32 + %mask_f16 = view.load %mask_view[%query_token3, %tail_key_token] : view<[%query_token_count]x[%bounded_key_value_token_count]xf16> -> f16 + %mask_f32 = scalar.extf %mask_f16 : f16 to f32 + %score = scalar.addf %raw_score, %mask_f32 : f32 + scf.yield %score : f32 + } else { + scf.yield %negative_large : f32 + } + %masked_scores = vector.from_elements %masked_score0, %masked_score1, %masked_score2, %masked_score3 : vector<4xf32> + %block_max = kernel.subgroup.reduce %masked_scores : vector<4xf32> + %next_max = vector.maxnumf %current_max, %block_max : vector<4xf32> + %score_delta = vector.subf %masked_scores, %next_max : vector<4xf32> + %raw_probability = vector.expf %score_delta : vector<4xf32> + %probability = vector.select %tail_valid, %raw_probability, %c0_f32x4 : vector<4xf32> + %block_sum = kernel.subgroup.reduce %probability : vector<4xf32> + %old_delta = vector.subf %current_max, %next_max : vector<4xf32> + %old_scale = vector.expf %old_delta : vector<4xf32> + %scaled_current_sum = vector.mulf %current_sum, %old_scale : vector<4xf32> + %next_sum = vector.addf %scaled_current_sum, %block_sum : vector<4xf32> + %probability_f16 = vector.fptrunc %probability : vector<4xf32> to vector<4xf16> + %probability0 = vector.extract %probability_f16[0] : vector<4xf16> -> f16 + %probability1 = vector.extract %probability_f16[1] : vector<4xf16> -> f16 + %probability2 = vector.extract %probability_f16[2] : vector<4xf16> -> f16 + %probability3 = vector.extract %probability_f16[3] : vector<4xf16> -> f16 + view.store %probability0, %probability_stage_view[%query_row0, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability1, %probability_stage_view[%query_row1, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability2, %probability_stage_view[%query_row2, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + view.store %probability3, %probability_stage_view[%query_row3, %lane] : f16, view<16x64xf16, %probability_transposed_layout> + kernel.barrier scope(workgroup) ordering(acq_rel) + // Reuse product scratch for V after every probability is resident in its + // disjoint LDS tile. + scf.for %load_iteration = [%c0 to %c16 step %c1] unroll { + %linear = index.madd %load_iteration, %c256, %workitem : index + %tail_value_row = index.div %linear, %c128 : index + %tail_value_channel = index.rem %linear, %c128 : index + %tail_value_valid = index.cmp ult, %tail_value_row, %tail_key_count : index + %tail_value = scf.if %tail_value_valid -> (f16) { + %tail_value_token0 = index.add %tail_key_origin, %tail_value_row : index + %tail_value_token = index.assume %tail_value_token0 [lt(%tail_value_token0, %bounded_key_value_token_count)] : index + %global_value_channel = index.add %key_value_head_base, %tail_value_channel : index + %loaded = view.load %value_view[%tail_value_token, %global_value_channel] : view<[%bounded_key_value_token_count]x[%key_value_width]xf16> -> f16 + scf.yield %loaded : f16 + } else { + scf.yield %c0_f16 : f16 + } + view.store %tail_value, %tail_key_value_stage_view[%tail_value_row, %tail_value_channel] : f16, view<32x128xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + // Keep tail V staging separate from the 2 KiB product exchange, then use + // the same two 64-channel phases as the aligned path. + %old_scale_f16 = vector.fptrunc %old_scale : vector<4xf32> to vector<4xf16> + %old_scale0_scalar = vector.extract %old_scale_f16[0] : vector<4xf16> -> f16 + %old_scale1_scalar = vector.extract %old_scale_f16[1] : vector<4xf16> -> f16 + %old_scale2_scalar = vector.extract %old_scale_f16[2] : vector<4xf16> -> f16 + %old_scale3_scalar = vector.extract %old_scale_f16[3] : vector<4xf16> -> f16 + %old_scale0 = vector.splat %old_scale0_scalar : vector<4xf16> + %old_scale1 = vector.splat %old_scale1_scalar : vector<4xf16> + %old_scale2 = vector.splat %old_scale2_scalar : vector<4xf16> + %old_scale3 = vector.splat %old_scale3_scalar : vector<4xf16> + %scaled_current_output0 = vector.mulf %current_output0, %old_scale0 : vector<4xf16> + %scaled_current_output1 = vector.mulf %current_output1, %old_scale1 : vector<4xf16> + %scaled_current_output2 = vector.mulf %current_output2, %old_scale2 : vector<4xf16> + %scaled_current_output3 = vector.mulf %current_output3, %old_scale3 : vector<4xf16> + %next_output0, %next_output1, %next_output2, %next_output3 = scf.for %output_tile = [%c0 to %c2 step %c1](%tile_output0 = %scaled_current_output0 : vector<4xf16>, %tile_output1 = %scaled_current_output1 : vector<4xf16>, %tile_output2 = %scaled_current_output2 : vector<4xf16>, %tile_output3 = %scaled_current_output3 : vector<4xf16>) -> (vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) unroll { + %output_tile_channel = index.mul %output_tile, %c64 : index + %value_channel = index.add %output_tile_channel, %subgroup_product_channel : index + %tail_product_init = vector.fragment %c0_f16x8 shape [%m, %n] : vector<8xf16> + %tail_product_fragment = scf.for %key_tile = [%c0 to %c32 step %c16](%product_accumulator = %tail_product_init : vector<8xf16>) -> (vector<8xf16>) unroll { + %probability_fragment = vector.fragment.load %probability_stage_view[%c0, %key_tile] shape [%m, %k] : view<16x64xf16, %probability_transposed_layout> -> vector<16xf16> + %value_fragment = vector.fragment.load %tail_key_value_stage_view[%key_tile, %value_channel] shape [%k, %n] : view<32x128xf16> -> vector<16xf16> + %next_product_accumulator = vector.mma %probability_fragment, %value_fragment, %product_accumulator : vector<16xf16>, vector<16xf16>, vector<8xf16> + scf.yield %next_product_accumulator : vector<8xf16> + } + vector.fragment.store %tail_product_fragment, %product_stage_view[%c0, %subgroup_product_channel] shape [%m, %n] : vector<8xf16>, view<16x64xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + %owns_output_tile = index.cmp eq, %lane_output_tile, %output_tile : index + %updated_output0, %updated_output1, %updated_output2, %updated_output3 = scf.if %owns_output_tile -> (vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16>) { + %block_output0 = vector.load %product_stage_view[%query_row0, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output1 = vector.load %product_stage_view[%query_row1, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output2 = vector.load %product_stage_view[%query_row2, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %block_output3 = vector.load %product_stage_view[%query_row3, %lane_product_channel] : view<16x64xf16> -> vector<4xf16> + %updated_tile_output0 = vector.addf %tile_output0, %block_output0 : vector<4xf16> + %updated_tile_output1 = vector.addf %tile_output1, %block_output1 : vector<4xf16> + %updated_tile_output2 = vector.addf %tile_output2, %block_output2 : vector<4xf16> + %updated_tile_output3 = vector.addf %tile_output3, %block_output3 : vector<4xf16> + scf.yield %updated_tile_output0, %updated_tile_output1, %updated_tile_output2, %updated_tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } else { + scf.yield %tile_output0, %tile_output1, %tile_output2, %tile_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %updated_output0, %updated_output1, %updated_output2, %updated_output3 : vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + scf.yield %next_max, %next_sum, %next_output0, %next_output1, %next_output2, %next_output3 : vector<4xf32>, vector<4xf32>, vector<4xf16>, vector<4xf16>, vector<4xf16>, vector<4xf16> + } + // Normalize and publish the lane-owned packets. Query-tail rows never + // participate in the output store. + scf.if %lane_has_output { + %sum0_scalar = vector.extract %final_sum[0] : vector<4xf32> -> f32 + %sum1_scalar = vector.extract %final_sum[1] : vector<4xf32> -> f32 + %sum2_scalar = vector.extract %final_sum[2] : vector<4xf32> -> f32 + %sum3_scalar = vector.extract %final_sum[3] : vector<4xf32> -> f32 + %inverse_sum0_f32 = scalar.divf %c1_f32, %sum0_scalar : f32 + %inverse_sum1_f32 = scalar.divf %c1_f32, %sum1_scalar : f32 + %inverse_sum2_f32 = scalar.divf %c1_f32, %sum2_scalar : f32 + %inverse_sum3_f32 = scalar.divf %c1_f32, %sum3_scalar : f32 + %inverse_sum0_f16 = scalar.fptrunc %inverse_sum0_f32 : f32 to f16 + %inverse_sum1_f16 = scalar.fptrunc %inverse_sum1_f32 : f32 to f16 + %inverse_sum2_f16 = scalar.fptrunc %inverse_sum2_f32 : f32 to f16 + %inverse_sum3_f16 = scalar.fptrunc %inverse_sum3_f32 : f32 to f16 + %inverse_sum0 = vector.splat %inverse_sum0_f16 : vector<4xf16> + %inverse_sum1 = vector.splat %inverse_sum1_f16 : vector<4xf16> + %inverse_sum2 = vector.splat %inverse_sum2_f16 : vector<4xf16> + %inverse_sum3 = vector.splat %inverse_sum3_f16 : vector<4xf16> + %normalized0_f16 = vector.mulf %final_output0, %inverse_sum0 : vector<4xf16> + %normalized1_f16 = vector.mulf %final_output1, %inverse_sum1 : vector<4xf16> + %normalized2_f16 = vector.mulf %final_output2, %inverse_sum2 : vector<4xf16> + %normalized3_f16 = vector.mulf %final_output3, %inverse_sum3 : vector<4xf16> + %normalized0 = vector.extf %normalized0_f16 : vector<4xf16> to vector<4xf32> + %normalized1 = vector.extf %normalized1_f16 : vector<4xf16> to vector<4xf32> + %normalized2 = vector.extf %normalized2_f16 : vector<4xf16> to vector<4xf32> + %normalized3 = vector.extf %normalized3_f16 : vector<4xf16> to vector<4xf32> + scf.if %query_valid0 { + vector.store %normalized0, %output_view[%query_token0, %query_head, %lane_output_channel] : vector<4xf32>, view<[%query_token_count]x[%query_head_count]x128xf32> + } + scf.if %query_valid1 { + vector.store %normalized1, %output_view[%query_token1, %query_head, %lane_output_channel] : vector<4xf32>, view<[%query_token_count]x[%query_head_count]x128xf32> + } + scf.if %query_valid2 { + vector.store %normalized2, %output_view[%query_token2, %query_head, %lane_output_channel] : vector<4xf32>, view<[%query_token_count]x[%query_head_count]x128xf32> + } + scf.if %query_valid3 { + vector.store %normalized3, %output_view[%query_token3, %query_head, %lane_output_channel] : vector<4xf32>, view<[%query_token_count]x[%query_head_count]x128xf32> + } + } + kernel.return +} + +// Test-only causal-mask construction for the production 14-token witness. +// Finite F16 minima retain exact zero probabilities without requiring infinity +// support from synthetic tensor generators. Selective linking drops this +// helper from production roots. +kernel.def target(@qwen3_moe_attention_gfx11_wave64) @qwen3_moe_flash_attention_test_make_causal_mask() { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%c1, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%mask: buffer) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c14 = index.constant 14 : index + %c196 = index.constant 196 : index + %c0_offset = index.constant 0 : offset + %c0_f16 = scalar.constant 0.0 : f16 + %negative_f16 = scalar.constant -65504.0 : f16 + %workitem = kernel.workitem.id : index + %in_bounds = index.cmp ult, %workitem, %c196 : index + %mask_noalias = buffer.assume.noalias %mask : buffer + %mask_view = buffer.view %mask_noalias[%c0_offset] : buffer -> view<14x14xf16> + scf.if %in_bounds { + %row = index.div %workitem, %c14 : index + %column = index.rem %workitem, %c14 : index + %row_limit = index.add %row, %c1 : index + %is_visible = index.cmp ult, %column, %row_limit : index + %value = scf.select %is_visible, %c0_f16, %negative_f16 : f16 + view.store %value, %mask_view[%row, %column] : f16, view<14x14xf16> + } + kernel.return +} + +// Test-only row extraction for comparing one multirow attention dispatch with +// independent one-row dispatches over identical data. Selective linking drops +// this helper from production roots. +kernel.def target(@qwen3_moe_attention_gfx11_wave64) @qwen3_moe_flash_attention_test_extract_row(%query_token_count: index, %context_count: index, %source_row: index) { + %query_head_count = config.get @qwen3_moe.attention.query_head_count : index + %context_capacity = config.get @qwen3_moe.attention.test.context_capacity : index + %c128 = index.constant 128 : index + %c255 = index.constant 255 : index + %c256 = index.constant 256 : index + %c1 = index.constant 1 : index + %query_element_count = index.mul %query_head_count, %c128 : index + %element_count = index.add %query_element_count, %context_capacity : index + %padded_element_count = index.add %element_count, %c255 : index + %workgroup_count = index.div %padded_element_count, %c256 : index + kernel.launch.config workgroups(%workgroup_count, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%query_token_count: index, %context_count: index, %source_row: index, %source_query: buffer, %source_mask: buffer, %target_query: buffer, %target_mask: buffer) { + %c0 = index.constant 0 : index + %c0_offset = index.constant 0 : offset + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %query_head_count = config.get @qwen3_moe.attention.query_head_count : index + %context_capacity = config.get @qwen3_moe.attention.test.context_capacity : index + %bounded_query_token_count = index.assume %query_token_count [range(%query_token_count, 1, 2048)] : index + %bounded_context_count = index.assume %context_count [range(%context_count, 1, 32768), le(%context_count, %context_capacity)] : index + %bounded_source_row, %source_query_token_count = index.assume %source_row, %bounded_query_token_count [range(%source_row, 0, 2047), lt(%source_row, %bounded_query_token_count)] : index, index + %workgroup = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %linear = index.madd %workgroup, %c256, %workitem : index + %query_element_count = index.mul %query_head_count, %c128 : index + %element_count = index.add %query_element_count, %bounded_context_count : index + %in_bounds = index.cmp ult, %linear, %element_count : index + %source_query_noalias, %source_mask_noalias, %target_query_noalias, %target_mask_noalias = buffer.assume.noalias %source_query, %source_mask, %target_query, %target_mask : buffer, buffer, buffer, buffer + %source_query_view = buffer.view %source_query_noalias[%c0_offset] : buffer -> view<[%source_query_token_count]x[%query_head_count]x128xf32> + %source_mask_view = buffer.view %source_mask_noalias[%c0_offset] : buffer -> view<[%source_query_token_count]x[%bounded_context_count]xf16> + %target_query_view = buffer.view %target_query_noalias[%c0_offset] : buffer -> view<1x[%query_head_count]x128xf32> + %target_mask_view = buffer.view %target_mask_noalias[%c0_offset] : buffer -> view<1x[%bounded_context_count]xf16> + scf.if %in_bounds { + %is_query_element = index.cmp ult, %linear, %query_element_count : index + scf.if %is_query_element { + %query_head = index.div %linear, %c128 : index + %query_channel = index.rem %linear, %c128 : index + %value = view.load %source_query_view[%bounded_source_row, %query_head, %query_channel] : view<[%source_query_token_count]x[%query_head_count]x128xf32> -> f32 + view.store %value, %target_query_view[%c0, %query_head, %query_channel] : f32, view<1x[%query_head_count]x128xf32> + } else { + %mask_column = index.sub %linear, %query_element_count : index + %value = view.load %source_mask_view[%bounded_source_row, %mask_column] : view<[%source_query_token_count]x[%bounded_context_count]xf16> -> f16 + view.store %value, %target_mask_view[%c0, %mask_column] : f16, view<1x[%bounded_context_count]xf16> + } + } + kernel.return +} + +// The mask selects the first KV row exactly. QK, F16 probability conversion, +// GQA addressing, and the P*V path all execute, while the expected result is +// the first V row and remains auditable as an iota. +check.case public @qwen3_moe_flash_attention_f32_f16_wmma_selected_row_case { + %query_token_count = check.literal value(1) : index + %key_value_token_count = check.literal value(64) : index + %query = check.generate.fill value(1.0) : tensor<1x1x128xf32> + %key = check.generate.fill value(1.0) : tensor<64x1x128xf16> + %value = check.generate.iota offset(0.0) step(0.125) : tensor<64x1x128xf16> + %mask = check.generate.iota offset(0.0) step(-10000.0) : tensor<1x64xf16> + %output = check.generate.fill value(-1.0) : tensor<1x1x128xf32> + %expected = check.generate.iota offset(0.0) step(0.125) : tensor<1x1x128xf32> + kernel.launch @qwen3_moe_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<1x1x128xf32>, tensor<64x1x128xf16>, tensor<64x1x128xf16>, tensor<1x64xf16>, tensor<1x1x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<1x1x128xf32> + check.return +} + +// Four query rows select four distinct KV rows. The 33-element mask period +// maps row-major index row*32+column to zero exactly when row == column for +// rows zero through three. The expected output is therefore the first four +// rows of V, preserving an auditable iota while distinguishing every component +// in the four-row per-wave ownership path. +check.case public @qwen3_moe_flash_attention_f32_f16_wmma_multirow_selected_case { + %query_token_count = check.literal value(4) : index + %key_value_token_count = check.literal value(32) : index + %query = check.generate.fill value(1.0) : tensor<4x1x128xf32> + %key = check.generate.fill value(1.0) : tensor<32x1x128xf16> + %value = check.generate.iota offset(0.0) step(0.125) : tensor<32x1x128xf16> + %mask = check.generate.iota offset(0.0) step(-10000.0) period(33) : tensor<4x32xf16> + %output = check.generate.fill value(-1.0) : tensor<4x1x128xf32> + %expected = check.generate.iota offset(0.0) step(0.125) : tensor<4x1x128xf32> + kernel.launch @qwen3_moe_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<4x1x128xf32>, tensor<32x1x128xf16>, tensor<32x1x128xf16>, tensor<4x32xf16>, tensor<4x1x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<4x1x128xf32> + check.return +} + +// A unit mask step gives every row the same geometric softmax distribution, +// shifted to begin at the row-matched KV entry. Values advance by 1/4096, so +// each KV row adds exactly 1/32 and the infinite-series weighted row offset is +// 1/(32*(e-1)). Terms that wrap at period 33 are below F16 significance. This +// exercises four independent online-softmax and P*V states while retaining a +// compact closed-form expected iota. +check.case public @qwen3_moe_flash_attention_f32_f16_wmma_multirow_online_case { + %query_token_count = check.literal value(4) : index + %key_value_token_count = check.literal value(32) : index + %query = check.generate.fill value(1.0) : tensor<4x1x128xf32> + %key = check.generate.fill value(1.0) : tensor<32x1x128xf16> + %value = check.generate.iota offset(0.0) step(0.000244140625) : tensor<32x1x128xf16> + %mask = check.generate.iota offset(0.0) step(-1.0) period(33) : tensor<4x32xf16> + %output = check.generate.fill value(-1.0) : tensor<4x1x128xf32> + %expected = check.generate.iota offset(0.01818677224124075) step(0.000244140625) : tensor<4x1x128xf32> + kernel.launch @qwen3_moe_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<4x1x128xf32>, tensor<32x1x128xf16>, tensor<32x1x128xf16>, tensor<4x32xf16>, tensor<4x1x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<4x1x128xf32> + check.return +} + +// Compare the first four rows of one fourteen-row launch with four one-row +// launches over identical, nonuniform data and the production causal mask. +// Fourteen rows retain the exact dynamic dense-row stride, ownership, and +// query-validity pressure, while the one-row path is the proven containment. +// Row extraction only reshapes bindings and performs no attention arithmetic. +check.case public @qwen3_moe_flash_attention_f32_f16_wmma_multirow_differential_case { + %query_token_count = check.literal value(14) : index + %c1 = check.literal value(1) : index + %key_value_token_count = check.literal value(14) : index + %row0 = check.literal value(0) : index + %row1 = check.literal value(1) : index + %row2 = check.literal value(2) : index + %row3 = check.literal value(3) : index + %query_seed = check.param.seed base(5858425849414763858) count(1) : i64 + %key_seed = check.param.seed base(5858425849313057073) count(1) : i64 + %value_seed = check.param.seed base(5858425849497340977) count(1) : i64 + %query = check.generate.random.uniform seed(%query_seed) range(-1.0 to 1.0) : tensor<14x32x128xf32> + %key = check.generate.random.uniform seed(%key_seed) range(-1.0 to 1.0) : tensor<14x4x128xf16> + %value = check.generate.random.uniform seed(%value_seed) range(-1.0 to 1.0) : tensor<14x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<14x14xf16> + %actual = check.generate.fill value(-1.0) : tensor<14x32x128xf32> + %row_query = check.generate.fill value(0.0) : tensor<1x32x128xf32> + %row_mask = check.generate.fill value(0.0) : tensor<1x14xf16> + %expected0 = check.generate.fill value(0.0) : tensor<1x32x128xf32> + %expected1 = check.generate.fill value(0.0) : tensor<1x32x128xf32> + %expected2 = check.generate.fill value(0.0) : tensor<1x32x128xf32> + %expected3 = check.generate.fill value(0.0) : tensor<1x32x128xf32> + %actual0 = check.generate.fill value(-1.0) : tensor<1x32x128xf32> + %actual1 = check.generate.fill value(-1.0) : tensor<1x32x128xf32> + %actual2 = check.generate.fill value(-1.0) : tensor<1x32x128xf32> + %actual3 = check.generate.fill value(-1.0) : tensor<1x32x128xf32> + kernel.launch @qwen3_moe_flash_attention_test_make_causal_mask(%mask) : (tensor<14x14xf16>) + kernel.launch @qwen3_moe_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %actual) : [index, index](index, index, tensor<14x32x128xf32>, tensor<14x4x128xf16>, tensor<14x4x128xf16>, tensor<14x14xf16>, tensor<14x32x128xf32>) + kernel.launch @qwen3_moe_flash_attention_test_extract_row[%query_token_count, %key_value_token_count, %row0](%query_token_count, %key_value_token_count, %row0, %query, %mask, %row_query, %row_mask) : [index, index, index](index, index, index, tensor<14x32x128xf32>, tensor<14x14xf16>, tensor<1x32x128xf32>, tensor<1x14xf16>) + kernel.launch @qwen3_moe_flash_attention_f32_f16_wmma[%c1, %key_value_token_count](%c1, %key_value_token_count, %row_query, %key, %value, %row_mask, %expected0) : [index, index](index, index, tensor<1x32x128xf32>, tensor<14x4x128xf16>, tensor<14x4x128xf16>, tensor<1x14xf16>, tensor<1x32x128xf32>) + kernel.launch @qwen3_moe_flash_attention_test_extract_row[%query_token_count, %key_value_token_count, %row0](%query_token_count, %key_value_token_count, %row0, %actual, %mask, %actual0, %row_mask) : [index, index, index](index, index, index, tensor<14x32x128xf32>, tensor<14x14xf16>, tensor<1x32x128xf32>, tensor<1x14xf16>) + kernel.launch @qwen3_moe_flash_attention_test_extract_row[%query_token_count, %key_value_token_count, %row1](%query_token_count, %key_value_token_count, %row1, %query, %mask, %row_query, %row_mask) : [index, index, index](index, index, index, tensor<14x32x128xf32>, tensor<14x14xf16>, tensor<1x32x128xf32>, tensor<1x14xf16>) + kernel.launch @qwen3_moe_flash_attention_f32_f16_wmma[%c1, %key_value_token_count](%c1, %key_value_token_count, %row_query, %key, %value, %row_mask, %expected1) : [index, index](index, index, tensor<1x32x128xf32>, tensor<14x4x128xf16>, tensor<14x4x128xf16>, tensor<1x14xf16>, tensor<1x32x128xf32>) + kernel.launch @qwen3_moe_flash_attention_test_extract_row[%query_token_count, %key_value_token_count, %row1](%query_token_count, %key_value_token_count, %row1, %actual, %mask, %actual1, %row_mask) : [index, index, index](index, index, index, tensor<14x32x128xf32>, tensor<14x14xf16>, tensor<1x32x128xf32>, tensor<1x14xf16>) + kernel.launch @qwen3_moe_flash_attention_test_extract_row[%query_token_count, %key_value_token_count, %row2](%query_token_count, %key_value_token_count, %row2, %query, %mask, %row_query, %row_mask) : [index, index, index](index, index, index, tensor<14x32x128xf32>, tensor<14x14xf16>, tensor<1x32x128xf32>, tensor<1x14xf16>) + kernel.launch @qwen3_moe_flash_attention_f32_f16_wmma[%c1, %key_value_token_count](%c1, %key_value_token_count, %row_query, %key, %value, %row_mask, %expected2) : [index, index](index, index, tensor<1x32x128xf32>, tensor<14x4x128xf16>, tensor<14x4x128xf16>, tensor<1x14xf16>, tensor<1x32x128xf32>) + kernel.launch @qwen3_moe_flash_attention_test_extract_row[%query_token_count, %key_value_token_count, %row2](%query_token_count, %key_value_token_count, %row2, %actual, %mask, %actual2, %row_mask) : [index, index, index](index, index, index, tensor<14x32x128xf32>, tensor<14x14xf16>, tensor<1x32x128xf32>, tensor<1x14xf16>) + kernel.launch @qwen3_moe_flash_attention_test_extract_row[%query_token_count, %key_value_token_count, %row3](%query_token_count, %key_value_token_count, %row3, %query, %mask, %row_query, %row_mask) : [index, index, index](index, index, index, tensor<14x32x128xf32>, tensor<14x14xf16>, tensor<1x32x128xf32>, tensor<1x14xf16>) + kernel.launch @qwen3_moe_flash_attention_f32_f16_wmma[%c1, %key_value_token_count](%c1, %key_value_token_count, %row_query, %key, %value, %row_mask, %expected3) : [index, index](index, index, tensor<1x32x128xf32>, tensor<14x4x128xf16>, tensor<14x4x128xf16>, tensor<1x14xf16>, tensor<1x32x128xf32>) + kernel.launch @qwen3_moe_flash_attention_test_extract_row[%query_token_count, %key_value_token_count, %row3](%query_token_count, %key_value_token_count, %row3, %actual, %mask, %actual3, %row_mask) : [index, index, index](index, index, index, tensor<14x32x128xf32>, tensor<14x14xf16>, tensor<1x32x128xf32>, tensor<1x14xf16>) + check.expect.close actual(%actual0) expected(%expected0) atol(0.001) rtol(0.001) nan(same) : tensor<1x32x128xf32> + check.expect.close actual(%actual1) expected(%expected1) atol(0.001) rtol(0.001) nan(same) : tensor<1x32x128xf32> + check.expect.close actual(%actual2) expected(%expected2) atol(0.001) rtol(0.001) nan(same) : tensor<1x32x128xf32> + check.expect.close actual(%actual3) expected(%expected3) atol(0.001) rtol(0.001) nan(same) : tensor<1x32x128xf32> + check.return +} + +// Seventeen rows force a partial second query tile. Equal scores and constant +// values make every valid output exactly two while still exercising online +// normalization and the query-tail guards. +check.case public @qwen3_moe_flash_attention_f32_f16_wmma_query_tail_case { + %query_token_count = check.literal value(17) : index + %key_value_token_count = check.literal value(64) : index + %query = check.generate.fill value(1.0) : tensor<17x1x128xf32> + %key = check.generate.fill value(1.0) : tensor<64x1x128xf16> + %value = check.generate.fill value(2.0) : tensor<64x1x128xf16> + %mask = check.generate.fill value(0.0) : tensor<17x64xf16> + %output = check.generate.fill value(-1.0) : tensor<17x1x128xf32> + %expected = check.generate.fill value(2.0) : tensor<17x1x128xf32> + kernel.launch @qwen3_moe_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<17x1x128xf32>, tensor<64x1x128xf16>, tensor<64x1x128xf16>, tensor<17x64xf16>, tensor<17x1x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<17x1x128xf32> + check.return +} + +// Sixty-five KV rows force one masked cleanup tile after a full WMMA block. +// The mask selects only that final row, whose iota values begin at 1024, so +// omitting the cleanup path cannot accidentally satisfy the check. +check.case public @qwen3_moe_flash_attention_f32_f16_wmma_key_value_tail_case { + %query_token_count = check.literal value(1) : index + %key_value_token_count = check.literal value(65) : index + %query = check.generate.fill value(1.0) : tensor<1x1x128xf32> + %key = check.generate.fill value(1.0) : tensor<65x1x128xf16> + %value = check.generate.iota offset(0.0) step(0.125) : tensor<65x1x128xf16> + %mask = check.generate.iota offset(-64000.0) step(1000.0) : tensor<1x65xf16> + %output = check.generate.fill value(-1.0) : tensor<1x1x128xf32> + %expected = check.generate.iota offset(1024.0) step(0.125) : tensor<1x1x128xf32> + kernel.launch @qwen3_moe_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<1x1x128xf32>, tensor<65x1x128xf16>, tensor<65x1x128xf16>, tensor<1x65xf16>, tensor<1x1x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<1x1x128xf32> + check.return +} + +// The first 64-row block contains the selected value while every mask in the +// second block rounds to negative infinity. This preserves a finite online +// softmax state while exercising the workgroup-wide block-pruning path. +check.case public @qwen3_moe_flash_attention_f32_f16_wmma_pruned_block_case { + %query_token_count = check.literal value(1) : index + %key_value_token_count = check.literal value(128) : index + %query = check.generate.fill value(1.0) : tensor<1x1x128xf32> + %key = check.generate.fill value(1.0) : tensor<128x1x128xf16> + %value = check.generate.iota offset(0.0) step(0.125) : tensor<128x1x128xf16> + %mask = check.generate.iota offset(0.0) step(-2000.0) : tensor<1x128xf16> + %output = check.generate.fill value(-1.0) : tensor<1x1x128xf32> + %expected = check.generate.iota offset(0.0) step(0.125) : tensor<1x1x128xf32> + kernel.launch @qwen3_moe_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<1x1x128xf32>, tensor<128x1x128xf16>, tensor<128x1x128xf16>, tensor<1x128xf16>, tensor<1x1x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.001) rtol(0.001) nan(same) : tensor<1x1x128xf32> + check.return +} + +check.case public @qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case { + %query_token_count = check.param.choice values([1, 32, 64, 128, 192, 255, 256, 257, 384, 511, 512, 513, 768, 1023, 1024, 1025, 1280, 1536, 1792, 2048]) name("query_token_count") : index + %key_value_token_count = check.param.choice values([64, 128, 192, 255, 256, 257, 384, 511, 512, 513, 768, 1023, 1024, 1025, 1280, 1536, 1792, 2048, 32768]) name("key_value_token_count") : index + %query = check.generate.fill value(0.0) : tensor<[%query_token_count]x32x128xf32> + %key = check.generate.fill value(0.0) : tensor<[%key_value_token_count]x4x128xf16> + %value = check.generate.fill value(0.0) : tensor<[%key_value_token_count]x4x128xf16> + %mask = check.generate.fill value(0.0) : tensor<[%query_token_count]x[%key_value_token_count]xf16> + %output = check.generate.fill value(1.0) : tensor<[%query_token_count]x32x128xf32> + %expected = check.generate.fill value(0.0) : tensor<[%query_token_count]x32x128xf32> + kernel.launch @qwen3_moe_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<[%query_token_count]x32x128xf32>, tensor<[%key_value_token_count]x4x128xf16>, tensor<[%key_value_token_count]x4x128xf16>, tensor<[%query_token_count]x[%key_value_token_count]xf16>, tensor<[%query_token_count]x32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%query_token_count]x32x128xf32> + check.return +} + +// Four finite 64-row blocks followed by four negative-infinity blocks model +// the aggregate computed/pruned work ratio of 512-token causal prefill while +// keeping every workgroup's path identical for a stable microbenchmark. +check.case public @qwen3_moe_flash_attention_f32_f16_wmma_pruned_half_benchmark_case { + %query_token_count = check.literal value(512) : index + %key_value_token_count = check.literal value(512) : index + %query = check.generate.fill value(0.0) : tensor<512x32x128xf32> + %key = check.generate.fill value(0.0) : tensor<512x4x128xf16> + %value = check.generate.fill value(0.0) : tensor<512x4x128xf16> + %mask = check.generate.iota offset(0.0) step(-300.0) period(512) : tensor<512x512xf16> + %output = check.generate.fill value(1.0) : tensor<512x32x128xf32> + %expected = check.generate.fill value(0.0) : tensor<512x32x128xf32> + kernel.launch @qwen3_moe_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<512x32x128xf32>, tensor<512x4x128xf16>, tensor<512x4x128xf16>, tensor<512x512xf16>, tensor<512x32x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<512x32x128xf32> + check.return +} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_selected_row_case> @qwen3_moe_flash_attention_f32_f16_wmma_selected_row + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_query_tail_case> @qwen3_moe_flash_attention_f32_f16_wmma_query_tail + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_key_value_tail_case> @qwen3_moe_flash_attention_f32_f16_wmma_key_value_tail + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_pruned_block_case> @qwen3_moe_flash_attention_f32_f16_wmma_pruned_block + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_pruned_half_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_prefill_512_pruned_half + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_decode_256 {key_value_token_count = 256, query_token_count = 1} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_decode_2048 {key_value_token_count = 2048, query_token_count = 1} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_decode_32768 {key_value_token_count = 32768, query_token_count = 1} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_prefill_32 {key_value_token_count = 256, query_token_count = 32} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_prefill_128 {key_value_token_count = 256, query_token_count = 128} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_prefill_512 {key_value_token_count = 512, query_token_count = 512} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_prefill_512_context_1024 {key_value_token_count = 1024, query_token_count = 512} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_prefill_512_context_1536 {key_value_token_count = 1536, query_token_count = 512} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_prefill_512_context_2048 {key_value_token_count = 2048, query_token_count = 512} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_prefill_1024 {key_value_token_count = 1024, query_token_count = 1024} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_prefill_2048 {key_value_token_count = 2048, query_token_count = 2048} + +// These aligned and boundary-adjacent self-attention shapes expose launch or +// tail cliffs that a powers-of-two-only benchmark would hide. They are +// measurement witnesses for one shape-specialized kernel, not routing buckets. +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_64 {key_value_token_count = 64, query_token_count = 64} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_128 {key_value_token_count = 128, query_token_count = 128} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_192 {key_value_token_count = 192, query_token_count = 192} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_255 {key_value_token_count = 255, query_token_count = 255} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_256 {key_value_token_count = 256, query_token_count = 256} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_257 {key_value_token_count = 257, query_token_count = 257} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_384 {key_value_token_count = 384, query_token_count = 384} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_511 {key_value_token_count = 511, query_token_count = 511} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_512 {key_value_token_count = 512, query_token_count = 512} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_513 {key_value_token_count = 513, query_token_count = 513} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_768 {key_value_token_count = 768, query_token_count = 768} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_1023 {key_value_token_count = 1023, query_token_count = 1023} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_1024 {key_value_token_count = 1024, query_token_count = 1024} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_1025 {key_value_token_count = 1025, query_token_count = 1025} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_1280 {key_value_token_count = 1280, query_token_count = 1280} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_1536 {key_value_token_count = 1536, query_token_count = 1536} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_1792 {key_value_token_count = 1792, query_token_count = 1792} + +check.benchmark<@qwen3_moe_flash_attention_f32_f16_wmma_benchmark_case> @qwen3_moe_flash_attention_f32_f16_wmma_self_2048 {key_value_token_count = 2048, query_token_count = 2048} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/model_config.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/model_config.loom new file mode 100644 index 000000000000..59f6b68a28cd --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/model_config.loom @@ -0,0 +1,18 @@ +// Qwen model-wide configuration shared across independently linked kernels. +// +// Layer-local storage formats and schedule choices remain with their owning +// kernels. These values describe the semantic model contract and therefore +// have one symbol definition regardless of how many kernels consume them. +config.decl @qwen3_moe.model.hidden_size : %value: index where [range(%value, 128, 32768), mul(%value, 128)] + +config.decl @qwen3_moe.model.rms_epsilon : f32 + +config.decl @qwen3_moe.attention.head_size : %value: index where [range(%value, 4, 1024), mul(%value, 4)] + +config.decl @qwen3_moe.attention.query_size : %value: index where [range(%value, 1, 262144)] + +config.decl @qwen3_moe.attention.key_value_size : %value: index where [range(%value, 1, 262144)] + +config.decl @qwen3_moe.router.expert_count : %value: index where [range(%value, 32, 512), mul(%value, 32)] + +config.decl @qwen3_moe.router.route_count : %value: index where [range(%value, 1, 32)] diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_next_q8.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_next_q8.loom new file mode 100644 index 000000000000..d6e4dea43918 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_next_q8.loom @@ -0,0 +1,59 @@ +// Shared decode completion protocol for routed-down projections that publish +// the normalized Q8_1 x4 row consumed by the next projection boundary. Each +// x workgroup must publish one output tile before applying this body and must +// provide both its tile width and the number of target subgroups participating +// in the row reduction. The last arrival acquires the complete residual row, +// normalizes and packs it, then resets the reusable completion word after every +// Q8 store. +// +// This outer device template deliberately applies the RMSNorm/Q8 device +// template. Keeping that composition authored here ensures the compiler +// preserves the eventual kernel ancestor through nested template selection. +template.decl @qwen3_moe.rmsnorm_quantize_q8_1_x4.body(%arg0: i1, %arg1: index, %arg2: index, %arg3: index, %arg4: buffer, %arg5: buffer, %arg6: buffer, %arg7: buffer) + +template.decl @qwen3_moe.routed_down.next_q8_completion(%output_channels_per_workgroup0: index, %reduction_subgroup_count0: index, %token_count: index, %output_size: index, %output: buffer, %norm_weight: buffer, %completion_counter: buffer, %next_q8_output: buffer) + +template.def<@qwen3_moe.routed_down.next_q8_completion> device @qwen3_moe_routed_down_next_q8_completion(%output_channels_per_workgroup0: index, %reduction_subgroup_count0: index, %token_count: index, %output_size: index, %output: buffer, %norm_weight: buffer, %completion_counter: buffer, %next_q8_output: buffer) { + %output_channels_per_workgroup = index.assume %output_channels_per_workgroup0 [range(%output_channels_per_workgroup0, 1, 8)] : index + %reduction_subgroup_count = index.assume %reduction_subgroup_count0 [range(%reduction_subgroup_count0, 1, 8)] : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 1)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 128, 32768), mul(%output_size, 128)] : index + %token0 = kernel.workgroup.id : index + %token = index.assume %token0 [lt(%token0, %bounded_token_count)] : index + %c0 = index.constant 0 : index + %workitem = kernel.workitem.id : index + %output_tile_count = index.div %bounded_output_size, %output_channels_per_workgroup : index + %is_arrival_workitem = index.cmp eq, %workitem, %c0 : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c0_offset = index.constant 0 : offset + %counter_scratch_bytes = index.constant 4 : offset + %output_noalias, %norm_weight_noalias, %completion_counter_noalias, %next_q8_output_noalias = buffer.assume.noalias %output, %norm_weight, %completion_counter, %next_q8_output : buffer, buffer, buffer, buffer + %completion_counter_aligned = buffer.assume.alignment %completion_counter_noalias {minimum_alignment = 16} : buffer + %completion_counter_view = buffer.view %completion_counter_aligned[%c0_offset] : buffer -> view<1xi32> + %counter_scratch = buffer.alloca align(4) %counter_scratch_bytes : buffer + %counter_scratch_view = buffer.view %counter_scratch[%c0_offset] : buffer -> view<1xi32> + // Publish every producer's residual stores before the leader advances one + // workgroup arrival. The last arrival then acquires the complete row. + kernel.barrier scope(workgroup) ordering(release) + scf.if %is_arrival_workitem { + %old_counter = view.atomic.rmw %c1_i32, %completion_counter_view[%c0] {ordering = acq_rel, scope = device} : i32, view<1xi32> -> i32 + view.store %old_counter, %counter_scratch_view[%c0] : i32, view<1xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %old_counter = view.load %counter_scratch_view[%c0] : view<1xi32> -> i32 + %output_tile_count_i32 = index.cast %output_tile_count : index to i32 + %last_output_tile_i32 = scalar.subi %output_tile_count_i32, %c1_i32 : i32 + %negative_output_tile_count_i32 = scalar.subi %c0_i32, %output_tile_count_i32 : i32 + %is_last_output_tile = scalar.cmpi eq, %old_counter, %last_output_tile_i32 : i32 + scf.if %is_last_output_tile { + kernel.barrier scope(workgroup) ordering(acquire) + %publish_normalized = scalar.constant false : i1 + template.apply<@qwen3_moe.rmsnorm_quantize_q8_1_x4.body>(%publish_normalized, %reduction_subgroup_count, %bounded_token_count, %token, %output_noalias, %norm_weight_noalias, %next_q8_output_noalias, %next_q8_output_noalias) : (i1, index, index, index, buffer, buffer, buffer, buffer) + kernel.barrier scope(workgroup) ordering(release) + scf.if %is_arrival_workitem { + view.atomic.reduce %negative_output_tile_count_i32, %completion_counter_view[%c0] {ordering = release, scope = device} : i32, view<1xi32> + } + } + template.return +} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_q4k.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_q4k.loom new file mode 100644 index 000000000000..3e8c57a96b5b --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_q4k.loom @@ -0,0 +1,399 @@ +// Fuses the Qwen3 MoE Q4_K down projection with route weighting, top-8 +// reduction, and residual publication. The input contains one Q8_1 x4 row per +// [token, route], while weights remain in GGUF's raw +// [expert, output, K / 256, 144 bytes] layout. +// +// Route IDs retain an independent physical stride because llama.cpp selects a +// top-8 view from its 128-entry argsort storage. Normalized route weights are +// compact [token, route]. The output enters containing the residual and is +// updated in place, so the fused boundary never materializes an unweighted or +// route-indexed [token, route, hidden] down-projection tensor. +// +// One wave owns one output channel and contracts all selected routes in +// registers. Four eight-lane cohorts contract four independent routes while +// walking every Q4_K block. Qwen's eight-route, three-block decode shape uses +// all 32 lanes across two route batches without addressing synthetic blocks. +template.decl @qwen3_moe.routed_down.next_q8_completion(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: buffer, %arg5: buffer, %arg6: buffer, %arg7: buffer) + +template.decl @qwen3_moe.routed_down.q4k_q8_1_x4.body(%publish_output: i1, %token_count: index, %token: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index, %q8_input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer) + +target.decl @qwen3_moe_attention_prepare_gfx11_wave32 + +config.decl @qwen3_moe.routed_down.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +config.decl @qwen3_moe.routed_down.route_count : %value: index where [range(%value, 1, 8)] + +config.decl @qwen3_moe.routed_down.expert_count : %value: index where [range(%value, 1, 512)] + +config.decl @qwen3_moe.routed_down.output_size : %value: index where [range(%value, 1, 4096)] + +config.decl @qwen3_moe.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +kernel.decl @ggml_quantize_q8_1_x4_f32(%token_count$26: index, %input_size$27: index) launch(%token_count$28: index, %input_size$29: index, %input: buffer, %output: buffer) + +func.decl @qwen3_moe_q4k_q8_1_x4_cohort_row_lane(%input_size: index, %weight: buffer, %weight_row_byte_base: offset, %q8_input: buffer, %q8_row_byte_base: offset, %cohort_lane: index) -> (f32) + +kernel.decl @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4(%token_count$39: index) launch(%token_count$40: index, %input: buffer, %weight: buffer, %q8_output: buffer) + +// Shared Q4_K contraction and residual publication. Each lane consumes both +// nibbles of a packed code load before the subgroup reduces the row. +template.def<@qwen3_moe.routed_down.q4k_q8_1_x4.body> device @qwen3_moe_routed_down_q4k_q8_1_x4_body(%publish_output: i1, %token_count: index, %token: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index, %q8_input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_token, %body_token_count = index.assume %token, %bounded_token_count [lt(%token, %bounded_token_count)] : index, index + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %bounded_route_count = index.assume %route_count [range(%route_count, 1, 8)] : index + %bounded_route_id_stride = index.assume %route_id_stride [range(%route_id_stride, 1, 512)] : index + %bounded_expert_count = index.assume %expert_count [range(%expert_count, 1, 512)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 262144)] : index + %channel_tile = kernel.workgroup.id : index + %subgroup0 = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c128 = index.constant 128 : index + %c144 = index.constant 144 : index + %c256 = index.constant 256 : index + %c1_byte = index.constant 1 : offset + %c0_i32 = scalar.constant 0 : i32 + %c0_f32 = scalar.constant 0.0 : f32 + %c0_offset = index.constant 0 : offset + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %cohort0 = index.div %lane, %c8 : index + %cohort = index.assume %cohort0 [range(%cohort0, 0, 3)] : index + %cohort_lane0 = index.rem %lane, %c8 : index + %cohort_lane = index.assume %cohort_lane0 [range(%cohort_lane0, 0, 7)] : index + %channel_base = index.mul %channel_tile, %c8 : index + %channel = index.add %channel_base, %subgroup : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %lane_i32 = index.cast %lane : index to i32 + %is_lane_zero = scalar.cmpi eq, %lane_i32, %c0_i32 : i32 + %q4_block_count = index.div %bounded_input_size, %c256 : index + %weight_row_byte_count = index.mul %q4_block_count, %c144 : index + %q8_group_count = index.div %bounded_input_size, %c128 : index + %q8_row_byte_count = index.mul %q8_group_count, %c144 : index + %padded_route_count = index.add %bounded_route_count, %c3 : index + %route_batch_count = index.div %padded_route_count, %c4 : index + %q8_noalias, %route_id_noalias, %route_weight_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %q8_input, %route_ids, %route_weights, %weight, %output : buffer, buffer, buffer, buffer, buffer + %route_id_view = buffer.view %route_id_noalias[%c0_offset] : buffer -> view<[%body_token_count]x[%bounded_route_id_stride]xi32> + %route_weight_view = buffer.view %route_weight_noalias[%c0_offset] : buffer -> view<[%body_token_count]x[%bounded_route_count]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%body_token_count]x[%bounded_output_size]xf32> + %lane_has_route = index.cmp ult, %lane, %bounded_route_count : index + %lane_within_route_stride = index.cmp ult, %lane, %bounded_route_id_stride : index + %loads_active_route = scalar.andi %publish_output, %lane_has_route : i1 + %loads_route_metadata = scalar.andi %loads_active_route, %lane_within_route_stride : i1 + %lane_expert_i32, %lane_route_weight = scf.if %loads_route_metadata -> (i32, f32) { + %route_lane0, %metadata_route_count = index.assume %lane, %bounded_route_count [lt(%lane, %bounded_route_count)] : index, index + %route_lane, %metadata_route_id_stride = index.assume %route_lane0, %bounded_route_id_stride [lt(%route_lane0, %bounded_route_id_stride)] : index, index + %expert_i32 = view.load %route_id_view[%bounded_token, %route_lane] : view<[%body_token_count]x[%bounded_route_id_stride]xi32> -> i32 + %route_weight = view.load %route_weight_view[%bounded_token, %route_lane0] : view<[%body_token_count]x[%bounded_route_count]xf32> -> f32 + scf.yield %expert_i32, %route_weight : i32, f32 + } else { + scf.yield %c0_i32, %c0_f32 : i32, f32 + } + %computes_channel = scalar.andi %publish_output, %valid_channel : i1 + %routed_lane_sum = scf.if %computes_channel -> (f32) { + %sum = scf.for %route_batch = [%c0 to %route_batch_count step %c1](%route_acc = %c0_f32 : f32) -> (f32) unroll { + %route_batch_base = index.mul %route_batch, %c4 : index + %route0 = index.add %route_batch_base, %cohort : index + %route = index.assume %route0 [range(%route0, 0, 7)] : index + %active_route = index.cmp ult, %route, %bounded_route_count : index + %weighted_lane = scf.if %active_route -> (f32) { + %bounded_route, %body_route_count = index.assume %route, %bounded_route_count [lt(%route, %bounded_route_count)] : index, index + %route_i32 = index.cast %bounded_route : index to i32 + %expert_i32 = kernel.subgroup.broadcast %lane_expert_i32 from %route_i32 : i32, i32 + %expert0 = index.cast %expert_i32 : i32 to index + %expert1 = index.assume %expert0 [range(%expert0, 0, 511)] : index + %expert, %weight_expert_count = index.assume %expert1, %bounded_expert_count [lt(%expert1, %bounded_expert_count)] : index, index + %expert_output_base = index.mul %expert, %bounded_output_size : index + %expert_channel = index.add %expert_output_base, %channel : index + %weight_row_byte_index = index.mul %expert_channel, %weight_row_byte_count : index + %weight_row_byte_base = index.scale %weight_row_byte_index, %c1_byte : index, offset -> offset + %q8_row_base0 = index.mul %bounded_token, %body_route_count : index + %q8_row = index.add %q8_row_base0, %bounded_route : index + %q8_row_byte_index = index.mul %q8_row, %q8_row_byte_count : index + %q8_row_byte_base = index.scale %q8_row_byte_index, %c1_byte : index, offset -> offset + %route_lane_sum = func.call @qwen3_moe_q4k_q8_1_x4_cohort_row_lane(%bounded_input_size, %weight_noalias, %weight_row_byte_base, %q8_noalias, %q8_row_byte_base, %cohort_lane) : (index, buffer, offset, buffer, offset, index) -> (f32) + %route_weight = kernel.subgroup.broadcast %lane_route_weight from %route_i32 : f32, i32 + %weighted = scalar.mulf %route_lane_sum, %route_weight : f32 + scf.yield %weighted : f32 + } else { + scf.yield %c0_f32 : f32 + } + %next = scalar.addf %route_acc, %weighted_lane : f32 + scf.yield %next : f32 + } + scf.yield %sum : f32 + } else { + scf.yield %c0_f32 : f32 + } + %routed_sum = kernel.subgroup.reduce %routed_lane_sum : f32 + %writes_active_channel = scalar.andi %publish_output, %valid_channel : i1 + %writes_output = scalar.andi %writes_active_channel, %is_lane_zero : i1 + scf.if %writes_output { + %bounded_channel, %body_output_size = index.assume %channel, %bounded_output_size [lt(%channel, %bounded_output_size)] : index, index + %residual = view.load %output_view[%bounded_token, %bounded_channel] : view<[%body_token_count]x[%bounded_output_size]xf32> -> f32 + %result = scalar.addf %residual, %routed_sum : f32 + view.store %result, %output_view[%bounded_token, %bounded_channel] : f32, view<[%body_token_count]x[%bounded_output_size]xf32> + } + template.return +} + +kernel.def @qwen3_moe_routed_down_q4k_q8_1_x4(%token_count: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index) { + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %configured_output_size = config.get @qwen3_moe.routed_down.output_size : index + %c8 = index.constant 8 : index + %c7 = index.constant 7 : index + %c1 = index.constant 1 : index + %workgroup_size = index.constant 256 : index + %padded_output_size = index.add %configured_output_size, %c7 : index + %output_tiles = index.div %padded_output_size, %c8 : index + kernel.launch.config workgroups(%output_tiles, %token_capacity, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index, %q8_input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer) { + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %configured_input_size0 = config.get @qwen3_moe.routed_down.input_size : index + %configured_route_count0 = config.get @qwen3_moe.routed_down.route_count : index + %configured_expert_count0 = config.get @qwen3_moe.routed_down.expert_count : index + %configured_output_size0 = config.get @qwen3_moe.routed_down.output_size : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %bounded_input_size, %configured_input_size = index.assume %input_size, %configured_input_size0 [range(%input_size, 256, 32768), mul(%input_size, 256), eq(%input_size, %configured_input_size0)] : index, index + %bounded_route_count, %configured_route_count = index.assume %route_count, %configured_route_count0 [range(%route_count, 1, 8), eq(%route_count, %configured_route_count0)] : index, index + %bounded_route_id_stride = index.assume %route_id_stride [range(%route_id_stride, 1, 512)] : index + %bounded_expert_count, %configured_expert_count = index.assume %expert_count, %configured_expert_count0 [range(%expert_count, 1, 512), eq(%expert_count, %configured_expert_count0)] : index, index + %bounded_output_size, %configured_output_size = index.assume %output_size, %configured_output_size0 [range(%output_size, 1, 4096), eq(%output_size, %configured_output_size0)] : index, index + %token0 = kernel.workgroup.id : index + %c0 = index.constant 0 : index + %valid_token = index.cmp ult, %token0, %bounded_token_count : index + %safe_token0 = scf.select %valid_token, %token0, %c0 : index + %safe_token, %body_token_count = index.assume %safe_token0, %bounded_token_count [lt(%safe_token0, %bounded_token_count)] : index, index + template.apply<@qwen3_moe.routed_down.q4k_q8_1_x4.body>(%valid_token, %body_token_count, %safe_token, %configured_input_size, %configured_route_count, %bounded_route_id_stride, %configured_expert_count, %configured_output_size, %q8_input, %route_ids, %route_weights, %weight, %output) : (i1, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +// Decode-only route that also publishes the normalized Q8_1 x4 row consumed by +// the next projection boundary. +kernel.def target(@qwen3_moe_attention_prepare_gfx11_wave32) @qwen3_moe_routed_down_q4k_q8_1_x4_next_q8(%token_count: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index) { + %configured_output_size = config.get @qwen3_moe.routed_down.output_size : index + %c8 = index.constant 8 : index + %c7 = index.constant 7 : index + %c1 = index.constant 1 : index + %workgroup_size = index.constant 256 : index + %padded_output_size = index.add %configured_output_size, %c7 : index + %output_tiles = index.div %padded_output_size, %c8 : index + kernel.launch.config workgroups(%output_tiles, %c1, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index, %q8_input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer, %norm_weight: buffer, %completion_counter: buffer, %next_q8_output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 1)] : index + %configured_input_size0 = config.get @qwen3_moe.routed_down.input_size : index + %configured_route_count0 = config.get @qwen3_moe.routed_down.route_count : index + %configured_expert_count0 = config.get @qwen3_moe.routed_down.expert_count : index + %configured_output_size0 = config.get @qwen3_moe.routed_down.output_size : index + %bounded_input_size, %configured_input_size = index.assume %input_size, %configured_input_size0 [range(%input_size, 256, 32768), mul(%input_size, 256), eq(%input_size, %configured_input_size0)] : index, index + %bounded_route_count, %configured_route_count = index.assume %route_count, %configured_route_count0 [range(%route_count, 1, 8), eq(%route_count, %configured_route_count0)] : index, index + %bounded_route_id_stride = index.assume %route_id_stride [range(%route_id_stride, 1, 512)] : index + %bounded_expert_count, %configured_expert_count = index.assume %expert_count, %configured_expert_count0 [range(%expert_count, 1, 512), eq(%expert_count, %configured_expert_count0)] : index, index + %bounded_output_size, %configured_output_size = index.assume %output_size, %configured_output_size0 [range(%output_size, 128, 4096), mul(%output_size, 128), eq(%output_size, %configured_output_size0)] : index, index + %c8 = index.constant 8 : index + %publishes_output = scalar.constant true : i1 + %body_token = index.constant 0 : index + template.apply<@qwen3_moe.routed_down.q4k_q8_1_x4.body>(%publishes_output, %bounded_token_count, %body_token, %configured_input_size, %configured_route_count, %bounded_route_id_stride, %configured_expert_count, %configured_output_size, %q8_input, %route_ids, %route_weights, %weight, %output) : (i1, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer) + template.apply<@qwen3_moe.routed_down.next_q8_completion>(%c8, %c8, %bounded_token_count, %configured_output_size, %output, %norm_weight, %completion_counter, %next_q8_output) : (index, index, index, index, buffer, buffer, buffer, buffer) + kernel.return +} + +// The production hidden width requires all 256 output tiles to arrive before +// normalization. Compare both residual and packed Q8 publication with the +// ordinary two-dispatch composition, then reuse the same completion word. +check.case public @qwen3_moe_routed_down_q4k_q8_1_x4_next_q8_differential_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(2) : index + %route_id_stride = check.literal value(4) : index + %expert_count = check.literal value(2) : index + %output_size = check.literal value(2048) : index + %routed_row_count = check.literal value(2) : index + %routed_input = check.generate.fill value(0.00390625) : tensor<2x768xf32> + %q8_input = check.generate.fill value(0) : tensor<2x864xi8> + %route_ids = check.generate.iota offset(0) step(1) period(2) : tensor<1x4xi32> + %route_weights = check.generate.fill value(0.5) : tensor<1x2xf32> + %weight = check.generate.iota offset(-72) step(1) period(144) : tensor<2x2048x3x144xi8> + %norm_weight = check.generate.iota offset(-1.0) step(0.0009765625) : tensor<2048xf32> + %expected_output = check.generate.iota offset(-0.5) step(0.00048828125) period(2048) : tensor<1x2048xf32> + %actual_output0 = check.generate.iota offset(-0.5) step(0.00048828125) period(2048) : tensor<1x2048xf32> + %actual_output1 = check.generate.iota offset(-0.5) step(0.00048828125) period(2048) : tensor<1x2048xf32> + %expected_q8 = check.generate.fill value(0) : tensor<2304xi8> + %actual_q8_0 = check.generate.fill value(1) : tensor<2304xi8> + %actual_q8_1 = check.generate.fill value(1) : tensor<2304xi8> + %completion_counter = check.generate.fill value(0) : tensor<1xi32> + %expected_counter = check.generate.fill value(0) : tensor<1xi32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%routed_row_count, %input_size](%routed_row_count, %input_size, %routed_input, %q8_input) : [index, index](index, index, tensor<2x768xf32>, tensor<2x864xi8>) + kernel.launch @qwen3_moe_routed_down_q4k_q8_1_x4[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %expected_output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<2x864xi8>, tensor<1x4xi32>, tensor<1x2xf32>, tensor<2x2048x3x144xi8>, tensor<1x2048xf32>) + kernel.launch @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4[%token_count](%token_count, %expected_output, %norm_weight, %expected_q8) : [index](index, tensor<1x2048xf32>, tensor<2048xf32>, tensor<2304xi8>) + kernel.launch @qwen3_moe_routed_down_q4k_q8_1_x4_next_q8[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %actual_output0, %norm_weight, %completion_counter, %actual_q8_0) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<2x864xi8>, tensor<1x4xi32>, tensor<1x2xf32>, tensor<2x2048x3x144xi8>, tensor<1x2048xf32>, tensor<2048xf32>, tensor<1xi32>, tensor<2304xi8>) + kernel.launch @qwen3_moe_routed_down_q4k_q8_1_x4_next_q8[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %actual_output1, %norm_weight, %completion_counter, %actual_q8_1) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<2x864xi8>, tensor<1x4xi32>, tensor<1x2xf32>, tensor<2x2048x3x144xi8>, tensor<1x2048xf32>, tensor<2048xf32>, tensor<1xi32>, tensor<2304xi8>) + check.expect.close actual(%actual_output0) expected(%expected_output) atol(9.9999999999999995e-07) rtol(9.9999999999999995e-07) nan(same) : tensor<1x2048xf32> + check.expect.close actual(%actual_output1) expected(%expected_output) atol(9.9999999999999995e-07) rtol(9.9999999999999995e-07) nan(same) : tensor<1x2048xf32> + check.expect.equal actual(%actual_q8_0) expected(%expected_q8) : tensor<2304xi8> + check.expect.equal actual(%actual_q8_1) expected(%expected_q8) : tensor<2304xi8> + check.expect.equal actual(%completion_counter) expected(%expected_counter) : tensor<1xi32> + check.return +} + +check.case public @qwen3_moe_routed_down_q4k_q8_1_x4_next_q8_benchmark_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(8) : index + %route_id_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(2048) : index + %q8_input = check.generate.fill value(0) : tensor<1x8x864xi8> + %route_ids = check.generate.iota offset(0) step(1) period(128) : tensor<1x8xi32> + %route_weights = check.generate.fill value(0.125) : tensor<1x8xf32> + %weight = check.generate.fill value(0) : tensor<128x2048x3x144xi8> + %output = check.generate.fill value(1.0) : tensor<1x2048xf32> + %norm_weight = check.generate.fill value(1.0) : tensor<2048xf32> + %completion_counter = check.generate.fill value(0) : tensor<1xi32> + %next_q8_output = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @qwen3_moe_routed_down_q4k_q8_1_x4_next_q8[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %output, %norm_weight, %completion_counter, %next_q8_output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<1x8x864xi8>, tensor<1x8xi32>, tensor<1x8xf32>, tensor<128x2048x3x144xi8>, tensor<1x2048xf32>, tensor<2048xf32>, tensor<1xi32>, tensor<2304xi8>) + check.return +} + +check.case public @qwen3_moe_routed_down_q4k_q8_1_x4_next_q8_composed_benchmark_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(8) : index + %route_id_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(2048) : index + %q8_input = check.generate.fill value(0) : tensor<1x8x864xi8> + %route_ids = check.generate.iota offset(0) step(1) period(128) : tensor<1x8xi32> + %route_weights = check.generate.fill value(0.125) : tensor<1x8xf32> + %weight = check.generate.fill value(0) : tensor<128x2048x3x144xi8> + %output = check.generate.fill value(1.0) : tensor<1x2048xf32> + %norm_weight = check.generate.fill value(1.0) : tensor<2048xf32> + %next_q8_output = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @qwen3_moe_routed_down_q4k_q8_1_x4[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<1x8x864xi8>, tensor<1x8xi32>, tensor<1x8xf32>, tensor<128x2048x3x144xi8>, tensor<1x2048xf32>) + kernel.launch @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4[%token_count](%token_count, %output, %norm_weight, %next_q8_output) : [index](index, tensor<1x2048xf32>, tensor<2048xf32>, tensor<2304xi8>) + check.return +} + +// Two selected experts exercise the noncompact route-ID stride, normalized +// weighted reduction, in-place residual update, the odd Q4_K block tail at +// K=768, and the output tile tail. Uniform 0xaa bytes decode to a deterministic +// nonzero row. +check.case public @qwen3_moe_routed_down_q4k_q8_1_x4_nonzero_residual_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(2) : index + %route_id_stride = check.literal value(4) : index + %expert_count = check.literal value(2) : index + %output_size = check.literal value(9) : index + %routed_input = check.generate.fill value(0.00390625) : tensor<2x768xf32> + %q8_input = check.generate.fill value(0) : tensor<2x864xi8> + %route_ids = check.generate.iota offset(0) step(1) period(2) : tensor<1x4xi32> + %route_weights = check.generate.fill value(0.5) : tensor<1x2xf32> + %weight = check.generate.fill value(-86) : tensor<2x9x3x144xi8> + %output = check.generate.fill value(1.0) : tensor<1x9xf32> + %expected = check.generate.fill value(-58.0394287109375) : tensor<1x9xf32> + %routed_row_count = check.literal value(2) : index + kernel.launch @ggml_quantize_q8_1_x4_f32[%routed_row_count, %input_size](%routed_row_count, %input_size, %routed_input, %q8_input) : [index, index](index, index, tensor<2x768xf32>, tensor<2x864xi8>) + kernel.launch @qwen3_moe_routed_down_q4k_q8_1_x4[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<2x864xi8>, tensor<1x4xi32>, tensor<1x2xf32>, tensor<2x9x3x144xi8>, tensor<1x9xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.25) rtol(9.9999999999999995e-07) nan(same) : tensor<1x9xf32> + check.return +} + +// Two tokens use different route-weight sums so token, route, packed-row, and +// output addressing cannot collapse to the single-token case. +check.case public @qwen3_moe_routed_down_q4k_q8_1_x4_nonzero_prefill_case { + %token_count = check.literal value(2) : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(2) : index + %route_id_stride = check.literal value(4) : index + %expert_count = check.literal value(2) : index + %output_size = check.literal value(1) : index + %routed_row_count = check.literal value(4) : index + %routed_input = check.generate.fill value(0.00390625) : tensor<2x2x768xf32> + %q8_input = check.generate.fill value(0) : tensor<2x2x864xi8> + %route_ids = check.generate.iota offset(0) step(1) period(2) : tensor<2x4xi32> + %route_weights = check.generate.iota offset(0.25) step(0.25) period(4) : tensor<2x2xf32> + %weight = check.generate.fill value(-86) : tensor<2x1x3x144xi8> + %output = check.generate.fill value(1.0) : tensor<2x1xf32> + %expected = check.generate.iota offset(-43.279571533203125) step(-59.0394287109375) : tensor<2x1xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%routed_row_count, %input_size](%routed_row_count, %input_size, %routed_input, %q8_input) : [index, index](index, index, tensor<2x2x768xf32>, tensor<2x2x864xi8>) + kernel.launch @qwen3_moe_routed_down_q4k_q8_1_x4[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<2x2x864xi8>, tensor<2x4xi32>, tensor<2x2xf32>, tensor<2x1x3x144xi8>, tensor<2x1xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.25) rtol(9.9999999999999995e-07) nan(same) : tensor<2x1xf32> + check.return +} + +check.case public @qwen3_moe_routed_down_q4k_q8_1_x4_benchmark_case { + %token_count = check.param.choice values([1, 2, 4, 8, 16, 17, 32, 63, 128, 129, 512]) name("token_count") : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(8) : index + %route_id_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(2048) : index + %q8_input = check.generate.fill value(0) : tensor<[%token_count]x8x864xi8> + %route_ids = check.generate.iota offset(0) step(1) period(127) : tensor<[%token_count]x128xi32> + %route_weights = check.generate.fill value(0.125) : tensor<[%token_count]x8xf32> + %weight = check.generate.fill value(0) : tensor<128x2048x3x144xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %expected = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + kernel.launch @qwen3_moe_routed_down_q4k_q8_1_x4[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<[%token_count]x8x864xi8>, tensor<[%token_count]x128xi32>, tensor<[%token_count]x8xf32>, tensor<128x2048x3x144xi8>, tensor<[%token_count]x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x2048xf32> + check.return +} + +check.case public @qwen3_moe_routed_down_q4k_q8_1_x4_pipeline_benchmark_case { + %token_count = check.param.choice values([1, 2, 4, 8, 16, 17, 32, 63, 128, 129, 512]) name("token_count") : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(8) : index + %route_id_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(2048) : index + // Each 768-element route contains six complete Q8_1 x4 groups, so packing + // the eight contiguous routes as one 6,144-element token row preserves the + // exact per-route physical layout without a derived testbench scalar. + %routed_input_size = check.literal value(6144) : index + %routed_input = check.generate.fill value(0.0) : tensor<[%token_count]x8x768xf32> + %q8_input = check.generate.fill value(1) : tensor<[%token_count]x8x864xi8> + %route_ids = check.generate.iota offset(0) step(1) period(127) : tensor<[%token_count]x128xi32> + %route_weights = check.generate.fill value(0.125) : tensor<[%token_count]x8xf32> + %weight = check.generate.fill value(0) : tensor<128x2048x3x144xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %expected = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %routed_input_size](%token_count, %routed_input_size, %routed_input, %q8_input) : [index, index](index, index, tensor<[%token_count]x8x768xf32>, tensor<[%token_count]x8x864xi8>) + kernel.launch @qwen3_moe_routed_down_q4k_q8_1_x4[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<[%token_count]x8x864xi8>, tensor<[%token_count]x128xi32>, tensor<[%token_count]x8xf32>, tensor<128x2048x3x144xi8>, tensor<[%token_count]x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x2048xf32> + check.return +} + +check.benchmark<@qwen3_moe_routed_down_q4k_q8_1_x4_nonzero_residual_case> @qwen3_moe_routed_down_q4k_q8_1_x4_small + +check.benchmark<@qwen3_moe_routed_down_q4k_q8_1_x4_benchmark_case> @qwen3_moe_routed_down_q4k_q8_1_x4_decode {token_count = 1} + +check.benchmark<@qwen3_moe_routed_down_q4k_q8_1_x4_benchmark_case> @qwen3_moe_routed_down_q4k_q8_1_x4_small_batch_2 {token_count = 2} + +check.benchmark<@qwen3_moe_routed_down_q4k_q8_1_x4_benchmark_case> @qwen3_moe_routed_down_q4k_q8_1_x4_small_batch_4 {token_count = 4} + +check.benchmark<@qwen3_moe_routed_down_q4k_q8_1_x4_benchmark_case> @qwen3_moe_routed_down_q4k_q8_1_x4_small_batch_8 {token_count = 8} + +check.benchmark<@qwen3_moe_routed_down_q4k_q8_1_x4_benchmark_case> @qwen3_moe_routed_down_q4k_q8_1_x4_small_batch_16 {token_count = 16} + +check.benchmark<@qwen3_moe_routed_down_q4k_q8_1_x4_benchmark_case> @qwen3_moe_routed_down_q4k_q8_1_x4_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_routed_down_q4k_q8_1_x4_benchmark_case> @qwen3_moe_routed_down_q4k_q8_1_x4_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_routed_down_q4k_q8_1_x4_benchmark_case> @qwen3_moe_routed_down_q4k_q8_1_x4_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_routed_down_q4k_q8_1_x4_pipeline_benchmark_case> @qwen3_moe_routed_down_q4k_q8_1_x4_pipeline_decode {token_count = 1} + +check.benchmark<@qwen3_moe_routed_down_q4k_q8_1_x4_pipeline_benchmark_case> @qwen3_moe_routed_down_q4k_q8_1_x4_pipeline_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_routed_down_q4k_q8_1_x4_pipeline_benchmark_case> @qwen3_moe_routed_down_q4k_q8_1_x4_pipeline_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_routed_down_q4k_q8_1_x4_pipeline_benchmark_case> @qwen3_moe_routed_down_q4k_q8_1_x4_pipeline_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_routed_down_q4k_q8_1_x4_next_q8_benchmark_case> @qwen3_moe_routed_down_q4k_q8_1_x4_next_q8_decode + +check.benchmark<@qwen3_moe_routed_down_q4k_q8_1_x4_next_q8_composed_benchmark_case> @qwen3_moe_routed_down_q4k_q8_1_x4_next_q8_composed_decode diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_q6k.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_q6k.loom new file mode 100644 index 000000000000..1bd11cdcf6a1 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_q6k.loom @@ -0,0 +1,747 @@ +// Fuses the Qwen3 MoE Q6_K down projection with route weighting, top-8 +// reduction, and residual publication. The input contains one Q8_1 x4 row per +// [token, route], while weights remain in GGUF's raw +// [expert, output, K / 256, 210 bytes] layout. +// +// Route IDs retain an independent physical stride because llama.cpp selects a +// top-8 view from its 128-entry argsort storage. Normalized route weights are +// compact [token, route]. The output enters containing the residual and is +// updated in place, so the fused boundary never materializes an unweighted or +// route-indexed [token, route, hidden] down-projection tensor. +// +// The Q8_1 provider assigns one output channel to a wave and reduces all routes +// in registers. The direct-F32 provider mirrors llama.cpp's gfx1151 Vulkan +// decode algorithm: four wave64 subgroups compute four channels, each subgroup +// contracts four raw Q6_K blocks at a time, and route weighting and residual +// publication remain fused. A grouped prefill provider can reuse the same raw +// Q6_K row primitive while staging each expert row across multiple tokens. +template.decl @qwen3_moe.rmsnorm_quantize_q8_1_x4.body(%arg0: i1, %arg1: index, %arg2: index, %arg3: index, %arg4: buffer, %arg5: buffer, %arg6: buffer, %arg7: buffer) + +template.decl @qwen3_moe.routed_down.next_q8_completion(%arg0: index, %arg1: index, %arg2: index, %arg3: index, %arg4: buffer, %arg5: buffer, %arg6: buffer, %arg7: buffer) + +template.decl @qwen3_moe.routed_down.q6k_f32.body(%publish_output: i1, %subgroup_count: index, %block_step: index, %scale_stage_bytes: offset, %token_count: index, %token: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index, %input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer) + +template.decl @qwen3_moe.routed_down.q6k_f32.pair_body(%token_count: index, %token: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index, %input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer) + +template.decl @qwen3_moe.routed_down.q6k_q8_1_x4.body(%publish_output: i1, %token_count: index, %token: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index, %q8_input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer) + +amdgpu.target @qwen3_moe_routed_down_q6k_gfx11_wave64 {subgroup_size = 64} + +target.decl @qwen3_moe_attention_prepare_gfx11_wave32 + +config.decl @qwen3_moe.routed_down.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +config.decl @qwen3_moe.routed_down.route_count : %value: index where [range(%value, 1, 8)] + +config.decl @qwen3_moe.routed_down.expert_count : %value: index where [range(%value, 1, 512)] + +config.decl @qwen3_moe.routed_down.output_size : %value: index where [range(%value, 1, 4096)] + +config.decl @qwen3_moe.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +kernel.decl @ggml_quantize_q8_1_x4_f32(%token_count$62: index, %input_size$63: index) launch(%token_count$64: index, %input_size$65: index, %input: buffer, %output: buffer) + +func.decl @ggml_q6k_q8_1_x4_row_lane(%input_size: index, %weight: buffer, %weight_row_byte_base: offset, %q8_input: buffer, %q8_row_byte_base: offset, %lane: index) -> (f32) + +func.decl @ggml_q6k_stage_f32_scales(%input_size: index, %row: index, %block: index, %frame_count: index, %frame: index, %lane: index, %weight: buffer, %scale_stage: buffer) + +func.decl @ggml_q6k_load_f32_block(%token_count: index, %input_size: index, %token: index, %block: index, %lane: index, %input: buffer) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) + +func.decl @ggml_q6k_f32_block_row(%input_size: index, %row: index, %block: index, %frame_count: index, %frame: index, %lane: index, %weight: buffer, %scale_stage: buffer, %input0: vector<4xf32>, %input1: vector<4xf32>, %input2: vector<4xf32>, %input3: vector<4xf32>) -> (f32) + +kernel.decl @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4(%token_count$106: index) launch(%token_count$107: index, %input: buffer, %weight: buffer, %q8_output: buffer) + +// Shared Q6_K contraction and residual publication used by ordinary and +// completion-fused exports with identical launch geometry. +template.def<@qwen3_moe.routed_down.q6k_q8_1_x4.body> device @qwen3_moe_routed_down_q6k_q8_1_x4_body(%publish_output: i1, %token_count: index, %token: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index, %q8_input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_token, %body_token_count = index.assume %token, %bounded_token_count [lt(%token, %bounded_token_count)] : index, index + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %bounded_route_count = index.assume %route_count [range(%route_count, 1, 8)] : index + %bounded_route_id_stride = index.assume %route_id_stride [range(%route_id_stride, 1, 512)] : index + %bounded_expert_count = index.assume %expert_count [range(%expert_count, 1, 512)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 262144)] : index + %channel_tile = kernel.workgroup.id : index + %subgroup0 = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c8 = index.constant 8 : index + %c128 = index.constant 128 : index + %c144_bytes = index.constant 144 : offset + %c210_bytes = index.constant 210 : offset + %c256 = index.constant 256 : index + %c0_i32 = scalar.constant 0 : i32 + %c0_f32 = scalar.constant 0.0 : f32 + %c0_offset = index.constant 0 : offset + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %channel_base = index.mul %channel_tile, %c8 : index + %channel = index.add %channel_base, %subgroup : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %lane_i32 = index.cast %lane : index to i32 + %is_lane_zero = scalar.cmpi eq, %lane_i32, %c0_i32 : i32 + %q6_block_count = index.div %bounded_input_size, %c256 : index + %weight_row_bytes = index.scale %q6_block_count, %c210_bytes : index, offset -> offset + %q8_group_count = index.div %bounded_input_size, %c128 : index + %q8_row_bytes = index.scale %q8_group_count, %c144_bytes : index, offset -> offset + %q8_noalias, %route_id_noalias, %route_weight_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %q8_input, %route_ids, %route_weights, %weight, %output : buffer, buffer, buffer, buffer, buffer + %route_id_view = buffer.view %route_id_noalias[%c0_offset] : buffer -> view<[%body_token_count]x[%bounded_route_id_stride]xi32> + %route_weight_view = buffer.view %route_weight_noalias[%c0_offset] : buffer -> view<[%body_token_count]x[%bounded_route_count]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%body_token_count]x[%bounded_output_size]xf32> + %lane_has_route = index.cmp ult, %lane, %bounded_route_count : index + %lane_within_route_stride = index.cmp ult, %lane, %bounded_route_id_stride : index + %loads_active_route = scalar.andi %publish_output, %lane_has_route : i1 + %loads_route_metadata = scalar.andi %loads_active_route, %lane_within_route_stride : i1 + %lane_expert_i32, %lane_route_weight = scf.if %loads_route_metadata -> (i32, f32) { + %route_lane0, %metadata_route_count = index.assume %lane, %bounded_route_count [lt(%lane, %bounded_route_count)] : index, index + %route_lane, %metadata_route_id_stride = index.assume %route_lane0, %bounded_route_id_stride [lt(%route_lane0, %bounded_route_id_stride)] : index, index + %expert_i32 = view.load %route_id_view[%bounded_token, %route_lane] : view<[%body_token_count]x[%bounded_route_id_stride]xi32> -> i32 + %route_weight = view.load %route_weight_view[%bounded_token, %route_lane0] : view<[%body_token_count]x[%bounded_route_count]xf32> -> f32 + scf.yield %expert_i32, %route_weight : i32, f32 + } else { + scf.yield %c0_i32, %c0_f32 : i32, f32 + } + %computes_channel = scalar.andi %publish_output, %valid_channel : i1 + %routed_lane_sum = scf.if %computes_channel -> (f32) { + %sum = scf.for %route = [%c0 to %bounded_route_count step %c1](%route_acc = %c0_f32 : f32) -> (f32) unroll { + %route_i32 = index.cast %route : index to i32 + %expert_i32 = kernel.subgroup.broadcast %lane_expert_i32 from %route_i32 : i32, i32 + %expert0 = index.cast %expert_i32 : i32 to index + %expert1 = index.assume %expert0 [range(%expert0, 0, 511)] : index + %expert, %weight_expert_count = index.assume %expert1, %bounded_expert_count [lt(%expert1, %bounded_expert_count)] : index, index + %expert_output_base = index.mul %expert, %bounded_output_size : index + %expert_channel = index.add %expert_output_base, %channel : index + %weight_row_byte_base = index.scale %expert_channel, %weight_row_bytes : index, offset -> offset + %q8_row_base0 = index.mul %bounded_token, %bounded_route_count : index + %q8_row = index.add %q8_row_base0, %route : index + %q8_row_byte_base = index.scale %q8_row, %q8_row_bytes : index, offset -> offset + %lane_acc = func.call @ggml_q6k_q8_1_x4_row_lane(%bounded_input_size, %weight_noalias, %weight_row_byte_base, %q8_noalias, %q8_row_byte_base, %lane) : (index, buffer, offset, buffer, offset, index) -> (f32) + %route_weight = kernel.subgroup.broadcast %lane_route_weight from %route_i32 : f32, i32 + %weighted_lane = scalar.mulf %lane_acc, %route_weight : f32 + %next = scalar.addf %route_acc, %weighted_lane : f32 + scf.yield %next : f32 + } + scf.yield %sum : f32 + } else { + scf.yield %c0_f32 : f32 + } + %routed_sum = kernel.subgroup.reduce %routed_lane_sum : f32 + %writes_active_channel = scalar.andi %publish_output, %valid_channel : i1 + %writes_output = scalar.andi %writes_active_channel, %is_lane_zero : i1 + scf.if %writes_output { + %bounded_channel, %body_output_size = index.assume %channel, %bounded_output_size [lt(%channel, %bounded_output_size)] : index, index + %residual = view.load %output_view[%bounded_token, %bounded_channel] : view<[%body_token_count]x[%bounded_output_size]xf32> -> f32 + %result = scalar.addf %residual, %routed_sum : f32 + view.store %result, %output_view[%bounded_token, %bounded_channel] : f32, view<[%body_token_count]x[%bounded_output_size]xf32> + } + template.return +} + +kernel.def @qwen3_moe_routed_down_q6k_q8_1_x4(%token_count: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index) { + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %configured_output_size = config.get @qwen3_moe.routed_down.output_size : index + %c8 = index.constant 8 : index + %c7 = index.constant 7 : index + %c1 = index.constant 1 : index + %workgroup_size = index.constant 256 : index + %padded_output_size = index.add %configured_output_size, %c7 : index + %output_tiles = index.div %padded_output_size, %c8 : index + kernel.launch.config workgroups(%output_tiles, %token_capacity, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index, %q8_input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer) { + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %configured_input_size0 = config.get @qwen3_moe.routed_down.input_size : index + %configured_route_count0 = config.get @qwen3_moe.routed_down.route_count : index + %configured_expert_count0 = config.get @qwen3_moe.routed_down.expert_count : index + %configured_output_size0 = config.get @qwen3_moe.routed_down.output_size : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %bounded_input_size, %configured_input_size = index.assume %input_size, %configured_input_size0 [range(%input_size, 256, 32768), mul(%input_size, 256), eq(%input_size, %configured_input_size0)] : index, index + %bounded_route_count, %configured_route_count = index.assume %route_count, %configured_route_count0 [range(%route_count, 1, 8), eq(%route_count, %configured_route_count0)] : index, index + %bounded_route_id_stride = index.assume %route_id_stride [range(%route_id_stride, 1, 512)] : index + %bounded_expert_count, %configured_expert_count = index.assume %expert_count, %configured_expert_count0 [range(%expert_count, 1, 512), eq(%expert_count, %configured_expert_count0)] : index, index + %bounded_output_size, %configured_output_size = index.assume %output_size, %configured_output_size0 [range(%output_size, 1, 4096), eq(%output_size, %configured_output_size0)] : index, index + %token0 = kernel.workgroup.id : index + %c0 = index.constant 0 : index + %valid_token = index.cmp ult, %token0, %bounded_token_count : index + %safe_token0 = scf.select %valid_token, %token0, %c0 : index + %safe_token, %body_token_count = index.assume %safe_token0, %bounded_token_count [lt(%safe_token0, %bounded_token_count)] : index, index + template.apply<@qwen3_moe.routed_down.q6k_q8_1_x4.body>(%valid_token, %body_token_count, %safe_token, %configured_input_size, %configured_route_count, %bounded_route_id_stride, %configured_expert_count, %configured_output_size, %q8_input, %route_ids, %route_weights, %weight, %output) : (i1, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +// Decode-only route that also publishes the normalized Q8_1 x4 row consumed by +// the next projection boundary. +kernel.def target(@qwen3_moe_attention_prepare_gfx11_wave32) @qwen3_moe_routed_down_q6k_q8_1_x4_next_q8(%token_count: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index) { + %configured_output_size = config.get @qwen3_moe.routed_down.output_size : index + %c8 = index.constant 8 : index + %c7 = index.constant 7 : index + %c1 = index.constant 1 : index + %workgroup_size = index.constant 256 : index + %padded_output_size = index.add %configured_output_size, %c7 : index + %output_tiles = index.div %padded_output_size, %c8 : index + kernel.launch.config workgroups(%output_tiles, %c1, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index, %q8_input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer, %norm_weight: buffer, %completion_counter: buffer, %next_q8_output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 1)] : index + %configured_input_size0 = config.get @qwen3_moe.routed_down.input_size : index + %configured_route_count0 = config.get @qwen3_moe.routed_down.route_count : index + %configured_expert_count0 = config.get @qwen3_moe.routed_down.expert_count : index + %configured_output_size0 = config.get @qwen3_moe.routed_down.output_size : index + %bounded_input_size, %configured_input_size = index.assume %input_size, %configured_input_size0 [range(%input_size, 256, 32768), mul(%input_size, 256), eq(%input_size, %configured_input_size0)] : index, index + %bounded_route_count, %configured_route_count = index.assume %route_count, %configured_route_count0 [range(%route_count, 1, 8), eq(%route_count, %configured_route_count0)] : index, index + %bounded_route_id_stride = index.assume %route_id_stride [range(%route_id_stride, 1, 512)] : index + %bounded_expert_count, %configured_expert_count = index.assume %expert_count, %configured_expert_count0 [range(%expert_count, 1, 512), eq(%expert_count, %configured_expert_count0)] : index, index + %bounded_output_size, %configured_output_size = index.assume %output_size, %configured_output_size0 [range(%output_size, 128, 4096), mul(%output_size, 128), eq(%output_size, %configured_output_size0)] : index, index + %c8 = index.constant 8 : index + %publishes_output = scalar.constant true : i1 + %body_token = index.constant 0 : index + template.apply<@qwen3_moe.routed_down.q6k_q8_1_x4.body>(%publishes_output, %bounded_token_count, %body_token, %configured_input_size, %configured_route_count, %bounded_route_id_stride, %configured_expert_count, %configured_output_size, %q8_input, %route_ids, %route_weights, %weight, %output) : (i1, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer) + template.apply<@qwen3_moe.routed_down.next_q8_completion>(%c8, %c8, %bounded_token_count, %configured_output_size, %output, %norm_weight, %completion_counter, %next_q8_output) : (index, index, index, index, buffer, buffer, buffer, buffer) + kernel.return +} + +// The production hidden width requires all 256 output tiles to arrive before +// normalization. Compare both residual and packed Q8 publication with the +// ordinary two-dispatch composition, then reuse the same completion word. +check.case public @qwen3_moe_routed_down_q6k_q8_1_x4_next_q8_differential_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(2) : index + %route_id_stride = check.literal value(4) : index + %expert_count = check.literal value(2) : index + %output_size = check.literal value(2048) : index + %routed_row_count = check.literal value(2) : index + %routed_input = check.generate.fill value(0.00390625) : tensor<2x768xf32> + %q8_input = check.generate.fill value(0) : tensor<2x864xi8> + %route_ids = check.generate.iota offset(0) step(1) period(2) : tensor<1x4xi32> + %route_weights = check.generate.fill value(0.5) : tensor<1x2xf32> + %weight = check.generate.fill value(-86) : tensor<2x2048x3x210xi8> + %norm_weight = check.generate.iota offset(-1.0) step(0.0009765625) : tensor<2048xf32> + %expected_output = check.generate.iota offset(-0.5) step(0.00048828125) period(2048) : tensor<1x2048xf32> + %actual_output0 = check.generate.iota offset(-0.5) step(0.00048828125) period(2048) : tensor<1x2048xf32> + %actual_output1 = check.generate.iota offset(-0.5) step(0.00048828125) period(2048) : tensor<1x2048xf32> + %expected_q8 = check.generate.fill value(0) : tensor<2304xi8> + %actual_q8_0 = check.generate.fill value(1) : tensor<2304xi8> + %actual_q8_1 = check.generate.fill value(1) : tensor<2304xi8> + %completion_counter = check.generate.fill value(0) : tensor<1xi32> + %expected_counter = check.generate.fill value(0) : tensor<1xi32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%routed_row_count, %input_size](%routed_row_count, %input_size, %routed_input, %q8_input) : [index, index](index, index, tensor<2x768xf32>, tensor<2x864xi8>) + kernel.launch @qwen3_moe_routed_down_q6k_q8_1_x4[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %expected_output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<2x864xi8>, tensor<1x4xi32>, tensor<1x2xf32>, tensor<2x2048x3x210xi8>, tensor<1x2048xf32>) + kernel.launch @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4[%token_count](%token_count, %expected_output, %norm_weight, %expected_q8) : [index](index, tensor<1x2048xf32>, tensor<2048xf32>, tensor<2304xi8>) + kernel.launch @qwen3_moe_routed_down_q6k_q8_1_x4_next_q8[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %actual_output0, %norm_weight, %completion_counter, %actual_q8_0) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<2x864xi8>, tensor<1x4xi32>, tensor<1x2xf32>, tensor<2x2048x3x210xi8>, tensor<1x2048xf32>, tensor<2048xf32>, tensor<1xi32>, tensor<2304xi8>) + kernel.launch @qwen3_moe_routed_down_q6k_q8_1_x4_next_q8[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %actual_output1, %norm_weight, %completion_counter, %actual_q8_1) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<2x864xi8>, tensor<1x4xi32>, tensor<1x2xf32>, tensor<2x2048x3x210xi8>, tensor<1x2048xf32>, tensor<2048xf32>, tensor<1xi32>, tensor<2304xi8>) + check.expect.close actual(%actual_output0) expected(%expected_output) atol(0.0) rtol(0.0) nan(same) : tensor<1x2048xf32> + check.expect.close actual(%actual_output1) expected(%expected_output) atol(0.0) rtol(0.0) nan(same) : tensor<1x2048xf32> + check.expect.equal actual(%actual_q8_0) expected(%expected_q8) : tensor<2304xi8> + check.expect.equal actual(%actual_q8_1) expected(%expected_q8) : tensor<2304xi8> + check.expect.equal actual(%completion_counter) expected(%expected_counter) : tensor<1xi32> + check.return +} + +check.case public @qwen3_moe_routed_down_q6k_q8_1_x4_next_q8_benchmark_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(8) : index + %route_id_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(2048) : index + %q8_input = check.generate.fill value(0) : tensor<1x8x864xi8> + %route_ids = check.generate.iota offset(0) step(1) period(128) : tensor<1x8xi32> + %route_weights = check.generate.fill value(0.125) : tensor<1x8xf32> + %weight = check.generate.fill value(0) : tensor<128x2048x3x210xi8> + %output = check.generate.fill value(1.0) : tensor<1x2048xf32> + %norm_weight = check.generate.fill value(1.0) : tensor<2048xf32> + %completion_counter = check.generate.fill value(0) : tensor<1xi32> + %next_q8_output = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @qwen3_moe_routed_down_q6k_q8_1_x4_next_q8[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %output, %norm_weight, %completion_counter, %next_q8_output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<1x8x864xi8>, tensor<1x8xi32>, tensor<1x8xf32>, tensor<128x2048x3x210xi8>, tensor<1x2048xf32>, tensor<2048xf32>, tensor<1xi32>, tensor<2304xi8>) + check.return +} + +check.case public @qwen3_moe_routed_down_q6k_q8_1_x4_next_q8_composed_benchmark_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(8) : index + %route_id_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(2048) : index + %q8_input = check.generate.fill value(0) : tensor<1x8x864xi8> + %route_ids = check.generate.iota offset(0) step(1) period(128) : tensor<1x8xi32> + %route_weights = check.generate.fill value(0.125) : tensor<1x8xf32> + %weight = check.generate.fill value(0) : tensor<128x2048x3x210xi8> + %output = check.generate.fill value(1.0) : tensor<1x2048xf32> + %norm_weight = check.generate.fill value(1.0) : tensor<2048xf32> + %next_q8_output = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @qwen3_moe_routed_down_q6k_q8_1_x4[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<1x8x864xi8>, tensor<1x8xi32>, tensor<1x8xf32>, tensor<128x2048x3x210xi8>, tensor<1x2048xf32>) + kernel.launch @qwen3_moe_attention_rmsnorm_quantize_q8_1_x4[%token_count](%token_count, %output, %norm_weight, %next_q8_output) : [index](index, tensor<1x2048xf32>, tensor<2048xf32>, tensor<2304xi8>) + check.return +} + +// Shared direct-F32 decode schedule. Every subgroup owns one scale frame and +// one output channel. Providers specialize how many Q6_K blocks a subgroup +// covers per pass from the target subgroup width. +template.def<@qwen3_moe.routed_down.q6k_f32.body> device @qwen3_moe_routed_down_q6k_f32_body(%publish_output: i1, %subgroup_count: index, %block_step: index, %scale_stage_bytes: offset, %token_count: index, %token: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index, %input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_token, %body_token_count = index.assume %token, %bounded_token_count [lt(%token, %bounded_token_count)] : index, index + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %bounded_route_count = index.assume %route_count [range(%route_count, 1, 8)] : index + %bounded_route_id_stride = index.assume %route_id_stride [range(%route_id_stride, 1, 512)] : index + %bounded_expert_count = index.assume %expert_count [range(%expert_count, 1, 512)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 262144)] : index + %channel_tile = kernel.workgroup.id : index + %subgroup0 = kernel.subgroup.id : index + %lane0 = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c16 = index.constant 16 : index + %c210 = index.constant 210 : index + %c256 = index.constant 256 : index + %c0_i32 = scalar.constant 0 : i32 + %c0_f32 = scalar.constant 0.0 : f32 + %c0_offset = index.constant 0 : offset + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %lane = index.assume %lane0 [range(%lane0, 0, 63)] : index + %channel_tile_base = index.mul %channel_tile, %subgroup_count : index + %channel = index.add %channel_tile_base, %subgroup : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %safe_channel = scf.select %valid_channel, %channel, %c0 : index + %lane_i32 = index.cast %lane : index to i32 + %is_lane_zero = scalar.cmpi eq, %lane_i32, %c0_i32 : i32 + %assignment_count = index.mul %body_token_count, %bounded_route_count : index + %q6_block_count = index.div %bounded_input_size, %c256 : index + %weight_row_byte_count = index.mul %q6_block_count, %c210 : index + %cohort = index.div %lane, %c16 : index + %input_noalias, %route_id_noalias, %route_weight_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input, %route_ids, %route_weights, %weight, %output : buffer, buffer, buffer, buffer, buffer + %route_id_view = buffer.view %route_id_noalias[%c0_offset] : buffer -> view<[%body_token_count]x[%bounded_route_id_stride]xi32> + %route_weight_view = buffer.view %route_weight_noalias[%c0_offset] : buffer -> view<[%body_token_count]x[%bounded_route_count]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%body_token_count]x[%bounded_output_size]xf32> + %scale_stage = buffer.alloca align(16) %scale_stage_bytes : buffer + %routed_sum0 = scf.if %publish_output -> (f32) { + %active_sum = scf.for %route = [%c0 to %bounded_route_count step %c1](%route_acc = %c0_f32 : f32) -> (f32) unroll { + // Route metadata is uniform across all channels. Keeping the access in + // the unrolled route body exposes that uniformity directly to target + // lowering instead of materializing a dynamic subgroup broadcast. + %bounded_route, %body_route_count = index.assume %route, %bounded_route_count [lt(%route, %bounded_route_count)] : index, index + %route_for_stride, %body_route_id_stride = index.assume %bounded_route, %bounded_route_id_stride [lt(%bounded_route, %bounded_route_id_stride)] : index, index + %expert_i32 = view.load %route_id_view[%bounded_token, %route_for_stride] : view<[%body_token_count]x[%bounded_route_id_stride]xi32> -> i32 + %route_weight = view.load %route_weight_view[%bounded_token, %bounded_route] : view<[%body_token_count]x[%bounded_route_count]xf32> -> f32 + %expert0 = index.cast %expert_i32 : i32 to index + %expert1 = index.assume %expert0 [range(%expert0, 0, 511)] : index + %expert, %weight_expert_count = index.assume %expert1, %bounded_expert_count [lt(%expert1, %bounded_expert_count)] : index, index + %expert_output_base = index.mul %expert, %bounded_output_size : index + %expert_channel = index.add %expert_output_base, %safe_channel : index + %input_row_base = index.mul %bounded_token, %bounded_route_count : index + %input_row = index.add %input_row_base, %bounded_route : index + %lane_sum = scf.for %block_base = [%c0 to %q6_block_count step %block_step](%block_acc = %c0_f32 : f32) -> (f32) { + %block = index.add %block_base, %cohort : index + func.call @ggml_q6k_stage_f32_scales(%bounded_input_size, %expert_channel, %block, %subgroup_count, %subgroup, %lane, %weight_noalias, %scale_stage) : (index, index, index, index, index, index, buffer, buffer) + %input0, %input1, %input2, %input3 = func.call @ggml_q6k_load_f32_block(%assignment_count, %bounded_input_size, %input_row, %block, %lane, %input_noalias) : (index, index, index, index, index, buffer) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) + %contribution = func.call @ggml_q6k_f32_block_row(%bounded_input_size, %expert_channel, %block, %subgroup_count, %subgroup, %lane, %weight_noalias, %scale_stage, %input0, %input1, %input2, %input3) : (index, index, index, index, index, index, buffer, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) -> (f32) + %next = scalar.addf %block_acc, %contribution : f32 + scf.yield %next : f32 + } + %route_sum = kernel.subgroup.reduce %lane_sum : f32 + %weighted = scalar.mulf %route_sum, %route_weight : f32 + %next = scalar.addf %route_acc, %weighted : f32 + scf.yield %next : f32 + } + scf.yield %active_sum : f32 + } else { + scf.yield %c0_f32 : f32 + } + %routed_sum = scf.select %valid_channel, %routed_sum0, %c0_f32 : f32 + %writes_active_channel = scalar.andi %publish_output, %valid_channel : i1 + %writes_output = scalar.andi %writes_active_channel, %is_lane_zero : i1 + scf.if %writes_output { + %bounded_channel, %body_output_size = index.assume %channel, %bounded_output_size [lt(%channel, %bounded_output_size)] : index, index + %residual = view.load %output_view[%bounded_token, %bounded_channel] : view<[%body_token_count]x[%bounded_output_size]xf32> -> f32 + %result = scalar.addf %residual, %routed_sum : f32 + view.store %result, %output_view[%bounded_token, %bounded_channel] : f32, view<[%body_token_count]x[%bounded_output_size]xf32> + } + template.return +} + +// Decode schedule matching the Vulkan oracle's two adjacent output rows per +// wave. Four wave64 subgroups publish eight channels per workgroup. Keeping the +// route loop rolled prevents its eight iterations from multiplying the two-row +// contraction's register pressure and instruction footprint. +template.def<@qwen3_moe.routed_down.q6k_f32.pair_body> device @qwen3_moe_routed_down_q6k_f32_pair_body(%token_count: index, %token: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index, %input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 1)] : index + %bounded_token, %body_token_count = index.assume %token, %bounded_token_count [lt(%token, %bounded_token_count)] : index, index + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %bounded_route_count = index.assume %route_count [range(%route_count, 1, 8)] : index + %bounded_route_id_stride = index.assume %route_id_stride [range(%route_id_stride, 1, 512)] : index + %bounded_expert_count = index.assume %expert_count [range(%expert_count, 1, 512)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 2, 4096)] : index + %pair_tile = kernel.workgroup.id : index + %subgroup0 = kernel.subgroup.id : index + %lane0 = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c256 = index.constant 256 : index + %c0_f32 = scalar.constant 0.0 : f32 + %c0_offset = index.constant 0 : offset + %scale_stage_bytes = index.constant 2048 : offset + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 3)] : index + %lane = index.assume %lane0 [range(%lane0, 0, 63)] : index + %pair_tile_base = index.mul %pair_tile, %c4 : index + %pair = index.add %pair_tile_base, %subgroup : index + %row00 = index.mul %pair, %c2 : index + %row0, %body_output_size = index.assume %row00, %bounded_output_size [lt(%row00, %bounded_output_size)] : index, index + %row1 = index.add %row0, %c1 : index + %row1_valid = index.cmp ult, %row1, %body_output_size : index + %safe_row1 = scf.select %row1_valid, %row1, %c0 : index + %assignment_count = index.mul %body_token_count, %bounded_route_count : index + %q6_block_count = index.div %bounded_input_size, %c256 : index + %cohort = index.div %lane, %c16 : index + %scale_frame0 = index.mul %subgroup, %c2 : index + %scale_frame1 = index.add %scale_frame0, %c1 : index + %input_noalias, %route_id_noalias, %route_weight_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input, %route_ids, %route_weights, %weight, %output : buffer, buffer, buffer, buffer, buffer + %route_id_view = buffer.view %route_id_noalias[%c0_offset] : buffer -> view<[%body_token_count]x[%bounded_route_id_stride]xi32> + %route_weight_view = buffer.view %route_weight_noalias[%c0_offset] : buffer -> view<[%body_token_count]x[%bounded_route_count]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%body_token_count]x[%body_output_size]xf32> + %scale_stage = buffer.alloca align(16) %scale_stage_bytes : buffer + %routed_sum0, %routed_sum1 = scf.for %route = [%c0 to %bounded_route_count step %c1](%route_acc0 = %c0_f32 : f32, %route_acc1 = %c0_f32 : f32) -> (f32, f32) { + %bounded_route, %body_route_count = index.assume %route, %bounded_route_count [lt(%route, %bounded_route_count)] : index, index + %route_for_stride, %body_route_id_stride = index.assume %bounded_route, %bounded_route_id_stride [lt(%bounded_route, %bounded_route_id_stride)] : index, index + %expert_i32 = view.load %route_id_view[%bounded_token, %route_for_stride] : view<[%body_token_count]x[%bounded_route_id_stride]xi32> -> i32 + %route_weight = view.load %route_weight_view[%bounded_token, %bounded_route] : view<[%body_token_count]x[%bounded_route_count]xf32> -> f32 + %expert0 = index.cast %expert_i32 : i32 to index + %expert1 = index.assume %expert0 [range(%expert0, 0, 511)] : index + %expert, %body_expert_count = index.assume %expert1, %bounded_expert_count [lt(%expert1, %bounded_expert_count)] : index, index + %expert_output_base = index.mul %expert, %body_output_size : index + %expert_row0 = index.add %expert_output_base, %row0 : index + %expert_row1 = index.add %expert_output_base, %safe_row1 : index + %input_row_base = index.mul %bounded_token, %body_route_count : index + %input_row = index.add %input_row_base, %bounded_route : index + %lane_sum0, %lane_sum1 = scf.for %block_base = [%c0 to %q6_block_count step %c4](%block_acc0 = %c0_f32 : f32, %block_acc1 = %c0_f32 : f32) -> (f32, f32) { + %block = index.add %block_base, %cohort : index + func.call @ggml_q6k_stage_f32_scales(%bounded_input_size, %expert_row0, %block, %c8, %scale_frame0, %lane, %weight_noalias, %scale_stage) : (index, index, index, index, index, index, buffer, buffer) + %input00, %input01, %input02, %input03 = func.call @ggml_q6k_load_f32_block(%assignment_count, %bounded_input_size, %input_row, %block, %lane, %input_noalias) : (index, index, index, index, index, buffer) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) + %contribution0 = func.call @ggml_q6k_f32_block_row(%bounded_input_size, %expert_row0, %block, %c8, %scale_frame0, %lane, %weight_noalias, %scale_stage, %input00, %input01, %input02, %input03) : (index, index, index, index, index, index, buffer, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) -> (f32) + %contribution1 = scf.if %row1_valid -> (f32) { + func.call @ggml_q6k_stage_f32_scales(%bounded_input_size, %expert_row1, %block, %c8, %scale_frame1, %lane, %weight_noalias, %scale_stage) : (index, index, index, index, index, index, buffer, buffer) + // Reloading the shared activations lets the first row's vectors die + // before the second scale stage. Carrying them across that stage + // increases register pressure and is slower on gfx1151. + %input10, %input11, %input12, %input13 = func.call @ggml_q6k_load_f32_block(%assignment_count, %bounded_input_size, %input_row, %block, %lane, %input_noalias) : (index, index, index, index, index, buffer) -> (vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) + %row1_contribution = func.call @ggml_q6k_f32_block_row(%bounded_input_size, %expert_row1, %block, %c8, %scale_frame1, %lane, %weight_noalias, %scale_stage, %input10, %input11, %input12, %input13) : (index, index, index, index, index, index, buffer, buffer, vector<4xf32>, vector<4xf32>, vector<4xf32>, vector<4xf32>) -> (f32) + scf.yield %row1_contribution : f32 + } else { + scf.yield %c0_f32 : f32 + } + %next0 = scalar.addf %block_acc0, %contribution0 : f32 + %next1 = scalar.addf %block_acc1, %contribution1 : f32 + scf.yield %next0, %next1 : f32, f32 + } + %route_sum0 = kernel.subgroup.reduce %lane_sum0 : f32 + %route_sum1 = kernel.subgroup.reduce %lane_sum1 : f32 + %weighted0 = scalar.mulf %route_sum0, %route_weight : f32 + %weighted1 = scalar.mulf %route_sum1, %route_weight : f32 + %next0 = scalar.addf %route_acc0, %weighted0 : f32 + %next1 = scalar.addf %route_acc1, %weighted1 : f32 + scf.yield %next0, %next1 : f32, f32 + } + %is_lane_zero = index.cmp eq, %lane, %c0 : index + scf.if %is_lane_zero { + %residual0 = view.load %output_view[%bounded_token, %row0] : view<[%body_token_count]x[%body_output_size]xf32> -> f32 + %result0 = scalar.addf %residual0, %routed_sum0 : f32 + view.store %result0, %output_view[%bounded_token, %row0] : f32, view<[%body_token_count]x[%body_output_size]xf32> + scf.if %row1_valid { + %residual1 = view.load %output_view[%bounded_token, %row1] : view<[%body_token_count]x[%body_output_size]xf32> -> f32 + %result1 = scalar.addf %residual1, %routed_sum1 : f32 + view.store %result1, %output_view[%bounded_token, %row1] : f32, view<[%body_token_count]x[%body_output_size]xf32> + } + } + template.return +} + +// Wave64 matches llama.cpp's Vulkan subgroup width and contracts four Q6_K +// blocks per subgroup pass. +kernel.def target(@qwen3_moe_routed_down_q6k_gfx11_wave64) @qwen3_moe_routed_down_q6k_f32_wave64(%token_count: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index) { + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %configured_output_size = config.get @qwen3_moe.routed_down.output_size : index + %c1 = index.constant 1 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c256 = index.constant 256 : index + %padded_output_size = index.add %configured_output_size, %c3 : index + %output_tiles = index.div %padded_output_size, %c4 : index + kernel.launch.config workgroups(%output_tiles, %token_capacity, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index, %input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer) { + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %configured_input_size0 = config.get @qwen3_moe.routed_down.input_size : index + %configured_route_count0 = config.get @qwen3_moe.routed_down.route_count : index + %configured_expert_count0 = config.get @qwen3_moe.routed_down.expert_count : index + %configured_output_size0 = config.get @qwen3_moe.routed_down.output_size : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %bounded_input_size, %configured_input_size = index.assume %input_size, %configured_input_size0 [range(%input_size, 256, 32768), mul(%input_size, 256), eq(%input_size, %configured_input_size0)] : index, index + %bounded_route_count, %configured_route_count = index.assume %route_count, %configured_route_count0 [range(%route_count, 1, 8), eq(%route_count, %configured_route_count0)] : index, index + %bounded_route_id_stride = index.assume %route_id_stride [range(%route_id_stride, 1, 512)] : index + %bounded_expert_count, %configured_expert_count = index.assume %expert_count, %configured_expert_count0 [range(%expert_count, 1, 512), eq(%expert_count, %configured_expert_count0)] : index, index + %bounded_output_size, %configured_output_size = index.assume %output_size, %configured_output_size0 [range(%output_size, 1, 4096), eq(%output_size, %configured_output_size0)] : index, index + %token0 = kernel.workgroup.id : index + %c0 = index.constant 0 : index + %valid_token = index.cmp ult, %token0, %bounded_token_count : index + %safe_token0 = scf.select %valid_token, %token0, %c0 : index + %safe_token, %body_token_count = index.assume %safe_token0, %bounded_token_count [lt(%safe_token0, %bounded_token_count)] : index, index + %c4 = index.constant 4 : index + %scale_stage_bytes = index.constant 1024 : offset + template.apply<@qwen3_moe.routed_down.q6k_f32.body>(%valid_token, %c4, %c4, %scale_stage_bytes, %body_token_count, %safe_token, %configured_input_size, %configured_route_count, %bounded_route_id_stride, %configured_expert_count, %configured_output_size, %input, %route_ids, %route_weights, %weight, %output) : (i1, index, index, offset, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +// Decode-only direct-F32 route that publishes the normalized Q8_1 x4 row +// consumed by the next projection boundary. Each wave64 subgroup owns two +// adjacent output channels, so four subgroups publish eight channels while the +// completion epilogue still reduces four subgroup sums. +kernel.def target(@qwen3_moe_routed_down_q6k_gfx11_wave64) @qwen3_moe_routed_down_q6k_f32_wave64_next_q8(%token_count: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index) { + %configured_output_size = config.get @qwen3_moe.routed_down.output_size : index + %c1 = index.constant 1 : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c256 = index.constant 256 : index + %padded_output_size = index.add %configured_output_size, %c7 : index + %output_tiles = index.div %padded_output_size, %c8 : index + kernel.launch.config workgroups(%output_tiles, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %input_size: index, %route_count: index, %route_id_stride: index, %expert_count: index, %output_size: index, %input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer, %norm_weight: buffer, %completion_counter: buffer, %next_q8_output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 1)] : index + %configured_input_size0 = config.get @qwen3_moe.routed_down.input_size : index + %configured_route_count0 = config.get @qwen3_moe.routed_down.route_count : index + %configured_expert_count0 = config.get @qwen3_moe.routed_down.expert_count : index + %configured_output_size0 = config.get @qwen3_moe.routed_down.output_size : index + %bounded_input_size, %configured_input_size = index.assume %input_size, %configured_input_size0 [range(%input_size, 256, 32768), mul(%input_size, 256), eq(%input_size, %configured_input_size0)] : index, index + %bounded_route_count, %configured_route_count = index.assume %route_count, %configured_route_count0 [range(%route_count, 1, 8), eq(%route_count, %configured_route_count0)] : index, index + %bounded_route_id_stride = index.assume %route_id_stride [range(%route_id_stride, 1, 512)] : index + %bounded_expert_count, %configured_expert_count = index.assume %expert_count, %configured_expert_count0 [range(%expert_count, 1, 512), eq(%expert_count, %configured_expert_count0)] : index, index + %bounded_output_size, %configured_output_size = index.assume %output_size, %configured_output_size0 [range(%output_size, 128, 4096), mul(%output_size, 128), eq(%output_size, %configured_output_size0)] : index, index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %body_token = index.constant 0 : index + template.apply<@qwen3_moe.routed_down.q6k_f32.pair_body>(%bounded_token_count, %body_token, %configured_input_size, %configured_route_count, %bounded_route_id_stride, %configured_expert_count, %configured_output_size, %input, %route_ids, %route_weights, %weight, %output) : (index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer) + template.apply<@qwen3_moe.routed_down.next_q8_completion>(%c8, %c4, %bounded_token_count, %configured_output_size, %output, %norm_weight, %completion_counter, %next_q8_output) : (index, index, index, index, buffer, buffer, buffer, buffer) + kernel.return +} + +// Two selected experts exercise the noncompact route-ID stride, normalized +// weighted reduction, in-place residual update, K=768, and the output tile +// tail. Uniform 0xaa Q6_K bytes decode to a deterministic nonzero row with +// negative signed group scales. +check.case public @qwen3_moe_routed_down_q6k_q8_1_x4_nonzero_residual_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(2) : index + %route_id_stride = check.literal value(4) : index + %expert_count = check.literal value(2) : index + %output_size = check.literal value(9) : index + %routed_input = check.generate.fill value(0.00390625) : tensor<2x768xf32> + %q8_input = check.generate.fill value(0) : tensor<2x864xi8> + %route_ids = check.generate.iota offset(0) step(1) period(2) : tensor<1x4xi32> + %route_weights = check.generate.fill value(0.5) : tensor<1x2xf32> + %weight = check.generate.fill value(-86) : tensor<2x9x3x210xi8> + %output = check.generate.fill value(1.0) : tensor<1x9xf32> + %expected = check.generate.fill value(135.31431579589844) : tensor<1x9xf32> + %routed_row_count = check.literal value(2) : index + kernel.launch @ggml_quantize_q8_1_x4_f32[%routed_row_count, %input_size](%routed_row_count, %input_size, %routed_input, %q8_input) : [index, index](index, index, tensor<2x768xf32>, tensor<2x864xi8>) + kernel.launch @qwen3_moe_routed_down_q6k_q8_1_x4[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<2x864xi8>, tensor<1x4xi32>, tensor<1x2xf32>, tensor<2x9x3x210xi8>, tensor<1x9xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.25) rtol(9.9999999999999995e-07) nan(same) : tensor<1x9xf32> + check.return +} + +// Two tokens use different route-weight sums so token, route, packed-row, and +// output addressing cannot collapse to the single-token case. +check.case public @qwen3_moe_routed_down_q6k_q8_1_x4_nonzero_prefill_case { + %token_count = check.literal value(2) : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(2) : index + %route_id_stride = check.literal value(4) : index + %expert_count = check.literal value(2) : index + %output_size = check.literal value(1) : index + %routed_row_count = check.literal value(4) : index + %routed_input = check.generate.fill value(0.00390625) : tensor<2x2x768xf32> + %q8_input = check.generate.fill value(0) : tensor<2x2x864xi8> + %route_ids = check.generate.iota offset(0) step(1) period(2) : tensor<2x4xi32> + %route_weights = check.generate.iota offset(0.25) step(0.25) period(4) : tensor<2x2xf32> + %weight = check.generate.fill value(-86) : tensor<2x1x3x210xi8> + %output = check.generate.fill value(1.0) : tensor<2x1xf32> + %expected = check.generate.iota offset(101.73573684692383) step(134.31431579589844) : tensor<2x1xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%routed_row_count, %input_size](%routed_row_count, %input_size, %routed_input, %q8_input) : [index, index](index, index, tensor<2x2x768xf32>, tensor<2x2x864xi8>) + kernel.launch @qwen3_moe_routed_down_q6k_q8_1_x4[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<2x2x864xi8>, tensor<2x4xi32>, tensor<2x2xf32>, tensor<2x1x3x210xi8>, tensor<2x1xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.25) rtol(9.9999999999999995e-07) nan(same) : tensor<2x1xf32> + check.return +} + +// The exact-representable routed inputs make direct-F32 and Q8_1 contraction +// comparable while preserving noncompact route IDs and an output-channel tail. +check.case public @qwen3_moe_routed_down_q6k_f32_wave64_differential_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(2) : index + %route_id_stride = check.literal value(4) : index + %expert_count = check.literal value(2) : index + %output_size = check.literal value(9) : index + %routed_row_count = check.literal value(2) : index + %routed_input = check.generate.fill value(0.00390625) : tensor<2x768xf32> + %q8_input = check.generate.fill value(0) : tensor<2x864xi8> + %route_ids = check.generate.iota offset(0) step(1) period(2) : tensor<1x4xi32> + %route_weights = check.generate.fill value(0.5) : tensor<1x2xf32> + %weight = check.generate.fill value(-86) : tensor<2x9x3x210xi8> + %expected = check.generate.fill value(1.0) : tensor<1x9xf32> + %actual_wave64 = check.generate.fill value(1.0) : tensor<1x9xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%routed_row_count, %input_size](%routed_row_count, %input_size, %routed_input, %q8_input) : [index, index](index, index, tensor<2x768xf32>, tensor<2x864xi8>) + kernel.launch @qwen3_moe_routed_down_q6k_q8_1_x4[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %expected) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<2x864xi8>, tensor<1x4xi32>, tensor<1x2xf32>, tensor<2x9x3x210xi8>, tensor<1x9xf32>) + kernel.launch @qwen3_moe_routed_down_q6k_f32_wave64[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %routed_input, %route_ids, %route_weights, %weight, %actual_wave64) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<2x768xf32>, tensor<1x4xi32>, tensor<1x2xf32>, tensor<2x9x3x210xi8>, tensor<1x9xf32>) + check.expect.close actual(%actual_wave64) expected(%expected) atol(0.25) rtol(0.01) nan(same) : tensor<1x9xf32> + check.return +} + +// Wave64 reference for the row producer used by the fused completion path. +kernel.def target(@qwen3_moe_routed_down_q6k_gfx11_wave64) @qwen3_moe_rmsnorm_quantize_q8_1_x4_wave64_production_check() { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%c1, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%input: buffer, %norm_weight: buffer, %q8_output: buffer) { + %publish_normalized = scalar.constant false : i1 + %reduction_subgroup_count = index.constant 4 : index + %token_count = index.constant 1 : index + %token = index.constant 0 : index + template.apply<@qwen3_moe.rmsnorm_quantize_q8_1_x4.body>(%publish_normalized, %reduction_subgroup_count, %token_count, %token, %input, %norm_weight, %q8_output, %q8_output) : (i1, index, index, index, buffer, buffer, buffer, buffer) + kernel.return +} + +// The production hidden width requires every direct-F32 output tile to arrive +// before normalization. Nonzero Q6_K data compares the paired-row producer +// against the ordinary wave64 producer and standalone normalization, then +// reuses the completion word so stale completion state is observable. +check.case public @qwen3_moe_routed_down_q6k_f32_wave64_next_q8_differential_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(8) : index + %route_id_stride = check.literal value(8) : index + %expert_count = check.literal value(8) : index + %output_size = check.literal value(2048) : index + %input = check.generate.iota offset(-0.5) step(0.0001627604166666667) period(6144) : tensor<8x768xf32> + %route_ids = check.generate.iota offset(0) step(1) period(8) : tensor<8xi32> + %route_weights = check.generate.iota offset(0.02777777777777778) step(0.02777777777777778) : tensor<8xf32> + %weight = check.generate.fill value(-86) : tensor<8x2048x3x210xi8> + %norm_weight = check.generate.iota offset(-1.0) step(0.0009765625) : tensor<2048xf32> + %expected_output = check.generate.iota offset(-0.5) step(0.00048828125) period(2048) : tensor<2048xf32> + %actual_output0 = check.generate.iota offset(-0.5) step(0.00048828125) period(2048) : tensor<2048xf32> + %actual_output1 = check.generate.iota offset(-0.5) step(0.00048828125) period(2048) : tensor<2048xf32> + %expected_q8 = check.generate.fill value(0) : tensor<2304xi8> + %actual_q8_0 = check.generate.fill value(1) : tensor<2304xi8> + %actual_q8_1 = check.generate.fill value(1) : tensor<2304xi8> + %completion_counter = check.generate.fill value(0) : tensor<1xi32> + %expected_counter = check.generate.fill value(0) : tensor<1xi32> + kernel.launch @qwen3_moe_routed_down_q6k_f32_wave64[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %input, %route_ids, %route_weights, %weight, %expected_output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<8x768xf32>, tensor<8xi32>, tensor<8xf32>, tensor<8x2048x3x210xi8>, tensor<2048xf32>) + kernel.launch @qwen3_moe_rmsnorm_quantize_q8_1_x4_wave64_production_check(%expected_output, %norm_weight, %expected_q8) : (tensor<2048xf32>, tensor<2048xf32>, tensor<2304xi8>) + kernel.launch @qwen3_moe_routed_down_q6k_f32_wave64_next_q8[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %input, %route_ids, %route_weights, %weight, %actual_output0, %norm_weight, %completion_counter, %actual_q8_0) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<8x768xf32>, tensor<8xi32>, tensor<8xf32>, tensor<8x2048x3x210xi8>, tensor<2048xf32>, tensor<2048xf32>, tensor<1xi32>, tensor<2304xi8>) + kernel.launch @qwen3_moe_routed_down_q6k_f32_wave64_next_q8[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %input, %route_ids, %route_weights, %weight, %actual_output1, %norm_weight, %completion_counter, %actual_q8_1) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<8x768xf32>, tensor<8xi32>, tensor<8xf32>, tensor<8x2048x3x210xi8>, tensor<2048xf32>, tensor<2048xf32>, tensor<1xi32>, tensor<2304xi8>) + check.expect.close actual(%actual_output0) expected(%expected_output) atol(0.0) rtol(0.0) nan(same) : tensor<2048xf32> + check.expect.close actual(%actual_output1) expected(%expected_output) atol(0.0) rtol(0.0) nan(same) : tensor<2048xf32> + check.expect.equal actual(%actual_q8_0) expected(%expected_q8) : tensor<2304xi8> + check.expect.equal actual(%actual_q8_1) expected(%expected_q8) : tensor<2304xi8> + check.expect.equal actual(%completion_counter) expected(%expected_counter) : tensor<1xi32> + check.return +} + +check.case public @qwen3_moe_routed_down_q6k_q8_1_x4_benchmark_case { + %token_count = check.param.choice values([1, 2, 4, 8, 16, 17, 32, 63, 128, 129, 512]) name("token_count") : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(8) : index + %route_id_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(2048) : index + %q8_input = check.generate.fill value(0) : tensor<[%token_count]x8x864xi8> + %route_ids = check.generate.iota offset(0) step(1) period(127) : tensor<[%token_count]x128xi32> + %route_weights = check.generate.fill value(0.125) : tensor<[%token_count]x8xf32> + %weight = check.generate.fill value(0) : tensor<128x2048x3x210xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %expected = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + kernel.launch @qwen3_moe_routed_down_q6k_q8_1_x4[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<[%token_count]x8x864xi8>, tensor<[%token_count]x128xi32>, tensor<[%token_count]x8xf32>, tensor<128x2048x3x210xi8>, tensor<[%token_count]x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x2048xf32> + check.return +} + +check.case public @qwen3_moe_routed_down_q6k_q8_1_x4_pipeline_benchmark_case { + %token_count = check.param.choice values([1, 2, 4, 8, 16, 17, 32, 63, 128, 129, 512]) name("token_count") : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(8) : index + %route_id_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(2048) : index + // Each 768-element route contains six complete Q8_1 x4 groups, so packing + // the eight contiguous routes as one 6,144-element token row preserves the + // exact per-route physical layout without a derived testbench scalar. + %routed_input_size = check.literal value(6144) : index + %routed_input = check.generate.fill value(0.0) : tensor<[%token_count]x8x768xf32> + %q8_input = check.generate.fill value(1) : tensor<[%token_count]x8x864xi8> + %route_ids = check.generate.iota offset(0) step(1) period(127) : tensor<[%token_count]x128xi32> + %route_weights = check.generate.fill value(0.125) : tensor<[%token_count]x8xf32> + %weight = check.generate.fill value(0) : tensor<128x2048x3x210xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %expected = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %routed_input_size](%token_count, %routed_input_size, %routed_input, %q8_input) : [index, index](index, index, tensor<[%token_count]x8x768xf32>, tensor<[%token_count]x8x864xi8>) + kernel.launch @qwen3_moe_routed_down_q6k_q8_1_x4[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<[%token_count]x8x864xi8>, tensor<[%token_count]x128xi32>, tensor<[%token_count]x8xf32>, tensor<128x2048x3x210xi8>, tensor<[%token_count]x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x2048xf32> + check.return +} + +check.case public @qwen3_moe_routed_down_q6k_f32_wave64_benchmark_case { + %token_count = check.param.choice values([1, 2, 4, 8, 16, 17, 32, 63, 128, 129, 512]) name("token_count") : index + %input_size = check.literal value(768) : index + %route_count = check.literal value(8) : index + %route_id_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(2048) : index + %routed_input = check.generate.fill value(0.0) : tensor<[%token_count]x8x768xf32> + %route_ids = check.generate.iota offset(0) step(1) period(127) : tensor<[%token_count]x128xi32> + %route_weights = check.generate.fill value(0.125) : tensor<[%token_count]x8xf32> + %weight = check.generate.fill value(0) : tensor<128x2048x3x210xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %expected = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + kernel.launch @qwen3_moe_routed_down_q6k_f32_wave64[%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_id_stride, %expert_count, %output_size, %routed_input, %route_ids, %route_weights, %weight, %output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<[%token_count]x8x768xf32>, tensor<[%token_count]x128xi32>, tensor<[%token_count]x8xf32>, tensor<128x2048x3x210xi8>, tensor<[%token_count]x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x2048xf32> + check.return +} + +check.benchmark<@qwen3_moe_routed_down_q6k_q8_1_x4_nonzero_residual_case> @qwen3_moe_routed_down_q6k_q8_1_x4_small + +check.benchmark<@qwen3_moe_routed_down_q6k_q8_1_x4_benchmark_case> @qwen3_moe_routed_down_q6k_q8_1_x4_decode {token_count = 1} + +check.benchmark<@qwen3_moe_routed_down_q6k_q8_1_x4_benchmark_case> @qwen3_moe_routed_down_q6k_q8_1_x4_small_batch_2 {token_count = 2} + +check.benchmark<@qwen3_moe_routed_down_q6k_q8_1_x4_benchmark_case> @qwen3_moe_routed_down_q6k_q8_1_x4_small_batch_4 {token_count = 4} + +check.benchmark<@qwen3_moe_routed_down_q6k_q8_1_x4_benchmark_case> @qwen3_moe_routed_down_q6k_q8_1_x4_small_batch_8 {token_count = 8} + +check.benchmark<@qwen3_moe_routed_down_q6k_q8_1_x4_benchmark_case> @qwen3_moe_routed_down_q6k_q8_1_x4_small_batch_16 {token_count = 16} + +check.benchmark<@qwen3_moe_routed_down_q6k_q8_1_x4_benchmark_case> @qwen3_moe_routed_down_q6k_q8_1_x4_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_routed_down_q6k_q8_1_x4_benchmark_case> @qwen3_moe_routed_down_q6k_q8_1_x4_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_routed_down_q6k_q8_1_x4_benchmark_case> @qwen3_moe_routed_down_q6k_q8_1_x4_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_routed_down_q6k_q8_1_x4_pipeline_benchmark_case> @qwen3_moe_routed_down_q6k_q8_1_x4_pipeline_decode {token_count = 1} + +check.benchmark<@qwen3_moe_routed_down_q6k_q8_1_x4_pipeline_benchmark_case> @qwen3_moe_routed_down_q6k_q8_1_x4_pipeline_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_routed_down_q6k_q8_1_x4_pipeline_benchmark_case> @qwen3_moe_routed_down_q6k_q8_1_x4_pipeline_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_routed_down_q6k_q8_1_x4_pipeline_benchmark_case> @qwen3_moe_routed_down_q6k_q8_1_x4_pipeline_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_routed_down_q6k_f32_wave64_benchmark_case> @qwen3_moe_routed_down_q6k_f32_wave64_decode {token_count = 1} + +check.benchmark<@qwen3_moe_routed_down_q6k_f32_wave64_benchmark_case> @qwen3_moe_routed_down_q6k_f32_wave64_small_batch_2 {token_count = 2} + +check.benchmark<@qwen3_moe_routed_down_q6k_f32_wave64_benchmark_case> @qwen3_moe_routed_down_q6k_f32_wave64_small_batch_4 {token_count = 4} + +check.benchmark<@qwen3_moe_routed_down_q6k_f32_wave64_benchmark_case> @qwen3_moe_routed_down_q6k_f32_wave64_small_batch_8 {token_count = 8} + +check.benchmark<@qwen3_moe_routed_down_q6k_f32_wave64_benchmark_case> @qwen3_moe_routed_down_q6k_f32_wave64_small_batch_16 {token_count = 16} + +check.benchmark<@qwen3_moe_routed_down_q6k_q8_1_x4_next_q8_benchmark_case> @qwen3_moe_routed_down_q6k_q8_1_x4_next_q8_decode + +check.benchmark<@qwen3_moe_routed_down_q6k_q8_1_x4_next_q8_composed_benchmark_case> @qwen3_moe_routed_down_q6k_q8_1_x4_next_q8_composed_decode diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_quantized_f16_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_quantized_f16_wmma.loom new file mode 100644 index 000000000000..f315e43ad239 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_quantized_f16_wmma.loom @@ -0,0 +1,887 @@ +// Expert-grouped Q4_K and Q6_K down projections for gfx11 prefill. +// +// Each two-wave workgroup decodes 64 output channels for 32 routed rows of one +// expert. Raw GGUF weights and FP16 routed SwiGLU rows are staged as FP16 WMMA +// operands. Results remain FP16 and are scattered into compact +// [token, route, hidden] order without collisions. A following reduction +// widens each route, applies its normalized weight, and accumulates the +// residual in FP32. +// +// Top-k routing selects an expert at most once per token. The expert table +// therefore contains at most token_count assignments per expert, while each +// assignment ordinal still identifies the compact [token, route] activation +// and route-weight row. +template.decl @qwen3_moe.routed_down.quantized_prefill.body(%weight_format: index, %token_count: index, %input: buffer, %expert_table: buffer, %weight: buffer, %output: buffer) + +amdgpu.target @qwen3_moe_routed_down_gfx11_wave64 {subgroup_size = 64} + +config.decl @qwen3_moe.routed_down.input_size : %value: index where [range(%value, 256, 32768), mul(%value, 256)] + +config.decl @qwen3_moe.routed_down.route_count : %value: index where [range(%value, 1, 8)] + +config.decl @qwen3_moe.routed_down.expert_count : %value: index where [range(%value, 1, 512)] + +config.decl @qwen3_moe.routed_down.output_size : %value: index where [range(%value, 1, 4096)] + +kernel.decl @ggml_quantize_q8_1_x4_f32(%token_count$10: index, %input_size$11: index) launch(%token_count$12: index, %input_size$13: index, %input: buffer, %output: buffer) + +kernel.decl @qwen3_moe_build_expert_table(%token_count$16: index, %route_count$17: index, %route_stride$18: index, %expert_count$19: index) launch(%token_count$20: index, %route_count$21: index, %route_stride$22: index, %expert_count$23: index, %route_ids: buffer, %expert_table: buffer) + +kernel.decl @qwen3_moe_build_expert_partition_table(%token_count$26: index, %route_count$27: index, %expert_count$28: index) launch(%token_count$29: index, %route_count$30: index, %expert_count$31: index, %expert_table: buffer, %partition_table: buffer) + +kernel.decl @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma(%token_count$34: index) launch(%token_count$35: index, %input: buffer, %expert_table: buffer, %partition_table: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) + +kernel.decl @qwen3_moe_routed_gate_up_swiglu_q4k_q8(%token_count$42: index, %route_count$43: index, %route_stride$44: index, %expert_count$45: index, %output_size$46: index) launch(%token_count$47: index, %route_count$48: index, %route_stride$49: index, %expert_count$50: index, %output_size$51: index, %q8_input: buffer, %route_ids: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) + +func.decl @qwen3_moe_q4k_wmma_vector4(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group: index, %packet: index) -> (vector<4xf16>) + +func.decl @qwen3_moe_q4k_wmma_load_header(%weight: buffer, %row_byte_base: offset, %q4_block: index) -> (vector<4xi32>) + +func.decl @qwen3_moe_q4k_wmma_load_code(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group_pair: index, %packet: index) -> (vector<1xi32>) + +func.decl @qwen3_moe_q4k_wmma_vector4_from_header_code(%q4_group: index, %header_words: vector<4xi32>, %q_word: vector<1xi32>) -> (vector<4xf16>) + +func.decl @ggml_q6k_f16_vector4(%weight: buffer, %weight_row_byte_base: offset, %q6_block: index, %q6_group: index, %packet: index) -> (vector<4xf16>) + +kernel.decl @qwen3_moe_routed_down_q4k_q8_1_x4(%token_count$83: index, %input_size$84: index, %route_count$85: index, %route_id_stride$86: index, %expert_count$87: index, %output_size$88: index) launch(%token_count$89: index, %input_size$90: index, %route_count$91: index, %route_id_stride$92: index, %expert_count$93: index, %output_size$94: index, %q8_input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer) + +kernel.decl @qwen3_moe_routed_down_q6k_q8_1_x4(%token_count$100: index, %input_size$101: index, %route_count$102: index, %route_id_stride$103: index, %expert_count$104: index, %output_size$105: index) launch(%token_count$106: index, %input_size$107: index, %route_count$108: index, %route_id_stride$109: index, %expert_count$110: index, %output_size$111: index, %q8_input: buffer, %route_ids: buffer, %route_weights: buffer, %weight: buffer, %output: buffer) + +// Acquires the packed words shared by one four-group half of a Q6_K block. +// Each QL word supplies two groups and the QH word supplies all four. +func.def inline @qwen3_moe_q6k_wmma_load_half_codes(%weight: buffer, %weight_row_byte_base: offset, %q6_block: index, %q6_half: index, %packet: index) -> (vector<1xi32>, vector<1xi32>, vector<1xi32>) { + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %block_bytes = index.constant 210 : offset + %qh_byte_add = index.constant 128 : offset + %bounded_half = index.assume %q6_half [range(%q6_half, 0, 1)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q6_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %weight_row_byte_base, %block_byte_add : offset + %qh_byte_base = index.add %block_byte_base, %qh_byte_add : offset + %ql_view = buffer.view %weight[%block_byte_base] : buffer -> view<32xi32> + %qh_view = buffer.view %weight[%qh_byte_base] : buffer -> view<16xi32> + %ql_half_word_base = index.mul %bounded_half, %c16 : index + %ql0_word_index = index.add %ql_half_word_base, %bounded_packet : index + %ql1_word_base = index.add %ql_half_word_base, %c8 : index + %ql1_word_index = index.add %ql1_word_base, %bounded_packet : index + %qh_half_word_base = index.mul %bounded_half, %c8 : index + %qh_word_index = index.add %qh_half_word_base, %bounded_packet : index + %ql0_word = vector.load %ql_view[%ql0_word_index] : view<32xi32> -> vector<1xi32> + %ql1_word = vector.load %ql_view[%ql1_word_index] : view<32xi32> -> vector<1xi32> + %qh_word = vector.load %qh_view[%qh_word_index] : view<16xi32> -> vector<1xi32> + func.return %ql0_word, %ql1_word, %qh_word : vector<1xi32>, vector<1xi32>, vector<1xi32> +} + +// Decodes four adjacent values after the surrounding schedule has selected +// the scale and retained the packed code words at their natural lifetimes. +func.def inline @qwen3_moe_q6k_wmma_vector4_from_scale_codes(%q6_group: index, %scale_i8: i8, %d_f16: f16, %ql_word: vector<1xi32>, %qh_word: vector<1xi32>) -> (vector<4xf16>) { + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c4_i32v = vector.constant 4 : vector<1xi32> + %nibble_mask = vector.constant 252645135 : vector<1xi32> + %high_mask = vector.constant 50529027 : vector<1xi32> + %c32_f32v = vector.constant 32.0 : vector<4xf32> + %bounded_group = index.assume %q6_group [range(%q6_group, 0, 7)] : index + %group_in_half = index.rem %bounded_group, %c4 : index + %nibble = index.div %group_in_half, %c2 : index + %nibble_shift_index = index.mul %nibble, %c4 : index + %nibble_shift_i32 = index.cast %nibble_shift_index : index to i32 + %nibble_shift = vector.splat %nibble_shift_i32 : vector<1xi32> + %qh_shift_index = index.mul %group_in_half, %c2 : index + %qh_shift_i32 = index.cast %qh_shift_index : index to i32 + %qh_shift = vector.splat %qh_shift_i32 : vector<1xi32> + %ql_shifted = vector.shrui %ql_word, %nibble_shift : vector<1xi32> + %ql = vector.andi %ql_shifted, %nibble_mask : vector<1xi32> + %qh_shifted = vector.shrui %qh_word, %qh_shift : vector<1xi32> + %qh_low = vector.andi %qh_shifted, %high_mask : vector<1xi32> + %qh = vector.shli %qh_low, %c4_i32v : vector<1xi32> + %code = vector.ori %ql, %qh : vector<1xi32> + %code_i8 = vector.bitcast %code : vector<1xi32> to vector<4xi8> + %code_f32 = vector.uitofp %code_i8 : vector<4xi8> to vector<4xf32> + %centered = vector.subf %code_f32, %c32_f32v : vector<4xf32> + %scale = scalar.sitofp %scale_i8 : i8 to f32 + %d = scalar.extf %d_f16 : f16 to f32 + %combined_scale = scalar.mulf %scale, %d : f32 + %combined_scale_vector = vector.splat %combined_scale : vector<4xf32> + %values_f32 = vector.mulf %centered, %combined_scale_vector : vector<4xf32> + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +// Loads only the group scale and block multiplier while reusing packed words +// retained by an enclosing four-group half-block schedule. +func.def inline @qwen3_moe_q6k_wmma_vector4_from_half_codes(%weight: buffer, %weight_row_byte_base: offset, %q6_block: index, %q6_group: index, %packet: index, %ql0_word: vector<1xi32>, %ql1_word: vector<1xi32>, %qh_word: vector<1xi32>) -> (vector<4xf16>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %block_bytes = index.constant 210 : offset + %scale_byte_add = index.constant 192 : offset + %d_byte_add = index.constant 208 : offset + %bounded_group = index.assume %q6_group [range(%q6_group, 0, 7)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q6_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %weight_row_byte_base, %block_byte_add : offset + %scale_byte_base = index.add %block_byte_base, %scale_byte_add : offset + %d_byte_base = index.add %block_byte_base, %d_byte_add : offset + %scale_view = buffer.view %weight[%scale_byte_base] : buffer -> view<16xi8> + %d_view = buffer.view %weight[%d_byte_base] : buffer -> view<1xf16> + %group_in_half = index.rem %bounded_group, %c4 : index + %ql_side = index.rem %group_in_half, %c2 : index + %uses_ql1 = index.cmp eq, %ql_side, %c1 : index + %ql_word = scf.select %uses_ql1, %ql1_word, %ql0_word : vector<1xi32> + %scale_packet_half = index.div %bounded_packet, %c4 : index + %scale_group_base = index.mul %bounded_group, %c2 : index + %scale_index = index.add %scale_group_base, %scale_packet_half : index + %scale_i8 = view.load %scale_view[%scale_index] : view<16xi8> -> i8 + %d_f16 = view.load %d_view[%c0] : view<1xf16> -> f16 + %values = func.call @qwen3_moe_q6k_wmma_vector4_from_scale_codes(%bounded_group, %scale_i8, %d_f16, %ql_word, %qh_word) : (index, i8, f16, vector<1xi32>, vector<1xi32>) -> (vector<4xf16>) + func.return %values : vector<4xf16> +} + +// Shared raw-quantized matrix schedule. Entry points pass a literal weight +// format so linking and JIT specialization erase the inactive packed decoder. +template.def<@qwen3_moe.routed_down.quantized_prefill.body> device @qwen3_moe_routed_down_quantized_f16_wmma_body(%weight_format: index, %token_count: index, %input: buffer, %expert_table: buffer, %weight: buffer, %output: buffer) { + %input_size = config.get @qwen3_moe.routed_down.input_size : index + %route_count = config.get @qwen3_moe.routed_down.route_count : index + %expert_count = config.get @qwen3_moe.routed_down.expert_count : index + %output_size = config.get @qwen3_moe.routed_down.output_size : index + %bounded_route_count = index.assume %route_count [range(%route_count, 1, 8)] : index + %bounded_expert_count = index.assume %expert_count [range(%expert_count, 1, 128)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 4096)] : index + %channel_tile = kernel.workgroup.id : index + %route_tile = kernel.workgroup.id : index + %expert = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 1)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c6 = index.constant 6 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c40 = index.constant 40 : index + %c48 = index.constant 48 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %c4_bytes = index.constant 4 : offset + %q4_block_bytes = index.constant 144 : offset + %q6_block_bytes = index.constant 210 : offset + %weight_stage_bytes = index.constant 5120 : offset + %activation_stage_bytes = index.constant 2560 : offset + %route_stage_bytes = index.constant 128 : offset + %wave_result_stage_bytes = index.constant 512 : offset + %result_stage_bytes = index.constant 1024 : offset + %c0_i32 = scalar.constant 0 : i32 + %cn1_i32 = scalar.constant -1 : i32 + %c0_i32x1 = vector.constant 0 : vector<1xi32> + %c0_i32x4 = vector.constant 0 : vector<4xi32> + %c0_f16x4 = vector.constant 0.0 : vector<4xf16> + %zero_accumulator = vector.constant 0.0 : vector<8xf16> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %assignment_count = index.mul %token_count, %bounded_route_count : index + %assignment_table_byte_base = index.scale %bounded_expert_count, %c4_bytes : index, offset -> offset + %is_q4 = index.cmp eq, %weight_format, %c4 : index + %is_q6 = index.cmp eq, %weight_format, %c6 : index + %quant_block_bytes = scf.select %is_q6, %q6_block_bytes, %q4_block_bytes : offset + %quant_block_count = index.div %input_size, %c256 : index + %weight_row_bytes = index.scale %quant_block_count, %quant_block_bytes : index, offset -> offset + %weight_expert_bytes = index.scale %bounded_output_size, %weight_row_bytes : index, offset -> offset + %input_noalias, %expert_table_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input, %expert_table, %weight, %output : buffer, buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%assignment_count]x[%input_size]xf16> + %count_view = buffer.view %expert_table_noalias[%c0_offset] : buffer -> view<[%bounded_expert_count]xi32> + %assignment_view = buffer.view %expert_table_noalias[%assignment_table_byte_base] : buffer -> view<[%bounded_expert_count]x[%token_count]xi32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%assignment_count]x[%bounded_output_size]xf16> + %weight_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %activation_stage = buffer.alloca align(16) %activation_stage_bytes : buffer + %route_stage = buffer.alloca align(16) %route_stage_bytes : buffer + %result_stage = buffer.alloca align(16) %result_stage_bytes : buffer + %weight_stage_view = buffer.view %weight_stage[%c0_offset] : buffer -> view<64x40xf16> + %activation_stage_physical_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x40xf16> + %activation_fragment_layout = encoding.layout.strided [1, %c40] : encoding + %activation_fragment_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x32xf16, %activation_fragment_layout> + %route_stage_view = buffer.view %route_stage[%c0_offset] : buffer -> view<32xi32> + %wave_result_stage_offset = index.scale %subgroup, %wave_result_stage_bytes : index, offset -> offset + %result_fragment_layout = encoding.layout.strided [1, %c16] : encoding + %result_fragment_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf16, %result_fragment_layout> + %result_physical_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf16> + %channel_tile_base = index.mul %channel_tile, %c64 : index + %initial_route_tile_base = index.mul %route_tile, %c32 : index + // The launch workload fixes the interleaved route partitions for this exact + // command-program specialization. + %padded_token_count = index.add %token_count, %c63 : index + %route_partition_count = index.div %padded_token_count, %c64 : index + %route_partition_step = index.mul %route_partition_count, %c32 : index + %bounded_expert, %table_expert_count = index.assume %expert, %bounded_expert_count [lt(%expert, %bounded_expert_count)] : index, index + %is_workitem_zero = index.cmp eq, %workitem, %c0 : index + %lane_expert_route_count = scf.if %is_workitem_zero -> (i32) { + %loaded = view.load %count_view[%bounded_expert] : view<[%bounded_expert_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %c0_i32 : i32 + } + %expert_route_count_reduced = kernel.workgroup.reduce %lane_expert_route_count : i32 + %expert_route_count_i32 = kernel.subgroup.broadcast.first %expert_route_count_reduced : i32 + %expert_route_count0 = index.cast %expert_route_count_i32 : i32 to index + %expert_route_count = index.assume %expert_route_count0 [range(%expert_route_count0, 0, 2048)] : index + scf.for %route_tile_base = [%initial_route_tile_base to %expert_route_count step %route_partition_step] { + // Snapshot this expert's compact assignment map once, then reuse it for + // every K group and output-channel tile. + %loads_route = index.cmp ult, %workitem, %c32 : index + scf.if %loads_route { + %local_route = index.assume %workitem [range(%workitem, 0, 31)] : index + %assignment_ordinal = index.add %route_tile_base, %local_route : index + %valid_row = index.cmp ult, %assignment_ordinal, %expert_route_count : index + %assignment_i32 = scf.if %valid_row -> (i32) { + %bounded_assignment_ordinal, %table_token_count = index.assume %assignment_ordinal, %token_count [lt(%assignment_ordinal, %token_count)] : index, index + %loaded = view.load %assignment_view[%bounded_expert, %bounded_assignment_ordinal] : view<[%bounded_expert_count]x[%token_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %cn1_i32 : i32 + } + view.store %assignment_i32, %route_stage_view[%local_route] : i32, view<32xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %load_packet = index.rem %workitem, %c8 : index + %load_k = index.mul %load_packet, %c4 : index + %load_row0 = index.div %workitem, %c8 : index + %load_row = index.assume %load_row0 [range(%load_row0, 0, 15)] : index + %expert_byte_base = index.scale %bounded_expert, %weight_expert_bytes : index, offset -> offset + %subgroup_channel_add = index.mul %subgroup, %c32 : index + %subgroup_channel1 = index.add %subgroup_channel_add, %c16 : index + %init00 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf16> + %init01 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf16> + %init10 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf16> + %init11 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf16> + %result00, %result01, %result10, %result11 = scf.for %quant_block = [%c0 to %quant_block_count step %c1](%block_acc00 = %init00 : vector<8xf16>, %block_acc01 = %init01 : vector<8xf16>, %block_acc10 = %init10 : vector<8xf16>, %block_acc11 = %init11 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16>) { + // Four lanes cover the 64 output rows owned by one load packet. Snapshot + // each Q4_K block header once and retain it across all eight quant groups. + // The Q6 specialization erases this loop with %is_q4=false. + %header0, %header1, %header2, %header3 = scf.for %header_row_offset = [%c0 to %c64 step %c16](%prior_header0 = %c0_i32x4 : vector<4xi32>, %prior_header1 = %c0_i32x4 : vector<4xi32>, %prior_header2 = %c0_i32x4 : vector<4xi32>, %prior_header3 = %c0_i32x4 : vector<4xi32>) -> (vector<4xi32>, vector<4xi32>, vector<4xi32>, vector<4xi32>) unroll { + %header_local_row0 = index.add %load_row, %header_row_offset : index + %header_local_row = index.assume %header_local_row0 [range(%header_local_row0, 0, 63)] : index + %header_channel = index.add %channel_tile_base, %header_local_row : index + %valid_header_channel = index.cmp ult, %header_channel, %bounded_output_size : index + %loads_q4_header = scalar.andi %is_q4, %valid_header_channel : i1 + %loaded_header = scf.if %loads_q4_header -> (vector<4xi32>) { + %header_channel_byte_add = index.scale %header_channel, %weight_row_bytes : index, offset -> offset + %header_row_byte_base = index.add %expert_byte_base, %header_channel_byte_add : offset + %header_words = func.call @qwen3_moe_q4k_wmma_load_header(%weight_noalias, %header_row_byte_base, %quant_block) : (buffer, offset, index) -> (vector<4xi32>) + scf.yield %header_words : vector<4xi32> + } else { + scf.yield %c0_i32x4 : vector<4xi32> + } + %updates_header0 = index.cmp eq, %header_row_offset, %c0 : index + %updates_header1 = index.cmp eq, %header_row_offset, %c16 : index + %updates_header2 = index.cmp eq, %header_row_offset, %c32 : index + %updates_header3 = index.cmp eq, %header_row_offset, %c48 : index + %next_header0 = scf.select %updates_header0, %loaded_header, %prior_header0 : vector<4xi32> + %next_header1 = scf.select %updates_header1, %loaded_header, %prior_header1 : vector<4xi32> + %next_header2 = scf.select %updates_header2, %loaded_header, %prior_header2 : vector<4xi32> + %next_header3 = scf.select %updates_header3, %loaded_header, %prior_header3 : vector<4xi32> + scf.yield %next_header0, %next_header1, %next_header2, %next_header3 : vector<4xi32>, vector<4xi32>, vector<4xi32>, vector<4xi32> + } + %quant_group_outer_count = scf.select %is_q4, %c4, %c2 : index + %groups_per_outer = scf.select %is_q4, %c2, %c4 : index + %block_result00, %block_result01, %block_result10, %block_result11 = scf.for %quant_group_outer = [%c0 to %quant_group_outer_count step %c1](%acc00 = %block_acc00 : vector<8xf16>, %acc01 = %block_acc01 : vector<8xf16>, %acc10 = %block_acc10 : vector<8xf16>, %acc11 = %block_acc11 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16>) { + // Q4_K retains one packed code word across its adjacent low/high group + // pair. The Q6 specialization erases this acquisition path. + %q_word0, %q_word1, %q_word2, %q_word3 = scf.for %code_row_offset = [%c0 to %c64 step %c16](%prior_q_word0 = %c0_i32x1 : vector<1xi32>, %prior_q_word1 = %c0_i32x1 : vector<1xi32>, %prior_q_word2 = %c0_i32x1 : vector<1xi32>, %prior_q_word3 = %c0_i32x1 : vector<1xi32>) -> (vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>) unroll { + %code_local_row0 = index.add %load_row, %code_row_offset : index + %code_local_row = index.assume %code_local_row0 [range(%code_local_row0, 0, 63)] : index + %code_channel = index.add %channel_tile_base, %code_local_row : index + %valid_code_channel = index.cmp ult, %code_channel, %bounded_output_size : index + %loads_q4_code = scalar.andi %is_q4, %valid_code_channel : i1 + %loaded_q_word = scf.if %loads_q4_code -> (vector<1xi32>) { + %q4_group_pair = index.assume %quant_group_outer [range(%quant_group_outer, 0, 3)] : index + %code_channel_byte_add = index.scale %code_channel, %weight_row_bytes : index, offset -> offset + %code_row_byte_base = index.add %expert_byte_base, %code_channel_byte_add : offset + %q_word = func.call @qwen3_moe_q4k_wmma_load_code(%weight_noalias, %code_row_byte_base, %quant_block, %q4_group_pair, %load_packet) : (buffer, offset, index, index, index) -> (vector<1xi32>) + scf.yield %q_word : vector<1xi32> + } else { + scf.yield %c0_i32x1 : vector<1xi32> + } + %updates_q_word0 = index.cmp eq, %code_row_offset, %c0 : index + %updates_q_word1 = index.cmp eq, %code_row_offset, %c16 : index + %updates_q_word2 = index.cmp eq, %code_row_offset, %c32 : index + %updates_q_word3 = index.cmp eq, %code_row_offset, %c48 : index + %next_q_word0 = scf.select %updates_q_word0, %loaded_q_word, %prior_q_word0 : vector<1xi32> + %next_q_word1 = scf.select %updates_q_word1, %loaded_q_word, %prior_q_word1 : vector<1xi32> + %next_q_word2 = scf.select %updates_q_word2, %loaded_q_word, %prior_q_word2 : vector<1xi32> + %next_q_word3 = scf.select %updates_q_word3, %loaded_q_word, %prior_q_word3 : vector<1xi32> + scf.yield %next_q_word0, %next_q_word1, %next_q_word2, %next_q_word3 : vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32> + } + // Four Q6_K groups in one half-block share a QH word and consume two + // nibbles from each of two QL words. Retain the three packed words for + // all four groups. The Q4 specialization erases this acquisition path. + %q6_ql00, %q6_ql10, %q6_qh0, %q6_ql01, %q6_ql11, %q6_qh1, %q6_ql02, %q6_ql12, %q6_qh2, %q6_ql03, %q6_ql13, %q6_qh3 = scf.for %q6_code_row_offset = [%c0 to %c64 step %c16](%prior_q6_ql00 = %c0_i32x1 : vector<1xi32>, %prior_q6_ql10 = %c0_i32x1 : vector<1xi32>, %prior_q6_qh0 = %c0_i32x1 : vector<1xi32>, %prior_q6_ql01 = %c0_i32x1 : vector<1xi32>, %prior_q6_ql11 = %c0_i32x1 : vector<1xi32>, %prior_q6_qh1 = %c0_i32x1 : vector<1xi32>, %prior_q6_ql02 = %c0_i32x1 : vector<1xi32>, %prior_q6_ql12 = %c0_i32x1 : vector<1xi32>, %prior_q6_qh2 = %c0_i32x1 : vector<1xi32>, %prior_q6_ql03 = %c0_i32x1 : vector<1xi32>, %prior_q6_ql13 = %c0_i32x1 : vector<1xi32>, %prior_q6_qh3 = %c0_i32x1 : vector<1xi32>) -> (vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>) unroll { + %q6_code_local_row0 = index.add %load_row, %q6_code_row_offset : index + %q6_code_local_row = index.assume %q6_code_local_row0 [range(%q6_code_local_row0, 0, 63)] : index + %q6_code_channel = index.add %channel_tile_base, %q6_code_local_row : index + %valid_q6_code_channel = index.cmp ult, %q6_code_channel, %bounded_output_size : index + %loads_q6_code = scalar.andi %is_q6, %valid_q6_code_channel : i1 + %loaded_q6_ql0, %loaded_q6_ql1, %loaded_q6_qh = scf.if %loads_q6_code -> (vector<1xi32>, vector<1xi32>, vector<1xi32>) { + %q6_half = index.assume %quant_group_outer [range(%quant_group_outer, 0, 1)] : index + %q6_code_channel_byte_add = index.scale %q6_code_channel, %weight_row_bytes : index, offset -> offset + %q6_code_row_byte_base = index.add %expert_byte_base, %q6_code_channel_byte_add : offset + %ql0_word, %ql1_word, %qh_word = func.call @qwen3_moe_q6k_wmma_load_half_codes(%weight_noalias, %q6_code_row_byte_base, %quant_block, %q6_half, %load_packet) : (buffer, offset, index, index, index) -> (vector<1xi32>, vector<1xi32>, vector<1xi32>) + scf.yield %ql0_word, %ql1_word, %qh_word : vector<1xi32>, vector<1xi32>, vector<1xi32> + } else { + scf.yield %c0_i32x1, %c0_i32x1, %c0_i32x1 : vector<1xi32>, vector<1xi32>, vector<1xi32> + } + %updates_q6_code0 = index.cmp eq, %q6_code_row_offset, %c0 : index + %updates_q6_code1 = index.cmp eq, %q6_code_row_offset, %c16 : index + %updates_q6_code2 = index.cmp eq, %q6_code_row_offset, %c32 : index + %updates_q6_code3 = index.cmp eq, %q6_code_row_offset, %c48 : index + %next_q6_ql00 = scf.select %updates_q6_code0, %loaded_q6_ql0, %prior_q6_ql00 : vector<1xi32> + %next_q6_ql10 = scf.select %updates_q6_code0, %loaded_q6_ql1, %prior_q6_ql10 : vector<1xi32> + %next_q6_qh0 = scf.select %updates_q6_code0, %loaded_q6_qh, %prior_q6_qh0 : vector<1xi32> + %next_q6_ql01 = scf.select %updates_q6_code1, %loaded_q6_ql0, %prior_q6_ql01 : vector<1xi32> + %next_q6_ql11 = scf.select %updates_q6_code1, %loaded_q6_ql1, %prior_q6_ql11 : vector<1xi32> + %next_q6_qh1 = scf.select %updates_q6_code1, %loaded_q6_qh, %prior_q6_qh1 : vector<1xi32> + %next_q6_ql02 = scf.select %updates_q6_code2, %loaded_q6_ql0, %prior_q6_ql02 : vector<1xi32> + %next_q6_ql12 = scf.select %updates_q6_code2, %loaded_q6_ql1, %prior_q6_ql12 : vector<1xi32> + %next_q6_qh2 = scf.select %updates_q6_code2, %loaded_q6_qh, %prior_q6_qh2 : vector<1xi32> + %next_q6_ql03 = scf.select %updates_q6_code3, %loaded_q6_ql0, %prior_q6_ql03 : vector<1xi32> + %next_q6_ql13 = scf.select %updates_q6_code3, %loaded_q6_ql1, %prior_q6_ql13 : vector<1xi32> + %next_q6_qh3 = scf.select %updates_q6_code3, %loaded_q6_qh, %prior_q6_qh3 : vector<1xi32> + scf.yield %next_q6_ql00, %next_q6_ql10, %next_q6_qh0, %next_q6_ql01, %next_q6_ql11, %next_q6_qh1, %next_q6_ql02, %next_q6_ql12, %next_q6_qh2, %next_q6_ql03, %next_q6_ql13, %next_q6_qh3 : vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32> + } + %outer_result00, %outer_result01, %outer_result10, %outer_result11 = scf.for %group_within_outer = [%c0 to %groups_per_outer step %c1](%group_acc00 = %acc00 : vector<8xf16>, %group_acc01 = %acc01 : vector<8xf16>, %group_acc10 = %acc10 : vector<8xf16>, %group_acc11 = %acc11 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16>) { + %quant_group_base = index.mul %quant_group_outer, %groups_per_outer : index + %quant_group0 = index.add %quant_group_base, %group_within_outer : index + %quant_group = index.assume %quant_group0 [range(%quant_group0, 0, 7)] : index + %block_k_base = index.mul %quant_block, %c256 : index + %group_k_add = index.mul %quant_group, %c32 : index + %k_origin = index.add %block_k_base, %group_k_add : index + scf.for %row_offset = [%c0 to %c64 step %c16] unroll { + %local_row0 = index.add %load_row, %row_offset : index + %local_row = index.assume %local_row0 [range(%local_row0, 0, 63)] : index + %selects_row0 = index.cmp eq, %row_offset, %c0 : index + %selects_row2 = index.cmp eq, %row_offset, %c32 : index + %selects_low_row_pair = index.cmp ult, %row_offset, %c32 : index + %selected_header01 = scf.select %selects_row0, %header0, %header1 : vector<4xi32> + %selected_header23 = scf.select %selects_row2, %header2, %header3 : vector<4xi32> + %selected_header = scf.select %selects_low_row_pair, %selected_header01, %selected_header23 : vector<4xi32> + %selected_q_word01 = scf.select %selects_row0, %q_word0, %q_word1 : vector<1xi32> + %selected_q_word23 = scf.select %selects_row2, %q_word2, %q_word3 : vector<1xi32> + %selected_q_word = scf.select %selects_low_row_pair, %selected_q_word01, %selected_q_word23 : vector<1xi32> + %selected_q6_ql0_01 = scf.select %selects_row0, %q6_ql00, %q6_ql01 : vector<1xi32> + %selected_q6_ql0_23 = scf.select %selects_row2, %q6_ql02, %q6_ql03 : vector<1xi32> + %selected_q6_ql0 = scf.select %selects_low_row_pair, %selected_q6_ql0_01, %selected_q6_ql0_23 : vector<1xi32> + %selected_q6_ql1_01 = scf.select %selects_row0, %q6_ql10, %q6_ql11 : vector<1xi32> + %selected_q6_ql1_23 = scf.select %selects_row2, %q6_ql12, %q6_ql13 : vector<1xi32> + %selected_q6_ql1 = scf.select %selects_low_row_pair, %selected_q6_ql1_01, %selected_q6_ql1_23 : vector<1xi32> + %selected_q6_qh01 = scf.select %selects_row0, %q6_qh0, %q6_qh1 : vector<1xi32> + %selected_q6_qh23 = scf.select %selects_row2, %q6_qh2, %q6_qh3 : vector<1xi32> + %selected_q6_qh = scf.select %selects_low_row_pair, %selected_q6_qh01, %selected_q6_qh23 : vector<1xi32> + %channel = index.add %channel_tile_base, %local_row : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %weight_values = scf.if %valid_channel -> (vector<4xf16>) { + %decoded = scf.if %is_q6 -> (vector<4xf16>) { + %channel_byte_add = index.scale %channel, %weight_row_bytes : index, offset -> offset + %row_byte_base = index.add %expert_byte_base, %channel_byte_add : offset + %q6_values = func.call @qwen3_moe_q6k_wmma_vector4_from_half_codes(%weight_noalias, %row_byte_base, %quant_block, %quant_group, %load_packet, %selected_q6_ql0, %selected_q6_ql1, %selected_q6_qh) : (buffer, offset, index, index, index, vector<1xi32>, vector<1xi32>, vector<1xi32>) -> (vector<4xf16>) + scf.yield %q6_values : vector<4xf16> + } else { + %q4_values = func.call @qwen3_moe_q4k_wmma_vector4_from_header_code(%quant_group, %selected_header, %selected_q_word) : (index, vector<4xi32>, vector<1xi32>) -> (vector<4xf16>) + scf.yield %q4_values : vector<4xf16> + } + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + %is_activation_row = index.cmp ult, %local_row, %c32 : index + vector.store %weight_values, %weight_stage_view[%local_row, %load_k] : vector<4xf16>, view<64x40xf16> + scf.if %is_activation_row { + %activation_row = index.assume %local_row [range(%local_row, 0, 31)] : index + %assignment_i32 = view.load %route_stage_view[%activation_row] : view<32xi32> -> i32 + %valid_assignment = scalar.cmpi sge, %assignment_i32, %c0_i32 : i32 + %activation_values = scf.if %valid_assignment -> (vector<4xf16>) { + %assignment0 = index.cast %assignment_i32 : i32 to index + %assignment = index.assume %assignment0 [range(%assignment0, 0, 16383)] : index + %bounded_assignment, %bounded_assignment_count = index.assume %assignment, %assignment_count [lt(%assignment, %assignment_count)] : index, index + %input_k = index.add %k_origin, %load_k : index + %loaded = vector.load %input_view[%bounded_assignment, %input_k] : view<[%assignment_count]x[%input_size]xf16> -> vector<4xf16> + scf.yield %loaded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %activation_values, %activation_stage_physical_view[%activation_row, %load_k] : vector<4xf16>, view<32x40xf16> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %next00, %next01, %next10, %next11 = scf.for %k_half = [%c0 to %c32 step %c16](%half_acc00 = %group_acc00 : vector<8xf16>, %half_acc01 = %group_acc01 : vector<8xf16>, %half_acc10 = %group_acc10 : vector<8xf16>, %half_acc11 = %group_acc11 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16>) unroll { + %lhs0 = vector.fragment.load %weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %lhs1 = vector.fragment.load %weight_stage_view[%subgroup_channel1, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %rhs0 = vector.fragment.load %activation_fragment_view[%k_half, %c0] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %rhs1 = vector.fragment.load %activation_fragment_view[%k_half, %c16] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %half_next00 = vector.mma %lhs0, %rhs0, %half_acc00 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %half_next01 = vector.mma %lhs0, %rhs1, %half_acc01 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %half_next10 = vector.mma %lhs1, %rhs0, %half_acc10 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %half_next11 = vector.mma %lhs1, %rhs1, %half_acc11 : vector<16xf16>, vector<16xf16>, vector<8xf16> + scf.yield %half_next00, %half_next01, %half_next10, %half_next11 : vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %next00, %next01, %next10, %next11 : vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16> + } + scf.yield %outer_result00, %outer_result01, %outer_result10, %outer_result11 : vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16> + } + scf.yield %block_result00, %block_result01, %block_result10, %block_result11 : vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16> + } + // WMMA produces [channel][route] fragments. Transpose through wave-private + // LDS so each lane publishes four adjacent channels for one assignment. + %publish_route0 = index.div %lane, %c4 : index + %publish_route = index.assume %publish_route0 [range(%publish_route0, 0, 15)] : index + %publish_packet0 = index.rem %lane, %c4 : index + %publish_packet = index.assume %publish_packet0 [range(%publish_packet0, 0, 3)] : index + %publish_channel_add = index.mul %publish_packet, %c4 : index + %local_route1 = index.add %c16, %publish_route : index + %assignment0_i32 = view.load %route_stage_view[%publish_route] : view<32xi32> -> i32 + %assignment1_i32 = view.load %route_stage_view[%local_route1] : view<32xi32> -> i32 + %assignment0_nonnegative = scalar.cmpi sge, %assignment0_i32, %c0_i32 : i32 + %assignment1_nonnegative = scalar.cmpi sge, %assignment1_i32, %c0_i32 : i32 + %safe_assignment0_i32 = scf.select %assignment0_nonnegative, %assignment0_i32, %c0_i32 : i32 + %safe_assignment1_i32 = scf.select %assignment1_nonnegative, %assignment1_i32, %c0_i32 : i32 + %safe_assignment0_0 = index.cast %safe_assignment0_i32 : i32 to index + %safe_assignment1_0 = index.cast %safe_assignment1_i32 : i32 to index + %safe_assignment0 = index.assume %safe_assignment0_0 [range(%safe_assignment0_0, 0, 16383)] : index + %safe_assignment1 = index.assume %safe_assignment1_0 [range(%safe_assignment1_0, 0, 16383)] : index + %bounded_assignment0, %bounded_assignment_count0 = index.assume %safe_assignment0, %assignment_count [lt(%safe_assignment0, %assignment_count)] : index, index + %bounded_assignment1, %bounded_assignment_count1 = index.assume %safe_assignment1, %assignment_count [lt(%safe_assignment1, %assignment_count)] : index, index + %subgroup_channel_base = index.add %channel_tile_base, %subgroup_channel_add : index + %channel0 = index.add %subgroup_channel_base, %publish_channel_add : index + %channel1_base = index.add %subgroup_channel_base, %c16 : index + %channel1 = index.add %channel1_base, %publish_channel_add : index + %valid_channel0 = index.cmp ult, %channel0, %bounded_output_size : index + %valid_channel1 = index.cmp ult, %channel1, %bounded_output_size : index + %writes00 = scalar.andi %assignment0_nonnegative, %valid_channel0 : i1 + %writes01 = scalar.andi %assignment1_nonnegative, %valid_channel0 : i1 + %writes10 = scalar.andi %assignment0_nonnegative, %valid_channel1 : i1 + %writes11 = scalar.andi %assignment1_nonnegative, %valid_channel1 : i1 + vector.fragment.store %result00, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<8xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes00 { + %values = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf16> -> vector<4xf16> + %mask = vector.mask.range [%channel0 to %bounded_output_size step %c1] : index -> vector<4xi1> + vector.store.mask %values, %output_view[%bounded_assignment0, %channel0], %mask : vector<4xf16>, view<[%assignment_count]x[%bounded_output_size]xf16>, vector<4xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result01, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<8xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes01 { + %values = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf16> -> vector<4xf16> + %mask = vector.mask.range [%channel0 to %bounded_output_size step %c1] : index -> vector<4xi1> + vector.store.mask %values, %output_view[%bounded_assignment1, %channel0], %mask : vector<4xf16>, view<[%assignment_count]x[%bounded_output_size]xf16>, vector<4xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result10, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<8xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes10 { + %values = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf16> -> vector<4xf16> + %mask = vector.mask.range [%channel1 to %bounded_output_size step %c1] : index -> vector<4xi1> + vector.store.mask %values, %output_view[%bounded_assignment0, %channel1], %mask : vector<4xf16>, view<[%assignment_count]x[%bounded_output_size]xf16>, vector<4xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result11, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<8xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes11 { + %values = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf16> -> vector<4xf16> + %mask = vector.mask.range [%channel1 to %bounded_output_size step %c1] : index -> vector<4xi1> + vector.store.mask %values, %output_view[%bounded_assignment1, %channel1], %mask : vector<4xf16>, view<[%assignment_count]x[%bounded_output_size]xf16>, vector<4xi1> + } + // Route, operand, and result stages are reused by the next concentrated + // routing partition. All waves must finish publication before reuse. + kernel.barrier scope(workgroup) ordering(acq_rel) + } + template.return +} + +// Q4_K and Q6_K retain separate entry points while sharing the complete matrix +// schedule. This keeps format routing outside the hot kernel. +kernel.def target(@qwen3_moe_routed_down_gfx11_wave64) @qwen3_moe_routed_down_q4k_f16_wmma_grouped(%token_count: index) { + %expert_count = config.get @qwen3_moe.routed_down.expert_count : index + %output_size = config.get @qwen3_moe.routed_down.output_size : index + %c1 = index.constant 1 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %padded_token_count = index.add %token_count, %c63 : index + %route_tiles = index.div %padded_token_count, %c64 : index + kernel.launch.config workgroups(%output_tiles, %route_tiles, %expert_count) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %expert_table: buffer, %weight: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %q4 = index.constant 4 : index + template.apply<@qwen3_moe.routed_down.quantized_prefill.body>(%q4, %token_count, %input, %expert_table, %weight, %output) : (index, index, buffer, buffer, buffer, buffer) + kernel.return +} + +kernel.def target(@qwen3_moe_routed_down_gfx11_wave64) @qwen3_moe_routed_down_q6k_f16_wmma_grouped(%token_count: index) { + %expert_count = config.get @qwen3_moe.routed_down.expert_count : index + %output_size = config.get @qwen3_moe.routed_down.output_size : index + %c1 = index.constant 1 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %padded_token_count = index.add %token_count, %c63 : index + %route_tiles = index.div %padded_token_count, %c64 : index + kernel.launch.config workgroups(%output_tiles, %route_tiles, %expert_count) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %expert_table: buffer, %weight: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %q6 = index.constant 6 : index + template.apply<@qwen3_moe.routed_down.quantized_prefill.body>(%q6, %token_count, %input, %expert_table, %weight, %output) : (index, index, buffer, buffer, buffer, buffer) + kernel.return +} + +// Reduces compact routed projections into the residual in logical token order. +// +// One workitem owns four adjacent output channels. All top-k routes are +// accumulated in FP32 before the residual is loaded and published once. The +// preceding grouped projection is the only producer of each FP16 route row, so +// this boundary requires no global atomics. +kernel.def target(@qwen3_moe_routed_down_gfx11_wave64) @qwen3_moe_routed_down_weighted_reduce_f16_f32(%token_count: index) { + %output_size = config.get @qwen3_moe.routed_down.output_size : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %padded_output_size = index.add %output_size, %c256 : index + %rounded_output_size = index.sub %padded_output_size, %c1 : index + %output_tiles = index.div %rounded_output_size, %c256 : index + kernel.launch.config workgroups(%output_tiles, %token_count, %c1) workgroup_size(%c64, %c1, %c1) : index +} launch(%token_count: index, %route_weights: buffer, %routed_output: buffer, %residual_input: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %route_count = config.get @qwen3_moe.routed_down.route_count : index + %output_size = config.get @qwen3_moe.routed_down.output_size : index + %bounded_route_count = index.assume %route_count [range(%route_count, 1, 8)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 4096)] : index + %channel_tile = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %c0_f16x4 = vector.constant 0.0 : vector<4xf16> + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %assignment_count = index.mul %token_count, %bounded_route_count : index + %channel_tile_base = index.mul %channel_tile, %c256 : index + %lane_channel_add = index.mul %lane, %c4 : index + %channel = index.add %channel_tile_base, %lane_channel_add : index + %active_token = index.cmp ult, %token0, %token_count : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %publishes_output = scalar.andi %active_token, %valid_channel : i1 + %route_weights_noalias, %routed_output_noalias, %residual_input_noalias, %output_noalias = buffer.assume.noalias %route_weights, %routed_output, %residual_input, %output : buffer, buffer, buffer, buffer + %route_weight_view = buffer.view %route_weights_noalias[%c0_offset] : buffer -> view<[%token_count]x[%bounded_route_count]xf32> + %routed_output_view = buffer.view %routed_output_noalias[%c0_offset] : buffer -> view<[%assignment_count]x[%bounded_output_size]xf16> + %residual_input_view = buffer.view %residual_input_noalias[%c0_offset] : buffer -> view<[%token_count]x[%bounded_output_size]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%token_count]x[%bounded_output_size]xf32> + scf.if %publishes_output { + %token, %output_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %mask = vector.mask.range [%channel to %bounded_output_size step %c1] : index -> vector<4xi1> + %weighted_sum = scf.for %route = [%c0 to %bounded_route_count step %c1](%sum = %c0_f32x4 : vector<4xf32>) -> (vector<4xf32>) unroll { + %assignment = index.madd %token, %bounded_route_count, %route : index + %routed = vector.load.mask %routed_output_view[%assignment, %channel], %mask, %c0_f16x4 : view<[%assignment_count]x[%bounded_output_size]xf16>, vector<4xi1>, vector<4xf16> + %wide = vector.extf %routed : vector<4xf16> to vector<4xf32> + %route_weight = view.load %route_weight_view[%token, %route] : view<[%token_count]x[%bounded_route_count]xf32> -> f32 + %route_weight_x4 = vector.splat %route_weight : vector<4xf32> + %weighted = vector.mulf %wide, %route_weight_x4 : vector<4xf32> + %next = vector.addf %sum, %weighted : vector<4xf32> + scf.yield %next : vector<4xf32> + } + %residual = vector.load.mask %residual_input_view[%token, %channel], %mask, %c0_f32x4 : view<[%token_count]x[%bounded_output_size]xf32>, vector<4xi1>, vector<4xf32> + %result = vector.addf %residual, %weighted_sum : vector<4xf32> + vector.store.mask %result, %output_view[%token, %channel], %mask : vector<4xf32>, view<[%token_count]x[%bounded_output_size]xf32>, vector<4xi1> + } + kernel.return +} + +// Complete expert-MLP coverage crosses route partitions and output tails. The +// established Q8_1 gate/up and Q4_K down providers form an independent +// reference while the production path keeps the routed SwiGLU handoff in F16. +// The complete-chain tolerance covers both independent quantized accumulation +// regimes; the focused gate/up and down cases constrain each boundary to 1%. +check.case public @qwen3_moe_expert_mlp_q4k_f16_handoff_differential_case { + %token_count = check.literal value(67) : index + %gate_input_size = check.literal value(512) : index + %gate_output_size = check.literal value(256) : index + %down_output_size = check.literal value(33) : index + %routed_row_count = check.literal value(134) : index + %route_count = check.literal value(2) : index + %route_stride = check.literal value(4) : index + %expert_count = check.literal value(4) : index + %input = check.generate.fill value(0.00390625) : tensor<67x512xf32> + %q8_input = check.generate.fill value(0) : tensor<67x576xi8> + %route_ids = check.generate.iota offset(0) step(1) period(4) : tensor<67x4xi32> + %route_weights = check.generate.fill value(0.5) : tensor<67x2xf32> + %expert_table = check.generate.fill value(-1) : tensor<272xi32> + %partition_table = check.generate.fill value(-1) : tensor<10xi32> + %gate_weight = check.generate.fill value(34) : tensor<4x256x2x144xi8> + %up_weight = check.generate.fill value(35) : tensor<4x256x2x144xi8> + %down_weight = check.generate.fill value(-86) : tensor<4x33x1x144xi8> + %reference_gate_output = check.generate.fill value(1.0) : tensor<67x2x256xf32> + %q8_down_input = check.generate.fill value(0) : tensor<67x2x288xi8> + %expected_output = check.generate.fill value(1.0) : tensor<67x33xf32> + %actual_gate_output = check.generate.fill value(1.0) : tensor<67x2x256xf16> + %actual_routed_output = check.generate.fill value(0.0) : tensor<67x2x33xf16> + %actual_output = check.generate.fill value(1.0) : tensor<67x33xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %gate_input_size](%token_count, %gate_input_size, %input, %q8_input) : [index, index](index, index, tensor<67x512xf32>, tensor<67x576xi8>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8[%token_count, %route_count, %route_stride, %expert_count, %gate_output_size](%token_count, %route_count, %route_stride, %expert_count, %gate_output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %reference_gate_output) : [index, index, index, index, index](index, index, index, index, index, tensor<67x576xi8>, tensor<67x4xi32>, tensor<4x256x2x144xi8>, tensor<4x256x2x144xi8>, tensor<67x2x256xf32>) + kernel.launch @ggml_quantize_q8_1_x4_f32[%routed_row_count, %gate_output_size](%routed_row_count, %gate_output_size, %reference_gate_output, %q8_down_input) : [index, index](index, index, tensor<67x2x256xf32>, tensor<67x2x288xi8>) + kernel.launch @qwen3_moe_routed_down_q4k_q8_1_x4[%token_count, %gate_output_size, %route_count, %route_stride, %expert_count, %down_output_size](%token_count, %gate_output_size, %route_count, %route_stride, %expert_count, %down_output_size, %q8_down_input, %route_ids, %route_weights, %down_weight, %expected_output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<67x2x288xi8>, tensor<67x4xi32>, tensor<67x2xf32>, tensor<4x33x1x144xi8>, tensor<67x33xf32>) + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<67x4xi32>, tensor<272xi32>) + kernel.launch @qwen3_moe_build_expert_partition_table[%token_count, %route_count, %expert_count](%token_count, %route_count, %expert_count, %expert_table, %partition_table) : [index, index, index](index, index, index, tensor<272xi32>, tensor<10xi32>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma[%token_count](%token_count, %input, %expert_table, %partition_table, %gate_weight, %up_weight, %actual_gate_output) : [index](index, tensor<67x512xf32>, tensor<272xi32>, tensor<10xi32>, tensor<4x256x2x144xi8>, tensor<4x256x2x144xi8>, tensor<67x2x256xf16>) + kernel.launch @qwen3_moe_routed_down_q4k_f16_wmma_grouped[%token_count](%token_count, %actual_gate_output, %expert_table, %down_weight, %actual_routed_output) : [index](index, tensor<67x2x256xf16>, tensor<272xi32>, tensor<4x33x1x144xi8>, tensor<67x2x33xf16>) + kernel.launch @qwen3_moe_routed_down_weighted_reduce_f16_f32[%token_count](%token_count, %route_weights, %actual_routed_output, %actual_output, %actual_output) : [index](index, tensor<67x2xf32>, tensor<67x2x33xf16>, tensor<67x33xf32>, tensor<67x33xf32>) + check.expect.close actual(%actual_output) expected(%expected_output) atol(0.5) rtol(0.05) nan(same) : tensor<67x33xf32> + check.return +} + +// Q6_K shares the production matrix schedule but has an independent packed +// decoder and Q8_1 reference provider. +check.case public @qwen3_moe_expert_mlp_q6k_f16_handoff_differential_case { + %token_count = check.literal value(67) : index + %gate_input_size = check.literal value(512) : index + %gate_output_size = check.literal value(256) : index + %down_output_size = check.literal value(33) : index + %routed_row_count = check.literal value(134) : index + %route_count = check.literal value(2) : index + %route_stride = check.literal value(4) : index + %expert_count = check.literal value(4) : index + %input = check.generate.fill value(0.00390625) : tensor<67x512xf32> + %q8_input = check.generate.fill value(0) : tensor<67x576xi8> + %route_ids = check.generate.iota offset(0) step(1) period(4) : tensor<67x4xi32> + %route_weights = check.generate.fill value(0.5) : tensor<67x2xf32> + %expert_table = check.generate.fill value(-1) : tensor<272xi32> + %partition_table = check.generate.fill value(-1) : tensor<10xi32> + %gate_weight = check.generate.fill value(34) : tensor<4x256x2x144xi8> + %up_weight = check.generate.fill value(35) : tensor<4x256x2x144xi8> + %down_weight = check.generate.fill value(-86) : tensor<4x33x1x210xi8> + %reference_gate_output = check.generate.fill value(1.0) : tensor<67x2x256xf32> + %q8_down_input = check.generate.fill value(0) : tensor<67x2x288xi8> + %expected_output = check.generate.fill value(1.0) : tensor<67x33xf32> + %actual_gate_output = check.generate.fill value(1.0) : tensor<67x2x256xf16> + %actual_routed_output = check.generate.fill value(0.0) : tensor<67x2x33xf16> + %actual_output = check.generate.fill value(1.0) : tensor<67x33xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %gate_input_size](%token_count, %gate_input_size, %input, %q8_input) : [index, index](index, index, tensor<67x512xf32>, tensor<67x576xi8>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8[%token_count, %route_count, %route_stride, %expert_count, %gate_output_size](%token_count, %route_count, %route_stride, %expert_count, %gate_output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %reference_gate_output) : [index, index, index, index, index](index, index, index, index, index, tensor<67x576xi8>, tensor<67x4xi32>, tensor<4x256x2x144xi8>, tensor<4x256x2x144xi8>, tensor<67x2x256xf32>) + kernel.launch @ggml_quantize_q8_1_x4_f32[%routed_row_count, %gate_output_size](%routed_row_count, %gate_output_size, %reference_gate_output, %q8_down_input) : [index, index](index, index, tensor<67x2x256xf32>, tensor<67x2x288xi8>) + kernel.launch @qwen3_moe_routed_down_q6k_q8_1_x4[%token_count, %gate_output_size, %route_count, %route_stride, %expert_count, %down_output_size](%token_count, %gate_output_size, %route_count, %route_stride, %expert_count, %down_output_size, %q8_down_input, %route_ids, %route_weights, %down_weight, %expected_output) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<67x2x288xi8>, tensor<67x4xi32>, tensor<67x2xf32>, tensor<4x33x1x210xi8>, tensor<67x33xf32>) + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<67x4xi32>, tensor<272xi32>) + kernel.launch @qwen3_moe_build_expert_partition_table[%token_count, %route_count, %expert_count](%token_count, %route_count, %expert_count, %expert_table, %partition_table) : [index, index, index](index, index, index, tensor<272xi32>, tensor<10xi32>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma[%token_count](%token_count, %input, %expert_table, %partition_table, %gate_weight, %up_weight, %actual_gate_output) : [index](index, tensor<67x512xf32>, tensor<272xi32>, tensor<10xi32>, tensor<4x256x2x144xi8>, tensor<4x256x2x144xi8>, tensor<67x2x256xf16>) + kernel.launch @qwen3_moe_routed_down_q6k_f16_wmma_grouped[%token_count](%token_count, %actual_gate_output, %expert_table, %down_weight, %actual_routed_output) : [index](index, tensor<67x2x256xf16>, tensor<272xi32>, tensor<4x33x1x210xi8>, tensor<67x2x33xf16>) + kernel.launch @qwen3_moe_routed_down_weighted_reduce_f16_f32[%token_count](%token_count, %route_weights, %actual_routed_output, %actual_output, %actual_output) : [index](index, tensor<67x2xf32>, tensor<67x2x33xf16>, tensor<67x33xf32>, tensor<67x33xf32>) + check.expect.close actual(%actual_output) expected(%expected_output) atol(0.5) rtol(0.05) nan(same) : tensor<67x33xf32> + check.return +} + +// Differential coverage crosses token, route, expert, channel, and output-tail +// boundaries. The direct Q8_1 provider supplies an independent raw-Q4_K +// reference while the grouped provider exercises collision-free FP16 +// publication and a separate weighted FP32 residual reduction. +check.case public @qwen3_moe_routed_down_q4k_f16_wmma_differential_case { + %token_count = check.literal value(67) : index + %input_size = check.literal value(512) : index + %route_count = check.literal value(2) : index + %route_stride = check.literal value(4) : index + %expert_count = check.literal value(4) : index + %output_size = check.literal value(33) : index + %routed_row_count = check.literal value(134) : index + %input_f32 = check.generate.fill value(0.00390625) : tensor<67x2x512xf32> + %input_f16 = check.generate.fill value(0.00390625) : tensor<67x2x512xf16> + %q8_input = check.generate.fill value(0) : tensor<67x2x576xi8> + %route_ids = check.generate.iota offset(0) step(1) period(4) : tensor<67x4xi32> + %route_weights = check.generate.fill value(0.5) : tensor<67x2xf32> + %expert_table = check.generate.fill value(-1) : tensor<272xi32> + %weight = check.generate.fill value(-86) : tensor<4x33x2x144xi8> + %expected = check.generate.fill value(1.0) : tensor<67x33xf32> + %routed_output = check.generate.fill value(0.0) : tensor<67x2x33xf16> + %actual = check.generate.fill value(1.0) : tensor<67x33xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%routed_row_count, %input_size](%routed_row_count, %input_size, %input_f32, %q8_input) : [index, index](index, index, tensor<67x2x512xf32>, tensor<67x2x576xi8>) + kernel.launch @qwen3_moe_routed_down_q4k_q8_1_x4[%token_count, %input_size, %route_count, %route_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %expected) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<67x2x576xi8>, tensor<67x4xi32>, tensor<67x2xf32>, tensor<4x33x2x144xi8>, tensor<67x33xf32>) + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<67x4xi32>, tensor<272xi32>) + kernel.launch @qwen3_moe_routed_down_q4k_f16_wmma_grouped[%token_count](%token_count, %input_f16, %expert_table, %weight, %routed_output) : [index](index, tensor<67x2x512xf16>, tensor<272xi32>, tensor<4x33x2x144xi8>, tensor<67x2x33xf16>) + kernel.launch @qwen3_moe_routed_down_weighted_reduce_f16_f32[%token_count](%token_count, %route_weights, %routed_output, %actual, %actual) : [index](index, tensor<67x2xf32>, tensor<67x2x33xf16>, tensor<67x33xf32>, tensor<67x33xf32>) + check.expect.close actual(%actual) expected(%expected) atol(0.5) rtol(0.01) nan(same) : tensor<67x33xf32> + check.return +} + +// Q6_K follows the identical routed matrix schedule while selecting its +// independent packed-row decoder. The same irregular shape prevents format +// specialization from weakening routing or tail coverage. +check.case public @qwen3_moe_routed_down_q6k_f16_wmma_differential_case { + %token_count = check.literal value(67) : index + %input_size = check.literal value(512) : index + %route_count = check.literal value(2) : index + %route_stride = check.literal value(4) : index + %expert_count = check.literal value(4) : index + %output_size = check.literal value(33) : index + %routed_row_count = check.literal value(134) : index + %input_f32 = check.generate.fill value(0.00390625) : tensor<67x2x512xf32> + %input_f16 = check.generate.fill value(0.00390625) : tensor<67x2x512xf16> + %q8_input = check.generate.fill value(0) : tensor<67x2x576xi8> + %route_ids = check.generate.iota offset(0) step(1) period(4) : tensor<67x4xi32> + %route_weights = check.generate.fill value(0.5) : tensor<67x2xf32> + %expert_table = check.generate.fill value(-1) : tensor<272xi32> + %weight = check.generate.fill value(-86) : tensor<4x33x2x210xi8> + %expected = check.generate.fill value(1.0) : tensor<67x33xf32> + %routed_output = check.generate.fill value(0.0) : tensor<67x2x33xf16> + %actual = check.generate.fill value(1.0) : tensor<67x33xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%routed_row_count, %input_size](%routed_row_count, %input_size, %input_f32, %q8_input) : [index, index](index, index, tensor<67x2x512xf32>, tensor<67x2x576xi8>) + kernel.launch @qwen3_moe_routed_down_q6k_q8_1_x4[%token_count, %input_size, %route_count, %route_stride, %expert_count, %output_size](%token_count, %input_size, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %route_weights, %weight, %expected) : [index, index, index, index, index, index](index, index, index, index, index, index, tensor<67x2x576xi8>, tensor<67x4xi32>, tensor<67x2xf32>, tensor<4x33x2x210xi8>, tensor<67x33xf32>) + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<67x4xi32>, tensor<272xi32>) + kernel.launch @qwen3_moe_routed_down_q6k_f16_wmma_grouped[%token_count](%token_count, %input_f16, %expert_table, %weight, %routed_output) : [index](index, tensor<67x2x512xf16>, tensor<272xi32>, tensor<4x33x2x210xi8>, tensor<67x2x33xf16>) + kernel.launch @qwen3_moe_routed_down_weighted_reduce_f16_f32[%token_count](%token_count, %route_weights, %routed_output, %actual, %actual) : [index](index, tensor<67x2xf32>, tensor<67x2x33xf16>, tensor<67x33xf32>, tensor<67x33xf32>) + check.expect.close actual(%actual) expected(%expected) atol(0.5) rtol(0.01) nan(same) : tensor<67x33xf32> + check.return +} + +check.case public @qwen3_moe_routed_down_q4k_f16_wmma_benchmark_case { + %token_count = check.param.choice values([1, 2, 4, 8, 16, 17, 32, 63, 128, 129, 512]) name("token_count") : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.iota offset(0) step(1) period(127) : tensor<[%token_count]x128xi32> + %route_weights = check.generate.fill value(0.125) : tensor<[%token_count]x8xf32> + %expert_table = check.generate.fill value(-1) : tensor<65664xi32> + %input = check.generate.fill value(0.0) : tensor<[%token_count]x8x768xf16> + %weight = check.generate.fill value(0) : tensor<128x2048x3x144xi8> + %routed_output = check.generate.fill value(0.0) : tensor<[%token_count]x8x2048xf16> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %expected = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<[%token_count]x128xi32>, tensor<65664xi32>) + kernel.launch @qwen3_moe_routed_down_q4k_f16_wmma_grouped[%token_count](%token_count, %input, %expert_table, %weight, %routed_output) : [index](index, tensor<[%token_count]x8x768xf16>, tensor<65664xi32>, tensor<128x2048x3x144xi8>, tensor<[%token_count]x8x2048xf16>) + kernel.launch @qwen3_moe_routed_down_weighted_reduce_f16_f32[%token_count](%token_count, %route_weights, %routed_output, %output, %output) : [index](index, tensor<[%token_count]x8xf32>, tensor<[%token_count]x8x2048xf16>, tensor<[%token_count]x2048xf32>, tensor<[%token_count]x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x2048xf32> + check.return +} + +check.case public @qwen3_moe_routed_down_q6k_f16_wmma_benchmark_case { + %token_count = check.param.choice values([1, 2, 4, 8, 16, 17, 32, 63, 128, 129, 512]) name("token_count") : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.iota offset(0) step(1) period(127) : tensor<[%token_count]x128xi32> + %route_weights = check.generate.fill value(0.125) : tensor<[%token_count]x8xf32> + %expert_table = check.generate.fill value(-1) : tensor<65664xi32> + %input = check.generate.fill value(0.0) : tensor<[%token_count]x8x768xf16> + %weight = check.generate.fill value(0) : tensor<128x2048x3x210xi8> + %routed_output = check.generate.fill value(0.0) : tensor<[%token_count]x8x2048xf16> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %expected = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<[%token_count]x128xi32>, tensor<65664xi32>) + kernel.launch @qwen3_moe_routed_down_q6k_f16_wmma_grouped[%token_count](%token_count, %input, %expert_table, %weight, %routed_output) : [index](index, tensor<[%token_count]x8x768xf16>, tensor<65664xi32>, tensor<128x2048x3x210xi8>, tensor<[%token_count]x8x2048xf16>) + kernel.launch @qwen3_moe_routed_down_weighted_reduce_f16_f32[%token_count](%token_count, %route_weights, %routed_output, %output, %output) : [index](index, tensor<[%token_count]x8xf32>, tensor<[%token_count]x8x2048xf16>, tensor<[%token_count]x2048xf32>, tensor<[%token_count]x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x2048xf32> + check.return +} + +// Maximally diverse decode-batch controls. Compact route storage makes the +// flattened 0..127 iota assign every M=16 route to a different expert. +check.case public @qwen3_moe_routed_down_q4k_f16_wmma_diverse_benchmark_case { + %token_count = check.param.choice values([1, 2, 4, 8, 16]) name("token_count") : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.iota offset(0) step(1) period(128) : tensor<[%token_count]x8xi32> + %route_weights = check.generate.fill value(0.125) : tensor<[%token_count]x8xf32> + %expert_table = check.generate.fill value(-1) : tensor<65664xi32> + %input = check.generate.fill value(0.0) : tensor<[%token_count]x8x768xf16> + %weight = check.generate.fill value(0) : tensor<128x2048x3x144xi8> + %routed_output = check.generate.fill value(0.0) : tensor<[%token_count]x8x2048xf16> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %expected = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<[%token_count]x8xi32>, tensor<65664xi32>) + kernel.launch @qwen3_moe_routed_down_q4k_f16_wmma_grouped[%token_count](%token_count, %input, %expert_table, %weight, %routed_output) : [index](index, tensor<[%token_count]x8x768xf16>, tensor<65664xi32>, tensor<128x2048x3x144xi8>, tensor<[%token_count]x8x2048xf16>) + kernel.launch @qwen3_moe_routed_down_weighted_reduce_f16_f32[%token_count](%token_count, %route_weights, %routed_output, %output, %output) : [index](index, tensor<[%token_count]x8xf32>, tensor<[%token_count]x8x2048xf16>, tensor<[%token_count]x2048xf32>, tensor<[%token_count]x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x2048xf32> + check.return +} + +check.case public @qwen3_moe_routed_down_q6k_f16_wmma_diverse_benchmark_case { + %token_count = check.param.choice values([1, 2, 4, 8, 16]) name("token_count") : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %route_ids = check.generate.iota offset(0) step(1) period(128) : tensor<[%token_count]x8xi32> + %route_weights = check.generate.fill value(0.125) : tensor<[%token_count]x8xf32> + %expert_table = check.generate.fill value(-1) : tensor<65664xi32> + %input = check.generate.fill value(0.0) : tensor<[%token_count]x8x768xf16> + %weight = check.generate.fill value(0) : tensor<128x2048x3x210xi8> + %routed_output = check.generate.fill value(0.0) : tensor<[%token_count]x8x2048xf16> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %expected = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<[%token_count]x8xi32>, tensor<65664xi32>) + kernel.launch @qwen3_moe_routed_down_q6k_f16_wmma_grouped[%token_count](%token_count, %input, %expert_table, %weight, %routed_output) : [index](index, tensor<[%token_count]x8x768xf16>, tensor<65664xi32>, tensor<128x2048x3x210xi8>, tensor<[%token_count]x8x2048xf16>) + kernel.launch @qwen3_moe_routed_down_weighted_reduce_f16_f32[%token_count](%token_count, %route_weights, %routed_output, %output, %output) : [index](index, tensor<[%token_count]x8xf32>, tensor<[%token_count]x8x2048xf16>, tensor<[%token_count]x2048xf32>, tensor<[%token_count]x2048xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x2048xf32> + check.return +} + +check.benchmark<@qwen3_moe_routed_down_q4k_f16_wmma_differential_case> @qwen3_moe_routed_down_q4k_f16_wmma_small + +check.benchmark<@qwen3_moe_routed_down_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_down_q4k_f16_wmma_decode {token_count = 1} + +check.benchmark<@qwen3_moe_routed_down_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_down_q4k_f16_wmma_small_batch_2 {token_count = 2} + +check.benchmark<@qwen3_moe_routed_down_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_down_q4k_f16_wmma_small_batch_4 {token_count = 4} + +check.benchmark<@qwen3_moe_routed_down_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_down_q4k_f16_wmma_small_batch_8 {token_count = 8} + +check.benchmark<@qwen3_moe_routed_down_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_down_q4k_f16_wmma_small_batch_16 {token_count = 16} + +check.benchmark<@qwen3_moe_routed_down_q4k_f16_wmma_diverse_benchmark_case> @qwen3_moe_routed_down_q4k_f16_wmma_diverse_decode {token_count = 1} + +check.benchmark<@qwen3_moe_routed_down_q4k_f16_wmma_diverse_benchmark_case> @qwen3_moe_routed_down_q4k_f16_wmma_diverse_small_batch_2 {token_count = 2} + +check.benchmark<@qwen3_moe_routed_down_q4k_f16_wmma_diverse_benchmark_case> @qwen3_moe_routed_down_q4k_f16_wmma_diverse_small_batch_4 {token_count = 4} + +check.benchmark<@qwen3_moe_routed_down_q4k_f16_wmma_diverse_benchmark_case> @qwen3_moe_routed_down_q4k_f16_wmma_diverse_small_batch_8 {token_count = 8} + +check.benchmark<@qwen3_moe_routed_down_q4k_f16_wmma_diverse_benchmark_case> @qwen3_moe_routed_down_q4k_f16_wmma_diverse_small_batch_16 {token_count = 16} + +check.benchmark<@qwen3_moe_routed_down_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_down_q4k_f16_wmma_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_routed_down_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_down_q4k_f16_wmma_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_routed_down_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_down_q4k_f16_wmma_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_routed_down_q6k_f16_wmma_differential_case> @qwen3_moe_routed_down_q6k_f16_wmma_small + +check.benchmark<@qwen3_moe_routed_down_q6k_f16_wmma_benchmark_case> @qwen3_moe_routed_down_q6k_f16_wmma_decode {token_count = 1} + +check.benchmark<@qwen3_moe_routed_down_q6k_f16_wmma_benchmark_case> @qwen3_moe_routed_down_q6k_f16_wmma_small_batch_2 {token_count = 2} + +check.benchmark<@qwen3_moe_routed_down_q6k_f16_wmma_benchmark_case> @qwen3_moe_routed_down_q6k_f16_wmma_small_batch_4 {token_count = 4} + +check.benchmark<@qwen3_moe_routed_down_q6k_f16_wmma_benchmark_case> @qwen3_moe_routed_down_q6k_f16_wmma_small_batch_8 {token_count = 8} + +check.benchmark<@qwen3_moe_routed_down_q6k_f16_wmma_benchmark_case> @qwen3_moe_routed_down_q6k_f16_wmma_small_batch_16 {token_count = 16} + +check.benchmark<@qwen3_moe_routed_down_q6k_f16_wmma_diverse_benchmark_case> @qwen3_moe_routed_down_q6k_f16_wmma_diverse_decode {token_count = 1} + +check.benchmark<@qwen3_moe_routed_down_q6k_f16_wmma_diverse_benchmark_case> @qwen3_moe_routed_down_q6k_f16_wmma_diverse_small_batch_2 {token_count = 2} + +check.benchmark<@qwen3_moe_routed_down_q6k_f16_wmma_diverse_benchmark_case> @qwen3_moe_routed_down_q6k_f16_wmma_diverse_small_batch_4 {token_count = 4} + +check.benchmark<@qwen3_moe_routed_down_q6k_f16_wmma_diverse_benchmark_case> @qwen3_moe_routed_down_q6k_f16_wmma_diverse_small_batch_8 {token_count = 8} + +check.benchmark<@qwen3_moe_routed_down_q6k_f16_wmma_diverse_benchmark_case> @qwen3_moe_routed_down_q6k_f16_wmma_diverse_small_batch_16 {token_count = 16} + +check.benchmark<@qwen3_moe_routed_down_q6k_f16_wmma_benchmark_case> @qwen3_moe_routed_down_q6k_f16_wmma_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_routed_down_q6k_f16_wmma_benchmark_case> @qwen3_moe_routed_down_q6k_f16_wmma_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_routed_down_q6k_f16_wmma_benchmark_case> @qwen3_moe_routed_down_q6k_f16_wmma_prefill_512 {token_count = 512} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_f32.loom new file mode 100644 index 000000000000..cf71ac51b464 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_f32.loom @@ -0,0 +1,160 @@ +// Fused grouped-prefill routed residual reduction and next-layer RMSNorm. +// +// One 256-workitem workgroup, comprising eight wave32 subgroups, owns one +// complete 2048-channel token row. Each workitem retains eight adjacent FP32 +// residual results across the workgroup RMS reduction, then publishes both +// the updated hidden state and the learned-weight normalized input consumed +// by the next layer. +amdgpu.target @qwen3_moe_routed_down_next_norm_gfx11_wave32 {subgroup_size = 32} + +config.decl @qwen3_moe.model.hidden_size : %value: index where [range(%value, 128, 32768), mul(%value, 128)] + +config.decl @qwen3_moe.model.rms_epsilon : f32 + +config.decl @qwen3_moe.routed_down.route_count : %value: index where [range(%value, 1, 8)] + +config.decl @qwen3_moe.routed_down.output_size : %value: index where [range(%value, 1, 4096)] + +kernel.decl @qwen3_moe_routed_down_weighted_reduce_f16_f32(%token_count$4: index) launch(%token_count$5: index, %route_weights: buffer, %routed_output: buffer, %residual_input: buffer, %output: buffer) + +kernel.decl @qwen3_moe_rmsnorm_f32(%token_count$9: index) launch(%token_count$10: index, %input: buffer, %weight: buffer, %output: buffer) + +kernel.def target(@qwen3_moe_routed_down_next_norm_gfx11_wave32) @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32(%token_count: index) { + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%token_count, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %route_weights: buffer, %routed_output: buffer, %hidden_state: buffer, %next_norm_weight: buffer, %next_projection_input: buffer) where [range(%token_count, 1, 2048)] { + %route_count0 = config.get @qwen3_moe.routed_down.route_count : index + %output_size0 = config.get @qwen3_moe.routed_down.output_size : index + %hidden_size0 = config.get @qwen3_moe.model.hidden_size : index + %epsilon = config.get @qwen3_moe.model.rms_epsilon : f32 + %route_count = index.assume %route_count0 [range(%route_count0, 8, 8)] : index + %output_size, %hidden_size = index.assume %output_size0, %hidden_size0 [range(%output_size0, 2048, 2048), range(%hidden_size0, 2048, 2048), eq(%output_size0, %hidden_size0)] : index, index + %token0 = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 255)] : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c8 = index.constant 8 : index + %channel = index.mul %workitem, %c8 : index + %scratch_bytes = index.constant 32 : offset + %c0_offset = index.constant 0 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %c0_f32x8 = vector.constant 0.0 : vector<8xf32> + %active_token = index.cmp ult, %token0, %token_count : index + %safe_token0 = scf.select %active_token, %token0, %c0 : index + %token, %launch_token_count = index.assume %safe_token0, %token_count [lt(%safe_token0, %token_count)] : index, index + %hidden_size_i32 = index.cast %hidden_size : index to i32 + %hidden_size_f32 = scalar.sitofp %hidden_size_i32 : i32 to f32 + %assignment_count = index.mul %launch_token_count, %route_count : index + %route_weights_noalias, %routed_output_noalias, %hidden_state_noalias, %next_norm_weight_noalias, %next_projection_input_noalias = buffer.assume.noalias %route_weights, %routed_output, %hidden_state, %next_norm_weight, %next_projection_input : buffer, buffer, buffer, buffer, buffer + %route_weights_view = buffer.view %route_weights_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%route_count]xf32> + %routed_output_view = buffer.view %routed_output_noalias[%c0_offset] : buffer -> view<[%assignment_count]x[%output_size]xf16> + %hidden_state_view = buffer.view %hidden_state_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%hidden_size]xf32> + %next_norm_weight_view = buffer.view %next_norm_weight_noalias[%c0_offset] : buffer -> view<[%hidden_size]xf32> + %next_projection_input_view = buffer.view %next_projection_input_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%hidden_size]xf32> + %scratch = buffer.alloca align(16) %scratch_bytes : buffer + %scratch_view = buffer.view %scratch[%c0_offset] : buffer -> view<8xf32> + + // The predicate depends only on the workgroup ID, so every workitem either + // executes the complete barrier-bearing row reduction or skips it together. + scf.if %active_token { + %weighted_sum = scf.for %route = [%c0 to %route_count step %c1](%sum = %c0_f32x8 : vector<8xf32>) -> (vector<8xf32>) unroll { + %assignment = index.madd %token, %route_count, %route : index + %routed_f16 = vector.load %routed_output_view[%assignment, %channel] : view<[%assignment_count]x[%output_size]xf16> -> vector<8xf16> + %routed = vector.extf %routed_f16 : vector<8xf16> to vector<8xf32> + %route_weight = view.load %route_weights_view[%token, %route] : view<[%launch_token_count]x[%route_count]xf32> -> f32 + %route_weight_x8 = vector.splat %route_weight : vector<8xf32> + %weighted = vector.mulf %routed, %route_weight_x8 : vector<8xf32> + %next = vector.addf %sum, %weighted : vector<8xf32> + scf.yield %next : vector<8xf32> + } + %residual = vector.load %hidden_state_view[%token, %channel] : view<[%launch_token_count]x[%hidden_size]xf32> -> vector<8xf32> + %result = vector.addf %residual, %weighted_sum : vector<8xf32> + %squares = vector.mulf %result, %result : vector<8xf32> + %thread_sum = vector.reduce %squares, %c0_f32 : vector<8xf32>, f32 + %subgroup_sum = kernel.subgroup.reduce %thread_sum : f32 + + %is_subgroup_leader = index.cmp eq, %lane, %c0 : index + scf.if %is_subgroup_leader { + view.store %subgroup_sum, %scratch_view[%subgroup] : f32, view<8xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %is_reduction_subgroup = index.cmp eq, %subgroup, %c0 : index + %is_reduction_lane = index.cmp ult, %lane, %c8 : index + %loads_subgroup_sum = scalar.andi %is_reduction_subgroup, %is_reduction_lane : i1 + %subgroup_partial = scf.if %loads_subgroup_sum -> (f32) { + %value = view.load %scratch_view[%lane] : view<8xf32> -> f32 + scf.yield %value : f32 + } else { + scf.yield %c0_f32 : f32 + } + %row_sum = kernel.subgroup.reduce %subgroup_partial : f32 + %writes_scale = scalar.andi %is_reduction_subgroup, %is_subgroup_leader : i1 + scf.if %writes_scale { + %mean = scalar.divf %row_sum, %hidden_size_f32 : f32 + %biased_mean = scalar.addf %mean, %epsilon : f32 + %scale = scalar.rsqrtf %biased_mean : f32 + view.store %scale, %scratch_view[%c0] : f32, view<8xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + + %scale = view.load %scratch_view[%c0] : view<8xf32> -> f32 + %scale_x8 = vector.splat %scale : vector<8xf32> + %learned_weight = vector.load %next_norm_weight_view[%channel] : view<[%hidden_size]xf32> -> vector<8xf32> + %normalized = vector.mulf %result, %scale_x8 : vector<8xf32> + %next_projection = vector.mulf %normalized, %learned_weight : vector<8xf32> + vector.store %result, %hidden_state_view[%token, %channel] : vector<8xf32>, view<[%launch_token_count]x[%hidden_size]xf32> + vector.store %next_projection, %next_projection_input_view[%token, %channel] : vector<8xf32>, view<[%launch_token_count]x[%hidden_size]xf32> + } + kernel.return +} + +// The exact Prefill-512 shape compares both published tensors against the +// ordinary weighted reduction followed by the canonical RMSNorm. +check.case public @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32_differential_case { + %token_count = check.literal value(512) : index + %route_weights = check.generate.iota offset(0.03125) step(0.015625) period(8) : tensor<512x8xf32> + %routed_output = check.generate.iota offset(-0.25) step(0.015625) period(31) : tensor<512x8x2048xf16> + %expected_hidden_state = check.generate.iota offset(-0.5) step(0.0078125) period(17) : tensor<512x2048xf32> + %actual_hidden_state = check.generate.iota offset(-0.5) step(0.0078125) period(17) : tensor<512x2048xf32> + %next_norm_weight = check.generate.iota offset(0.5) step(0.0078125) period(29) : tensor<2048xf32> + %expected_next_projection_input = check.generate.fill value(0.0) : tensor<512x2048xf32> + %actual_next_projection_input = check.generate.fill value(1.0) : tensor<512x2048xf32> + kernel.launch @qwen3_moe_routed_down_weighted_reduce_f16_f32[%token_count](%token_count, %route_weights, %routed_output, %expected_hidden_state, %expected_hidden_state) : [index](index, tensor<512x8xf32>, tensor<512x8x2048xf16>, tensor<512x2048xf32>, tensor<512x2048xf32>) + kernel.launch @qwen3_moe_rmsnorm_f32[%token_count](%token_count, %expected_hidden_state, %next_norm_weight, %expected_next_projection_input) : [index](index, tensor<512x2048xf32>, tensor<2048xf32>, tensor<512x2048xf32>) + kernel.launch @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32[%token_count](%token_count, %route_weights, %routed_output, %actual_hidden_state, %next_norm_weight, %actual_next_projection_input) : [index](index, tensor<512x8xf32>, tensor<512x8x2048xf16>, tensor<512x2048xf32>, tensor<2048xf32>, tensor<512x2048xf32>) + check.expect.close actual(%actual_hidden_state) expected(%expected_hidden_state) atol(0.0001) rtol(0.0001) nan(same) : tensor<512x2048xf32> + check.expect.close actual(%actual_next_projection_input) expected(%expected_next_projection_input) atol(0.001) rtol(0.001) nan(same) : tensor<512x2048xf32> + check.return +} + +check.case public @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32_fused_benchmark_case { + %token_count = check.literal value(512) : index + %route_weights = check.generate.fill value(0.125) : tensor<512x8xf32> + %routed_output = check.generate.fill value(0.5) : tensor<512x8x2048xf16> + %hidden_state = check.generate.fill value(1.0) : tensor<512x2048xf32> + %next_norm_weight = check.generate.fill value(1.0) : tensor<2048xf32> + %next_projection_input = check.generate.fill value(0.0) : tensor<512x2048xf32> + kernel.launch @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32[%token_count](%token_count, %route_weights, %routed_output, %hidden_state, %next_norm_weight, %next_projection_input) : [index](index, tensor<512x8xf32>, tensor<512x8x2048xf16>, tensor<512x2048xf32>, tensor<2048xf32>, tensor<512x2048xf32>) + check.return +} + +check.case public @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32_composed_benchmark_case { + %token_count = check.literal value(512) : index + %route_weights = check.generate.fill value(0.125) : tensor<512x8xf32> + %routed_output = check.generate.fill value(0.5) : tensor<512x8x2048xf16> + %hidden_state = check.generate.fill value(1.0) : tensor<512x2048xf32> + %next_norm_weight = check.generate.fill value(1.0) : tensor<2048xf32> + %next_projection_input = check.generate.fill value(0.0) : tensor<512x2048xf32> + kernel.launch @qwen3_moe_routed_down_weighted_reduce_f16_f32[%token_count](%token_count, %route_weights, %routed_output, %hidden_state, %hidden_state) : [index](index, tensor<512x8xf32>, tensor<512x8x2048xf16>, tensor<512x2048xf32>, tensor<512x2048xf32>) + kernel.launch @qwen3_moe_rmsnorm_f32[%token_count](%token_count, %hidden_state, %next_norm_weight, %next_projection_input) : [index](index, tensor<512x2048xf32>, tensor<2048xf32>, tensor<512x2048xf32>) + check.return +} + +check.benchmark<@qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32_fused_benchmark_case> @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32_fused_prefill_512 + +check.benchmark<@qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32_composed_benchmark_case> @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32_composed_prefill_512 diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_q8_1_x4.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_q8_1_x4.loom new file mode 100644 index 000000000000..9d452beccb71 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_down_weighted_reduce_next_rmsnorm_q8_1_x4.loom @@ -0,0 +1,269 @@ +// Fused grouped-prefill routed residual reduction, next-layer RMSNorm, and +// GGML Q8_1 x4 publication. +// +// One 256-workitem workgroup, comprising eight wave32 subgroups, owns one +// complete 2048-channel token row. Each workitem retains eight adjacent +// residual results across the RMS reduction, updates hidden state, and +// quantizes the learned-weight normalized values directly into the physical +// row consumed by quantized QKV. This is the interlayer producer for the +// quantized attention schedule; no F32 normalized row is materialized. +// +// Four neighboring workitems own one logical 32-element Q8_1 block. Each +// workitem contributes two packed four-value words. LDS preserves the same +// eight-word max and sum reduction order as the standalone GGML packer so the +// fused and decomposed paths produce identical metadata and payload bytes. +amdgpu.target @qwen3_moe_routed_down_next_norm_q8_gfx11_wave32 {subgroup_size = 32} + +config.decl @qwen3_moe.model.hidden_size : %value: index where [range(%value, 128, 32768), mul(%value, 128)] + +config.decl @qwen3_moe.model.rms_epsilon : f32 + +config.decl @qwen3_moe.routed_down.route_count : %value: index where [range(%value, 1, 8)] + +config.decl @qwen3_moe.routed_down.output_size : %value: index where [range(%value, 1, 4096)] + +config.decl @qwen3_moe.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +kernel.decl @qwen3_moe_routed_down_weighted_reduce_f16_f32(%token_count$5: index) launch(%token_count$6: index, %route_weights: buffer, %routed_output: buffer, %residual_input: buffer, %output: buffer) + +kernel.decl @qwen3_moe_rmsnorm_f32(%token_count$10: index) launch(%token_count$11: index, %input: buffer, %weight: buffer, %output: buffer) + +kernel.decl @ggml_quantize_q8_1_x4_f32(%token_count$15: index, %input_size$16: index) launch(%token_count$17: index, %input_size$18: index, %input: buffer, %output: buffer) + +kernel.def target(@qwen3_moe_routed_down_next_norm_q8_gfx11_wave32) @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_q8_1_x4(%token_count: index) { + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%token_capacity, %c1, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %route_weights: buffer, %routed_output: buffer, %hidden_state: buffer, %next_norm_weight: buffer, %next_projection_input: buffer) { + %route_count0 = config.get @qwen3_moe.routed_down.route_count : index + %output_size0 = config.get @qwen3_moe.routed_down.output_size : index + %hidden_size0 = config.get @qwen3_moe.model.hidden_size : index + %epsilon = config.get @qwen3_moe.model.rms_epsilon : f32 + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048), le(%token_count, %token_capacity)] : index + %route_count = index.assume %route_count0 [range(%route_count0, 8, 8)] : index + %output_size, %hidden_size = index.assume %output_size0, %hidden_size0 [range(%output_size0, 2048, 2048), range(%hidden_size0, 2048, 2048), eq(%output_size0, %hidden_size0)] : index, index + %token0 = kernel.workgroup.id : index + %workitem0 = kernel.workitem.id : index + %workitem = index.assume %workitem0 [range(%workitem0, 0, 255)] : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c128 = index.constant 128 : index + %group_bytes = index.constant 144 : offset + %payload_byte_add = index.constant 16 : offset + %scratch_d_byte_add = index.constant 2048 : offset + %scratch_bytes = index.constant 2304 : offset + %c0_offset = index.constant 0 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %c1_f32 = scalar.constant 1.0 : f32 + %c127 = scalar.constant 127.0 : f32 + %c0_f32x8 = vector.constant 0.0 : vector<8xf32> + %active_token = index.cmp ult, %token0, %bounded_token_count : index + %safe_token0 = scf.select %active_token, %token0, %c0 : index + %token, %launch_token_count = index.assume %safe_token0, %bounded_token_count [lt(%safe_token0, %bounded_token_count)] : index, index + %hidden_size_i32 = index.cast %hidden_size : index to i32 + %hidden_size_f32 = scalar.sitofp %hidden_size_i32 : i32 to f32 + %assignment_count = index.mul %launch_token_count, %route_count : index + %physical_group_count = index.div %hidden_size, %c128 : index + %row_bytes = index.scale %physical_group_count, %group_bytes : index, offset -> offset + %token_output_byte_base = index.scale %token, %row_bytes : index, offset -> offset + %route_weights_noalias, %routed_output_noalias, %hidden_state_noalias, %next_norm_weight_noalias, %next_projection_input_noalias = buffer.assume.noalias %route_weights, %routed_output, %hidden_state, %next_norm_weight, %next_projection_input : buffer, buffer, buffer, buffer, buffer + %route_weights_view = buffer.view %route_weights_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%route_count]xf32> + %routed_output_view = buffer.view %routed_output_noalias[%c0_offset] : buffer -> view<[%assignment_count]x[%output_size]xf16> + %hidden_state_view = buffer.view %hidden_state_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%hidden_size]xf32> + %next_norm_weight_view = buffer.view %next_norm_weight_noalias[%c0_offset] : buffer -> view<[%hidden_size]xf32> + %scratch = buffer.alloca align(16) %scratch_bytes : buffer + %scratch_values = buffer.view %scratch[%c0_offset] : buffer -> view<512xf32> + %scratch_d = buffer.view %scratch[%scratch_d_byte_add] : buffer -> view<64xf32> + // The token predicate is workgroup-uniform, keeping every RMSNorm and Q8_1 + // barrier either active for the whole workgroup or skipped by all workitems. + scf.if %active_token { + %channel = index.mul %workitem, %c8 : index + %weighted_sum = scf.for %route = [%c0 to %route_count step %c1](%sum = %c0_f32x8 : vector<8xf32>) -> (vector<8xf32>) unroll { + %assignment = index.madd %token, %route_count, %route : index + %routed_f16 = vector.load %routed_output_view[%assignment, %channel] : view<[%assignment_count]x[%output_size]xf16> -> vector<8xf16> + %routed = vector.extf %routed_f16 : vector<8xf16> to vector<8xf32> + %route_weight = view.load %route_weights_view[%token, %route] : view<[%launch_token_count]x[%route_count]xf32> -> f32 + %route_weight_x8 = vector.splat %route_weight : vector<8xf32> + %weighted = vector.mulf %routed, %route_weight_x8 : vector<8xf32> + %next = vector.addf %sum, %weighted : vector<8xf32> + scf.yield %next : vector<8xf32> + } + %residual = vector.load %hidden_state_view[%token, %channel] : view<[%launch_token_count]x[%hidden_size]xf32> -> vector<8xf32> + %result = vector.addf %residual, %weighted_sum : vector<8xf32> + %squares = vector.mulf %result, %result : vector<8xf32> + %thread_sum = vector.reduce %squares, %c0_f32 : vector<8xf32>, f32 + %subgroup_sum = kernel.subgroup.reduce %thread_sum : f32 + + %is_subgroup_leader = index.cmp eq, %lane, %c0 : index + scf.if %is_subgroup_leader { + view.store %subgroup_sum, %scratch_values[%subgroup] : f32, view<512xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %is_reduction_subgroup = index.cmp eq, %subgroup, %c0 : index + %is_reduction_lane = index.cmp ult, %lane, %c8 : index + %loads_subgroup_sum = scalar.andi %is_reduction_subgroup, %is_reduction_lane : i1 + %subgroup_partial = scf.if %loads_subgroup_sum -> (f32) { + %value = view.load %scratch_values[%lane] : view<512xf32> -> f32 + scf.yield %value : f32 + } else { + scf.yield %c0_f32 : f32 + } + %row_sum = kernel.subgroup.reduce %subgroup_partial : f32 + %writes_scale = scalar.andi %is_reduction_subgroup, %is_subgroup_leader : i1 + scf.if %writes_scale { + %mean = scalar.divf %row_sum, %hidden_size_f32 : f32 + %biased_mean = scalar.addf %mean, %epsilon : f32 + %scale = scalar.rsqrtf %biased_mean : f32 + view.store %scale, %scratch_values[%c0] : f32, view<512xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + + %row_scale = view.load %scratch_values[%c0] : view<512xf32> -> f32 + // All waves must capture the row scale before scratch is reused for the + // independent Q8 block reductions. + kernel.barrier scope(workgroup) ordering(acq_rel) + %row_scale_vector = vector.splat %row_scale : vector<8xf32> + %learned_weight = vector.load %next_norm_weight_view[%channel] : view<[%hidden_size]xf32> -> vector<8xf32> + %normalized0 = vector.mulf %result, %row_scale_vector : vector<8xf32> + %normalized = vector.mulf %normalized0, %learned_weight : vector<8xf32> + vector.store %result, %hidden_state_view[%token, %channel] : vector<8xf32>, view<[%launch_token_count]x[%hidden_size]xf32> + + %low_values = vector.slice %normalized[0] : vector<8xf32> -> vector<4xf32> + %high_values = vector.slice %normalized[4] : vector<8xf32> -> vector<4xf32> + %low_absolute_values = vector.absf %low_values : vector<4xf32> + %high_absolute_values = vector.absf %high_values : vector<4xf32> + %low_max = vector.reduce %low_absolute_values, %c0_f32 : vector<4xf32>, f32 + %high_max = vector.reduce %high_absolute_values, %c0_f32 : vector<4xf32>, f32 + %partial_base = index.mul %workitem, %c2 : index + %partial_high = index.add %partial_base, %c1 : index + view.store %low_max, %scratch_values[%partial_base] : f32, view<512xf32> + view.store %high_max, %scratch_values[%partial_high] : f32, view<512xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + + %block_in_row = index.div %workitem, %c4 : index + %workitem_in_block = index.rem %workitem, %c4 : index + %is_block_leader = index.cmp eq, %workitem_in_block, %c0 : index + scf.if %is_block_leader { + %block_partial_base = index.mul %block_in_row, %c8 : index + %block_maxima = vector.load %scratch_values[%block_partial_base] : view<512xf32> -> vector<8xf32> + %amax = vector.reduce %block_maxima, %c0_f32 : vector<8xf32>, f32 + %d = scalar.divf %amax, %c127 : f32 + view.store %d, %scratch_d[%block_in_row] : f32, view<64xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + + %d = view.load %scratch_d[%block_in_row] : view<64xf32> -> f32 + %d_nonzero = scalar.cmpf one, %d, %c0_f32 : f32 + %d_inverse = scf.if %d_nonzero -> (f32) { + %inverse = scalar.divf %c1_f32, %d : f32 + scf.yield %inverse : f32 + } else { + scf.yield %c0_f32 : f32 + } + %d_inverse_vector = vector.splat %d_inverse : vector<8xf32> + %scaled_values = vector.mulf %normalized, %d_inverse_vector : vector<8xf32> + %rounded_values = vector.roundf %scaled_values : vector<8xf32> + %quantized_values = vector.fptosi %rounded_values : vector<8xf32> to vector<8xi8> + %packed_words = vector.bitcast %quantized_values : vector<8xi8> to vector<2xi32> + %physical_group = index.div %block_in_row, %c4 : index + %block_in_group = index.rem %block_in_row, %c4 : index + %group_byte_add = index.scale %physical_group, %group_bytes : index, offset -> offset + %group_byte_offset = index.add %token_output_byte_base, %group_byte_add : offset + %payload_byte_offset = index.add %group_byte_offset, %payload_byte_add : offset + %group_ds = buffer.view %next_projection_input_noalias[%group_byte_offset] : buffer -> view<8xf16> + %group_qs = buffer.view %next_projection_input_noalias[%payload_byte_offset] : buffer -> view<32xi32> + %block_word_base = index.mul %block_in_group, %c8 : index + %workitem_word_add = index.mul %workitem_in_block, %c2 : index + %packed_word_index0 = index.add %block_word_base, %workitem_word_add : index + %packed_word_index = index.assume %packed_word_index0 [range(%packed_word_index0, 0, 31)] : index + vector.store %packed_words, %group_qs[%packed_word_index] : vector<2xi32>, view<32xi32> + + %low_rounded_values = vector.slice %rounded_values[0] : vector<8xf32> -> vector<4xf32> + %high_rounded_values = vector.slice %rounded_values[4] : vector<8xf32> -> vector<4xf32> + %low_quantized_sum = vector.reduce %low_rounded_values, %c0_f32 : vector<4xf32>, f32 + %high_quantized_sum = vector.reduce %high_rounded_values, %c0_f32 : vector<4xf32>, f32 + view.store %low_quantized_sum, %scratch_values[%partial_base] : f32, view<512xf32> + view.store %high_quantized_sum, %scratch_values[%partial_high] : f32, view<512xf32> + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.if %is_block_leader { + %block_partial_base = index.mul %block_in_row, %c8 : index + %block_sums = vector.load %scratch_values[%block_partial_base] : view<512xf32> -> vector<8xf32> + %quantized_sum = vector.reduce %block_sums, %c0_f32 : vector<8xf32>, f32 + %s = scalar.mulf %quantized_sum, %d : f32 + %d_f16 = scalar.fptrunc %d : f32 to f16 + %s_f16 = scalar.fptrunc %s : f32 to f16 + %ds_index = index.mul %block_in_group, %c2 : index + %s_index = index.add %ds_index, %c1 : index + view.store %d_f16, %group_ds[%ds_index] : f16, view<8xf16> + view.store %s_f16, %group_ds[%s_index] : f16, view<8xf16> + } + } + kernel.return +} + +// Fourteen independent token rows lock the complete interlayer contract. The +// fused producer must match ordinary weighted reduction for hidden state and +// RMSNorm followed by the generic GGML packer for every Q8 metadata and +// payload byte. +check.case public @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_q8_1_x4_differential_case { + %token_count = check.literal value(14) : index + %hidden_size = check.literal value(2048) : index + %route_weights = check.generate.iota offset(0.03125) step(0.015625) period(8) : tensor<14x8xf32> + %routed_seed = check.param.seed base(5858425849430430008) count(1) : i64 + %routed_output = check.generate.random.uniform seed(%routed_seed) range(-0.5 to 0.5) : tensor<14x8x2048xf16> + %expected_hidden_state = check.generate.iota offset(-0.5) step(0.0078125) period(31) : tensor<14x2048xf32> + %actual_hidden_state = check.generate.iota offset(-0.5) step(0.0078125) period(31) : tensor<14x2048xf32> + %next_norm_weight = check.generate.iota offset(0.5) step(0.0078125) period(29) : tensor<2048xf32> + %normalized = check.generate.fill value(0.0) : tensor<14x2048xf32> + %expected_q8 = check.generate.fill value(0) : tensor<14x2304xi8> + %actual_q8 = check.generate.fill value(1) : tensor<14x2304xi8> + kernel.launch @qwen3_moe_routed_down_weighted_reduce_f16_f32[%token_count](%token_count, %route_weights, %routed_output, %expected_hidden_state, %expected_hidden_state) : [index](index, tensor<14x8xf32>, tensor<14x8x2048xf16>, tensor<14x2048xf32>, tensor<14x2048xf32>) + kernel.launch @qwen3_moe_rmsnorm_f32[%token_count](%token_count, %expected_hidden_state, %next_norm_weight, %normalized) : [index](index, tensor<14x2048xf32>, tensor<2048xf32>, tensor<14x2048xf32>) + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %hidden_size](%token_count, %hidden_size, %normalized, %expected_q8) : [index, index](index, index, tensor<14x2048xf32>, tensor<14x2304xi8>) + kernel.launch @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_q8_1_x4[%token_count](%token_count, %route_weights, %routed_output, %actual_hidden_state, %next_norm_weight, %actual_q8) : [index](index, tensor<14x8xf32>, tensor<14x8x2048xf16>, tensor<14x2048xf32>, tensor<2048xf32>, tensor<14x2304xi8>) + check.expect.close actual(%actual_hidden_state) expected(%expected_hidden_state) atol(0.0001) rtol(0.0001) nan(same) : tensor<14x2048xf32> + check.expect.equal actual(%actual_q8) expected(%expected_q8) : tensor<14x2304xi8> + check.return +} + +check.case public @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_q8_1_x4_fused_benchmark_case { + %token_count = check.param.choice values([14, 32, 128, 512]) name("token_count") : index + %route_weights = check.generate.fill value(0.125) : tensor<[%token_count]x8xf32> + %routed_output = check.generate.fill value(0.0) : tensor<[%token_count]x8x2048xf16> + %hidden_state = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %next_norm_weight = check.generate.fill value(1.0) : tensor<2048xf32> + %next_projection_input = check.generate.fill value(0) : tensor<[%token_count]x2304xi8> + kernel.launch @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_q8_1_x4[%token_count](%token_count, %route_weights, %routed_output, %hidden_state, %next_norm_weight, %next_projection_input) : [index](index, tensor<[%token_count]x8xf32>, tensor<[%token_count]x8x2048xf16>, tensor<[%token_count]x2048xf32>, tensor<2048xf32>, tensor<[%token_count]x2304xi8>) + check.return +} + +check.case public @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_q8_1_x4_composed_benchmark_case { + %token_count = check.param.choice values([14, 32, 128, 512]) name("token_count") : index + %hidden_size = check.literal value(2048) : index + %route_weights = check.generate.fill value(0.125) : tensor<[%token_count]x8xf32> + %routed_output = check.generate.fill value(0.0) : tensor<[%token_count]x8x2048xf16> + %hidden_state = check.generate.fill value(1.0) : tensor<[%token_count]x2048xf32> + %next_norm_weight = check.generate.fill value(1.0) : tensor<2048xf32> + %normalized = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %next_projection_input = check.generate.fill value(0) : tensor<[%token_count]x2304xi8> + kernel.launch @qwen3_moe_routed_down_weighted_reduce_f16_f32[%token_count](%token_count, %route_weights, %routed_output, %hidden_state, %hidden_state) : [index](index, tensor<[%token_count]x8xf32>, tensor<[%token_count]x8x2048xf16>, tensor<[%token_count]x2048xf32>, tensor<[%token_count]x2048xf32>) + kernel.launch @qwen3_moe_rmsnorm_f32[%token_count](%token_count, %hidden_state, %next_norm_weight, %normalized) : [index](index, tensor<[%token_count]x2048xf32>, tensor<2048xf32>, tensor<[%token_count]x2048xf32>) + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %hidden_size](%token_count, %hidden_size, %normalized, %next_projection_input) : [index, index](index, index, tensor<[%token_count]x2048xf32>, tensor<[%token_count]x2304xi8>) + check.return +} + +check.benchmark<@qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_q8_1_x4_fused_benchmark_case> @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_q8_1_x4_fused_prefill_14 {token_count = 14} + +check.benchmark<@qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_q8_1_x4_fused_benchmark_case> @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_q8_1_x4_fused_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_q8_1_x4_composed_benchmark_case> @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_q8_1_x4_composed_prefill_14 {token_count = 14} + +check.benchmark<@qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_q8_1_x4_composed_benchmark_case> @qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_q8_1_x4_composed_prefill_512 {token_count = 512} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_gate_up_swiglu_q4k.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_gate_up_swiglu_q4k.loom new file mode 100644 index 000000000000..d466b3e6d0fb --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_gate_up_swiglu_q4k.loom @@ -0,0 +1,1665 @@ +// Fuses the Qwen3 MoE gate and up expert projections with the SwiGLU +// epilogue. Weights are consumed directly from GGUF's raw Q4_K layout: +// [expert][output channel][K / 256][144 bytes]. Activations use GGML's Q8_1 +// x4 layout produced by the shared quantization kernel. +// +// Route IDs have a logical route_count and an independent physical +// route_stride. This admits llama.cpp's [token][128] argsort storage while +// operating on only the selected top-8 entries, and also admits compact route +// buffers produced by a future fused router. +// +// Two schedules share this contract. Small token counts use one wave per +// [token, route, output channel]. Larger token counts first build a compact +// assignment table per expert, then group up to 32 selected rows with 32 +// output channels so raw weights and quantized activations can be reused +// through workgroup memory. +template.decl @qwen3_moe.routed_gate_up.q4k_q8.body(%publish_output: i1, %token_count: index, %token: index, %route_count: index, %route: index, %route_stride: index, %expert_count: index, %output_size: index, %channel: index, %lane: index, %q8_input: buffer, %route_ids: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) + +template.decl @qwen3_moe.routed_gate_up.quantize_q8_1_x4.subgroup_body(%publish_output: i1, %group_count0: index, %group0: index, %input: buffer, %output: buffer) + +config.decl @qwen3_moe.routed_gate_up.input_size : %value: index where [range(%value, 512, 32768), mul(%value, 512)] + +config.decl @qwen3_moe.routed_gate_up.expert_count : %value: index where [range(%value, 1, 512)] + +config.decl @qwen3_moe.routed_gate_up.route_count : %value: index where [range(%value, 1, 8)] + +config.decl @qwen3_moe.routed_gate_up.output_size : %value: index where [range(%value, 1, 4096)] + +config.decl @qwen3_moe.workload.token_capacity : %value: index where [range(%value, 1, 2048)] + +kernel.decl @ggml_quantize_q8_1_x4_f32(%token_count$25: index, %input_size$26: index) launch(%token_count$27: index, %input_size$28: index, %input: buffer, %output: buffer) + +func.decl @ggml_q8_1_x4_block(%q8_input: buffer, %row_byte_base: offset, %q8_block: index) -> (vector<32xi8>, f32, f32) + +// Decodes one 16-element Q4_K chunk while preserving its unsigned nibbles for +// AMDGPU's mixed u8*s8 dot4 instruction. The returned scale terms are shared +// by every routed activation row consuming the same weight chunk. +func.def inline @qwen3_moe_q4k_chunk_local(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group: index, %q4_half: index) -> (vector<16xi8>, f32, f32) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %block_bytes = index.constant 144 : offset + %header_byte_count = index.constant 4 : index + %code_byte_add = index.constant 16 : offset + %c4_i32 = scalar.constant 4 : i32 + %c6_i32 = scalar.constant 6 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c63_i32 = scalar.constant 63 : i32 + %nibble_mask = vector.constant 252645135 : vector<4xi32> + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + %bounded_half = index.assume %q4_half [range(%q4_half, 0, 1)] : index + %block_byte_add = index.scale %q4_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_byte_add : offset + %header_view = buffer.view %weight[%block_byte_base] : buffer -> view<4xi32> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<32xi32> + %header_words = vector.load %header_view[0] : view<4xi32> -> vector<4xi32> + %header_halves = vector.bitcast %header_words : vector<4xi32> to vector<8xf16> + %d_f16 = vector.extract %header_halves[0] : vector<8xf16> -> f16 + %dmin_f16 = vector.extract %header_halves[1] : vector<8xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %header_bytes = vector.bitcast %header_words : vector<4xi32> to vector<16xi8> + %iqs_group = index.mul %bounded_group, %c8 : index + %iqs_half = index.mul %bounded_half, %c4 : index + %iqs = index.add %iqs_group, %iqs_half : index + %qs_page0 = index.div %iqs, %c16 : index + %qs_page = index.mul %qs_page0, %c8 : index + %qs_lane = index.rem %iqs, %c8 : index + %qs_index0 = index.add %qs_page, %qs_lane : index + %qs_index = index.assume %qs_index0 [range(%qs_index0, 0, 28)] : index + %iqs_mod16 = index.rem %iqs, %c16 : index + %nibble_page = index.div %iqs_mod16, %c8 : index + %nibble_shift_index = index.mul %nibble_page, %c4 : index + %nibble_shift_i32 = index.cast %nibble_shift_index : index to i32 + %nibble_shift = vector.splat %nibble_shift_i32 : vector<4xi32> + %packed_codes = vector.load %code_view[%qs_index] : view<32xi32> -> vector<4xi32> + %shifted_codes = vector.shrui %packed_codes, %nibble_shift : vector<4xi32> + %masked_codes = vector.andi %shifted_codes, %nibble_mask : vector<4xi32> + %q4_values = vector.bitcast %masked_codes : vector<4xi32> to vector<16xi8> + %is_low_group = index.cmp ult, %bounded_group, %c4 : index + %scale, %minimum = scf.if %is_low_group -> (i32, i32) { + %scale_index = index.add %bounded_group, %header_byte_count : index + %minimum_index0 = index.add %bounded_group, %c4 : index + %minimum_index = index.add %minimum_index0, %header_byte_count : index + %scale_i8 = vector.extract %header_bytes[%scale_index] : vector<16xi8> -> i8 + %minimum_i8 = vector.extract %header_bytes[%minimum_index] : vector<16xi8> -> i8 + %scale_u8 = scalar.extui %scale_i8 : i8 to i32 + %minimum_u8 = scalar.extui %minimum_i8 : i8 to i32 + %scale_low6 = scalar.andi %scale_u8, %c63_i32 : i32 + %minimum_low6 = scalar.andi %minimum_u8, %c63_i32 : i32 + scf.yield %scale_low6, %minimum_low6 : i32, i32 + } else { + %packed_index0 = index.add %bounded_group, %c4 : index + %packed_index = index.add %packed_index0, %header_byte_count : index + %scale_high_index = index.sub %bounded_group, %c4 : index + %scale_high_header_index = index.add %scale_high_index, %header_byte_count : index + %minimum_high_header_index = index.add %bounded_group, %header_byte_count : index + %packed_i8 = vector.extract %header_bytes[%packed_index] : vector<16xi8> -> i8 + %scale_high_i8 = vector.extract %header_bytes[%scale_high_header_index] : vector<16xi8> -> i8 + %minimum_high_i8 = vector.extract %header_bytes[%minimum_high_header_index] : vector<16xi8> -> i8 + %packed = scalar.extui %packed_i8 : i8 to i32 + %scale_high = scalar.extui %scale_high_i8 : i8 to i32 + %minimum_high = scalar.extui %minimum_high_i8 : i8 to i32 + %scale_low4 = scalar.andi %packed, %c15_i32 : i32 + %minimum_low4 = scalar.shrui %packed, %c4_i32 : i32 + %scale_high2_raw = scalar.shrui %scale_high, %c6_i32 : i32 + %minimum_high2_raw = scalar.shrui %minimum_high, %c6_i32 : i32 + %scale_high2 = scalar.shli %scale_high2_raw, %c4_i32 : i32 + %minimum_high2 = scalar.shli %minimum_high2_raw, %c4_i32 : i32 + %scale_value = scalar.ori %scale_low4, %scale_high2 : i32 + %minimum_value = scalar.ori %minimum_low4, %minimum_high2 : i32 + scf.yield %scale_value, %minimum_value : i32, i32 + } + %scale_f32 = scalar.uitofp %scale : i32 to f32 + %minimum_f32 = scalar.uitofp %minimum : i32 to f32 + %d_scale = scalar.mulf %d, %scale_f32 : f32 + %dmin_scale = scalar.mulf %dmin, %minimum_f32 : f32 + func.return %q4_values, %d_scale, %dmin_scale : vector<16xi8>, f32, f32 +} + +// Global-weight variant of the Q4_K chunk decoder. Scalar scale-byte loads are +// legal and cheaper for the one-wave schedule, while the LDS schedule uses the +// packed-header variant above because AMDGPU has no local i8 load descriptor. +func.def inline @qwen3_moe_q4k_chunk_global(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group: index, %q4_half: index) -> (vector<16xi8>, f32, f32) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %block_bytes = index.constant 144 : offset + %scale_byte_add = index.constant 4 : offset + %code_byte_add = index.constant 16 : offset + %c4_i32 = scalar.constant 4 : i32 + %c6_i32 = scalar.constant 6 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c63_i32 = scalar.constant 63 : i32 + %nibble_mask = vector.constant 252645135 : vector<4xi32> + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + %bounded_half = index.assume %q4_half [range(%q4_half, 0, 1)] : index + %block_byte_add = index.scale %q4_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %scale_byte_base = index.add %block_byte_base, %scale_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_byte_add : offset + %half_view = buffer.view %weight[%block_byte_base] : buffer -> view<2xf16> + %scale_view = buffer.view %weight[%scale_byte_base] : buffer -> view<12xi8> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<32xi32> + %d_f16 = view.load %half_view[0] : view<2xf16> -> f16 + %dmin_f16 = view.load %half_view[1] : view<2xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %iqs_group = index.mul %bounded_group, %c8 : index + %iqs_half = index.mul %bounded_half, %c4 : index + %iqs = index.add %iqs_group, %iqs_half : index + %qs_page0 = index.div %iqs, %c16 : index + %qs_page = index.mul %qs_page0, %c8 : index + %qs_lane = index.rem %iqs, %c8 : index + %qs_index0 = index.add %qs_page, %qs_lane : index + %qs_index = index.assume %qs_index0 [range(%qs_index0, 0, 28)] : index + %iqs_mod16 = index.rem %iqs, %c16 : index + %nibble_page = index.div %iqs_mod16, %c8 : index + %nibble_shift_index = index.mul %nibble_page, %c4 : index + %nibble_shift_i32 = index.cast %nibble_shift_index : index to i32 + %nibble_shift = vector.splat %nibble_shift_i32 : vector<4xi32> + %packed_codes = vector.load %code_view[%qs_index] : view<32xi32> -> vector<4xi32> + %shifted_codes = vector.shrui %packed_codes, %nibble_shift : vector<4xi32> + %masked_codes = vector.andi %shifted_codes, %nibble_mask : vector<4xi32> + %q4_values = vector.bitcast %masked_codes : vector<4xi32> to vector<16xi8> + %is_low_group = index.cmp ult, %bounded_group, %c4 : index + %scale, %minimum = scf.if %is_low_group -> (i32, i32) { + %minimum_index = index.add %bounded_group, %c4 : index + %scale_i8 = view.load %scale_view[%bounded_group] : view<12xi8> -> i8 + %minimum_i8 = view.load %scale_view[%minimum_index] : view<12xi8> -> i8 + %scale_u8 = scalar.extui %scale_i8 : i8 to i32 + %minimum_u8 = scalar.extui %minimum_i8 : i8 to i32 + %scale_low6 = scalar.andi %scale_u8, %c63_i32 : i32 + %minimum_low6 = scalar.andi %minimum_u8, %c63_i32 : i32 + scf.yield %scale_low6, %minimum_low6 : i32, i32 + } else { + %packed_index = index.add %bounded_group, %c4 : index + %scale_high_index = index.sub %bounded_group, %c4 : index + %packed_i8 = view.load %scale_view[%packed_index] : view<12xi8> -> i8 + %scale_high_i8 = view.load %scale_view[%scale_high_index] : view<12xi8> -> i8 + %minimum_high_i8 = view.load %scale_view[%bounded_group] : view<12xi8> -> i8 + %packed = scalar.extui %packed_i8 : i8 to i32 + %scale_high = scalar.extui %scale_high_i8 : i8 to i32 + %minimum_high = scalar.extui %minimum_high_i8 : i8 to i32 + %scale_low4 = scalar.andi %packed, %c15_i32 : i32 + %minimum_low4 = scalar.shrui %packed, %c4_i32 : i32 + %scale_high2_raw = scalar.shrui %scale_high, %c6_i32 : i32 + %minimum_high2_raw = scalar.shrui %minimum_high, %c6_i32 : i32 + %scale_high2 = scalar.shli %scale_high2_raw, %c4_i32 : i32 + %minimum_high2 = scalar.shli %minimum_high2_raw, %c4_i32 : i32 + %scale_value = scalar.ori %scale_low4, %scale_high2 : i32 + %minimum_value = scalar.ori %minimum_low4, %minimum_high2 : i32 + scf.yield %scale_value, %minimum_value : i32, i32 + } + %scale_f32 = scalar.uitofp %scale : i32 to f32 + %minimum_f32 = scalar.uitofp %minimum : i32 to f32 + %d_scale = scalar.mulf %d, %scale_f32 : f32 + %dmin_scale = scalar.mulf %dmin, %minimum_f32 : f32 + func.return %q4_values, %d_scale, %dmin_scale : vector<16xi8>, f32, f32 +} + +// Decodes one Q4_K scale/minimum pair from a header loaded as three packed +// scale words. Keeping the complete header in registers lets paired-group +// contractions share both the header and packed-code load. +func.def inline @qwen3_moe_q4k_scale_from_header(%scale0: i32, %scale1: i32, %scale2: i32, %q4_group: index) -> (i32, i32) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c48_i32 = scalar.constant 48 : i32 + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + %is_low_group = index.cmp ult, %bounded_group, %c4 : index + %scale_lane = index.rem %bounded_group, %c4 : index + %scale_shift_index = index.mul %scale_lane, %c8 : index + %scale_shift = index.cast %scale_shift_index : index to i32 + %high_shift = scalar.addi %scale_shift, %c2_i32 : i32 + %minimum_shift = scalar.addi %scale_shift, %c4_i32 : i32 + %selected_scale_source = scf.select %is_low_group, %scale0, %scale2 : i32 + %selected_minimum_source = scf.select %is_low_group, %scale1, %scale2 : i32 + %selected_scale_high_shift = scf.select %is_low_group, %scale_shift, %high_shift : i32 + %selected_minimum_low_shift = scf.select %is_low_group, %scale_shift, %minimum_shift : i32 + %scale_low0 = scalar.shrui %selected_scale_source, %scale_shift : i32 + %scale_low = scalar.andi %scale_low0, %c15_i32 : i32 + %scale_high0 = scalar.shrui %scale0, %selected_scale_high_shift : i32 + %scale_high = scalar.andi %scale_high0, %c48_i32 : i32 + %scale = scalar.ori %scale_low, %scale_high : i32 + %minimum_low0 = scalar.shrui %selected_minimum_source, %selected_minimum_low_shift : i32 + %minimum_low = scalar.andi %minimum_low0, %c15_i32 : i32 + %minimum_high0 = scalar.shrui %scale1, %selected_scale_high_shift : i32 + %minimum_high = scalar.andi %minimum_high0, %c48_i32 : i32 + %minimum = scalar.ori %minimum_low, %minimum_high : i32 + func.return %scale, %minimum : i32, i32 +} + +// Decodes two adjacent Q4_K groups from the low and high nibbles of one +// packed-code load. The caller supplies the already-loaded 16-byte header. +func.def inline @qwen3_moe_q4k_chunk_pair_global(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group_pair: index, %q4_half: index, %header_words: vector<4xi32>) -> (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) { + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %block_bytes = index.constant 144 : offset + %code_byte_add = index.constant 16 : offset + %c4_i32 = scalar.constant 4 : i32 + %nibble_mask = vector.constant 252645135 : vector<4xi32> + %bounded_pair = index.assume %q4_group_pair [range(%q4_group_pair, 0, 3)] : index + %bounded_half = index.assume %q4_half [range(%q4_half, 0, 1)] : index + %block_byte_add = index.scale %q4_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_byte_add : offset + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<32xi32> + %header_halves = vector.bitcast %header_words : vector<4xi32> to vector<8xf16> + %d_f16 = vector.extract %header_halves[0] : vector<8xf16> -> f16 + %dmin_f16 = vector.extract %header_halves[1] : vector<8xf16> -> f16 + %scale0 = vector.extract %header_words[1] : vector<4xi32> -> i32 + %scale1 = vector.extract %header_words[2] : vector<4xi32> -> i32 + %scale2 = vector.extract %header_words[3] : vector<4xi32> -> i32 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %pair_code_base = index.mul %bounded_pair, %c8 : index + %half_code_add = index.mul %bounded_half, %c4 : index + %code_index0 = index.add %pair_code_base, %half_code_add : index + %code_index = index.assume %code_index0 [range(%code_index0, 0, 28)] : index + %packed_codes = vector.load %code_view[%code_index] : view<32xi32> -> vector<4xi32> + %low_codes = vector.andi %packed_codes, %nibble_mask : vector<4xi32> + %c4_i32v = vector.splat %c4_i32 : vector<4xi32> + %high_shifted = vector.shrui %packed_codes, %c4_i32v : vector<4xi32> + %high_codes = vector.andi %high_shifted, %nibble_mask : vector<4xi32> + %q4_low = vector.bitcast %low_codes : vector<4xi32> to vector<16xi8> + %q4_high = vector.bitcast %high_codes : vector<4xi32> to vector<16xi8> + %low_group = index.mul %bounded_pair, %c2 : index + %high_group = index.add %low_group, %c1 : index + %low_scale, %low_minimum = func.call @qwen3_moe_q4k_scale_from_header(%scale0, %scale1, %scale2, %low_group) : (i32, i32, i32, index) -> (i32, i32) + %high_scale, %high_minimum = func.call @qwen3_moe_q4k_scale_from_header(%scale0, %scale1, %scale2, %high_group) : (i32, i32, i32, index) -> (i32, i32) + %low_scale_f32 = scalar.uitofp %low_scale : i32 to f32 + %low_minimum_f32 = scalar.uitofp %low_minimum : i32 to f32 + %high_scale_f32 = scalar.uitofp %high_scale : i32 to f32 + %high_minimum_f32 = scalar.uitofp %high_minimum : i32 to f32 + %low_d_scale = scalar.mulf %d, %low_scale_f32 : f32 + %low_dmin_scale = scalar.mulf %dmin, %low_minimum_f32 : f32 + %high_d_scale = scalar.mulf %d, %high_scale_f32 : f32 + %high_dmin_scale = scalar.mulf %dmin, %high_minimum_f32 : f32 + func.return %q4_low, %low_d_scale, %low_dmin_scale, %q4_high, %high_d_scale, %high_dmin_scale : vector<16xi8>, f32, f32, vector<16xi8>, f32, f32 +} + +// Contracts one decoded Q4_K half-group with a Q8_1 half-block. Both +// half-groups apply half of the Q8_1 block-sum correction. +func.def inline @qwen3_moe_q4k_q8_1_dot(%q4_values: vector<16xi8>, %d_scale: f32, %dmin_scale: f32, %q8_values: vector<16xi8>, %q8_d: f32, %q8_s: f32) -> (f32) { + %c0_i32 = scalar.constant 0 : i32 + %c0_i32v = vector.constant 0 : vector<4xi32> + %half_f32 = scalar.constant 0.5 : f32 + %partial_dots = vector.dot4i %q4_values, %q8_values, %c0_i32v : vector<16xi8>, vector<16xi8>, vector<4xi32> + %q_sum = vector.reduce %partial_dots, %c0_i32 : vector<4xi32>, i32 + %q_sum_f32 = scalar.sitofp %q_sum : i32 to f32 + %scaled_dot0 = scalar.mulf %q8_d, %d_scale : f32 + %scaled_dot = scalar.mulf %scaled_dot0, %q_sum_f32 : f32 + %q8_half_sum = scalar.mulf %q8_s, %half_f32 : f32 + %minimum_correction = scalar.mulf %dmin_scale, %q8_half_sum : f32 + %contribution = scalar.subf %scaled_dot, %minimum_correction : f32 + func.return %contribution : f32 +} + +// Contracts one decoded Q4_K half-group against four routed Q8_1 rows. The +// leading row axis stays explicit through dot4 and scale correction so the +// grouped schedule can reuse each weight decode without scalarizing the rows. +func.def inline @qwen3_moe_q4k_q8_1_dot4_rows(%q4_values: vector<16xi8>, %d_scale: f32, %dmin_scale: f32, %q8_values: vector<4x16xi8>, %q8_d: vector<4xf32>, %q8_s: vector<4xf32>) -> (vector<4xf32>) { + %q8_values0 = vector.extract %q8_values[0] : vector<4x16xi8> -> vector<16xi8> + %q8_values1 = vector.extract %q8_values[1] : vector<4x16xi8> -> vector<16xi8> + %q8_values2 = vector.extract %q8_values[2] : vector<4x16xi8> -> vector<16xi8> + %q8_values3 = vector.extract %q8_values[3] : vector<4x16xi8> -> vector<16xi8> + %q8_d0 = vector.extract %q8_d[0] : vector<4xf32> -> f32 + %q8_d1 = vector.extract %q8_d[1] : vector<4xf32> -> f32 + %q8_d2 = vector.extract %q8_d[2] : vector<4xf32> -> f32 + %q8_d3 = vector.extract %q8_d[3] : vector<4xf32> -> f32 + %q8_s0 = vector.extract %q8_s[0] : vector<4xf32> -> f32 + %q8_s1 = vector.extract %q8_s[1] : vector<4xf32> -> f32 + %q8_s2 = vector.extract %q8_s[2] : vector<4xf32> -> f32 + %q8_s3 = vector.extract %q8_s[3] : vector<4xf32> -> f32 + %contribution0 = func.call @qwen3_moe_q4k_q8_1_dot(%q4_values, %d_scale, %dmin_scale, %q8_values0, %q8_d0, %q8_s0) : (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) -> (f32) + %contribution1 = func.call @qwen3_moe_q4k_q8_1_dot(%q4_values, %d_scale, %dmin_scale, %q8_values1, %q8_d1, %q8_s1) : (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) -> (f32) + %contribution2 = func.call @qwen3_moe_q4k_q8_1_dot(%q4_values, %d_scale, %dmin_scale, %q8_values2, %q8_d2, %q8_s2) : (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) -> (f32) + %contribution3 = func.call @qwen3_moe_q4k_q8_1_dot(%q4_values, %d_scale, %dmin_scale, %q8_values3, %q8_d3, %q8_s3) : (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) -> (f32) + %contribution = vector.from_elements %contribution0, %contribution1, %contribution2, %contribution3 : vector<4xf32> + func.return %contribution : vector<4xf32> +} + +// Convenience wrapper used by the one-wave schedule. +func.def inline @qwen3_moe_q4k_q8_1_chunk(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group: index, %q4_half: index, %q8_values: vector<16xi8>, %q8_d: f32, %q8_s: f32) -> (f32) { + %q4_values, %d_scale, %dmin_scale = func.call @qwen3_moe_q4k_chunk_global(%weight, %row_byte_base, %q4_block, %q4_group, %q4_half) : (buffer, offset, index, index, index) -> (vector<16xi8>, f32, f32) + %contribution = func.call @qwen3_moe_q4k_q8_1_dot(%q4_values, %d_scale, %dmin_scale, %q8_values, %q8_d, %q8_s) : (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) -> (f32) + func.return %contribution : f32 +} + +// Computes one lane's partial for a complete Q4_K row. Dense and routed +// providers choose output ownership independently, then share this exact +// packed-row contraction before reducing across a wave32 subgroup. +func.def inline @qwen3_moe_q4k_q8_1_x4_row_lane(%input_size: index, %weight: buffer, %weight_row_byte_base: offset, %q8_input: buffer, %q8_row_byte_base: offset, %lane: index) -> (f32) { + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %bounded_lane = index.assume %lane [range(%lane, 0, 31)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c128 = index.constant 128 : index + %c144 = index.constant 144 : index + %c256 = index.constant 256 : index + %c511 = index.constant 511 : index + %c512 = index.constant 512 : index + %q8_group_bytes = index.constant 144 : offset + %q8_payload_byte_add = index.constant 16 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %q4_block_count = index.div %bounded_input_size, %c256 : index + %q8_group_count = index.div %bounded_input_size, %c128 : index + %padded_input_size = index.add %bounded_input_size, %c511 : index + %iteration_count = index.div %padded_input_size, %c512 : index + %q4_group0 = index.div %bounded_lane, %c2 : index + %q4_group1 = index.rem %q4_group0, %c8 : index + %q4_group = index.assume %q4_group1 [range(%q4_group1, 0, 7)] : index + %q4_half0 = index.rem %bounded_lane, %c2 : index + %q4_half = index.assume %q4_half0 [range(%q4_half0, 0, 1)] : index + %lane_q4_block = index.div %bounded_lane, %c16 : index + %lane_q8_group = index.div %bounded_lane, %c8 : index + %q8_inner_block0 = index.div %bounded_lane, %c2 : index + %q8_inner_block1 = index.rem %q8_inner_block0, %c4 : index + %q8_inner_block = index.assume %q8_inner_block1 [range(%q8_inner_block1, 0, 3)] : index + %q8_inner_word_base = index.mul %q8_inner_block, %c8 : index + %q8_half_word_add = index.mul %q4_half, %c4 : index + %q8_word_index0 = index.add %q8_inner_word_base, %q8_half_word_add : index + %q8_word_index = index.assume %q8_word_index0 [range(%q8_word_index0, 0, 28)] : index + %q8_ds_index = index.mul %q8_inner_block, %c2 : index + %q8_s_index = index.add %q8_ds_index, %c1 : index + %sum = scf.for %iteration = [%c0 to %iteration_count step %c1](%iteration_acc = %c0_f32 : f32) -> (f32) unroll { + %iteration_q4_block = index.mul %iteration, %c2 : index + %q4_block0 = index.add %iteration_q4_block, %lane_q4_block : index + %valid_q4_block = index.cmp ult, %q4_block0, %q4_block_count : index + %contribution = scf.if %valid_q4_block -> (f32) { + %q4_block, %bounded_q4_block_count = index.assume %q4_block0, %q4_block_count [lt(%q4_block0, %q4_block_count)] : index, index + %iteration_q8_group = index.mul %iteration, %c4 : index + %q8_group0 = index.add %iteration_q8_group, %lane_q8_group : index + %q8_group, %bounded_q8_group_count = index.assume %q8_group0, %q8_group_count [lt(%q8_group0, %q8_group_count)] : index, index + %q8_group_byte_add = index.scale %q8_group, %q8_group_bytes : index, offset -> offset + %q8_group_byte_base = index.add %q8_row_byte_base, %q8_group_byte_add : offset + %q8_payload_byte_base = index.add %q8_group_byte_base, %q8_payload_byte_add : offset + %q8_ds_view = buffer.view %q8_input[%q8_group_byte_base] : buffer -> view<8xf16> + %q8_words_view = buffer.view %q8_input[%q8_payload_byte_base] : buffer -> view<32xi32> + %q8_d_f16 = view.load %q8_ds_view[%q8_ds_index] : view<8xf16> -> f16 + %q8_s_f16 = view.load %q8_ds_view[%q8_s_index] : view<8xf16> -> f16 + %q8_d = scalar.extf %q8_d_f16 : f16 to f32 + %q8_s = scalar.extf %q8_s_f16 : f16 to f32 + %q8_words = vector.load %q8_words_view[%q8_word_index] : view<32xi32> -> vector<4xi32> + %q8_values = vector.bitcast %q8_words : vector<4xi32> to vector<16xi8> + %dot = func.call @qwen3_moe_q4k_q8_1_chunk(%weight, %weight_row_byte_base, %q4_block, %q4_group, %q4_half, %q8_values, %q8_d, %q8_s) : (buffer, offset, index, index, index, vector<16xi8>, f32, f32) -> (f32) + scf.yield %dot : f32 + } else { + scf.yield %c0_f32 : f32 + } + %next = scalar.addf %iteration_acc, %contribution : f32 + scf.yield %next : f32 + } + func.return %sum : f32 +} + +// Computes one lane's two adjacent-group contributions within one Q4_K block. +// Eight block lanes consume both nibbles from every packed code load and cover +// the block's full 256-element input extent. +func.def inline @qwen3_moe_q4k_q8_1_x4_paired_block_lane(%input_size: index, %weight: buffer, %weight_row_byte_base: offset, %q8_input: buffer, %q8_row_byte_base: offset, %q4_block: index, %block_lane: index) -> (f32) { + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %bounded_q4_block0 = index.assume %q4_block [range(%q4_block, 0, 127)] : index + %bounded_block_lane = index.assume %block_lane [range(%block_lane, 0, 7)] : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %q4_block_bytes = index.constant 144 : offset + %q8_group_bytes = index.constant 144 : offset + %q8_payload_byte_add = index.constant 16 : offset + %q4_block_count = index.div %bounded_input_size, %c256 : index + %q8_group_count = index.div %bounded_input_size, %c128 : index + %bounded_q4_block, %bounded_q4_block_count = index.assume %bounded_q4_block0, %q4_block_count [lt(%bounded_q4_block0, %q4_block_count)] : index, index + %q4_group_pair0 = index.div %bounded_block_lane, %c2 : index + %q4_group_pair = index.assume %q4_group_pair0 [range(%q4_group_pair0, 0, 3)] : index + %q4_half0 = index.rem %bounded_block_lane, %c2 : index + %q4_half = index.assume %q4_half0 [range(%q4_half0, 0, 1)] : index + %q8_group_in_block0 = index.div %q4_group_pair, %c2 : index + %q8_group_in_block = index.assume %q8_group_in_block0 [range(%q8_group_in_block0, 0, 1)] : index + %pair_in_q8_group0 = index.rem %q4_group_pair, %c2 : index + %pair_in_q8_group = index.assume %pair_in_q8_group0 [range(%pair_in_q8_group0, 0, 1)] : index + %q8_low_inner_block0 = index.mul %pair_in_q8_group, %c2 : index + %q8_low_inner_block = index.assume %q8_low_inner_block0 [range(%q8_low_inner_block0, 0, 2)] : index + %q8_high_inner_block0 = index.add %q8_low_inner_block, %c1 : index + %q8_high_inner_block = index.assume %q8_high_inner_block0 [range(%q8_high_inner_block0, 1, 3)] : index + %q8_half_word_add = index.mul %q4_half, %c4 : index + %q8_low_inner_word_base = index.mul %q8_low_inner_block, %c8 : index + %q8_low_word_index0 = index.add %q8_low_inner_word_base, %q8_half_word_add : index + %q8_low_word_index = index.assume %q8_low_word_index0 [range(%q8_low_word_index0, 0, 20)] : index + %q8_high_inner_word_base = index.mul %q8_high_inner_block, %c8 : index + %q8_high_word_index0 = index.add %q8_high_inner_word_base, %q8_half_word_add : index + %q8_high_word_index = index.assume %q8_high_word_index0 [range(%q8_high_word_index0, 8, 28)] : index + %q8_low_ds_index0 = index.mul %q8_low_inner_block, %c2 : index + %q8_low_ds_index = index.assume %q8_low_ds_index0 [range(%q8_low_ds_index0, 0, 4)] : index + %q8_block_group_base = index.mul %bounded_q4_block, %c2 : index + %q8_group0 = index.add %q8_block_group_base, %q8_group_in_block : index + %q8_group, %bounded_q8_group_count = index.assume %q8_group0, %q8_group_count [lt(%q8_group0, %q8_group_count)] : index, index + %q8_group_byte_add = index.scale %q8_group, %q8_group_bytes : index, offset -> offset + %q8_group_byte_base = index.add %q8_row_byte_base, %q8_group_byte_add : offset + %q8_payload_byte_base = index.add %q8_group_byte_base, %q8_payload_byte_add : offset + %q8_ds_view = buffer.view %q8_input[%q8_group_byte_base] : buffer -> view<8xf16> + %q8_words_view = buffer.view %q8_input[%q8_payload_byte_base] : buffer -> view<32xi32> + %q8_ds = vector.load %q8_ds_view[%q8_low_ds_index] : view<8xf16> -> vector<4xf16> + %q8_low_d_f16 = vector.extract %q8_ds[0] : vector<4xf16> -> f16 + %q8_low_s_f16 = vector.extract %q8_ds[1] : vector<4xf16> -> f16 + %q8_high_d_f16 = vector.extract %q8_ds[2] : vector<4xf16> -> f16 + %q8_high_s_f16 = vector.extract %q8_ds[3] : vector<4xf16> -> f16 + %q8_low_d = scalar.extf %q8_low_d_f16 : f16 to f32 + %q8_low_s = scalar.extf %q8_low_s_f16 : f16 to f32 + %q8_high_d = scalar.extf %q8_high_d_f16 : f16 to f32 + %q8_high_s = scalar.extf %q8_high_s_f16 : f16 to f32 + %q8_low_words = vector.load %q8_words_view[%q8_low_word_index] : view<32xi32> -> vector<4xi32> + %q8_high_words = vector.load %q8_words_view[%q8_high_word_index] : view<32xi32> -> vector<4xi32> + %q8_low_values = vector.bitcast %q8_low_words : vector<4xi32> to vector<16xi8> + %q8_high_values = vector.bitcast %q8_high_words : vector<4xi32> to vector<16xi8> + %q4_block_byte_add = index.scale %bounded_q4_block, %q4_block_bytes : index, offset -> offset + %q4_block_byte_base = index.add %weight_row_byte_base, %q4_block_byte_add : offset + %q4_header_view = buffer.view %weight[%q4_block_byte_base] : buffer -> view<4xi32> + %q4_header_words = vector.load %q4_header_view[0] : view<4xi32> -> vector<4xi32> + %q4_low, %low_d_scale, %low_dmin_scale, %q4_high, %high_d_scale, %high_dmin_scale = func.call @qwen3_moe_q4k_chunk_pair_global(%weight, %weight_row_byte_base, %bounded_q4_block, %q4_group_pair, %q4_half, %q4_header_words) : (buffer, offset, index, index, index, vector<4xi32>) -> (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) + %low = func.call @qwen3_moe_q4k_q8_1_dot(%q4_low, %low_d_scale, %low_dmin_scale, %q8_low_values, %q8_low_d, %q8_low_s) : (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) -> (f32) + %high = func.call @qwen3_moe_q4k_q8_1_dot(%q4_high, %high_d_scale, %high_dmin_scale, %q8_high_values, %q8_high_d, %q8_high_s) : (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) -> (f32) + %pair = scalar.addf %low, %high : f32 + func.return %pair : f32 +} + +// Computes one lane's partial while consuming both nibbles of each packed +// Q4_K code load. Eight lanes cover one 256-element Q4_K block, so a wave32 +// advances through four blocks per iteration and reduces the adjacent-group +// contributions together. +func.def inline @qwen3_moe_q4k_q8_1_x4_paired_row_lane(%input_size: index, %weight: buffer, %weight_row_byte_base: offset, %q8_input: buffer, %q8_row_byte_base: offset, %lane: index) -> (f32) { + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %bounded_lane = index.assume %lane [range(%lane, 0, 31)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c256 = index.constant 256 : index + %c1023 = index.constant 1023 : index + %c1024 = index.constant 1024 : index + %c0_f32 = scalar.constant 0.0 : f32 + %q4_block_count = index.div %bounded_input_size, %c256 : index + %padded_input_size = index.add %bounded_input_size, %c1023 : index + %iteration_count = index.div %padded_input_size, %c1024 : index + %lane_q4_block = index.div %bounded_lane, %c8 : index + %block_lane0 = index.rem %bounded_lane, %c8 : index + %block_lane = index.assume %block_lane0 [range(%block_lane0, 0, 7)] : index + %sum = scf.for %iteration = [%c0 to %iteration_count step %c1](%iteration_acc = %c0_f32 : f32) -> (f32) unroll { + %iteration_q4_block = index.mul %iteration, %c4 : index + %q4_block0 = index.add %iteration_q4_block, %lane_q4_block : index + %valid_q4_block = index.cmp ult, %q4_block0, %q4_block_count : index + %contribution = scf.if %valid_q4_block -> (f32) { + %q4_block, %bounded_q4_block_count = index.assume %q4_block0, %q4_block_count [lt(%q4_block0, %q4_block_count)] : index, index + %pair = func.call @qwen3_moe_q4k_q8_1_x4_paired_block_lane(%bounded_input_size, %weight, %weight_row_byte_base, %q8_input, %q8_row_byte_base, %q4_block, %block_lane) : (index, buffer, offset, buffer, offset, index, index) -> (f32) + scf.yield %pair : f32 + } else { + scf.yield %c0_f32 : f32 + } + %next = scalar.addf %iteration_acc, %contribution : f32 + scf.yield %next : f32 + } + func.return %sum : f32 +} + +// Computes one lane's partial for a row owned by an eight-lane cohort. Every +// cohort lane consumes both nibbles for one adjacent Q4_K group pair while the +// cohort walks all blocks in the row. This schedule lets separate cohorts +// contract independent routed rows concurrently. +func.def inline @qwen3_moe_q4k_q8_1_x4_cohort_row_lane(%input_size: index, %weight: buffer, %weight_row_byte_base: offset, %q8_input: buffer, %q8_row_byte_base: offset, %cohort_lane: index) -> (f32) { + %bounded_input_size = index.assume %input_size [range(%input_size, 256, 32768), mul(%input_size, 256)] : index + %bounded_cohort_lane = index.assume %cohort_lane [range(%cohort_lane, 0, 7)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + %c0_f32 = scalar.constant 0.0 : f32 + %q4_block_count = index.div %bounded_input_size, %c256 : index + %sum = scf.for %q4_block0 = [%c0 to %q4_block_count step %c1](%block_acc = %c0_f32 : f32) -> (f32) unroll { + %q4_block, %bounded_q4_block_count = index.assume %q4_block0, %q4_block_count [lt(%q4_block0, %q4_block_count)] : index, index + %pair = func.call @qwen3_moe_q4k_q8_1_x4_paired_block_lane(%bounded_input_size, %weight, %weight_row_byte_base, %q8_input, %q8_row_byte_base, %q4_block, %bounded_cohort_lane) : (index, buffer, offset, buffer, offset, index, index) -> (f32) + %next = scalar.addf %block_acc, %pair : f32 + scf.yield %next : f32 + } + func.return %sum : f32 +} + +// Decodes the packed queue descriptor shared by expert-grouped projections. +// The producer stores a 7-bit expert ordinal, a 6-bit 32-row partition +// ordinal, and a 5-bit row count minus one. +func.def inline @qwen3_moe_unpack_expert_partition_descriptor(%descriptor: i32) -> (index, index, index) { + %c1_i32 = scalar.constant 1 : i32 + %c5_i32 = scalar.constant 5 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c13_i32 = scalar.constant 13 : i32 + %c31_i32 = scalar.constant 31 : i32 + %c63_i32 = scalar.constant 63 : i32 + %c127_i32 = scalar.constant 127 : i32 + %expert_i32 = scalar.andi %descriptor, %c127_i32 : i32 + %partition_shifted_i32 = scalar.shrui %descriptor, %c7_i32 : i32 + %partition_i32 = scalar.andi %partition_shifted_i32, %c63_i32 : i32 + %route_tile_base_i32 = scalar.shli %partition_i32, %c5_i32 : i32 + %row_count_shifted_i32 = scalar.shrui %descriptor, %c13_i32 : i32 + %row_count_minus_one_i32 = scalar.andi %row_count_shifted_i32, %c31_i32 : i32 + %partition_row_count_i32 = scalar.addi %row_count_minus_one_i32, %c1_i32 : i32 + %expert0 = index.cast %expert_i32 : i32 to index + %expert = index.assume %expert0 [range(%expert0, 0, 127)] : index + %route_tile_base0 = index.cast %route_tile_base_i32 : i32 to index + %route_tile_base = index.assume %route_tile_base0 [range(%route_tile_base0, 0, 2016)] : index + %partition_row_count0 = index.cast %partition_row_count_i32 : i32 to index + %partition_row_count = index.assume %partition_row_count0 [range(%partition_row_count0, 1, 32)] : index + func.return %expert, %route_tile_base, %partition_row_count : index, index, index +} + +// Builds the transient expert table consumed by the grouped projection. The +// table packs [expert_count] route counts followed by +// [expert_count][token_count] compact assignment ordinals. Top-k routing +// selects each expert at most once per token, so token_count entries are +// sufficient for every expert even though there are route_count assignments +// per token. +kernel.def @qwen3_moe_build_expert_table(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index) { + %configured_expert_count = config.get @qwen3_moe.routed_gate_up.expert_count : index + %c1 = index.constant 1 : index + %workgroup_size = index.constant 256 : index + kernel.launch.config workgroups(%configured_expert_count, %c1, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index, %route_ids: buffer, %expert_table: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %configured_route_count0 = config.get @qwen3_moe.routed_gate_up.route_count : index + %bounded_route_count, %configured_route_count = index.assume %route_count, %configured_route_count0 [range(%route_count, 1, 8), eq(%route_count, %configured_route_count0)] : index, index + %bounded_route_stride = index.assume %route_stride [range(%route_stride, 1, 512)] : index + %configured_expert_count0 = config.get @qwen3_moe.routed_gate_up.expert_count : index + %bounded_expert_count, %configured_expert_count = index.assume %expert_count, %configured_expert_count0 [range(%expert_count, 1, 512), eq(%expert_count, %configured_expert_count0)] : index, index + %expert0 = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %c0 = index.constant 0 : index + %workgroup_size = index.constant 256 : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c0_offset = index.constant 0 : offset + %c4_bytes = index.constant 4 : offset + %assignment_count = index.mul %bounded_token_count, %bounded_route_count : index + %expert, %table_expert_count = index.assume %expert0, %bounded_expert_count [lt(%expert0, %bounded_expert_count)] : index, index + %assignment_table_byte_base = index.scale %table_expert_count, %c4_bytes : index, offset -> offset + %route_view = buffer.view %route_ids[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_route_stride]xi32> + %count_view = buffer.view %expert_table[%c0_offset] : buffer -> view<[%table_expert_count]xi32> + %assignment_view = buffer.view %expert_table[%assignment_table_byte_base] : buffer -> view<[%table_expert_count]x[%bounded_token_count]xi32> + %expert_route_count = scf.for %block_base = [%c0 to %assignment_count step %workgroup_size](%matched_base = %c0_i32 : i32) -> (i32) { + %assignment = index.add %block_base, %lane : index + %in_range = index.cmp ult, %assignment, %assignment_count : index + %route_expert_i32 = scf.if %in_range -> (i32) { + %token0 = index.div %assignment, %configured_route_count : index + %route0 = index.rem %assignment, %configured_route_count : index + %token, %route_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + // The route stride is the physical row width and must contain every + // logical top-k route. + %route, %route_row_stride = index.assume %route0, %bounded_route_stride [lt(%route0, %bounded_route_stride)] : index, index + %loaded = view.load %route_view[%token, %route] : view<[%bounded_token_count]x[%bounded_route_stride]xi32> -> i32 + scf.yield %loaded : i32 + } else { + %cn1_i32 = scalar.constant -1 : i32 + scf.yield %cn1_i32 : i32 + } + %route_expert0 = index.cast %route_expert_i32 : i32 to index + %route_expert = index.assume %route_expert0 [range(%route_expert0, -1, 511)] : index + %matches = index.cmp eq, %route_expert, %expert : index + %match_i32 = scf.if %matches -> (i32) { + scf.yield %c1_i32 : i32 + } else { + scf.yield %c0_i32 : i32 + } + %block_prefix = kernel.workgroup.scan %match_i32 {direction = forward, mode = exclusive} : i32 + %block_match_count_reduced = kernel.workgroup.reduce %match_i32 : i32 + %block_match_count = kernel.subgroup.broadcast.first %block_match_count_reduced : i32 + scf.if %matches { + %match_ordinal_i32 = scalar.addi %matched_base, %block_prefix : i32 + %match_ordinal0 = index.cast %match_ordinal_i32 : i32 to index + %match_ordinal = index.assume %match_ordinal0 [range(%match_ordinal0, 0, 2047)] : index + // Top-k route IDs are unique within a token, so one expert can own at + // most token_count assignments. + %bounded_match_ordinal, %table_token_count = index.assume %match_ordinal, %bounded_token_count [lt(%match_ordinal, %bounded_token_count)] : index, index + %assignment_i32 = index.cast %assignment : index to i32 + view.store %assignment_i32, %assignment_view[%expert, %bounded_match_ordinal] : i32, view<[%table_expert_count]x[%bounded_token_count]xi32> + } + %next_matched_base = scalar.addi %matched_base, %block_match_count : i32 + scf.yield %next_matched_base : i32 + } + %is_lane_zero = index.cmp eq, %lane, %c0 : index + scf.if %is_lane_zero { + view.store %expert_route_count, %count_view[%expert] : i32, view<[%table_expert_count]xi32> + } + kernel.return +} + +// Compacts expert assignment counts into exact 32-row projection partitions. +// +// One lane owns each expert. A workgroup scan assigns deterministic descriptor +// offsets, then each lane publishes an 18-bit descriptor containing the expert, +// 32-row partition ordinal, and tail row count. The consumer can launch a +// distribution-independent grid without serializing a concentrated expert +// inside one workgroup. +kernel.def @qwen3_moe_build_expert_partition_table(%token_count: index, %route_count: index, %expert_count: index) { + %c1 = index.constant 1 : index + %workgroup_size = index.constant 128 : index + kernel.launch.config workgroups(%c1, %c1, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %route_count: index, %expert_count: index, %expert_table: buffer, %partition_table: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_route_count = index.assume %route_count [range(%route_count, 1, 8)] : index + %bounded_expert_count = index.assume %expert_count [range(%expert_count, 1, 128)] : index + %lane = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c13_i32 = scalar.constant 13 : i32 + %c0_offset = index.constant 0 : offset + %c4_bytes = index.constant 4 : offset + %assignment_count = index.mul %bounded_token_count, %bounded_route_count : index + %rounded_assignment_count = index.add %assignment_count, %c31 : index + %assignment_partition_count = index.div %rounded_assignment_count, %c32 : index + %maximum_partition_count = index.add %assignment_partition_count, %bounded_expert_count : index + %count_view = buffer.view %expert_table[%c0_offset] : buffer -> view<[%bounded_expert_count]xi32> + %partition_count_view = buffer.view %partition_table[%c0_offset] : buffer -> view<1xi32> + %partition_descriptor_view = buffer.view %partition_table[%c4_bytes] : buffer -> view<[%maximum_partition_count]xi32> + %has_expert = index.cmp ult, %lane, %bounded_expert_count : index + %expert_assignment_count_i32 = scf.if %has_expert -> (i32) { + %expert, %table_expert_count = index.assume %lane, %bounded_expert_count [lt(%lane, %bounded_expert_count)] : index, index + %loaded = view.load %count_view[%expert] : view<[%bounded_expert_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %c0_i32 : i32 + } + %expert_assignment_count0 = index.cast %expert_assignment_count_i32 : i32 to index + %expert_assignment_count = index.assume %expert_assignment_count0 [range(%expert_assignment_count0, 0, 2048)] : index + %rounded_expert_assignment_count = index.add %expert_assignment_count, %c31 : index + %expert_partition_count = index.div %rounded_expert_assignment_count, %c32 : index + %expert_partition_count_i32 = index.cast %expert_partition_count : index to i32 + %expert_partition_base_i32 = kernel.workgroup.scan %expert_partition_count_i32 {direction = forward, mode = exclusive} : i32 + %partition_count_i32 = kernel.workgroup.reduce %expert_partition_count_i32 : i32 + %partition_count = index.cast %partition_count_i32 : i32 to index + %bounded_partition_count, %table_partition_capacity = index.assume %partition_count, %maximum_partition_count [lt(%partition_count, %maximum_partition_count)] : index, index + %expert_partition_base0 = index.cast %expert_partition_base_i32 : i32 to index + %expert_partition_base = index.assume %expert_partition_base0 [range(%expert_partition_base0, 0, 639)] : index + scf.if %has_expert { + %expert, %table_expert_count = index.assume %lane, %bounded_expert_count [lt(%lane, %bounded_expert_count)] : index, index + %expert_i32 = index.cast %expert : index to i32 + scf.for %partition = [%c0 to %expert_partition_count step %c1] { + %descriptor_ordinal0 = index.add %expert_partition_base, %partition : index + %descriptor_ordinal, %descriptor_count = index.assume %descriptor_ordinal0, %bounded_partition_count [lt(%descriptor_ordinal0, %bounded_partition_count)] : index, index + %table_descriptor_ordinal, %table_descriptor_capacity = index.assume %descriptor_ordinal, %maximum_partition_count [lt(%descriptor_ordinal, %maximum_partition_count)] : index, index + %partition_remainder = index.rem %expert_assignment_count, %c32 : index + %has_partial_tail = index.cmp ne, %partition_remainder, %c0 : index + %partition_row_count = scf.if %has_partial_tail -> (index) { + %next_partition = index.add %partition, %c1 : index + %is_tail_partition = index.cmp eq, %next_partition, %expert_partition_count : index + %tail_row_count = scf.if %is_tail_partition -> (index) { + scf.yield %partition_remainder : index + } else { + scf.yield %c32 : index + } + scf.yield %tail_row_count : index + } else { + scf.yield %c32 : index + } + %partition_i32 = index.cast %partition : index to i32 + %partition_row_count_i32 = index.cast %partition_row_count : index to i32 + %packed_partition = scalar.shli %partition_i32, %c7_i32 : i32 + %partition_row_count_minus_one = scalar.subi %partition_row_count_i32, %c1_i32 : i32 + %packed_row_count = scalar.shli %partition_row_count_minus_one, %c13_i32 : i32 + %packed_expert_partition = scalar.ori %expert_i32, %packed_partition : i32 + %packed_descriptor = scalar.ori %packed_expert_partition, %packed_row_count : i32 + view.store %packed_descriptor, %partition_descriptor_view[%table_descriptor_ordinal] : i32, view<[%maximum_partition_count]xi32> + } + } + %is_lane_zero = index.cmp eq, %lane, %c0 : index + scf.if %is_lane_zero { + view.store %partition_count_i32, %partition_count_view[%c0] : i32, view<1xi32> + } + kernel.return +} + +// Small-token body for routed Q4_K gate/up projections. One wave computes one +// [token, route, output channel]. Callers own launch geometry and may append a +// producer epilogue after the F32 SwiGLU value is visible. +template.def<@qwen3_moe.routed_gate_up.q4k_q8.body> device @qwen3_moe_routed_gate_up_swiglu_q4k_q8_body(%publish_output: i1, %token_count: index, %token: index, %route_count: index, %route: index, %route_stride: index, %expert_count: index, %output_size: index, %channel: index, %lane: index, %q8_input: buffer, %route_ids: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) { + %input_size = config.get @qwen3_moe.routed_gate_up.input_size : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 512)] : index + %bounded_route_count, %bounded_route_stride = index.assume %route_count, %route_stride [range(%route_count, 1, 8), range(%route_stride, 1, 128), le(%route_count, %route_stride)] : index, index + %bounded_expert_count = index.assume %expert_count [range(%expert_count, 1, 128)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 4096)] : index + %bounded_token, %body_token_count = index.assume %token, %bounded_token_count [lt(%token, %bounded_token_count)] : index, index + %bounded_channel, %body_output_size = index.assume %channel, %bounded_output_size [lt(%channel, %bounded_output_size)] : index, index + %bounded_route, %body_route_count = index.assume %route, %bounded_route_count [lt(%route, %bounded_route_count)] : index, index + %bounded_lane = index.assume %lane [range(%lane, 0, 31)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %c1023 = index.constant 1023 : index + %c1024 = index.constant 1024 : index + %q4_block_bytes = index.constant 144 : index + %q4_block_bytes_offset = index.constant 144 : offset + %q8_group_bytes = index.constant 144 : offset + %q8_payload_byte_add = index.constant 16 : offset + %c1_byte = index.constant 1 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %c0_i32 = scalar.constant 0 : i32 + %c0_offset = index.constant 0 : offset + %q4_block_count = index.div %input_size, %c256 : index + %weight_row_bytes = index.mul %q4_block_count, %q4_block_bytes : index + %weight_expert_bytes = index.mul %bounded_output_size, %weight_row_bytes : index + %q8_group_count = index.div %input_size, %c128 : index + %q8_bytes_per_token = index.mul %q8_group_count, %q4_block_bytes : index + %padded_input_size = index.add %input_size, %c1023 : index + %iteration_count = index.div %padded_input_size, %c1024 : index + %q8_noalias, %route_ids_noalias, %gate_noalias, %up_noalias, %output_noalias = buffer.assume.noalias %q8_input, %route_ids, %gate_weight, %up_weight, %output : buffer, buffer, buffer, buffer, buffer + %route_ids_view = buffer.view %route_ids_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_route_stride]xi32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%bounded_route_count]x[%bounded_output_size]xf32> + %route_index, %route_row_stride = index.assume %bounded_route, %bounded_route_stride [lt(%bounded_route, %bounded_route_stride)] : index, index + %expert_i32 = view.load %route_ids_view[%bounded_token, %route_index] : view<[%bounded_token_count]x[%bounded_route_stride]xi32> -> i32 + %expert0 = index.cast %expert_i32 : i32 to index + %expert = index.assume %expert0 [range(%expert0, 0, 127)] : index + %expert_byte_base = index.mul %expert, %weight_expert_bytes : index + %channel_byte_add = index.mul %bounded_channel, %weight_row_bytes : index + %row_byte_index = index.add %expert_byte_base, %channel_byte_add : index + %row_byte_base = index.scale %row_byte_index, %c1_byte : index, offset -> offset + %q8_token_byte_index = index.mul %bounded_token, %q8_bytes_per_token : index + %q8_token_byte_base = index.scale %q8_token_byte_index, %c1_byte : index, offset -> offset + %q4_group_pair0 = index.div %bounded_lane, %c2 : index + %q4_group_pair1 = index.rem %q4_group_pair0, %c4 : index + %q4_group_pair = index.assume %q4_group_pair1 [range(%q4_group_pair1, 0, 3)] : index + %q4_half0 = index.rem %bounded_lane, %c2 : index + %q4_half = index.assume %q4_half0 [range(%q4_half0, 0, 1)] : index + %lane_q4_block = index.div %bounded_lane, %c8 : index + %q8_group_in_block0 = index.div %q4_group_pair, %c2 : index + %q8_group_in_block = index.assume %q8_group_in_block0 [range(%q8_group_in_block0, 0, 1)] : index + %pair_in_q8_group0 = index.rem %q4_group_pair, %c2 : index + %pair_in_q8_group = index.assume %pair_in_q8_group0 [range(%pair_in_q8_group0, 0, 1)] : index + %q8_low_inner_block0 = index.mul %pair_in_q8_group, %c2 : index + %q8_low_inner_block = index.assume %q8_low_inner_block0 [range(%q8_low_inner_block0, 0, 2)] : index + %q8_high_inner_block0 = index.add %q8_low_inner_block, %c1 : index + %q8_high_inner_block = index.assume %q8_high_inner_block0 [range(%q8_high_inner_block0, 1, 3)] : index + %q8_half_word_add = index.mul %q4_half, %c4 : index + %q8_low_inner_word_base = index.mul %q8_low_inner_block, %c8 : index + %q8_low_word_index0 = index.add %q8_low_inner_word_base, %q8_half_word_add : index + %q8_low_word_index = index.assume %q8_low_word_index0 [range(%q8_low_word_index0, 0, 20)] : index + %q8_high_inner_word_base = index.mul %q8_high_inner_block, %c8 : index + %q8_high_word_index0 = index.add %q8_high_inner_word_base, %q8_half_word_add : index + %q8_high_word_index = index.assume %q8_high_word_index0 [range(%q8_high_word_index0, 8, 28)] : index + %q8_low_ds_index0 = index.mul %q8_low_inner_block, %c2 : index + %q8_low_ds_index = index.assume %q8_low_ds_index0 [range(%q8_low_ds_index0, 0, 4)] : index + %gate_acc, %up_acc = scf.for %iteration = [%c0 to %iteration_count step %c1](%gate_iter = %c0_f32 : f32, %up_iter = %c0_f32 : f32) -> (f32, f32) unroll { + %iteration_q4_block = index.mul %iteration, %c4 : index + %q4_block0 = index.add %iteration_q4_block, %lane_q4_block : index + %valid_q4_block = index.cmp ult, %q4_block0, %q4_block_count : index + %gate_contribution, %up_contribution = scf.if %valid_q4_block -> (f32, f32) { + %q4_block, %bounded_q4_block_count = index.assume %q4_block0, %q4_block_count [lt(%q4_block0, %q4_block_count)] : index, index + %q8_block_group_base = index.mul %q4_block, %c2 : index + %q8_group0 = index.add %q8_block_group_base, %q8_group_in_block : index + %q8_group, %bounded_q8_group_count = index.assume %q8_group0, %q8_group_count [lt(%q8_group0, %q8_group_count)] : index, index + %q8_group_byte_add = index.scale %q8_group, %q8_group_bytes : index, offset -> offset + %q8_group_byte_base = index.add %q8_token_byte_base, %q8_group_byte_add : offset + %q8_payload_byte_base = index.add %q8_group_byte_base, %q8_payload_byte_add : offset + %q8_ds_view = buffer.view %q8_noalias[%q8_group_byte_base] : buffer -> view<8xf16> + %q8_words_view = buffer.view %q8_noalias[%q8_payload_byte_base] : buffer -> view<32xi32> + %q8_ds = vector.load %q8_ds_view[%q8_low_ds_index] : view<8xf16> -> vector<4xf16> + %q8_low_d_f16 = vector.extract %q8_ds[0] : vector<4xf16> -> f16 + %q8_low_s_f16 = vector.extract %q8_ds[1] : vector<4xf16> -> f16 + %q8_high_d_f16 = vector.extract %q8_ds[2] : vector<4xf16> -> f16 + %q8_high_s_f16 = vector.extract %q8_ds[3] : vector<4xf16> -> f16 + %q8_low_d = scalar.extf %q8_low_d_f16 : f16 to f32 + %q8_low_s = scalar.extf %q8_low_s_f16 : f16 to f32 + %q8_high_d = scalar.extf %q8_high_d_f16 : f16 to f32 + %q8_high_s = scalar.extf %q8_high_s_f16 : f16 to f32 + %q8_low_words = vector.load %q8_words_view[%q8_low_word_index] : view<32xi32> -> vector<4xi32> + %q8_high_words = vector.load %q8_words_view[%q8_high_word_index] : view<32xi32> -> vector<4xi32> + %q8_low_values = vector.bitcast %q8_low_words : vector<4xi32> to vector<16xi8> + %q8_high_values = vector.bitcast %q8_high_words : vector<4xi32> to vector<16xi8> + %q4_block_byte_add = index.scale %q4_block, %q4_block_bytes_offset : index, offset -> offset + %q4_block_byte_base = index.add %row_byte_base, %q4_block_byte_add : offset + %gate_header_view = buffer.view %gate_noalias[%q4_block_byte_base] : buffer -> view<4xi32> + %up_header_view = buffer.view %up_noalias[%q4_block_byte_base] : buffer -> view<4xi32> + %gate_header_words = vector.load %gate_header_view[0] : view<4xi32> -> vector<4xi32> + %up_header_words = vector.load %up_header_view[0] : view<4xi32> -> vector<4xi32> + %gate_q4_low, %gate_low_d_scale, %gate_low_dmin_scale, %gate_q4_high, %gate_high_d_scale, %gate_high_dmin_scale = func.call @qwen3_moe_q4k_chunk_pair_global(%gate_noalias, %row_byte_base, %q4_block, %q4_group_pair, %q4_half, %gate_header_words) : (buffer, offset, index, index, index, vector<4xi32>) -> (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) + %up_q4_low, %up_low_d_scale, %up_low_dmin_scale, %up_q4_high, %up_high_d_scale, %up_high_dmin_scale = func.call @qwen3_moe_q4k_chunk_pair_global(%up_noalias, %row_byte_base, %q4_block, %q4_group_pair, %q4_half, %up_header_words) : (buffer, offset, index, index, index, vector<4xi32>) -> (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) + %gate_low = func.call @qwen3_moe_q4k_q8_1_dot(%gate_q4_low, %gate_low_d_scale, %gate_low_dmin_scale, %q8_low_values, %q8_low_d, %q8_low_s) : (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) -> (f32) + %gate_high = func.call @qwen3_moe_q4k_q8_1_dot(%gate_q4_high, %gate_high_d_scale, %gate_high_dmin_scale, %q8_high_values, %q8_high_d, %q8_high_s) : (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) -> (f32) + %up_low = func.call @qwen3_moe_q4k_q8_1_dot(%up_q4_low, %up_low_d_scale, %up_low_dmin_scale, %q8_low_values, %q8_low_d, %q8_low_s) : (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) -> (f32) + %up_high = func.call @qwen3_moe_q4k_q8_1_dot(%up_q4_high, %up_high_d_scale, %up_high_dmin_scale, %q8_high_values, %q8_high_d, %q8_high_s) : (vector<16xi8>, f32, f32, vector<16xi8>, f32, f32) -> (f32) + %gate_pair = scalar.addf %gate_low, %gate_high : f32 + %up_pair = scalar.addf %up_low, %up_high : f32 + scf.yield %gate_pair, %up_pair : f32, f32 + } else { + scf.yield %c0_f32, %c0_f32 : f32, f32 + } + %gate_next = scalar.addf %gate_iter, %gate_contribution : f32 + %up_next = scalar.addf %up_iter, %up_contribution : f32 + scf.yield %gate_next, %up_next : f32, f32 + } + %gate_dot = kernel.subgroup.reduce %gate_acc : f32 + %up_dot = kernel.subgroup.reduce %up_acc : f32 + %lane_i32 = index.cast %bounded_lane : index to i32 + %is_lane_zero = scalar.cmpi eq, %lane_i32, %c0_i32 : i32 + %writes_output = scalar.andi %publish_output, %is_lane_zero : i1 + scf.if %writes_output { + %gate_silu = scalar.siluf %gate_dot : f32 + %result = scalar.mulf %gate_silu, %up_dot : f32 + view.store %result, %output_view[%bounded_token, %bounded_route, %bounded_channel] : f32, view<[%bounded_token_count]x[%bounded_route_count]x[%bounded_output_size]xf32> + } + template.return +} + +// Small-token schedule for routed Q4_K gate/up projections. Each workgroup +// packs four independent channel waves for one [token, route] pair. Eight +// lanes consume both nibbles of the packed codes for one 256-element block, so +// each wave advances through K in 1024-element stripes. Gate and up share each +// paired Q8_1 load before independent raw-weight dot products and a fused +// SwiGLU epilogue. +kernel.def @qwen3_moe_routed_gate_up_swiglu_q4k_q8(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index, %output_size: index) { + %configured_route_count = config.get @qwen3_moe.routed_gate_up.route_count : index + %configured_output_size = config.get @qwen3_moe.routed_gate_up.output_size : index + %unit = index.constant 1 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %workgroup_size = index.constant 128 : index + %padded_output_size = index.add %configured_output_size, %c3 : index + %channel_workgroup_count = index.div %padded_output_size, %c4 : index + kernel.launch.config workgroups(%channel_workgroup_count, %configured_route_count, %token_count) workgroup_size(%workgroup_size, %unit, %unit) : index +} launch(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index, %output_size: index, %q8_input: buffer, %route_ids: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) { + %configured_route_count0 = config.get @qwen3_moe.routed_gate_up.route_count : index + %configured_expert_count0 = config.get @qwen3_moe.routed_gate_up.expert_count : index + %configured_output_size0 = config.get @qwen3_moe.routed_gate_up.output_size : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 512)] : index + %bounded_route_count, %configured_route_count = index.assume %route_count, %configured_route_count0 [range(%route_count, 1, 8), eq(%route_count, %configured_route_count0)] : index, index + %bounded_route_stride = index.assume %route_stride [range(%route_stride, 1, 128)] : index + %bounded_expert_count, %configured_expert_count = index.assume %expert_count, %configured_expert_count0 [range(%expert_count, 1, 128), eq(%expert_count, %configured_expert_count0)] : index, index + %bounded_output_size, %configured_output_size = index.assume %output_size, %configured_output_size0 [range(%output_size, 1, 4096), eq(%output_size, %configured_output_size0)] : index, index + %token0 = kernel.workgroup.id : index + %channel_workgroup = kernel.workgroup.id : index + %route0 = kernel.workgroup.id : index + %subgroup = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %valid_token = index.cmp ult, %token0, %bounded_token_count : index + %safe_token0 = scf.select %valid_token, %token0, %c0 : index + %safe_token, %body_token_count = index.assume %safe_token0, %bounded_token_count [lt(%safe_token0, %bounded_token_count)] : index, index + %channel_base = index.mul %channel_workgroup, %c4 : index + %channel0 = index.add %channel_base, %subgroup : index + %valid_channel = index.cmp ult, %channel0, %bounded_output_size : index + %safe_channel0 = scf.select %valid_channel, %channel0, %c0 : index + %safe_channel, %body_output_size = index.assume %safe_channel0, %bounded_output_size [lt(%safe_channel0, %bounded_output_size)] : index, index + %route, %body_route_count = index.assume %route0, %bounded_route_count [lt(%route0, %bounded_route_count)] : index, index + %publishes_output = scalar.andi %valid_token, %valid_channel : i1 + template.apply<@qwen3_moe.routed_gate_up.q4k_q8.body>(%publishes_output, %body_token_count, %safe_token, %body_route_count, %route, %bounded_route_stride, %bounded_expert_count, %body_output_size, %safe_channel, %lane, %q8_input, %route_ids, %gate_weight, %up_weight, %output) : (i1, index, index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer) + kernel.return +} + +// Packs one Q8_1 x4 physical group with one selected subgroup. Each lane owns +// four adjacent values and contributes its maximum and sum to one of four +// eight-lane logical blocks. One-hot vectors let ordinary subgroup reductions +// compute all four block aggregates without LDS or workgroup barriers. +template.def<@qwen3_moe.routed_gate_up.quantize_q8_1_x4.subgroup_body> device @qwen3_moe_routed_gate_up_quantize_q8_1_x4_subgroup_body(%publish_output: i1, %group_count0: index, %group0: index, %input: buffer, %output: buffer) { + %group_count, %group = index.assume %group_count0, %group0 [range(%group_count0, 1, 524288), lt(%group0, %group_count0)] : index, index + %lane0 = kernel.subgroup.lane.id : index + %lane = index.assume %lane0 [range(%lane0, 0, 31)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c128 = index.constant 128 : index + %group_bytes = index.constant 144 : offset + %payload_byte_add = index.constant 16 : offset + %c0_f32 = scalar.constant 0.0 : f32 + %c1_f32 = scalar.constant 1.0 : f32 + %c127_f32 = scalar.constant 127.0 : f32 + %c0_f32x4 = vector.constant 0.0 : vector<4xf32> + %c0_offset = index.constant 0 : offset + %launched_element_count = index.mul %group_count, %c128 : index + %group_element_base = index.mul %group, %c128 : index + %lane_element_add = index.mul %lane, %c4 : index + %input_index = index.add %group_element_base, %lane_element_add : index + %input_noalias, %output_noalias = buffer.assume.noalias %input, %output : buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%launched_element_count]xf32> + %input_values = vector.load %input_view[%input_index] : view<[%launched_element_count]xf32> -> vector<4xf32> + %absolute_values = vector.absf %input_values : vector<4xf32> + %thread_maximum = vector.reduce %absolute_values, %c0_f32 : vector<4xf32>, f32 + %block0 = index.div %lane, %c8 : index + %block = index.assume %block0 [range(%block0, 0, 3)] : index + %block_maxima = vector.insert %thread_maximum into %c0_f32x4[%block] : f32, vector<4xf32> + %reduced_maxima = kernel.subgroup.reduce %block_maxima : vector<4xf32> + %amax = vector.extract %reduced_maxima[%block] : vector<4xf32> -> f32 + %d = scalar.divf %amax, %c127_f32 : f32 + %d_nonzero = scalar.cmpf one, %d, %c0_f32 : f32 + %d_inverse = scf.if %d_nonzero -> (f32) { + %inverse = scalar.divf %c1_f32, %d : f32 + scf.yield %inverse : f32 + } else { + scf.yield %c0_f32 : f32 + } + %d_inverse_vector = vector.splat %d_inverse : vector<4xf32> + %scaled_values = vector.mulf %input_values, %d_inverse_vector : vector<4xf32> + %rounded_values = vector.roundf %scaled_values : vector<4xf32> + %quantized_values = vector.fptosi %rounded_values : vector<4xf32> to vector<4xi8> + %packed_word = vector.bitcast %quantized_values : vector<4xi8> to vector<1xi32> + %group_byte_offset = index.scale %group, %group_bytes : index, offset -> offset + %payload_byte_offset = index.add %group_byte_offset, %payload_byte_add : offset + %group_ds = buffer.view %output_noalias[%group_byte_offset] : buffer -> view<8xf16> + %group_qs = buffer.view %output_noalias[%payload_byte_offset] : buffer -> view<32xi32> + scf.if %publish_output { + vector.store %packed_word, %group_qs[%lane] : vector<1xi32>, view<32xi32> + } + %thread_sum = vector.reduce %rounded_values, %c0_f32 : vector<4xf32>, f32 + %block_sums0 = vector.insert %thread_sum into %c0_f32x4[%block] : f32, vector<4xf32> + %reduced_sums = kernel.subgroup.reduce %block_sums0 : vector<4xf32> + %quantized_sum = vector.extract %reduced_sums[%block] : vector<4xf32> -> f32 + %s = scalar.mulf %quantized_sum, %d : f32 + %word_in_block = index.rem %lane, %c8 : index + %is_block_leader = index.cmp eq, %word_in_block, %c0 : index + %publishes_metadata = scalar.andi %publish_output, %is_block_leader : i1 + scf.if %publishes_metadata { + %d_f16 = scalar.fptrunc %d : f32 to f16 + %s_f16 = scalar.fptrunc %s : f32 to f16 + %ds_index = index.mul %block, %c2 : index + %s_index = index.add %ds_index, %c1 : index + view.store %d_f16, %group_ds[%ds_index] : f16, view<8xf16> + view.store %s_f16, %group_ds[%s_index] : f16, view<8xf16> + } + template.return +} + +// Decode producer that publishes both the ordinary F32 SwiGLU rows and their +// packed Q8_1 x4 representation. Four independent channel waves share one +// workgroup and publish one completion arrival. The last workgroup to complete +// a 128-channel physical group packs all four logical Q8_1 blocks with one +// selected subgroup, then resets the counter before returning. +kernel.def @qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index, %output_size: index) { + %configured_route_count = config.get @qwen3_moe.routed_gate_up.route_count : index + %configured_output_size = config.get @qwen3_moe.routed_gate_up.output_size : index + %unit = index.constant 1 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %workgroup_size = index.constant 128 : index + %padded_output_size = index.add %configured_output_size, %c3 : index + %channel_workgroup_count = index.div %padded_output_size, %c4 : index + kernel.launch.config workgroups(%channel_workgroup_count, %configured_route_count, %unit) workgroup_size(%workgroup_size, %unit, %unit) : index +} launch(%token_count: index, %route_count: index, %route_stride: index, %expert_count: index, %output_size: index, %q8_input: buffer, %route_ids: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer, %completion_counters: buffer, %next_q8_output: buffer) { + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 1)] : index + %configured_route_count0 = config.get @qwen3_moe.routed_gate_up.route_count : index + %configured_expert_count0 = config.get @qwen3_moe.routed_gate_up.expert_count : index + %configured_output_size0 = config.get @qwen3_moe.routed_gate_up.output_size : index + %bounded_route_count, %configured_route_count = index.assume %route_count, %configured_route_count0 [range(%route_count, 1, 8), eq(%route_count, %configured_route_count0)] : index, index + %bounded_route_stride = index.assume %route_stride [range(%route_stride, 1, 128)] : index + %bounded_expert_count, %configured_expert_count = index.assume %expert_count, %configured_expert_count0 [range(%expert_count, 1, 128), eq(%expert_count, %configured_expert_count0)] : index, index + %bounded_output_size, %configured_output_size = index.assume %output_size, %configured_output_size0 [range(%output_size, 128, 4096), mul(%output_size, 128), eq(%output_size, %configured_output_size0)] : index, index + %q8_input_noalias, %route_ids_noalias, %gate_weight_noalias, %up_weight_noalias, %output_noalias, %completion_counters_noalias, %next_q8_output_noalias = buffer.assume.noalias %q8_input, %route_ids, %gate_weight, %up_weight, %output, %completion_counters, %next_q8_output : buffer, buffer, buffer, buffer, buffer, buffer, buffer + %publishes_swiglu = scalar.constant true : i1 + %body_token = index.constant 0 : index + %channel_workgroup = kernel.workgroup.id : index + %route = kernel.workgroup.id : index + %token = kernel.workgroup.id : index + %subgroup = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %workitem = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c128 = index.constant 128 : index + %c0_i32 = scalar.constant 0 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c0_offset = index.constant 0 : offset + %counter_scratch_bytes = index.constant 4 : offset + %channel_base = index.mul %channel_workgroup, %c4 : index + %channel = index.add %channel_base, %subgroup : index + template.apply<@qwen3_moe.routed_gate_up.q4k_q8.body>(%publishes_swiglu, %bounded_token_count, %body_token, %bounded_route_count, %route, %bounded_route_stride, %bounded_expert_count, %bounded_output_size, %channel, %lane, %q8_input_noalias, %route_ids_noalias, %gate_weight_noalias, %up_weight_noalias, %output_noalias) : (i1, index, index, index, index, index, index, index, index, index, buffer, buffer, buffer, buffer, buffer) + %physical_group_count = index.div %bounded_output_size, %c128 : index + %row_count = index.mul %bounded_token_count, %bounded_route_count : index + %completion_counter_count = index.mul %row_count, %physical_group_count : index + %token_row_base = index.mul %token, %bounded_route_count : index + %row = index.add %token_row_base, %route : index + %row_group_base = index.mul %row, %physical_group_count : index + %group_in_row = index.div %channel_base, %c128 : index + %counter_index0 = index.add %row_group_base, %group_in_row : index + %counter_index, %bounded_completion_counter_count = index.assume %counter_index0, %completion_counter_count [lt(%counter_index0, %completion_counter_count)] : index, index + %completion_counters_aligned = buffer.assume.alignment %completion_counters_noalias {minimum_alignment = 16} : buffer + %completion_counters_view = buffer.view %completion_counters_aligned[%c0_offset] : buffer -> view<[%bounded_completion_counter_count]xi32> + %counter_scratch = buffer.alloca align(4) %counter_scratch_bytes : buffer + %counter_scratch_view = buffer.view %counter_scratch[%c0_offset] : buffer -> view<1xi32> + %is_arrival_lane = index.cmp eq, %workitem, %c0 : index + // Publish every producer lane's SwiGLU store before the leader advances one + // workgroup arrival. The last arrival then acquires the physical group. + kernel.barrier scope(workgroup) ordering(release) + scf.if %is_arrival_lane { + %old_counter = view.atomic.rmw %c4_i32, %completion_counters_view[%counter_index] {ordering = acq_rel, scope = device} : i32, view<[%bounded_completion_counter_count]xi32> -> i32 + view.store %old_counter, %counter_scratch_view[%c0] : i32, view<1xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %old_counter = view.load %counter_scratch_view[%c0] : view<1xi32> -> i32 + %group_size_i32 = index.cast %c128 : index to i32 + %last_arrival_i32 = scalar.subi %group_size_i32, %c4_i32 : i32 + %negative_group_size_i32 = scalar.subi %c0_i32, %group_size_i32 : i32 + %is_last_arrival = scalar.cmpi eq, %old_counter, %last_arrival_i32 : i32 + scf.if %is_last_arrival { + kernel.barrier scope(workgroup) ordering(acquire) + %publish_output = scalar.constant true : i1 + %is_quantize_subgroup = index.cmp eq, %subgroup, %c0 : index + scf.if %is_quantize_subgroup { + template.apply<@qwen3_moe.routed_gate_up.quantize_q8_1_x4.subgroup_body>(%publish_output, %bounded_completion_counter_count, %counter_index, %output_noalias, %next_q8_output_noalias) : (i1, index, index, buffer, buffer) + } + kernel.barrier scope(workgroup) ordering(release) + scf.if %is_arrival_lane { + view.atomic.reduce %negative_group_size_i32, %completion_counters_view[%counter_index] {ordering = release, scope = device} : i32, view<[%bounded_completion_counter_count]xi32> + } + } + kernel.return +} + +// Large-token schedule for routed Q4_K gate/up projections. Each workgroup +// owns one expert, 32 output channels, and up to 32 selected rows. It +// reads those rows from the compact expert table, stages one K=256 slice of +// both raw Q4_K projections and Q8_1 activations in LDS, and reuses each +// decoded weight chunk across four rows before a fused SwiGLU epilogue. +kernel.def @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped(%token_count: index, %route_count: index, %expert_count: index, %output_size: index) { + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %configured_expert_count = config.get @qwen3_moe.routed_gate_up.expert_count : index + %configured_output_size = config.get @qwen3_moe.routed_gate_up.output_size : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %workgroup_size = index.constant 256 : index + %padded_output_size = index.add %configured_output_size, %c31 : index + %output_tiles = index.div %padded_output_size, %c32 : index + %padded_token_count = index.add %token_capacity, %c31 : index + %route_tiles = index.div %padded_token_count, %c32 : index + %route_partitions = index.min %route_tiles, %c4 : index + kernel.launch.config workgroups(%output_tiles, %route_partitions, %configured_expert_count) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %route_count: index, %expert_count: index, %output_size: index, %q8_input: buffer, %expert_table: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) { + %input_size = config.get @qwen3_moe.routed_gate_up.input_size : index + %token_capacity = config.get @qwen3_moe.workload.token_capacity : index + %configured_route_count0 = config.get @qwen3_moe.routed_gate_up.route_count : index + %configured_expert_count0 = config.get @qwen3_moe.routed_gate_up.expert_count : index + %configured_output_size0 = config.get @qwen3_moe.routed_gate_up.output_size : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 512), le(%token_count, %token_capacity)] : index + %bounded_route_count, %configured_route_count = index.assume %route_count, %configured_route_count0 [range(%route_count, 1, 8), eq(%route_count, %configured_route_count0)] : index, index + %bounded_expert_count, %configured_expert_count = index.assume %expert_count, %configured_expert_count0 [range(%expert_count, 1, 128), eq(%expert_count, %configured_expert_count0)] : index, index + %bounded_output_size, %configured_output_size = index.assume %output_size, %configured_output_size0 [range(%output_size, 1, 4096), eq(%output_size, %configured_output_size0)] : index, index + %channel_tile = kernel.workgroup.id : index + %route_tile = kernel.workgroup.id : index + %expert0 = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c32 = index.constant 32 : index + %c36 = index.constant 36 : index + %c72 = index.constant 72 : index + %c128 = index.constant 128 : index + %c256 = index.constant 256 : index + %gate_stage_word_count = index.constant 1152 : index + %q8_stage_word_count = index.constant 2304 : index + %c0_i32 = scalar.constant 0 : i32 + %c0_f32_rows = vector.constant 0.0 : vector<4xf32> + %c0_offset = index.constant 0 : offset + %c4_bytes = index.constant 4 : offset + %q4_block_bytes = index.constant 144 : offset + %q8_row_bytes = index.constant 288 : offset + %row_id_bytes = index.constant 128 : offset + %weight_stage_bytes = index.constant 4608 : offset + %q8_stage_bytes = index.constant 9216 : offset + %assignment_count = index.mul %bounded_token_count, %bounded_route_count : index + %expert, %table_expert_count = index.assume %expert0, %bounded_expert_count [lt(%expert0, %bounded_expert_count)] : index, index + %assignment_table_byte_base = index.scale %table_expert_count, %c4_bytes : index, offset -> offset + %q4_block_count = index.div %input_size, %c256 : index + %output_route_count = index.mul %bounded_token_count, %bounded_route_count : index + %q8_noalias, %expert_table_noalias, %gate_noalias, %up_noalias, %output_noalias = buffer.assume.noalias %q8_input, %expert_table, %gate_weight, %up_weight, %output : buffer, buffer, buffer, buffer, buffer + %q8_words = buffer.view %q8_noalias[%c0_offset] : buffer -> view<[%bounded_token_count]x[%q4_block_count]x72xi32> + %count_view = buffer.view %expert_table_noalias[%c0_offset] : buffer -> view<[%table_expert_count]xi32> + %assignment_view = buffer.view %expert_table_noalias[%assignment_table_byte_base] : buffer -> view<[%table_expert_count]x[%bounded_token_count]xi32> + %gate_words = buffer.view %gate_noalias[%c0_offset] : buffer -> view<[%bounded_expert_count]x[%bounded_output_size]x[%q4_block_count]x36xi32> + %up_words = buffer.view %up_noalias[%c0_offset] : buffer -> view<[%bounded_expert_count]x[%bounded_output_size]x[%q4_block_count]x36xi32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%output_route_count]x[%bounded_output_size]xf32> + %row_ids = buffer.alloca align(16) %row_id_bytes : buffer + %gate_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %up_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %q8_stage = buffer.alloca align(16) %q8_stage_bytes : buffer + %row_ids_view = buffer.view %row_ids[%c0_offset] : buffer -> view<32xi32> + %gate_stage_words = buffer.view %gate_stage[%c0_offset] : buffer -> view<1152xi32> + %up_stage_words = buffer.view %up_stage[%c0_offset] : buffer -> view<1152xi32> + %q8_stage_words = buffer.view %q8_stage[%c0_offset] : buffer -> view<2304xi32> + %channel_base = index.mul %channel_tile, %c32 : index + %channel0 = index.rem %lane, %c32 : index + %channel = index.add %channel_base, %channel0 : index + %row_base0 = index.div %lane, %c32 : index + %row_base = index.assume %row_base0 [range(%row_base0, 0, 7)] : index + %initial_tile_base = index.mul %route_tile, %c32 : index + // Four interleaved route partitions cover the first 128 rows. Smaller + // token counts execute at most one iteration; larger expert populations + // continue in 128-row strides without increasing the launch grid. + %route_partition_step = index.constant 128 : index + %is_lane_zero = index.cmp eq, %lane, %c0 : index + %lane_expert_route_count = scf.if %is_lane_zero -> (i32) { + %loaded = view.load %count_view[%expert] : view<[%table_expert_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %c0_i32 : i32 + } + %expert_route_count_reduced = kernel.workgroup.reduce %lane_expert_route_count : i32 + %expert_route_count_i32 = kernel.subgroup.broadcast.first %expert_route_count_reduced : i32 + %expert_route_count0 = index.cast %expert_route_count_i32 : i32 to index + %expert_route_count = index.assume %expert_route_count0 [range(%expert_route_count0, 0, 4096)] : index + scf.for %tile_base = [%initial_tile_base to %expert_route_count step %route_partition_step] { + %tile_base_i32 = index.cast %tile_base : index to i32 + %remaining_rows_i32 = scalar.subi %expert_route_count_i32, %tile_base_i32 : i32 + %remaining_rows0 = index.cast %remaining_rows_i32 : i32 to index + %remaining_rows = index.assume %remaining_rows0 [range(%remaining_rows0, 1, 4096)] : index + %has_full_tile = index.cmp uge, %remaining_rows, %c32 : index + %tile_row_count = scf.if %has_full_tile -> (index) { + scf.yield %c32 : index + } else { + scf.yield %remaining_rows : index + } + %loads_row = index.cmp ult, %lane, %tile_row_count : index + scf.if %loads_row { + %local_row = index.assume %lane [range(%lane, 0, 31)] : index + %expert_assignment_ordinal0 = index.add %tile_base, %local_row : index + %expert_assignment_ordinal, %table_token_count = index.assume %expert_assignment_ordinal0, %bounded_token_count [lt(%expert_assignment_ordinal0, %bounded_token_count)] : index, index + %assignment_i32 = view.load %assignment_view[%expert, %expert_assignment_ordinal] : view<[%table_expert_count]x[%bounded_token_count]xi32> -> i32 + view.store %assignment_i32, %row_ids_view[%local_row] : i32, view<32xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %gate_acc, %up_acc = scf.for %q4_block = [%c0 to %q4_block_count step %c1](%gate_block_acc = %c0_f32_rows : vector<4xf32>, %up_block_acc = %c0_f32_rows : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>) { + %bounded_q4_block, %body_q4_block_count = index.assume %q4_block, %q4_block_count [lt(%q4_block, %q4_block_count)] : index, index + scf.for %stage_word = [%lane to %gate_stage_word_count step %c256] { + %local_channel = index.div %stage_word, %c36 : index + %word_in_block0 = index.rem %stage_word, %c36 : index + %word_in_block = index.assume %word_in_block0 [range(%word_in_block0, 0, 35)] : index + %global_channel = index.add %channel_base, %local_channel : index + %valid_channel = index.cmp ult, %global_channel, %bounded_output_size : index + %gate_word, %up_word = scf.if %valid_channel -> (i32, i32) { + %bounded_global_channel, %table_output_size = index.assume %global_channel, %bounded_output_size [lt(%global_channel, %bounded_output_size)] : index, index + %gate_loaded = view.load %gate_words[%expert, %bounded_global_channel, %bounded_q4_block, %word_in_block] : view<[%bounded_expert_count]x[%bounded_output_size]x[%q4_block_count]x36xi32> -> i32 + %up_loaded = view.load %up_words[%expert, %bounded_global_channel, %bounded_q4_block, %word_in_block] : view<[%bounded_expert_count]x[%bounded_output_size]x[%q4_block_count]x36xi32> -> i32 + scf.yield %gate_loaded, %up_loaded : i32, i32 + } else { + scf.yield %c0_i32, %c0_i32 : i32, i32 + } + view.store %gate_word, %gate_stage_words[%stage_word] : i32, view<1152xi32> + view.store %up_word, %up_stage_words[%stage_word] : i32, view<1152xi32> + } + scf.for %stage_word = [%lane to %q8_stage_word_count step %c256] { + %local_row = index.div %stage_word, %c72 : index + %word_in_row0 = index.rem %stage_word, %c72 : index + %word_in_row = index.assume %word_in_row0 [range(%word_in_row0, 0, 71)] : index + %valid_row = index.cmp ult, %local_row, %tile_row_count : index + %q8_word = scf.if %valid_row -> (i32) { + %bounded_local_row = index.assume %local_row [range(%local_row, 0, 31)] : index + %assignment_i32 = view.load %row_ids_view[%bounded_local_row] : view<32xi32> -> i32 + %assignment0 = index.cast %assignment_i32 : i32 to index + %assignment = index.assume %assignment0 [range(%assignment0, 0, 4095)] : index + %bounded_assignment, %bounded_assignment_count = index.assume %assignment, %assignment_count [lt(%assignment, %assignment_count)] : index, index + %token0 = index.div %bounded_assignment, %configured_route_count : index + %token, %table_token_count = index.assume %token0, %bounded_token_count [lt(%token0, %bounded_token_count)] : index, index + %loaded = view.load %q8_words[%token, %bounded_q4_block, %word_in_row] : view<[%bounded_token_count]x[%q4_block_count]x72xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %c0_i32 : i32 + } + view.store %q8_word, %q8_stage_words[%stage_word] : i32, view<2304xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %channel_byte_base = index.scale %channel0, %q4_block_bytes : index, offset -> offset + %after_groups_gate, %after_groups_up = scf.for %q4_group = [%c0 to %c8 step %c1](%gate_group_acc = %gate_block_acc : vector<4xf32>, %up_group_acc = %up_block_acc : vector<4xf32>) -> (vector<4xf32>, vector<4xf32>) unroll { + %row1 = index.add %row_base, %c8 : index + %row2_add = index.constant 16 : index + %row2 = index.add %row_base, %row2_add : index + %row3_add = index.constant 24 : index + %row3 = index.add %row_base, %row3_add : index + %row0_byte_base = index.scale %row_base, %q8_row_bytes : index, offset -> offset + %row1_byte_base = index.scale %row1, %q8_row_bytes : index, offset -> offset + %row2_byte_base = index.scale %row2, %q8_row_bytes : index, offset -> offset + %row3_byte_base = index.scale %row3, %q8_row_bytes : index, offset -> offset + %q8_values0, %q8_d0, %q8_s0 = func.call @ggml_q8_1_x4_block(%q8_stage, %row0_byte_base, %q4_group) : (buffer, offset, index) -> (vector<32xi8>, f32, f32) + %q8_values1, %q8_d1, %q8_s1 = func.call @ggml_q8_1_x4_block(%q8_stage, %row1_byte_base, %q4_group) : (buffer, offset, index) -> (vector<32xi8>, f32, f32) + %q8_values2, %q8_d2, %q8_s2 = func.call @ggml_q8_1_x4_block(%q8_stage, %row2_byte_base, %q4_group) : (buffer, offset, index) -> (vector<32xi8>, f32, f32) + %q8_values3, %q8_d3, %q8_s3 = func.call @ggml_q8_1_x4_block(%q8_stage, %row3_byte_base, %q4_group) : (buffer, offset, index) -> (vector<32xi8>, f32, f32) + %q8_low0 = vector.slice %q8_values0[0] : vector<32xi8> -> vector<16xi8> + %q8_low1 = vector.slice %q8_values1[0] : vector<32xi8> -> vector<16xi8> + %q8_low2 = vector.slice %q8_values2[0] : vector<32xi8> -> vector<16xi8> + %q8_low3 = vector.slice %q8_values3[0] : vector<32xi8> -> vector<16xi8> + %q8_high0 = vector.slice %q8_values0[16] : vector<32xi8> -> vector<16xi8> + %q8_high1 = vector.slice %q8_values1[16] : vector<32xi8> -> vector<16xi8> + %q8_high2 = vector.slice %q8_values2[16] : vector<32xi8> -> vector<16xi8> + %q8_high3 = vector.slice %q8_values3[16] : vector<32xi8> -> vector<16xi8> + %q8_rows_low0 = vector.constant 0 : vector<4x16xi8> + %q8_rows_low1 = vector.insert %q8_low0 into %q8_rows_low0[0] : vector<16xi8>, vector<4x16xi8> + %q8_rows_low2 = vector.insert %q8_low1 into %q8_rows_low1[1] : vector<16xi8>, vector<4x16xi8> + %q8_rows_low3 = vector.insert %q8_low2 into %q8_rows_low2[2] : vector<16xi8>, vector<4x16xi8> + %q8_rows_low = vector.insert %q8_low3 into %q8_rows_low3[3] : vector<16xi8>, vector<4x16xi8> + %q8_rows_high0 = vector.constant 0 : vector<4x16xi8> + %q8_rows_high1 = vector.insert %q8_high0 into %q8_rows_high0[0] : vector<16xi8>, vector<4x16xi8> + %q8_rows_high2 = vector.insert %q8_high1 into %q8_rows_high1[1] : vector<16xi8>, vector<4x16xi8> + %q8_rows_high3 = vector.insert %q8_high2 into %q8_rows_high2[2] : vector<16xi8>, vector<4x16xi8> + %q8_rows_high = vector.insert %q8_high3 into %q8_rows_high3[3] : vector<16xi8>, vector<4x16xi8> + %q8_d = vector.from_elements %q8_d0, %q8_d1, %q8_d2, %q8_d3 : vector<4xf32> + %q8_s = vector.from_elements %q8_s0, %q8_s1, %q8_s2, %q8_s3 : vector<4xf32> + %gate_q4_low, %gate_d_scale_low, %gate_dmin_scale_low = func.call @qwen3_moe_q4k_chunk_local(%gate_stage, %channel_byte_base, %c0, %q4_group, %c0) : (buffer, offset, index, index, index) -> (vector<16xi8>, f32, f32) + %gate_q4_high, %gate_d_scale_high, %gate_dmin_scale_high = func.call @qwen3_moe_q4k_chunk_local(%gate_stage, %channel_byte_base, %c0, %q4_group, %c1) : (buffer, offset, index, index, index) -> (vector<16xi8>, f32, f32) + %up_q4_low, %up_d_scale_low, %up_dmin_scale_low = func.call @qwen3_moe_q4k_chunk_local(%up_stage, %channel_byte_base, %c0, %q4_group, %c0) : (buffer, offset, index, index, index) -> (vector<16xi8>, f32, f32) + %up_q4_high, %up_d_scale_high, %up_dmin_scale_high = func.call @qwen3_moe_q4k_chunk_local(%up_stage, %channel_byte_base, %c0, %q4_group, %c1) : (buffer, offset, index, index, index) -> (vector<16xi8>, f32, f32) + %gate_low = func.call @qwen3_moe_q4k_q8_1_dot4_rows(%gate_q4_low, %gate_d_scale_low, %gate_dmin_scale_low, %q8_rows_low, %q8_d, %q8_s) : (vector<16xi8>, f32, f32, vector<4x16xi8>, vector<4xf32>, vector<4xf32>) -> (vector<4xf32>) + %gate_high = func.call @qwen3_moe_q4k_q8_1_dot4_rows(%gate_q4_high, %gate_d_scale_high, %gate_dmin_scale_high, %q8_rows_high, %q8_d, %q8_s) : (vector<16xi8>, f32, f32, vector<4x16xi8>, vector<4xf32>, vector<4xf32>) -> (vector<4xf32>) + %up_low = func.call @qwen3_moe_q4k_q8_1_dot4_rows(%up_q4_low, %up_d_scale_low, %up_dmin_scale_low, %q8_rows_low, %q8_d, %q8_s) : (vector<16xi8>, f32, f32, vector<4x16xi8>, vector<4xf32>, vector<4xf32>) -> (vector<4xf32>) + %up_high = func.call @qwen3_moe_q4k_q8_1_dot4_rows(%up_q4_high, %up_d_scale_high, %up_dmin_scale_high, %q8_rows_high, %q8_d, %q8_s) : (vector<16xi8>, f32, f32, vector<4x16xi8>, vector<4xf32>, vector<4xf32>) -> (vector<4xf32>) + %gate_pair = vector.addf %gate_low, %gate_high : vector<4xf32> + %up_pair = vector.addf %up_low, %up_high : vector<4xf32> + %gate_next = vector.addf %gate_group_acc, %gate_pair : vector<4xf32> + %up_next = vector.addf %up_group_acc, %up_pair : vector<4xf32> + scf.yield %gate_next, %up_next : vector<4xf32>, vector<4xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %after_groups_gate, %after_groups_up : vector<4xf32>, vector<4xf32> + } + %gate_silu = vector.siluf %gate_acc : vector<4xf32> + %result_rows = vector.mulf %gate_silu, %up_acc : vector<4xf32> + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + scf.for %row_variant = [%c0 to %c4 step %c1] unroll { + %row_add = index.mul %row_variant, %c8 : index + %local_row = index.add %row_base, %row_add : index + %valid_row = index.cmp ult, %local_row, %tile_row_count : index + %writes_output = scalar.andi %valid_channel, %valid_row : i1 + scf.if %writes_output { + %bounded_local_row = index.assume %local_row [range(%local_row, 0, 31)] : index + %assignment_i32 = view.load %row_ids_view[%bounded_local_row] : view<32xi32> -> i32 + %assignment0 = index.cast %assignment_i32 : i32 to index + %assignment = index.assume %assignment0 [range(%assignment0, 0, 4095)] : index + %bounded_assignment, %bounded_assignment_count = index.assume %assignment, %assignment_count [lt(%assignment, %assignment_count)] : index, index + %result = vector.extract %result_rows[%row_variant] : vector<4xf32> -> f32 + view.store %result, %output_view[%bounded_assignment, %channel] : f32, view<[%output_route_count]x[%bounded_output_size]xf32> + } + } + } + kernel.return +} + +check.case public @qwen3_moe_routed_gate_up_swiglu_q4k_q8_nonzero_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(2048) : index + %route_count = check.literal value(1) : index + %route_stride = check.literal value(1) : index + %expert_count = check.literal value(1) : index + %output_size = check.literal value(1) : index + %input = check.generate.fill value(0.00390625) : tensor<1x2048xf32> + %q8_input = check.generate.fill value(0) : tensor<1x2304xi8> + %route_ids = check.generate.fill value(0) : tensor<1xi32> + %gate_weight = check.generate.fill value(85) : tensor<8x144xi8> + %up_weight = check.generate.fill value(-86) : tensor<8x144xi8> + %output = check.generate.fill value(0.0) : tensor<1xf32> + %expected = check.generate.fill value(-9024645.0) : tensor<1xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<1x2048xf32>, tensor<1x2304xi8>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8[%token_count, %route_count, %route_stride, %expert_count, %output_size](%token_count, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %output) : [index, index, index, index, index](index, index, index, index, index, tensor<1x2304xi8>, tensor<1xi32>, tensor<8x144xi8>, tensor<8x144xi8>, tensor<1xf32>) + check.expect.close actual(%output) expected(%expected) atol(16.0) rtol(9.9999999999999995e-07) nan(same) : tensor<1xf32> + check.return +} + +// Exact decode topology with compact route IDs and every physical output +// group populated. Two fused invocations prove counter reuse in addition to +// comparing the F32 and packed outputs with the ordinary composition. +check.case public @qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8_differential_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(2048) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(768) : index + %routed_row_count = check.literal value(8) : index + %input = check.generate.fill value(0.00390625) : tensor<1x2048xf32> + %q8_input = check.generate.fill value(0) : tensor<2304xi8> + %route_ids = check.generate.iota offset(0) step(1) period(128) : tensor<8xi32> + %gate_weight = check.generate.iota offset(-72) step(1) period(144) : tensor<128x768x8x144xi8> + %up_weight = check.generate.iota offset(-71) step(1) period(144) : tensor<128x768x8x144xi8> + %expected_output = check.generate.fill value(0.0) : tensor<8x768xf32> + %expected_q8 = check.generate.fill value(0) : tensor<8x864xi8> + %actual_output0 = check.generate.fill value(1.0) : tensor<8x768xf32> + %actual_q8_0 = check.generate.fill value(1) : tensor<8x864xi8> + %actual_output1 = check.generate.fill value(2.0) : tensor<8x768xf32> + %actual_q8_1 = check.generate.fill value(2) : tensor<8x864xi8> + %completion_counters = check.generate.fill value(0) : tensor<48xi32> + %expected_counters = check.generate.fill value(0) : tensor<48xi32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<1x2048xf32>, tensor<2304xi8>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8[%token_count, %route_count, %route_stride, %expert_count, %output_size](%token_count, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %expected_output) : [index, index, index, index, index](index, index, index, index, index, tensor<2304xi8>, tensor<8xi32>, tensor<128x768x8x144xi8>, tensor<128x768x8x144xi8>, tensor<8x768xf32>) + kernel.launch @ggml_quantize_q8_1_x4_f32[%routed_row_count, %output_size](%routed_row_count, %output_size, %expected_output, %expected_q8) : [index, index](index, index, tensor<8x768xf32>, tensor<8x864xi8>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8[%token_count, %route_count, %route_stride, %expert_count, %output_size](%token_count, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %actual_output0, %completion_counters, %actual_q8_0) : [index, index, index, index, index](index, index, index, index, index, tensor<2304xi8>, tensor<8xi32>, tensor<128x768x8x144xi8>, tensor<128x768x8x144xi8>, tensor<8x768xf32>, tensor<48xi32>, tensor<8x864xi8>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8[%token_count, %route_count, %route_stride, %expert_count, %output_size](%token_count, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %actual_output1, %completion_counters, %actual_q8_1) : [index, index, index, index, index](index, index, index, index, index, tensor<2304xi8>, tensor<8xi32>, tensor<128x768x8x144xi8>, tensor<128x768x8x144xi8>, tensor<8x768xf32>, tensor<48xi32>, tensor<8x864xi8>) + check.expect.close actual(%actual_output0) expected(%expected_output) atol(16.0) rtol(9.9999999999999995e-07) nan(same) : tensor<8x768xf32> + check.expect.close actual(%actual_output1) expected(%expected_output) atol(16.0) rtol(9.9999999999999995e-07) nan(same) : tensor<8x768xf32> + check.expect.equal actual(%actual_q8_0) expected(%expected_q8) : tensor<8x864xi8> + check.expect.equal actual(%actual_q8_1) expected(%expected_q8) : tensor<8x864xi8> + check.expect.equal actual(%completion_counters) expected(%expected_counters) : tensor<48xi32> + check.return +} + +check.case public @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_differential_case { + %token_count = check.literal value(64) : index + %input_size = check.literal value(512) : index + %route_count = check.literal value(2) : index + %route_stride = check.literal value(4) : index + %expert_count = check.literal value(4) : index + %output_size = check.literal value(32) : index + %input = check.generate.fill value(0.00390625) : tensor<64x512xf32> + %q8_input = check.generate.fill value(0) : tensor<64x576xi8> + // The physical row retains four argsort entries while the logical top-k view + // selects experts 0 and 1. Both experts span two 32-row grouped tiles. + %route_ids = check.generate.iota offset(0) step(1) period(4) : tensor<64x4xi32> + %expert_table = check.generate.fill value(-1) : tensor<260xi32> + %gate_weight = check.generate.iota offset(-72) step(1) period(144) : tensor<4x32x2x144xi8> + %up_weight = check.generate.iota offset(-71) step(1) period(144) : tensor<4x32x2x144xi8> + %expected = check.generate.fill value(0.0) : tensor<64x2x32xf32> + %actual = check.generate.fill value(1.0) : tensor<64x2x32xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<64x512xf32>, tensor<64x576xi8>) + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<64x4xi32>, tensor<260xi32>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8[%token_count, %route_count, %route_stride, %expert_count, %output_size](%token_count, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %expected) : [index, index, index, index, index](index, index, index, index, index, tensor<64x576xi8>, tensor<64x4xi32>, tensor<4x32x2x144xi8>, tensor<4x32x2x144xi8>, tensor<64x2x32xf32>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped[%token_count, %route_count, %expert_count, %output_size](%token_count, %route_count, %expert_count, %output_size, %q8_input, %expert_table, %gate_weight, %up_weight, %actual) : [index, index, index, index](index, index, index, index, tensor<64x576xi8>, tensor<260xi32>, tensor<4x32x2x144xi8>, tensor<4x32x2x144xi8>, tensor<64x2x32xf32>) + check.expect.close actual(%actual) expected(%expected) atol(0.25) rtol(9.9999999999999995e-07) nan(same) : tensor<64x2x32xf32> + check.return +} + +check.case public @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_tail_case { + %token_count = check.literal value(37) : index + %input_size = check.literal value(512) : index + %route_count = check.literal value(3) : index + %route_stride = check.literal value(5) : index + %expert_count = check.literal value(4) : index + %output_size = check.literal value(33) : index + %input = check.generate.fill value(0.00390625) : tensor<37x512xf32> + %q8_input = check.generate.fill value(0) : tensor<37x576xi8> + // Three unique selected experts rotate through a five-entry physical row. + // This leaves the second route tile empty and the second channel tile with + // only one live output channel. + %route_ids = check.generate.iota offset(0) step(1) period(4) : tensor<37x5xi32> + %expert_table = check.generate.fill value(-1) : tensor<152xi32> + %gate_weight = check.generate.iota offset(-72) step(1) period(144) : tensor<4x33x2x144xi8> + %up_weight = check.generate.iota offset(-71) step(1) period(144) : tensor<4x33x2x144xi8> + %expected = check.generate.fill value(0.0) : tensor<37x3x33xf32> + %actual = check.generate.fill value(1.0) : tensor<37x3x33xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<37x512xf32>, tensor<37x576xi8>) + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<37x5xi32>, tensor<152xi32>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8[%token_count, %route_count, %route_stride, %expert_count, %output_size](%token_count, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %expected) : [index, index, index, index, index](index, index, index, index, index, tensor<37x576xi8>, tensor<37x5xi32>, tensor<4x33x2x144xi8>, tensor<4x33x2x144xi8>, tensor<37x3x33xf32>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped[%token_count, %route_count, %expert_count, %output_size](%token_count, %route_count, %expert_count, %output_size, %q8_input, %expert_table, %gate_weight, %up_weight, %actual) : [index, index, index, index](index, index, index, index, tensor<37x576xi8>, tensor<152xi32>, tensor<4x33x2x144xi8>, tensor<4x33x2x144xi8>, tensor<37x3x33xf32>) + check.expect.close actual(%actual) expected(%expected) atol(0.25) rtol(9.9999999999999995e-07) nan(same) : tensor<37x3x33xf32> + check.return +} + +// Forces one route partition to process a second 128-row band. This models a +// maximally skewed router while retaining a noncompact physical route stride. +check.case public @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_partition_stride_case { + %token_count = check.literal value(129) : index + %input_size = check.literal value(512) : index + %route_count = check.literal value(1) : index + %route_stride = check.literal value(3) : index + %expert_count = check.literal value(2) : index + %output_size = check.literal value(1) : index + %input = check.generate.fill value(0.00390625) : tensor<129x512xf32> + %q8_input = check.generate.fill value(0) : tensor<129x576xi8> + %route_ids = check.generate.fill value(0) : tensor<129x3xi32> + %expert_table = check.generate.fill value(-1) : tensor<260xi32> + %gate_weight = check.generate.fill value(85) : tensor<2x1x2x144xi8> + %up_weight = check.generate.fill value(-86) : tensor<2x1x2x144xi8> + %expected = check.generate.fill value(0.0) : tensor<129x1x1xf32> + %actual = check.generate.fill value(1.0) : tensor<129x1x1xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<129x512xf32>, tensor<129x576xi8>) + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<129x3xi32>, tensor<260xi32>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8[%token_count, %route_count, %route_stride, %expert_count, %output_size](%token_count, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %expected) : [index, index, index, index, index](index, index, index, index, index, tensor<129x576xi8>, tensor<129x3xi32>, tensor<2x1x2x144xi8>, tensor<2x1x2x144xi8>, tensor<129x1x1xf32>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped[%token_count, %route_count, %expert_count, %output_size](%token_count, %route_count, %expert_count, %output_size, %q8_input, %expert_table, %gate_weight, %up_weight, %actual) : [index, index, index, index](index, index, index, index, tensor<129x576xi8>, tensor<260xi32>, tensor<2x1x2x144xi8>, tensor<2x1x2x144xi8>, tensor<129x1x1xf32>) + check.expect.close actual(%actual) expected(%expected) atol(0.25) rtol(9.9999999999999995e-07) nan(same) : tensor<129x1x1xf32> + check.return +} + +check.case public @qwen3_moe_routed_gate_up_swiglu_q4k_q8_benchmark_case { + %token_count = check.param.choice values([1, 2, 4, 8, 16, 17, 32, 63, 128, 129, 512]) name("token_count") : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(768) : index + %q8_input = check.generate.fill value(0) : tensor<[%token_count]x2304xi8> + %route_ids = check.generate.iota offset(0) step(1) period(127) : tensor<[%token_count]x128xi32> + %gate_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %up_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x8x768xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x8x768xf32> + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8[%token_count, %route_count, %route_stride, %expert_count, %output_size](%token_count, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %output) : [index, index, index, index, index](index, index, index, index, index, tensor<[%token_count]x2304xi8>, tensor<[%token_count]x128xi32>, tensor<128x768x8x144xi8>, tensor<128x768x8x144xi8>, tensor<[%token_count]x8x768xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x8x768xf32> + check.return +} + +check.case public @qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8_benchmark_case { + %token_count = check.literal value(1) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(768) : index + %q8_input = check.generate.fill value(0) : tensor<2304xi8> + %route_ids = check.generate.iota offset(0) step(1) period(128) : tensor<8xi32> + %gate_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %up_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<8x768xf32> + %completion_counters = check.generate.fill value(0) : tensor<48xi32> + %next_q8_output = check.generate.fill value(1) : tensor<8x864xi8> + %expected_output = check.generate.fill value(0.0) : tensor<8x768xf32> + %expected_q8 = check.generate.fill value(0) : tensor<8x864xi8> + %expected_counters = check.generate.fill value(0) : tensor<48xi32> + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8[%token_count, %route_count, %route_stride, %expert_count, %output_size](%token_count, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %output, %completion_counters, %next_q8_output) : [index, index, index, index, index](index, index, index, index, index, tensor<2304xi8>, tensor<8xi32>, tensor<128x768x8x144xi8>, tensor<128x768x8x144xi8>, tensor<8x768xf32>, tensor<48xi32>, tensor<8x864xi8>) + check.expect.close actual(%output) expected(%expected_output) atol(0.0) rtol(0.0) nan(same) : tensor<8x768xf32> + check.expect.equal actual(%next_q8_output) expected(%expected_q8) : tensor<8x864xi8> + check.expect.equal actual(%completion_counters) expected(%expected_counters) : tensor<48xi32> + check.return +} + +check.case public @qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8_composed_benchmark_case { + %token_count = check.literal value(1) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(768) : index + %routed_row_count = check.literal value(8) : index + %q8_input = check.generate.fill value(0) : tensor<2304xi8> + %route_ids = check.generate.iota offset(0) step(1) period(128) : tensor<8xi32> + %gate_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %up_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<8x768xf32> + %next_q8_output = check.generate.fill value(1) : tensor<8x864xi8> + %expected_output = check.generate.fill value(0.0) : tensor<8x768xf32> + %expected_q8 = check.generate.fill value(0) : tensor<8x864xi8> + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8[%token_count, %route_count, %route_stride, %expert_count, %output_size](%token_count, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %output) : [index, index, index, index, index](index, index, index, index, index, tensor<2304xi8>, tensor<8xi32>, tensor<128x768x8x144xi8>, tensor<128x768x8x144xi8>, tensor<8x768xf32>) + kernel.launch @ggml_quantize_q8_1_x4_f32[%routed_row_count, %output_size](%routed_row_count, %output_size, %output, %next_q8_output) : [index, index](index, index, tensor<8x768xf32>, tensor<8x864xi8>) + check.expect.close actual(%output) expected(%expected_output) atol(0.0) rtol(0.0) nan(same) : tensor<8x768xf32> + check.expect.equal actual(%next_q8_output) expected(%expected_q8) : tensor<8x864xi8> + check.return +} + +check.case public @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_benchmark_case { + %token_count = check.param.choice values([1, 2, 4, 8, 16, 17, 32, 63, 128, 129, 512]) name("token_count") : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(768) : index + %q8_input = check.generate.fill value(0) : tensor<[%token_count]x2304xi8> + // A period of 127 rotates each physical 128-entry row by one expert while + // keeping the first eight logical routes distinct within every token. + %route_ids = check.generate.iota offset(0) step(1) period(127) : tensor<[%token_count]x128xi32> + // One packed transient buffer holds 128 counts followed by room for 512 + // assignments per expert. The production allocation uses the exact token + // count; this fixed test capacity permits one parameterized benchmark case. + %expert_table = check.generate.fill value(-1) : tensor<65664xi32> + %gate_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %up_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x8x768xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x8x768xf32> + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<[%token_count]x128xi32>, tensor<65664xi32>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped[%token_count, %route_count, %expert_count, %output_size](%token_count, %route_count, %expert_count, %output_size, %q8_input, %expert_table, %gate_weight, %up_weight, %output) : [index, index, index, index](index, index, index, index, tensor<[%token_count]x2304xi8>, tensor<65664xi32>, tensor<128x768x8x144xi8>, tensor<128x768x8x144xi8>, tensor<[%token_count]x8x768xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x8x768xf32> + check.return +} + +// Maximally diverse decode-batch control. Compact route storage makes the +// flattened 0..127 iota assign every M=16 route to a different expert. +check.case public @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_diverse_benchmark_case { + %token_count = check.param.choice values([1, 2, 4, 8, 16]) name("token_count") : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(768) : index + %q8_input = check.generate.fill value(0) : tensor<[%token_count]x2304xi8> + %route_ids = check.generate.iota offset(0) step(1) period(128) : tensor<[%token_count]x8xi32> + %expert_table = check.generate.fill value(-1) : tensor<65664xi32> + %gate_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %up_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x8x768xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x8x768xf32> + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<[%token_count]x8xi32>, tensor<65664xi32>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped[%token_count, %route_count, %expert_count, %output_size](%token_count, %route_count, %expert_count, %output_size, %q8_input, %expert_table, %gate_weight, %up_weight, %output) : [index, index, index, index](index, index, index, index, tensor<[%token_count]x2304xi8>, tensor<65664xi32>, tensor<128x768x8x144xi8>, tensor<128x768x8x144xi8>, tensor<[%token_count]x8x768xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x8x768xf32> + check.return +} + +check.case public @qwen3_moe_routed_gate_up_swiglu_q4k_q8_pipeline_benchmark_case { + %token_count = check.param.choice values([1, 2, 4, 8, 16, 32, 128, 512]) name("token_count") : index + %input_size = check.literal value(2048) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(768) : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %q8_input = check.generate.fill value(1) : tensor<[%token_count]x2304xi8> + %route_ids = check.generate.iota offset(0) step(1) period(127) : tensor<[%token_count]x128xi32> + %gate_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %up_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x8x768xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x8x768xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<[%token_count]x2048xf32>, tensor<[%token_count]x2304xi8>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8[%token_count, %route_count, %route_stride, %expert_count, %output_size](%token_count, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %output) : [index, index, index, index, index](index, index, index, index, index, tensor<[%token_count]x2304xi8>, tensor<[%token_count]x128xi32>, tensor<128x768x8x144xi8>, tensor<128x768x8x144xi8>, tensor<[%token_count]x8x768xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x8x768xf32> + check.return +} + +check.case public @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_pipeline_benchmark_case { + %token_count = check.param.choice values([1, 2, 4, 8, 16, 17, 32, 63, 128, 129, 512]) name("token_count") : index + %input_size = check.literal value(2048) : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(128) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(768) : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %q8_input = check.generate.fill value(1) : tensor<[%token_count]x2304xi8> + %route_ids = check.generate.iota offset(0) step(1) period(127) : tensor<[%token_count]x128xi32> + %expert_table = check.generate.fill value(-1) : tensor<65664xi32> + %gate_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %up_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x8x768xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x8x768xf32> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<[%token_count]x2048xf32>, tensor<[%token_count]x2304xi8>) + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<[%token_count]x128xi32>, tensor<65664xi32>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped[%token_count, %route_count, %expert_count, %output_size](%token_count, %route_count, %expert_count, %output_size, %q8_input, %expert_table, %gate_weight, %up_weight, %output) : [index, index, index, index](index, index, index, index, tensor<[%token_count]x2304xi8>, tensor<65664xi32>, tensor<128x768x8x144xi8>, tensor<128x768x8x144xi8>, tensor<[%token_count]x8x768xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x8x768xf32> + check.return +} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_nonzero_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_small + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_decode {token_count = 1} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_small_batch_2 {token_count = 2} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8_decode + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8_composed_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_1_x4_next_q8_composed_decode + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_small_batch_4 {token_count = 4} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_small_batch_8 {token_count = 8} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_small_batch_16 {token_count = 16} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_prefill_17 {token_count = 17} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_prefill_63 {token_count = 63} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_prefill_129 {token_count = 129} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_prefill_17 {token_count = 17} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_decode {token_count = 1} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_small_batch_2 {token_count = 2} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_small_batch_4 {token_count = 4} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_small_batch_8 {token_count = 8} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_small_batch_16 {token_count = 16} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_diverse_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_diverse_decode {token_count = 1} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_diverse_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_diverse_small_batch_2 {token_count = 2} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_diverse_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_diverse_small_batch_4 {token_count = 4} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_diverse_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_diverse_small_batch_8 {token_count = 8} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_diverse_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_diverse_small_batch_16 {token_count = 16} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_prefill_63 {token_count = 63} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_prefill_129 {token_count = 129} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_pipeline_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_pipeline_decode {token_count = 1} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_pipeline_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_pipeline_small_batch_8 {token_count = 8} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_pipeline_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_pipeline_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_pipeline_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_pipeline_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_pipeline_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_pipeline_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_pipeline_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_pipeline_prefill_17 {token_count = 17} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_pipeline_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_pipeline_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_pipeline_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_pipeline_prefill_63 {token_count = 63} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_pipeline_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_pipeline_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_pipeline_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_pipeline_prefill_129 {token_count = 129} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_pipeline_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_q8_grouped_pipeline_prefill_512 {token_count = 512} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_linear_q4k_f16_wmma.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_linear_q4k_f16_wmma.loom new file mode 100644 index 000000000000..45d2192f3bef --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/routed_linear_q4k_f16_wmma.loom @@ -0,0 +1,949 @@ +// Gfx11 routed-expert projection matching llama.cpp Vulkan's matrix path. +// +// Each two-wave workgroup computes 64 output channels for 32 routed rows of one +// expert. The waves share padded FP16 operand stages, each owns 32 output +// channels and four WMMA accumulators, and scatters one transposed accumulator +// fragment at a time through a compact route map. Gate and up invoke this same +// projection independently before the separate SwiGLU epilogue. +// +// The raw weights remain [expert][output channel][K / 256][144 bytes]. +// No persistent repacking or expanded-weight allocation is required. +amdgpu.target @qwen3_moe_gfx11_wave64 {subgroup_size = 64} + +amdgpu.target @qwen3_moe_gfx11_wave32 {subgroup_size = 32} + +config.decl @qwen3_moe.routed_gate_up.input_size : %value: index where [range(%value, 512, 32768), mul(%value, 512)] + +// Top-k is a model hyperparameter and is specialized with the kernel. Keeping +// it out of the dynamic workload lets address arithmetic and assignment decode +// fold to the exact model contract. +config.decl @qwen3_moe.routed_gate_up.route_count : %value: index where [range(%value, 1, 8)] + +config.decl @qwen3_moe.routed_gate_up.expert_count : %value: index where [range(%value, 1, 512)] + +config.decl @qwen3_moe.routed_gate_up.output_size : %value: index where [range(%value, 1, 4096)] + +kernel.decl @ggml_quantize_q8_1_x4_f32(%token_count$4: index, %input_size$5: index) launch(%token_count$6: index, %input_size$7: index, %input: buffer, %output: buffer) + +kernel.decl @qwen3_moe_build_expert_table(%token_count$10: index, %route_count$11: index, %route_stride$12: index, %expert_count$13: index) launch(%token_count$14: index, %route_count$15: index, %route_stride$16: index, %expert_count$17: index, %route_ids: buffer, %expert_table: buffer) + +kernel.decl @qwen3_moe_build_expert_partition_table(%token_count$20: index, %route_count$21: index, %expert_count$22: index) launch(%token_count$23: index, %route_count$24: index, %expert_count$25: index, %expert_table: buffer, %partition_table: buffer) + +func.decl @qwen3_moe_unpack_expert_partition_descriptor(%descriptor: i32) -> (index, index, index) + +kernel.decl @qwen3_moe_routed_gate_up_swiglu_q4k_q8(%token_count$32: index, %route_count$33: index, %route_stride$34: index, %expert_count$35: index, %output_size$36: index) launch(%token_count$37: index, %route_count$38: index, %route_stride$39: index, %expert_count$40: index, %output_size$41: index, %q8_input: buffer, %route_ids: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) + +func.decl @qwen3_moe_q4k_scale_from_header(%scale0: i32, %scale1: i32, %scale2: i32, %q4_group: index) -> (i32, i32) + +// Acquires the packed code word shared by one adjacent Q4_K group pair. +func.def inline @qwen3_moe_q4k_wmma_load_code(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group_pair: index, %packet: index) -> (vector<1xi32>) { + %c8 = index.constant 8 : index + %block_bytes = index.constant 144 : offset + %code_offset = index.constant 16 : offset + %bounded_group_pair = index.assume %q4_group_pair [range(%q4_group_pair, 0, 3)] : index + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %q4_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<32xi32> + %q_page = index.mul %bounded_group_pair, %c8 : index + %q_word_index0 = index.add %q_page, %bounded_packet : index + %q_word_index = index.assume %q_word_index0 [range(%q_word_index0, 0, 31)] : index + %q_word = vector.load %code_view[%q_word_index] : view<32xi32> -> vector<1xi32> + func.return %q_word : vector<1xi32> +} + +// Decodes the four adjacent Q4_K values owned by one load packet from an +// already-loaded block header and packed code word. Matrix schedules choose +// the lifetime of both immutable packets. +func.def inline @qwen3_moe_q4k_wmma_vector4_from_header_code(%q4_group: index, %header_words: vector<4xi32>, %q_word: vector<1xi32>) -> (vector<4xf16>) { + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %q4_mask = vector.constant 252645135 : vector<1xi32> + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + %header_halves = vector.bitcast %header_words : vector<4xi32> to vector<8xf16> + %d_f16 = vector.extract %header_halves[0] : vector<8xf16> -> f16 + %dmin_f16 = vector.extract %header_halves[1] : vector<8xf16> -> f16 + %scale0 = vector.extract %header_words[1] : vector<4xi32> -> i32 + %scale1 = vector.extract %header_words[2] : vector<4xi32> -> i32 + %scale2 = vector.extract %header_words[3] : vector<4xi32> -> i32 + %d = scalar.extf %d_f16 : f16 to f32 + %dmin = scalar.extf %dmin_f16 : f16 to f32 + %scale, %minimum = func.call @qwen3_moe_q4k_scale_from_header(%scale0, %scale1, %scale2, %bounded_group) : (i32, i32, i32, index) -> (i32, i32) + %scale_f32 = scalar.uitofp %scale : i32 to f32 + %minimum_f32 = scalar.uitofp %minimum : i32 to f32 + %d_scale = scalar.mulf %d, %scale_f32 : f32 + %minimum_scale = scalar.mulf %dmin, %minimum_f32 : f32 + %q_half = index.rem %bounded_group, %c2 : index + %q_shift_index = index.mul %q_half, %c4 : index + %q_shift_i32 = index.cast %q_shift_index : index to i32 + %q_shift = vector.splat %q_shift_i32 : vector<1xi32> + %shifted_q = vector.shrui %q_word, %q_shift : vector<1xi32> + %masked_q = vector.andi %shifted_q, %q4_mask : vector<1xi32> + %q_i8 = vector.bitcast %masked_q : vector<1xi32> to vector<4xi8> + %q_f32 = vector.uitofp %q_i8 : vector<4xi8> to vector<4xf32> + // Form adjacent FP16 lanes from fused FP32 affine expressions. AMDGPU maps + // this natural shape to packed mixlo/mixhi instructions where available. + %negative_minimum_scale = scalar.negf %minimum_scale : f32 + %q0 = vector.extract %q_f32[0] : vector<4xf32> -> f32 + %q1 = vector.extract %q_f32[1] : vector<4xf32> -> f32 + %q2 = vector.extract %q_f32[2] : vector<4xf32> -> f32 + %q3 = vector.extract %q_f32[3] : vector<4xf32> -> f32 + %value0 = scalar.fmaf %q0, %d_scale, %negative_minimum_scale : f32 + %value1 = scalar.fmaf %q1, %d_scale, %negative_minimum_scale : f32 + %value2 = scalar.fmaf %q2, %d_scale, %negative_minimum_scale : f32 + %value3 = scalar.fmaf %q3, %d_scale, %negative_minimum_scale : f32 + %half0 = scalar.fptrunc %value0 : f32 to f16 + %half1 = scalar.fptrunc %value1 : f32 to f16 + %half2 = scalar.fptrunc %value2 : f32 to f16 + %half3 = scalar.fptrunc %value3 : f32 to f16 + %result = vector.from_elements %half0, %half1, %half2, %half3 : vector<4xf16> + func.return %result : vector<4xf16> +} + +// Decodes one group when its caller has retained only the Q4_K block header. +func.def inline @qwen3_moe_q4k_wmma_vector4_from_header(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group: index, %packet: index, %header_words: vector<4xi32>) -> (vector<4xf16>) { + %c2 = index.constant 2 : index + %bounded_group = index.assume %q4_group [range(%q4_group, 0, 7)] : index + %q4_group_pair = index.div %bounded_group, %c2 : index + %q_word = func.call @qwen3_moe_q4k_wmma_load_code(%weight, %row_byte_base, %q4_block, %q4_group_pair, %packet) : (buffer, offset, index, index, index) -> (vector<1xi32>) + %values = func.call @qwen3_moe_q4k_wmma_vector4_from_header_code(%bounded_group, %header_words, %q_word) : (index, vector<4xi32>, vector<1xi32>) -> (vector<4xf16>) + func.return %values : vector<4xf16> +} + +// Acquires one naturally aligned Q4_K block header as a single 16-byte packet. +func.def inline @qwen3_moe_q4k_wmma_load_header(%weight: buffer, %row_byte_base: offset, %q4_block: index) -> (vector<4xi32>) { + %c0 = index.constant 0 : index + %block_bytes = index.constant 144 : offset + %block_byte_add = index.scale %q4_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %header_view = buffer.view %weight[%block_byte_base] : buffer -> view<4xi32> + %header_words = vector.load %header_view[%c0] : view<4xi32> -> vector<4xi32> + func.return %header_words : vector<4xi32> +} + +// Acquires one block header before decoding the selected four-value group +// packet. Matrix schedules that span several groups call the two operations +// separately so the header lifetime matches their complete block loop. +func.def inline @qwen3_moe_q4k_wmma_vector4(%weight: buffer, %row_byte_base: offset, %q4_block: index, %q4_group: index, %packet: index) -> (vector<4xf16>) { + %header_words = func.call @qwen3_moe_q4k_wmma_load_header(%weight, %row_byte_base, %q4_block) : (buffer, offset, index) -> (vector<4xi32>) + %values = func.call @qwen3_moe_q4k_wmma_vector4_from_header(%weight, %row_byte_base, %q4_block, %q4_group, %packet, %header_words) : (buffer, offset, index, index, index, vector<4xi32>) -> (vector<4xf16>) + func.return %values : vector<4xf16> +} + +kernel.def target(@qwen3_moe_gfx11_wave64) @qwen3_moe_routed_linear_q4k_f16_wmma(%token_count: index) { + %expert_count = config.get @qwen3_moe.routed_gate_up.expert_count : index + %output_size = config.get @qwen3_moe.routed_gate_up.output_size : index + %c1 = index.constant 1 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c128 = index.constant 128 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + %padded_token_count = index.add %token_count, %c63 : index + %route_tiles = index.div %padded_token_count, %c64 : index + kernel.launch.config workgroups(%output_tiles, %route_tiles, %expert_count) workgroup_size(%c128, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %expert_table: buffer, %weight: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %input_size = config.get @qwen3_moe.routed_gate_up.input_size : index + %route_count = config.get @qwen3_moe.routed_gate_up.route_count : index + %expert_count = config.get @qwen3_moe.routed_gate_up.expert_count : index + %output_size = config.get @qwen3_moe.routed_gate_up.output_size : index + %bounded_route_count = index.assume %route_count [range(%route_count, 1, 8)] : index + %bounded_expert_count = index.assume %expert_count [range(%expert_count, 1, 128)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 4096)] : index + %channel_tile = kernel.workgroup.id : index + %route_tile = kernel.workgroup.id : index + %expert = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 1)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c40 = index.constant 40 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %c4_bytes = index.constant 4 : offset + %q4_block_bytes = index.constant 144 : offset + %weight_stage_bytes = index.constant 5120 : offset + %activation_stage_bytes = index.constant 2560 : offset + %route_stage_bytes = index.constant 128 : offset + %wave_result_stage_bytes = index.constant 512 : offset + %result_stage_bytes = index.constant 1024 : offset + %c0_i32 = scalar.constant 0 : i32 + %cn1_i32 = scalar.constant -1 : i32 + %c0_f16x4 = vector.constant 0.0 : vector<4xf16> + %zero_accumulator = vector.constant 0.0 : vector<8xf16> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %assignment_count = index.mul %token_count, %bounded_route_count : index + %assignment_table_byte_base = index.scale %bounded_expert_count, %c4_bytes : index, offset -> offset + %q4_block_count = index.div %input_size, %c256 : index + %weight_row_bytes = index.scale %q4_block_count, %q4_block_bytes : index, offset -> offset + %weight_expert_bytes = index.scale %bounded_output_size, %weight_row_bytes : index, offset -> offset + %output_row_count = index.mul %token_count, %bounded_route_count : index + %input_noalias, %expert_table_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input, %expert_table, %weight, %output : buffer, buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%token_count]x[%input_size]xf32> + %count_view = buffer.view %expert_table_noalias[%c0_offset] : buffer -> view<[%bounded_expert_count]xi32> + %assignment_view = buffer.view %expert_table_noalias[%assignment_table_byte_base] : buffer -> view<[%bounded_expert_count]x[%token_count]xi32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%output_row_count]x[%bounded_output_size]xf32> + %weight_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %activation_stage = buffer.alloca align(16) %activation_stage_bytes : buffer + %route_stage = buffer.alloca align(16) %route_stage_bytes : buffer + %result_stage = buffer.alloca align(16) %result_stage_bytes : buffer + %weight_stage_view = buffer.view %weight_stage[%c0_offset] : buffer -> view<64x40xf16> + %activation_stage_physical_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x40xf16> + %activation_fragment_layout = encoding.layout.strided [1, %c40] : encoding + %activation_fragment_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x32xf16, %activation_fragment_layout> + %route_stage_view = buffer.view %route_stage[%c0_offset] : buffer -> view<32xi32> + %wave_result_stage_offset = index.scale %subgroup, %wave_result_stage_bytes : index, offset -> offset + %result_fragment_layout = encoding.layout.strided [1, %c16] : encoding + %result_fragment_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf16, %result_fragment_layout> + %result_physical_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf16> + %channel_tile_base = index.mul %channel_tile, %c64 : index + %initial_route_tile_base = index.mul %route_tile, %c32 : index + %padded_token_count = index.add %token_count, %c63 : index + %route_partition_count = index.div %padded_token_count, %c64 : index + %route_partition_step = index.mul %route_partition_count, %c32 : index + %bounded_expert, %table_expert_count = index.assume %expert, %bounded_expert_count [lt(%expert, %bounded_expert_count)] : index, index + %is_workitem_zero = index.cmp eq, %workitem, %c0 : index + %lane_expert_route_count = scf.if %is_workitem_zero -> (i32) { + %loaded = view.load %count_view[%bounded_expert] : view<[%bounded_expert_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %c0_i32 : i32 + } + %expert_route_count_reduced = kernel.workgroup.reduce %lane_expert_route_count : i32 + %expert_route_count_i32 = kernel.subgroup.broadcast.first %expert_route_count_reduced : i32 + %expert_route_count0 = index.cast %expert_route_count_i32 : i32 to index + %expert_route_count = index.assume %expert_route_count0 [range(%expert_route_count0, 0, 2048)] : index + // Route partitions are distributed across the launch grid and continue in + // uniform strides for concentrated routing. Balanced Qwen prefill gives each + // expert 32, 64, or 128 rows at 512, 1024, or 2048 tokens. Concentrated + // experts can consume the full token count without changing the launch + // geometry. + scf.for %route_tile_base = [%initial_route_tile_base to %expert_route_count step %route_partition_step] { + // The first wave snapshots the compact route map once. Every K tile then + // reuses these 32 entries while the full workgroup cooperatively fills the + // operand stages. + %loads_route = index.cmp ult, %workitem, %c32 : index + scf.if %loads_route { + %local_route = index.assume %workitem [range(%workitem, 0, 31)] : index + %assignment_ordinal = index.add %route_tile_base, %local_route : index + %valid_row = index.cmp ult, %assignment_ordinal, %expert_route_count : index + %assignment_i32 = scf.if %valid_row -> (i32) { + %bounded_assignment_ordinal, %table_token_count = index.assume %assignment_ordinal, %token_count [lt(%assignment_ordinal, %token_count)] : index, index + %loaded = view.load %assignment_view[%bounded_expert, %bounded_assignment_ordinal] : view<[%bounded_expert_count]x[%token_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %cn1_i32 : i32 + } + view.store %assignment_i32, %route_stage_view[%local_route] : i32, view<32xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %load_packet = index.rem %workitem, %c8 : index + %load_k = index.mul %load_packet, %c4 : index + %load_row0 = index.div %workitem, %c8 : index + %load_row = index.assume %load_row0 [range(%load_row0, 0, 15)] : index + %expert_byte_base = index.scale %bounded_expert, %weight_expert_bytes : index, offset -> offset + %subgroup_channel_add = index.mul %subgroup, %c32 : index + %subgroup_channel1 = index.add %subgroup_channel_add, %c16 : index + %init00 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf16> + %init01 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf16> + %init10 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf16> + %init11 = vector.fragment %zero_accumulator shape [%m, %n] : vector<8xf16> + %result00, %result01, %result10, %result11 = scf.for %q4_block = [%c0 to %q4_block_count step %c1](%block_acc00 = %init00 : vector<8xf16>, %block_acc01 = %init01 : vector<8xf16>, %block_acc10 = %init10 : vector<8xf16>, %block_acc11 = %init11 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16>) { + %block_result00, %block_result01, %block_result10, %block_result11 = scf.for %q4_group = [%c0 to %c8 step %c1](%acc00 = %block_acc00 : vector<8xf16>, %acc01 = %block_acc01 : vector<8xf16>, %acc10 = %block_acc10 : vector<8xf16>, %acc11 = %block_acc11 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16>) { + %block_k_base = index.mul %q4_block, %c256 : index + %group_k_add = index.mul %q4_group, %c32 : index + %k_origin = index.add %block_k_base, %group_k_add : index + scf.for %row_offset = [%c0 to %c64 step %c16] unroll { + %local_row0 = index.add %load_row, %row_offset : index + %local_row = index.assume %local_row0 [range(%local_row0, 0, 63)] : index + %channel = index.add %channel_tile_base, %local_row : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %weight_values = scf.if %valid_channel -> (vector<4xf16>) { + %channel_byte_add = index.scale %channel, %weight_row_bytes : index, offset -> offset + %row_byte_base = index.add %expert_byte_base, %channel_byte_add : offset + %decoded = func.call @qwen3_moe_q4k_wmma_vector4(%weight_noalias, %row_byte_base, %q4_block, %q4_group, %load_packet) : (buffer, offset, index, index, index) -> (vector<4xf16>) + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + %is_activation_row = index.cmp ult, %local_row, %c32 : index + vector.store %weight_values, %weight_stage_view[%local_row, %load_k] : vector<4xf16>, view<64x40xf16> + scf.if %is_activation_row { + %activation_row = index.assume %local_row [range(%local_row, 0, 31)] : index + %assignment_i32 = view.load %route_stage_view[%activation_row] : view<32xi32> -> i32 + %valid_assignment = scalar.cmpi sge, %assignment_i32, %c0_i32 : i32 + %activation_values = scf.if %valid_assignment -> (vector<4xf16>) { + %assignment0 = index.cast %assignment_i32 : i32 to index + %assignment = index.assume %assignment0 [range(%assignment0, 0, 16383)] : index + %bounded_assignment, %bounded_assignment_count = index.assume %assignment, %assignment_count [lt(%assignment, %assignment_count)] : index, index + %token0 = index.div %bounded_assignment, %bounded_route_count : index + %token, %input_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %input_k = index.add %k_origin, %load_k : index + %loaded = vector.load %input_view[%token, %input_k] : view<[%token_count]x[%input_size]xf32> -> vector<4xf32> + %converted = vector.fptrunc %loaded : vector<4xf32> to vector<4xf16> + scf.yield %converted : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %activation_values, %activation_stage_physical_view[%activation_row, %load_k] : vector<4xf16>, view<32x40xf16> + } + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %next00, %next01, %next10, %next11 = scf.for %k_half = [%c0 to %c32 step %c16](%half_acc00 = %acc00 : vector<8xf16>, %half_acc01 = %acc01 : vector<8xf16>, %half_acc10 = %acc10 : vector<8xf16>, %half_acc11 = %acc11 : vector<8xf16>) -> (vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16>) unroll { + %lhs0 = vector.fragment.load %weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %lhs1 = vector.fragment.load %weight_stage_view[%subgroup_channel1, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %rhs0 = vector.fragment.load %activation_fragment_view[%k_half, %c0] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %rhs1 = vector.fragment.load %activation_fragment_view[%k_half, %c16] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %half_next00 = vector.mma %lhs0, %rhs0, %half_acc00 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %half_next01 = vector.mma %lhs0, %rhs1, %half_acc01 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %half_next10 = vector.mma %lhs1, %rhs0, %half_acc10 : vector<16xf16>, vector<16xf16>, vector<8xf16> + %half_next11 = vector.mma %lhs1, %rhs1, %half_acc11 : vector<16xf16>, vector<16xf16>, vector<8xf16> + scf.yield %half_next00, %half_next01, %half_next10, %half_next11 : vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %next00, %next01, %next10, %next11 : vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16> + } + scf.yield %block_result00, %block_result01, %block_result10, %block_result11 : vector<8xf16>, vector<8xf16>, vector<8xf16>, vector<8xf16> + } + // WMMA produces [channel][route] fragments. Store each through a + // transposed, wave-private LDS slice so every lane can scatter four + // contiguous channels to one routed output row. This uses the full wave + // for conversion and publication instead of serializing 16 channels onto + // each of 16 lanes. Subgroup-scoped fences suffice because waves never + // access each other's result slice. + %publish_route0 = index.div %lane, %c4 : index + %publish_route = index.assume %publish_route0 [range(%publish_route0, 0, 15)] : index + %publish_packet0 = index.rem %lane, %c4 : index + %publish_packet = index.assume %publish_packet0 [range(%publish_packet0, 0, 3)] : index + %publish_channel_add = index.mul %publish_packet, %c4 : index + %local_route1 = index.add %c16, %publish_route : index + %assignment0_i32 = view.load %route_stage_view[%publish_route] : view<32xi32> -> i32 + %assignment1_i32 = view.load %route_stage_view[%local_route1] : view<32xi32> -> i32 + %assignment0_nonnegative = scalar.cmpi sge, %assignment0_i32, %c0_i32 : i32 + %assignment1_nonnegative = scalar.cmpi sge, %assignment1_i32, %c0_i32 : i32 + %safe_assignment0_i32 = scf.if %assignment0_nonnegative -> (i32) { + scf.yield %assignment0_i32 : i32 + } else { + scf.yield %c0_i32 : i32 + } + %safe_assignment1_i32 = scf.if %assignment1_nonnegative -> (i32) { + scf.yield %assignment1_i32 : i32 + } else { + scf.yield %c0_i32 : i32 + } + %safe_assignment0_0 = index.cast %safe_assignment0_i32 : i32 to index + %safe_assignment1_0 = index.cast %safe_assignment1_i32 : i32 to index + %safe_assignment0 = index.assume %safe_assignment0_0 [range(%safe_assignment0_0, 0, 16383)] : index + %safe_assignment1 = index.assume %safe_assignment1_0 [range(%safe_assignment1_0, 0, 16383)] : index + %bounded_assignment0, %bounded_assignment_count0 = index.assume %safe_assignment0, %assignment_count [lt(%safe_assignment0, %assignment_count)] : index, index + %bounded_assignment1, %bounded_assignment_count1 = index.assume %safe_assignment1, %assignment_count [lt(%safe_assignment1, %assignment_count)] : index, index + %subgroup_channel_base = index.add %channel_tile_base, %subgroup_channel_add : index + %channel0 = index.add %subgroup_channel_base, %publish_channel_add : index + %channel1_base = index.add %subgroup_channel_base, %c16 : index + %channel1 = index.add %channel1_base, %publish_channel_add : index + %valid_channel0 = index.cmp ult, %channel0, %bounded_output_size : index + %valid_channel1 = index.cmp ult, %channel1, %bounded_output_size : index + %writes00 = scalar.andi %assignment0_nonnegative, %valid_channel0 : i1 + %writes01 = scalar.andi %assignment1_nonnegative, %valid_channel0 : i1 + %writes10 = scalar.andi %assignment0_nonnegative, %valid_channel1 : i1 + %writes11 = scalar.andi %assignment1_nonnegative, %valid_channel1 : i1 + vector.fragment.store %result00, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<8xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes00 { + %values = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf16> -> vector<4xf16> + %wide = vector.extf %values : vector<4xf16> to vector<4xf32> + %mask = vector.mask.range [%channel0 to %bounded_output_size step %c1] : index -> vector<4xi1> + vector.store.mask %wide, %output_view[%bounded_assignment0, %channel0], %mask : vector<4xf32>, view<[%output_row_count]x[%bounded_output_size]xf32>, vector<4xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result01, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<8xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes01 { + %values = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf16> -> vector<4xf16> + %wide = vector.extf %values : vector<4xf16> to vector<4xf32> + %mask = vector.mask.range [%channel0 to %bounded_output_size step %c1] : index -> vector<4xi1> + vector.store.mask %wide, %output_view[%bounded_assignment1, %channel0], %mask : vector<4xf32>, view<[%output_row_count]x[%bounded_output_size]xf32>, vector<4xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result10, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<8xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes10 { + %values = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf16> -> vector<4xf16> + %wide = vector.extf %values : vector<4xf16> to vector<4xf32> + %mask = vector.mask.range [%channel1 to %bounded_output_size step %c1] : index -> vector<4xi1> + vector.store.mask %wide, %output_view[%bounded_assignment0, %channel1], %mask : vector<4xf32>, view<[%output_row_count]x[%bounded_output_size]xf32>, vector<4xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %result11, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<8xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes11 { + %values = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf16> -> vector<4xf16> + %wide = vector.extf %values : vector<4xf16> to vector<4xf32> + %mask = vector.mask.range [%channel1 to %bounded_output_size step %c1] : index -> vector<4xi1> + vector.store.mask %wide, %output_view[%bounded_assignment1, %channel1], %mask : vector<4xf32>, view<[%output_row_count]x[%bounded_output_size]xf32>, vector<4xi1> + } + // Route and result stages are reused by the next concentrated-routing + // partition. All waves must finish publication before either stage changes. + kernel.barrier scope(workgroup) ordering(acq_rel) + } + kernel.return +} + +// Fuses gate and up projection through SwiGLU for one routed expert tile. +// +// Eight waves split the 64 output channels and 32 routed rows into 16x16 result +// tiles. Each wave carries one gate and one up fragment while the workgroup +// shares one routed activation tile across both contractions. The final +// fragments meet in a wave-private LDS slice and publish the activated product +// directly, avoiding both full-size projection intermediates. +// +// The SwiGLU product is rounded to FP16 at its sole publication point because +// the grouped-down contraction consumes that exact FP16 WMMA operand. Keeping +// the transient in FP16 avoids a widen-store-load-truncate round trip and +// halves its global-memory footprint. +kernel.def target(@qwen3_moe_gfx11_wave32) @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma(%token_count: index) { + %route_count = config.get @qwen3_moe.routed_gate_up.route_count : index + %expert_count = config.get @qwen3_moe.routed_gate_up.expert_count : index + %output_size = config.get @qwen3_moe.routed_gate_up.output_size : index + %c1 = index.constant 1 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c63 = index.constant 63 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %padded_output_size = index.add %output_size, %c63 : index + %output_tiles = index.div %padded_output_size, %c64 : index + // The exact partition count lives in device memory. The rounding bound keeps + // it below assignment_partitions + expert_count, so launching the larger + // term lets every workgroup consume at most two strided descriptors. + %assignment_count = index.mul %token_count, %route_count : index + %rounded_assignment_count = index.add %assignment_count, %c31 : index + %assignment_partition_count = index.div %rounded_assignment_count, %c32 : index + %has_more_assignment_partitions = index.cmp ugt, %assignment_partition_count, %expert_count : index + %launch_partition_count = scf.select %has_more_assignment_partitions, %assignment_partition_count, %expert_count : index + kernel.launch.config workgroups(%output_tiles, %launch_partition_count, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %expert_table: buffer, %partition_table: buffer, %gate_weight: buffer, %up_weight: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %input_size = config.get @qwen3_moe.routed_gate_up.input_size : index + %route_count = config.get @qwen3_moe.routed_gate_up.route_count : index + %expert_count = config.get @qwen3_moe.routed_gate_up.expert_count : index + %output_size = config.get @qwen3_moe.routed_gate_up.output_size : index + %bounded_route_count = index.assume %route_count [range(%route_count, 1, 8)] : index + %bounded_expert_count = index.assume %expert_count [range(%expert_count, 1, 128)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 4096)] : index + %channel_tile = kernel.workgroup.id : index + %partition_ordinal = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %subgroup0 = kernel.subgroup.id : index + %subgroup = index.assume %subgroup0 [range(%subgroup0, 0, 7)] : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c31 = index.constant 31 : index + %c32 = index.constant 32 : index + %c40 = index.constant 40 : index + %c64 = index.constant 64 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %c4_bytes = index.constant 4 : offset + %q4_block_bytes = index.constant 144 : offset + %weight_stage_bytes = index.constant 5120 : offset + %activation_stage_bytes = index.constant 2560 : offset + %route_stage_bytes = index.constant 128 : offset + %wave_result_stage_bytes = index.constant 512 : offset + %result_stage_bytes = index.constant 4096 : offset + %c0_i32 = scalar.constant 0 : i32 + %cn1_i32 = scalar.constant -1 : i32 + %c0_i32x1 = vector.constant 0 : vector<1xi32> + %c0_i32x4 = vector.constant 0 : vector<4xi32> + %c0_f16x4 = vector.constant 0.0 : vector<4xf16> + %zero_accumulator = vector.constant 0.0 : vector<16xf16> + %m = index.constant 16 : index + %n = index.constant 16 : index + %k = index.constant 16 : index + %assignment_count = index.mul %token_count, %bounded_route_count : index + %rounded_assignment_count = index.add %assignment_count, %c31 : index + %assignment_partition_count = index.div %rounded_assignment_count, %c32 : index + %maximum_partition_count = index.add %assignment_partition_count, %bounded_expert_count : index + %has_more_assignment_partitions = index.cmp ugt, %assignment_partition_count, %bounded_expert_count : index + %launch_partition_count = scf.select %has_more_assignment_partitions, %assignment_partition_count, %bounded_expert_count : index + %assignment_table_byte_base = index.scale %bounded_expert_count, %c4_bytes : index, offset -> offset + %q4_block_count = index.div %input_size, %c256 : index + %weight_row_bytes = index.scale %q4_block_count, %q4_block_bytes : index, offset -> offset + %weight_expert_bytes = index.scale %bounded_output_size, %weight_row_bytes : index, offset -> offset + %output_row_count = index.mul %token_count, %bounded_route_count : index + %input_noalias, %expert_table_noalias, %partition_table_noalias, %gate_weight_noalias, %up_weight_noalias, %output_noalias = buffer.assume.noalias %input, %expert_table, %partition_table, %gate_weight, %up_weight, %output : buffer, buffer, buffer, buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%token_count]x[%input_size]xf32> + %assignment_view = buffer.view %expert_table_noalias[%assignment_table_byte_base] : buffer -> view<[%bounded_expert_count]x[%token_count]xi32> + %partition_count_view = buffer.view %partition_table_noalias[%c0_offset] : buffer -> view<1xi32> + %partition_descriptor_view = buffer.view %partition_table_noalias[%c4_bytes] : buffer -> view<[%maximum_partition_count]xi32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%output_row_count]x[%bounded_output_size]xf16> + %weight_stage = buffer.alloca align(16) %weight_stage_bytes : buffer + %activation_stage = buffer.alloca align(16) %activation_stage_bytes : buffer + %route_stage = buffer.alloca align(16) %route_stage_bytes : buffer + %result_stage = buffer.alloca align(16) %result_stage_bytes : buffer + %weight_stage_view = buffer.view %weight_stage[%c0_offset] : buffer -> view<64x40xf16> + %activation_stage_physical_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x40xf16> + %activation_fragment_layout = encoding.layout.strided [1, %c40] : encoding + %activation_fragment_view = buffer.view %activation_stage[%c0_offset] : buffer -> view<32x32xf16, %activation_fragment_layout> + %route_stage_view = buffer.view %route_stage[%c0_offset] : buffer -> view<32xi32> + %wave_result_stage_offset = index.scale %subgroup, %wave_result_stage_bytes : index, offset -> offset + %result_fragment_layout = encoding.layout.strided [1, %c16] : encoding + %result_fragment_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf16, %result_fragment_layout> + %result_physical_view = buffer.view %result_stage[%wave_result_stage_offset] : buffer -> view<16x16xf16> + %channel_tile_base = index.mul %channel_tile, %c64 : index + %is_workitem_zero = index.cmp eq, %workitem, %c0 : index + %lane_partition_count_i32 = scf.if %is_workitem_zero -> (i32) { + %loaded = view.load %partition_count_view[%c0] : view<1xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %c0_i32 : i32 + } + %partition_count_reduced = kernel.workgroup.reduce %lane_partition_count_i32 : i32 + %partition_count_i32 = kernel.subgroup.broadcast.first %partition_count_reduced : i32 + %partition_count0 = index.cast %partition_count_i32 : i32 to index + %partition_count, %partition_capacity = index.assume %partition_count0, %maximum_partition_count [lt(%partition_count0, %maximum_partition_count)] : index, index + scf.for %active_partition = [%partition_ordinal to %partition_count step %launch_partition_count] { + %descriptor_ordinal, %descriptor_count = index.assume %active_partition, %partition_count [lt(%active_partition, %partition_count)] : index, index + %table_descriptor_ordinal, %table_descriptor_capacity = index.assume %descriptor_ordinal, %maximum_partition_count [lt(%descriptor_ordinal, %maximum_partition_count)] : index, index + %lane_descriptor_i32 = scf.if %is_workitem_zero -> (i32) { + %loaded = view.load %partition_descriptor_view[%table_descriptor_ordinal] : view<[%maximum_partition_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %c0_i32 : i32 + } + %descriptor_reduced_i32 = kernel.workgroup.reduce %lane_descriptor_i32 : i32 + %descriptor_i32 = kernel.subgroup.broadcast.first %descriptor_reduced_i32 : i32 + %expert, %route_tile_base, %partition_row_count = func.call @qwen3_moe_unpack_expert_partition_descriptor(%descriptor_i32) : (i32) -> (index, index, index) + %bounded_expert, %table_expert_count = index.assume %expert, %bounded_expert_count [lt(%expert, %bounded_expert_count)] : index, index + %loads_route = index.cmp ult, %workitem, %c32 : index + scf.if %loads_route { + %local_route = index.assume %workitem [range(%workitem, 0, 31)] : index + %assignment_ordinal = index.add %route_tile_base, %local_route : index + %valid_row = index.cmp ult, %local_route, %partition_row_count : index + %assignment_i32 = scf.if %valid_row -> (i32) { + %bounded_assignment_ordinal, %table_token_count = index.assume %assignment_ordinal, %token_count [lt(%assignment_ordinal, %token_count)] : index, index + %loaded = view.load %assignment_view[%bounded_expert, %bounded_assignment_ordinal] : view<[%bounded_expert_count]x[%token_count]xi32> -> i32 + scf.yield %loaded : i32 + } else { + scf.yield %cn1_i32 : i32 + } + view.store %assignment_i32, %route_stage_view[%local_route] : i32, view<32xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %load_packet = index.rem %workitem, %c8 : index + %load_k = index.mul %load_packet, %c4 : index + %load_row0 = index.div %workitem, %c8 : index + %load_row = index.assume %load_row0 [range(%load_row0, 0, 31)] : index + %expert_byte_base = index.scale %bounded_expert, %weight_expert_bytes : index, offset -> offset + %channel_subgroup = index.rem %subgroup, %c4 : index + %route_subgroup = index.div %subgroup, %c4 : index + %subgroup_channel_add = index.mul %channel_subgroup, %c16 : index + %subgroup_route_add = index.mul %route_subgroup, %c16 : index + %init_gate = vector.fragment %zero_accumulator shape [%m, %n] : vector<16xf16> + %init_up = vector.fragment %zero_accumulator shape [%m, %n] : vector<16xf16> + %gate_result, %up_result = scf.for %q4_block = [%c0 to %q4_block_count step %c1](%block_gate = %init_gate : vector<16xf16>, %block_up = %init_up : vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) { + // Two lanes cover the 64 output rows owned by one load packet. Retain + // both projection headers across the complete eight-group block. + %gate_header0, %gate_header1, %up_header0, %up_header1 = scf.for %header_row_offset = [%c0 to %c64 step %c32](%prior_gate_header0 = %c0_i32x4 : vector<4xi32>, %prior_gate_header1 = %c0_i32x4 : vector<4xi32>, %prior_up_header0 = %c0_i32x4 : vector<4xi32>, %prior_up_header1 = %c0_i32x4 : vector<4xi32>) -> (vector<4xi32>, vector<4xi32>, vector<4xi32>, vector<4xi32>) unroll { + %header_local_row0 = index.add %load_row, %header_row_offset : index + %header_local_row = index.assume %header_local_row0 [range(%header_local_row0, 0, 63)] : index + %header_channel = index.add %channel_tile_base, %header_local_row : index + %valid_header_channel = index.cmp ult, %header_channel, %bounded_output_size : index + %loaded_gate_header, %loaded_up_header = scf.if %valid_header_channel -> (vector<4xi32>, vector<4xi32>) { + %header_channel_byte_add = index.scale %header_channel, %weight_row_bytes : index, offset -> offset + %header_row_byte_base = index.add %expert_byte_base, %header_channel_byte_add : offset + %gate_header = func.call @qwen3_moe_q4k_wmma_load_header(%gate_weight_noalias, %header_row_byte_base, %q4_block) : (buffer, offset, index) -> (vector<4xi32>) + %up_header = func.call @qwen3_moe_q4k_wmma_load_header(%up_weight_noalias, %header_row_byte_base, %q4_block) : (buffer, offset, index) -> (vector<4xi32>) + scf.yield %gate_header, %up_header : vector<4xi32>, vector<4xi32> + } else { + scf.yield %c0_i32x4, %c0_i32x4 : vector<4xi32>, vector<4xi32> + } + %updates_header0 = index.cmp eq, %header_row_offset, %c0 : index + %next_gate_header0 = scf.select %updates_header0, %loaded_gate_header, %prior_gate_header0 : vector<4xi32> + %next_gate_header1 = scf.select %updates_header0, %prior_gate_header1, %loaded_gate_header : vector<4xi32> + %next_up_header0 = scf.select %updates_header0, %loaded_up_header, %prior_up_header0 : vector<4xi32> + %next_up_header1 = scf.select %updates_header0, %prior_up_header1, %loaded_up_header : vector<4xi32> + scf.yield %next_gate_header0, %next_gate_header1, %next_up_header0, %next_up_header1 : vector<4xi32>, vector<4xi32>, vector<4xi32>, vector<4xi32> + } + %next_block_gate, %next_block_up = scf.for %q4_group_pair = [%c0 to %c4 step %c1](%pair_gate_acc = %block_gate : vector<16xf16>, %pair_up_acc = %block_up : vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) { + // Each code word supplies the low and high nibbles for one adjacent + // group pair. Retain both projections' words for exactly those uses. + %gate_q_word0, %gate_q_word1, %up_q_word0, %up_q_word1 = scf.for %code_row_offset = [%c0 to %c64 step %c32](%prior_gate_q_word0 = %c0_i32x1 : vector<1xi32>, %prior_gate_q_word1 = %c0_i32x1 : vector<1xi32>, %prior_up_q_word0 = %c0_i32x1 : vector<1xi32>, %prior_up_q_word1 = %c0_i32x1 : vector<1xi32>) -> (vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32>) unroll { + %code_local_row0 = index.add %load_row, %code_row_offset : index + %code_local_row = index.assume %code_local_row0 [range(%code_local_row0, 0, 63)] : index + %code_channel = index.add %channel_tile_base, %code_local_row : index + %valid_code_channel = index.cmp ult, %code_channel, %bounded_output_size : index + %loaded_gate_q_word, %loaded_up_q_word = scf.if %valid_code_channel -> (vector<1xi32>, vector<1xi32>) { + %bounded_q4_group_pair = index.assume %q4_group_pair [range(%q4_group_pair, 0, 3)] : index + %code_channel_byte_add = index.scale %code_channel, %weight_row_bytes : index, offset -> offset + %code_row_byte_base = index.add %expert_byte_base, %code_channel_byte_add : offset + %gate_q_word = func.call @qwen3_moe_q4k_wmma_load_code(%gate_weight_noalias, %code_row_byte_base, %q4_block, %bounded_q4_group_pair, %load_packet) : (buffer, offset, index, index, index) -> (vector<1xi32>) + %up_q_word = func.call @qwen3_moe_q4k_wmma_load_code(%up_weight_noalias, %code_row_byte_base, %q4_block, %bounded_q4_group_pair, %load_packet) : (buffer, offset, index, index, index) -> (vector<1xi32>) + scf.yield %gate_q_word, %up_q_word : vector<1xi32>, vector<1xi32> + } else { + scf.yield %c0_i32x1, %c0_i32x1 : vector<1xi32>, vector<1xi32> + } + %updates_q_word0 = index.cmp eq, %code_row_offset, %c0 : index + %next_gate_q_word0 = scf.select %updates_q_word0, %loaded_gate_q_word, %prior_gate_q_word0 : vector<1xi32> + %next_gate_q_word1 = scf.select %updates_q_word0, %prior_gate_q_word1, %loaded_gate_q_word : vector<1xi32> + %next_up_q_word0 = scf.select %updates_q_word0, %loaded_up_q_word, %prior_up_q_word0 : vector<1xi32> + %next_up_q_word1 = scf.select %updates_q_word0, %prior_up_q_word1, %loaded_up_q_word : vector<1xi32> + scf.yield %next_gate_q_word0, %next_gate_q_word1, %next_up_q_word0, %next_up_q_word1 : vector<1xi32>, vector<1xi32>, vector<1xi32>, vector<1xi32> + } + %next_pair_gate, %next_pair_up = scf.for %group_within_pair = [%c0 to %c2 step %c1](%group_gate = %pair_gate_acc : vector<16xf16>, %group_up = %pair_up_acc : vector<16xf16>) -> (vector<16xf16>, vector<16xf16>) { + %q4_group_base = index.mul %q4_group_pair, %c2 : index + %q4_group0 = index.add %q4_group_base, %group_within_pair : index + %q4_group = index.assume %q4_group0 [range(%q4_group0, 0, 7)] : index + %block_k_base = index.mul %q4_block, %c256 : index + %group_k_add = index.mul %q4_group, %c32 : index + %k_origin = index.add %block_k_base, %group_k_add : index + scf.for %row_offset = [%c0 to %c64 step %c32] unroll { + %local_row0 = index.add %load_row, %row_offset : index + %local_row = index.assume %local_row0 [range(%local_row0, 0, 63)] : index + %selects_first_row = index.cmp eq, %row_offset, %c0 : index + %selected_gate_header = scf.select %selects_first_row, %gate_header0, %gate_header1 : vector<4xi32> + %selected_gate_q_word = scf.select %selects_first_row, %gate_q_word0, %gate_q_word1 : vector<1xi32> + %channel = index.add %channel_tile_base, %local_row : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %gate_values = scf.if %valid_channel -> (vector<4xf16>) { + %decoded = func.call @qwen3_moe_q4k_wmma_vector4_from_header_code(%q4_group, %selected_gate_header, %selected_gate_q_word) : (index, vector<4xi32>, vector<1xi32>) -> (vector<4xf16>) + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %gate_values, %weight_stage_view[%local_row, %load_k] : vector<4xf16>, view<64x40xf16> + } + %assignment_i32 = view.load %route_stage_view[%load_row] : view<32xi32> -> i32 + %valid_assignment = scalar.cmpi sge, %assignment_i32, %c0_i32 : i32 + %activation_values = scf.if %valid_assignment -> (vector<4xf16>) { + %assignment0 = index.cast %assignment_i32 : i32 to index + %assignment = index.assume %assignment0 [range(%assignment0, 0, 16383)] : index + %bounded_assignment, %bounded_assignment_count = index.assume %assignment, %assignment_count [lt(%assignment, %assignment_count)] : index, index + %token0 = index.div %bounded_assignment, %bounded_route_count : index + %token, %input_token_count = index.assume %token0, %token_count [lt(%token0, %token_count)] : index, index + %input_k = index.add %k_origin, %load_k : index + %loaded = vector.load %input_view[%token, %input_k] : view<[%token_count]x[%input_size]xf32> -> vector<4xf32> + %converted = vector.fptrunc %loaded : vector<4xf32> to vector<4xf16> + scf.yield %converted : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %activation_values, %activation_stage_physical_view[%load_row, %load_k] : vector<4xf16>, view<32x40xf16> + kernel.barrier scope(workgroup) ordering(acq_rel) + %gate_next = scf.for %k_half = [%c0 to %c32 step %c16](%half_gate = %group_gate : vector<16xf16>) -> (vector<16xf16>) unroll { + %lhs = vector.fragment.load %weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %rhs = vector.fragment.load %activation_fragment_view[%k_half, %subgroup_route_add] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %next = vector.mma %lhs, %rhs, %half_gate : vector<16xf16>, vector<16xf16>, vector<16xf16> + scf.yield %next : vector<16xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.for %row_offset = [%c0 to %c64 step %c32] unroll { + %local_row0 = index.add %load_row, %row_offset : index + %local_row = index.assume %local_row0 [range(%local_row0, 0, 63)] : index + %selects_first_row = index.cmp eq, %row_offset, %c0 : index + %selected_up_header = scf.select %selects_first_row, %up_header0, %up_header1 : vector<4xi32> + %selected_up_q_word = scf.select %selects_first_row, %up_q_word0, %up_q_word1 : vector<1xi32> + %channel = index.add %channel_tile_base, %local_row : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %up_values = scf.if %valid_channel -> (vector<4xf16>) { + %decoded = func.call @qwen3_moe_q4k_wmma_vector4_from_header_code(%q4_group, %selected_up_header, %selected_up_q_word) : (index, vector<4xi32>, vector<1xi32>) -> (vector<4xf16>) + scf.yield %decoded : vector<4xf16> + } else { + scf.yield %c0_f16x4 : vector<4xf16> + } + vector.store %up_values, %weight_stage_view[%local_row, %load_k] : vector<4xf16>, view<64x40xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %up_next = scf.for %k_half = [%c0 to %c32 step %c16](%half_up = %group_up : vector<16xf16>) -> (vector<16xf16>) unroll { + %lhs = vector.fragment.load %weight_stage_view[%subgroup_channel_add, %k_half] shape [%m, %k] : view<64x40xf16> -> vector<16xf16> + %rhs = vector.fragment.load %activation_fragment_view[%k_half, %subgroup_route_add] shape [%k, %n] : view<32x32xf16, %activation_fragment_layout> -> vector<16xf16> + %next = vector.mma %lhs, %rhs, %half_up : vector<16xf16>, vector<16xf16>, vector<16xf16> + scf.yield %next : vector<16xf16> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.yield %gate_next, %up_next : vector<16xf16>, vector<16xf16> + } + scf.yield %next_pair_gate, %next_pair_up : vector<16xf16>, vector<16xf16> + } + scf.yield %next_block_gate, %next_block_up : vector<16xf16>, vector<16xf16> + } + %publish_route0 = index.div %lane, %c2 : index + %publish_route = index.assume %publish_route0 [range(%publish_route0, 0, 15)] : index + %publish_packet0 = index.rem %lane, %c2 : index + %publish_packet = index.assume %publish_packet0 [range(%publish_packet0, 0, 1)] : index + %publish_channel_add = index.mul %publish_packet, %c8 : index + %local_route = index.add %subgroup_route_add, %publish_route : index + %assignment_i32 = view.load %route_stage_view[%local_route] : view<32xi32> -> i32 + %assignment_nonnegative = scalar.cmpi sge, %assignment_i32, %c0_i32 : i32 + %safe_assignment_i32 = scf.if %assignment_nonnegative -> (i32) { + scf.yield %assignment_i32 : i32 + } else { + scf.yield %c0_i32 : i32 + } + %safe_assignment0 = index.cast %safe_assignment_i32 : i32 to index + %safe_assignment = index.assume %safe_assignment0 [range(%safe_assignment0, 0, 16383)] : index + %bounded_assignment, %bounded_assignment_count = index.assume %safe_assignment, %assignment_count [lt(%safe_assignment, %assignment_count)] : index, index + %subgroup_channel_base = index.add %channel_tile_base, %subgroup_channel_add : index + %channel = index.add %subgroup_channel_base, %publish_channel_add : index + %valid_channel = index.cmp ult, %channel, %bounded_output_size : index + %writes = scalar.andi %assignment_nonnegative, %valid_channel : i1 + vector.fragment.store %gate_result, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<16xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + %gate_values = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf16> -> vector<8xf16> + %gate_wide = vector.extf %gate_values : vector<8xf16> to vector<8xf32> + %activated = vector.siluf %gate_wide : vector<8xf32> + kernel.barrier scope(subgroup) ordering(acq_rel) + vector.fragment.store %up_result, %result_fragment_view[%c0, %c0] shape [%m, %n] : vector<16xf16>, view<16x16xf16, %result_fragment_layout> + kernel.barrier scope(subgroup) ordering(acq_rel) + scf.if %writes { + %up_values = vector.load %result_physical_view[%publish_route, %publish_channel_add] : view<16x16xf16> -> vector<8xf16> + %up_wide = vector.extf %up_values : vector<8xf16> to vector<8xf32> + %wide_values = vector.mulf %activated, %up_wide : vector<8xf32> + %values = vector.fptrunc %wide_values : vector<8xf32> to vector<8xf16> + %mask = vector.mask.range [%channel to %bounded_output_size step %c1] : index -> vector<8xi1> + vector.store.mask %values, %output_view[%bounded_assignment, %channel], %mask : vector<8xf16>, view<[%output_row_count]x[%bounded_output_size]xf16>, vector<8xi1> + } + kernel.barrier scope(subgroup) ordering(acq_rel) + } + kernel.return +} + +// Applies the model's SwiGLU epilogue to gate and up projections already +// scattered into logical [token][route][output channel] order. +kernel.def target(@qwen3_moe_gfx11_wave64) @qwen3_moe_routed_swiglu_f32(%token_count: index) { + %c1 = index.constant 1 : index + %workgroup_size = index.constant 256 : index + %rounding = index.constant 255 : index + %output_size = config.get @qwen3_moe.routed_gate_up.output_size : index + %route_count = config.get @qwen3_moe.routed_gate_up.route_count : index + %row_count = index.mul %token_count, %route_count : index + %element_count = index.mul %row_count, %output_size : index + %rounded_count = index.add %element_count, %rounding : index + %workgroup_count = index.div %rounded_count, %workgroup_size : index + kernel.launch.config workgroups(%workgroup_count, %c1, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %gate: buffer, %up: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %route_count = config.get @qwen3_moe.routed_gate_up.route_count : index + %output_size = config.get @qwen3_moe.routed_gate_up.output_size : index + %bounded_route_count = index.assume %route_count [range(%route_count, 1, 8)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 4096)] : index + %workgroup = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %workgroup_size = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %row_count = index.mul %token_count, %bounded_route_count : index + %element_count = index.mul %row_count, %bounded_output_size : index + %element = index.madd %workgroup, %workgroup_size, %lane : index + %in_bounds = index.cmp ult, %element, %element_count : index + %gate_noalias, %up_noalias, %output_noalias = buffer.assume.noalias %gate, %up, %output : buffer, buffer, buffer + scf.if %in_bounds { + %bounded_element, %view_element_count = index.assume %element, %element_count [lt(%element, %element_count)] : index, index + %gate_view = buffer.view %gate_noalias[%c0_offset] : buffer -> view<[%view_element_count]xf32> + %up_view = buffer.view %up_noalias[%c0_offset] : buffer -> view<[%view_element_count]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%view_element_count]xf32> + %gate_value = view.load %gate_view[%bounded_element] : view<[%view_element_count]xf32> -> f32 + %up_value = view.load %up_view[%bounded_element] : view<[%view_element_count]xf32> -> f32 + %activated = scalar.siluf %gate_value : f32 + %result = scalar.mulf %activated, %up_value : f32 + view.store %result, %output_view[%bounded_element] : f32, view<[%view_element_count]xf32> + } + kernel.return +} + +// Reference FP16 publication path for checking the fused gate/up provider at +// its actual model boundary. Gate and up projections remain independently +// materialized so this schedule does not share the fused kernel's matrix or +// publication implementation. +kernel.def target(@qwen3_moe_gfx11_wave64) @qwen3_moe_routed_swiglu_f16(%token_count: index) { + %c1 = index.constant 1 : index + %workgroup_size = index.constant 256 : index + %rounding = index.constant 255 : index + %output_size = config.get @qwen3_moe.routed_gate_up.output_size : index + %route_count = config.get @qwen3_moe.routed_gate_up.route_count : index + %row_count = index.mul %token_count, %route_count : index + %element_count = index.mul %row_count, %output_size : index + %rounded_count = index.add %element_count, %rounding : index + %workgroup_count = index.div %rounded_count, %workgroup_size : index + kernel.launch.config workgroups(%workgroup_count, %c1, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %gate: buffer, %up: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %route_count = config.get @qwen3_moe.routed_gate_up.route_count : index + %output_size = config.get @qwen3_moe.routed_gate_up.output_size : index + %bounded_route_count = index.assume %route_count [range(%route_count, 1, 8)] : index + %bounded_output_size = index.assume %output_size [range(%output_size, 1, 4096)] : index + %workgroup = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %workgroup_size = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %row_count = index.mul %token_count, %bounded_route_count : index + %element_count = index.mul %row_count, %bounded_output_size : index + %element = index.madd %workgroup, %workgroup_size, %lane : index + %in_bounds = index.cmp ult, %element, %element_count : index + %gate_noalias, %up_noalias, %output_noalias = buffer.assume.noalias %gate, %up, %output : buffer, buffer, buffer + scf.if %in_bounds { + %bounded_element, %view_element_count = index.assume %element, %element_count [lt(%element, %element_count)] : index, index + %gate_view = buffer.view %gate_noalias[%c0_offset] : buffer -> view<[%view_element_count]xf32> + %up_view = buffer.view %up_noalias[%c0_offset] : buffer -> view<[%view_element_count]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%view_element_count]xf16> + %gate_value = view.load %gate_view[%bounded_element] : view<[%view_element_count]xf32> -> f32 + %up_value = view.load %up_view[%bounded_element] : view<[%view_element_count]xf32> -> f32 + %activated = scalar.siluf %gate_value : f32 + %wide_result = scalar.mulf %activated, %up_value : f32 + %result = scalar.fptrunc %wide_result : f32 to f16 + view.store %result, %output_view[%bounded_element] : f16, view<[%view_element_count]xf16> + } + kernel.return +} + +// The constant input is exactly representable in Q8_1 and FP16. Comparing +// against the established integer-dot path therefore isolates Q4_K decode, +// routing, FP16 matrix accumulation, tail handling, and the SwiGLU boundary. +check.case public @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_differential_case { + %token_count = check.literal value(67) : index + %input_size = check.literal value(512) : index + %route_count = check.literal value(2) : index + %route_stride = check.literal value(4) : index + %expert_count = check.literal value(4) : index + %output_size = check.literal value(33) : index + %input = check.generate.fill value(0.00390625) : tensor<67x512xf32> + %q8_input = check.generate.fill value(0) : tensor<67x576xi8> + %route_ids = check.generate.iota offset(0) step(1) period(4) : tensor<67x4xi32> + %expert_table = check.generate.fill value(-1) : tensor<272xi32> + %partition_table = check.generate.fill value(-1) : tensor<10xi32> + %gate_weight = check.generate.fill value(34) : tensor<4x33x2x144xi8> + %up_weight = check.generate.fill value(35) : tensor<4x33x2x144xi8> + %gate_projection = check.generate.fill value(0.0) : tensor<67x2x33xf32> + %up_projection = check.generate.fill value(0.0) : tensor<67x2x33xf32> + %expected = check.generate.fill value(0.0) : tensor<67x2x33xf32> + %actual = check.generate.fill value(1.0) : tensor<67x2x33xf32> + %actual_f16 = check.generate.fill value(1.0) : tensor<67x2x33xf16> + %fused_actual = check.generate.fill value(1.0) : tensor<67x2x33xf16> + kernel.launch @ggml_quantize_q8_1_x4_f32[%token_count, %input_size](%token_count, %input_size, %input, %q8_input) : [index, index](index, index, tensor<67x512xf32>, tensor<67x576xi8>) + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<67x4xi32>, tensor<272xi32>) + kernel.launch @qwen3_moe_build_expert_partition_table[%token_count, %route_count, %expert_count](%token_count, %route_count, %expert_count, %expert_table, %partition_table) : [index, index, index](index, index, index, tensor<272xi32>, tensor<10xi32>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_q8[%token_count, %route_count, %route_stride, %expert_count, %output_size](%token_count, %route_count, %route_stride, %expert_count, %output_size, %q8_input, %route_ids, %gate_weight, %up_weight, %expected) : [index, index, index, index, index](index, index, index, index, index, tensor<67x576xi8>, tensor<67x4xi32>, tensor<4x33x2x144xi8>, tensor<4x33x2x144xi8>, tensor<67x2x33xf32>) + kernel.launch @qwen3_moe_routed_linear_q4k_f16_wmma[%token_count](%token_count, %input, %expert_table, %gate_weight, %gate_projection) : [index](index, tensor<67x512xf32>, tensor<272xi32>, tensor<4x33x2x144xi8>, tensor<67x2x33xf32>) + kernel.launch @qwen3_moe_routed_linear_q4k_f16_wmma[%token_count](%token_count, %input, %expert_table, %up_weight, %up_projection) : [index](index, tensor<67x512xf32>, tensor<272xi32>, tensor<4x33x2x144xi8>, tensor<67x2x33xf32>) + kernel.launch @qwen3_moe_routed_swiglu_f32[%token_count](%token_count, %gate_projection, %up_projection, %actual) : [index](index, tensor<67x2x33xf32>, tensor<67x2x33xf32>, tensor<67x2x33xf32>) + kernel.launch @qwen3_moe_routed_swiglu_f16[%token_count](%token_count, %gate_projection, %up_projection, %actual_f16) : [index](index, tensor<67x2x33xf32>, tensor<67x2x33xf32>, tensor<67x2x33xf16>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma[%token_count](%token_count, %input, %expert_table, %partition_table, %gate_weight, %up_weight, %fused_actual) : [index](index, tensor<67x512xf32>, tensor<272xi32>, tensor<10xi32>, tensor<4x33x2x144xi8>, tensor<4x33x2x144xi8>, tensor<67x2x33xf16>) + check.expect.close actual(%actual) expected(%expected) atol(0.01) rtol(0.01) nan(same) : tensor<67x2x33xf32> + check.expect.close actual(%fused_actual) expected(%actual_f16) atol(0.01) rtol(0.01) nan(same) : tensor<67x2x33xf16> + check.return +} + +check.case public @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_benchmark_case { + %token_count = check.param.choice values([1, 2, 4, 8, 16, 17, 32, 63, 128, 129, 512, 1024, 2048]) name("token_count") : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(768) : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + // The compact iota exactly matches the oracle's balanced ring: + // expert(token, route) = (8 * token + route) % 128. + %route_ids = check.generate.iota offset(0) step(1) period(128) : tensor<[%token_count]x8xi32> + // One count per expert followed by 2048 assignment slots per expert. + %expert_table = check.generate.fill value(-1) : tensor<262272xi32> + %gate_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %up_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %gate_projection = check.generate.fill value(1.0) : tensor<[%token_count]x8x768xf32> + %up_projection = check.generate.fill value(1.0) : tensor<[%token_count]x8x768xf32> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x8x768xf32> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x8x768xf32> + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<[%token_count]x8xi32>, tensor<262272xi32>) + kernel.launch @qwen3_moe_routed_linear_q4k_f16_wmma[%token_count](%token_count, %input, %expert_table, %gate_weight, %gate_projection) : [index](index, tensor<[%token_count]x2048xf32>, tensor<262272xi32>, tensor<128x768x8x144xi8>, tensor<[%token_count]x8x768xf32>) + kernel.launch @qwen3_moe_routed_linear_q4k_f16_wmma[%token_count](%token_count, %input, %expert_table, %up_weight, %up_projection) : [index](index, tensor<[%token_count]x2048xf32>, tensor<262272xi32>, tensor<128x768x8x144xi8>, tensor<[%token_count]x8x768xf32>) + kernel.launch @qwen3_moe_routed_swiglu_f32[%token_count](%token_count, %gate_projection, %up_projection, %output) : [index](index, tensor<[%token_count]x8x768xf32>, tensor<[%token_count]x8x768xf32>, tensor<[%token_count]x8x768xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x8x768xf32> + check.return +} + +check.case public @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_benchmark_case { + %token_count = check.param.choice values([1, 2, 4, 8, 16, 17, 32, 63, 128, 129, 512, 1024, 2048]) name("token_count") : index + %route_count = check.literal value(8) : index + %route_stride = check.literal value(8) : index + %expert_count = check.literal value(128) : index + %output_size = check.literal value(768) : index + %input = check.generate.fill value(0.0) : tensor<[%token_count]x2048xf32> + %route_ids = check.generate.iota offset(0) step(1) period(128) : tensor<[%token_count]x8xi32> + // One count per expert followed by 2048 assignment slots per expert. + %expert_table = check.generate.fill value(-1) : tensor<262272xi32> + // One exact count followed by at most 640 packed descriptors. + %partition_table = check.generate.fill value(-1) : tensor<641xi32> + %gate_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %up_weight = check.generate.fill value(0) : tensor<128x768x8x144xi8> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x8x768xf16> + %expected = check.generate.fill value(0.0) : tensor<[%token_count]x8x768xf16> + kernel.launch @qwen3_moe_build_expert_table[%token_count, %route_count, %route_stride, %expert_count](%token_count, %route_count, %route_stride, %expert_count, %route_ids, %expert_table) : [index, index, index, index](index, index, index, index, tensor<[%token_count]x8xi32>, tensor<262272xi32>) + kernel.launch @qwen3_moe_build_expert_partition_table[%token_count, %route_count, %expert_count](%token_count, %route_count, %expert_count, %expert_table, %partition_table) : [index, index, index](index, index, index, tensor<262272xi32>, tensor<641xi32>) + kernel.launch @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma[%token_count](%token_count, %input, %expert_table, %partition_table, %gate_weight, %up_weight, %output) : [index](index, tensor<[%token_count]x2048xf32>, tensor<262272xi32>, tensor<641xi32>, tensor<128x768x8x144xi8>, tensor<128x768x8x144xi8>, tensor<[%token_count]x8x768xf16>) + check.expect.close actual(%output) expected(%expected) atol(0.0) rtol(0.0) nan(same) : tensor<[%token_count]x8x768xf16> + check.return +} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_differential_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_differential + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_decode {token_count = 1} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_small_batch_2 {token_count = 2} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_small_batch_4 {token_count = 4} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_small_batch_8 {token_count = 8} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_small_batch_16 {token_count = 16} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_prefill_17 {token_count = 17} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_prefill_63 {token_count = 63} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_prefill_129 {token_count = 129} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_prefill_1024 {token_count = 1024} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_prefill_2048 {token_count = 2048} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_decode {token_count = 1} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_small_batch_2 {token_count = 2} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_small_batch_4 {token_count = 4} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_small_batch_8 {token_count = 8} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_small_batch_16 {token_count = 16} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_prefill_17 {token_count = 17} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_prefill_63 {token_count = 63} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_prefill_129 {token_count = 129} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_prefill_1024 {token_count = 1024} + +check.benchmark<@qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_benchmark_case> @qwen3_moe_routed_gate_up_swiglu_q4k_f16_wmma_fused_prefill_2048 {token_count = 2048} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/router_projection_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/router_projection_f32.loom new file mode 100644 index 000000000000..be653b373509 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/router_projection_f32.loom @@ -0,0 +1,266 @@ +// Dense F32 projection from hidden activations to MoE router logits. +// +// Two matrix-vector schedules cover the measured shape classes. Decode assigns +// one expert row to each wave64, matching the llama.cpp Vulkan geometry and +// maximizing independent waves for one token. Prefill assigns four adjacent +// expert rows to each wave32 and reuses every activation packet across their +// dot products. The latter reduces activation traffic once token parallelism +// already fills the device. +// +// gfx1151 uses descriptor-backed MUBUF accesses for the decode schedule. The +// generic gfx11 provider retains global addressing where descriptor setup costs +// more than it saves. +template.decl @qwen3_moe.router_projection.storage(%input: buffer, %weight: buffer, %output: buffer) -> (buffer, buffer, buffer) + +amdgpu.target @qwen3_moe_router_projection_gfx11_wave32 {subgroup_size = 32} + +amdgpu.target @qwen3_moe_router_projection_gfx11_wave64 {subgroup_size = 64} + +amdgpu.target @qwen3_moe_router_projection_gfx1151 + +config.decl @qwen3_moe.model.hidden_size : %value: index where [range(%value, 128, 32768), mul(%value, 128)] + +config.decl @qwen3_moe.router.expert_count : %value: index where [range(%value, 32, 512), mul(%value, 32)] + +// Selects the target's preferred storage contract without exposing addressing +// policy at the projection call site. +template.def<@qwen3_moe.router_projection.storage> target(@qwen3_moe_router_projection_gfx1151) priority(20) @qwen3_moe_router_projection_descriptor_storage(%input: buffer, %weight: buffer, %output: buffer) -> (buffer, buffer, buffer) { + %input_descriptor = buffer.assume.memory_space %input : buffer + %weight_descriptor = buffer.assume.memory_space %weight : buffer + %output_descriptor = buffer.assume.memory_space %output : buffer + template.return %input_descriptor, %weight_descriptor, %output_descriptor : buffer, buffer, buffer +} + +template.def<@qwen3_moe.router_projection.storage> priority(1) @qwen3_moe_router_projection_global_storage(%input: buffer, %weight: buffer, %output: buffer) -> (buffer, buffer, buffer) { + template.return %input, %weight, %output : buffer, buffer, buffer +} + +// Scalar differential oracle. It is reachable only from check cases and keeps +// production validation independent of the packetized wave schedule. +kernel.def target(@qwen3_moe_router_projection_gfx11_wave32) @qwen3_moe_router_projection_f32_reference(%token_count: index) { + %expert_count = config.get @qwen3_moe.router.expert_count : index + %c1 = index.constant 1 : index + kernel.launch.config workgroups(%expert_count, %token_count, %c1) workgroup_size(%c1, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %hidden_size0 = config.get @qwen3_moe.model.hidden_size : index + %expert_count0 = config.get @qwen3_moe.router.expert_count : index + %hidden_size, %expert_count = index.assume %hidden_size0, %expert_count0 [range(%hidden_size0, 128, 32768), mul(%hidden_size0, 128), range(%expert_count0, 32, 512), mul(%expert_count0, 32)] : index, index + %expert0 = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c0_f32 = scalar.constant 0.0 : f32 + %c0_offset = index.constant 0 : offset + %expert = index.assume %expert0 [lt(%expert0, %expert_count)] : index + %valid_token = index.cmp ult, %token0, %token_count : index + %safe_token0 = scf.select %valid_token, %token0, %c0 : index + %token, %launch_token_count = index.assume %safe_token0, %token_count [lt(%safe_token0, %token_count)] : index, index + %input_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input, %weight, %output : buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%hidden_size]xf32> + %weight_view = buffer.view %weight_noalias[%c0_offset] : buffer -> view<[%expert_count]x[%hidden_size]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%expert_count]xf32> + %sum = scf.for %channel = [%c0 to %hidden_size step %c1](%accumulator = %c0_f32 : f32) -> (f32) { + %input_value = view.load %input_view[%token, %channel] : view<[%launch_token_count]x[%hidden_size]xf32> -> f32 + %weight_value = view.load %weight_view[%expert, %channel] : view<[%expert_count]x[%hidden_size]xf32> -> f32 + %next_accumulator = scalar.fmaf %input_value, %weight_value, %accumulator : f32 + scf.yield %next_accumulator : f32 + } + scf.if %valid_token { + view.store %sum, %output_view[%token, %expert] : f32, view<[%launch_token_count]x[%expert_count]xf32> + } + kernel.return +} + +// One wave64 owns one output row. This schedule is selected for decode. +kernel.def target(@qwen3_moe_router_projection_gfx11_wave64) @qwen3_moe_router_projection_f32_one_row_wave64(%token_count: index) { + %expert_count = config.get @qwen3_moe.router.expert_count : index + %c1 = index.constant 1 : index + %wave_size = target.subgroup.size : index + kernel.launch.config workgroups(%expert_count, %token_count, %c1) workgroup_size(%wave_size, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %hidden_size0 = config.get @qwen3_moe.model.hidden_size : index + %expert_count0 = config.get @qwen3_moe.router.expert_count : index + %hidden_size, %expert_count = index.assume %hidden_size0, %expert_count0 [range(%hidden_size0, 128, 32768), mul(%hidden_size0, 128), range(%expert_count0, 32, 512), mul(%expert_count0, 32)] : index, index + %expert0 = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c256 = index.constant 256 : index + %c1024 = index.constant 1024 : index + %c0_i32 = scalar.constant 0 : i32 + %c0_f32 = scalar.constant 0.0 : f32 + %c0_offset = index.constant 0 : offset + %valid_token = index.cmp ult, %token0, %token_count : index + %safe_token0 = scf.select %valid_token, %token0, %c0 : index + %token, %launch_token_count = index.assume %safe_token0, %token_count [lt(%safe_token0, %token_count)] : index, index + %expert = index.assume %expert0 [lt(%expert0, %expert_count)] : index + %lane_channel = index.mul %lane, %c4 : index + %input_storage, %weight_storage, %output_storage = template.apply<@qwen3_moe.router_projection.storage>(%input, %weight, %output) : (buffer, buffer, buffer) -> (buffer, buffer, buffer) + %input_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input_storage, %weight_storage, %output_storage : buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%hidden_size]xf32> + %weight_view = buffer.view %weight_noalias[%c0_offset] : buffer -> view<[%expert_count]x[%hidden_size]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%expert_count]xf32> + %full_channel_limit = index.sub %hidden_size, %c3 : index + %unroll_remainder = index.rem %hidden_size, %c1024 : index + %uses_unrolled_schedule = index.cmp eq, %unroll_remainder, %c0 : index + %lane_sum = scf.if %uses_unrolled_schedule -> (f32) { + %unrolled_sum = scf.for %channel = [%lane_channel to %full_channel_limit step %c256](%accumulator = %c0_f32 : f32) -> (f32) unroll(%c4) schedule(interleaved) { + %input_values = vector.load %input_view[%token, %channel] : view<[%launch_token_count]x[%hidden_size]xf32> -> vector<4xf32> + %weight_values = vector.load %weight_view[%expert, %channel] : view<[%expert_count]x[%hidden_size]xf32> -> vector<4xf32> + %next_accumulator = vector.dotf %input_values, %weight_values, %accumulator : vector<4xf32>, vector<4xf32>, f32 + scf.yield %next_accumulator : f32 + } + scf.yield %unrolled_sum : f32 + } else { + %general_sum = scf.for %channel = [%lane_channel to %full_channel_limit step %c256](%accumulator = %c0_f32 : f32) -> (f32) { + %input_values = vector.load %input_view[%token, %channel] : view<[%launch_token_count]x[%hidden_size]xf32> -> vector<4xf32> + %weight_values = vector.load %weight_view[%expert, %channel] : view<[%expert_count]x[%hidden_size]xf32> -> vector<4xf32> + %next_accumulator = vector.dotf %input_values, %weight_values, %accumulator : vector<4xf32>, vector<4xf32>, f32 + scf.yield %next_accumulator : f32 + } + scf.yield %general_sum : f32 + } + %sum = kernel.subgroup.reduce %lane_sum : f32 + %lane_i32 = index.cast %lane : index to i32 + %writes_output = scalar.cmpi eq, %lane_i32, %c0_i32 : i32 + %publishes_output = scalar.andi %valid_token, %writes_output : i1 + scf.if %publishes_output { + view.store %sum, %output_view[%token, %expert] : f32, view<[%launch_token_count]x[%expert_count]xf32> + } + kernel.return +} + +// One wave32 owns four adjacent output rows and reuses each input packet across +// their contractions. This schedule is selected once token parallelism can +// populate the device independently. +kernel.def target(@qwen3_moe_router_projection_gfx11_wave32) @qwen3_moe_router_projection_f32_four_row_wave32(%token_count: index) { + %expert_count = config.get @qwen3_moe.router.expert_count : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %wave_size = target.subgroup.size : index + %expert_tiles = index.div %expert_count, %c4 : index + kernel.launch.config workgroups(%expert_tiles, %token_count, %c1) workgroup_size(%wave_size, %c1, %c1) : index +} launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) where [range(%token_count, 1, 2048)] { + %hidden_size0 = config.get @qwen3_moe.model.hidden_size : index + %expert_count0 = config.get @qwen3_moe.router.expert_count : index + %hidden_size, %expert_count = index.assume %hidden_size0, %expert_count0 [range(%hidden_size0, 128, 32768), mul(%hidden_size0, 128), range(%expert_count0, 32, 512), mul(%expert_count0, 32)] : index, index + %expert_tile = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %lane = kernel.subgroup.lane.id : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c128 = index.constant 128 : index + %c1024 = index.constant 1024 : index + %c0_i32 = scalar.constant 0 : i32 + %c0_f32 = scalar.constant 0.0 : f32 + %c0_offset = index.constant 0 : offset + %c0 = index.constant 0 : index + %valid_token = index.cmp ult, %token0, %token_count : index + %safe_token0 = scf.select %valid_token, %token0, %c0 : index + %token, %launch_token_count = index.assume %safe_token0, %token_count [lt(%safe_token0, %token_count)] : index, index + %expert_base = index.mul %expert_tile, %c4 : index + %expert1 = index.add %expert_base, %c1 : index + %expert2 = index.add %expert_base, %c2 : index + %expert3 = index.add %expert_base, %c3 : index + %lane_channel = index.mul %lane, %c4 : index + %c0_f32x4 = vector.splat %c0_f32 : vector<4xf32> + %input_noalias, %weight_noalias, %output_noalias = buffer.assume.noalias %input, %weight, %output : buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%hidden_size]xf32> + %weight_view = buffer.view %weight_noalias[%c0_offset] : buffer -> view<[%expert_count]x[%hidden_size]xf32> + %output_view = buffer.view %output_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%expert_count]xf32> + %full_channel_limit = index.sub %hidden_size, %c3 : index + %unroll_remainder = index.rem %hidden_size, %c1024 : index + %uses_unrolled_schedule = index.cmp eq, %unroll_remainder, %c0 : index + %unroll_count = scf.select %uses_unrolled_schedule, %c4, %c1 : index + %lane_sums = scf.for %channel = [%lane_channel to %full_channel_limit step %c128](%accumulators = %c0_f32x4 : vector<4xf32>) -> (vector<4xf32>) unroll(%unroll_count) schedule(interleaved) { + %input_values = vector.load %input_view[%token, %channel] : view<[%launch_token_count]x[%hidden_size]xf32> -> vector<4xf32> + %weight_values0 = vector.load %weight_view[%expert_base, %channel] : view<[%expert_count]x[%hidden_size]xf32> -> vector<4xf32> + %weight_values1 = vector.load %weight_view[%expert1, %channel] : view<[%expert_count]x[%hidden_size]xf32> -> vector<4xf32> + %weight_values2 = vector.load %weight_view[%expert2, %channel] : view<[%expert_count]x[%hidden_size]xf32> -> vector<4xf32> + %weight_values3 = vector.load %weight_view[%expert3, %channel] : view<[%expert_count]x[%hidden_size]xf32> -> vector<4xf32> + %accumulator0 = vector.extract %accumulators[0] : vector<4xf32> -> f32 + %accumulator1 = vector.extract %accumulators[1] : vector<4xf32> -> f32 + %accumulator2 = vector.extract %accumulators[2] : vector<4xf32> -> f32 + %accumulator3 = vector.extract %accumulators[3] : vector<4xf32> -> f32 + %next0 = vector.dotf %input_values, %weight_values0, %accumulator0 : vector<4xf32>, vector<4xf32>, f32 + %next1 = vector.dotf %input_values, %weight_values1, %accumulator1 : vector<4xf32>, vector<4xf32>, f32 + %next2 = vector.dotf %input_values, %weight_values2, %accumulator2 : vector<4xf32>, vector<4xf32>, f32 + %next3 = vector.dotf %input_values, %weight_values3, %accumulator3 : vector<4xf32>, vector<4xf32>, f32 + %next_accumulators = vector.from_elements %next0, %next1, %next2, %next3 : vector<4xf32> + scf.yield %next_accumulators : vector<4xf32> + } + %sums = kernel.subgroup.reduce %lane_sums : vector<4xf32> + %lane_i32 = index.cast %lane : index to i32 + %writes_output = scalar.cmpi eq, %lane_i32, %c0_i32 : i32 + %publishes_output = scalar.andi %valid_token, %writes_output : i1 + scf.if %publishes_output { + vector.store %sums, %output_view[%token, %expert_base] : vector<4xf32>, view<[%launch_token_count]x[%expert_count]xf32> + } + kernel.return +} + +// Nonuniform inputs and weights compare both production schedules against an +// independently structured scalar contraction. The shape crosses every lane +// and all eight production output tiles. +check.case public @qwen3_moe_router_projection_f32_differential_case { + %token_count = check.literal value(2) : index + %input = check.generate.iota offset(-0.25) step(0.0078125) period(17) : tensor<2x512xf32> + %weight = check.generate.iota offset(-0.5) step(0.015625) period(31) : tensor<32x512xf32> + %expected = check.generate.fill value(0.0) : tensor<2x32xf32> + %decode_actual = check.generate.fill value(1.0) : tensor<2x32xf32> + %prefill_actual = check.generate.fill value(1.0) : tensor<2x32xf32> + kernel.launch @qwen3_moe_router_projection_f32_reference[%token_count](%token_count, %input, %weight, %expected) : [index](index, tensor<2x512xf32>, tensor<32x512xf32>, tensor<2x32xf32>) + kernel.launch @qwen3_moe_router_projection_f32_one_row_wave64[%token_count](%token_count, %input, %weight, %decode_actual) : [index](index, tensor<2x512xf32>, tensor<32x512xf32>, tensor<2x32xf32>) + kernel.launch @qwen3_moe_router_projection_f32_four_row_wave32[%token_count](%token_count, %input, %weight, %prefill_actual) : [index](index, tensor<2x512xf32>, tensor<32x512xf32>, tensor<2x32xf32>) + check.expect.close actual(%decode_actual) expected(%expected) atol(0.0001) rtol(0.0001) nan(same) : tensor<2x32xf32> + check.expect.close actual(%prefill_actual) expected(%expected) atol(0.0001) rtol(0.0001) nan(same) : tensor<2x32xf32> + check.return +} + +// Binary-exact values give every production row the analytic result 0.125, +// retaining correctness checks in each measured runtime token bucket. Both +// schedule cases expose the full shape sweep so routing decisions remain +// directly measurable. +check.case public @qwen3_moe_router_projection_f32_decode_benchmark_case { + %token_count = check.param.choice values([1, 32, 128, 512]) name("token_count") : index + %input = check.generate.fill value(0.00390625) : tensor<[%token_count]x2048xf32> + %weight = check.generate.fill value(0.015625) : tensor<128x2048xf32> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x128xf32> + %expected = check.generate.fill value(0.125) : tensor<[%token_count]x128xf32> + kernel.launch @qwen3_moe_router_projection_f32_one_row_wave64[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<[%token_count]x2048xf32>, tensor<128x2048xf32>, tensor<[%token_count]x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0001) rtol(0.0001) nan(same) : tensor<[%token_count]x128xf32> + check.return +} + +check.case public @qwen3_moe_router_projection_f32_prefill_benchmark_case { + %token_count = check.param.choice values([1, 32, 128, 512]) name("token_count") : index + %input = check.generate.fill value(0.00390625) : tensor<[%token_count]x2048xf32> + %weight = check.generate.fill value(0.015625) : tensor<128x2048xf32> + %output = check.generate.fill value(1.0) : tensor<[%token_count]x128xf32> + %expected = check.generate.fill value(0.125) : tensor<[%token_count]x128xf32> + kernel.launch @qwen3_moe_router_projection_f32_four_row_wave32[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<[%token_count]x2048xf32>, tensor<128x2048xf32>, tensor<[%token_count]x128xf32>) + check.expect.close actual(%output) expected(%expected) atol(0.0001) rtol(0.0001) nan(same) : tensor<[%token_count]x128xf32> + check.return +} + +check.benchmark<@qwen3_moe_router_projection_f32_differential_case> @qwen3_moe_router_projection_f32_differential + +check.benchmark<@qwen3_moe_router_projection_f32_decode_benchmark_case> @qwen3_moe_router_projection_f32_decode {token_count = 1} + +check.benchmark<@qwen3_moe_router_projection_f32_prefill_benchmark_case> @qwen3_moe_router_projection_f32_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_router_projection_f32_prefill_benchmark_case> @qwen3_moe_router_projection_f32_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_router_projection_f32_prefill_benchmark_case> @qwen3_moe_router_projection_f32_prefill_512 {token_count = 512} + +check.benchmark<@qwen3_moe_router_projection_f32_prefill_benchmark_case> @qwen3_moe_router_projection_f32_prefill_schedule_decode {token_count = 1} + +check.benchmark<@qwen3_moe_router_projection_f32_decode_benchmark_case> @qwen3_moe_router_projection_f32_decode_schedule_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_router_projection_f32_decode_benchmark_case> @qwen3_moe_router_projection_f32_decode_schedule_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_router_projection_f32_decode_benchmark_case> @qwen3_moe_router_projection_f32_decode_schedule_prefill_512 {token_count = 512} diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/router_projection_top8_fused_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/router_projection_top8_fused_f32.loom new file mode 100644 index 000000000000..a2fa68657229 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/router_projection_top8_fused_f32.loom @@ -0,0 +1,235 @@ +// Fused decode router projection and deterministic normalized top-8 routing. +// +// Each wave64 computes a target-selected adjacent expert tile. A device-scope completion +// sequence publishes the row across workgroups, and the last arrival invokes +// the same row-selection contract as the standalone router. The completion +// counter is reset after route publication so reusable command buffers can +// issue the kernel repeatedly against the same storage. +// +// This schedule is decode-only. Prefill has enough token parallelism to keep +// projection and routing independent without paying completion atomics. +template.decl @qwen3_moe.router.fused.experts_per_wave() -> (index) + +template.decl @qwen3_moe.router.fused.projection(%hidden_size: index, %expert_count: index, %token_count: index, %token: index, %expert_tile: index, %lane: index, %writes_projection: i1, %input: buffer, %weight: buffer, %logits: buffer) + +template.decl @qwen3_moe.router.top8.row(%arg0: i1, %arg1: index, %arg2: index, %arg3: index, %arg4: buffer, %arg5: buffer, %arg6: buffer) + +template.decl @qwen3_moe.router_projection.storage(%arg0: buffer, %arg1: buffer, %arg2: buffer) -> (buffer, buffer, buffer) + +amdgpu.target @qwen3_moe_router_fused_gfx11_wave64 {subgroup_size = 64} + +amdgpu.target @qwen3_moe_router_fused_gfx1151 + +config.decl @qwen3_moe.model.hidden_size : %value: index where [range(%value, 128, 32768), mul(%value, 128)] + +config.decl @qwen3_moe.router.expert_count : %value: index where [range(%value, 32, 512), mul(%value, 32)] + +kernel.decl @qwen3_moe_router_projection_f32_four_row_wave32(%token_count$26: index) launch(%token_count$27: index, %input: buffer, %weight: buffer, %output: buffer) + +kernel.decl @qwen3_moe_router_top8_f32(%token_count$31: index, %route_id_stride$32: index) launch(%token_count$33: index, %route_id_stride$34: index, %logits: buffer, %route_ids: buffer, %route_weights: buffer) + +template.def<@qwen3_moe.router.fused.experts_per_wave> target(@qwen3_moe_router_fused_gfx1151) requires [#target.subgroup.size<64>] priority(20) @qwen3_moe_router_fused_two_experts_per_wave() -> (index) { + %c2 = index.constant 2 : index + template.return %c2 : index +} + +template.def<@qwen3_moe.router.fused.experts_per_wave> requires [#target.subgroup.size<64>] priority(1) @qwen3_moe_router_fused_four_experts_per_wave() -> (index) { + %c4 = index.constant 4 : index + template.return %c4 : index +} + +template.def<@qwen3_moe.router.fused.projection> device target(@qwen3_moe_router_fused_gfx1151) requires [#target.subgroup.size<64>] priority(20) @qwen3_moe_router_fused_two_expert_projection(%hidden_size: index, %expert_count: index, %token_count: index, %token: index, %expert_tile: index, %lane: index, %writes_projection: i1, %input: buffer, %weight: buffer, %logits: buffer) { + %c0_offset = index.constant 0 : offset + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c256 = index.constant 256 : index + %c0_f32 = scalar.constant 0.0 : f32 + %expert_base = index.mul %expert_tile, %c2 : index + %expert1 = index.add %expert_base, %c1 : index + %lane_channel = index.mul %lane, %c4 : index + %input_view = buffer.view %input[%c0_offset] : buffer -> view<[%token_count]x[%hidden_size]xf32> + %weight_view = buffer.view %weight[%c0_offset] : buffer -> view<[%expert_count]x[%hidden_size]xf32> + %logits_view = buffer.view %logits[%c0_offset] : buffer -> view<[%token_count]x[%expert_count]xf32> + %c0_f32x2 = vector.splat %c0_f32 : vector<2xf32> + %lane_sums = scf.for %channel = [%lane_channel to %hidden_size step %c256](%accumulators = %c0_f32x2 : vector<2xf32>) -> (vector<2xf32>) unroll(%c4) schedule(interleaved) { + %input_values = vector.load %input_view[%token, %channel] : view<[%token_count]x[%hidden_size]xf32> -> vector<4xf32> + %weight_values0 = vector.load %weight_view[%expert_base, %channel] : view<[%expert_count]x[%hidden_size]xf32> -> vector<4xf32> + %weight_values1 = vector.load %weight_view[%expert1, %channel] : view<[%expert_count]x[%hidden_size]xf32> -> vector<4xf32> + %accumulator0 = vector.extract %accumulators[0] : vector<2xf32> -> f32 + %accumulator1 = vector.extract %accumulators[1] : vector<2xf32> -> f32 + %next0 = vector.dotf %input_values, %weight_values0, %accumulator0 : vector<4xf32>, vector<4xf32>, f32 + %next1 = vector.dotf %input_values, %weight_values1, %accumulator1 : vector<4xf32>, vector<4xf32>, f32 + %next_accumulators = vector.from_elements %next0, %next1 : vector<2xf32> + scf.yield %next_accumulators : vector<2xf32> + } + %sums = kernel.subgroup.reduce %lane_sums : vector<2xf32> + scf.if %writes_projection { + vector.store %sums, %logits_view[%token, %expert_base] : vector<2xf32>, view<[%token_count]x[%expert_count]xf32> + } + template.return +} + +template.def<@qwen3_moe.router.fused.projection> device requires [#target.subgroup.size<64>] priority(1) @qwen3_moe_router_fused_four_expert_projection(%hidden_size: index, %expert_count: index, %token_count: index, %token: index, %expert_tile: index, %lane: index, %writes_projection: i1, %input: buffer, %weight: buffer, %logits: buffer) { + %c0_offset = index.constant 0 : offset + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c256 = index.constant 256 : index + %c0_f32 = scalar.constant 0.0 : f32 + %expert_base = index.mul %expert_tile, %c4 : index + %expert1 = index.add %expert_base, %c1 : index + %expert2 = index.add %expert_base, %c2 : index + %expert3 = index.add %expert_base, %c3 : index + %lane_channel = index.mul %lane, %c4 : index + %input_view = buffer.view %input[%c0_offset] : buffer -> view<[%token_count]x[%hidden_size]xf32> + %weight_view = buffer.view %weight[%c0_offset] : buffer -> view<[%expert_count]x[%hidden_size]xf32> + %logits_view = buffer.view %logits[%c0_offset] : buffer -> view<[%token_count]x[%expert_count]xf32> + %c0_f32x4 = vector.splat %c0_f32 : vector<4xf32> + %lane_sums = scf.for %channel = [%lane_channel to %hidden_size step %c256](%accumulators = %c0_f32x4 : vector<4xf32>) -> (vector<4xf32>) unroll(%c4) schedule(interleaved) { + %input_values = vector.load %input_view[%token, %channel] : view<[%token_count]x[%hidden_size]xf32> -> vector<4xf32> + %weight_values0 = vector.load %weight_view[%expert_base, %channel] : view<[%expert_count]x[%hidden_size]xf32> -> vector<4xf32> + %weight_values1 = vector.load %weight_view[%expert1, %channel] : view<[%expert_count]x[%hidden_size]xf32> -> vector<4xf32> + %weight_values2 = vector.load %weight_view[%expert2, %channel] : view<[%expert_count]x[%hidden_size]xf32> -> vector<4xf32> + %weight_values3 = vector.load %weight_view[%expert3, %channel] : view<[%expert_count]x[%hidden_size]xf32> -> vector<4xf32> + %accumulator0 = vector.extract %accumulators[0] : vector<4xf32> -> f32 + %accumulator1 = vector.extract %accumulators[1] : vector<4xf32> -> f32 + %accumulator2 = vector.extract %accumulators[2] : vector<4xf32> -> f32 + %accumulator3 = vector.extract %accumulators[3] : vector<4xf32> -> f32 + %next0 = vector.dotf %input_values, %weight_values0, %accumulator0 : vector<4xf32>, vector<4xf32>, f32 + %next1 = vector.dotf %input_values, %weight_values1, %accumulator1 : vector<4xf32>, vector<4xf32>, f32 + %next2 = vector.dotf %input_values, %weight_values2, %accumulator2 : vector<4xf32>, vector<4xf32>, f32 + %next3 = vector.dotf %input_values, %weight_values3, %accumulator3 : vector<4xf32>, vector<4xf32>, f32 + %next_accumulators = vector.from_elements %next0, %next1, %next2, %next3 : vector<4xf32> + scf.yield %next_accumulators : vector<4xf32> + } + %sums = kernel.subgroup.reduce %lane_sums : vector<4xf32> + scf.if %writes_projection { + vector.store %sums, %logits_view[%token, %expert_base] : vector<4xf32>, view<[%token_count]x[%expert_count]xf32> + } + template.return +} + +kernel.def target(@qwen3_moe_router_fused_gfx11_wave64) @qwen3_moe_router_projection_top8_fused_decode_f32(%token_count: index, %route_id_stride: index) { + %expert_count = config.get @qwen3_moe.router.expert_count : index + %c1 = index.constant 1 : index + %wave_size = target.subgroup.size : index + %experts_per_wave = template.apply<@qwen3_moe.router.fused.experts_per_wave>() pure : () -> (index) + %expert_tiles = index.div %expert_count, %experts_per_wave : index + kernel.launch.config workgroups(%expert_tiles, %c1, %c1) workgroup_size(%wave_size, %c1, %c1) : index +} launch(%token_count: index, %route_id_stride: index, %input: buffer, %weight: buffer, %logits: buffer, %completion_counter: buffer, %route_ids: buffer, %route_weights: buffer) { + %hidden_size0 = config.get @qwen3_moe.model.hidden_size : index + %expert_count0 = config.get @qwen3_moe.router.expert_count : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 1)] : index + %bounded_route_id_stride = index.assume %route_id_stride [range(%route_id_stride, 1, 512)] : index + %hidden_size, %expert_count = index.assume %hidden_size0, %expert_count0 [range(%hidden_size0, 128, 32768), mul(%hidden_size0, 128), range(%expert_count0, 64, 512), mul(%expert_count0, 64)] : index, index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %expert_tile_count0 = kernel.workgroup.count : index + %expert_tile_count = index.assume %expert_tile_count0 [range(%expert_tile_count0, 16, 256)] : index + %expert_tile0 = kernel.workgroup.id : index + %expert_tile = index.assume %expert_tile0 [lt(%expert_tile0, %expert_tile_count)] : index + %lane = kernel.subgroup.lane.id : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %publishes_row = scalar.constant true : i1 + %c0_offset = index.constant 0 : offset + %counter_scratch_bytes = index.constant 4 : offset + %token = index.assume %c0 [lt(%c0, %bounded_token_count)] : index + %input_storage, %weight_storage, %logits_storage = template.apply<@qwen3_moe.router_projection.storage>(%input, %weight, %logits) : (buffer, buffer, buffer) -> (buffer, buffer, buffer) + %input_noalias, %weight_noalias, %logits_noalias, %completion_counter_noalias, %route_ids_noalias, %route_weights_noalias = buffer.assume.noalias %input_storage, %weight_storage, %logits_storage, %completion_counter, %route_ids, %route_weights : buffer, buffer, buffer, buffer, buffer, buffer + %completion_counter_aligned = buffer.assume.alignment %completion_counter_noalias {minimum_alignment = 16} : buffer + %completion_counter_view = buffer.view %completion_counter_aligned[%c0_offset] : buffer -> view<1xi32> + %counter_scratch = buffer.alloca align(4) %counter_scratch_bytes : buffer + %counter_scratch_view = buffer.view %counter_scratch[%c0_offset] : buffer -> view<1xi32> + + %lane_i32 = index.cast %lane : index to i32 + %writes_projection = scalar.cmpi eq, %lane_i32, %c0_i32 : i32 + template.apply<@qwen3_moe.router.fused.projection>(%hidden_size, %expert_count, %bounded_token_count, %token, %expert_tile, %lane, %writes_projection, %input_noalias, %weight_noalias, %logits_noalias) : (index, index, index, index, index, index, i1, buffer, buffer, buffer) + + // Every projection store precedes its workgroup's release. The last arrival + // acquires all preceding releases before any lane loads the complete row. + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.if %writes_projection { + %old_counter = view.atomic.rmw %c1_i32, %completion_counter_view[%c0] {ordering = acq_rel, scope = device} : i32, view<1xi32> -> i32 + view.store %old_counter, %counter_scratch_view[%c0] : i32, view<1xi32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %old_counter = view.load %counter_scratch_view[%c0] : view<1xi32> -> i32 + %expert_tile_count_i32 = index.cast %expert_tile_count : index to i32 + %last_expert_tile_i32 = scalar.subi %expert_tile_count_i32, %c1_i32 : i32 + %negative_expert_tile_count_i32 = scalar.subi %c0_i32, %expert_tile_count_i32 : i32 + %is_last_projection = scalar.cmpi eq, %old_counter, %last_expert_tile_i32 : i32 + scf.if %is_last_projection { + template.apply<@qwen3_moe.router.top8.row>(%publishes_row, %token, %bounded_token_count, %bounded_route_id_stride, %logits_noalias, %route_ids_noalias, %route_weights_noalias) : (i1, index, index, index, buffer, buffer, buffer) + // The counter cannot become reusable until all route stores complete. + kernel.barrier scope(workgroup) ordering(acq_rel) + scf.if %writes_projection { + view.atomic.reduce %negative_expert_tile_count_i32, %completion_counter_view[%c0] {ordering = release, scope = device} : i32, view<1xi32> + } + } + kernel.return +} + +// Compare the fused output against the production composition, then invoke the +// fused route again against the same counter to make reset correctness visible. +check.case public @qwen3_moe_router_projection_top8_fused_differential_case { + %token_count = check.literal value(1) : index + %route_id_stride = check.literal value(8) : index + %input = check.generate.iota offset(-0.25) step(0.0078125) period(17) : tensor<1x2048xf32> + %weight = check.generate.iota offset(-0.5) step(0.015625) period(31) : tensor<128x2048xf32> + %expected_logits = check.generate.fill value(0.0) : tensor<1x128xf32> + %expected_route_ids = check.generate.fill value(-1) : tensor<1x8xi32> + %expected_route_weights = check.generate.fill value(0.0) : tensor<1x8xf32> + %actual_logits0 = check.generate.fill value(1.0) : tensor<1x128xf32> + %actual_route_ids0 = check.generate.fill value(-1) : tensor<1x8xi32> + %actual_route_weights0 = check.generate.fill value(0.0) : tensor<1x8xf32> + %actual_logits1 = check.generate.fill value(1.0) : tensor<1x128xf32> + %actual_route_ids1 = check.generate.fill value(-1) : tensor<1x8xi32> + %actual_route_weights1 = check.generate.fill value(0.0) : tensor<1x8xf32> + %completion_counter = check.generate.fill value(0) : tensor<1xi32> + %expected_counter = check.generate.fill value(0) : tensor<1xi32> + kernel.launch @qwen3_moe_router_projection_f32_four_row_wave32[%token_count](%token_count, %input, %weight, %expected_logits) : [index](index, tensor<1x2048xf32>, tensor<128x2048xf32>, tensor<1x128xf32>) + kernel.launch @qwen3_moe_router_top8_f32[%token_count, %route_id_stride](%token_count, %route_id_stride, %expected_logits, %expected_route_ids, %expected_route_weights) : [index, index](index, index, tensor<1x128xf32>, tensor<1x8xi32>, tensor<1x8xf32>) + kernel.launch @qwen3_moe_router_projection_top8_fused_decode_f32[%token_count, %route_id_stride](%token_count, %route_id_stride, %input, %weight, %actual_logits0, %completion_counter, %actual_route_ids0, %actual_route_weights0) : [index, index](index, index, tensor<1x2048xf32>, tensor<128x2048xf32>, tensor<1x128xf32>, tensor<1xi32>, tensor<1x8xi32>, tensor<1x8xf32>) + kernel.launch @qwen3_moe_router_projection_top8_fused_decode_f32[%token_count, %route_id_stride](%token_count, %route_id_stride, %input, %weight, %actual_logits1, %completion_counter, %actual_route_ids1, %actual_route_weights1) : [index, index](index, index, tensor<1x2048xf32>, tensor<128x2048xf32>, tensor<1x128xf32>, tensor<1xi32>, tensor<1x8xi32>, tensor<1x8xf32>) + check.expect.close actual(%actual_logits0) expected(%expected_logits) atol(0.001) rtol(0.001) nan(same) : tensor<1x128xf32> + check.expect.close actual(%actual_logits1) expected(%expected_logits) atol(0.001) rtol(0.001) nan(same) : tensor<1x128xf32> + check.expect.equal actual(%actual_route_ids0) expected(%expected_route_ids) : tensor<1x8xi32> + check.expect.equal actual(%actual_route_ids1) expected(%expected_route_ids) : tensor<1x8xi32> + check.expect.close actual(%actual_route_weights0) expected(%expected_route_weights) atol(0.0001) rtol(0.0001) nan(same) : tensor<1x8xf32> + check.expect.close actual(%actual_route_weights1) expected(%expected_route_weights) atol(0.0001) rtol(0.0001) nan(same) : tensor<1x8xf32> + check.expect.equal actual(%completion_counter) expected(%expected_counter) : tensor<1xi32> + check.return +} + +check.case public @qwen3_moe_router_projection_top8_fused_benchmark_case { + %token_count = check.literal value(1) : index + %route_id_stride = check.literal value(8) : index + %input = check.generate.fill value(0.00390625) : tensor<1x2048xf32> + %weight = check.generate.fill value(0.015625) : tensor<128x2048xf32> + %logits = check.generate.fill value(0.0) : tensor<1x128xf32> + %completion_counter = check.generate.fill value(0) : tensor<1xi32> + %route_ids = check.generate.fill value(-1) : tensor<1x8xi32> + %route_weights = check.generate.fill value(0.0) : tensor<1x8xf32> + kernel.launch @qwen3_moe_router_projection_top8_fused_decode_f32[%token_count, %route_id_stride](%token_count, %route_id_stride, %input, %weight, %logits, %completion_counter, %route_ids, %route_weights) : [index, index](index, index, tensor<1x2048xf32>, tensor<128x2048xf32>, tensor<1x128xf32>, tensor<1xi32>, tensor<1x8xi32>, tensor<1x8xf32>) + check.return +} + +check.case public @qwen3_moe_router_projection_top8_composed_benchmark_case { + %token_count = check.literal value(1) : index + %route_id_stride = check.literal value(8) : index + %input = check.generate.fill value(0.00390625) : tensor<1x2048xf32> + %weight = check.generate.fill value(0.015625) : tensor<128x2048xf32> + %logits = check.generate.fill value(0.0) : tensor<1x128xf32> + %route_ids = check.generate.fill value(-1) : tensor<1x8xi32> + %route_weights = check.generate.fill value(0.0) : tensor<1x8xf32> + kernel.launch @qwen3_moe_router_projection_f32_four_row_wave32[%token_count](%token_count, %input, %weight, %logits) : [index](index, tensor<1x2048xf32>, tensor<128x2048xf32>, tensor<1x128xf32>) + kernel.launch @qwen3_moe_router_top8_f32[%token_count, %route_id_stride](%token_count, %route_id_stride, %logits, %route_ids, %route_weights) : [index, index](index, index, tensor<1x128xf32>, tensor<1x8xi32>, tensor<1x8xf32>) + check.return +} + +check.benchmark<@qwen3_moe_router_projection_top8_fused_benchmark_case> @qwen3_moe_router_projection_top8_fused_decode + +check.benchmark<@qwen3_moe_router_projection_top8_composed_benchmark_case> @qwen3_moe_router_projection_top8_composed_decode diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/router_top8_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/router_top8_f32.loom new file mode 100644 index 000000000000..ae4113298146 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/qwen_moe/qwen3_moe/router_top8_f32.loom @@ -0,0 +1,228 @@ +// Deterministic Qwen MoE routing from F32 expert logits. +// +// One wave owns one token row. Lanes load adjacent expert packets, repeatedly +// select the largest remaining logit, and break equal-logit ties in favor of +// the lower expert ordinal. Route IDs use an explicit physical row stride so +// the same kernel can publish either compact `[token][route]` rows or the +// `[token][expert]` argsort storage exposed by GGML views. +// +// Qwen applies a full expert softmax, selects the top-k probabilities, and +// renormalizes those selected probabilities. The full-softmax denominator +// cancels during renormalization: +// +// (exp(x_i) / sum_all) / sum_topk(exp(x_j) / sum_all) +// = exp(x_i) / sum_topk(exp(x_j)) +// +// Selection therefore operates directly on logits and only the selected +// values reach the exponential. This preserves the model contract while +// avoiding exponentials for experts that cannot contribute to the result. +template.decl @qwen3_moe.router.top8.row(%valid_token: i1, %token: index, %token_count: index, %route_id_stride: index, %logits: buffer, %route_ids: buffer, %route_weights: buffer) + +amdgpu.target @qwen3_moe_router_gfx11_wave64 {subgroup_size = 64} + +config.decl @qwen3_moe.router.expert_count : %value: index where [range(%value, 32, 512), mul(%value, 32)] + +config.decl @qwen3_moe.router.route_count : %value: index where [range(%value, 1, 32)] + +// Selects and normalizes one logical router row. The caller owns mapping a +// subgroup to a safe physical row and tells the helper whether that row should +// publish. Keeping row semantics independent of launch geometry lets fused +// producers consume this exact tie-breaking and normalization contract. +template.def<@qwen3_moe.router.top8.row> device requires [#target.subgroup.size<64>] @qwen3_moe_router_top8_row_f32(%valid_token: i1, %token: index, %token_count: index, %route_id_stride: index, %logits: buffer, %route_ids: buffer, %route_weights: buffer) { + %expert_count0 = config.get @qwen3_moe.router.expert_count : index + %route_count0 = config.get @qwen3_moe.router.route_count : index + %bounded_token_count = index.assume %token_count [range(%token_count, 1, 2048)] : index + %bounded_route_id_stride0 = index.assume %route_id_stride [range(%route_id_stride, 1, 512)] : index + %expert_count, %route_count, %bounded_route_id_stride = index.assume %expert_count0, %route_count0, %bounded_route_id_stride0 [range(%expert_count0, 32, 512), mul(%expert_count0, 32), range(%route_count0, 1, 32), le(%route_count0, %expert_count0), le(%route_count0, %bounded_route_id_stride0)] : index, index, index + %safe_token, %launch_token_count = index.assume %token, %bounded_token_count [lt(%token, %bounded_token_count)] : index, index + %lane = kernel.subgroup.lane.id : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %wave_size = target.subgroup.size : index + %negative_large = scalar.constant -3.4028234663852886e+38 : f32 + %largest_i32 = scalar.constant 2147483647 : i32 + %c0_offset = index.constant 0 : offset + %experts_per_lane0 = index.div %expert_count, %wave_size : index + %experts_per_lane = index.assume %experts_per_lane0 [range(%experts_per_lane0, 1, 16)] : index + %lane_expert_base = index.mul %lane, %experts_per_lane : index + %route_id_storage_count = index.mul %launch_token_count, %bounded_route_id_stride : index + %route_weight_storage_count = index.mul %launch_token_count, %route_count : index + %logits_noalias, %route_ids_noalias, %route_weights_noalias = buffer.assume.noalias %logits, %route_ids, %route_weights : buffer, buffer, buffer + %logits_view = buffer.view %logits_noalias[%c0_offset] : buffer -> view<[%launch_token_count]x[%expert_count]xf32> + %route_ids_view = buffer.view %route_ids_noalias[%c0_offset] : buffer -> view<[%route_id_storage_count]xi32> + %route_weights_view = buffer.view %route_weights_noalias[%c0_offset] : buffer -> view<[%route_weight_storage_count]xf32> + %initial_logits = vector.load %logits_view[%safe_token, %lane_expert_base] : view<[%launch_token_count]x[%expert_count]xf32> -> vector<[%experts_per_lane]xf32> + // Seed the per-lane argmax with this lane's first expert id instead of the + // 0x7FFFFFFF sentinel. If every candidate is unordered (NaN) or below + // -FLT_MAX, `ogt`/`oeq` never fire and the old seed was published verbatim + // into route_ids, where the consumer turned it into a wild weight address + // (engine#123). lane_expert_base is always < expert_count for a real lane. + %lane_expert_base_i32 = index.cast %lane_expert_base : index to i32 + %remaining_final, %selected_logit = scf.for %route = [%c0 to %route_count step %c1](%remaining_logits = %initial_logits : vector<[%experts_per_lane]xf32>, %lane_selected_logit = %negative_large : f32) -> (vector<[%experts_per_lane]xf32>, f32) { + %local_value, %local_id = scf.for %slot = [%c0 to %experts_per_lane step %c1](%best_value = %negative_large : f32, %best_id = %lane_expert_base_i32 : i32) -> (f32, i32) unroll { + %candidate_value = vector.extract %remaining_logits[%slot] : vector<[%experts_per_lane]xf32> -> f32 + %candidate_expert = index.add %lane_expert_base, %slot : index + %candidate_id = index.cast %candidate_expert : index to i32 + %is_greater = scalar.cmpf ogt, %candidate_value, %best_value : f32 + %is_equal = scalar.cmpf oeq, %candidate_value, %best_value : f32 + %is_lower_id = scalar.cmpi ult, %candidate_id, %best_id : i32 + %is_lower_tie = scalar.andi %is_equal, %is_lower_id : i1 + %is_better = scalar.ori %is_greater, %is_lower_tie : i1 + %next_value, %next_id = scf.if %is_better -> (f32, i32) { + scf.yield %candidate_value, %candidate_id : f32, i32 + } else { + scf.yield %best_value, %best_id : f32, i32 + } + scf.yield %next_value, %next_id : f32, i32 + } + %winner_value = kernel.subgroup.reduce %local_value : f32 + %matches_winner = scalar.cmpf oeq, %local_value, %winner_value : f32 + %winner_lane_mask = kernel.subgroup.vote.ballot %matches_winner : i1 -> i64 + %nonzero_winner_lane_mask = scalar.assume %winner_lane_mask [ne(%winner_lane_mask, 0)] : i64 + %winner_lane_i64 = scalar.cttzi %nonzero_winner_lane_mask : i64 + %winner_lane0 = index.cast %winner_lane_i64 : i64 to index + %winner_lane = index.assume %winner_lane0 [range(%winner_lane0, 0, 63)] : index + %local_expert0 = index.cast %local_id : i32 to index + %local_expert = index.assume %local_expert0 [range(%local_expert0, 0, 511)] : index + %winner_slot = index.rem %local_expert, %experts_per_lane : index + %owns_winner = index.cmp eq, %lane, %winner_lane : index + %next_remaining_logits = scf.if %owns_winner -> (vector<[%experts_per_lane]xf32>) { + %removed = vector.insert %negative_large into %remaining_logits[%winner_slot] : f32, vector<[%experts_per_lane]xf32> + scf.yield %removed : vector<[%experts_per_lane]xf32> + } else { + scf.yield %remaining_logits : vector<[%experts_per_lane]xf32> + } + %publishes_route_id = scalar.andi %valid_token, %owns_winner : i1 + scf.if %publishes_route_id { + %route_id_token_base = index.mul %safe_token, %bounded_route_id_stride : index + %route_id_index = index.add %route_id_token_base, %route : index + view.store %local_id, %route_ids_view[%route_id_index] : i32, view<[%route_id_storage_count]xi32> + } + %lane_publishes_route = index.cmp eq, %lane, %route : index + %publishes_route = scalar.andi %valid_token, %lane_publishes_route : i1 + %next_lane_selected_logit = scf.if %publishes_route -> (f32) { + scf.yield %winner_value : f32 + } else { + scf.yield %lane_selected_logit : f32 + } + scf.yield %next_remaining_logits, %next_lane_selected_logit : vector<[%experts_per_lane]xf32>, f32 + } + %selected_max = kernel.subgroup.reduce %selected_logit : f32 + %selected_delta = scalar.subf %selected_logit, %selected_max : f32 + %unnormalized_weight = scalar.expf %selected_delta : f32 + %selected_sum = kernel.subgroup.reduce %unnormalized_weight : f32 + %route_weight = scalar.divf %unnormalized_weight, %selected_sum : f32 + // engine#123: when no lane of the row found an ordered candidate, the row's router logits + // are entirely NaN. The argmax seed is a finite -FLT_MAX, so the softmax below still + // produces a plausible uniform 1/route_count, which routes the token to a wrong-but-valid + // expert and hides the corruption completely. Publish the row's own NaN instead, so the + // corruption reaches the host logits and is caught loudly rather than silently. + %row_had_ordered = scalar.cmpf ogt, %selected_max, %negative_large : f32 + %raw_candidate = vector.extract %initial_logits[%c0] : vector<[%experts_per_lane]xf32> -> f32 + %published_weight = scf.select %row_had_ordered, %route_weight, %raw_candidate : f32 + %lane_publishes_weight = index.cmp ult, %lane, %route_count : index + %publishes_weight = scalar.andi %valid_token, %lane_publishes_weight : i1 + scf.if %publishes_weight { + %route_weight_token_base = index.mul %safe_token, %route_count : index + %route_weight_index = index.add %route_weight_token_base, %lane : index + view.store %published_weight, %route_weights_view[%route_weight_index] : f32, view<[%route_weight_storage_count]xf32> + } + template.return +} + +kernel.def target(@qwen3_moe_router_gfx11_wave64) @qwen3_moe_router_top8_f32(%token_count: index, %route_id_stride: index) { + %c1 = index.constant 1 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %subgroup_size = target.subgroup.size : index + %workgroup_size = index.mul %subgroup_size, %c4 : index + %padded_token_count = index.add %token_count, %c3 : index + %workgroup_count = index.div %padded_token_count, %c4 : index + kernel.launch.config workgroups(%workgroup_count, %c1, %c1) workgroup_size(%workgroup_size, %c1, %c1) : index +} launch(%token_count: index, %route_id_stride: index, %logits: buffer, %route_ids: buffer, %route_weights: buffer) where [range(%token_count, 1, 2048), range(%route_id_stride, 1, 512)] { + %token_workgroup = kernel.workgroup.id : index + %subgroup = kernel.subgroup.id : index + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %token_base = index.mul %token_workgroup, %c4 : index + %token0 = index.add %token_base, %subgroup : index + %valid_token = index.cmp ult, %token0, %token_count : index + %safe_token = scf.select %valid_token, %token0, %c0 : index + template.apply<@qwen3_moe.router.top8.row>(%valid_token, %safe_token, %token_count, %route_id_stride, %logits, %route_ids, %route_weights) : (i1, index, index, index, buffer, buffer, buffer) + kernel.return +} + +// Builds the deterministic route-ID oracle for two physical rows with eight +// published IDs and eight untouched padding slots per row. +kernel.def @qwen3_moe_router_top8_wide_stride_reference() { + %c1 = index.constant 1 : index + %c32 = index.constant 32 : index + kernel.launch.config workgroups(%c1, %c1, %c1) workgroup_size(%c32, %c1, %c1) : index +} launch(%route_ids: buffer) { + %element = kernel.workitem.id : index + %c7 = index.constant 7 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c0_offset = index.constant 0 : offset + %slot = index.rem %element, %c16 : index + %publishes_id = index.cmp ult, %slot, %c8 : index + %route_ids_view = buffer.view %route_ids[%c0_offset] : buffer -> view<32xi32> + scf.if %publishes_id { + %scaled_slot = index.mul %slot, %c8 : index + %route_id_index = index.add %scaled_slot, %c7 : index + %route_id = index.cast %route_id_index : index to i32 + view.store %route_id, %route_ids_view[%element] : i32, view<32xi32> + } + kernel.return +} + +// Period-eight logits place sixteen equal maxima in different lane-local +// slots. Repeated selection must remove each winner and retain the lower-ID +// half of that tie set. Two rows leave two inactive waves in the four-wave +// production workgroup and exercise guarded tail publication. A second call +// proves that a wider physical route-ID stride preserves its padding. +check.case public @qwen3_moe_router_top8_f32_repeated_maxima_case { + %token_count = check.literal value(2) : index + %route_id_stride = check.literal value(8) : index + %wide_route_id_stride = check.literal value(16) : index + %logits = check.generate.iota offset(0.0) step(1.0) period(8) : tensor<2x128xf32> + %route_ids = check.generate.fill value(-1) : tensor<2x8xi32> + %route_weights = check.generate.fill value(0.0) : tensor<2x8xf32> + %wide_route_ids = check.generate.fill value(-1) : tensor<2x16xi32> + %wide_route_weights = check.generate.fill value(0.0) : tensor<2x8xf32> + %expected_ids = check.generate.iota offset(7) step(8) period(8) : tensor<2x8xi32> + %expected_wide_ids = check.generate.fill value(-1) : tensor<2x16xi32> + %expected_weights = check.generate.fill value(0.125) : tensor<2x8xf32> + kernel.launch @qwen3_moe_router_top8_f32[%token_count, %route_id_stride](%token_count, %route_id_stride, %logits, %route_ids, %route_weights) : [index, index](index, index, tensor<2x128xf32>, tensor<2x8xi32>, tensor<2x8xf32>) + kernel.launch @qwen3_moe_router_top8_wide_stride_reference(%expected_wide_ids) : (tensor<2x16xi32>) + kernel.launch @qwen3_moe_router_top8_f32[%token_count, %wide_route_id_stride](%token_count, %wide_route_id_stride, %logits, %wide_route_ids, %wide_route_weights) : [index, index](index, index, tensor<2x128xf32>, tensor<2x16xi32>, tensor<2x8xf32>) + check.expect.equal actual(%route_ids) expected(%expected_ids) : tensor<2x8xi32> + check.expect.close actual(%route_weights) expected(%expected_weights) atol(1.0000000000000001e-05) rtol(1.0000000000000001e-05) nan(same) : tensor<2x8xf32> + check.expect.equal actual(%wide_route_ids) expected(%expected_wide_ids) : tensor<2x16xi32> + check.expect.close actual(%wide_route_weights) expected(%expected_weights) atol(1.0000000000000001e-05) rtol(1.0000000000000001e-05) nan(same) : tensor<2x8xf32> + check.return +} + +check.case public @qwen3_moe_router_top8_f32_benchmark_case { + %token_count = check.param.choice values([1, 32, 128, 512]) name("token_count") : index + %route_id_stride = check.literal value(8) : index + %logits = check.generate.iota offset(0.0) step(1.0) period(8) : tensor<[%token_count]x128xf32> + %route_ids = check.generate.fill value(-1) : tensor<[%token_count]x8xi32> + %route_weights = check.generate.fill value(0.0) : tensor<[%token_count]x8xf32> + %expected_ids = check.generate.iota offset(7) step(8) period(8) : tensor<[%token_count]x8xi32> + %expected_weights = check.generate.fill value(0.125) : tensor<[%token_count]x8xf32> + kernel.launch @qwen3_moe_router_top8_f32[%token_count, %route_id_stride](%token_count, %route_id_stride, %logits, %route_ids, %route_weights) : [index, index](index, index, tensor<[%token_count]x128xf32>, tensor<[%token_count]x8xi32>, tensor<[%token_count]x8xf32>) + check.expect.equal actual(%route_ids) expected(%expected_ids) : tensor<[%token_count]x8xi32> + check.expect.close actual(%route_weights) expected(%expected_weights) atol(1.0000000000000001e-05) rtol(1.0000000000000001e-05) nan(same) : tensor<[%token_count]x8xf32> + check.return +} + +check.benchmark<@qwen3_moe_router_top8_f32_repeated_maxima_case> @qwen3_moe_router_top8_f32_repeated_maxima + +check.benchmark<@qwen3_moe_router_top8_f32_benchmark_case> @qwen3_moe_router_top8_f32_decode {token_count = 1} + +check.benchmark<@qwen3_moe_router_top8_f32_benchmark_case> @qwen3_moe_router_top8_f32_prefill_32 {token_count = 32} + +check.benchmark<@qwen3_moe_router_top8_f32_benchmark_case> @qwen3_moe_router_top8_f32_prefill_128 {token_count = 128} + +check.benchmark<@qwen3_moe_router_top8_f32_benchmark_case> @qwen3_moe_router_top8_f32_prefill_512 {token_count = 512} diff --git a/ggml/src/ggml-hrx/loom-jit.cpp b/ggml/src/ggml-hrx/loom-jit.cpp new file mode 100644 index 000000000000..38698b026518 --- /dev/null +++ b/ggml/src/ggml-hrx/loom-jit.cpp @@ -0,0 +1,1068 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +#include "loom-jit.h" + +#include "loomc/launch_config.h" +#include "loomc/loomc.h" +#include "loomc/sanitizer.h" +#include "loomc/target/amdgpu.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +ggml_hrx_loom_jit_compile_result::~ggml_hrx_loom_jit_compile_result() { + reset(); +} + +ggml_hrx_loom_jit_compile_result::ggml_hrx_loom_jit_compile_result(ggml_hrx_loom_jit_compile_result && other) noexcept { + *this = std::move(other); +} + +ggml_hrx_loom_jit_compile_result & ggml_hrx_loom_jit_compile_result::operator=( + ggml_hrx_loom_jit_compile_result && other) noexcept { + if (this == &other) { + return *this; + } + reset(); + hsaco_data = std::exchange(other.hsaco_data, nullptr); + hsaco_size = std::exchange(other.hsaco_size, 0); + manifest_json = std::exchange(other.manifest_json, nullptr); + manifest_json_size = std::exchange(other.manifest_json_size, 0); + compile_report_json = std::exchange(other.compile_report_json, nullptr); + compile_report_json_size = std::exchange(other.compile_report_json_size, 0); + final_module_text = std::exchange(other.final_module_text, nullptr); + final_module_text_size = std::exchange(other.final_module_text_size, 0); + launch_config = std::exchange(other.launch_config, {}); + return *this; +} + +void ggml_hrx_loom_jit_compile_result::reset() { + hrx_host_allocator_t allocator = hrx_host_allocator_system(); + hrx_host_allocator_free(allocator, hsaco_data); + hrx_host_allocator_free(allocator, manifest_json); + hrx_host_allocator_free(allocator, compile_report_json); + hrx_host_allocator_free(allocator, final_module_text); + hsaco_data = nullptr; + hsaco_size = 0; + manifest_json = nullptr; + manifest_json_size = 0; + compile_report_json = nullptr; + compile_report_json_size = 0; + final_module_text = nullptr; + final_module_text_size = 0; + launch_config = {}; +} + +struct ggml_hrx_loom_jit_amdgpu { + loomc_target_environment_t * target_environment = nullptr; + loomc_context_t * context = nullptr; + loomc_target_profile_t * target_profile = nullptr; + loomc_compiler_t * compiler = nullptr; + loomc_pass_program_t * pass_program = nullptr; + loomc_amdgpu_runtime_global_flags_t runtime_globals = LOOMC_AMDGPU_RUNTIME_GLOBAL_NONE; +}; + +namespace { + +// LoomC currently accepts concrete workload values for launch evaluation but +// not for compilation. Until that becomes one operation in LoomC, specialize +// the pinned textual kernel source before parsing so body optimization sees +// the same exact facts as launch evaluation. The public kernel ABI is retained: +// scalar arguments remain present but their uses in both kernel regions are +// replaced by backend-authored constants. This is deliberately local to the +// HRX backend and its emitted module text is always available for inspection. +static bool ggml_hrx_loom_specialize_workload_text(const std::string & input, + const char * root_symbol, + const int64_t * workload_arguments, + size_t workload_argument_count, + std::string & output, + std::string & error) { + output = input; + if (workload_argument_count == 0) { + return true; + } + std::string symbol = root_symbol ? root_symbol : ""; + if (!symbol.empty() && symbol.front() == '@') { + symbol.erase(symbol.begin()); + } + const std::string marker = "@" + symbol + "("; + const size_t symbol_position = output.find(marker); + if (symbol_position == std::string::npos) { + error = "workload specialization cannot find root " + marker; + return false; + } + const size_t parameter_begin = symbol_position + marker.size(); + const size_t parameter_end = output.find(')', parameter_begin); + if (parameter_end == std::string::npos) { + error = "workload specialization cannot parse root parameters"; + return false; + } + const std::string parameters = output.substr(parameter_begin, parameter_end - parameter_begin); + std::vector names; + size_t cursor = 0; + while (names.size() < workload_argument_count) { + const size_t percent = parameters.find('%', cursor); + if (percent == std::string::npos) { + break; + } + size_t name_end = percent + 1; + while (name_end < parameters.size() && + (std::isalnum(static_cast(parameters[name_end])) || parameters[name_end] == '_')) { + ++name_end; + } + const size_t colon = parameters.find(':', name_end); + if (colon == std::string::npos) { + break; + } + const size_t type_begin = parameters.find_first_not_of(" \t", colon + 1); + if (type_begin != std::string::npos && parameters.compare(type_begin, 5, "index") == 0) { + names.push_back(parameters.substr(percent + 1, name_end - percent - 1)); + } + cursor = name_end; + } + if (names.size() != workload_argument_count) { + error = "workload specialization argument count does not match root index parameters"; + return false; + } + + auto matching_brace = [&](size_t open) -> size_t { + size_t depth = 0; + for (size_t i = open; i < output.size(); ++i) { + if (output[i] == '{') { + ++depth; + } + if (output[i] == '}' && --depth == 0) { + return i; + } + } + return std::string::npos; + }; + std::vector> regions; + const size_t config_open = output.find('{', parameter_end); + const size_t config_close = config_open == std::string::npos ? std::string::npos : matching_brace(config_open); + if (config_close == std::string::npos) { + error = "workload specialization cannot find kernel config region"; + return false; + } + regions.push_back({ config_open, config_close }); + const size_t launch = output.find("launch(", config_close + 1); + const size_t launch_open = launch == std::string::npos ? std::string::npos : output.find('{', launch); + const size_t launch_close = launch_open == std::string::npos ? std::string::npos : matching_brace(launch_open); + if (launch_close == std::string::npos) { + error = "workload specialization cannot find kernel launch region"; + return false; + } + regions.push_back({ launch_open, launch_close }); + + for (auto region = regions.rbegin(); region != regions.rend(); ++region) { + std::string body = output.substr(region->first + 1, region->second - region->first - 1); + std::string prefix; + for (size_t i = 0; i < names.size(); ++i) { + const std::string original = "%" + names[i]; + const std::string specialized = "%ggml_hrx_specialized_" + names[i]; + size_t use = 0; + while ((use = body.find(original, use)) != std::string::npos) { + const size_t after = use + original.size(); + if (after == body.size() || + (!std::isalnum(static_cast(body[after])) && body[after] != '_')) { + body.replace(use, original.size(), specialized); + use += specialized.size(); + } else { + use = after; + } + } + prefix += "\n " + specialized + " = index.constant " + std::to_string(workload_arguments[i]) + " : index"; + } + output.replace(region->first + 1, region->second - region->first - 1, prefix + body); + } + return true; +} + +template class LoomHandle { + public: + LoomHandle() = default; + LoomHandle(const LoomHandle &) = delete; + LoomHandle & operator=(const LoomHandle &) = delete; + + ~LoomHandle() { reset(); } + + T * get() const { return value_; } + + T ** out() { + reset(); + return &value_; + } + + void reset(T * value = nullptr) { + if (value_) { + Release(value_); + } + value_ = value; + } + + private: + T * value_ = nullptr; +}; + +using LoomWorkspace = LoomHandle; +using LoomSource = LoomHandle; +using LoomModule = LoomHandle; +using LoomResult = LoomHandle; +using LoomLinkIndexBuilder = LoomHandle; +using LoomLinkIndex = LoomHandle; +using LoomLinker = LoomHandle; +using LoomLaunchConfigProgram = LoomHandle; + +struct HrxLoomJitDeleter { + void operator()(ggml_hrx_loom_jit_amdgpu * jit) const { ggml_hrx_loom_jit_amdgpu_release(jit); } +}; + +hrx_status_t ggml_hrx_loom_jit_make_status(hrx_status_code_t code, const char * message) { + return hrx_make_status(code, message ? message : "GGML HRX Loom JIT failure"); +} + +hrx_status_code_t ggml_hrx_loom_jit_status_code_from_loom(loomc_status_code_t code) { + switch (code) { + case LOOMC_STATUS_OK: + return HRX_STATUS_OK; + case LOOMC_STATUS_CANCELLED: + return HRX_STATUS_CANCELLED; + case LOOMC_STATUS_UNKNOWN: + return HRX_STATUS_UNKNOWN; + case LOOMC_STATUS_INVALID_ARGUMENT: + return HRX_STATUS_INVALID_ARGUMENT; + case LOOMC_STATUS_DEADLINE_EXCEEDED: + return HRX_STATUS_DEADLINE_EXCEEDED; + case LOOMC_STATUS_NOT_FOUND: + return HRX_STATUS_NOT_FOUND; + case LOOMC_STATUS_ALREADY_EXISTS: + return HRX_STATUS_ALREADY_EXISTS; + case LOOMC_STATUS_PERMISSION_DENIED: + return HRX_STATUS_PERMISSION_DENIED; + case LOOMC_STATUS_RESOURCE_EXHAUSTED: + return HRX_STATUS_OUT_OF_MEMORY; + case LOOMC_STATUS_FAILED_PRECONDITION: + return HRX_STATUS_FAILED_PRECONDITION; + case LOOMC_STATUS_ABORTED: + return HRX_STATUS_ABORTED; + case LOOMC_STATUS_OUT_OF_RANGE: + return HRX_STATUS_OUT_OF_RANGE; + case LOOMC_STATUS_UNIMPLEMENTED: + return HRX_STATUS_UNIMPLEMENTED; + case LOOMC_STATUS_INTERNAL: + return HRX_STATUS_INTERNAL; + case LOOMC_STATUS_UNAVAILABLE: + return HRX_STATUS_UNAVAILABLE; + case LOOMC_STATUS_DATA_LOSS: + return HRX_STATUS_DATA_LOSS; + case LOOMC_STATUS_UNAUTHENTICATED: + return HRX_STATUS_PERMISSION_DENIED; + case LOOMC_STATUS_DEFERRED: + return HRX_STATUS_UNAVAILABLE; + case LOOMC_STATUS_INCOMPATIBLE: + return HRX_STATUS_FAILED_PRECONDITION; + case LOOMC_STATUS_CODE_MASK: + return HRX_STATUS_INTERNAL; + } + return HRX_STATUS_INTERNAL; +} + +std::string ggml_hrx_loom_jit_format_status(loomc_status_t status) { + if (loomc_status_is_ok(status)) { + return "OK"; + } + loomc_host_size_t length = 0; + loomc_status_format(status, 0, nullptr, &length); + std::unique_ptr buffer(new (std::nothrow) char[length + 1]()); + if (!buffer) { + char fallback[4096] = { 0 }; + loomc_host_size_t fallback_length = 0; + loomc_status_format(status, sizeof(fallback), fallback, &fallback_length); + if (fallback_length >= sizeof(fallback)) { + fallback_length = sizeof(fallback) - 1; + } + return std::string(fallback, fallback_length); + } + loomc_host_size_t actual_length = 0; + loomc_status_format(status, length + 1, buffer.get(), &actual_length); + return std::string(buffer.get(), actual_length); +} + +void ggml_hrx_loom_jit_spam_failure(const char * context, const std::string & message) { + std::fprintf(stderr, "HRX Loom JIT %s failed: %s\n", context ? context : "operation", message.c_str()); + std::fflush(stderr); +} + +hrx_status_t ggml_hrx_loom_jit_status_from_loom(loomc_status_t status, const char * context) { + if (loomc_status_is_ok(status)) { + return hrx_ok_status(); + } + const hrx_status_code_t hrx_code = ggml_hrx_loom_jit_status_code_from_loom(loomc_status_code(status)); + const std::string formatted_status = ggml_hrx_loom_jit_format_status(status); + loomc_status_free(status); + std::string message = std::string(context ? context : "loomc") + ": " + formatted_status; + ggml_hrx_loom_jit_spam_failure(context, message); + return ggml_hrx_loom_jit_make_status(hrx_code, message.c_str()); +} + +hrx_status_t ggml_hrx_loom_jit_status_from_result(const loomc_result_t * result, const char * context) { + if (result && loomc_result_succeeded(result)) { + return hrx_ok_status(); + } + std::string message = context ? context : "loomc result failed"; + if (result) { + const loomc_host_size_t diagnostic_count = loomc_result_diagnostic_count(result); + for (loomc_host_size_t i = 0; i < diagnostic_count; ++i) { + const loomc_diagnostic_t * diagnostic = loomc_result_diagnostic_at(result, i); + if (!diagnostic) { + continue; + } + message += "\n diagnostic["; + message += std::to_string(static_cast(i)); + message += "] "; + message.append(diagnostic->code.data, diagnostic->code.size); + message += ": "; + message.append(diagnostic->message.data, diagnostic->message.size); + if (diagnostic->range.start_line || diagnostic->range.start_column) { + message += " @ "; + message += std::to_string(diagnostic->range.start_line); + message += ":"; + message += std::to_string(diagnostic->range.start_column); + } + } + } + ggml_hrx_loom_jit_spam_failure(context, message); + return ggml_hrx_loom_jit_make_status(HRX_STATUS_FAILED_PRECONDITION, message.c_str()); +} + +void * ggml_hrx_loom_jit_malloc_copy(const void * data, size_t size, bool nul_terminate) { + if (!data || size == 0) { + return nullptr; + } + const size_t alloc_size = nul_terminate ? size + 1 : size; + void * result = nullptr; + hrx_status_t status = hrx_host_allocator_malloc_uninitialized(hrx_host_allocator_system(), alloc_size, &result); + if (!hrx_status_is_ok(status)) { + hrx_status_ignore(status); + return nullptr; + } + std::memcpy(result, data, size); + if (nul_terminate) { + static_cast(result)[size] = 0; + } + return result; +} + +#if defined(LOOMC_ARTIFACT_ROLE_COMPILE_REPORT) +using ggml_hrx_loom_jit_artifact_selector_t = loomc_string_view_t; +#define GGML_HRX_LOOM_ARTIFACT_COMPILE_REPORT loomc_make_cstring_view(LOOMC_ARTIFACT_ROLE_COMPILE_REPORT) +#define GGML_HRX_LOOM_ARTIFACT_MODULE loomc_make_cstring_view(LOOMC_ARTIFACT_ROLE_MODULE) +#define GGML_HRX_LOOM_ARTIFACT_LAUNCH_CONFIG loomc_make_cstring_view(LOOMC_ARTIFACT_ROLE_LAUNCH_CONFIG) +#define GGML_HRX_LOOM_ARTIFACT_KERNEL loomc_make_cstring_view(LOOMC_ARTIFACT_ROLE_KERNEL) +#define GGML_HRX_LOOM_ARTIFACT_MANIFEST loomc_make_cstring_view(LOOMC_ARTIFACT_ROLE_ARTIFACT_MANIFEST) +static bool ggml_hrx_loom_jit_artifact_matches(const loomc_artifact_t * artifact, + ggml_hrx_loom_jit_artifact_selector_t selector) { + return loomc_string_view_equal(artifact->role, selector); +} +#else +using ggml_hrx_loom_jit_artifact_selector_t = loomc_artifact_kind_t; +#define GGML_HRX_LOOM_ARTIFACT_COMPILE_REPORT LOOMC_ARTIFACT_KIND_REPORT +#define GGML_HRX_LOOM_ARTIFACT_MODULE LOOMC_ARTIFACT_KIND_MODULE +#define GGML_HRX_LOOM_ARTIFACT_LAUNCH_CONFIG LOOMC_ARTIFACT_KIND_LAUNCH_CONFIG +#define GGML_HRX_LOOM_ARTIFACT_KERNEL LOOMC_ARTIFACT_KIND_EXECUTABLE +#define GGML_HRX_LOOM_ARTIFACT_MANIFEST LOOMC_ARTIFACT_KIND_REPORT +static bool ggml_hrx_loom_jit_artifact_matches(const loomc_artifact_t * artifact, + ggml_hrx_loom_jit_artifact_selector_t selector) { + return artifact->kind == selector; +} +#endif + +const loomc_artifact_t * ggml_hrx_loom_jit_find_artifact(const loomc_result_t * result, + ggml_hrx_loom_jit_artifact_selector_t selector, + loomc_string_view_t format) { + for (loomc_host_size_t i = 0; i < loomc_result_artifact_count(result); ++i) { + const loomc_artifact_t * artifact = loomc_result_artifact_at(result, i); + if (!artifact) { + continue; + } + if (ggml_hrx_loom_jit_artifact_matches(artifact, selector) && + loomc_string_view_equal(artifact->format, format)) { + return artifact; + } + } + return nullptr; +} + +hrx_status_t ggml_hrx_loom_jit_copy_artifact_bytes(const loomc_artifact_t * artifact, + void ** out_data, + size_t * out_size, + bool nul_terminate) { + if (out_data) { + *out_data = nullptr; + } + if (out_size) { + *out_size = 0; + } + if (!artifact || !out_data || !out_size) { + return hrx_ok_status(); + } + + loomc_byte_span_t contents = loomc_byte_span_empty(); + const bool contents_borrowed = loomc_byte_sequence_try_get_contiguous_span(artifact->contents, &contents); + if (!contents_borrowed) { + loomc_status_t status = loomc_byte_sequence_clone(artifact->contents, loomc_allocator_system(), &contents); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "copy Loom artifact"); + } + } + + void * copy = ggml_hrx_loom_jit_malloc_copy(contents.data, contents.data_length, nul_terminate); + if (!contents_borrowed) { + loomc_allocator_free(loomc_allocator_system(), const_cast(contents.data)); + } + if (!copy) { + return ggml_hrx_loom_jit_make_status(HRX_STATUS_OUT_OF_MEMORY, "failed to copy Loom artifact"); + } + *out_data = copy; + *out_size = contents.data_length; + return hrx_ok_status(); +} + +hrx_status_t ggml_hrx_loom_jit_evaluate_launch_config(const loomc_artifact_t * artifact, + const char * root_symbol, + const int64_t * workload_arguments, + size_t workload_argument_count, + ggml_hrx_loom_jit_launch_config * out_launch_config) { + if (!out_launch_config) { + return ggml_hrx_loom_jit_make_status(HRX_STATUS_INVALID_ARGUMENT, "out_launch_config is required"); + } + if (!artifact) { + return ggml_hrx_loom_jit_make_status(HRX_STATUS_NOT_FOUND, "Loom did not return a launch-config artifact"); + } + + LoomLaunchConfigProgram program; + loomc_status_t status = loomc_launch_config_program_load(artifact, loomc_allocator_system(), program.out()); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "load Loom launch config program"); + } + + std::string export_name = root_symbol ? root_symbol : ""; + if (!export_name.empty() && export_name.front() == '@') { + export_name.erase(export_name.begin()); + } + loomc_launch_config_function_t function = loomc_launch_config_function_invalid(); + status = loomc_launch_config_program_lookup_function(program.get(), loomc_make_cstring_view(export_name.c_str()), + &function); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "find Loom launch config function"); + } + + std::vector workload_bits; + workload_bits.reserve(workload_argument_count); + for (size_t i = 0; i < workload_argument_count; ++i) { + workload_bits.push_back(static_cast(workload_arguments[i])); + } + + loomc_launch_config_t launch_config = {}; + launch_config.type = LOOMC_STRUCTURE_TYPE_LAUNCH_CONFIG; + launch_config.structure_size = sizeof(launch_config); + + status = loomc_launch_config_program_invoke(program.get(), function, + workload_bits.empty() ? nullptr : workload_bits.data(), + workload_bits.size(), &launch_config); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "invoke Loom launch config function"); + } + if (!launch_config.workgroup_count.x || !launch_config.workgroup_count.y || !launch_config.workgroup_count.z || + !launch_config.workgroup_size.x || !launch_config.workgroup_size.y || !launch_config.workgroup_size.z) { + return ggml_hrx_loom_jit_make_status(HRX_STATUS_FAILED_PRECONDITION, + "Loom launch config did not provide required workgroup count and size"); + } + + out_launch_config->fields = 0; + out_launch_config->workgroup_count[0] = launch_config.workgroup_count.x; + out_launch_config->workgroup_count[1] = launch_config.workgroup_count.y; + out_launch_config->workgroup_count[2] = launch_config.workgroup_count.z; + out_launch_config->workgroup_size[0] = launch_config.workgroup_size.x; + out_launch_config->workgroup_size[1] = launch_config.workgroup_size.y; + out_launch_config->workgroup_size[2] = launch_config.workgroup_size.z; + out_launch_config->subgroup_size = launch_config.subgroup_size; + out_launch_config->workgroup_storage_bytes = launch_config.workgroup_storage_bytes; + out_launch_config->workload_argument_count = workload_argument_count; + return hrx_ok_status(); +} + +hrx_status_t ggml_hrx_loom_jit_parse_sanitizer_checks(const char * value, loomc_sanitizer_checks_t * out_checks) { + *out_checks = 0; + if (!value || !value[0] || std::strcmp(value, "0") == 0 || std::strcmp(value, "none") == 0) { + return hrx_ok_status(); + } + if (std::strcmp(value, "access") == 0 || std::strcmp(value, "asan") == 0) { + *out_checks = LOOMC_SANITIZER_CHECKS_ASAN_LIKE; + return hrx_ok_status(); + } + if (std::strcmp(value, "value") == 0) { + *out_checks = LOOMC_SANITIZER_CHECK_VALUE; + return hrx_ok_status(); + } + if (std::strcmp(value, "operation") == 0) { + *out_checks = LOOMC_SANITIZER_CHECK_OPERATION; + return hrx_ok_status(); + } + if (std::strcmp(value, "ubsan") == 0) { + *out_checks = LOOMC_SANITIZER_CHECKS_UBSAN_LIKE; + return hrx_ok_status(); + } + if (std::strcmp(value, "all") == 0) { + *out_checks = LOOMC_SANITIZER_CHECK_ACCESS | LOOMC_SANITIZER_CHECK_VALUE | LOOMC_SANITIZER_CHECK_OPERATION; + return hrx_ok_status(); + } + char message[256] = { 0 }; + std::snprintf(message, sizeof(message), "unsupported GGML_HRX_LOOM_SANITIZER '%s'", value); + return ggml_hrx_loom_jit_make_status(HRX_STATUS_INVALID_ARGUMENT, message); +} + +hrx_status_t ggml_hrx_loom_jit_parse_sanitizer_reporting(const char * value, + loomc_sanitizer_reporting_mode_t * out_reporting_mode) { + *out_reporting_mode = LOOMC_SANITIZER_REPORTING_MODE_REPORT_ONLY; + if (!value || !value[0] || std::strcmp(value, "report-only") == 0 || std::strcmp(value, "report_only") == 0 || + std::strcmp(value, "report") == 0) { + return hrx_ok_status(); + } + if (std::strcmp(value, "default") == 0) { + *out_reporting_mode = LOOMC_SANITIZER_REPORTING_MODE_DEFAULT; + return hrx_ok_status(); + } + if (std::strcmp(value, "trap") == 0) { + *out_reporting_mode = LOOMC_SANITIZER_REPORTING_MODE_TRAP; + return hrx_ok_status(); + } + char message[256] = { 0 }; + std::snprintf(message, sizeof(message), "unsupported GGML_HRX_LOOM_SANITIZER_REPORTING '%s'", value); + return ggml_hrx_loom_jit_make_status(HRX_STATUS_INVALID_ARGUMENT, message); +} + +loomc_amdgpu_runtime_global_flags_t ggml_hrx_loom_jit_runtime_globals(loomc_sanitizer_checks_t sanitizer_checks) { + if (!sanitizer_checks) { + return LOOMC_AMDGPU_RUNTIME_GLOBAL_NONE; + } + loomc_amdgpu_runtime_global_flags_t runtime_globals = LOOMC_AMDGPU_RUNTIME_GLOBAL_FEEDBACK_CONFIG; + if (sanitizer_checks & LOOMC_SANITIZER_CHECK_ACCESS) { + runtime_globals |= LOOMC_AMDGPU_RUNTIME_GLOBAL_ASAN_CONFIG; + } + return runtime_globals; +} + +} // namespace + +hrx_status_t ggml_hrx_loom_jit_amdgpu_create(const ggml_hrx_loom_jit_amdgpu_options * options, + ggml_hrx_loom_jit_amdgpu ** out_jit) { + if (!out_jit) { + return ggml_hrx_loom_jit_make_status(HRX_STATUS_INVALID_ARGUMENT, "out_jit must not be NULL"); + } + *out_jit = nullptr; + if (!options || !options->processor || options->processor[0] == 0) { + return ggml_hrx_loom_jit_make_status(HRX_STATUS_INVALID_ARGUMENT, + "valid ggml_hrx_loom_jit_amdgpu_options_t with processor is required"); + } + + std::unique_ptr jit(new (std::nothrow) ggml_hrx_loom_jit_amdgpu()); + if (!jit) { + return ggml_hrx_loom_jit_make_status(HRX_STATUS_OUT_OF_MEMORY, "failed to allocate GGML HRX Loom JIT"); + } + + LoomResult result; + loomc_status_t status = loomc_target_environment_create_amdgpu(loomc_allocator_system(), &jit->target_environment); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "create AMDGPU target environment"); + } + + loomc_context_target_options_t target_options = {}; + target_options.type = LOOMC_STRUCTURE_TYPE_CONTEXT_TARGET_OPTIONS; + target_options.structure_size = sizeof(target_options); + target_options.target_environment = jit->target_environment; + loomc_context_options_t context_options = {}; + context_options.type = LOOMC_STRUCTURE_TYPE_CONTEXT_OPTIONS; + context_options.structure_size = sizeof(context_options); + context_options.next = &target_options; + status = loomc_context_create(&context_options, loomc_allocator_system(), &jit->context); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "create Loom context"); + } + + loomc_amdgpu_profile_options_t profile_options = {}; + profile_options.type = LOOMC_STRUCTURE_TYPE_AMDGPU_PROFILE_OPTIONS; + profile_options.structure_size = sizeof(profile_options); + profile_options.identifier = loomc_make_cstring_view(options->identifier); + profile_options.identity.target = loomc_make_cstring_view(options->processor); + status = loomc_target_profile_create_amdgpu(jit->target_environment, &profile_options, loomc_allocator_system(), + &jit->target_profile); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "create AMDGPU target profile"); + } + status = loomc_compiler_create(jit->context, nullptr, loomc_allocator_system(), &jit->compiler); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "create Loom compiler"); + } + + loomc_sanitizer_options_t sanitizer_options = {}; + sanitizer_options.type = LOOMC_STRUCTURE_TYPE_SANITIZER_OPTIONS; + sanitizer_options.structure_size = sizeof(sanitizer_options); + sanitizer_options.next = nullptr; + hrx_status_t sanitizer_status = + ggml_hrx_loom_jit_parse_sanitizer_checks(options->sanitizer, &sanitizer_options.checks); + if (!hrx_status_is_ok(sanitizer_status)) { + return sanitizer_status; + } + if (sanitizer_options.checks) { + hrx_status_t sanitizer_reporting_status = ggml_hrx_loom_jit_parse_sanitizer_reporting( + options->sanitizer_reporting, &sanitizer_options.reporting_mode); + if (!hrx_status_is_ok(sanitizer_reporting_status)) { + return sanitizer_reporting_status; + } + } + jit->runtime_globals = ggml_hrx_loom_jit_runtime_globals(sanitizer_options.checks); + loomc_target_pipeline_options_t pipeline_options = {}; + pipeline_options.type = LOOMC_STRUCTURE_TYPE_TARGET_PIPELINE_OPTIONS; + pipeline_options.structure_size = sizeof(pipeline_options); + pipeline_options.next = sanitizer_options.checks ? static_cast(&sanitizer_options) : nullptr; + pipeline_options.identifier = loomc_make_cstring_view("ggml-hrx-amdgpu-jit-prepared-low"); + pipeline_options.kind = LOOMC_TARGET_PIPELINE_KIND_PREPARED_LOW; + pipeline_options.control_flow_lowering = LOOMC_TARGET_CONTROL_FLOW_LOWERING_CFG; + pipeline_options.source_to_low_max_errors = 20; + status = loomc_pass_program_create_from_target_pipeline(jit->context, &pipeline_options, loomc_allocator_system(), + &jit->pass_program, result.out()); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "create target pass program"); + } + if (!loomc_result_succeeded(result.get())) { + return ggml_hrx_loom_jit_status_from_result(result.get(), "target pass program failed"); + } + + *out_jit = jit.release(); + return hrx_ok_status(); +} + +void ggml_hrx_loom_jit_amdgpu_release(ggml_hrx_loom_jit_amdgpu * jit) { + if (!jit) { + return; + } + loomc_pass_program_release(jit->pass_program); + loomc_compiler_release(jit->compiler); + loomc_target_profile_release(jit->target_profile); + loomc_context_release(jit->context); + loomc_target_environment_release(jit->target_environment); + delete jit; +} + +hrx_status_t ggml_hrx_loom_jit_amdgpu_compile(ggml_hrx_loom_jit_amdgpu * jit, + const ggml_hrx_loom_jit_compile_options * options, + ggml_hrx_loom_jit_compile_result * out_result) { + if (!out_result) { + return ggml_hrx_loom_jit_make_status(HRX_STATUS_INVALID_ARGUMENT, "out_result must not be NULL"); + } + out_result->reset(); + if (!jit || !options || !options->source_data || options->source_size == 0 || !options->root_symbol || + options->root_symbol[0] == 0) { + return ggml_hrx_loom_jit_make_status( + HRX_STATUS_INVALID_ARGUMENT, "valid GGML HRX Loom JIT compile options with source and root are required"); + } + if (options->config_binding_count > 0 && !options->config_bindings) { + return ggml_hrx_loom_jit_make_status(HRX_STATUS_INVALID_ARGUMENT, + "GGML HRX Loom JIT config binding count requires config bindings"); + } + if (options->workload_argument_count > 0 && !options->workload_arguments) { + return ggml_hrx_loom_jit_make_status(HRX_STATUS_INVALID_ARGUMENT, + "GGML HRX Loom JIT workload argument count requires workload arguments"); + } + if (options->dependency_count > 0 && !options->dependencies) { + return ggml_hrx_loom_jit_make_status(HRX_STATUS_INVALID_ARGUMENT, + "GGML HRX Loom JIT dependency count requires dependencies"); + } + for (size_t i = 0; i < options->config_binding_count; ++i) { + if (!options->config_bindings[i].key || !options->config_bindings[i].value) { + return ggml_hrx_loom_jit_make_status(HRX_STATUS_INVALID_ARGUMENT, + "GGML HRX Loom JIT config binding keys and values must not be NULL"); + } + } + + std::unique_ptr config_bindings; + if (options->config_binding_count > 0) { + config_bindings.reset(new (std::nothrow) loomc_config_binding_t[options->config_binding_count]()); + if (!config_bindings) { + return ggml_hrx_loom_jit_make_status(HRX_STATUS_OUT_OF_MEMORY, + "failed to allocate GGML HRX Loom JIT config bindings"); + } + for (size_t i = 0; i < options->config_binding_count; ++i) { + config_bindings[i].key = loomc_make_cstring_view(options->config_bindings[i].key); + config_bindings[i].value = loomc_make_cstring_view(options->config_bindings[i].value); + } + } + + LoomWorkspace workspace; + LoomSource source; + LoomModule module; + LoomResult result; + std::string specialized_source; + if (options->source_format == GGML_HRX_LOOM_JIT_SOURCE_FORMAT_TEXT && options->dependency_count == 0 && + options->workload_argument_count > 0) { + std::string specialization_error; + const std::string source_text(static_cast(options->source_data), options->source_size); + if (!ggml_hrx_loom_specialize_workload_text(source_text, options->root_symbol, options->workload_arguments, + options->workload_argument_count, specialized_source, + specialization_error)) { + return ggml_hrx_loom_jit_make_status(HRX_STATUS_FAILED_PRECONDITION, specialization_error.c_str()); + } + } + + loomc_status_t status = loomc_workspace_create(nullptr, loomc_allocator_system(), workspace.out()); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "create Loom workspace"); + } + + loomc_source_options_t source_options = {}; + source_options.type = LOOMC_STRUCTURE_TYPE_SOURCE_OPTIONS; + source_options.structure_size = sizeof(source_options); + source_options.format = options->source_format == GGML_HRX_LOOM_JIT_SOURCE_FORMAT_BYTECODE ? + LOOMC_SOURCE_FORMAT_BYTECODE : + LOOMC_SOURCE_FORMAT_TEXT; + source_options.identifier = loomc_make_cstring_view(options->source_identifier); + source_options.contents = specialized_source.empty() ? + loomc_make_byte_span(options->source_data, options->source_size) : + loomc_make_byte_span(specialized_source.data(), specialized_source.size()); + source_options.storage = LOOMC_SOURCE_STORAGE_BORROWED; + status = loomc_source_create(&source_options, loomc_allocator_system(), source.out()); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "create Loom source"); + } + + LoomLinkIndexBuilder link_index_builder; + status = loomc_link_index_builder_create(jit->context, nullptr, loomc_allocator_system(), link_index_builder.out()); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "create Loom link index builder"); + } + loomc_link_index_source_options_t link_source_options = {}; + link_source_options.provider_name = loomc_make_cstring_view(options->source_identifier); + link_source_options.role = LOOMC_LINK_PROVIDER_ROLE_INPUT; + status = loomc_link_index_builder_add_source(link_index_builder.get(), source.get(), &link_source_options, nullptr); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "index Loom source"); + } + std::vector dependency_sources; + dependency_sources.reserve(options->dependency_count); + for (size_t i = 0; i < options->dependency_count; ++i) { + const ggml_hrx_loom_jit_source & dependency = options->dependencies[i]; + if (!dependency.source_data || dependency.source_size == 0 || !dependency.source_identifier) { + for (loomc_source_t * dependency_source : dependency_sources) { + loomc_source_release(dependency_source); + } + return ggml_hrx_loom_jit_make_status(HRX_STATUS_INVALID_ARGUMENT, "invalid Loom JIT dependency source"); + } + loomc_source_options_t dependency_options = {}; + dependency_options.type = LOOMC_STRUCTURE_TYPE_SOURCE_OPTIONS; + dependency_options.structure_size = sizeof(dependency_options); + dependency_options.format = dependency.source_format == GGML_HRX_LOOM_JIT_SOURCE_FORMAT_BYTECODE ? + LOOMC_SOURCE_FORMAT_BYTECODE : + LOOMC_SOURCE_FORMAT_TEXT; + dependency_options.identifier = loomc_make_cstring_view(dependency.source_identifier); + dependency_options.contents = loomc_make_byte_span(dependency.source_data, dependency.source_size); + dependency_options.storage = LOOMC_SOURCE_STORAGE_BORROWED; + loomc_source_t * dependency_source = nullptr; + status = loomc_source_create(&dependency_options, loomc_allocator_system(), &dependency_source); + if (!loomc_status_is_ok(status)) { + for (loomc_source_t * retained_source : dependency_sources) { + loomc_source_release(retained_source); + } + return ggml_hrx_loom_jit_status_from_loom(status, "create Loom dependency source"); + } + dependency_sources.push_back(dependency_source); + loomc_link_index_source_options_t dependency_link_options = {}; + dependency_link_options.provider_name = loomc_make_cstring_view(dependency.source_identifier); + dependency_link_options.role = LOOMC_LINK_PROVIDER_ROLE_INPUT; + status = loomc_link_index_builder_add_source(link_index_builder.get(), dependency_source, + &dependency_link_options, nullptr); + if (!loomc_status_is_ok(status)) { + for (loomc_source_t * retained_source : dependency_sources) { + loomc_source_release(retained_source); + } + return ggml_hrx_loom_jit_status_from_loom(status, "index Loom dependency source"); + } + } + LoomLinkIndex link_index; + status = loomc_link_index_builder_finish(link_index_builder.get(), link_index.out(), result.out()); + for (loomc_source_t * dependency_source : dependency_sources) { + loomc_source_release(dependency_source); + } + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "finish Loom link index"); + } + if (!loomc_result_succeeded(result.get())) { + return ggml_hrx_loom_jit_status_from_result(result.get(), "Loom source indexing failed"); + } + result.reset(); + + LoomLinker linker; + status = loomc_linker_create(jit->context, nullptr, loomc_allocator_system(), linker.out()); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "create Loom linker"); + } + + // BUILD-authored kernel recipes first archive their primary source and + // libraries, then compile roots from that linked module. Preserve that + // composition exactly. In particular, config.decl operations are not + // callable dependency edges and disappear if raw source libraries are + // fed directly to a selective link. Bytecode sources with workload + // arguments also use this path so the linked module can be serialized + // back to text for the current workload specialization pass. + LoomSource archived_source; + LoomSource specialized_archive_source; + std::string specialized_archive_text; + const bool needs_archive_source = + options->dependency_count > 0 || + (options->source_format == GGML_HRX_LOOM_JIT_SOURCE_FORMAT_BYTECODE && options->workload_argument_count > 0); + if (needs_archive_source) { + LoomModule archive_module; + loomc_link_options_t archive_options = {}; + archive_options.type = LOOMC_STRUCTURE_TYPE_LINK_OPTIONS; + archive_options.structure_size = sizeof(archive_options); + archive_options.link_index = link_index.get(); + archive_options.module_name = loomc_make_cstring_view(options->module_name); + archive_options.flags = LOOMC_LINK_FLAG_STRIP_TEST_SYMBOLS; + status = loomc_link_module(linker.get(), workspace.get(), &archive_options, archive_module.out(), result.out()); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "archive Loom kernel module"); + } + if (!loomc_result_succeeded(result.get())) { + return ggml_hrx_loom_jit_status_from_result(result.get(), "Loom kernel archive linking failed"); + } + result.reset(); + loomc_module_serialize_options_t serialize_options = {}; + serialize_options.type = LOOMC_STRUCTURE_TYPE_MODULE_SERIALIZE_OPTIONS; + serialize_options.structure_size = sizeof(serialize_options); + serialize_options.format = LOOMC_SOURCE_FORMAT_TEXT; + serialize_options.identifier = loomc_make_cstring_view(options->source_identifier); + status = loomc_module_serialize_to_source(archive_module.get(), &serialize_options, loomc_allocator_system(), + archived_source.out()); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "serialize Loom kernel archive"); + } + loomc_source_t * archive_index_source = archived_source.get(); + if (options->workload_argument_count > 0) { + const loomc_byte_span_t archive_contents = loomc_source_contents(archived_source.get()); + const std::string archive_text(reinterpret_cast(archive_contents.data), + archive_contents.data_length); + std::string specialization_error; + if (!ggml_hrx_loom_specialize_workload_text(archive_text, options->root_symbol, options->workload_arguments, + options->workload_argument_count, specialized_archive_text, + specialization_error)) { + return ggml_hrx_loom_jit_make_status(HRX_STATUS_FAILED_PRECONDITION, specialization_error.c_str()); + } + loomc_source_options_t specialized_options = {}; + specialized_options.type = LOOMC_STRUCTURE_TYPE_SOURCE_OPTIONS; + specialized_options.structure_size = sizeof(specialized_options); + specialized_options.format = LOOMC_SOURCE_FORMAT_TEXT; + specialized_options.identifier = loomc_make_cstring_view(options->source_identifier); + specialized_options.contents = + loomc_make_byte_span(specialized_archive_text.data(), specialized_archive_text.size()); + specialized_options.storage = LOOMC_SOURCE_STORAGE_BORROWED; + status = + loomc_source_create(&specialized_options, loomc_allocator_system(), specialized_archive_source.out()); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "create specialized Loom kernel archive"); + } + archive_index_source = specialized_archive_source.get(); + } + LoomLinkIndexBuilder archive_index_builder; + status = loomc_link_index_builder_create(jit->context, nullptr, loomc_allocator_system(), + archive_index_builder.out()); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "create Loom kernel archive index"); + } + loomc_link_index_source_options_t archive_source_options = {}; + archive_source_options.provider_name = loomc_make_cstring_view(options->source_identifier); + archive_source_options.role = LOOMC_LINK_PROVIDER_ROLE_INPUT; + status = loomc_link_index_builder_add_source(archive_index_builder.get(), archive_index_source, + &archive_source_options, nullptr); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "index Loom kernel archive"); + } + status = loomc_link_index_builder_finish(archive_index_builder.get(), link_index.out(), result.out()); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "finish Loom kernel archive index"); + } + if (!loomc_result_succeeded(result.get())) { + return ggml_hrx_loom_jit_status_from_result(result.get(), "Loom kernel archive indexing failed"); + } + result.reset(); + } + const loomc_target_specialization_t specialization = { + loomc_make_cstring_view(options->root_symbol), + jit->target_profile, + }; + loomc_target_specialization_options_t compile_target_options = {}; + compile_target_options.type = LOOMC_STRUCTURE_TYPE_TARGET_SPECIALIZATION_OPTIONS; + compile_target_options.structure_size = sizeof(compile_target_options); + compile_target_options.specializations = &specialization; + compile_target_options.specialization_count = 1; + + const loomc_string_view_t root_symbols[] = { loomc_make_cstring_view(options->root_symbol) }; + loomc_link_options_t link_options = {}; + link_options.type = LOOMC_STRUCTURE_TYPE_LINK_OPTIONS; + link_options.structure_size = sizeof(link_options); + link_options.next = &compile_target_options; + link_options.mode = LOOMC_LINK_MODE_LINK; + link_options.link_index = link_index.get(); + link_options.module_name = loomc_make_cstring_view(options->module_name); + link_options.root_symbols = root_symbols; + link_options.root_symbol_count = 1; + link_options.flags = LOOMC_LINK_FLAG_STRIP_TEST_SYMBOLS; + link_options.config.bindings = config_bindings.get(); + link_options.config.binding_count = options->config_binding_count; + status = loomc_link_module(linker.get(), workspace.get(), &link_options, module.out(), result.out()); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "link Loom root"); + } + if (!loomc_result_succeeded(result.get())) { + return ggml_hrx_loom_jit_status_from_result(result.get(), "Loom root linking failed"); + } + result.reset(); + + loomc_compile_options_t compile_options = {}; + compile_options.type = LOOMC_STRUCTURE_TYPE_COMPILE_OPTIONS; + compile_options.structure_size = sizeof(compile_options); + compile_options.next = &compile_target_options; + compile_options.module_name = loomc_make_cstring_view(options->module_name); + compile_options.artifact_flags = LOOMC_COMPILE_ARTIFACT_FLAG_MODULE_TEXT | LOOMC_COMPILE_ARTIFACT_FLAG_REPORT_JSON; + if (options->evaluate_launch_config) { + compile_options.artifact_flags |= LOOMC_COMPILE_ARTIFACT_FLAG_LAUNCH_CONFIG; + } + compile_options.config_flags = LOOMC_CONFIG_POLICY_FLAG_REQUIRE_RESOLVED; + status = loomc_compile_module(jit->compiler, workspace.get(), jit->pass_program, module.get(), &compile_options, + loomc_allocator_system(), result.out()); + if (!loomc_status_is_ok(status)) { + return ggml_hrx_loom_jit_status_from_loom(status, "compile Loom module"); + } + if (!loomc_result_succeeded(result.get())) { + return ggml_hrx_loom_jit_status_from_result(result.get(), "Loom compilation failed"); + } + + hrx_status_t hrx_status = hrx_ok_status(); + const loomc_artifact_t * compile_report = + ggml_hrx_loom_jit_find_artifact(result.get(), GGML_HRX_LOOM_ARTIFACT_COMPILE_REPORT, + loomc_make_cstring_view(LOOMC_ARTIFACT_FORMAT_JSON)); + hrx_status = ggml_hrx_loom_jit_copy_artifact_bytes(compile_report, + reinterpret_cast(&out_result->compile_report_json), + &out_result->compile_report_json_size, true); + if (!hrx_status_is_ok(hrx_status)) { + return hrx_status; + } + const loomc_artifact_t * final_module = + ggml_hrx_loom_jit_find_artifact(result.get(), GGML_HRX_LOOM_ARTIFACT_MODULE, + loomc_make_cstring_view(LOOMC_ARTIFACT_FORMAT_LOOM_TEXT)); + hrx_status = + ggml_hrx_loom_jit_copy_artifact_bytes(final_module, reinterpret_cast(&out_result->final_module_text), + &out_result->final_module_text_size, true); + if (!hrx_status_is_ok(hrx_status)) { + return hrx_status; + } + if (options->evaluate_launch_config) { + const char * launch_config_symbol = options->launch_config_symbol; + if (launch_config_symbol == nullptr || launch_config_symbol[0] == '\0') { + launch_config_symbol = options->root_symbol; + } + const loomc_artifact_t * launch_config = + ggml_hrx_loom_jit_find_artifact(result.get(), GGML_HRX_LOOM_ARTIFACT_LAUNCH_CONFIG, + loomc_make_cstring_view(LOOMC_ARTIFACT_FORMAT_LOOM_BYTECODE)); + hrx_status = ggml_hrx_loom_jit_evaluate_launch_config( + launch_config, launch_config_symbol, + options->workload_argument_count == 0 ? nullptr : options->workload_arguments, + options->workload_argument_count, &out_result->launch_config); + if (!hrx_status_is_ok(hrx_status)) { + return hrx_status; + } + } + result.reset(); + + loomc_amdgpu_emit_options_t amdgpu_options = {}; + amdgpu_options.type = LOOMC_STRUCTURE_TYPE_AMDGPU_EMIT_OPTIONS; + amdgpu_options.structure_size = sizeof(amdgpu_options); + amdgpu_options.next = nullptr; + amdgpu_options.runtime_globals = jit->runtime_globals; + const loomc_option_entry_t emit_entries[] = { + { + loomc_make_cstring_view(LOOMC_EMIT_OPTION_KEY_IDENTIFIER), + loomc_make_cstring_view(options->artifact_identifier), + }, + }; + loomc_option_dict_t option_dict = {}; + option_dict.type = LOOMC_STRUCTURE_TYPE_OPTION_DICT; + option_dict.structure_size = sizeof(option_dict); + option_dict.next = &amdgpu_options; + option_dict.entries = emit_entries; + option_dict.entry_count = options->artifact_identifier ? 1 : 0; + loomc_artifact_manifest_options_t manifest_options = {}; + manifest_options.type = LOOMC_STRUCTURE_TYPE_ARTIFACT_MANIFEST_OPTIONS; + manifest_options.structure_size = sizeof(manifest_options); + manifest_options.next = &option_dict; + manifest_options.mode = LOOMC_ARTIFACT_MANIFEST_MODE_DETAILS; + loomc_compile_report_options_t report_options = {}; + report_options.type = LOOMC_STRUCTURE_TYPE_COMPILE_REPORT_OPTIONS; + report_options.structure_size = sizeof(report_options); + report_options.next = &manifest_options; + report_options.mode = LOOMC_COMPILE_REPORT_MODE_DETAILS; + loomc_emit_options_t emit_options = {}; + emit_options.type = LOOMC_STRUCTURE_TYPE_EMIT_OPTIONS; + emit_options.structure_size = sizeof(emit_options); + emit_options.next = &report_options; + emit_options.artifact_format = loomc_make_cstring_view(LOOMC_ARTIFACT_FORMAT_AMDGPU_HSACO); + emit_options.identifier = loomc_make_cstring_view(options->artifact_identifier); + emit_options.artifact_flags = LOOMC_EMIT_ARTIFACT_FLAG_PRIMARY; + status = loomc_emit_module(jit->target_environment, workspace.get(), module.get(), &emit_options, + loomc_allocator_system(), result.out()); + if (!loomc_status_is_ok(status)) { + out_result->reset(); + return ggml_hrx_loom_jit_status_from_loom(status, "emit AMDGPU HSACO"); + } + if (!loomc_result_succeeded(result.get())) { + hrx_status = ggml_hrx_loom_jit_status_from_result(result.get(), "AMDGPU HSACO emission failed"); + out_result->reset(); + return hrx_status; + } + + const loomc_artifact_t * hsaco = + ggml_hrx_loom_jit_find_artifact(result.get(), GGML_HRX_LOOM_ARTIFACT_KERNEL, + loomc_make_cstring_view(LOOMC_ARTIFACT_FORMAT_AMDGPU_HSACO)); + if (!hsaco) { + out_result->reset(); + return ggml_hrx_loom_jit_make_status(HRX_STATUS_NOT_FOUND, "Loom did not return an AMDGPU HSACO artifact"); + } + hrx_status = ggml_hrx_loom_jit_copy_artifact_bytes(hsaco, &out_result->hsaco_data, &out_result->hsaco_size, false); + if (hrx_status_is_ok(hrx_status)) { + const loomc_artifact_t * report = + ggml_hrx_loom_jit_find_artifact(result.get(), GGML_HRX_LOOM_ARTIFACT_COMPILE_REPORT, + loomc_make_cstring_view(LOOMC_ARTIFACT_FORMAT_COMPILE_REPORT_JSON)); + hrx_status = + ggml_hrx_loom_jit_copy_artifact_bytes(report, reinterpret_cast(&out_result->compile_report_json), + &out_result->compile_report_json_size, true); + } + if (hrx_status_is_ok(hrx_status)) { + const loomc_artifact_t * manifest = ggml_hrx_loom_jit_find_artifact( + result.get(), GGML_HRX_LOOM_ARTIFACT_MANIFEST, + loomc_make_cstring_view(LOOMC_ARTIFACT_FORMAT_ARTIFACT_MANIFEST_JSON)); + hrx_status = ggml_hrx_loom_jit_copy_artifact_bytes( + manifest, reinterpret_cast(&out_result->manifest_json), &out_result->manifest_json_size, true); + } + + if (!hrx_status_is_ok(hrx_status)) { + out_result->reset(); + } + return hrx_status; +} diff --git a/ggml/src/ggml-hrx/loom-jit.h b/ggml/src/ggml-hrx/loom-jit.h new file mode 100644 index 000000000000..4dd05f73519c --- /dev/null +++ b/ggml/src/ggml-hrx/loom-jit.h @@ -0,0 +1,95 @@ +// Copyright 2026 The HRX Authors +// SPDX-License-Identifier: Apache-2.0 + +#pragma once + +#include "hrx_runtime.h" + +#include +#include +#include + +struct ggml_hrx_loom_jit_amdgpu; + +enum class ggml_hrx_loom_jit_source_format { + Text, + Bytecode, +}; +inline constexpr auto GGML_HRX_LOOM_JIT_SOURCE_FORMAT_TEXT = ggml_hrx_loom_jit_source_format::Text; +inline constexpr auto GGML_HRX_LOOM_JIT_SOURCE_FORMAT_BYTECODE = ggml_hrx_loom_jit_source_format::Bytecode; + +struct ggml_hrx_loom_jit_amdgpu_options { + const char * processor = nullptr; + const char * identifier = nullptr; + const char * sanitizer = nullptr; + const char * sanitizer_reporting = nullptr; +}; + +struct ggml_hrx_loom_jit_config_binding { + const char * key = nullptr; + const char * value = nullptr; +}; + +struct ggml_hrx_loom_jit_source { + const void * source_data = nullptr; + size_t source_size = 0; + ggml_hrx_loom_jit_source_format source_format = ggml_hrx_loom_jit_source_format::Text; + const char * source_identifier = nullptr; +}; + +struct ggml_hrx_loom_jit_launch_config { + std::array workgroup_count = {}; + std::array workgroup_size = {}; + uint32_t subgroup_size = 0; + uint64_t workgroup_storage_bytes = 0; + size_t workload_argument_count = 0; + uint32_t fields = 0; +}; + +struct ggml_hrx_loom_jit_compile_options { + const void * source_data = nullptr; + size_t source_size = 0; + ggml_hrx_loom_jit_source_format source_format = ggml_hrx_loom_jit_source_format::Text; + const char * source_identifier = nullptr; + const char * root_symbol = nullptr; + const char * launch_config_symbol = nullptr; + const char * module_name = nullptr; + const char * artifact_identifier = nullptr; + const ggml_hrx_loom_jit_source * dependencies = nullptr; + size_t dependency_count = 0; + const ggml_hrx_loom_jit_config_binding * config_bindings = nullptr; + size_t config_binding_count = 0; + const int64_t * workload_arguments = nullptr; + size_t workload_argument_count = 0; + bool evaluate_launch_config = false; +}; + +struct ggml_hrx_loom_jit_compile_result { + ggml_hrx_loom_jit_compile_result() = default; + ~ggml_hrx_loom_jit_compile_result(); + ggml_hrx_loom_jit_compile_result(const ggml_hrx_loom_jit_compile_result &) = delete; + ggml_hrx_loom_jit_compile_result & operator=(const ggml_hrx_loom_jit_compile_result &) = delete; + ggml_hrx_loom_jit_compile_result(ggml_hrx_loom_jit_compile_result && other) noexcept; + ggml_hrx_loom_jit_compile_result & operator=(ggml_hrx_loom_jit_compile_result && other) noexcept; + + void reset(); + + void * hsaco_data = nullptr; + size_t hsaco_size = 0; + char * manifest_json = nullptr; + size_t manifest_json_size = 0; + char * compile_report_json = nullptr; + size_t compile_report_json_size = 0; + char * final_module_text = nullptr; + size_t final_module_text_size = 0; + ggml_hrx_loom_jit_launch_config launch_config; +}; + +hrx_status_t ggml_hrx_loom_jit_amdgpu_create(const ggml_hrx_loom_jit_amdgpu_options * options, + ggml_hrx_loom_jit_amdgpu ** out_jit); + +void ggml_hrx_loom_jit_amdgpu_release(ggml_hrx_loom_jit_amdgpu * jit); + +hrx_status_t ggml_hrx_loom_jit_amdgpu_compile(ggml_hrx_loom_jit_amdgpu * jit, + const ggml_hrx_loom_jit_compile_options * options, + ggml_hrx_loom_jit_compile_result * out_result); diff --git a/ggml/src/ggml-hrx/runtime/command-program-executor.cpp b/ggml/src/ggml-hrx/runtime/command-program-executor.cpp new file mode 100644 index 000000000000..4b98d0c5d9aa --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/command-program-executor.cpp @@ -0,0 +1,1997 @@ +#include "command-program-executor.h" + +#include "dispatch/command-program-diagnostics.h" +#include "dispatch/command-program-resolver.h" +#include "ggml-impl.h" +#include "hrx-interop-utils.h" +#include "runtime/graph-record-order.h" +#include "runtime/hrx-sleeping-wait.h" +#include "runtime/kernel-executable-cache.h" +#include "runtime/transient-arena.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { + +GraphReplayStreamState::~GraphReplayStreamState() { + clear(); +} + +void GraphReplayStreamState::clear() { + for (PendingHostWriteback & writeback : pending_host_writebacks_) { + if (writeback.retained_buffer != nullptr) { + hrx_buffer_release(writeback.retained_buffer); + writeback.retained_buffer = nullptr; + } + } + pending_host_writebacks_.clear(); +} + +void GraphReplayStreamState::mark_stream_synchronized() { + for (PendingHostWriteback & writeback : pending_host_writebacks_) { + if (writeback.host_destination != nullptr && writeback.mapped_source != nullptr && writeback.size > 0) { + std::memcpy(writeback.host_destination, writeback.mapped_source, writeback.size); + } + if (writeback.retained_buffer != nullptr) { + hrx_buffer_release(writeback.retained_buffer); + writeback.retained_buffer = nullptr; + } + } + pending_host_writebacks_.clear(); +} + +void GraphReplayStreamState::add_host_writeback(void * host_destination, + const void * mapped_source, + size_t size, + hrx_buffer_t buffer) { + if (buffer != nullptr) { + hrx_buffer_retain(buffer); + } + pending_host_writebacks_.push_back({ host_destination, mapped_source, size, buffer }); +} + +PreparedProgramConstantBuffer::~PreparedProgramConstantBuffer() { + if (buffer != nullptr) { + hrx_buffer_release(buffer); + } +} + +PreparedProgramConstantBuffer::PreparedProgramConstantBuffer(PreparedProgramConstantBuffer && other) noexcept : + value(other.value), + name(std::move(other.name)), + buffer(std::exchange(other.buffer, nullptr)), + size(other.size) { + other.size = 0; +} + +PreparedProgramConstantBuffer & PreparedProgramConstantBuffer::operator=( + PreparedProgramConstantBuffer && other) noexcept { + if (this != &other) { + if (buffer != nullptr) { + hrx_buffer_release(buffer); + } + value = other.value; + name = std::move(other.name); + buffer = std::exchange(other.buffer, nullptr); + size = other.size; + other.size = 0; + } + return *this; +} + +RecordedCommandGraph::~RecordedCommandGraph() { + if (exec != nullptr) { + hrx_graph_exec_release(exec); + } + if (graph != nullptr) { + hrx_graph_release(graph); + } + for (hrx_buffer_t buffer : retained_buffers) { + hrx_buffer_release(buffer); + } + for (hrx_executable_t executable : retained_executables) { + hrx_executable_release(executable); + } +} + +RecordedCommandGraph::RecordedCommandGraph(RecordedCommandGraph && other) noexcept : + graph(std::exchange(other.graph, nullptr)), + exec(std::exchange(other.exec, nullptr)), + retained_buffers(std::move(other.retained_buffers)), + retained_executables(std::move(other.retained_executables)), + bound_transient_arena_allocation_id(other.bound_transient_arena_allocation_id), + dispatch_count(other.dispatch_count), + status(std::move(other.status)) { + other.bound_transient_arena_allocation_id = kInvalidTransientArenaAllocationId; + other.dispatch_count = 0; +} + +RecordedCommandGraph & RecordedCommandGraph::operator=(RecordedCommandGraph && other) noexcept { + if (this != &other) { + if (exec != nullptr) { + hrx_graph_exec_release(exec); + } + if (graph != nullptr) { + hrx_graph_release(graph); + } + for (hrx_buffer_t buffer : retained_buffers) { + hrx_buffer_release(buffer); + } + for (hrx_executable_t executable : retained_executables) { + hrx_executable_release(executable); + } + graph = std::exchange(other.graph, nullptr); + exec = std::exchange(other.exec, nullptr); + retained_buffers = std::move(other.retained_buffers); + retained_executables = std::move(other.retained_executables); + bound_transient_arena_allocation_id = other.bound_transient_arena_allocation_id; + dispatch_count = other.dispatch_count; + status = std::move(other.status); + other.bound_transient_arena_allocation_id = kInvalidTransientArenaAllocationId; + other.dispatch_count = 0; + } + return *this; +} + +namespace { + +static const char * status_first_error(const Status & status) { + return status.errors().empty() ? "" : status.errors().front().c_str(); +} + +static bool environment_flag_enabled(const char * name) { + const char * value = std::getenv(name); + return value != nullptr && value[0] != '\0' && !(value[0] == '0' && value[1] == '\0'); +} + +static size_t environment_size_value(const char * name, size_t fallback) { + const char * value = std::getenv(name); + if (value == nullptr || value[0] == '\0') { + return fallback; + } + return static_cast(std::strtoull(value, nullptr, 0)); +} + +static void debug_serial_log(const char * event, const char * phase, const std::string & detail) { + std::fprintf(stderr, "hrx debug serial: %s phase=%s", event, phase); + if (!detail.empty()) { + std::fprintf(stderr, " %s", detail.c_str()); + } + std::fprintf(stderr, "\n"); + std::fflush(stderr); +} + +class DebugSerialExecutionTrace { + public: + explicit DebugSerialExecutionTrace(const CommandProgramExecutionContext & context) : + context_(context), + enabled_(debug_serial_command_execution_enabled()) {} + + DebugSerialExecutionTrace(const CommandProgramExecutionContext & context, const CommandProgram & commands) : + DebugSerialExecutionTrace(context) { + if (enabled_) { + std::ostringstream out; + out << "main_commands=" << commands.commands.size() + << " init_commands=" << commands.initialization_commands.size() + << " transient_arena=" << commands.transients.arena_size; + program_detail_ = out.str(); + } + } + + DebugSerialExecutionTrace(const CommandProgramExecutionContext & context, + const PreparedCommandProgram & commands) : + DebugSerialExecutionTrace(context) { + if (enabled_) { + std::ostringstream out; + out << "main_commands=" << commands.commands.size() + << " init_commands=" << commands.initialization_commands.size(); + program_detail_ = out.str(); + } + } + + bool enabled() const { + return enabled_; + } + + void log_program_begin() const { + log("begin", "command-program", program_detail_); + } + + void log_program_end(bool success) const { + log(success ? "end" : "failed", "command-program", program_detail_); + } + + bool sync_program_phase(const char * phase) const { + return sync(phase, program_detail_); + } + + bool sync(const char * phase, const std::string & detail = {}) const { + if (!enabled_) { + return true; + } + log("sync-begin", phase, detail); + if (ErrorResult error = take_status(hrx_stream_synchronize(context_.stream))) { + GGML_LOG_ERROR("HRX debug serial sync failed phase=%s %s: %s\n", phase, detail.c_str(), error->c_str()); + return false; + } + log("sync-end", phase, detail); + return true; + } + + bool should_trace_program(size_t command_count) const { + const char * value = std::getenv("GGML_HRX_DEBUG_SERIAL_PROGRAM_COMMANDS"); + return value == nullptr || value[0] == '\0' || command_count == std::strtoull(value, nullptr, 0); + } + + bool should_trace_command(size_t program_command_count, size_t index) const { + if (!should_trace_program(program_command_count)) { + return false; + } + const size_t start = environment_size_value("GGML_HRX_DEBUG_SERIAL_START_COMMAND", 0); + const size_t end = environment_size_value("GGML_HRX_DEBUG_SERIAL_END_COMMAND", + std::numeric_limits::max()); + return index >= start && index <= end; + } + + void log(const char * event, const char * phase, const std::string & detail = {}) const { + if (enabled_) { + debug_serial_log(event, phase, detail); + } + } + + static std::string command_detail(const char * list_kind, + size_t list_size, + size_t index, + const PreparedCommand & command) { + std::ostringstream out; + out << "list=" << list_kind + << " index=" << index + << " list_commands=" << list_size + << " ordinal=" << command.ordinal + << " kind=" << command_kind_name(command.kind); + if (command.kind == CommandKind::Kernel) { + out << " kernel_id=" << command.kernel.specialization.kernel_id + << " bindings=" << command.kernel.bindings.size(); + } + return out.str(); + } + + static std::string list_detail(const char * list_kind, size_t list_size, size_t program_command_count) { + std::ostringstream out; + out << "list=" << list_kind + << " list_commands=" << list_size + << " main_commands=" << program_command_count; + return out.str(); + } + + static std::string host_staging_detail(const HostStagingBuffer & staging) { + std::ostringstream out; + out << "value=" << staging.value << " length=" << staging.length; + return out.str(); + } + + private: + const CommandProgramExecutionContext & context_; + bool enabled_ = false; + std::string program_detail_; +}; + +#define HRX_DEBUG_SERIAL_SYNC(trace, phase) \ + do { \ + if (!(trace).sync_program_phase(phase)) { \ + return false; \ + } \ + } while (0) + +static bool resource_access_writes(ResourceAccess access) { + return access == ResourceAccess::Write || access == ResourceAccess::ReadWrite; +} + +static Status command_program_metadata_context_valid(const CommandProgramExecutionContext & context) { + Status status; + if (context.target == nullptr) { + status.log("missing HRX target"); + return status; + } + if (context.corpus == nullptr) { + status.log("missing HRX kernel corpus"); + return status; + } + return status; +} + +static Status command_program_preparation_context_valid(const CommandProgramExecutionContext & context) { + Status status; + if (context.device == nullptr) { + status.log("missing HRX device"); + return status; + } + if (context.kernel_executables == nullptr) { + status.log("missing HRX kernel executable cache"); + return status; + } + if (context.host_transfers == nullptr) { + status.log("missing HRX host transfer manager"); + return status; + } + if (context.host_weights == nullptr) { + status.log("missing HRX host weight cache"); + return status; + } + return status; +} + +static Status command_program_transient_context_valid(const CommandProgramExecutionContext & context, + const CommandProgram & commands) { + Status status; + if (commands.transients.arena_size == 0) { + return status; + } + if (context.transient_arena == nullptr) { + status.log("missing HRX transient arena"); + return status; + } + if (context.stream == nullptr) { + status.log("missing HRX stream for transient arena"); + return status; + } + return status; +} + +static bool prepared_program_has_constant(const PreparedCommandProgram & prepared, ValueId value) { + for (const PreparedProgramConstantBuffer & constant : prepared.program_constants) { + if (constant.value == value) { + return true; + } + } + return false; +} + +static bool prepared_execution_context_valid(const CommandProgramExecutionContext & context) { + if (context.stream == nullptr) { + GGML_LOG_ERROR("%s: missing HRX stream\n", __func__); + return false; + } + return true; +} + +static Status ensure_transient_arena(const CommandProgramExecutionContext & context, + const CommandProgram & commands, + TransientArenaAllocationRef & allocation) { + allocation = {}; + Status status = command_program_transient_context_valid(context, commands); + if (!status.success()) { + return status; + } + if (commands.transients.arena_size == 0) { + return status; + } + status = context.transient_arena->ensure_capacity(context.device, context.stream, commands.transients.arena_size); + if (!status.success()) { + return status; + } + allocation = context.transient_arena->current_allocation(); + return status; +} + +static Status initialize_command_program_constants(const CommandProgramExecutionContext & context, + const CommandProgram & commands, + const TransientArenaAllocationRef & allocation, + const PreparedCommandProgram & prepared) { + Status status; + if (commands.constant_initializations.empty()) { + return status; + } + for (const ConstantInitialization & initialization : commands.constant_initializations) { + if (prepared_program_has_constant(prepared, initialization.value)) { + continue; + } + if (allocation.buffer == nullptr) { + status.log("command program has constant initialization %s without a transient arena allocation", + initialization.name.c_str()); + continue; + } + const TransientAllocation * transient = find_transient_allocation(commands.transients, initialization.value); + if (transient == nullptr) { + status.log("constant initialization %s references missing transient value %d", initialization.name.c_str(), + initialization.value.value); + continue; + } + if (initialization.offset > transient->size || + initialization.data.size() > transient->size - initialization.offset) { + status.log("constant initialization %s is outside transient allocation length %zu", + initialization.name.c_str(), transient->size); + continue; + } + // TODO: Track initialized transient arena allocation ids so constants are not transferred every invocation. + if (context.host_transfers == nullptr) { + status.log("constant initialization %s requires a host transfer manager", initialization.name.c_str()); + continue; + } + Status upload_status = context.host_transfers->upload_synchronous( + context.stream, initialization.data.data(), allocation.buffer, + transient->arena_offset + initialization.offset, initialization.data.size()); + if (!upload_status.success()) { + status.log("failed to upload constant initialization %s", initialization.name.c_str()); + status.append(upload_status); + } + } + return status; +} + +static Status initialize_command_program_completion_counters(const CommandProgramExecutionContext & context, + const CommandProgram & commands, + const TransientArenaAllocationRef & allocation) { + Status status; + if (commands.completion_counters.byte_count == 0) { + return status; + } + if (allocation.buffer == nullptr) { + status.log("command program has completion counters without a transient arena allocation"); + return status; + } + if (commands.completion_counters.arena_offset > commands.transients.arena_size || + commands.completion_counters.byte_count > + commands.transients.arena_size - commands.completion_counters.arena_offset) { + status.log("completion counter initialization is outside transient arena length %zu", + commands.transients.arena_size); + return status; + } + const uint32_t zero_pattern = 0; + if (ErrorResult error = take_status( + hrx_stream_fill_buffer(context.stream, allocation.buffer, commands.completion_counters.arena_offset, + commands.completion_counters.byte_count, &zero_pattern, sizeof(zero_pattern)))) { + status.log("failed to initialize completion counters: %s", error->c_str()); + } + return status; +} + +static std::string format_resolved_command_context(const ResolvedCommand & command) { + std::ostringstream out; + out << "command " << command.ordinal << " kind=" << command_kind_name(command.kind) + << " kernel_id=" << command.kernel.kernel_id << " bindings=" << command.bindings.size(); + return out.str(); +} + +static std::string format_prepared_command_context(const PreparedCommand & command) { + std::ostringstream out; + out << "command " << command.ordinal << " kind=" << command_kind_name(command.kind); + if (command.kind == CommandKind::Kernel) { + out << " kernel_id=" << command.kernel.specialization.kernel_id + << " bindings=" << command.kernel.bindings.size(); + } + return out.str(); +} + +static Dispatch build_dispatch(const ResolvedCommand & command) { + Dispatch dispatch; + dispatch.kernel = command.kernel; + dispatch.bindings.reserve(command.bindings.size()); + for (const ResolvedCommandBinding & binding : command.bindings) { + DispatchBinding dispatch_binding; + dispatch_binding.value = binding.binding.value; + dispatch_binding.offset = binding.binding.offset; + dispatch_binding.length = binding.binding.length; + dispatch_binding.layout = binding.binding.layout; + dispatch_binding.source_type = binding.binding.source_type; + dispatch_binding.input_size = binding.binding.input_size; + dispatch_binding.output_size = binding.binding.output_size; + dispatch_binding.source_length = binding.binding.source_length; + dispatch.bindings.push_back(std::move(dispatch_binding)); + } + return dispatch; +} + +struct GraphValueAccess { + bool read = false; + bool write = false; + bool native_read = false; + bool transformed_read = false; + bool layout_conflict = false; + std::string layout = kNativeWeightLayout; + ggml_type source_type = GGML_TYPE_COUNT; + int64_t input_size = 0; + int64_t output_size = 0; + size_t source_length = 0; + size_t materialized_length = 0; +}; + +static std::unordered_map collect_graph_value_access(const CommandProgram & commands) { + std::unordered_map access_by_value; + auto append_command_list_access = [&](const std::vector & command_list) { + for (const Command & command : command_list) { + for (const CommandBinding & binding : command.bindings) { + if (binding.origin != CommandBindingOrigin::GraphValue) { + continue; + } + GraphValueAccess & access = access_by_value[binding.value.value]; + switch (binding.access) { + case ResourceAccess::Read: + access.read = true; + break; + case ResourceAccess::Write: + access.write = true; + break; + case ResourceAccess::ReadWrite: + access.read = true; + access.write = true; + break; + } + if (binding.access == ResourceAccess::Write) { + continue; + } + if (binding.layout == kNativeWeightLayout) { + access.native_read = true; + if (access.transformed_read) { + access.layout_conflict = true; + } + continue; + } + if (binding.offset != 0 || binding.source_length == 0 || binding.length == 0 || + binding.source_type == GGML_TYPE_COUNT || binding.input_size <= 0 || binding.output_size <= 0) { + access.layout_conflict = true; + continue; + } + if (!access.transformed_read) { + access.transformed_read = true; + access.layout = binding.layout; + access.source_type = binding.source_type; + access.input_size = binding.input_size; + access.output_size = binding.output_size; + access.source_length = binding.source_length; + access.materialized_length = binding.length; + } else if (access.layout != binding.layout || access.source_type != binding.source_type || + access.input_size != binding.input_size || access.output_size != binding.output_size || + access.source_length != binding.source_length || + access.materialized_length != binding.length) { + access.layout_conflict = true; + } + if (access.native_read) { + access.layout_conflict = true; + } + } + } + }; + append_command_list_access(commands.initialization_commands); + append_command_list_access(commands.commands); + return access_by_value; +} + +static CommandProgramBindings materialize_host_bindings(const CommandProgramExecutionContext & context, + const CommandProgram & commands, + const CommandProgramBindings & bindings, + PreparedCommandProgram & prepared) { + std::vector materialized; + Status status; + materialized.reserve(bindings.bindings().size()); + const std::unordered_map access_by_value = collect_graph_value_access(commands); + for (const CommandProgramBinding & binding : bindings.bindings()) { + const auto found_access = access_by_value.find(binding.value.value); + const GraphValueAccess access = + found_access != access_by_value.end() ? found_access->second : GraphValueAccess{}; + const bool materialize_weight = binding.weight && access.read && !access.write && + (binding.requires_materialization() || access.transformed_read); + if (binding.length == 0 || (!binding.requires_materialization() && !materialize_weight)) { + materialized.push_back(binding); + continue; + } + if (materialize_weight) { + if (access.layout_conflict) { + status.log("weight value %d has conflicting resident layout requests", binding.value.value); + materialized.push_back(binding); + continue; + } + HostWeightSource source; + source.host_data = binding.host_data; + source.device_buffer = binding.host_data == nullptr ? binding.buffer : nullptr; + source.identity = binding.identity; + source.generation = binding.generation; + source.capacity = binding.capacity; + source.offset = binding.offset; + source.length = access.transformed_read ? access.source_length : binding.length; + source.materialized_length = access.transformed_read ? access.materialized_length : binding.length; + source.layout = access.layout; + source.source_type = access.source_type; + source.input_size = access.input_size; + source.output_size = access.output_size; + if (source.length > binding.length) { + status.log("weight value %d layout source length %zu exceeds runtime length %zu", binding.value.value, + source.length, binding.length); + materialized.push_back(binding); + continue; + } + HostWeightAcquireResult resident = + context.host_weights->acquire(context.device, context.stream, *context.host_transfers, source); + if (!resident.valid()) { + status.log("materialize host weight value %d failed", binding.value.value); + status.append(resident.status); + materialized.push_back(binding); + continue; + } + CommandProgramBinding device_binding = binding; + device_binding.buffer = resident.lease.buffer(); + device_binding.host_data = nullptr; + device_binding.offset = 0; + device_binding.length = resident.lease.length(); + device_binding.capacity = resident.lease.length(); + materialized.push_back(device_binding); + prepared.resident_host_weights.push_back(std::move(resident.lease)); + continue; + } + + if (binding.host_data == nullptr) { + materialized.push_back(binding); + continue; + } + + HostStagingBuffer staging; + Status allocation_status = allocate_host_staging_buffer(context.device, binding.length, staging); + if (!allocation_status.success()) { + status.log("allocate host staging for value %d failed", binding.value.value); + status.append(allocation_status); + materialized.push_back(binding); + continue; + } + staging.value = binding.value.value; + staging.host_data = static_cast(binding.host_data) + binding.offset; + staging.upload = access.read; + staging.download = access.write && !binding.graph_input; + if (staging.upload && context.host_buffers != nullptr) { + staging.source_host_buffer = context.host_buffers->find(staging.host_data, staging.length); + } + CommandProgramBinding device_binding = binding; + device_binding.buffer = staging.buffer; + device_binding.host_data = nullptr; + device_binding.offset = 0; + device_binding.capacity = binding.length; + materialized.push_back(device_binding); + prepared.host_staging.push_back(std::move(staging)); + } + return CommandProgramBindings::from_bindings(std::move(materialized), status); +} + +struct ProgramConstantImage { + ValueId value; + std::string name; + std::vector data; + bool read = false; +}; + +static Status collect_program_constant_images(const CommandProgram & commands, + std::vector & images) { + Status status; + std::unordered_map image_by_value; + for (const ConstantInitialization & initialization : commands.constant_initializations) { + const TransientAllocation * allocation = find_transient_allocation(commands.transients, initialization.value); + if (allocation == nullptr) { + status.log("constant initialization %s references missing transient value %d", initialization.name.c_str(), + initialization.value.value); + continue; + } + if (initialization.offset > allocation->size || + initialization.data.size() > allocation->size - initialization.offset) { + status.log("constant initialization %s is outside transient allocation length %zu", + initialization.name.c_str(), allocation->size); + continue; + } + + ProgramConstantImage * image = nullptr; + const auto found = image_by_value.find(initialization.value.value); + if (found == image_by_value.end()) { + ProgramConstantImage next; + next.value = initialization.value; + next.name = initialization.name; + next.data.resize(allocation->size); + image_by_value.emplace(initialization.value.value, images.size()); + images.push_back(std::move(next)); + image = &images.back(); + } else { + image = &images[found->second]; + } + + std::copy(initialization.data.begin(), initialization.data.end(), image->data.begin() + initialization.offset); + } + return status; +} + +static Status validate_program_constant_access(const CommandProgram & commands, + std::vector & images) { + Status status; + std::unordered_map image_by_value; + for (size_t i = 0; i < images.size(); ++i) { + image_by_value.emplace(images[i].value.value, i); + } + + auto validate_command_list = [&](const std::vector & command_list) { + for (const Command & command : command_list) { + for (const CommandBinding & binding : command.bindings) { + const auto found = image_by_value.find(binding.value.value); + if (found == image_by_value.end()) { + continue; + } + ProgramConstantImage & image = images[found->second]; + if (resource_access_writes(binding.access)) { + status.log( + "constant initialization %s is written by command %u; prepared constant buffer copy " + "support is required", + image.name.c_str(), command.ordinal); + continue; + } + image.read = true; + } + } + }; + validate_command_list(commands.initialization_commands); + validate_command_list(commands.commands); + return status; +} + +static PreparedProgramConstantBuffer make_program_constant_buffer(ValueId value, + std::string name, + hrx_buffer_t buffer, + size_t size) { + PreparedProgramConstantBuffer result; + result.value = value; + result.name = std::move(name); + result.buffer = buffer; + result.size = size; + return result; +} + +static Status bind_prepared_command_list_program_constants(const PreparedCommandProgram & prepared, + std::vector & prepared_commands) { + Status status; + for (PreparedCommand & command : prepared_commands) { + for (PreparedCommandBinding & binding : command.kernel.bindings) { + for (const PreparedProgramConstantBuffer & constant : prepared.program_constants) { + if (binding.binding.origin != CommandBindingOrigin::Transient || + binding.binding.value != constant.value) { + continue; + } + if (binding.binding.offset > constant.size || + binding.binding.length > constant.size - binding.binding.offset) { + status.log("%s is outside prepared constant %s length %zu", + format_command_binding(binding.binding).c_str(), constant.name.c_str(), constant.size); + continue; + } + binding.ref = { constant.buffer, binding.binding.offset, binding.binding.length }; + binding.binding.origin = CommandBindingOrigin::ProgramConstant; + } + } + } + return status; +} + +static Status prepare_program_constant_buffers(const CommandProgramExecutionContext & context, + const CommandProgram & commands, + PreparedCommandProgram & prepared) { + Status status; + if (commands.constant_initializations.empty()) { + return status; + } + + std::vector images; + status.append(collect_program_constant_images(commands, images)); + status.append(validate_program_constant_access(commands, images)); + if (!status.success()) { + return status; + } + if (images.empty()) { + return status; + } + if (context.device == nullptr) { + status.log("missing HRX device for prepared constants"); + return status; + } + if (context.stream == nullptr) { + status.log("missing HRX stream for prepared constants"); + return status; + } + + hrx_buffer_params_t params = { + HRX_MEMORY_TYPE_DEVICE_LOCAL, + HRX_MEMORY_ACCESS_ALL, + HRX_BUFFER_USAGE_DEFAULT, + 0, + }; + for (const ProgramConstantImage & image : images) { + if (!image.read) { + continue; + } + hrx_buffer_t buffer = nullptr; + if (ErrorResult error = take_status(hrx_allocator_allocate_buffer(hrx_device_allocator(context.device), params, + image.data.size(), &buffer))) { + status.log("allocate prepared constant %s: %s", image.name.c_str(), error->c_str()); + continue; + } + if (context.host_transfers == nullptr) { + hrx_buffer_release(buffer); + status.log("prepared constant %s requires a host transfer manager", image.name.c_str()); + continue; + } + Status upload_status = + context.host_transfers->upload_synchronous(context.stream, image.data.data(), buffer, 0, image.data.size()); + if (!upload_status.success()) { + hrx_buffer_release(buffer); + status.log("upload prepared constant %s", image.name.c_str()); + status.append(upload_status); + continue; + } + prepared.program_constants.push_back( + make_program_constant_buffer(image.value, image.name, buffer, image.data.size())); + } + if (!status.success()) { + return status; + } + + status.append(bind_prepared_command_list_program_constants(prepared, prepared.initialization_commands)); + status.append(bind_prepared_command_list_program_constants(prepared, prepared.commands)); + return status; +} + +static Status rebind_prepared_host_staging(const CommandProgramExecutionContext & context, + const CommandProgramBindings & bindings, + PreparedCommandProgram & prepared) { + Status status; + for (HostStagingBuffer & staging : prepared.host_staging) { + const CommandProgramBinding * binding = bindings.find(ValueId(staging.value)); + if (binding == nullptr || binding->host_data == nullptr || binding->length != staging.length || + binding->offset > binding->capacity || binding->length > binding->capacity - binding->offset) { + status.log("live host binding does not match prepared value %d", staging.value); + continue; + } + staging.host_data = static_cast(binding->host_data) + binding->offset; + staging.source_host_buffer = HostBufferRef{}; + if (staging.upload && context.host_buffers != nullptr) { + staging.source_host_buffer = context.host_buffers->find(staging.host_data, staging.length); + } + } + return status; +} + +static Status upload_prepared_host_staging(const CommandProgramExecutionContext & context, + const PreparedCommandProgram & prepared) { + Status status; + if (prepared.host_staging.empty()) { + return status; + } + if (context.host_transfers == nullptr) { + status.log("missing HRX host transfer manager"); + return status; + } + const DebugSerialExecutionTrace debug(context); + for (const HostStagingBuffer & staging : prepared.host_staging) { + if (!staging.upload) { + continue; + } + // A host download staged by an earlier replay is only published to host + // memory by mark_stream_synchronized(); add_host_writeback merely queues + // the memcpy. Uploading the same host memory before that has happened + // would read the previous contents. The split scheduler computes splits + // back to back and only synchronizes a backend when it inserts a + // cross-backend input copy, so nothing else guarantees the ordering. + if (context.graph_replay_state != nullptr && context.graph_replay_state->has_pending_host_writebacks()) { + static thread_local WaitHistory writeback_wait_history; + if (ErrorResult error = + take_status(stream_synchronize_sleeping(context.stream, writeback_wait_history))) { + status.log("synchronize HRX stream before host-staging upload: %s", error->c_str()); + return status; + } + context.graph_replay_state->mark_stream_synchronized(); + } + const std::string detail = debug.enabled() ? DebugSerialExecutionTrace::host_staging_detail(staging) : + std::string(); + if (!debug.sync("host-staging-upload-buffer-pre", detail)) { + status.log("debug serial sync before host-staging upload failed for value %d", staging.value); + return status; + } + Status upload_status; + if (staging.source_host_buffer.valid()) { + if (ErrorResult error = take_status(hrx_stream_copy_buffer(context.stream, staging.source_host_buffer.buffer(), + staging.source_host_buffer.offset(), + staging.buffer, 0, staging.length))) { + upload_status.log("HRX host staging buffer copy failed: %s", error->c_str()); + } + } else { + upload_status = + context.host_transfers->upload_async(context.stream, staging.host_data, staging.buffer, 0, staging.length); + } + status.append(upload_status); + if (!debug.sync("host-staging-upload-buffer-post", detail)) { + status.log("debug serial sync after host-staging upload failed for value %d", staging.value); + return status; + } + } + return status; +} + +static void mark_graph_replay_synchronized(const CommandProgramExecutionContext & context) { + if (context.graph_replay_state != nullptr) { + context.graph_replay_state->mark_stream_synchronized(); + } +} + +static Status enqueue_stream_execution_barrier(const CommandProgramExecutionContext & context, const char * label) { + Status status; + if (ErrorResult error = take_status(hrx_stream_execution_barrier(context.stream))) { + status.log("%s failed: %s", label, error->c_str()); + } + return status; +} + +static Status flush_stream_commands(const CommandProgramExecutionContext & context, const char * label) { + Status status; + if (ErrorResult error = take_status(hrx_stream_flush(context.stream))) { + status.log("%s failed: %s", label, error->c_str()); + } + return status; +} + +static Status wait_stream_commands(const CommandProgramExecutionContext & context, const char * label, + const void * work) { + Status status; + if (ErrorResult error = take_status(stream_wait_sleeping(context.stream, wait_history_for(work)))) { + status.log("%s failed: %s", label, error->c_str()); + } + return status; +} + +static Status download_prepared_host_staging(const CommandProgramExecutionContext & context, + const PreparedCommandProgram & prepared, + bool insert_barrier = true) { + Status status; + + bool has_download = false; + for (const HostStagingBuffer & staging : prepared.host_staging) { + has_download = has_download || staging.download; + } + if (!has_download) { + return status; + } + if (context.graph_replay_state == nullptr) { + status.log("missing HRX stream completion state"); + return status; + } + if (insert_barrier) { + Status barrier_status = enqueue_stream_execution_barrier(context, "insert HRX host download barrier"); + if (!barrier_status.success()) { + return barrier_status; + } + } + for (const HostStagingBuffer & staging : prepared.host_staging) { + if (!staging.download) { + continue; + } + hrx_buffer_t download_buffer = nullptr; + void * download_data = nullptr; + Status allocation_status = + allocate_mapped_host_staging_buffer(context.device, staging.length, download_buffer, download_data); + if (!allocation_status.success()) { + status.log("allocate HRX host download staging buffer for value %d failed", staging.value); + status.append(allocation_status); + return status; + } + if (ErrorResult error = take_status(hrx_stream_copy_buffer(context.stream, staging.buffer, 0, + download_buffer, 0, staging.length))) { + status.log("HRX host download staging copy failed for value %d: %s", staging.value, error->c_str()); + hrx_buffer_release(download_buffer); + return status; + } + context.graph_replay_state->add_host_writeback(staging.host_data, download_data, staging.length, download_buffer); + hrx_buffer_release(download_buffer); + } + return status; +} + +static PreparedCommand make_prepared_command_shape(const ResolvedCommand & command) { + PreparedCommand prepared; + prepared.ordinal = command.ordinal; + prepared.kind = command.kind; + prepared.kernel.specialization = command.kernel; + prepared.kernel.bindings.reserve(command.bindings.size()); + for (const ResolvedCommandBinding & binding : command.bindings) { + prepared.kernel.bindings.push_back({ + binding.binding, + { binding.ref.buffer, binding.ref.offset, binding.ref.length }, + }); + } + return prepared; +} + +static Status prepare_kernel_command(const CommandProgramExecutionContext & context, + const ResolvedCommand & command, + PreparedCommand & prepared, + KernelExecutableRef & executable_ref) { + Status status; + const std::string command_context = format_resolved_command_context(command); + if (command.kind != CommandKind::Kernel) { + status.log("unsupported command kind in %s", command_context.c_str()); + return status; + } + Dispatch dispatch = build_dispatch(command); + + KernelResolveResult resolved = + resolve_kernel_definition(*context.corpus, context.target, dispatch.kernel.kernel_id); + if (!resolved.found()) { + status.log("%s: %s", command_context.c_str(), + format_kernel_resolve_error(resolved, dispatch.kernel.kernel_id).c_str()); + return status; + } + + prepared = make_prepared_command_shape(command); + executable_ref = context.kernel_executables->get_or_compile( + { context.device, context.target }, *resolved.definition, dispatch, prepared.kernel.constants); + if (!executable_ref.valid()) { + status.log("failed to prepare %s", command_context.c_str()); + return status; + } + return status; +} + +static bool execute_prepared_kernel_command(const CommandProgramExecutionContext & context, + const PreparedCommand & command) { + const std::string command_context = format_prepared_command_context(command); + if (command.kind != CommandKind::Kernel) { + GGML_LOG_ERROR("%s: unsupported command kind in %s\n", __func__, command_context.c_str()); + return false; + } + if (command.kernel.executable == nullptr) { + GGML_LOG_ERROR("%s: missing kernel executable for %s\n", __func__, command_context.c_str()); + return false; + } + + std::vector refs; + refs.reserve(command.kernel.bindings.size()); + for (const PreparedCommandBinding & binding : command.kernel.bindings) { + refs.push_back({ binding.ref.buffer, binding.ref.offset, binding.ref.length }); + } + + const KernelExecutable & executable = *command.kernel.executable; + hrx_dispatch_config_t config = { + { executable.launch.workgroup_count[0], executable.launch.workgroup_count[1], + executable.launch.workgroup_count[2] }, + { executable.launch.workgroup_size[0], executable.launch.workgroup_size[1], + executable.launch.workgroup_size[2] }, + executable.launch.subgroup_size, + }; + if (ErrorResult error = take_status(hrx_stream_dispatch( + context.stream, executable.executable, executable.export_ordinal, &config, command.kernel.constants.data(), + command.kernel.constants.size(), refs.data(), refs.size(), 0))) { + GGML_LOG_ERROR("%s: failed to execute %s: %s\n", __func__, command_context.c_str(), error->c_str()); + return false; + } + return true; +} + +static void prepare_command_list(const CommandProgramExecutionContext & context, + const std::vector & commands, + std::vector & prepared_commands, + std::vector & executable_refs, + Status & status) { + prepared_commands.reserve(commands.size()); + executable_refs.reserve(commands.size()); + for (const ResolvedCommand & command : commands) { + PreparedCommand prepared_command; + KernelExecutableRef executable_ref; + Status command_status = prepare_kernel_command(context, command, prepared_command, executable_ref); + if (command_status.success()) { + prepared_commands.push_back(std::move(prepared_command)); + executable_refs.push_back(std::move(executable_ref)); + } else { + status.append(command_status); + } + } +} + +static void materialize_command_list_executables(const CommandProgramExecutionContext & context, + std::vector & prepared_commands, + const std::vector & executable_refs, + Status & status) { + for (size_t i = 0; i < prepared_commands.size(); ++i) { + PreparedCommand & command = prepared_commands[i]; + command.kernel.executable = context.kernel_executables->materialize( + { context.device, context.target }, executable_refs[i], command.kernel.constants); + if (command.kernel.executable == nullptr) { + status.log("failed to prepare %s", format_prepared_command_context(command).c_str()); + } + } +} + +static bool bind_prepared_command_list_transients(const CommandProgram & commands, + const TransientArenaAllocationRef & transient_allocation, + std::vector & prepared_commands) { + for (PreparedCommand & command : prepared_commands) { + for (PreparedCommandBinding & binding : command.kernel.bindings) { + if (binding.binding.origin != CommandBindingOrigin::Transient) { + continue; + } + const TransientAllocation * allocation = + find_transient_allocation(commands.transients, binding.binding.value); + if (allocation == nullptr) { + GGML_LOG_ERROR("%s: %s has no transient allocation\n", __func__, + format_command_binding(binding.binding).c_str()); + return false; + } + if (binding.binding.offset > allocation->size || + binding.binding.length > allocation->size - binding.binding.offset || + commands.transients.arena_size > transient_allocation.capacity) { + GGML_LOG_ERROR("%s: %s is outside transient arena\n", __func__, + format_command_binding(binding.binding).c_str()); + return false; + } + binding.ref = { + transient_allocation.buffer, + allocation->arena_offset + binding.binding.offset, + binding.binding.length, + }; + } + } + return true; +} + +static bool execute_prepared_command_list_serial(const CommandProgramExecutionContext & context, + const std::vector & commands, + const char * list_kind, + size_t program_command_count) { + const DebugSerialExecutionTrace debug(context); + const bool trace_list = debug.enabled() && debug.should_trace_program(program_command_count); + if (trace_list) { + debug.log("begin", "command-list", + DebugSerialExecutionTrace::list_detail(list_kind, commands.size(), program_command_count)); + } + for (size_t i = 0; i < commands.size(); ++i) { + const PreparedCommand & command = commands[i]; + const bool trace_command = debug.enabled() && debug.should_trace_command(program_command_count, i); + const std::string command_detail = trace_command ? + DebugSerialExecutionTrace::command_detail(list_kind, commands.size(), i, + command) : + std::string(); + if (!debug.sync("command-pre-dispatch", command_detail)) { + return false; + } + if (trace_command) { + debug.log("begin", "command-dispatch", command_detail); + } + if (!execute_prepared_kernel_command(context, command)) { + return false; + } + if (trace_command) { + debug.log("end", "command-dispatch", command_detail); + } + if (!debug.sync("command-post-dispatch", command_detail)) { + return false; + } + } + if (trace_list) { + debug.log("end", "command-list", + DebugSerialExecutionTrace::list_detail(list_kind, commands.size(), program_command_count)); + } + return true; +} + +static bool execute_prepared_command_list(const CommandProgramExecutionContext & context, + const std::vector & commands) { + return execute_prepared_command_list_serial(context, commands, "main", commands.size()); +} + +struct GraphResourceState { + size_t end = 0; + hrx_graph_node_t last_writer = nullptr; + std::vector readers; +}; + +class GraphDependencyPlanner { + public: + std::vector dependencies(const std::vector & bindings) const { + std::vector result; + for (const PreparedCommandBinding & binding : bindings) { + collect_dependencies(binding.ref.buffer, binding.ref.offset, binding.ref.length, + resource_access_writes(binding.binding.access), result); + } + return result; + } + + void record(hrx_graph_node_t node, const std::vector & bindings) { + for (const PreparedCommandBinding & binding : bindings) { + update(node, binding.ref.buffer, binding.ref.offset, binding.ref.length, + resource_access_writes(binding.binding.access)); + } + } + + void record(hrx_graph_node_t node, hrx_buffer_t buffer, size_t offset, size_t length, bool writes) { + update(node, buffer, offset, length, writes); + } + + private: + using ResourceMap = std::map; + + static size_t range_end(size_t offset, size_t length) { + return length > std::numeric_limits::max() - offset ? std::numeric_limits::max() : + offset + length; + } + + static void append_unique(std::vector & nodes, hrx_graph_node_t node) { + if (node != nullptr && std::find(nodes.begin(), nodes.end(), node) == nodes.end()) { + nodes.push_back(node); + } + } + + static void split(ResourceMap & ranges, size_t offset) { + auto upper = ranges.upper_bound(offset); + if (upper == ranges.begin()) { + return; + } + auto current = std::prev(upper); + if (offset <= current->first || offset >= current->second.end) { + return; + } + GraphResourceState right = current->second; + current->second.end = offset; + ranges.emplace(offset, std::move(right)); + } + + void collect_dependencies(hrx_buffer_t buffer, + size_t offset, + size_t length, + bool writes, + std::vector & result) const { + if (buffer == nullptr || length == 0) { + return; + } + const auto resource = resources_.find(buffer); + if (resource == resources_.end()) { + return; + } + const size_t end = range_end(offset, length); + auto range = resource->second.upper_bound(offset); + if (range != resource->second.begin()) { + --range; + if (range->second.end <= offset) { + ++range; + } + } + for (; range != resource->second.end() && range->first < end; ++range) { + append_unique(result, range->second.last_writer); + if (writes) { + for (hrx_graph_node_t reader : range->second.readers) { + append_unique(result, reader); + } + } + } + } + + void update(hrx_graph_node_t node, hrx_buffer_t buffer, size_t offset, size_t length, bool writes) { + if (buffer == nullptr || length == 0) { + return; + } + const size_t end = range_end(offset, length); + ResourceMap & ranges = resources_[buffer]; + split(ranges, offset); + split(ranges, end); + + size_t cursor = offset; + auto range = ranges.lower_bound(offset); + while (cursor < end) { + if (range == ranges.end() || range->first > cursor) { + const size_t gap_end = range == ranges.end() ? end : std::min(end, range->first); + GraphResourceState state; + state.end = gap_end; + if (writes) { + state.last_writer = node; + } else { + state.readers.push_back(node); + } + ranges.emplace(cursor, std::move(state)); + cursor = gap_end; + continue; + } + + if (writes) { + range->second.last_writer = node; + range->second.readers.clear(); + } else { + append_unique(range->second.readers, node); + } + cursor = range->second.end; + ++range; + } + } + + std::unordered_map resources_; +}; + +static Status record_completion_counter_fill(hrx_graph_t graph, + GraphDependencyPlanner & dependencies, + const CommandProgram & commands, + const TransientArenaAllocationRef & allocation) { + Status status; + if (commands.completion_counters.byte_count == 0) { + return status; + } + if (allocation.buffer == nullptr) { + status.log("command program has completion counters without a transient arena allocation"); + return status; + } + if (commands.completion_counters.arena_offset > commands.transients.arena_size || + commands.completion_counters.byte_count > + commands.transients.arena_size - commands.completion_counters.arena_offset) { + status.log("completion counter graph fill is outside transient arena length %zu", + commands.transients.arena_size); + return status; + } + + hrx_graph_fill_buffer_node_attrs_t attrs = { + { allocation.buffer, commands.completion_counters.arena_offset, commands.completion_counters.byte_count }, + 0, + sizeof(uint32_t), + }; + hrx_graph_node_t node = nullptr; + if (ErrorResult error = take_status(hrx_graph_add_fill_buffer_node(graph, nullptr, 0, &attrs, &node))) { + status.log("record completion counter fill: %s", error->c_str()); + return status; + } + dependencies.record(node, allocation.buffer, commands.completion_counters.arena_offset, + commands.completion_counters.byte_count, true); + return status; +} + +static Status record_prepared_kernel_command(hrx_graph_t graph, + GraphDependencyPlanner & dependencies, + const PreparedCommand & command) { + Status status; + const std::string command_context = format_prepared_command_context(command); + if (command.kind != CommandKind::Kernel) { + status.log("unsupported command kind in %s", command_context.c_str()); + return status; + } + if (command.kernel.executable == nullptr) { + status.log("missing kernel executable for %s", command_context.c_str()); + return status; + } + + std::vector refs; + refs.reserve(command.kernel.bindings.size()); + for (const PreparedCommandBinding & binding : command.kernel.bindings) { + if (binding.ref.buffer == nullptr) { + status.log("%s has unbound buffer in %s", format_command_binding(binding.binding).c_str(), + command_context.c_str()); + continue; + } + refs.push_back({ binding.ref.buffer, binding.ref.offset, binding.ref.length }); + } + if (!status.success()) { + return status; + } + + const KernelExecutable & executable = *command.kernel.executable; + hrx_graph_kernel_node_attrs_t attrs = { + executable.executable, + executable.export_ordinal, + { + { executable.launch.workgroup_count[0], executable.launch.workgroup_count[1], + executable.launch.workgroup_count[2] }, + { executable.launch.workgroup_size[0], executable.launch.workgroup_size[1], + executable.launch.workgroup_size[2] }, + executable.launch.subgroup_size, + }, + command.kernel.constants.data(), + command.kernel.constants.size(), + refs.data(), + refs.size(), + 0, + }; + const std::vector dependency_nodes = dependencies.dependencies(command.kernel.bindings); + hrx_graph_node_t node = nullptr; + if (ErrorResult error = take_status( + hrx_graph_add_kernel_node(graph, dependency_nodes.data(), dependency_nodes.size(), &attrs, &node))) { + status.log("record %s: %s", command_context.c_str(), error->c_str()); + return status; + } + dependencies.record(node, command.kernel.bindings); + return status; +} + +static Status record_prepared_command_list(hrx_graph_t graph, + GraphDependencyPlanner & dependencies, + const std::vector & commands, + size_t & dispatch_count) { + Status status; + for (const PreparedCommand & command : commands) { + Status command_status = record_prepared_kernel_command(graph, dependencies, command); + if (!command_status.success()) { + status.append(command_status); + return status; + } + ++dispatch_count; + } + return status; +} + +static void retain_prepared_command_list_resources(const std::vector & commands, + std::unordered_map & retained_buffers, + std::unordered_map & retained_executables, + RecordedCommandGraph & recorded) { + for (const PreparedCommand & command : commands) { + if (command.kernel.executable != nullptr && command.kernel.executable->executable != nullptr && + retained_executables.emplace(command.kernel.executable->executable, true).second) { + hrx_executable_retain(command.kernel.executable->executable); + recorded.retained_executables.push_back(command.kernel.executable->executable); + } + for (const PreparedCommandBinding & binding : command.kernel.bindings) { + if (binding.ref.buffer != nullptr && retained_buffers.emplace(binding.ref.buffer, true).second) { + hrx_buffer_retain(binding.ref.buffer); + recorded.retained_buffers.push_back(binding.ref.buffer); + } + } + } +} + +static void retain_prepared_command_graph_resources(const PreparedCommandProgram & prepared, + const TransientArenaAllocationRef & allocation, + RecordedCommandGraph & recorded) { + std::unordered_map retained_buffers; + std::unordered_map retained_executables; + if (allocation.buffer != nullptr) { + hrx_buffer_retain(allocation.buffer); + recorded.retained_buffers.push_back(allocation.buffer); + retained_buffers.emplace(allocation.buffer, true); + } + retain_prepared_command_list_resources(prepared.initialization_commands, retained_buffers, retained_executables, + recorded); + retain_prepared_command_list_resources(prepared.commands, retained_buffers, retained_executables, recorded); +} + +static RecordedCommandGraph record_prepared_command_graph(const CommandProgramExecutionContext & context, + const CommandProgram & commands, + const PreparedCommandProgram & prepared, + const TransientArenaAllocationRef & allocation) { + RecordedCommandGraph recorded; + if (context.device == nullptr) { + recorded.status.log("missing HRX device for graph replay"); + return recorded; + } + + hrx_graph_t graph = nullptr; + if (ErrorResult error = take_status(hrx_graph_create(context.device, 0, &graph))) { + recorded.status.log("create HRX graph replay: %s", error->c_str()); + return recorded; + } + recorded.graph = graph; + + GraphDependencyPlanner dependencies; + recorded.status.append(record_completion_counter_fill(recorded.graph, dependencies, commands, allocation)); + if (!recorded.status.success()) { + return recorded; + } + recorded.status.append(record_prepared_command_list(recorded.graph, dependencies, prepared.initialization_commands, + recorded.dispatch_count)); + if (!recorded.status.success()) { + return recorded; + } + const std::vector ordered = order_prepared_commands_for_graph(prepared.commands); // graph-record-order.h + recorded.status.append( + record_prepared_command_list(recorded.graph, dependencies, ordered.empty() ? prepared.commands : ordered, recorded.dispatch_count)); + if (!recorded.status.success()) { + return recorded; + } + + // Unretained HRX command buffers require graph-lifetime resources. + retain_prepared_command_graph_resources(prepared, allocation, recorded); + + hrx_graph_exec_t exec = nullptr; + if (ErrorResult error = take_status(hrx_graph_instantiate(recorded.graph, 0, &exec))) { + recorded.status.log("instantiate HRX graph replay: %s", error->c_str()); + return recorded; + } + recorded.exec = exec; + recorded.bound_transient_arena_allocation_id = prepared.bound_transient_arena_allocation_id; + return recorded; +} + +static bool execute_prepared_kernel_command_via_graph(const CommandProgramExecutionContext & context, + const PreparedCommand & command) { + hrx_graph_t graph = nullptr; + if (take_status(hrx_graph_create(context.device, 0, &graph))) { + return false; + } + GraphDependencyPlanner dependencies; + Status status = record_prepared_kernel_command(graph, dependencies, command); + hrx_graph_exec_t exec = nullptr; + if (status.success()) { + if (ErrorResult error = take_status(hrx_graph_instantiate(graph, 0, &exec))) { + status.log("instantiate diagnostic graph: %s", error->c_str()); + } + } + if (status.success()) { + if (ErrorResult error = take_status(hrx_graph_exec_launch(exec, context.stream))) { + status.log("launch diagnostic graph: %s", error->c_str()); + } + } + if (status.success()) { + if (ErrorResult error = take_status(hrx_stream_synchronize(context.stream))) { + status.log("synchronize diagnostic graph: %s", error->c_str()); + } + } + if (exec != nullptr) { + hrx_graph_exec_release(exec); + } + hrx_graph_release(graph); + return status.success(); +} + +static bool execute_prepared_command_prefix_via_graph(const CommandProgramExecutionContext & context, + const std::vector & commands, + size_t prefix_count) { + if (prefix_count == 0 || prefix_count > commands.size()) { + return prefix_count == 0; + } + + hrx_graph_t graph = nullptr; + if (take_status(hrx_graph_create(context.device, 0, &graph))) { + return false; + } + GraphDependencyPlanner dependencies; + Status status; + for (size_t i = 0; i < prefix_count && status.success(); ++i) { + status.append(record_prepared_kernel_command(graph, dependencies, commands[i])); + } + + hrx_graph_exec_t exec = nullptr; + if (status.success()) { + if (ErrorResult error = take_status(hrx_graph_instantiate(graph, 0, &exec))) { + status.log("instantiate diagnostic prefix graph: %s", error->c_str()); + } + } + if (status.success()) { + if (ErrorResult error = take_status(hrx_graph_exec_launch(exec, context.stream))) { + status.log("launch diagnostic prefix graph: %s", error->c_str()); + } + } + if (status.success()) { + if (ErrorResult error = take_status(hrx_stream_synchronize(context.stream))) { + status.log("synchronize diagnostic prefix graph: %s", error->c_str()); + } + } + if (exec != nullptr) { + hrx_graph_exec_release(exec); + } + hrx_graph_release(graph); + return status.success(); +} + +} // namespace + +bool debug_serial_command_execution_enabled() { + return environment_flag_enabled("GGML_HRX_DEBUG_SERIAL_EXECUTION"); +} + +PreparedCommandProgram prepare_command_program(const CommandProgramExecutionContext & context, + const CommandProgram & commands, + const CommandProgramBindings & bindings) { + PreparedCommandProgram prepared; + prepared.status = command_program_metadata_context_valid(context); + if (!prepared.status.success()) { + return prepared; + } + + const VerificationResult verification = verify_command_program(commands, *context.corpus, context.target); + if (!verification.valid()) { + prepared.status.append(verification.status); + return prepared; + } + if (!bindings.valid()) { + prepared.status.append(bindings.status); + return prepared; + } + + TransientArenaAllocationRef transient_allocation; + prepared.status = ensure_transient_arena(context, commands, transient_allocation); + if (!prepared.status.success()) { + return prepared; + } + + prepared.status = command_program_preparation_context_valid(context); + if (!prepared.status.success()) { + return prepared; + } + + const CommandProgramBindings materialized_bindings = + materialize_host_bindings(context, commands, bindings, prepared); + if (!materialized_bindings.valid()) { + prepared.status.append(materialized_bindings.status); + return prepared; + } + + const TransientArenaAllocationRef * transient_allocation_ptr = + commands.transients.arena_size == 0 ? nullptr : &transient_allocation; + const ResolvedCommandProgram resolved = + resolve_command_program_bindings(commands, materialized_bindings, transient_allocation_ptr); + if (!resolved.valid()) { + prepared.status.append(resolved.status); + return prepared; + } + std::vector initialization_executable_refs; + std::vector command_executable_refs; + prepare_command_list(context, resolved.initialization_commands, prepared.initialization_commands, + initialization_executable_refs, prepared.status); + prepare_command_list(context, resolved.commands, prepared.commands, command_executable_refs, prepared.status); + materialize_command_list_executables(context, prepared.initialization_commands, initialization_executable_refs, + prepared.status); + materialize_command_list_executables(context, prepared.commands, command_executable_refs, prepared.status); + prepared.bound_transient_arena_allocation_id = transient_allocation.allocation_id; + if (prepared.status.success()) { + prepared.status.append(prepare_program_constant_buffers(context, commands, prepared)); + } + return prepared; +} + +bool bind_prepared_command_program_transients(const CommandProgram & commands, + const TransientArenaAllocationRef & transient_allocation, + PreparedCommandProgram & prepared) { + if (!prepared.valid()) { + return false; + } + if (commands.transients.arena_size == 0) { + prepared.bound_transient_arena_allocation_id = kInvalidTransientArenaAllocationId; + return true; + } + if (transient_allocation.buffer == nullptr || + transient_allocation.allocation_id == kInvalidTransientArenaAllocationId) { + GGML_LOG_ERROR("%s: missing transient arena allocation\n", __func__); + return false; + } + if (prepared.bound_transient_arena_allocation_id == transient_allocation.allocation_id) { + return true; + } + if (!bind_prepared_command_list_transients(commands, transient_allocation, prepared.initialization_commands) || + !bind_prepared_command_list_transients(commands, transient_allocation, prepared.commands)) { + return false; + } + prepared.bound_transient_arena_allocation_id = transient_allocation.allocation_id; + return true; +} + +bool bind_and_execute_prepared_command_program(const CommandProgramExecutionContext & context, + const CommandProgram & commands, + const CommandProgramBindings & bindings, + PreparedCommandProgram & prepared) { + if (!prepared.valid()) { + return execute_prepared_command_program(context, prepared); + } + const DebugSerialExecutionTrace debug(context, commands); + debug.log_program_begin(); + HRX_DEBUG_SERIAL_SYNC(debug, "host-staging-rebind-pre"); + Status rebind_status = rebind_prepared_host_staging(context, bindings, prepared); + if (!rebind_status.success()) { + GGML_LOG_ERROR("%s: %s\n", __func__, status_first_error(rebind_status)); + return false; + } + HRX_DEBUG_SERIAL_SYNC(debug, "host-staging-rebind-post"); + if (commands.transients.arena_size == 0) { + HRX_DEBUG_SERIAL_SYNC(debug, "constant-initialization-pre"); + Status status = initialize_command_program_constants(context, commands, {}, prepared); + if (!status.success()) { + GGML_LOG_ERROR("%s: %s\n", __func__, status_first_error(status)); + return false; + } + HRX_DEBUG_SERIAL_SYNC(debug, "constant-initialization-post"); + HRX_DEBUG_SERIAL_SYNC(debug, "completion-counter-initialization-pre"); + status = initialize_command_program_completion_counters(context, commands, {}); + if (!status.success()) { + GGML_LOG_ERROR("%s: %s\n", __func__, status_first_error(status)); + return false; + } + HRX_DEBUG_SERIAL_SYNC(debug, "completion-counter-initialization-post"); + HRX_DEBUG_SERIAL_SYNC(debug, "transient-bind-pre"); + if (!bind_prepared_command_program_transients(commands, {}, prepared)) { + return false; + } + HRX_DEBUG_SERIAL_SYNC(debug, "transient-bind-post"); + const bool success = execute_prepared_command_program(context, prepared); + debug.log_program_end(success); + return success; + } + + Status status = command_program_transient_context_valid(context, commands); + if (!status.success()) { + GGML_LOG_ERROR("%s: %s\n", __func__, status_first_error(status)); + return false; + } + + HRX_DEBUG_SERIAL_SYNC(debug, "transient-arena-acquire-pre"); + TransientArena::AllocationLease lease = context.transient_arena->acquire_allocation_lease(); + status = lease.ensure_capacity(context.device, context.stream, commands.transients.arena_size); + if (!status.success()) { + GGML_LOG_ERROR("%s: %s\n", __func__, status_first_error(status)); + return false; + } + HRX_DEBUG_SERIAL_SYNC(debug, "transient-arena-acquire-post"); + HRX_DEBUG_SERIAL_SYNC(debug, "transient-bind-pre"); + if (!bind_prepared_command_program_transients(commands, lease.current_allocation(), prepared)) { + return false; + } + HRX_DEBUG_SERIAL_SYNC(debug, "transient-bind-post"); + HRX_DEBUG_SERIAL_SYNC(debug, "constant-initialization-pre"); + status = initialize_command_program_constants(context, commands, lease.current_allocation(), prepared); + if (!status.success()) { + GGML_LOG_ERROR("%s: %s\n", __func__, status_first_error(status)); + return false; + } + HRX_DEBUG_SERIAL_SYNC(debug, "constant-initialization-post"); + HRX_DEBUG_SERIAL_SYNC(debug, "completion-counter-initialization-pre"); + status = initialize_command_program_completion_counters(context, commands, lease.current_allocation()); + if (!status.success()) { + GGML_LOG_ERROR("%s: %s\n", __func__, status_first_error(status)); + return false; + } + HRX_DEBUG_SERIAL_SYNC(debug, "completion-counter-initialization-post"); + const bool success = execute_prepared_command_program(context, prepared); + debug.log_program_end(success); + return success; +} + +RecordedCommandGraphExecutionResult bind_and_launch_recorded_command_graph( + const CommandProgramExecutionContext & context, + const CommandProgram & commands, + const CommandProgramBindings & bindings, + PreparedCommandProgram & prepared, + RecordedCommandGraph & recorded) { + RecordedCommandGraphExecutionResult result; + result.event = HrxGraphReplayEvent::Ineligible; + + if (!prepared.valid()) { + result.status.append(prepared.status); + if (result.status.success()) { + result.status.log("invalid prepared command program"); + } + return result; + } + if (!prepared_execution_context_valid(context)) { + result.status.log("missing HRX stream"); + result.event = HrxGraphReplayEvent::BuildFailed; + return result; + } + + Status rebind_status = rebind_prepared_host_staging(context, bindings, prepared); + if (!rebind_status.success()) { + result.status.append(rebind_status); + result.event = HrxGraphReplayEvent::BuildFailed; + return result; + } + + TransientArenaAllocationRef transient_allocation; + TransientArena::AllocationLease lease; + if (commands.transients.arena_size == 0) { + if (!bind_prepared_command_program_transients(commands, {}, prepared)) { + result.status.log("bind transient-free prepared command program failed"); + result.event = HrxGraphReplayEvent::BuildFailed; + return result; + } + } else { + Status status = command_program_transient_context_valid(context, commands); + if (!status.success()) { + result.status.append(status); + result.event = HrxGraphReplayEvent::BuildFailed; + return result; + } + lease = context.transient_arena->acquire_allocation_lease(); + status = lease.ensure_capacity(context.device, context.stream, commands.transients.arena_size); + if (!status.success()) { + result.status.append(status); + result.event = HrxGraphReplayEvent::BuildFailed; + return result; + } + transient_allocation = lease.current_allocation(); + if (!bind_prepared_command_program_transients(commands, transient_allocation, prepared)) { + result.status.log("bind prepared command program transients for graph replay failed"); + result.event = HrxGraphReplayEvent::BuildFailed; + return result; + } + } + + const bool had_recorded = recorded.valid(); + + const char * diagnostic_kernel = std::getenv("GGML_HRX_DIAGNOSTIC_GRAPH_KERNEL_ID"); + if (diagnostic_kernel != nullptr && diagnostic_kernel[0] != '\0') { + const uint64_t diagnostic_graph_kernel_id = std::strtoull(diagnostic_kernel, nullptr, 0); + Status diagnostic_upload_status = upload_prepared_host_staging(context, prepared); + if (!diagnostic_upload_status.success()) { + result.status.append(diagnostic_upload_status); + result.event = HrxGraphReplayEvent::LaunchFailed; + return result; + } + auto execute_diagnostic_list = [&](const std::vector & command_list) { + for (const PreparedCommand & command : command_list) { + const bool ok = (diagnostic_graph_kernel_id == 0 || + command.kernel.specialization.kernel_id == diagnostic_graph_kernel_id) ? + execute_prepared_kernel_command_via_graph(context, command) : + execute_prepared_kernel_command(context, command); + if (!ok) { + return false; + } + } + return true; + }; + if (!execute_diagnostic_list(prepared.initialization_commands) || !execute_diagnostic_list(prepared.commands)) { + result.status.log("diagnostic mixed graph execution failed"); + result.event = HrxGraphReplayEvent::LaunchFailed; + return result; + } + result.success = true; + return result; + } + + const char * diagnostic_program_size = std::getenv("GGML_HRX_DIAGNOSTIC_GRAPH_PROGRAM_COMMANDS"); + const char * diagnostic_prefix = std::getenv("GGML_HRX_DIAGNOSTIC_GRAPH_PREFIX_COMMANDS"); + if (diagnostic_program_size != nullptr && diagnostic_program_size[0] != '\0' && diagnostic_prefix != nullptr && + diagnostic_prefix[0] != '\0') { + const size_t target_size = std::strtoull(diagnostic_program_size, nullptr, 0); + const size_t prefix_size = std::strtoull(diagnostic_prefix, nullptr, 0); + Status upload_status = upload_prepared_host_staging(context, prepared); + if (!upload_status.success()) { + result.status.append(upload_status); + result.event = HrxGraphReplayEvent::LaunchFailed; + return result; + } + + bool ok = execute_prepared_command_list(context, prepared.initialization_commands); + if (ok && prepared.commands.size() == target_size) { + ok = execute_prepared_command_prefix_via_graph(context, prepared.commands, prefix_size); + for (size_t i = prefix_size; ok && i < prepared.commands.size(); ++i) { + ok = execute_prepared_kernel_command(context, prepared.commands[i]); + } + } else if (ok) { + ok = execute_prepared_command_list(context, prepared.commands); + } + if (!ok) { + result.status.log("diagnostic prefix graph execution failed"); + result.event = HrxGraphReplayEvent::LaunchFailed; + return result; + } + Status download_status = download_prepared_host_staging(context, prepared); + if (!download_status.success()) { + result.status.append(download_status); + result.event = HrxGraphReplayEvent::LaunchFailed; + return result; + } + result.success = true; + return result; + } + + result.transient_allocation_changed = + had_recorded && recorded.bound_transient_arena_allocation_id != prepared.bound_transient_arena_allocation_id; + if (!had_recorded || result.transient_allocation_changed) { + result.event = result.transient_allocation_changed ? HrxGraphReplayEvent::RebuildTransient : + HrxGraphReplayEvent::MissBuild; + const uint64_t build_start_ns = hrx_graph_replay_now_ns(); + RecordedCommandGraph rebuilt = record_prepared_command_graph(context, commands, prepared, transient_allocation); + result.build_ns = hrx_graph_replay_now_ns() - build_start_ns; + if (!rebuilt.valid()) { + result.status.append(rebuilt.status); + result.event = HrxGraphReplayEvent::BuildFailed; + return result; + } + recorded = std::move(rebuilt); + } else { + result.event = HrxGraphReplayEvent::Hit; + } + + const uint64_t launch_start_ns = hrx_graph_replay_now_ns(); + Status upload_status = upload_prepared_host_staging(context, prepared); + if (!upload_status.success()) { + result.launch_ns = hrx_graph_replay_now_ns() - launch_start_ns; + result.status.append(upload_status); + result.event = HrxGraphReplayEvent::LaunchFailed; + return result; + } + if (ErrorResult error = take_status(hrx_graph_exec_launch(recorded.exec, context.stream))) { + result.launch_ns = hrx_graph_replay_now_ns() - launch_start_ns; + result.status.log("launch HRX graph replay: %s", error->c_str()); + result.event = HrxGraphReplayEvent::LaunchFailed; + return result; + } + Status replay_flush_status = flush_stream_commands(context, "flush HRX graph replay commands"); + if (!replay_flush_status.success()) { + result.launch_ns = hrx_graph_replay_now_ns() - launch_start_ns; + result.status.append(replay_flush_status); + result.event = HrxGraphReplayEvent::LaunchFailed; + return result; + } + Status replay_wait_status = wait_stream_commands(context, "wait for HRX graph replay commands", recorded.exec); + if (!replay_wait_status.success()) { + result.launch_ns = hrx_graph_replay_now_ns() - launch_start_ns; + result.status.append(replay_wait_status); + result.event = HrxGraphReplayEvent::LaunchFailed; + return result; + } + const char * diagnostic_sync_size = std::getenv("GGML_HRX_DIAGNOSTIC_GRAPH_SYNC_PROGRAM_COMMANDS"); + const bool diagnostic_sync = std::getenv("GGML_HRX_DIAGNOSTIC_GRAPH_SYNC") != nullptr || + (diagnostic_sync_size != nullptr && diagnostic_sync_size[0] != '\0' && + recorded.dispatch_count == std::strtoull(diagnostic_sync_size, nullptr, 0)); + if (diagnostic_sync) { + if (ErrorResult error = take_status(hrx_stream_synchronize(context.stream))) { + result.launch_ns = hrx_graph_replay_now_ns() - launch_start_ns; + result.status.log("synchronize diagnostic graph replay: %s", error->c_str()); + result.event = HrxGraphReplayEvent::LaunchFailed; + return result; + } + mark_graph_replay_synchronized(context); + } + Status download_status = download_prepared_host_staging(context, prepared, false); + if (!download_status.success()) { + result.launch_ns = hrx_graph_replay_now_ns() - launch_start_ns; + result.status.append(download_status); + result.event = HrxGraphReplayEvent::LaunchFailed; + return result; + } + result.launch_ns = hrx_graph_replay_now_ns() - launch_start_ns; + result.dispatch_count = recorded.dispatch_count; + result.success = true; + return result; +} + +bool execute_prepared_command_program(const CommandProgramExecutionContext & context, + const PreparedCommandProgram & commands) { + if (!commands.valid()) { + GGML_LOG_ERROR("%s: invalid HRX prepared command program: %s\n", __func__, status_first_error(commands.status)); + return false; + } + if (!prepared_execution_context_valid(context)) { + return false; + } + const DebugSerialExecutionTrace debug(context, commands); + HRX_DEBUG_SERIAL_SYNC(debug, "host-staging-upload-pre"); + Status upload_status = upload_prepared_host_staging(context, commands); + if (!upload_status.success()) { + GGML_LOG_ERROR("%s: %s\n", __func__, status_first_error(upload_status)); + return false; + } + HRX_DEBUG_SERIAL_SYNC(debug, "host-staging-upload-post"); + if (!execute_prepared_command_list_serial(context, commands.initialization_commands, "init", commands.commands.size()) || + !execute_prepared_command_list_serial(context, commands.commands, "main", commands.commands.size())) { + return false; + } + HRX_DEBUG_SERIAL_SYNC(debug, "host-staging-download-pre"); + Status download_status = download_prepared_host_staging(context, commands); + if (!download_status.success()) { + GGML_LOG_ERROR("%s: %s\n", __func__, status_first_error(download_status)); + return false; + } + HRX_DEBUG_SERIAL_SYNC(debug, "host-staging-download-post"); + // Uploads from registered host buffers are device-timeline copies that read + // the host memory when they execute. ggml writes the next graph's inputs + // into that same memory as soon as graph_compute returns, so wait for this + // program here, as the graph replay path does after its launch. + bool reads_registered_host_memory = false; + for (const HostStagingBuffer & staging : commands.host_staging) { + reads_registered_host_memory = + reads_registered_host_memory || (staging.upload && staging.source_host_buffer.valid()); + } + if (reads_registered_host_memory) { + static thread_local WaitHistory registered_upload_wait_history; + if (ErrorResult error = + take_status(stream_synchronize_sleeping(context.stream, registered_upload_wait_history))) { + GGML_LOG_ERROR("%s: wait for HRX command program reading host memory: %s\n", __func__, error->c_str()); + return false; + } + mark_graph_replay_synchronized(context); + } + return true; +} + +bool execute_command_program(const CommandProgramExecutionContext & context, + const CommandProgram & commands, + const CommandProgramBindings & bindings) { + PreparedCommandProgram prepared = prepare_command_program(context, commands, bindings); + return bind_and_execute_prepared_command_program(context, commands, bindings, prepared); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/command-program-executor.h b/ggml/src/ggml-hrx/runtime/command-program-executor.h new file mode 100644 index 000000000000..437c5531ae02 --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/command-program-executor.h @@ -0,0 +1,179 @@ +#pragma once + +#include "dispatch/command-program-bindings.h" +#include "dispatch/command-program-resolver.h" +#include "dispatch/command-program.h" +#include "kernel-corpus/kernel-corpus.h" +#include "runtime/graph-replay.h" +#include "runtime/host-memory.h" + +#include +#include +#include +#include +#include + +typedef struct hrx_device_s * hrx_device_t; +typedef struct hrx_stream_s * hrx_stream_t; +typedef struct hrx_buffer_s * hrx_buffer_t; +typedef struct hrx_executable_s * hrx_executable_t; +typedef struct hrx_graph_s * hrx_graph_t; +typedef struct hrx_graph_exec_s * hrx_graph_exec_t; +struct ggml_hrx_loom_jit_amdgpu; + +namespace ggml::hrx { + +class KernelExecutableCache; +struct KernelExecutable; +class TransientArena; +class HostBufferRegistry; + +struct PendingHostWriteback { + void * host_destination = nullptr; + const void * mapped_source = nullptr; + size_t size = 0; + hrx_buffer_t retained_buffer = nullptr; +}; + +struct GraphReplayStreamState { + GraphReplayStreamState() = default; + ~GraphReplayStreamState(); + + GraphReplayStreamState(const GraphReplayStreamState &) = delete; + GraphReplayStreamState & operator=(const GraphReplayStreamState &) = delete; + + void mark_stream_synchronized(); + void add_host_writeback(void * host_destination, const void * mapped_source, size_t size, hrx_buffer_t buffer); + void clear(); + + // True while a host download staged by an earlier replay has not been + // published to host memory yet (see mark_stream_synchronized). + bool has_pending_host_writebacks() const { return !pending_host_writebacks_.empty(); } + + private: + std::vector pending_host_writebacks_; +}; + +struct CommandProgramExecutionContext { + hrx_device_t device = nullptr; + hrx_stream_t stream = nullptr; + const char * target = nullptr; + const KernelCorpus * corpus = nullptr; + KernelExecutableCache * kernel_executables = nullptr; + TransientArena * transient_arena = nullptr; + HostTransferManager * host_transfers = nullptr; + HostWeightCache * host_weights = nullptr; + HostBufferRegistry * host_buffers = nullptr; + GraphReplayStreamState * graph_replay_state = nullptr; +}; + +struct PreparedCommandBinding { + CommandBinding binding; + ResolvedBufferRef ref; +}; + +struct PreparedKernelCommand { + KernelSpecialization specialization; + std::shared_ptr executable; + std::vector constants; + std::vector bindings; +}; + +struct PreparedCommand { + uint32_t ordinal = 0; + CommandKind kind = CommandKind::Invalid; + PreparedKernelCommand kernel; +}; + +struct PreparedProgramConstantBuffer { + ValueId value; + std::string name; + hrx_buffer_t buffer = nullptr; + size_t size = 0; + + PreparedProgramConstantBuffer() = default; + ~PreparedProgramConstantBuffer(); + + PreparedProgramConstantBuffer(PreparedProgramConstantBuffer && other) noexcept; + PreparedProgramConstantBuffer & operator=(PreparedProgramConstantBuffer && other) noexcept; + + PreparedProgramConstantBuffer(const PreparedProgramConstantBuffer &) = delete; + PreparedProgramConstantBuffer & operator=(const PreparedProgramConstantBuffer &) = delete; +}; + +struct PreparedCommandProgram { + std::vector initialization_commands; + std::vector commands; + std::vector host_staging; + std::vector resident_host_weights; + std::vector program_constants; + Status status; + uint64_t bound_transient_arena_allocation_id = kInvalidTransientArenaAllocationId; + + bool valid() const { return status.success(); } +}; + +struct RecordedCommandGraph { + hrx_graph_t graph = nullptr; + hrx_graph_exec_t exec = nullptr; + std::vector retained_buffers; + std::vector retained_executables; + uint64_t bound_transient_arena_allocation_id = kInvalidTransientArenaAllocationId; + size_t dispatch_count = 0; + Status status; + + RecordedCommandGraph() = default; + ~RecordedCommandGraph(); + + RecordedCommandGraph(RecordedCommandGraph && other) noexcept; + RecordedCommandGraph & operator=(RecordedCommandGraph && other) noexcept; + + RecordedCommandGraph(const RecordedCommandGraph &) = delete; + RecordedCommandGraph & operator=(const RecordedCommandGraph &) = delete; + + bool valid() const { return status.success() && exec != nullptr; } +}; + +struct RecordedCommandGraphExecutionResult { + bool success = false; + Status status; + HrxGraphReplayEvent event = HrxGraphReplayEvent::Disabled; + std::string ineligible_reason; + size_t dispatch_count = 0; + bool transient_allocation_changed = false; + uint64_t build_ns = 0; + uint64_t launch_ns = 0; + + uint64_t total_ns() const { return build_ns + launch_ns; } +}; + +PreparedCommandProgram prepare_command_program(const CommandProgramExecutionContext & context, + const CommandProgram & commands, + const CommandProgramBindings & bindings); + +bool execute_prepared_command_program(const CommandProgramExecutionContext & context, + const PreparedCommandProgram & commands); + +bool bind_prepared_command_program_transients(const CommandProgram & commands, + const TransientArenaAllocationRef & transient_allocation, + PreparedCommandProgram & prepared); + +bool bind_and_execute_prepared_command_program(const CommandProgramExecutionContext & context, + const CommandProgram & commands, + const CommandProgramBindings & bindings, + PreparedCommandProgram & prepared); + +RecordedCommandGraphExecutionResult bind_and_launch_recorded_command_graph( + const CommandProgramExecutionContext & context, + const CommandProgram & commands, + const CommandProgramBindings & bindings, + PreparedCommandProgram & prepared, + RecordedCommandGraph & recorded); + +bool debug_serial_command_execution_enabled(); + +bool execute_command_program(const CommandProgramExecutionContext & context, + const CommandProgram & commands, + const CommandProgramBindings & bindings); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/graph-executor.cpp b/ggml/src/ggml-hrx/runtime/graph-executor.cpp new file mode 100644 index 000000000000..be928a554f0f --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/graph-executor.cpp @@ -0,0 +1,148 @@ +#include "graph-executor.h" + +#include "backend-buffer-binding.h" +#include "ggml-impl.h" +#include "runtime/graph-program-cache-limit.h" +#include "runtime/kernel-executable-cache.h" +#include "runtime/prepared-command-program-cache.h" +#include "runtime/transient-arena.h" + +#include +#include + +namespace ggml::hrx { + +GraphExecutor::GraphExecutor(ggml_backend_hrx_context & context) : context_(context) {} + +Status GraphExecutor::context_valid_for_graph_programs() const { + Status status; + if (context_.device == nullptr) { + status.log("missing HRX device context"); + } else if (context_.device->architecture.empty()) { + status.log("missing HRX target"); + } + return status; +} + +Status GraphExecutor::context_valid_for_execution() const { + return context_valid_for_graph_programs(); +} + +GraphSupportResult GraphExecutor::can_execute(const ggml_cgraph & graph) const { + GraphSupportResult result; + if (graph.n_nodes == 0) { + result.supported = true; + return result; + } + result.status = context_valid_for_graph_programs(); + if (!result.status.success()) { + return result; + } + const KernelCorpus & corpus = get_qwen_kernel_corpus(); + const GraphProgramSupportResult support = + context_.graph_programs.check_support(graph, corpus, context_.device->architecture); + result.supported = support.supported; + result.status.append(support.status); + return result; +} + +CommandProgramBindings GraphExecutor::bind_external_value_buffers(const GraphProgramMatch & match) const { + std::vector bindings; + Status status; + bindings.reserve(match.external_bindings.size()); + for (const GraphProgramExternalBinding & external : match.external_bindings) { + ValueBufferBinding value_binding; + CommandProgramBinding binding; + binding.value = external.value; + if (ggml_backend_hrx_resolve_value_buffer(external.tensor, value_binding)) { + binding.buffer = value_binding.buffer; + binding.host_data = value_binding.host_data; + binding.offset = value_binding.offset; + binding.length = value_binding.length; + binding.identity = value_binding.identity; + binding.generation = value_binding.generation; + binding.capacity = value_binding.capacity; + binding.weight = value_binding.weight; + binding.empty_value = ggml_nbytes(external.tensor) == 0; + binding.graph_input = (external.tensor->flags & GGML_TENSOR_FLAG_INPUT) != 0 || + (external.tensor->view_src != nullptr && + (external.tensor->view_src->flags & GGML_TENSOR_FLAG_INPUT) != 0); + } else { + status.log("external value %d is not bound", external.value.value); + } + bindings.push_back(binding); + } + return CommandProgramBindings::from_bindings(std::move(bindings), status); +} + +GraphExecutionResult GraphExecutor::execute(const ggml_cgraph & graph) const { + GraphExecutionResult result; + if (graph.n_nodes == 0) { + result.code = GGML_STATUS_SUCCESS; + return result; + } + result.status = context_valid_for_execution(); + if (!result.status.success()) { + return result; + } + + const KernelCorpus & corpus = get_qwen_kernel_corpus(); + const uint64_t builds_before = context_.graph_programs.stats().builds; + GraphProgramLookup lookup = context_.graph_programs.get_or_build(graph, corpus, context_.device->architecture); + if (lookup.valid()) { + trim_graph_programs(context_, lookup.program, builds_before); // graph-program-cache-limit.h + } + if (!lookup.valid()) { + result.status.append(lookup.status); + result.status.append(lookup.match.status); + if (result.status.success()) { + result.status.log("build HRX graph program failed"); + } + return result; + } + + const bool use_graph_prepared = lookup.program->can_use_prepared_fast_path(graph); + GraphProgramMatch binding_match = std::move(lookup.match); + if (use_graph_prepared && lookup.program->has_prepared_program()) { + binding_match = lookup.program->match_host_staging_graph(graph); + if (!binding_match.valid()) { + result.status.append(binding_match.status); + return result; + } + } + + CommandProgramBindings bindings = bind_external_value_buffers(binding_match); + if (!bindings.valid()) { + result.status.append(bindings.status); + return result; + } + const CommandProgramExecutionContext execution_context = { + context_.device->device, + context_.stream, + context_.device->architecture.c_str(), + &corpus, + &context_.kernel_executables, + &context_.transient_arena, + &context_.host_transfers, + &context_.host_weights, + &context_.device->host_buffers, + &context_.graph_replay_state, + }; + const PreparedCommandProgramCacheExecutionResult execution = + use_graph_prepared ? lookup.program->execute_with_result(execution_context, bindings) : + context_.prepared_programs.execute_with_result(execution_context, lookup.program->uid(), + lookup.program->command_shape_hash(), + lookup.program->commands(), bindings); + if (!execution.success) { + result.status.append(execution.status); + if (result.status.success()) { + result.status.log("execute HRX command program failed"); + } + return result; + } + + result.code = GGML_STATUS_SUCCESS; + return result; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/graph-executor.h b/ggml/src/ggml-hrx/runtime/graph-executor.h new file mode 100644 index 000000000000..a2759ea71875 --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/graph-executor.h @@ -0,0 +1,43 @@ +#pragma once + +#include "backend-context.h" +#include "dispatch/command-program-bindings.h" +#include "ggml.h" +#include "runtime/graph-program-cache.h" +#include "status.h" + +struct ggml_cgraph; + +namespace ggml::hrx { + +struct GraphSupportResult { + bool supported = false; + Status status; + + bool success() const { return supported && status.success(); } +}; + +struct GraphExecutionResult { + enum ggml_status code = GGML_STATUS_FAILED; + Status status; + + bool success() const { return code == GGML_STATUS_SUCCESS && status.success(); } +}; + +class GraphExecutor { + public: + explicit GraphExecutor(ggml_backend_hrx_context & context); + + GraphSupportResult can_execute(const ggml_cgraph & graph) const; + GraphExecutionResult execute(const ggml_cgraph & graph) const; + + private: + Status context_valid_for_graph_programs() const; + Status context_valid_for_execution() const; + + CommandProgramBindings bind_external_value_buffers(const GraphProgramMatch & match) const; + + ggml_backend_hrx_context & context_; +}; + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/graph-program-cache-limit.cpp b/ggml/src/ggml-hrx/runtime/graph-program-cache-limit.cpp new file mode 100644 index 000000000000..42613f0f03a7 --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/graph-program-cache-limit.cpp @@ -0,0 +1,67 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "graph-program-cache-limit.h" + +#include "backend-context.h" +#include "ggml-impl.h" +#include "runtime/hrx-sleeping-wait.h" + +#include +#include + +namespace ggml::hrx { + +size_t graph_program_cache_limit() { + static const size_t limit = [] { + const char * value = std::getenv("GGML_HRX_GRAPH_PROGRAM_CACHE"); + if (value == nullptr || *value == '\0') { + return size_t{ 64 }; + } + char * end = nullptr; + const unsigned long long parsed = std::strtoull(value, &end, 10); + return end != nullptr && *end == '\0' ? static_cast(parsed) : size_t{ 64 }; + }(); + return limit; +} + +bool trim_graph_programs(ggml_backend_hrx_context & context, GraphProgram * current, uint64_t builds_before) { + const size_t limit = graph_program_cache_limit(); + if (limit == 0 || current == nullptr || context.graph_programs.stats().builds == builds_before) { + return true; + } + const size_t held = context.graph_programs.size(); + if (held <= limit) { + return true; + } + static thread_local WaitHistory history; + hrx_status_t status = stream_synchronize_sleeping(context.stream, history); + if (!hrx_status_is_ok(status)) { + GGML_LOG_ERROR("%s: stream wait failed; graph programs kept\n", __func__); + hrx_status_ignore(status); + return false; + } + context.graph_replay_state.mark_stream_synchronized(); + context.graph_programs.retain_only(current); + context.prepared_programs.clear(); + // one line per flush: it should come every `limit` new shapes, never during steady decode + static std::atomic flushes{ 0 }; + GGML_LOG_WARN("%s: HRX graph program cache over its limit (%zu programs, limit %zu): kept the current one, " + "flush %llu\n", + __func__, held, limit, static_cast(++flushes)); + return true; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/graph-program-cache-limit.h b/ggml/src/ggml-hrx/runtime/graph-program-cache-limit.h new file mode 100644 index 000000000000..8903284723a5 --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/graph-program-cache-limit.h @@ -0,0 +1,43 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +// A bound on the graph programs one HRX backend keeps. Every new graph shape (each new prompt or +// ubatch length) builds a GraphProgram that retains its HRX buffers, recorded graph and transient +// arena binding, and GraphProgramCache never evicted: a llama-server answering prompts of varying +// length grew about 1 GiB of GTT per request (ZAYA1-8B, 2026-10-01) until the box ran out of memory. + +#include +#include + +struct ggml_backend_hrx_context; + +namespace ggml::hrx { + +class GraphProgram; + +// GGML_HRX_GRAPH_PROGRAM_CACHE: programs kept per backend (default 64; 0 keeps them all, the old +// behaviour). +size_t graph_program_cache_limit(); + +// After a lookup: if it built a new program (the cache's build count moved past builds_before) and +// the cache now holds more than the limit, wait for the stream (no command of a dropped program may +// still be running), then drop every cached program except `current`, and the prepared programs. +// A lookup that reused a program, by uid or by structure, never drops anything, so steady decode +// does not churn. Returns false only if the stream wait failed; the caches are then kept. +bool trim_graph_programs(ggml_backend_hrx_context & context, GraphProgram * current, uint64_t builds_before); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/graph-program-cache.cpp b/ggml/src/ggml-hrx/runtime/graph-program-cache.cpp new file mode 100644 index 000000000000..06d38d78fccb --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/graph-program-cache.cpp @@ -0,0 +1,866 @@ +#include "graph-program-cache.h" + +#include "dispatch/command-program-dump.h" +#include "dispatch/dispatch-scheduler.h" +#include "ggml-impl.h" +#include "ggml.h" + +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static const ggml_tensor * tensor_storage_root(const ggml_tensor * tensor) { + while (tensor != nullptr && tensor->view_src != nullptr) { + tensor = tensor->view_src; + } + return tensor; +} + +static size_t tensor_storage_offset(const ggml_tensor * tensor) { + return tensor != nullptr && tensor->view_src != nullptr ? tensor->view_offs : 0; +} + +static bool tensor_storage_relative_offset(const ggml_tensor * source, const ggml_tensor * tensor, size_t & offset) { + if (source == nullptr || tensor == nullptr || tensor_storage_root(source) != tensor_storage_root(tensor)) { + return false; + } + const size_t source_offset = tensor_storage_offset(source); + const size_t tensor_offset = tensor_storage_offset(tensor); + if (tensor_offset < source_offset) { + return false; + } + const size_t relative_offset = tensor_offset - source_offset; + if (ggml_nbytes(source) == 0) { + // A zero-byte (empty) source can only contain a zero-byte tensor; the byte + // offset of a degenerate view into an empty tensor is meaningless (#95). + offset = relative_offset; + return ggml_nbytes(tensor) == 0; + } + if (relative_offset > ggml_nbytes(source) || ggml_nbytes(tensor) > ggml_nbytes(source) - relative_offset) { + return false; + } + offset = relative_offset; + return true; +} + +static bool tensor_alias_matches(const ValueMap & values, + const Value & value, + const ggml_tensor * tensor, + const std::vector & tensor_by_value) { + const bool tensor_alias = tensor->view_src != nullptr; + const bool value_alias = value.alias_source.value >= 0; + if (!value_alias) { + return !tensor_alias || value.kind == ValueKind::External; + } + if (!tensor_alias) { + return true; + } + + const Value * source_value = values.find(value.alias_source); + if (source_value == nullptr || value.alias_source.value < 0 || + static_cast(value.alias_source.value) >= tensor_by_value.size()) { + return false; + } + const ggml_tensor * source_tensor = tensor_by_value[static_cast(value.alias_source.value)]; + if (source_tensor == nullptr) { + return tensor_alias && value.storage_offset == tensor->view_offs; + } + size_t relative_offset = 0; + if (!tensor_storage_relative_offset(source_tensor, tensor, relative_offset) || + value.storage_offset < source_value->storage_offset) { + return false; + } + return relative_offset == value.storage_offset - source_value->storage_offset; +} + +static bool tensor_metadata_matches(const ValueMap & values, + const Value & value, + const ggml_tensor * tensor, + const std::vector & tensor_by_value) { + if (tensor == nullptr || value.type != tensor->type || value.element_count != ggml_nelements(tensor) || + value.byte_count != ggml_nbytes(tensor) || value.contiguous != ggml_is_contiguous(tensor)) { + return false; + } + if (!tensor_alias_matches(values, value, tensor, tensor_by_value)) { + return false; + } + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (value.ne[i] != tensor->ne[i] || value.nb[i] != tensor->nb[i]) { + return false; + } + } + return true; +} + +static std::string format_tensor_metadata(const ggml_tensor * tensor) { + if (tensor == nullptr) { + return "null"; + } + std::ostringstream out; + out << ggml_type_name(tensor->type) << " ne=[" << tensor->ne[0] << ',' << tensor->ne[1] << ',' << tensor->ne[2] + << ',' << tensor->ne[3] << "] nb=[" << tensor->nb[0] << ',' << tensor->nb[1] << ',' << tensor->nb[2] << ',' + << tensor->nb[3] << "] elements=" << ggml_nelements(tensor) << " bytes=" << ggml_nbytes(tensor) + << " contiguous=" << (ggml_is_contiguous(tensor) ? 1 : 0); + if (tensor->view_src != nullptr) { + out << " view_offs=" << tensor->view_offs; + } + return out.str(); +} + +static std::string format_value_metadata(const Value & value) { + std::ostringstream out; + out << (value.kind == ValueKind::External ? "external " : "transient ") << ggml_type_name(value.type) << " ne=[" + << value.ne[0] << ',' << value.ne[1] << ',' << value.ne[2] << ',' << value.ne[3] << "] nb=[" << value.nb[0] + << ',' << value.nb[1] << ',' << value.nb[2] << ',' << value.nb[3] << "] elements=" << value.element_count + << " bytes=" << value.byte_count << " contiguous=" << (value.contiguous ? 1 : 0); + if (value.alias_source.value >= 0) { + out << " storage_offset=" << value.storage_offset; + } + return out.str(); +} + +static bool graph_node_params_match(const GraphNode & cached_node, const ggml_tensor * current_node) { + if (current_node == nullptr) { + return false; + } + return op_params_equivalent(cached_node.op, cached_node.params, *current_node); +} + +static bool environment_flag_enabled(const char * name) { + const char * value = std::getenv(name); + return value != nullptr && value[0] != '\0' && value[0] != '0'; +} + +static void apply_graph_replay_result(PreparedCommandProgramCacheExecutionResult & result, + const RecordedCommandGraphExecutionResult & replay) { + result.graph_replay_event = replay.event; + result.graph_replay_ineligible_reason = replay.ineligible_reason; + result.graph_replay_build_ns = replay.build_ns; + result.graph_replay_launch_ns = replay.launch_ns; + result.graph_replay_total_ns = replay.total_ns(); + result.graph_replay_dispatches = replay.dispatch_count; + result.graph_replay_transient_allocation_changed = replay.transient_allocation_changed; +} + +static bool graph_replay_should_fallback(HrxGraphReplayEvent event) { + return event == HrxGraphReplayEvent::Ineligible || event == HrxGraphReplayEvent::BuildFailed; +} + +static void collect_command_graph_values(const std::vector & commands, std::vector & values) { + for (const Command & command : commands) { + for (const CommandBinding & binding : command.bindings) { + if (binding.origin == CommandBindingOrigin::GraphValue && binding.value.value >= 0 && + static_cast(binding.value.value) < values.size()) { + values[static_cast(binding.value.value)] = 1; + } + } + } +} + +static std::vector collect_command_graph_values(const CommandProgram & commands, size_t value_count) { + std::vector values(value_count, 0); + collect_command_graph_values(commands.initialization_commands, values); + collect_command_graph_values(commands.commands, values); + return values; +} + +static bool can_skip_external_binding(const Value & value, + const CommandProgram & commands, + size_t value_count, + std::optional> & command_graph_values) { + if (value.element_count != 0 && value.byte_count != 0) { + return false; + } + if (!command_graph_values.has_value()) { + command_graph_values = collect_command_graph_values(commands, value_count); + } + return (*command_graph_values)[static_cast(value.id.value)] == 0; +} + +class TensorValueIndex { + public: + explicit TensorValueIndex(size_t maximum_size) : maximum_size_(maximum_size) {} + ~TensorValueIndex() { ggml_hash_set_free(&tensors_); } + + TensorValueIndex(const TensorValueIndex &) = delete; + TensorValueIndex & operator=(const TensorValueIndex &) = delete; + + std::optional find(const ggml_tensor * tensor) const { + if (tensors_.size == 0) { + return std::nullopt; + } + const size_t index = ggml_hash_find(&tensors_, tensor); + if (index == GGML_HASHSET_FULL || !ggml_bitset_get(tensors_.used, index)) { + return std::nullopt; + } + return values_[index]; + } + + void emplace(const ggml_tensor * tensor, int32_t value) { + if (tensors_.size == 0) { + GGML_ASSERT(maximum_size_ <= SIZE_MAX / 2); + tensors_ = ggml_hash_set_new(2 * maximum_size_); + values_.resize(tensors_.size); + } + const size_t index = ggml_hash_insert(&tensors_, const_cast(tensor)); + GGML_ASSERT(index != GGML_HASHSET_ALREADY_EXISTS); + values_[index] = value; + } + + private: + size_t maximum_size_; + ggml_hash_set tensors_ = {}; + std::vector values_; +}; + +static Status bind_current_value(const ValueMap & values, + ValueId expected, + const ggml_tensor * tensor, + std::vector & tensor_by_value, + TensorValueIndex & value_by_tensor, + const char * role, + size_t node_index) { + Status status; + const Value * value = values.find(expected); + if (value == nullptr || expected.value < 0 || static_cast(expected.value) >= tensor_by_value.size()) { + status.log("node %zu %s references missing cached value %d", node_index, role, expected.value); + return status; + } + if (tensor == nullptr) { + status.log("node %zu %s value %d maps to a null tensor", node_index, role, expected.value); + return status; + } + + const ggml_tensor * existing_tensor = tensor_by_value[static_cast(expected.value)]; + if (existing_tensor != nullptr && existing_tensor != tensor) { + status.log("node %zu %s value %d maps to multiple current tensors", node_index, role, expected.value); + return status; + } + + if (existing_tensor == tensor) { + // Both maps are populated together when a value is first seen. + return status; + } + + const auto existing_value = value_by_tensor.find(tensor); + if (existing_value.has_value() && *existing_value != expected.value) { + status.log("node %zu %s tensor maps to cached values %d and %d", node_index, role, *existing_value, + expected.value); + return status; + } + + if (existing_tensor == nullptr && !existing_value.has_value() && + !tensor_metadata_matches(values, *value, tensor, tensor_by_value)) { + status.log("node %zu %s value %d metadata does not match current tensor: cached %s current %s", node_index, + role, expected.value, format_value_metadata(*value).c_str(), format_tensor_metadata(tensor).c_str()); + return status; + } + + tensor_by_value[static_cast(expected.value)] = tensor; + if (!existing_value.has_value()) { + value_by_tensor.emplace(tensor, expected.value); + } + return status; +} + +static std::string command_program_shape_key(const CommandProgram & commands) { + std::ostringstream out; + out << "hrx-command-program-v1|commands=" << commands.commands.size(); + for (const Command & command : commands.commands) { + out << "|ordinal=" << command.ordinal << "|kind=" << static_cast(command.kind) + << "|kernel=" << command.kernel.kernel_id; + for (const auto & parameter : command.kernel.integer_parameters) { + out << "|ip:" << parameter.first << '=' << parameter.second; + } + for (const auto & parameter : command.kernel.compile_parameters) { + out << "|cp:" << parameter.first << '=' << parameter.second; + } + out << "|bindings=" << command.bindings.size(); + for (const CommandBinding & binding : command.bindings) { + out << "|b:" << binding.name << ':' << binding.value.value << ':' << static_cast(binding.origin) << ':' + << binding.offset << ':' << binding.length << ':' << static_cast(binding.access) << ':' + << binding.layout << ':' << static_cast(binding.source_type) << ':' << binding.input_size << ':' + << binding.output_size << ':' << binding.source_length; + } + out << "|deps=" << command.dependencies.size(); + for (const uint32_t dependency : command.dependencies) { + out << ':' << dependency; + } + } + out << "|transients=" << commands.transients.allocations.size() << "|arena=" << commands.transients.arena_size + << "|arena_alignment=" << commands.transients.arena_alignment; + for (const TransientAllocation & allocation : commands.transients.allocations) { + out << "|t:" << allocation.value.value << ':' << allocation.arena_offset << ':' << allocation.size << ':' + << allocation.alignment; + } + return out.str(); +} + +} // namespace + +GraphProgram::GraphProgram(uint64_t uid, + std::string target, + std::unique_ptr graph, + std::unique_ptr commands, + std::string command_shape) : + uid_(uid), + target_(std::move(target)), + graph_(std::move(graph)), + commands_(std::move(commands)), + command_shape_(std::move(command_shape)), + command_shape_hash_(command_program_shape_hash(command_shape_)) {} + +const GraphProgramExternalSlot * GraphProgram::find_external_slot(ValueId value) const { + const auto found = external_slot_by_value_.find(value.value); + if (found == external_slot_by_value_.end() || found->second >= external_slots_.size()) { + return nullptr; + } + return &external_slots_[found->second]; +} + +const ggml_tensor * GraphProgram::resolve_external_slot(const ggml_cgraph & graph, + const GraphProgramExternalSlot & slot, + Status & status) const { + if (slot.node_index >= static_cast(graph.n_nodes)) { + status.log("external value %d references node slot %zu but current graph has %d nodes", slot.value.value, + slot.node_index, graph.n_nodes); + return nullptr; + } + const ggml_tensor * node = graph.nodes[slot.node_index]; + if (node == nullptr) { + status.log("external value %d references null node slot %zu", slot.value.value, slot.node_index); + return nullptr; + } + if (slot.kind == GraphProgramExternalSlotKind::Node) { + return node; + } + if (slot.source_index < 0 || slot.source_index >= GGML_MAX_SRC) { + status.log("external value %d references invalid source slot %d", slot.value.value, slot.source_index); + return nullptr; + } + const ggml_tensor * source = node->src[slot.source_index]; + if (source == nullptr) { + status.log("external value %d references null source slot %zu:%d", slot.value.value, slot.node_index, + slot.source_index); + return nullptr; + } + return source; +} + +GraphProgramMatch GraphProgram::match_trusted_graph(const ggml_cgraph & current_graph, bool bind_external) const { + GraphProgramMatch result; + if (graph_ == nullptr || commands_ == nullptr) { + result.status.log("missing cached HRX graph program"); + return result; + } + if (graph_->nodes().size() != static_cast(current_graph.n_nodes)) { + result.status.log("cached graph has %zu nodes but current graph has %d", graph_->nodes().size(), + current_graph.n_nodes); + return result; + } + if (!graph_->nodes().empty()) { + const ggml_tensor * first = current_graph.nodes[0]; + const ggml_tensor * last = current_graph.nodes[current_graph.n_nodes - 1]; + if (first == nullptr || last == nullptr) { + result.status.log("current graph has null sentinel nodes"); + return result; + } + if (first->op != graph_->nodes().front().op || last->op != graph_->nodes().back().op) { + result.status.log("current graph sentinel ops do not match cached HRX graph"); + return result; + } + } + if (!bind_external) { + return result; + } + result.external_bindings.reserve(external_slots_.size()); + for (const GraphProgramExternalSlot & slot : external_slots_) { + const ggml_tensor * tensor = resolve_external_slot(current_graph, slot, result.status); + if (tensor == nullptr) { + return result; + } + result.external_bindings.push_back({ slot.value, tensor }); + } + return result; +} + +GraphProgramMatch GraphProgram::match_host_staging_graph(const ggml_cgraph & current_graph) const { + GraphProgramMatch result; + std::lock_guard lock(prepared_mutex_); + if (!has_prepared_) { + result.status.log("missing prepared HRX command program"); + return result; + } + result.external_bindings.reserve(prepared_.host_staging.size()); + for (const HostStagingBuffer & staging : prepared_.host_staging) { + const GraphProgramExternalSlot * slot = find_external_slot(ValueId(staging.value)); + if (slot == nullptr) { + result.status.log("prepared host staging value %d has no external graph slot", staging.value); + return result; + } + const ggml_tensor * tensor = resolve_external_slot(current_graph, *slot, result.status); + if (tensor == nullptr) { + return result; + } + result.external_bindings.push_back({ ValueId(staging.value), tensor }); + } + return result; +} + +Status GraphProgram::capture_external_slots(const ggml_cgraph & graph, const GraphProgramMatch & match) { + Status status; + external_slots_.clear(); + external_slot_by_value_.clear(); + fast_path_nodes_ = graph.nodes; + external_slots_.reserve(match.external_bindings.size()); + for (const GraphProgramExternalBinding & binding : match.external_bindings) { + GraphProgramExternalSlot slot; + slot.value = binding.value; + bool found = false; + for (int i = 0; i < graph.n_nodes && !found; ++i) { + const ggml_tensor * node = graph.nodes[i]; + if (node == nullptr) { + continue; + } + if (node == binding.tensor) { + slot.kind = GraphProgramExternalSlotKind::Node; + slot.node_index = static_cast(i); + found = true; + break; + } + for (int j = 0; j < GGML_MAX_SRC; ++j) { + if (node->src[j] == binding.tensor) { + slot.kind = GraphProgramExternalSlotKind::Source; + slot.node_index = static_cast(i); + slot.source_index = j; + found = true; + break; + } + } + } + if (!found) { + status.log("external value %d has no current graph slot", binding.value.value); + continue; + } + external_slot_by_value_[binding.value.value] = external_slots_.size(); + external_slots_.push_back(slot); + } + return status; +} + +bool GraphProgram::has_prepared_program() const { + std::lock_guard lock(prepared_mutex_); + return has_prepared_; +} + +bool GraphProgram::can_use_prepared_fast_path(const ggml_cgraph & graph) const { + // The UID prevents graph-arena address reuse from passing the fast path. + return graph.uid == uid_ && graph.nodes == fast_path_nodes_; +} + +PreparedCommandProgramCacheStats GraphProgram::prepared_stats() const { + std::lock_guard lock(prepared_mutex_); + return prepared_stats_; +} + +PreparedCommandProgramCacheExecutionResult GraphProgram::execute_with_result( + const CommandProgramExecutionContext & context, + const CommandProgramBindings & bindings) { + PreparedCommandProgramCacheExecutionResult result; + if (commands_ == nullptr || !commands_->valid() || !bindings.valid()) { + result.graph_replay_event = HrxGraphReplayEvent::Ineligible; + result.graph_replay_ineligible_reason = "invalid_graph_program"; + if (commands_ == nullptr) { + result.status.log("missing cached HRX command program"); + } + result.status.append(bindings.status); + return result; + } + + std::lock_guard lock(prepared_mutex_); + if (!has_prepared_) { + prepared_ = prepare_command_program(context, *commands_, bindings); + if (!prepared_.valid()) { + result.status.append(prepared_.status); + return result; + } + has_prepared_ = true; + ++prepared_stats_.builds; + } else { + ++prepared_stats_.hits; + } + + if (environment_flag_enabled("GGML_HRX_DISABLE_GRAPH_REPLAY") || debug_serial_command_execution_enabled()) { + result.graph_replay_event = HrxGraphReplayEvent::Disabled; + result.graph_replay_ineligible_reason = debug_serial_command_execution_enabled() ? + "debug_serial_execution" : + "disabled_by_environment"; + result.success = bind_and_execute_prepared_command_program(context, *commands_, bindings, prepared_); + if (!result.success) { + result.status.log("execute cached HRX command program failed"); + } + return result; + } + + if (environment_flag_enabled("GGML_HRX_REBUILD_GRAPH_REPLAY")) { + recorded_ = {}; + } + + const RecordedCommandGraphExecutionResult replay = + bind_and_launch_recorded_command_graph(context, *commands_, bindings, prepared_, recorded_); + apply_graph_replay_result(result, replay); + if (replay.success) { + result.success = true; + return result; + } + if (!graph_replay_should_fallback(replay.event)) { + result.status.append(replay.status); + if (result.status.success()) { + result.status.log("execute cached HRX graph replay failed"); + } + return result; + } + result.success = bind_and_execute_prepared_command_program(context, *commands_, bindings, prepared_); + if (!result.success) { + result.status.log("execute cached HRX command program failed"); + } + return result; +} + +GraphProgramMatch GraphProgram::match_current_graph(const ggml_cgraph & current_graph) const { + GraphProgramMatch result; + if (graph_ == nullptr) { + result.status.log("missing cached HRX graph"); + return result; + } + if (commands_ == nullptr) { + result.status.log("missing cached HRX command program"); + return result; + } + if (graph_->nodes().size() != static_cast(current_graph.n_nodes)) { + result.status.log("cached graph has %zu nodes but current graph has %d", graph_->nodes().size(), + current_graph.n_nodes); + return result; + } + + const ValueMap & values = graph_->values(); + std::vector tensor_by_value(values.size(), nullptr); + TensorValueIndex value_by_tensor(values.size()); + + for (size_t node_index = 0; node_index < graph_->nodes().size(); ++node_index) { + const GraphNode & cached_node = graph_->nodes()[node_index]; + const ggml_tensor * current_node = current_graph.nodes[node_index]; + if (current_node == nullptr) { + result.status.log("current graph node %zu is null", node_index); + return result; + } + if (cached_node.op != current_node->op) { + result.status.log("node %zu cached op %s does not match current op %s", node_index, + ggml_op_name(cached_node.op), ggml_op_name(current_node->op)); + return result; + } + if (!graph_node_params_match(cached_node, current_node)) { + result.status.log("node %zu cached op params do not match current graph", node_index); + return result; + } + + size_t input_index = 0; + for (const ggml_tensor * source : current_node->src) { + if (source == nullptr) { + continue; + } + if (input_index >= cached_node.inputs.size()) { + result.status.log("node %zu has more inputs than the cached graph", node_index); + return result; + } + Status status = bind_current_value(values, cached_node.inputs[input_index], source, tensor_by_value, + value_by_tensor, "input", node_index); + if (!status.success()) { + result.status.append(status); + return result; + } + ++input_index; + } + if (input_index != cached_node.inputs.size()) { + result.status.log("node %zu has %zu inputs but cached graph has %zu", node_index, input_index, + cached_node.inputs.size()); + return result; + } + Status status = bind_current_value(values, cached_node.output, current_node, tensor_by_value, value_by_tensor, + "output", node_index); + if (!status.success()) { + result.status.append(status); + return result; + } + } + + std::optional> command_graph_values; + const std::vector external_ids = values.external_value_ids(); + result.external_bindings.reserve(external_ids.size()); + for (const ValueId id : external_ids) { + const Value * value = values.find(id); + if (value == nullptr) { + result.status.log("external value %d is missing from the cached graph", id.value); + return result; + } + if (id.value < 0 || static_cast(id.value) >= tensor_by_value.size() || + tensor_by_value[static_cast(id.value)] == nullptr) { + result.status.log("external value %d is missing from the current graph", id.value); + return result; + } + if (can_skip_external_binding(*value, *commands_, values.size(), command_graph_values)) { + continue; + } + result.external_bindings.push_back({ id, tensor_by_value[static_cast(id.value)] }); + } + return result; +} + +bool GraphProgramCache::can_execute(const ggml_cgraph & graph, + const KernelCorpus & corpus, + const std::string & target) const { + return check_support(graph, corpus, target).supported; +} + +GraphProgramSupportResult GraphProgramCache::check_support(const ggml_cgraph & graph, + const KernelCorpus & corpus, + const std::string & target) const { + GraphProgramSupportResult result; + if (graph.n_nodes == 0) { + result.supported = true; + return result; + } + GraphImportResult imported = import_ggml_graph(graph); + if (!imported.valid()) { + result.status.append(imported.status); + return result; + } + std::unique_ptr program = + build_program_from_imported(graph.uid, std::move(imported.graph), corpus, target, result.status); + result.supported = program != nullptr && result.status.success(); + return result; +} + +GraphProgramLookup GraphProgramCache::build_from_imported(const ggml_cgraph & graph, + Graph && imported_graph, + const KernelCorpus & corpus, + const std::string & target) { + GraphProgramLookup result; + std::unique_ptr program = + build_program_from_imported(graph.uid, std::move(imported_graph), corpus, target, result.status); + if (program == nullptr) { + return result; + } + + GraphProgramMatch match = program->match_current_graph(graph); + if (!match.valid()) { + result.status.append(match.status); + return result; + } + Status slot_status = program->capture_external_slots(graph, match); + if (!slot_status.success()) { + result.status.append(slot_status); + return result; + } + + if (graph.uid == 0) { + result.uncached_program = std::move(program); + result.program = result.uncached_program.get(); + result.match = std::move(match); + return result; + } + + GraphProgram * cached_program = program.get(); + { + std::lock_guard lock(mutex_); + validated_matches_.clear(); + programs_[graph.uid] = std::move(program); + cached_program = programs_[graph.uid].get(); + last_program_ = cached_program; + ++stats_.builds; + } + result.program = cached_program; + result.match = std::move(match); + return result; +} + +GraphProgramLookup GraphProgramCache::get_or_build(const ggml_cgraph & graph, + const KernelCorpus & corpus, + const std::string & target) { + GraphProgramLookup result; + if (graph.uid != 0) { + const bool disable_fast_path = environment_flag_enabled("GGML_HRX_DISABLE_GRAPH_UID_FAST_PATH"); + const bool validate_fast_path = environment_flag_enabled("GGML_HRX_VALIDATE_GRAPH_UID_CACHE"); + const bool allow_structural_reuse = std::getenv("GGML_HRX_DUMP_COMMAND_PROGRAM_DIR") == nullptr; + GraphProgram * cached_program = nullptr; + { + std::lock_guard lock(mutex_); + if (!disable_fast_path && last_program_ != nullptr && last_program_->uid() == graph.uid && + last_program_->target() == target) { + cached_program = last_program_; + } else { + const auto found = programs_.find(graph.uid); + if (found != programs_.end() && found->second->target() == target) { + cached_program = found->second.get(); + last_program_ = cached_program; + } + if (cached_program == nullptr && !disable_fast_path && allow_structural_reuse) { + const auto matched = validated_matches_.find(graph.uid); + if (matched != validated_matches_.end() && matched->second.nodes == graph.nodes && + matched->second.program->target() == target) { + cached_program = matched->second.program; + } + } + } + } + if (cached_program != nullptr) { + const bool bind_external = + !cached_program->has_prepared_program() || !cached_program->can_use_prepared_fast_path(graph); + GraphProgramMatch match = disable_fast_path || validate_fast_path ? + cached_program->match_current_graph(graph) : + cached_program->match_trusted_graph(graph, bind_external); + if (match.valid() && validate_fast_path && !disable_fast_path) { + GraphProgramMatch trusted_match = cached_program->match_trusted_graph(graph, bind_external); + if (!trusted_match.valid()) { + result.status.append(trusted_match.status); + return result; + } + } + if (match.valid()) { + result.program = cached_program; + result.match = std::move(match); + { + std::lock_guard lock(mutex_); + ++stats_.hits; + } + return result; + } + if (validate_fast_path) { + result.status.append(match.status); + return result; + } + std::lock_guard lock(mutex_); + validated_matches_.erase(graph.uid); + } + + // Physical bindings have a separate prepared-program cache key. + if (allow_structural_reuse) { + std::lock_guard lock(mutex_); + for (const auto & entry : programs_) { + GraphProgram * candidate = entry.second.get(); + if (candidate == cached_program || candidate->target() != target) { + continue; + } + GraphProgramMatch match = candidate->match_current_graph(graph); + if (!match.valid()) { + continue; + } + last_program_ = candidate; + if (!disable_fast_path) { + if (validated_matches_.size() >= 128) { + validated_matches_.clear(); + } + validated_matches_[graph.uid] = { candidate, graph.nodes }; + } + ++stats_.hits; + result.program = candidate; + result.match = std::move(match); + return result; + } + } + } + + GraphImportResult imported = import_ggml_graph(graph); + if (!imported.valid()) { + result.status.append(imported.status); + return result; + } + return build_from_imported(graph, std::move(imported.graph), corpus, target); +} + +GraphProgramCacheStats GraphProgramCache::stats() const { + std::lock_guard lock(mutex_); + GraphProgramCacheStats stats = stats_; + for (const auto & entry : programs_) { + const PreparedCommandProgramCacheStats prepared = entry.second->prepared_stats(); + stats.prepared_program_builds += prepared.builds; + stats.prepared_program_hits += prepared.hits; + } + return stats; +} + +size_t GraphProgramCache::size() const { + std::lock_guard lock(mutex_); + return programs_.size(); +} + +void GraphProgramCache::retain_only(const GraphProgram * keep) { + std::lock_guard lock(mutex_); + for (auto it = programs_.begin(); it != programs_.end();) { + it = it->second.get() == keep ? std::next(it) : programs_.erase(it); + } + validated_matches_.clear(); + last_program_ = programs_.empty() ? nullptr : programs_.begin()->second.get(); +} + +void GraphProgramCache::clear() { + std::lock_guard lock(mutex_); + validated_matches_.clear(); + programs_.clear(); + last_program_ = nullptr; +} + +std::unique_ptr GraphProgramCache::build_program_from_imported(uint64_t uid, + Graph && imported_graph, + const KernelCorpus & corpus, + const std::string & target, + Status & errors) const { + DispatchScheduler scheduler; + if (!scheduler.schedule_graph(imported_graph, { target })) { + errors.append(scheduler.plan().status); + return nullptr; + } + + CommandProgram commands = build_command_program(imported_graph, scheduler.plan(), corpus, target); + if (!commands.valid()) { + errors.append(commands.status); + return nullptr; + } + + std::string command_shape = command_program_shape_key(commands); + Status dump_status = dump_command_program_kernels_if_requested(commands, corpus, target, command_shape); + for (const std::string & error : dump_status.errors()) { + GGML_LOG_WARN("ggml_hrx: %s\n", error.c_str()); + } + return std::make_unique(uid, target, std::make_unique(std::move(imported_graph)), + std::make_unique(std::move(commands)), + std::move(command_shape)); +} + +bool can_execute_standalone_op_as_graph(const ggml_tensor * op, const std::string & target) { + if (op == nullptr) { + return false; + } + Graph graph; + std::vector inputs; + for (const ggml_tensor * source : op->src) { + if (source == nullptr) { + continue; + } + inputs.push_back(graph.values().get_or_add_tensor_value(source, ValueKind::External)); + } + const ValueId output = graph.values().get_or_add_tensor_value(op, ValueKind::External); + GraphNode & node = graph.add_node(op->op, output, std::move(inputs)); + node.params = import_op_params(*op); + if (!graph.build_index().success()) { + return false; + } + return DispatchScheduler::supports_node(graph, &node, { target }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/graph-program-cache.h b/ggml/src/ggml-hrx/runtime/graph-program-cache.h new file mode 100644 index 000000000000..7e570040f58a --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/graph-program-cache.h @@ -0,0 +1,173 @@ +#pragma once + +#include "dispatch/command-program.h" +#include "graph/graph.h" +#include "kernel-corpus/kernel-corpus.h" +#include "runtime/prepared-command-program-cache.h" +#include "status.h" + +#include +#include +#include +#include +#include +#include + +struct ggml_cgraph; +struct ggml_tensor; + +namespace ggml::hrx { + +struct GraphProgramExternalBinding { + ValueId value; + const ggml_tensor * tensor = nullptr; +}; + +enum class GraphProgramExternalSlotKind { + Node, + Source, +}; + +struct GraphProgramExternalSlot { + ValueId value; + GraphProgramExternalSlotKind kind = GraphProgramExternalSlotKind::Node; + size_t node_index = 0; + int32_t source_index = -1; +}; + +struct GraphProgramMatch { + std::vector external_bindings; + Status status; + + bool valid() const { return status.success(); } +}; + +class GraphProgram { + public: + GraphProgram(uint64_t uid, + std::string target, + std::unique_ptr graph, + std::unique_ptr commands, + std::string command_shape); + + uint64_t uid() const { return uid_; } + + const std::string & target() const { return target_; } + + const std::string & command_shape() const { return command_shape_; } + + uint64_t command_shape_hash() const { return command_shape_hash_; } + + const Graph & graph() const { return *graph_; } + + Graph & graph() { return *graph_; } + + const CommandProgram & commands() const { return *commands_; } + + CommandProgram & commands() { return *commands_; } + + GraphProgramMatch match_current_graph(const ggml_cgraph & graph) const; + GraphProgramMatch match_trusted_graph(const ggml_cgraph & graph, bool bind_external = true) const; + GraphProgramMatch match_host_staging_graph(const ggml_cgraph & graph) const; + + Status capture_external_slots(const ggml_cgraph & graph, const GraphProgramMatch & match); + + bool has_prepared_program() const; + bool can_use_prepared_fast_path(const ggml_cgraph & graph) const; + + PreparedCommandProgramCacheStats prepared_stats() const; + + PreparedCommandProgramCacheExecutionResult execute_with_result(const CommandProgramExecutionContext & context, + const CommandProgramBindings & bindings); + + private: + const GraphProgramExternalSlot * find_external_slot(ValueId value) const; + const ggml_tensor * resolve_external_slot(const ggml_cgraph & graph, + const GraphProgramExternalSlot & slot, + Status & status) const; + + uint64_t uid_ = 0; + std::string target_; + std::unique_ptr graph_; + std::unique_ptr commands_; + std::string command_shape_; + uint64_t command_shape_hash_ = 0; + + std::vector external_slots_; + std::unordered_map external_slot_by_value_; + const ggml_tensor * const * fast_path_nodes_ = nullptr; + mutable std::mutex prepared_mutex_; + PreparedCommandProgram prepared_; + RecordedCommandGraph recorded_; + PreparedCommandProgramCacheStats prepared_stats_; + bool has_prepared_ = false; +}; + +struct GraphProgramCacheStats { + uint64_t builds = 0; + uint64_t hits = 0; + uint64_t prepared_program_builds = 0; + uint64_t prepared_program_hits = 0; +}; + +struct GraphProgramLookup { + GraphProgram * program = nullptr; + std::unique_ptr uncached_program; + GraphProgramMatch match; + Status status; + + bool valid() const { return program != nullptr && status.success() && match.valid(); } +}; + +struct GraphProgramSupportResult { + bool supported = false; + Status status; + + bool valid() const { return supported && status.success(); } +}; + +class GraphProgramCache { + public: + bool can_execute(const ggml_cgraph & graph, const KernelCorpus & corpus, const std::string & target) const; + + GraphProgramSupportResult check_support(const ggml_cgraph & graph, + const KernelCorpus & corpus, + const std::string & target) const; + + GraphProgramLookup get_or_build(const ggml_cgraph & graph, const KernelCorpus & corpus, const std::string & target); + + GraphProgramCacheStats stats() const; + + size_t size() const; + // drop every cached program except `keep` (graph-program-cache-limit.h) + void retain_only(const GraphProgram * keep); + + void clear(); + + private: + struct ValidatedGraphMatch { + GraphProgram * program = nullptr; + const ggml_tensor * const * nodes = nullptr; + }; + + GraphProgramLookup build_from_imported(const ggml_cgraph & graph, + Graph && imported_graph, + const KernelCorpus & corpus, + const std::string & target); + + std::unique_ptr build_program_from_imported(uint64_t uid, + Graph && imported_graph, + const KernelCorpus & corpus, + const std::string & target, + Status & errors) const; + + mutable std::mutex mutex_; + std::unordered_map> programs_; + std::unordered_map validated_matches_; + GraphProgram * last_program_ = nullptr; + GraphProgramCacheStats stats_; +}; + +bool can_execute_standalone_op_as_graph(const ggml_tensor * op, const std::string & target); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/graph-record-order.cpp b/ggml/src/ggml-hrx/runtime/graph-record-order.cpp new file mode 100644 index 000000000000..774a1fc017d0 --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/graph-record-order.cpp @@ -0,0 +1,299 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "runtime/graph-record-order.h" + +#include "ggml-impl.h" + +#include +#include +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +constexpr size_t kMaxTransientExtent = size_t{ 16 } << 20; + +bool level_order_enabled() { + static const bool enabled = [] { + const char * value = std::getenv("GGML_HRX_GRAPH_ORDER"); + return value == nullptr || std::strcmp(value, "program") != 0; + }(); + return enabled; +} + +bool log_enabled() { + static const bool enabled = [] { + const char * value = std::getenv("GGML_HRX_LOG_GRAPH_ORDER"); + return value != nullptr && value[0] != '\0' && std::strcmp(value, "0") != 0; + }(); + return enabled; +} + +bool access_writes(ResourceAccess access) { + return access == ResourceAccess::Write || access == ResourceAccess::ReadWrite; +} + +// Memory-range hazards between commands, in program order: a command depends on the last writer of +// every byte it touches and, when it writes, on every reader since that writer. +class RangeHazards { + public: + void collect(hrx_buffer_t buffer, size_t offset, size_t length, bool writes, std::vector & deps) const { + if (buffer == nullptr || length == 0) { + return; + } + const auto found = buffers_.find(buffer); + if (found == buffers_.end()) { + return; + } + const size_t end = range_end(offset, length); + auto range = found->second.upper_bound(offset); + if (range != found->second.begin()) { + --range; + if (range->second.end <= offset) { + ++range; + } + } + for (; range != found->second.end() && range->first < end; ++range) { + if (range->second.last_writer >= 0) { + deps.push_back(static_cast(range->second.last_writer)); + } + if (writes) { + deps.insert(deps.end(), range->second.readers.begin(), range->second.readers.end()); + } + } + } + + void update(uint32_t command, hrx_buffer_t buffer, size_t offset, size_t length, bool writes) { + if (buffer == nullptr || length == 0) { + return; + } + const size_t end = range_end(offset, length); + Ranges & ranges = buffers_[buffer]; + split(ranges, offset); + split(ranges, end); + size_t cursor = offset; + auto range = ranges.lower_bound(offset); + while (cursor < end) { + if (range == ranges.end() || range->first > cursor) { + Segment segment; + segment.end = range == ranges.end() ? end : std::min(end, range->first); + if (writes) { + segment.last_writer = static_cast(command); + } else { + segment.readers.push_back(command); + } + const size_t next = segment.end; + range = std::next(ranges.emplace(cursor, std::move(segment)).first); + cursor = next; + continue; + } + if (writes) { + range->second.last_writer = static_cast(command); + range->second.readers.clear(); + } else if (range->second.readers.empty() || range->second.readers.back() != command) { + range->second.readers.push_back(command); + } + cursor = range->second.end; + ++range; + } + } + + private: + struct Segment { + size_t end = 0; + int32_t last_writer = -1; + std::vector readers; + }; + using Ranges = std::map; + + static size_t range_end(size_t offset, size_t length) { + return length > std::numeric_limits::max() - offset ? std::numeric_limits::max() : + offset + length; + } + + static void split(Ranges & ranges, size_t offset) { + auto upper = ranges.upper_bound(offset); + if (upper == ranges.begin()) { + return; + } + auto current = std::prev(upper); + if (offset <= current->first || offset >= current->second.end) { + return; + } + Segment right = current->second; + current->second.end = offset; + ranges.emplace(offset, std::move(right)); + } + + std::unordered_map buffers_; +}; + +// Barriers HRX records for |order|: one before a command with a dependency recorded since the last barrier. +size_t count_barriers(const std::vector> & deps, const std::vector & order) { + std::vector window_epoch(deps.size(), 0); + uint32_t epoch = 1; + size_t barriers = 0; + for (uint32_t command : order) { + for (uint32_t dep : deps[command]) { + if (window_epoch[dep] == epoch) { + ++barriers; + ++epoch; + break; + } + } + window_epoch[command] = epoch; + } + return barriers; +} + +// Dependency-level (as-soon-as-possible) order: barriers = longest dependency chain. Logged for reference. +std::vector asap_order(const std::vector> & deps) { + std::vector level(deps.size(), 0), order(deps.size()); + for (uint32_t c = 0; c < deps.size(); ++c) { + order[c] = c; + for (uint32_t dep : deps[c]) { + level[c] = std::max(level[c], level[dep] + 1); + } + } + std::stable_sort(order.begin(), order.end(), [&](uint32_t lhs, uint32_t rhs) { return level[lhs] < level[rhs]; }); + return order; +} + +// Program order, except that when the next command needs a barrier, a command from the next +// kLookahead unscheduled ones that is ready (all its dependencies recorded) and needs no barrier +// (none recorded since the last barrier) is recorded first. Every hazard pair keeps its program order +// because a command is only taken once all of its dependencies are recorded. +std::vector lookahead_order(const std::vector> & deps) { + constexpr size_t kLookahead = 8; + const size_t n = deps.size(); + std::vector scheduled(n, 0); + std::vector window_epoch(n, 0); + std::vector order; + order.reserve(n); + uint32_t epoch = 1; + size_t first = 0; + auto ready = [&](uint32_t c) { + for (uint32_t dep : deps[c]) { + if (!scheduled[dep]) { + return false; + } + } + return true; + }; + auto in_window = [&](uint32_t c) { + for (uint32_t dep : deps[c]) { + if (window_epoch[dep] == epoch) { + return true; + } + } + return false; + }; + while (order.size() < n) { + while (scheduled[first]) { + ++first; + } + uint32_t pick = static_cast(first); + bool found = false; + size_t examined = 0; + for (size_t c = first; c < n && examined < kLookahead; ++c) { + if (scheduled[c]) { + continue; + } + ++examined; + if (ready(static_cast(c)) && !in_window(static_cast(c))) { + pick = static_cast(c); + found = true; + break; + } + } + if (!found && in_window(pick)) { + ++epoch; // barrier + } + scheduled[pick] = 1; + window_epoch[pick] = epoch; + order.push_back(pick); + } + return order; +} + +} // namespace + +std::vector order_prepared_commands_for_graph(const std::vector & commands) { + const bool enabled = level_order_enabled(); + if ((!enabled && !log_enabled()) || commands.size() < 2) { + return {}; + } + + std::vector> deps(commands.size()); + RangeHazards hazards; + for (uint32_t c = 0; c < commands.size(); ++c) { + for (const PreparedCommandBinding & binding : commands[c].kernel.bindings) { + hazards.collect(binding.ref.buffer, binding.ref.offset, binding.ref.length, + access_writes(binding.binding.access), deps[c]); + } + std::sort(deps[c].begin(), deps[c].end()); + deps[c].erase(std::unique(deps[c].begin(), deps[c].end()), deps[c].end()); + deps[c].erase(std::remove(deps[c].begin(), deps[c].end(), c), deps[c].end()); + for (const PreparedCommandBinding & binding : commands[c].kernel.bindings) { + hazards.update(c, binding.ref.buffer, binding.ref.offset, binding.ref.length, + access_writes(binding.binding.access)); + } + } + + std::vector program_order(commands.size()); + for (uint32_t c = 0; c < commands.size(); ++c) { + program_order[c] = c; + } + const std::vector level_order = lookahead_order(deps); + // Reordering changes which kernels overlap. Where it saves no barrier it is not used: Qwen3-0.6B's + // 512-token prompt program has 254 barriers in either order and its pp512 read 1.1% lower reordered. + const size_t program_barriers = count_barriers(deps, program_order); + const size_t level_barriers = count_barriers(deps, level_order); + const uint32_t levels = commands.empty() ? 0u : static_cast(count_barriers(deps, asap_order(deps)) + 1); + // Prompt programs keep program order too: their barriers are a small share of the long kernels. + // They are told apart by their transient arena extent: decode programs measured 0.2-6.5 MB, 512-token + // prompt programs 12-335 MB. + size_t transient_extent = 0; + for (const PreparedCommand & command : commands) { + for (const PreparedCommandBinding & binding : command.kernel.bindings) { + if (binding.binding.origin == CommandBindingOrigin::Transient) { + transient_extent = std::max(transient_extent, binding.ref.offset + binding.ref.length); + } + } + } + const bool use_level = enabled && level_barriers < program_barriers && transient_extent <= kMaxTransientExtent; + if (log_enabled()) { + GGML_LOG_WARN("ggml-hrx graph order: %zu commands, %u levels, barriers %zu in program order, %zu in lookahead order, transient extent %zu (%s used)\n", + commands.size(), levels, + program_barriers, level_barriers, transient_extent, use_level ? "lookahead order" : "program order"); + } + if (!use_level) { + return {}; + } + std::vector ordered; + ordered.reserve(commands.size()); + for (uint32_t c : level_order) { + ordered.push_back(commands[c]); + } + return ordered; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/graph-record-order.h b/ggml/src/ggml-hrx/runtime/graph-record-order.h new file mode 100644 index 000000000000..fc3ae7b6fa60 --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/graph-record-order.h @@ -0,0 +1,45 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Order in which a command program's kernels are recorded into an HRX graph. +// +// HRX puts an execution barrier before a recorded kernel when one of its dependencies was recorded after +// the last barrier. Recording in program order therefore costs a barrier at almost every kernel of a +// dependent chain, even where the program has independent kernels that could share one barrier window +// (the q and kv projections of one input, the router next to a shared expert). Here a kernel that would +// need a barrier first lets a ready, barrier-free kernel from the next eight in program order go ahead. +// The dependencies are the same memory-range hazards the graph recorder uses, and a kernel is only moved +// once all of its dependencies are recorded, so every hazard pair keeps its program order. +// +// Measured on decode programs this reaches the dependency-level minimum (GLM-4.7-Flash 1351 -> 1211 +// barriers, Qwen3.6-35B-A3B 1003 -> 973, gpt-oss-20b 628 -> 556 on its multi-token programs) while +// moving kernels only a few places; full dependency-level order gave the same counts but Qwen3.6-35B +// decode read 0.8% lower. Programs where reordering saves no barrier, and prompt programs (transient +// arena extent over 16 MiB), keep program order. +// GGML_HRX_GRAPH_ORDER=program records in program order. GGML_HRX_LOG_GRAPH_ORDER=1 logs the barrier +// count of both orders for every recorded graph. + +#pragma once + +#include "runtime/command-program-executor.h" + +#include + +namespace ggml::hrx { + +// Returns |commands| in recording order, or an empty vector when program order should be used. +std::vector order_prepared_commands_for_graph(const std::vector & commands); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/graph-replay.h b/ggml/src/ggml-hrx/runtime/graph-replay.h new file mode 100644 index 000000000000..3bcd3abdcbbe --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/graph-replay.h @@ -0,0 +1,44 @@ +#pragma once + +#include +#include + +namespace ggml::hrx { + +enum class HrxGraphReplayEvent { + Disabled, + Ineligible, + MissBuild, + Hit, + RebuildTransient, + BuildFailed, + LaunchFailed, +}; + +inline const char * hrx_graph_replay_event_name(HrxGraphReplayEvent event) { + switch (event) { + case HrxGraphReplayEvent::Disabled: + return "disabled"; + case HrxGraphReplayEvent::Ineligible: + return "ineligible"; + case HrxGraphReplayEvent::MissBuild: + return "miss_build"; + case HrxGraphReplayEvent::Hit: + return "hit"; + case HrxGraphReplayEvent::RebuildTransient: + return "rebuild_transient"; + case HrxGraphReplayEvent::BuildFailed: + return "build_failed"; + case HrxGraphReplayEvent::LaunchFailed: + return "launch_failed"; + } + return "unknown"; +} + +inline uint64_t hrx_graph_replay_now_ns() { + return static_cast( + std::chrono::duration_cast(std::chrono::steady_clock::now().time_since_epoch()) + .count()); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/host-buffer-registry.cpp b/ggml/src/ggml-hrx/runtime/host-buffer-registry.cpp new file mode 100644 index 000000000000..02b591fcfc9c --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/host-buffer-registry.cpp @@ -0,0 +1,71 @@ +#include "host-buffer-registry.h" + +#include "hrx_runtime.h" + +#include +#include + +namespace ggml::hrx { + +HostBufferRef::HostBufferRef(hrx_buffer_t buffer, size_t offset) : buffer_(buffer), offset_(offset) {} + +HostBufferRef::~HostBufferRef() { + if (buffer_ != nullptr) { + hrx_buffer_release(buffer_); + } +} + +HostBufferRef::HostBufferRef(HostBufferRef && other) noexcept : + buffer_(std::exchange(other.buffer_, nullptr)), + offset_(other.offset_) { + other.offset_ = 0; +} + +HostBufferRef & HostBufferRef::operator=(HostBufferRef && other) noexcept { + if (this != &other) { + if (buffer_ != nullptr) { + hrx_buffer_release(buffer_); + } + buffer_ = std::exchange(other.buffer_, nullptr); + offset_ = other.offset_; + other.offset_ = 0; + } + return *this; +} + +void HostBufferRegistry::add(hrx_buffer_t buffer, void * base, size_t size) { + if (buffer == nullptr || base == nullptr || size == 0) { + return; + } + std::lock_guard lock(mutex_); + entries_.push_back({ buffer, reinterpret_cast(base), size }); +} + +void HostBufferRegistry::remove(hrx_buffer_t buffer) { + std::lock_guard lock(mutex_); + entries_.erase(std::remove_if(entries_.begin(), entries_.end(), + [buffer](const Entry & entry) { return entry.buffer == buffer; }), + entries_.end()); +} + +HostBufferRef HostBufferRegistry::find(const void * data, size_t size) const { + if (data == nullptr) { + return {}; + } + const uintptr_t address = reinterpret_cast(data); + std::lock_guard lock(mutex_); + for (const Entry & entry : entries_) { + if (address < entry.base) { + continue; + } + const size_t offset = static_cast(address - entry.base); + if (offset > entry.size || size > entry.size - offset) { + continue; + } + hrx_buffer_retain(entry.buffer); + return HostBufferRef(entry.buffer, offset); + } + return {}; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/host-buffer-registry.h b/ggml/src/ggml-hrx/runtime/host-buffer-registry.h new file mode 100644 index 000000000000..30ce0eba2ec9 --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/host-buffer-registry.h @@ -0,0 +1,56 @@ +#pragma once + +#include +#include +#include +#include + +typedef struct hrx_buffer_s * hrx_buffer_t; + +namespace ggml::hrx { + +class HostBufferRef { + public: + HostBufferRef() = default; + ~HostBufferRef(); + + HostBufferRef(HostBufferRef && other) noexcept; + HostBufferRef & operator=(HostBufferRef && other) noexcept; + + HostBufferRef(const HostBufferRef &) = delete; + HostBufferRef & operator=(const HostBufferRef &) = delete; + + bool valid() const { return buffer_ != nullptr; } + + hrx_buffer_t buffer() const { return buffer_; } + + size_t offset() const { return offset_; } + + private: + hrx_buffer_t buffer_ = nullptr; + size_t offset_ = 0; + + HostBufferRef(hrx_buffer_t buffer, size_t offset); + friend class HostBufferRegistry; +}; + +// Tracks mapped HRX host buffers so pointer-based GGML transfers can remain stream ordered and handle based. +class HostBufferRegistry { + public: + void add(hrx_buffer_t buffer, void * base, size_t size); + void remove(hrx_buffer_t buffer); + + HostBufferRef find(const void * data, size_t size) const; + + private: + struct Entry { + hrx_buffer_t buffer = nullptr; + uintptr_t base = 0; + size_t size = 0; + }; + + mutable std::mutex mutex_; + std::vector entries_; +}; + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/host-memory-ternary.cpp b/ggml/src/ggml-hrx/runtime/host-memory-ternary.cpp new file mode 100644 index 000000000000..f2e906cd0576 --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/host-memory-ternary.cpp @@ -0,0 +1,120 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "runtime/host-memory-ternary.h" + +#include "dispatch/ternary-q4-0.h" +#include "ggml-quants.h" + +#include +#include +#include +#include + +namespace ggml::hrx { + +namespace { + +// One 256-value block: four Q4_0 blocks per 128-value group, two groups. Returns false if a nibble is not +// 7, 8 or 9 or the blocks of a group that hold a nonzero value disagree on the scale. An all-zero block +// (every nibble 8) decodes to zeros whatever its scale, so it does not take part in that check; a group +// with no nonzero value gets scale 0. +bool pack_block(const block_q4_0 * in, uint8_t * out) { + uint16_t scales[2] = { 0, 0 }; + std::memset(out + 4, 0, 64); + for (int group = 0; group < 2; ++group) { + const block_q4_0 * g = in + 4 * group; + bool have_scale = false; + for (int b = 0; b < 4; ++b) { + bool nonzero = false; + for (int j = 0; j < QK4_0 / 2; ++j) { + const int lo = g[b].qs[j] & 0x0F; + const int hi = g[b].qs[j] >> 4; + if (lo < 7 || lo > 9 || hi < 7 || hi > 9) { + return false; + } + nonzero = nonzero || lo != 8 || hi != 8; + // values j and j + 16 of Q4_0 block b: index within the 256-value block + const int i_lo = 128 * group + 32 * b + j; + const int i_hi = i_lo + 16; + out[4 + i_lo / 4] |= static_cast((lo - 7) << (2 * (i_lo % 4))); + out[4 + i_hi / 4] |= static_cast((hi - 7) << (2 * (i_hi % 4))); + } + if (!nonzero) { + continue; + } + uint16_t d; + std::memcpy(&d, &g[b].d, sizeof(d)); + if (have_scale && d != scales[group]) { + return false; + } + scales[group] = d; + have_scale = true; + } + } + std::memcpy(out, scales, sizeof(scales)); + return true; +} + +} // namespace + +Status materialize_ternary_q4_0_k128(const HostWeightSource & source, std::vector & output) { + Status status; + if (source.source_type != GGML_TYPE_Q4_0 || source.input_size <= 0 || source.input_size % 256 != 0 || + source.output_size <= 0) { + status.log("layout %s requires Q4_0 with K divisible by 256", source.layout.c_str()); + return status; + } + static_assert(sizeof(block_q4_0) == 18); + const size_t blocks = static_cast(source.input_size / 256); + const size_t rows = static_cast(source.output_size); + const size_t in_bytes = rows * blocks * 8 * sizeof(block_q4_0); + const size_t out_bytes = ternary_q4_0_k128_bytes(source.input_size, source.output_size); + if (source.length != in_bytes || source.materialized_length != out_bytes) { + status.log("layout %s has inconsistent source/materialized lengths", source.layout.c_str()); + return status; + } + output.resize(out_bytes); + const auto * input = + reinterpret_cast(static_cast(source.host_data) + source.offset); + const size_t units = rows * blocks; + const size_t threads = + std::max(1, std::min(std::thread::hardware_concurrency(), (units + 4095) / 4096)); + std::atomic ok{ true }; + std::vector pool; + for (size_t t = 0; t < threads; ++t) { + pool.emplace_back([&, t] { + for (size_t u = units * t / threads; u < units * (t + 1) / threads && ok.load(std::memory_order_relaxed); + ++u) { + if (!pack_block(input + 8 * u, output.data() + 68 * u)) { + ok.store(false, std::memory_order_relaxed); + } + } + }); + } + for (std::thread & thread : pool) { + thread.join(); + } + if (!ok.load()) { + output.clear(); + status.log( + "layout %s: weight is not exact ternary Q4_0 (a nibble outside 7..9 or differing scales in a " + "128-value group); unset GGML_HRX_TERNARY_Q4_0 for this model", + source.layout.c_str()); + } + return status; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/host-memory-ternary.h b/ggml/src/ggml-hrx/runtime/host-memory-ternary.h new file mode 100644 index 000000000000..4f474414bd2e --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/host-memory-ternary.h @@ -0,0 +1,32 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Upload-time repack of exact-ternary Q4_0 weights into the packed ternary layout (dispatch/ternary-q4-0.h). +#pragma once + +#include "dispatch/ternary-q4-0.h" +#include "runtime/host-memory.h" +#include "status.h" + +#include +#include + +namespace ggml::hrx { + +// Verifies that every Q4_0 nibble is 7, 8 or 9 and that the four blocks of each 128-value group share one +// scale, then writes the packed layout. Fails (and writes nothing) on the first block that is not ternary. +Status materialize_ternary_q4_0_k128(const HostWeightSource & source, std::vector & output); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/host-memory.cpp b/ggml/src/ggml-hrx/runtime/host-memory.cpp new file mode 100644 index 000000000000..98b5064b8ae2 --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/host-memory.cpp @@ -0,0 +1,1593 @@ +#include "host-memory.h" + +#include "ggml-quants.h" +#include "hrx-interop-utils.h" +#include "hrx_runtime.h" +#include "runtime/host-memory-ternary.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr size_t kMaxInlineUploadBytes = 63 * 1024; +static constexpr size_t kLargeHostUploadBytes = 1024 * 1024; + +struct SymmetricI4Block { + std::array scales = {}; + std::array, 8> payloads = {}; +}; + +static bool checked_multiply(size_t lhs, size_t rhs, size_t & result) { + if (rhs != 0 && lhs > std::numeric_limits::max() / rhs) { + return false; + } + result = lhs * rhs; + return true; +} + +template static bool run_worker_threads(size_t thread_count, Work && work) { + std::atomic worker_failed = false; + std::vector workers; + auto join_workers = [&]() { + for (std::thread & worker : workers) { + if (worker.joinable()) { + worker.join(); + } + } + }; + + try { + workers.reserve(thread_count); + for (size_t thread = 0; thread < thread_count; ++thread) { + workers.emplace_back([&, thread]() { + try { + work(thread); + } catch (...) { + worker_failed.store(true, std::memory_order_relaxed); + } + }); + } + } catch (...) { + join_workers(); + return false; + } + join_workers(); + return !worker_failed.load(std::memory_order_relaxed); +} + +static void dequantize_block(ggml_type type, const uint8_t * source, float * destination) { + switch (type) { + case GGML_TYPE_Q4_K: + dequantize_row_q4_K(reinterpret_cast(source), destination, QK_K); + break; + case GGML_TYPE_Q5_K: + dequantize_row_q5_K(reinterpret_cast(source), destination, QK_K); + break; + case GGML_TYPE_IQ4_XS: + dequantize_row_iq4_xs(reinterpret_cast(source), destination, QK_K); + break; + case GGML_TYPE_Q6_K: + dequantize_row_q6_K(reinterpret_cast(source), destination, QK_K); + break; + default: + break; + } +} + +static bool quantize_symmetric_value(float value, float scale, int quant_min, int quant_max, int & quantized) { + if (!std::isfinite(value) || !std::isfinite(scale) || scale <= 0.0f) { + return false; + } + float quotient = value / scale; + if (std::isnan(quotient)) { + return false; + } + quotient = std::clamp(quotient, static_cast(quant_min), static_cast(quant_max)); + quantized = static_cast(std::nearbyint(quotient)); + return true; +} + +static bool fit_symmetric_scale(const float * values, + size_t value_count, + int quant_min, + int quant_max, + float & encoded_scale) { + float positive_max = 0.0f; + float negative_max = 0.0f; + bool nonzero = false; + for (size_t element = 0; element < value_count; ++element) { + if (!std::isfinite(values[element])) { + return false; + } + nonzero = nonzero || values[element] != 0.0f; + positive_max = std::max(positive_max, values[element]); + negative_max = std::max(negative_max, -values[element]); + } + if (!nonzero) { + encoded_scale = 0.0f; + return true; + } + + float scale = std::max(positive_max / quant_max, negative_max / -quant_min); + if (!std::isfinite(scale) || scale <= 0.0f) { + return false; + } + for (int iteration = 0; iteration < 2; ++iteration) { + double numerator = 0.0; + double denominator = 0.0; + for (size_t element = 0; element < value_count; ++element) { + int quantized = 0; + if (!quantize_symmetric_value(values[element], scale, quant_min, quant_max, quantized)) { + return false; + } + numerator += static_cast(values[element]) * quantized; + denominator += static_cast(quantized) * quantized; + } + if (denominator > 0.0) { + scale = static_cast(numerator / denominator); + if (!std::isfinite(scale) || scale <= 0.0f) { + return false; + } + } + } + + encoded_scale = ggml_fp16_to_fp32(ggml_fp32_to_fp16(scale)); + return std::isfinite(encoded_scale) && encoded_scale > 0.0f; +} + +static bool fit_shared_symmetric_scale(const std::array, 4> & values, + size_t row_count, + size_t group, + int quant_min, + int quant_max, + bool multistart, + float & encoded_scale) { + float positive_max = 0.0f; + float negative_max = 0.0f; + bool nonzero = false; + for (size_t row = 0; row < row_count; ++row) { + const float * group_values = values[row].data() + group * 32; + for (size_t element = 0; element < 32; ++element) { + if (!std::isfinite(group_values[element])) { + return false; + } + nonzero = nonzero || group_values[element] != 0.0f; + positive_max = std::max(positive_max, group_values[element]); + negative_max = std::max(negative_max, -group_values[element]); + } + } + if (!nonzero) { + encoded_scale = 0.0f; + return true; + } + + const float initial_scale = std::max(positive_max / quant_max, negative_max / -quant_min); + if (!std::isfinite(initial_scale) || initial_scale <= 0.0f) { + return false; + } + + auto refit = [&](float & scale) { + const int iteration_count = multistart ? 4 : 2; + for (int iteration = 0; iteration < iteration_count; ++iteration) { + double numerator = 0.0; + double denominator = 0.0; + for (size_t row = 0; row < row_count; ++row) { + const float * group_values = values[row].data() + group * 32; + for (size_t element = 0; element < 32; ++element) { + int quantized = 0; + if (!quantize_symmetric_value(group_values[element], scale, quant_min, quant_max, quantized)) { + return false; + } + numerator += static_cast(group_values[element]) * quantized; + denominator += static_cast(quantized) * quantized; + } + } + if (denominator > 0.0) { + scale = static_cast(numerator / denominator); + if (!std::isfinite(scale) || scale <= 0.0f) { + return false; + } + } + } + return true; + }; + + if (!multistart) { + float scale = initial_scale; + if (!refit(scale)) { + return false; + } + encoded_scale = ggml_fp16_to_fp32(ggml_fp32_to_fp16(scale)); + return std::isfinite(encoded_scale) && encoded_scale > 0.0f; + } + + constexpr std::array kInitialScaleRatios = { 0.50f, 0.60f, 0.70f, 0.80f, 0.90f, 1.00f, 1.10f }; + double best_error = std::numeric_limits::infinity(); + float best_scale = 0.0f; + for (float ratio : kInitialScaleRatios) { + float scale = initial_scale * ratio; + if (!refit(scale)) { + return false; + } + scale = ggml_fp16_to_fp32(ggml_fp32_to_fp16(scale)); + double error = 0.0; + for (size_t row = 0; row < row_count; ++row) { + const float * group_values = values[row].data() + group * 32; + for (size_t element = 0; element < 32; ++element) { + int quantized = 0; + if (!quantize_symmetric_value(group_values[element], scale, quant_min, quant_max, quantized)) { + return false; + } + const double delta = + static_cast(group_values[element]) - static_cast(quantized) * scale; + error += delta * delta; + } + } + if (error < best_error) { + best_error = error; + best_scale = scale; + } + } + + encoded_scale = best_scale; + return std::isfinite(encoded_scale) && encoded_scale > 0.0f; +} + +static bool quantize_symmetric_i4_k64(const float * values, SymmetricI4Block & output) { + for (size_t pair = 0; pair < output.scales.size() / 2; ++pair) { + const float * pair_values = values + pair * 64; + float scale = 0.0f; + if (!fit_symmetric_scale(pair_values, 64, -8, 7, scale)) { + return false; + } + if (scale == 0.0f) { + continue; + } + + const ggml_fp16_t encoded_scale = ggml_fp32_to_fp16(scale); + output.scales[pair * 2] = encoded_scale; + output.scales[pair * 2 + 1] = encoded_scale; + for (size_t half = 0; half < 2; ++half) { + const size_t logical_group = pair * 2 + half; + const float * group_values = pair_values + half * 32; + for (size_t element_pair = 0; element_pair < 16; ++element_pair) { + int low = 0; + int high = 0; + if (!quantize_symmetric_value(group_values[element_pair * 2], scale, -8, 7, low) || + !quantize_symmetric_value(group_values[element_pair * 2 + 1], scale, -8, 7, high)) { + return false; + } + output.payloads[logical_group][element_pair] = + static_cast((low & 0x0F) | ((high & 0x0F) << 4)); + } + } + } + return true; +} + +static bool quantize_symmetric_i4_k32(const float * values, SymmetricI4Block & output) { + for (size_t group = 0; group < output.scales.size(); ++group) { + const float * group_values = values + group * 32; + float scale = 0.0f; + if (!fit_symmetric_scale(group_values, 32, -8, 7, scale)) { + return false; + } + if (scale == 0.0f) { + continue; + } + output.scales[group] = ggml_fp32_to_fp16(scale); + for (size_t element_pair = 0; element_pair < 16; ++element_pair) { + int low = 0; + int high = 0; + if (!quantize_symmetric_value(group_values[element_pair * 2], scale, -8, 7, low) || + !quantize_symmetric_value(group_values[element_pair * 2 + 1], scale, -8, 7, high)) { + return false; + } + output.payloads[group][element_pair] = static_cast((low & 0x0F) | ((high & 0x0F) << 4)); + } + } + return true; +} + +using SymmetricI4Quantize = bool (*)(const float *, SymmetricI4Block &); + +static Status materialize_symmetric_i4_row64(const HostWeightSource & source, + SymmetricI4Quantize quantize, + std::vector & output) { + Status status; + if (source.source_type != GGML_TYPE_Q4_K && source.source_type != GGML_TYPE_Q5_K && + source.source_type != GGML_TYPE_Q6_K && source.source_type != GGML_TYPE_IQ4_XS) { + status.log("layout %s does not support GGML type %d", source.layout.c_str(), + static_cast(source.source_type)); + return status; + } + if (source.input_size <= 0 || source.input_size % QK_K != 0 || source.output_size <= 0 || + source.output_size % 64 != 0) { + status.log("layout %s requires K divisible by %d and rows divisible by 64, got K=%lld rows=%lld", + source.layout.c_str(), QK_K, static_cast(source.input_size), + static_cast(source.output_size)); + return status; + } + + const size_t row_bytes = ggml_row_size(source.source_type, source.input_size); + size_t expected_source_bytes = 0; + if (!checked_multiply(row_bytes, static_cast(source.output_size), expected_source_bytes) || + source.length != expected_source_bytes) { + status.log("layout %s source length %zu does not match expected %zu", source.layout.c_str(), source.length, + expected_source_bytes); + return status; + } + + constexpr size_t materialized_row_bytes = 144; + constexpr size_t row_group = 64; + const size_t logical_row_count = static_cast(source.output_size); + const size_t padded_row_count = (logical_row_count + row_group - 1) / row_group * row_group; + const size_t block_count = static_cast(source.input_size / QK_K); + size_t expected_output_bytes = 0; + if (!checked_multiply(padded_row_count, block_count, expected_output_bytes) || + !checked_multiply(expected_output_bytes, materialized_row_bytes, expected_output_bytes) || + source.materialized_length != expected_output_bytes) { + status.log("layout %s materialized length %zu does not match expected %zu", source.layout.c_str(), + source.materialized_length, expected_output_bytes); + return status; + } + + output.assign(expected_output_bytes, 0); + const auto * source_bytes = static_cast(source.host_data) + source.offset; + constexpr size_t field_count = 9; + constexpr size_t field_bytes = 16; + constexpr size_t block_out_bytes = field_count * field_bytes; + const size_t source_block_bytes = ggml_type_size(source.source_type); + const size_t thread_count = std::min(logical_row_count, std::max(1u, std::thread::hardware_concurrency())); + std::atomic conversion_failed = false; + + auto convert_rows = [&](size_t row_begin, size_t row_end) { + std::array decoded = {}; + for (size_t row = row_begin; row < row_end; ++row) { + if (conversion_failed.load(std::memory_order_relaxed)) { + return; + } + const size_t group = row / row_group; + const size_t lane = row % row_group; + const uint8_t * source_row = source_bytes + row * row_bytes; + for (size_t block = 0; block < block_count; ++block) { + dequantize_block(source.source_type, source_row + block * source_block_bytes, decoded.data()); + SymmetricI4Block converted; + if (!quantize(decoded.data(), converted)) { + conversion_failed.store(true, std::memory_order_relaxed); + return; + } + const size_t group_base = (group * block_count + block) * row_group * block_out_bytes; + uint8_t * header = output.data() + group_base + lane * field_bytes; + for (size_t group = 0; group < converted.scales.size(); ++group) { + std::memcpy(header + group * sizeof(ggml_fp16_t), &converted.scales[group], sizeof(ggml_fp16_t)); + } + for (size_t logical_group = 0; logical_group < converted.payloads.size(); ++logical_group) { + uint8_t * payload = + output.data() + group_base + (1 + logical_group) * row_group * field_bytes + lane * field_bytes; + std::memcpy(payload, converted.payloads[logical_group].data(), field_bytes); + } + } + } + }; + + if (!run_worker_threads(thread_count, [&](size_t thread) { + const size_t row_begin = logical_row_count * thread / thread_count; + const size_t row_end = logical_row_count * (thread + 1) / thread_count; + convert_rows(row_begin, row_end); + })) { + output.clear(); + status.log("layout %s host materialization worker failed", source.layout.c_str()); + return status; + } + if (conversion_failed.load(std::memory_order_relaxed)) { + output.clear(); + status.log("layout %s cannot represent non-finite or out-of-range symmetric weights", source.layout.c_str()); + } + return status; +} + +static Status materialize_symmetric_i4_k32_row64(const HostWeightSource & source, std::vector & output) { + return materialize_symmetric_i4_row64(source, quantize_symmetric_i4_k32, output); +} + +static Status materialize_symmetric_i4_k64_row64(const HostWeightSource & source, std::vector & output) { + return materialize_symmetric_i4_row64(source, quantize_symmetric_i4_k64, output); +} + +static Status materialize_symmetric_k32_eightgroups_shared4(const HostWeightSource & source, + int quant_bits, + std::vector & output) { + Status status; + if (source.source_type != GGML_TYPE_Q4_K && source.source_type != GGML_TYPE_Q5_K && + source.source_type != GGML_TYPE_Q6_K && source.source_type != GGML_TYPE_IQ4_XS) { + status.log("layout %s does not support GGML type %d", source.layout.c_str(), + static_cast(source.source_type)); + return status; + } + if (source.input_size <= 0 || source.input_size % QK_K != 0 || source.output_size <= 0) { + status.log("layout %s requires K divisible by %d and positive rows, got K=%lld rows=%lld", + source.layout.c_str(), QK_K, static_cast(source.input_size), + static_cast(source.output_size)); + return status; + } + + constexpr size_t shared_rows = 4; + const size_t logical_rows = static_cast(source.output_size); + if (logical_rows > std::numeric_limits::max() - 255) { + status.log("layout %s row count overflows row-group sizing", source.layout.c_str()); + return status; + } + const size_t materialized_row_bytes = 4 + 8 * 32 * static_cast(quant_bits) / 8; + const size_t unpadded_bytes = + logical_rows * static_cast(source.input_size / QK_K) * materialized_row_bytes; + const size_t dynamic_row_group = ((logical_rows + 255) / 256) * 32; + const size_t row_group = + quant_bits == 4 && unpadded_bytes < size_t{ 16 } * 1024 * 1024 ? 32 : + quant_bits == 4 && unpadded_bytes <= size_t{ 32 } * 1024 * 1024 && source.output_size > source.input_size ? 96 : + dynamic_row_group; + if (logical_rows > std::numeric_limits::max() - (row_group - 1)) { + status.log("layout %s row count overflows row-group padding", source.layout.c_str()); + return status; + } + const size_t physical_rows = (logical_rows + row_group - 1) / row_group * row_group; + + const size_t row_bytes = ggml_row_size(source.source_type, source.input_size); + size_t expected_source_bytes = 0; + if (!checked_multiply(row_bytes, logical_rows, expected_source_bytes) || source.length != expected_source_bytes) { + status.log("layout %s source length %zu does not match expected %zu", source.layout.c_str(), source.length, + expected_source_bytes); + return status; + } + + if (quant_bits != 2 && quant_bits != 4) { + status.log("layout %s requires 2-bit or 4-bit symmetric materialization", source.layout.c_str()); + return status; + } + + const int quant_min = -(1 << (quant_bits - 1)); + const int quant_max = (1 << (quant_bits - 1)) - 1; + const size_t field_bytes = 32 * static_cast(quant_bits) / 8; + const size_t encoded_row_bytes = 4 + 8 * field_bytes; + const size_t block_count = static_cast(source.input_size / QK_K); + size_t expected_output_bytes = 0; + if (!checked_multiply(physical_rows, block_count, expected_output_bytes) || + !checked_multiply(expected_output_bytes, encoded_row_bytes, expected_output_bytes) || + source.materialized_length != expected_output_bytes) { + status.log("layout %s materialized length %zu does not match expected %zu", source.layout.c_str(), + source.materialized_length, expected_output_bytes); + return status; + } + + output.assign(expected_output_bytes, uint8_t{ 0 }); + const auto * source_bytes = static_cast(source.host_data) + source.offset; + const size_t scale_plane_bytes = row_group / shared_rows * 16; + const size_t payload_plane_bytes = row_group * field_bytes; + const size_t payload_block_bytes = 8 * payload_plane_bytes; + const size_t block_out_bytes = scale_plane_bytes + payload_block_bytes; + const size_t row_group_bytes = block_count * block_out_bytes; + const size_t source_block_bytes = ggml_type_size(source.source_type); + const size_t logical_cohorts = (logical_rows + shared_rows - 1) / shared_rows; + const size_t thread_count = std::min(logical_cohorts, std::max(1u, std::thread::hardware_concurrency())); + std::atomic conversion_failed = false; + + auto convert_cohorts = [&](size_t cohort_begin, size_t cohort_end) { + std::array, shared_rows> decoded = {}; + for (size_t cohort = cohort_begin; cohort < cohort_end; ++cohort) { + if (conversion_failed.load(std::memory_order_relaxed)) { + return; + } + const size_t first_row = cohort * shared_rows; + const size_t row_count = std::min(shared_rows, logical_rows - first_row); + const size_t row_group_id = first_row / row_group; + const size_t cohort_lane = first_row % row_group / shared_rows; + for (size_t block = 0; block < block_count; ++block) { + for (size_t row = 0; row < row_count; ++row) { + const uint8_t * source_row = source_bytes + (first_row + row) * row_bytes; + dequantize_block(source.source_type, source_row + block * source_block_bytes, decoded[row].data()); + } + + const size_t group_base = row_group_id * row_group_bytes; + const size_t payload_block_base = group_base + block * payload_block_bytes; + const size_t scale_region = group_base + block_count * payload_block_bytes; + uint8_t * header = output.data() + scale_region + block * scale_plane_bytes + cohort_lane * 16; + std::array scales = {}; + for (size_t group = 0; group < scales.size(); ++group) { + float scale = 0.0f; + const bool multistart = source.layout == kSymmetricI4K32EightGroupsShared4MultistartLayout; + if (!fit_shared_symmetric_scale(decoded, row_count, group, quant_min, quant_max, multistart, + scale)) { + conversion_failed.store(true, std::memory_order_relaxed); + return; + } + const ggml_fp16_t encoded_scale = ggml_fp32_to_fp16(scale); + std::memcpy(header + group * sizeof(encoded_scale), &encoded_scale, sizeof(encoded_scale)); + scales[group] = scale; + } + + for (size_t row = 0; row < row_count; ++row) { + const size_t row_lane = (first_row + row) % row_group; + for (size_t group = 0; group < scales.size(); ++group) { + uint8_t * payload = + output.data() + payload_block_base + group * payload_plane_bytes + row_lane * field_bytes; + const float * group_values = decoded[row].data() + group * 32; + const float scale = scales[group]; + if (scale == 0.0f) { + continue; + } + if (quant_bits == 4) { + for (size_t element_pair = 0; element_pair < 16; ++element_pair) { + int low = 0; + int high = 0; + if (!quantize_symmetric_value(group_values[element_pair * 2], scale, quant_min, + quant_max, low) || + !quantize_symmetric_value(group_values[element_pair * 2 + 1], scale, quant_min, + quant_max, high)) { + conversion_failed.store(true, std::memory_order_relaxed); + return; + } + payload[element_pair] = static_cast((low & 0x0F) | ((high & 0x0F) << 4)); + } + } else { + for (size_t element_quartet = 0; element_quartet < 8; ++element_quartet) { + uint8_t packed = 0; + for (size_t element = 0; element < 4; ++element) { + const size_t index = element_quartet * 4 + element; + int quantized = 0; + if (!quantize_symmetric_value(group_values[index], scale, quant_min, quant_max, + quantized)) { + conversion_failed.store(true, std::memory_order_relaxed); + return; + } + packed |= static_cast((quantized & 0x03) << (element * 2)); + } + payload[element_quartet] = packed; + } + } + } + } + } + } + }; + + if (!run_worker_threads(thread_count, [&](size_t thread) { + const size_t cohort_begin = logical_cohorts * thread / thread_count; + const size_t cohort_end = logical_cohorts * (thread + 1) / thread_count; + convert_cohorts(cohort_begin, cohort_end); + })) { + output.clear(); + status.log("layout %s host materialization worker failed", source.layout.c_str()); + return status; + } + if (conversion_failed.load(std::memory_order_relaxed)) { + output.clear(); + status.log("layout %s cannot represent non-finite or out-of-range symmetric weights", source.layout.c_str()); + } + return status; +} + +static Status materialize_symmetric_i4_k32_eightgroups_shared4(const HostWeightSource & source, + std::vector & output) { + return materialize_symmetric_k32_eightgroups_shared4(source, 4, output); +} + +static Status materialize_symmetric_i5_k32(const HostWeightSource & source, std::vector & output) { + Status status; + if (source.source_type != GGML_TYPE_Q5_K) { + status.log("layout %s does not support GGML type %d", source.layout.c_str(), + static_cast(source.source_type)); + return status; + } + if (source.input_size <= 0 || source.input_size % QK_K != 0 || source.output_size <= 0 || + source.output_size % 64 != 0) { + status.log("layout %s requires K divisible by %d and rows divisible by 64, got K=%lld rows=%lld", + source.layout.c_str(), QK_K, static_cast(source.input_size), + static_cast(source.output_size)); + return status; + } + + const size_t row_bytes = ggml_row_size(GGML_TYPE_Q5_K, source.input_size); + const size_t row_count = static_cast(source.output_size); + const size_t block_count = static_cast(source.input_size / QK_K); + size_t expected_bytes = 0; + if (!checked_multiply(row_bytes, row_count, expected_bytes) || source.length != expected_bytes || + source.materialized_length != expected_bytes) { + status.log("layout %s has inconsistent source/materialized lengths for K=%lld rows=%lld", source.layout.c_str(), + static_cast(source.input_size), static_cast(source.output_size)); + return status; + } + + output.assign(expected_bytes, uint8_t{ 0 }); + const auto * source_bytes = static_cast(source.host_data) + source.offset; + const size_t thread_count = + std::min(row_count, static_cast(std::max(1u, std::thread::hardware_concurrency()))); + std::atomic conversion_failed = false; + + auto convert_rows = [&](size_t row_begin, size_t row_end) { + std::array values = {}; + std::array, 8> codes = {}; + for (size_t row = row_begin; row < row_end; ++row) { + const uint8_t * source_row = source_bytes + row * row_bytes; + uint8_t * output_row = output.data() + row * row_bytes; + for (size_t block = 0; block < block_count; ++block) { + dequantize_row_q5_K(reinterpret_cast(source_row) + block, values.data(), QK_K); + uint8_t * record = output_row + block * sizeof(block_q5_K); + for (size_t group = 0; group < 8; ++group) { + const float * group_values = values.data() + group * 32; + float positive_max = 0.0f; + float negative_max = 0.0f; + for (size_t element = 0; element < 32; ++element) { + const float value = group_values[element]; + if (!std::isfinite(value)) { + conversion_failed.store(true, std::memory_order_relaxed); + return; + } + positive_max = std::max(positive_max, value); + negative_max = std::max(negative_max, -value); + } + float scale = std::max(positive_max / 15.0f, negative_max / 16.0f); + if (scale == 0.0f) { + scale = 1.0f; + } + for (int iteration = 0; iteration < 2; ++iteration) { + double numerator = 0.0; + double denominator = 0.0; + for (size_t element = 0; element < 32; ++element) { + int quantized = 0; + if (!quantize_symmetric_value(group_values[element], scale, -16, 15, quantized)) { + conversion_failed.store(true, std::memory_order_relaxed); + return; + } + numerator += static_cast(group_values[element]) * quantized; + denominator += static_cast(quantized) * quantized; + } + if (denominator > 0.0) { + scale = static_cast(numerator / denominator); + } + if (!std::isfinite(scale) || scale <= 0.0f) { + conversion_failed.store(true, std::memory_order_relaxed); + return; + } + } + + const ggml_fp16_t encoded_scale = ggml_fp32_to_fp16(scale); + std::memcpy(record + group * sizeof(encoded_scale), &encoded_scale, sizeof(encoded_scale)); + scale = ggml_fp16_to_fp32(encoded_scale); + for (size_t element = 0; element < 32; ++element) { + int quantized = 0; + if (!quantize_symmetric_value(group_values[element], scale, -16, 15, quantized)) { + conversion_failed.store(true, std::memory_order_relaxed); + return; + } + codes[group][element] = static_cast(quantized) & uint8_t{ 31 }; + } + } + + uint8_t * high = record + 16; + uint8_t * packed = record + 48; + for (size_t element = 0; element < 32; ++element) { + uint8_t high_byte = 0; + for (size_t group = 0; group < 8; ++group) { + high_byte |= static_cast(((codes[group][element] >> 4) & 1) << group); + } + high[element] = high_byte; + } + for (size_t pair = 0; pair < 4; ++pair) { + const size_t group0 = pair * 2; + const size_t group1 = group0 + 1; + for (size_t element = 0; element < 32; ++element) { + packed[pair * 32 + element] = + static_cast((codes[group0][element] & 15) | ((codes[group1][element] & 15) << 4)); + } + } + } + } + }; + + if (!run_worker_threads(thread_count, [&](size_t thread) { + convert_rows(row_count * thread / thread_count, row_count * (thread + 1) / thread_count); + })) { + output.clear(); + status.log("layout %s host materialization worker failed", source.layout.c_str()); + return status; + } + if (conversion_failed.load(std::memory_order_relaxed)) { + output.clear(); + status.log("layout %s cannot represent non-finite symmetric weights", source.layout.c_str()); + } + return status; +} + +static Status materialize_symmetric_i8_k256_row64(const HostWeightSource & source, std::vector & output) { + Status status; + if (source.source_type != GGML_TYPE_Q5_K) { + status.log("layout %s does not support GGML type %d", source.layout.c_str(), + static_cast(source.source_type)); + return status; + } + if (source.input_size <= 0 || source.input_size % QK_K != 0 || source.output_size <= 0 || + source.output_size % 64 != 0) { + status.log("layout %s requires K divisible by %d and rows divisible by 64, got K=%lld rows=%lld", + source.layout.c_str(), QK_K, static_cast(source.input_size), + static_cast(source.output_size)); + return status; + } + + const size_t row_bytes = ggml_row_size(GGML_TYPE_Q5_K, source.input_size); + const size_t row_count = static_cast(source.output_size); + const size_t block_count = static_cast(source.input_size / QK_K); + constexpr size_t row_group = 64; + constexpr size_t scale_plane_bytes = row_group * sizeof(ggml_fp16_t); + constexpr size_t field_bytes = row_group * 64; + constexpr size_t record_bytes = scale_plane_bytes + 4 * field_bytes; + constexpr size_t row_record_bytes = record_bytes / row_group; + + size_t expected_source_bytes = 0; + size_t expected_output_bytes = 0; + if (!checked_multiply(row_bytes, row_count, expected_source_bytes) || source.length != expected_source_bytes || + !checked_multiply(row_count, block_count, expected_output_bytes) || + !checked_multiply(expected_output_bytes, row_record_bytes, expected_output_bytes) || + source.materialized_length != expected_output_bytes) { + status.log("layout %s has inconsistent source/materialized lengths for K=%lld rows=%lld", source.layout.c_str(), + static_cast(source.input_size), static_cast(source.output_size)); + return status; + } + + output.resize(expected_output_bytes); + const auto * source_bytes = static_cast(source.host_data) + source.offset; + const size_t thread_count = + std::min(row_count, static_cast(std::max(1u, std::thread::hardware_concurrency()))); + std::atomic conversion_failed = false; + + auto convert_rows = [&](size_t row_begin, size_t row_end) { + std::array values = {}; + for (size_t row = row_begin; row < row_end; ++row) { + const size_t row_group_id = row / row_group; + const size_t lane = row % row_group; + const uint8_t * source_row = source_bytes + row * row_bytes; + for (size_t block = 0; block < block_count; ++block) { + dequantize_row_q5_K(reinterpret_cast(source_row) + block, values.data(), QK_K); + + float positive_max = 0.0f; + float negative_max = 0.0f; + for (float value : values) { + if (!std::isfinite(value)) { + conversion_failed.store(true, std::memory_order_relaxed); + return; + } + positive_max = std::max(positive_max, value); + negative_max = std::max(negative_max, -value); + } + float scale = std::max(positive_max, negative_max) / 127.0f; + if (scale == 0.0f) { + scale = 1.0f; + } + for (int iteration = 0; iteration < 2; ++iteration) { + double numerator = 0.0; + double denominator = 0.0; + for (float value : values) { + int quantized = 0; + if (!quantize_symmetric_value(value, scale, -127, 127, quantized)) { + conversion_failed.store(true, std::memory_order_relaxed); + return; + } + numerator += static_cast(value) * quantized; + denominator += static_cast(quantized) * quantized; + } + if (denominator > 0.0) { + scale = static_cast(numerator / denominator); + } + if (!std::isfinite(scale) || scale <= 0.0f) { + conversion_failed.store(true, std::memory_order_relaxed); + return; + } + } + + const ggml_fp16_t encoded_scale = ggml_fp32_to_fp16(scale); + scale = ggml_fp16_to_fp32(encoded_scale); + const size_t record_base = (row_group_id * block_count + block) * record_bytes; + std::memcpy(output.data() + record_base + lane * sizeof(encoded_scale), &encoded_scale, + sizeof(encoded_scale)); + for (size_t chunk = 0; chunk < 4; ++chunk) { + auto * payload = reinterpret_cast(output.data() + record_base + scale_plane_bytes + + chunk * field_bytes + lane * 64); + for (size_t element = 0; element < 64; ++element) { + int quantized = 0; + if (!quantize_symmetric_value(values[chunk * 64 + element], scale, -127, 127, quantized)) { + conversion_failed.store(true, std::memory_order_relaxed); + return; + } + payload[element] = static_cast(quantized); + } + } + } + } + }; + + if (!run_worker_threads(thread_count, [&](size_t thread) { + convert_rows(row_count * thread / thread_count, row_count * (thread + 1) / thread_count); + })) { + output.clear(); + status.log("layout %s host materialization worker failed", source.layout.c_str()); + return status; + } + if (conversion_failed.load(std::memory_order_relaxed)) { + output.clear(); + status.log("layout %s cannot represent non-finite symmetric weights", source.layout.c_str()); + } + return status; +} + +// Transpose native Q4_K records in 16-byte fields. No scale or code is +// reconstructed: applying the inverse permutation restores every source byte. +static Status materialize_q4_k_packed_k256_row64(const HostWeightSource & source, std::vector & output) { + Status status; + if (source.source_type != GGML_TYPE_Q4_K || source.input_size <= 0 || source.input_size % QK_K != 0 || + source.output_size <= 0 || source.output_size % 64 != 0) { + status.log("layout %s requires Q4_K, K divisible by %d, and rows divisible by 64", source.layout.c_str(), + QK_K); + return status; + } + + constexpr size_t row_group = 64; + constexpr size_t field_bytes = 16; + constexpr size_t field_count = sizeof(block_q4_K) / field_bytes; + static_assert(sizeof(block_q4_K) == 144); + const size_t blocks = static_cast(source.input_size / QK_K); + const size_t rows = static_cast(source.output_size); + size_t bytes = 0; + if (!checked_multiply(ggml_row_size(GGML_TYPE_Q4_K, source.input_size), rows, bytes) || + source.length != bytes || source.materialized_length != bytes) { + status.log("layout %s has inconsistent source/materialized lengths", source.layout.c_str()); + return status; + } + + output.resize(bytes); + const auto * input = static_cast(source.host_data) + source.offset; + const size_t groups = rows / row_group; + const size_t thread_count = std::min(groups, std::max(1u, std::thread::hardware_concurrency())); + if (!run_worker_threads(thread_count, [&](size_t thread) { + const size_t begin = groups * thread / thread_count; + const size_t end = groups * (thread + 1) / thread_count; + for (size_t group = begin; group < end; ++group) { + for (size_t block = 0; block < blocks; ++block) { + for (size_t field = 0; field < field_count; ++field) { + for (size_t lane = 0; lane < row_group; ++lane) { + const size_t src = (((group * row_group + lane) * blocks + block) * field_count + field) * + field_bytes; + const size_t dst = (((group * blocks + block) * field_count + field) * row_group + lane) * + field_bytes; + std::memcpy(output.data() + dst, input + src, field_bytes); + } + } + } + } + })) { + output.clear(); + status.log("layout %s host materialization worker failed", source.layout.c_str()); + } + return status; +} + +static Status materialize_q6_k_i8_k32_row64(const HostWeightSource & source, std::vector & output) { + Status status; + if (source.source_type != GGML_TYPE_Q6_K) { + status.log("layout %s does not support GGML type %d", source.layout.c_str(), + static_cast(source.source_type)); + return status; + } + if (source.input_size <= 0 || source.input_size % QK_K != 0 || source.output_size <= 0 || + source.output_size % 64 != 0) { + status.log("layout %s requires K divisible by %d and rows divisible by 64, got K=%lld rows=%lld", + source.layout.c_str(), QK_K, static_cast(source.input_size), + static_cast(source.output_size)); + return status; + } + + static_assert(sizeof(block_q6_K) == 210); + constexpr size_t row_group = 64; + constexpr size_t group_count = 8; + constexpr size_t group_elements = 32; + constexpr size_t d_plane_bytes = row_group * sizeof(ggml_fp16_t); + constexpr size_t scale_plane_bytes = group_count * row_group * 2; + constexpr size_t code_plane_bytes = group_count * row_group * group_elements; + constexpr size_t tile_block_bytes = d_plane_bytes + scale_plane_bytes + code_plane_bytes; + static_assert(tile_block_bytes == row_group * 274); + + const size_t block_count = static_cast(source.input_size / QK_K); + const size_t row_bytes = block_count * sizeof(block_q6_K); + size_t expected_source_bytes = 0; + size_t expected_output_bytes = 0; + if (!checked_multiply(row_bytes, static_cast(source.output_size), expected_source_bytes) || + !checked_multiply(static_cast(source.output_size), block_count, expected_output_bytes) || + !checked_multiply(expected_output_bytes, size_t{ 274 }, expected_output_bytes) || + source.length != expected_source_bytes || source.materialized_length != expected_output_bytes) { + status.log("layout %s source/materialized lengths %zu/%zu do not match expected %zu/%zu", source.layout.c_str(), + source.length, source.materialized_length, expected_source_bytes, expected_output_bytes); + return status; + } + + output.resize(expected_output_bytes); + const auto * source_bytes = static_cast(source.host_data) + source.offset; + const size_t output_group_count = static_cast(source.output_size) / row_group; + const size_t thread_count = std::min(output_group_count, std::max(1u, std::thread::hardware_concurrency())); + + auto convert_groups = [&](size_t group_begin, size_t group_end) { + for (size_t output_group = group_begin; output_group < group_end; ++output_group) { + for (size_t block = 0; block < block_count; ++block) { + const size_t tile_base = (output_group * block_count + block) * tile_block_bytes; + uint8_t * d_plane = output.data() + tile_base; + int8_t * scale_plane = reinterpret_cast(d_plane + d_plane_bytes); + int8_t * code_plane = scale_plane + scale_plane_bytes; + for (size_t lane = 0; lane < row_group; ++lane) { + const size_t row = output_group * row_group + lane; + const auto * source_block = reinterpret_cast(source_bytes + row * row_bytes + + block * sizeof(block_q6_K)); + std::memcpy(d_plane + lane * sizeof(ggml_fp16_t), &source_block->d, sizeof(ggml_fp16_t)); + for (size_t quant_group = 0; quant_group < group_count; ++quant_group) { + int8_t * scales = scale_plane + (quant_group * row_group + lane) * 2; + scales[0] = source_block->scales[quant_group * 2]; + scales[1] = source_block->scales[quant_group * 2 + 1]; + int8_t * codes = code_plane + (quant_group * row_group + lane) * group_elements; + const size_t half128 = quant_group / 4; + const size_t group_in_half = quant_group % 4; + for (size_t element = 0; element < group_elements; ++element) { + const size_t low_index = half128 * 64 + (group_in_half % 2) * 32 + element; + const size_t high_index = half128 * 32 + element; + const uint8_t low = + static_cast((source_block->ql[low_index] >> ((group_in_half / 2) * 4)) & 0x0F); + const uint8_t high = + static_cast((source_block->qh[high_index] >> (group_in_half * 2)) & 0x03); + codes[element] = static_cast((low | (high << 4)) - 32); + } + } + } + } + } + }; + + if (!run_worker_threads(thread_count, [&](size_t thread) { + const size_t group_begin = output_group_count * thread / thread_count; + const size_t group_end = output_group_count * (thread + 1) / thread_count; + convert_groups(group_begin, group_end); + })) { + output.clear(); + status.log("layout %s host materialization worker failed", source.layout.c_str()); + } + return status; +} + +static Status materialize_q6_k_packed_k256_row64_scalerow(const HostWeightSource & source, + std::vector & output) { + Status status; + if (source.source_type != GGML_TYPE_Q6_K) { + status.log("layout %s does not support GGML type %d", source.layout.c_str(), + static_cast(source.source_type)); + return status; + } + if (source.input_size <= 0 || source.input_size % QK_K != 0 || source.output_size <= 0 || + source.output_size % 64 != 0) { + status.log("layout %s requires K divisible by %d and rows divisible by 64, got K=%lld rows=%lld", + source.layout.c_str(), QK_K, static_cast(source.input_size), + static_cast(source.output_size)); + return status; + } + + static_assert(sizeof(block_q6_K) == 210); + constexpr size_t row_group = 64; + constexpr size_t half_count = 2; + constexpr size_t raw_field_count = 3; + constexpr size_t raw_field_bytes = 32; + constexpr size_t scale_group_count = 8; + constexpr size_t d_plane_bytes = row_group * sizeof(ggml_fp16_t); + constexpr size_t scale_plane_bytes = scale_group_count * row_group * half_count; + constexpr size_t raw_plane_bytes = half_count * raw_field_count * row_group * raw_field_bytes; + constexpr size_t tile_block_bytes = d_plane_bytes + scale_plane_bytes + raw_plane_bytes; + static_assert(tile_block_bytes == row_group * sizeof(block_q6_K)); + + const size_t block_count = static_cast(source.input_size / QK_K); + const size_t row_bytes = block_count * sizeof(block_q6_K); + size_t expected_source_bytes = 0; + if (!checked_multiply(row_bytes, static_cast(source.output_size), expected_source_bytes) || + source.length != expected_source_bytes || source.materialized_length != expected_source_bytes) { + status.log("layout %s source/materialized lengths %zu/%zu do not match expected %zu", source.layout.c_str(), + source.length, source.materialized_length, expected_source_bytes); + return status; + } + + output.resize(expected_source_bytes); + const auto * source_bytes = static_cast(source.host_data) + source.offset; + const size_t output_group_count = static_cast(source.output_size) / row_group; + const size_t thread_count = std::min(output_group_count, std::max(1u, std::thread::hardware_concurrency())); + + auto convert_groups = [&](size_t group_begin, size_t group_end) { + for (size_t output_group = group_begin; output_group < group_end; ++output_group) { + for (size_t block = 0; block < block_count; ++block) { + const size_t tile_base = (output_group * block_count + block) * tile_block_bytes; + uint8_t * d_plane = output.data() + tile_base; + int8_t * scale_plane = reinterpret_cast(d_plane + d_plane_bytes); + uint8_t * raw_plane = reinterpret_cast(scale_plane + scale_plane_bytes); + for (size_t lane = 0; lane < row_group; ++lane) { + const size_t row = output_group * row_group + lane; + const auto * source_block = reinterpret_cast(source_bytes + row * row_bytes + + block * sizeof(block_q6_K)); + std::memcpy(d_plane + lane * sizeof(ggml_fp16_t), &source_block->d, sizeof(ggml_fp16_t)); + for (size_t quant_group = 0; quant_group < scale_group_count; ++quant_group) { + for (size_t half = 0; half < half_count; ++half) { + const size_t destination_index = + (lane * half_count + half) * scale_group_count + quant_group; + scale_plane[destination_index] = source_block->scales[quant_group * half_count + half]; + } + } + for (size_t half = 0; half < half_count; ++half) { + uint8_t * ql0 = raw_plane + ((half * raw_field_count) * row_group + lane) * raw_field_bytes; + uint8_t * ql1 = ql0 + row_group * raw_field_bytes; + uint8_t * qh = ql1 + row_group * raw_field_bytes; + std::memcpy(ql0, source_block->ql + half * 64, raw_field_bytes); + std::memcpy(ql1, source_block->ql + half * 64 + raw_field_bytes, raw_field_bytes); + std::memcpy(qh, source_block->qh + half * raw_field_bytes, raw_field_bytes); + } + } + } + } + }; + + if (!run_worker_threads(thread_count, [&](size_t thread) { + const size_t group_begin = output_group_count * thread / thread_count; + const size_t group_end = output_group_count * (thread + 1) / thread_count; + convert_groups(group_begin, group_end); + })) { + output.clear(); + status.log("layout %s host materialization worker failed", source.layout.c_str()); + } + return status; +} + +static Status materialize_q6_k_symmetric_i2_packed_k256_row64_scalerow(const HostWeightSource & source, + std::vector & output) { + Status status; + if (source.source_type != GGML_TYPE_Q6_K || source.input_size <= 0 || source.input_size % QK_K != 0 || + source.output_size <= 0 || source.output_size % 64 != 0 || + static_cast(source.output_size) > std::numeric_limits::max() - 255) { + status.log( + "layout %s requires Q6_K, K divisible by %d, and rows divisible by 64, got type=%d K=%lld " + "rows=%lld", + source.layout.c_str(), QK_K, static_cast(source.source_type), + static_cast(source.input_size), static_cast(source.output_size)); + return status; + } + + const size_t logical_rows = static_cast(source.output_size); + const size_t row_group = ((logical_rows + 255) / 256) * 32; + const size_t physical_rows = (logical_rows + row_group - 1) / row_group * row_group; + const size_t block_count = static_cast(source.input_size / QK_K); + size_t symmetric_bytes = 0; + size_t packed_bytes = 0; + if (!checked_multiply(physical_rows, block_count, symmetric_bytes) || + !checked_multiply(symmetric_bytes, size_t{ 68 }, symmetric_bytes) || + !checked_multiply(logical_rows, block_count, packed_bytes) || + !checked_multiply(packed_bytes, sizeof(block_q6_K), packed_bytes) || + symmetric_bytes > std::numeric_limits::max() - packed_bytes || source.length != packed_bytes || + source.materialized_length != symmetric_bytes + packed_bytes) { + status.log("layout %s source/materialized lengths %zu/%zu are inconsistent with K=%lld rows=%lld", + source.layout.c_str(), source.length, source.materialized_length, + static_cast(source.input_size), static_cast(source.output_size)); + return status; + } + + HostWeightSource symmetric_source = source; + symmetric_source.layout = kSymmetricI2K32EightGroupsShared4Layout; + symmetric_source.materialized_length = symmetric_bytes; + status = materialize_symmetric_k32_eightgroups_shared4(symmetric_source, 2, output); + if (!status.success()) { + return status; + } + output.resize(source.materialized_length); + + HostWeightSource packed_source = source; + packed_source.layout = kQ6KPackedK256Row64ScaleRowLayout; + packed_source.materialized_length = packed_bytes; + std::vector packed; + status = materialize_q6_k_packed_k256_row64_scalerow(packed_source, packed); + if (!status.success()) { + output.clear(); + return status; + } + std::memcpy(output.data() + symmetric_bytes, packed.data(), packed.size()); + return status; +} + +static Status materialize_weight(const HostWeightSource & source, + const void *& upload_data, + size_t & upload_size, + std::vector & transformed) { + Status status; + if (source.layout == kNativeWeightLayout) { + upload_data = static_cast(source.host_data) + source.offset; + upload_size = source.length; + return status; + } + if (source.layout == kSymmetricI4K32EightGroupsShared4Layout || + source.layout == kSymmetricI4K32EightGroupsShared4MultistartLayout) { + status = materialize_symmetric_i4_k32_eightgroups_shared4(source, transformed); + if (status.success()) { + upload_data = transformed.data(); + upload_size = transformed.size(); + } + return status; + } + if (source.layout == kSymmetricI4K32Row64Layout) { + status = materialize_symmetric_i4_k32_row64(source, transformed); + if (status.success()) { + upload_data = transformed.data(); + upload_size = transformed.size(); + } + return status; + } + if (source.layout == kSymmetricI4K64Row64Layout) { + status = materialize_symmetric_i4_k64_row64(source, transformed); + if (status.success()) { + upload_data = transformed.data(); + upload_size = transformed.size(); + } + return status; + } + if (source.layout == kQ5KSymmetricI5K32Layout) { + status = materialize_symmetric_i5_k32(source, transformed); + if (status.success()) { + upload_data = transformed.data(); + upload_size = transformed.size(); + } + return status; + } + if (source.layout == kQ5KSymmetricI8K256Row64Layout) { + status = materialize_symmetric_i8_k256_row64(source, transformed); + if (status.success()) { + upload_data = transformed.data(); + upload_size = transformed.size(); + } + return status; + } + if (source.layout == kTernaryQ40K128Layout) { + status = materialize_ternary_q4_0_k128(source, transformed); + if (status.success()) { + upload_data = transformed.data(); + upload_size = transformed.size(); + } + return status; + } + if (source.layout == kQ4KPackedK256Row64Layout) { + status = materialize_q4_k_packed_k256_row64(source, transformed); + if (status.success()) { + upload_data = transformed.data(); + upload_size = transformed.size(); + } + return status; + } + if (source.layout == kQ6KI8K32Row64Layout) { + status = materialize_q6_k_i8_k32_row64(source, transformed); + if (status.success()) { + upload_data = transformed.data(); + upload_size = transformed.size(); + } + return status; + } + if (source.layout == kQ6KPackedK256Row64ScaleRowLayout) { + status = materialize_q6_k_packed_k256_row64_scalerow(source, transformed); + if (status.success()) { + upload_data = transformed.data(); + upload_size = transformed.size(); + } + return status; + } + if (source.layout == kQ6KSymmetricI2PackedK256Row64ScaleRowLayout) { + status = materialize_q6_k_symmetric_i2_packed_k256_row64_scalerow(source, transformed); + if (status.success()) { + upload_data = transformed.data(); + upload_size = transformed.size(); + } + return status; + } + status.log("unknown host weight layout %s", source.layout.c_str()); + return status; +} + +static Status allocate_device_buffer(hrx_device_t device, size_t size, hrx_buffer_t & buffer) { + Status status; + if (device == nullptr) { + status.log("missing HRX device for host memory allocation"); + return status; + } + if (size == 0) { + status.log("cannot allocate an empty host memory device buffer"); + return status; + } + hrx_buffer_params_t params = { + HRX_MEMORY_TYPE_DEVICE_LOCAL, + HRX_MEMORY_ACCESS_ALL, + HRX_BUFFER_USAGE_DEFAULT, + 0, + }; + if (ErrorResult error = + take_status(hrx_allocator_allocate_buffer(hrx_device_allocator(device), params, size, &buffer))) { + status.log("allocate host memory device buffer: %s", error->c_str()); + } + return status; +} + +} // namespace + +Status allocate_mapped_host_staging_buffer(hrx_device_t device, size_t size, hrx_buffer_t & buffer, void *& host_data) { + Status status; + buffer = nullptr; + host_data = nullptr; + if (device == nullptr) { + status.log("missing HRX device for host staging allocation"); + return status; + } + if (size == 0) { + status.log("cannot allocate an empty host staging buffer"); + return status; + } + hrx_buffer_params_t params = { + HRX_MEMORY_TYPE_HOST_LOCAL | HRX_MEMORY_TYPE_DEVICE_VISIBLE, + HRX_MEMORY_ACCESS_ALL, + HRX_BUFFER_USAGE_DEFAULT | HRX_BUFFER_USAGE_MAPPING_SCOPED | HRX_BUFFER_USAGE_MAPPING_PERSISTENT, + 0, + }; + if (ErrorResult error = + take_status(hrx_allocator_allocate_buffer(hrx_device_allocator(device), params, size, &buffer))) { + status.log("allocate mapped HRX host staging buffer: %s", error->c_str()); + return status; + } + if (ErrorResult error = take_status(hrx_buffer_map(buffer, HRX_MAP_READ | HRX_MAP_WRITE, 0, size, &host_data))) { + status.log("map HRX host staging buffer: %s", error->c_str()); + hrx_buffer_release(buffer); + buffer = nullptr; + } + return status; +} + +Status HostTransferManager::upload_synchronous(hrx_stream_t stream, + const void * host_source, + hrx_buffer_t destination, + size_t offset, + size_t size) { + Status status; + if (size == 0) { + return status; + } + if (stream == nullptr || host_source == nullptr || destination == nullptr) { + status.log("invalid HRX host upload"); + return status; + } + hrx_device_t device = nullptr; + if (ErrorResult error = take_status(hrx_stream_get_device(stream, &device))) { + status.log("query HRX upload device failed: %s", error->c_str()); + return status; + } + if (ErrorResult error = take_status(hrx_stream_synchronize(stream))) { + status.log("synchronize before HRX host upload failed: %s", error->c_str()); + return status; + } + if (ErrorResult error = take_status(hrx_synchronous_h2d(device, host_source, destination, offset, size))) { + status.log("synchronous HRX host upload failed: %s", error->c_str()); + return status; + } + std::lock_guard lock(mutex_); + ++stats_.uploads; + stats_.upload_bytes += size; + return status; +} + +Status HostTransferManager::upload_async(hrx_stream_t stream, + const void * host_source, + hrx_buffer_t destination, + size_t offset, + size_t size) { + Status status; + if (size == 0) { + return status; + } + if (stream == nullptr || host_source == nullptr || destination == nullptr) { + status.log("invalid HRX host upload"); + return status; + } + + // Inline stream updates are intended for small payloads. Use bulk H2D for + // large host-staged tensors to avoid hundreds of update commands. + if (size >= kLargeHostUploadBytes) { + return upload_synchronous(stream, host_source, destination, offset, size); + } + + const uint8_t * host_bytes = static_cast(host_source); + size_t uploaded = 0; + while (uploaded < size) { + const size_t remaining = size - uploaded; + const size_t chunk_size = remaining < kMaxInlineUploadBytes ? remaining : kMaxInlineUploadBytes; + if (ErrorResult error = take_status( + hrx_stream_update_buffer(stream, host_bytes + uploaded, chunk_size, destination, offset + uploaded))) { + status.log("HRX async host upload failed: %s", error->c_str()); + return status; + } + uploaded += chunk_size; + } + + std::lock_guard lock(mutex_); + ++stats_.uploads; + stats_.upload_bytes += size; + return status; +} + +Status HostTransferManager::download_synchronous(hrx_stream_t stream, + hrx_buffer_t source, + size_t offset, + void * host_destination, + size_t size) { + Status status; + if (size == 0) { + return status; + } + if (stream == nullptr || source == nullptr || host_destination == nullptr) { + status.log("invalid HRX host download"); + return status; + } + hrx_device_t device = nullptr; + if (ErrorResult error = take_status(hrx_stream_get_device(stream, &device))) { + status.log("query HRX download device failed: %s", error->c_str()); + return status; + } + if (ErrorResult error = take_status(hrx_stream_synchronize(stream))) { + status.log("synchronize before HRX host download failed: %s", error->c_str()); + return status; + } + if (ErrorResult error = take_status(hrx_synchronous_d2h(device, source, offset, host_destination, size))) { + status.log("synchronous HRX host download failed: %s", error->c_str()); + return status; + } + std::lock_guard lock(mutex_); + ++stats_.downloads; + stats_.download_bytes += size; + return status; +} + +HostTransferStats HostTransferManager::stats() const { + std::lock_guard lock(mutex_); + return stats_; +} + +void HostTransferManager::clear() { + std::lock_guard lock(mutex_); + stats_ = {}; +} + +struct HostWeightLease::Entry { + ~Entry() { + if (buffer != nullptr) { + hrx_buffer_release(buffer); + } + } + + hrx_buffer_t buffer = nullptr; + size_t length = 0; + std::string layout; +}; + +HostWeightLease::HostWeightLease(std::shared_ptr entry) : entry_(std::move(entry)) {} + +bool HostWeightLease::valid() const { + return entry_ != nullptr && entry_->buffer != nullptr; +} + +hrx_buffer_t HostWeightLease::buffer() const { + return valid() ? entry_->buffer : nullptr; +} + +size_t HostWeightLease::length() const { + return entry_ != nullptr ? entry_->length : 0; +} + +const std::string & HostWeightLease::layout() const { + static const std::string empty; + return entry_ != nullptr ? entry_->layout : empty; +} + +HostWeightCache::~HostWeightCache() { + clear(); +} + +size_t HostWeightCache::SourceKeyHash::operator()(const SourceKey & key) const { + uint64_t hash = UINT64_C(1469598103934665603); + auto mix = [&](uint64_t value) { + hash ^= value; + hash *= UINT64_C(1099511628211); + }; + mix(key.identity); + mix(key.generation); + mix(static_cast(key.capacity)); + mix(static_cast(key.offset)); + mix(static_cast(key.length)); + mix(static_cast(key.materialized_length)); + for (unsigned char byte : key.layout) { + mix(byte); + } + mix(static_cast(key.source_type)); + mix(static_cast(key.input_size)); + mix(static_cast(key.output_size)); + return static_cast(hash); +} + +HostWeightAcquireResult HostWeightCache::acquire(hrx_device_t device, + hrx_stream_t stream, + HostTransferManager & transfers, + const HostWeightSource & source) { + HostWeightAcquireResult result; + if (stream == nullptr) { + result.status.log("host weight residency requires an HRX stream"); + return result; + } + const bool has_host_source = source.host_data != nullptr; + const bool has_device_source = source.device_buffer != nullptr; + if (has_host_source == has_device_source || source.identity == 0 || source.generation == 0 || source.length == 0 || + source.offset > source.capacity || source.length > source.capacity - source.offset) { + result.status.log("invalid host weight source"); + return result; + } + if (source.layout.empty()) { + result.status.log("host weight source has no layout"); + return result; + } + + const size_t materialized_length = + source.layout == kNativeWeightLayout ? source.length : source.materialized_length; + if (materialized_length == 0) { + result.status.log("host weight source has an empty materialized layout"); + return result; + } + + const SourceKey key{ source.identity, source.generation, source.capacity, source.offset, + source.length, materialized_length, source.layout, source.source_type, + source.input_size, source.output_size }; + { + std::lock_guard lock(mutex_); + const auto found = entries_.find(key); + if (found != entries_.end()) { + ++stats_.hits; + result.lease = HostWeightLease(found->second); + return result; + } + } + + HostWeightSource materialization_source = source; + std::vector canonical; + if (has_device_source) { + if (source.layout == kNativeWeightLayout) { + result.status.log("native device weights do not require host materialization"); + return result; + } + canonical.resize(source.length); + result.status = transfers.download_synchronous(stream, source.device_buffer, source.offset, canonical.data(), + canonical.size()); + if (!result.status.success()) { + return result; + } + materialization_source.host_data = canonical.data(); + materialization_source.device_buffer = nullptr; + materialization_source.capacity = canonical.size(); + materialization_source.offset = 0; + } + + const void * upload_data = nullptr; + size_t upload_size = 0; + std::vector transformed; + result.status = materialize_weight(materialization_source, upload_data, upload_size, transformed); + if (!result.status.success()) { + return result; + } + if (upload_size != materialized_length) { + result.status.log("host weight layout %s produced %zu bytes, expected %zu", source.layout.c_str(), upload_size, + materialized_length); + return result; + } + + auto entry = std::make_shared(); + entry->length = materialized_length; + entry->layout = source.layout; + result.status = allocate_device_buffer(device, materialized_length, entry->buffer); + if (!result.status.success()) { + return result; + } + result.status = transfers.upload_synchronous(stream, upload_data, entry->buffer, 0, upload_size); + if (!result.status.success()) { + return result; + } + + { + std::lock_guard lock(mutex_); + const auto inserted = entries_.emplace(key, entry); + if (!inserted.second) { + ++stats_.hits; + result.lease = HostWeightLease(inserted.first->second); + return result; + } + ++stats_.misses; + stats_.allocation_count = entries_.size(); + stats_.resident_bytes += materialized_length; + } + result.lease = HostWeightLease(std::move(entry)); + return result; +} + +HostWeightCacheStats HostWeightCache::stats() const { + std::lock_guard lock(mutex_); + return stats_; +} + +void HostWeightCache::clear() { + std::lock_guard lock(mutex_); + entries_.clear(); + stats_ = {}; +} + +HostStagingBuffer::~HostStagingBuffer() { + clear(); +} + +HostStagingBuffer::HostStagingBuffer(HostStagingBuffer && other) noexcept { + *this = std::move(other); +} + +HostStagingBuffer & HostStagingBuffer::operator=(HostStagingBuffer && other) noexcept { + if (this == &other) { + return *this; + } + clear(); + buffer = other.buffer; + host_data = other.host_data; + source_host_buffer = std::move(other.source_host_buffer); + value = other.value; + length = other.length; + upload = other.upload; + download = other.download; + other.buffer = nullptr; + other.host_data = nullptr; + other.value = -1; + other.length = 0; + other.upload = false; + other.download = false; + return *this; +} + +void HostStagingBuffer::clear() { + if (buffer != nullptr) { + hrx_buffer_release(buffer); + buffer = nullptr; + } + host_data = nullptr; + source_host_buffer = HostBufferRef{}; + value = -1; + length = 0; + upload = false; + download = false; +} + +Status allocate_host_staging_buffer(hrx_device_t device, size_t size, HostStagingBuffer & staging) { + staging.clear(); + Status status = allocate_device_buffer(device, size, staging.buffer); + if (status.success()) { + staging.length = size; + } + return status; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/host-memory.h b/ggml/src/ggml-hrx/runtime/host-memory.h new file mode 100644 index 000000000000..cd6db3589a6f --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/host-memory.h @@ -0,0 +1,173 @@ +#pragma once + +#include "dispatch/dispatch.h" +#include "ggml.h" +#include "runtime/host-buffer-registry.h" +#include "status.h" + +#include +#include +#include +#include +#include +#include + +typedef struct hrx_buffer_s * hrx_buffer_t; +typedef struct hrx_device_s * hrx_device_t; +typedef struct hrx_stream_s * hrx_stream_t; + +namespace ggml::hrx { + +struct HostTransferStats { + uint64_t uploads = 0; + uint64_t downloads = 0; + size_t upload_bytes = 0; + size_t download_bytes = 0; +}; + +class HostTransferManager { + public: + Status upload_synchronous(hrx_stream_t stream, + const void * host_source, + hrx_buffer_t destination, + size_t offset, + size_t size); + Status upload_async(hrx_stream_t stream, + const void * host_source, + hrx_buffer_t destination, + size_t offset, + size_t size); + Status download_synchronous(hrx_stream_t stream, + hrx_buffer_t source, + size_t offset, + void * host_destination, + size_t size); + + HostTransferStats stats() const; + void clear(); + + private: + mutable std::mutex mutex_; + HostTransferStats stats_; +}; + +struct HostWeightSource { + const void * host_data = nullptr; + hrx_buffer_t device_buffer = nullptr; + uint64_t identity = 0; + uint64_t generation = 0; + size_t capacity = 0; + size_t offset = 0; + size_t length = 0; + size_t materialized_length = 0; + std::string layout = kNativeWeightLayout; + ggml_type source_type = GGML_TYPE_COUNT; + int64_t input_size = 0; + int64_t output_size = 0; +}; + +struct HostWeightCacheStats { + uint64_t hits = 0; + uint64_t misses = 0; + size_t allocation_count = 0; + size_t resident_bytes = 0; +}; + +class HostWeightLease { + public: + HostWeightLease() = default; + + bool valid() const; + hrx_buffer_t buffer() const; + size_t length() const; + const std::string & layout() const; + + private: + struct Entry; + std::shared_ptr entry_; + + explicit HostWeightLease(std::shared_ptr entry); + friend class HostWeightCache; +}; + +struct HostWeightAcquireResult { + HostWeightLease lease; + Status status; + + bool valid() const { return lease.valid() && status.success(); } +}; + +class HostWeightCache { + public: + HostWeightCache() = default; + ~HostWeightCache(); + + HostWeightCache(const HostWeightCache &) = delete; + HostWeightCache & operator=(const HostWeightCache &) = delete; + + HostWeightAcquireResult acquire(hrx_device_t device, + hrx_stream_t stream, + HostTransferManager & transfers, + const HostWeightSource & source); + HostWeightCacheStats stats() const; + void clear(); + + private: + struct SourceKey { + uint64_t identity = 0; + uint64_t generation = 0; + size_t capacity = 0; + size_t offset = 0; + size_t length = 0; + size_t materialized_length = 0; + std::string layout; + ggml_type source_type = GGML_TYPE_COUNT; + int64_t input_size = 0; + int64_t output_size = 0; + + bool operator==(const SourceKey & other) const { + return identity == other.identity && generation == other.generation && capacity == other.capacity && + offset == other.offset && length == other.length && + materialized_length == other.materialized_length && layout == other.layout && + source_type == other.source_type && input_size == other.input_size && + output_size == other.output_size; + } + }; + + struct SourceKeyHash { + size_t operator()(const SourceKey & key) const; + }; + + mutable std::mutex mutex_; + std::unordered_map, SourceKeyHash> entries_; + HostWeightCacheStats stats_; +}; + +struct HostStagingBuffer { + HostStagingBuffer() = default; + ~HostStagingBuffer(); + + HostStagingBuffer(HostStagingBuffer && other) noexcept; + HostStagingBuffer & operator=(HostStagingBuffer && other) noexcept; + + HostStagingBuffer(const HostStagingBuffer &) = delete; + HostStagingBuffer & operator=(const HostStagingBuffer &) = delete; + + hrx_buffer_t buffer = nullptr; + void * host_data = nullptr; + HostBufferRef source_host_buffer; + int32_t value = -1; + size_t length = 0; + bool upload = false; + bool download = false; + + void clear(); +}; + +Status allocate_host_staging_buffer(hrx_device_t device, size_t size, HostStagingBuffer & staging); +Status allocate_mapped_host_staging_buffer(hrx_device_t device, + size_t size, + hrx_buffer_t & buffer, + void *& host_data); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/hrx-sleeping-wait.cpp b/ggml/src/ggml-hrx/runtime/hrx-sleeping-wait.cpp new file mode 100644 index 000000000000..f7eff52df06d --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/hrx-sleeping-wait.cpp @@ -0,0 +1,137 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// A blocking hrx_stream_wait goes to ROCr's hsa_signal_wait_scacquire on the stream timeline. The +// timeline advances with every completed command (about 900 per decode token of a 27B model), and +// each advance returns the wait from the KFD event ioctl without sleeping, so the waiting thread +// runs a full CPU core (measured: 35% user + 65% system, 30 context switches in 6 s) for the whole +// token. On a dense 27B that is ~10 W of package power and the difference between a 74 C and a +// 90-93 C APU during decode. +// +// When a wait is expected to be long (the shortest of the last four waits at this call site is over +// 20 ms, so one long prefill wait does not make the next decode token oversleep), sleep through 80% +// of it and leave the rest to the normal wait. Shorter waits are left as they were: sleeping +// through 3-10 ms tokens cost small models 2-5% of decode (the host work between tokens then runs +// on a core that has clocked down), and they were never the thermal problem. +// +// The estimate must not feed on its own sleep. A wait that is still asleep when the work finishes +// records the sleep, not the work, so the next estimate was 80% of a stale one: after a 512-token +// prompt graph (98-123 ms waits on Qwen3-4B) the 13 ms decode tokens slept 98, 78, 63, 50, ... ms, +// an 0.8x decay over ~10 tokens (measured 2026-10-04). Two guards: a wait that finds the work +// already done when it wakes forgets the history instead of recording the sleep (the next wait +// blocks and measures the work), and the graph replay wait keeps one history per replayed graph +// (wait_history_for), so prompt graph waits never set the sleep of a decode graph. +// +// ONEBIT_HRX_BLOCKING_WAIT=1 turns the sleep off. + +#include "hrx-sleeping-wait.h" + +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { + +namespace { + +using Clock = std::chrono::steady_clock; + +constexpr int64_t kSleepAboveNs = 20000000; + +bool blocking_wait() { + static const bool blocking = [] { + const char * value = std::getenv("ONEBIT_HRX_BLOCKING_WAIT"); + return value != nullptr && value[0] == '1'; + }(); + return blocking; +} + +void sleep_ns(int64_t ns) { + // Timer slack defaults to 50 us per thread; keep the wake-up close to the budget. + thread_local bool slack_set = [] { + prctl(PR_SET_TIMERSLACK, 1000UL, 0, 0, 0); + return true; + }(); + (void) slack_set; + const timespec interval = { static_cast(ns / 1000000000), static_cast(ns % 1000000000) }; + nanosleep(&interval, nullptr); +} + +int64_t expected_ns(const WaitHistory & history) { + if (history.count == 0) { + return 0; + } + const uint32_t n = std::min(history.count, history.duration_ns.size()); + return *std::min_element(history.duration_ns.begin(), history.duration_ns.begin() + n); +} + +} // namespace + +hrx_status_t stream_wait_sleeping(hrx_stream_t stream, WaitHistory & history) { + if (blocking_wait() || stream == nullptr) { + return hrx_stream_wait(stream); + } + const auto start = Clock::now(); + const int64_t expected = expected_ns(history); + bool overslept = false; + if (expected > kSleepAboveNs) { + bool complete = false; + hrx_status_t status = hrx_stream_query(stream, &complete); + if (!hrx_status_is_ok(status)) { + return status; + } + if (!complete) { + sleep_ns(expected * 4 / 5); + status = hrx_stream_query(stream, &complete); + if (!hrx_status_is_ok(status)) { + return status; + } + overslept = complete; + } + } + hrx_status_t status = hrx_stream_wait(stream); + if (hrx_status_is_ok(status)) { + if (overslept) { + history.count = 0; + } else { + history.duration_ns[history.count % history.duration_ns.size()] = + std::chrono::duration_cast(Clock::now() - start).count(); + history.count++; + } + } + return status; +} + +WaitHistory & wait_history_for(const void * key) { + // Bounded: a replayed graph is re-recorded (new key) when the transient arena moves. + thread_local std::unordered_map histories; + if (histories.size() >= 256 && histories.find(key) == histories.end()) { + histories.clear(); + } + return histories[key]; +} + +hrx_status_t stream_synchronize_sleeping(hrx_stream_t stream, WaitHistory & history) { + hrx_status_t status = hrx_stream_flush(stream); + if (!hrx_status_is_ok(status)) { + return status; + } + return stream_wait_sleeping(stream, history); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/hrx-sleeping-wait.h b/ggml/src/ggml-hrx/runtime/hrx-sleeping-wait.h new file mode 100644 index 000000000000..fe26969cb52c --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/hrx-sleeping-wait.h @@ -0,0 +1,42 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include "hrx_runtime.h" + +#include +#include + +namespace ggml::hrx { + +// Durations of the last few waits at one call site; see hrx-sleeping-wait.cpp. +struct WaitHistory { + std::array duration_ns{}; + uint32_t count = 0; +}; + +// hrx_stream_wait that sleeps through most of a long expected wait instead of spending it in the +// HSA signal wait, which keeps a CPU core busy. See hrx-sleeping-wait.cpp. +hrx_status_t stream_wait_sleeping(hrx_stream_t stream, WaitHistory & history); + +// This thread's WaitHistory for the work identified by `key` (a replayed graph), so waits on +// different work do not share one estimate. +WaitHistory & wait_history_for(const void * key); + +// hrx_stream_synchronize (flush, then wait) with the sleeping wait. +hrx_status_t stream_synchronize_sleeping(hrx_stream_t stream, WaitHistory & history); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/kernel-executable-cache.cpp b/ggml/src/ggml-hrx/runtime/kernel-executable-cache.cpp new file mode 100644 index 000000000000..a11db6154882 --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/kernel-executable-cache.cpp @@ -0,0 +1,419 @@ +#include "kernel-executable-cache.h" + +#include "ggml-impl.h" +#include "hrx-interop-utils.h" +#include "hip/hip-kernel-loader.h" +#include "hip/hip-kernel-registry.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static ggml_hrx_loom_jit_source_format to_jit_source_format(KernelSourceFormat format) { + switch (format) { + case KERNEL_SOURCE_FORMAT_TEXT: + return GGML_HRX_LOOM_JIT_SOURCE_FORMAT_TEXT; + case KERNEL_SOURCE_FORMAT_BINARY: + return GGML_HRX_LOOM_JIT_SOURCE_FORMAT_BYTECODE; + } + return GGML_HRX_LOOM_JIT_SOURCE_FORMAT_TEXT; +} + +static void append_u32(std::vector & bytes, uint32_t value) { + const size_t offset = bytes.size(); + bytes.resize(offset + sizeof(value)); + std::memcpy(bytes.data() + offset, &value, sizeof(value)); +} + +static bool pack_kernel_constants(const KernelDefinition & definition, + const Dispatch & dispatch, + std::vector & constants) { + constants.clear(); + for (const KernelScalarDefinition & parameter : definition.launch_parameters) { + const char * name = parameter.name != nullptr ? parameter.name : ""; + const char * type = parameter.type != nullptr ? parameter.type : ""; + const auto item = dispatch.kernel.integer_parameters.find(name); + if (item == dispatch.kernel.integer_parameters.end() || std::strcmp(type, "index") != 0 || item->second < 0 || + static_cast(item->second) > std::numeric_limits::max()) { + constants.clear(); + GGML_LOG_ERROR("%s: invalid launch scalar %s for %s\n", __func__, name, + kernel_definition_name(definition).c_str()); + return false; + } + append_u32(constants, static_cast(item->second)); + } + return true; +} + +static std::string kernel_executable_key(const KernelDefinition & definition, + const Dispatch & dispatch, + const char * target) { + std::ostringstream out; + out << (target != nullptr ? target : "") << '|' << definition.source_digest << '|' << definition.symbol + << "|recipe=" << definition.compile_recipe.mode; + for (const KernelScalarDefinition & parameter : definition.workload_parameters) { + const char * name = parameter.name != nullptr ? parameter.name : ""; + const auto item = dispatch.kernel.integer_parameters.find(name); + out << '|' << name << '='; + if (item == dispatch.kernel.integer_parameters.end()) { + out << ""; + } else { + out << item->second; + } + } + for (const KernelCompileConfig & config : definition.compile_config) { + out << '|' << (config.key != nullptr ? config.key : "") << '=' << (config.value != nullptr ? config.value : ""); + } + for (const auto & config : dispatch.kernel.compile_parameters) { + out << '|' << config.first << '=' << config.second; + } + return out.str(); +} + +static bool build_compile_request(const KernelDefinition & definition, + const Dispatch & dispatch, + LoomKernelCompileRequest & request) { + if (definition.compile_recipe.primary_sources.empty()) { + GGML_LOG_ERROR("%s: kernel %s has no primary source\n", __func__, kernel_definition_name(definition).c_str()); + return false; + } + const KernelSourceRef & primary_source = definition.compile_recipe.primary_sources.front(); + const KernelSource * source = primary_source.contents; + if (source == nullptr) { + GGML_LOG_ERROR("%s: missing embedded source for %s\n", __func__, primary_source.path); + return false; + } + + request.source_data = source->source.data; + request.source_size = source->source.length; + request.source_format = to_jit_source_format(source->source.format); + request.source_identifier = primary_source.path != nullptr ? primary_source.path : ""; + request.symbol = definition.symbol != nullptr ? definition.symbol : ""; + request.launch_config_symbol = definition.name != nullptr ? definition.name : ""; + + request.dependencies.reserve(definition.compile_recipe.library_sources.size()); + for (const KernelSourceRef & dependency_ref : definition.compile_recipe.library_sources) { + const KernelSource * dependency = dependency_ref.contents; + if (dependency == nullptr) { + GGML_LOG_ERROR("%s: missing embedded dependency for %s\n", __func__, dependency_ref.path); + return false; + } + request.dependencies.push_back({ + dependency->source.data, + dependency->source.length, + to_jit_source_format(dependency->source.format), + dependency_ref.path, + }); + } + + std::map merged_configs; + for (const KernelCompileConfig & config : definition.compile_config) { + merged_configs[config.key != nullptr ? config.key : ""] = config.value != nullptr ? config.value : ""; + } + for (const auto & config : dispatch.kernel.compile_parameters) { + merged_configs[config.first] = config.second; + } + request.config_storage.reserve(merged_configs.size()); + for (const auto & config : merged_configs) { + request.config_storage.push_back(config); + } + + request.workload.reserve(definition.workload_parameters.size()); + for (const KernelScalarDefinition & parameter : definition.workload_parameters) { + const char * name = parameter.name != nullptr ? parameter.name : ""; + const char * type = parameter.type != nullptr ? parameter.type : ""; + const auto item = dispatch.kernel.integer_parameters.find(name); + if (item == dispatch.kernel.integer_parameters.end() || std::strcmp(type, "index") != 0) { + GGML_LOG_ERROR("%s: invalid workload scalar %s for %s\n", __func__, name, + kernel_definition_name(definition).c_str()); + return false; + } + request.workload.push_back(item->second); + } + return true; +} + +static std::shared_ptr load_kernel_executable(const KernelExecutablePrepareContext & context, + const KernelDefinition & definition, + const Dispatch & dispatch, + const std::vector & constants, + const std::string & key, + ggml_hrx_loom_jit_compile_result & compiled, + std::string & error_message) { + if (context.device == nullptr) { + error_message = "missing HRX device"; + GGML_LOG_ERROR("%s: load %s: %s\n", __func__, key.c_str(), error_message.c_str()); + return nullptr; + } + if (context.target == nullptr) { + error_message = "missing HRX target"; + GGML_LOG_ERROR("%s: load %s: %s\n", __func__, key.c_str(), error_message.c_str()); + return nullptr; + } + + auto executable = std::make_shared(); + executable->launch = compiled.launch_config; + if (ErrorResult error = + take_status(hrx_executable_load_data(context.device, compiled.hsaco_data, compiled.hsaco_size, "amdgpu", + context.target, &executable->executable))) { + error_message = "load " + key + ": " + *error; + GGML_LOG_ERROR("%s: %s\n", __func__, error_message.c_str()); + return nullptr; + } + if (ErrorResult error = take_status(hrx_executable_lookup_export_by_name(executable->executable, definition.name, + &executable->export_ordinal))) { + error_message = "lookup " + key + ": " + *error; + GGML_LOG_ERROR("%s: %s\n", __func__, error_message.c_str()); + return nullptr; + } + if (ErrorResult error = take_status( + hrx_executable_export_info(executable->executable, executable->export_ordinal, &executable->export_info))) { + error_message = "inspect " + key + ": " + *error; + GGML_LOG_ERROR("%s: %s\n", __func__, error_message.c_str()); + return nullptr; + } + if (executable->export_info.binding_count != dispatch.bindings.size() || + executable->export_info.constant_byte_length != constants.size() || + executable->export_info.parameter_count != dispatch.bindings.size() + definition.launch_parameters.size()) { + error_message = "compiled ABI does not match manifest for " + key; + GGML_LOG_ERROR("%s: %s\n", __func__, error_message.c_str()); + return nullptr; + } + if (executable->launch.workgroup_count[0] == 0 || executable->launch.workgroup_size[0] == 0) { + error_message = "compiled launch geometry is empty for " + key; + GGML_LOG_ERROR("%s: %s\n", __func__, error_message.c_str()); + return nullptr; + } + return executable; +} + +} // namespace + +class KernelExecutableCacheEntry { + public: + KernelExecutableCacheEntry(std::string key, + const KernelDefinition & definition, + const Dispatch & dispatch, + LoomCompiledKernelRef compiled_ref) : + key(std::move(key)), + definition(&definition), + dispatch(dispatch), + compiled_ref(std::move(compiled_ref)) {} + + std::mutex mutex; + std::condition_variable complete; + std::string key; + const KernelDefinition * definition = nullptr; + Dispatch dispatch; + LoomCompiledKernelRef compiled_ref; + std::atomic> executable; + std::string error; + + enum class LoadState { + Unloaded, + Loading, + Loaded, + Failed, + }; + + std::atomic load_state = LoadState::Unloaded; +}; + +KernelExecutable::~KernelExecutable() { + if (executable != nullptr) { + hrx_executable_release(executable); + } +} + +KernelExecutableCache::KernelExecutableCache(LoomJitMode mode) : mode_(mode), mode_is_forced_(true) {} + +KernelExecutableCache::~KernelExecutableCache() { + clear(); +} + +bool KernelExecutableCache::ensure_jit_locked(const char * target, std::string & error_message) { + if (target == nullptr || target[0] == '\0') { + error_message = "missing HRX target"; + GGML_LOG_ERROR("%s: %s\n", __func__, error_message.c_str()); + return false; + } + if (jit_ != nullptr) { + if (target_ != target) { + error_message = + "HRX kernel executable cache target mismatch: existing " + target_ + ", requested " + target; + GGML_LOG_ERROR("%s: %s\n", __func__, error_message.c_str()); + return false; + } + return true; + } + + jit_ = mode_is_forced_ ? create_loom_jit(target, mode_, error_message) : create_loom_jit(target, error_message); + if (jit_ == nullptr) { + if (error_message.empty()) { + error_message = "create Loom JIT failed"; + } + return false; + } + target_ = target; + return true; +} + +KernelExecutableRef KernelExecutableCache::get_or_compile(const KernelExecutablePrepareContext & context, + const KernelDefinition & definition, + const Dispatch & dispatch, + std::vector & constants) { + KernelExecutableRef ref; + if (!pack_kernel_constants(definition, dispatch, constants)) { + return ref; + } + + const std::string key = kernel_executable_key(definition, dispatch, context.target); + LoomKernelCompileRequest request; + LoomCompiledKernelRef compiled_ref; + { + std::lock_guard lock(mutex_); + const auto found = cache_.find(key); + if (found != cache_.end()) { + ref.entry = found->second; + return ref; + } + if (is_hip_kernel_definition(definition)) { + auto entry = std::make_shared( + key, definition, dispatch, make_hip_compiled_kernel(key, definition, dispatch, context.target)); + cache_.emplace(key, entry); + ref.entry = std::move(entry); + return ref; + } + std::string error_message; + if (!ensure_jit_locked(context.target, error_message)) { + return ref; + } + if (!build_compile_request(definition, dispatch, request)) { + return ref; + } + + compiled_ref = jit_->compile(key, std::move(request)); + auto entry = std::make_shared(key, definition, dispatch, compiled_ref); + cache_.emplace(key, entry); + ref.entry = std::move(entry); + } + return ref; +} + +std::shared_ptr KernelExecutableCache::materialize(const KernelExecutablePrepareContext & context, + const KernelExecutableRef & ref, + const std::vector & constants) { + if (!ref.valid()) { + return nullptr; + } + + KernelExecutableCacheEntry & entry = *ref.entry; + KernelExecutableCacheEntry::LoadState load_state = entry.load_state.load(std::memory_order_acquire); + if (load_state == KernelExecutableCacheEntry::LoadState::Loaded) { + return entry.executable.load(std::memory_order_acquire); + } + if (load_state == KernelExecutableCacheEntry::LoadState::Failed) { + std::lock_guard entry_lock(entry.mutex); + GGML_LOG_ERROR("%s: %s\n", __func__, entry.error.c_str()); + return nullptr; + } + + KernelExecutableCacheEntry::LoadState expected = KernelExecutableCacheEntry::LoadState::Unloaded; + if (!entry.load_state.compare_exchange_strong(expected, KernelExecutableCacheEntry::LoadState::Loading, + std::memory_order_acq_rel, std::memory_order_acquire)) { + std::unique_lock entry_lock(entry.mutex); + entry.complete.wait(entry_lock, [&] { + const KernelExecutableCacheEntry::LoadState current = entry.load_state.load(std::memory_order_acquire); + return current == KernelExecutableCacheEntry::LoadState::Loaded || + current == KernelExecutableCacheEntry::LoadState::Failed; + }); + if (entry.load_state.load(std::memory_order_acquire) == KernelExecutableCacheEntry::LoadState::Failed) { + GGML_LOG_ERROR("%s: %s\n", __func__, entry.error.c_str()); + return nullptr; + } + return entry.executable.load(std::memory_order_acquire); + } + + LoomCompiledKernelRef compiled_ref; + { + std::lock_guard entry_lock(entry.mutex); + compiled_ref = entry.compiled_ref; + } + + std::string error_message; + if (compiled_ref == nullptr || !compiled_ref->resolve()) { + error_message = compiled_ref ? compiled_ref->error_message() : "missing compiled kernel"; + GGML_LOG_ERROR("%s: %s\n", __func__, error_message.c_str()); + + { + std::lock_guard entry_lock(entry.mutex); + entry.error = error_message; + entry.load_state.store(KernelExecutableCacheEntry::LoadState::Failed, std::memory_order_release); + } + entry.complete.notify_all(); + + std::lock_guard lock(mutex_); + const auto found = cache_.find(entry.key); + if (found != cache_.end() && found->second == ref.entry) { + cache_.erase(found); + } + return nullptr; + } + + ggml_hrx_loom_jit_compile_result compiled = compiled_ref->take_result(); + std::shared_ptr executable = load_kernel_executable( + context, *entry.definition, entry.dispatch, constants, entry.key, compiled, error_message); + + if (executable != nullptr) { + compiled.reset(); + { + std::lock_guard entry_lock(entry.mutex); + entry.compiled_ref.reset(); + entry.executable.store(executable, std::memory_order_release); + entry.load_state.store(KernelExecutableCacheEntry::LoadState::Loaded, std::memory_order_release); + } + } else { + { + std::lock_guard entry_lock(entry.mutex); + entry.error = std::move(error_message); + entry.load_state.store(KernelExecutableCacheEntry::LoadState::Failed, std::memory_order_release); + } + } + entry.complete.notify_all(); + + if (executable == nullptr) { + std::lock_guard lock(mutex_); + const auto found = cache_.find(entry.key); + if (found != cache_.end() && found->second == ref.entry) { + cache_.erase(found); + } + } + return executable; +} + +std::shared_ptr KernelExecutableCache::prepare(const KernelExecutablePrepareContext & context, + const KernelDefinition & definition, + const Dispatch & dispatch, + std::vector & constants) { + const KernelExecutableRef ref = get_or_compile(context, definition, dispatch, constants); + return materialize(context, ref, constants); +} + +void KernelExecutableCache::clear() { + if (jit_ != nullptr) { + jit_->clear(); + } + std::lock_guard lock(mutex_); + cache_.clear(); + jit_.reset(); + target_.clear(); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/kernel-executable-cache.h b/ggml/src/ggml-hrx/runtime/kernel-executable-cache.h new file mode 100644 index 000000000000..2ae130d9918e --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/kernel-executable-cache.h @@ -0,0 +1,72 @@ +#pragma once + +#include "dispatch/dispatch.h" +#include "hrx_runtime.h" +#include "kernel-corpus/kernel-corpus.h" +#include "runtime/loom-kernel-jit.h" + +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { + +class KernelExecutableCacheEntry; + +struct KernelExecutable { + ~KernelExecutable(); + + hrx_executable_t executable = nullptr; + uint32_t export_ordinal = 0; + hrx_executable_export_info_t export_info = {}; + ggml_hrx_loom_jit_launch_config launch; +}; + +struct KernelExecutablePrepareContext { + hrx_device_t device = nullptr; + const char * target = nullptr; +}; + +struct KernelExecutableRef { + std::shared_ptr entry; + + bool valid() const { return entry != nullptr; } +}; + +class KernelExecutableCache { + public: + KernelExecutableCache() = default; + explicit KernelExecutableCache(LoomJitMode mode); + ~KernelExecutableCache(); + + KernelExecutableRef get_or_compile(const KernelExecutablePrepareContext & context, + const KernelDefinition & definition, + const Dispatch & dispatch, + std::vector & constants); + + std::shared_ptr materialize(const KernelExecutablePrepareContext & context, + const KernelExecutableRef & ref, + const std::vector & constants); + + std::shared_ptr prepare(const KernelExecutablePrepareContext & context, + const KernelDefinition & definition, + const Dispatch & dispatch, + std::vector & constants); + + void clear(); + + private: + bool ensure_jit_locked(const char * target, std::string & error_message); + + std::mutex mutex_; + std::unordered_map> cache_; + std::unique_ptr jit_; + std::string target_; + LoomJitMode mode_ = LoomJitMode::Async; + bool mode_is_forced_ = false; +}; + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/loom-jit-disk-cache.cpp b/ggml/src/ggml-hrx/runtime/loom-jit-disk-cache.cpp new file mode 100644 index 000000000000..22355551ee4a --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/loom-jit-disk-cache.cpp @@ -0,0 +1,244 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "loom-jit-disk-cache.h" + +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { + +namespace { + +constexpr uint32_t kMagic = 0x434a4231; // "1BJC" +constexpr uint32_t kVersion = 1; + +std::mutex g_mutex; +std::string g_target; + +// Two FNV-1a lanes with different offset bases give a 128-bit key. +struct KeyHash { + uint64_t a = 0xcbf29ce484222325ull; + uint64_t b = 0x84222325cbf29ce4ull; + + void bytes(const void * data, size_t size) { + const auto * p = static_cast(data); + for (size_t i = 0; i < size; ++i) { + a = (a ^ p[i]) * 0x100000001b3ull; + b = (b ^ p[i]) * 0x100000001b3ull; + b ^= b >> 29; + } + } + void field(const void * data, size_t size) { // length-prefixed, so fields cannot run together + const uint64_t n = size; + bytes(&n, sizeof n); + bytes(data, size); + } + void str(const std::string & s) { field(s.data(), s.size()); } + std::string hex() const { + char out[33]; + std::snprintf(out, sizeof out, "%016llx%016llx", (unsigned long long) a, (unsigned long long) b); + return out; + } +}; + +bool enabled() { + static const bool on = [] { + const char * v = std::getenv("GGML_HRX_JIT_CACHE"); + if (v != nullptr && (std::strcmp(v, "0") == 0 || std::strcmp(v, "off") == 0)) { + return false; + } + const char * s = std::getenv("GGML_HRX_LOOM_SANITIZER"); + return s == nullptr || s[0] == '\0'; + }(); + return on; +} + +// The library holding this code (and the statically linked Loom compiler): path, size, mtime. +const std::string & library_identity() { + static const std::string id = [] { + Dl_info info = {}; + struct stat st = {}; + if (dladdr(reinterpret_cast(&loom_jit_disk_cache_set_target), &info) == 0 || + info.dli_fname == nullptr || stat(info.dli_fname, &st) != 0) { + return std::string(); + } + return std::string(info.dli_fname) + ":" + std::to_string(st.st_size) + ":" + std::to_string(st.st_mtime); + }(); + return id; +} + +const std::filesystem::path & cache_dir() { + static const std::filesystem::path dir = [] { + if (const char * d = std::getenv("GGML_HRX_JIT_CACHE_DIR"); d != nullptr && d[0] != '\0') { + return std::filesystem::path(d); + } + const char * home = std::getenv("HOME"); + return home ? std::filesystem::path(home) / ".cache" / "1bit" / "hrx-jit" : std::filesystem::path(); + }(); + return dir; +} + +std::string key_for(const LoomKernelCompileRequest & request) { + KeyHash h; + const uint32_t version = kVersion; + h.field(&version, sizeof version); + h.str(library_identity()); + { + std::lock_guard lock(g_mutex); + h.str(g_target); + } + h.field(request.source_data, request.source_size); + h.field(&request.source_format, sizeof request.source_format); + h.str(request.source_identifier); + h.str(request.symbol); + h.str(request.launch_config_symbol); + for (const auto & dep : request.dependencies) { + h.field(dep.source_data, dep.source_size); + h.field(&dep.source_format, sizeof dep.source_format); + h.str(dep.source_identifier ? dep.source_identifier : ""); + } + std::vector> configs = request.config_storage; + std::sort(configs.begin(), configs.end()); + for (const auto & c : configs) { + h.str(c.first); + h.str(c.second); + } + h.field(request.workload.data(), request.workload.size() * sizeof(int64_t)); + return h.hex(); +} + +bool read_exact(FILE * f, void * out, size_t n) { return n == 0 || std::fread(out, 1, n, f) == n; } + +bool read_blob(FILE * f, void ** out, size_t * out_size, bool nul_terminate) { + uint64_t n = 0; + if (!read_exact(f, &n, sizeof n) || n > (1ull << 30)) { + return false; + } + if (n == 0) { + *out = nullptr; + *out_size = 0; + return true; + } + void * p = nullptr; + if (!hrx_status_is_ok(hrx_host_allocator_malloc(hrx_host_allocator_system(), n + (nul_terminate ? 1 : 0), &p))) { + return false; + } + if (!read_exact(f, p, n)) { + hrx_host_allocator_free(hrx_host_allocator_system(), p); + return false; + } + if (nul_terminate) { + static_cast(p)[n] = '\0'; + } + *out = p; + *out_size = n; + return true; +} + +} // namespace + +void loom_jit_disk_cache_set_target(const char * target) { + std::lock_guard lock(g_mutex); + g_target = target ? target : ""; +} + +bool loom_jit_disk_cache_load(const LoomKernelCompileRequest & request, ggml_hrx_loom_jit_compile_result & compiled) { + if (!enabled() || cache_dir().empty() || library_identity().empty()) { + return false; + } + const std::filesystem::path path = cache_dir() / (key_for(request) + ".bin"); + FILE * f = std::fopen(path.c_str(), "rb"); + if (f == nullptr) { + return false; + } + ggml_hrx_loom_jit_compile_result result; + uint32_t magic = 0, version = 0; + auto & lc = result.launch_config; + uint64_t workload_argument_count = 0; + bool ok = read_exact(f, &magic, sizeof magic) && read_exact(f, &version, sizeof version) && magic == kMagic && + version == kVersion && read_exact(f, lc.workgroup_count.data(), sizeof(uint32_t) * 3) && + read_exact(f, lc.workgroup_size.data(), sizeof(uint32_t) * 3) && + read_exact(f, &lc.subgroup_size, sizeof lc.subgroup_size) && + read_exact(f, &lc.workgroup_storage_bytes, sizeof lc.workgroup_storage_bytes) && + read_exact(f, &workload_argument_count, sizeof workload_argument_count) && + read_exact(f, &lc.fields, sizeof lc.fields); + if (ok) { + lc.workload_argument_count = static_cast(workload_argument_count); + size_t manifest_size = 0; + ok = read_blob(f, &result.hsaco_data, &result.hsaco_size, false) && + read_blob(f, reinterpret_cast(&result.manifest_json), &manifest_size, true) && + result.hsaco_size > 0; + result.manifest_json_size = manifest_size; + } + std::fclose(f); + if (!ok) { + return false; // a damaged entry is recompiled and overwritten + } + compiled = std::move(result); + return true; +} + +void loom_jit_disk_cache_store(const LoomKernelCompileRequest & request, + const ggml_hrx_loom_jit_compile_result & compiled) { + if (!enabled() || cache_dir().empty() || library_identity().empty() || compiled.hsaco_data == nullptr || + compiled.hsaco_size == 0) { + return; + } + std::error_code ec; + std::filesystem::create_directories(cache_dir(), ec); + const std::string key = key_for(request); + const std::filesystem::path path = cache_dir() / (key + ".bin"); + const std::filesystem::path tmp = + cache_dir() / (key + ".tmp." + std::to_string(reinterpret_cast(&compiled))); + FILE * f = std::fopen(tmp.c_str(), "wb"); + if (f == nullptr) { + return; + } + const auto & lc = compiled.launch_config; + const uint64_t workload_argument_count = lc.workload_argument_count; + const uint64_t hsaco_size = compiled.hsaco_size; + const uint64_t manifest_size = compiled.manifest_json ? compiled.manifest_json_size : 0; + bool ok = std::fwrite(&kMagic, sizeof kMagic, 1, f) == 1 && std::fwrite(&kVersion, sizeof kVersion, 1, f) == 1 && + std::fwrite(lc.workgroup_count.data(), sizeof(uint32_t), 3, f) == 3 && + std::fwrite(lc.workgroup_size.data(), sizeof(uint32_t), 3, f) == 3 && + std::fwrite(&lc.subgroup_size, sizeof lc.subgroup_size, 1, f) == 1 && + std::fwrite(&lc.workgroup_storage_bytes, sizeof lc.workgroup_storage_bytes, 1, f) == 1 && + std::fwrite(&workload_argument_count, sizeof workload_argument_count, 1, f) == 1 && + std::fwrite(&lc.fields, sizeof lc.fields, 1, f) == 1 && + std::fwrite(&hsaco_size, sizeof hsaco_size, 1, f) == 1 && + std::fwrite(compiled.hsaco_data, 1, compiled.hsaco_size, f) == compiled.hsaco_size && + std::fwrite(&manifest_size, sizeof manifest_size, 1, f) == 1 && + (manifest_size == 0 || std::fwrite(compiled.manifest_json, 1, manifest_size, f) == manifest_size); + ok = (std::fclose(f) == 0) && ok; + if (ok) { + std::filesystem::rename(tmp, path, ec); // atomic: readers see a whole entry or none + } + if (!ok || ec) { + std::filesystem::remove(tmp, ec); + } +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/loom-jit-disk-cache.h b/ggml/src/ggml-hrx/runtime/loom-jit-disk-cache.h new file mode 100644 index 000000000000..a438d3c99e5c --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/loom-jit-disk-cache.h @@ -0,0 +1,42 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Disk cache for Loom JIT results. HRX compiles a kernel the first time each specialization is +// dispatched (every new compile parameter or workload value), and the compiled code only lives in +// the process, so every server start pays the compiles again (~1 s for each new prompt-length +// remainder on ZAYA1-8B). This keeps each result under $GGML_HRX_JIT_CACHE_DIR (default +// ~/.cache/1bit/hrx-jit), keyed on every compile input plus the identity of the library that +// holds the compiler, so a rebuilt or updated HRX never reuses old code. +// +// GGML_HRX_JIT_CACHE=0 turns it off. It is also off while a Loom sanitizer is enabled. + +#pragma once + +#include "loom-jit.h" +#include "loom-kernel-jit.h" + +namespace ggml::hrx { + +// Records the JIT target (e.g. gfx1151); part of every key. +void loom_jit_disk_cache_set_target(const char * target); + +// Fills `compiled` from the cache; false when there is no usable entry. +bool loom_jit_disk_cache_load(const LoomKernelCompileRequest & request, ggml_hrx_loom_jit_compile_result & compiled); + +// Stores a successful compile; failures are ignored (the cache is only an accelerator). +void loom_jit_disk_cache_store(const LoomKernelCompileRequest & request, + const ggml_hrx_loom_jit_compile_result & compiled); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/loom-kernel-jit.cpp b/ggml/src/ggml-hrx/runtime/loom-kernel-jit.cpp new file mode 100644 index 000000000000..e6dda0b0b413 --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/loom-kernel-jit.cpp @@ -0,0 +1,316 @@ +#include "loom-kernel-jit.h" +#include "loom-jit-disk-cache.h" + +#include "ggml-impl.h" +#include "hrx-interop-utils.h" + +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +class LoomAmdgpuJit { + public: + LoomAmdgpuJit() = default; + LoomAmdgpuJit(const LoomAmdgpuJit &) = delete; + LoomAmdgpuJit & operator=(const LoomAmdgpuJit &) = delete; + + ~LoomAmdgpuJit() { reset(); } + + bool create(const char * target, std::string & error_message) { + reset(); + ggml_hrx_loom_jit_amdgpu_options options = {}; + options.processor = target; + loom_jit_disk_cache_set_target(target); // 1bit + options.identifier = target; + options.sanitizer = std::getenv("GGML_HRX_LOOM_SANITIZER"); + options.sanitizer_reporting = std::getenv("GGML_HRX_LOOM_SANITIZER_REPORTING"); + if (ErrorResult error = take_status(ggml_hrx_loom_jit_amdgpu_create(&options, &jit_))) { + error_message = "create Loom JIT: " + *error; + GGML_LOG_ERROR("%s: %s\n", __func__, error_message.c_str()); + return false; + } + return true; + } + + ggml_hrx_loom_jit_amdgpu * get() const { return jit_; } + + private: + void reset() { + if (jit_ != nullptr) { + ggml_hrx_loom_jit_amdgpu_release(jit_); + jit_ = nullptr; + } + } + + ggml_hrx_loom_jit_amdgpu * jit_ = nullptr; +}; + +static size_t default_worker_count() { + const unsigned int hardware_threads = std::thread::hardware_concurrency(); + if (hardware_threads == 0) { + return 1; + } + return std::min(hardware_threads, 4); +} + +static bool compile_kernel(ggml_hrx_loom_jit_amdgpu * jit, + const LoomKernelCompileRequest & request, + const std::string & key, + ggml_hrx_loom_jit_compile_result & compiled, + std::string & error_message) { + if (loom_jit_disk_cache_load(request, compiled)) { // 1bit: compiled in an earlier process + return true; + } + std::vector configs; + configs.reserve(request.config_storage.size()); + for (const auto & config : request.config_storage) { + configs.push_back({ config.first.c_str(), config.second.c_str() }); + } + + ggml_hrx_loom_jit_compile_options compile_options = {}; + compile_options.source_data = request.source_data; + compile_options.source_size = request.source_size; + compile_options.source_format = request.source_format; + compile_options.source_identifier = request.source_identifier.c_str(); + compile_options.root_symbol = request.symbol.c_str(); + compile_options.launch_config_symbol = request.launch_config_symbol.c_str(); + compile_options.module_name = request.symbol.c_str(); + compile_options.artifact_identifier = request.symbol.c_str(); + compile_options.dependencies = request.dependencies.data(); + compile_options.dependency_count = request.dependencies.size(); + compile_options.config_bindings = configs.data(); + compile_options.config_binding_count = configs.size(); + compile_options.workload_arguments = request.workload.data(); + compile_options.workload_argument_count = request.workload.size(); + compile_options.evaluate_launch_config = true; + + if (ErrorResult error = take_status(ggml_hrx_loom_jit_amdgpu_compile(jit, &compile_options, &compiled))) { + error_message = "compile " + key + ": " + *error; + GGML_LOG_ERROR("%s: %s\n", __func__, error_message.c_str()); + return false; + } + loom_jit_disk_cache_store(request, compiled); // 1bit + return true; +} + +static bool is_disabled_value(const char * value) { + if (value == nullptr) { + return false; + } + return std::strcmp(value, "0") == 0 || std::strcmp(value, "false") == 0 || std::strcmp(value, "FALSE") == 0 || + std::strcmp(value, "off") == 0 || std::strcmp(value, "OFF") == 0; +} + +} // namespace + +class LoomSyncJit final : public LoomJit { + public: + LoomSyncJit(const char * target, std::string & error_message) { valid_ = jit_.create(target, error_message); } + + LoomCompiledKernelRef compile(std::string key, LoomKernelCompileRequest request) override { + auto compiled_ref = std::make_shared(std::move(key), std::move(request)); + if (!valid_) { + compiled_ref->complete({}, false, "Loom JIT is not initialized"); + return compiled_ref; + } + compile_ref(*compiled_ref, jit_.get()); + return compiled_ref; + } + + bool async_enabled() const override { return false; } + + private: + static void compile_ref(LoomCompiledKernel & compiled_ref, ggml_hrx_loom_jit_amdgpu * jit) { + ggml_hrx_loom_jit_compile_result compiled; + std::string error_message; + const bool success = compile_kernel(jit, compiled_ref.request(), compiled_ref.key(), compiled, error_message); + compiled_ref.complete(std::move(compiled), success, std::move(error_message)); + } + + LoomAmdgpuJit jit_; + bool valid_ = false; +}; + +class LoomAsyncJit final : public LoomJit { + public: + LoomAsyncJit(const char * target, std::string & error_message) : target_(target != nullptr ? target : "") { + if (target_.empty()) { + error_message = "missing HRX target"; + return; + } + valid_ = true; + } + + ~LoomAsyncJit() override { clear(); } + + LoomCompiledKernelRef compile(std::string key, LoomKernelCompileRequest request) override { + auto compiled_ref = std::make_shared(std::move(key), std::move(request)); + if (!valid_) { + compiled_ref->complete({}, false, "Loom JIT is not initialized"); + return compiled_ref; + } + + { + std::lock_guard lock(mutex_); + start_workers_locked(); + pending_.push_back(compiled_ref); + } + work_available_.notify_one(); + return compiled_ref; + } + + void clear() override { stop_workers(); } + + bool async_enabled() const override { return true; } + + private: + void start_workers_locked() { + if (!workers_.empty()) { + return; + } + shutdown_ = false; + const size_t count = default_worker_count(); + workers_.reserve(count); + for (size_t i = 0; i < count; ++i) { + workers_.emplace_back([this] { worker_loop(); }); + } + } + + void stop_workers() { + { + std::lock_guard lock(mutex_); + if (workers_.empty()) { + return; + } + shutdown_ = true; + } + work_available_.notify_all(); + for (std::thread & worker : workers_) { + if (worker.joinable()) { + worker.join(); + } + } + { + std::lock_guard lock(mutex_); + pending_.clear(); + workers_.clear(); + shutdown_ = false; + } + } + + void worker_loop() { + LoomAmdgpuJit worker_jit; + std::string jit_error; + const bool jit_ready = worker_jit.create(target_.c_str(), jit_error); + for (;;) { + LoomCompiledKernelRef compiled_ref; + { + std::unique_lock lock(mutex_); + work_available_.wait(lock, [&] { return shutdown_ || !pending_.empty(); }); + if (shutdown_ && pending_.empty()) { + return; + } + compiled_ref = std::move(pending_.front()); + pending_.pop_front(); + } + + if (!jit_ready) { + compiled_ref->complete({}, false, jit_error); + continue; + } + compile_ref(*compiled_ref, worker_jit.get()); + } + } + + static void compile_ref(LoomCompiledKernel & compiled_ref, ggml_hrx_loom_jit_amdgpu * jit) { + ggml_hrx_loom_jit_compile_result compiled; + std::string error_message; + const bool success = compile_kernel(jit, compiled_ref.request(), compiled_ref.key(), compiled, error_message); + compiled_ref.complete(std::move(compiled), success, std::move(error_message)); + } + + const std::string target_; + bool valid_ = false; + std::mutex mutex_; + std::condition_variable work_available_; + std::deque pending_; + std::vector workers_; + bool shutdown_ = false; +}; + +bool loom_async_jit_enabled_from_environment() { + static const bool enabled = [] { + const char * value = std::getenv("GGML_HRX_ASYNC_JIT"); + return !is_disabled_value(value); + }(); + return enabled; +} + +std::unique_ptr create_loom_jit(const char * target, LoomJitMode mode, std::string & error_message) { + if (target == nullptr || target[0] == '\0') { + error_message = "missing HRX target"; + GGML_LOG_ERROR("%s: %s\n", __func__, error_message.c_str()); + return nullptr; + } + if (mode == LoomJitMode::Async) { + auto jit = std::make_unique(target, error_message); + if (!error_message.empty()) { + GGML_LOG_ERROR("%s: %s\n", __func__, error_message.c_str()); + return nullptr; + } + return jit; + } + auto jit = std::make_unique(target, error_message); + if (!error_message.empty()) { + return nullptr; + } + return jit; +} + +std::unique_ptr create_loom_jit(const char * target, std::string & error_message) { + return create_loom_jit(target, loom_async_jit_enabled_from_environment() ? LoomJitMode::Async : LoomJitMode::Sync, + error_message); +} + +LoomCompiledKernel::LoomCompiledKernel(std::string key, LoomKernelCompileRequest request) : + key_(std::move(key)), + request_(std::move(request)) {} + +bool LoomCompiledKernel::resolve() const { + const State current = state(); + if (current != State::Pending) { + return current == State::Succeeded; + } + + std::unique_lock lock(mutex_); + complete_.wait(lock, [&] { return state() != State::Pending; }); + return state() == State::Succeeded; +} + +std::string LoomCompiledKernel::error_message() const { + std::lock_guard lock(mutex_); + return error_; +} + +ggml_hrx_loom_jit_compile_result LoomCompiledKernel::take_result() { + std::lock_guard lock(mutex_); + return std::move(compiled_); +} + +void LoomCompiledKernel::complete(ggml_hrx_loom_jit_compile_result compiled, bool success, std::string error) { + { + std::lock_guard lock(mutex_); + compiled_ = std::move(compiled); + error_ = std::move(error); + state_.store(success ? State::Succeeded : State::Failed, std::memory_order_release); + } + complete_.notify_all(); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/loom-kernel-jit.h b/ggml/src/ggml-hrx/runtime/loom-kernel-jit.h new file mode 100644 index 000000000000..3903e2d88821 --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/loom-kernel-jit.h @@ -0,0 +1,96 @@ +#pragma once + +#include "loom-jit.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { + +struct LoomKernelCompileRequest { + const void * source_data = nullptr; + size_t source_size = 0; + ggml_hrx_loom_jit_source_format source_format = GGML_HRX_LOOM_JIT_SOURCE_FORMAT_TEXT; + std::string source_identifier; + std::string symbol; + std::string launch_config_symbol; + std::vector dependencies; + std::vector> config_storage; + std::vector workload; +}; + +class LoomCompiledKernel { + public: + LoomCompiledKernel(std::string key, LoomKernelCompileRequest request); + + LoomCompiledKernel(const LoomCompiledKernel &) = delete; + LoomCompiledKernel & operator=(const LoomCompiledKernel &) = delete; + + const std::string & key() const { return key_; } + + bool resolve() const; + std::string error_message() const; + ggml_hrx_loom_jit_compile_result take_result(); + + private: + friend class LoomSyncJit; + friend class LoomAsyncJit; + friend struct HipCodeObjectLoader; + + const LoomKernelCompileRequest & request() const { return request_; } + + void complete(ggml_hrx_loom_jit_compile_result compiled, bool success, std::string error); + + enum class State { + Pending, + Succeeded, + Failed, + }; + + State state() const { return state_.load(std::memory_order_acquire); } + + mutable std::mutex mutex_; + mutable std::condition_variable complete_; + std::string key_; + LoomKernelCompileRequest request_; + ggml_hrx_loom_jit_compile_result compiled_; + std::string error_; + std::atomic state_ = State::Pending; +}; + +using LoomCompiledKernelRef = std::shared_ptr; + +class LoomJit { + public: + virtual ~LoomJit() = default; + + LoomJit(const LoomJit &) = delete; + LoomJit & operator=(const LoomJit &) = delete; + + virtual LoomCompiledKernelRef compile(std::string key, LoomKernelCompileRequest request) = 0; + + virtual void clear() {} + + virtual bool async_enabled() const = 0; + + protected: + LoomJit() = default; +}; + +enum class LoomJitMode { + Sync, + Async, +}; + +std::unique_ptr create_loom_jit(const char * target, LoomJitMode mode, std::string & error_message); +std::unique_ptr create_loom_jit(const char * target, std::string & error_message); +bool loom_async_jit_enabled_from_environment(); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/prepared-command-program-cache.cpp b/ggml/src/ggml-hrx/runtime/prepared-command-program-cache.cpp new file mode 100644 index 000000000000..ad81b1a23afc --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/prepared-command-program-cache.cpp @@ -0,0 +1,228 @@ +#include "prepared-command-program-cache.h" + +#include +#include +#include +#include + +namespace ggml::hrx { +namespace { + +static void mix_hash(uint64_t & hash, uint64_t value) { + hash ^= value; + hash *= UINT64_C(1099511628211); +} + +static uint64_t hash_text(const char * text) { + uint64_t hash = UINT64_C(1469598103934665603); + if (text == nullptr) { + return hash; + } + while (*text != 0) { + mix_hash(hash, static_cast(*text)); + ++text; + } + return hash; +} + +static void apply_graph_replay_result(PreparedCommandProgramCacheExecutionResult & result, + const RecordedCommandGraphExecutionResult & replay) { + result.graph_replay_event = replay.event; + result.graph_replay_ineligible_reason = replay.ineligible_reason; + result.graph_replay_build_ns = replay.build_ns; + result.graph_replay_launch_ns = replay.launch_ns; + result.graph_replay_total_ns = replay.total_ns(); + result.graph_replay_dispatches = replay.dispatch_count; + result.graph_replay_transient_allocation_changed = replay.transient_allocation_changed; +} + +static bool graph_replay_should_fallback(HrxGraphReplayEvent event) { + return event == HrxGraphReplayEvent::Ineligible || event == HrxGraphReplayEvent::BuildFailed; +} + +} // namespace + +uint64_t command_program_shape_hash(const std::string & command_shape) { + uint64_t hash = UINT64_C(1469598103934665603); + for (const char c : command_shape) { + mix_hash(hash, static_cast(c)); + } + return hash; +} + +size_t PreparedCommandProgramCache::KeyHash::operator()(const Key & key) const { + uint64_t hash = UINT64_C(1469598103934665603); + mix_hash(hash, key.graph_uid); + mix_hash(hash, key.target_hash); + mix_hash(hash, key.command_shape_hash); + mix_hash(hash, key.bindings_hash); + return static_cast(hash); +} + +PreparedCommandProgramCache::Key PreparedCommandProgramCache::cache_key(uint64_t graph_uid, + const CommandProgramExecutionContext & context, + uint64_t command_shape_hash, + const CommandProgramBindings & bindings) const { + return { + graph_uid, + hash_text(context.target), + command_shape_hash, + command_program_bindings_hash(bindings).value, + }; +} + +bool PreparedCommandProgramCache::execute(const CommandProgramExecutionContext & context, + uint64_t graph_uid, + const std::string & command_shape, + const CommandProgram & commands, + const CommandProgramBindings & bindings) { + return execute_with_result(context, graph_uid, command_shape, commands, bindings).success; +} + +PreparedCommandProgramCacheExecutionResult PreparedCommandProgramCache::execute_with_result( + const CommandProgramExecutionContext & context, + uint64_t graph_uid, + const std::string & command_shape, + const CommandProgram & commands, + const CommandProgramBindings & bindings) { + return execute_with_result(context, graph_uid, command_program_shape_hash(command_shape), commands, bindings); +} + +PreparedCommandProgramCacheExecutionResult PreparedCommandProgramCache::execute_with_result( + const CommandProgramExecutionContext & context, + uint64_t graph_uid, + uint64_t command_shape_hash, + const CommandProgram & commands, + const CommandProgramBindings & bindings) { + PreparedCommandProgramCacheExecutionResult result; + if (graph_uid == 0 || !commands.valid() || !bindings.valid()) { + result.graph_replay_event = HrxGraphReplayEvent::Ineligible; + result.graph_replay_ineligible_reason = "uncached_graph"; + PreparedCommandProgram prepared = prepare_command_program(context, commands, bindings); + if (!prepared.valid()) { + result.status.append(prepared.status); + return result; + } + result.success = bind_and_execute_prepared_command_program(context, commands, bindings, prepared); + if (!result.success) { + result.status.log("execute uncached HRX command program failed"); + } + return result; + } + + const Key key = cache_key(graph_uid, context, command_shape_hash, bindings); + std::shared_ptr entry; + bool created_entry = false; + { + std::lock_guard lock(mutex_); + auto found = programs_.find(key); + if (found == programs_.end()) { + entry = std::make_shared(); + programs_.emplace(key, entry); + created_entry = true; + } else { + entry = found->second; + } + } + + std::lock_guard entry_lock(entry->mutex); + if (entry->has_program && entry->program.valid()) { + record_hit(); + if (debug_serial_command_execution_enabled()) { + result.graph_replay_event = HrxGraphReplayEvent::Disabled; + result.graph_replay_ineligible_reason = "debug_serial_execution"; + result.success = bind_and_execute_prepared_command_program(context, commands, bindings, entry->program); + if (!result.success) { + result.status.log("execute cached HRX command program failed"); + } + return result; + } + const RecordedCommandGraphExecutionResult replay = + bind_and_launch_recorded_command_graph(context, commands, bindings, entry->program, entry->recorded); + apply_graph_replay_result(result, replay); + if (replay.success) { + result.success = true; + return result; + } + if (!graph_replay_should_fallback(replay.event)) { + result.status.append(replay.status); + if (result.status.success()) { + result.status.log("execute cached HRX graph replay failed"); + } + return result; + } + result.success = bind_and_execute_prepared_command_program(context, commands, bindings, entry->program); + if (!result.success) { + result.status.log("execute cached HRX command program failed"); + } + return result; + } + + PreparedCommandProgram prepared = prepare_command_program(context, commands, bindings); + if (!prepared.valid()) { + if (created_entry) { + std::lock_guard lock(mutex_); + const auto found = programs_.find(key); + if (found != programs_.end() && found->second == entry && !entry->has_program) { + programs_.erase(found); + } + } + result.status.append(prepared.status); + return result; + } + entry->program = std::move(prepared); + entry->has_program = true; + record_build(); + + if (debug_serial_command_execution_enabled()) { + result.graph_replay_event = HrxGraphReplayEvent::Disabled; + result.graph_replay_ineligible_reason = "debug_serial_execution"; + result.success = bind_and_execute_prepared_command_program(context, commands, bindings, entry->program); + if (!result.success) { + result.status.log("execute prepared HRX command program failed"); + } + return result; + } + + const RecordedCommandGraphExecutionResult replay = + bind_and_launch_recorded_command_graph(context, commands, bindings, entry->program, entry->recorded); + apply_graph_replay_result(result, replay); + if (replay.success) { + result.success = true; + return result; + } + if (!graph_replay_should_fallback(replay.event)) { + result.status.append(replay.status); + if (result.status.success()) { + result.status.log("execute prepared HRX graph replay failed"); + } + return result; + } + result.success = bind_and_execute_prepared_command_program(context, commands, bindings, entry->program); + if (!result.success) { + result.status.log("execute prepared HRX command program failed"); + } + return result; +} + +PreparedCommandProgramCacheStats PreparedCommandProgramCache::stats() const { + std::lock_guard lock(mutex_); + return stats_; +} + +void PreparedCommandProgramCache::clear() { + std::lock_guard lock(mutex_); + programs_.clear(); +} + +void PreparedCommandProgramCache::record_build() { + std::lock_guard lock(mutex_); + ++stats_.builds; +} + +void PreparedCommandProgramCache::record_hit() { + std::lock_guard lock(mutex_); + ++stats_.hits; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/prepared-command-program-cache.h b/ggml/src/ggml-hrx/runtime/prepared-command-program-cache.h new file mode 100644 index 000000000000..069ffd624480 --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/prepared-command-program-cache.h @@ -0,0 +1,97 @@ +#pragma once + +#include "command-program-executor.h" +#include "dispatch/command-program-bindings.h" +#include "dispatch/command-program.h" +#include "runtime/graph-replay.h" + +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { + +uint64_t command_program_shape_hash(const std::string & command_shape); + +struct PreparedCommandProgramCacheStats { + uint64_t builds = 0; + uint64_t hits = 0; +}; + +struct PreparedCommandProgramCacheExecutionResult { + bool success = false; + Status status; + HrxGraphReplayEvent graph_replay_event = HrxGraphReplayEvent::Disabled; + std::string graph_replay_ineligible_reason; + uint64_t graph_replay_build_ns = 0; + uint64_t graph_replay_launch_ns = 0; + uint64_t graph_replay_total_ns = 0; + size_t graph_replay_dispatches = 0; + bool graph_replay_transient_allocation_changed = false; +}; + +class PreparedCommandProgramCache { + public: + bool execute(const CommandProgramExecutionContext & context, + uint64_t graph_uid, + const std::string & command_shape, + const CommandProgram & commands, + const CommandProgramBindings & bindings); + + PreparedCommandProgramCacheExecutionResult execute_with_result(const CommandProgramExecutionContext & context, + uint64_t graph_uid, + const std::string & command_shape, + const CommandProgram & commands, + const CommandProgramBindings & bindings); + + PreparedCommandProgramCacheExecutionResult execute_with_result(const CommandProgramExecutionContext & context, + uint64_t graph_uid, + uint64_t command_shape_hash, + const CommandProgram & commands, + const CommandProgramBindings & bindings); + + PreparedCommandProgramCacheStats stats() const; + + void clear(); + + private: + struct Key { + uint64_t graph_uid = 0; + uint64_t target_hash = 0; + uint64_t command_shape_hash = 0; + uint64_t bindings_hash = 0; + + bool operator==(const Key & other) const { + return graph_uid == other.graph_uid && target_hash == other.target_hash && + command_shape_hash == other.command_shape_hash && bindings_hash == other.bindings_hash; + } + }; + + struct KeyHash { + size_t operator()(const Key & key) const; + }; + + Key cache_key(uint64_t graph_uid, + const CommandProgramExecutionContext & context, + uint64_t command_shape_hash, + const CommandProgramBindings & bindings) const; + + struct Entry { + std::mutex mutex; + PreparedCommandProgram program; + RecordedCommandGraph recorded; + bool has_program = false; + }; + + void record_build(); + void record_hit(); + + mutable std::mutex mutex_; + std::unordered_map, KeyHash> programs_; + PreparedCommandProgramCacheStats stats_; +}; + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/transient-arena.cpp b/ggml/src/ggml-hrx/runtime/transient-arena.cpp new file mode 100644 index 000000000000..a12f89e0ab38 --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/transient-arena.cpp @@ -0,0 +1,135 @@ +#include "transient-arena.h" + +#include "hrx-interop-utils.h" +#include "hrx_runtime.h" + +#include +#include + +namespace ggml::hrx { + +TransientArena::AllocationLease::AllocationLease(TransientArena & arena, std::unique_lock lock) : + arena_(&arena), + lock_(std::move(lock)) {} + +Status TransientArena::AllocationLease::ensure_capacity(hrx_device_t device, + hrx_stream_t stream, + size_t required_size) { + Status status; + if (arena_ == nullptr || !lock_.owns_lock()) { + status.log("missing transient arena allocation lease"); + return status; + } + return arena_->ensure_capacity_locked(device, stream, required_size); +} + +TransientArenaAllocationRef TransientArena::AllocationLease::current_allocation() const { + if (arena_ == nullptr || !lock_.owns_lock()) { + return {}; + } + return arena_->current_allocation_locked(); +} + +TransientArena::~TransientArena() { + clear(); +} + +void TransientArena::clear() { + std::lock_guard lock(mutex_); + if (buffer_ != nullptr) { + hrx_buffer_release(buffer_); + buffer_ = nullptr; + } + for (hrx_buffer_t buffer : diagnostic_retired_buffers_) { + hrx_buffer_release(buffer); + } + diagnostic_retired_buffers_.clear(); + allocation_capacity_ = 0; + allocation_id_ = kInvalidTransientArenaAllocationId; +} + +uint64_t TransientArena::next_allocation_id() { + const uint64_t id = next_allocation_id_++; + if (next_allocation_id_ == kInvalidTransientArenaAllocationId) { + ++next_allocation_id_; + } + return id; +} + +Status TransientArena::ensure_capacity(hrx_device_t device, hrx_stream_t stream, size_t required_size) { + std::lock_guard lock(mutex_); + return ensure_capacity_locked(device, stream, required_size); +} + +TransientArena::AllocationLease TransientArena::acquire_allocation_lease() { + return AllocationLease(*this, std::unique_lock(mutex_)); +} + +Status TransientArena::ensure_capacity_locked(hrx_device_t device, hrx_stream_t stream, size_t required_size) { + Status status; + const bool diagnostic_fresh = std::getenv("GGML_HRX_DIAGNOSTIC_FRESH_TRANSIENT_ARENA") != nullptr; + if (required_size == 0 || (!diagnostic_fresh && allocation_capacity_ >= required_size)) { + return status; + } + if (device == nullptr) { + status.log("missing HRX device for transient arena allocation"); + return status; + } + if (stream == nullptr) { + status.log("missing HRX stream for transient arena allocation"); + return status; + } + if (buffer_ != nullptr) { + if (diagnostic_fresh) { + diagnostic_retired_buffers_.push_back(buffer_); + } else { + if (ErrorResult error = take_status(hrx_stream_synchronize(stream))) { + status.log("synchronize before growing transient arena: %s", error->c_str()); + return status; + } + hrx_buffer_release(buffer_); + } + buffer_ = nullptr; + allocation_capacity_ = 0; + allocation_id_ = kInvalidTransientArenaAllocationId; + } + + hrx_buffer_params_t params = { + HRX_MEMORY_TYPE_DEVICE_LOCAL, + HRX_MEMORY_ACCESS_ALL, + HRX_BUFFER_USAGE_DEFAULT, + 0, + }; + hrx_buffer_t allocation = nullptr; + if (ErrorResult error = take_status( + hrx_allocator_allocate_buffer(hrx_device_allocator(device), params, required_size, &allocation))) { + status.log("allocate transient arena: %s", error->c_str()); + return status; + } + + buffer_ = allocation; + allocation_capacity_ = required_size; + allocation_id_ = next_allocation_id(); + return status; +} + +TransientArenaAllocationRef TransientArena::current_allocation() const { + std::lock_guard lock(mutex_); + return current_allocation_locked(); +} + +size_t TransientArena::capacity() const { + std::lock_guard lock(mutex_); + return allocation_capacity_; +} + +uint64_t TransientArena::allocation_id() const { + std::lock_guard lock(mutex_); + return allocation_id_; +} + +TransientArenaAllocationRef TransientArena::current_allocation_locked() const { + return { buffer_, allocation_capacity_, allocation_id_ }; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/runtime/transient-arena.h b/ggml/src/ggml-hrx/runtime/transient-arena.h new file mode 100644 index 000000000000..a50639962c7c --- /dev/null +++ b/ggml/src/ggml-hrx/runtime/transient-arena.h @@ -0,0 +1,68 @@ +#pragma once + +#include "dispatch/command-program-resolver.h" +#include "status.h" + +#include +#include +#include +#include + +typedef struct hrx_device_s * hrx_device_t; +typedef struct hrx_stream_s * hrx_stream_t; + +namespace ggml::hrx { + +class TransientArena { + public: + class AllocationLease { + public: + AllocationLease() = default; + AllocationLease(AllocationLease &&) noexcept = default; + AllocationLease & operator=(AllocationLease &&) noexcept = default; + + AllocationLease(const AllocationLease &) = delete; + AllocationLease & operator=(const AllocationLease &) = delete; + + Status ensure_capacity(hrx_device_t device, hrx_stream_t stream, size_t required_size); + TransientArenaAllocationRef current_allocation() const; + + private: + friend class TransientArena; + + AllocationLease(TransientArena & arena, std::unique_lock lock); + + TransientArena * arena_ = nullptr; + std::unique_lock lock_; + }; + + TransientArena() = default; + ~TransientArena(); + + TransientArena(const TransientArena &) = delete; + TransientArena & operator=(const TransientArena &) = delete; + + Status ensure_capacity(hrx_device_t device, hrx_stream_t stream, size_t required_size); + AllocationLease acquire_allocation_lease(); + void clear(); + + TransientArenaAllocationRef current_allocation() const; + + size_t capacity() const; + + uint64_t allocation_id() const; + + private: + uint64_t next_allocation_id(); + Status ensure_capacity_locked(hrx_device_t device, hrx_stream_t stream, size_t required_size); + TransientArenaAllocationRef current_allocation_locked() const; + + mutable std::mutex mutex_; + hrx_buffer_t buffer_ = nullptr; + std::vector diagnostic_retired_buffers_; + size_t allocation_capacity_ = 0; + uint64_t allocation_id_ = kInvalidTransientArenaAllocationId; + uint64_t next_allocation_id_ = 1; +}; + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/status.h b/ggml/src/ggml-hrx/status.h new file mode 100644 index 000000000000..6d6aafc8ab6e --- /dev/null +++ b/ggml/src/ggml-hrx/status.h @@ -0,0 +1,72 @@ +#pragma once + +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx { + +class [[nodiscard]] Status { + public: + Status() = default; + + Status(const Status &) = delete; + Status & operator=(const Status &) = delete; + + Status(Status &&) noexcept = default; + Status & operator=(Status &&) noexcept = default; + + bool success() const { return messages_.empty(); } + + const std::vector & errors() const { return messages_; } + + void append(const Status & other) { + messages_.insert(messages_.end(), other.messages_.begin(), other.messages_.end()); + } + + void log(const char * format, ...) { + va_list args; + va_start(args, format); + log_va(format, args); + va_end(args); + } + + private: + void push(std::string message) { messages_.push_back(std::move(message)); } + + void log_va(const char * format, va_list args) { + if (format == nullptr) { + push("failed to format error message"); + return; + } + + char stack[256]; + va_list args_copy; + va_copy(args_copy, args); + const int written = std::vsnprintf(stack, sizeof(stack), format, args_copy); + va_end(args_copy); + if (written < 0) { + push("failed to format error message"); + return; + } + if (static_cast(written) < sizeof(stack)) { + push(std::string(stack, static_cast(written))); + return; + } + + std::vector buffer(static_cast(written) + 1); + const int rewritten = std::vsnprintf(buffer.data(), buffer.size(), format, args); + if (rewritten < 0) { + push("failed to format error message"); + return; + } + push(std::string(buffer.data(), static_cast(rewritten))); + } + + std::vector messages_; +}; + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/tools/analyze-graph.cpp b/ggml/src/ggml-hrx/tools/analyze-graph.cpp new file mode 100644 index 000000000000..8ae0b7febfb2 --- /dev/null +++ b/ggml/src/ggml-hrx/tools/analyze-graph.cpp @@ -0,0 +1,159 @@ +#include "dispatch/command-program-diagnostics.h" +#include "dispatch/command-program.h" +#include "dispatch/dispatch-scheduler.h" +#include "graph/graph-diagnostics.h" +#include "kernel-corpus/kernel-corpus.h" + +#include +#include +#include +#include +#include +#include + +namespace { + +struct Options { + std::filesystem::path graph; + std::filesystem::path output_directory; + std::string target; +}; + +static void print_usage(const char * program) { + std::cerr << "usage: " << program << " --graph [--target ] [--out ]\n"; +} + +static bool parse_options(int argc, char ** argv, Options & options) { + for (int i = 1; i < argc; ++i) { + const std::string arg = argv[i]; + if (arg == "--help" || arg == "-h") { + print_usage(argv[0]); + return false; + } + if (arg == "--graph" && i + 1 < argc) { + options.graph = argv[++i]; + continue; + } + if (arg == "--target" && i + 1 < argc) { + options.target = argv[++i]; + continue; + } + if (arg == "--out" && i + 1 < argc) { + options.output_directory = argv[++i]; + continue; + } + std::cerr << "unknown or incomplete argument: " << arg << '\n'; + print_usage(argv[0]); + return false; + } + if (options.graph.empty()) { + std::cerr << "--graph is required\n"; + print_usage(argv[0]); + return false; + } + return true; +} + +static std::string read_file(const std::filesystem::path & path) { + std::ifstream input(path, std::ios::binary); + if (!input) { + return {}; + } + return { std::istreambuf_iterator(input), std::istreambuf_iterator() }; +} + +static void write_file(const std::filesystem::path & path, const std::string & contents) { + std::filesystem::create_directories(path.parent_path()); + std::ofstream output(path, std::ios::binary | std::ios::trunc); + if (!output) { + throw std::runtime_error("cannot create " + path.string()); + } + output << contents; + if (contents.empty() || contents.back() != '\n') { + output << '\n'; + } +} + +static void write_output(const Options & options, const std::string & name, const std::string & contents) { + if (options.output_directory.empty()) { + std::cout << "== " << name << " ==\n" << contents; + if (contents.empty() || contents.back() != '\n') { + std::cout << '\n'; + } + return; + } + write_file(options.output_directory / name, contents); +} + +} // namespace + +int main(int argc, char ** argv) { + Options options; + if (!parse_options(argc, argv, options)) { + return 1; + } + + const std::string contents = read_file(options.graph); + if (contents.empty()) { + std::cerr << "cannot read graph snapshot: " << options.graph << '\n'; + return 1; + } + + ggml::hrx::GraphSnapshotLoadResult snapshot = ggml::hrx::load_graph_snapshot_json(contents); + if (!snapshot.valid()) { + for (const std::string & error : snapshot.status.errors()) { + std::cerr << error << '\n'; + } + return 1; + } + + if (options.target.empty()) { + options.target = snapshot.target; + } + if (options.target.empty()) { + std::cerr << "target is missing from both --target and graph snapshot\n"; + return 1; + } + + write_output(options, "graph.txt", + ggml::hrx::format_graph_snapshot_text(snapshot.graph, options.target, snapshot.uid)); + write_output(options, "graph.json", + ggml::hrx::serialize_graph_snapshot_json(snapshot.graph, options.target, snapshot.uid)); + + ggml::hrx::DispatchScheduler scheduler; + ggml::hrx::DispatchScheduleDiagnostics diagnostics; + const bool scheduled = scheduler.schedule_graph(snapshot.graph, { options.target }, &diagnostics); + write_output(options, "schedule.txt", + ggml::hrx::format_schedule_diagnostics_text(snapshot.graph, scheduler.plan(), diagnostics)); + write_output(options, "schedule.json", + ggml::hrx::serialize_schedule_diagnostics_json(snapshot.graph, scheduler.plan(), diagnostics)); + if (!scheduled) { + write_output(options, "unmatched.txt", + ggml::hrx::format_schedule_diagnostics_text(snapshot.graph, scheduler.plan(), diagnostics)); + write_output(options, "unmatched.json", + ggml::hrx::serialize_schedule_diagnostics_json(snapshot.graph, scheduler.plan(), diagnostics)); + return 2; + } + + const ggml::hrx::KernelCorpus & corpus = ggml::hrx::get_qwen_kernel_corpus(); + ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(snapshot.graph, scheduler.plan(), corpus, options.target); + write_output(options, "commands.txt", ggml::hrx::format_command_program(commands)); + if (!commands.valid()) { + for (const std::string & error : commands.status.errors()) { + std::cerr << error << '\n'; + } + return 3; + } + + const ggml::hrx::VerificationResult verification = + ggml::hrx::verify_command_program(commands, corpus, options.target); + if (!verification.valid()) { + for (const std::string & error : verification.status.errors()) { + std::cerr << error << '\n'; + } + return 4; + } + + return 0; +} diff --git a/ggml/src/ggml-hrx/tools/benchmarks/analyze-model-fusion-adjacency.py b/ggml/src/ggml-hrx/tools/benchmarks/analyze-model-fusion-adjacency.py new file mode 100644 index 000000000000..00fa5edbe3b0 --- /dev/null +++ b/ggml/src/ggml-hrx/tools/benchmarks/analyze-model-fusion-adjacency.py @@ -0,0 +1,750 @@ +#!/usr/bin/env python3 +# +# Analyze HRX command-program adjacency for model-scoped Loom benchmarks. + +from __future__ import annotations + +import argparse +import hashlib +import json +from collections import Counter, defaultdict +from dataclasses import dataclass +from pathlib import Path +from typing import Any + + +TRANSIENT_ORIGIN = "Transient" +READ_ACCESS = "Read" +WRITE_ACCESSES = {"Write", "ReadWrite"} + + +@dataclass(frozen=True) +class BindingKey: + value: int + offset: int + length: int + + +@dataclass(frozen=True) +class Producer: + ordinal: int + kernel: str + binding: str + key: BindingKey + + +@dataclass(frozen=True) +class Consumer: + ordinal: int + kernel: str + binding: str + + +def fail(message: str) -> None: + raise SystemExit(message) + + +def load_json(path: Path) -> Any: + with path.open("r", encoding="utf-8") as f: + return json.load(f) + + +def write_json(path: Path, data: Any) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", encoding="utf-8") as f: + json.dump(data, f, indent=2, sort_keys=True) + f.write("\n") + + +def kernel_symbol(kernel: str) -> str: + return kernel.split(":")[-1] + + +def command_shape_key(command: dict[str, Any]) -> str: + shape_data = { + "kernel": command.get("kernel"), + "integer_parameters": command.get("integer_parameters", {}), + "compile_parameters": command.get("compile_parameters", {}), + "binding_lengths": [binding.get("length") for binding in command.get("bindings", [])], + } + encoded = json.dumps(shape_data, sort_keys=True, separators=(",", ":")).encode("utf-8") + return hashlib.sha256(encoded).hexdigest()[:16] + + +def load_programs(dump_dir: Path) -> list[dict[str, Any]]: + paths = sorted(dump_dir.glob("program-*/program.json")) + if not paths: + fail(f"no program.json files found under {dump_dir}") + programs = [] + for path in paths: + program = load_json(path) + program["path"] = str(path) + programs.append(program) + return programs + + +def load_commands(dump_dir: Path) -> list[dict[str, Any]]: + commands = [] + for program in load_programs(dump_dir): + for command in program.get("commands", []): + command = dict(command) + command["program"] = { + "directory": Path(program["path"]).parent.name, + "dump_id": program.get("dump_id"), + "shape_hash": program.get("shape_hash"), + "target": program.get("target"), + } + commands.append(command) + return commands + + +def shape_key_order(commands: list[dict[str, Any]]) -> list[str]: + keys = [] + seen = set() + for command in commands: + key = command_shape_key(command) + if key not in seen: + seen.add(key) + keys.append(key) + return keys + + +def shape_counts(commands: list[dict[str, Any]]) -> Counter[str]: + return Counter(command_shape_key(command) for command in commands) + + +def shape_metrics_from_summary(summary: dict[str, Any]) -> dict[str, dict[str, Any]]: + by_benchmark = {} + for kernel in summary.get("kernels", []): + for shape in kernel.get("shapes", []): + benchmark = shape.get("benchmark") + if benchmark: + by_benchmark[benchmark] = shape + return by_benchmark + + +def map_shape_keys_to_dispatches(commands: list[dict[str, Any]], + manifest: dict[str, Any], + summary: dict[str, Any]) -> tuple[dict[str, dict[str, Any]], list[str]]: + warnings = [] + keys = shape_key_order(commands) + dispatches = manifest.get("dispatches", []) + if len(keys) != len(dispatches): + warnings.append(f"shape count mismatch: dump has {len(keys)} unique shapes, manifest has {len(dispatches)} dispatches") + + summary_by_benchmark = shape_metrics_from_summary(summary) + shape_info = {} + for index, key in enumerate(keys[:len(dispatches)]): + dispatch = dispatches[index] + command = next(command for command in commands if command_shape_key(command) == key) + if command.get("kernel") != dispatch.get("kernel"): + warnings.append( + f"shape {index} kernel mismatch: dump has {command.get('kernel')}, manifest has {dispatch.get('kernel')}" + ) + benchmark = dispatch.get("benchmark") + metric = summary_by_benchmark.get(benchmark, {}) + shape_info[key] = { + "benchmark": benchmark, + "kernel": dispatch.get("kernel"), + "count": dispatch.get("count", 0), + "metric_ns": metric.get("metric_ns"), + "weighted_ns": metric.get("weighted_ns"), + "state": metric.get("state", "missing_summary"), + "error": metric.get("error"), + } + return shape_info, warnings + + +def command_metric_ns(command: dict[str, Any], shape_info: dict[str, dict[str, Any]]) -> float: + info = shape_info.get(command_shape_key(command), {}) + metric = info.get("metric_ns") + if metric is None: + return 0.0 + return float(metric) + + +def binding_key(binding: dict[str, Any]) -> BindingKey: + return BindingKey(int(binding["value"]), int(binding.get("offset", 0)), int(binding.get("length", 0))) + + +def read_bindings(command: dict[str, Any]) -> list[dict[str, Any]]: + return [ + binding + for binding in command.get("bindings", []) + if binding.get("origin") == TRANSIENT_ORIGIN and binding.get("access") == READ_ACCESS + ] + + +def write_bindings(command: dict[str, Any]) -> list[dict[str, Any]]: + return [ + binding + for binding in command.get("bindings", []) + if binding.get("origin") == TRANSIENT_ORIGIN and binding.get("access") in WRITE_ACCESSES + ] + + +def exact_transient_edges(commands: list[dict[str, Any]]) -> tuple[list[dict[str, Any]], dict[Producer, list[Consumer]]]: + active: dict[BindingKey, Producer] = {} + edges = [] + fanout: dict[Producer, list[Consumer]] = defaultdict(list) + for command in sorted(commands, key=lambda item: int(item["ordinal"])): + consumer_ordinal = int(command["ordinal"]) + for binding in read_bindings(command): + key = binding_key(binding) + producer = active.get(key) + if producer is None: + continue + consumer = Consumer(consumer_ordinal, command["kernel"], str(binding.get("name", ""))) + edges.append( + { + "producer_ordinal": producer.ordinal, + "producer_kernel": producer.kernel, + "producer_binding": producer.binding, + "consumer_ordinal": consumer.ordinal, + "consumer_kernel": consumer.kernel, + "consumer_binding": consumer.binding, + "value": key.value, + "offset": key.offset, + "length": key.length, + } + ) + fanout[producer].append(consumer) + for binding in write_bindings(command): + key = binding_key(binding) + active[key] = Producer(consumer_ordinal, command["kernel"], str(binding.get("name", "")), key) + return edges, fanout + + +def summarize_edges(edges: list[dict[str, Any]]) -> list[dict[str, Any]]: + counts: Counter[tuple[str, str]] = Counter() + examples: dict[tuple[str, str], list[dict[str, Any]]] = defaultdict(list) + for edge in edges: + key = (edge["producer_kernel"], edge["consumer_kernel"]) + counts[key] += 1 + if len(examples[key]) < 5: + examples[key].append(edge) + return [ + { + "producer_kernel": producer, + "consumer_kernel": consumer, + "edge_count": count, + "examples": examples[(producer, consumer)], + } + for (producer, consumer), count in counts.most_common() + ] + + +def summarize_fanout(fanout: dict[Producer, list[Consumer]]) -> list[dict[str, Any]]: + counts: Counter[tuple[str, tuple[str, ...]]] = Counter() + examples: dict[tuple[str, tuple[str, ...]], list[dict[str, Any]]] = defaultdict(list) + for producer, consumers in fanout.items(): + consumer_kernels = tuple(sorted(consumer.kernel for consumer in consumers)) + key = (producer.kernel, consumer_kernels) + counts[key] += 1 + if len(examples[key]) < 5: + examples[key].append( + { + "producer_ordinal": producer.ordinal, + "producer_binding": producer.binding, + "consumers": [ + { + "ordinal": consumer.ordinal, + "kernel": consumer.kernel, + "binding": consumer.binding, + } + for consumer in consumers + ], + } + ) + rows = [] + for (producer_kernel, consumer_kernels), count in counts.most_common(): + rows.append( + { + "producer_kernel": producer_kernel, + "consumer_kernels": list(consumer_kernels), + "producer_count": count, + "examples": examples[(producer_kernel, consumer_kernels)], + } + ) + return rows + + +def producers_by_exact_key(commands: list[dict[str, Any]]) -> dict[tuple[int, BindingKey], Producer]: + active: dict[BindingKey, Producer] = {} + result = {} + for command in sorted(commands, key=lambda item: int(item["ordinal"])): + ordinal = int(command["ordinal"]) + for binding in read_bindings(command): + key = binding_key(binding) + producer = active.get(key) + if producer is not None: + result[(ordinal, key)] = producer + for binding in write_bindings(command): + key = binding_key(binding) + active[key] = Producer(ordinal, command["kernel"], str(binding.get("name", "")), key) + return result + + +def producer_for_read(command: dict[str, Any], + binding_name: str, + read_producers: dict[tuple[int, BindingKey], Producer]) -> Producer | None: + ordinal = int(command["ordinal"]) + for binding in read_bindings(command): + if binding.get("name") == binding_name: + return read_producers.get((ordinal, binding_key(binding))) + return None + + +def output_consumers(command: dict[str, Any], fanout: dict[Producer, list[Consumer]], binding_name: str = "output") -> list[Consumer]: + ordinal = int(command["ordinal"]) + for binding in write_bindings(command): + if binding.get("name") != binding_name: + continue + producer = Producer(ordinal, command["kernel"], str(binding.get("name", "")), binding_key(binding)) + return fanout.get(producer, []) + return [] + + +def command_by_ordinal(commands: list[dict[str, Any]]) -> dict[int, dict[str, Any]]: + return {int(command["ordinal"]): command for command in commands} + + +def compile_parameters(command: dict[str, Any]) -> dict[str, Any]: + return command.get("compile_parameters", {}) + + +def integer_parameters(command: dict[str, Any]) -> dict[str, Any]: + return command.get("integer_parameters", {}) + + +def op_name(command: dict[str, Any]) -> str: + if command["kernel"] != "loom_libs:ggml_binary_f32": + return "" + op = str(compile_parameters(command).get("ggml.binary_f32.op", "")) + names = { + "0": "add", + "1": "sub", + "2": "mul", + "3": "div", + "4": "swiglu", + "5": "geglu", + "6": "reglu", + "7": "geglu_erf", + "8": "geglu_quick", + } + return names.get(op, f"op_{op}") + + +def add_candidate(candidates: list[dict[str, Any]], + name: str, + kind: str, + occurrences: list[dict[str, Any]], + current_chain_ns: float, + removable_ns: float, + notes: list[str]) -> None: + candidates.append( + { + "name": name, + "kind": kind, + "occurrence_count": len(occurrences), + "current_chain_ns": current_chain_ns, + "removable_standalone_ns": removable_ns, + "examples": occurrences[:8], + "notes": notes, + } + ) + + +def classify_candidates(commands: list[dict[str, Any]], + shape_info: dict[str, dict[str, Any]], + fanout: dict[Producer, list[Consumer]]) -> list[dict[str, Any]]: + by_ordinal = command_by_ordinal(commands) + read_producers = producers_by_exact_key(commands) + candidates: list[dict[str, Any]] = [] + + attention_occurrences = [] + attention_chain_ns = 0.0 + attention_removable_ns = 0.0 + for command in commands: + if command["kernel"] != "loom_libs:ggml_flash_attention_decode_split_f32_f16_wmma_next_q8": + continue + ordinal = int(command["ordinal"]) + rope = by_ordinal.get(ordinal - 3) + rope_set_rows = by_ordinal.get(ordinal - 2) + set_rows = by_ordinal.get(ordinal - 1) + q_projection = producer_for_read(rope, "input", read_producers) if rope else None + k_projection = producer_for_read(rope_set_rows, "input", read_producers) if rope_set_rows else None + v_projection = producer_for_read(set_rows, "rows", read_producers) if set_rows else None + matched = ( + rope is not None + and rope_set_rows is not None + and set_rows is not None + and rope["kernel"] == "loom_libs:ggml_rope_f32" + and rope_set_rows["kernel"] == "loom_libs:ggml_rope_set_rows_f32" + and set_rows["kernel"] == "loom_libs:ggml_set_rows" + and q_projection is not None + and k_projection is not None + and v_projection is not None + and q_projection.kernel == "loom_libs:ggml_mul_mat_f32_f32_decode_wave64" + and k_projection.kernel == "loom_libs:ggml_mul_mat_f32_f32_decode_wave64" + and v_projection.kernel == "loom_libs:ggml_mul_mat_f32_f32_decode_wave64" + ) + if not matched: + continue + occurrence_ns = sum(command_metric_ns(item, shape_info) for item in (rope, rope_set_rows, set_rows)) + attention_chain_ns += occurrence_ns + command_metric_ns(command, shape_info) + attention_removable_ns += occurrence_ns + attention_occurrences.append( + { + "attention_ordinal": ordinal, + "q_projection_ordinal": q_projection.ordinal, + "q_rope_ordinal": int(rope["ordinal"]), + "k_projection_ordinal": k_projection.ordinal, + "k_rope_set_rows_ordinal": int(rope_set_rows["ordinal"]), + "v_projection_ordinal": v_projection.ordinal, + "v_set_rows_ordinal": int(set_rows["ordinal"]), + "postprocess_ns": occurrence_ns, + "head": { + "query_heads": compile_parameters(command).get("ggml.flash_attention.query_head_count"), + "key_value_heads": compile_parameters(command).get("ggml.flash_attention.key_value_head_count"), + "qk_head_size": compile_parameters(command).get("ggml.flash_attention.qk_head_size"), + "value_head_size": compile_parameters(command).get("ggml.flash_attention.value_head_size"), + }, + } + ) + add_candidate( + candidates, + "decode attention projection postprocess", + "attention_qkv_decode", + attention_occurrences, + attention_chain_ns, + attention_removable_ns, + [ + "Current decode path uses standalone Q rope, K rope+cache write, and V cache write before split flash attention.", + "Existing llm_attention_qkv fused projection matchers require token_count > 1, so this looks like decode-route coverage is missing.", + ], + ) + + swiglu_occurrences = [] + swiglu_chain_ns = 0.0 + swiglu_removable_ns = 0.0 + residual_occurrences = [] + residual_chain_ns = 0.0 + residual_removable_ns = 0.0 + tail_occurrences = [] + tail_chain_ns = 0.0 + tail_removable_ns = 0.0 + for command in commands: + if command["kernel"] != "loom_libs:ggml_binary_f32": + continue + ordinal = int(command["ordinal"]) + current_ns = command_metric_ns(command, shape_info) + consumers = output_consumers(command, fanout) + lhs = producer_for_read(command, "lhs", read_producers) + rhs = producer_for_read(command, "rhs", read_producers) + op = op_name(command) + if op == "swiglu": + down_projection = consumers[0] if len(consumers) == 1 else None + occurrence = { + "binary_ordinal": ordinal, + "lhs_producer": lhs.ordinal if lhs else None, + "rhs_producer": rhs.ordinal if rhs else None, + "consumer": down_projection.ordinal if down_projection else None, + "element_count": integer_parameters(command).get("element_count"), + "binary_ns": current_ns, + } + swiglu_occurrences.append(occurrence) + swiglu_removable_ns += current_ns + swiglu_chain_ns += current_ns + if lhs is not None: + swiglu_chain_ns += command_metric_ns(by_ordinal[lhs.ordinal], shape_info) + if rhs is not None: + swiglu_chain_ns += command_metric_ns(by_ordinal[rhs.ordinal], shape_info) + if down_projection is not None: + swiglu_chain_ns += command_metric_ns(by_ordinal[down_projection.ordinal], shape_info) + elif op == "add": + consumer_kernels = [consumer.kernel for consumer in consumers] + occurrence = { + "binary_ordinal": ordinal, + "lhs_producer": lhs.ordinal if lhs else None, + "rhs_producer": rhs.ordinal if rhs else None, + "consumers": [ + { + "ordinal": consumer.ordinal, + "kernel": consumer.kernel, + "binding": consumer.binding, + } + for consumer in consumers + ], + "element_count": integer_parameters(command).get("element_count"), + "binary_ns": current_ns, + } + if "loom_libs:ggml_rmsnorm_binary_f32" in consumer_kernels: + residual_occurrences.append(occurrence) + residual_removable_ns += current_ns + residual_chain_ns += current_ns + for consumer in consumers: + if consumer.kernel == "loom_libs:ggml_rmsnorm_binary_f32": + residual_chain_ns += command_metric_ns(by_ordinal[consumer.ordinal], shape_info) + if any(consumer.kernel == "loom_libs:ggml_rmsnorm_binary_f32" for consumer in consumers) and ( + (lhs and lhs.kernel == "hrx:ggml_gather_add_f32") or (rhs and rhs.kernel == "hrx:ggml_gather_add_f32") + ): + tail_occurrences.append(occurrence) + tail_removable_ns += current_ns + tail_chain_ns += current_ns + for consumer in consumers: + tail_chain_ns += command_metric_ns(by_ordinal[consumer.ordinal], shape_info) + + add_candidate( + candidates, + "decode SwiGLU MLP", + "swiglu_decode", + swiglu_occurrences, + swiglu_chain_ns, + swiglu_removable_ns, + [ + "Pattern is two decode matmuls feeding binary_f32 op=4, then a down projection.", + "This is separate from existing binary_swiglu_symmetric_i4_k32 coverage.", + ], + ) + add_candidate( + candidates, + "residual add next RMS norm", + "add_next_rmsnorm_f32", + residual_occurrences, + residual_chain_ns, + residual_removable_ns, + [ + "Pattern is binary_f32 op=0 feeding rmsnorm_binary_f32.", + "Some add outputs have multiple consumers, so the fused kernel must still publish the residual/add output when it is reused.", + ], + ) + add_candidate( + candidates, + "final output tail add norm", + "output_tail", + tail_occurrences, + tail_chain_ns, + tail_removable_ns, + [ + "Low-count final path involving gather_add, binary add, RMS norm, and output projection.", + ], + ) + + matmul_add_norm_occurrences = [] + matmul_add_norm_chain_ns = 0.0 + for command in commands: + if command["kernel"] != "loom_libs:ggml_add_rmsnorm_binary_symmetric_i4_k32": + continue + lhs = producer_for_read(command, "lhs", read_producers) + if lhs is None or lhs.kernel != "loom_libs:ggml_mul_mat_f32_f32_decode_wave64": + continue + ordinal = int(command["ordinal"]) + current_ns = command_metric_ns(command, shape_info) + lhs_ns = command_metric_ns(by_ordinal[lhs.ordinal], shape_info) + matmul_add_norm_chain_ns += current_ns + lhs_ns + matmul_add_norm_occurrences.append( + { + "add_rmsnorm_ordinal": ordinal, + "matmul_ordinal": lhs.ordinal, + "hidden_size": compile_parameters(command).get("ggml.add_rmsnorm_binary_symmetric_i4.hidden_size"), + "chain_ns": current_ns + lhs_ns, + } + ) + add_candidate( + candidates, + "decode matmul add RMS norm", + "matmul_add_next_rmsnorm_decode", + matmul_add_norm_occurrences, + matmul_add_norm_chain_ns, + 0.0, + [ + "Pattern is decode matmul feeding add_rmsnorm_binary_symmetric_i4.", + "Savings depend on a decode matmul epilogue or fused next-rmsnorm route, so removable standalone time is not estimated here.", + ], + ) + + total_removable = sum(candidate["removable_standalone_ns"] for candidate in candidates) + for candidate in candidates: + candidate["current_chain_ms"] = candidate["current_chain_ns"] / 1_000_000.0 + candidate["removable_standalone_ms"] = candidate["removable_standalone_ns"] / 1_000_000.0 + candidate["removable_share_of_candidates"] = ( + candidate["removable_standalone_ns"] / total_removable * 100.0 if total_removable else 0.0 + ) + return sorted(candidates, key=lambda item: (item["removable_standalone_ns"], item["current_chain_ns"]), reverse=True) + + +def sequential_triples(commands: list[dict[str, Any]], limit: int = 30) -> list[dict[str, Any]]: + counts: Counter[tuple[str, str, str]] = Counter() + ordered = sorted(commands, key=lambda item: int(item["ordinal"])) + for first, second, third in zip(ordered, ordered[1:], ordered[2:], strict=False): + counts[(first["kernel"], second["kernel"], third["kernel"])] += 1 + return [ + {"kernels": list(kernels), "count": count} + for kernels, count in counts.most_common(limit) + ] + + +def build_report(commands: list[dict[str, Any]], + manifest: dict[str, Any], + summary: dict[str, Any]) -> dict[str, Any]: + shape_info, warnings = map_shape_keys_to_dispatches(commands, manifest, summary) + edges, fanout = exact_transient_edges(commands) + candidates = classify_candidates(commands, shape_info, fanout) + total_weighted_ns = float(summary.get("total_weighted_ns", 0.0)) + for candidate in candidates: + candidate["removable_share_of_total_modelled_kernel_time"] = ( + candidate["removable_standalone_ns"] / total_weighted_ns * 100.0 if total_weighted_ns else 0.0 + ) + candidate["current_chain_share_of_total_modelled_kernel_time"] = ( + candidate["current_chain_ns"] / total_weighted_ns * 100.0 if total_weighted_ns else 0.0 + ) + return { + "schema": "ggml-hrx-model-fusion-adjacency-v1", + "model": manifest.get("model"), + "scenario": manifest.get("scenario"), + "command_count": len(commands), + "dispatch_count": manifest.get("dispatch_count"), + "generated_count": manifest.get("generated_count"), + "summary_metric": summary.get("metric_description"), + "total_weighted_ns": total_weighted_ns, + "warnings": warnings, + "candidate_fusions": candidates, + "exact_transient_edge_summary": summarize_edges(edges), + "fanout_summary": summarize_fanout(fanout), + "sequential_triples": sequential_triples(commands), + } + + +def md_escape(text: Any) -> str: + return str(text).replace("|", "\\|") + + +def format_ms(ns: float) -> str: + return f"{ns / 1_000_000.0:.3f}" + + +def write_markdown(path: Path, report: dict[str, Any]) -> None: + lines = [ + "# HRX Model Fusion Adjacency", + "", + f"Model: `{report.get('model')}`", + f"Scenario: `{report.get('scenario')}`", + f"Commands: `{report.get('command_count')}`", + f"Generated dispatch shapes: `{report.get('generated_count')} / {report.get('dispatch_count')}`", + f"Timing metric: `{report.get('summary_metric')}`", + f"Total weighted modeled kernel time: `{format_ms(float(report.get('total_weighted_ns', 0.0)))} ms`", + "", + ] + if report.get("warnings"): + lines.append("## Warnings") + lines.append("") + for warning in report["warnings"]: + lines.append(f"- {warning}") + lines.append("") + + lines.extend( + [ + "## Candidate Fusions", + "", + "| Candidate | Occurrences | Current chain ms | Removable standalone ms | Current chain share | Removable share | Notes |", + "| --- | ---: | ---: | ---: | ---: | ---: | --- |", + ] + ) + for candidate in report["candidate_fusions"]: + notes = "
".join(md_escape(note) for note in candidate.get("notes", [])) + lines.append( + f"| {md_escape(candidate['name'])} | {candidate['occurrence_count']} | " + f"{candidate['current_chain_ms']:.3f} | {candidate['removable_standalone_ms']:.3f} | " + f"{candidate['current_chain_share_of_total_modelled_kernel_time']:.2f}% | " + f"{candidate['removable_share_of_total_modelled_kernel_time']:.2f}% | {notes} |" + ) + + lines.extend( + [ + "", + "## Exact Transient Producer-Consumer Edges", + "", + "| Producer | Consumer | Edges | Example ordinals |", + "| --- | --- | ---: | --- |", + ] + ) + for edge in report["exact_transient_edge_summary"][:30]: + examples = ", ".join( + f"{example['producer_ordinal']}:{example['producer_binding']}->{example['consumer_ordinal']}:{example['consumer_binding']}" + for example in edge.get("examples", [])[:3] + ) + lines.append( + f"| `{edge['producer_kernel']}` | `{edge['consumer_kernel']}` | {edge['edge_count']} | {md_escape(examples)} |" + ) + + lines.extend( + [ + "", + "## Fanout Patterns", + "", + "| Producer | Consumers | Produced values |", + "| --- | --- | ---: |", + ] + ) + for row in report["fanout_summary"][:25]: + consumers = ", ".join(f"`{kernel}`" for kernel in row["consumer_kernels"]) + lines.append(f"| `{row['producer_kernel']}` | {consumers} | {row['producer_count']} |") + + lines.extend( + [ + "", + "## Top Sequential Triples", + "", + "| Kernels | Count |", + "| --- | ---: |", + ] + ) + for triple in report["sequential_triples"][:20]: + kernels = " -> ".join(f"`{kernel}`" for kernel in triple["kernels"]) + lines.append(f"| {kernels} | {triple['count']} |") + + lines.extend( + [ + "", + "## Candidate Examples", + "", + ] + ) + for candidate in report["candidate_fusions"]: + lines.append(f"### {candidate['name']}") + lines.append("") + if not candidate.get("examples"): + lines.append("No matching occurrences.") + lines.append("") + continue + lines.append("```json") + lines.append(json.dumps(candidate["examples"][:3], indent=2, sort_keys=True)) + lines.append("```") + lines.append("") + + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text("\n".join(lines).rstrip() + "\n", encoding="utf-8") + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--dump-dir", type=Path, required=True, help="HRX command-program dump directory.") + parser.add_argument("--scenario-manifest", type=Path, required=True, help="Generated model scenario manifest.") + parser.add_argument("--summary-json", type=Path, required=True, help="Summarized Loom benchmark JSON.") + parser.add_argument("--output-json", type=Path, required=True, help="Fusion adjacency JSON output.") + parser.add_argument("--output-md", type=Path, required=True, help="Fusion adjacency Markdown output.") + args = parser.parse_args() + + commands = load_commands(args.dump_dir) + manifest = load_json(args.scenario_manifest) + summary = load_json(args.summary_json) + report = build_report(commands, manifest, summary) + write_json(args.output_json, report) + write_markdown(args.output_md, report) + print(f"wrote {args.output_json}") + print(f"wrote {args.output_md}") + + +if __name__ == "__main__": + main() diff --git a/ggml/src/ggml-hrx/tools/benchmarks/generate-model-benchmarks.py b/ggml/src/ggml-hrx/tools/benchmarks/generate-model-benchmarks.py new file mode 100755 index 000000000000..a94bea2d966f --- /dev/null +++ b/ggml/src/ggml-hrx/tools/benchmarks/generate-model-benchmarks.py @@ -0,0 +1,1030 @@ +#!/usr/bin/env python3 +# +# Generate model-scoped Loom benchmarks from HRX command program dumps. + +from __future__ import annotations + +import argparse +import hashlib +import json +import re +from collections import Counter +from dataclasses import dataclass +from pathlib import Path +from typing import Any + + +SCRIPT_DIR = Path(__file__).resolve().parent +HRX_DIR = SCRIPT_DIR.parents[1] +BENCHMARK_DIR = HRX_DIR / "benchmarks" +KERNEL_CORPUS_DIR = HRX_DIR / "kernel-corpus" / "kernels" +KERNEL_TARGET_RE = re.compile( + r"kernel\.def\s+target\(@(?P[A-Za-z0-9_.$-]+)\)" + r"(?:\s+export\(\"[^\"]+\"\))?\s+@(?P[A-Za-z0-9_.$-]+)\s*\(" +) + + +@dataclass +class ScenarioBenchmarks: + scenario: str + command_count: int + dispatch_count: int + generated_count: int + kernel_counts: dict[str, int] + dispatches: list[dict[str, Any]] + cases: list[str] + used_exports: dict[str, dict[str, Any]] + + +def fail(message: str) -> None: + raise SystemExit(message) + + +def load_json(path: Path) -> Any: + with path.open("r", encoding="utf-8") as f: + return json.load(f) + + +def read_text(path: Path) -> str: + return path.read_text(encoding="utf-8") + + +def write_json(path: Path, data: Any) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", encoding="utf-8") as f: + json.dump(data, f, indent=2, sort_keys=True) + f.write("\n") + + +def sanitize_symbol(text: str) -> str: + text = text.split(":")[-1] + text = re.sub(r"[^A-Za-z0-9_]+", "_", text) + text = re.sub(r"_+", "_", text).strip("_") + if not text: + return "unknown" + if text[0].isdigit(): + return "k_" + text + return text + + +def int_param(command: dict[str, Any], name: str, default: int | None = None) -> int: + params = command.get("integer_parameters", {}) + if name in params: + return int(params[name]) + if default is not None: + return default + fail(f"command {command.get('ordinal')} {command.get('kernel')} is missing integer parameter {name}") + + +def config_int(command: dict[str, Any], name: str, default: int | None = None) -> int: + params = command.get("compile_parameters", {}) + if name in params: + return int(params[name]) + if default is not None: + return default + fail(f"command {command.get('ordinal')} {command.get('kernel')} is missing compile parameter {name}") + + +def binding_length(command: dict[str, Any], name: str) -> int: + for binding in command.get("bindings", []): + if binding.get("name") == name: + return int(binding["length"]) + fail(f"command {command.get('ordinal')} {command.get('kernel')} is missing binding {name}") + + +def ceil_div(value: int, divisor: int) -> int: + return (value + divisor - 1) // divisor + + +def tensor_type_for_format(format_value: int) -> str | None: + if format_value == 16: + return "f16" + if format_value == 30: + return "bf16" + if format_value == 32: + return "f32" + return None + + +def fill_tensor(name: str, value: str, shape: str, dtype: str) -> str: + return f" %{name} = check.generate.fill value({value}) : tensor<{shape}x{dtype}>" + + +def iota_tensor(name: str, offset: str, step: str, shape: str, dtype: str, period: int | None = None) -> str: + period_text = "" if period is None else f" period({period})" + return f" %{name} = check.generate.iota offset({offset}) step({step}){period_text} : tensor<{shape}x{dtype}>" + + +def case_binary(symbol: str, command: dict[str, Any]) -> str: + count = int_param(command, "element_count") + return "\n".join( + [ + f"check.case public @{symbol}_case {{", + f" %element_count = check.literal value({count}) : index", + fill_tensor("lhs", "2.0", f"{count}", "f32"), + fill_tensor("rhs", "3.0", f"{count}", "f32"), + fill_tensor("output", "0.0", f"{count}", "f32"), + f" kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<{count}xf32>, tensor<{count}xf32>, tensor<{count}xf32>)", + " check.return", + "}", + ] + ) + + +def case_rmsnorm_binary(symbol: str, command: dict[str, Any]) -> str: + token_count = int_param(command, "token_count") + hidden_size = config_int(command, "ggml.rmsnorm_binary_f32.hidden_size") + shape = f"{token_count}x{hidden_size}" + return "\n".join( + [ + f"check.case public @{symbol}_case {{", + f" %token_count = check.literal value({token_count}) : index", + fill_tensor("input", "2.0", shape, "f32"), + fill_tensor("rhs", "3.0", f"{hidden_size}", "f32"), + fill_tensor("output", "0.0", shape, "f32"), + f" kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<{shape}xf32>, tensor<{hidden_size}xf32>, tensor<{shape}xf32>)", + " check.return", + "}", + ] + ) + + +def case_mul_mat(symbol: str, command: dict[str, Any], kernel: str) -> str | None: + token_count = int_param(command, "token_count") + if kernel == "ggml_mul_mat_f32_f32_decode_wave64": + input_size = int_param(command, "input_size") + output_size = int_param(command, "output_size") + accumulation = 0 + weight_format = config_int(command, "ggml.mul_mat_f32_f32_decode.weight_format") + else: + input_size = config_int(command, "ggml.mul_mat.input_size") + output_size = config_int(command, "ggml.mul_mat.output_size") + accumulation = config_int(command, "ggml.mul_mat.output_accumulation", 0) + weight_format = config_int(command, "ggml.mul_mat.weight_format") + weight_type = tensor_type_for_format(weight_format) + if weight_type is None: + return None + shape_in = f"{token_count}x{input_size}" + shape_weight = f"{output_size}x{input_size}" + shape_out = f"{token_count}x{output_size}" + output_init = 1.0 if accumulation else 0.0 + if kernel == "ggml_mul_mat_f32_f32_decode_wave64": + launch_params = "%token_count, %input_size, %output_size" + workload_params = "%token_count, %input_size, %output_size" + workload_types = "[index, index, index]" + launch_types = "index, index, index" + literals = [ + f" %token_count = check.literal value({token_count}) : index", + f" %input_size = check.literal value({input_size}) : index", + f" %output_size = check.literal value({output_size}) : index", + ] + else: + launch_params = "%token_count" + workload_params = "%token_count" + workload_types = "[index]" + launch_types = "index" + literals = [f" %token_count = check.literal value({token_count}) : index"] + return "\n".join( + [ + f"check.case public @{symbol}_case {{", + *literals, + fill_tensor("input", "1.0", shape_in, "f32"), + fill_tensor("weight", "1.0", shape_weight, weight_type), + fill_tensor("output", f"{output_init:.1f}", shape_out, "f32"), + f" kernel.launch @{kernel}[{workload_params}]({launch_params}, %input, %weight, %output) : {workload_types}({launch_types}, tensor<{shape_in}xf32>, tensor<{shape_weight}x{weight_type}>, tensor<{shape_out}xf32>)", + " check.return", + "}", + ] + ) + + +def case_mul_mat_add_decode(symbol: str, command: dict[str, Any], kernel: str) -> str | None: + token_count = int_param(command, "token_count") + input_size = int_param(command, "input_size") + output_size = int_param(command, "output_size") + weight_format = config_int(command, "ggml.mul_mat_f32_f32_decode.weight_format") + weight_type = tensor_type_for_format(weight_format) + if weight_type is None: + return None + + shape_in = f"{token_count}x{input_size}" + shape_weight = f"{output_size}x{input_size}" + shape_out = f"{token_count}x{output_size}" + return "\n".join( + [ + f"check.case public @{symbol}_case {{", + f" %token_count = check.literal value({token_count}) : index", + f" %input_size = check.literal value({input_size}) : index", + f" %output_size = check.literal value({output_size}) : index", + fill_tensor("input", "1.0", shape_in, "f32"), + fill_tensor("weight", "1.0", shape_weight, weight_type), + fill_tensor("residual_input", "0.25", shape_out, "f32"), + fill_tensor("residual_output", "0.0", shape_out, "f32"), + f" kernel.launch @{kernel}[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %residual_input, %residual_output) : [index, index, index](index, index, index, tensor<{shape_in}xf32>, tensor<{shape_weight}x{weight_type}>, tensor<{shape_out}xf32>, tensor<{shape_out}xf32>)", + " check.return", + "}", + ] + ) + + +def case_mul_mat_postops(symbol: str, command: dict[str, Any], kernel: str) -> str | None: + token_count = int_param(command, "token_count") + input_size = config_int(command, "ggml.mul_mat_postops.input_size") + output_size = config_int(command, "ggml.mul_mat_postops.output_size") + weight_format = config_int(command, "ggml.mul_mat_postops.weight_format") + weight_type = tensor_type_for_format(weight_format) + if weight_type is None: + return None + + has_bias = "bias" in kernel + has_residual = "add" in kernel + has_rmsnorm = "next_rmsnorm" in kernel + shape_in = f"{token_count}x{input_size}" + shape_weight = f"{output_size}x{input_size}" + shape_out = f"{token_count}x{output_size}" + completion_count = (token_count + 31) // 32 + + tensors = [ + fill_tensor("input", "1.0", shape_in, "f32"), + fill_tensor("weight", "1.0", shape_weight, weight_type), + ] + launch_values = ["%token_count", "%input", "%weight"] + launch_types = ["index", f"tensor<{shape_in}xf32>", f"tensor<{shape_weight}x{weight_type}>"] + + if has_bias: + tensors.append(fill_tensor("bias", "0.5", f"{output_size}", "f32")) + launch_values.append("%bias") + launch_types.append(f"tensor<{output_size}xf32>") + if has_residual: + tensors.append(fill_tensor("residual_input", "0.25", shape_out, "f32")) + tensors.append(fill_tensor("residual_output", "0.0", shape_out, "f32")) + launch_values.extend(["%residual_input", "%residual_output"]) + launch_types.extend([f"tensor<{shape_out}xf32>", f"tensor<{shape_out}xf32>"]) + else: + tensors.append(fill_tensor("output", "0.0", shape_out, "f32")) + launch_values.append("%output") + launch_types.append(f"tensor<{shape_out}xf32>") + + if has_rmsnorm: + tensors.append(fill_tensor("norm_weight", "1.0", f"{output_size}", "f32")) + tensors.append(fill_tensor("normalized_output", "0.0", shape_out, "f32")) + tensors.append(fill_tensor("completion_counters", "0", f"{completion_count}", "i32")) + launch_values.extend(["%norm_weight", "%normalized_output", "%completion_counters"]) + launch_types.extend([f"tensor<{output_size}xf32>", f"tensor<{shape_out}xf32>", f"tensor<{completion_count}xi32>"]) + + return "\n".join( + [ + f"check.case public @{symbol}_case {{", + f" %token_count = check.literal value({token_count}) : index", + *tensors, + f" kernel.launch @{kernel}[%token_count]({', '.join(launch_values)}) : [index]({', '.join(launch_types)})", + " check.return", + "}", + ] + ) + + +def case_swiglu(symbol: str, command: dict[str, Any], kernel: str) -> str | None: + token_count = int_param(command, "token_count") + input_size = config_int(command, "ggml.mul_mat_swiglu.input_size") + output_size = config_int(command, "ggml.mul_mat_swiglu.output_size") + gate_format = config_int(command, "ggml.mul_mat_swiglu.gate_weight_format") + up_format = config_int(command, "ggml.mul_mat_swiglu.up_weight_format") + gate_type = tensor_type_for_format(gate_format) + up_type = tensor_type_for_format(up_format) + if gate_type is None or up_type is None: + return None + shape_in = f"{token_count}x{input_size}" + shape_weight = f"{output_size}x{input_size}" + shape_out = f"{token_count}x{output_size}" + return "\n".join( + [ + f"check.case public @{symbol}_case {{", + f" %token_count = check.literal value({token_count}) : index", + fill_tensor("input", "0.0", shape_in, "f32"), + fill_tensor("gate_weight", "1.0", shape_weight, gate_type), + fill_tensor("up_weight", "1.0", shape_weight, up_type), + fill_tensor("output", "1.0", shape_out, "f32"), + f" kernel.launch @{kernel}[%token_count](%token_count, %input, %gate_weight, %up_weight, %output) : [index](index, tensor<{shape_in}xf32>, tensor<{shape_weight}x{gate_type}>, tensor<{shape_weight}x{up_type}>, tensor<{shape_out}xf32>)", + " check.return", + "}", + ] + ) + + +def case_llm_attention_q_matmul_rope(symbol: str, command: dict[str, Any], kernel: str) -> str | None: + token_count = int_param(command, "token_count") + input_size = config_int(command, "llm.attention_qkv.input_size") + output_size = config_int(command, "llm.attention_qkv.output_size") + weight_format = config_int(command, "llm.attention_qkv.weight_format") + head_size = config_int(command, "llm.attention_qkv.head_size") + head_count = config_int(command, "llm.attention_qkv.head_count") + weight_type = tensor_type_for_format(weight_format) + if weight_type is None: + return None + shape_in = f"{token_count}x{input_size}" + shape_weight = f"{output_size}x{input_size}" + shape_out = f"{token_count}x{head_count}x{head_size}" + half_head_size = head_size // 2 + return "\n".join( + [ + f"check.case public @{symbol}_case {{", + f" %token_count = check.literal value({token_count}) : index", + fill_tensor("input", "1.0", shape_in, "f32"), + fill_tensor("weight", "1.0", shape_weight, weight_type), + fill_tensor("positions", "0", f"{token_count}", "i32"), + fill_tensor("theta", "0.0", f"{half_head_size}", "f32"), + fill_tensor("freq_factors", "0.0", f"{half_head_size}", "f32"), + fill_tensor("output", "0.0", shape_out, "f32"), + f" kernel.launch @{kernel}[%token_count](%token_count, %input, %weight, %positions, %theta, %freq_factors, %output) : [index](index, tensor<{shape_in}xf32>, tensor<{shape_weight}x{weight_type}>, tensor<{token_count}xi32>, tensor<{half_head_size}xf32>, tensor<{half_head_size}xf32>, tensor<{shape_out}xf32>)", + " check.return", + "}", + ] + ) + + +def case_llm_attention_k_matmul_rope_set_rows(symbol: str, command: dict[str, Any], kernel: str) -> str | None: + token_count = int_param(command, "token_count") + input_size = config_int(command, "llm.attention_qkv.input_size") + output_size = config_int(command, "llm.attention_qkv.output_size") + weight_format = config_int(command, "llm.attention_qkv.weight_format") + head_size = config_int(command, "llm.attention_qkv.head_size") + head_count = config_int(command, "llm.attention_qkv.head_count") + cache_rows = config_int(command, "llm.attention_qkv.cache_row_count") + cache_format = config_int(command, "llm.attention_qkv.cache_output_format") + weight_type = tensor_type_for_format(weight_format) + cache_type = tensor_type_for_format(cache_format) + if weight_type is None or cache_type is None: + return None + shape_in = f"{token_count}x{input_size}" + shape_weight = f"{output_size}x{input_size}" + cache_shape = f"{cache_rows}x{head_count}x{head_size}" + half_head_size = head_size // 2 + return "\n".join( + [ + f"check.case public @{symbol}_case {{", + f" %token_count = check.literal value({token_count}) : index", + fill_tensor("input", "1.0", shape_in, "f32"), + fill_tensor("weight", "1.0", shape_weight, weight_type), + fill_tensor("positions", "0", f"{token_count}", "i32"), + iota_tensor("indices", "0", "1", f"{token_count}", "i64", period=cache_rows), + fill_tensor("theta", "0.0", f"{half_head_size}", "f32"), + fill_tensor("freq_factors", "0.0", f"{half_head_size}", "f32"), + fill_tensor("cache", "0.0", cache_shape, cache_type), + f" kernel.launch @{kernel}[%token_count](%token_count, %input, %weight, %positions, %indices, %theta, %freq_factors, %cache) : [index](index, tensor<{shape_in}xf32>, tensor<{shape_weight}x{weight_type}>, tensor<{token_count}xi32>, tensor<{token_count}xi64>, tensor<{half_head_size}xf32>, tensor<{half_head_size}xf32>, tensor<{cache_shape}x{cache_type}>)", + " check.return", + "}", + ] + ) + + +def case_llm_attention_v_matmul_set_rows(symbol: str, command: dict[str, Any], kernel: str) -> str | None: + token_count = int_param(command, "token_count") + input_size = config_int(command, "llm.attention_qkv.input_size") + output_size = config_int(command, "llm.attention_qkv.output_size") + weight_format = config_int(command, "llm.attention_qkv.weight_format") + cache_rows = config_int(command, "llm.attention_qkv.cache_row_count") + cache_format = config_int(command, "llm.attention_qkv.cache_output_format") + weight_type = tensor_type_for_format(weight_format) + cache_type = tensor_type_for_format(cache_format) + if weight_type is None or cache_type is None: + return None + shape_in = f"{token_count}x{input_size}" + shape_weight = f"{output_size}x{input_size}" + cache_shape = f"{cache_rows}x{output_size}" + return "\n".join( + [ + f"check.case public @{symbol}_case {{", + f" %token_count = check.literal value({token_count}) : index", + fill_tensor("input", "1.0", shape_in, "f32"), + fill_tensor("weight", "1.0", shape_weight, weight_type), + iota_tensor("indices", "0", "1", f"{token_count}", "i64", period=cache_rows), + fill_tensor("cache", "0.0", cache_shape, cache_type), + f" kernel.launch @{kernel}[%token_count](%token_count, %input, %weight, %indices, %cache) : [index](index, tensor<{shape_in}xf32>, tensor<{shape_weight}x{weight_type}>, tensor<{token_count}xi64>, tensor<{cache_shape}x{cache_type}>)", + " check.return", + "}", + ] + ) + + +def case_flash_attention(symbol: str, command: dict[str, Any], kernel: str, prefix: str) -> str: + query_tokens = int_param(command, "query_token_count") + kv_tokens = int_param(command, "key_value_token_count") + if prefix == "ggml": + query_heads = config_int(command, "ggml.flash_attention.query_head_count") + kv_heads = config_int(command, "ggml.flash_attention.key_value_head_count") + legacy_head_size = config_int(command, "ggml.flash_attention.head_size", 0) + qk_head_size = config_int(command, "ggml.flash_attention.qk_head_size", legacy_head_size) + value_head_size = config_int(command, "ggml.flash_attention.value_head_size", qk_head_size) + else: + query_heads = config_int(command, f"{prefix}.attention.query_head_count") + kv_heads = config_int(command, f"{prefix}.attention.key_value_head_count") + qk_head_size = 128 + value_head_size = 128 + query_shape = f"{query_tokens}x{query_heads}x{qk_head_size}" + key_shape = f"{kv_tokens}x{kv_heads}x{qk_head_size}" + value_shape = f"{kv_tokens}x{kv_heads}x{value_head_size}" + output_shape = f"{query_tokens}x{query_heads}x{value_head_size}" + mask_shape = f"{query_tokens}x{kv_tokens}" + if prefix == "ggml": + launch = ( + f" kernel.launch @{kernel}[%query_token_count, %key_value_token_count]" + f"(%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %query, %output)" + f" : [index, index](index, index, tensor<{query_shape}xf32>, tensor<{key_shape}xf16>, " + f"tensor<{value_shape}xf16>, tensor<{mask_shape}xf16>, tensor<{query_shape}xf32>, " + f"tensor<{output_shape}xf32>)" + ) + else: + launch = ( + f" kernel.launch @{kernel}[%query_token_count, %key_value_token_count]" + f"(%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output)" + f" : [index, index](index, index, tensor<{query_shape}xf32>, tensor<{key_shape}xf16>, " + f"tensor<{value_shape}xf16>, tensor<{mask_shape}xf16>, tensor<{output_shape}xf32>)" + ) + return "\n".join( + [ + f"check.case public @{symbol}_case {{", + f" %query_token_count = check.literal value({query_tokens}) : index", + f" %key_value_token_count = check.literal value({kv_tokens}) : index", + fill_tensor("query", "0.0", query_shape, "f32"), + fill_tensor("key", "0.0", key_shape, "f16"), + fill_tensor("value", "0.0", value_shape, "f16"), + fill_tensor("mask", "0.0", mask_shape, "f16"), + fill_tensor("output", "1.0", output_shape, "f32"), + launch, + " check.return", + "}", + ] + ) + + +def case_flash_attention_decode_split(symbol: str, command: dict[str, Any], kernel: str) -> str: + kv_tokens = int_param(command, "key_value_token_count") + kv_capacity = config_int(command, "ggml.flash_attention.decode.key_value_token_capacity") + query_heads = config_int(command, "ggml.flash_attention.query_head_count") + kv_heads = config_int(command, "ggml.flash_attention.key_value_head_count") + legacy_head_size = config_int(command, "ggml.flash_attention.head_size", 0) + qk_head_size = config_int(command, "ggml.flash_attention.qk_head_size", legacy_head_size) + value_head_size = config_int(command, "ggml.flash_attention.value_head_size", qk_head_size) + partial_blocks = ceil_div(kv_capacity, 64) + query_shape = f"{query_heads}x{qk_head_size}" + key_shape = f"{kv_tokens}x{kv_heads}x{qk_head_size}" + value_shape = f"{kv_tokens}x{kv_heads}x{value_head_size}" + mask_shape = f"{kv_tokens}" + partial_shape = f"{kv_heads}x{partial_blocks}x16" + partial_output_shape = f"{kv_heads}x{partial_blocks}x16x{value_head_size}" + completion_shape = f"{kv_heads}" + output_shape = f"{query_heads}x{value_head_size}" + tensors = [ + fill_tensor("query", "0.0", query_shape, "f32"), + fill_tensor("key", "0.0", key_shape, "f16"), + fill_tensor("value", "0.0", value_shape, "f16"), + fill_tensor("mask", "0.0", mask_shape, "f16"), + fill_tensor("partial_max", "0.0", partial_shape, "f32"), + fill_tensor("partial_sum", "0.0", partial_shape, "f32"), + fill_tensor("partial_output", "0.0", partial_output_shape, "f16"), + fill_tensor("completion_counter", "0", completion_shape, "i32"), + fill_tensor("output", "1.0", output_shape, "f32"), + ] + launch_values = [ + "%key_value_token_count", + "%query", + "%key", + "%value", + "%mask", + "%partial_max", + "%partial_sum", + "%partial_output", + "%completion_counter", + "%output", + ] + launch_types = [ + "index", + f"tensor<{query_shape}xf32>", + f"tensor<{key_shape}xf16>", + f"tensor<{value_shape}xf16>", + f"tensor<{mask_shape}xf16>", + f"tensor<{partial_shape}xf32>", + f"tensor<{partial_shape}xf32>", + f"tensor<{partial_output_shape}xf16>", + f"tensor<{completion_shape}xi32>", + f"tensor<{output_shape}xf32>", + ] + if kernel.endswith("_next_q8"): + next_q8_bytes = binding_length(command, "next_q8_output") + tensors.append(fill_tensor("next_q8_output", "0", f"{next_q8_bytes}", "i8")) + launch_values.append("%next_q8_output") + launch_types.append(f"tensor<{next_q8_bytes}xi8>") + return "\n".join( + [ + f"check.case public @{symbol}_case {{", + f" %key_value_token_count = check.literal value({kv_tokens}) : index", + *tensors, + f" kernel.launch @{kernel}[%key_value_token_count]({', '.join(launch_values)}) : [index]({', '.join(launch_types)})", + " check.return", + "}", + ] + ) + + +def case_rope(symbol: str, command: dict[str, Any], kernel: str, prefix: str) -> str: + token_count = int_param(command, "token_count") + head_size = config_int(command, f"{prefix}.head_size") + head_count = config_int(command, f"{prefix}.head_count") + half_head_size = head_size // 2 + data_shape = f"{token_count}x{head_count}x{head_size}" + return "\n".join( + [ + f"check.case public @{symbol}_case {{", + f" %token_count = check.literal value({token_count}) : index", + fill_tensor("positions", "0", f"{token_count}", "i32"), + fill_tensor("input", "0.0", data_shape, "f32"), + fill_tensor("theta", "0.0", f"{half_head_size}", "f32"), + fill_tensor("freq_factors", "0.0", f"{half_head_size}", "f32"), + fill_tensor("output", "1.0", data_shape, "f32"), + f" kernel.launch @{kernel}[%token_count](%token_count, %positions, %input, %theta, %freq_factors, %output) : [index](index, tensor<{token_count}xi32>, tensor<{data_shape}xf32>, tensor<{half_head_size}xf32>, tensor<{half_head_size}xf32>, tensor<{data_shape}xf32>)", + " check.return", + "}", + ] + ) + + +def case_rope_set_rows(symbol: str, command: dict[str, Any]) -> str | None: + token_count = int_param(command, "token_count") + cache_rows = int_param(command, "cache_row_count") + head_size = config_int(command, "ggml.rope_set_rows_f32.head_size") + head_count = config_int(command, "ggml.rope_set_rows_f32.head_count") + output_format = config_int(command, "ggml.rope_set_rows_f32.output_format") + output_type = tensor_type_for_format(output_format) + if output_type is None: + return None + half_head_size = head_size // 2 + input_shape = f"{token_count}x{head_count}x{head_size}" + cache_shape = f"{cache_rows}x{head_count}x{head_size}" + return "\n".join( + [ + f"check.case public @{symbol}_case {{", + f" %token_count = check.literal value({token_count}) : index", + f" %cache_row_count = check.literal value({cache_rows}) : index", + fill_tensor("positions", "0", f"{token_count}", "i32"), + iota_tensor("indices", "0", "1", f"{token_count}", "i64", period=cache_rows), + fill_tensor("input", "0.0", input_shape, "f32"), + fill_tensor("theta", "0.0", f"{half_head_size}", "f32"), + fill_tensor("freq_factors", "0.0", f"{half_head_size}", "f32"), + fill_tensor("cache", "0.0", cache_shape, output_type), + f" kernel.launch @ggml_rope_set_rows_f32[%token_count, %cache_row_count](%token_count, %cache_row_count, %positions, %indices, %input, %theta, %freq_factors, %cache) : [index, index](index, index, tensor<{token_count}xi32>, tensor<{token_count}xi64>, tensor<{input_shape}xf32>, tensor<{half_head_size}xf32>, tensor<{half_head_size}xf32>, tensor<{cache_shape}x{output_type}>)", + " check.return", + "}", + ] + ) + + +def case_set_rows(symbol: str, command: dict[str, Any]) -> str | None: + token_count = int_param(command, "token_count") + cache_rows = int_param(command, "cache_row_count") + hidden_size = int_param(command, "hidden_size") + input_format = config_int(command, "ggml.set_rows.input_format") + output_format = config_int(command, "ggml.set_rows.output_format") + input_type = tensor_type_for_format(input_format) + output_type = tensor_type_for_format(output_format) + if input_type is None or output_type is None: + return None + rows_shape = f"{token_count}x{hidden_size}" + cache_shape = f"{cache_rows}x{hidden_size}" + return "\n".join( + [ + f"check.case public @{symbol}_case {{", + f" %token_count = check.literal value({token_count}) : index", + f" %cache_row_count = check.literal value({cache_rows}) : index", + f" %hidden_size = check.literal value({hidden_size}) : index", + fill_tensor("rows", "0.0", rows_shape, input_type), + iota_tensor("indices", "0", "1", f"{token_count}", "i64", period=cache_rows), + fill_tensor("cache", "0.0", cache_shape, output_type), + f" kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<{rows_shape}x{input_type}>, tensor<{token_count}xi64>, tensor<{cache_shape}x{output_type}>)", + " check.return", + "}", + ] + ) + + +def case_get_rows(symbol: str, command: dict[str, Any]) -> str | None: + token_count = int_param(command, "token_count") + row_count = int_param(command, "row_count") + hidden_size = int_param(command, "hidden_size") + weight_format = config_int(command, "ggml.get_rows_f32.weight_format") + weight_type = tensor_type_for_format(weight_format) + if weight_type is None: + return None + output_shape = f"{token_count}x{hidden_size}" + weight_shape = f"{row_count}x{hidden_size}" + return "\n".join( + [ + f"check.case public @{symbol}_case {{", + f" %token_count = check.literal value({token_count}) : index", + f" %row_count = check.literal value({row_count}) : index", + f" %hidden_size = check.literal value({hidden_size}) : index", + fill_tensor("token_ids", "0", f"{token_count}", "i32"), + fill_tensor("weight", "0.0", weight_shape, weight_type), + fill_tensor("output", "1.0", output_shape, "f32"), + f" kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<{token_count}xi32>, tensor<{weight_shape}x{weight_type}>, tensor<{output_shape}xf32>)", + " check.return", + "}", + ] + ) + + +def case_gather_add(symbol: str, command: dict[str, Any]) -> str: + source_tokens = int_param(command, "source_token_count") + output_tokens = int_param(command, "output_token_count") + hidden_size = int_param(command, "hidden_size") + source_shape = f"{source_tokens}x{hidden_size}" + output_shape = f"{output_tokens}x{hidden_size}" + return "\n".join( + [ + f"check.case public @{symbol}_case {{", + f" %source_token_count = check.literal value({source_tokens}) : index", + f" %output_token_count = check.literal value({output_tokens}) : index", + f" %hidden_size = check.literal value({hidden_size}) : index", + fill_tensor("attention", "0.0", source_shape, "f32"), + fill_tensor("residual", "0.0", source_shape, "f32"), + fill_tensor("output_ids", "0", f"{output_tokens}", "i32"), + fill_tensor("output", "1.0", output_shape, "f32"), + f" kernel.launch @ggml_gather_add_f32[%source_token_count, %output_token_count, %hidden_size](%source_token_count, %output_token_count, %hidden_size, %attention, %residual, %output_ids, %output) : [index, index, index](index, index, index, tensor<{source_shape}xf32>, tensor<{source_shape}xf32>, tensor<{output_tokens}xi32>, tensor<{output_shape}xf32>)", + " check.return", + "}", + ] + ) + + +def case_generic(symbol: str, command: dict[str, Any], export: dict[str, Any]) -> str | None: + integer_parameters = command.get("integer_parameters", {}) + binding_lengths = { + binding.get("name"): int(binding["length"]) + for binding in command.get("bindings", []) + if binding.get("name") is not None and "length" in binding + } + workload_parameters = list(export.get("workload_parameters", [])) + launch_parameters = list(export.get("launch_parameters", [])) + bindings = list(export.get("bindings", [])) + + scalar_names: list[str] = [] + for parameter in [*workload_parameters, *launch_parameters]: + name = parameter["name"] + if name not in scalar_names: + scalar_names.append(name) + missing_scalars = [name for name in scalar_names if name not in integer_parameters] + missing_bindings = [name for name in bindings if name not in binding_lengths] + if missing_scalars or missing_bindings: + return None + + literals = [ + f" %{name} = check.literal value({int(integer_parameters[name])}) : index" + for name in scalar_names + ] + tensors = [ + fill_tensor(name, "0", f"{binding_lengths[name]}", "i8") + for name in bindings + ] + workload_values = ", ".join(f"%{parameter['name']}" for parameter in workload_parameters) + launch_values = [f"%{parameter['name']}" for parameter in launch_parameters] + launch_values.extend(f"%{name}" for name in bindings) + workload_types = ", ".join("index" for _ in workload_parameters) + launch_types = ["index" for _ in launch_parameters] + launch_types.extend(f"tensor<{binding_lengths[name]}xi8>" for name in bindings) + return "\n".join( + [ + f"check.case public @{symbol}_case {{", + *literals, + *tensors, + f" kernel.launch @{export['symbol']}[{workload_values}]({', '.join(launch_values)}) : [{workload_types}]({', '.join(launch_types)})", + " check.return", + "}", + ] + ) + + +def render_case(symbol: str, command: dict[str, Any], export: dict[str, Any]) -> str | None: + kernel = command["kernel"].split(":")[-1] + case: str | None + if kernel == "ggml_binary_f32": + case = case_binary(symbol, command) + elif kernel == "ggml_rmsnorm_binary_f32": + case = case_rmsnorm_binary(symbol, command) + elif kernel in ("ggml_mul_mat_f32_f32_wmma", "ggml_mul_mat_f32_f32_decode_wave64"): + case = case_mul_mat(symbol, command, kernel) + elif kernel == "ggml_mul_mat_add_f32_f32_decode_wave64": + case = case_mul_mat_add_decode(symbol, command, kernel) + elif kernel in ( + "ggml_mul_mat_bias_f32_f32_wmma", + "ggml_mul_mat_add_f32_f32_wmma", + "ggml_mul_mat_bias_add_f32_f32_wmma", + "ggml_mul_mat_add_next_rmsnorm_f32_f32_wmma", + "ggml_mul_mat_bias_add_next_rmsnorm_f32_f32_wmma", + ): + case = case_mul_mat_postops(symbol, command, kernel) + elif kernel in ("ggml_mul_mat_swiglu_f32_f32_wmma", "ggml_mul_mat_swiglu_f32_f32_decode_wave64"): + case = case_swiglu(symbol, command, kernel) + elif kernel in ("llm_attention_q_matmul_rope_f32_f32_wmma", "llm_attention_q_matmul_rope_decode_f32_f32"): + case = case_llm_attention_q_matmul_rope(symbol, command, kernel) + elif kernel in ( + "llm_attention_k_matmul_rope_set_rows_f32_f32_wmma", + "llm_attention_k_matmul_rope_set_rows_decode_f32_f32", + ): + case = case_llm_attention_k_matmul_rope_set_rows(symbol, command, kernel) + elif kernel in ("llm_attention_v_matmul_set_rows_f32_f32_wmma", "llm_attention_v_matmul_set_rows_decode_f32_f32"): + case = case_llm_attention_v_matmul_set_rows(symbol, command, kernel) + elif kernel == "qwen3_moe_flash_attention_f32_f16_wmma": + case = case_flash_attention(symbol, command, kernel, "qwen3_moe") + elif kernel == "ggml_flash_attention_f32_f16_wmma": + case = case_flash_attention(symbol, command, kernel, "ggml") + elif kernel in ("ggml_flash_attention_decode_split_f32_f16_wmma", "ggml_flash_attention_decode_split_f32_f16_wmma_next_q8"): + case = case_flash_attention_decode_split(symbol, command, kernel) + elif kernel == "ggml_rope_f32": + case = case_rope(symbol, command, kernel, "ggml.rope_f32") + elif kernel == "ggml_rope_set_rows_f32": + case = case_rope_set_rows(symbol, command) + elif kernel == "ggml_set_rows": + case = case_set_rows(symbol, command) + elif kernel == "ggml_get_rows_f32": + case = case_get_rows(symbol, command) + elif kernel == "ggml_gather_add_f32": + case = case_gather_add(symbol, command) + else: + case = None + if case is not None: + return case + return case_generic(symbol, command, export) + + +def loom_target_for_export(corpus_dir: Path, export: dict[str, Any]) -> str: + source = corpus_dir / export["source"] + if not source.is_file(): + return "" + text = read_text(source) + symbol = export["symbol"] + for match in KERNEL_TARGET_RE.finditer(text): + if match.group("symbol") == symbol: + return match.group("target") + return "" + + +def load_exports() -> dict[str, dict[str, Any]]: + exports: dict[str, dict[str, Any]] = {} + for manifest_path in sorted(KERNEL_CORPUS_DIR.glob("*/manifest.json")): + manifest = load_json(manifest_path) + corpus_dir = manifest_path.parent + for export in manifest.get("exports", []): + name = export.get("name") + if not name: + continue + export = dict(export) + export["corpus_dir"] = str(corpus_dir.relative_to(HRX_DIR)) + export["loom_target"] = loom_target_for_export(corpus_dir, export) + previous = exports.get(name) + if previous is None or previous.get("target_selector") and not export.get("target_selector"): + exports[name] = export + return exports + + +def render_decl(export: dict[str, Any]) -> str: + symbol = export["symbol"] + workload = ", ".join(f"%{param['name']}: {param['type']}" for param in export.get("workload_parameters", [])) + launch_items = [f"%{param['name']}: {param['type']}" for param in export.get("launch_parameters", [])] + launch_items += [f"%{name}: buffer" for name in export.get("bindings", [])] + launch = ", ".join(launch_items) + target = export.get("loom_target", "") + target_attr = "" if not target else f" target(@{target})" + return f"kernel.decl{target_attr} @{symbol}({workload}) launch({launch})" + + +def command_shape_key(command: dict[str, Any]) -> str: + shape_data = { + "kernel": command.get("kernel"), + "integer_parameters": command.get("integer_parameters", {}), + "compile_parameters": command.get("compile_parameters", {}), + "binding_lengths": [binding.get("length") for binding in command.get("bindings", [])], + } + encoded = json.dumps(shape_data, sort_keys=True, separators=(",", ":")).encode("utf-8") + return hashlib.sha256(encoded).hexdigest()[:16] + + +def load_program_commands(dump_dir: Path) -> list[dict[str, Any]]: + program_paths = sorted(dump_dir.glob("program-*/program.json")) + if not program_paths: + fail(f"no program.json files found under {dump_dir}") + commands: list[dict[str, Any]] = [] + for program_path in program_paths: + program = load_json(program_path) + for command in program.get("commands", []): + command = dict(command) + command["program"] = { + "directory": program_path.parent.name, + "dump_id": program.get("dump_id"), + "shape_hash": program.get("shape_hash"), + "target": program.get("target"), + } + commands.append(command) + return commands + + +def source_paths_for_export(export: dict[str, Any]) -> list[str]: + recipe = export.get("compile_recipe", {}) + paths = [] + for source in recipe.get("primary_sources", []): + paths.append(source) + for source in recipe.get("library_sources", []): + paths.append(source) + return list(dict.fromkeys(paths)) + + +def primary_source_paths_for_export(export: dict[str, Any]) -> list[str]: + return list(dict.fromkeys(export.get("compile_recipe", {}).get("primary_sources", []))) + + +def library_source_paths_for_export(export: dict[str, Any]) -> list[str]: + return list(dict.fromkeys(export.get("compile_recipe", {}).get("library_sources", []))) + + +def generate_scenario(model_slug: str, + scenario: str, + commands: list[dict[str, Any]], + exports: dict[str, dict[str, Any]]) -> ScenarioBenchmarks: + dispatch_entries: list[dict[str, Any]] = [] + cases: list[str] = [] + used_exports: dict[str, dict[str, Any]] = {} + shape_keys = [command_shape_key(command) for command in commands] + shape_counts = Counter(shape_keys) + shape_commands: dict[str, list[dict[str, Any]]] = {} + for command, shape_key in zip(commands, shape_keys, strict=True): + shape_commands.setdefault(shape_key, []).append(command) + per_kernel_counts = Counter(command["kernel"] for command in commands) + + for shape_key, represented_commands in shape_commands.items(): + command = represented_commands[0] + kernel_name = command["kernel"] + export_name = kernel_name.split(":")[-1] + export = exports.get(export_name) + if export is None: + continue + else: + used_exports[export_name] = export + symbol = f"{model_slug}_{scenario}_{len(dispatch_entries):03d}_{sanitize_symbol(export_name)}" + try: + case = render_case(symbol, command, export) + except SystemExit as exc: + case = None + print(f"skipped {kernel_name}: {exc}") + else: + if case is None: + print(f"skipped {kernel_name}: unsupported template") + if case is not None: + cases.append(case) + benchmark_symbol = "@" + symbol + else: + continue + + dispatch_entries.append( + { + "benchmark": benchmark_symbol, + "kernel": kernel_name, + "symbol": export["symbol"], + "count": shape_counts[shape_key], + "integer_parameters": command.get("integer_parameters", {}), + "compile_parameters": command.get("compile_parameters", {}), + "workload_parameters": export.get("workload_parameters", []), + "sources": source_paths_for_export(export), + "primary_sources": primary_source_paths_for_export(export), + "library_sources": library_source_paths_for_export(export), + "corpus_dir": export.get("corpus_dir"), + } + ) + + return ScenarioBenchmarks( + scenario=scenario, + command_count=len(commands), + dispatch_count=len(shape_commands), + generated_count=len(dispatch_entries), + kernel_counts=dict(sorted(per_kernel_counts.items())), + dispatches=dispatch_entries, + cases=cases, + used_exports=used_exports, + ) + + +def write_model_loom(model_file: Path, + model_slug: str, + scenarios: list[ScenarioBenchmarks], + used_exports: dict[str, dict[str, Any]]) -> None: + scenario_names = ", ".join(scenario.scenario for scenario in scenarios) + targets = sorted({ + export["loom_target"] + for export in used_exports.values() + if export.get("loom_target") + }) + decls = [render_decl(export) for _, export in sorted(used_exports.items())] + lines = [ + f"// Generated by tools/benchmarks/generate-model-benchmarks.py for {model_slug} scenarios: {scenario_names}.", + "// Regenerate from an HRX command program dump rather than editing by hand.", + "", + *(f"target.decl @{target}" for target in targets), + "", + *decls, + "", + "", + ] + for scenario in scenarios: + lines.append(f"// Scenario: {scenario.scenario}") + lines.append("") + for case in scenario.cases: + lines.append(case) + lines.append("") + symbol_match = re.match(r"check\.case public @(.+?)_case", case) + if symbol_match: + benchmark_symbol = symbol_match.group(1) + lines.append(f"check.benchmark<@{benchmark_symbol}_case> @{benchmark_symbol}") + lines.append("") + + model_file.parent.mkdir(parents=True, exist_ok=True) + model_file.write_text("\n".join(lines).rstrip() + "\n", encoding="utf-8") + + +def write_manifest(manifest_file: Path, model_file: Path, model_slug: str, scenario: ScenarioBenchmarks) -> None: + write_json( + manifest_file, + { + "schema": "ggml-hrx-model-loom-benchmarks-v2", + "model": model_slug, + "scenario": scenario.scenario, + "loom_source": str(model_file.relative_to(HRX_DIR)), + "command_count": scenario.command_count, + "dispatch_count": scenario.dispatch_count, + "generated_count": scenario.generated_count, + "kernel_counts": scenario.kernel_counts, + "dispatches": scenario.dispatches, + }, + ) + + +def parse_scenario_dump(value: str) -> tuple[str, Path]: + scenario, separator, dump_dir = value.partition("=") + if not separator or not scenario or not dump_dir: + fail(f"invalid --scenario-dump value {value!r}; expected =") + return scenario, Path(dump_dir) + + +def scenario_dumps_from_args(args: argparse.Namespace) -> list[tuple[str, Path]]: + if args.scenario_dump: + if args.scenario is not None or args.dump_dir is not None: + fail("--scenario-dump cannot be combined with --scenario or --dump-dir") + if args.output_manifest is not None: + fail("--output-manifest cannot be used with multiple --scenario-dump values") + return [parse_scenario_dump(value) for value in args.scenario_dump] + if args.scenario is None or args.dump_dir is None: + fail("either --scenario and --dump-dir, or one or more --scenario-dump values, are required") + return [(args.scenario, args.dump_dir)] + + +def generate(args: argparse.Namespace) -> None: + exports = load_exports() + model_slug = args.model + model_file = args.output_loom or (BENCHMARK_DIR / "loom" / f"{model_slug}.loom") + scenario_dumps = scenario_dumps_from_args(args) + scenarios: list[ScenarioBenchmarks] = [] + used_exports: dict[str, dict[str, Any]] = {} + for scenario, dump_dir in scenario_dumps: + commands = load_program_commands(dump_dir) + generated = generate_scenario(model_slug, scenario, commands, exports) + scenarios.append(generated) + used_exports.update(generated.used_exports) + + write_model_loom(model_file, model_slug, scenarios, used_exports) + for scenario in scenarios: + manifest_file = args.output_manifest or (BENCHMARK_DIR / "loom" / f"{model_slug}.{scenario.scenario}.json") + write_manifest(manifest_file, model_file, model_slug, scenario) + print(f"wrote {manifest_file}") + + print(f"wrote {model_file}") + for scenario in scenarios: + print( + f"{scenario.scenario}: commands={scenario.command_count} " + f"dispatches={scenario.dispatch_count} generated={scenario.generated_count}" + ) + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--model", required=True, help="Model slug, for example llama32_3b_f16.") + parser.add_argument("--scenario", help="Scenario slug, for example pp512.") + parser.add_argument("--dump-dir", type=Path, help="Directory containing HRX program-*/program.json dumps.") + parser.add_argument( + "--scenario-dump", + action="append", + help="Scenario and dump directory as =. May be repeated.", + ) + parser.add_argument("--output-loom", type=Path, help="Generated model Loom file.") + parser.add_argument("--output-manifest", type=Path, help="Generated scenario manifest.") + generate(parser.parse_args()) + + +if __name__ == "__main__": + main() diff --git a/ggml/src/ggml-hrx/tools/benchmarks/loom/qwen38_27b_udq4kxl.work.loom b/ggml/src/ggml-hrx/tools/benchmarks/loom/qwen38_27b_udq4kxl.work.loom new file mode 100644 index 000000000000..d57a9f34cabf --- /dev/null +++ b/ggml/src/ggml-hrx/tools/benchmarks/loom/qwen38_27b_udq4kxl.work.loom @@ -0,0 +1,2008 @@ +// Generated by tools/benchmarks/generate-model-benchmarks.py for qwen38_27b_udq4kxl scenarios: tg_c5. +// Regenerate from an HRX command program dump rather than editing by hand. + +target.decl @ggml_binary_f32_gfx11_wave64 +target.decl @ggml_binary_swiglu_i4_gfx11_wave32 +target.decl @ggml_copy_f32_gfx11_wave64 +target.decl @ggml_flash_attention_gfx11_wave64 +target.decl @ggml_get_rows_f32_gfx11_wave64 +target.decl @ggml_mul_mat_f32_f32_decode_gfx11_wave64 +target.decl @ggml_rmsnorm_binary_gfx11_wave32 +target.decl @ggml_rmsnorm_gfx11_wave32 +target.decl @ggml_scale_bias_f32_gfx11_wave64 +target.decl @ggml_set_rows_gfx11_wave64 +target.decl @qwen3_moe_dense_gfx11_wave64 + +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_add_rmsnorm_binary_symmetric_i4_k32(%token_count: index) launch(%token_count: index, %lhs: buffer, %rhs: buffer, %residual_output: buffer, %weight: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_binary_f32_gfx11_wave64) @ggml_binary_f32(%element_count: index) launch(%element_count: index, %lhs: buffer, %rhs: buffer, %output: buffer) +kernel.decl target(@ggml_binary_swiglu_i4_gfx11_wave32) @ggml_binary_swiglu_symmetric_i4_k32() launch(%lhs: buffer, %rhs: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_copy_f32_gfx11_wave64) @ggml_concat_dim0_f32() launch(%lhs: buffer, %rhs: buffer, %output: buffer) +kernel.decl target(@ggml_copy_f32_gfx11_wave64) @ggml_copy_f32(%element_count: index) launch(%element_count: index, %source: buffer, %output: buffer) +kernel.decl @ggml_fill_negative_f32(%element_count: index) launch(%element_count: index, %output: buffer) +kernel.decl target(@ggml_flash_attention_gfx11_wave64) @ggml_flash_attention_f32_f16_wmma(%query_token_count: index, %key_value_token_count: index) launch(%query_token_count: index, %key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %gate: buffer, %output: buffer) +kernel.decl target(@ggml_get_rows_f32_gfx11_wave64) @ggml_get_rows_f32(%token_count: index, %row_count: index, %hidden_size: index) launch(%token_count: index, %row_count: index, %hidden_size: index, %token_ids: buffer, %weight: buffer, %output: buffer) +kernel.decl target(@ggml_mul_mat_f32_f32_decode_gfx11_wave64) @ggml_mul_mat_f32_f32_decode_wave64(%token_count: index, %input_size: index, %output_size: index) launch(%token_count: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %output: buffer) +kernel.decl target(@qwen3_moe_dense_gfx11_wave64) @ggml_mul_mat_q4_k_f32_adjacent_dual_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %first_weight: buffer, %second_weight: buffer, %first_output: buffer, %second_output: buffer) +kernel.decl @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) +kernel.decl @ggml_mul_mat_q6_k_packed_selected_refine_token1(%token_count: index, %candidate_count: index) launch(%token_count: index, %candidate_count: index, %input: buffer, %weight: buffer, %candidates: buffer, %exact_output: buffer, %output: buffer) +kernel.decl @ggml_mul_mat_q6_k_packed_token1_f16_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) +kernel.decl @ggml_mul_mat_q6_k_symmetric_i2_scan_token1() launch(%weight: buffer, %output: buffer, %qact: buffer, %scales: buffer) +kernel.decl @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma() launch(%first_weight: buffer, %second_weight: buffer, %first_output: buffer, %second_output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma() launch(%weight: buffer, %input: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer, %partial: buffer, %completion_counters: buffer) +kernel.decl @ggml_mul_mat_symmetric_i4_lowrow_wmma() launch(%weight: buffer, %input: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl @ggml_quantize_f32_symmetric_i4_k32() launch(%input: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_rmsnorm_binary_f32(%token_count: index) launch(%token_count: index, %input: buffer, %rhs: buffer, %output: buffer) +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_rmsnorm_binary_symmetric_i4_k32(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %raw_gate: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_rmsnorm_gfx11_wave32) @ggml_rmsnorm_mul_rope_f32() launch(%input: buffer, %weight: buffer, %positions: buffer, %output: buffer) +kernel.decl target(@ggml_scale_bias_f32_gfx11_wave64) @ggml_scale_bias_f32(%element_count: index) launch(%element_count: index, %input: buffer, %output: buffer) +kernel.decl target(@ggml_set_rows_gfx11_wave64) @ggml_set_rows(%token_count: index, %cache_row_count: index, %hidden_size: index) launch(%token_count: index, %cache_row_count: index, %hidden_size: index, %rows: buffer, %indices: buffer, %cache: buffer) +kernel.decl @ggml_top_k128_f32_reduce_gather_register(%element_count: index) launch(%element_count: index, %partial_values: buffer, %partial_ids: buffer, %candidate_output: buffer, %value_output: buffer) +kernel.decl @ggml_top_k8_f32_partitions_register(%element_count: index) launch(%element_count: index, %values: buffer, %partial_values: buffer, %partial_ids: buffer) +kernel.decl @llm_gated_delta_net_f32_wmma_head128() launch(%q: buffer, %k: buffer, %v: buffer, %g: buffer, %beta: buffer, %state_in: buffer, %dst: buffer, %rms_scales: buffer) +kernel.decl @llm_gated_delta_net_f32_wmma_head128_inplace() launch(%q: buffer, %k: buffer, %v: buffer, %g: buffer, %beta: buffer, %state_inout: buffer, %dst: buffer, %rms_scales: buffer) +kernel.decl @llm_gated_delta_net_f32_wmma_head128_snapshot() launch(%q: buffer, %k: buffer, %v: buffer, %g: buffer, %beta: buffer, %state_in: buffer, %snapshot_cache: buffer, %dst: buffer, %rms_scales: buffer) +kernel.decl @llm_gated_delta_net_projection_epilogue_f32() launch(%alpha_raw: buffer, %beta_raw: buffer, %bias: buffer, %a_scale: buffer, %gate_dst: buffer, %beta_dst: buffer) +kernel.decl @llm_ssm_conv_dconv4_silu_rollback_f32() launch(%state: buffer, %x: buffer, %filter: buffer, %output: buffer, %cache0: buffer, %cache1: buffer, %cache2: buffer, %cache3: buffer, %cache4: buffer) +kernel.decl target(@qwen3_moe_dense_gfx11_wave64) @qwen3_moe_dense_linear_q6k_f16_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) + + +// Scenario: tg_c5 + +check.case public @qwen38_27b_udq4kxl_tg_c5_000_ggml_get_rows_f32_case { + %token_count = check.literal value(2) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<8xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<8xi8>, tensor<715161600xi8>, tensor<40960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_000_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_tg_c5_000_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_001_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<40960xi8>, tensor<20480xi8>, tensor<40960xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_001_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_001_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + %partial = check.generate.fill value(0) : tensor<81920xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<40960xi8>, tensor<81920xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>, tensor<81920xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_tg_c5_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_003_ggml_mul_mat_q4_k_f32_adjacent_dual_wmma_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<40960xi8> + %first_weight = check.generate.fill value(0) : tensor<138240xi8> + %second_weight = check.generate.fill value(0) : tensor<138240xi8> + %first_output = check.generate.fill value(0) : tensor<384xi8> + %second_output = check.generate.fill value(0) : tensor<384xi8> + kernel.launch @ggml_mul_mat_q4_k_f32_adjacent_dual_wmma[%token_count](%token_count, %input, %first_weight, %second_weight, %first_output, %second_output) : [index](index, tensor<40960xi8>, tensor<138240xi8>, tensor<138240xi8>, tensor<384xi8>, tensor<384xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_003_ggml_mul_mat_q4_k_f32_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_tg_c5_003_ggml_mul_mat_q4_k_f32_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_004_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<49152xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<40960xi8>, tensor<49152xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_004_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_004_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_005_ggml_scale_bias_f32_case { + %element_count = check.literal value(30720) : index + %input = check.generate.fill value(0) : tensor<122880xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @ggml_scale_bias_f32[%element_count](%element_count, %input, %output) : [index](index, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_005_ggml_scale_bias_f32_case> @qwen38_27b_udq4kxl_tg_c5_005_ggml_scale_bias_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_006_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(20) : index + %hidden_size = check.literal value(30720) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<20x30720xf32> + %output = check.generate.fill value(1.0) : tensor<1x30720xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<20x30720xf32>, tensor<1x30720xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_006_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_tg_c5_006_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_007_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<81920xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<81920xi8>, tensor<163840xi8>, tensor<81920xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_007_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_tg_c5_007_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_008_ggml_scale_bias_f32_case { + %element_count = check.literal value(786432) : index + %input = check.generate.fill value(0) : tensor<3145728xi8> + %output = check.generate.fill value(0) : tensor<3145728xi8> + kernel.launch @ggml_scale_bias_f32[%element_count](%element_count, %input, %output) : [index](index, tensor<3145728xi8>, tensor<3145728xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_008_ggml_scale_bias_f32_case> @qwen38_27b_udq4kxl_tg_c5_008_ggml_scale_bias_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_009_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(20) : index + %hidden_size = check.literal value(786432) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<20x786432xf32> + %output = check.generate.fill value(1.0) : tensor<1x786432xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<20x786432xf32>, tensor<1x786432xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_009_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_tg_c5_009_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_010_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<384xi8> + %beta_raw = check.generate.fill value(0) : tensor<384xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<384xi8> + %beta_dst = check.generate.fill value(0) : tensor<384xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<384xi8>, tensor<384xi8>, tensor<192xi8>, tensor<192xi8>, tensor<384xi8>, tensor<384xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_010_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_tg_c5_010_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_011_llm_gated_delta_net_f32_wmma_head128_snapshot_case { + %q = check.generate.fill value(0) : tensor<49152xi8> + %k = check.generate.fill value(0) : tensor<49152xi8> + %v = check.generate.fill value(0) : tensor<65536xi8> + %g = check.generate.fill value(0) : tensor<384xi8> + %beta = check.generate.fill value(0) : tensor<384xi8> + %state_in = check.generate.fill value(0) : tensor<3145728xi8> + %snapshot_cache = check.generate.fill value(0) : tensor<15728640xi8> + %dst = check.generate.fill value(0) : tensor<49152xi8> + %rms_scales = check.generate.fill value(0) : tensor<384xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_snapshot[](%q, %k, %v, %g, %beta, %state_in, %snapshot_cache, %dst, %rms_scales) : [](tensor<49152xi8>, tensor<49152xi8>, tensor<65536xi8>, tensor<384xi8>, tensor<384xi8>, tensor<3145728xi8>, tensor<15728640xi8>, tensor<49152xi8>, tensor<384xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_011_llm_gated_delta_net_f32_wmma_head128_snapshot_case> @qwen38_27b_udq4kxl_tg_c5_011_llm_gated_delta_net_f32_wmma_head128_snapshot + +check.case public @qwen38_27b_udq4kxl_tg_c5_012_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(96) : index + %input = check.generate.fill value(0) : tensor<49152xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<49152xi8> + %output = check.generate.fill value(0) : tensor<49152xi8> + %quantized_values = check.generate.fill value(0) : tensor<6144xi8> + %scales = check.generate.fill value(0) : tensor<1536xi8> + %sums = check.generate.fill value(0) : tensor<1536xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<49152xi8>, tensor<512xi8>, tensor<49152xi8>, tensor<49152xi8>, tensor<6144xi8>, tensor<1536xi8>, tensor<1536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_012_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_012_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_013_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<49152xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<6144xi8> + %scales = check.generate.fill value(0) : tensor<1536xi8> + %sums = check.generate.fill value(0) : tensor<1536xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<49152xi8>, tensor<40960xi8>, tensor<6144xi8>, tensor<1536xi8>, tensor<1536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_013_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_013_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_014_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(2) : index + %lhs = check.generate.fill value(0) : tensor<40960xi8> + %rhs = check.generate.fill value(0) : tensor<40960xi8> + %residual_output = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<40960xi8>, tensor<40960xi8>, tensor<40960xi8>, tensor<20480xi8>, tensor<40960xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_014_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_014_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_015_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<139264xi8> + %second_output = check.generate.fill value(0) : tensor<139264xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<139264xi8>, tensor<139264xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_015_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_tg_c5_015_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_016_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<139264xi8> + %rhs = check.generate.fill value(0) : tensor<139264xi8> + %output = check.generate.fill value(0) : tensor<139264xi8> + %quantized_values = check.generate.fill value(0) : tensor<17408xi8> + %scales = check.generate.fill value(0) : tensor<4352xi8> + %sums = check.generate.fill value(0) : tensor<4352xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<139264xi8>, tensor<139264xi8>, tensor<139264xi8>, tensor<17408xi8>, tensor<4352xi8>, tensor<4352xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_016_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_016_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_017_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<139264xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<17408xi8> + %scales = check.generate.fill value(0) : tensor<4352xi8> + %sums = check.generate.fill value(0) : tensor<4352xi8> + %partial = check.generate.fill value(0) : tensor<40960xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<139264xi8>, tensor<40960xi8>, tensor<17408xi8>, tensor<4352xi8>, tensor<4352xi8>, tensor<40960xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_017_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_tg_c5_017_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_018_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<98304xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<40960xi8>, tensor<98304xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_018_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_018_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_019_qwen3_moe_dense_linear_q6k_f16_wmma_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<4300800xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + kernel.launch @qwen3_moe_dense_linear_q6k_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<40960xi8>, tensor<4300800xi8>, tensor<8192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_019_qwen3_moe_dense_linear_q6k_f16_wmma_case> @qwen38_27b_udq4kxl_tg_c5_019_qwen3_moe_dense_linear_q6k_f16_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_020_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<40960xi8>, tensor<8192xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_020_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_020_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_021_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<97280xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<32xi8> + %output = check.generate.fill value(0) : tensor<49152xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<97280xi8>, tensor<1024xi8>, tensor<32xi8>, tensor<49152xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_021_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_tg_c5_021_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_022_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<8192xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<32xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<8192xi8>, tensor<1024xi8>, tensor<32xi8>, tensor<8192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_022_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_tg_c5_022_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_023_ggml_set_rows_case { + %token_count = check.literal value(2) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<2x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<2xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<2x1024xf32>, tensor<2xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_023_ggml_set_rows_case> @qwen38_27b_udq4kxl_tg_c5_023_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_tg_c5_024_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(2) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<2x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<2x256xf16> + %output = check.generate.fill value(1.0) : tensor<2x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<2x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<2x256xf16>, tensor<2x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_024_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_tg_c5_024_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_025_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<49152xi8> + %quantized_values = check.generate.fill value(0) : tensor<6144xi8> + %scales = check.generate.fill value(0) : tensor<1536xi8> + %sums = check.generate.fill value(0) : tensor<1536xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<49152xi8>, tensor<6144xi8>, tensor<1536xi8>, tensor<1536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_025_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_025_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_026_ggml_binary_f32_case { + %element_count = check.literal value(10240) : index + %lhs = check.generate.fill value(2.0) : tensor<10240xf32> + %rhs = check.generate.fill value(3.0) : tensor<10240xf32> + %output = check.generate.fill value(0.0) : tensor<10240xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<10240xf32>, tensor<10240xf32>, tensor<10240xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_026_ggml_binary_f32_case> @qwen38_27b_udq4kxl_tg_c5_026_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_027_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(2.0) : tensor<2x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<2x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<2x5120xf32>, tensor<5120xf32>, tensor<2x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_027_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_tg_c5_027_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_028_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(2) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<2x5120xf32> + %output = check.generate.fill value(1.0) : tensor<1x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<2x5120xf32>, tensor<1x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_028_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_tg_c5_028_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_029_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<1042944000xi8> + %output = check.generate.fill value(0) : tensor<993280xi8> + kernel.launch @ggml_mul_mat_q6_k_packed_token1_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<20480xi8>, tensor<1042944000xi8>, tensor<993280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_029_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case> @qwen38_27b_udq4kxl_tg_c5_029_ggml_mul_mat_q6_k_packed_token1_f16_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_030_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<40960xi8> + %rhs = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<40960xi8>, tensor<40960xi8>, tensor<81920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_030_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_tg_c5_030_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_031_qwen3_moe_dense_linear_q6k_f16_wmma_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<43008000xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + kernel.launch @qwen3_moe_dense_linear_q6k_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<81920xi8>, tensor<43008000xi8>, tensor<40960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_031_qwen3_moe_dense_linear_q6k_f16_wmma_case> @qwen38_27b_udq4kxl_tg_c5_031_qwen3_moe_dense_linear_q6k_f16_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_032_ggml_get_rows_f32_case { + %token_count = check.literal value(5) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<20xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<20xi8>, tensor<715161600xi8>, tensor<102400xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_032_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_tg_c5_032_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_033_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(5) : index + %input = check.generate.fill value(0) : tensor<102400xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<102400xi8>, tensor<20480xi8>, tensor<102400xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_033_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_033_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_034_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<102400xi8> + %output = check.generate.fill value(0) : tensor<204800xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + %partial = check.generate.fill value(0) : tensor<204800xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<102400xi8>, tensor<204800xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>, tensor<204800xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_034_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_tg_c5_034_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_035_ggml_mul_mat_q4_k_f32_adjacent_dual_wmma_case { + %token_count = check.literal value(5) : index + %input = check.generate.fill value(0) : tensor<102400xi8> + %first_weight = check.generate.fill value(0) : tensor<138240xi8> + %second_weight = check.generate.fill value(0) : tensor<138240xi8> + %first_output = check.generate.fill value(0) : tensor<960xi8> + %second_output = check.generate.fill value(0) : tensor<960xi8> + kernel.launch @ggml_mul_mat_q4_k_f32_adjacent_dual_wmma[%token_count](%token_count, %input, %first_weight, %second_weight, %first_output, %second_output) : [index](index, tensor<102400xi8>, tensor<138240xi8>, tensor<138240xi8>, tensor<960xi8>, tensor<960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_035_ggml_mul_mat_q4_k_f32_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_tg_c5_035_ggml_mul_mat_q4_k_f32_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_036_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<102400xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<102400xi8>, tensor<122880xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_036_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_036_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_037_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<204800xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<204800xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<204800xi8>, tensor<163840xi8>, tensor<204800xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_037_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_tg_c5_037_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_038_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<960xi8> + %beta_raw = check.generate.fill value(0) : tensor<960xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<960xi8> + %beta_dst = check.generate.fill value(0) : tensor<960xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<960xi8>, tensor<960xi8>, tensor<192xi8>, tensor<192xi8>, tensor<960xi8>, tensor<960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_038_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_tg_c5_038_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_039_llm_gated_delta_net_f32_wmma_head128_snapshot_case { + %q = check.generate.fill value(0) : tensor<172032xi8> + %k = check.generate.fill value(0) : tensor<172032xi8> + %v = check.generate.fill value(0) : tensor<188416xi8> + %g = check.generate.fill value(0) : tensor<960xi8> + %beta = check.generate.fill value(0) : tensor<960xi8> + %state_in = check.generate.fill value(0) : tensor<3145728xi8> + %snapshot_cache = check.generate.fill value(0) : tensor<53477376xi8> + %dst = check.generate.fill value(0) : tensor<122880xi8> + %rms_scales = check.generate.fill value(0) : tensor<960xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_snapshot[](%q, %k, %v, %g, %beta, %state_in, %snapshot_cache, %dst, %rms_scales) : [](tensor<172032xi8>, tensor<172032xi8>, tensor<188416xi8>, tensor<960xi8>, tensor<960xi8>, tensor<3145728xi8>, tensor<53477376xi8>, tensor<122880xi8>, tensor<960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_039_llm_gated_delta_net_f32_wmma_head128_snapshot_case> @qwen38_27b_udq4kxl_tg_c5_039_llm_gated_delta_net_f32_wmma_head128_snapshot + +check.case public @qwen38_27b_udq4kxl_tg_c5_040_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(240) : index + %input = check.generate.fill value(0) : tensor<122880xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<122880xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + %quantized_values = check.generate.fill value(0) : tensor<15360xi8> + %scales = check.generate.fill value(0) : tensor<3840xi8> + %sums = check.generate.fill value(0) : tensor<3840xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<122880xi8>, tensor<512xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<15360xi8>, tensor<3840xi8>, tensor<3840xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_040_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_040_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_041_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<122880xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + %quantized_values = check.generate.fill value(0) : tensor<15360xi8> + %scales = check.generate.fill value(0) : tensor<3840xi8> + %sums = check.generate.fill value(0) : tensor<3840xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<122880xi8>, tensor<102400xi8>, tensor<15360xi8>, tensor<3840xi8>, tensor<3840xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_041_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_041_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_042_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(5) : index + %lhs = check.generate.fill value(0) : tensor<102400xi8> + %rhs = check.generate.fill value(0) : tensor<102400xi8> + %residual_output = check.generate.fill value(0) : tensor<102400xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<102400xi8>, tensor<102400xi8>, tensor<102400xi8>, tensor<20480xi8>, tensor<102400xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_042_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_042_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_043_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<348160xi8> + %second_output = check.generate.fill value(0) : tensor<348160xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<348160xi8>, tensor<348160xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_043_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_tg_c5_043_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_044_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<348160xi8> + %rhs = check.generate.fill value(0) : tensor<348160xi8> + %output = check.generate.fill value(0) : tensor<348160xi8> + %quantized_values = check.generate.fill value(0) : tensor<43520xi8> + %scales = check.generate.fill value(0) : tensor<10880xi8> + %sums = check.generate.fill value(0) : tensor<10880xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<348160xi8>, tensor<348160xi8>, tensor<348160xi8>, tensor<43520xi8>, tensor<10880xi8>, tensor<10880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_044_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_044_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_045_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<348160xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + %quantized_values = check.generate.fill value(0) : tensor<43520xi8> + %scales = check.generate.fill value(0) : tensor<10880xi8> + %sums = check.generate.fill value(0) : tensor<10880xi8> + %partial = check.generate.fill value(0) : tensor<102400xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<348160xi8>, tensor<102400xi8>, tensor<43520xi8>, tensor<10880xi8>, tensor<10880xi8>, tensor<102400xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_045_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_tg_c5_045_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_046_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<102400xi8> + %output = check.generate.fill value(0) : tensor<245760xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<102400xi8>, tensor<245760xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_046_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_046_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_047_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(5) : index + %input = check.generate.fill value(0) : tensor<102400xi8> + %weight = check.generate.fill value(0) : tensor<5611520xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<102400xi8>, tensor<5611520xi8>, tensor<20480xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_047_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_tg_c5_047_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_048_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<102400xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<102400xi8>, tensor<20480xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_048_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_048_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_049_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<244736xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<80xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<244736xi8>, tensor<1024xi8>, tensor<80xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_049_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_tg_c5_049_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_050_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<80xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<20480xi8>, tensor<1024xi8>, tensor<80xi8>, tensor<20480xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_050_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_tg_c5_050_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_051_ggml_set_rows_case { + %token_count = check.literal value(5) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<5x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<5xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<5x1024xf32>, tensor<5xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_051_ggml_set_rows_case> @qwen38_27b_udq4kxl_tg_c5_051_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_tg_c5_052_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(5) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<5x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<5x256xf16> + %output = check.generate.fill value(1.0) : tensor<5x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<5x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<5x256xf16>, tensor<5x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_052_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_tg_c5_052_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_053_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<122880xi8> + %quantized_values = check.generate.fill value(0) : tensor<15360xi8> + %scales = check.generate.fill value(0) : tensor<3840xi8> + %sums = check.generate.fill value(0) : tensor<3840xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<122880xi8>, tensor<15360xi8>, tensor<3840xi8>, tensor<3840xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_053_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_053_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_054_ggml_binary_f32_case { + %element_count = check.literal value(25600) : index + %lhs = check.generate.fill value(2.0) : tensor<25600xf32> + %rhs = check.generate.fill value(3.0) : tensor<25600xf32> + %output = check.generate.fill value(0.0) : tensor<25600xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<25600xf32>, tensor<25600xf32>, tensor<25600xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_054_ggml_binary_f32_case> @qwen38_27b_udq4kxl_tg_c5_054_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_055_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(5) : index + %input = check.generate.fill value(2.0) : tensor<5x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<5x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<5x5120xf32>, tensor<5120xf32>, tensor<5x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_055_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_tg_c5_055_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_056_ggml_get_rows_f32_case { + %token_count = check.literal value(5) : index + %row_count = check.literal value(5) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<5xi32> + %weight = check.generate.fill value(0.0) : tensor<5x5120xf32> + %output = check.generate.fill value(1.0) : tensor<5x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<5xi32>, tensor<5x5120xf32>, tensor<5x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_056_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_tg_c5_056_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_057_qwen3_moe_dense_linear_q6k_f16_wmma_case { + %token_count = check.literal value(5) : index + %input = check.generate.fill value(0) : tensor<102400xi8> + %weight = check.generate.fill value(0) : tensor<1042944000xi8> + %output = check.generate.fill value(0) : tensor<4966400xi8> + kernel.launch @qwen3_moe_dense_linear_q6k_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<102400xi8>, tensor<1042944000xi8>, tensor<4966400xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_057_qwen3_moe_dense_linear_q6k_f16_wmma_case> @qwen38_27b_udq4kxl_tg_c5_057_qwen3_moe_dense_linear_q6k_f16_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_058_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<102400xi8> + %rhs = check.generate.fill value(0) : tensor<102400xi8> + %output = check.generate.fill value(0) : tensor<204800xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<102400xi8>, tensor<102400xi8>, tensor<204800xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_058_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_tg_c5_058_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_059_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(5) : index + %input = check.generate.fill value(0) : tensor<204800xi8> + %weight = check.generate.fill value(0) : tensor<56115200xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<204800xi8>, tensor<56115200xi8>, tensor<102400xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_059_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_tg_c5_059_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_060_ggml_get_rows_f32_case { + %token_count = check.literal value(2) : index + %row_count = check.literal value(2) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<2xi32> + %weight = check.generate.fill value(0.0) : tensor<2x5120xf32> + %output = check.generate.fill value(1.0) : tensor<2x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<2xi32>, tensor<2x5120xf32>, tensor<2x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_060_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_tg_c5_060_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_061_qwen3_moe_dense_linear_q6k_f16_wmma_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<1042944000xi8> + %output = check.generate.fill value(0) : tensor<1986560xi8> + kernel.launch @qwen3_moe_dense_linear_q6k_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<40960xi8>, tensor<1042944000xi8>, tensor<1986560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_061_qwen3_moe_dense_linear_q6k_f16_wmma_case> @qwen38_27b_udq4kxl_tg_c5_061_qwen3_moe_dense_linear_q6k_f16_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_062_ggml_get_rows_f32_case { + %token_count = check.literal value(16) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<64xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<64xi8>, tensor<715161600xi8>, tensor<327680xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_062_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_tg_c5_062_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_063_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(16) : index + %input = check.generate.fill value(0) : tensor<327680xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<327680xi8>, tensor<20480xi8>, tensor<327680xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_063_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_063_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_064_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<327680xi8> + %output = check.generate.fill value(0) : tensor<655360xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + %partial = check.generate.fill value(0) : tensor<655360xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<327680xi8>, tensor<655360xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>, tensor<655360xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_064_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_tg_c5_064_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_065_ggml_mul_mat_q4_k_f32_adjacent_dual_wmma_case { + %token_count = check.literal value(16) : index + %input = check.generate.fill value(0) : tensor<327680xi8> + %first_weight = check.generate.fill value(0) : tensor<138240xi8> + %second_weight = check.generate.fill value(0) : tensor<138240xi8> + %first_output = check.generate.fill value(0) : tensor<3072xi8> + %second_output = check.generate.fill value(0) : tensor<3072xi8> + kernel.launch @ggml_mul_mat_q4_k_f32_adjacent_dual_wmma[%token_count](%token_count, %input, %first_weight, %second_weight, %first_output, %second_output) : [index](index, tensor<327680xi8>, tensor<138240xi8>, tensor<138240xi8>, tensor<3072xi8>, tensor<3072xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_065_ggml_mul_mat_q4_k_f32_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_tg_c5_065_ggml_mul_mat_q4_k_f32_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_066_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<327680xi8> + %output = check.generate.fill value(0) : tensor<393216xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<327680xi8>, tensor<393216xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_066_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_066_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_067_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<655360xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<655360xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<655360xi8>, tensor<163840xi8>, tensor<655360xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_067_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_tg_c5_067_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_068_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<3072xi8> + %beta_raw = check.generate.fill value(0) : tensor<3072xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<3072xi8> + %beta_dst = check.generate.fill value(0) : tensor<3072xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<3072xi8>, tensor<3072xi8>, tensor<192xi8>, tensor<192xi8>, tensor<3072xi8>, tensor<3072xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_068_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_tg_c5_068_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_069_llm_gated_delta_net_f32_wmma_head128_case { + %q = check.generate.fill value(0) : tensor<417792xi8> + %k = check.generate.fill value(0) : tensor<417792xi8> + %v = check.generate.fill value(0) : tensor<434176xi8> + %g = check.generate.fill value(0) : tensor<2112xi8> + %beta = check.generate.fill value(0) : tensor<2112xi8> + %state_in = check.generate.fill value(0) : tensor<3145728xi8> + %dst = check.generate.fill value(0) : tensor<3416064xi8> + %rms_scales = check.generate.fill value(0) : tensor<2112xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128[](%q, %k, %v, %g, %beta, %state_in, %dst, %rms_scales) : [](tensor<417792xi8>, tensor<417792xi8>, tensor<434176xi8>, tensor<2112xi8>, tensor<2112xi8>, tensor<3145728xi8>, tensor<3416064xi8>, tensor<2112xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_069_llm_gated_delta_net_f32_wmma_head128_case> @qwen38_27b_udq4kxl_tg_c5_069_llm_gated_delta_net_f32_wmma_head128 + +check.case public @qwen38_27b_udq4kxl_tg_c5_070_ggml_copy_f32_case { + %element_count = check.literal value(786432) : index + %source = check.generate.fill value(0) : tensor<3145728xi8> + %output = check.generate.fill value(0) : tensor<3145728xi8> + kernel.launch @ggml_copy_f32[%element_count](%element_count, %source, %output) : [index](index, tensor<3145728xi8>, tensor<3145728xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_070_ggml_copy_f32_case> @qwen38_27b_udq4kxl_tg_c5_070_ggml_copy_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_071_llm_gated_delta_net_f32_wmma_head128_inplace_case { + %q = check.generate.fill value(0) : tensor<8192xi8> + %k = check.generate.fill value(0) : tensor<8192xi8> + %v = check.generate.fill value(0) : tensor<24576xi8> + %g = check.generate.fill value(0) : tensor<192xi8> + %beta = check.generate.fill value(0) : tensor<192xi8> + %state_inout = check.generate.fill value(0) : tensor<3145728xi8> + %dst = check.generate.fill value(0) : tensor<24576xi8> + %rms_scales = check.generate.fill value(0) : tensor<192xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_inplace[](%q, %k, %v, %g, %beta, %state_inout, %dst, %rms_scales) : [](tensor<8192xi8>, tensor<8192xi8>, tensor<24576xi8>, tensor<192xi8>, tensor<192xi8>, tensor<3145728xi8>, tensor<24576xi8>, tensor<192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_071_llm_gated_delta_net_f32_wmma_head128_inplace_case> @qwen38_27b_udq4kxl_tg_c5_071_llm_gated_delta_net_f32_wmma_head128_inplace + +check.case public @qwen38_27b_udq4kxl_tg_c5_072_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(768) : index + %input = check.generate.fill value(0) : tensor<393216xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<393216xi8> + %output = check.generate.fill value(0) : tensor<393216xi8> + %quantized_values = check.generate.fill value(0) : tensor<49152xi8> + %scales = check.generate.fill value(0) : tensor<12288xi8> + %sums = check.generate.fill value(0) : tensor<12288xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<393216xi8>, tensor<512xi8>, tensor<393216xi8>, tensor<393216xi8>, tensor<49152xi8>, tensor<12288xi8>, tensor<12288xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_072_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_072_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_073_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<393216xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + %quantized_values = check.generate.fill value(0) : tensor<49152xi8> + %scales = check.generate.fill value(0) : tensor<12288xi8> + %sums = check.generate.fill value(0) : tensor<12288xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<393216xi8>, tensor<327680xi8>, tensor<49152xi8>, tensor<12288xi8>, tensor<12288xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_073_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_073_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_074_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(16) : index + %lhs = check.generate.fill value(0) : tensor<327680xi8> + %rhs = check.generate.fill value(0) : tensor<327680xi8> + %residual_output = check.generate.fill value(0) : tensor<327680xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<327680xi8>, tensor<327680xi8>, tensor<327680xi8>, tensor<20480xi8>, tensor<327680xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_074_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_074_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_075_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<1114112xi8> + %second_output = check.generate.fill value(0) : tensor<1114112xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<1114112xi8>, tensor<1114112xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_075_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_tg_c5_075_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_076_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<1114112xi8> + %rhs = check.generate.fill value(0) : tensor<1114112xi8> + %output = check.generate.fill value(0) : tensor<1114112xi8> + %quantized_values = check.generate.fill value(0) : tensor<139264xi8> + %scales = check.generate.fill value(0) : tensor<34816xi8> + %sums = check.generate.fill value(0) : tensor<34816xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<1114112xi8>, tensor<1114112xi8>, tensor<1114112xi8>, tensor<139264xi8>, tensor<34816xi8>, tensor<34816xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_076_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_076_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_077_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<1114112xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + %quantized_values = check.generate.fill value(0) : tensor<139264xi8> + %scales = check.generate.fill value(0) : tensor<34816xi8> + %sums = check.generate.fill value(0) : tensor<34816xi8> + %partial = check.generate.fill value(0) : tensor<327680xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<1114112xi8>, tensor<327680xi8>, tensor<139264xi8>, tensor<34816xi8>, tensor<34816xi8>, tensor<327680xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_077_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_tg_c5_077_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_078_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<327680xi8> + %output = check.generate.fill value(0) : tensor<786432xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<327680xi8>, tensor<786432xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_078_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_078_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_079_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(16) : index + %input = check.generate.fill value(0) : tensor<327680xi8> + %weight = check.generate.fill value(0) : tensor<5611520xi8> + %output = check.generate.fill value(0) : tensor<65536xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<327680xi8>, tensor<5611520xi8>, tensor<65536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_079_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_tg_c5_079_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_080_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<327680xi8> + %output = check.generate.fill value(0) : tensor<65536xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<327680xi8>, tensor<65536xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_080_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_080_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_081_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<785408xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<256xi8> + %output = check.generate.fill value(0) : tensor<393216xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<785408xi8>, tensor<1024xi8>, tensor<256xi8>, tensor<393216xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_081_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_tg_c5_081_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_082_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<65536xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<256xi8> + %output = check.generate.fill value(0) : tensor<65536xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<65536xi8>, tensor<1024xi8>, tensor<256xi8>, tensor<65536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_082_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_tg_c5_082_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_083_ggml_set_rows_case { + %token_count = check.literal value(16) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<16x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<16xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<16x1024xf32>, tensor<16xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_083_ggml_set_rows_case> @qwen38_27b_udq4kxl_tg_c5_083_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_tg_c5_084_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(16) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<16x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<16x256xf16> + %output = check.generate.fill value(1.0) : tensor<16x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<16x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<16x256xf16>, tensor<16x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_084_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_tg_c5_084_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_085_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<393216xi8> + %quantized_values = check.generate.fill value(0) : tensor<49152xi8> + %scales = check.generate.fill value(0) : tensor<12288xi8> + %sums = check.generate.fill value(0) : tensor<12288xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<393216xi8>, tensor<49152xi8>, tensor<12288xi8>, tensor<12288xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_085_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_085_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_086_ggml_binary_f32_case { + %element_count = check.literal value(81920) : index + %lhs = check.generate.fill value(2.0) : tensor<81920xf32> + %rhs = check.generate.fill value(3.0) : tensor<81920xf32> + %output = check.generate.fill value(0.0) : tensor<81920xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<81920xf32>, tensor<81920xf32>, tensor<81920xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_086_ggml_binary_f32_case> @qwen38_27b_udq4kxl_tg_c5_086_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_087_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(16) : index + %input = check.generate.fill value(2.0) : tensor<16x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<16x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<16x5120xf32>, tensor<5120xf32>, tensor<16x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_087_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_tg_c5_087_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_088_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<327680xi8> + %rhs = check.generate.fill value(0) : tensor<327680xi8> + %output = check.generate.fill value(0) : tensor<655360xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<327680xi8>, tensor<327680xi8>, tensor<655360xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_088_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_tg_c5_088_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_089_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(16) : index + %input = check.generate.fill value(0) : tensor<655360xi8> + %weight = check.generate.fill value(0) : tensor<56115200xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<655360xi8>, tensor<56115200xi8>, tensor<327680xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_089_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_tg_c5_089_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_090_ggml_get_rows_f32_case { + %token_count = check.literal value(4) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<16xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<16xi8>, tensor<715161600xi8>, tensor<81920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_090_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_tg_c5_090_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_091_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<81920xi8>, tensor<20480xi8>, tensor<81920xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_091_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_091_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_092_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<163840xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + %partial = check.generate.fill value(0) : tensor<163840xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<81920xi8>, tensor<163840xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>, tensor<163840xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_092_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_tg_c5_092_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_093_ggml_mul_mat_q4_k_f32_adjacent_dual_wmma_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(0) : tensor<81920xi8> + %first_weight = check.generate.fill value(0) : tensor<138240xi8> + %second_weight = check.generate.fill value(0) : tensor<138240xi8> + %first_output = check.generate.fill value(0) : tensor<768xi8> + %second_output = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @ggml_mul_mat_q4_k_f32_adjacent_dual_wmma[%token_count](%token_count, %input, %first_weight, %second_weight, %first_output, %second_output) : [index](index, tensor<81920xi8>, tensor<138240xi8>, tensor<138240xi8>, tensor<768xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_093_ggml_mul_mat_q4_k_f32_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_tg_c5_093_ggml_mul_mat_q4_k_f32_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_094_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<98304xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<81920xi8>, tensor<98304xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_094_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_094_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_095_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<163840xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<163840xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<163840xi8>, tensor<163840xi8>, tensor<163840xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_095_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_tg_c5_095_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_096_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<768xi8> + %beta_raw = check.generate.fill value(0) : tensor<768xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<768xi8> + %beta_dst = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<768xi8>, tensor<768xi8>, tensor<192xi8>, tensor<192xi8>, tensor<768xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_096_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_tg_c5_096_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_097_llm_gated_delta_net_f32_wmma_head128_snapshot_case { + %q = check.generate.fill value(0) : tensor<131072xi8> + %k = check.generate.fill value(0) : tensor<131072xi8> + %v = check.generate.fill value(0) : tensor<147456xi8> + %g = check.generate.fill value(0) : tensor<768xi8> + %beta = check.generate.fill value(0) : tensor<768xi8> + %state_in = check.generate.fill value(0) : tensor<3145728xi8> + %snapshot_cache = check.generate.fill value(0) : tensor<40894464xi8> + %dst = check.generate.fill value(0) : tensor<98304xi8> + %rms_scales = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_snapshot[](%q, %k, %v, %g, %beta, %state_in, %snapshot_cache, %dst, %rms_scales) : [](tensor<131072xi8>, tensor<131072xi8>, tensor<147456xi8>, tensor<768xi8>, tensor<768xi8>, tensor<3145728xi8>, tensor<40894464xi8>, tensor<98304xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_097_llm_gated_delta_net_f32_wmma_head128_snapshot_case> @qwen38_27b_udq4kxl_tg_c5_097_llm_gated_delta_net_f32_wmma_head128_snapshot + +check.case public @qwen38_27b_udq4kxl_tg_c5_098_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(192) : index + %input = check.generate.fill value(0) : tensor<98304xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<98304xi8> + %output = check.generate.fill value(0) : tensor<98304xi8> + %quantized_values = check.generate.fill value(0) : tensor<12288xi8> + %scales = check.generate.fill value(0) : tensor<3072xi8> + %sums = check.generate.fill value(0) : tensor<3072xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<98304xi8>, tensor<512xi8>, tensor<98304xi8>, tensor<98304xi8>, tensor<12288xi8>, tensor<3072xi8>, tensor<3072xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_098_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_098_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_099_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<98304xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<12288xi8> + %scales = check.generate.fill value(0) : tensor<3072xi8> + %sums = check.generate.fill value(0) : tensor<3072xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<98304xi8>, tensor<81920xi8>, tensor<12288xi8>, tensor<3072xi8>, tensor<3072xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_099_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_099_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_100_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(4) : index + %lhs = check.generate.fill value(0) : tensor<81920xi8> + %rhs = check.generate.fill value(0) : tensor<81920xi8> + %residual_output = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<81920xi8>, tensor<81920xi8>, tensor<81920xi8>, tensor<20480xi8>, tensor<81920xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_100_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_100_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_101_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<278528xi8> + %second_output = check.generate.fill value(0) : tensor<278528xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<278528xi8>, tensor<278528xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_101_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_tg_c5_101_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_102_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<278528xi8> + %rhs = check.generate.fill value(0) : tensor<278528xi8> + %output = check.generate.fill value(0) : tensor<278528xi8> + %quantized_values = check.generate.fill value(0) : tensor<34816xi8> + %scales = check.generate.fill value(0) : tensor<8704xi8> + %sums = check.generate.fill value(0) : tensor<8704xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<278528xi8>, tensor<278528xi8>, tensor<278528xi8>, tensor<34816xi8>, tensor<8704xi8>, tensor<8704xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_102_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_102_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_103_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<278528xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<34816xi8> + %scales = check.generate.fill value(0) : tensor<8704xi8> + %sums = check.generate.fill value(0) : tensor<8704xi8> + %partial = check.generate.fill value(0) : tensor<81920xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<278528xi8>, tensor<81920xi8>, tensor<34816xi8>, tensor<8704xi8>, tensor<8704xi8>, tensor<81920xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_103_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_tg_c5_103_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_104_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<196608xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<81920xi8>, tensor<196608xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_104_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_104_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_105_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<5611520xi8> + %output = check.generate.fill value(0) : tensor<16384xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<81920xi8>, tensor<5611520xi8>, tensor<16384xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_105_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_tg_c5_105_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_106_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<16384xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<81920xi8>, tensor<16384xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_106_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_106_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_107_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<195584xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<64xi8> + %output = check.generate.fill value(0) : tensor<98304xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<195584xi8>, tensor<1024xi8>, tensor<64xi8>, tensor<98304xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_107_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_tg_c5_107_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_108_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<16384xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<64xi8> + %output = check.generate.fill value(0) : tensor<16384xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<16384xi8>, tensor<1024xi8>, tensor<64xi8>, tensor<16384xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_108_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_tg_c5_108_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_109_ggml_set_rows_case { + %token_count = check.literal value(4) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<4x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<4xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<4x1024xf32>, tensor<4xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_109_ggml_set_rows_case> @qwen38_27b_udq4kxl_tg_c5_109_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_tg_c5_110_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(4) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<4x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<4x256xf16> + %output = check.generate.fill value(1.0) : tensor<4x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<4x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<4x256xf16>, tensor<4x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_110_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_tg_c5_110_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_111_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<98304xi8> + %quantized_values = check.generate.fill value(0) : tensor<12288xi8> + %scales = check.generate.fill value(0) : tensor<3072xi8> + %sums = check.generate.fill value(0) : tensor<3072xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<98304xi8>, tensor<12288xi8>, tensor<3072xi8>, tensor<3072xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_111_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_111_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_112_ggml_binary_f32_case { + %element_count = check.literal value(20480) : index + %lhs = check.generate.fill value(2.0) : tensor<20480xf32> + %rhs = check.generate.fill value(3.0) : tensor<20480xf32> + %output = check.generate.fill value(0.0) : tensor<20480xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<20480xf32>, tensor<20480xf32>, tensor<20480xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_112_ggml_binary_f32_case> @qwen38_27b_udq4kxl_tg_c5_112_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_113_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(2.0) : tensor<4x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<4x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<4x5120xf32>, tensor<5120xf32>, tensor<4x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_113_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_tg_c5_113_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_114_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(4) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<4x5120xf32> + %output = check.generate.fill value(1.0) : tensor<1x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<4x5120xf32>, tensor<1x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_114_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_tg_c5_114_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_115_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<81920xi8> + %rhs = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<163840xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<81920xi8>, tensor<81920xi8>, tensor<163840xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_115_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_tg_c5_115_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_116_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(0) : tensor<163840xi8> + %weight = check.generate.fill value(0) : tensor<56115200xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<163840xi8>, tensor<56115200xi8>, tensor<81920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_116_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_tg_c5_116_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_117_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<4xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<4xi8>, tensor<715161600xi8>, tensor<20480xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_117_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_tg_c5_117_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_118_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(2.0) : tensor<1x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<1x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<1x5120xf32>, tensor<5120xf32>, tensor<1x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_118_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_tg_c5_118_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_119_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<20480xi8> + %rhs = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<20480xi8>, tensor<20480xi8>, tensor<40960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_119_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_tg_c5_119_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_120_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(10240) : index + %output_size = check.literal value(5120) : index + %input = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<43008000xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<40960xi8>, tensor<43008000xi8>, tensor<20480xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_120_ggml_mul_mat_f32_f32_decode_wave64_case> @qwen38_27b_udq4kxl_tg_c5_120_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @qwen38_27b_udq4kxl_tg_c5_121_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<20480xi8>, tensor<20480xi8>, tensor<20480xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_121_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_121_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_122_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<49152xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<20480xi8>, tensor<49152xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_122_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_122_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_123_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(5120) : index + %output_size = check.literal value(1024) : index + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<4300800xi8> + %output = check.generate.fill value(0) : tensor<4096xi8> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<20480xi8>, tensor<4300800xi8>, tensor<4096xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_123_ggml_mul_mat_f32_f32_decode_wave64_case> @qwen38_27b_udq4kxl_tg_c5_123_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @qwen38_27b_udq4kxl_tg_c5_124_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<4096xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<20480xi8>, tensor<4096xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_124_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_124_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_125_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<48128xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<16xi8> + %output = check.generate.fill value(0) : tensor<24576xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<48128xi8>, tensor<1024xi8>, tensor<16xi8>, tensor<24576xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_125_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_tg_c5_125_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_126_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<4096xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<16xi8> + %output = check.generate.fill value(0) : tensor<4096xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<4096xi8>, tensor<1024xi8>, tensor<16xi8>, tensor<4096xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_126_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_tg_c5_126_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_127_ggml_set_rows_case { + %token_count = check.literal value(1) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<1x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<1xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<1x1024xf32>, tensor<1xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_127_ggml_set_rows_case> @qwen38_27b_udq4kxl_tg_c5_127_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_tg_c5_128_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(1) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<1x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<1x256xf16> + %output = check.generate.fill value(1.0) : tensor<1x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<1x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<1x256xf16>, tensor<1x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_128_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_tg_c5_128_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_129_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<24576xi8> + %quantized_values = check.generate.fill value(0) : tensor<3072xi8> + %scales = check.generate.fill value(0) : tensor<768xi8> + %sums = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<24576xi8>, tensor<3072xi8>, tensor<768xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_129_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_129_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_130_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<24576xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<3072xi8> + %scales = check.generate.fill value(0) : tensor<768xi8> + %sums = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<24576xi8>, tensor<20480xi8>, tensor<3072xi8>, tensor<768xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_130_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_130_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_131_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(1) : index + %lhs = check.generate.fill value(0) : tensor<20480xi8> + %rhs = check.generate.fill value(0) : tensor<20480xi8> + %residual_output = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<20480xi8>, tensor<20480xi8>, tensor<20480xi8>, tensor<20480xi8>, tensor<20480xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_131_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_131_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_132_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<69632xi8> + %second_output = check.generate.fill value(0) : tensor<69632xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<69632xi8>, tensor<69632xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_132_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_tg_c5_132_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_133_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<69632xi8> + %rhs = check.generate.fill value(0) : tensor<69632xi8> + %output = check.generate.fill value(0) : tensor<69632xi8> + %quantized_values = check.generate.fill value(0) : tensor<8704xi8> + %scales = check.generate.fill value(0) : tensor<2176xi8> + %sums = check.generate.fill value(0) : tensor<2176xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<69632xi8>, tensor<69632xi8>, tensor<69632xi8>, tensor<8704xi8>, tensor<2176xi8>, tensor<2176xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_133_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_133_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_134_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<69632xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<8704xi8> + %scales = check.generate.fill value(0) : tensor<2176xi8> + %sums = check.generate.fill value(0) : tensor<2176xi8> + %partial = check.generate.fill value(0) : tensor<20480xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<69632xi8>, tensor<20480xi8>, tensor<8704xi8>, tensor<2176xi8>, tensor<2176xi8>, tensor<20480xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_134_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_tg_c5_134_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_135_ggml_binary_f32_case { + %element_count = check.literal value(5120) : index + %lhs = check.generate.fill value(2.0) : tensor<5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<5120xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<5120xf32>, tensor<5120xf32>, tensor<5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_135_ggml_binary_f32_case> @qwen38_27b_udq4kxl_tg_c5_135_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_136_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(1) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<1x5120xf32> + %output = check.generate.fill value(1.0) : tensor<1x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<1x5120xf32>, tensor<1x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_136_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_tg_c5_136_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_137_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<20480xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_137_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_137_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_138_ggml_mul_mat_q6_k_symmetric_i2_scan_token1_case { + %weight = check.generate.fill value(0) : tensor<1380659200xi8> + %output = check.generate.fill value(0) : tensor<993280xi8> + %qact = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_mul_mat_q6_k_symmetric_i2_scan_token1[](%weight, %output, %qact, %scales) : [](tensor<1380659200xi8>, tensor<993280xi8>, tensor<2560xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_138_ggml_mul_mat_q6_k_symmetric_i2_scan_token1_case> @qwen38_27b_udq4kxl_tg_c5_138_ggml_mul_mat_q6_k_symmetric_i2_scan_token1 + +check.case public @qwen38_27b_udq4kxl_tg_c5_139_ggml_top_k8_f32_partitions_register_case { + %element_count = check.literal value(248320) : index + %values = check.generate.fill value(0) : tensor<993280xi8> + %partial_values = check.generate.fill value(0) : tensor<4096xi8> + %partial_ids = check.generate.fill value(0) : tensor<4096xi8> + kernel.launch @ggml_top_k8_f32_partitions_register[%element_count](%element_count, %values, %partial_values, %partial_ids) : [index](index, tensor<993280xi8>, tensor<4096xi8>, tensor<4096xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_139_ggml_top_k8_f32_partitions_register_case> @qwen38_27b_udq4kxl_tg_c5_139_ggml_top_k8_f32_partitions_register + +check.case public @qwen38_27b_udq4kxl_tg_c5_140_ggml_top_k128_f32_reduce_gather_register_case { + %element_count = check.literal value(248320) : index + %partial_values = check.generate.fill value(0) : tensor<4096xi8> + %partial_ids = check.generate.fill value(0) : tensor<4096xi8> + %candidate_output = check.generate.fill value(0) : tensor<512xi8> + %value_output = check.generate.fill value(0) : tensor<512xi8> + kernel.launch @ggml_top_k128_f32_reduce_gather_register[%element_count](%element_count, %partial_values, %partial_ids, %candidate_output, %value_output) : [index](index, tensor<4096xi8>, tensor<4096xi8>, tensor<512xi8>, tensor<512xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_140_ggml_top_k128_f32_reduce_gather_register_case> @qwen38_27b_udq4kxl_tg_c5_140_ggml_top_k128_f32_reduce_gather_register + +check.case public @qwen38_27b_udq4kxl_tg_c5_141_ggml_fill_negative_f32_case { + %element_count = check.literal value(248320) : index + %output = check.generate.fill value(0) : tensor<993280xi8> + kernel.launch @ggml_fill_negative_f32[%element_count](%element_count, %output) : [index](index, tensor<993280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_141_ggml_fill_negative_f32_case> @qwen38_27b_udq4kxl_tg_c5_141_ggml_fill_negative_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_142_ggml_mul_mat_q6_k_packed_selected_refine_token1_case { + %token_count = check.literal value(1) : index + %candidate_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<1380659200xi8> + %candidates = check.generate.fill value(0) : tensor<256xi8> + %exact_output = check.generate.fill value(0) : tensor<256xi8> + %output = check.generate.fill value(0) : tensor<993280xi8> + kernel.launch @ggml_mul_mat_q6_k_packed_selected_refine_token1[%token_count, %candidate_count](%token_count, %candidate_count, %input, %weight, %candidates, %exact_output, %output) : [index, index](index, index, tensor<20480xi8>, tensor<1380659200xi8>, tensor<256xi8>, tensor<256xi8>, tensor<993280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_142_ggml_mul_mat_q6_k_packed_selected_refine_token1_case> @qwen38_27b_udq4kxl_tg_c5_142_ggml_mul_mat_q6_k_packed_selected_refine_token1 + +check.case public @qwen38_27b_udq4kxl_tg_c5_143_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + %partial = check.generate.fill value(0) : tensor<40960xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<20480xi8>, tensor<40960xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>, tensor<40960xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_143_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_tg_c5_143_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_144_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(5120) : index + %output_size = check.literal value(48) : index + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<138240xi8> + %output = check.generate.fill value(0) : tensor<192xi8> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<20480xi8>, tensor<138240xi8>, tensor<192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_144_ggml_mul_mat_f32_f32_decode_wave64_case> @qwen38_27b_udq4kxl_tg_c5_144_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @qwen38_27b_udq4kxl_tg_c5_145_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<24576xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<20480xi8>, tensor<24576xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_145_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_tg_c5_145_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_tg_c5_146_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<40960xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<40960xi8>, tensor<163840xi8>, tensor<40960xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_146_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_tg_c5_146_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_147_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<192xi8> + %beta_raw = check.generate.fill value(0) : tensor<192xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<192xi8> + %beta_dst = check.generate.fill value(0) : tensor<192xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<192xi8>, tensor<192xi8>, tensor<192xi8>, tensor<192xi8>, tensor<192xi8>, tensor<192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_147_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_tg_c5_147_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_148_llm_gated_delta_net_f32_wmma_head128_inplace_case { + %q = check.generate.fill value(0) : tensor<8192xi8> + %k = check.generate.fill value(0) : tensor<8192xi8> + %v = check.generate.fill value(0) : tensor<24576xi8> + %g = check.generate.fill value(0) : tensor<192xi8> + %beta = check.generate.fill value(0) : tensor<192xi8> + %state_inout = check.generate.fill value(0) : tensor<3145728xi8> + %dst = check.generate.fill value(0) : tensor<24576xi8> + %rms_scales = check.generate.fill value(0) : tensor<192xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_inplace[](%q, %k, %v, %g, %beta, %state_inout, %dst, %rms_scales) : [](tensor<8192xi8>, tensor<8192xi8>, tensor<24576xi8>, tensor<192xi8>, tensor<192xi8>, tensor<3145728xi8>, tensor<24576xi8>, tensor<192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_148_llm_gated_delta_net_f32_wmma_head128_inplace_case> @qwen38_27b_udq4kxl_tg_c5_148_llm_gated_delta_net_f32_wmma_head128_inplace + +check.case public @qwen38_27b_udq4kxl_tg_c5_149_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(48) : index + %input = check.generate.fill value(0) : tensor<24576xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<24576xi8> + %output = check.generate.fill value(0) : tensor<24576xi8> + %quantized_values = check.generate.fill value(0) : tensor<3072xi8> + %scales = check.generate.fill value(0) : tensor<768xi8> + %sums = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<24576xi8>, tensor<512xi8>, tensor<24576xi8>, tensor<24576xi8>, tensor<3072xi8>, tensor<768xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_149_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_tg_c5_149_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_150_ggml_get_rows_f32_case { + %token_count = check.literal value(4) : index + %row_count = check.literal value(4) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<4xi32> + %weight = check.generate.fill value(0.0) : tensor<4x5120xf32> + %output = check.generate.fill value(1.0) : tensor<4x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<4xi32>, tensor<4x5120xf32>, tensor<4x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_150_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_tg_c5_150_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_tg_c5_151_qwen3_moe_dense_linear_q6k_f16_wmma_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<1042944000xi8> + %output = check.generate.fill value(0) : tensor<3973120xi8> + kernel.launch @qwen3_moe_dense_linear_q6k_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<81920xi8>, tensor<1042944000xi8>, tensor<3973120xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_tg_c5_151_qwen3_moe_dense_linear_q6k_f16_wmma_case> @qwen38_27b_udq4kxl_tg_c5_151_qwen3_moe_dense_linear_q6k_f16_wmma diff --git a/ggml/src/ggml-hrx/tools/benchmarks/loom/qwen38_27b_udq4kxl_recurrent.work.loom b/ggml/src/ggml-hrx/tools/benchmarks/loom/qwen38_27b_udq4kxl_recurrent.work.loom new file mode 100644 index 000000000000..2d1444805422 --- /dev/null +++ b/ggml/src/ggml-hrx/tools/benchmarks/loom/qwen38_27b_udq4kxl_recurrent.work.loom @@ -0,0 +1,425 @@ +// Generated by tools/benchmarks/generate-model-benchmarks.py for qwen38_27b_udq4kxl_recurrent scenarios: tg_c5. +// Regenerate from an HRX command program dump rather than editing by hand. + +target.decl @ggml_binary_f32_gfx11_wave64 +target.decl @ggml_binary_swiglu_i4_gfx11_wave32 +target.decl @ggml_copy_f32_gfx11_wave64 +target.decl @ggml_flash_attention_gfx11_wave64 +target.decl @ggml_get_rows_f32_gfx11_wave64 +target.decl @ggml_rmsnorm_binary_gfx11_wave32 +target.decl @ggml_rmsnorm_gfx11_wave32 +target.decl @ggml_set_rows_gfx11_wave64 +target.decl @qwen3_moe_dense_gfx11_wave64 + +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_add_rmsnorm_binary_symmetric_i4_k32(%token_count: index) launch(%token_count: index, %lhs: buffer, %rhs: buffer, %residual_output: buffer, %weight: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_binary_f32_gfx11_wave64) @ggml_binary_f32(%element_count: index) launch(%element_count: index, %lhs: buffer, %rhs: buffer, %output: buffer) +kernel.decl target(@ggml_binary_swiglu_i4_gfx11_wave32) @ggml_binary_swiglu_symmetric_i4_k32() launch(%lhs: buffer, %rhs: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_copy_f32_gfx11_wave64) @ggml_concat_dim0_f32() launch(%lhs: buffer, %rhs: buffer, %output: buffer) +kernel.decl target(@ggml_flash_attention_gfx11_wave64) @ggml_flash_attention_f32_f16_wmma(%query_token_count: index, %key_value_token_count: index) launch(%query_token_count: index, %key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %gate: buffer, %output: buffer) +kernel.decl target(@ggml_get_rows_f32_gfx11_wave64) @ggml_get_rows_f32(%token_count: index, %row_count: index, %hidden_size: index) launch(%token_count: index, %row_count: index, %hidden_size: index, %token_ids: buffer, %weight: buffer, %output: buffer) +kernel.decl target(@qwen3_moe_dense_gfx11_wave64) @ggml_mul_mat_q4_k_f32_adjacent_dual_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %first_weight: buffer, %second_weight: buffer, %first_output: buffer, %second_output: buffer) +kernel.decl @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma() launch(%first_weight: buffer, %second_weight: buffer, %first_output: buffer, %second_output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma() launch(%weight: buffer, %input: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer, %partial: buffer, %completion_counters: buffer) +kernel.decl @ggml_mul_mat_symmetric_i4_lowrow_wmma() launch(%weight: buffer, %input: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl @ggml_quantize_f32_symmetric_i4_k32() launch(%input: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_rmsnorm_binary_f32(%token_count: index) launch(%token_count: index, %input: buffer, %rhs: buffer, %output: buffer) +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_rmsnorm_binary_symmetric_i4_k32(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %raw_gate: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_rmsnorm_gfx11_wave32) @ggml_rmsnorm_mul_rope_f32() launch(%input: buffer, %weight: buffer, %positions: buffer, %output: buffer) +kernel.decl target(@ggml_set_rows_gfx11_wave64) @ggml_set_rows(%token_count: index, %cache_row_count: index, %hidden_size: index) launch(%token_count: index, %cache_row_count: index, %hidden_size: index, %rows: buffer, %indices: buffer, %cache: buffer) +kernel.decl @llm_gated_delta_net_f32_wmma_head128_snapshot() launch(%q: buffer, %k: buffer, %v: buffer, %g: buffer, %beta: buffer, %state_in: buffer, %snapshot_cache: buffer, %dst: buffer, %rms_scales: buffer) +kernel.decl @llm_gated_delta_net_projection_epilogue_f32() launch(%alpha_raw: buffer, %beta_raw: buffer, %bias: buffer, %a_scale: buffer, %gate_dst: buffer, %beta_dst: buffer) +kernel.decl @llm_ssm_conv_dconv4_silu_rollback_f32() launch(%state: buffer, %x: buffer, %filter: buffer, %output: buffer, %cache0: buffer, %cache1: buffer, %cache2: buffer, %cache3: buffer, %cache4: buffer) +kernel.decl target(@qwen3_moe_dense_gfx11_wave64) @qwen3_moe_dense_linear_q6k_f16_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) + + +// Scenario: tg_c5 + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_000_ggml_get_rows_f32_case { + %token_count = check.literal value(2) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<8xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<8xi8>, tensor<715161600xi8>, tensor<40960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_000_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_000_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_001_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<40960xi8>, tensor<20480xi8>, tensor<40960xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_001_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_001_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + %partial = check.generate.fill value(0) : tensor<81920xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<40960xi8>, tensor<81920xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>, tensor<81920xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_003_ggml_mul_mat_q4_k_f32_adjacent_dual_wmma_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<40960xi8> + %first_weight = check.generate.fill value(0) : tensor<138240xi8> + %second_weight = check.generate.fill value(0) : tensor<138240xi8> + %first_output = check.generate.fill value(0) : tensor<384xi8> + %second_output = check.generate.fill value(0) : tensor<384xi8> + kernel.launch @ggml_mul_mat_q4_k_f32_adjacent_dual_wmma[%token_count](%token_count, %input, %first_weight, %second_weight, %first_output, %second_output) : [index](index, tensor<40960xi8>, tensor<138240xi8>, tensor<138240xi8>, tensor<384xi8>, tensor<384xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_003_ggml_mul_mat_q4_k_f32_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_003_ggml_mul_mat_q4_k_f32_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_004_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<49152xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<40960xi8>, tensor<49152xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_004_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_004_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_005_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(20) : index + %hidden_size = check.literal value(30720) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<20x30720xf32> + %output = check.generate.fill value(1.0) : tensor<1x30720xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<20x30720xf32>, tensor<1x30720xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_005_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_005_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_006_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<81920xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<81920xi8>, tensor<163840xi8>, tensor<81920xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_006_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_006_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_007_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(20) : index + %hidden_size = check.literal value(786432) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<20x786432xf32> + %output = check.generate.fill value(1.0) : tensor<1x786432xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<20x786432xf32>, tensor<1x786432xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_007_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_007_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_008_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<384xi8> + %beta_raw = check.generate.fill value(0) : tensor<384xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<384xi8> + %beta_dst = check.generate.fill value(0) : tensor<384xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<384xi8>, tensor<384xi8>, tensor<192xi8>, tensor<192xi8>, tensor<384xi8>, tensor<384xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_008_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_008_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_009_llm_gated_delta_net_f32_wmma_head128_snapshot_case { + %q = check.generate.fill value(0) : tensor<49152xi8> + %k = check.generate.fill value(0) : tensor<49152xi8> + %v = check.generate.fill value(0) : tensor<65536xi8> + %g = check.generate.fill value(0) : tensor<384xi8> + %beta = check.generate.fill value(0) : tensor<384xi8> + %state_in = check.generate.fill value(0) : tensor<3145728xi8> + %snapshot_cache = check.generate.fill value(0) : tensor<15728640xi8> + %dst = check.generate.fill value(0) : tensor<49152xi8> + %rms_scales = check.generate.fill value(0) : tensor<384xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_snapshot[](%q, %k, %v, %g, %beta, %state_in, %snapshot_cache, %dst, %rms_scales) : [](tensor<49152xi8>, tensor<49152xi8>, tensor<65536xi8>, tensor<384xi8>, tensor<384xi8>, tensor<3145728xi8>, tensor<15728640xi8>, tensor<49152xi8>, tensor<384xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_009_llm_gated_delta_net_f32_wmma_head128_snapshot_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_009_llm_gated_delta_net_f32_wmma_head128_snapshot + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_010_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(96) : index + %input = check.generate.fill value(0) : tensor<49152xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<49152xi8> + %output = check.generate.fill value(0) : tensor<49152xi8> + %quantized_values = check.generate.fill value(0) : tensor<6144xi8> + %scales = check.generate.fill value(0) : tensor<1536xi8> + %sums = check.generate.fill value(0) : tensor<1536xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<49152xi8>, tensor<512xi8>, tensor<49152xi8>, tensor<49152xi8>, tensor<6144xi8>, tensor<1536xi8>, tensor<1536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_010_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_010_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_011_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<49152xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<6144xi8> + %scales = check.generate.fill value(0) : tensor<1536xi8> + %sums = check.generate.fill value(0) : tensor<1536xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<49152xi8>, tensor<40960xi8>, tensor<6144xi8>, tensor<1536xi8>, tensor<1536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_011_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_011_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_012_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(2) : index + %lhs = check.generate.fill value(0) : tensor<40960xi8> + %rhs = check.generate.fill value(0) : tensor<40960xi8> + %residual_output = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<40960xi8>, tensor<40960xi8>, tensor<40960xi8>, tensor<20480xi8>, tensor<40960xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_012_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_012_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_013_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<139264xi8> + %second_output = check.generate.fill value(0) : tensor<139264xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<139264xi8>, tensor<139264xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_013_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_013_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_014_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<139264xi8> + %rhs = check.generate.fill value(0) : tensor<139264xi8> + %output = check.generate.fill value(0) : tensor<139264xi8> + %quantized_values = check.generate.fill value(0) : tensor<17408xi8> + %scales = check.generate.fill value(0) : tensor<4352xi8> + %sums = check.generate.fill value(0) : tensor<4352xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<139264xi8>, tensor<139264xi8>, tensor<139264xi8>, tensor<17408xi8>, tensor<4352xi8>, tensor<4352xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_014_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_014_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_015_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<139264xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<17408xi8> + %scales = check.generate.fill value(0) : tensor<4352xi8> + %sums = check.generate.fill value(0) : tensor<4352xi8> + %partial = check.generate.fill value(0) : tensor<40960xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<139264xi8>, tensor<40960xi8>, tensor<17408xi8>, tensor<4352xi8>, tensor<4352xi8>, tensor<40960xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_015_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_015_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_016_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<98304xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<40960xi8>, tensor<98304xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_016_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_016_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_017_qwen3_moe_dense_linear_q6k_f16_wmma_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<4300800xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + kernel.launch @qwen3_moe_dense_linear_q6k_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<40960xi8>, tensor<4300800xi8>, tensor<8192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_017_qwen3_moe_dense_linear_q6k_f16_wmma_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_017_qwen3_moe_dense_linear_q6k_f16_wmma + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_018_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<40960xi8>, tensor<8192xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_018_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_018_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_019_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<97280xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<32xi8> + %output = check.generate.fill value(0) : tensor<49152xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<97280xi8>, tensor<1024xi8>, tensor<32xi8>, tensor<49152xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_019_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_019_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_020_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<8192xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<32xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<8192xi8>, tensor<1024xi8>, tensor<32xi8>, tensor<8192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_020_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_020_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_021_ggml_set_rows_case { + %token_count = check.literal value(2) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<2x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<2xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<2x1024xf32>, tensor<2xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_021_ggml_set_rows_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_021_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_022_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(2) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<2x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<2x256xf16> + %output = check.generate.fill value(1.0) : tensor<2x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<2x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<2x256xf16>, tensor<2x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_022_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_022_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_023_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<49152xi8> + %quantized_values = check.generate.fill value(0) : tensor<6144xi8> + %scales = check.generate.fill value(0) : tensor<1536xi8> + %sums = check.generate.fill value(0) : tensor<1536xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<49152xi8>, tensor<6144xi8>, tensor<1536xi8>, tensor<1536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_023_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_023_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_024_ggml_binary_f32_case { + %element_count = check.literal value(10240) : index + %lhs = check.generate.fill value(2.0) : tensor<10240xf32> + %rhs = check.generate.fill value(3.0) : tensor<10240xf32> + %output = check.generate.fill value(0.0) : tensor<10240xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<10240xf32>, tensor<10240xf32>, tensor<10240xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_024_ggml_binary_f32_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_024_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_025_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(2.0) : tensor<2x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<2x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<2x5120xf32>, tensor<5120xf32>, tensor<2x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_025_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_025_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_026_ggml_get_rows_f32_case { + %token_count = check.literal value(2) : index + %row_count = check.literal value(2) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<2xi32> + %weight = check.generate.fill value(0.0) : tensor<2x5120xf32> + %output = check.generate.fill value(1.0) : tensor<2x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<2xi32>, tensor<2x5120xf32>, tensor<2x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_026_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_026_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_027_qwen3_moe_dense_linear_q6k_f16_wmma_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<1042944000xi8> + %output = check.generate.fill value(0) : tensor<1986560xi8> + kernel.launch @qwen3_moe_dense_linear_q6k_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<40960xi8>, tensor<1042944000xi8>, tensor<1986560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_027_qwen3_moe_dense_linear_q6k_f16_wmma_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_027_qwen3_moe_dense_linear_q6k_f16_wmma + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_028_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<40960xi8> + %rhs = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<40960xi8>, tensor<40960xi8>, tensor<81920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_028_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_028_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_recurrent_tg_c5_029_qwen3_moe_dense_linear_q6k_f16_wmma_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<43008000xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + kernel.launch @qwen3_moe_dense_linear_q6k_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<81920xi8>, tensor<43008000xi8>, tensor<40960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_recurrent_tg_c5_029_qwen3_moe_dense_linear_q6k_f16_wmma_case> @qwen38_27b_udq4kxl_recurrent_tg_c5_029_qwen3_moe_dense_linear_q6k_f16_wmma diff --git a/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_draft_programs.tg_c1.json b/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_draft_programs.tg_c1.json new file mode 100644 index 000000000000..40e1175cfb47 --- /dev/null +++ b/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_draft_programs.tg_c1.json @@ -0,0 +1,5573 @@ +{ + "command_count": 8673, + "dispatch_count": 181, + "dispatches": [ + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_000_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "2", + "ggml.get_rows_f32.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 4, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 248320, + "token_count": 2 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_001_ggml_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 4, + "integer_parameters": { + "token_count": 2 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "10240", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_003_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "48", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_004_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_005_ggml_scale_bias_f32", + "compile_parameters": { + "ggml.scale.bias": "0", + "ggml.scale.scale": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": { + "element_count": 30720 + }, + "kernel": "loom_libs:ggml_scale_bias_f32", + "library_sources": [], + "primary_sources": [ + "ops/scale_bias_f32.loom" + ], + "sources": [ + "ops/scale_bias_f32.loom" + ], + "symbol": "ggml_scale_bias_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_006_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "30720", + "ggml.get_rows_f32.token_capacity": "1", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 384, + "integer_parameters": { + "hidden_size": 30720, + "row_count": 20, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_007_llm_ssm_conv_dconv4_silu_rollback_f32", + "compile_parameters": { + "llm.ssm_conv.rollback.cache_count": "5", + "llm.ssm_conv.rollback.d_inner": "10240", + "llm.ssm_conv.rollback.n_s": "1", + "llm.ssm_conv.rollback.n_t": "2", + "llm.ssm_conv.rollback.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:llm_ssm_conv_dconv4_silu_rollback_f32", + "library_sources": [], + "primary_sources": [ + "ops/ssm_conv_f32.loom" + ], + "sources": [ + "ops/ssm_conv_f32.loom" + ], + "symbol": "llm_ssm_conv_dconv4_silu_rollback_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_008_ggml_scale_bias_f32", + "compile_parameters": { + "ggml.scale.bias": "0", + "ggml.scale.scale": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": { + "element_count": 786432 + }, + "kernel": "loom_libs:ggml_scale_bias_f32", + "library_sources": [], + "primary_sources": [ + "ops/scale_bias_f32.loom" + ], + "sources": [ + "ops/scale_bias_f32.loom" + ], + "symbol": "ggml_scale_bias_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_009_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "786432", + "ggml.get_rows_f32.token_capacity": "1", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 384, + "integer_parameters": { + "hidden_size": 786432, + "row_count": 20, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_010_llm_gated_delta_net_projection_epilogue_f32", + "compile_parameters": { + "llm.gated_delta_net.epilogue_element_count": "96", + "llm.gated_delta_net.epilogue_head_count": "48", + "llm.gated_delta_net.epilogue_workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_projection_epilogue_f32", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_projection_epilogue_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_011_llm_gated_delta_net_f32_wmma_head128_snapshot", + "compile_parameters": { + "llm.gated_delta_net.head_count": "48", + "llm.gated_delta_net.head_width": "128", + "llm.gated_delta_net.l2_epsilon": "9.99999997e-07", + "llm.gated_delta_net.qk_stride1": "128", + "llm.gated_delta_net.qk_stride2": "10240", + "llm.gated_delta_net.qk_stride3": "20480", + "llm.gated_delta_net.query_head_count": "16", + "llm.gated_delta_net.query_sequence_ratio": "1", + "llm.gated_delta_net.scalar_stride1": "1", + "llm.gated_delta_net.scalar_stride2": "48", + "llm.gated_delta_net.scalar_stride3": "96", + "llm.gated_delta_net.sequence_count": "1", + "llm.gated_delta_net.snapshot_stride": "3145728", + "llm.gated_delta_net.token_count": "2", + "llm.gated_delta_net.value_stride1": "128", + "llm.gated_delta_net.value_stride2": "10240", + "llm.gated_delta_net.value_stride3": "20480", + "llm.gated_delta_net.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_f32_wmma_head128_snapshot", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_f32_wmma_head128_snapshot", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_012_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.hidden_size": "128", + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": { + "token_count": 96 + }, + "kernel": "loom_libs:ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_013_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 130, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_014_ggml_add_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.add_rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.add_rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 256, + "integer_parameters": { + "token_count": 2 + }, + "kernel": "loom_libs:ggml_add_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_add_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_015_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 130, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_016_ggml_binary_swiglu_symmetric_i4_k32", + "compile_parameters": { + "ggml.binary_swiglu_symmetric_i4.input_size": "17408", + "ggml.binary_swiglu_symmetric_i4.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 130, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_binary_swiglu_symmetric_i4_k32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_swiglu_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_017_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 130, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_018_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "12288", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_019_ggml_mul_mat_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat.input_size": "5120", + "ggml.mul_mat.output_accumulation": "0", + "ggml.mul_mat.output_size": "1024", + "ggml.mul_mat.output_unary_op": "23", + "ggml.mul_mat.weight_format": "6", + "ggml.workload.token_capacity": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": { + "token_count": 2 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_020_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "1024", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_021_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "512", + "ggml.rmsnorm_mul_rope.input_stride2": "12288", + "ggml.rmsnorm_mul_rope.input_stride3": "24576", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "24", + "ggml.rmsnorm_mul_rope.ne2": "2", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "6144", + "ggml.rmsnorm_mul_rope.output_stride3": "12288", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_022_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "256", + "ggml.rmsnorm_mul_rope.input_stride2": "1024", + "ggml.rmsnorm_mul_rope.input_stride3": "2048", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "4", + "ggml.rmsnorm_mul_rope.ne2": "2", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "1024", + "ggml.rmsnorm_mul_rope.output_stride3": "2048", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_023_ggml_set_rows", + "compile_parameters": { + "ggml.set_rows.hidden_capacity": "1024", + "ggml.set_rows.input_format": "32", + "ggml.set_rows.output_format": "16", + "ggml.set_rows.token_capacity": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 68, + "integer_parameters": { + "cache_row_count": 262144, + "hidden_size": 1024, + "token_count": 2 + }, + "kernel": "loom_libs:ggml_set_rows", + "library_sources": [], + "primary_sources": [ + "ops/set_rows.loom" + ], + "sources": [ + "ops/set_rows.loom" + ], + "symbol": "ggml_set_rows", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_024_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.apply_gate": "1", + "ggml.flash_attention.gate_stride_head": "512", + "ggml.flash_attention.gate_stride_token": "12288", + "ggml.flash_attention.head_size": "256", + "ggml.flash_attention.key_value_head_count": "4", + "ggml.flash_attention.query_head_count": "24" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": { + "key_value_token_count": 256, + "query_token_count": 2 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_025_ggml_quantize_f32_symmetric_i4_k32", + "compile_parameters": { + "ggml.quantize_symmetric_i4_k32.input_size": "6144", + "ggml.quantize_symmetric_i4_k32.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_quantize_f32_symmetric_i4_k32", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_quantize_f32_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_026_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 4, + "integer_parameters": { + "element_count": 10240 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_027_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "5120", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 8, + "integer_parameters": { + "token_count": 2 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_028_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "1", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 2, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_029_ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "248320" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 4, + "integer_parameters": { + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_030_ggml_concat_dim0_f32", + "compile_parameters": { + "ggml.concat_dim0_f32.lhs_width": "5120", + "ggml.concat_dim0_f32.rhs_width": "5120", + "ggml.concat_dim0_f32.row_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_concat_dim0_f32", + "library_sources": [], + "primary_sources": [ + "ops/copy_f32.loom" + ], + "sources": [ + "ops/copy_f32.loom" + ], + "symbol": "ggml_concat_dim0_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_031_ggml_mul_mat_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat.input_size": "10240", + "ggml.mul_mat.output_accumulation": "0", + "ggml.mul_mat.output_size": "5120", + "ggml.mul_mat.output_unary_op": "23", + "ggml.mul_mat.weight_format": "6", + "ggml.workload.token_capacity": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "token_count": 2 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_032_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "4", + "ggml.get_rows_f32.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 3, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 248320, + "token_count": 4 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_033_ggml_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 3, + "integer_parameters": { + "token_count": 4 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_034_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "10240", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_035_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "48", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_036_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_037_llm_ssm_conv_dconv4_silu_rollback_f32", + "compile_parameters": { + "llm.ssm_conv.rollback.cache_count": "5", + "llm.ssm_conv.rollback.d_inner": "10240", + "llm.ssm_conv.rollback.n_s": "1", + "llm.ssm_conv.rollback.n_t": "4", + "llm.ssm_conv.rollback.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:llm_ssm_conv_dconv4_silu_rollback_f32", + "library_sources": [], + "primary_sources": [ + "ops/ssm_conv_f32.loom" + ], + "sources": [ + "ops/ssm_conv_f32.loom" + ], + "symbol": "llm_ssm_conv_dconv4_silu_rollback_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_038_llm_gated_delta_net_projection_epilogue_f32", + "compile_parameters": { + "llm.gated_delta_net.epilogue_element_count": "192", + "llm.gated_delta_net.epilogue_head_count": "48", + "llm.gated_delta_net.epilogue_workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_projection_epilogue_f32", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_projection_epilogue_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_039_llm_gated_delta_net_f32_wmma_head128_snapshot", + "compile_parameters": { + "llm.gated_delta_net.head_count": "48", + "llm.gated_delta_net.head_width": "128", + "llm.gated_delta_net.l2_epsilon": "9.99999997e-07", + "llm.gated_delta_net.qk_stride1": "128", + "llm.gated_delta_net.qk_stride2": "10240", + "llm.gated_delta_net.qk_stride3": "40960", + "llm.gated_delta_net.query_head_count": "16", + "llm.gated_delta_net.query_sequence_ratio": "1", + "llm.gated_delta_net.scalar_stride1": "1", + "llm.gated_delta_net.scalar_stride2": "48", + "llm.gated_delta_net.scalar_stride3": "192", + "llm.gated_delta_net.sequence_count": "1", + "llm.gated_delta_net.snapshot_stride": "3145728", + "llm.gated_delta_net.token_count": "4", + "llm.gated_delta_net.value_stride1": "128", + "llm.gated_delta_net.value_stride2": "10240", + "llm.gated_delta_net.value_stride3": "40960", + "llm.gated_delta_net.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_f32_wmma_head128_snapshot", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_f32_wmma_head128_snapshot", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_040_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.hidden_size": "128", + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": { + "token_count": 192 + }, + "kernel": "loom_libs:ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_041_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 129, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_042_ggml_add_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.add_rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.add_rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 255, + "integer_parameters": { + "token_count": 4 + }, + "kernel": "loom_libs:ggml_add_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_add_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_043_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 129, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_044_ggml_binary_swiglu_symmetric_i4_k32", + "compile_parameters": { + "ggml.binary_swiglu_symmetric_i4.input_size": "17408", + "ggml.binary_swiglu_symmetric_i4.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 129, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_binary_swiglu_symmetric_i4_k32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_swiglu_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_045_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 129, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_046_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "12288", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 33, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_047_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "1024" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 33, + "integer_parameters": { + "token_count": 4 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_048_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "1024", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 33, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_049_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "512", + "ggml.rmsnorm_mul_rope.input_stride2": "12288", + "ggml.rmsnorm_mul_rope.input_stride3": "49152", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "24", + "ggml.rmsnorm_mul_rope.ne2": "4", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "6144", + "ggml.rmsnorm_mul_rope.output_stride3": "24576", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 33, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_050_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "256", + "ggml.rmsnorm_mul_rope.input_stride2": "1024", + "ggml.rmsnorm_mul_rope.input_stride3": "4096", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "4", + "ggml.rmsnorm_mul_rope.ne2": "4", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "1024", + "ggml.rmsnorm_mul_rope.output_stride3": "4096", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 33, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_051_ggml_set_rows", + "compile_parameters": { + "ggml.set_rows.hidden_capacity": "1024", + "ggml.set_rows.input_format": "32", + "ggml.set_rows.output_format": "16", + "ggml.set_rows.token_capacity": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 66, + "integer_parameters": { + "cache_row_count": 262144, + "hidden_size": 1024, + "token_count": 4 + }, + "kernel": "loom_libs:ggml_set_rows", + "library_sources": [], + "primary_sources": [ + "ops/set_rows.loom" + ], + "sources": [ + "ops/set_rows.loom" + ], + "symbol": "ggml_set_rows", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_052_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.apply_gate": "1", + "ggml.flash_attention.gate_stride_head": "512", + "ggml.flash_attention.gate_stride_token": "12288", + "ggml.flash_attention.head_size": "256", + "ggml.flash_attention.key_value_head_count": "4", + "ggml.flash_attention.query_head_count": "24" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 33, + "integer_parameters": { + "key_value_token_count": 256, + "query_token_count": 4 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_053_ggml_quantize_f32_symmetric_i4_k32", + "compile_parameters": { + "ggml.quantize_symmetric_i4_k32.input_size": "6144", + "ggml.quantize_symmetric_i4_k32.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 33, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_quantize_f32_symmetric_i4_k32", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_quantize_f32_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_054_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 3, + "integer_parameters": { + "element_count": 20480 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_055_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "5120", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 5, + "integer_parameters": { + "token_count": 4 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_056_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "4", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 4, + "token_count": 4 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_057_ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "248320" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 4 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_058_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "5", + "ggml.get_rows_f32.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 248320, + "token_count": 5 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_059_ggml_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "token_count": 5 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_060_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "10240", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_061_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "48", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_062_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_063_llm_ssm_conv_dconv4_silu_rollback_f32", + "compile_parameters": { + "llm.ssm_conv.rollback.cache_count": "5", + "llm.ssm_conv.rollback.d_inner": "10240", + "llm.ssm_conv.rollback.n_s": "1", + "llm.ssm_conv.rollback.n_t": "5", + "llm.ssm_conv.rollback.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_ssm_conv_dconv4_silu_rollback_f32", + "library_sources": [], + "primary_sources": [ + "ops/ssm_conv_f32.loom" + ], + "sources": [ + "ops/ssm_conv_f32.loom" + ], + "symbol": "llm_ssm_conv_dconv4_silu_rollback_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_064_llm_gated_delta_net_projection_epilogue_f32", + "compile_parameters": { + "llm.gated_delta_net.epilogue_element_count": "240", + "llm.gated_delta_net.epilogue_head_count": "48", + "llm.gated_delta_net.epilogue_workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_projection_epilogue_f32", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_projection_epilogue_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_065_llm_gated_delta_net_f32_wmma_head128_snapshot", + "compile_parameters": { + "llm.gated_delta_net.head_count": "48", + "llm.gated_delta_net.head_width": "128", + "llm.gated_delta_net.l2_epsilon": "9.99999997e-07", + "llm.gated_delta_net.qk_stride1": "128", + "llm.gated_delta_net.qk_stride2": "10240", + "llm.gated_delta_net.qk_stride3": "51200", + "llm.gated_delta_net.query_head_count": "16", + "llm.gated_delta_net.query_sequence_ratio": "1", + "llm.gated_delta_net.scalar_stride1": "1", + "llm.gated_delta_net.scalar_stride2": "48", + "llm.gated_delta_net.scalar_stride3": "240", + "llm.gated_delta_net.sequence_count": "1", + "llm.gated_delta_net.snapshot_stride": "3145728", + "llm.gated_delta_net.token_count": "5", + "llm.gated_delta_net.value_stride1": "128", + "llm.gated_delta_net.value_stride2": "10240", + "llm.gated_delta_net.value_stride3": "51200", + "llm.gated_delta_net.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_f32_wmma_head128_snapshot", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_f32_wmma_head128_snapshot", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_066_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.hidden_size": "128", + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "token_count": 240 + }, + "kernel": "loom_libs:ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_067_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_068_ggml_add_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.add_rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.add_rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 128, + "integer_parameters": { + "token_count": 5 + }, + "kernel": "loom_libs:ggml_add_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_add_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_069_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_070_ggml_binary_swiglu_symmetric_i4_k32", + "compile_parameters": { + "ggml.binary_swiglu_symmetric_i4.input_size": "17408", + "ggml.binary_swiglu_symmetric_i4.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_binary_swiglu_symmetric_i4_k32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_swiglu_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_071_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_072_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "12288", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_073_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "1024" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": { + "token_count": 5 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_074_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "1024", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_075_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "512", + "ggml.rmsnorm_mul_rope.input_stride2": "12288", + "ggml.rmsnorm_mul_rope.input_stride3": "61440", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "24", + "ggml.rmsnorm_mul_rope.ne2": "5", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "6144", + "ggml.rmsnorm_mul_rope.output_stride3": "30720", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_076_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "256", + "ggml.rmsnorm_mul_rope.input_stride2": "1024", + "ggml.rmsnorm_mul_rope.input_stride3": "5120", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "4", + "ggml.rmsnorm_mul_rope.ne2": "5", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "1024", + "ggml.rmsnorm_mul_rope.output_stride3": "5120", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_077_ggml_set_rows", + "compile_parameters": { + "ggml.set_rows.hidden_capacity": "1024", + "ggml.set_rows.input_format": "32", + "ggml.set_rows.output_format": "16", + "ggml.set_rows.token_capacity": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": { + "cache_row_count": 262144, + "hidden_size": 1024, + "token_count": 5 + }, + "kernel": "loom_libs:ggml_set_rows", + "library_sources": [], + "primary_sources": [ + "ops/set_rows.loom" + ], + "sources": [ + "ops/set_rows.loom" + ], + "symbol": "ggml_set_rows", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_078_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.apply_gate": "1", + "ggml.flash_attention.gate_stride_head": "512", + "ggml.flash_attention.gate_stride_token": "12288", + "ggml.flash_attention.head_size": "256", + "ggml.flash_attention.key_value_head_count": "4", + "ggml.flash_attention.query_head_count": "24" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": { + "key_value_token_count": 256, + "query_token_count": 5 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_079_ggml_quantize_f32_symmetric_i4_k32", + "compile_parameters": { + "ggml.quantize_symmetric_i4_k32.input_size": "6144", + "ggml.quantize_symmetric_i4_k32.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_quantize_f32_symmetric_i4_k32", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_quantize_f32_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_080_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "element_count": 25600 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_081_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "5120", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 4, + "integer_parameters": { + "token_count": 5 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_082_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "5", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 5, + "token_count": 5 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_083_ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "248320" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 5 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_084_ggml_concat_dim0_f32", + "compile_parameters": { + "ggml.concat_dim0_f32.lhs_width": "5120", + "ggml.concat_dim0_f32.rhs_width": "5120", + "ggml.concat_dim0_f32.row_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_concat_dim0_f32", + "library_sources": [], + "primary_sources": [ + "ops/copy_f32.loom" + ], + "sources": [ + "ops/copy_f32.loom" + ], + "symbol": "ggml_concat_dim0_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_085_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "10240", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "5120" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 5 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_086_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "3", + "ggml.get_rows_f32.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 248320, + "token_count": 3 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_087_ggml_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "token_count": 3 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_088_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "10240", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_089_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "48", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_090_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_091_llm_ssm_conv_dconv4_silu_rollback_f32", + "compile_parameters": { + "llm.ssm_conv.rollback.cache_count": "5", + "llm.ssm_conv.rollback.d_inner": "10240", + "llm.ssm_conv.rollback.n_s": "1", + "llm.ssm_conv.rollback.n_t": "3", + "llm.ssm_conv.rollback.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_ssm_conv_dconv4_silu_rollback_f32", + "library_sources": [], + "primary_sources": [ + "ops/ssm_conv_f32.loom" + ], + "sources": [ + "ops/ssm_conv_f32.loom" + ], + "symbol": "llm_ssm_conv_dconv4_silu_rollback_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_092_llm_gated_delta_net_projection_epilogue_f32", + "compile_parameters": { + "llm.gated_delta_net.epilogue_element_count": "144", + "llm.gated_delta_net.epilogue_head_count": "48", + "llm.gated_delta_net.epilogue_workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_projection_epilogue_f32", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_projection_epilogue_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_093_llm_gated_delta_net_f32_wmma_head128_snapshot", + "compile_parameters": { + "llm.gated_delta_net.head_count": "48", + "llm.gated_delta_net.head_width": "128", + "llm.gated_delta_net.l2_epsilon": "9.99999997e-07", + "llm.gated_delta_net.qk_stride1": "128", + "llm.gated_delta_net.qk_stride2": "10240", + "llm.gated_delta_net.qk_stride3": "30720", + "llm.gated_delta_net.query_head_count": "16", + "llm.gated_delta_net.query_sequence_ratio": "1", + "llm.gated_delta_net.scalar_stride1": "1", + "llm.gated_delta_net.scalar_stride2": "48", + "llm.gated_delta_net.scalar_stride3": "144", + "llm.gated_delta_net.sequence_count": "1", + "llm.gated_delta_net.snapshot_stride": "3145728", + "llm.gated_delta_net.token_count": "3", + "llm.gated_delta_net.value_stride1": "128", + "llm.gated_delta_net.value_stride2": "10240", + "llm.gated_delta_net.value_stride3": "30720", + "llm.gated_delta_net.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_f32_wmma_head128_snapshot", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_f32_wmma_head128_snapshot", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_094_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.hidden_size": "128", + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "token_count": 144 + }, + "kernel": "loom_libs:ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_095_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_096_ggml_add_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.add_rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.add_rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 128, + "integer_parameters": { + "token_count": 3 + }, + "kernel": "loom_libs:ggml_add_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_add_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_097_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_098_ggml_binary_swiglu_symmetric_i4_k32", + "compile_parameters": { + "ggml.binary_swiglu_symmetric_i4.input_size": "17408", + "ggml.binary_swiglu_symmetric_i4.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_binary_swiglu_symmetric_i4_k32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_swiglu_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_099_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_100_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "12288", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_101_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "1024" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": { + "token_count": 3 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_102_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "1024", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_103_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "512", + "ggml.rmsnorm_mul_rope.input_stride2": "12288", + "ggml.rmsnorm_mul_rope.input_stride3": "36864", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "24", + "ggml.rmsnorm_mul_rope.ne2": "3", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "6144", + "ggml.rmsnorm_mul_rope.output_stride3": "18432", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_104_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "256", + "ggml.rmsnorm_mul_rope.input_stride2": "1024", + "ggml.rmsnorm_mul_rope.input_stride3": "3072", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "4", + "ggml.rmsnorm_mul_rope.ne2": "3", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "1024", + "ggml.rmsnorm_mul_rope.output_stride3": "3072", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_105_ggml_set_rows", + "compile_parameters": { + "ggml.set_rows.hidden_capacity": "1024", + "ggml.set_rows.input_format": "32", + "ggml.set_rows.output_format": "16", + "ggml.set_rows.token_capacity": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": { + "cache_row_count": 262144, + "hidden_size": 1024, + "token_count": 3 + }, + "kernel": "loom_libs:ggml_set_rows", + "library_sources": [], + "primary_sources": [ + "ops/set_rows.loom" + ], + "sources": [ + "ops/set_rows.loom" + ], + "symbol": "ggml_set_rows", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_106_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.apply_gate": "1", + "ggml.flash_attention.gate_stride_head": "512", + "ggml.flash_attention.gate_stride_token": "12288", + "ggml.flash_attention.head_size": "256", + "ggml.flash_attention.key_value_head_count": "4", + "ggml.flash_attention.query_head_count": "24" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": { + "key_value_token_count": 256, + "query_token_count": 3 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_107_ggml_quantize_f32_symmetric_i4_k32", + "compile_parameters": { + "ggml.quantize_symmetric_i4_k32.input_size": "6144", + "ggml.quantize_symmetric_i4_k32.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_quantize_f32_symmetric_i4_k32", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_quantize_f32_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_108_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "element_count": 15360 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_109_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "5120", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 4, + "integer_parameters": { + "token_count": 3 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_110_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "3", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 3, + "token_count": 3 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_111_ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "248320" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 3 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_112_ggml_concat_dim0_f32", + "compile_parameters": { + "ggml.concat_dim0_f32.lhs_width": "5120", + "ggml.concat_dim0_f32.rhs_width": "5120", + "ggml.concat_dim0_f32.row_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_concat_dim0_f32", + "library_sources": [], + "primary_sources": [ + "ops/copy_f32.loom" + ], + "sources": [ + "ops/copy_f32.loom" + ], + "symbol": "ggml_concat_dim0_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_113_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "10240", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "5120" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 3 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_114_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "16", + "ggml.get_rows_f32.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 248320, + "token_count": 16 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_115_ggml_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "token_count": 16 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_116_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "10240", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_117_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "48", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_118_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_119_llm_ssm_conv_dconv4_silu_rollback_f32", + "compile_parameters": { + "llm.ssm_conv.rollback.cache_count": "5", + "llm.ssm_conv.rollback.d_inner": "10240", + "llm.ssm_conv.rollback.n_s": "1", + "llm.ssm_conv.rollback.n_t": "16", + "llm.ssm_conv.rollback.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_ssm_conv_dconv4_silu_rollback_f32", + "library_sources": [], + "primary_sources": [ + "ops/ssm_conv_f32.loom" + ], + "sources": [ + "ops/ssm_conv_f32.loom" + ], + "symbol": "llm_ssm_conv_dconv4_silu_rollback_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_120_llm_gated_delta_net_projection_epilogue_f32", + "compile_parameters": { + "llm.gated_delta_net.epilogue_element_count": "768", + "llm.gated_delta_net.epilogue_head_count": "48", + "llm.gated_delta_net.epilogue_workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_projection_epilogue_f32", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_projection_epilogue_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_121_llm_gated_delta_net_f32_wmma_head128", + "compile_parameters": { + "llm.gated_delta_net.head_count": "48", + "llm.gated_delta_net.head_width": "128", + "llm.gated_delta_net.l2_epsilon": "9.99999997e-07", + "llm.gated_delta_net.qk_stride1": "128", + "llm.gated_delta_net.qk_stride2": "10240", + "llm.gated_delta_net.qk_stride3": "163840", + "llm.gated_delta_net.query_head_count": "16", + "llm.gated_delta_net.query_sequence_ratio": "1", + "llm.gated_delta_net.scalar_stride1": "1", + "llm.gated_delta_net.scalar_stride2": "48", + "llm.gated_delta_net.scalar_stride3": "768", + "llm.gated_delta_net.sequence_count": "1", + "llm.gated_delta_net.token_count": "11", + "llm.gated_delta_net.value_stride1": "128", + "llm.gated_delta_net.value_stride2": "10240", + "llm.gated_delta_net.value_stride3": "163840", + "llm.gated_delta_net.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_f32_wmma_head128", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_f32_wmma_head128", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_122_ggml_copy_f32", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 288, + "integer_parameters": { + "element_count": 786432 + }, + "kernel": "loom_libs:ggml_copy_f32", + "library_sources": [], + "primary_sources": [ + "ops/copy_f32.loom" + ], + "sources": [ + "ops/copy_f32.loom" + ], + "symbol": "ggml_copy_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_123_llm_gated_delta_net_f32_wmma_head128_inplace", + "compile_parameters": { + "llm.gated_delta_net.head_count": "48", + "llm.gated_delta_net.head_width": "128", + "llm.gated_delta_net.l2_epsilon": "9.99999997e-07", + "llm.gated_delta_net.qk_stride1": "128", + "llm.gated_delta_net.qk_stride2": "10240", + "llm.gated_delta_net.qk_stride3": "163840", + "llm.gated_delta_net.query_head_count": "16", + "llm.gated_delta_net.query_sequence_ratio": "1", + "llm.gated_delta_net.scalar_stride1": "1", + "llm.gated_delta_net.scalar_stride2": "48", + "llm.gated_delta_net.scalar_stride3": "768", + "llm.gated_delta_net.sequence_count": "1", + "llm.gated_delta_net.token_count": "1", + "llm.gated_delta_net.value_stride1": "128", + "llm.gated_delta_net.value_stride2": "10240", + "llm.gated_delta_net.value_stride3": "163840", + "llm.gated_delta_net.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 240, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_f32_wmma_head128_inplace", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_f32_wmma_head128_inplace", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_124_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.hidden_size": "128", + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "token_count": 768 + }, + "kernel": "loom_libs:ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_125_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_126_ggml_add_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.add_rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.add_rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 128, + "integer_parameters": { + "token_count": 16 + }, + "kernel": "loom_libs:ggml_add_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_add_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_127_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_128_ggml_binary_swiglu_symmetric_i4_k32", + "compile_parameters": { + "ggml.binary_swiglu_symmetric_i4.input_size": "17408", + "ggml.binary_swiglu_symmetric_i4.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_binary_swiglu_symmetric_i4_k32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_swiglu_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_129_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_130_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "12288", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_131_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "1024" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": { + "token_count": 16 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_132_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "1024", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_133_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "512", + "ggml.rmsnorm_mul_rope.input_stride2": "12288", + "ggml.rmsnorm_mul_rope.input_stride3": "196608", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "24", + "ggml.rmsnorm_mul_rope.ne2": "16", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "6144", + "ggml.rmsnorm_mul_rope.output_stride3": "98304", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_134_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "256", + "ggml.rmsnorm_mul_rope.input_stride2": "1024", + "ggml.rmsnorm_mul_rope.input_stride3": "16384", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "4", + "ggml.rmsnorm_mul_rope.ne2": "16", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "1024", + "ggml.rmsnorm_mul_rope.output_stride3": "16384", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_135_ggml_set_rows", + "compile_parameters": { + "ggml.set_rows.hidden_capacity": "1024", + "ggml.set_rows.input_format": "32", + "ggml.set_rows.output_format": "16", + "ggml.set_rows.token_capacity": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": { + "cache_row_count": 262144, + "hidden_size": 1024, + "token_count": 16 + }, + "kernel": "loom_libs:ggml_set_rows", + "library_sources": [], + "primary_sources": [ + "ops/set_rows.loom" + ], + "sources": [ + "ops/set_rows.loom" + ], + "symbol": "ggml_set_rows", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_136_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.apply_gate": "1", + "ggml.flash_attention.gate_stride_head": "512", + "ggml.flash_attention.gate_stride_token": "12288", + "ggml.flash_attention.head_size": "256", + "ggml.flash_attention.key_value_head_count": "4", + "ggml.flash_attention.query_head_count": "24" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": { + "key_value_token_count": 256, + "query_token_count": 16 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_137_ggml_quantize_f32_symmetric_i4_k32", + "compile_parameters": { + "ggml.quantize_symmetric_i4_k32.input_size": "6144", + "ggml.quantize_symmetric_i4_k32.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_quantize_f32_symmetric_i4_k32", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_quantize_f32_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_138_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "element_count": 81920 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_139_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "5120", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 4, + "integer_parameters": { + "token_count": 16 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_140_ggml_concat_dim0_f32", + "compile_parameters": { + "ggml.concat_dim0_f32.lhs_width": "5120", + "ggml.concat_dim0_f32.rhs_width": "5120", + "ggml.concat_dim0_f32.row_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_concat_dim0_f32", + "library_sources": [], + "primary_sources": [ + "ops/copy_f32.loom" + ], + "sources": [ + "ops/copy_f32.loom" + ], + "symbol": "ggml_concat_dim0_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_141_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "10240", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "5120" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 16 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_142_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "1", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 4, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_143_ggml_concat_dim0_f32", + "compile_parameters": { + "ggml.concat_dim0_f32.lhs_width": "5120", + "ggml.concat_dim0_f32.rhs_width": "5120", + "ggml.concat_dim0_f32.row_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_concat_dim0_f32", + "library_sources": [], + "primary_sources": [ + "ops/copy_f32.loom" + ], + "sources": [ + "ops/copy_f32.loom" + ], + "symbol": "ggml_concat_dim0_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_144_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "10240", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "5120" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 4 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_145_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "1", + "ggml.get_rows_f32.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 3, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 248320, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_146_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "5120", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 7, + "integer_parameters": { + "token_count": 1 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_147_ggml_concat_dim0_f32", + "compile_parameters": { + "ggml.concat_dim0_f32.lhs_width": "5120", + "ggml.concat_dim0_f32.rhs_width": "5120", + "ggml.concat_dim0_f32.row_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_concat_dim0_f32", + "library_sources": [], + "primary_sources": [ + "ops/copy_f32.loom" + ], + "sources": [ + "ops/copy_f32.loom" + ], + "symbol": "ggml_concat_dim0_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_148_ggml_mul_mat_f32_f32_decode_wave64", + "compile_parameters": { + "ggml.mul_mat_f32_f32_decode.output_capacity": "5120", + "ggml.mul_mat_f32_f32_decode.token_capacity": "1", + "ggml.mul_mat_f32_f32_decode.weight_format": "6" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "input_size": 10240, + "output_size": 5120, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_149_ggml_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 3, + "integer_parameters": { + "token_count": 1 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_150_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "12288", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 18, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_151_ggml_mul_mat_f32_f32_decode_wave64", + "compile_parameters": { + "ggml.mul_mat_f32_f32_decode.output_capacity": "1024", + "ggml.mul_mat_f32_f32_decode.token_capacity": "1", + "ggml.mul_mat_f32_f32_decode.weight_format": "6" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 18, + "integer_parameters": { + "input_size": 5120, + "output_size": 1024, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_152_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "1024", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 18, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_153_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "512", + "ggml.rmsnorm_mul_rope.input_stride2": "12288", + "ggml.rmsnorm_mul_rope.input_stride3": "12288", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "24", + "ggml.rmsnorm_mul_rope.ne2": "1", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "6144", + "ggml.rmsnorm_mul_rope.output_stride3": "6144", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 18, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_154_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "256", + "ggml.rmsnorm_mul_rope.input_stride2": "1024", + "ggml.rmsnorm_mul_rope.input_stride3": "1024", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "4", + "ggml.rmsnorm_mul_rope.ne2": "1", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "1024", + "ggml.rmsnorm_mul_rope.output_stride3": "1024", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 18, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_155_ggml_set_rows", + "compile_parameters": { + "ggml.set_rows.hidden_capacity": "1024", + "ggml.set_rows.input_format": "32", + "ggml.set_rows.output_format": "16", + "ggml.set_rows.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 36, + "integer_parameters": { + "cache_row_count": 262144, + "hidden_size": 1024, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_set_rows", + "library_sources": [], + "primary_sources": [ + "ops/set_rows.loom" + ], + "sources": [ + "ops/set_rows.loom" + ], + "symbol": "ggml_set_rows", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_156_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.apply_gate": "1", + "ggml.flash_attention.gate_stride_head": "512", + "ggml.flash_attention.gate_stride_token": "12288", + "ggml.flash_attention.head_size": "256", + "ggml.flash_attention.key_value_head_count": "4", + "ggml.flash_attention.query_head_count": "24" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 18, + "integer_parameters": { + "key_value_token_count": 256, + "query_token_count": 1 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_157_ggml_quantize_f32_symmetric_i4_k32", + "compile_parameters": { + "ggml.quantize_symmetric_i4_k32.input_size": "6144", + "ggml.quantize_symmetric_i4_k32.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 18, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_quantize_f32_symmetric_i4_k32", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_quantize_f32_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_158_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 66, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_159_ggml_add_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.add_rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.add_rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 129, + "integer_parameters": { + "token_count": 1 + }, + "kernel": "loom_libs:ggml_add_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_add_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_160_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 66, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_161_ggml_binary_swiglu_symmetric_i4_k32", + "compile_parameters": { + "ggml.binary_swiglu_symmetric_i4.input_size": "17408", + "ggml.binary_swiglu_symmetric_i4.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 66, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_binary_swiglu_symmetric_i4_k32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_swiglu_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_162_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 66, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_163_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 3, + "integer_parameters": { + "element_count": 5120 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_164_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "1", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 1, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_165_ggml_quantize_f32_symmetric_i4_k32", + "compile_parameters": { + "ggml.quantize_symmetric_i4_k32.input_size": "5120", + "ggml.quantize_symmetric_i4_k32.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_quantize_f32_symmetric_i4_k32", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_quantize_f32_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_166_ggml_mul_mat_q6_k_symmetric_i2_scan_token1", + "compile_parameters": { + "ggml.mul_mat_q6_k_shortlist.input_size": "5120", + "ggml.mul_mat_q6_k_shortlist.output_size": "248320" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_q6_k_symmetric_i2_scan_token1", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_symmetric_i2_scan_token1", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_167_ggml_top_k8_f32_partitions_register", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "element_count": 248320 + }, + "kernel": "loom_libs:ggml_top_k8_f32_partitions_register", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_top_k8_f32_partitions_register", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_168_ggml_top_k128_f32_reduce_gather_register", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "element_count": 248320 + }, + "kernel": "loom_libs:ggml_top_k128_f32_reduce_gather_register", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_top_k128_f32_reduce_gather_register", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_169_ggml_fill_negative_f32", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "element_count": 248320 + }, + "kernel": "loom_libs:ggml_fill_negative_f32", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_fill_negative_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_170_ggml_mul_mat_q6_k_packed_selected_refine_token1", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "248320", + "ggml.mul_mat_q6_k_packed.weight_offset": "337715200" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "candidate_count": 64, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_packed_selected_refine_token1", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_packed_selected_refine_token1", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "candidate_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_171_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "10240", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_172_ggml_mul_mat_f32_f32_decode_wave64", + "compile_parameters": { + "ggml.mul_mat_f32_f32_decode.output_capacity": "48", + "ggml.mul_mat_f32_f32_decode.token_capacity": "1", + "ggml.mul_mat_f32_f32_decode.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": { + "input_size": 5120, + "output_size": 48, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_173_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_174_llm_ssm_conv_dconv4_silu_rollback_f32", + "compile_parameters": { + "llm.ssm_conv.rollback.cache_count": "5", + "llm.ssm_conv.rollback.d_inner": "10240", + "llm.ssm_conv.rollback.n_s": "1", + "llm.ssm_conv.rollback.n_t": "1", + "llm.ssm_conv.rollback.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_ssm_conv_dconv4_silu_rollback_f32", + "library_sources": [], + "primary_sources": [ + "ops/ssm_conv_f32.loom" + ], + "sources": [ + "ops/ssm_conv_f32.loom" + ], + "symbol": "llm_ssm_conv_dconv4_silu_rollback_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_175_llm_gated_delta_net_projection_epilogue_f32", + "compile_parameters": { + "llm.gated_delta_net.epilogue_element_count": "48", + "llm.gated_delta_net.epilogue_head_count": "48", + "llm.gated_delta_net.epilogue_workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_projection_epilogue_f32", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_projection_epilogue_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_176_llm_gated_delta_net_f32_wmma_head128_inplace", + "compile_parameters": { + "llm.gated_delta_net.head_count": "48", + "llm.gated_delta_net.head_width": "128", + "llm.gated_delta_net.l2_epsilon": "9.99999997e-07", + "llm.gated_delta_net.qk_stride1": "128", + "llm.gated_delta_net.qk_stride2": "10240", + "llm.gated_delta_net.qk_stride3": "10240", + "llm.gated_delta_net.query_head_count": "16", + "llm.gated_delta_net.query_sequence_ratio": "1", + "llm.gated_delta_net.scalar_stride1": "1", + "llm.gated_delta_net.scalar_stride2": "48", + "llm.gated_delta_net.scalar_stride3": "48", + "llm.gated_delta_net.sequence_count": "1", + "llm.gated_delta_net.token_count": "1", + "llm.gated_delta_net.value_stride1": "128", + "llm.gated_delta_net.value_stride2": "10240", + "llm.gated_delta_net.value_stride3": "10240", + "llm.gated_delta_net.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_f32_wmma_head128_inplace", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_f32_wmma_head128_inplace", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_177_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.hidden_size": "128", + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "token_count": 48 + }, + "kernel": "loom_libs:ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_178_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "2", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 2, + "token_count": 2 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_draft_tg_c1_179_ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "248320" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 2 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + } + ], + "generated_count": 180, + "kernel_counts": { + "loom_libs:ggml_add_rmsnorm_binary_symmetric_i4_k32": 1024, + "loom_libs:ggml_binary_f32": 16, + "loom_libs:ggml_binary_swiglu_symmetric_i4_k32": 520, + "loom_libs:ggml_concat_dim0_f32": 8, + "loom_libs:ggml_copy_f32": 288, + "loom_libs:ggml_fill_negative_f32": 1, + "loom_libs:ggml_flash_attention_f32_f16_wmma": 136, + "loom_libs:ggml_get_rows_f32": 793, + "loom_libs:ggml_mul_mat_f32_f32_decode_wave64": 116, + "loom_libs:ggml_mul_mat_f32_f32_wmma": 36, + "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma": 88, + "loom_libs:ggml_mul_mat_q6_k_packed_selected_refine_token1": 2, + "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma": 8, + "loom_libs:ggml_mul_mat_q6_k_symmetric_i2_scan_token1": 1, + "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma": 856, + "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma": 904, + "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma": 1176, + "loom_libs:ggml_quantize_f32_symmetric_i4_k32": 137, + "loom_libs:ggml_rmsnorm_binary_f32": 32, + "loom_libs:ggml_rmsnorm_binary_symmetric_i4_k32": 16, + "loom_libs:ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32": 384, + "loom_libs:ggml_rmsnorm_mul_rope_f32": 272, + "loom_libs:ggml_scale_bias_f32": 192, + "loom_libs:ggml_select_symmetric_i4_k32_groups": 1, + "loom_libs:ggml_set_rows": 272, + "loom_libs:ggml_top_k128_f32_reduce_gather_register": 1, + "loom_libs:ggml_top_k8_f32_partitions_register": 1, + "loom_libs:llm_gated_delta_net_f32_wmma_head128": 48, + "loom_libs:llm_gated_delta_net_f32_wmma_head128_inplace": 288, + "loom_libs:llm_gated_delta_net_f32_wmma_head128_snapshot": 288, + "loom_libs:llm_gated_delta_net_projection_epilogue_f32": 384, + "loom_libs:llm_ssm_conv_dconv4_silu_rollback_f32": 384 + }, + "loom_source": "tools/benchmarks/loom/v2_draft_programs.tg_c1.work.loom", + "model": "qwen38_27b_udq4kxl_draft", + "scenario": "tg_c1", + "schema": "ggml-hrx-model-loom-benchmarks-v2" +} diff --git a/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_draft_programs.tg_c1.work.loom b/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_draft_programs.tg_c1.work.loom new file mode 100644 index 000000000000..6c3225e9b420 --- /dev/null +++ b/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_draft_programs.tg_c1.work.loom @@ -0,0 +1,2377 @@ +// Generated by tools/benchmarks/generate-model-benchmarks.py for qwen38_27b_udq4kxl_draft scenarios: tg_c1. +// Regenerate from an HRX command program dump rather than editing by hand. + +target.decl @ggml_binary_f32_gfx11_wave64 +target.decl @ggml_binary_swiglu_i4_gfx11_wave32 +target.decl @ggml_copy_f32_gfx11_wave64 +target.decl @ggml_flash_attention_gfx11_wave64 +target.decl @ggml_get_rows_f32_gfx11_wave64 +target.decl @ggml_mul_mat_f32_f32_decode_gfx11_wave64 +target.decl @ggml_mul_mat_gfx11_wave64 +target.decl @ggml_rmsnorm_binary_gfx11_wave32 +target.decl @ggml_rmsnorm_gfx11_wave32 +target.decl @ggml_scale_bias_f32_gfx11_wave64 +target.decl @ggml_set_rows_gfx11_wave64 + +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_add_rmsnorm_binary_symmetric_i4_k32(%token_count: index) launch(%token_count: index, %lhs: buffer, %rhs: buffer, %residual_output: buffer, %weight: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_binary_f32_gfx11_wave64) @ggml_binary_f32(%element_count: index) launch(%element_count: index, %lhs: buffer, %rhs: buffer, %output: buffer) +kernel.decl target(@ggml_binary_swiglu_i4_gfx11_wave32) @ggml_binary_swiglu_symmetric_i4_k32() launch(%lhs: buffer, %rhs: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_copy_f32_gfx11_wave64) @ggml_concat_dim0_f32() launch(%lhs: buffer, %rhs: buffer, %output: buffer) +kernel.decl target(@ggml_copy_f32_gfx11_wave64) @ggml_copy_f32(%element_count: index) launch(%element_count: index, %source: buffer, %output: buffer) +kernel.decl @ggml_fill_negative_f32(%element_count: index) launch(%element_count: index, %output: buffer) +kernel.decl target(@ggml_flash_attention_gfx11_wave64) @ggml_flash_attention_f32_f16_wmma(%query_token_count: index, %key_value_token_count: index) launch(%query_token_count: index, %key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %gate: buffer, %output: buffer) +kernel.decl target(@ggml_get_rows_f32_gfx11_wave64) @ggml_get_rows_f32(%token_count: index, %row_count: index, %hidden_size: index) launch(%token_count: index, %row_count: index, %hidden_size: index, %token_ids: buffer, %weight: buffer, %output: buffer) +kernel.decl target(@ggml_mul_mat_f32_f32_decode_gfx11_wave64) @ggml_mul_mat_f32_f32_decode_wave64(%token_count: index, %input_size: index, %output_size: index) launch(%token_count: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %output: buffer) +kernel.decl target(@ggml_mul_mat_gfx11_wave64) @ggml_mul_mat_f32_f32_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) +kernel.decl @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) +kernel.decl @ggml_mul_mat_q6_k_packed_selected_refine_token1(%token_count: index, %candidate_count: index) launch(%token_count: index, %candidate_count: index, %input: buffer, %weight: buffer, %candidates: buffer, %exact_output: buffer, %output: buffer) +kernel.decl @ggml_mul_mat_q6_k_packed_token1_f16_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) +kernel.decl @ggml_mul_mat_q6_k_symmetric_i2_scan_token1() launch(%weight: buffer, %output: buffer, %qact: buffer, %scales: buffer, %selected_mask: buffer) +kernel.decl @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma() launch(%first_weight: buffer, %second_weight: buffer, %first_output: buffer, %second_output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma() launch(%weight: buffer, %input: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer, %partial: buffer, %completion_counters: buffer) +kernel.decl @ggml_mul_mat_symmetric_i4_lowrow_wmma() launch(%weight: buffer, %input: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl @ggml_quantize_f32_symmetric_i4_k32() launch(%input: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_rmsnorm_binary_f32(%token_count: index) launch(%token_count: index, %input: buffer, %rhs: buffer, %output: buffer) +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_rmsnorm_binary_symmetric_i4_k32(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %raw_gate: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_rmsnorm_gfx11_wave32) @ggml_rmsnorm_mul_rope_f32() launch(%input: buffer, %weight: buffer, %positions: buffer, %output: buffer) +kernel.decl target(@ggml_scale_bias_f32_gfx11_wave64) @ggml_scale_bias_f32(%element_count: index) launch(%element_count: index, %input: buffer, %output: buffer) +kernel.decl @ggml_select_symmetric_i4_k32_groups() launch(%input: buffer, %selected_groups: buffer) +kernel.decl target(@ggml_set_rows_gfx11_wave64) @ggml_set_rows(%token_count: index, %cache_row_count: index, %hidden_size: index) launch(%token_count: index, %cache_row_count: index, %hidden_size: index, %rows: buffer, %indices: buffer, %cache: buffer) +kernel.decl @ggml_top_k128_f32_reduce_gather_register(%element_count: index) launch(%element_count: index, %partial_values: buffer, %partial_ids: buffer, %candidate_output: buffer, %value_output: buffer) +kernel.decl @ggml_top_k8_f32_partitions_register(%element_count: index) launch(%element_count: index, %values: buffer, %partial_values: buffer, %partial_ids: buffer) +kernel.decl @llm_gated_delta_net_f32_wmma_head128() launch(%q: buffer, %k: buffer, %v: buffer, %g: buffer, %beta: buffer, %state_in: buffer, %dst: buffer, %rms_scales: buffer) +kernel.decl @llm_gated_delta_net_f32_wmma_head128_inplace() launch(%q: buffer, %k: buffer, %v: buffer, %g: buffer, %beta: buffer, %state_inout: buffer, %dst: buffer, %rms_scales: buffer) +kernel.decl @llm_gated_delta_net_f32_wmma_head128_snapshot() launch(%q: buffer, %k: buffer, %v: buffer, %g: buffer, %beta: buffer, %state_in: buffer, %snapshot_cache: buffer, %dst: buffer, %rms_scales: buffer) +kernel.decl @llm_gated_delta_net_projection_epilogue_f32() launch(%alpha_raw: buffer, %beta_raw: buffer, %bias: buffer, %a_scale: buffer, %gate_dst: buffer, %beta_dst: buffer) +kernel.decl @llm_ssm_conv_dconv4_silu_rollback_f32() launch(%state: buffer, %x: buffer, %filter: buffer, %output: buffer, %cache0: buffer, %cache1: buffer, %cache2: buffer, %cache3: buffer, %cache4: buffer) + + +// Scenario: tg_c1 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_000_ggml_get_rows_f32_case { + %token_count = check.literal value(2) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<8xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<8xi8>, tensor<715161600xi8>, tensor<40960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_000_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_000_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_001_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<40960xi8>, tensor<20480xi8>, tensor<40960xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_001_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_001_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + %partial = check.generate.fill value(0) : tensor<81920xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<40960xi8>, tensor<81920xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>, tensor<81920xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_003_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<168960xi8> + %second_weight = check.generate.fill value(0) : tensor<168960xi8> + %first_output = check.generate.fill value(0) : tensor<384xi8> + %second_output = check.generate.fill value(0) : tensor<384xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<168960xi8>, tensor<168960xi8>, tensor<384xi8>, tensor<384xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_003_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_003_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_004_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<49152xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<40960xi8>, tensor<49152xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_004_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_004_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_005_ggml_scale_bias_f32_case { + %element_count = check.literal value(30720) : index + %input = check.generate.fill value(0) : tensor<122880xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @ggml_scale_bias_f32[%element_count](%element_count, %input, %output) : [index](index, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_005_ggml_scale_bias_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_005_ggml_scale_bias_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_006_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(20) : index + %hidden_size = check.literal value(30720) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<20x30720xf32> + %output = check.generate.fill value(1.0) : tensor<1x30720xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<20x30720xf32>, tensor<1x30720xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_006_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_006_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_007_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<81920xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<81920xi8>, tensor<163840xi8>, tensor<81920xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_007_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_007_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_008_ggml_scale_bias_f32_case { + %element_count = check.literal value(786432) : index + %input = check.generate.fill value(0) : tensor<3145728xi8> + %output = check.generate.fill value(0) : tensor<3145728xi8> + kernel.launch @ggml_scale_bias_f32[%element_count](%element_count, %input, %output) : [index](index, tensor<3145728xi8>, tensor<3145728xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_008_ggml_scale_bias_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_008_ggml_scale_bias_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_009_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(20) : index + %hidden_size = check.literal value(786432) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<20x786432xf32> + %output = check.generate.fill value(1.0) : tensor<1x786432xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<20x786432xf32>, tensor<1x786432xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_009_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_009_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_010_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<384xi8> + %beta_raw = check.generate.fill value(0) : tensor<384xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<384xi8> + %beta_dst = check.generate.fill value(0) : tensor<384xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<384xi8>, tensor<384xi8>, tensor<192xi8>, tensor<192xi8>, tensor<384xi8>, tensor<384xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_010_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_010_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_011_llm_gated_delta_net_f32_wmma_head128_snapshot_case { + %q = check.generate.fill value(0) : tensor<49152xi8> + %k = check.generate.fill value(0) : tensor<49152xi8> + %v = check.generate.fill value(0) : tensor<65536xi8> + %g = check.generate.fill value(0) : tensor<384xi8> + %beta = check.generate.fill value(0) : tensor<384xi8> + %state_in = check.generate.fill value(0) : tensor<3145728xi8> + %snapshot_cache = check.generate.fill value(0) : tensor<15728640xi8> + %dst = check.generate.fill value(0) : tensor<49152xi8> + %rms_scales = check.generate.fill value(0) : tensor<384xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_snapshot[](%q, %k, %v, %g, %beta, %state_in, %snapshot_cache, %dst, %rms_scales) : [](tensor<49152xi8>, tensor<49152xi8>, tensor<65536xi8>, tensor<384xi8>, tensor<384xi8>, tensor<3145728xi8>, tensor<15728640xi8>, tensor<49152xi8>, tensor<384xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_011_llm_gated_delta_net_f32_wmma_head128_snapshot_case> @qwen38_27b_udq4kxl_draft_tg_c1_011_llm_gated_delta_net_f32_wmma_head128_snapshot + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_012_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(96) : index + %input = check.generate.fill value(0) : tensor<49152xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<49152xi8> + %output = check.generate.fill value(0) : tensor<49152xi8> + %quantized_values = check.generate.fill value(0) : tensor<6144xi8> + %scales = check.generate.fill value(0) : tensor<1536xi8> + %sums = check.generate.fill value(0) : tensor<1536xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<49152xi8>, tensor<512xi8>, tensor<49152xi8>, tensor<49152xi8>, tensor<6144xi8>, tensor<1536xi8>, tensor<1536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_012_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_012_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_013_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<49152xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<6144xi8> + %scales = check.generate.fill value(0) : tensor<1536xi8> + %sums = check.generate.fill value(0) : tensor<1536xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<49152xi8>, tensor<40960xi8>, tensor<6144xi8>, tensor<1536xi8>, tensor<1536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_013_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_013_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_014_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(2) : index + %lhs = check.generate.fill value(0) : tensor<40960xi8> + %rhs = check.generate.fill value(0) : tensor<40960xi8> + %residual_output = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<40960xi8>, tensor<40960xi8>, tensor<40960xi8>, tensor<20480xi8>, tensor<40960xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_014_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_014_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_015_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<139264xi8> + %second_output = check.generate.fill value(0) : tensor<139264xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<139264xi8>, tensor<139264xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_015_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_015_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_016_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<139264xi8> + %rhs = check.generate.fill value(0) : tensor<139264xi8> + %output = check.generate.fill value(0) : tensor<139264xi8> + %quantized_values = check.generate.fill value(0) : tensor<17408xi8> + %scales = check.generate.fill value(0) : tensor<4352xi8> + %sums = check.generate.fill value(0) : tensor<4352xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<139264xi8>, tensor<139264xi8>, tensor<139264xi8>, tensor<17408xi8>, tensor<4352xi8>, tensor<4352xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_016_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_016_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_017_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<139264xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<17408xi8> + %scales = check.generate.fill value(0) : tensor<4352xi8> + %sums = check.generate.fill value(0) : tensor<4352xi8> + %partial = check.generate.fill value(0) : tensor<40960xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<139264xi8>, tensor<40960xi8>, tensor<17408xi8>, tensor<4352xi8>, tensor<4352xi8>, tensor<40960xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_017_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_017_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_018_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<98304xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<40960xi8>, tensor<98304xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_018_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_018_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_019_ggml_mul_mat_f32_f32_wmma_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<4300800xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + kernel.launch @ggml_mul_mat_f32_f32_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<40960xi8>, tensor<4300800xi8>, tensor<8192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_019_ggml_mul_mat_f32_f32_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_019_ggml_mul_mat_f32_f32_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_020_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<40960xi8>, tensor<8192xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_020_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_020_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_021_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<97280xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<32xi8> + %output = check.generate.fill value(0) : tensor<49152xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<97280xi8>, tensor<1024xi8>, tensor<32xi8>, tensor<49152xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_021_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_021_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_022_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<8192xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<32xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<8192xi8>, tensor<1024xi8>, tensor<32xi8>, tensor<8192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_022_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_022_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_023_ggml_set_rows_case { + %token_count = check.literal value(2) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<2x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<2xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<2x1024xf32>, tensor<2xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_023_ggml_set_rows_case> @qwen38_27b_udq4kxl_draft_tg_c1_023_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_024_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(2) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<2x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<2x256xf16> + %output = check.generate.fill value(1.0) : tensor<2x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<2x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<2x256xf16>, tensor<2x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_024_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_024_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_025_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<49152xi8> + %quantized_values = check.generate.fill value(0) : tensor<6144xi8> + %scales = check.generate.fill value(0) : tensor<1536xi8> + %sums = check.generate.fill value(0) : tensor<1536xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<49152xi8>, tensor<6144xi8>, tensor<1536xi8>, tensor<1536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_025_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_025_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_026_ggml_binary_f32_case { + %element_count = check.literal value(10240) : index + %lhs = check.generate.fill value(2.0) : tensor<10240xf32> + %rhs = check.generate.fill value(3.0) : tensor<10240xf32> + %output = check.generate.fill value(0.0) : tensor<10240xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<10240xf32>, tensor<10240xf32>, tensor<10240xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_026_ggml_binary_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_026_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_027_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(2.0) : tensor<2x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<2x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<2x5120xf32>, tensor<5120xf32>, tensor<2x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_027_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_027_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_028_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(2) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<2x5120xf32> + %output = check.generate.fill value(1.0) : tensor<1x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<2x5120xf32>, tensor<1x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_028_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_028_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_029_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<1042944000xi8> + %output = check.generate.fill value(0) : tensor<993280xi8> + kernel.launch @ggml_mul_mat_q6_k_packed_token1_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<20480xi8>, tensor<1042944000xi8>, tensor<993280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_029_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_029_ggml_mul_mat_q6_k_packed_token1_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_030_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<40960xi8> + %rhs = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<40960xi8>, tensor<40960xi8>, tensor<81920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_030_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_030_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_031_ggml_mul_mat_f32_f32_wmma_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<43008000xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + kernel.launch @ggml_mul_mat_f32_f32_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<81920xi8>, tensor<43008000xi8>, tensor<40960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_031_ggml_mul_mat_f32_f32_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_031_ggml_mul_mat_f32_f32_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_032_ggml_get_rows_f32_case { + %token_count = check.literal value(4) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<16xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<16xi8>, tensor<715161600xi8>, tensor<81920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_032_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_032_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_033_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<81920xi8>, tensor<20480xi8>, tensor<81920xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_033_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_033_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_034_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<163840xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + %partial = check.generate.fill value(0) : tensor<163840xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<81920xi8>, tensor<163840xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>, tensor<163840xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_034_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_034_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_035_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<168960xi8> + %second_weight = check.generate.fill value(0) : tensor<168960xi8> + %first_output = check.generate.fill value(0) : tensor<768xi8> + %second_output = check.generate.fill value(0) : tensor<768xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<168960xi8>, tensor<168960xi8>, tensor<768xi8>, tensor<768xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_035_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_035_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_036_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<98304xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<81920xi8>, tensor<98304xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_036_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_036_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_037_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<163840xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<163840xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<163840xi8>, tensor<163840xi8>, tensor<163840xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_037_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_037_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_038_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<768xi8> + %beta_raw = check.generate.fill value(0) : tensor<768xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<768xi8> + %beta_dst = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<768xi8>, tensor<768xi8>, tensor<192xi8>, tensor<192xi8>, tensor<768xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_038_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_038_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_039_llm_gated_delta_net_f32_wmma_head128_snapshot_case { + %q = check.generate.fill value(0) : tensor<131072xi8> + %k = check.generate.fill value(0) : tensor<131072xi8> + %v = check.generate.fill value(0) : tensor<147456xi8> + %g = check.generate.fill value(0) : tensor<768xi8> + %beta = check.generate.fill value(0) : tensor<768xi8> + %state_in = check.generate.fill value(0) : tensor<3145728xi8> + %snapshot_cache = check.generate.fill value(0) : tensor<40894464xi8> + %dst = check.generate.fill value(0) : tensor<98304xi8> + %rms_scales = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_snapshot[](%q, %k, %v, %g, %beta, %state_in, %snapshot_cache, %dst, %rms_scales) : [](tensor<131072xi8>, tensor<131072xi8>, tensor<147456xi8>, tensor<768xi8>, tensor<768xi8>, tensor<3145728xi8>, tensor<40894464xi8>, tensor<98304xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_039_llm_gated_delta_net_f32_wmma_head128_snapshot_case> @qwen38_27b_udq4kxl_draft_tg_c1_039_llm_gated_delta_net_f32_wmma_head128_snapshot + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_040_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(192) : index + %input = check.generate.fill value(0) : tensor<98304xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<98304xi8> + %output = check.generate.fill value(0) : tensor<98304xi8> + %quantized_values = check.generate.fill value(0) : tensor<12288xi8> + %scales = check.generate.fill value(0) : tensor<3072xi8> + %sums = check.generate.fill value(0) : tensor<3072xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<98304xi8>, tensor<512xi8>, tensor<98304xi8>, tensor<98304xi8>, tensor<12288xi8>, tensor<3072xi8>, tensor<3072xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_040_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_040_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_041_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<98304xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<12288xi8> + %scales = check.generate.fill value(0) : tensor<3072xi8> + %sums = check.generate.fill value(0) : tensor<3072xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<98304xi8>, tensor<81920xi8>, tensor<12288xi8>, tensor<3072xi8>, tensor<3072xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_041_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_041_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_042_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(4) : index + %lhs = check.generate.fill value(0) : tensor<81920xi8> + %rhs = check.generate.fill value(0) : tensor<81920xi8> + %residual_output = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<81920xi8>, tensor<81920xi8>, tensor<81920xi8>, tensor<20480xi8>, tensor<81920xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_042_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_042_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_043_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<278528xi8> + %second_output = check.generate.fill value(0) : tensor<278528xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<278528xi8>, tensor<278528xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_043_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_043_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_044_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<278528xi8> + %rhs = check.generate.fill value(0) : tensor<278528xi8> + %output = check.generate.fill value(0) : tensor<278528xi8> + %quantized_values = check.generate.fill value(0) : tensor<34816xi8> + %scales = check.generate.fill value(0) : tensor<8704xi8> + %sums = check.generate.fill value(0) : tensor<8704xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<278528xi8>, tensor<278528xi8>, tensor<278528xi8>, tensor<34816xi8>, tensor<8704xi8>, tensor<8704xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_044_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_044_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_045_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<278528xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<34816xi8> + %scales = check.generate.fill value(0) : tensor<8704xi8> + %sums = check.generate.fill value(0) : tensor<8704xi8> + %partial = check.generate.fill value(0) : tensor<81920xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<278528xi8>, tensor<81920xi8>, tensor<34816xi8>, tensor<8704xi8>, tensor<8704xi8>, tensor<81920xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_045_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_045_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_046_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<196608xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<81920xi8>, tensor<196608xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_046_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_046_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_047_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<5611520xi8> + %output = check.generate.fill value(0) : tensor<16384xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<81920xi8>, tensor<5611520xi8>, tensor<16384xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_047_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_047_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_048_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<16384xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<81920xi8>, tensor<16384xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_048_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_048_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_049_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<195584xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<64xi8> + %output = check.generate.fill value(0) : tensor<98304xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<195584xi8>, tensor<1024xi8>, tensor<64xi8>, tensor<98304xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_049_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_049_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_050_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<16384xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<64xi8> + %output = check.generate.fill value(0) : tensor<16384xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<16384xi8>, tensor<1024xi8>, tensor<64xi8>, tensor<16384xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_050_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_050_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_051_ggml_set_rows_case { + %token_count = check.literal value(4) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<4x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<4xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<4x1024xf32>, tensor<4xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_051_ggml_set_rows_case> @qwen38_27b_udq4kxl_draft_tg_c1_051_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_052_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(4) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<4x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<4x256xf16> + %output = check.generate.fill value(1.0) : tensor<4x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<4x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<4x256xf16>, tensor<4x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_052_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_052_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_053_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<98304xi8> + %quantized_values = check.generate.fill value(0) : tensor<12288xi8> + %scales = check.generate.fill value(0) : tensor<3072xi8> + %sums = check.generate.fill value(0) : tensor<3072xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<98304xi8>, tensor<12288xi8>, tensor<3072xi8>, tensor<3072xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_053_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_053_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_054_ggml_binary_f32_case { + %element_count = check.literal value(20480) : index + %lhs = check.generate.fill value(2.0) : tensor<20480xf32> + %rhs = check.generate.fill value(3.0) : tensor<20480xf32> + %output = check.generate.fill value(0.0) : tensor<20480xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<20480xf32>, tensor<20480xf32>, tensor<20480xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_054_ggml_binary_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_054_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_055_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(2.0) : tensor<4x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<4x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<4x5120xf32>, tensor<5120xf32>, tensor<4x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_055_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_055_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_056_ggml_get_rows_f32_case { + %token_count = check.literal value(4) : index + %row_count = check.literal value(4) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<4xi32> + %weight = check.generate.fill value(0.0) : tensor<4x5120xf32> + %output = check.generate.fill value(1.0) : tensor<4x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<4xi32>, tensor<4x5120xf32>, tensor<4x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_056_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_056_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_057_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<1042944000xi8> + %output = check.generate.fill value(0) : tensor<3973120xi8> + kernel.launch @ggml_mul_mat_q6_k_packed_token1_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<81920xi8>, tensor<1042944000xi8>, tensor<3973120xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_057_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_057_ggml_mul_mat_q6_k_packed_token1_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_058_ggml_get_rows_f32_case { + %token_count = check.literal value(5) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<20xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<20xi8>, tensor<715161600xi8>, tensor<102400xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_058_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_058_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_059_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(5) : index + %input = check.generate.fill value(0) : tensor<102400xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<102400xi8>, tensor<20480xi8>, tensor<102400xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_059_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_059_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_060_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<102400xi8> + %output = check.generate.fill value(0) : tensor<204800xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + %partial = check.generate.fill value(0) : tensor<204800xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<102400xi8>, tensor<204800xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>, tensor<204800xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_060_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_060_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_061_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<168960xi8> + %second_weight = check.generate.fill value(0) : tensor<168960xi8> + %first_output = check.generate.fill value(0) : tensor<960xi8> + %second_output = check.generate.fill value(0) : tensor<960xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<168960xi8>, tensor<168960xi8>, tensor<960xi8>, tensor<960xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_061_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_061_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_062_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<102400xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<102400xi8>, tensor<122880xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_062_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_062_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_063_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<204800xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<204800xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<204800xi8>, tensor<163840xi8>, tensor<204800xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_063_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_063_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_064_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<960xi8> + %beta_raw = check.generate.fill value(0) : tensor<960xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<960xi8> + %beta_dst = check.generate.fill value(0) : tensor<960xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<960xi8>, tensor<960xi8>, tensor<192xi8>, tensor<192xi8>, tensor<960xi8>, tensor<960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_064_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_064_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_065_llm_gated_delta_net_f32_wmma_head128_snapshot_case { + %q = check.generate.fill value(0) : tensor<172032xi8> + %k = check.generate.fill value(0) : tensor<172032xi8> + %v = check.generate.fill value(0) : tensor<188416xi8> + %g = check.generate.fill value(0) : tensor<960xi8> + %beta = check.generate.fill value(0) : tensor<960xi8> + %state_in = check.generate.fill value(0) : tensor<3145728xi8> + %snapshot_cache = check.generate.fill value(0) : tensor<53477376xi8> + %dst = check.generate.fill value(0) : tensor<122880xi8> + %rms_scales = check.generate.fill value(0) : tensor<960xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_snapshot[](%q, %k, %v, %g, %beta, %state_in, %snapshot_cache, %dst, %rms_scales) : [](tensor<172032xi8>, tensor<172032xi8>, tensor<188416xi8>, tensor<960xi8>, tensor<960xi8>, tensor<3145728xi8>, tensor<53477376xi8>, tensor<122880xi8>, tensor<960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_065_llm_gated_delta_net_f32_wmma_head128_snapshot_case> @qwen38_27b_udq4kxl_draft_tg_c1_065_llm_gated_delta_net_f32_wmma_head128_snapshot + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_066_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(240) : index + %input = check.generate.fill value(0) : tensor<122880xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<122880xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + %quantized_values = check.generate.fill value(0) : tensor<15360xi8> + %scales = check.generate.fill value(0) : tensor<3840xi8> + %sums = check.generate.fill value(0) : tensor<3840xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<122880xi8>, tensor<512xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<15360xi8>, tensor<3840xi8>, tensor<3840xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_066_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_066_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_067_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<122880xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + %quantized_values = check.generate.fill value(0) : tensor<15360xi8> + %scales = check.generate.fill value(0) : tensor<3840xi8> + %sums = check.generate.fill value(0) : tensor<3840xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<122880xi8>, tensor<102400xi8>, tensor<15360xi8>, tensor<3840xi8>, tensor<3840xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_067_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_067_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_068_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(5) : index + %lhs = check.generate.fill value(0) : tensor<102400xi8> + %rhs = check.generate.fill value(0) : tensor<102400xi8> + %residual_output = check.generate.fill value(0) : tensor<102400xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<102400xi8>, tensor<102400xi8>, tensor<102400xi8>, tensor<20480xi8>, tensor<102400xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_068_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_068_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_069_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<348160xi8> + %second_output = check.generate.fill value(0) : tensor<348160xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<348160xi8>, tensor<348160xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_069_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_069_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_070_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<348160xi8> + %rhs = check.generate.fill value(0) : tensor<348160xi8> + %output = check.generate.fill value(0) : tensor<348160xi8> + %quantized_values = check.generate.fill value(0) : tensor<43520xi8> + %scales = check.generate.fill value(0) : tensor<10880xi8> + %sums = check.generate.fill value(0) : tensor<10880xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<348160xi8>, tensor<348160xi8>, tensor<348160xi8>, tensor<43520xi8>, tensor<10880xi8>, tensor<10880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_070_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_070_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_071_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<348160xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + %quantized_values = check.generate.fill value(0) : tensor<43520xi8> + %scales = check.generate.fill value(0) : tensor<10880xi8> + %sums = check.generate.fill value(0) : tensor<10880xi8> + %partial = check.generate.fill value(0) : tensor<102400xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<348160xi8>, tensor<102400xi8>, tensor<43520xi8>, tensor<10880xi8>, tensor<10880xi8>, tensor<102400xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_071_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_071_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_072_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<102400xi8> + %output = check.generate.fill value(0) : tensor<245760xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<102400xi8>, tensor<245760xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_072_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_072_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_073_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(5) : index + %input = check.generate.fill value(0) : tensor<102400xi8> + %weight = check.generate.fill value(0) : tensor<5611520xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<102400xi8>, tensor<5611520xi8>, tensor<20480xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_073_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_073_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_074_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<102400xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<102400xi8>, tensor<20480xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_074_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_074_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_075_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<244736xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<80xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<244736xi8>, tensor<1024xi8>, tensor<80xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_075_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_075_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_076_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<80xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<20480xi8>, tensor<1024xi8>, tensor<80xi8>, tensor<20480xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_076_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_076_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_077_ggml_set_rows_case { + %token_count = check.literal value(5) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<5x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<5xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<5x1024xf32>, tensor<5xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_077_ggml_set_rows_case> @qwen38_27b_udq4kxl_draft_tg_c1_077_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_078_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(5) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<5x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<5x256xf16> + %output = check.generate.fill value(1.0) : tensor<5x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<5x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<5x256xf16>, tensor<5x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_078_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_078_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_079_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<122880xi8> + %quantized_values = check.generate.fill value(0) : tensor<15360xi8> + %scales = check.generate.fill value(0) : tensor<3840xi8> + %sums = check.generate.fill value(0) : tensor<3840xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<122880xi8>, tensor<15360xi8>, tensor<3840xi8>, tensor<3840xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_079_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_079_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_080_ggml_binary_f32_case { + %element_count = check.literal value(25600) : index + %lhs = check.generate.fill value(2.0) : tensor<25600xf32> + %rhs = check.generate.fill value(3.0) : tensor<25600xf32> + %output = check.generate.fill value(0.0) : tensor<25600xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<25600xf32>, tensor<25600xf32>, tensor<25600xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_080_ggml_binary_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_080_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_081_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(5) : index + %input = check.generate.fill value(2.0) : tensor<5x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<5x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<5x5120xf32>, tensor<5120xf32>, tensor<5x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_081_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_081_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_082_ggml_get_rows_f32_case { + %token_count = check.literal value(5) : index + %row_count = check.literal value(5) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<5xi32> + %weight = check.generate.fill value(0.0) : tensor<5x5120xf32> + %output = check.generate.fill value(1.0) : tensor<5x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<5xi32>, tensor<5x5120xf32>, tensor<5x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_082_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_082_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_083_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case { + %token_count = check.literal value(5) : index + %input = check.generate.fill value(0) : tensor<102400xi8> + %weight = check.generate.fill value(0) : tensor<1042944000xi8> + %output = check.generate.fill value(0) : tensor<4966400xi8> + kernel.launch @ggml_mul_mat_q6_k_packed_token1_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<102400xi8>, tensor<1042944000xi8>, tensor<4966400xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_083_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_083_ggml_mul_mat_q6_k_packed_token1_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_084_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<102400xi8> + %rhs = check.generate.fill value(0) : tensor<102400xi8> + %output = check.generate.fill value(0) : tensor<204800xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<102400xi8>, tensor<102400xi8>, tensor<204800xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_084_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_084_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_085_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(5) : index + %input = check.generate.fill value(0) : tensor<204800xi8> + %weight = check.generate.fill value(0) : tensor<56115200xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<204800xi8>, tensor<56115200xi8>, tensor<102400xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_085_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_085_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_086_ggml_get_rows_f32_case { + %token_count = check.literal value(3) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<12xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<61440xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<12xi8>, tensor<715161600xi8>, tensor<61440xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_086_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_086_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_087_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(3) : index + %input = check.generate.fill value(0) : tensor<61440xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<61440xi8> + %quantized_values = check.generate.fill value(0) : tensor<7680xi8> + %scales = check.generate.fill value(0) : tensor<1920xi8> + %sums = check.generate.fill value(0) : tensor<1920xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<61440xi8>, tensor<20480xi8>, tensor<61440xi8>, tensor<7680xi8>, tensor<1920xi8>, tensor<1920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_087_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_087_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_088_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<61440xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + %quantized_values = check.generate.fill value(0) : tensor<7680xi8> + %scales = check.generate.fill value(0) : tensor<1920xi8> + %sums = check.generate.fill value(0) : tensor<1920xi8> + %partial = check.generate.fill value(0) : tensor<122880xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<61440xi8>, tensor<122880xi8>, tensor<7680xi8>, tensor<1920xi8>, tensor<1920xi8>, tensor<122880xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_088_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_088_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_089_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<168960xi8> + %second_weight = check.generate.fill value(0) : tensor<168960xi8> + %first_output = check.generate.fill value(0) : tensor<576xi8> + %second_output = check.generate.fill value(0) : tensor<576xi8> + %quantized_values = check.generate.fill value(0) : tensor<7680xi8> + %scales = check.generate.fill value(0) : tensor<1920xi8> + %sums = check.generate.fill value(0) : tensor<1920xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<168960xi8>, tensor<168960xi8>, tensor<576xi8>, tensor<576xi8>, tensor<7680xi8>, tensor<1920xi8>, tensor<1920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_089_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_089_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_090_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<61440xi8> + %output = check.generate.fill value(0) : tensor<73728xi8> + %quantized_values = check.generate.fill value(0) : tensor<7680xi8> + %scales = check.generate.fill value(0) : tensor<1920xi8> + %sums = check.generate.fill value(0) : tensor<1920xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<61440xi8>, tensor<73728xi8>, tensor<7680xi8>, tensor<1920xi8>, tensor<1920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_090_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_090_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_091_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<122880xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<122880xi8>, tensor<163840xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_091_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_091_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_092_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<576xi8> + %beta_raw = check.generate.fill value(0) : tensor<576xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<576xi8> + %beta_dst = check.generate.fill value(0) : tensor<576xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<576xi8>, tensor<576xi8>, tensor<192xi8>, tensor<192xi8>, tensor<576xi8>, tensor<576xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_092_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_092_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_093_llm_gated_delta_net_f32_wmma_head128_snapshot_case { + %q = check.generate.fill value(0) : tensor<90112xi8> + %k = check.generate.fill value(0) : tensor<90112xi8> + %v = check.generate.fill value(0) : tensor<106496xi8> + %g = check.generate.fill value(0) : tensor<576xi8> + %beta = check.generate.fill value(0) : tensor<576xi8> + %state_in = check.generate.fill value(0) : tensor<3145728xi8> + %snapshot_cache = check.generate.fill value(0) : tensor<28311552xi8> + %dst = check.generate.fill value(0) : tensor<73728xi8> + %rms_scales = check.generate.fill value(0) : tensor<576xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_snapshot[](%q, %k, %v, %g, %beta, %state_in, %snapshot_cache, %dst, %rms_scales) : [](tensor<90112xi8>, tensor<90112xi8>, tensor<106496xi8>, tensor<576xi8>, tensor<576xi8>, tensor<3145728xi8>, tensor<28311552xi8>, tensor<73728xi8>, tensor<576xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_093_llm_gated_delta_net_f32_wmma_head128_snapshot_case> @qwen38_27b_udq4kxl_draft_tg_c1_093_llm_gated_delta_net_f32_wmma_head128_snapshot + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_094_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(144) : index + %input = check.generate.fill value(0) : tensor<73728xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<73728xi8> + %output = check.generate.fill value(0) : tensor<73728xi8> + %quantized_values = check.generate.fill value(0) : tensor<9216xi8> + %scales = check.generate.fill value(0) : tensor<2304xi8> + %sums = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<73728xi8>, tensor<512xi8>, tensor<73728xi8>, tensor<73728xi8>, tensor<9216xi8>, tensor<2304xi8>, tensor<2304xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_094_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_094_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_095_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<73728xi8> + %output = check.generate.fill value(0) : tensor<61440xi8> + %quantized_values = check.generate.fill value(0) : tensor<9216xi8> + %scales = check.generate.fill value(0) : tensor<2304xi8> + %sums = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<73728xi8>, tensor<61440xi8>, tensor<9216xi8>, tensor<2304xi8>, tensor<2304xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_095_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_095_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_096_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(3) : index + %lhs = check.generate.fill value(0) : tensor<61440xi8> + %rhs = check.generate.fill value(0) : tensor<61440xi8> + %residual_output = check.generate.fill value(0) : tensor<61440xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<61440xi8> + %quantized_values = check.generate.fill value(0) : tensor<7680xi8> + %scales = check.generate.fill value(0) : tensor<1920xi8> + %sums = check.generate.fill value(0) : tensor<1920xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<61440xi8>, tensor<61440xi8>, tensor<61440xi8>, tensor<20480xi8>, tensor<61440xi8>, tensor<7680xi8>, tensor<1920xi8>, tensor<1920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_096_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_096_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_097_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<208896xi8> + %second_output = check.generate.fill value(0) : tensor<208896xi8> + %quantized_values = check.generate.fill value(0) : tensor<7680xi8> + %scales = check.generate.fill value(0) : tensor<1920xi8> + %sums = check.generate.fill value(0) : tensor<1920xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<208896xi8>, tensor<208896xi8>, tensor<7680xi8>, tensor<1920xi8>, tensor<1920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_097_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_097_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_098_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<208896xi8> + %rhs = check.generate.fill value(0) : tensor<208896xi8> + %output = check.generate.fill value(0) : tensor<208896xi8> + %quantized_values = check.generate.fill value(0) : tensor<26112xi8> + %scales = check.generate.fill value(0) : tensor<6528xi8> + %sums = check.generate.fill value(0) : tensor<6528xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<208896xi8>, tensor<208896xi8>, tensor<208896xi8>, tensor<26112xi8>, tensor<6528xi8>, tensor<6528xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_098_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_098_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_099_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<208896xi8> + %output = check.generate.fill value(0) : tensor<61440xi8> + %quantized_values = check.generate.fill value(0) : tensor<26112xi8> + %scales = check.generate.fill value(0) : tensor<6528xi8> + %sums = check.generate.fill value(0) : tensor<6528xi8> + %partial = check.generate.fill value(0) : tensor<61440xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<208896xi8>, tensor<61440xi8>, tensor<26112xi8>, tensor<6528xi8>, tensor<6528xi8>, tensor<61440xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_099_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_099_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_100_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<61440xi8> + %output = check.generate.fill value(0) : tensor<147456xi8> + %quantized_values = check.generate.fill value(0) : tensor<7680xi8> + %scales = check.generate.fill value(0) : tensor<1920xi8> + %sums = check.generate.fill value(0) : tensor<1920xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<61440xi8>, tensor<147456xi8>, tensor<7680xi8>, tensor<1920xi8>, tensor<1920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_100_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_100_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_101_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(3) : index + %input = check.generate.fill value(0) : tensor<61440xi8> + %weight = check.generate.fill value(0) : tensor<5611520xi8> + %output = check.generate.fill value(0) : tensor<12288xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<61440xi8>, tensor<5611520xi8>, tensor<12288xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_101_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_101_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_102_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<61440xi8> + %output = check.generate.fill value(0) : tensor<12288xi8> + %quantized_values = check.generate.fill value(0) : tensor<7680xi8> + %scales = check.generate.fill value(0) : tensor<1920xi8> + %sums = check.generate.fill value(0) : tensor<1920xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<61440xi8>, tensor<12288xi8>, tensor<7680xi8>, tensor<1920xi8>, tensor<1920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_102_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_102_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_103_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<146432xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<48xi8> + %output = check.generate.fill value(0) : tensor<73728xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<146432xi8>, tensor<1024xi8>, tensor<48xi8>, tensor<73728xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_103_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_103_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_104_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<12288xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<48xi8> + %output = check.generate.fill value(0) : tensor<12288xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<12288xi8>, tensor<1024xi8>, tensor<48xi8>, tensor<12288xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_104_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_104_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_105_ggml_set_rows_case { + %token_count = check.literal value(3) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<3x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<3xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<3x1024xf32>, tensor<3xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_105_ggml_set_rows_case> @qwen38_27b_udq4kxl_draft_tg_c1_105_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_106_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(3) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<3x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<3x256xf16> + %output = check.generate.fill value(1.0) : tensor<3x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<3x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<3x256xf16>, tensor<3x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_106_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_106_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_107_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<73728xi8> + %quantized_values = check.generate.fill value(0) : tensor<9216xi8> + %scales = check.generate.fill value(0) : tensor<2304xi8> + %sums = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<73728xi8>, tensor<9216xi8>, tensor<2304xi8>, tensor<2304xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_107_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_107_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_108_ggml_binary_f32_case { + %element_count = check.literal value(15360) : index + %lhs = check.generate.fill value(2.0) : tensor<15360xf32> + %rhs = check.generate.fill value(3.0) : tensor<15360xf32> + %output = check.generate.fill value(0.0) : tensor<15360xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<15360xf32>, tensor<15360xf32>, tensor<15360xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_108_ggml_binary_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_108_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_109_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(3) : index + %input = check.generate.fill value(2.0) : tensor<3x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<3x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<3x5120xf32>, tensor<5120xf32>, tensor<3x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_109_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_109_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_110_ggml_get_rows_f32_case { + %token_count = check.literal value(3) : index + %row_count = check.literal value(3) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<3xi32> + %weight = check.generate.fill value(0.0) : tensor<3x5120xf32> + %output = check.generate.fill value(1.0) : tensor<3x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<3xi32>, tensor<3x5120xf32>, tensor<3x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_110_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_110_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_111_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case { + %token_count = check.literal value(3) : index + %input = check.generate.fill value(0) : tensor<61440xi8> + %weight = check.generate.fill value(0) : tensor<1042944000xi8> + %output = check.generate.fill value(0) : tensor<2979840xi8> + kernel.launch @ggml_mul_mat_q6_k_packed_token1_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<61440xi8>, tensor<1042944000xi8>, tensor<2979840xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_111_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_111_ggml_mul_mat_q6_k_packed_token1_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_112_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<61440xi8> + %rhs = check.generate.fill value(0) : tensor<61440xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<61440xi8>, tensor<61440xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_112_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_112_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_113_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(3) : index + %input = check.generate.fill value(0) : tensor<122880xi8> + %weight = check.generate.fill value(0) : tensor<56115200xi8> + %output = check.generate.fill value(0) : tensor<61440xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<122880xi8>, tensor<56115200xi8>, tensor<61440xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_113_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_113_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_114_ggml_get_rows_f32_case { + %token_count = check.literal value(16) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<64xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<64xi8>, tensor<715161600xi8>, tensor<327680xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_114_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_114_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_115_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(16) : index + %input = check.generate.fill value(0) : tensor<327680xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<327680xi8>, tensor<20480xi8>, tensor<327680xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_115_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_115_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_116_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<327680xi8> + %output = check.generate.fill value(0) : tensor<655360xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + %partial = check.generate.fill value(0) : tensor<655360xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<327680xi8>, tensor<655360xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>, tensor<655360xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_116_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_116_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_117_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<168960xi8> + %second_weight = check.generate.fill value(0) : tensor<168960xi8> + %first_output = check.generate.fill value(0) : tensor<3072xi8> + %second_output = check.generate.fill value(0) : tensor<3072xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<168960xi8>, tensor<168960xi8>, tensor<3072xi8>, tensor<3072xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_117_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_117_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_118_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<327680xi8> + %output = check.generate.fill value(0) : tensor<393216xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<327680xi8>, tensor<393216xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_118_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_118_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_119_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<655360xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<655360xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<655360xi8>, tensor<163840xi8>, tensor<655360xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_119_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_119_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_120_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<3072xi8> + %beta_raw = check.generate.fill value(0) : tensor<3072xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<3072xi8> + %beta_dst = check.generate.fill value(0) : tensor<3072xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<3072xi8>, tensor<3072xi8>, tensor<192xi8>, tensor<192xi8>, tensor<3072xi8>, tensor<3072xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_120_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_120_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_121_llm_gated_delta_net_f32_wmma_head128_case { + %q = check.generate.fill value(0) : tensor<417792xi8> + %k = check.generate.fill value(0) : tensor<417792xi8> + %v = check.generate.fill value(0) : tensor<434176xi8> + %g = check.generate.fill value(0) : tensor<2112xi8> + %beta = check.generate.fill value(0) : tensor<2112xi8> + %state_in = check.generate.fill value(0) : tensor<3145728xi8> + %dst = check.generate.fill value(0) : tensor<3416064xi8> + %rms_scales = check.generate.fill value(0) : tensor<2112xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128[](%q, %k, %v, %g, %beta, %state_in, %dst, %rms_scales) : [](tensor<417792xi8>, tensor<417792xi8>, tensor<434176xi8>, tensor<2112xi8>, tensor<2112xi8>, tensor<3145728xi8>, tensor<3416064xi8>, tensor<2112xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_121_llm_gated_delta_net_f32_wmma_head128_case> @qwen38_27b_udq4kxl_draft_tg_c1_121_llm_gated_delta_net_f32_wmma_head128 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_122_ggml_copy_f32_case { + %element_count = check.literal value(786432) : index + %source = check.generate.fill value(0) : tensor<3145728xi8> + %output = check.generate.fill value(0) : tensor<3145728xi8> + kernel.launch @ggml_copy_f32[%element_count](%element_count, %source, %output) : [index](index, tensor<3145728xi8>, tensor<3145728xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_122_ggml_copy_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_122_ggml_copy_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_123_llm_gated_delta_net_f32_wmma_head128_inplace_case { + %q = check.generate.fill value(0) : tensor<8192xi8> + %k = check.generate.fill value(0) : tensor<8192xi8> + %v = check.generate.fill value(0) : tensor<24576xi8> + %g = check.generate.fill value(0) : tensor<192xi8> + %beta = check.generate.fill value(0) : tensor<192xi8> + %state_inout = check.generate.fill value(0) : tensor<3145728xi8> + %dst = check.generate.fill value(0) : tensor<24576xi8> + %rms_scales = check.generate.fill value(0) : tensor<192xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_inplace[](%q, %k, %v, %g, %beta, %state_inout, %dst, %rms_scales) : [](tensor<8192xi8>, tensor<8192xi8>, tensor<24576xi8>, tensor<192xi8>, tensor<192xi8>, tensor<3145728xi8>, tensor<24576xi8>, tensor<192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_123_llm_gated_delta_net_f32_wmma_head128_inplace_case> @qwen38_27b_udq4kxl_draft_tg_c1_123_llm_gated_delta_net_f32_wmma_head128_inplace + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_124_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(768) : index + %input = check.generate.fill value(0) : tensor<393216xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<393216xi8> + %output = check.generate.fill value(0) : tensor<393216xi8> + %quantized_values = check.generate.fill value(0) : tensor<49152xi8> + %scales = check.generate.fill value(0) : tensor<12288xi8> + %sums = check.generate.fill value(0) : tensor<12288xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<393216xi8>, tensor<512xi8>, tensor<393216xi8>, tensor<393216xi8>, tensor<49152xi8>, tensor<12288xi8>, tensor<12288xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_124_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_124_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_125_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<393216xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + %quantized_values = check.generate.fill value(0) : tensor<49152xi8> + %scales = check.generate.fill value(0) : tensor<12288xi8> + %sums = check.generate.fill value(0) : tensor<12288xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<393216xi8>, tensor<327680xi8>, tensor<49152xi8>, tensor<12288xi8>, tensor<12288xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_125_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_125_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_126_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(16) : index + %lhs = check.generate.fill value(0) : tensor<327680xi8> + %rhs = check.generate.fill value(0) : tensor<327680xi8> + %residual_output = check.generate.fill value(0) : tensor<327680xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<327680xi8>, tensor<327680xi8>, tensor<327680xi8>, tensor<20480xi8>, tensor<327680xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_126_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_126_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_127_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<1114112xi8> + %second_output = check.generate.fill value(0) : tensor<1114112xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<1114112xi8>, tensor<1114112xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_127_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_127_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_128_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<1114112xi8> + %rhs = check.generate.fill value(0) : tensor<1114112xi8> + %output = check.generate.fill value(0) : tensor<1114112xi8> + %quantized_values = check.generate.fill value(0) : tensor<139264xi8> + %scales = check.generate.fill value(0) : tensor<34816xi8> + %sums = check.generate.fill value(0) : tensor<34816xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<1114112xi8>, tensor<1114112xi8>, tensor<1114112xi8>, tensor<139264xi8>, tensor<34816xi8>, tensor<34816xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_128_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_128_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_129_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<1114112xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + %quantized_values = check.generate.fill value(0) : tensor<139264xi8> + %scales = check.generate.fill value(0) : tensor<34816xi8> + %sums = check.generate.fill value(0) : tensor<34816xi8> + %partial = check.generate.fill value(0) : tensor<327680xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<1114112xi8>, tensor<327680xi8>, tensor<139264xi8>, tensor<34816xi8>, tensor<34816xi8>, tensor<327680xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_129_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_129_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_130_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<327680xi8> + %output = check.generate.fill value(0) : tensor<786432xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<327680xi8>, tensor<786432xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_130_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_130_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_131_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(16) : index + %input = check.generate.fill value(0) : tensor<327680xi8> + %weight = check.generate.fill value(0) : tensor<5611520xi8> + %output = check.generate.fill value(0) : tensor<65536xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<327680xi8>, tensor<5611520xi8>, tensor<65536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_131_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_131_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_132_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<327680xi8> + %output = check.generate.fill value(0) : tensor<65536xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<327680xi8>, tensor<65536xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_132_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_132_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_133_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<785408xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<256xi8> + %output = check.generate.fill value(0) : tensor<393216xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<785408xi8>, tensor<1024xi8>, tensor<256xi8>, tensor<393216xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_133_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_133_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_134_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<65536xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<256xi8> + %output = check.generate.fill value(0) : tensor<65536xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<65536xi8>, tensor<1024xi8>, tensor<256xi8>, tensor<65536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_134_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_134_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_135_ggml_set_rows_case { + %token_count = check.literal value(16) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<16x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<16xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<16x1024xf32>, tensor<16xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_135_ggml_set_rows_case> @qwen38_27b_udq4kxl_draft_tg_c1_135_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_136_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(16) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<16x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<16x256xf16> + %output = check.generate.fill value(1.0) : tensor<16x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<16x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<16x256xf16>, tensor<16x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_136_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_136_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_137_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<393216xi8> + %quantized_values = check.generate.fill value(0) : tensor<49152xi8> + %scales = check.generate.fill value(0) : tensor<12288xi8> + %sums = check.generate.fill value(0) : tensor<12288xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<393216xi8>, tensor<49152xi8>, tensor<12288xi8>, tensor<12288xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_137_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_137_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_138_ggml_binary_f32_case { + %element_count = check.literal value(81920) : index + %lhs = check.generate.fill value(2.0) : tensor<81920xf32> + %rhs = check.generate.fill value(3.0) : tensor<81920xf32> + %output = check.generate.fill value(0.0) : tensor<81920xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<81920xf32>, tensor<81920xf32>, tensor<81920xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_138_ggml_binary_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_138_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_139_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(16) : index + %input = check.generate.fill value(2.0) : tensor<16x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<16x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<16x5120xf32>, tensor<5120xf32>, tensor<16x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_139_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_139_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_140_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<327680xi8> + %rhs = check.generate.fill value(0) : tensor<327680xi8> + %output = check.generate.fill value(0) : tensor<655360xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<327680xi8>, tensor<327680xi8>, tensor<655360xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_140_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_140_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_141_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(16) : index + %input = check.generate.fill value(0) : tensor<655360xi8> + %weight = check.generate.fill value(0) : tensor<56115200xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<655360xi8>, tensor<56115200xi8>, tensor<327680xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_141_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_141_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_142_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(4) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<4x5120xf32> + %output = check.generate.fill value(1.0) : tensor<1x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<4x5120xf32>, tensor<1x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_142_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_142_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_143_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<81920xi8> + %rhs = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<163840xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<81920xi8>, tensor<81920xi8>, tensor<163840xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_143_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_143_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_144_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(0) : tensor<163840xi8> + %weight = check.generate.fill value(0) : tensor<56115200xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<163840xi8>, tensor<56115200xi8>, tensor<81920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_144_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_144_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_145_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<4xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<4xi8>, tensor<715161600xi8>, tensor<20480xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_145_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_145_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_146_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(2.0) : tensor<1x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<1x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<1x5120xf32>, tensor<5120xf32>, tensor<1x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_146_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_146_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_147_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<20480xi8> + %rhs = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<20480xi8>, tensor<20480xi8>, tensor<40960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_147_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_147_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_148_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(10240) : index + %output_size = check.literal value(5120) : index + %input = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<43008000xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<40960xi8>, tensor<43008000xi8>, tensor<20480xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_148_ggml_mul_mat_f32_f32_decode_wave64_case> @qwen38_27b_udq4kxl_draft_tg_c1_148_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_149_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<20480xi8>, tensor<20480xi8>, tensor<20480xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_149_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_149_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_150_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<49152xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<20480xi8>, tensor<49152xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_150_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_150_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_151_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(5120) : index + %output_size = check.literal value(1024) : index + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<4300800xi8> + %output = check.generate.fill value(0) : tensor<4096xi8> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<20480xi8>, tensor<4300800xi8>, tensor<4096xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_151_ggml_mul_mat_f32_f32_decode_wave64_case> @qwen38_27b_udq4kxl_draft_tg_c1_151_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_152_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<4096xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<20480xi8>, tensor<4096xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_152_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_152_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_153_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<48128xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<16xi8> + %output = check.generate.fill value(0) : tensor<24576xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<48128xi8>, tensor<1024xi8>, tensor<16xi8>, tensor<24576xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_153_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_153_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_154_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<4096xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<16xi8> + %output = check.generate.fill value(0) : tensor<4096xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<4096xi8>, tensor<1024xi8>, tensor<16xi8>, tensor<4096xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_154_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_154_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_155_ggml_set_rows_case { + %token_count = check.literal value(1) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<1x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<1xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<1x1024xf32>, tensor<1xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_155_ggml_set_rows_case> @qwen38_27b_udq4kxl_draft_tg_c1_155_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_156_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(1) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<1x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<1x256xf16> + %output = check.generate.fill value(1.0) : tensor<1x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<1x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<1x256xf16>, tensor<1x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_156_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_156_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_157_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<24576xi8> + %quantized_values = check.generate.fill value(0) : tensor<3072xi8> + %scales = check.generate.fill value(0) : tensor<768xi8> + %sums = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<24576xi8>, tensor<3072xi8>, tensor<768xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_157_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_157_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_158_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<24576xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<3072xi8> + %scales = check.generate.fill value(0) : tensor<768xi8> + %sums = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<24576xi8>, tensor<20480xi8>, tensor<3072xi8>, tensor<768xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_158_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_158_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_159_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(1) : index + %lhs = check.generate.fill value(0) : tensor<20480xi8> + %rhs = check.generate.fill value(0) : tensor<20480xi8> + %residual_output = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<20480xi8>, tensor<20480xi8>, tensor<20480xi8>, tensor<20480xi8>, tensor<20480xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_159_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_159_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_160_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<69632xi8> + %second_output = check.generate.fill value(0) : tensor<69632xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<69632xi8>, tensor<69632xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_160_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_160_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_161_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<69632xi8> + %rhs = check.generate.fill value(0) : tensor<69632xi8> + %output = check.generate.fill value(0) : tensor<69632xi8> + %quantized_values = check.generate.fill value(0) : tensor<8704xi8> + %scales = check.generate.fill value(0) : tensor<2176xi8> + %sums = check.generate.fill value(0) : tensor<2176xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<69632xi8>, tensor<69632xi8>, tensor<69632xi8>, tensor<8704xi8>, tensor<2176xi8>, tensor<2176xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_161_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_161_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_162_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<69632xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<8704xi8> + %scales = check.generate.fill value(0) : tensor<2176xi8> + %sums = check.generate.fill value(0) : tensor<2176xi8> + %partial = check.generate.fill value(0) : tensor<20480xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<69632xi8>, tensor<20480xi8>, tensor<8704xi8>, tensor<2176xi8>, tensor<2176xi8>, tensor<20480xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_162_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_162_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_163_ggml_binary_f32_case { + %element_count = check.literal value(5120) : index + %lhs = check.generate.fill value(2.0) : tensor<5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<5120xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<5120xf32>, tensor<5120xf32>, tensor<5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_163_ggml_binary_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_163_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_164_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(1) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<1x5120xf32> + %output = check.generate.fill value(1.0) : tensor<1x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<1x5120xf32>, tensor<1x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_164_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_164_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_165_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<20480xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_165_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_165_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_166_ggml_mul_mat_q6_k_symmetric_i2_scan_token1_case { + %weight = check.generate.fill value(0) : tensor<1380659200xi8> + %output = check.generate.fill value(0) : tensor<993280xi8> + %qact = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %selected_mask = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_mul_mat_q6_k_symmetric_i2_scan_token1[](%weight, %output, %qact, %scales, %selected_mask) : [](tensor<1380659200xi8>, tensor<993280xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_166_ggml_mul_mat_q6_k_symmetric_i2_scan_token1_case> @qwen38_27b_udq4kxl_draft_tg_c1_166_ggml_mul_mat_q6_k_symmetric_i2_scan_token1 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_167_ggml_top_k8_f32_partitions_register_case { + %element_count = check.literal value(248320) : index + %values = check.generate.fill value(0) : tensor<993280xi8> + %partial_values = check.generate.fill value(0) : tensor<4096xi8> + %partial_ids = check.generate.fill value(0) : tensor<4096xi8> + kernel.launch @ggml_top_k8_f32_partitions_register[%element_count](%element_count, %values, %partial_values, %partial_ids) : [index](index, tensor<993280xi8>, tensor<4096xi8>, tensor<4096xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_167_ggml_top_k8_f32_partitions_register_case> @qwen38_27b_udq4kxl_draft_tg_c1_167_ggml_top_k8_f32_partitions_register + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_168_ggml_top_k128_f32_reduce_gather_register_case { + %element_count = check.literal value(248320) : index + %partial_values = check.generate.fill value(0) : tensor<4096xi8> + %partial_ids = check.generate.fill value(0) : tensor<4096xi8> + %candidate_output = check.generate.fill value(0) : tensor<512xi8> + %value_output = check.generate.fill value(0) : tensor<512xi8> + kernel.launch @ggml_top_k128_f32_reduce_gather_register[%element_count](%element_count, %partial_values, %partial_ids, %candidate_output, %value_output) : [index](index, tensor<4096xi8>, tensor<4096xi8>, tensor<512xi8>, tensor<512xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_168_ggml_top_k128_f32_reduce_gather_register_case> @qwen38_27b_udq4kxl_draft_tg_c1_168_ggml_top_k128_f32_reduce_gather_register + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_169_ggml_fill_negative_f32_case { + %element_count = check.literal value(248320) : index + %output = check.generate.fill value(0) : tensor<993280xi8> + kernel.launch @ggml_fill_negative_f32[%element_count](%element_count, %output) : [index](index, tensor<993280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_169_ggml_fill_negative_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_169_ggml_fill_negative_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_170_ggml_mul_mat_q6_k_packed_selected_refine_token1_case { + %token_count = check.literal value(1) : index + %candidate_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<1380659200xi8> + %candidates = check.generate.fill value(0) : tensor<256xi8> + %exact_output = check.generate.fill value(0) : tensor<256xi8> + %output = check.generate.fill value(0) : tensor<993280xi8> + kernel.launch @ggml_mul_mat_q6_k_packed_selected_refine_token1[%token_count, %candidate_count](%token_count, %candidate_count, %input, %weight, %candidates, %exact_output, %output) : [index, index](index, index, tensor<20480xi8>, tensor<1380659200xi8>, tensor<256xi8>, tensor<256xi8>, tensor<993280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_170_ggml_mul_mat_q6_k_packed_selected_refine_token1_case> @qwen38_27b_udq4kxl_draft_tg_c1_170_ggml_mul_mat_q6_k_packed_selected_refine_token1 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_171_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + %partial = check.generate.fill value(0) : tensor<40960xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<20480xi8>, tensor<40960xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>, tensor<40960xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_171_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_171_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_172_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(5120) : index + %output_size = check.literal value(48) : index + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<138240xi8> + %output = check.generate.fill value(0) : tensor<192xi8> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<20480xi8>, tensor<138240xi8>, tensor<192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_172_ggml_mul_mat_f32_f32_decode_wave64_case> @qwen38_27b_udq4kxl_draft_tg_c1_172_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_173_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<24576xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<20480xi8>, tensor<24576xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_173_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_173_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_174_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<40960xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<40960xi8>, tensor<163840xi8>, tensor<40960xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_174_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_174_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_175_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<192xi8> + %beta_raw = check.generate.fill value(0) : tensor<192xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<192xi8> + %beta_dst = check.generate.fill value(0) : tensor<192xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<192xi8>, tensor<192xi8>, tensor<192xi8>, tensor<192xi8>, tensor<192xi8>, tensor<192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_175_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_175_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_176_llm_gated_delta_net_f32_wmma_head128_inplace_case { + %q = check.generate.fill value(0) : tensor<8192xi8> + %k = check.generate.fill value(0) : tensor<8192xi8> + %v = check.generate.fill value(0) : tensor<24576xi8> + %g = check.generate.fill value(0) : tensor<192xi8> + %beta = check.generate.fill value(0) : tensor<192xi8> + %state_inout = check.generate.fill value(0) : tensor<3145728xi8> + %dst = check.generate.fill value(0) : tensor<24576xi8> + %rms_scales = check.generate.fill value(0) : tensor<192xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_inplace[](%q, %k, %v, %g, %beta, %state_inout, %dst, %rms_scales) : [](tensor<8192xi8>, tensor<8192xi8>, tensor<24576xi8>, tensor<192xi8>, tensor<192xi8>, tensor<3145728xi8>, tensor<24576xi8>, tensor<192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_176_llm_gated_delta_net_f32_wmma_head128_inplace_case> @qwen38_27b_udq4kxl_draft_tg_c1_176_llm_gated_delta_net_f32_wmma_head128_inplace + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_177_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(48) : index + %input = check.generate.fill value(0) : tensor<24576xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<24576xi8> + %output = check.generate.fill value(0) : tensor<24576xi8> + %quantized_values = check.generate.fill value(0) : tensor<3072xi8> + %scales = check.generate.fill value(0) : tensor<768xi8> + %sums = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<24576xi8>, tensor<512xi8>, tensor<24576xi8>, tensor<24576xi8>, tensor<3072xi8>, tensor<768xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_177_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_draft_tg_c1_177_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_178_ggml_get_rows_f32_case { + %token_count = check.literal value(2) : index + %row_count = check.literal value(2) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<2xi32> + %weight = check.generate.fill value(0.0) : tensor<2x5120xf32> + %output = check.generate.fill value(1.0) : tensor<2x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<2xi32>, tensor<2x5120xf32>, tensor<2x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_178_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_draft_tg_c1_178_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_draft_tg_c1_179_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<1042944000xi8> + %output = check.generate.fill value(0) : tensor<1986560xi8> + kernel.launch @ggml_mul_mat_q6_k_packed_token1_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<40960xi8>, tensor<1042944000xi8>, tensor<1986560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_draft_tg_c1_179_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case> @qwen38_27b_udq4kxl_draft_tg_c1_179_ggml_mul_mat_q6_k_packed_token1_f16_wmma diff --git a/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_mtp4_companion_all.tg_c5.json b/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_mtp4_companion_all.tg_c5.json new file mode 100644 index 000000000000..558d85759a79 --- /dev/null +++ b/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_mtp4_companion_all.tg_c5.json @@ -0,0 +1,5593 @@ +{ + "command_count": 8673, + "dispatch_count": 181, + "dispatches": [ + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_000_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "2", + "ggml.get_rows_f32.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 4, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 248320, + "token_count": 2 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_001_ggml_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 4, + "integer_parameters": { + "token_count": 2 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "10240", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_003_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "48", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_004_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_005_ggml_scale_bias_f32", + "compile_parameters": { + "ggml.scale.bias": "0", + "ggml.scale.scale": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": { + "element_count": 30720 + }, + "kernel": "loom_libs:ggml_scale_bias_f32", + "library_sources": [], + "primary_sources": [ + "ops/scale_bias_f32.loom" + ], + "sources": [ + "ops/scale_bias_f32.loom" + ], + "symbol": "ggml_scale_bias_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_006_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "30720", + "ggml.get_rows_f32.token_capacity": "1", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 384, + "integer_parameters": { + "hidden_size": 30720, + "row_count": 20, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_007_llm_ssm_conv_dconv4_silu_rollback_f32", + "compile_parameters": { + "llm.ssm_conv.rollback.cache_count": "5", + "llm.ssm_conv.rollback.d_inner": "10240", + "llm.ssm_conv.rollback.n_s": "1", + "llm.ssm_conv.rollback.n_t": "2", + "llm.ssm_conv.rollback.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:llm_ssm_conv_dconv4_silu_rollback_f32", + "library_sources": [], + "primary_sources": [ + "ops/ssm_conv_f32.loom" + ], + "sources": [ + "ops/ssm_conv_f32.loom" + ], + "symbol": "llm_ssm_conv_dconv4_silu_rollback_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_008_ggml_scale_bias_f32", + "compile_parameters": { + "ggml.scale.bias": "0", + "ggml.scale.scale": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": { + "element_count": 786432 + }, + "kernel": "loom_libs:ggml_scale_bias_f32", + "library_sources": [], + "primary_sources": [ + "ops/scale_bias_f32.loom" + ], + "sources": [ + "ops/scale_bias_f32.loom" + ], + "symbol": "ggml_scale_bias_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_009_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "786432", + "ggml.get_rows_f32.token_capacity": "1", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 384, + "integer_parameters": { + "hidden_size": 786432, + "row_count": 20, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_010_llm_gated_delta_net_projection_epilogue_f32", + "compile_parameters": { + "llm.gated_delta_net.epilogue_element_count": "96", + "llm.gated_delta_net.epilogue_head_count": "48", + "llm.gated_delta_net.epilogue_workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_projection_epilogue_f32", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_projection_epilogue_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_011_llm_gated_delta_net_f32_wmma_head128_snapshot", + "compile_parameters": { + "llm.gated_delta_net.head_count": "48", + "llm.gated_delta_net.head_width": "128", + "llm.gated_delta_net.l2_epsilon": "9.99999997e-07", + "llm.gated_delta_net.qk_stride1": "128", + "llm.gated_delta_net.qk_stride2": "10240", + "llm.gated_delta_net.qk_stride3": "20480", + "llm.gated_delta_net.query_head_count": "16", + "llm.gated_delta_net.query_sequence_ratio": "1", + "llm.gated_delta_net.scalar_stride1": "1", + "llm.gated_delta_net.scalar_stride2": "48", + "llm.gated_delta_net.scalar_stride3": "96", + "llm.gated_delta_net.sequence_count": "1", + "llm.gated_delta_net.snapshot_stride": "3145728", + "llm.gated_delta_net.token_count": "2", + "llm.gated_delta_net.value_stride1": "128", + "llm.gated_delta_net.value_stride2": "10240", + "llm.gated_delta_net.value_stride3": "20480", + "llm.gated_delta_net.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_f32_wmma_head128_snapshot", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_f32_wmma_head128_snapshot", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_012_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.hidden_size": "128", + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": { + "token_count": 96 + }, + "kernel": "loom_libs:ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_013_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 130, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_014_ggml_add_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.add_rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.add_rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 256, + "integer_parameters": { + "token_count": 2 + }, + "kernel": "loom_libs:ggml_add_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_add_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_015_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 130, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_016_ggml_binary_swiglu_symmetric_i4_k32", + "compile_parameters": { + "ggml.binary_swiglu_symmetric_i4.input_size": "17408", + "ggml.binary_swiglu_symmetric_i4.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 130, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_binary_swiglu_symmetric_i4_k32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_swiglu_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_017_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 130, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_018_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "12288", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_019_ggml_mul_mat_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat.input_size": "5120", + "ggml.mul_mat.output_accumulation": "0", + "ggml.mul_mat.output_size": "1024", + "ggml.mul_mat.output_unary_op": "23", + "ggml.mul_mat.weight_format": "6", + "ggml.workload.token_capacity": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": { + "token_count": 2 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_020_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "1024", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_021_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "512", + "ggml.rmsnorm_mul_rope.input_stride2": "12288", + "ggml.rmsnorm_mul_rope.input_stride3": "24576", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "24", + "ggml.rmsnorm_mul_rope.ne2": "2", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "6144", + "ggml.rmsnorm_mul_rope.output_stride3": "12288", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_022_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "256", + "ggml.rmsnorm_mul_rope.input_stride2": "1024", + "ggml.rmsnorm_mul_rope.input_stride3": "2048", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "4", + "ggml.rmsnorm_mul_rope.ne2": "2", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "1024", + "ggml.rmsnorm_mul_rope.output_stride3": "2048", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_023_ggml_set_rows", + "compile_parameters": { + "ggml.set_rows.hidden_capacity": "1024", + "ggml.set_rows.input_format": "32", + "ggml.set_rows.output_format": "16", + "ggml.set_rows.token_capacity": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 68, + "integer_parameters": { + "cache_row_count": 262144, + "hidden_size": 1024, + "token_count": 2 + }, + "kernel": "loom_libs:ggml_set_rows", + "library_sources": [], + "primary_sources": [ + "ops/set_rows.loom" + ], + "sources": [ + "ops/set_rows.loom" + ], + "symbol": "ggml_set_rows", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_024_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.apply_gate": "1", + "ggml.flash_attention.gate_stride_head": "512", + "ggml.flash_attention.gate_stride_token": "12288", + "ggml.flash_attention.head_size": "256", + "ggml.flash_attention.key_value_head_count": "4", + "ggml.flash_attention.query_head_count": "24" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": { + "key_value_token_count": 256, + "query_token_count": 2 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_025_ggml_quantize_f32_symmetric_i4_k32", + "compile_parameters": { + "ggml.quantize_symmetric_i4_k32.input_size": "6144", + "ggml.quantize_symmetric_i4_k32.token_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_quantize_f32_symmetric_i4_k32", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_quantize_f32_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_026_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 4, + "integer_parameters": { + "element_count": 10240 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_027_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "5120", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 8, + "integer_parameters": { + "token_count": 2 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_028_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "1", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 2, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_029_ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "248320" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 4, + "integer_parameters": { + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_030_ggml_concat_dim0_f32", + "compile_parameters": { + "ggml.concat_dim0_f32.lhs_width": "5120", + "ggml.concat_dim0_f32.rhs_width": "5120", + "ggml.concat_dim0_f32.row_count": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_concat_dim0_f32", + "library_sources": [], + "primary_sources": [ + "ops/copy_f32.loom" + ], + "sources": [ + "ops/copy_f32.loom" + ], + "symbol": "ggml_concat_dim0_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_031_ggml_mul_mat_f32_f32_wmma", + "compile_parameters": { + "ggml.mul_mat.input_size": "10240", + "ggml.mul_mat.output_accumulation": "0", + "ggml.mul_mat.output_size": "5120", + "ggml.mul_mat.output_unary_op": "23", + "ggml.mul_mat.weight_format": "6", + "ggml.workload.token_capacity": "2" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "token_count": 2 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_wmma", + "library_sources": [ + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_wmma.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_wmma.loom", + "motifs/mul_mat_f32_f32_wmma_core.loom", + "motifs/dequant.loom", + "motifs/unary_f32_apply.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_032_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "4", + "ggml.get_rows_f32.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 3, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 248320, + "token_count": 4 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_033_ggml_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 3, + "integer_parameters": { + "token_count": 4 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_034_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "10240", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_035_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "48", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_036_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_037_llm_ssm_conv_dconv4_silu_rollback_f32", + "compile_parameters": { + "llm.ssm_conv.rollback.cache_count": "5", + "llm.ssm_conv.rollback.d_inner": "10240", + "llm.ssm_conv.rollback.n_s": "1", + "llm.ssm_conv.rollback.n_t": "4", + "llm.ssm_conv.rollback.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:llm_ssm_conv_dconv4_silu_rollback_f32", + "library_sources": [], + "primary_sources": [ + "ops/ssm_conv_f32.loom" + ], + "sources": [ + "ops/ssm_conv_f32.loom" + ], + "symbol": "llm_ssm_conv_dconv4_silu_rollback_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_038_llm_gated_delta_net_projection_epilogue_f32", + "compile_parameters": { + "llm.gated_delta_net.epilogue_element_count": "192", + "llm.gated_delta_net.epilogue_head_count": "48", + "llm.gated_delta_net.epilogue_workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_projection_epilogue_f32", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_projection_epilogue_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_039_llm_gated_delta_net_f32_wmma_head128_snapshot", + "compile_parameters": { + "llm.gated_delta_net.head_count": "48", + "llm.gated_delta_net.head_width": "128", + "llm.gated_delta_net.l2_epsilon": "9.99999997e-07", + "llm.gated_delta_net.qk_stride1": "128", + "llm.gated_delta_net.qk_stride2": "10240", + "llm.gated_delta_net.qk_stride3": "40960", + "llm.gated_delta_net.query_head_count": "16", + "llm.gated_delta_net.query_sequence_ratio": "1", + "llm.gated_delta_net.scalar_stride1": "1", + "llm.gated_delta_net.scalar_stride2": "48", + "llm.gated_delta_net.scalar_stride3": "192", + "llm.gated_delta_net.sequence_count": "1", + "llm.gated_delta_net.snapshot_stride": "3145728", + "llm.gated_delta_net.token_count": "4", + "llm.gated_delta_net.value_stride1": "128", + "llm.gated_delta_net.value_stride2": "10240", + "llm.gated_delta_net.value_stride3": "40960", + "llm.gated_delta_net.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_f32_wmma_head128_snapshot", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_f32_wmma_head128_snapshot", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_040_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.hidden_size": "128", + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": { + "token_count": 192 + }, + "kernel": "loom_libs:ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_041_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 129, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_042_ggml_add_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.add_rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.add_rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 255, + "integer_parameters": { + "token_count": 4 + }, + "kernel": "loom_libs:ggml_add_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_add_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_043_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 129, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_044_ggml_binary_swiglu_symmetric_i4_k32", + "compile_parameters": { + "ggml.binary_swiglu_symmetric_i4.input_size": "17408", + "ggml.binary_swiglu_symmetric_i4.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 129, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_binary_swiglu_symmetric_i4_k32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_swiglu_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_045_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 129, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_046_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "12288", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 33, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_047_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "1024" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 33, + "integer_parameters": { + "token_count": 4 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_048_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "1024", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 33, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_049_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "512", + "ggml.rmsnorm_mul_rope.input_stride2": "12288", + "ggml.rmsnorm_mul_rope.input_stride3": "49152", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "24", + "ggml.rmsnorm_mul_rope.ne2": "4", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "6144", + "ggml.rmsnorm_mul_rope.output_stride3": "24576", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 33, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_050_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "256", + "ggml.rmsnorm_mul_rope.input_stride2": "1024", + "ggml.rmsnorm_mul_rope.input_stride3": "4096", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "4", + "ggml.rmsnorm_mul_rope.ne2": "4", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "1024", + "ggml.rmsnorm_mul_rope.output_stride3": "4096", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 33, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_051_ggml_set_rows", + "compile_parameters": { + "ggml.set_rows.hidden_capacity": "1024", + "ggml.set_rows.input_format": "32", + "ggml.set_rows.output_format": "16", + "ggml.set_rows.token_capacity": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 66, + "integer_parameters": { + "cache_row_count": 262144, + "hidden_size": 1024, + "token_count": 4 + }, + "kernel": "loom_libs:ggml_set_rows", + "library_sources": [], + "primary_sources": [ + "ops/set_rows.loom" + ], + "sources": [ + "ops/set_rows.loom" + ], + "symbol": "ggml_set_rows", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_052_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.apply_gate": "1", + "ggml.flash_attention.gate_stride_head": "512", + "ggml.flash_attention.gate_stride_token": "12288", + "ggml.flash_attention.head_size": "256", + "ggml.flash_attention.key_value_head_count": "4", + "ggml.flash_attention.query_head_count": "24" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 33, + "integer_parameters": { + "key_value_token_count": 256, + "query_token_count": 4 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_053_ggml_quantize_f32_symmetric_i4_k32", + "compile_parameters": { + "ggml.quantize_symmetric_i4_k32.input_size": "6144", + "ggml.quantize_symmetric_i4_k32.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 33, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_quantize_f32_symmetric_i4_k32", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_quantize_f32_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_054_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 3, + "integer_parameters": { + "element_count": 20480 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_055_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "5120", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 5, + "integer_parameters": { + "token_count": 4 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_056_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "4", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 4, + "token_count": 4 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_057_ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "248320" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 4 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_058_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "5", + "ggml.get_rows_f32.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 248320, + "token_count": 5 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_059_ggml_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "token_count": 5 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_060_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "10240", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_061_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "48", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_062_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_063_llm_ssm_conv_dconv4_silu_rollback_f32", + "compile_parameters": { + "llm.ssm_conv.rollback.cache_count": "5", + "llm.ssm_conv.rollback.d_inner": "10240", + "llm.ssm_conv.rollback.n_s": "1", + "llm.ssm_conv.rollback.n_t": "5", + "llm.ssm_conv.rollback.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_ssm_conv_dconv4_silu_rollback_f32", + "library_sources": [], + "primary_sources": [ + "ops/ssm_conv_f32.loom" + ], + "sources": [ + "ops/ssm_conv_f32.loom" + ], + "symbol": "llm_ssm_conv_dconv4_silu_rollback_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_064_llm_gated_delta_net_projection_epilogue_f32", + "compile_parameters": { + "llm.gated_delta_net.epilogue_element_count": "240", + "llm.gated_delta_net.epilogue_head_count": "48", + "llm.gated_delta_net.epilogue_workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_projection_epilogue_f32", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_projection_epilogue_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_065_llm_gated_delta_net_f32_wmma_head128_snapshot", + "compile_parameters": { + "llm.gated_delta_net.head_count": "48", + "llm.gated_delta_net.head_width": "128", + "llm.gated_delta_net.l2_epsilon": "9.99999997e-07", + "llm.gated_delta_net.qk_stride1": "128", + "llm.gated_delta_net.qk_stride2": "10240", + "llm.gated_delta_net.qk_stride3": "51200", + "llm.gated_delta_net.query_head_count": "16", + "llm.gated_delta_net.query_sequence_ratio": "1", + "llm.gated_delta_net.scalar_stride1": "1", + "llm.gated_delta_net.scalar_stride2": "48", + "llm.gated_delta_net.scalar_stride3": "240", + "llm.gated_delta_net.sequence_count": "1", + "llm.gated_delta_net.snapshot_stride": "3145728", + "llm.gated_delta_net.token_count": "5", + "llm.gated_delta_net.value_stride1": "128", + "llm.gated_delta_net.value_stride2": "10240", + "llm.gated_delta_net.value_stride3": "51200", + "llm.gated_delta_net.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_f32_wmma_head128_snapshot", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_f32_wmma_head128_snapshot", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_066_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.hidden_size": "128", + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "token_count": 240 + }, + "kernel": "loom_libs:ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_067_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_068_ggml_add_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.add_rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.add_rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 128, + "integer_parameters": { + "token_count": 5 + }, + "kernel": "loom_libs:ggml_add_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_add_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_069_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_070_ggml_binary_swiglu_symmetric_i4_k32", + "compile_parameters": { + "ggml.binary_swiglu_symmetric_i4.input_size": "17408", + "ggml.binary_swiglu_symmetric_i4.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_binary_swiglu_symmetric_i4_k32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_swiglu_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_071_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_072_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "12288", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_073_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "1024" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": { + "token_count": 5 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_074_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "1024", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_075_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "512", + "ggml.rmsnorm_mul_rope.input_stride2": "12288", + "ggml.rmsnorm_mul_rope.input_stride3": "61440", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "24", + "ggml.rmsnorm_mul_rope.ne2": "5", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "6144", + "ggml.rmsnorm_mul_rope.output_stride3": "30720", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_076_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "256", + "ggml.rmsnorm_mul_rope.input_stride2": "1024", + "ggml.rmsnorm_mul_rope.input_stride3": "5120", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "4", + "ggml.rmsnorm_mul_rope.ne2": "5", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "1024", + "ggml.rmsnorm_mul_rope.output_stride3": "5120", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_077_ggml_set_rows", + "compile_parameters": { + "ggml.set_rows.hidden_capacity": "1024", + "ggml.set_rows.input_format": "32", + "ggml.set_rows.output_format": "16", + "ggml.set_rows.token_capacity": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": { + "cache_row_count": 262144, + "hidden_size": 1024, + "token_count": 5 + }, + "kernel": "loom_libs:ggml_set_rows", + "library_sources": [], + "primary_sources": [ + "ops/set_rows.loom" + ], + "sources": [ + "ops/set_rows.loom" + ], + "symbol": "ggml_set_rows", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_078_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.apply_gate": "1", + "ggml.flash_attention.gate_stride_head": "512", + "ggml.flash_attention.gate_stride_token": "12288", + "ggml.flash_attention.head_size": "256", + "ggml.flash_attention.key_value_head_count": "4", + "ggml.flash_attention.query_head_count": "24" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": { + "key_value_token_count": 256, + "query_token_count": 5 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_079_ggml_quantize_f32_symmetric_i4_k32", + "compile_parameters": { + "ggml.quantize_symmetric_i4_k32.input_size": "6144", + "ggml.quantize_symmetric_i4_k32.token_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_quantize_f32_symmetric_i4_k32", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_quantize_f32_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_080_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "element_count": 25600 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_081_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "5120", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 4, + "integer_parameters": { + "token_count": 5 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_082_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "5", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 5, + "token_count": 5 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_083_ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "248320" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 5 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_084_ggml_concat_dim0_f32", + "compile_parameters": { + "ggml.concat_dim0_f32.lhs_width": "5120", + "ggml.concat_dim0_f32.rhs_width": "5120", + "ggml.concat_dim0_f32.row_count": "5" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_concat_dim0_f32", + "library_sources": [], + "primary_sources": [ + "ops/copy_f32.loom" + ], + "sources": [ + "ops/copy_f32.loom" + ], + "symbol": "ggml_concat_dim0_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_085_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "10240", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "5120" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 5 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_086_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "3", + "ggml.get_rows_f32.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 248320, + "token_count": 3 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_087_ggml_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "token_count": 3 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_088_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "10240", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_089_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "48", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_090_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_091_llm_ssm_conv_dconv4_silu_rollback_f32", + "compile_parameters": { + "llm.ssm_conv.rollback.cache_count": "5", + "llm.ssm_conv.rollback.d_inner": "10240", + "llm.ssm_conv.rollback.n_s": "1", + "llm.ssm_conv.rollback.n_t": "3", + "llm.ssm_conv.rollback.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_ssm_conv_dconv4_silu_rollback_f32", + "library_sources": [], + "primary_sources": [ + "ops/ssm_conv_f32.loom" + ], + "sources": [ + "ops/ssm_conv_f32.loom" + ], + "symbol": "llm_ssm_conv_dconv4_silu_rollback_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_092_llm_gated_delta_net_projection_epilogue_f32", + "compile_parameters": { + "llm.gated_delta_net.epilogue_element_count": "144", + "llm.gated_delta_net.epilogue_head_count": "48", + "llm.gated_delta_net.epilogue_workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_projection_epilogue_f32", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_projection_epilogue_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_093_llm_gated_delta_net_f32_wmma_head128_snapshot", + "compile_parameters": { + "llm.gated_delta_net.head_count": "48", + "llm.gated_delta_net.head_width": "128", + "llm.gated_delta_net.l2_epsilon": "9.99999997e-07", + "llm.gated_delta_net.qk_stride1": "128", + "llm.gated_delta_net.qk_stride2": "10240", + "llm.gated_delta_net.qk_stride3": "30720", + "llm.gated_delta_net.query_head_count": "16", + "llm.gated_delta_net.query_sequence_ratio": "1", + "llm.gated_delta_net.scalar_stride1": "1", + "llm.gated_delta_net.scalar_stride2": "48", + "llm.gated_delta_net.scalar_stride3": "144", + "llm.gated_delta_net.sequence_count": "1", + "llm.gated_delta_net.snapshot_stride": "3145728", + "llm.gated_delta_net.token_count": "3", + "llm.gated_delta_net.value_stride1": "128", + "llm.gated_delta_net.value_stride2": "10240", + "llm.gated_delta_net.value_stride3": "30720", + "llm.gated_delta_net.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_f32_wmma_head128_snapshot", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_f32_wmma_head128_snapshot", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_094_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.hidden_size": "128", + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "token_count": 144 + }, + "kernel": "loom_libs:ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_095_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_096_ggml_add_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.add_rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.add_rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 128, + "integer_parameters": { + "token_count": 3 + }, + "kernel": "loom_libs:ggml_add_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_add_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_097_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_098_ggml_binary_swiglu_symmetric_i4_k32", + "compile_parameters": { + "ggml.binary_swiglu_symmetric_i4.input_size": "17408", + "ggml.binary_swiglu_symmetric_i4.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_binary_swiglu_symmetric_i4_k32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_swiglu_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_099_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_100_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "12288", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_101_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "1024" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": { + "token_count": 3 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_102_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "1024", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_103_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "512", + "ggml.rmsnorm_mul_rope.input_stride2": "12288", + "ggml.rmsnorm_mul_rope.input_stride3": "36864", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "24", + "ggml.rmsnorm_mul_rope.ne2": "3", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "6144", + "ggml.rmsnorm_mul_rope.output_stride3": "18432", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_104_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "256", + "ggml.rmsnorm_mul_rope.input_stride2": "1024", + "ggml.rmsnorm_mul_rope.input_stride3": "3072", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "4", + "ggml.rmsnorm_mul_rope.ne2": "3", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "1024", + "ggml.rmsnorm_mul_rope.output_stride3": "3072", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_105_ggml_set_rows", + "compile_parameters": { + "ggml.set_rows.hidden_capacity": "1024", + "ggml.set_rows.input_format": "32", + "ggml.set_rows.output_format": "16", + "ggml.set_rows.token_capacity": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": { + "cache_row_count": 262144, + "hidden_size": 1024, + "token_count": 3 + }, + "kernel": "loom_libs:ggml_set_rows", + "library_sources": [], + "primary_sources": [ + "ops/set_rows.loom" + ], + "sources": [ + "ops/set_rows.loom" + ], + "symbol": "ggml_set_rows", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_106_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.apply_gate": "1", + "ggml.flash_attention.gate_stride_head": "512", + "ggml.flash_attention.gate_stride_token": "12288", + "ggml.flash_attention.head_size": "256", + "ggml.flash_attention.key_value_head_count": "4", + "ggml.flash_attention.query_head_count": "24" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": { + "key_value_token_count": 256, + "query_token_count": 3 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_107_ggml_quantize_f32_symmetric_i4_k32", + "compile_parameters": { + "ggml.quantize_symmetric_i4_k32.input_size": "6144", + "ggml.quantize_symmetric_i4_k32.token_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_quantize_f32_symmetric_i4_k32", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_quantize_f32_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_108_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "element_count": 15360 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_109_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "5120", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 4, + "integer_parameters": { + "token_count": 3 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_110_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "3", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 3, + "token_count": 3 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_111_ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "248320" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 3 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_112_ggml_concat_dim0_f32", + "compile_parameters": { + "ggml.concat_dim0_f32.lhs_width": "5120", + "ggml.concat_dim0_f32.rhs_width": "5120", + "ggml.concat_dim0_f32.row_count": "3" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_concat_dim0_f32", + "library_sources": [], + "primary_sources": [ + "ops/copy_f32.loom" + ], + "sources": [ + "ops/copy_f32.loom" + ], + "symbol": "ggml_concat_dim0_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_113_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "10240", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "5120" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 3 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_114_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "16", + "ggml.get_rows_f32.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 248320, + "token_count": 16 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_115_ggml_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "token_count": 16 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_116_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "10240", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_117_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "48", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_118_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_119_llm_ssm_conv_dconv4_silu_rollback_f32", + "compile_parameters": { + "llm.ssm_conv.rollback.cache_count": "5", + "llm.ssm_conv.rollback.d_inner": "10240", + "llm.ssm_conv.rollback.n_s": "1", + "llm.ssm_conv.rollback.n_t": "16", + "llm.ssm_conv.rollback.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_ssm_conv_dconv4_silu_rollback_f32", + "library_sources": [], + "primary_sources": [ + "ops/ssm_conv_f32.loom" + ], + "sources": [ + "ops/ssm_conv_f32.loom" + ], + "symbol": "llm_ssm_conv_dconv4_silu_rollback_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_120_llm_gated_delta_net_projection_epilogue_f32", + "compile_parameters": { + "llm.gated_delta_net.epilogue_element_count": "768", + "llm.gated_delta_net.epilogue_head_count": "48", + "llm.gated_delta_net.epilogue_workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_projection_epilogue_f32", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_projection_epilogue_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_121_llm_gated_delta_net_f32_wmma_head128", + "compile_parameters": { + "llm.gated_delta_net.head_count": "48", + "llm.gated_delta_net.head_width": "128", + "llm.gated_delta_net.l2_epsilon": "9.99999997e-07", + "llm.gated_delta_net.qk_stride1": "128", + "llm.gated_delta_net.qk_stride2": "10240", + "llm.gated_delta_net.qk_stride3": "163840", + "llm.gated_delta_net.query_head_count": "16", + "llm.gated_delta_net.query_sequence_ratio": "1", + "llm.gated_delta_net.scalar_stride1": "1", + "llm.gated_delta_net.scalar_stride2": "48", + "llm.gated_delta_net.scalar_stride3": "768", + "llm.gated_delta_net.sequence_count": "1", + "llm.gated_delta_net.token_count": "11", + "llm.gated_delta_net.value_stride1": "128", + "llm.gated_delta_net.value_stride2": "10240", + "llm.gated_delta_net.value_stride3": "163840", + "llm.gated_delta_net.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_f32_wmma_head128", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_f32_wmma_head128", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_122_ggml_copy_f32", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 288, + "integer_parameters": { + "element_count": 786432 + }, + "kernel": "loom_libs:ggml_copy_f32", + "library_sources": [], + "primary_sources": [ + "ops/copy_f32.loom" + ], + "sources": [ + "ops/copy_f32.loom" + ], + "symbol": "ggml_copy_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_123_llm_gated_delta_net_f32_wmma_head128_inplace", + "compile_parameters": { + "llm.gated_delta_net.head_count": "48", + "llm.gated_delta_net.head_width": "128", + "llm.gated_delta_net.l2_epsilon": "9.99999997e-07", + "llm.gated_delta_net.qk_stride1": "128", + "llm.gated_delta_net.qk_stride2": "10240", + "llm.gated_delta_net.qk_stride3": "163840", + "llm.gated_delta_net.query_head_count": "16", + "llm.gated_delta_net.query_sequence_ratio": "1", + "llm.gated_delta_net.scalar_stride1": "1", + "llm.gated_delta_net.scalar_stride2": "48", + "llm.gated_delta_net.scalar_stride3": "768", + "llm.gated_delta_net.sequence_count": "1", + "llm.gated_delta_net.token_count": "1", + "llm.gated_delta_net.value_stride1": "128", + "llm.gated_delta_net.value_stride2": "10240", + "llm.gated_delta_net.value_stride3": "163840", + "llm.gated_delta_net.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 240, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_f32_wmma_head128_inplace", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_f32_wmma_head128_inplace", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_124_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.hidden_size": "128", + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "token_count": 768 + }, + "kernel": "loom_libs:ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_125_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_126_ggml_add_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.add_rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.add_rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 128, + "integer_parameters": { + "token_count": 16 + }, + "kernel": "loom_libs:ggml_add_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_add_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_127_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_128_ggml_binary_swiglu_symmetric_i4_k32", + "compile_parameters": { + "ggml.binary_swiglu_symmetric_i4.input_size": "17408", + "ggml.binary_swiglu_symmetric_i4.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_binary_swiglu_symmetric_i4_k32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_swiglu_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_129_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 65, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_130_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "12288", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_131_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "1024" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": { + "token_count": 16 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_132_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "1024", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_133_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "512", + "ggml.rmsnorm_mul_rope.input_stride2": "12288", + "ggml.rmsnorm_mul_rope.input_stride3": "196608", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "24", + "ggml.rmsnorm_mul_rope.ne2": "16", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "6144", + "ggml.rmsnorm_mul_rope.output_stride3": "98304", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_134_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "256", + "ggml.rmsnorm_mul_rope.input_stride2": "1024", + "ggml.rmsnorm_mul_rope.input_stride3": "16384", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "4", + "ggml.rmsnorm_mul_rope.ne2": "16", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "1024", + "ggml.rmsnorm_mul_rope.output_stride3": "16384", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_135_ggml_set_rows", + "compile_parameters": { + "ggml.set_rows.hidden_capacity": "1024", + "ggml.set_rows.input_format": "32", + "ggml.set_rows.output_format": "16", + "ggml.set_rows.token_capacity": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 34, + "integer_parameters": { + "cache_row_count": 262144, + "hidden_size": 1024, + "token_count": 16 + }, + "kernel": "loom_libs:ggml_set_rows", + "library_sources": [], + "primary_sources": [ + "ops/set_rows.loom" + ], + "sources": [ + "ops/set_rows.loom" + ], + "symbol": "ggml_set_rows", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_136_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.apply_gate": "1", + "ggml.flash_attention.gate_stride_head": "512", + "ggml.flash_attention.gate_stride_token": "12288", + "ggml.flash_attention.head_size": "256", + "ggml.flash_attention.key_value_head_count": "4", + "ggml.flash_attention.query_head_count": "24" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": { + "key_value_token_count": 256, + "query_token_count": 16 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_137_ggml_quantize_f32_symmetric_i4_k32", + "compile_parameters": { + "ggml.quantize_symmetric_i4_k32.input_size": "6144", + "ggml.quantize_symmetric_i4_k32.token_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 17, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_quantize_f32_symmetric_i4_k32", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_quantize_f32_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_138_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "element_count": 81920 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_139_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "5120", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 4, + "integer_parameters": { + "token_count": 16 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_140_ggml_concat_dim0_f32", + "compile_parameters": { + "ggml.concat_dim0_f32.lhs_width": "5120", + "ggml.concat_dim0_f32.rhs_width": "5120", + "ggml.concat_dim0_f32.row_count": "16" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_concat_dim0_f32", + "library_sources": [], + "primary_sources": [ + "ops/copy_f32.loom" + ], + "sources": [ + "ops/copy_f32.loom" + ], + "symbol": "ggml_concat_dim0_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_141_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "10240", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "5120" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 16 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_142_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "1", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 4, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_143_ggml_concat_dim0_f32", + "compile_parameters": { + "ggml.concat_dim0_f32.lhs_width": "5120", + "ggml.concat_dim0_f32.rhs_width": "5120", + "ggml.concat_dim0_f32.row_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_concat_dim0_f32", + "library_sources": [], + "primary_sources": [ + "ops/copy_f32.loom" + ], + "sources": [ + "ops/copy_f32.loom" + ], + "symbol": "ggml_concat_dim0_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_144_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "10240", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "5120" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 4 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_145_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "1", + "ggml.get_rows_f32.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 3, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 248320, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_146_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "5120", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 7, + "integer_parameters": { + "token_count": 1 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_147_ggml_concat_dim0_f32", + "compile_parameters": { + "ggml.concat_dim0_f32.lhs_width": "5120", + "ggml.concat_dim0_f32.rhs_width": "5120", + "ggml.concat_dim0_f32.row_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_concat_dim0_f32", + "library_sources": [], + "primary_sources": [ + "ops/copy_f32.loom" + ], + "sources": [ + "ops/copy_f32.loom" + ], + "symbol": "ggml_concat_dim0_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_148_ggml_mul_mat_f32_f32_decode_wave64", + "compile_parameters": { + "ggml.mul_mat_f32_f32_decode.output_capacity": "5120", + "ggml.mul_mat_f32_f32_decode.token_capacity": "1", + "ggml.mul_mat_f32_f32_decode.weight_format": "6" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "input_size": 10240, + "output_size": 5120, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_149_ggml_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 3, + "integer_parameters": { + "token_count": 1 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_150_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "12288", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 18, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_151_ggml_mul_mat_f32_f32_decode_wave64", + "compile_parameters": { + "ggml.mul_mat_f32_f32_decode.output_capacity": "1024", + "ggml.mul_mat_f32_f32_decode.token_capacity": "1", + "ggml.mul_mat_f32_f32_decode.weight_format": "6" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 18, + "integer_parameters": { + "input_size": 5120, + "output_size": 1024, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_152_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "1024", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 18, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_153_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "512", + "ggml.rmsnorm_mul_rope.input_stride2": "12288", + "ggml.rmsnorm_mul_rope.input_stride3": "12288", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "24", + "ggml.rmsnorm_mul_rope.ne2": "1", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "6144", + "ggml.rmsnorm_mul_rope.output_stride3": "6144", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 18, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_154_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "256", + "ggml.rmsnorm_mul_rope.input_stride2": "1024", + "ggml.rmsnorm_mul_rope.input_stride3": "1024", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "4", + "ggml.rmsnorm_mul_rope.ne2": "1", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "1024", + "ggml.rmsnorm_mul_rope.output_stride3": "1024", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 18, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_155_ggml_set_rows", + "compile_parameters": { + "ggml.set_rows.hidden_capacity": "1024", + "ggml.set_rows.input_format": "32", + "ggml.set_rows.output_format": "16", + "ggml.set_rows.token_capacity": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 36, + "integer_parameters": { + "cache_row_count": 262144, + "hidden_size": 1024, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_set_rows", + "library_sources": [], + "primary_sources": [ + "ops/set_rows.loom" + ], + "sources": [ + "ops/set_rows.loom" + ], + "symbol": "ggml_set_rows", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_156_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.apply_gate": "1", + "ggml.flash_attention.gate_stride_head": "512", + "ggml.flash_attention.gate_stride_token": "12288", + "ggml.flash_attention.head_size": "256", + "ggml.flash_attention.key_value_head_count": "4", + "ggml.flash_attention.query_head_count": "24" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 18, + "integer_parameters": { + "key_value_token_count": 256, + "query_token_count": 1 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_157_ggml_quantize_f32_symmetric_i4_k32", + "compile_parameters": { + "ggml.quantize_symmetric_i4_k32.input_size": "6144", + "ggml.quantize_symmetric_i4_k32.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 18, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_quantize_f32_symmetric_i4_k32", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_quantize_f32_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_158_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 66, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_159_ggml_add_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.add_rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.add_rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 129, + "integer_parameters": { + "token_count": 1 + }, + "kernel": "loom_libs:ggml_add_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_add_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_160_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 66, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_161_ggml_binary_swiglu_symmetric_i4_k32", + "compile_parameters": { + "ggml.binary_swiglu_symmetric_i4.input_size": "17408", + "ggml.binary_swiglu_symmetric_i4.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 66, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_binary_swiglu_symmetric_i4_k32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_swiglu_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_162_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 66, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_163_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 3, + "integer_parameters": { + "element_count": 5120 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_164_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "1", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 1, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_165_ggml_quantize_f32_symmetric_i4_k32", + "compile_parameters": { + "ggml.quantize_symmetric_i4_k32.input_size": "5120", + "ggml.quantize_symmetric_i4_k32.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_quantize_f32_symmetric_i4_k32", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_quantize_f32_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_166_ggml_select_symmetric_i4_k32_groups", + "compile_parameters": { + "ggml.mul_mat_q6_k_shortlist.input_size": "5120", + "ggml.mul_mat_q6_k_shortlist.selected_group_count": "96" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_select_symmetric_i4_k32_groups", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_select_symmetric_i4_k32_groups", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_167_ggml_mul_mat_q6_k_symmetric_i2_scan_token1", + "compile_parameters": { + "ggml.mul_mat_q6_k_shortlist.input_size": "5120", + "ggml.mul_mat_q6_k_shortlist.output_size": "248320" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_q6_k_symmetric_i2_scan_token1", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_symmetric_i2_scan_token1", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_168_ggml_top_k8_f32_partitions_register", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "element_count": 248320 + }, + "kernel": "loom_libs:ggml_top_k8_f32_partitions_register", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_top_k8_f32_partitions_register", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_169_ggml_top_k128_f32_reduce_gather_register", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "element_count": 248320 + }, + "kernel": "loom_libs:ggml_top_k128_f32_reduce_gather_register", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_top_k128_f32_reduce_gather_register", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_170_ggml_fill_negative_f32", + "compile_parameters": {}, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "element_count": 248320 + }, + "kernel": "loom_libs:ggml_fill_negative_f32", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_fill_negative_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_171_ggml_mul_mat_q6_k_packed_selected_refine_token1", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "248320", + "ggml.mul_mat_q6_k_packed.weight_offset": "337715200" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 2, + "integer_parameters": { + "candidate_count": 64, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_packed_selected_refine_token1", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_packed_selected_refine_token1", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "candidate_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_172_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "10240", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_173_ggml_mul_mat_f32_f32_decode_wave64", + "compile_parameters": { + "ggml.mul_mat_f32_f32_decode.output_capacity": "48", + "ggml.mul_mat_f32_f32_decode.token_capacity": "1", + "ggml.mul_mat_f32_f32_decode.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 96, + "integer_parameters": { + "input_size": 5120, + "output_size": 48, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + "library_sources": [ + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/mul_mat_f32_f32_decode.loom" + ], + "sources": [ + "ops/mul_mat_f32_f32_decode.loom", + "motifs/dequant.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_mul_mat_f32_f32_decode_wave64", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_174_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "1" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_175_llm_ssm_conv_dconv4_silu_rollback_f32", + "compile_parameters": { + "llm.ssm_conv.rollback.cache_count": "5", + "llm.ssm_conv.rollback.d_inner": "10240", + "llm.ssm_conv.rollback.n_s": "1", + "llm.ssm_conv.rollback.n_t": "1", + "llm.ssm_conv.rollback.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_ssm_conv_dconv4_silu_rollback_f32", + "library_sources": [], + "primary_sources": [ + "ops/ssm_conv_f32.loom" + ], + "sources": [ + "ops/ssm_conv_f32.loom" + ], + "symbol": "llm_ssm_conv_dconv4_silu_rollback_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_176_llm_gated_delta_net_projection_epilogue_f32", + "compile_parameters": { + "llm.gated_delta_net.epilogue_element_count": "48", + "llm.gated_delta_net.epilogue_head_count": "48", + "llm.gated_delta_net.epilogue_workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_projection_epilogue_f32", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_projection_epilogue_f32", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_177_llm_gated_delta_net_f32_wmma_head128_inplace", + "compile_parameters": { + "llm.gated_delta_net.head_count": "48", + "llm.gated_delta_net.head_width": "128", + "llm.gated_delta_net.l2_epsilon": "9.99999997e-07", + "llm.gated_delta_net.qk_stride1": "128", + "llm.gated_delta_net.qk_stride2": "10240", + "llm.gated_delta_net.qk_stride3": "10240", + "llm.gated_delta_net.query_head_count": "16", + "llm.gated_delta_net.query_sequence_ratio": "1", + "llm.gated_delta_net.scalar_stride1": "1", + "llm.gated_delta_net.scalar_stride2": "48", + "llm.gated_delta_net.scalar_stride3": "48", + "llm.gated_delta_net.sequence_count": "1", + "llm.gated_delta_net.token_count": "1", + "llm.gated_delta_net.value_stride1": "128", + "llm.gated_delta_net.value_stride2": "10240", + "llm.gated_delta_net.value_stride3": "10240", + "llm.gated_delta_net.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_f32_wmma_head128_inplace", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_f32_wmma_head128_inplace", + "workload_parameters": [] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_178_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.hidden_size": "128", + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "token_count": 48 + }, + "kernel": "loom_libs:ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_179_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "2", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 2, + "token_count": 2 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@qwen38_27b_udq4kxl_companion_tg_c5_180_ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "248320" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 2 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + } + ], + "generated_count": 181, + "kernel_counts": { + "loom_libs:ggml_add_rmsnorm_binary_symmetric_i4_k32": 1024, + "loom_libs:ggml_binary_f32": 16, + "loom_libs:ggml_binary_swiglu_symmetric_i4_k32": 520, + "loom_libs:ggml_concat_dim0_f32": 8, + "loom_libs:ggml_copy_f32": 288, + "loom_libs:ggml_fill_negative_f32": 1, + "loom_libs:ggml_flash_attention_f32_f16_wmma": 136, + "loom_libs:ggml_get_rows_f32": 793, + "loom_libs:ggml_mul_mat_f32_f32_decode_wave64": 116, + "loom_libs:ggml_mul_mat_f32_f32_wmma": 36, + "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma": 88, + "loom_libs:ggml_mul_mat_q6_k_packed_selected_refine_token1": 2, + "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma": 8, + "loom_libs:ggml_mul_mat_q6_k_symmetric_i2_scan_token1": 1, + "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma": 856, + "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma": 904, + "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma": 1176, + "loom_libs:ggml_quantize_f32_symmetric_i4_k32": 137, + "loom_libs:ggml_rmsnorm_binary_f32": 32, + "loom_libs:ggml_rmsnorm_binary_symmetric_i4_k32": 16, + "loom_libs:ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32": 384, + "loom_libs:ggml_rmsnorm_mul_rope_f32": 272, + "loom_libs:ggml_scale_bias_f32": 192, + "loom_libs:ggml_select_symmetric_i4_k32_groups": 1, + "loom_libs:ggml_set_rows": 272, + "loom_libs:ggml_top_k128_f32_reduce_gather_register": 1, + "loom_libs:ggml_top_k8_f32_partitions_register": 1, + "loom_libs:llm_gated_delta_net_f32_wmma_head128": 48, + "loom_libs:llm_gated_delta_net_f32_wmma_head128_inplace": 288, + "loom_libs:llm_gated_delta_net_f32_wmma_head128_snapshot": 288, + "loom_libs:llm_gated_delta_net_projection_epilogue_f32": 384, + "loom_libs:llm_ssm_conv_dconv4_silu_rollback_f32": 384 + }, + "loom_source": "tools/benchmarks/loom/v2_mtp4_companion_all.work.loom", + "model": "qwen38_27b_udq4kxl_companion", + "scenario": "tg_c5", + "schema": "ggml-hrx-model-loom-benchmarks-v2" +} diff --git a/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_mtp4_companion_all.work.loom b/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_mtp4_companion_all.work.loom new file mode 100644 index 000000000000..841905c9608e --- /dev/null +++ b/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_mtp4_companion_all.work.loom @@ -0,0 +1,2386 @@ +// Generated by tools/benchmarks/generate-model-benchmarks.py for qwen38_27b_udq4kxl_companion scenarios: tg_c5. +// Regenerate from an HRX command program dump rather than editing by hand. + +target.decl @ggml_binary_f32_gfx11_wave64 +target.decl @ggml_binary_swiglu_i4_gfx11_wave32 +target.decl @ggml_copy_f32_gfx11_wave64 +target.decl @ggml_flash_attention_gfx11_wave64 +target.decl @ggml_get_rows_f32_gfx11_wave64 +target.decl @ggml_mul_mat_f32_f32_decode_gfx11_wave64 +target.decl @ggml_mul_mat_gfx11_wave64 +target.decl @ggml_rmsnorm_binary_gfx11_wave32 +target.decl @ggml_rmsnorm_gfx11_wave32 +target.decl @ggml_scale_bias_f32_gfx11_wave64 +target.decl @ggml_set_rows_gfx11_wave64 + +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_add_rmsnorm_binary_symmetric_i4_k32(%token_count: index) launch(%token_count: index, %lhs: buffer, %rhs: buffer, %residual_output: buffer, %weight: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_binary_f32_gfx11_wave64) @ggml_binary_f32(%element_count: index) launch(%element_count: index, %lhs: buffer, %rhs: buffer, %output: buffer) +kernel.decl target(@ggml_binary_swiglu_i4_gfx11_wave32) @ggml_binary_swiglu_symmetric_i4_k32() launch(%lhs: buffer, %rhs: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_copy_f32_gfx11_wave64) @ggml_concat_dim0_f32() launch(%lhs: buffer, %rhs: buffer, %output: buffer) +kernel.decl target(@ggml_copy_f32_gfx11_wave64) @ggml_copy_f32(%element_count: index) launch(%element_count: index, %source: buffer, %output: buffer) +kernel.decl @ggml_fill_negative_f32(%element_count: index) launch(%element_count: index, %output: buffer) +kernel.decl target(@ggml_flash_attention_gfx11_wave64) @ggml_flash_attention_f32_f16_wmma(%query_token_count: index, %key_value_token_count: index) launch(%query_token_count: index, %key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %gate: buffer, %output: buffer) +kernel.decl target(@ggml_get_rows_f32_gfx11_wave64) @ggml_get_rows_f32(%token_count: index, %row_count: index, %hidden_size: index) launch(%token_count: index, %row_count: index, %hidden_size: index, %token_ids: buffer, %weight: buffer, %output: buffer) +kernel.decl target(@ggml_mul_mat_f32_f32_decode_gfx11_wave64) @ggml_mul_mat_f32_f32_decode_wave64(%token_count: index, %input_size: index, %output_size: index) launch(%token_count: index, %input_size: index, %output_size: index, %input: buffer, %weight: buffer, %output: buffer) +kernel.decl target(@ggml_mul_mat_gfx11_wave64) @ggml_mul_mat_f32_f32_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) +kernel.decl @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) +kernel.decl @ggml_mul_mat_q6_k_packed_selected_refine_token1(%token_count: index, %candidate_count: index) launch(%token_count: index, %candidate_count: index, %input: buffer, %weight: buffer, %candidates: buffer, %exact_output: buffer, %output: buffer) +kernel.decl @ggml_mul_mat_q6_k_packed_token1_f16_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) +kernel.decl @ggml_mul_mat_q6_k_symmetric_i2_scan_token1() launch(%weight: buffer, %output: buffer, %qact: buffer, %scales: buffer, %selected_mask: buffer) +kernel.decl @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma() launch(%first_weight: buffer, %second_weight: buffer, %first_output: buffer, %second_output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma() launch(%weight: buffer, %input: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer, %partial: buffer, %completion_counters: buffer) +kernel.decl @ggml_mul_mat_symmetric_i4_lowrow_wmma() launch(%weight: buffer, %input: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl @ggml_quantize_f32_symmetric_i4_k32() launch(%input: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_rmsnorm_binary_f32(%token_count: index) launch(%token_count: index, %input: buffer, %rhs: buffer, %output: buffer) +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_rmsnorm_binary_symmetric_i4_k32(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %raw_gate: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_rmsnorm_gfx11_wave32) @ggml_rmsnorm_mul_rope_f32() launch(%input: buffer, %weight: buffer, %positions: buffer, %output: buffer) +kernel.decl target(@ggml_scale_bias_f32_gfx11_wave64) @ggml_scale_bias_f32(%element_count: index) launch(%element_count: index, %input: buffer, %output: buffer) +kernel.decl @ggml_select_symmetric_i4_k32_groups() launch(%input: buffer, %selected_mask: buffer) +kernel.decl target(@ggml_set_rows_gfx11_wave64) @ggml_set_rows(%token_count: index, %cache_row_count: index, %hidden_size: index) launch(%token_count: index, %cache_row_count: index, %hidden_size: index, %rows: buffer, %indices: buffer, %cache: buffer) +kernel.decl @ggml_top_k128_f32_reduce_gather_register(%element_count: index) launch(%element_count: index, %partial_values: buffer, %partial_ids: buffer, %candidate_output: buffer, %value_output: buffer) +kernel.decl @ggml_top_k8_f32_partitions_register(%element_count: index) launch(%element_count: index, %values: buffer, %partial_values: buffer, %partial_ids: buffer) +kernel.decl @llm_gated_delta_net_f32_wmma_head128() launch(%q: buffer, %k: buffer, %v: buffer, %g: buffer, %beta: buffer, %state_in: buffer, %dst: buffer, %rms_scales: buffer) +kernel.decl @llm_gated_delta_net_f32_wmma_head128_inplace() launch(%q: buffer, %k: buffer, %v: buffer, %g: buffer, %beta: buffer, %state_inout: buffer, %dst: buffer, %rms_scales: buffer) +kernel.decl @llm_gated_delta_net_f32_wmma_head128_snapshot() launch(%q: buffer, %k: buffer, %v: buffer, %g: buffer, %beta: buffer, %state_in: buffer, %snapshot_cache: buffer, %dst: buffer, %rms_scales: buffer) +kernel.decl @llm_gated_delta_net_projection_epilogue_f32() launch(%alpha_raw: buffer, %beta_raw: buffer, %bias: buffer, %a_scale: buffer, %gate_dst: buffer, %beta_dst: buffer) +kernel.decl @llm_ssm_conv_dconv4_silu_rollback_f32() launch(%state: buffer, %x: buffer, %filter: buffer, %output: buffer, %cache0: buffer, %cache1: buffer, %cache2: buffer, %cache3: buffer, %cache4: buffer) + + +// Scenario: tg_c5 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_000_ggml_get_rows_f32_case { + %token_count = check.literal value(2) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<8xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<8xi8>, tensor<715161600xi8>, tensor<40960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_000_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_000_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_001_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<40960xi8>, tensor<20480xi8>, tensor<40960xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_001_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_001_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + %partial = check.generate.fill value(0) : tensor<81920xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<40960xi8>, tensor<81920xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>, tensor<81920xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_003_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<168960xi8> + %second_weight = check.generate.fill value(0) : tensor<168960xi8> + %first_output = check.generate.fill value(0) : tensor<384xi8> + %second_output = check.generate.fill value(0) : tensor<384xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<168960xi8>, tensor<168960xi8>, tensor<384xi8>, tensor<384xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_003_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_003_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_004_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<49152xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<40960xi8>, tensor<49152xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_004_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_004_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_005_ggml_scale_bias_f32_case { + %element_count = check.literal value(30720) : index + %input = check.generate.fill value(0) : tensor<122880xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @ggml_scale_bias_f32[%element_count](%element_count, %input, %output) : [index](index, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_005_ggml_scale_bias_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_005_ggml_scale_bias_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_006_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(20) : index + %hidden_size = check.literal value(30720) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<20x30720xf32> + %output = check.generate.fill value(1.0) : tensor<1x30720xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<20x30720xf32>, tensor<1x30720xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_006_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_006_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_007_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<81920xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<81920xi8>, tensor<163840xi8>, tensor<81920xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_007_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_007_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_008_ggml_scale_bias_f32_case { + %element_count = check.literal value(786432) : index + %input = check.generate.fill value(0) : tensor<3145728xi8> + %output = check.generate.fill value(0) : tensor<3145728xi8> + kernel.launch @ggml_scale_bias_f32[%element_count](%element_count, %input, %output) : [index](index, tensor<3145728xi8>, tensor<3145728xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_008_ggml_scale_bias_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_008_ggml_scale_bias_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_009_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(20) : index + %hidden_size = check.literal value(786432) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<20x786432xf32> + %output = check.generate.fill value(1.0) : tensor<1x786432xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<20x786432xf32>, tensor<1x786432xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_009_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_009_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_010_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<384xi8> + %beta_raw = check.generate.fill value(0) : tensor<384xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<384xi8> + %beta_dst = check.generate.fill value(0) : tensor<384xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<384xi8>, tensor<384xi8>, tensor<192xi8>, tensor<192xi8>, tensor<384xi8>, tensor<384xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_010_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_010_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_011_llm_gated_delta_net_f32_wmma_head128_snapshot_case { + %q = check.generate.fill value(0) : tensor<49152xi8> + %k = check.generate.fill value(0) : tensor<49152xi8> + %v = check.generate.fill value(0) : tensor<65536xi8> + %g = check.generate.fill value(0) : tensor<384xi8> + %beta = check.generate.fill value(0) : tensor<384xi8> + %state_in = check.generate.fill value(0) : tensor<3145728xi8> + %snapshot_cache = check.generate.fill value(0) : tensor<15728640xi8> + %dst = check.generate.fill value(0) : tensor<49152xi8> + %rms_scales = check.generate.fill value(0) : tensor<384xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_snapshot[](%q, %k, %v, %g, %beta, %state_in, %snapshot_cache, %dst, %rms_scales) : [](tensor<49152xi8>, tensor<49152xi8>, tensor<65536xi8>, tensor<384xi8>, tensor<384xi8>, tensor<3145728xi8>, tensor<15728640xi8>, tensor<49152xi8>, tensor<384xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_011_llm_gated_delta_net_f32_wmma_head128_snapshot_case> @qwen38_27b_udq4kxl_companion_tg_c5_011_llm_gated_delta_net_f32_wmma_head128_snapshot + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_012_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(96) : index + %input = check.generate.fill value(0) : tensor<49152xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<49152xi8> + %output = check.generate.fill value(0) : tensor<49152xi8> + %quantized_values = check.generate.fill value(0) : tensor<6144xi8> + %scales = check.generate.fill value(0) : tensor<1536xi8> + %sums = check.generate.fill value(0) : tensor<1536xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<49152xi8>, tensor<512xi8>, tensor<49152xi8>, tensor<49152xi8>, tensor<6144xi8>, tensor<1536xi8>, tensor<1536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_012_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_012_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_013_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<49152xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<6144xi8> + %scales = check.generate.fill value(0) : tensor<1536xi8> + %sums = check.generate.fill value(0) : tensor<1536xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<49152xi8>, tensor<40960xi8>, tensor<6144xi8>, tensor<1536xi8>, tensor<1536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_013_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_013_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_014_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(2) : index + %lhs = check.generate.fill value(0) : tensor<40960xi8> + %rhs = check.generate.fill value(0) : tensor<40960xi8> + %residual_output = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<40960xi8>, tensor<40960xi8>, tensor<40960xi8>, tensor<20480xi8>, tensor<40960xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_014_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_014_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_015_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<139264xi8> + %second_output = check.generate.fill value(0) : tensor<139264xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<139264xi8>, tensor<139264xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_015_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_015_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_016_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<139264xi8> + %rhs = check.generate.fill value(0) : tensor<139264xi8> + %output = check.generate.fill value(0) : tensor<139264xi8> + %quantized_values = check.generate.fill value(0) : tensor<17408xi8> + %scales = check.generate.fill value(0) : tensor<4352xi8> + %sums = check.generate.fill value(0) : tensor<4352xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<139264xi8>, tensor<139264xi8>, tensor<139264xi8>, tensor<17408xi8>, tensor<4352xi8>, tensor<4352xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_016_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_016_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_017_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<139264xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<17408xi8> + %scales = check.generate.fill value(0) : tensor<4352xi8> + %sums = check.generate.fill value(0) : tensor<4352xi8> + %partial = check.generate.fill value(0) : tensor<40960xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<139264xi8>, tensor<40960xi8>, tensor<17408xi8>, tensor<4352xi8>, tensor<4352xi8>, tensor<40960xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_017_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_017_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_018_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<98304xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<40960xi8>, tensor<98304xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_018_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_018_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_019_ggml_mul_mat_f32_f32_wmma_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<4300800xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + kernel.launch @ggml_mul_mat_f32_f32_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<40960xi8>, tensor<4300800xi8>, tensor<8192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_019_ggml_mul_mat_f32_f32_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_019_ggml_mul_mat_f32_f32_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_020_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + %quantized_values = check.generate.fill value(0) : tensor<5120xi8> + %scales = check.generate.fill value(0) : tensor<1280xi8> + %sums = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<40960xi8>, tensor<8192xi8>, tensor<5120xi8>, tensor<1280xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_020_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_020_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_021_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<97280xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<32xi8> + %output = check.generate.fill value(0) : tensor<49152xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<97280xi8>, tensor<1024xi8>, tensor<32xi8>, tensor<49152xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_021_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_021_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_022_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<8192xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<32xi8> + %output = check.generate.fill value(0) : tensor<8192xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<8192xi8>, tensor<1024xi8>, tensor<32xi8>, tensor<8192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_022_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_022_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_023_ggml_set_rows_case { + %token_count = check.literal value(2) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<2x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<2xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<2x1024xf32>, tensor<2xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_023_ggml_set_rows_case> @qwen38_27b_udq4kxl_companion_tg_c5_023_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_024_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(2) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<2x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<2x256xf16> + %output = check.generate.fill value(1.0) : tensor<2x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<2x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<2x256xf16>, tensor<2x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_024_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_024_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_025_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<49152xi8> + %quantized_values = check.generate.fill value(0) : tensor<6144xi8> + %scales = check.generate.fill value(0) : tensor<1536xi8> + %sums = check.generate.fill value(0) : tensor<1536xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<49152xi8>, tensor<6144xi8>, tensor<1536xi8>, tensor<1536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_025_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_025_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_026_ggml_binary_f32_case { + %element_count = check.literal value(10240) : index + %lhs = check.generate.fill value(2.0) : tensor<10240xf32> + %rhs = check.generate.fill value(3.0) : tensor<10240xf32> + %output = check.generate.fill value(0.0) : tensor<10240xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<10240xf32>, tensor<10240xf32>, tensor<10240xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_026_ggml_binary_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_026_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_027_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(2.0) : tensor<2x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<2x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<2x5120xf32>, tensor<5120xf32>, tensor<2x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_027_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_027_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_028_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(2) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<2x5120xf32> + %output = check.generate.fill value(1.0) : tensor<1x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<2x5120xf32>, tensor<1x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_028_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_028_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_029_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<1042944000xi8> + %output = check.generate.fill value(0) : tensor<993280xi8> + kernel.launch @ggml_mul_mat_q6_k_packed_token1_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<20480xi8>, tensor<1042944000xi8>, tensor<993280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_029_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_029_ggml_mul_mat_q6_k_packed_token1_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_030_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<40960xi8> + %rhs = check.generate.fill value(0) : tensor<40960xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<40960xi8>, tensor<40960xi8>, tensor<81920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_030_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_030_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_031_ggml_mul_mat_f32_f32_wmma_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<43008000xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + kernel.launch @ggml_mul_mat_f32_f32_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<81920xi8>, tensor<43008000xi8>, tensor<40960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_031_ggml_mul_mat_f32_f32_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_031_ggml_mul_mat_f32_f32_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_032_ggml_get_rows_f32_case { + %token_count = check.literal value(4) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<16xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<16xi8>, tensor<715161600xi8>, tensor<81920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_032_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_032_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_033_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<81920xi8>, tensor<20480xi8>, tensor<81920xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_033_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_033_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_034_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<163840xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + %partial = check.generate.fill value(0) : tensor<163840xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<81920xi8>, tensor<163840xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>, tensor<163840xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_034_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_034_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_035_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<168960xi8> + %second_weight = check.generate.fill value(0) : tensor<168960xi8> + %first_output = check.generate.fill value(0) : tensor<768xi8> + %second_output = check.generate.fill value(0) : tensor<768xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<168960xi8>, tensor<168960xi8>, tensor<768xi8>, tensor<768xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_035_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_035_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_036_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<98304xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<81920xi8>, tensor<98304xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_036_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_036_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_037_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<163840xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<163840xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<163840xi8>, tensor<163840xi8>, tensor<163840xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_037_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_037_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_038_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<768xi8> + %beta_raw = check.generate.fill value(0) : tensor<768xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<768xi8> + %beta_dst = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<768xi8>, tensor<768xi8>, tensor<192xi8>, tensor<192xi8>, tensor<768xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_038_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_038_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_039_llm_gated_delta_net_f32_wmma_head128_snapshot_case { + %q = check.generate.fill value(0) : tensor<131072xi8> + %k = check.generate.fill value(0) : tensor<131072xi8> + %v = check.generate.fill value(0) : tensor<147456xi8> + %g = check.generate.fill value(0) : tensor<768xi8> + %beta = check.generate.fill value(0) : tensor<768xi8> + %state_in = check.generate.fill value(0) : tensor<3145728xi8> + %snapshot_cache = check.generate.fill value(0) : tensor<40894464xi8> + %dst = check.generate.fill value(0) : tensor<98304xi8> + %rms_scales = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_snapshot[](%q, %k, %v, %g, %beta, %state_in, %snapshot_cache, %dst, %rms_scales) : [](tensor<131072xi8>, tensor<131072xi8>, tensor<147456xi8>, tensor<768xi8>, tensor<768xi8>, tensor<3145728xi8>, tensor<40894464xi8>, tensor<98304xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_039_llm_gated_delta_net_f32_wmma_head128_snapshot_case> @qwen38_27b_udq4kxl_companion_tg_c5_039_llm_gated_delta_net_f32_wmma_head128_snapshot + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_040_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(192) : index + %input = check.generate.fill value(0) : tensor<98304xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<98304xi8> + %output = check.generate.fill value(0) : tensor<98304xi8> + %quantized_values = check.generate.fill value(0) : tensor<12288xi8> + %scales = check.generate.fill value(0) : tensor<3072xi8> + %sums = check.generate.fill value(0) : tensor<3072xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<98304xi8>, tensor<512xi8>, tensor<98304xi8>, tensor<98304xi8>, tensor<12288xi8>, tensor<3072xi8>, tensor<3072xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_040_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_040_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_041_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<98304xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<12288xi8> + %scales = check.generate.fill value(0) : tensor<3072xi8> + %sums = check.generate.fill value(0) : tensor<3072xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<98304xi8>, tensor<81920xi8>, tensor<12288xi8>, tensor<3072xi8>, tensor<3072xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_041_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_041_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_042_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(4) : index + %lhs = check.generate.fill value(0) : tensor<81920xi8> + %rhs = check.generate.fill value(0) : tensor<81920xi8> + %residual_output = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<81920xi8>, tensor<81920xi8>, tensor<81920xi8>, tensor<20480xi8>, tensor<81920xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_042_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_042_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_043_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<278528xi8> + %second_output = check.generate.fill value(0) : tensor<278528xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<278528xi8>, tensor<278528xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_043_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_043_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_044_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<278528xi8> + %rhs = check.generate.fill value(0) : tensor<278528xi8> + %output = check.generate.fill value(0) : tensor<278528xi8> + %quantized_values = check.generate.fill value(0) : tensor<34816xi8> + %scales = check.generate.fill value(0) : tensor<8704xi8> + %sums = check.generate.fill value(0) : tensor<8704xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<278528xi8>, tensor<278528xi8>, tensor<278528xi8>, tensor<34816xi8>, tensor<8704xi8>, tensor<8704xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_044_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_044_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_045_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<278528xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<34816xi8> + %scales = check.generate.fill value(0) : tensor<8704xi8> + %sums = check.generate.fill value(0) : tensor<8704xi8> + %partial = check.generate.fill value(0) : tensor<81920xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<278528xi8>, tensor<81920xi8>, tensor<34816xi8>, tensor<8704xi8>, tensor<8704xi8>, tensor<81920xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_045_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_045_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_046_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<196608xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<81920xi8>, tensor<196608xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_046_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_046_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_047_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<5611520xi8> + %output = check.generate.fill value(0) : tensor<16384xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<81920xi8>, tensor<5611520xi8>, tensor<16384xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_047_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_047_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_048_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<16384xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<81920xi8>, tensor<16384xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_048_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_048_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_049_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<195584xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<64xi8> + %output = check.generate.fill value(0) : tensor<98304xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<195584xi8>, tensor<1024xi8>, tensor<64xi8>, tensor<98304xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_049_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_049_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_050_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<16384xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<64xi8> + %output = check.generate.fill value(0) : tensor<16384xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<16384xi8>, tensor<1024xi8>, tensor<64xi8>, tensor<16384xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_050_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_050_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_051_ggml_set_rows_case { + %token_count = check.literal value(4) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<4x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<4xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<4x1024xf32>, tensor<4xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_051_ggml_set_rows_case> @qwen38_27b_udq4kxl_companion_tg_c5_051_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_052_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(4) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<4x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<4x256xf16> + %output = check.generate.fill value(1.0) : tensor<4x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<4x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<4x256xf16>, tensor<4x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_052_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_052_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_053_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<98304xi8> + %quantized_values = check.generate.fill value(0) : tensor<12288xi8> + %scales = check.generate.fill value(0) : tensor<3072xi8> + %sums = check.generate.fill value(0) : tensor<3072xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<98304xi8>, tensor<12288xi8>, tensor<3072xi8>, tensor<3072xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_053_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_053_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_054_ggml_binary_f32_case { + %element_count = check.literal value(20480) : index + %lhs = check.generate.fill value(2.0) : tensor<20480xf32> + %rhs = check.generate.fill value(3.0) : tensor<20480xf32> + %output = check.generate.fill value(0.0) : tensor<20480xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<20480xf32>, tensor<20480xf32>, tensor<20480xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_054_ggml_binary_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_054_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_055_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(2.0) : tensor<4x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<4x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<4x5120xf32>, tensor<5120xf32>, tensor<4x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_055_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_055_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_056_ggml_get_rows_f32_case { + %token_count = check.literal value(4) : index + %row_count = check.literal value(4) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<4xi32> + %weight = check.generate.fill value(0.0) : tensor<4x5120xf32> + %output = check.generate.fill value(1.0) : tensor<4x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<4xi32>, tensor<4x5120xf32>, tensor<4x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_056_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_056_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_057_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<1042944000xi8> + %output = check.generate.fill value(0) : tensor<3973120xi8> + kernel.launch @ggml_mul_mat_q6_k_packed_token1_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<81920xi8>, tensor<1042944000xi8>, tensor<3973120xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_057_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_057_ggml_mul_mat_q6_k_packed_token1_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_058_ggml_get_rows_f32_case { + %token_count = check.literal value(5) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<20xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<20xi8>, tensor<715161600xi8>, tensor<102400xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_058_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_058_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_059_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(5) : index + %input = check.generate.fill value(0) : tensor<102400xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<102400xi8>, tensor<20480xi8>, tensor<102400xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_059_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_059_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_060_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<102400xi8> + %output = check.generate.fill value(0) : tensor<204800xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + %partial = check.generate.fill value(0) : tensor<204800xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<102400xi8>, tensor<204800xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>, tensor<204800xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_060_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_060_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_061_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<168960xi8> + %second_weight = check.generate.fill value(0) : tensor<168960xi8> + %first_output = check.generate.fill value(0) : tensor<960xi8> + %second_output = check.generate.fill value(0) : tensor<960xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<168960xi8>, tensor<168960xi8>, tensor<960xi8>, tensor<960xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_061_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_061_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_062_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<102400xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<102400xi8>, tensor<122880xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_062_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_062_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_063_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<204800xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<204800xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<204800xi8>, tensor<163840xi8>, tensor<204800xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_063_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_063_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_064_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<960xi8> + %beta_raw = check.generate.fill value(0) : tensor<960xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<960xi8> + %beta_dst = check.generate.fill value(0) : tensor<960xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<960xi8>, tensor<960xi8>, tensor<192xi8>, tensor<192xi8>, tensor<960xi8>, tensor<960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_064_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_064_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_065_llm_gated_delta_net_f32_wmma_head128_snapshot_case { + %q = check.generate.fill value(0) : tensor<172032xi8> + %k = check.generate.fill value(0) : tensor<172032xi8> + %v = check.generate.fill value(0) : tensor<188416xi8> + %g = check.generate.fill value(0) : tensor<960xi8> + %beta = check.generate.fill value(0) : tensor<960xi8> + %state_in = check.generate.fill value(0) : tensor<3145728xi8> + %snapshot_cache = check.generate.fill value(0) : tensor<53477376xi8> + %dst = check.generate.fill value(0) : tensor<122880xi8> + %rms_scales = check.generate.fill value(0) : tensor<960xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_snapshot[](%q, %k, %v, %g, %beta, %state_in, %snapshot_cache, %dst, %rms_scales) : [](tensor<172032xi8>, tensor<172032xi8>, tensor<188416xi8>, tensor<960xi8>, tensor<960xi8>, tensor<3145728xi8>, tensor<53477376xi8>, tensor<122880xi8>, tensor<960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_065_llm_gated_delta_net_f32_wmma_head128_snapshot_case> @qwen38_27b_udq4kxl_companion_tg_c5_065_llm_gated_delta_net_f32_wmma_head128_snapshot + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_066_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(240) : index + %input = check.generate.fill value(0) : tensor<122880xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<122880xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + %quantized_values = check.generate.fill value(0) : tensor<15360xi8> + %scales = check.generate.fill value(0) : tensor<3840xi8> + %sums = check.generate.fill value(0) : tensor<3840xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<122880xi8>, tensor<512xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<15360xi8>, tensor<3840xi8>, tensor<3840xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_066_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_066_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_067_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<122880xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + %quantized_values = check.generate.fill value(0) : tensor<15360xi8> + %scales = check.generate.fill value(0) : tensor<3840xi8> + %sums = check.generate.fill value(0) : tensor<3840xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<122880xi8>, tensor<102400xi8>, tensor<15360xi8>, tensor<3840xi8>, tensor<3840xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_067_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_067_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_068_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(5) : index + %lhs = check.generate.fill value(0) : tensor<102400xi8> + %rhs = check.generate.fill value(0) : tensor<102400xi8> + %residual_output = check.generate.fill value(0) : tensor<102400xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<102400xi8>, tensor<102400xi8>, tensor<102400xi8>, tensor<20480xi8>, tensor<102400xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_068_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_068_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_069_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<348160xi8> + %second_output = check.generate.fill value(0) : tensor<348160xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<348160xi8>, tensor<348160xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_069_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_069_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_070_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<348160xi8> + %rhs = check.generate.fill value(0) : tensor<348160xi8> + %output = check.generate.fill value(0) : tensor<348160xi8> + %quantized_values = check.generate.fill value(0) : tensor<43520xi8> + %scales = check.generate.fill value(0) : tensor<10880xi8> + %sums = check.generate.fill value(0) : tensor<10880xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<348160xi8>, tensor<348160xi8>, tensor<348160xi8>, tensor<43520xi8>, tensor<10880xi8>, tensor<10880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_070_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_070_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_071_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<348160xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + %quantized_values = check.generate.fill value(0) : tensor<43520xi8> + %scales = check.generate.fill value(0) : tensor<10880xi8> + %sums = check.generate.fill value(0) : tensor<10880xi8> + %partial = check.generate.fill value(0) : tensor<102400xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<348160xi8>, tensor<102400xi8>, tensor<43520xi8>, tensor<10880xi8>, tensor<10880xi8>, tensor<102400xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_071_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_071_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_072_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<102400xi8> + %output = check.generate.fill value(0) : tensor<245760xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<102400xi8>, tensor<245760xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_072_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_072_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_073_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(5) : index + %input = check.generate.fill value(0) : tensor<102400xi8> + %weight = check.generate.fill value(0) : tensor<5611520xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<102400xi8>, tensor<5611520xi8>, tensor<20480xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_073_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_073_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_074_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<102400xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<12800xi8> + %scales = check.generate.fill value(0) : tensor<3200xi8> + %sums = check.generate.fill value(0) : tensor<3200xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<102400xi8>, tensor<20480xi8>, tensor<12800xi8>, tensor<3200xi8>, tensor<3200xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_074_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_074_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_075_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<244736xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<80xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<244736xi8>, tensor<1024xi8>, tensor<80xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_075_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_075_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_076_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<80xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<20480xi8>, tensor<1024xi8>, tensor<80xi8>, tensor<20480xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_076_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_076_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_077_ggml_set_rows_case { + %token_count = check.literal value(5) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<5x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<5xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<5x1024xf32>, tensor<5xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_077_ggml_set_rows_case> @qwen38_27b_udq4kxl_companion_tg_c5_077_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_078_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(5) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<5x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<5x256xf16> + %output = check.generate.fill value(1.0) : tensor<5x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<5x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<5x256xf16>, tensor<5x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_078_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_078_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_079_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<122880xi8> + %quantized_values = check.generate.fill value(0) : tensor<15360xi8> + %scales = check.generate.fill value(0) : tensor<3840xi8> + %sums = check.generate.fill value(0) : tensor<3840xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<122880xi8>, tensor<15360xi8>, tensor<3840xi8>, tensor<3840xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_079_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_079_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_080_ggml_binary_f32_case { + %element_count = check.literal value(25600) : index + %lhs = check.generate.fill value(2.0) : tensor<25600xf32> + %rhs = check.generate.fill value(3.0) : tensor<25600xf32> + %output = check.generate.fill value(0.0) : tensor<25600xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<25600xf32>, tensor<25600xf32>, tensor<25600xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_080_ggml_binary_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_080_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_081_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(5) : index + %input = check.generate.fill value(2.0) : tensor<5x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<5x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<5x5120xf32>, tensor<5120xf32>, tensor<5x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_081_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_081_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_082_ggml_get_rows_f32_case { + %token_count = check.literal value(5) : index + %row_count = check.literal value(5) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<5xi32> + %weight = check.generate.fill value(0.0) : tensor<5x5120xf32> + %output = check.generate.fill value(1.0) : tensor<5x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<5xi32>, tensor<5x5120xf32>, tensor<5x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_082_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_082_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_083_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case { + %token_count = check.literal value(5) : index + %input = check.generate.fill value(0) : tensor<102400xi8> + %weight = check.generate.fill value(0) : tensor<1042944000xi8> + %output = check.generate.fill value(0) : tensor<4966400xi8> + kernel.launch @ggml_mul_mat_q6_k_packed_token1_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<102400xi8>, tensor<1042944000xi8>, tensor<4966400xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_083_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_083_ggml_mul_mat_q6_k_packed_token1_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_084_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<102400xi8> + %rhs = check.generate.fill value(0) : tensor<102400xi8> + %output = check.generate.fill value(0) : tensor<204800xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<102400xi8>, tensor<102400xi8>, tensor<204800xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_084_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_084_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_085_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(5) : index + %input = check.generate.fill value(0) : tensor<204800xi8> + %weight = check.generate.fill value(0) : tensor<56115200xi8> + %output = check.generate.fill value(0) : tensor<102400xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<204800xi8>, tensor<56115200xi8>, tensor<102400xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_085_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_085_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_086_ggml_get_rows_f32_case { + %token_count = check.literal value(3) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<12xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<61440xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<12xi8>, tensor<715161600xi8>, tensor<61440xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_086_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_086_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_087_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(3) : index + %input = check.generate.fill value(0) : tensor<61440xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<61440xi8> + %quantized_values = check.generate.fill value(0) : tensor<7680xi8> + %scales = check.generate.fill value(0) : tensor<1920xi8> + %sums = check.generate.fill value(0) : tensor<1920xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<61440xi8>, tensor<20480xi8>, tensor<61440xi8>, tensor<7680xi8>, tensor<1920xi8>, tensor<1920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_087_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_087_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_088_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<61440xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + %quantized_values = check.generate.fill value(0) : tensor<7680xi8> + %scales = check.generate.fill value(0) : tensor<1920xi8> + %sums = check.generate.fill value(0) : tensor<1920xi8> + %partial = check.generate.fill value(0) : tensor<122880xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<61440xi8>, tensor<122880xi8>, tensor<7680xi8>, tensor<1920xi8>, tensor<1920xi8>, tensor<122880xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_088_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_088_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_089_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<168960xi8> + %second_weight = check.generate.fill value(0) : tensor<168960xi8> + %first_output = check.generate.fill value(0) : tensor<576xi8> + %second_output = check.generate.fill value(0) : tensor<576xi8> + %quantized_values = check.generate.fill value(0) : tensor<7680xi8> + %scales = check.generate.fill value(0) : tensor<1920xi8> + %sums = check.generate.fill value(0) : tensor<1920xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<168960xi8>, tensor<168960xi8>, tensor<576xi8>, tensor<576xi8>, tensor<7680xi8>, tensor<1920xi8>, tensor<1920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_089_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_089_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_090_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<61440xi8> + %output = check.generate.fill value(0) : tensor<73728xi8> + %quantized_values = check.generate.fill value(0) : tensor<7680xi8> + %scales = check.generate.fill value(0) : tensor<1920xi8> + %sums = check.generate.fill value(0) : tensor<1920xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<61440xi8>, tensor<73728xi8>, tensor<7680xi8>, tensor<1920xi8>, tensor<1920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_090_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_090_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_091_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<122880xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<122880xi8>, tensor<163840xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_091_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_091_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_092_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<576xi8> + %beta_raw = check.generate.fill value(0) : tensor<576xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<576xi8> + %beta_dst = check.generate.fill value(0) : tensor<576xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<576xi8>, tensor<576xi8>, tensor<192xi8>, tensor<192xi8>, tensor<576xi8>, tensor<576xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_092_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_092_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_093_llm_gated_delta_net_f32_wmma_head128_snapshot_case { + %q = check.generate.fill value(0) : tensor<90112xi8> + %k = check.generate.fill value(0) : tensor<90112xi8> + %v = check.generate.fill value(0) : tensor<106496xi8> + %g = check.generate.fill value(0) : tensor<576xi8> + %beta = check.generate.fill value(0) : tensor<576xi8> + %state_in = check.generate.fill value(0) : tensor<3145728xi8> + %snapshot_cache = check.generate.fill value(0) : tensor<28311552xi8> + %dst = check.generate.fill value(0) : tensor<73728xi8> + %rms_scales = check.generate.fill value(0) : tensor<576xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_snapshot[](%q, %k, %v, %g, %beta, %state_in, %snapshot_cache, %dst, %rms_scales) : [](tensor<90112xi8>, tensor<90112xi8>, tensor<106496xi8>, tensor<576xi8>, tensor<576xi8>, tensor<3145728xi8>, tensor<28311552xi8>, tensor<73728xi8>, tensor<576xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_093_llm_gated_delta_net_f32_wmma_head128_snapshot_case> @qwen38_27b_udq4kxl_companion_tg_c5_093_llm_gated_delta_net_f32_wmma_head128_snapshot + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_094_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(144) : index + %input = check.generate.fill value(0) : tensor<73728xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<73728xi8> + %output = check.generate.fill value(0) : tensor<73728xi8> + %quantized_values = check.generate.fill value(0) : tensor<9216xi8> + %scales = check.generate.fill value(0) : tensor<2304xi8> + %sums = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<73728xi8>, tensor<512xi8>, tensor<73728xi8>, tensor<73728xi8>, tensor<9216xi8>, tensor<2304xi8>, tensor<2304xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_094_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_094_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_095_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<73728xi8> + %output = check.generate.fill value(0) : tensor<61440xi8> + %quantized_values = check.generate.fill value(0) : tensor<9216xi8> + %scales = check.generate.fill value(0) : tensor<2304xi8> + %sums = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<73728xi8>, tensor<61440xi8>, tensor<9216xi8>, tensor<2304xi8>, tensor<2304xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_095_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_095_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_096_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(3) : index + %lhs = check.generate.fill value(0) : tensor<61440xi8> + %rhs = check.generate.fill value(0) : tensor<61440xi8> + %residual_output = check.generate.fill value(0) : tensor<61440xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<61440xi8> + %quantized_values = check.generate.fill value(0) : tensor<7680xi8> + %scales = check.generate.fill value(0) : tensor<1920xi8> + %sums = check.generate.fill value(0) : tensor<1920xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<61440xi8>, tensor<61440xi8>, tensor<61440xi8>, tensor<20480xi8>, tensor<61440xi8>, tensor<7680xi8>, tensor<1920xi8>, tensor<1920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_096_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_096_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_097_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<208896xi8> + %second_output = check.generate.fill value(0) : tensor<208896xi8> + %quantized_values = check.generate.fill value(0) : tensor<7680xi8> + %scales = check.generate.fill value(0) : tensor<1920xi8> + %sums = check.generate.fill value(0) : tensor<1920xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<208896xi8>, tensor<208896xi8>, tensor<7680xi8>, tensor<1920xi8>, tensor<1920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_097_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_097_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_098_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<208896xi8> + %rhs = check.generate.fill value(0) : tensor<208896xi8> + %output = check.generate.fill value(0) : tensor<208896xi8> + %quantized_values = check.generate.fill value(0) : tensor<26112xi8> + %scales = check.generate.fill value(0) : tensor<6528xi8> + %sums = check.generate.fill value(0) : tensor<6528xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<208896xi8>, tensor<208896xi8>, tensor<208896xi8>, tensor<26112xi8>, tensor<6528xi8>, tensor<6528xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_098_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_098_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_099_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<208896xi8> + %output = check.generate.fill value(0) : tensor<61440xi8> + %quantized_values = check.generate.fill value(0) : tensor<26112xi8> + %scales = check.generate.fill value(0) : tensor<6528xi8> + %sums = check.generate.fill value(0) : tensor<6528xi8> + %partial = check.generate.fill value(0) : tensor<61440xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<208896xi8>, tensor<61440xi8>, tensor<26112xi8>, tensor<6528xi8>, tensor<6528xi8>, tensor<61440xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_099_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_099_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_100_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<61440xi8> + %output = check.generate.fill value(0) : tensor<147456xi8> + %quantized_values = check.generate.fill value(0) : tensor<7680xi8> + %scales = check.generate.fill value(0) : tensor<1920xi8> + %sums = check.generate.fill value(0) : tensor<1920xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<61440xi8>, tensor<147456xi8>, tensor<7680xi8>, tensor<1920xi8>, tensor<1920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_100_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_100_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_101_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(3) : index + %input = check.generate.fill value(0) : tensor<61440xi8> + %weight = check.generate.fill value(0) : tensor<5611520xi8> + %output = check.generate.fill value(0) : tensor<12288xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<61440xi8>, tensor<5611520xi8>, tensor<12288xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_101_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_101_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_102_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<61440xi8> + %output = check.generate.fill value(0) : tensor<12288xi8> + %quantized_values = check.generate.fill value(0) : tensor<7680xi8> + %scales = check.generate.fill value(0) : tensor<1920xi8> + %sums = check.generate.fill value(0) : tensor<1920xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<61440xi8>, tensor<12288xi8>, tensor<7680xi8>, tensor<1920xi8>, tensor<1920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_102_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_102_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_103_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<146432xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<48xi8> + %output = check.generate.fill value(0) : tensor<73728xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<146432xi8>, tensor<1024xi8>, tensor<48xi8>, tensor<73728xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_103_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_103_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_104_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<12288xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<48xi8> + %output = check.generate.fill value(0) : tensor<12288xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<12288xi8>, tensor<1024xi8>, tensor<48xi8>, tensor<12288xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_104_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_104_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_105_ggml_set_rows_case { + %token_count = check.literal value(3) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<3x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<3xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<3x1024xf32>, tensor<3xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_105_ggml_set_rows_case> @qwen38_27b_udq4kxl_companion_tg_c5_105_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_106_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(3) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<3x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<3x256xf16> + %output = check.generate.fill value(1.0) : tensor<3x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<3x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<3x256xf16>, tensor<3x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_106_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_106_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_107_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<73728xi8> + %quantized_values = check.generate.fill value(0) : tensor<9216xi8> + %scales = check.generate.fill value(0) : tensor<2304xi8> + %sums = check.generate.fill value(0) : tensor<2304xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<73728xi8>, tensor<9216xi8>, tensor<2304xi8>, tensor<2304xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_107_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_107_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_108_ggml_binary_f32_case { + %element_count = check.literal value(15360) : index + %lhs = check.generate.fill value(2.0) : tensor<15360xf32> + %rhs = check.generate.fill value(3.0) : tensor<15360xf32> + %output = check.generate.fill value(0.0) : tensor<15360xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<15360xf32>, tensor<15360xf32>, tensor<15360xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_108_ggml_binary_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_108_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_109_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(3) : index + %input = check.generate.fill value(2.0) : tensor<3x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<3x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<3x5120xf32>, tensor<5120xf32>, tensor<3x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_109_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_109_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_110_ggml_get_rows_f32_case { + %token_count = check.literal value(3) : index + %row_count = check.literal value(3) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<3xi32> + %weight = check.generate.fill value(0.0) : tensor<3x5120xf32> + %output = check.generate.fill value(1.0) : tensor<3x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<3xi32>, tensor<3x5120xf32>, tensor<3x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_110_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_110_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_111_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case { + %token_count = check.literal value(3) : index + %input = check.generate.fill value(0) : tensor<61440xi8> + %weight = check.generate.fill value(0) : tensor<1042944000xi8> + %output = check.generate.fill value(0) : tensor<2979840xi8> + kernel.launch @ggml_mul_mat_q6_k_packed_token1_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<61440xi8>, tensor<1042944000xi8>, tensor<2979840xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_111_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_111_ggml_mul_mat_q6_k_packed_token1_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_112_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<61440xi8> + %rhs = check.generate.fill value(0) : tensor<61440xi8> + %output = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<61440xi8>, tensor<61440xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_112_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_112_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_113_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(3) : index + %input = check.generate.fill value(0) : tensor<122880xi8> + %weight = check.generate.fill value(0) : tensor<56115200xi8> + %output = check.generate.fill value(0) : tensor<61440xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<122880xi8>, tensor<56115200xi8>, tensor<61440xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_113_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_113_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_114_ggml_get_rows_f32_case { + %token_count = check.literal value(16) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<64xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<64xi8>, tensor<715161600xi8>, tensor<327680xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_114_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_114_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_115_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(16) : index + %input = check.generate.fill value(0) : tensor<327680xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<327680xi8>, tensor<20480xi8>, tensor<327680xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_115_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_115_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_116_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<327680xi8> + %output = check.generate.fill value(0) : tensor<655360xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + %partial = check.generate.fill value(0) : tensor<655360xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<327680xi8>, tensor<655360xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>, tensor<655360xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_116_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_116_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_117_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<168960xi8> + %second_weight = check.generate.fill value(0) : tensor<168960xi8> + %first_output = check.generate.fill value(0) : tensor<3072xi8> + %second_output = check.generate.fill value(0) : tensor<3072xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<168960xi8>, tensor<168960xi8>, tensor<3072xi8>, tensor<3072xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_117_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_117_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_118_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<327680xi8> + %output = check.generate.fill value(0) : tensor<393216xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<327680xi8>, tensor<393216xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_118_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_118_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_119_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<655360xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<655360xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<655360xi8>, tensor<163840xi8>, tensor<655360xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_119_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_119_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_120_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<3072xi8> + %beta_raw = check.generate.fill value(0) : tensor<3072xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<3072xi8> + %beta_dst = check.generate.fill value(0) : tensor<3072xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<3072xi8>, tensor<3072xi8>, tensor<192xi8>, tensor<192xi8>, tensor<3072xi8>, tensor<3072xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_120_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_120_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_121_llm_gated_delta_net_f32_wmma_head128_case { + %q = check.generate.fill value(0) : tensor<417792xi8> + %k = check.generate.fill value(0) : tensor<417792xi8> + %v = check.generate.fill value(0) : tensor<434176xi8> + %g = check.generate.fill value(0) : tensor<2112xi8> + %beta = check.generate.fill value(0) : tensor<2112xi8> + %state_in = check.generate.fill value(0) : tensor<3145728xi8> + %dst = check.generate.fill value(0) : tensor<3416064xi8> + %rms_scales = check.generate.fill value(0) : tensor<2112xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128[](%q, %k, %v, %g, %beta, %state_in, %dst, %rms_scales) : [](tensor<417792xi8>, tensor<417792xi8>, tensor<434176xi8>, tensor<2112xi8>, tensor<2112xi8>, tensor<3145728xi8>, tensor<3416064xi8>, tensor<2112xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_121_llm_gated_delta_net_f32_wmma_head128_case> @qwen38_27b_udq4kxl_companion_tg_c5_121_llm_gated_delta_net_f32_wmma_head128 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_122_ggml_copy_f32_case { + %element_count = check.literal value(786432) : index + %source = check.generate.fill value(0) : tensor<3145728xi8> + %output = check.generate.fill value(0) : tensor<3145728xi8> + kernel.launch @ggml_copy_f32[%element_count](%element_count, %source, %output) : [index](index, tensor<3145728xi8>, tensor<3145728xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_122_ggml_copy_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_122_ggml_copy_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_123_llm_gated_delta_net_f32_wmma_head128_inplace_case { + %q = check.generate.fill value(0) : tensor<8192xi8> + %k = check.generate.fill value(0) : tensor<8192xi8> + %v = check.generate.fill value(0) : tensor<24576xi8> + %g = check.generate.fill value(0) : tensor<192xi8> + %beta = check.generate.fill value(0) : tensor<192xi8> + %state_inout = check.generate.fill value(0) : tensor<3145728xi8> + %dst = check.generate.fill value(0) : tensor<24576xi8> + %rms_scales = check.generate.fill value(0) : tensor<192xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_inplace[](%q, %k, %v, %g, %beta, %state_inout, %dst, %rms_scales) : [](tensor<8192xi8>, tensor<8192xi8>, tensor<24576xi8>, tensor<192xi8>, tensor<192xi8>, tensor<3145728xi8>, tensor<24576xi8>, tensor<192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_123_llm_gated_delta_net_f32_wmma_head128_inplace_case> @qwen38_27b_udq4kxl_companion_tg_c5_123_llm_gated_delta_net_f32_wmma_head128_inplace + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_124_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(768) : index + %input = check.generate.fill value(0) : tensor<393216xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<393216xi8> + %output = check.generate.fill value(0) : tensor<393216xi8> + %quantized_values = check.generate.fill value(0) : tensor<49152xi8> + %scales = check.generate.fill value(0) : tensor<12288xi8> + %sums = check.generate.fill value(0) : tensor<12288xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<393216xi8>, tensor<512xi8>, tensor<393216xi8>, tensor<393216xi8>, tensor<49152xi8>, tensor<12288xi8>, tensor<12288xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_124_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_124_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_125_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<393216xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + %quantized_values = check.generate.fill value(0) : tensor<49152xi8> + %scales = check.generate.fill value(0) : tensor<12288xi8> + %sums = check.generate.fill value(0) : tensor<12288xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<393216xi8>, tensor<327680xi8>, tensor<49152xi8>, tensor<12288xi8>, tensor<12288xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_125_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_125_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_126_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(16) : index + %lhs = check.generate.fill value(0) : tensor<327680xi8> + %rhs = check.generate.fill value(0) : tensor<327680xi8> + %residual_output = check.generate.fill value(0) : tensor<327680xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<327680xi8>, tensor<327680xi8>, tensor<327680xi8>, tensor<20480xi8>, tensor<327680xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_126_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_126_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_127_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<1114112xi8> + %second_output = check.generate.fill value(0) : tensor<1114112xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<1114112xi8>, tensor<1114112xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_127_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_127_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_128_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<1114112xi8> + %rhs = check.generate.fill value(0) : tensor<1114112xi8> + %output = check.generate.fill value(0) : tensor<1114112xi8> + %quantized_values = check.generate.fill value(0) : tensor<139264xi8> + %scales = check.generate.fill value(0) : tensor<34816xi8> + %sums = check.generate.fill value(0) : tensor<34816xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<1114112xi8>, tensor<1114112xi8>, tensor<1114112xi8>, tensor<139264xi8>, tensor<34816xi8>, tensor<34816xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_128_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_128_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_129_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<1114112xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + %quantized_values = check.generate.fill value(0) : tensor<139264xi8> + %scales = check.generate.fill value(0) : tensor<34816xi8> + %sums = check.generate.fill value(0) : tensor<34816xi8> + %partial = check.generate.fill value(0) : tensor<327680xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<1114112xi8>, tensor<327680xi8>, tensor<139264xi8>, tensor<34816xi8>, tensor<34816xi8>, tensor<327680xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_129_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_129_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_130_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<327680xi8> + %output = check.generate.fill value(0) : tensor<786432xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<327680xi8>, tensor<786432xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_130_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_130_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_131_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(16) : index + %input = check.generate.fill value(0) : tensor<327680xi8> + %weight = check.generate.fill value(0) : tensor<5611520xi8> + %output = check.generate.fill value(0) : tensor<65536xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<327680xi8>, tensor<5611520xi8>, tensor<65536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_131_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_131_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_132_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<327680xi8> + %output = check.generate.fill value(0) : tensor<65536xi8> + %quantized_values = check.generate.fill value(0) : tensor<40960xi8> + %scales = check.generate.fill value(0) : tensor<10240xi8> + %sums = check.generate.fill value(0) : tensor<10240xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<327680xi8>, tensor<65536xi8>, tensor<40960xi8>, tensor<10240xi8>, tensor<10240xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_132_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_132_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_133_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<785408xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<256xi8> + %output = check.generate.fill value(0) : tensor<393216xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<785408xi8>, tensor<1024xi8>, tensor<256xi8>, tensor<393216xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_133_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_133_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_134_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<65536xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<256xi8> + %output = check.generate.fill value(0) : tensor<65536xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<65536xi8>, tensor<1024xi8>, tensor<256xi8>, tensor<65536xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_134_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_134_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_135_ggml_set_rows_case { + %token_count = check.literal value(16) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<16x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<16xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<16x1024xf32>, tensor<16xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_135_ggml_set_rows_case> @qwen38_27b_udq4kxl_companion_tg_c5_135_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_136_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(16) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<16x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<16x256xf16> + %output = check.generate.fill value(1.0) : tensor<16x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<16x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<16x256xf16>, tensor<16x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_136_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_136_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_137_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<393216xi8> + %quantized_values = check.generate.fill value(0) : tensor<49152xi8> + %scales = check.generate.fill value(0) : tensor<12288xi8> + %sums = check.generate.fill value(0) : tensor<12288xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<393216xi8>, tensor<49152xi8>, tensor<12288xi8>, tensor<12288xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_137_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_137_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_138_ggml_binary_f32_case { + %element_count = check.literal value(81920) : index + %lhs = check.generate.fill value(2.0) : tensor<81920xf32> + %rhs = check.generate.fill value(3.0) : tensor<81920xf32> + %output = check.generate.fill value(0.0) : tensor<81920xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<81920xf32>, tensor<81920xf32>, tensor<81920xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_138_ggml_binary_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_138_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_139_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(16) : index + %input = check.generate.fill value(2.0) : tensor<16x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<16x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<16x5120xf32>, tensor<5120xf32>, tensor<16x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_139_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_139_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_140_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<327680xi8> + %rhs = check.generate.fill value(0) : tensor<327680xi8> + %output = check.generate.fill value(0) : tensor<655360xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<327680xi8>, tensor<327680xi8>, tensor<655360xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_140_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_140_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_141_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(16) : index + %input = check.generate.fill value(0) : tensor<655360xi8> + %weight = check.generate.fill value(0) : tensor<56115200xi8> + %output = check.generate.fill value(0) : tensor<327680xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<655360xi8>, tensor<56115200xi8>, tensor<327680xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_141_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_141_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_142_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(4) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<4x5120xf32> + %output = check.generate.fill value(1.0) : tensor<1x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<4x5120xf32>, tensor<1x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_142_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_142_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_143_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<81920xi8> + %rhs = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<163840xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<81920xi8>, tensor<81920xi8>, tensor<163840xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_143_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_143_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_144_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(0) : tensor<163840xi8> + %weight = check.generate.fill value(0) : tensor<56115200xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<163840xi8>, tensor<56115200xi8>, tensor<81920xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_144_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_144_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_145_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<4xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<4xi8>, tensor<715161600xi8>, tensor<20480xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_145_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_145_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_146_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(2.0) : tensor<1x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<1x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<1x5120xf32>, tensor<5120xf32>, tensor<1x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_146_ggml_rmsnorm_binary_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_146_ggml_rmsnorm_binary_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_147_ggml_concat_dim0_f32_case { + %lhs = check.generate.fill value(0) : tensor<20480xi8> + %rhs = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + kernel.launch @ggml_concat_dim0_f32[](%lhs, %rhs, %output) : [](tensor<20480xi8>, tensor<20480xi8>, tensor<40960xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_147_ggml_concat_dim0_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_147_ggml_concat_dim0_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_148_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(10240) : index + %output_size = check.literal value(5120) : index + %input = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<43008000xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<40960xi8>, tensor<43008000xi8>, tensor<20480xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_148_ggml_mul_mat_f32_f32_decode_wave64_case> @qwen38_27b_udq4kxl_companion_tg_c5_148_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_149_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<20480xi8>, tensor<20480xi8>, tensor<20480xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_149_ggml_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_149_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_150_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<49152xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<20480xi8>, tensor<49152xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_150_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_150_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_151_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(5120) : index + %output_size = check.literal value(1024) : index + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<4300800xi8> + %output = check.generate.fill value(0) : tensor<4096xi8> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<20480xi8>, tensor<4300800xi8>, tensor<4096xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_151_ggml_mul_mat_f32_f32_decode_wave64_case> @qwen38_27b_udq4kxl_companion_tg_c5_151_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_152_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<4096xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<20480xi8>, tensor<4096xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_152_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_152_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_153_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<48128xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<16xi8> + %output = check.generate.fill value(0) : tensor<24576xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<48128xi8>, tensor<1024xi8>, tensor<16xi8>, tensor<24576xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_153_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_153_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_154_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<4096xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<16xi8> + %output = check.generate.fill value(0) : tensor<4096xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<4096xi8>, tensor<1024xi8>, tensor<16xi8>, tensor<4096xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_154_ggml_rmsnorm_mul_rope_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_154_ggml_rmsnorm_mul_rope_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_155_ggml_set_rows_case { + %token_count = check.literal value(1) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<1x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<1xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<1x1024xf32>, tensor<1xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_155_ggml_set_rows_case> @qwen38_27b_udq4kxl_companion_tg_c5_155_ggml_set_rows + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_156_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(1) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<1x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<1x256xf16> + %output = check.generate.fill value(1.0) : tensor<1x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<1x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<1x256xf16>, tensor<1x24x256xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_156_ggml_flash_attention_f32_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_156_ggml_flash_attention_f32_f16_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_157_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<24576xi8> + %quantized_values = check.generate.fill value(0) : tensor<3072xi8> + %scales = check.generate.fill value(0) : tensor<768xi8> + %sums = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<24576xi8>, tensor<3072xi8>, tensor<768xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_157_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_157_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_158_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<24576xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<3072xi8> + %scales = check.generate.fill value(0) : tensor<768xi8> + %sums = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<24576xi8>, tensor<20480xi8>, tensor<3072xi8>, tensor<768xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_158_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_158_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_159_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(1) : index + %lhs = check.generate.fill value(0) : tensor<20480xi8> + %rhs = check.generate.fill value(0) : tensor<20480xi8> + %residual_output = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<20480xi8>, tensor<20480xi8>, tensor<20480xi8>, tensor<20480xi8>, tensor<20480xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_159_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_159_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_160_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<69632xi8> + %second_output = check.generate.fill value(0) : tensor<69632xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<69632xi8>, tensor<69632xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_160_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_160_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_161_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<69632xi8> + %rhs = check.generate.fill value(0) : tensor<69632xi8> + %output = check.generate.fill value(0) : tensor<69632xi8> + %quantized_values = check.generate.fill value(0) : tensor<8704xi8> + %scales = check.generate.fill value(0) : tensor<2176xi8> + %sums = check.generate.fill value(0) : tensor<2176xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<69632xi8>, tensor<69632xi8>, tensor<69632xi8>, tensor<8704xi8>, tensor<2176xi8>, tensor<2176xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_161_ggml_binary_swiglu_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_161_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_162_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<69632xi8> + %output = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<8704xi8> + %scales = check.generate.fill value(0) : tensor<2176xi8> + %sums = check.generate.fill value(0) : tensor<2176xi8> + %partial = check.generate.fill value(0) : tensor<20480xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<69632xi8>, tensor<20480xi8>, tensor<8704xi8>, tensor<2176xi8>, tensor<2176xi8>, tensor<20480xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_162_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_162_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_163_ggml_binary_f32_case { + %element_count = check.literal value(5120) : index + %lhs = check.generate.fill value(2.0) : tensor<5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<5120xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<5120xf32>, tensor<5120xf32>, tensor<5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_163_ggml_binary_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_163_ggml_binary_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_164_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(1) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<1x5120xf32> + %output = check.generate.fill value(1.0) : tensor<1x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<1x5120xf32>, tensor<1x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_164_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_164_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_165_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<20480xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<20480xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_165_ggml_quantize_f32_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_165_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_166_ggml_select_symmetric_i4_k32_groups_case { + %input = check.generate.fill value(0) : tensor<20480xi8> + %selected_mask = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_select_symmetric_i4_k32_groups[](%input, %selected_mask) : [](tensor<20480xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_166_ggml_select_symmetric_i4_k32_groups_case> @qwen38_27b_udq4kxl_companion_tg_c5_166_ggml_select_symmetric_i4_k32_groups + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_167_ggml_mul_mat_q6_k_symmetric_i2_scan_token1_case { + %weight = check.generate.fill value(0) : tensor<1380659200xi8> + %output = check.generate.fill value(0) : tensor<993280xi8> + %qact = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %selected_mask = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_mul_mat_q6_k_symmetric_i2_scan_token1[](%weight, %output, %qact, %scales, %selected_mask) : [](tensor<1380659200xi8>, tensor<993280xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_167_ggml_mul_mat_q6_k_symmetric_i2_scan_token1_case> @qwen38_27b_udq4kxl_companion_tg_c5_167_ggml_mul_mat_q6_k_symmetric_i2_scan_token1 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_168_ggml_top_k8_f32_partitions_register_case { + %element_count = check.literal value(248320) : index + %values = check.generate.fill value(0) : tensor<993280xi8> + %partial_values = check.generate.fill value(0) : tensor<4096xi8> + %partial_ids = check.generate.fill value(0) : tensor<4096xi8> + kernel.launch @ggml_top_k8_f32_partitions_register[%element_count](%element_count, %values, %partial_values, %partial_ids) : [index](index, tensor<993280xi8>, tensor<4096xi8>, tensor<4096xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_168_ggml_top_k8_f32_partitions_register_case> @qwen38_27b_udq4kxl_companion_tg_c5_168_ggml_top_k8_f32_partitions_register + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_169_ggml_top_k128_f32_reduce_gather_register_case { + %element_count = check.literal value(248320) : index + %partial_values = check.generate.fill value(0) : tensor<4096xi8> + %partial_ids = check.generate.fill value(0) : tensor<4096xi8> + %candidate_output = check.generate.fill value(0) : tensor<512xi8> + %value_output = check.generate.fill value(0) : tensor<512xi8> + kernel.launch @ggml_top_k128_f32_reduce_gather_register[%element_count](%element_count, %partial_values, %partial_ids, %candidate_output, %value_output) : [index](index, tensor<4096xi8>, tensor<4096xi8>, tensor<512xi8>, tensor<512xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_169_ggml_top_k128_f32_reduce_gather_register_case> @qwen38_27b_udq4kxl_companion_tg_c5_169_ggml_top_k128_f32_reduce_gather_register + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_170_ggml_fill_negative_f32_case { + %element_count = check.literal value(248320) : index + %output = check.generate.fill value(0) : tensor<993280xi8> + kernel.launch @ggml_fill_negative_f32[%element_count](%element_count, %output) : [index](index, tensor<993280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_170_ggml_fill_negative_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_170_ggml_fill_negative_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_171_ggml_mul_mat_q6_k_packed_selected_refine_token1_case { + %token_count = check.literal value(1) : index + %candidate_count = check.literal value(64) : index + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<1380659200xi8> + %candidates = check.generate.fill value(0) : tensor<256xi8> + %exact_output = check.generate.fill value(0) : tensor<256xi8> + %output = check.generate.fill value(0) : tensor<993280xi8> + kernel.launch @ggml_mul_mat_q6_k_packed_selected_refine_token1[%token_count, %candidate_count](%token_count, %candidate_count, %input, %weight, %candidates, %exact_output, %output) : [index, index](index, index, tensor<20480xi8>, tensor<1380659200xi8>, tensor<256xi8>, tensor<256xi8>, tensor<993280xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_171_ggml_mul_mat_q6_k_packed_selected_refine_token1_case> @qwen38_27b_udq4kxl_companion_tg_c5_171_ggml_mul_mat_q6_k_packed_selected_refine_token1 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_172_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + %partial = check.generate.fill value(0) : tensor<40960xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<20480xi8>, tensor<40960xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>, tensor<40960xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_172_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_172_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_173_ggml_mul_mat_f32_f32_decode_wave64_case { + %token_count = check.literal value(1) : index + %input_size = check.literal value(5120) : index + %output_size = check.literal value(48) : index + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<138240xi8> + %output = check.generate.fill value(0) : tensor<192xi8> + kernel.launch @ggml_mul_mat_f32_f32_decode_wave64[%token_count, %input_size, %output_size](%token_count, %input_size, %output_size, %input, %weight, %output) : [index, index, index](index, index, index, tensor<20480xi8>, tensor<138240xi8>, tensor<192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_173_ggml_mul_mat_f32_f32_decode_wave64_case> @qwen38_27b_udq4kxl_companion_tg_c5_173_ggml_mul_mat_f32_f32_decode_wave64 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_174_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<24576xi8> + %quantized_values = check.generate.fill value(0) : tensor<2560xi8> + %scales = check.generate.fill value(0) : tensor<640xi8> + %sums = check.generate.fill value(0) : tensor<640xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<20480xi8>, tensor<24576xi8>, tensor<2560xi8>, tensor<640xi8>, tensor<640xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_174_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_174_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_175_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<40960xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<40960xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<40960xi8>, tensor<163840xi8>, tensor<40960xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_175_llm_ssm_conv_dconv4_silu_rollback_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_175_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_176_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<192xi8> + %beta_raw = check.generate.fill value(0) : tensor<192xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<192xi8> + %beta_dst = check.generate.fill value(0) : tensor<192xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<192xi8>, tensor<192xi8>, tensor<192xi8>, tensor<192xi8>, tensor<192xi8>, tensor<192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_176_llm_gated_delta_net_projection_epilogue_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_176_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_177_llm_gated_delta_net_f32_wmma_head128_inplace_case { + %q = check.generate.fill value(0) : tensor<8192xi8> + %k = check.generate.fill value(0) : tensor<8192xi8> + %v = check.generate.fill value(0) : tensor<24576xi8> + %g = check.generate.fill value(0) : tensor<192xi8> + %beta = check.generate.fill value(0) : tensor<192xi8> + %state_inout = check.generate.fill value(0) : tensor<3145728xi8> + %dst = check.generate.fill value(0) : tensor<24576xi8> + %rms_scales = check.generate.fill value(0) : tensor<192xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_inplace[](%q, %k, %v, %g, %beta, %state_inout, %dst, %rms_scales) : [](tensor<8192xi8>, tensor<8192xi8>, tensor<24576xi8>, tensor<192xi8>, tensor<192xi8>, tensor<3145728xi8>, tensor<24576xi8>, tensor<192xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_177_llm_gated_delta_net_f32_wmma_head128_inplace_case> @qwen38_27b_udq4kxl_companion_tg_c5_177_llm_gated_delta_net_f32_wmma_head128_inplace + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_178_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(48) : index + %input = check.generate.fill value(0) : tensor<24576xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<24576xi8> + %output = check.generate.fill value(0) : tensor<24576xi8> + %quantized_values = check.generate.fill value(0) : tensor<3072xi8> + %scales = check.generate.fill value(0) : tensor<768xi8> + %sums = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<24576xi8>, tensor<512xi8>, tensor<24576xi8>, tensor<24576xi8>, tensor<3072xi8>, tensor<768xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_178_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @qwen38_27b_udq4kxl_companion_tg_c5_178_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_179_ggml_get_rows_f32_case { + %token_count = check.literal value(2) : index + %row_count = check.literal value(2) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<2xi32> + %weight = check.generate.fill value(0.0) : tensor<2x5120xf32> + %output = check.generate.fill value(1.0) : tensor<2x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<2xi32>, tensor<2x5120xf32>, tensor<2x5120xf32>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_179_ggml_get_rows_f32_case> @qwen38_27b_udq4kxl_companion_tg_c5_179_ggml_get_rows_f32 + +check.case public @qwen38_27b_udq4kxl_companion_tg_c5_180_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case { + %token_count = check.literal value(2) : index + %input = check.generate.fill value(0) : tensor<40960xi8> + %weight = check.generate.fill value(0) : tensor<1042944000xi8> + %output = check.generate.fill value(0) : tensor<1986560xi8> + kernel.launch @ggml_mul_mat_q6_k_packed_token1_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<40960xi8>, tensor<1042944000xi8>, tensor<1986560xi8>) + check.return +} + +check.benchmark<@qwen38_27b_udq4kxl_companion_tg_c5_180_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case> @qwen38_27b_udq4kxl_companion_tg_c5_180_ggml_mul_mat_q6_k_packed_token1_f16_wmma diff --git a/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_mtp4_program4.tg_c4.json b/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_mtp4_program4.tg_c4.json new file mode 100644 index 000000000000..57e4131da09b --- /dev/null +++ b/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_mtp4_program4.tg_c4.json @@ -0,0 +1,932 @@ +{ + "command_count": 965, + "dispatch_count": 28, + "dispatches": [ + { + "benchmark": "@v2_mtp4_program4_tg_c4_000_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "4", + "ggml.get_rows_f32.weight_format": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 248320, + "token_count": 4 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_001_ggml_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 4 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "10240", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_003_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "48", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_004_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_005_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "30720", + "ggml.get_rows_f32.token_capacity": "1", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "hidden_size": 30720, + "row_count": 20, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_006_llm_ssm_conv_dconv4_silu_rollback_f32", + "compile_parameters": { + "llm.ssm_conv.rollback.cache_count": "5", + "llm.ssm_conv.rollback.d_inner": "10240", + "llm.ssm_conv.rollback.n_s": "1", + "llm.ssm_conv.rollback.n_t": "4", + "llm.ssm_conv.rollback.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_ssm_conv_dconv4_silu_rollback_f32", + "library_sources": [], + "primary_sources": [ + "ops/ssm_conv_f32.loom" + ], + "sources": [ + "ops/ssm_conv_f32.loom" + ], + "symbol": "llm_ssm_conv_dconv4_silu_rollback_f32", + "workload_parameters": [] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_007_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "786432", + "ggml.get_rows_f32.token_capacity": "1", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "hidden_size": 786432, + "row_count": 20, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_008_llm_gated_delta_net_projection_epilogue_f32", + "compile_parameters": { + "llm.gated_delta_net.epilogue_element_count": "192", + "llm.gated_delta_net.epilogue_head_count": "48", + "llm.gated_delta_net.epilogue_workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_projection_epilogue_f32", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_projection_epilogue_f32", + "workload_parameters": [] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_009_llm_gated_delta_net_f32_wmma_head128_snapshot", + "compile_parameters": { + "llm.gated_delta_net.head_count": "48", + "llm.gated_delta_net.head_width": "128", + "llm.gated_delta_net.l2_epsilon": "9.99999997e-07", + "llm.gated_delta_net.qk_stride1": "128", + "llm.gated_delta_net.qk_stride2": "10240", + "llm.gated_delta_net.qk_stride3": "40960", + "llm.gated_delta_net.query_head_count": "16", + "llm.gated_delta_net.query_sequence_ratio": "1", + "llm.gated_delta_net.scalar_stride1": "1", + "llm.gated_delta_net.scalar_stride2": "48", + "llm.gated_delta_net.scalar_stride3": "192", + "llm.gated_delta_net.sequence_count": "1", + "llm.gated_delta_net.snapshot_stride": "3145728", + "llm.gated_delta_net.token_count": "4", + "llm.gated_delta_net.value_stride1": "128", + "llm.gated_delta_net.value_stride2": "10240", + "llm.gated_delta_net.value_stride3": "40960", + "llm.gated_delta_net.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": {}, + "kernel": "loom_libs:llm_gated_delta_net_f32_wmma_head128_snapshot", + "library_sources": [], + "primary_sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "sources": [ + "ops/gated_delta_net_f32_wmma.loom" + ], + "symbol": "llm_gated_delta_net_f32_wmma_head128_snapshot", + "workload_parameters": [] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_010_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "compile_parameters": { + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.hidden_size": "128", + "ggml.rmsnorm_gate_silu_mul_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 48, + "integer_parameters": { + "token_count": 192 + }, + "kernel": "loom_libs:ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_011_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "6144", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 64, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_012_ggml_add_rmsnorm_binary_symmetric_i4_k32", + "compile_parameters": { + "ggml.add_rmsnorm_binary_symmetric_i4.hidden_size": "5120", + "ggml.add_rmsnorm_binary_symmetric_i4.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 127, + "integer_parameters": { + "token_count": 4 + }, + "kernel": "loom_libs:ggml_add_rmsnorm_binary_symmetric_i4_k32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_add_rmsnorm_binary_symmetric_i4_k32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_013_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 64, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_014_ggml_binary_swiglu_symmetric_i4_k32", + "compile_parameters": { + "ggml.binary_swiglu_symmetric_i4.input_size": "17408", + "ggml.binary_swiglu_symmetric_i4.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 64, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_binary_swiglu_symmetric_i4_k32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_swiglu_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_015_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "17408", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 64, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_016_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "12288", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 16, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_017_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "1024" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 16, + "integer_parameters": { + "token_count": 4 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_i8_prepacked_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_018_ggml_mul_mat_symmetric_i4_lowrow_wmma", + "compile_parameters": { + "ggml.mul_mat.symmetric_i4.lowrow.input_size": "5120", + "ggml.mul_mat.symmetric_i4.lowrow.output_size": "1024", + "ggml.mul_mat.symmetric_i4.lowrow.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 16, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_mul_mat_symmetric_i4_lowrow_wmma", + "workload_parameters": [] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_019_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "512", + "ggml.rmsnorm_mul_rope.input_stride2": "12288", + "ggml.rmsnorm_mul_rope.input_stride3": "49152", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "24", + "ggml.rmsnorm_mul_rope.ne2": "4", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "6144", + "ggml.rmsnorm_mul_rope.output_stride3": "24576", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 16, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_020_ggml_rmsnorm_mul_rope_f32", + "compile_parameters": { + "ggml.rmsnorm_mul_rope.attn_factor": "1", + "ggml.rmsnorm_mul_rope.epsilon": "9.99999997e-07", + "ggml.rmsnorm_mul_rope.freq_base": "10000000", + "ggml.rmsnorm_mul_rope.freq_scale": "1", + "ggml.rmsnorm_mul_rope.hidden_size": "256", + "ggml.rmsnorm_mul_rope.input_stride1": "256", + "ggml.rmsnorm_mul_rope.input_stride2": "1024", + "ggml.rmsnorm_mul_rope.input_stride3": "4096", + "ggml.rmsnorm_mul_rope.mode": "40", + "ggml.rmsnorm_mul_rope.n_dims": "64", + "ggml.rmsnorm_mul_rope.ne1": "4", + "ggml.rmsnorm_mul_rope.ne2": "4", + "ggml.rmsnorm_mul_rope.ne3": "1", + "ggml.rmsnorm_mul_rope.output_stride1": "256", + "ggml.rmsnorm_mul_rope.output_stride2": "1024", + "ggml.rmsnorm_mul_rope.output_stride3": "4096", + "ggml.rmsnorm_mul_rope.section0": "11", + "ggml.rmsnorm_mul_rope.section1": "11", + "ggml.rmsnorm_mul_rope.section2": "10", + "ggml.rmsnorm_mul_rope.section3": "0", + "ggml.rmsnorm_mul_rope.workgroup_size": "256" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 16, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_rmsnorm_mul_rope_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom" + ], + "primary_sources": [ + "ops/rmsnorm_f32.loom" + ], + "sources": [ + "ops/rmsnorm_f32.loom", + "motifs/rmsnorm_f32.loom" + ], + "symbol": "ggml_rmsnorm_mul_rope_f32", + "workload_parameters": [] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_021_ggml_set_rows", + "compile_parameters": { + "ggml.set_rows.hidden_capacity": "1024", + "ggml.set_rows.input_format": "32", + "ggml.set_rows.output_format": "16", + "ggml.set_rows.token_capacity": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 32, + "integer_parameters": { + "cache_row_count": 262144, + "hidden_size": 1024, + "token_count": 4 + }, + "kernel": "loom_libs:ggml_set_rows", + "library_sources": [], + "primary_sources": [ + "ops/set_rows.loom" + ], + "sources": [ + "ops/set_rows.loom" + ], + "symbol": "ggml_set_rows", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "cache_row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_022_ggml_flash_attention_f32_f16_wmma", + "compile_parameters": { + "ggml.flash_attention.apply_gate": "1", + "ggml.flash_attention.gate_stride_head": "512", + "ggml.flash_attention.gate_stride_token": "12288", + "ggml.flash_attention.head_size": "256", + "ggml.flash_attention.key_value_head_count": "4", + "ggml.flash_attention.query_head_count": "24" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 16, + "integer_parameters": { + "key_value_token_count": 256, + "query_token_count": 4 + }, + "kernel": "loom_libs:ggml_flash_attention_f32_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "sources": [ + "ops/flash_attention_f32_f16_wmma.loom" + ], + "symbol": "ggml_flash_attention_f32_f16_wmma", + "workload_parameters": [ + { + "name": "query_token_count", + "type": "index" + }, + { + "name": "key_value_token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_023_ggml_quantize_f32_symmetric_i4_k32", + "compile_parameters": { + "ggml.quantize_symmetric_i4_k32.input_size": "6144", + "ggml.quantize_symmetric_i4_k32.token_count": "4" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 16, + "integer_parameters": {}, + "kernel": "loom_libs:ggml_quantize_f32_symmetric_i4_k32", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "sources": [ + "ops/mul_mat_swiglu_symmetric_i4_wmma.loom" + ], + "symbol": "ggml_quantize_f32_symmetric_i4_k32", + "workload_parameters": [] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_024_ggml_binary_f32", + "compile_parameters": { + "ggml.binary_f32.op": "0" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "element_count": 20480 + }, + "kernel": "loom_libs:ggml_binary_f32", + "library_sources": [ + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/binary_f32.loom" + ], + "sources": [ + "ops/binary_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_binary_f32", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_025_ggml_rmsnorm_binary_f32", + "compile_parameters": { + "ggml.rmsnorm_binary_f32.hidden_size": "5120", + "ggml.rmsnorm_binary_f32.op": "2", + "ggml.rmsnorm_binary_f32.rms_epsilon": "9.99999997e-07" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 4 + }, + "kernel": "loom_libs:ggml_rmsnorm_binary_f32", + "library_sources": [ + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "primary_sources": [ + "ops/rmsnorm_binary_f32.loom" + ], + "sources": [ + "ops/rmsnorm_binary_f32.loom", + "motifs/rmsnorm_f32.loom", + "motifs/binary_f32_apply.loom" + ], + "symbol": "ggml_rmsnorm_binary_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_026_ggml_get_rows_f32", + "compile_parameters": { + "ggml.get_rows_f32.hidden_capacity": "5120", + "ggml.get_rows_f32.token_capacity": "1", + "ggml.get_rows_f32.weight_format": "32" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "hidden_size": 5120, + "row_count": 4, + "token_count": 1 + }, + "kernel": "loom_libs:ggml_get_rows_f32", + "library_sources": [ + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "primary_sources": [ + "ops/get_rows_f32.loom" + ], + "sources": [ + "ops/get_rows_f32.loom", + "motifs/dequant.loom", + "motifs/publish_f32.loom", + "motifs/quantize_q8_1_x4.loom", + "motifs/q4_k_f16.loom", + "motifs/q6_k_f16.loom", + "motifs/q8_0_f16.loom", + "motifs/q8_1_f16.loom", + "motifs/f16_f16.loom", + "motifs/bf16_f16.loom", + "motifs/f32_f16.loom" + ], + "symbol": "ggml_get_rows_f32", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "hidden_size", + "type": "index" + } + ] + }, + { + "benchmark": "@v2_mtp4_program4_tg_c4_027_ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "compile_parameters": { + "ggml.mul_mat_q6_k_packed.input_size": "5120", + "ggml.mul_mat_q6_k_packed.output_accumulation": "0", + "ggml.mul_mat_q6_k_packed.output_size": "248320" + }, + "corpus_dir": "kernel-corpus/kernels/loom-libs", + "count": 1, + "integer_parameters": { + "token_count": 1 + }, + "kernel": "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "library_sources": [], + "primary_sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "sources": [ + "ops/mul_mat_q6_k_packed_f16_wmma.loom" + ], + "symbol": "ggml_mul_mat_q6_k_packed_token1_f16_wmma", + "workload_parameters": [ + { + "name": "token_count", + "type": "index" + } + ] + } + ], + "generated_count": 28, + "kernel_counts": { + "loom_libs:ggml_add_rmsnorm_binary_symmetric_i4_k32": 127, + "loom_libs:ggml_binary_f32": 1, + "loom_libs:ggml_binary_swiglu_symmetric_i4_k32": 64, + "loom_libs:ggml_flash_attention_f32_f16_wmma": 16, + "loom_libs:ggml_get_rows_f32": 98, + "loom_libs:ggml_mul_mat_q6_k_i8_prepacked_f16_wmma": 16, + "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma": 1, + "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma": 112, + "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma": 112, + "loom_libs:ggml_mul_mat_symmetric_i4_lowrow_wmma": 144, + "loom_libs:ggml_quantize_f32_symmetric_i4_k32": 16, + "loom_libs:ggml_rmsnorm_binary_f32": 1, + "loom_libs:ggml_rmsnorm_binary_symmetric_i4_k32": 1, + "loom_libs:ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32": 48, + "loom_libs:ggml_rmsnorm_mul_rope_f32": 32, + "loom_libs:ggml_set_rows": 32, + "loom_libs:llm_gated_delta_net_f32_wmma_head128_snapshot": 48, + "loom_libs:llm_gated_delta_net_projection_epilogue_f32": 48, + "loom_libs:llm_ssm_conv_dconv4_silu_rollback_f32": 48 + }, + "loom_source": "tools/benchmarks/loom/v2_mtp4_program4.work.loom", + "model": "v2_mtp4_program4", + "scenario": "tg_c4", + "schema": "ggml-hrx-model-loom-benchmarks-v2" +} diff --git a/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_mtp4_program4.work.loom b/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_mtp4_program4.work.loom new file mode 100644 index 000000000000..3b3e610abcbc --- /dev/null +++ b/ggml/src/ggml-hrx/tools/benchmarks/loom/v2_mtp4_program4.work.loom @@ -0,0 +1,402 @@ +// Generated by tools/benchmarks/generate-model-benchmarks.py for v2_mtp4_program4 scenarios: tg_c4. +// Regenerate from an HRX command program dump rather than editing by hand. + +target.decl @ggml_binary_f32_gfx11_wave64 +target.decl @ggml_binary_swiglu_i4_gfx11_wave32 +target.decl @ggml_flash_attention_gfx11_wave64 +target.decl @ggml_get_rows_f32_gfx11_wave64 +target.decl @ggml_rmsnorm_binary_gfx11_wave32 +target.decl @ggml_rmsnorm_gfx11_wave32 +target.decl @ggml_set_rows_gfx11_wave64 + +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_add_rmsnorm_binary_symmetric_i4_k32(%token_count: index) launch(%token_count: index, %lhs: buffer, %rhs: buffer, %residual_output: buffer, %weight: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_binary_f32_gfx11_wave64) @ggml_binary_f32(%element_count: index) launch(%element_count: index, %lhs: buffer, %rhs: buffer, %output: buffer) +kernel.decl target(@ggml_binary_swiglu_i4_gfx11_wave32) @ggml_binary_swiglu_symmetric_i4_k32() launch(%lhs: buffer, %rhs: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_flash_attention_gfx11_wave64) @ggml_flash_attention_f32_f16_wmma(%query_token_count: index, %key_value_token_count: index) launch(%query_token_count: index, %key_value_token_count: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %gate: buffer, %output: buffer) +kernel.decl target(@ggml_get_rows_f32_gfx11_wave64) @ggml_get_rows_f32(%token_count: index, %row_count: index, %hidden_size: index) launch(%token_count: index, %row_count: index, %hidden_size: index, %token_ids: buffer, %weight: buffer, %output: buffer) +kernel.decl @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) +kernel.decl @ggml_mul_mat_q6_k_packed_token1_f16_wmma(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer) +kernel.decl @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma() launch(%first_weight: buffer, %second_weight: buffer, %first_output: buffer, %second_output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma() launch(%weight: buffer, %input: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer, %partial: buffer, %completion_counters: buffer) +kernel.decl @ggml_mul_mat_symmetric_i4_lowrow_wmma() launch(%weight: buffer, %input: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl @ggml_quantize_f32_symmetric_i4_k32() launch(%input: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_rmsnorm_binary_f32(%token_count: index) launch(%token_count: index, %input: buffer, %rhs: buffer, %output: buffer) +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_rmsnorm_binary_symmetric_i4_k32(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_rmsnorm_binary_gfx11_wave32) @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32(%token_count: index) launch(%token_count: index, %input: buffer, %weight: buffer, %raw_gate: buffer, %output: buffer, %quantized_values: buffer, %scales: buffer, %sums: buffer) +kernel.decl target(@ggml_rmsnorm_gfx11_wave32) @ggml_rmsnorm_mul_rope_f32() launch(%input: buffer, %weight: buffer, %positions: buffer, %output: buffer) +kernel.decl target(@ggml_set_rows_gfx11_wave64) @ggml_set_rows(%token_count: index, %cache_row_count: index, %hidden_size: index) launch(%token_count: index, %cache_row_count: index, %hidden_size: index, %rows: buffer, %indices: buffer, %cache: buffer) +kernel.decl @llm_gated_delta_net_f32_wmma_head128_snapshot() launch(%q: buffer, %k: buffer, %v: buffer, %g: buffer, %beta: buffer, %state_in: buffer, %snapshot_cache: buffer, %dst: buffer, %rms_scales: buffer) +kernel.decl @llm_gated_delta_net_projection_epilogue_f32() launch(%alpha_raw: buffer, %beta_raw: buffer, %bias: buffer, %a_scale: buffer, %gate_dst: buffer, %beta_dst: buffer) +kernel.decl @llm_ssm_conv_dconv4_silu_rollback_f32() launch(%state: buffer, %x: buffer, %filter: buffer, %output: buffer, %cache0: buffer, %cache1: buffer, %cache2: buffer, %cache3: buffer, %cache4: buffer) + + +// Scenario: tg_c4 + +check.case public @v2_mtp4_program4_tg_c4_000_ggml_get_rows_f32_case { + %token_count = check.literal value(4) : index + %row_count = check.literal value(248320) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<16xi8> + %weight = check.generate.fill value(0) : tensor<715161600xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<16xi8>, tensor<715161600xi8>, tensor<81920xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_000_ggml_get_rows_f32_case> @v2_mtp4_program4_tg_c4_000_ggml_get_rows_f32 + +check.case public @v2_mtp4_program4_tg_c4_001_ggml_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<81920xi8>, tensor<20480xi8>, tensor<81920xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_001_ggml_rmsnorm_binary_symmetric_i4_k32_case> @v2_mtp4_program4_tg_c4_001_ggml_rmsnorm_binary_symmetric_i4_k32 + +check.case public @v2_mtp4_program4_tg_c4_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<27033600xi8> + %input = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<163840xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + %partial = check.generate.fill value(0) : tensor<163840xi8> + %completion_counters = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<27033600xi8>, tensor<81920xi8>, tensor<163840xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>, tensor<163840xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @v2_mtp4_program4_tg_c4_002_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @v2_mtp4_program4_tg_c4_003_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<168960xi8> + %second_weight = check.generate.fill value(0) : tensor<168960xi8> + %first_output = check.generate.fill value(0) : tensor<768xi8> + %second_output = check.generate.fill value(0) : tensor<768xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<168960xi8>, tensor<168960xi8>, tensor<768xi8>, tensor<768xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_003_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @v2_mtp4_program4_tg_c4_003_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @v2_mtp4_program4_tg_c4_004_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<98304xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<81920xi8>, tensor<98304xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_004_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @v2_mtp4_program4_tg_c4_004_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @v2_mtp4_program4_tg_c4_005_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(20) : index + %hidden_size = check.literal value(30720) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<20x30720xf32> + %output = check.generate.fill value(1.0) : tensor<1x30720xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<20x30720xf32>, tensor<1x30720xf32>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_005_ggml_get_rows_f32_case> @v2_mtp4_program4_tg_c4_005_ggml_get_rows_f32 + +check.case public @v2_mtp4_program4_tg_c4_006_llm_ssm_conv_dconv4_silu_rollback_f32_case { + %state = check.generate.fill value(0) : tensor<122880xi8> + %x = check.generate.fill value(0) : tensor<163840xi8> + %filter = check.generate.fill value(0) : tensor<163840xi8> + %output = check.generate.fill value(0) : tensor<163840xi8> + %cache0 = check.generate.fill value(0) : tensor<122880xi8> + %cache1 = check.generate.fill value(0) : tensor<122880xi8> + %cache2 = check.generate.fill value(0) : tensor<122880xi8> + %cache3 = check.generate.fill value(0) : tensor<122880xi8> + %cache4 = check.generate.fill value(0) : tensor<122880xi8> + kernel.launch @llm_ssm_conv_dconv4_silu_rollback_f32[](%state, %x, %filter, %output, %cache0, %cache1, %cache2, %cache3, %cache4) : [](tensor<122880xi8>, tensor<163840xi8>, tensor<163840xi8>, tensor<163840xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>, tensor<122880xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_006_llm_ssm_conv_dconv4_silu_rollback_f32_case> @v2_mtp4_program4_tg_c4_006_llm_ssm_conv_dconv4_silu_rollback_f32 + +check.case public @v2_mtp4_program4_tg_c4_007_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(20) : index + %hidden_size = check.literal value(786432) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<20x786432xf32> + %output = check.generate.fill value(1.0) : tensor<1x786432xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<20x786432xf32>, tensor<1x786432xf32>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_007_ggml_get_rows_f32_case> @v2_mtp4_program4_tg_c4_007_ggml_get_rows_f32 + +check.case public @v2_mtp4_program4_tg_c4_008_llm_gated_delta_net_projection_epilogue_f32_case { + %alpha_raw = check.generate.fill value(0) : tensor<768xi8> + %beta_raw = check.generate.fill value(0) : tensor<768xi8> + %bias = check.generate.fill value(0) : tensor<192xi8> + %a_scale = check.generate.fill value(0) : tensor<192xi8> + %gate_dst = check.generate.fill value(0) : tensor<768xi8> + %beta_dst = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @llm_gated_delta_net_projection_epilogue_f32[](%alpha_raw, %beta_raw, %bias, %a_scale, %gate_dst, %beta_dst) : [](tensor<768xi8>, tensor<768xi8>, tensor<192xi8>, tensor<192xi8>, tensor<768xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_008_llm_gated_delta_net_projection_epilogue_f32_case> @v2_mtp4_program4_tg_c4_008_llm_gated_delta_net_projection_epilogue_f32 + +check.case public @v2_mtp4_program4_tg_c4_009_llm_gated_delta_net_f32_wmma_head128_snapshot_case { + %q = check.generate.fill value(0) : tensor<131072xi8> + %k = check.generate.fill value(0) : tensor<131072xi8> + %v = check.generate.fill value(0) : tensor<147456xi8> + %g = check.generate.fill value(0) : tensor<768xi8> + %beta = check.generate.fill value(0) : tensor<768xi8> + %state_in = check.generate.fill value(0) : tensor<3145728xi8> + %snapshot_cache = check.generate.fill value(0) : tensor<40894464xi8> + %dst = check.generate.fill value(0) : tensor<98304xi8> + %rms_scales = check.generate.fill value(0) : tensor<768xi8> + kernel.launch @llm_gated_delta_net_f32_wmma_head128_snapshot[](%q, %k, %v, %g, %beta, %state_in, %snapshot_cache, %dst, %rms_scales) : [](tensor<131072xi8>, tensor<131072xi8>, tensor<147456xi8>, tensor<768xi8>, tensor<768xi8>, tensor<3145728xi8>, tensor<40894464xi8>, tensor<98304xi8>, tensor<768xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_009_llm_gated_delta_net_f32_wmma_head128_snapshot_case> @v2_mtp4_program4_tg_c4_009_llm_gated_delta_net_f32_wmma_head128_snapshot + +check.case public @v2_mtp4_program4_tg_c4_010_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case { + %token_count = check.literal value(192) : index + %input = check.generate.fill value(0) : tensor<98304xi8> + %weight = check.generate.fill value(0) : tensor<512xi8> + %raw_gate = check.generate.fill value(0) : tensor<98304xi8> + %output = check.generate.fill value(0) : tensor<98304xi8> + %quantized_values = check.generate.fill value(0) : tensor<12288xi8> + %scales = check.generate.fill value(0) : tensor<3072xi8> + %sums = check.generate.fill value(0) : tensor<3072xi8> + kernel.launch @ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32[%token_count](%token_count, %input, %weight, %raw_gate, %output, %quantized_values, %scales, %sums) : [index](index, tensor<98304xi8>, tensor<512xi8>, tensor<98304xi8>, tensor<98304xi8>, tensor<12288xi8>, tensor<3072xi8>, tensor<3072xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_010_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32_case> @v2_mtp4_program4_tg_c4_010_ggml_rmsnorm_gate_silu_mul_symmetric_i4_k32 + +check.case public @v2_mtp4_program4_tg_c4_011_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<16220160xi8> + %input = check.generate.fill value(0) : tensor<98304xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<12288xi8> + %scales = check.generate.fill value(0) : tensor<3072xi8> + %sums = check.generate.fill value(0) : tensor<3072xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<16220160xi8>, tensor<98304xi8>, tensor<81920xi8>, tensor<12288xi8>, tensor<3072xi8>, tensor<3072xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_011_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @v2_mtp4_program4_tg_c4_011_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @v2_mtp4_program4_tg_c4_012_ggml_add_rmsnorm_binary_symmetric_i4_k32_case { + %token_count = check.literal value(4) : index + %lhs = check.generate.fill value(0) : tensor<81920xi8> + %rhs = check.generate.fill value(0) : tensor<81920xi8> + %residual_output = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<20480xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_add_rmsnorm_binary_symmetric_i4_k32[%token_count](%token_count, %lhs, %rhs, %residual_output, %weight, %output, %quantized_values, %scales, %sums) : [index](index, tensor<81920xi8>, tensor<81920xi8>, tensor<81920xi8>, tensor<20480xi8>, tensor<81920xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_012_ggml_add_rmsnorm_binary_symmetric_i4_k32_case> @v2_mtp4_program4_tg_c4_012_ggml_add_rmsnorm_binary_symmetric_i4_k32 + +check.case public @v2_mtp4_program4_tg_c4_013_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case { + %first_weight = check.generate.fill value(0) : tensor<45957120xi8> + %second_weight = check.generate.fill value(0) : tensor<45957120xi8> + %first_output = check.generate.fill value(0) : tensor<278528xi8> + %second_output = check.generate.fill value(0) : tensor<278528xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma[](%first_weight, %second_weight, %first_output, %second_output, %quantized_values, %scales, %sums) : [](tensor<45957120xi8>, tensor<45957120xi8>, tensor<278528xi8>, tensor<278528xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_013_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma_case> @v2_mtp4_program4_tg_c4_013_ggml_mul_mat_symmetric_i4_lowrow_adjacent_dual_wmma + +check.case public @v2_mtp4_program4_tg_c4_014_ggml_binary_swiglu_symmetric_i4_k32_case { + %lhs = check.generate.fill value(0) : tensor<278528xi8> + %rhs = check.generate.fill value(0) : tensor<278528xi8> + %output = check.generate.fill value(0) : tensor<278528xi8> + %quantized_values = check.generate.fill value(0) : tensor<34816xi8> + %scales = check.generate.fill value(0) : tensor<8704xi8> + %sums = check.generate.fill value(0) : tensor<8704xi8> + kernel.launch @ggml_binary_swiglu_symmetric_i4_k32[](%lhs, %rhs, %output, %quantized_values, %scales, %sums) : [](tensor<278528xi8>, tensor<278528xi8>, tensor<278528xi8>, tensor<34816xi8>, tensor<8704xi8>, tensor<8704xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_014_ggml_binary_swiglu_symmetric_i4_k32_case> @v2_mtp4_program4_tg_c4_014_ggml_binary_swiglu_symmetric_i4_k32 + +check.case public @v2_mtp4_program4_tg_c4_015_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case { + %weight = check.generate.fill value(0) : tensor<45957120xi8> + %input = check.generate.fill value(0) : tensor<278528xi8> + %output = check.generate.fill value(0) : tensor<81920xi8> + %quantized_values = check.generate.fill value(0) : tensor<34816xi8> + %scales = check.generate.fill value(0) : tensor<8704xi8> + %sums = check.generate.fill value(0) : tensor<8704xi8> + %partial = check.generate.fill value(0) : tensor<81920xi8> + %completion_counters = check.generate.fill value(0) : tensor<1280xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums, %partial, %completion_counters) : [](tensor<45957120xi8>, tensor<278528xi8>, tensor<81920xi8>, tensor<34816xi8>, tensor<8704xi8>, tensor<8704xi8>, tensor<81920xi8>, tensor<1280xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_015_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma_case> @v2_mtp4_program4_tg_c4_015_ggml_mul_mat_symmetric_i4_lowrow_split_k2_wmma + +check.case public @v2_mtp4_program4_tg_c4_016_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<32440320xi8> + %input = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<196608xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<32440320xi8>, tensor<81920xi8>, tensor<196608xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_016_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @v2_mtp4_program4_tg_c4_016_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @v2_mtp4_program4_tg_c4_017_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(0) : tensor<81920xi8> + %weight = check.generate.fill value(0) : tensor<5611520xi8> + %output = check.generate.fill value(0) : tensor<16384xi8> + kernel.launch @ggml_mul_mat_q6_k_i8_prepacked_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<81920xi8>, tensor<5611520xi8>, tensor<16384xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_017_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma_case> @v2_mtp4_program4_tg_c4_017_ggml_mul_mat_q6_k_i8_prepacked_f16_wmma + +check.case public @v2_mtp4_program4_tg_c4_018_ggml_mul_mat_symmetric_i4_lowrow_wmma_case { + %weight = check.generate.fill value(0) : tensor<2703360xi8> + %input = check.generate.fill value(0) : tensor<81920xi8> + %output = check.generate.fill value(0) : tensor<16384xi8> + %quantized_values = check.generate.fill value(0) : tensor<10240xi8> + %scales = check.generate.fill value(0) : tensor<2560xi8> + %sums = check.generate.fill value(0) : tensor<2560xi8> + kernel.launch @ggml_mul_mat_symmetric_i4_lowrow_wmma[](%weight, %input, %output, %quantized_values, %scales, %sums) : [](tensor<2703360xi8>, tensor<81920xi8>, tensor<16384xi8>, tensor<10240xi8>, tensor<2560xi8>, tensor<2560xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_018_ggml_mul_mat_symmetric_i4_lowrow_wmma_case> @v2_mtp4_program4_tg_c4_018_ggml_mul_mat_symmetric_i4_lowrow_wmma + +check.case public @v2_mtp4_program4_tg_c4_019_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<195584xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<64xi8> + %output = check.generate.fill value(0) : tensor<98304xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<195584xi8>, tensor<1024xi8>, tensor<64xi8>, tensor<98304xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_019_ggml_rmsnorm_mul_rope_f32_case> @v2_mtp4_program4_tg_c4_019_ggml_rmsnorm_mul_rope_f32 + +check.case public @v2_mtp4_program4_tg_c4_020_ggml_rmsnorm_mul_rope_f32_case { + %input = check.generate.fill value(0) : tensor<16384xi8> + %weight = check.generate.fill value(0) : tensor<1024xi8> + %positions = check.generate.fill value(0) : tensor<64xi8> + %output = check.generate.fill value(0) : tensor<16384xi8> + kernel.launch @ggml_rmsnorm_mul_rope_f32[](%input, %weight, %positions, %output) : [](tensor<16384xi8>, tensor<1024xi8>, tensor<64xi8>, tensor<16384xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_020_ggml_rmsnorm_mul_rope_f32_case> @v2_mtp4_program4_tg_c4_020_ggml_rmsnorm_mul_rope_f32 + +check.case public @v2_mtp4_program4_tg_c4_021_ggml_set_rows_case { + %token_count = check.literal value(4) : index + %cache_row_count = check.literal value(262144) : index + %hidden_size = check.literal value(1024) : index + %rows = check.generate.fill value(0.0) : tensor<4x1024xf32> + %indices = check.generate.iota offset(0) step(1) period(262144) : tensor<4xi64> + %cache = check.generate.fill value(0.0) : tensor<262144x1024xf16> + kernel.launch @ggml_set_rows[%token_count, %cache_row_count, %hidden_size](%token_count, %cache_row_count, %hidden_size, %rows, %indices, %cache) : [index, index, index](index, index, index, tensor<4x1024xf32>, tensor<4xi64>, tensor<262144x1024xf16>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_021_ggml_set_rows_case> @v2_mtp4_program4_tg_c4_021_ggml_set_rows + +check.case public @v2_mtp4_program4_tg_c4_022_ggml_flash_attention_f32_f16_wmma_case { + %query_token_count = check.literal value(4) : index + %key_value_token_count = check.literal value(256) : index + %query = check.generate.fill value(0.0) : tensor<4x24x256xf32> + %key = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %value = check.generate.fill value(0.0) : tensor<256x4x256xf16> + %mask = check.generate.fill value(0.0) : tensor<4x256xf16> + %output = check.generate.fill value(1.0) : tensor<4x24x256xf32> + kernel.launch @ggml_flash_attention_f32_f16_wmma[%query_token_count, %key_value_token_count](%query_token_count, %key_value_token_count, %query, %key, %value, %mask, %output) : [index, index](index, index, tensor<4x24x256xf32>, tensor<256x4x256xf16>, tensor<256x4x256xf16>, tensor<4x256xf16>, tensor<4x24x256xf32>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_022_ggml_flash_attention_f32_f16_wmma_case> @v2_mtp4_program4_tg_c4_022_ggml_flash_attention_f32_f16_wmma + +check.case public @v2_mtp4_program4_tg_c4_023_ggml_quantize_f32_symmetric_i4_k32_case { + %input = check.generate.fill value(0) : tensor<98304xi8> + %quantized_values = check.generate.fill value(0) : tensor<12288xi8> + %scales = check.generate.fill value(0) : tensor<3072xi8> + %sums = check.generate.fill value(0) : tensor<3072xi8> + kernel.launch @ggml_quantize_f32_symmetric_i4_k32[](%input, %quantized_values, %scales, %sums) : [](tensor<98304xi8>, tensor<12288xi8>, tensor<3072xi8>, tensor<3072xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_023_ggml_quantize_f32_symmetric_i4_k32_case> @v2_mtp4_program4_tg_c4_023_ggml_quantize_f32_symmetric_i4_k32 + +check.case public @v2_mtp4_program4_tg_c4_024_ggml_binary_f32_case { + %element_count = check.literal value(20480) : index + %lhs = check.generate.fill value(2.0) : tensor<20480xf32> + %rhs = check.generate.fill value(3.0) : tensor<20480xf32> + %output = check.generate.fill value(0.0) : tensor<20480xf32> + kernel.launch @ggml_binary_f32[%element_count](%element_count, %lhs, %rhs, %output) : [index](index, tensor<20480xf32>, tensor<20480xf32>, tensor<20480xf32>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_024_ggml_binary_f32_case> @v2_mtp4_program4_tg_c4_024_ggml_binary_f32 + +check.case public @v2_mtp4_program4_tg_c4_025_ggml_rmsnorm_binary_f32_case { + %token_count = check.literal value(4) : index + %input = check.generate.fill value(2.0) : tensor<4x5120xf32> + %rhs = check.generate.fill value(3.0) : tensor<5120xf32> + %output = check.generate.fill value(0.0) : tensor<4x5120xf32> + kernel.launch @ggml_rmsnorm_binary_f32[%token_count](%token_count, %input, %rhs, %output) : [index](index, tensor<4x5120xf32>, tensor<5120xf32>, tensor<4x5120xf32>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_025_ggml_rmsnorm_binary_f32_case> @v2_mtp4_program4_tg_c4_025_ggml_rmsnorm_binary_f32 + +check.case public @v2_mtp4_program4_tg_c4_026_ggml_get_rows_f32_case { + %token_count = check.literal value(1) : index + %row_count = check.literal value(4) : index + %hidden_size = check.literal value(5120) : index + %token_ids = check.generate.fill value(0) : tensor<1xi32> + %weight = check.generate.fill value(0.0) : tensor<4x5120xf32> + %output = check.generate.fill value(1.0) : tensor<1x5120xf32> + kernel.launch @ggml_get_rows_f32[%token_count, %row_count, %hidden_size](%token_count, %row_count, %hidden_size, %token_ids, %weight, %output) : [index, index, index](index, index, index, tensor<1xi32>, tensor<4x5120xf32>, tensor<1x5120xf32>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_026_ggml_get_rows_f32_case> @v2_mtp4_program4_tg_c4_026_ggml_get_rows_f32 + +check.case public @v2_mtp4_program4_tg_c4_027_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case { + %token_count = check.literal value(1) : index + %input = check.generate.fill value(0) : tensor<20480xi8> + %weight = check.generate.fill value(0) : tensor<1042944000xi8> + %output = check.generate.fill value(0) : tensor<993280xi8> + kernel.launch @ggml_mul_mat_q6_k_packed_token1_f16_wmma[%token_count](%token_count, %input, %weight, %output) : [index](index, tensor<20480xi8>, tensor<1042944000xi8>, tensor<993280xi8>) + check.return +} + +check.benchmark<@v2_mtp4_program4_tg_c4_027_ggml_mul_mat_q6_k_packed_token1_f16_wmma_case> @v2_mtp4_program4_tg_c4_027_ggml_mul_mat_q6_k_packed_token1_f16_wmma diff --git a/ggml/src/ggml-hrx/tools/benchmarks/run-model-benchmarks.py b/ggml/src/ggml-hrx/tools/benchmarks/run-model-benchmarks.py new file mode 100755 index 000000000000..e93268dc2683 --- /dev/null +++ b/ggml/src/ggml-hrx/tools/benchmarks/run-model-benchmarks.py @@ -0,0 +1,494 @@ +#!/usr/bin/env python3 +# +# Run generated HRX model-scoped Loom benchmarks. + +from __future__ import annotations + +import argparse +import json +import os +import shutil +import signal +import subprocess +from pathlib import Path +from typing import Any, NamedTuple + + +SCRIPT_DIR = Path(__file__).resolve().parent +HRX_DIR = SCRIPT_DIR.parents[1] +REPO_DIR = HRX_DIR.parents[2] +BENCHMARK_DIR = HRX_DIR / "benchmarks" + + +def fail(message: str) -> None: + raise SystemExit(message) + + +def load_json(path: Path) -> Any: + with path.open("r", encoding="utf-8") as f: + return json.load(f) + + +def find_tool(env_name: str, binary_name: str, build_dir: Path | None, build_suffix: str) -> Path: + override = os.environ.get(env_name) + if override: + path = Path(override) + if path.is_file() and os.access(path, os.X_OK): + return path + fail(f"{env_name} is not executable: {path}") + if build_dir is not None: + matches = sorted(build_dir.glob(build_suffix)) + for match in matches: + if match.is_file() and os.access(match, os.X_OK): + return match + found = shutil.which(binary_name) + if found: + return Path(found) + fail(f"could not find {binary_name}; set {env_name} or pass --build-dir") + + +def source_path(entry: dict[str, Any], source: str) -> Path: + corpus_dir = entry.get("corpus_dir") + if not corpus_dir: + fail(f"{entry['kernel']} has no corpus_dir in manifest") + return HRX_DIR / corpus_dir / source + + +def manifest_entries(manifest: dict[str, Any]) -> list[dict[str, Any]]: + if "dispatches" in manifest: + return manifest.get("dispatches", []) + return manifest.get("benchmarks", []) + + +def entry_status(entry: dict[str, Any]) -> str: + status = entry.get("status") + if status is not None: + return str(status) + if entry.get("benchmark"): + return "generated" + return "unsupported" + + +def entry_count(entry: dict[str, Any]) -> int: + return int(entry.get("count", entry.get("shape_multiplicity", 1))) + + +def parse_bool(value: str) -> bool: + value = value.lower() + if value in ("1", "true", "yes", "on"): + return True + if value in ("0", "false", "no", "off"): + return False + raise argparse.ArgumentTypeError(f"expected true or false, got {value!r}") + + +class CommandResult(NamedTuple): + state: str + error: str | None + returncode: int | None + signal_name: str | None + stdout: Path + stderr: Path + command: Path + + +def format_signal(returncode: int | None) -> str | None: + if returncode is None or returncode >= 0: + return None + signum = -returncode + try: + return signal.Signals(signum).name + except ValueError: + return f"SIG{signum}" + + +def process_output_text(output: str | bytes | None) -> str: + if output is None: + return "" + if isinstance(output, bytes): + return output.decode("utf-8", errors="replace") + return output + + +def run_command(command: list[str], cwd: Path, timeout_sec: float | None, bench_dir: Path, stage: str) -> CommandResult: + stdout_path = bench_dir / f"{stage}.stdout.txt" + stderr_path = bench_dir / f"{stage}.stderr.txt" + command_path = bench_dir / f"{stage}.command.json" + command_record: dict[str, Any] = { + "argv": command, + "cwd": str(cwd), + "timeout_sec": timeout_sec, + } + try: + result = subprocess.run(command, cwd=str(cwd), capture_output=True, text=True, timeout=timeout_sec) + except subprocess.TimeoutExpired as exc: + stdout_path.write_text(process_output_text(exc.stdout), encoding="utf-8") + stderr_path.write_text(process_output_text(exc.stderr), encoding="utf-8") + command_record.update( + { + "state": "timeout", + "returncode": None, + "signal": None, + "timed_out": True, + } + ) + command_path.write_text(json.dumps(command_record, indent=2, sort_keys=True) + "\n", encoding="utf-8") + return CommandResult( + "timeout", + f"{stage} timed out after {exc.timeout} seconds", + None, + None, + stdout_path, + stderr_path, + command_path, + ) + + stdout_path.write_text(result.stdout, encoding="utf-8") + stderr_path.write_text(result.stderr, encoding="utf-8") + sig_name = format_signal(result.returncode) + command_record.update( + { + "state": "ok" if result.returncode == 0 else "failed", + "returncode": result.returncode, + "signal": sig_name, + "timed_out": False, + } + ) + command_path.write_text(json.dumps(command_record, indent=2, sort_keys=True) + "\n", encoding="utf-8") + if result.returncode == 0: + return CommandResult("ok", None, result.returncode, sig_name, stdout_path, stderr_path, command_path) + if sig_name: + return CommandResult( + "failed", + f"{stage} exited with status {result.returncode} ({sig_name})", + result.returncode, + sig_name, + stdout_path, + stderr_path, + command_path, + ) + return CommandResult( + "failed", + f"{stage} exited with status {result.returncode}", + result.returncode, + sig_name, + stdout_path, + stderr_path, + command_path, + ) + + +def command_result_json(result: CommandResult) -> dict[str, Any]: + return { + "state": result.state, + "error": result.error, + "returncode": result.returncode, + "signal": result.signal_name, + "stdout": str(result.stdout), + "stderr": str(result.stderr), + "command": str(result.command), + } + + +def find_matching_char(text: str, open_index: int, open_char: str, close_char: str) -> int: + depth = 0 + for index in range(open_index, len(text)): + char = text[index] + if char == open_char: + depth += 1 + elif char == close_char: + depth -= 1 + if depth == 0: + return index + raise ValueError(f"unmatched {open_char!r}") + + +def parse_root_index_parameters(source_text: str, symbol: str) -> list[str]: + marker = f"@{symbol}(" + start = source_text.find(marker) + if start < 0: + raise ValueError(f"could not find root @{symbol}") + params_start = start + len(marker) - 1 + params_end = find_matching_char(source_text, params_start, "(", ")") + params_text = source_text[params_start + 1 : params_end] + names = [] + for raw_param in params_text.split(","): + parts = raw_param.strip().split() + if len(parts) < 2 or parts[-1] != "index": + continue + name = parts[-2].rstrip(":") + if not name.startswith("%"): + continue + names.append(name[1:]) + return names + + +def replace_index_uses(region: str, name: str, constant_name: str) -> str: + token = "%" + name + output = [] + index = 0 + while True: + match = region.find(token, index) + if match < 0: + output.append(region[index:]) + return "".join(output) + before_ok = match == 0 or not (region[match - 1].isalnum() or region[match - 1] == "_") + end = match + len(token) + after_ok = end == len(region) or not (region[end].isalnum() or region[end] == "_") + output.append(region[index:match]) + if before_ok and after_ok: + output.append("%" + constant_name) + else: + output.append(token) + index = end + + +def specialize_region(region: str, workload_values: list[tuple[str, int]]) -> str: + body = region + constants = [] + for name, value in workload_values: + constant_name = f"ggml_hrx_specialized_{name}" + constants.append(f" %{constant_name} = index.constant {value} : index\n") + body = replace_index_uses(body, name, constant_name) + return "\n" + "".join(constants) + body + + +def specialize_workload_text(source_text: str, symbol: str, workload_values: list[tuple[str, int]]) -> str: + if not workload_values: + return source_text + marker = f"@{symbol}(" + root_start = source_text.find(marker) + if root_start < 0: + raise ValueError(f"could not find root @{symbol}") + params_start = root_start + len(marker) - 1 + params_end = find_matching_char(source_text, params_start, "(", ")") + + config_open = source_text.find("{", params_end) + if config_open < 0: + raise ValueError("could not find kernel config block") + config_close = find_matching_char(source_text, config_open, "{", "}") + + launch_marker = "launch(" + launch_start = source_text.find(launch_marker, config_close) + if launch_start < 0: + raise ValueError("could not find launch region") + launch_body_open = source_text.find("{", launch_start) + if launch_body_open < 0: + raise ValueError("could not find launch body") + launch_body_close = find_matching_char(source_text, launch_body_open, "{", "}") + + config_body = source_text[config_open + 1 : config_close] + launch_body = source_text[launch_body_open + 1 : launch_body_close] + specialized_config = specialize_region(config_body, workload_values) + specialized_launch = specialize_region(launch_body, workload_values) + + return ( + source_text[: config_open + 1] + + specialized_config + + source_text[config_close: launch_body_open + 1] + + specialized_launch + + source_text[launch_body_close:] + ) + + +def specialize_linked_source(entry: dict[str, Any], linked_source: Path, output_source: Path) -> bool: + symbol = entry.get("symbol") or entry["kernel"].split(":")[-1] + integer_parameters = entry.get("integer_parameters", {}) + workload_parameters = entry.get("workload_parameters") or [] + workload_names = [param["name"] for param in workload_parameters if param.get("type") == "index"] + source_text = linked_source.read_text(encoding="utf-8") + if not workload_names: + workload_names = parse_root_index_parameters(source_text, symbol) + workload_values = [] + for name in workload_names: + if name not in integer_parameters: + continue + workload_values.append((name, int(integer_parameters[name]))) + if not workload_values: + return False + output_source.write_text(specialize_workload_text(source_text, symbol, workload_values), encoding="utf-8") + return True + + +def select_entries(manifest: dict[str, Any], benchmark: str | None, include_unsupported: bool) -> list[dict[str, Any]]: + entries = manifest_entries(manifest) + if benchmark: + requested = benchmark if benchmark.startswith("@") else "@" + benchmark + entries = [entry for entry in entries if entry.get("benchmark") == requested] + if not entries: + fail(f"benchmark not found in manifest: {requested}") + if include_unsupported: + return entries + return [entry for entry in entries if entry_status(entry) == "generated"] + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--model", required=True, help="Model slug, for example llama32_3b_f16.") + parser.add_argument("--scenario", required=True, help="Scenario slug, for example pp512.") + parser.add_argument("--build-dir", type=Path, help="llama.cpp build directory.") + parser.add_argument("--output-dir", type=Path, help="Directory for linked sources, plans, and results.") + parser.add_argument("--manifest", type=Path, help="Generated model scenario manifest.") + parser.add_argument("--benchmark", help="Single benchmark symbol to run.") + parser.add_argument("--device", default=os.environ.get("DEVICE", "amdgpu"), help="HAL device. Defaults to amdgpu.") + parser.add_argument("--dry-run-only", action="store_true", help="Stop after benchmark planning.") + parser.add_argument("--list", action="store_true", help="List selected generated benchmarks and exit.") + parser.add_argument("--include-unsupported", action="store_true", help="Include unsupported manifest entries when listing.") + parser.add_argument("--batch-size", default="64") + parser.add_argument("--iterations", default="10") + parser.add_argument("--warmup-iterations", default="1") + parser.add_argument("--continue-on-failure", action="store_true", help="Record failed benchmarks and continue running the rest.") + parser.add_argument("--benchmark-timeout-sec", type=float, default=300.0, help="Timeout per benchmark tool invocation. Defaults to 300.") + parser.add_argument("--profile-final-batch", type=parse_bool, default=True, help="Pass --profile-final-batch to iree-benchmark-loom. Defaults to true.") + parser.add_argument("--input-ring-count", help="Optional --input-ring-count value for iree-benchmark-loom.") + parser.add_argument("--benchmark-extra-arg", action="append", default=[], help="Additional argument to pass to benchmark invocations. May be repeated.") + parser.add_argument("--runtime-specialization", action=argparse.BooleanOptionalAction, default=True, help="Specialize workload index parameters after loom-link. Defaults to enabled.") + args = parser.parse_args() + + build_dir = args.build_dir + if build_dir is not None and not build_dir.is_absolute(): + build_dir = REPO_DIR / build_dir + manifest_path = args.manifest or (BENCHMARK_DIR / "loom" / f"{args.model}.{args.scenario}.json") + manifest = load_json(manifest_path) + loom_source = HRX_DIR / manifest["loom_source"] + entries = select_entries(manifest, args.benchmark, args.include_unsupported) + + if args.list: + for entry in entries: + count = entry_count(entry) + print(f"{entry.get('benchmark') or '-'} {entry_status(entry)} {entry['kernel']} x{count}") + return + + entries = [entry for entry in entries if entry_status(entry) == "generated"] + if not entries: + fail("no generated benchmarks selected") + if args.output_dir is None: + fail("--output-dir is required unless --list is used") + + loom_link = find_tool("LOOM_LINK", "loom-link", build_dir, "**/loom-link") + iree_benchmark_loom = find_tool("IREE_BENCHMARK_LOOM", "iree-benchmark-loom", build_dir, "**/iree-benchmark-loom") + args.output_dir.mkdir(parents=True, exist_ok=True) + summary_path = args.output_dir / "results.jsonl" + with summary_path.open("w", encoding="utf-8") as summary: + for entry in entries: + benchmark = entry["benchmark"] + safe_name = benchmark[1:] + bench_dir = args.output_dir / safe_name + artifact_dir = bench_dir / "artifacts" + bench_dir.mkdir(parents=True, exist_ok=True) + artifact_dir.mkdir(parents=True, exist_ok=True) + linked_source = bench_dir / "linked.loom" + specialized_source = bench_dir / "linked.runtime-specialized.loom" + plan_output = bench_dir / "plan.json" + result_output = bench_dir / "results.json" + stage_results: dict[str, Any] = {} + config_args = [f"--config={name}={value}" for name, value in sorted(entry.get("compile_parameters", {}).items())] + + sources = entry.get("sources") + if sources is None: + sources = [*entry.get("primary_sources", []), *entry.get("library_sources", [])] + link_inputs = [str(source_path(entry, source)) for source in sources] + state = "ok" + error = None + link_result = run_command( + [ + str(loom_link), + str(loom_source), + *link_inputs, + "--mode=link", + "--to=text", + f"--root={benchmark}", + f"--output={linked_source}", + *config_args, + ], + REPO_DIR, + args.benchmark_timeout_sec, + bench_dir, + "loom-link", + ) + stage_results["loom-link"] = command_result_json(link_result) + state = link_result.state + error = link_result.error + if state == "ok": + benchmark_source = linked_source + runtime_specialized = False + if args.runtime_specialization: + try: + runtime_specialized = specialize_linked_source(entry, linked_source, specialized_source) + except ValueError as exc: + state = "failed" + error = f"runtime specialization failed: {exc}" + if runtime_specialized: + benchmark_source = specialized_source + else: + runtime_specialized = False + else: + benchmark_source = linked_source + runtime_specialized = False + if state == "ok": + dry_run = [ + str(iree_benchmark_loom), + str(benchmark_source), + f"--benchmark={benchmark}", + f"--device={args.device}", + "--dry-run", + "--measure=dispatch_complete", + "--output-format=jsonl", + f"--output={plan_output}", + *config_args, + ] + dry_run_result = run_command(dry_run, REPO_DIR, args.benchmark_timeout_sec, bench_dir, "dry-run") + stage_results["dry-run"] = command_result_json(dry_run_result) + state = dry_run_result.state + error = dry_run_result.error + if state == "ok" and not args.dry_run_only: + benchmark_command = [ + str(iree_benchmark_loom), + str(benchmark_source), + f"--benchmark={benchmark}", + f"--device={args.device}", + "--measure=dispatch_complete", + f"--batch-size={args.batch_size}", + f"--iterations={args.iterations}", + f"--warmup-iterations={args.warmup_iterations}", + f"--profile-final-batch={str(args.profile_final_batch).lower()}", + f"--artifact-bundle-dir={artifact_dir}", + "--output-format=jsonl", + f"--output={result_output}", + *config_args, + *args.benchmark_extra_arg, + ] + if args.input_ring_count is not None: + benchmark_command.append(f"--input-ring-count={args.input_ring_count}") + benchmark_result = run_command(benchmark_command, REPO_DIR, args.benchmark_timeout_sec, bench_dir, "benchmark") + stage_results["benchmark"] = command_result_json(benchmark_result) + state = benchmark_result.state + error = benchmark_result.error + summary.write( + json.dumps( + { + "benchmark": benchmark, + "kernel": entry["kernel"], + "count": entry_count(entry), + "linked_source": str(linked_source), + "benchmark_source": str(benchmark_source), + "runtime_specialization": runtime_specialized, + "plan": str(plan_output), + "results": str(result_output) if not args.dry_run_only else None, + "state": state, + "error": error, + "stages": stage_results, + }, + sort_keys=True, + ) + + "\n" + ) + summary.flush() + print(f"{state} {benchmark}") + if state != "ok" and not args.continue_on_failure: + fail(f"{benchmark}: {error}") + print(f"wrote {summary_path}") + + +if __name__ == "__main__": + main() diff --git a/ggml/src/ggml-hrx/tools/benchmarks/run-model-benchmarks.sh b/ggml/src/ggml-hrx/tools/benchmarks/run-model-benchmarks.sh new file mode 100755 index 000000000000..b5510cbdae83 --- /dev/null +++ b/ggml/src/ggml-hrx/tools/benchmarks/run-model-benchmarks.sh @@ -0,0 +1,5 @@ +#!/usr/bin/env bash +set -euo pipefail + +script_dir="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +exec python3 "${script_dir}/run-model-benchmarks.py" "$@" diff --git a/ggml/src/ggml-hrx/tools/benchmarks/summarize-model-benchmarks.py b/ggml/src/ggml-hrx/tools/benchmarks/summarize-model-benchmarks.py new file mode 100755 index 000000000000..b1ec1255977b --- /dev/null +++ b/ggml/src/ggml-hrx/tools/benchmarks/summarize-model-benchmarks.py @@ -0,0 +1,220 @@ +#!/usr/bin/env python3 +# +# Summarize generated HRX model-scoped Loom benchmark results. + +from __future__ import annotations + +import argparse +import json +from pathlib import Path +from typing import Any + + +def fail(message: str) -> None: + raise SystemExit(message) + + +def load_jsonl(path: Path) -> list[dict[str, Any]]: + rows = [] + with path.open("r", encoding="utf-8") as f: + for line in f: + if line.strip(): + rows.append(json.loads(line)) + return rows + + +def load_benchmark_result(path: Path) -> dict[str, Any] | None: + with path.open("r", encoding="utf-8", errors="replace") as f: + for line in f: + if not line.strip(): + continue + row = json.loads(line) + if row.get("row") == "benchmark": + return row.get("benchmark_result") + return None + + +def timing_from_result(result: dict[str, Any], metric: str) -> float | None: + if metric == "operation-p50": + return result.get("measurement", {}).get("operation_timing_ns", {}).get("p50") + if metric == "profile-p50": + return ( + result.get("profile_replay", {}) + .get("dispatch_timing", {}) + .get("dispatch_distribution", {}) + .get("duration_ns", {}) + .get("p50") + ) + fail(f"unsupported metric: {metric}") + + +def secondary_profile_p50(result: dict[str, Any]) -> float | None: + return ( + result.get("profile_replay", {}) + .get("dispatch_timing", {}) + .get("dispatch_distribution", {}) + .get("duration_ns", {}) + .get("p50") + ) + + +def summarize(rows: list[dict[str, Any]], metric: str) -> dict[str, Any]: + shape_rows = [] + for row in rows: + count = int(row.get("count", row.get("shape_multiplicity", 1))) + result_path = row.get("results") + shape = { + "benchmark": row.get("benchmark"), + "kernel": row.get("kernel"), + "count": count, + "state": row.get("state", "ok"), + "error": row.get("error"), + "metric_ns": None, + "profile_p50_ns": None, + "weighted_ns": None, + } + if shape["state"] != "ok": + shape_rows.append(shape) + continue + if not result_path: + shape["state"] = "missing_results" + shape["error"] = "runner did not record a results path" + shape_rows.append(shape) + continue + result_file = Path(result_path) + if not result_file.is_file(): + shape["state"] = "missing_results" + shape["error"] = f"results file not found: {result_file}" + shape_rows.append(shape) + continue + result = load_benchmark_result(result_file) + if result is None: + shape["state"] = "missing_benchmark_row" + shape["error"] = "results file did not contain a benchmark row" + shape_rows.append(shape) + continue + if result.get("state") != "ok": + shape["state"] = result.get("state", "failed") + shape["error"] = result.get("failure", {}).get("kind") or "benchmark did not complete successfully" + shape_rows.append(shape) + continue + metric_ns = timing_from_result(result, metric) + if metric_ns is None: + shape["state"] = "missing_metric" + shape["error"] = f"metric not found: {metric}" + shape_rows.append(shape) + continue + shape["metric_ns"] = metric_ns + shape["profile_p50_ns"] = secondary_profile_p50(result) + shape["weighted_ns"] = metric_ns * count + shape_rows.append(shape) + + kernels: dict[str, dict[str, Any]] = {} + for shape in shape_rows: + kernel = shape["kernel"] + entry = kernels.setdefault( + kernel, + { + "kernel": kernel, + "shape_count": 0, + "invocation_count": 0, + "measured_shape_count": 0, + "measured_invocation_count": 0, + "missing_or_failed_count": 0, + "missing_or_failed_invocations": 0, + "weighted_ns": 0.0, + "weighted_percent": 0.0, + "shapes": [], + }, + ) + entry["shape_count"] += 1 + entry["invocation_count"] += shape["count"] + entry["shapes"].append(shape) + if shape["weighted_ns"] is None: + entry["missing_or_failed_count"] += 1 + entry["missing_or_failed_invocations"] += shape["count"] + continue + entry["measured_shape_count"] += 1 + entry["measured_invocation_count"] += shape["count"] + entry["weighted_ns"] += shape["weighted_ns"] + + total_weighted_ns = sum(entry["weighted_ns"] for entry in kernels.values()) + for entry in kernels.values(): + if total_weighted_ns: + entry["weighted_percent"] = entry["weighted_ns"] / total_weighted_ns * 100.0 + kernel_rows = sorted(kernels.values(), key=lambda entry: entry["weighted_ns"], reverse=True) + return { + "schema": "ggml-hrx-model-loom-benchmark-summary-v2", + "metric": metric, + "metric_description": "operation_timing_ns.p50" if metric == "operation-p50" else "profile_replay dispatch_distribution duration_ns p50", + "shape_count": len(shape_rows), + "measured_shape_count": sum(1 for shape in shape_rows if shape["weighted_ns"] is not None), + "missing_or_failed_shape_count": sum(1 for shape in shape_rows if shape["weighted_ns"] is None), + "invocation_count": sum(shape["count"] for shape in shape_rows), + "measured_invocation_count": sum(shape["count"] for shape in shape_rows if shape["weighted_ns"] is not None), + "missing_or_failed_invocations": sum(shape["count"] for shape in shape_rows if shape["weighted_ns"] is None), + "total_weighted_ns": total_weighted_ns, + "kernels": kernel_rows, + } + + +def write_markdown(path: Path, summary: dict[str, Any]) -> None: + lines = [ + "# HRX Model Loom Benchmark Summary", + "", + f"Metric: `{summary['metric_description']}`", + "", + f"Measured shapes: {summary['measured_shape_count']} / {summary['shape_count']}", + f"Measured invocations: {summary['measured_invocation_count']} / {summary['invocation_count']}", + "", + "| Kernel | Shapes | Invocations | Measured Shapes | Weighted ms | Weighted Share | Missing/Failed Shapes |", + "| --- | ---: | ---: | ---: | ---: | ---: | ---: |", + ] + for kernel in summary["kernels"]: + weighted_ms = kernel["weighted_ns"] / 1_000_000.0 + lines.append( + f"| `{kernel['kernel']}` | {kernel['shape_count']} | {kernel['invocation_count']} | " + f"{kernel['measured_shape_count']} | {weighted_ms:.3f} | " + f"{kernel['weighted_percent']:.2f}% | {kernel['missing_or_failed_count']} |" + ) + lines.append("") + lines.append("## Missing Or Failed Shapes") + lines.append("") + lines.append("| Kernel | Benchmark | Invocations | State | Error |") + lines.append("| --- | --- | ---: | --- | --- |") + missing = False + for kernel in summary["kernels"]: + for shape in kernel["shapes"]: + if shape["weighted_ns"] is not None: + continue + missing = True + error = str(shape.get("error") or "").replace("|", "\\|") + lines.append( + f"| `{shape['kernel']}` | `{shape['benchmark']}` | {shape['count']} | " + f"`{shape['state']}` | {error} |" + ) + if not missing: + lines.append("| - | - | 0 | - | - |") + path.write_text("\n".join(lines) + "\n", encoding="utf-8") + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("results_jsonl", type=Path, help="Runner results.jsonl path.") + parser.add_argument("--metric", choices=["operation-p50", "profile-p50"], default="operation-p50") + parser.add_argument("--output-json", type=Path, help="Summary JSON output path.") + parser.add_argument("--output-md", type=Path, help="Summary Markdown output path.") + args = parser.parse_args() + + rows = load_jsonl(args.results_jsonl) + summary = summarize(rows, args.metric) + output_json = args.output_json or (args.results_jsonl.parent / "summary.json") + output_md = args.output_md or (args.results_jsonl.parent / "summary.md") + output_json.write_text(json.dumps(summary, indent=2, sort_keys=True) + "\n", encoding="utf-8") + write_markdown(output_md, summary) + print(f"wrote {output_json}") + print(f"wrote {output_md}") + + +if __name__ == "__main__": + main() diff --git a/ggml/src/ggml-hrx/tools/compile-kernel.cpp b/ggml/src/ggml-hrx/tools/compile-kernel.cpp new file mode 100644 index 000000000000..390688eba0d9 --- /dev/null +++ b/ggml/src/ggml-hrx/tools/compile-kernel.cpp @@ -0,0 +1,110 @@ +#include "../loom-jit.h" +#include "tool-utils.h" + +#include +#include +#include +#include +#include + +using ggml::hrx::tool::report_status; +using ggml::hrx::tool::write_file; + +int main(int argc, char ** argv) { + std::string target; + std::string source_path; + std::string root; + std::string output_path; + std::vector config_storage; + std::vector workload; + for (int i = 1; i < argc; ++i) { + const std::string argument = argv[i]; + if (argument == "--target" && i + 1 < argc) { + target = argv[++i]; + } else if (argument == "--source" && i + 1 < argc) { + source_path = argv[++i]; + } else if (argument == "--root" && i + 1 < argc) { + root = argv[++i]; + } else if (argument == "--output" && i + 1 < argc) { + output_path = argv[++i]; + } else if (argument == "--config" && i + 1 < argc) { + config_storage.emplace_back(argv[++i]); + } else if (argument == "--workload" && i + 1 < argc) { + workload.push_back(std::stoll(argv[++i])); + } else { + std::cerr << "unknown or incomplete argument: " << argument << '\n'; + return 2; + } + } + if (target.empty() || source_path.empty() || root.empty() || output_path.empty()) { + std::cerr << "usage: ggml-hrx-compile-kernel --target gfx... --source linked.loom --root symbol " + "--output dir [--config key=value] [--workload value]\n"; + return 2; + } + const std::string source = ggml::hrx::tool::read_file(source_path); + if (source.empty()) { + std::cerr << "cannot read Loom source: " << source_path << '\n'; + return 2; + } + std::vector configs; + std::vector config_keys; + std::vector config_values; + for (const std::string & item : config_storage) { + const size_t equals = item.find('='); + if (equals == std::string::npos) { + std::cerr << "invalid config binding: " << item << '\n'; + return 2; + } + config_keys.push_back(item.substr(0, equals)); + config_values.push_back(item.substr(equals + 1)); + } + for (size_t i = 0; i < config_keys.size(); ++i) { + configs.push_back({ config_keys[i].c_str(), config_values[i].c_str() }); + } + + ggml_hrx_loom_jit_amdgpu_options jit_options; + jit_options.processor = target.c_str(); + jit_options.identifier = target.c_str(); + ggml_hrx_loom_jit_amdgpu * jit = nullptr; + hrx_status_t status = ggml_hrx_loom_jit_amdgpu_create(&jit_options, &jit); + if (!report_status(status, "create JIT")) { + return 1; + } + + ggml_hrx_loom_jit_compile_options options; + options.source_data = source.data(); + options.source_size = source.size(); + options.source_format = GGML_HRX_LOOM_JIT_SOURCE_FORMAT_TEXT; + options.source_identifier = source_path.c_str(); + options.root_symbol = root.c_str(); + options.module_name = root.c_str(); + options.artifact_identifier = root.c_str(); + options.config_bindings = configs.data(); + options.config_binding_count = configs.size(); + options.workload_arguments = workload.data(); + options.workload_argument_count = workload.size(); + options.evaluate_launch_config = !workload.empty(); + ggml_hrx_loom_jit_compile_result result; + status = ggml_hrx_loom_jit_amdgpu_compile(jit, &options, &result); + if (!hrx_status_is_ok(status)) { + ggml_hrx_loom_jit_amdgpu_release(jit); + report_status(status, "compile kernel"); + return 1; + } + std::filesystem::create_directories(output_path); + write_file(std::filesystem::path(output_path) / "kernel.hsaco", result.hsaco_data, result.hsaco_size); + write_file(std::filesystem::path(output_path) / "manifest.json", result.manifest_json, result.manifest_json_size); + write_file(std::filesystem::path(output_path) / "compile-report.json", result.compile_report_json, + result.compile_report_json_size); + write_file(std::filesystem::path(output_path) / "final.loom", result.final_module_text, + result.final_module_text_size); + std::ofstream launch(std::filesystem::path(output_path) / "launch.txt", std::ios::trunc); + launch << "target=" << target << '\n' + << "root=" << root << '\n' + << "workgroups=" << result.launch_config.workgroup_count[0] << ',' << result.launch_config.workgroup_count[1] + << ',' << result.launch_config.workgroup_count[2] << '\n' + << "workgroup_size=" << result.launch_config.workgroup_size[0] << ',' + << result.launch_config.workgroup_size[1] << ',' << result.launch_config.workgroup_size[2] << '\n'; + ggml_hrx_loom_jit_amdgpu_release(jit); + return 0; +} diff --git a/ggml/src/ggml-hrx/tools/compile_qwen_kernel_corpus.py b/ggml/src/ggml-hrx/tools/compile_qwen_kernel_corpus.py new file mode 100755 index 000000000000..d92136fd3ca2 --- /dev/null +++ b/ggml/src/ggml-hrx/tools/compile_qwen_kernel_corpus.py @@ -0,0 +1,257 @@ +#!/usr/bin/env python3 +"""Compiles the pinned Qwen corpus using only BUILD.bazel-authored recipes.""" + +from __future__ import annotations + +import argparse +import json +import os +import pathlib +import re +import subprocess +import sys + + +BENCHMARK_RE = re.compile( + r"check\.benchmark<@(?P[A-Za-z0-9_]+)>\s+@(?P[A-Za-z0-9_]+)" + r"(?:\s*\{(?P[^}]*)\})?" +) +ATTR_RE = re.compile(r"(?P[A-Za-z0-9_.]+)\s*=\s*(?P-?[0-9]+)") +CASE_RE = re.compile(r"check\.case(?:\s+public)?\s+@(?P[A-Za-z0-9_]+)\s*\{") +LITERAL_RE = re.compile(r"%(?P[A-Za-z0-9_]+)\s*=\s*check\.literal\s+value\((?P-?[0-9]+)\)\s*:\s*index") +FUNC_CALL_RE = re.compile(r"func\.call\s+@(?P[A-Za-z0-9_]+)\s*\(") + + +def run(command: list[str], *, env: dict[str, str] | None = None) -> subprocess.CompletedProcess[str]: + return subprocess.run(command, text=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE, env=env, check=False) + + +def require_run(command: list[str], *, env: dict[str, str] | None = None) -> subprocess.CompletedProcess[str]: + result = run(command, env=env) + if result.returncode: + raise RuntimeError(f"command failed ({result.returncode}): {' '.join(command)}\n{result.stderr}") + return result + + +def main() -> int: + parser = argparse.ArgumentParser() + parser.add_argument("--manifest", type=pathlib.Path, required=True) + parser.add_argument("--corpus-dir", type=pathlib.Path, required=True) + parser.add_argument("--hrx-source", type=pathlib.Path, required=True) + parser.add_argument("--loom-link", type=pathlib.Path, required=True) + parser.add_argument("--benchmark-tool", type=pathlib.Path, required=True) + parser.add_argument("--compiler", type=pathlib.Path, required=True) + parser.add_argument("--target", choices=("gfx1100", "gfx1151"), required=True) + parser.add_argument("--output-dir", type=pathlib.Path, required=True) + parser.add_argument("--schedule", type=pathlib.Path, action="append", default=[], + help="materialized program.json whose exact specializations must compile") + args = parser.parse_args() + + manifest = json.loads(args.manifest.read_text(encoding="utf-8")) + revision = require_run(["git", "-C", str(args.hrx_source), "rev-parse", "HEAD"]).stdout.strip() + upstream_revision = manifest.get("upstream_revision", "local") + if upstream_revision != "local" and revision != upstream_revision: + raise RuntimeError(f"HRX checkout {revision} does not match corpus {upstream_revision}") + if require_run(["git", "-C", str(args.hrx_source), "status", "--porcelain"]).stdout.strip(): + raise RuntimeError("refusing to compile against a dirty pinned HRX checkout") + override = os.environ.get("HSA_OVERRIDE_GFX_VERSION") + if args.target == "gfx1100" and override != "11.0.0": + raise RuntimeError("gfx1100 validation on this gfx1151 host requires HSA_OVERRIDE_GFX_VERSION=11.0.0") + if args.target == "gfx1151" and override: + raise RuntimeError("native gfx1151 validation must run without HSA_OVERRIDE_GFX_VERSION") + + args.output_dir.mkdir(parents=True, exist_ok=True) + modules = {item["name"]: item for item in manifest["link_modules"]} + linked_dir = args.output_dir / "linked" + linked_dir.mkdir(exist_ok=True) + linked_paths: dict[str, pathlib.Path] = {} + + def link_module(name: str) -> pathlib.Path: + if name in linked_paths: + return linked_paths[name] + recipe = modules[name] + output = linked_dir / f"{name}.loom" + command = [str(args.loom_link), "--mode=merge", f"--output={output}"] + command.extend(str(args.corpus_dir / source) for source in recipe["srcs"]) + for library in recipe["libraries"]: + path = link_module(library[1:]) if library.startswith(":") else args.corpus_dir / library + command.append(f"--library={path}") + require_run(command) + linked_paths[name] = output + return output + + exports = {item["symbol"]: item for item in manifest["exports"]} + + def source_for_root(root: str) -> pathlib.Path: + source_name = exports[root]["source"] + direct_modules = [name for name, recipe in modules.items() if source_name in recipe["srcs"]] + if direct_modules: + return link_module(direct_modules[0]) + library_modules = [name for name, recipe in modules.items() if source_name in recipe["libraries"]] + if library_modules: + return link_module(library_modules[0]) + return args.corpus_dir / source_name + + results: list[dict[str, object]] = [] + compiled_keys: set[tuple[str, tuple[int, ...], tuple[str, ...]]] = set() + root_sources: dict[str, pathlib.Path] = {} + planned_invocation_count = 0 + resolved_invocation_count = 0 + for case in manifest["plan_cases"]: + source = link_module(case["link_module"]) if "link_module" in case else args.corpus_dir / case["source"] + plan_command = [str(args.benchmark_tool)] + for item in case["args"]: + if item.startswith("$(location"): + plan_command.append(str(source)) + else: + plan_command.append(item) + plan = require_run(plan_command) + case_dir = args.output_dir / "recipes" / case["name"] + case_dir.mkdir(parents=True, exist_ok=True) + (case_dir / "plan.jsonl").write_text(plan.stdout, encoding="utf-8") + (case_dir / "plan.stderr.txt").write_text(plan.stderr, encoding="utf-8") + plan_rows = [json.loads(line) for line in plan.stdout.splitlines() if line.strip()] + plan_rows = [row for row in plan_rows if row.get("row") == "plan"] + if not plan_rows: + raise RuntimeError(f"BUILD recipe {case['name']} produced no planner rows") + source_text = source.read_text(encoding="utf-8") + benchmark_defs = { + match.group("name"): (match.group("case"), + {item.group("name"): int(item.group("value")) for item in ATTR_RE.finditer(match.group("attrs") or "")}) + for match in BENCHMARK_RE.finditer(source.read_text(encoding="utf-8")) + } + case_literals: dict[str, dict[str, int]] = {} + case_calls: dict[str, list[str]] = {} + for match in CASE_RE.finditer(source_text): + depth = 1 + cursor = match.end() + while cursor < len(source_text) and depth: + depth += source_text[cursor] == "{" + depth -= source_text[cursor] == "}" + cursor += 1 + body = source_text[match.end():cursor - 1] + case_literals[match.group("name")] = { + item.group("name"): int(item.group("value")) for item in LITERAL_RE.finditer(body) + } + case_calls[match.group("name")] = [item.group("name") for item in FUNC_CALL_RE.finditer(body)] + configs = [item.removeprefix("--config=") for item in case["args"] if item.startswith("--config=")] + selects_benchmark = any(item.startswith("--benchmark=") for item in case["args"]) + owned_sources = set(modules[case["link_module"]]["srcs"]) if "link_module" in case else {case["source"]} + for row in plan_rows: + benchmark_case, attrs = benchmark_defs.get(row["benchmark"], (row["case"], {})) + concrete_values = dict(case_literals.get(benchmark_case, {})) + concrete_values.update(attrs) + expected_count = int(row.get("actual_invocation_count", 1)) + planned_invocation_count += expected_count + if row.get("actual_entry"): + roots = [row["actual_entry"]] + else: + roots = [root for root in case_calls.get(benchmark_case, []) if root in exports] + if len(roots) != expected_count: + raise RuntimeError( + f"BUILD recipe {case['name']} planner reports {expected_count} invocations for " + f"{benchmark_case}, but source resolves {len(roots)} exported calls: {roots}" + ) + resolved_invocation_count += len(roots) + for root in roots: + if root not in exports: + raise RuntimeError(f"BUILD recipe {case['name']} selected unmanifested kernel {root}") + if not selects_benchmark and exports[root]["source"] not in owned_sources: + continue + root_sources.setdefault(root, source) + parameter_names = [item["name"] for item in exports[root]["workload_parameters"]] + missing = [name for name in parameter_names if name not in concrete_values] + if missing: + raise RuntimeError(f"{case['name']} recipe omits concrete {missing} for {root}") + workload = tuple(concrete_values[name] for name in parameter_names) + key = (root, workload, tuple(configs)) + if key in compiled_keys: + continue + compiled_keys.add(key) + compile_dir = case_dir / f"{root}-{len(compiled_keys):03d}" + command = [str(args.compiler), "--target", args.target, "--source", str(source), + "--root", root, "--output", str(compile_dir)] + for config in configs: + command.extend(("--config", config)) + for value in workload: + command.extend(("--workload", str(value))) + compile_result = run(command, env=os.environ.copy()) + (case_dir / f"{root}-{len(compiled_keys):03d}.stdout.txt").write_text(compile_result.stdout, encoding="utf-8") + (case_dir / f"{root}-{len(compiled_keys):03d}.stderr.txt").write_text(compile_result.stderr, encoding="utf-8") + result = {"case": case["name"], "root": root, "workload": workload, + "configs": configs, "artifact_dir": str(compile_dir), + "status": "ok" if compile_result.returncode == 0 else "failed"} + results.append(result) + + schedule_requirement_count = 0 + schedule_unique_requirement_count = 0 + for schedule_path in args.schedule: + schedule = json.loads(schedule_path.read_text(encoding="utf-8")) + schedule_dir = args.output_dir / "schedules" / schedule_path.parent.name + schedule_dir.mkdir(parents=True, exist_ok=True) + for invocation in schedule.get("invocations", []): + for dispatch in invocation.get("dispatches", []): + specialization = dispatch["kernel"] + if specialization.get("execution") != "native": + raise RuntimeError(f"schedule {schedule_path} contains non-native dispatch {specialization['variant']}") + root = specialization["variant"] + if root not in exports: + raise RuntimeError(f"schedule {schedule_path} selects unmanifested kernel {root}") + schedule_requirement_count += 1 + parameters = specialization.get("parameters", {}) + parameter_names = [item["name"] for item in exports[root]["workload_parameters"]] + missing = [name for name in parameter_names if name not in parameters] + if missing: + raise RuntimeError(f"schedule {schedule_path} omits concrete {missing} for {root}") + workload = tuple(int(parameters[name]) for name in parameter_names) + configs = tuple(f"{name}={value}" for name, value in + sorted(specialization.get("compile_parameters", {}).items())) + key = (root, workload, configs) + if key in compiled_keys: + continue + schedule_unique_requirement_count += 1 + compiled_keys.add(key) + source = root_sources[root] if root in root_sources else source_for_root(root) + compile_dir = schedule_dir / f"{root}-{schedule_unique_requirement_count:03d}" + command = [str(args.compiler), "--target", args.target, "--source", str(source), + "--root", root, "--output", str(compile_dir)] + for config in configs: + command.extend(("--config", config)) + for value in workload: + command.extend(("--workload", str(value))) + compile_result = run(command, env=os.environ.copy()) + (schedule_dir / f"{root}-{schedule_unique_requirement_count:03d}.stdout.txt").write_text( + compile_result.stdout, encoding="utf-8") + (schedule_dir / f"{root}-{schedule_unique_requirement_count:03d}.stderr.txt").write_text( + compile_result.stderr, encoding="utf-8") + results.append({"case": f"schedule:{schedule_path}", "root": root, "workload": workload, + "configs": configs, "artifact_dir": str(compile_dir), + "status": "ok" if compile_result.returncode == 0 else "failed"}) + + summary = { + "schema": "ggml-hrx-qwen-compile-report-v1", + "target": args.target, + "hsa_override_gfx_version": override, + "hrx_revision": revision, + "corpus_digest": manifest.get("corpus_sha256", ""), + "recipe_digest": manifest.get("build_bazel_sha256", ""), + "plan_case_count": len(manifest["plan_cases"]), + "planned_invocation_count": planned_invocation_count, + "resolved_invocation_count": resolved_invocation_count, + "schedule_requirement_count": schedule_requirement_count, + "schedule_unique_requirement_count": schedule_unique_requirement_count, + "compile_count": len(results), + "failed_count": sum(item["status"] != "ok" for item in results), + "results": results, + } + (args.output_dir / "summary.json").write_text(json.dumps(summary, indent=2) + "\n", encoding="utf-8") + print(json.dumps({key: summary[key] for key in ("target", "plan_case_count", "compile_count", "failed_count")})) + return 0 if summary["failed_count"] == 0 and planned_invocation_count == resolved_invocation_count else 1 + + +if __name__ == "__main__": + try: + raise SystemExit(main()) + except (OSError, RuntimeError, ValueError, json.JSONDecodeError) as error: + print(f"compile corpus: {error}", file=sys.stderr) + raise SystemExit(2) diff --git a/ggml/src/ggml-hrx/tools/generate_iq1_loom.py b/ggml/src/ggml-hrx/tools/generate_iq1_loom.py new file mode 100644 index 000000000000..a7d5c2e61fa9 --- /dev/null +++ b/ggml/src/ggml-hrx/tools/generate_iq1_loom.py @@ -0,0 +1,589 @@ +#!/usr/bin/env python3 +# Copyright 2026 bong-water-water-bong +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# IQ1_S / IQ1_M grid lookups and decoders for motifs/dequant.loom (mode "dequant") and lane functions +# and grid fill for ops/kquant_decode_f32.loom (mode "kquant"), generated from ggml-common.h +# (iq1s_grid: MIT, the ggml authors). Every grid byte is -1, 0 or 1, so an entry is stored as a +# 16-bit code: 2 bits per value, value + 1. +# Usage: generate_iq1_loom.py ggml/src/ggml-common.h dequant|kquant > out.loom +# The output sits verbatim in kernel-corpus/kernels/loom-libs/motifs/dequant.loom (dequant) or +# kernel-corpus/kernels/loom-libs/ops/kquant_decode_f32.loom (kquant). +import re +import sys + +src = open(sys.argv[1]).read() +mode = sys.argv[2] + + +def grid(): + m = re.search(r'GGML_TABLE_BEGIN\(uint64_t, iq1s_grid, NGRID_IQ1S\)(.*?)GGML_TABLE_END', src, re.S) + assert m, "iq1s_grid" + v = [int(x, 16) for x in re.findall(r'0x[0-9a-fA-F]+', m.group(1))] + assert len(v) == 2048, len(v) + lv = {0xff: 0, 0x00: 1, 0x01: 2} + return [sum(lv[(e >> (8 * j)) & 255] << (2 * j) for j in range(8)) for e in v] + + +def lookup(fname, prefix, codes): + n = len(codes) // 32 + o = [f'// {prefix}: {len(codes)} grid entries as 16-bit codes (2 bits per value: value + 1)', + f'func.def inline @{fname}(%grid_index: i32) -> (i32) {{', + ' %c5_i32 = scalar.constant 5 : i32', ' %c31_i32 = scalar.constant 31 : i32'] + for k in range(1, n): + o.append(f' %chunk_id{k} = scalar.constant {k} : i32') + o += [' %chunk_i32 = scalar.shrui %grid_index, %c5_i32 : i32', + ' %lane_i32 = scalar.andi %grid_index, %c31_i32 : i32', + ' %codes = vector.from_elements %lane_i32 : vector<1xi32>'] + for k in range(n): + names = [] + for i in range(32): + nm = f'%{prefix}_{k}_{i}' + o.append(f' {nm} = scalar.constant {float(codes[32 * k + i])} : f32') + names.append(nm) + o.append(f' %{prefix}_{k} = vector.from_elements {", ".join(names)} : vector<32xf32>') + prev = f'%{prefix}_0' + for k in range(1, n): + o.append(f' %is_chunk{k} = scalar.cmpi eq, %chunk_i32, %chunk_id{k} : i32') + o.append(f' %sel{k} = scf.select %is_chunk{k}, %{prefix}_{k}, {prev} : vector<32xf32>') + prev = f'%sel{k}' + o += [f' %v = vector.table.lookup {prev}[%codes] : vector<32xf32>, vector<1xi32> -> vector<1xf32>', + ' %v_f32 = vector.extract %v[0] : vector<1xf32> -> f32', + ' %code = scalar.fptoui %v_f32 : f32 to i32', + ' func.return %code : i32', '}', ''] + return '\n'.join(o) + + +dequant_common = ''' +// four values 4 (p % 2) .. +3 of an IQ1 grid code: ((code >> 2j) & 3) - 1 + delta, times %scale +func.def inline @ggml_iq1_code_vector4(%code: i32, %half: i32, %delta: f32, %scale: f32) -> (vector<4xf32>) { + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c2v = scalar.constant 2 : i32 + %c4v = scalar.constant 4 : i32 + %c6v = scalar.constant 6 : i32 + %c0v = scalar.constant 0 : i32 + %base = scalar.muli %half, %c8_i32 : i32 + %cv = vector.splat %code : vector<4xi32> + %bv = vector.splat %base : vector<4xi32> + %sh0 = vector.from_elements %c0v, %c2v, %c4v, %c6v : vector<4xi32> + %sh = vector.addi %sh0, %bv : vector<4xi32> + %s = vector.shrui %cv, %sh : vector<4xi32> + %three = vector.splat %c3_i32 : vector<4xi32> + %one = vector.splat %c1_i32 : vector<4xi32> + %c = vector.andi %s, %three : vector<4xi32> + %t = vector.subi %c, %one : vector<4xi32> + %tf = vector.sitofp %t : vector<4xi32> to vector<4xf32> + %dv = vector.splat %delta : vector<4xf32> + %sv = vector.splat %scale : vector<4xf32> + %td = vector.addf %tf, %dv : vector<4xf32> + %result = vector.mulf %td, %sv : vector<4xf32> + func.return %result : vector<4xf32> +} + +// IQ1_S (50 bytes: d, qs[32], qh[8] u16; dequantize_row_iq1_s). Group g: qh[g] holds three 3-bit +// high index parts (slot l at bit 3l), a 3-bit scale at bit 12 and the delta sign at bit 15: +// value = d * (2 s + 1) * (grid + delta), delta = +-0.125. Packet p covers slot p / 2. +func.def inline @ggml_iq1s_f32_vector4(%weight: buffer, %row_byte_base: offset, %iq1_block: index, %iq1_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c17 = index.constant 17 : index + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c12_i32 = scalar.constant 12 : i32 + %c32768_i32 = scalar.constant 32768 : i32 + %c0_i32 = scalar.constant 0 : i32 + %pos_delta = scalar.constant 0.125 : f32 + %neg_delta = scalar.constant -0.125 : f32 + %block_bytes = index.constant 50 : offset + %block_byte_add = index.scale %iq1_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %hv = buffer.view %weight[%block_byte_base] : buffer -> view<25xf16> + %wv = buffer.view %weight[%block_byte_base] : buffer -> view<25xi16> + %bv = buffer.view %weight[%block_byte_base] : buffer -> view<50xi8> + %g = index.assume %iq1_group [range(%iq1_group, 0, 7)] : index + %p = index.assume %packet [range(%packet, 0, 7)] : index + %slot = index.div %p, %c2 : index + %half = index.rem %p, %c2 : index + %g4 = index.mul %g, %c4 : index + %qs_at0 = index.add %g4, %c2 : index + %qs_at = index.add %qs_at0, %slot : index + %qh_at = index.add %c17, %g : index + %d_f16 = view.load %hv[%c0] : view<25xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %qs_i8 = view.load %bv[%qs_at] : view<50xi8> -> i8 + %qs = scalar.extui %qs_i8 : i8 to i32 + %qh_i16 = view.load %wv[%qh_at] : view<25xi16> -> i16 + %qh = scalar.extui %qh_i16 : i16 to i32 + %slot_i32 = index.cast %slot : index to i32 + %hsh = scalar.muli %slot_i32, %c3_i32 : i32 + %hi0 = scalar.shrui %qh, %hsh : i32 + %hi1 = scalar.andi %hi0, %c7_i32 : i32 + %hi = scalar.shli %hi1, %c8_i32 : i32 + %gi = scalar.ori %qs, %hi : i32 + %s0 = scalar.shrui %qh, %c12_i32 : i32 + %s = scalar.andi %s0, %c7_i32 : i32 + %s2 = scalar.addi %s, %s : i32 + %s21 = scalar.addi %s2, %c1_i32 : i32 + %sf = scalar.uitofp %s21 : i32 to f32 + %scale = scalar.mulf %d, %sf : f32 + %neg_bit = scalar.andi %qh, %c32768_i32 : i32 + %neg = scalar.cmpi ne, %neg_bit, %c0_i32 : i32 + %delta = scf.select %neg, %neg_delta, %pos_delta : f32 + %code = func.call @ggml_iq1s_grid_code_i32(%gi) : (i32) -> (i32) + %half_i32 = index.cast %half : index to i32 + %result = func.call @ggml_iq1_code_vector4(%code, %half_i32, %delta, %scale) : (i32, i32, f32, f32) -> (vector<4xf32>) + func.return %result : vector<4xf32> +} + +func.def inline @ggml_iq1s_f16_vector4(%weight: buffer, %row_byte_base: offset, %iq1_block: index, %iq1_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_iq1s_f32_vector4(%weight, %row_byte_base, %iq1_block, %iq1_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +// IQ1_M's fp16 block scale: the top nibbles of its four u16 scale words (bytes 48..55). +func.def inline @ggml_iq1m_block_scale(%weight: buffer, %block_byte_base: offset) -> (f32) { + %wv = buffer.view %weight[%block_byte_base] : buffer -> view<28xi16> + %c24 = index.constant 24 : index + %c25 = index.constant 25 : index + %c26 = index.constant 26 : index + %c27 = index.constant 27 : index + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c12_i32 = scalar.constant 12 : i32 + %c240_i32 = scalar.constant 240 : i32 + %c3840_i32 = scalar.constant 3840 : i32 + %c61440_i32 = scalar.constant 61440 : i32 + %c255_i32 = scalar.constant 255 : i32 + %w0_i16 = view.load %wv[%c24] : view<28xi16> -> i16 + %w1_i16 = view.load %wv[%c25] : view<28xi16> -> i16 + %w2_i16 = view.load %wv[%c26] : view<28xi16> -> i16 + %w3_i16 = view.load %wv[%c27] : view<28xi16> -> i16 + %w0 = scalar.extui %w0_i16 : i16 to i32 + %w1 = scalar.extui %w1_i16 : i16 to i32 + %w2 = scalar.extui %w2_i16 : i16 to i32 + %w3 = scalar.extui %w3_i16 : i16 to i32 + %n0 = scalar.shrui %w0, %c12_i32 : i32 + %n1a = scalar.shrui %w1, %c8_i32 : i32 + %n1 = scalar.andi %n1a, %c240_i32 : i32 + %n2a = scalar.shrui %w2, %c4_i32 : i32 + %n2 = scalar.andi %n2a, %c3840_i32 : i32 + %n3 = scalar.andi %w3, %c61440_i32 : i32 + %u01 = scalar.ori %n0, %n1 : i32 + %u012 = scalar.ori %u01, %n2 : i32 + %u = scalar.ori %u012, %n3 : i32 + %lo = scalar.andi %u, %c255_i32 : i32 + %hi = scalar.shrui %u, %c8_i32 : i32 + %lo8 = scalar.trunci %lo : i32 to i8 + %hi8 = scalar.trunci %hi : i32 to i8 + %bytes = vector.from_elements %lo8, %hi8 : vector<2xi8> + %h = vector.bitcast %bytes : vector<2xi8> to vector<1xf16> + %h0 = vector.extract %h[0] : vector<1xf16> -> f16 + %f = scalar.extf %h0 : f16 to f32 + func.return %f : f32 +} + +// IQ1_M (56 bytes: qs[32], qh[16], scales[8]; dequantize_row_iq1_m). Slot l of group g: qh byte +// 2g + l / 2 holds the high index bits (bits 0..2 for even l, 4..6 for odd) and the delta sign +// (bit 3 / bit 7); the 3-bit scale of slots 0-1 / 2-3 sits at bit 6 (g % 2) / +3 of u16 scale +// word g / 2: value = d * (2 s + 1) * (grid + delta). +func.def inline @ggml_iq1m_f32_vector4(%weight: buffer, %row_byte_base: offset, %iq1_block: index, %iq1_group: index, %packet: index) -> (vector<4xf32>) { + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c24 = index.constant 24 : index + %c32 = index.constant 32 : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c6_i32 = scalar.constant 6 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c1792_i32 = scalar.constant 1792 : i32 + %pos_delta = scalar.constant 0.125 : f32 + %neg_delta = scalar.constant -0.125 : f32 + %block_bytes = index.constant 56 : offset + %block_byte_add = index.scale %iq1_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %wv = buffer.view %weight[%block_byte_base] : buffer -> view<28xi16> + %bv = buffer.view %weight[%block_byte_base] : buffer -> view<56xi8> + %g = index.assume %iq1_group [range(%iq1_group, 0, 7)] : index + %p = index.assume %packet [range(%packet, 0, 7)] : index + %slot = index.div %p, %c2 : index + %half = index.rem %p, %c2 : index + %pair = index.div %slot, %c2 : index + %odd = index.rem %slot, %c2 : index + %g4 = index.mul %g, %c4 : index + %qs_at = index.add %g4, %slot : index + %g2 = index.add %g, %g : index + %qh_at0 = index.add %c32, %g2 : index + %qh_at = index.add %qh_at0, %pair : index + %gh = index.div %g, %c2 : index + %gp = index.rem %g, %c2 : index + %sc_at = index.add %c24, %gh : index + %d = func.call @ggml_iq1m_block_scale(%weight, %block_byte_base) : (buffer, offset) -> (f32) + %qs_i8 = view.load %bv[%qs_at] : view<56xi8> -> i8 + %qs = scalar.extui %qs_i8 : i8 to i32 + %qh_i8 = view.load %bv[%qh_at] : view<56xi8> -> i8 + %qh = scalar.extui %qh_i8 : i8 to i32 + %odd_i32 = index.cast %odd : index to i32 + %odd4 = scalar.muli %odd_i32, %c4_i32 : i32 + %qh_n = scalar.shrui %qh, %odd4 : i32 + %hi0 = scalar.shli %qh_n, %c8_i32 : i32 + %hi = scalar.andi %hi0, %c1792_i32 : i32 + %gi = scalar.ori %qs, %hi : i32 + %neg_bit0 = scalar.shrui %qh_n, %c3_i32 : i32 + %neg_bit = scalar.andi %neg_bit0, %c1_i32 : i32 + %neg = scalar.cmpi ne, %neg_bit, %c0_i32 : i32 + %delta = scf.select %neg, %neg_delta, %pos_delta : f32 + %sc_i16 = view.load %wv[%sc_at] : view<28xi16> -> i16 + %sc = scalar.extui %sc_i16 : i16 to i32 + %gp_i32 = index.cast %gp : index to i32 + %pair_i32 = index.cast %pair : index to i32 + %ssh0 = scalar.muli %gp_i32, %c6_i32 : i32 + %ssh1 = scalar.muli %pair_i32, %c3_i32 : i32 + %ssh = scalar.addi %ssh0, %ssh1 : i32 + %s0 = scalar.shrui %sc, %ssh : i32 + %s = scalar.andi %s0, %c7_i32 : i32 + %s2 = scalar.addi %s, %s : i32 + %s21 = scalar.addi %s2, %c1_i32 : i32 + %sf = scalar.uitofp %s21 : i32 to f32 + %scale = scalar.mulf %d, %sf : f32 + %code = func.call @ggml_iq1s_grid_code_i32(%gi) : (i32) -> (i32) + %half_i32 = index.cast %half : index to i32 + %result = func.call @ggml_iq1_code_vector4(%code, %half_i32, %delta, %scale) : (i32, i32, f32, f32) -> (vector<4xf32>) + func.return %result : vector<4xf32> +} + +func.def inline @ggml_iq1m_f16_vector4(%weight: buffer, %row_byte_base: offset, %iq1_block: index, %iq1_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_iq1m_f32_vector4(%weight, %row_byte_base, %iq1_block, %iq1_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} +''' + + +def fill(fname, codes): + words = [codes[2 * w] | (codes[2 * w + 1] << 16) for w in range(len(codes) // 2)] + words = [w - (1 << 32) if w >= 1 << 31 else w for w in words] # i32 constants are signed + per = len(words) // 4 + o = [f'// Stages the IQ1 grid as {len(words)} words (two 16-bit codes each); subgroup %chunk (0..3) writes', + f'// words {per} %chunk .. +{per - 1}.', + f'func.def inline @{fname}(%grid: buffer, %chunk: index) {{', + ' %zero_offset = index.constant 0 : offset', + f' %gv = buffer.view %grid[%zero_offset] : buffer -> view<{len(words)}xi32>'] + for c in range(4): + o += [f' %k{c} = index.constant {c} : index', f' %is{c} = index.cmp eq, %chunk, %k{c} : index', f' scf.if %is{c} {{'] + for q in range(per // 4): + e = per * c + 4 * q + for i in range(4): + o.append(f' %v{c}_{q}_{i} = scalar.constant {words[e + i]} : i32') + o.append(f' %w{c}_{q} = vector.from_elements %v{c}_{q}_0, %v{c}_{q}_1, %v{c}_{q}_2, %v{c}_{q}_3 : vector<4xi32>') + o.append(f' %o{c}_{q} = index.constant {e} : index') + o.append(f' vector.store %w{c}_{q}, %gv[%o{c}_{q}] : vector<4xi32>, view<{len(words)}xi32>') + o.append(' }') + o += [' func.return', '}', ''] + return '\n'.join(o) + + +kquant_common = ''' +// IQ1_M's fp16 block scale (as motifs/dequant.loom's ggml_iq1m_block_scale). +func.def inline @ggml_kquant_iq1m_block_scale(%weight: buffer, %block_byte_base: offset) -> (f32) { + %wv = buffer.view %weight[%block_byte_base] : buffer -> view<28xi16> + %c24 = index.constant 24 : index + %c25 = index.constant 25 : index + %c26 = index.constant 26 : index + %c27 = index.constant 27 : index + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c12_i32 = scalar.constant 12 : i32 + %c240_i32 = scalar.constant 240 : i32 + %c3840_i32 = scalar.constant 3840 : i32 + %c61440_i32 = scalar.constant 61440 : i32 + %c255_i32 = scalar.constant 255 : i32 + %w0_i16 = view.load %wv[%c24] : view<28xi16> -> i16 + %w1_i16 = view.load %wv[%c25] : view<28xi16> -> i16 + %w2_i16 = view.load %wv[%c26] : view<28xi16> -> i16 + %w3_i16 = view.load %wv[%c27] : view<28xi16> -> i16 + %w0 = scalar.extui %w0_i16 : i16 to i32 + %w1 = scalar.extui %w1_i16 : i16 to i32 + %w2 = scalar.extui %w2_i16 : i16 to i32 + %w3 = scalar.extui %w3_i16 : i16 to i32 + %n0 = scalar.shrui %w0, %c12_i32 : i32 + %n1a = scalar.shrui %w1, %c8_i32 : i32 + %n1 = scalar.andi %n1a, %c240_i32 : i32 + %n2a = scalar.shrui %w2, %c4_i32 : i32 + %n2 = scalar.andi %n2a, %c3840_i32 : i32 + %n3 = scalar.andi %w3, %c61440_i32 : i32 + %u01 = scalar.ori %n0, %n1 : i32 + %u012 = scalar.ori %u01, %n2 : i32 + %u = scalar.ori %u012, %n3 : i32 + %lo = scalar.andi %u, %c255_i32 : i32 + %hi = scalar.shrui %u, %c8_i32 : i32 + %lo8 = scalar.trunci %lo : i32 to i8 + %hi8 = scalar.trunci %hi : i32 to i8 + %bytes = vector.from_elements %lo8, %hi8 : vector<2xi8> + %h = vector.bitcast %bytes : vector<2xi8> to vector<1xf16> + %h0 = vector.extract %h[0] : vector<1xf16> -> f16 + %f = scalar.extf %h0 : f16 to f32 + func.return %f : f32 +} + +// IQ1_S / IQ1_M: lane l16 owns group l16 / 2, slots 2 (l16 % 2) and +1: values 32 (l16 / 2) + +// 16 (l16 % 2) .. +15, one scale. The grid (2048 16-bit codes, 2 bits per value: value + 1) is +// staged in workgroup memory two codes per word. +func.def inline @ggml_kquant_iq1_code(%grid: buffer, %index: i32) -> (i32) { + %zero_offset = index.constant 0 : offset + %gv = buffer.view %grid[%zero_offset] : buffer -> view<1024xi32> + %c1_i32 = scalar.constant 1 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c65535_i32 = scalar.constant 65535 : i32 + %w_i32 = scalar.shrui %index, %c1_i32 : i32 + %w_x = index.cast %w_i32 : i32 to index + %w_b = index.assume %w_x [range(%w_x, 0, 1023)] : index + %word = view.load %gv[%w_b] : view<1024xi32> -> i32 + %odd = scalar.andi %index, %c1_i32 : i32 + %sh = scalar.shli %odd, %c4_i32 : i32 + %shifted = scalar.shrui %word, %sh : i32 + %code = scalar.andi %shifted, %c65535_i32 : i32 + func.return %code : i32 +} + +func.def inline @ggml_kquant_iq1_slot_values(%code: i32, %delta: f32) -> (vector<8xf32>) { + %c1_i32 = scalar.constant 1 : i32 + %s0 = scalar.constant 0 : i32 + %s2 = scalar.constant 2 : i32 + %s3 = scalar.constant 3 : i32 + %s4 = scalar.constant 4 : i32 + %s6 = scalar.constant 6 : i32 + %s8 = scalar.constant 8 : i32 + %s10 = scalar.constant 10 : i32 + %s12 = scalar.constant 12 : i32 + %s14 = scalar.constant 14 : i32 + %shift2 = vector.from_elements %s0, %s2, %s4, %s6, %s8, %s10, %s12, %s14 : vector<8xi32> + %three = vector.splat %s3 : vector<8xi32> + %one = vector.splat %c1_i32 : vector<8xi32> + %cv = vector.splat %code : vector<8xi32> + %cs = vector.shrui %cv, %shift2 : vector<8xi32> + %c = vector.andi %cs, %three : vector<8xi32> + %t = vector.subi %c, %one : vector<8xi32> + %tf = vector.sitofp %t : vector<8xi32> to vector<8xf32> + %dv = vector.splat %delta : vector<8xf32> + %v = vector.addf %tf, %dv : vector<8xf32> + func.return %v : vector<8xf32> +} + +// IQ1_S (50 bytes: d, qs[32], qh[8] u16): qh[g] = high index bits (slot l at bit 3l), scale s at +// bit 12, delta sign at bit 15: value = d * (2 s + 1) * (grid + delta). +func.def inline @ggml_kquant_iq1s_lane_parts(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c17 = index.constant 17 : index + %c32 = index.constant 32 : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c12_i32 = scalar.constant 12 : i32 + %c32768_i32 = scalar.constant 32768 : i32 + %pos_delta = scalar.constant 0.125 : f32 + %neg_delta = scalar.constant -0.125 : f32 + %block_bytes = index.constant 50 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %g = index.div %l, %c2 : index + %h = index.rem %l, %c2 : index + %sa = index.mul %h, %c2 : index + %sb = index.add %sa, %c1 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %hv = buffer.view %weight[%block_base] : buffer -> view<25xf16> + %wv = buffer.view %weight[%block_base] : buffer -> view<25xi16> + %bv = buffer.view %weight[%block_base] : buffer -> view<50xi8> + %d_f16 = view.load %hv[%c0] : view<25xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %qh_at = index.add %c17, %g : index + %qh_i16 = view.load %wv[%qh_at] : view<25xi16> -> i16 + %qh = scalar.extui %qh_i16 : i16 to i32 + %g4 = index.mul %g, %c4 : index + %qs0 = index.add %g4, %c2 : index + %qa_at = index.add %qs0, %sa : index + %qb_at = index.add %qs0, %sb : index + %qa_i8 = view.load %bv[%qa_at] : view<50xi8> -> i8 + %qb_i8 = view.load %bv[%qb_at] : view<50xi8> -> i8 + %qa = scalar.extui %qa_i8 : i8 to i32 + %qb = scalar.extui %qb_i8 : i8 to i32 + %sa_i32 = index.cast %sa : index to i32 + %sb_i32 = index.cast %sb : index to i32 + %sha = scalar.muli %sa_i32, %c3_i32 : i32 + %shb = scalar.muli %sb_i32, %c3_i32 : i32 + %ha0 = scalar.shrui %qh, %sha : i32 + %hb0 = scalar.shrui %qh, %shb : i32 + %ha1 = scalar.andi %ha0, %c7_i32 : i32 + %hb1 = scalar.andi %hb0, %c7_i32 : i32 + %ha = scalar.shli %ha1, %c8_i32 : i32 + %hb = scalar.shli %hb1, %c8_i32 : i32 + %ia = scalar.ori %qa, %ha : i32 + %ib = scalar.ori %qb, %hb : i32 + %code_a = func.call @ggml_kquant_iq1_code(%grid, %ia) : (buffer, i32) -> (i32) + %code_b = func.call @ggml_kquant_iq1_code(%grid, %ib) : (buffer, i32) -> (i32) + %s0 = scalar.shrui %qh, %c12_i32 : i32 + %s = scalar.andi %s0, %c7_i32 : i32 + %s2 = scalar.addi %s, %s : i32 + %s21 = scalar.addi %s2, %c1_i32 : i32 + %sf = scalar.uitofp %s21 : i32 to f32 + %scale = scalar.mulf %d, %sf : f32 + %neg_bit = scalar.andi %qh, %c32768_i32 : i32 + %neg = scalar.cmpi ne, %neg_bit, %c0_i32 : i32 + %delta = scf.select %neg, %neg_delta, %pos_delta : f32 + %va = func.call @ggml_kquant_iq1_slot_values(%code_a, %delta) : (i32, f32) -> (vector<8xf32>) + %vb = func.call @ggml_kquant_iq1_slot_values(%code_b, %delta) : (i32, f32) -> (vector<8xf32>) + %v = func.call @ggml_kquant_iq2_join16(%va, %vb) : (vector<8xf32>, vector<8xf32>) -> (vector<16xf32>) + %g32 = index.mul %g, %c32 : index + %h16 = index.mul %h, %c16 : index + %p0 = index.add %g32, %h16 : index + func.return %v, %scale, %p0 : vector<16xf32>, f32, index +} + +// IQ1_M (56 bytes: qs[32], qh[16], scales[8]): the lane's two slots share qh byte 2g + h (index +// high bits 0..2 / 4..6, delta signs bit 3 / 7) and the 3-bit scale at bit 6 (g % 2) + 3h of u16 +// scale word g / 2; d is the fp16 made of the four scale words' top nibbles. +func.def inline @ggml_kquant_iq1m_lane_parts(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c24 = index.constant 24 : index + %c32 = index.constant 32 : index + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c6_i32 = scalar.constant 6 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c128_i32 = scalar.constant 128 : i32 + %c1792_i32 = scalar.constant 1792 : i32 + %pos_delta = scalar.constant 0.125 : f32 + %neg_delta = scalar.constant -0.125 : f32 + %block_bytes = index.constant 56 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %g = index.div %l, %c2 : index + %h = index.rem %l, %c2 : index + %sa = index.mul %h, %c2 : index + %sb = index.add %sa, %c1 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %wv = buffer.view %weight[%block_base] : buffer -> view<28xi16> + %bv = buffer.view %weight[%block_base] : buffer -> view<56xi8> + %d = func.call @ggml_kquant_iq1m_block_scale(%weight, %block_base) : (buffer, offset) -> (f32) + %g4 = index.mul %g, %c4 : index + %qa_at = index.add %g4, %sa : index + %qb_at = index.add %g4, %sb : index + %qa_i8 = view.load %bv[%qa_at] : view<56xi8> -> i8 + %qb_i8 = view.load %bv[%qb_at] : view<56xi8> -> i8 + %qa = scalar.extui %qa_i8 : i8 to i32 + %qb = scalar.extui %qb_i8 : i8 to i32 + %g2 = index.add %g, %g : index + %qh_at0 = index.add %c32, %g2 : index + %qh_at = index.add %qh_at0, %h : index + %qh_i8 = view.load %bv[%qh_at] : view<56xi8> -> i8 + %qh = scalar.extui %qh_i8 : i8 to i32 + %ha0 = scalar.shli %qh, %c8_i32 : i32 + %ha = scalar.andi %ha0, %c1792_i32 : i32 + %hb0 = scalar.shli %qh, %c4_i32 : i32 + %hb = scalar.andi %hb0, %c1792_i32 : i32 + %ia = scalar.ori %qa, %ha : i32 + %ib = scalar.ori %qb, %hb : i32 + %code_a = func.call @ggml_kquant_iq1_code(%grid, %ia) : (buffer, i32) -> (i32) + %code_b = func.call @ggml_kquant_iq1_code(%grid, %ib) : (buffer, i32) -> (i32) + %na0 = scalar.andi %qh, %c8_i32 : i32 + %nb0 = scalar.andi %qh, %c128_i32 : i32 + %na = scalar.cmpi ne, %na0, %c0_i32 : i32 + %nb = scalar.cmpi ne, %nb0, %c0_i32 : i32 + %delta_a = scf.select %na, %neg_delta, %pos_delta : f32 + %delta_b = scf.select %nb, %neg_delta, %pos_delta : f32 + %gh = index.div %g, %c2 : index + %gp = index.rem %g, %c2 : index + %sc_at = index.add %c24, %gh : index + %sc_i16 = view.load %wv[%sc_at] : view<28xi16> -> i16 + %sc = scalar.extui %sc_i16 : i16 to i32 + %gp_i32 = index.cast %gp : index to i32 + %h_i32 = index.cast %h : index to i32 + %ssh0 = scalar.muli %gp_i32, %c6_i32 : i32 + %ssh1 = scalar.muli %h_i32, %c3_i32 : i32 + %ssh = scalar.addi %ssh0, %ssh1 : i32 + %s0 = scalar.shrui %sc, %ssh : i32 + %s = scalar.andi %s0, %c7_i32 : i32 + %s2 = scalar.addi %s, %s : i32 + %s21 = scalar.addi %s2, %c1_i32 : i32 + %sf = scalar.uitofp %s21 : i32 to f32 + %scale = scalar.mulf %d, %sf : f32 + %va = func.call @ggml_kquant_iq1_slot_values(%code_a, %delta_a) : (i32, f32) -> (vector<8xf32>) + %vb = func.call @ggml_kquant_iq1_slot_values(%code_b, %delta_b) : (i32, f32) -> (vector<8xf32>) + %v = func.call @ggml_kquant_iq2_join16(%va, %vb) : (vector<8xf32>, vector<8xf32>) -> (vector<16xf32>) + %g32 = index.mul %g, %c32 : index + %h16 = index.mul %h, %c16 : index + %p0 = index.add %g32, %h16 : index + func.return %v, %scale, %p0 : vector<16xf32>, f32, index +} +''' + + +def dot_and_weights(fmt): + return f''' +func.def inline @ggml_kquant_{fmt}_lane_dot(%grid: buffer, %weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) {{ + %zero_scalar = scalar.constant 0.0 : f32 + %x_block_bytes = index.constant 1024 : offset + %v, %scale, %p = func.call @ggml_kquant_{fmt}_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %x = vector.load %xv[%p] : view<256xf32> -> vector<16xf32> + %vx = vector.mulf %v, %x : vector<16xf32> + %sum = vector.reduce %vx, %zero_scalar : vector<16xf32>, f32 + %result = scalar.mulf %scale, %sum : f32 + func.return %result : f32 +}} + +func.def inline @ggml_kquant_{fmt}_lane_weights(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) {{ + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c12 = index.constant 12 : index + %v, %scale, %p0 = func.call @ggml_kquant_{fmt}_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %sv = vector.splat %scale : vector<16xf32> + %w = vector.mulf %v, %sv : vector<16xf32> + %p1 = index.add %p0, %c4 : index + %p2 = index.add %p0, %c8 : index + %p3 = index.add %p0, %c12 : index + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +}} +''' + + +if mode == "dequant": + out = lookup('ggml_iq1s_grid_code_i32', 'iq1s_code', grid()) + dequant_common +elif mode == "kquant": + out = fill('ggml_kquant_iq1s_grid_fill', grid()) + kquant_common + dot_and_weights('iq1s') + dot_and_weights('iq1m') +else: + sys.exit(f"unknown mode {mode}") +sys.stdout.write(out) diff --git a/ggml/src/ggml-hrx/tools/generate_iq2_dequant_loom.py b/ggml/src/ggml-hrx/tools/generate_iq2_dequant_loom.py new file mode 100755 index 000000000000..866d80badc1c --- /dev/null +++ b/ggml/src/ggml-hrx/tools/generate_iq2_dequant_loom.py @@ -0,0 +1,254 @@ +#!/usr/bin/env python3 +# Copyright 2026 bong-water-water-bong +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# Generates the Loom IQ2_XXS / IQ2_XS grid-code lookups and decoders for motifs/dequant.loom from +# ggml-common.h (iq2xxs_grid, iq2xs_grid: MIT, the ggml authors). Every grid byte is 8, 25 or 43, +# so an entry is stored as a 16-bit code: 2 bits per value, level index 0/1/2. +# Usage: generate_iq2_dequant_loom.py ggml/src/ggml-common.h > out.loom +# The output sits verbatim in kernel-corpus/kernels/loom-libs/motifs/dequant.loom. +import re +import sys +src = open(sys.argv[1]).read() + + +def grid(name, n): + m = re.search(r'GGML_TABLE_BEGIN\(uint64_t, ' + name + r', ' + str(n) + r'\)(.*?)GGML_TABLE_END', src, re.S) + assert m, name + v = [int(x, 16) for x in re.findall(r'0x[0-9a-fA-F]+', m.group(1))] + assert len(v) == n + lv = {8: 0, 25: 1, 43: 2} + return [sum(lv[(e >> (8 * j)) & 255] << (2 * j) for j in range(8)) for e in v] + + +def lookup(fname, prefix, codes): + n = len(codes) // 32 + o = [f'// {prefix}: {len(codes)} grid entries as 16-bit codes (2 bits per value: 0 = 8, 1 = 25, 2 = 43)', + f'func.def inline @{fname}(%grid_index: i32) -> (i32) {{', + ' %c5_i32 = scalar.constant 5 : i32', ' %c31_i32 = scalar.constant 31 : i32'] + for k in range(1, n): + o.append(f' %chunk_id{k} = scalar.constant {k} : i32') + o += [' %chunk_i32 = scalar.shrui %grid_index, %c5_i32 : i32', + ' %lane_i32 = scalar.andi %grid_index, %c31_i32 : i32', + ' %codes = vector.from_elements %lane_i32 : vector<1xi32>'] + for k in range(n): + names = [] + for i in range(32): + nm = f'%{prefix}_{k}_{i}' + o.append(f' {nm} = scalar.constant {float(codes[32 * k + i])} : f32') + names.append(nm) + o.append(f' %{prefix}_{k} = vector.from_elements {", ".join(names)} : vector<32xf32>') + prev = f'%{prefix}_0' + for k in range(1, n): + o.append(f' %is_chunk{k} = scalar.cmpi eq, %chunk_i32, %chunk_id{k} : i32') + o.append(f' %sel{k} = scf.select %is_chunk{k}, %{prefix}_{k}, {prev} : vector<32xf32>') + prev = f'%sel{k}' + o += [f' %v = vector.table.lookup {prev}[%codes] : vector<32xf32>, vector<1xi32> -> vector<1xf32>', + ' %v_f32 = vector.extract %v[0] : vector<1xf32> -> f32', + ' %code = scalar.fptoui %v_f32 : f32 to i32', + ' func.return %code : i32', '}', ''] + return '\n'.join(o) + + +common = ''' +// ksigns_iq2xs[i] (ggml-common.h) is i with its parity as bit 7. +func.def inline @ggml_iq2_signs8(%signs7: i32) -> (i32) { + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c7_i32 = scalar.constant 7 : i32 + %p4 = scalar.shrui %signs7, %c4_i32 : i32 + %x4 = scalar.xori %signs7, %p4 : i32 + %p2 = scalar.shrui %x4, %c2_i32 : i32 + %x2 = scalar.xori %x4, %p2 : i32 + %p1 = scalar.shrui %x2, %c1_i32 : i32 + %x1 = scalar.xori %x2, %p1 : i32 + %parity = scalar.andi %x1, %c1_i32 : i32 + %high = scalar.shli %parity, %c7_i32 : i32 + %signs8 = scalar.ori %signs7, %high : i32 + func.return %signs8 : i32 +} + +// value j (0..7) of a grid code: level (8, 25, 43) with sign bit j of %signs8, times %scale +func.def inline @ggml_iq2_code_value_f32(%code: i32, %j: i32, %signs8: i32, %scale: f32) -> (f32) { + %c0_i32 = scalar.constant 0 : i32 + %c1_i32 = scalar.constant 1 : i32 + %c3_i32 = scalar.constant 3 : i32 + %l8 = scalar.constant 8.0 : f32 + %l25 = scalar.constant 25.0 : f32 + %l43 = scalar.constant 43.0 : f32 + %two_j = scalar.addi %j, %j : i32 + %shifted = scalar.shrui %code, %two_j : i32 + %level = scalar.andi %shifted, %c3_i32 : i32 + %is0 = scalar.cmpi eq, %level, %c0_i32 : i32 + %is1 = scalar.cmpi eq, %level, %c1_i32 : i32 + %v01 = scf.select %is1, %l25, %l43 : f32 + %v = scf.select %is0, %l8, %v01 : f32 + %bit = scalar.shli %c1_i32, %j : i32 + %sign_mask = scalar.andi %signs8, %bit : i32 + %negative = scalar.cmpi ne, %sign_mask, %c0_i32 : i32 + %neg_v = scalar.negf %v : f32 + %signed = scf.select %negative, %neg_v, %v : f32 + %result = scalar.mulf %signed, %scale : f32 + func.return %result : f32 +} + +// four values 4 (p % 2) .. +3 of grid slot p / 2 +func.def inline @ggml_iq2_code_vector4(%code: i32, %half: i32, %signs8: i32, %scale: f32) -> (vector<4xf32>) { + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c3_i32 = scalar.constant 3 : i32 + %c4_i32 = scalar.constant 4 : i32 + %j0 = scalar.muli %half, %c4_i32 : i32 + %j1 = scalar.addi %j0, %c1_i32 : i32 + %j2 = scalar.addi %j0, %c2_i32 : i32 + %j3 = scalar.addi %j0, %c3_i32 : i32 + %v0 = func.call @ggml_iq2_code_value_f32(%code, %j0, %signs8, %scale) : (i32, i32, i32, f32) -> (f32) + %v1 = func.call @ggml_iq2_code_value_f32(%code, %j1, %signs8, %scale) : (i32, i32, i32, f32) -> (f32) + %v2 = func.call @ggml_iq2_code_value_f32(%code, %j2, %signs8, %scale) : (i32, i32, i32, f32) -> (f32) + %v3 = func.call @ggml_iq2_code_value_f32(%code, %j3, %signs8, %scale) : (i32, i32, i32, f32) -> (f32) + %result = vector.from_elements %v0, %v1, %v2, %v3 : vector<4xf32> + func.return %result : vector<4xf32> +} + +// IQ2_XXS (66 bytes: d, qs[32] u16; dequantize_row_iq2_xxs). Group g (ib32) is 8 bytes at 2 + 8g: +// four grid indices, then a u32 with four 7-bit sign groups and a 4-bit scale on top: +// value = d * (0.5 + (aux >> 28)) * 0.25 * grid * sign. Packet p covers slot p / 2, values 4 (p % 2) .. +func.def inline @ggml_iq2xxs_f32_vector4(%weight: buffer, %row_byte_base: offset, %iq2_block: index, %iq2_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c7_i32 = scalar.constant 7 : i32 + %c16_i32 = scalar.constant 16 : i32 + %c28_i32 = scalar.constant 28 : i32 + %c127_i32 = scalar.constant 127 : i32 + %c65535_i32 = scalar.constant 65535 : i32 + %c05_f32 = scalar.constant 0.5 : f32 + %c025_f32 = scalar.constant 0.25 : f32 + %block_bytes = index.constant 66 : offset + %block_byte_add = index.scale %iq2_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %hv = buffer.view %weight[%block_byte_base] : buffer -> view<33xf16> + %wv = buffer.view %weight[%block_byte_base] : buffer -> view<33xi16> + %bv = buffer.view %weight[%block_byte_base] : buffer -> view<66xi8> + %g = index.assume %iq2_group [range(%iq2_group, 0, 7)] : index + %p = index.assume %packet [range(%packet, 0, 7)] : index + %slot = index.div %p, %c2 : index + %half = index.rem %p, %c2 : index + %g8 = index.mul %g, %c4 : index + %gw = index.mul %g, %c4 : index + %idx_byte0 = index.add %g8, %g8 : index + %idx_byte1 = index.add %idx_byte0, %c2 : index + %idx_byte = index.add %idx_byte1, %slot : index + %aux_w0_0 = index.add %gw, %c3 : index + %aux_w1_0 = index.add %aux_w0_0, %c1 : index + %d_f16 = view.load %hv[%c0] : view<33xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %gi_i8 = view.load %bv[%idx_byte] : view<66xi8> -> i8 + %gi = scalar.extui %gi_i8 : i8 to i32 + %w0_i16 = view.load %wv[%aux_w0_0] : view<33xi16> -> i16 + %w1_i16 = view.load %wv[%aux_w1_0] : view<33xi16> -> i16 + %w0 = scalar.extui %w0_i16 : i16 to i32 + %w1 = scalar.extui %w1_i16 : i16 to i32 + %w1s = scalar.shli %w1, %c16_i32 : i32 + %aux = scalar.ori %w0, %w1s : i32 + %sc4 = scalar.shrui %aux, %c28_i32 : i32 + %sc_f = scalar.uitofp %sc4 : i32 to f32 + %sc_plus = scalar.addf %sc_f, %c05_f32 : f32 + %ds = scalar.mulf %d, %sc_plus : f32 + %scale = scalar.mulf %ds, %c025_f32 : f32 + %slot_i32 = index.cast %slot : index to i32 + %sshift = scalar.muli %slot_i32, %c7_i32 : i32 + %sgrp0 = scalar.shrui %aux, %sshift : i32 + %signs7 = scalar.andi %sgrp0, %c127_i32 : i32 + %signs8 = func.call @ggml_iq2_signs8(%signs7) : (i32) -> (i32) + %code = func.call @ggml_iq2xxs_grid_code_i32(%gi) : (i32) -> (i32) + %half_i32 = index.cast %half : index to i32 + %result = func.call @ggml_iq2_code_vector4(%code, %half_i32, %signs8, %scale) : (i32, i32, i32, f32) -> (vector<4xf32>) + func.return %result : vector<4xf32> +} + +func.def inline @ggml_iq2xxs_f16_vector4(%weight: buffer, %row_byte_base: offset, %iq2_block: index, %iq2_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_iq2xxs_f32_vector4(%weight, %row_byte_base, %iq2_block, %iq2_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} + +// IQ2_XS (74 bytes: d, qs[32] u16, scales[8]; dequantize_row_iq2_xs). Slot l of group g is +// q = qs[4g + l]: grid index q & 511, 7 sign bits q >> 9; the scale nibble (l / 2) of scales[g]: +// value = d * (0.5 + nibble) * 0.25 * grid * sign. +func.def inline @ggml_iq2xs_f32_vector4(%weight: buffer, %row_byte_base: offset, %iq2_block: index, %iq2_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c66 = index.constant 66 : index + %c4_i32 = scalar.constant 4 : i32 + %c9_i32 = scalar.constant 9 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c511_i32 = scalar.constant 511 : i32 + %c127_i32 = scalar.constant 127 : i32 + %c05_f32 = scalar.constant 0.5 : f32 + %c025_f32 = scalar.constant 0.25 : f32 + %block_bytes = index.constant 74 : offset + %block_byte_add = index.scale %iq2_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %hv = buffer.view %weight[%block_byte_base] : buffer -> view<37xf16> + %wv = buffer.view %weight[%block_byte_base] : buffer -> view<37xi16> + %bv = buffer.view %weight[%block_byte_base] : buffer -> view<74xi8> + %g = index.assume %iq2_group [range(%iq2_group, 0, 7)] : index + %p = index.assume %packet [range(%packet, 0, 7)] : index + %slot = index.div %p, %c2 : index + %half = index.rem %p, %c2 : index + %g4 = index.mul %g, %c4 : index + %q_at0 = index.add %g4, %slot : index + %q_at = index.add %q_at0, %c1 : index + %sc_at = index.add %c66, %g : index + %d_f16 = view.load %hv[%c0] : view<37xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %q_i16 = view.load %wv[%q_at] : view<37xi16> -> i16 + %q = scalar.extui %q_i16 : i16 to i32 + %gi = scalar.andi %q, %c511_i32 : i32 + %signs7_0 = scalar.shrui %q, %c9_i32 : i32 + %signs7 = scalar.andi %signs7_0, %c127_i32 : i32 + %sc_i8 = view.load %bv[%sc_at] : view<74xi8> -> i8 + %sc = scalar.extui %sc_i8 : i8 to i32 + %nib_sel = index.div %slot, %c2 : index + %nib_sel_i32 = index.cast %nib_sel : index to i32 + %nib_shift = scalar.muli %nib_sel_i32, %c4_i32 : i32 + %nib0 = scalar.shrui %sc, %nib_shift : i32 + %nib = scalar.andi %nib0, %c15_i32 : i32 + %nib_f = scalar.uitofp %nib : i32 to f32 + %nib_plus = scalar.addf %nib_f, %c05_f32 : f32 + %ds = scalar.mulf %d, %nib_plus : f32 + %scale = scalar.mulf %ds, %c025_f32 : f32 + %signs8 = func.call @ggml_iq2_signs8(%signs7) : (i32) -> (i32) + %code = func.call @ggml_iq2xs_grid_code_i32(%gi) : (i32) -> (i32) + %half_i32 = index.cast %half : index to i32 + %result = func.call @ggml_iq2_code_vector4(%code, %half_i32, %signs8, %scale) : (i32, i32, i32, f32) -> (vector<4xf32>) + func.return %result : vector<4xf32> +} + +func.def inline @ggml_iq2xs_f16_vector4(%weight: buffer, %row_byte_base: offset, %iq2_block: index, %iq2_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_iq2xs_f32_vector4(%weight, %row_byte_base, %iq2_block, %iq2_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} +''' +out = lookup('ggml_iq2xxs_grid_code_i32', 'iq2xxs_code', grid('iq2xxs_grid', 256)) + '\n' + \ + lookup('ggml_iq2xs_grid_code_i32', 'iq2xs_code', grid('iq2xs_grid', 512)) + common +sys.stdout.write(out) diff --git a/ggml/src/ggml-hrx/tools/generate_iq2_kquant_loom.py b/ggml/src/ggml-hrx/tools/generate_iq2_kquant_loom.py new file mode 100755 index 000000000000..498ad15d35c3 --- /dev/null +++ b/ggml/src/ggml-hrx/tools/generate_iq2_kquant_loom.py @@ -0,0 +1,307 @@ +#!/usr/bin/env python3 +# Copyright 2026 bong-water-water-bong +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# IQ2_XXS / IQ2_XS lane functions and grid fills for ops/kquant_decode_f32.loom, generated from +# ggml-common.h (iq2xxs_grid, iq2xs_grid: MIT, the ggml authors) as 16-bit codes (2 bits per value). +# Usage: generate_iq2_kquant_loom.py ggml/src/ggml-common.h > out.loom +# The output sits verbatim in kernel-corpus/kernels/loom-libs/ops/kquant_decode_f32.loom. +import re +import sys +src = open(sys.argv[1]).read() + + +def grid(name, n): + m = re.search(r'GGML_TABLE_BEGIN\(uint64_t, ' + name + r', ' + str(n) + r'\)(.*?)GGML_TABLE_END', src, re.S) + assert m, name + v = [int(x, 16) for x in re.findall(r'0x[0-9a-fA-F]+', m.group(1))] + lv = {8: 0, 25: 1, 43: 2} + return [sum(lv[(e >> (8 * j)) & 255] << (2 * j) for j in range(8)) for e in v] + + +def fill(fname, codes): + o = [f'func.def inline @{fname}(%grid: buffer, %chunk: index) {{', + ' %zero_offset = index.constant 0 : offset', + ' %gv = buffer.view %grid[%zero_offset] : buffer -> view<512xi32>'] + for c in range(len(codes) // 128): + o += [f' %k{c} = index.constant {c} : index', f' %is{c} = index.cmp eq, %chunk, %k{c} : index', f' scf.if %is{c} {{'] + for q in range(32): + e = 128 * c + 4 * q + for i in range(4): + o.append(f' %v{c}_{q}_{i} = scalar.constant {codes[e + i]} : i32') + o.append(f' %w{c}_{q} = vector.from_elements %v{c}_{q}_0, %v{c}_{q}_1, %v{c}_{q}_2, %v{c}_{q}_3 : vector<4xi32>') + o.append(f' %o{c}_{q} = index.constant {e} : index') + o.append(f' vector.store %w{c}_{q}, %gv[%o{c}_{q}] : vector<4xi32>, view<512xi32>') + o.append(' }') + o += [' func.return', '}', ''] + return '\n'.join(o) + + +common = ''' +// IQ2_XXS / IQ2_XS: a 256-value block is 8 groups of 32, each 4 grid slots of 8 values; lane l16 +// owns group l16 / 2, slots 2 (l16 % 2) and +1: values 32 (l16 / 2) + 16 (l16 % 2) .. +15, one +// scale. The grid (16-bit codes, 2 bits per value: level 8 + 17c + (c >> 1) = 8, 25, 43) is +// staged in workgroup memory like IQ3_S's; the sign byte is ksigns_iq2xs = 7 bits + parity. +func.def inline @ggml_kquant_iq2_slot_values(%code: i32, %signs7: i32) -> (vector<8xf32>) { + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c7_i32 = scalar.constant 7 : i32 + %p4 = scalar.shrui %signs7, %c4_i32 : i32 + %x4 = scalar.xori %signs7, %p4 : i32 + %p2 = scalar.shrui %x4, %c2_i32 : i32 + %x2 = scalar.xori %x4, %p2 : i32 + %p1 = scalar.shrui %x2, %c1_i32 : i32 + %x1 = scalar.xori %x2, %p1 : i32 + %parity = scalar.andi %x1, %c1_i32 : i32 + %high = scalar.shli %parity, %c7_i32 : i32 + %signs8 = scalar.ori %signs7, %high : i32 + %s0 = scalar.constant 0 : i32 + %s2 = scalar.constant 2 : i32 + %s4 = scalar.constant 4 : i32 + %s6 = scalar.constant 6 : i32 + %s8 = scalar.constant 8 : i32 + %s10 = scalar.constant 10 : i32 + %s12 = scalar.constant 12 : i32 + %s14 = scalar.constant 14 : i32 + %s3 = scalar.constant 3 : i32 + %s5 = scalar.constant 5 : i32 + %shift2 = vector.from_elements %s0, %s2, %s4, %s6, %s8, %s10, %s12, %s14 : vector<8xi32> + %shift1 = vector.from_elements %s0, %c1_i32, %s2, %s3, %s4, %s5, %s6, %c7_i32 : vector<8xi32> + %three = vector.splat %s3 : vector<8xi32> + %one = vector.splat %c1_i32 : vector<8xi32> + %c8v = vector.splat %s8 : vector<8xi32> + %c17_i32 = scalar.constant 17 : i32 + %c17v = vector.splat %c17_i32 : vector<8xi32> + %cv = vector.splat %code : vector<8xi32> + %cs = vector.shrui %cv, %shift2 : vector<8xi32> + %c = vector.andi %cs, %three : vector<8xi32> + %c17 = vector.muli %c, %c17v : vector<8xi32> + %chi = vector.shrui %c, %one : vector<8xi32> + %lv0 = vector.addi %c17, %chi : vector<8xi32> + %lv = vector.addi %lv0, %c8v : vector<8xi32> + %sv = vector.splat %signs8 : vector<8xi32> + %sb0 = vector.shrui %sv, %shift1 : vector<8xi32> + %sb = vector.andi %sb0, %one : vector<8xi32> + %sb2 = vector.shli %sb, %one : vector<8xi32> + %sgn = vector.subi %one, %sb2 : vector<8xi32> + %v = vector.muli %lv, %sgn : vector<8xi32> + %vf = vector.sitofp %v : vector<8xi32> to vector<8xf32> + func.return %vf : vector<8xf32> +} + +func.def inline @ggml_kquant_iq2_join16(%a: vector<8xf32>, %b: vector<8xf32>) -> (vector<16xf32>) { + %a0 = vector.extract %a[0] : vector<8xf32> -> f32 + %a1 = vector.extract %a[1] : vector<8xf32> -> f32 + %a2 = vector.extract %a[2] : vector<8xf32> -> f32 + %a3 = vector.extract %a[3] : vector<8xf32> -> f32 + %a4 = vector.extract %a[4] : vector<8xf32> -> f32 + %a5 = vector.extract %a[5] : vector<8xf32> -> f32 + %a6 = vector.extract %a[6] : vector<8xf32> -> f32 + %a7 = vector.extract %a[7] : vector<8xf32> -> f32 + %b0 = vector.extract %b[0] : vector<8xf32> -> f32 + %b1 = vector.extract %b[1] : vector<8xf32> -> f32 + %b2 = vector.extract %b[2] : vector<8xf32> -> f32 + %b3 = vector.extract %b[3] : vector<8xf32> -> f32 + %b4 = vector.extract %b[4] : vector<8xf32> -> f32 + %b5 = vector.extract %b[5] : vector<8xf32> -> f32 + %b6 = vector.extract %b[6] : vector<8xf32> -> f32 + %b7 = vector.extract %b[7] : vector<8xf32> -> f32 + %r = vector.from_elements %a0, %a1, %a2, %a3, %a4, %a5, %a6, %a7, %b0, %b1, %b2, %b3, %b4, %b5, %b6, %b7 : vector<16xf32> + func.return %r : vector<16xf32> +} + +// IQ2_XXS (66 bytes: d, qs[32] u16): group g = 8 bytes at 2 + 8g: slot indices (bytes 0..3), +// then u32 aux: four 7-bit sign groups, scale nibble in bits 28..31 (d * (0.5 + s) * 0.25). +func.def inline @ggml_kquant_iq2xxs_lane_parts(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c7_i32 = scalar.constant 7 : i32 + %c16_i32 = scalar.constant 16 : i32 + %c28_i32 = scalar.constant 28 : i32 + %c127_i32 = scalar.constant 127 : i32 + %c05 = scalar.constant 0.5 : f32 + %c025 = scalar.constant 0.25 : f32 + %block_bytes = index.constant 66 : offset + %zero_offset = index.constant 0 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %g = index.div %l, %c2 : index + %h = index.rem %l, %c2 : index + %s0 = index.mul %h, %c2 : index + %s1 = index.add %s0, %c1 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %hv = buffer.view %weight[%block_base] : buffer -> view<33xf16> + %wv = buffer.view %weight[%block_base] : buffer -> view<33xi16> + %bv = buffer.view %weight[%block_base] : buffer -> view<66xi8> + %gv = buffer.view %grid[%zero_offset] : buffer -> view<512xi32> + %d_f16 = view.load %hv[%c0] : view<33xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %g8 = index.mul %g, %c8 : index + %ib0 = index.add %g8, %c2 : index + %ia = index.add %ib0, %s0 : index + %ib = index.add %ib0, %s1 : index + %ia_i8 = view.load %bv[%ia] : view<66xi8> -> i8 + %ib_i8 = view.load %bv[%ib] : view<66xi8> -> i8 + %ga = scalar.extui %ia_i8 : i8 to i32 + %gb = scalar.extui %ib_i8 : i8 to i32 + %ga_x = index.cast %ga : i32 to index + %gb_x = index.cast %gb : i32 to index + %ga_b = index.assume %ga_x [range(%ga_x, 0, 255)] : index + %gb_b = index.assume %gb_x [range(%gb_x, 0, 255)] : index + %code_a = view.load %gv[%ga_b] : view<512xi32> -> i32 + %code_b = view.load %gv[%gb_b] : view<512xi32> -> i32 + %g4 = index.mul %g, %c4 : index + %w0_at = index.add %g4, %c3 : index + %w1_at = index.add %g4, %c4 : index + %w0_i16 = view.load %wv[%w0_at] : view<33xi16> -> i16 + %w1_i16 = view.load %wv[%w1_at] : view<33xi16> -> i16 + %w0 = scalar.extui %w0_i16 : i16 to i32 + %w1 = scalar.extui %w1_i16 : i16 to i32 + %w1s = scalar.shli %w1, %c16_i32 : i32 + %aux = scalar.ori %w0, %w1s : i32 + %sc4 = scalar.shrui %aux, %c28_i32 : i32 + %sc_f = scalar.uitofp %sc4 : i32 to f32 + %sc_p = scalar.addf %sc_f, %c05 : f32 + %ds = scalar.mulf %d, %sc_p : f32 + %scale = scalar.mulf %ds, %c025 : f32 + %s0_i32 = index.cast %s0 : index to i32 + %s1_i32 = index.cast %s1 : index to i32 + %sha = scalar.muli %s0_i32, %c7_i32 : i32 + %shb = scalar.muli %s1_i32, %c7_i32 : i32 + %sa0 = scalar.shrui %aux, %sha : i32 + %sb0 = scalar.shrui %aux, %shb : i32 + %sa = scalar.andi %sa0, %c127_i32 : i32 + %sb = scalar.andi %sb0, %c127_i32 : i32 + %va = func.call @ggml_kquant_iq2_slot_values(%code_a, %sa) : (i32, i32) -> (vector<8xf32>) + %vb = func.call @ggml_kquant_iq2_slot_values(%code_b, %sb) : (i32, i32) -> (vector<8xf32>) + %v = func.call @ggml_kquant_iq2_join16(%va, %vb) : (vector<8xf32>, vector<8xf32>) -> (vector<16xf32>) + %g32 = index.mul %g, %c32 : index + %h16 = index.mul %h, %c16 : index + %p0 = index.add %g32, %h16 : index + func.return %v, %scale, %p0 : vector<16xf32>, f32, index +} + +// IQ2_XS (74 bytes: d, qs[32] u16, scales[8]): slot q = qs[4g + s]: grid index q & 511, 7 sign +// bits q >> 9; scale nibble h of scales[g] (slots 2h, 2h + 1). +func.def inline @ggml_kquant_iq2xs_lane_parts(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c66 = index.constant 66 : index + %c4_i32 = scalar.constant 4 : i32 + %c9_i32 = scalar.constant 9 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c127_i32 = scalar.constant 127 : i32 + %c511_i32 = scalar.constant 511 : i32 + %c05 = scalar.constant 0.5 : f32 + %c025 = scalar.constant 0.25 : f32 + %block_bytes = index.constant 74 : offset + %zero_offset = index.constant 0 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %g = index.div %l, %c2 : index + %h = index.rem %l, %c2 : index + %s0 = index.mul %h, %c2 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %hv = buffer.view %weight[%block_base] : buffer -> view<37xf16> + %wv = buffer.view %weight[%block_base] : buffer -> view<37xi16> + %bv = buffer.view %weight[%block_base] : buffer -> view<74xi8> + %gv = buffer.view %grid[%zero_offset] : buffer -> view<512xi32> + %d_f16 = view.load %hv[%c0] : view<37xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %g4 = index.mul %g, %c4 : index + %qa_at0 = index.add %g4, %s0 : index + %qa_at = index.add %qa_at0, %c1 : index + %qb_at = index.add %qa_at, %c1 : index + %qa_i16 = view.load %wv[%qa_at] : view<37xi16> -> i16 + %qb_i16 = view.load %wv[%qb_at] : view<37xi16> -> i16 + %qa = scalar.extui %qa_i16 : i16 to i32 + %qb = scalar.extui %qb_i16 : i16 to i32 + %ga = scalar.andi %qa, %c511_i32 : i32 + %gb = scalar.andi %qb, %c511_i32 : i32 + %sa0 = scalar.shrui %qa, %c9_i32 : i32 + %sb0 = scalar.shrui %qb, %c9_i32 : i32 + %sa = scalar.andi %sa0, %c127_i32 : i32 + %sb = scalar.andi %sb0, %c127_i32 : i32 + %ga_x = index.cast %ga : i32 to index + %gb_x = index.cast %gb : i32 to index + %ga_b = index.assume %ga_x [range(%ga_x, 0, 511)] : index + %gb_b = index.assume %gb_x [range(%gb_x, 0, 511)] : index + %code_a = view.load %gv[%ga_b] : view<512xi32> -> i32 + %code_b = view.load %gv[%gb_b] : view<512xi32> -> i32 + %sc_at = index.add %c66, %g : index + %sc_i8 = view.load %bv[%sc_at] : view<74xi8> -> i8 + %sc = scalar.extui %sc_i8 : i8 to i32 + %h_i32 = index.cast %h : index to i32 + %nsh = scalar.muli %h_i32, %c4_i32 : i32 + %nib0 = scalar.shrui %sc, %nsh : i32 + %nib = scalar.andi %nib0, %c15_i32 : i32 + %nib_f = scalar.uitofp %nib : i32 to f32 + %nib_p = scalar.addf %nib_f, %c05 : f32 + %ds = scalar.mulf %d, %nib_p : f32 + %scale = scalar.mulf %ds, %c025 : f32 + %va = func.call @ggml_kquant_iq2_slot_values(%code_a, %sa) : (i32, i32) -> (vector<8xf32>) + %vb = func.call @ggml_kquant_iq2_slot_values(%code_b, %sb) : (i32, i32) -> (vector<8xf32>) + %v = func.call @ggml_kquant_iq2_join16(%va, %vb) : (vector<8xf32>, vector<8xf32>) -> (vector<16xf32>) + %g32 = index.mul %g, %c32 : index + %h16 = index.mul %h, %c16 : index + %p0 = index.add %g32, %h16 : index + func.return %v, %scale, %p0 : vector<16xf32>, f32, index +} +''' + + +def dot_and_weights(fmt): + return f''' +func.def inline @ggml_kquant_{fmt}_lane_dot(%grid: buffer, %weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) {{ + %zero_scalar = scalar.constant 0.0 : f32 + %x_block_bytes = index.constant 1024 : offset + %v, %scale, %p = func.call @ggml_kquant_{fmt}_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %x = vector.load %xv[%p] : view<256xf32> -> vector<16xf32> + %vx = vector.mulf %v, %x : vector<16xf32> + %sum = vector.reduce %vx, %zero_scalar : vector<16xf32>, f32 + %result = scalar.mulf %scale, %sum : f32 + func.return %result : f32 +}} + +func.def inline @ggml_kquant_{fmt}_lane_weights(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) {{ + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c12 = index.constant 12 : index + %v, %scale, %p0 = func.call @ggml_kquant_{fmt}_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %sv = vector.splat %scale : vector<16xf32> + %w = vector.mulf %v, %sv : vector<16xf32> + %p1 = index.add %p0, %c4 : index + %p2 = index.add %p0, %c8 : index + %p3 = index.add %p0, %c12 : index + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +}} +''' + + +out = fill('ggml_kquant_iq2xxs_grid_fill', grid('iq2xxs_grid', 256)) + fill('ggml_kquant_iq2xs_grid_fill', grid('iq2xs_grid', 512)) + common + dot_and_weights('iq2xxs') + dot_and_weights('iq2xs') +sys.stdout.write(out) diff --git a/ggml/src/ggml-hrx/tools/generate_iq2s_kquant_loom.py b/ggml/src/ggml-hrx/tools/generate_iq2s_kquant_loom.py new file mode 100644 index 000000000000..e0845c87990b --- /dev/null +++ b/ggml/src/ggml-hrx/tools/generate_iq2s_kquant_loom.py @@ -0,0 +1,229 @@ +#!/usr/bin/env python3 +# Copyright 2026 bong-water-water-bong +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# IQ2_S lane functions and grid fill for ops/kquant_decode_f32.loom, generated from ggml-common.h +# (iq2s_grid: MIT, the ggml authors). Every grid byte is 8, 25 or 43, so an entry (8 values) is a +# 16-bit code (2 bits per value, level 0/1/2); the 1024 codes are staged two per word (512 words, the +# kernels' 2 KiB grid buffer). +# Usage: generate_iq2s_kquant_loom.py ggml/src/ggml-common.h > out.loom +# The output sits verbatim in kernel-corpus/kernels/loom-libs/ops/kquant_decode_f32.loom. +import re +import sys + +src = open(sys.argv[1]).read() + + +def grid(): + m = re.search(r'GGML_TABLE_BEGIN\(uint64_t, iq2s_grid, 1024\)(.*?)GGML_TABLE_END', src, re.S) + assert m, "iq2s_grid" + v = [int(x, 16) for x in re.findall(r'0x[0-9a-fA-F]+', m.group(1))] + assert len(v) == 1024, len(v) + lv = {8: 0, 25: 1, 43: 2} + return [sum(lv[(e >> (8 * j)) & 255] << (2 * j) for j in range(8)) for e in v] + + +def fill(fname, codes): + words = [codes[2 * w] | (codes[2 * w + 1] << 16) for w in range(len(codes) // 2)] + words = [w - (1 << 32) if w >= 1 << 31 else w for w in words] # i32 constants are signed + per = len(words) // 4 + o = [f'// Stages the IQ2_S grid as {len(words)} words (two 16-bit codes each); subgroup %chunk (0..3) writes', + f'// words {per} %chunk .. +{per - 1}.', + f'func.def inline @{fname}(%grid: buffer, %chunk: index) {{', + ' %zero_offset = index.constant 0 : offset', + f' %gv = buffer.view %grid[%zero_offset] : buffer -> view<{len(words)}xi32>'] + for c in range(4): + o += [f' %k{c} = index.constant {c} : index', f' %is{c} = index.cmp eq, %chunk, %k{c} : index', f' scf.if %is{c} {{'] + for q in range(per // 4): + e = per * c + 4 * q + for i in range(4): + o.append(f' %v{c}_{q}_{i} = scalar.constant {words[e + i]} : i32') + o.append(f' %w{c}_{q} = vector.from_elements %v{c}_{q}_0, %v{c}_{q}_1, %v{c}_{q}_2, %v{c}_{q}_3 : vector<4xi32>') + o.append(f' %o{c}_{q} = index.constant {e} : index') + o.append(f' vector.store %w{c}_{q}, %gv[%o{c}_{q}] : vector<4xi32>, view<{len(words)}xi32>') + o.append(' }') + o += [' func.return', '}', ''] + return '\n'.join(o) + + +common = ''' +// IQ2_S: lane l16 owns group l16 / 2, slots 2 (l16 % 2) and +1 (values 32 (l16 / 2) + 16 (l16 % 2) .. +// +15), which share scale nibble l16 % 2 of scales[g]. Slot s: grid index qs[4 g + s] | (qh[g] << (8 - 2 s) +// & 0x300), a full sign byte signs[4 g + s] (no parity bit, unlike IQ2_XXS / IQ2_XS). +func.def inline @ggml_kquant_iq2s_code(%grid: buffer, %index: i32) -> (i32) { + %zero_offset = index.constant 0 : offset + %gv = buffer.view %grid[%zero_offset] : buffer -> view<512xi32> + %c1_i32 = scalar.constant 1 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c65535_i32 = scalar.constant 65535 : i32 + %w_i32 = scalar.shrui %index, %c1_i32 : i32 + %w_x = index.cast %w_i32 : i32 to index + %w_b = index.assume %w_x [range(%w_x, 0, 511)] : index + %word = view.load %gv[%w_b] : view<512xi32> -> i32 + %odd = scalar.andi %index, %c1_i32 : i32 + %sh = scalar.shli %odd, %c4_i32 : i32 + %shifted = scalar.shrui %word, %sh : i32 + %code = scalar.andi %shifted, %c65535_i32 : i32 + func.return %code : i32 +} + +func.def inline @ggml_kquant_iq2s_slot_values(%code: i32, %signs8: i32) -> (vector<8xf32>) { + %c1_i32 = scalar.constant 1 : i32 + %c7_i32 = scalar.constant 7 : i32 + %s0 = scalar.constant 0 : i32 + %s2 = scalar.constant 2 : i32 + %s3 = scalar.constant 3 : i32 + %s4 = scalar.constant 4 : i32 + %s5 = scalar.constant 5 : i32 + %s6 = scalar.constant 6 : i32 + %s8 = scalar.constant 8 : i32 + %s10 = scalar.constant 10 : i32 + %s12 = scalar.constant 12 : i32 + %s14 = scalar.constant 14 : i32 + %shift2 = vector.from_elements %s0, %s2, %s4, %s6, %s8, %s10, %s12, %s14 : vector<8xi32> + %shift1 = vector.from_elements %s0, %c1_i32, %s2, %s3, %s4, %s5, %s6, %c7_i32 : vector<8xi32> + %three = vector.splat %s3 : vector<8xi32> + %one = vector.splat %c1_i32 : vector<8xi32> + %c8v = vector.splat %s8 : vector<8xi32> + %c17_i32 = scalar.constant 17 : i32 + %c17v = vector.splat %c17_i32 : vector<8xi32> + %cv = vector.splat %code : vector<8xi32> + %cs = vector.shrui %cv, %shift2 : vector<8xi32> + %c = vector.andi %cs, %three : vector<8xi32> + %c17 = vector.muli %c, %c17v : vector<8xi32> + %chi = vector.shrui %c, %one : vector<8xi32> + %lv0 = vector.addi %c17, %chi : vector<8xi32> + %lv = vector.addi %lv0, %c8v : vector<8xi32> + %sv = vector.splat %signs8 : vector<8xi32> + %sb0 = vector.shrui %sv, %shift1 : vector<8xi32> + %sb = vector.andi %sb0, %one : vector<8xi32> + %sb2 = vector.shli %sb, %one : vector<8xi32> + %sgn = vector.subi %one, %sb2 : vector<8xi32> + %v = vector.muli %lv, %sgn : vector<8xi32> + %vf = vector.sitofp %v : vector<8xi32> to vector<8xf32> + func.return %vf : vector<8xf32> +} + +// IQ2_S (82 bytes: d, qs[32] grid low bits, signs[32], qh[8], scales[8]): value = d * (0.5 + nibble) * 0.25 +// * grid * sign. +func.def inline @ggml_kquant_iq2s_lane_parts(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c4 = index.constant 4 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c34 = index.constant 34 : index + %c66 = index.constant 66 : index + %c74 = index.constant 74 : index + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c15_i32 = scalar.constant 15 : i32 + %c768_i32 = scalar.constant 768 : i32 + %c05 = scalar.constant 0.5 : f32 + %c025 = scalar.constant 0.25 : f32 + %block_bytes = index.constant 82 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %g = index.div %l, %c2 : index + %h = index.rem %l, %c2 : index + %sa = index.mul %h, %c2 : index + %sb = index.add %sa, %c1 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %hv = buffer.view %weight[%block_base] : buffer -> view<41xf16> + %bv = buffer.view %weight[%block_base] : buffer -> view<82xi8> + %d_f16 = view.load %hv[%c0] : view<41xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %g4 = index.mul %g, %c4 : index + %qs0 = index.add %g4, %c2 : index + %qa_at = index.add %qs0, %sa : index + %qb_at = index.add %qs0, %sb : index + %sg0 = index.add %g4, %c34 : index + %sa_at = index.add %sg0, %sa : index + %sb_at = index.add %sg0, %sb : index + %qh_at = index.add %c66, %g : index + %sc_at = index.add %c74, %g : index + %qa_i8 = view.load %bv[%qa_at] : view<82xi8> -> i8 + %qb_i8 = view.load %bv[%qb_at] : view<82xi8> -> i8 + %sga_i8 = view.load %bv[%sa_at] : view<82xi8> -> i8 + %sgb_i8 = view.load %bv[%sb_at] : view<82xi8> -> i8 + %qh_i8 = view.load %bv[%qh_at] : view<82xi8> -> i8 + %sc_i8 = view.load %bv[%sc_at] : view<82xi8> -> i8 + %qa = scalar.extui %qa_i8 : i8 to i32 + %qb = scalar.extui %qb_i8 : i8 to i32 + %sga = scalar.extui %sga_i8 : i8 to i32 + %sgb = scalar.extui %sgb_i8 : i8 to i32 + %qh = scalar.extui %qh_i8 : i8 to i32 + %sc = scalar.extui %sc_i8 : i8 to i32 + %sa_i32 = index.cast %sa : index to i32 + %sb_i32 = index.cast %sb : index to i32 + %sa2 = scalar.muli %sa_i32, %c2_i32 : i32 + %sb2 = scalar.muli %sb_i32, %c2_i32 : i32 + %sha = scalar.subi %c8_i32, %sa2 : i32 + %shb = scalar.subi %c8_i32, %sb2 : i32 + %ha0 = scalar.shli %qh, %sha : i32 + %hb0 = scalar.shli %qh, %shb : i32 + %ha = scalar.andi %ha0, %c768_i32 : i32 + %hb = scalar.andi %hb0, %c768_i32 : i32 + %ia = scalar.ori %qa, %ha : i32 + %ib = scalar.ori %qb, %hb : i32 + %code_a = func.call @ggml_kquant_iq2s_code(%grid, %ia) : (buffer, i32) -> (i32) + %code_b = func.call @ggml_kquant_iq2s_code(%grid, %ib) : (buffer, i32) -> (i32) + %h_i32 = index.cast %h : index to i32 + %nsh = scalar.muli %h_i32, %c4_i32 : i32 + %nib0 = scalar.shrui %sc, %nsh : i32 + %nib = scalar.andi %nib0, %c15_i32 : i32 + %nib_f = scalar.uitofp %nib : i32 to f32 + %nib_p = scalar.addf %nib_f, %c05 : f32 + %ds = scalar.mulf %d, %nib_p : f32 + %scale = scalar.mulf %ds, %c025 : f32 + %va = func.call @ggml_kquant_iq2s_slot_values(%code_a, %sga) : (i32, i32) -> (vector<8xf32>) + %vb = func.call @ggml_kquant_iq2s_slot_values(%code_b, %sgb) : (i32, i32) -> (vector<8xf32>) + %v = func.call @ggml_kquant_iq2_join16(%va, %vb) : (vector<8xf32>, vector<8xf32>) -> (vector<16xf32>) + %g32 = index.mul %g, %c32 : index + %h16 = index.mul %h, %c16 : index + %p0 = index.add %g32, %h16 : index + func.return %v, %scale, %p0 : vector<16xf32>, f32, index +} + +func.def inline @ggml_kquant_iq2s_lane_dot(%grid: buffer, %weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %zero_scalar = scalar.constant 0.0 : f32 + %x_block_bytes = index.constant 1024 : offset + %v, %scale, %p = func.call @ggml_kquant_iq2s_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %x = vector.load %xv[%p] : view<256xf32> -> vector<16xf32> + %vx = vector.mulf %v, %x : vector<16xf32> + %sum = vector.reduce %vx, %zero_scalar : vector<16xf32>, f32 + %result = scalar.mulf %scale, %sum : f32 + func.return %result : f32 +} + +func.def inline @ggml_kquant_iq2s_lane_weights(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c12 = index.constant 12 : index + %v, %scale, %p0 = func.call @ggml_kquant_iq2s_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %sv = vector.splat %scale : vector<16xf32> + %w = vector.mulf %v, %sv : vector<16xf32> + %p1 = index.add %p0, %c4 : index + %p2 = index.add %p0, %c8 : index + %p3 = index.add %p0, %c12 : index + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} +''' + +sys.stdout.write(fill('ggml_kquant_iq2s_grid_fill', grid()) + common) diff --git a/ggml/src/ggml-hrx/tools/generate_iq3xxs_loom.py b/ggml/src/ggml-hrx/tools/generate_iq3xxs_loom.py new file mode 100644 index 000000000000..6e1dd2e6e456 --- /dev/null +++ b/ggml/src/ggml-hrx/tools/generate_iq3xxs_loom.py @@ -0,0 +1,390 @@ +#!/usr/bin/env python3 +# Copyright 2026 bong-water-water-bong +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# IQ3_XXS grid lookups and decoders for motifs/dequant.loom (mode "dequant") and lane functions and +# grid fill for ops/kquant_decode_f32.loom (mode "kquant"), generated from ggml-common.h (iq3xxs_grid: +# MIT, the ggml authors). Every grid byte is one of 4, 12, 20, 28, 36, 44, 52, 62, so an entry (4 +# values) is stored as a 12-bit code: 3 bits per value, level L = 0..7, value 4 + 8 L (+ 2 for L = 7). +# Usage: generate_iq3xxs_loom.py ggml/src/ggml-common.h dequant|kquant > out.loom +# The output sits verbatim in kernel-corpus/kernels/loom-libs/motifs/dequant.loom (dequant) or +# kernel-corpus/kernels/loom-libs/ops/kquant_decode_f32.loom (kquant). +import re +import sys + +src = open(sys.argv[1]).read() +mode = sys.argv[2] +LEVELS = {4: 0, 12: 1, 20: 2, 28: 3, 36: 4, 44: 5, 52: 6, 62: 7} + + +def grid(): + m = re.search(r'GGML_TABLE_BEGIN\(uint32_t, iq3xxs_grid, 256\)(.*?)GGML_TABLE_END', src, re.S) + assert m, "iq3xxs_grid" + v = [int(x, 16) for x in re.findall(r'0x[0-9a-fA-F]+', m.group(1))] + assert len(v) == 256, len(v) + return [sum(LEVELS[(e >> (8 * j)) & 255] << (3 * j) for j in range(4)) for e in v] + + +def lookup(fname, prefix, codes): + n = len(codes) // 32 + o = [f'// {prefix}: {len(codes)} grid entries as 12-bit codes (3 bits per value: level L, value 4 + 8 L, 62 for L = 7)', + f'func.def inline @{fname}(%grid_index: i32) -> (i32) {{', + ' %c5_i32 = scalar.constant 5 : i32', ' %c31_i32 = scalar.constant 31 : i32'] + for k in range(1, n): + o.append(f' %chunk_id{k} = scalar.constant {k} : i32') + o += [' %chunk_i32 = scalar.shrui %grid_index, %c5_i32 : i32', + ' %lane_i32 = scalar.andi %grid_index, %c31_i32 : i32', + ' %codes = vector.from_elements %lane_i32 : vector<1xi32>'] + for k in range(n): + names = [] + for i in range(32): + nm = f'%{prefix}_{k}_{i}' + o.append(f' {nm} = scalar.constant {float(codes[32 * k + i])} : f32') + names.append(nm) + o.append(f' %{prefix}_{k} = vector.from_elements {", ".join(names)} : vector<32xf32>') + prev = f'%{prefix}_0' + for k in range(1, n): + o.append(f' %is_chunk{k} = scalar.cmpi eq, %chunk_i32, %chunk_id{k} : i32') + o.append(f' %sel{k} = scf.select %is_chunk{k}, %{prefix}_{k}, {prev} : vector<32xf32>') + prev = f'%sel{k}' + o += [f' %v = vector.table.lookup {prev}[%codes] : vector<32xf32>, vector<1xi32> -> vector<1xf32>', + ' %v_f32 = vector.extract %v[0] : vector<1xf32> -> f32', + ' %code = scalar.fptoui %v_f32 : f32 to i32', + ' func.return %code : i32', '}', ''] + return '\n'.join(o) + + +dequant_common = ''' +// the four values of a 12-bit IQ3_XXS grid code, with sign bits 4 half .. +3 of %signs8, times %scale +func.def inline @ggml_iq3xxs_code_vector4(%code: i32, %half: i32, %signs8: i32, %scale: f32) -> (vector<4xf32>) { + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c8_i32 = scalar.constant 8 : i32 + %s0 = scalar.constant 0 : i32 + %s3 = scalar.constant 3 : i32 + %s6 = scalar.constant 6 : i32 + %s9 = scalar.constant 9 : i32 + %shift3 = vector.from_elements %s0, %s3, %s6, %s9 : vector<4xi32> + %cv = vector.splat %code : vector<4xi32> + %lvs = vector.shrui %cv, %shift3 : vector<4xi32> + %seven = vector.splat %c7_i32 : vector<4xi32> + %lv = vector.andi %lvs, %seven : vector<4xi32> + %eight = vector.splat %c8_i32 : vector<4xi32> + %four = vector.splat %c4_i32 : vector<4xi32> + %lv8 = vector.muli %lv, %eight : vector<4xi32> + %base = vector.addi %lv8, %four : vector<4xi32> + // level 7 is the only one with all three bits set: bump = 2 (L & L >> 1 & L >> 2 & 1) + %one7 = vector.splat %c1_i32 : vector<4xi32> + %two7 = vector.splat %c2_i32 : vector<4xi32> + %l1 = vector.shrui %lv, %one7 : vector<4xi32> + %l2 = vector.shrui %lv, %two7 : vector<4xi32> + %a01 = vector.andi %lv, %l1 : vector<4xi32> + %a012 = vector.andi %a01, %l2 : vector<4xi32> + %is7 = vector.andi %a012, %one7 : vector<4xi32> + %bump = vector.shli %is7, %one7 : vector<4xi32> + %mag = vector.addi %base, %bump : vector<4xi32> + %sbase = scalar.muli %half, %c4_i32 : i32 + %c1v = scalar.constant 1 : i32 + %c2v = scalar.constant 2 : i32 + %c3v = scalar.constant 3 : i32 + %sh0 = vector.from_elements %s0, %c1v, %c2v, %c3v : vector<4xi32> + %sbv = vector.splat %sbase : vector<4xi32> + %sh = vector.addi %sh0, %sbv : vector<4xi32> + %sv = vector.splat %signs8 : vector<4xi32> + %sb0 = vector.shrui %sv, %sh : vector<4xi32> + %one = vector.splat %c1_i32 : vector<4xi32> + %sb = vector.andi %sb0, %one : vector<4xi32> + %sb2 = vector.shli %sb, %one : vector<4xi32> + %sgn = vector.subi %one, %sb2 : vector<4xi32> + %signed = vector.muli %mag, %sgn : vector<4xi32> + %f = vector.sitofp %signed : vector<4xi32> to vector<4xf32> + %scv = vector.splat %scale : vector<4xf32> + %result = vector.mulf %f, %scv : vector<4xf32> + func.return %result : vector<4xf32> +} + +// IQ3_XXS (98 bytes: d, qs[64] grid indices, 8 x u32 scales and signs; dequantize_row_iq3_xxs). Group g +// has grid indices qs[8 g .. 8 g + 7] (bytes 2 + 8 g ..) and aux = u32 at byte 66 + 4 g: four 7-bit sign +// groups (slot l at bit 7 l) and a 4-bit scale on top: value = d * (0.5 + (aux >> 28)) * 0.5 * grid * sign. +// Packet p covers slot p / 2, grid entry 2 (p / 2) + p % 2 (its four values, sign bits 4 (p % 2) ..). +func.def inline @ggml_iq3xxs_f32_vector4(%weight: buffer, %row_byte_base: offset, %iq3_block: index, %iq3_group: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c8 = index.constant 8 : index + %c33 = index.constant 33 : index + %c7_i32 = scalar.constant 7 : i32 + %c16_i32 = scalar.constant 16 : i32 + %c28_i32 = scalar.constant 28 : i32 + %c127_i32 = scalar.constant 127 : i32 + %c05_f32 = scalar.constant 0.5 : f32 + %block_bytes = index.constant 98 : offset + %block_byte_add = index.scale %iq3_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %hv = buffer.view %weight[%block_byte_base] : buffer -> view<49xf16> + %wv = buffer.view %weight[%block_byte_base] : buffer -> view<49xi16> + %bv = buffer.view %weight[%block_byte_base] : buffer -> view<98xi8> + %g = index.assume %iq3_group [range(%iq3_group, 0, 7)] : index + %p = index.assume %packet [range(%packet, 0, 7)] : index + %slot = index.div %p, %c2 : index + %half = index.rem %p, %c2 : index + %g8 = index.mul %g, %c8 : index + %idx0 = index.add %g8, %c2 : index + %idx_at = index.add %idx0, %p : index + %g2 = index.add %g, %g : index + %aux_w0 = index.add %c33, %g2 : index + %aux_w1 = index.add %aux_w0, %c1 : index + %d_f16 = view.load %hv[%c0] : view<49xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %gi_i8 = view.load %bv[%idx_at] : view<98xi8> -> i8 + %gi = scalar.extui %gi_i8 : i8 to i32 + %w0_i16 = view.load %wv[%aux_w0] : view<49xi16> -> i16 + %w1_i16 = view.load %wv[%aux_w1] : view<49xi16> -> i16 + %w0 = scalar.extui %w0_i16 : i16 to i32 + %w1 = scalar.extui %w1_i16 : i16 to i32 + %w1s = scalar.shli %w1, %c16_i32 : i32 + %aux = scalar.ori %w0, %w1s : i32 + %sc4 = scalar.shrui %aux, %c28_i32 : i32 + %sc_f = scalar.uitofp %sc4 : i32 to f32 + %sc_plus = scalar.addf %sc_f, %c05_f32 : f32 + %ds = scalar.mulf %d, %sc_plus : f32 + %scale = scalar.mulf %ds, %c05_f32 : f32 + %slot_i32 = index.cast %slot : index to i32 + %sshift = scalar.muli %slot_i32, %c7_i32 : i32 + %sgrp = scalar.shrui %aux, %sshift : i32 + %signs7 = scalar.andi %sgrp, %c127_i32 : i32 + %signs8 = func.call @ggml_iq2_signs8(%signs7) : (i32) -> (i32) + %code = func.call @ggml_iq3xxs_grid_code_i32(%gi) : (i32) -> (i32) + %half_i32 = index.cast %half : index to i32 + %result = func.call @ggml_iq3xxs_code_vector4(%code, %half_i32, %signs8, %scale) : (i32, i32, i32, f32) -> (vector<4xf32>) + func.return %result : vector<4xf32> +} + +func.def inline @ggml_iq3xxs_f16_vector4(%weight: buffer, %row_byte_base: offset, %iq3_block: index, %iq3_group: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_iq3xxs_f32_vector4(%weight, %row_byte_base, %iq3_block, %iq3_group, %packet) : (buffer, offset, index, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} +''' + + +def fill(fname, codes): + per = len(codes) // 4 + o = [f'// Stages the IQ3_XXS grid, one 12-bit code per word; subgroup %chunk (0..3) writes words {per} %chunk .. +{per - 1}.', + f'func.def inline @{fname}(%grid: buffer, %chunk: index) {{', + ' %zero_offset = index.constant 0 : offset', + ' %gv = buffer.view %grid[%zero_offset] : buffer -> view<512xi32>'] + for c in range(4): + o += [f' %k{c} = index.constant {c} : index', f' %is{c} = index.cmp eq, %chunk, %k{c} : index', f' scf.if %is{c} {{'] + for q in range(per // 4): + e = per * c + 4 * q + for i in range(4): + o.append(f' %v{c}_{q}_{i} = scalar.constant {codes[e + i]} : i32') + o.append(f' %w{c}_{q} = vector.from_elements %v{c}_{q}_0, %v{c}_{q}_1, %v{c}_{q}_2, %v{c}_{q}_3 : vector<4xi32>') + o.append(f' %o{c}_{q} = index.constant {e} : index') + o.append(f' vector.store %w{c}_{q}, %gv[%o{c}_{q}] : vector<4xi32>, view<512xi32>') + o.append(' }') + o += [' func.return', '}', ''] + return '\n'.join(o) + + +kquant_common = ''' +// IQ3_XXS: a 256-value block is 8 groups of 32, each 4 slots of 8 values (two grid entries of 4); lane l16 +// owns group l16 / 2, slots 2 (l16 % 2) and +1: values 32 (l16 / 2) + 16 (l16 % 2) .. +15, one scale. +// The grid (12-bit codes) is staged in workgroup memory like IQ2's; signs are ksigns_iq2xs (7 bits + parity). +func.def inline @ggml_kquant_iq3xxs_slot_values(%code_a: i32, %code_b: i32, %signs7: i32) -> (vector<8xf32>) { + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c4_i32 = scalar.constant 4 : i32 + %c7_i32 = scalar.constant 7 : i32 + %c12_i32 = scalar.constant 12 : i32 + %p4 = scalar.shrui %signs7, %c4_i32 : i32 + %x4 = scalar.xori %signs7, %p4 : i32 + %p2 = scalar.shrui %x4, %c2_i32 : i32 + %x2 = scalar.xori %x4, %p2 : i32 + %p1 = scalar.shrui %x2, %c1_i32 : i32 + %x1 = scalar.xori %x2, %p1 : i32 + %parity = scalar.andi %x1, %c1_i32 : i32 + %high = scalar.shli %parity, %c7_i32 : i32 + %signs8 = scalar.ori %signs7, %high : i32 + %b_hi = scalar.shli %code_b, %c12_i32 : i32 + %both = scalar.ori %code_a, %b_hi : i32 + %s0 = scalar.constant 0 : i32 + %s3 = scalar.constant 3 : i32 + %s5 = scalar.constant 5 : i32 + %s6 = scalar.constant 6 : i32 + %s8 = scalar.constant 8 : i32 + %s9 = scalar.constant 9 : i32 + %s12 = scalar.constant 12 : i32 + %s15 = scalar.constant 15 : i32 + %s18 = scalar.constant 18 : i32 + %s21 = scalar.constant 21 : i32 + %shift3 = vector.from_elements %s0, %s3, %s6, %s9, %s12, %s15, %s18, %s21 : vector<8xi32> + %shift1 = vector.from_elements %s0, %c1_i32, %c2_i32, %s3, %c4_i32, %s5, %s6, %c7_i32 : vector<8xi32> + %cv = vector.splat %both : vector<8xi32> + %lvs = vector.shrui %cv, %shift3 : vector<8xi32> + %seven = vector.splat %c7_i32 : vector<8xi32> + %lv = vector.andi %lvs, %seven : vector<8xi32> + %eight = vector.splat %s8 : vector<8xi32> + %four = vector.splat %c4_i32 : vector<8xi32> + %lv8 = vector.muli %lv, %eight : vector<8xi32> + %base = vector.addi %lv8, %four : vector<8xi32> + // level 7 is the only one with all three bits set: bump = 2 (L & L >> 1 & L >> 2 & 1) + %one7 = vector.splat %c1_i32 : vector<8xi32> + %two7 = vector.splat %c2_i32 : vector<8xi32> + %l1 = vector.shrui %lv, %one7 : vector<8xi32> + %l2 = vector.shrui %lv, %two7 : vector<8xi32> + %a01 = vector.andi %lv, %l1 : vector<8xi32> + %a012 = vector.andi %a01, %l2 : vector<8xi32> + %is7 = vector.andi %a012, %one7 : vector<8xi32> + %bump = vector.shli %is7, %one7 : vector<8xi32> + %mag = vector.addi %base, %bump : vector<8xi32> + %sv = vector.splat %signs8 : vector<8xi32> + %sb0 = vector.shrui %sv, %shift1 : vector<8xi32> + %one = vector.splat %c1_i32 : vector<8xi32> + %sb = vector.andi %sb0, %one : vector<8xi32> + %sb2 = vector.shli %sb, %one : vector<8xi32> + %sgn = vector.subi %one, %sb2 : vector<8xi32> + %v = vector.muli %mag, %sgn : vector<8xi32> + %vf = vector.sitofp %v : vector<8xi32> to vector<8xf32> + func.return %vf : vector<8xf32> +} + +// IQ3_XXS (98 bytes): lane slots s = 2h, 2h + 1 of group g use grid indices at bytes 2 + 8 g + 2 s (+1) +// and sign group s of aux = u32 at byte 66 + 4 g, scale d * (0.5 + (aux >> 28)) * 0.5. +func.def inline @ggml_kquant_iq3xxs_lane_parts(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, f32, index) { + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c3 = index.constant 3 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c16 = index.constant 16 : index + %c32 = index.constant 32 : index + %c33 = index.constant 33 : index + %c7_i32 = scalar.constant 7 : i32 + %c16_i32 = scalar.constant 16 : i32 + %c28_i32 = scalar.constant 28 : i32 + %c127_i32 = scalar.constant 127 : i32 + %c05 = scalar.constant 0.5 : f32 + %block_bytes = index.constant 98 : offset + %zero_offset = index.constant 0 : offset + %l = index.assume %lane16 [range(%lane16, 0, 15)] : index + %g = index.div %l, %c2 : index + %h = index.rem %l, %c2 : index + %s0 = index.mul %h, %c2 : index + %s1 = index.add %s0, %c1 : index + %block_add = index.scale %block, %block_bytes : index, offset -> offset + %block_base = index.add %row_base, %block_add : offset + %hv = buffer.view %weight[%block_base] : buffer -> view<49xf16> + %wv = buffer.view %weight[%block_base] : buffer -> view<49xi16> + %bv = buffer.view %weight[%block_base] : buffer -> view<98xi8> + %gv = buffer.view %grid[%zero_offset] : buffer -> view<512xi32> + %d_f16 = view.load %hv[%c0] : view<49xf16> -> f16 + %d = scalar.extf %d_f16 : f16 to f32 + %g8 = index.mul %g, %c8 : index + %i0 = index.add %g8, %c2 : index + %h4 = index.mul %h, %c4 : index + %ia = index.add %i0, %h4 : index + %ib = index.add %ia, %c1 : index + %ic = index.add %ia, %c2 : index + %id = index.add %ia, %c3 : index + %ia_i8 = view.load %bv[%ia] : view<98xi8> -> i8 + %ib_i8 = view.load %bv[%ib] : view<98xi8> -> i8 + %ic_i8 = view.load %bv[%ic] : view<98xi8> -> i8 + %id_i8 = view.load %bv[%id] : view<98xi8> -> i8 + %ga = scalar.extui %ia_i8 : i8 to i32 + %gb = scalar.extui %ib_i8 : i8 to i32 + %gc = scalar.extui %ic_i8 : i8 to i32 + %gd = scalar.extui %id_i8 : i8 to i32 + %ga_x = index.cast %ga : i32 to index + %gb_x = index.cast %gb : i32 to index + %gc_x = index.cast %gc : i32 to index + %gd_x = index.cast %gd : i32 to index + %ga_b = index.assume %ga_x [range(%ga_x, 0, 255)] : index + %gb_b = index.assume %gb_x [range(%gb_x, 0, 255)] : index + %gc_b = index.assume %gc_x [range(%gc_x, 0, 255)] : index + %gd_b = index.assume %gd_x [range(%gd_x, 0, 255)] : index + %code_a = view.load %gv[%ga_b] : view<512xi32> -> i32 + %code_b = view.load %gv[%gb_b] : view<512xi32> -> i32 + %code_c = view.load %gv[%gc_b] : view<512xi32> -> i32 + %code_d = view.load %gv[%gd_b] : view<512xi32> -> i32 + %g2 = index.add %g, %g : index + %w0_at = index.add %c33, %g2 : index + %w1_at = index.add %w0_at, %c1 : index + %w0_i16 = view.load %wv[%w0_at] : view<49xi16> -> i16 + %w1_i16 = view.load %wv[%w1_at] : view<49xi16> -> i16 + %w0 = scalar.extui %w0_i16 : i16 to i32 + %w1 = scalar.extui %w1_i16 : i16 to i32 + %w1s = scalar.shli %w1, %c16_i32 : i32 + %aux = scalar.ori %w0, %w1s : i32 + %sc4 = scalar.shrui %aux, %c28_i32 : i32 + %sc_f = scalar.uitofp %sc4 : i32 to f32 + %sc_p = scalar.addf %sc_f, %c05 : f32 + %ds = scalar.mulf %d, %sc_p : f32 + %scale = scalar.mulf %ds, %c05 : f32 + %s0_i32 = index.cast %s0 : index to i32 + %s1_i32 = index.cast %s1 : index to i32 + %sha = scalar.muli %s0_i32, %c7_i32 : i32 + %shb = scalar.muli %s1_i32, %c7_i32 : i32 + %sa0 = scalar.shrui %aux, %sha : i32 + %sb0 = scalar.shrui %aux, %shb : i32 + %sa = scalar.andi %sa0, %c127_i32 : i32 + %sb = scalar.andi %sb0, %c127_i32 : i32 + %va = func.call @ggml_kquant_iq3xxs_slot_values(%code_a, %code_b, %sa) : (i32, i32, i32) -> (vector<8xf32>) + %vb = func.call @ggml_kquant_iq3xxs_slot_values(%code_c, %code_d, %sb) : (i32, i32, i32) -> (vector<8xf32>) + %v = func.call @ggml_kquant_iq2_join16(%va, %vb) : (vector<8xf32>, vector<8xf32>) -> (vector<16xf32>) + %g32 = index.mul %g, %c32 : index + %h16 = index.mul %h, %c16 : index + %p0 = index.add %g32, %h16 : index + func.return %v, %scale, %p0 : vector<16xf32>, f32, index +} + +func.def inline @ggml_kquant_iq3xxs_lane_dot(%grid: buffer, %weight: buffer, %input: buffer, %row_base: offset, %block: index, %lane16: index) -> (f32) { + %zero_scalar = scalar.constant 0.0 : f32 + %x_block_bytes = index.constant 1024 : offset + %v, %scale, %p = func.call @ggml_kquant_iq3xxs_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %x_base = index.scale %block, %x_block_bytes : index, offset -> offset + %xv = buffer.view %input[%x_base] : buffer -> view<256xf32> + %x = vector.load %xv[%p] : view<256xf32> -> vector<16xf32> + %vx = vector.mulf %v, %x : vector<16xf32> + %sum = vector.reduce %vx, %zero_scalar : vector<16xf32>, f32 + %result = scalar.mulf %scale, %sum : f32 + func.return %result : f32 +} + +func.def inline @ggml_kquant_iq3xxs_lane_weights(%grid: buffer, %weight: buffer, %row_base: offset, %block: index, %lane16: index) -> (vector<16xf32>, index, index, index, index) { + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c12 = index.constant 12 : index + %v, %scale, %p0 = func.call @ggml_kquant_iq3xxs_lane_parts(%grid, %weight, %row_base, %block, %lane16) : (buffer, buffer, offset, index, index) -> (vector<16xf32>, f32, index) + %sv = vector.splat %scale : vector<16xf32> + %w = vector.mulf %v, %sv : vector<16xf32> + %p1 = index.add %p0, %c4 : index + %p2 = index.add %p0, %c8 : index + %p3 = index.add %p0, %c12 : index + func.return %w, %p0, %p1, %p2, %p3 : vector<16xf32>, index, index, index, index +} +''' + +if mode == "dequant": + out = lookup('ggml_iq3xxs_grid_code_i32', 'iq3xxs_code', grid()) + dequant_common +elif mode == "kquant": + out = fill('ggml_kquant_iq3xxs_grid_fill', grid()) + kquant_common +else: + sys.exit(f"unknown mode {mode}") +sys.stdout.write(out) diff --git a/ggml/src/ggml-hrx/tools/generate_kernel_corpus.py b/ggml/src/ggml-hrx/tools/generate_kernel_corpus.py new file mode 100644 index 000000000000..4c5826126847 --- /dev/null +++ b/ggml/src/ggml-hrx/tools/generate_kernel_corpus.py @@ -0,0 +1,666 @@ +#!/usr/bin/env python3 + +import argparse +import hashlib +import json +import pathlib +import re +import subprocess +import sys +import tempfile +from typing import Dict, Iterable, List, Tuple + + +DEFAULT_KERNEL_FAMILY = "qwen3_moe" + +SOURCE_ARRAY_TEMPLATE = """static const unsigned char {symbol}[] = {{ +{bytes} +}}; +static constexpr size_t {symbol}Size = {size}; +""" + +DEPENDENCY_TABLE_TEMPLATE = """static const KernelSourceSpan {dependency_table}[] = {{ +{dependencies} +}}; +""" + +SOURCE_RECORD_TEMPLATE = """static const KernelSource {record} = {{ + {{ reinterpret_cast({source_symbol}), {source_symbol}Size, {source_format} }}, + {dependency_table}, + {dependency_count}, +}}; +""" + +CORPUS_ARRAY_TEMPLATE = """static {type} {symbol}[] = {{ +{values} +}}; +""" + +KERNEL_RECORD_TEMPLATE = """ {{ + {family}, + {name}, + kernel_catalog_id({family}, {name}), + {source}, + {dependencies}, + {symbol}, + "amdgpu", + {target_selector}, + {{ nullptr, 0 }}, + {scalar_parameters}, + {bindings}, + {source_digest}, + {workload_parameters}, + {launch_parameters}, + {{ + {compile_mode}, + {link_module}, + {primary_sources}, + {library_sources}, + }}, + }},""" + +SOURCE_DATA_TEMPLATE = """{source_arrays} +{dependency_tables} +{source_records} +static const KernelSourceRecordEntry kKernelSourceRecords[] = {{ +{lookup_entries} +}}; +""" + +CORPUS_DATA_TEMPLATE = """{kernel_arrays} +static const KernelDefinition kQwenKernelDefinitions[] = {{ +{kernel_records} +}}; + +static const KernelCorpus kQwenKernelCorpus = {{ + "ggml-hrx-kernel-corpus-v2", + {upstream_revision}, + {corpus_digest}, + {recipe_digest}, + {plan_case_count}, + {{ kQwenKernelDefinitions, {kernel_count} }}, +}}; +""" + +CATALOG_DATA_TEMPLATE = """struct KernelCatalogEntry {{ + const char * family; + const char * name; +}}; + +static constexpr KernelCatalogEntry kKernelCatalogEntries[] = {{ +{kernel_entries} +}}; + +constexpr bool kernel_catalog_entry_exists(const char * family, const char * name) {{ + for (const KernelCatalogEntry & known : kKernelCatalogEntries) {{ + if (kernel_catalog_name_equal(family, known.family) && kernel_catalog_name_equal(name, known.name)) {{ + return true; + }} + }} + return false; +}} +""" + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Generate embedded kernel corpus data includes.") + parser.add_argument("--source-output", type=pathlib.Path, required=True) + parser.add_argument("--corpus-output", type=pathlib.Path, required=True) + parser.add_argument("--catalog-output", type=pathlib.Path, required=True) + parser.add_argument("--manifest", type=pathlib.Path, action="append", required=True) + parser.add_argument("--corpus-dir", type=pathlib.Path, action="append", required=True) + parser.add_argument("--source-format", choices=("text", "binary"), default="text") + parser.add_argument("--loom-link", type=pathlib.Path) + parser.add_argument("--loom-format", type=pathlib.Path) + parser.add_argument("--depfile", type=pathlib.Path) + return parser.parse_args() + + +def read_text(path: pathlib.Path) -> str: + try: + return path.read_text() + except OSError as exc: + raise RuntimeError(f"failed to read {path}: {exc}") from exc + + +def read_bytes(path: pathlib.Path) -> bytes: + try: + return path.read_bytes() + except OSError as exc: + raise RuntimeError(f"failed to read {path}: {exc}") from exc + + +def run_tool(command: List[str]) -> subprocess.CompletedProcess[str]: + return subprocess.run(command, text=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE, check=False) + + +def require_tool(command: List[str]) -> None: + result = run_tool(command) + if result.returncode: + raise RuntimeError(f"command failed ({result.returncode}): {' '.join(command)}\n{result.stderr}") + + +def convert_source_to_bytecode(source_path: pathlib.Path, loom_link: pathlib.Path, loom_format: pathlib.Path) -> bytes: + with tempfile.TemporaryDirectory(prefix="ggml-hrx-loom-") as temp_dir_name: + temp_dir = pathlib.Path(temp_dir_name) + stripped = temp_dir / "stripped.loom" + bytecode = temp_dir / "stripped.loombc" + require_tool([ + str(loom_link), + "--verify=false", + "--mode=merge", + "--strip-check", + "--to=text", + f"--output={stripped}", + str(source_path), + ]) + format_result = run_tool([ + str(loom_format), + "--from=text", + "--to=bc", + f"--output={bytecode}", + str(stripped), + ]) + if format_result.returncode: + require_tool([ + str(loom_link), + "--verify=false", + "--mode=merge", + "--strip-check", + "--to=bc", + f"--output={bytecode}", + str(source_path), + ]) + return read_bytes(bytecode) + + +def sanitize_symbol(path: str, index: int) -> str: + stem = re.sub(r"[^0-9A-Za-z_]", "_", path) + if not stem or stem[0].isdigit(): + stem = f"_{stem}" + return f"kernel_source_{index}_{stem}" + + +def format_byte_array(data: bytes) -> str: + if not data: + return "" + lines = [] + for offset in range(0, len(data), 16): + chunk = data[offset : offset + 16] + lines.append(" " + ", ".join(f"0x{byte:02x}" for byte in chunk) + ",") + return "\n".join(lines) + "\n" + + +def escape_cpp_string(text: str) -> str: + return text.replace("\\", "\\\\").replace('"', '\\"') + + +def cpp_string(text: str) -> str: + return f"\"{escape_cpp_string(text)}\"" + + +def resource_access_value(access: str) -> str: + if access == "read": + return "ResourceAccess::Read" + if access == "write": + return "ResourceAccess::Write" + if access == "read_write": + return "ResourceAccess::ReadWrite" + raise RuntimeError(f"invalid kernel binding access metadata: {access}") + + +def span_initializer(symbol: str, count: int) -> str: + if count == 0: + return "{ nullptr, 0 }" + return "{ " + symbol + ", " + str(count) + " }" + + +def typed_array(symbol: str, value_type: str, values: List[str]) -> Tuple[str, str]: + if not values: + return "", "{ nullptr, 0 }" + array = CORPUS_ARRAY_TEMPLATE.format( + type=value_type, + symbol=symbol, + values="\n".join(" " + value + "," for value in values), + ) + return array, span_initializer(symbol, len(values)) + + +def string_array(symbol: str, items: Iterable[str]) -> Tuple[str, str]: + return typed_array(symbol, "const char * const", [cpp_string(item) for item in items]) + + +def source_ref_array(symbol: str, items: Iterable[str], source_records: Dict[str, str]) -> Tuple[str, str]: + values = [ + "{ " + cpp_string(item) + ", &" + source_records[item] + " }" + for item in items + ] + return typed_array(symbol, "const KernelSourceRef", values) + + +def scalar_array(symbol: str, items: Iterable[dict]) -> Tuple[str, str]: + values = [ + "{ " + cpp_string(item["name"]) + ", " + cpp_string(item["type"]) + " }" + for item in items + ] + return typed_array(symbol, "const KernelScalarDefinition", values) + + +def binding_array(symbol: str, names: List[str], access: List[str]) -> Tuple[str, str]: + if len(names) != len(access): + raise RuntimeError("kernel binding access metadata has the wrong arity") + values = [ + "{ " + cpp_string(name) + ", " + resource_access_value(access_value) + " }" + for name, access_value in zip(names, access) + ] + return typed_array(symbol, "const KernelBindingDefinition", values) + + +def scalar_parameter_names(workload_parameters: List[dict], launch_parameters: List[dict]) -> List[str]: + names: List[str] = [] + for parameter in [*workload_parameters, *launch_parameters]: + name = parameter["name"] + if name not in names: + names.append(name) + return names + + +def depfile_escape(path: pathlib.Path) -> str: + return str(path).replace("\\", "\\\\").replace(" ", "\\ ") + + +def collect_sources(manifest: dict) -> Tuple[List[str], Dict[str, List[str]]]: + source_dependencies: Dict[str, List[str]] = {} + all_sources = set() + + for export in manifest.get("exports", []): + recipe = export.get("compile_recipe", {}) + primary_sources = recipe.get("primary_sources", []) + library_sources = recipe.get("library_sources", []) + if len(primary_sources) != 1: + raise RuntimeError("expected each export compile_recipe to have exactly one primary source") + + primary = primary_sources[0] + dependencies = list(library_sources) + previous = source_dependencies.get(primary) + if previous is not None and previous != dependencies: + raise RuntimeError(f"conflicting dependency list for {primary}") + source_dependencies[primary] = dependencies + all_sources.add(primary) + all_sources.update(dependencies) + + return sorted(all_sources), source_dependencies + + +def sha256(data: bytes) -> str: + return hashlib.sha256(data).hexdigest() + + +def json_digest(value: object) -> str: + return sha256(json.dumps(value, sort_keys=True, separators=(",", ":")).encode("utf-8")) + + +def complete_manifest_source_metadata(manifest: dict, source_paths: Dict[str, pathlib.Path]) -> dict: + completed = dict(manifest) + files = [] + for file in manifest.get("files", []): + path = str(file["path"]) + file_record = {key: value for key, value in file.items() if key not in ("owner", "sha256", "size")} + data = read_bytes(source_paths[path]) + file_record["sha256"] = sha256(data) + files.append(file_record) + completed["files"] = files + corpus_digest_input = { + "files": [{"path": file["path"], "sha256": file["sha256"]} for file in files], + "exports": [ + { + "family": export.get("family", DEFAULT_KERNEL_FAMILY), + "name": export["name"], + "source": export["source"], + "symbol": export["symbol"], + "target_selector": export.get("target_selector", ""), + "compile_recipe": export["compile_recipe"], + } + for export in completed.get("exports", []) + ], + } + recipe_digest_input = { + "link_modules": completed.get("link_modules", []), + "plan_cases": completed.get("plan_cases", []), + } + completed["corpus_sha256"] = json_digest(corpus_digest_input) + completed["build_bazel_sha256"] = json_digest(recipe_digest_input) + return completed + + +def merge_unique_text(values: Iterable[str]) -> str: + unique = list(dict.fromkeys(value for value in values if value)) + return ";".join(unique) + + +def require_unique_export_variants(exports: List[dict]) -> None: + variants: Dict[Tuple[str, str, str], dict] = {} + for export in exports: + key = ( + str(export.get("family", DEFAULT_KERNEL_FAMILY)), + str(export["name"]), + str(export.get("target_selector", "")), + ) + previous = variants.get(key) + if previous is not None: + family, name, target = key + label = target or "default" + raise RuntimeError(f"kernel corpus repeats target variant {family}:{name}@{label}") + variants[key] = export + + +def merge_manifests(manifest_paths: List[pathlib.Path], corpus_dirs: List[pathlib.Path]) -> Tuple[dict, Dict[str, pathlib.Path]]: + if len(manifest_paths) != len(corpus_dirs): + raise RuntimeError("--manifest and --corpus-dir must be provided the same number of times") + if not manifest_paths: + raise RuntimeError("at least one --manifest is required") + + loaded = [(manifest_path, corpus_dir, json.loads(read_text(manifest_path))) + for manifest_path, corpus_dir in zip(manifest_paths, corpus_dirs)] + if len(loaded) == 1: + manifest_path, corpus_dir, manifest = loaded[0] + source_paths = {file["path"]: corpus_dir / file["path"] for file in manifest.get("files", [])} + return complete_manifest_source_metadata(manifest, source_paths), source_paths + + files = [] + file_digests: Dict[str, str] = {} + source_paths: Dict[str, pathlib.Path] = {} + exports = [] + link_modules = [] + plan_cases = [] + metadata = [] + + for manifest_path, corpus_dir, manifest in loaded: + manifest_source_paths = {file["path"]: corpus_dir / file["path"] for file in manifest.get("files", [])} + manifest = complete_manifest_source_metadata(manifest, manifest_source_paths) + metadata.append({ + "manifest": str(manifest_path), + "upstream_revision": manifest.get("upstream_revision", ""), + "corpus_sha256": manifest.get("corpus_sha256", ""), + "build_bazel_sha256": manifest.get("build_bazel_sha256", ""), + }) + for file in manifest.get("files", []): + path = str(file["path"]) + digest = str(file["sha256"]) + previous = file_digests.get(path) + if previous is not None: + if previous != digest: + raise RuntimeError(f"manifest file digest conflict for {path}: {previous} vs {digest}") + continue + file_digests[path] = digest + source_paths[path] = corpus_dir / path + files.append(dict(file)) + exports.extend(dict(export) for export in manifest.get("exports", [])) + link_modules.extend(dict(module) for module in manifest.get("link_modules", [])) + plan_cases.extend(dict(case) for case in manifest.get("plan_cases", [])) + + require_unique_export_variants(exports) + + corpus_digest_input = { + "files": [{"path": file["path"], "sha256": file["sha256"]} for file in files], + "exports": [ + { + "family": export.get("family", DEFAULT_KERNEL_FAMILY), + "name": export["name"], + "source": export["source"], + "symbol": export["symbol"], + "target_selector": export.get("target_selector", ""), + "compile_recipe": export["compile_recipe"], + } + for export in exports + ], + } + recipe_digest_input = { + "manifests": metadata, + "link_modules": link_modules, + "plan_cases": plan_cases, + } + + return { + "schema": "ggml-hrx-kernel-corpus-v2", + "upstream_revision": merge_unique_text(str(manifest.get("upstream_revision", "")) for _, _, manifest in loaded), + "corpus_sha256": json_digest(corpus_digest_input), + "build_bazel_sha256": json_digest(recipe_digest_input), + "files": files, + "exports": exports, + "link_modules": link_modules, + "plan_cases": plan_cases, + }, source_paths + + +def kernel_catalog_id(family: str, name: str) -> int: + hash_value = 1469598103934665603 + for byte in family.encode("utf-8"): + hash_value ^= byte + hash_value = (hash_value * 1099511628211) & 0xFFFFFFFFFFFFFFFF + hash_value ^= 0 + hash_value = (hash_value * 1099511628211) & 0xFFFFFFFFFFFFFFFF + for byte in name.encode("utf-8"): + hash_value ^= byte + hash_value = (hash_value * 1099511628211) & 0xFFFFFFFFFFFFFFFF + return hash_value + + +def generate_corpus_records(manifest: dict, source_records: Dict[str, str], source_digests: Dict[str, str]) -> Tuple[str, str, int]: + arrays = [] + records = [] + exports = manifest.get("exports", []) + for index, export in enumerate(exports): + recipe = export["compile_recipe"] + dependencies = list(export["compile_dependencies"]) + library_sources = list(recipe["library_sources"]) + if dependencies != library_sources: + raise RuntimeError("legacy dependency closure disagrees with compile recipe") + workload_parameters = list(export["workload_parameters"]) + launch_parameters = list(export["launch_parameters"]) + + dependencies_array, dependencies_span = string_array(f"kKernelDependencies{index}", dependencies) + scalar_array_text, scalar_span = string_array( + f"kKernelScalarParameters{index}", + scalar_parameter_names(workload_parameters, launch_parameters), + ) + bindings_array, bindings_span = binding_array( + f"kKernelBindings{index}", + list(export["bindings"]), + list(export.get("binding_access", [])), + ) + workload_array, workload_span = scalar_array(f"kKernelWorkloadParameters{index}", workload_parameters) + launch_array, launch_span = scalar_array(f"kKernelLaunchParameters{index}", launch_parameters) + primary_array, primary_span = source_ref_array(f"kKernelPrimarySources{index}", recipe["primary_sources"], source_records) + library_array, library_span = source_ref_array(f"kKernelLibrarySources{index}", library_sources, source_records) + arrays.extend( + item for item in [ + dependencies_array, + scalar_array_text, + bindings_array, + workload_array, + launch_array, + primary_array, + library_array, + ] if item + ) + records.append( + KERNEL_RECORD_TEMPLATE.format( + family=cpp_string(export.get("family", DEFAULT_KERNEL_FAMILY)), + name=cpp_string(export["name"]), + symbol=cpp_string(export["symbol"]), + target_selector=cpp_string(export.get("target_selector", "")), + source=cpp_string(export["source"]), + dependencies=dependencies_span, + source_digest=cpp_string(source_digests[export["source"]]), + scalar_parameters=scalar_span, + bindings=bindings_span, + workload_parameters=workload_span, + launch_parameters=launch_span, + compile_mode=cpp_string(recipe["mode"]), + link_module=cpp_string(recipe.get("link_module", "")), + primary_sources=primary_span, + library_sources=library_span, + ) + ) + return "\n".join(arrays), "\n".join(records), len(exports) + + +def generate_catalog_verifier(manifest: dict) -> str: + kernel_entries = sorted(set( + (export.get("family", DEFAULT_KERNEL_FAMILY), export["name"]) + for export in manifest.get("exports", []) + )) + ids: Dict[int, Tuple[str, str]] = {} + for family, kernel_name in kernel_entries: + catalog_id = kernel_catalog_id(family, kernel_name) + if catalog_id in ids: + previous_family, previous_name = ids[catalog_id] + raise RuntimeError( + "kernel catalog id collision: " + f"{previous_family}/{previous_name} and {family}/{kernel_name}") + ids[catalog_id] = (family, kernel_name) + return CATALOG_DATA_TEMPLATE.format( + kernel_entries="\n".join( + " { " + cpp_string(family) + ", " + cpp_string(kernel_name) + " }," + for family, kernel_name in kernel_entries + ), + ) + + +def generate_includes(args: argparse.Namespace, manifest: dict, source_paths: Dict[str, pathlib.Path]) -> Tuple[str, str, str, List[pathlib.Path], int]: + if args.source_format == "binary" and (args.loom_link is None or args.loom_format is None): + raise RuntimeError("binary source format requires --loom-link and --loom-format") + + sources, source_dependencies = collect_sources(manifest) + source_digests = {str(file["path"]): str(file["sha256"]) for file in manifest.get("files", [])} + source_bytes: Dict[str, bytes] = {} + source_format = "KERNEL_SOURCE_FORMAT_BINARY" if args.source_format == "binary" else "KERNEL_SOURCE_FORMAT_TEXT" + input_files: List[pathlib.Path] = [] + + for export in manifest.get("exports", []): + recipe = export.get("compile_recipe", {}) + for primary in recipe.get("primary_sources", []): + if primary not in sources: + raise RuntimeError(f"export primary source is not embedded: {primary}") + + for source in sources: + if source not in source_digests: + raise RuntimeError(f"embedded source is missing from manifest file table: {source}") + path = source_paths[source] + input_files.append(path) + data = read_bytes(path) + source_digests[source] = sha256(data) + if args.source_format == "binary": + data = convert_source_to_bytecode(path, args.loom_link, args.loom_format) + source_bytes[source] = data + + source_symbols: Dict[str, str] = {} + source_arrays = [] + for index, source in enumerate(sources): + symbol = sanitize_symbol(source, index) + source_symbols[source] = symbol + data = source_bytes[source] + source_arrays.append( + SOURCE_ARRAY_TEMPLATE.format( + symbol=symbol, + bytes=format_byte_array(data), + size=len(data), + ) + ) + + dependency_tables = [] + source_record_definitions = [] + source_record_symbols = {} + lookup_entries = [] + for index, source in enumerate(sources): + dependencies = source_dependencies.get(source, []) + record = f"kernel_source_record_{index}" + source_record_symbols[source] = record + if dependencies: + dependency_table = f"kernel_source_dependencies_{index}" + entries = [] + for dependency in dependencies: + symbol = source_symbols[dependency] + entries.append( + f" {{ reinterpret_cast({symbol}), {symbol}Size, {source_format} }}," + ) + dependency_tables.append( + DEPENDENCY_TABLE_TEMPLATE.format( + dependency_table=dependency_table, + dependencies="\n".join(entries), + ) + ) + else: + dependency_table = "nullptr" + + source_record_definitions.append( + SOURCE_RECORD_TEMPLATE.format( + record=record, + source_symbol=source_symbols[source], + source_format=source_format, + dependency_table=dependency_table, + dependency_count=len(dependencies), + ) + ) + lookup_entries.append(f" {{ {cpp_string(source)}, &{record} }},") + + kernel_arrays, kernel_records, kernel_count = generate_corpus_records(manifest, source_record_symbols, source_digests) + return ( + SOURCE_DATA_TEMPLATE.format( + source_arrays="\n".join(source_arrays), + dependency_tables="\n".join(dependency_tables), + source_records="\n".join(source_record_definitions), + lookup_entries="\n".join(lookup_entries), + ), + CORPUS_DATA_TEMPLATE.format( + kernel_arrays=kernel_arrays, + kernel_records=kernel_records, + upstream_revision=cpp_string(manifest["upstream_revision"]), + corpus_digest=cpp_string(manifest["corpus_sha256"]), + recipe_digest=cpp_string(manifest["build_bazel_sha256"]), + plan_case_count=len(manifest["plan_cases"]), + kernel_count=kernel_count, + ), + generate_catalog_verifier(manifest), + input_files, + sum(len(data) for data in source_bytes.values()), + ) + + +def write_depfile(path: pathlib.Path, outputs: List[pathlib.Path], inputs: List[pathlib.Path]) -> None: + targets = " ".join(depfile_escape(output_path) for output_path in outputs) + entries = [depfile_escape(input_path) for input_path in inputs] + path.write_text(f"{targets}: {' '.join(entries)}\n") + + +def main() -> int: + args = parse_args() + try: + manifest, source_paths = merge_manifests(args.manifest, args.corpus_dir) + source_include, corpus_include, catalog_include, input_files, byte_count = generate_includes( + args, manifest, source_paths) + args.source_output.parent.mkdir(parents=True, exist_ok=True) + args.corpus_output.parent.mkdir(parents=True, exist_ok=True) + args.catalog_output.parent.mkdir(parents=True, exist_ok=True) + args.source_output.write_text(source_include) + args.corpus_output.write_text(corpus_include) + args.catalog_output.write_text(catalog_include) + if args.depfile is not None: + args.depfile.parent.mkdir(parents=True, exist_ok=True) + write_depfile(args.depfile, [args.source_output, args.corpus_output, args.catalog_output], + [*args.manifest, *input_files]) + except Exception as exc: + print(f"generate_kernel_corpus.py: {exc}", file=sys.stderr) + return 1 + + print( + f"embedded kernel corpus: source_files={len(input_files)} source_bytes={byte_count} " + f"source_format={args.source_format}", + file=sys.stderr, + ) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/ggml/src/ggml-hrx/tools/tool-utils.h b/ggml/src/ggml-hrx/tools/tool-utils.h new file mode 100644 index 000000000000..c01e089ceab4 --- /dev/null +++ b/ggml/src/ggml-hrx/tools/tool-utils.h @@ -0,0 +1,61 @@ +#pragma once + +#include "../hrx-interop-utils.h" + +#include +#include +#include +#include +#include +#include + +namespace ggml::hrx::tool { + +inline std::string read_file(const std::filesystem::path & path) { + std::ifstream input(path, std::ios::binary); + if (!input) { + return {}; + } + return { std::istreambuf_iterator(input), std::istreambuf_iterator() }; +} + +inline void write_file(const std::filesystem::path & path, const void * data, size_t size) { + std::ofstream output(path, std::ios::binary | std::ios::trunc); + if (!output) { + throw std::runtime_error("cannot create " + path.string()); + } + output.write(static_cast(data), static_cast(size)); + if (!output) { + throw std::runtime_error("cannot write " + path.string()); + } +} + +inline void write_file(const std::filesystem::path & path, const std::string & contents) { + std::ofstream output(path, std::ios::binary | std::ios::trunc); + if (!output) { + throw std::runtime_error("cannot create " + path.string()); + } + output << contents; + if (contents.empty() || contents.back() != '\n') { + output << '\n'; + } + if (!output) { + throw std::runtime_error("cannot write " + path.string()); + } +} + +inline void check_status(hrx_status_t status, const std::string & operation) { + if (ErrorResult error = take_status(status)) { + throw std::runtime_error(operation + ": " + *error); + } +} + +inline bool report_status(hrx_status_t status, const std::string & operation) { + if (ErrorResult error = take_status(status)) { + std::cerr << operation << ": " << *error << '\n'; + return false; + } + return true; +} + +} // namespace ggml::hrx::tool diff --git a/ggml/src/ggml-prism-quants.c b/ggml/src/ggml-prism-quants.c new file mode 100644 index 000000000000..b6c368e105ec --- /dev/null +++ b/ggml/src/ggml-prism-quants.c @@ -0,0 +1,243 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// Reference codecs for PrismML's PQ2_0 and PTQ1_0 (ggml-prism.h). The encoders and the PTQ1_0 decoder follow +// quantize_row_pq2_0_ref, quantize_row_ptq1_0_ref and dequantize_row_ptq1_0 in PrismML's llama.cpp fork +// (ggml/src/ggml-quants.c, branch prism, 87268f775), which carries this notice: +// +// MIT License +// +// Copyright (c) 2023-2026 The ggml authors +// +// Permission is hereby granted, free of charge, to any person obtaining a copy +// of this software and associated documentation files (the "Software"), to deal +// in the Software without restriction, including without limitation the rights +// to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +// copies of the Software, and to permit persons to whom the Software is +// furnished to do so, subject to the following conditions: +// +// The above copyright notice and this permission notice shall be included in all +// copies or substantial portions of the Software. +// +// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +// IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +// FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +// AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +// LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +// OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +// SOFTWARE. + +#include "ggml-prism.h" +#include "ggml-impl.h" + +#include +#include +#include +#include + +// ------------------------------------------------------------------------------------------------ PQ2_0 + +void quantize_row_pq2_0_ref(const float * GGML_RESTRICT x, block_pq2_0 * GGML_RESTRICT y, int64_t k) { + assert(k % QK_PQ2_0 == 0); + const int64_t nb = k / QK_PQ2_0; + + for (int64_t i = 0; i < nb; i++) { + float amax = 0.0f; + for (int j = 0; j < QK_PQ2_0; j++) { + const float a = fabsf(x[i * QK_PQ2_0 + j]); + if (a > amax) { + amax = a; + } + } + const float d = amax; + const float id = d > 0.0f ? 1.0f / d : 0.0f; + + y[i].d = GGML_FP32_TO_FP16(d); + memset(y[i].qs, 0, sizeof(y[i].qs)); + + // round(w / d) clamped to -1..2, stored + 1 + for (int j = 0; j < QK_PQ2_0; ++j) { + int q = (int) roundf(x[i * QK_PQ2_0 + j] * id) + 1; + q = q < 0 ? 0 : (q > 3 ? 3 : q); + y[i].qs[j / 4] |= (uint8_t) (q << (2 * (j % 4))); + } + } +} + +void dequantize_row_pq2_0(const block_pq2_0 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k) { + assert(k % QK_PQ2_0 == 0); + const int64_t nb = k / QK_PQ2_0; + + for (int64_t i = 0; i < nb; i++) { + const float d = GGML_FP16_TO_FP32(x[i].d); + for (int j = 0; j < QK_PQ2_0; ++j) { + const int q = (x[i].qs[j / 4] >> (2 * (j % 4))) & 3; + y[i * QK_PQ2_0 + j] = (float) (q - 1) * d; + } + } +} + +size_t quantize_pq2_0(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrow, int64_t n_per_row, + const float * imatrix) { + (void) imatrix; + const size_t row_size = ggml_row_size(GGML_TYPE_PQ2_0, n_per_row); + char * qrow = (char *) dst; + for (int64_t row = 0; row < nrow; ++row) { + quantize_row_pq2_0_ref(src, (block_pq2_0 *) qrow, n_per_row); + src += n_per_row; + qrow += row_size; + } + return nrow * row_size; +} + +// ------------------------------------------------------------------------------------------------ PTQ1_0 + +// The qs bytes are filled in stages of 32, 16 and 8 bytes (TQ1_0's 32-then-16 generalised to 24 bytes): +// a stage of c bytes holds 5c values, value n * c + m in trit n of byte m. With 24 bytes only the 16 and +// 8 stages are used. +static const size_t ptq1_0_stages[3] = { 32, 16, 8 }; + +void quantize_row_ptq1_0_ref(const float * GGML_RESTRICT x, block_ptq1_0 * GGML_RESTRICT y, int64_t k) { + assert(k % QK_PTQ1_0 == 0); + const int64_t nb = k / QK_PTQ1_0; + + for (int64_t i = 0; i < nb; i++) { + float amax = 0.0f; + for (int j = 0; j < QK_PTQ1_0; j++) { + amax = MAX(amax, fabsf(x[j])); + } + const float d = amax; + const float id = d ? 1.0f / d : 0.0f; + + y[i].d = GGML_FP32_TO_FP16(d); + + size_t j = 0; + for (size_t s = 0; s < 3; ++s) { + const size_t c = ptq1_0_stages[s]; + for (; j + c <= sizeof(y->qs); j += c) { + for (size_t m = 0; m < c; ++m) { + uint8_t q = 0; + for (size_t n = 0; n < 5; ++n) { + const int xi = lroundf(x[m + n * c] * id) + 1; // -1, 0, 1 -> 0, 1, 2 + q *= 3; + q += xi; + } + // ceiling division (243 == 3^5) + q = ((uint16_t) q * 256 + (243 - 1)) / 243; + y[i].qs[j + m] = q; + } + x += 5 * c; + } + } + for (size_t h = 0; h < sizeof(y->qh); ++h) { + uint8_t q = 0; + for (size_t m = 0; m < 4; ++m) { + const int xi = lroundf(x[h + m * sizeof(y->qh)] * id) + 1; + q *= 3; + q += xi; + } + // the first value goes to the most significant trit + q *= 3; + q = ((uint16_t) q * 256 + (243 - 1)) / 243; + y[i].qh[h] = q; + } + x += 4 * sizeof(y->qh); + } +} + +void ggml_ptq1_0_trits(const block_ptq1_0 * GGML_RESTRICT x, int8_t * GGML_RESTRICT t) { + static const uint8_t pow3[6] = { 1, 3, 9, 27, 81, 243 }; + + size_t j = 0; + for (size_t s = 0; s < 3; ++s) { + const size_t c = ptq1_0_stages[s]; + for (; j + c <= sizeof(x->qs); j += c) { + for (size_t n = 0; n < 5; ++n) { + for (size_t m = 0; m < c; ++m) { + const uint8_t q = x->qs[j + m] * pow3[n]; + const int16_t xi = ((uint16_t) q * 3) >> 8; + *t++ = (int8_t) (xi - 1); + } + } + } + } + for (size_t n = 0; n < 4; ++n) { + for (size_t h = 0; h < sizeof(x->qh); ++h) { + const uint8_t q = x->qh[h] * pow3[n]; + const int16_t xi = ((uint16_t) q * 3) >> 8; + *t++ = (int8_t) (xi - 1); + } + } +} + +void dequantize_row_ptq1_0(const block_ptq1_0 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k) { + assert(k % QK_PTQ1_0 == 0); + const int64_t nb = k / QK_PTQ1_0; + + int8_t t[QK_PTQ1_0]; + for (int64_t i = 0; i < nb; ++i) { + const float d = GGML_FP16_TO_FP32(x[i].d); + ggml_ptq1_0_trits(&x[i], t); + for (int j = 0; j < QK_PTQ1_0; ++j) { + y[i * QK_PTQ1_0 + j] = (float) t[j] * d; + } + } +} + +size_t quantize_ptq1_0(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrow, int64_t n_per_row, + const float * imatrix) { + (void) imatrix; // the trits come from the weights themselves + const size_t row_size = ggml_row_size(GGML_TYPE_PTQ1_0, n_per_row); + char * qrow = (char *) dst; + for (int64_t row = 0; row < nrow; ++row) { + quantize_row_ptq1_0_ref(src, (block_ptq1_0 *) qrow, n_per_row); + src += n_per_row; + qrow += row_size; + } + return nrow * row_size; +} + +// ------------------------------------------------------------------------------------------------ validation + +static bool prism_validate_fp16(ggml_half h, size_t i) { + const uint16_t f = h; + if ((f & 0x7c00) == 0x7c00) { + fprintf(stderr, "ggml_validate_row_data: found %s value at block %zu\n", (f & 0x03ff) ? "nan" : "inf", i); + return false; + } + return true; +} + +bool ggml_prism_validate_row_data(enum ggml_type type, const void * data, size_t nbytes) { + if (type == GGML_TYPE_PQ2_0) { + const block_pq2_0 * q = (const block_pq2_0 *) data; + for (size_t i = 0; i < nbytes / sizeof(block_pq2_0); ++i) { + if (!prism_validate_fp16(q[i].d, i)) { + return false; + } + } + return true; + } + if (type == GGML_TYPE_PTQ1_0) { + const block_ptq1_0 * q = (const block_ptq1_0 *) data; + for (size_t i = 0; i < nbytes / sizeof(block_ptq1_0); ++i) { + if (!prism_validate_fp16(q[i].d, i)) { + return false; + } + } + return true; + } + return false; +} diff --git a/ggml/src/ggml-prism.h b/ggml/src/ggml-prism.h new file mode 100644 index 000000000000..832995f70aab --- /dev/null +++ b/ggml/src/ggml-prism.h @@ -0,0 +1,68 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// PrismML's group-128 types PQ2_0 (ggml type 142) and PTQ1_0 (ggml type 143), as their llama.cpp fork +// (https://github.com/PrismML-Eng/llama.cpp, branch prism, MIT License, Copyright (c) 2023-2026 The ggml +// authors) defines them, so that their published GGUFs (Ternary-Bonsai-2-*) load as they are. The type ids +// and block layouts must stay identical to theirs; see ggml-prism-quants.c for the codecs. +// +// PQ2_0 (34 bytes / 128 values, 2.125 bpw): d (fp16), qs[32]. Value j is code c = (qs[j / 4] >> 2 (j % 4)) & 3, +// w = (c - 1) * d, so c = 0, 1, 2, 3 mean -1, 0, +1, +2 (the same codec as Q2_0, one scale per 128). +// PTQ1_0 (28 bytes / 128 values, 1.75 bpw): qs[24], qh[2], d (fp16). Trits t in {-1, 0, +1}, w = t * d, +// stored base 3 as in TQ1_0: a byte b holds trits n = 0..4 (resp. 0..3 for qh) read as +// ((uint8_t)(b * 3^n) * 3) >> 8, minus 1. Value order: qs[0..15] give values n * 16 + m (byte m, +// trit n) for 0..79, qs[16..23] give 80 + n * 8 + m, and qh[0..1] give 120 + n * 2 + h. +#pragma once + +#include "ggml-quants.h" + +#ifdef __cplusplus +extern "C" { +#endif + +#define QK_PQ2_0 128 +typedef struct { + ggml_half d; // scale + uint8_t qs[QK_PQ2_0 / 4]; // 2-bit codes, value j at bits 2 (j % 4) of byte j / 4 +} block_pq2_0; +static_assert(sizeof(block_pq2_0) == sizeof(ggml_half) + QK_PQ2_0 / 4, "wrong pq2_0 block size/padding"); + +#define QK_PTQ1_0 128 +typedef struct { + uint8_t qs[(QK_PTQ1_0 - 4 * QK_PTQ1_0 / 64) / 5]; // 24 bytes, 5 trits each -> 120 values + uint8_t qh[QK_PTQ1_0 / 64]; // 2 bytes, 4 trits each -> 8 values + ggml_half d; // scale +} block_ptq1_0; +static_assert(sizeof(block_ptq1_0) == sizeof(ggml_half) + QK_PTQ1_0 / 64 + (QK_PTQ1_0 - 4 * QK_PTQ1_0 / 64) / 5, + "wrong ptq1_0 block size/padding"); + +GGML_API void quantize_row_pq2_0_ref(const float * GGML_RESTRICT x, block_pq2_0 * GGML_RESTRICT y, int64_t k); +GGML_API void quantize_row_ptq1_0_ref(const float * GGML_RESTRICT x, block_ptq1_0 * GGML_RESTRICT y, int64_t k); + +GGML_API void dequantize_row_pq2_0(const block_pq2_0 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); +GGML_API void dequantize_row_ptq1_0(const block_ptq1_0 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); + +GGML_API size_t quantize_pq2_0(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix); +GGML_API size_t quantize_ptq1_0(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix); + +// The 128 trits of one PTQ1_0 block in value order (-1, 0, +1). +GGML_API void ggml_ptq1_0_trits(const block_ptq1_0 * GGML_RESTRICT x, int8_t * GGML_RESTRICT t); + +// ggml_validate_row_data for the two types: every scale must be a finite fp16. +GGML_API bool ggml_prism_validate_row_data(enum ggml_type type, const void * data, size_t nbytes); + +#ifdef __cplusplus +} +#endif diff --git a/ggml/src/ggml-quants.c b/ggml/src/ggml-quants.c index 1ebc50a763f1..2e2213a55826 100644 --- a/ggml/src/ggml-quants.c +++ b/ggml/src/ggml-quants.c @@ -2,6 +2,7 @@ #include "ggml-common.h" #include "ggml-quants.h" +#include "ggml-prism.h" #include "ggml-impl.h" #include "ggml-cpu/ggml-cpu-impl.h" #include "ggml-cpu.h" @@ -5650,6 +5651,9 @@ bool ggml_validate_row_data(enum ggml_type type, const void * data, size_t nbyte VALIDATE_ROW_DATA_D_F16_IMPL(block_iq4_nl, data, nb); } break; + case GGML_TYPE_PQ2_0: + case GGML_TYPE_PTQ1_0: + return ggml_prism_validate_row_data(type, data, nbytes); case GGML_TYPE_I8: case GGML_TYPE_I16: case GGML_TYPE_I32: diff --git a/ggml/src/ggml-vulkan/ggml-vulkan-dmabuf.inc b/ggml/src/ggml-vulkan/ggml-vulkan-dmabuf.inc new file mode 100644 index 000000000000..3c59b67c5eb0 --- /dev/null +++ b/ggml/src/ggml-vulkan/ggml-vulkan-dmabuf.inc @@ -0,0 +1,86 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Included by ggml-vulkan.cpp. Imports a dma-buf (device memory another API exported, for example HRX) as a +// Vulkan device buffer with no copy. Reached through the proc "ggml_backend_dev_buffer_from_dmabuf". + +#include + +static ggml_backend_buffer_t ggml_backend_vk_dev_buffer_from_dmabuf(ggml_backend_dev_t dev, int fd, size_t offset, size_t size) { + ggml_backend_vk_device_context * ctx = (ggml_backend_vk_device_context *)dev->context; + vk_device device = ggml_vk_get_device(ctx->device); + const off_t fd_size = lseek(fd, 0, SEEK_END); + if (!device->external_memory_dma_buf || fd_size <= 0 || offset + size > (size_t) fd_size || size > device->max_buffer_size) { + return nullptr; + } + + vk::BufferUsageFlags usage_flags = vk::BufferUsageFlagBits::eStorageBuffer | vk::BufferUsageFlagBits::eTransferSrc | vk::BufferUsageFlagBits::eTransferDst; + vk::MemoryAllocateFlags mem_flags {}; + if (device->buffer_device_address) { + usage_flags |= vk::BufferUsageFlagBits::eShaderDeviceAddress; + mem_flags |= vk::MemoryAllocateFlagBits::eDeviceAddress; + } + vk::ExternalMemoryBufferCreateInfo external_bci { vk::ExternalMemoryHandleTypeFlagBits::eDmaBufEXT }; + vk::BufferCreateInfo bci { vk::BufferCreateFlags(), size, usage_flags, vk::SharingMode::eExclusive, 0, nullptr }; + bci.setPNext(&external_bci); + + vk_buffer buf = std::make_shared(); + buf->buffer = device->device.createBuffer(bci); + const vk::MemoryRequirements mem_req = device->device.getBufferMemoryRequirements(buf->buffer); + const vk::MemoryFdPropertiesKHR fd_props = device->device.getMemoryFdPropertiesKHR(vk::ExternalMemoryHandleTypeFlagBits::eDmaBufEXT, fd); + const vk::PhysicalDeviceMemoryProperties mem_props = device->physical_device.getMemoryProperties(); + const uint32_t bits = fd_props.memoryTypeBits & mem_req.memoryTypeBits; + // the memory type only describes CPU mappings here: GPU access to exported device memory runs at full speed + uint32_t type = UINT32_MAX; + for (uint32_t i = 0; i < mem_props.memoryTypeCount && type == UINT32_MAX; i++) { + if ((bits & (1u << i)) && !(mem_props.memoryTypes[i].propertyFlags & vk::MemoryPropertyFlagBits::eHostCached)) { + type = i; + } + } + for (uint32_t i = 0; i < mem_props.memoryTypeCount && type == UINT32_MAX; i++) { + if (bits & (1u << i)) { + type = i; + } + } + if (type == UINT32_MAX || offset % mem_req.alignment != 0) { + device->device.destroyBuffer(buf->buffer); + return nullptr; + } + + const int import_fd = dup(fd); // on success Vulkan owns it + vk::MemoryAllocateFlagsInfo mem_flags_info { mem_flags }; + vk::ImportMemoryFdInfoKHR import_info { vk::ExternalMemoryHandleTypeFlagBits::eDmaBufEXT, import_fd }; + import_info.setPNext(&mem_flags_info); + try { + buf->device_memory = device->device.allocateMemory({ (vk::DeviceSize) fd_size, type, &import_info }); + } catch (const vk::SystemError & e) { + GGML_LOG_WARN("ggml_vulkan: dma-buf import failed (%s)\n", e.what()); + close(import_fd); + device->device.destroyBuffer(buf->buffer); + return nullptr; + } + device->device.bindBufferMemory(buf->buffer, buf->device_memory, offset); + // keep ggml on GPU-side paths: never map the imported memory on the host + buf->memory_property_flags = vk::MemoryPropertyFlagBits::eDeviceLocal; + buf->ptr = nullptr; + buf->device = device; + buf->size = size; + if (device->buffer_device_address) { + buf->bda_addr = device->device.getBufferAddress(vk::BufferDeviceAddressInfo(buf->buffer)); + } + + ggml_backend_vk_buffer_context * bufctx = new ggml_backend_vk_buffer_context(device, std::move(buf), device->name); + return ggml_backend_buffer_init(ggml_backend_vk_device_get_buffer_type(dev), ggml_backend_vk_buffer_interface, bufctx, size); +} diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 510fb6892b37..45aa11747340 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -760,6 +760,7 @@ struct vk_device_struct { uint64_t suballocation_block_size; uint64_t min_imported_host_pointer_alignment; bool external_memory_host {}; + bool external_memory_dma_buf {}; // [1bit] VK_KHR_external_memory_fd + VK_EXT_external_memory_dma_buf (ggml-vulkan-dmabuf.inc) bool fp16; bool bf16; bool pipeline_robustness; @@ -6101,6 +6102,8 @@ static vk_device ggml_vk_get_device(size_t idx) { device->memory_priority = true; } else if (strcmp("VK_EXT_external_memory_host", properties.extensionName) == 0) { device->external_memory_host = true; + } else if (strcmp("VK_EXT_external_memory_dma_buf", properties.extensionName) == 0) { + device->external_memory_dma_buf = true; #if defined(VK_EXT_shader_64bit_indexing) } else if (strcmp("VK_EXT_shader_64bit_indexing", properties.extensionName) == 0) { device->shader_64b_indexing = true; @@ -6451,6 +6454,11 @@ static vk_device ggml_vk_get_device(size_t idx) { device_extensions.push_back("VK_EXT_external_memory_host"); } + if (device->external_memory_dma_buf) { + device_extensions.push_back("VK_KHR_external_memory_fd"); + device_extensions.push_back("VK_EXT_external_memory_dma_buf"); + } + #if defined(VK_EXT_shader_64bit_indexing) VkPhysicalDeviceShader64BitIndexingFeaturesEXT shader_64bit_indexing_features {}; shader_64bit_indexing_features.sType = VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_SHADER_64_BIT_INDEXING_FEATURES_EXT; @@ -18436,11 +18444,28 @@ static ggml_backend_dev_t ggml_backend_vk_reg_get_device(ggml_backend_reg_t reg, return devices[device]; } +// [1bit] zero-copy sharing with another API over dma-buf (Linux) +#ifdef __linux__ +#include "ggml-vulkan-dmabuf.inc" +#endif + +static void * ggml_backend_vk_reg_get_proc_address(ggml_backend_reg_t reg, const char * name) { + UNUSED(reg); +#ifdef __linux__ + if (strcmp(name, "ggml_backend_dev_buffer_from_dmabuf") == 0) { + return (void *) ggml_backend_vk_dev_buffer_from_dmabuf; + } +#else + UNUSED(name); +#endif + return NULL; +} + static const struct ggml_backend_reg_i ggml_backend_vk_reg_i = { /* .get_name = */ ggml_backend_vk_reg_get_name, /* .get_device_count = */ ggml_backend_vk_reg_get_device_count, /* .get_device = */ ggml_backend_vk_reg_get_device, - /* .get_proc_address = */ NULL, + /* .get_proc_address = */ ggml_backend_vk_reg_get_proc_address, }; ggml_backend_reg_t ggml_backend_vk_reg() { diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index 59191c663eb0..a11d17e6d869 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -9,6 +9,7 @@ // FIXME: required here for quantization functions #include "ggml-quants.h" +#include "ggml-prism.h" #ifdef GGML_USE_CPU_HBM #include @@ -681,6 +682,22 @@ static const struct ggml_type_traits type_traits[GGML_TYPE_COUNT] = { .to_float = (ggml_to_float_t) dequantize_row_q1_0, .from_float_ref = (ggml_from_float_t) quantize_row_q1_0_ref, }, + [GGML_TYPE_PQ2_0] = { + .type_name = "pq2_0", + .blck_size = QK_PQ2_0, + .type_size = sizeof(block_pq2_0), + .is_quantized = true, + .to_float = (ggml_to_float_t) dequantize_row_pq2_0, + .from_float_ref = (ggml_from_float_t) quantize_row_pq2_0_ref, + }, + [GGML_TYPE_PTQ1_0] = { + .type_name = "ptq1_0", + .blck_size = QK_PTQ1_0, + .type_size = sizeof(block_ptq1_0), + .is_quantized = true, + .to_float = (ggml_to_float_t) dequantize_row_ptq1_0, + .from_float_ref = (ggml_from_float_t) quantize_row_ptq1_0_ref, + }, [GGML_TYPE_Q2_0] = { .type_name = "q2_0", .blck_size = QK2_0, @@ -1434,6 +1451,8 @@ enum ggml_type ggml_ftype_to_ggml_type(enum ggml_ftype ftype) { case GGML_FTYPE_MOSTLY_Q4_1: wtype = GGML_TYPE_Q4_1; break; case GGML_FTYPE_MOSTLY_Q1_0: wtype = GGML_TYPE_Q1_0; break; case GGML_FTYPE_MOSTLY_Q2_0: wtype = GGML_TYPE_Q2_0; break; + case GGML_FTYPE_MOSTLY_PQ2_0: wtype = GGML_TYPE_PQ2_0; break; + case GGML_FTYPE_MOSTLY_PTQ1_0: wtype = GGML_TYPE_PTQ1_0; break; case GGML_FTYPE_MOSTLY_Q5_0: wtype = GGML_TYPE_Q5_0; break; case GGML_FTYPE_MOSTLY_Q5_1: wtype = GGML_TYPE_Q5_1; break; case GGML_FTYPE_MOSTLY_Q8_0: wtype = GGML_TYPE_Q8_0; break; @@ -6397,10 +6416,12 @@ struct ggml_tensor * ggml_dsv4_hc_comb( // ggml_dsv4_hc_pre -struct ggml_tensor * ggml_dsv4_hc_pre( +static struct ggml_tensor * ggml_dsv4_hc_pre_impl( struct ggml_context * ctx, struct ggml_tensor * x, - struct ggml_tensor * weights) { + struct ggml_tensor * weights, + float scale, + bool gated) { GGML_ASSERT(x->type == GGML_TYPE_F32); GGML_ASSERT(weights->type == GGML_TYPE_F32); @@ -6410,13 +6431,22 @@ struct ggml_tensor * ggml_dsv4_hc_pre( GGML_ASSERT(hc > 0); GGML_ASSERT(x->ne[3] == 1); - GGML_ASSERT(weights->ne[0] == hc); - GGML_ASSERT(weights->ne[1] == n_tokens); - GGML_ASSERT(weights->ne[2] == 1); + if (gated) { + GGML_ASSERT(weights->ne[0] == n_embd); + GGML_ASSERT(weights->ne[1] == hc); + GGML_ASSERT(weights->ne[2] == n_tokens); + } else { + GGML_ASSERT(weights->ne[0] == hc); + GGML_ASSERT(weights->ne[1] == n_tokens); + GGML_ASSERT(weights->ne[2] == 1); + } GGML_ASSERT(weights->ne[3] == 1); struct ggml_tensor * result = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, n_embd, n_tokens); + ggml_set_op_params_f32(result, 0, scale); + ggml_set_op_params_i32(result, 1, gated ? 1 : 0); + result->op = GGML_OP_DSV4_HC_PRE; result->src[0] = x; result->src[1] = weights; @@ -6424,6 +6454,21 @@ struct ggml_tensor * ggml_dsv4_hc_pre( return result; } +struct ggml_tensor * ggml_dsv4_hc_pre( + struct ggml_context * ctx, + struct ggml_tensor * x, + struct ggml_tensor * weights) { + return ggml_dsv4_hc_pre_impl(ctx, x, weights, 1.0f, false); +} + +struct ggml_tensor * ggml_dsv4_hc_pre_gated( + struct ggml_context * ctx, + struct ggml_tensor * x, + struct ggml_tensor * gate, + float scale) { + return ggml_dsv4_hc_pre_impl(ctx, x, gate, scale, true); +} + // ggml_dsv4_hc_post struct ggml_tensor * ggml_dsv4_hc_post( @@ -6435,7 +6480,6 @@ struct ggml_tensor * ggml_dsv4_hc_post( GGML_ASSERT(x->type == GGML_TYPE_F32); GGML_ASSERT(residual->type == GGML_TYPE_F32); GGML_ASSERT(post->type == GGML_TYPE_F32); - GGML_ASSERT(comb->type == GGML_TYPE_F32); const int64_t n_embd = x->ne[0]; const int64_t n_tokens = x->ne[1]; @@ -6454,10 +6498,13 @@ struct ggml_tensor * ggml_dsv4_hc_post( GGML_ASSERT(post->ne[2] == 1); GGML_ASSERT(post->ne[3] == 1); - GGML_ASSERT(comb->ne[0] == hc); - GGML_ASSERT(comb->ne[1] == hc); - GGML_ASSERT(comb->ne[2] == n_tokens); - GGML_ASSERT(comb->ne[3] == 1); + if (comb) { + GGML_ASSERT(comb->type == GGML_TYPE_F32); + GGML_ASSERT(comb->ne[0] == hc); + GGML_ASSERT(comb->ne[1] == hc); + GGML_ASSERT(comb->ne[2] == n_tokens); + GGML_ASSERT(comb->ne[3] == 1); + } struct ggml_tensor * result = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd, hc, n_tokens); @@ -7939,6 +7986,8 @@ size_t ggml_quantize_chunk( switch (type) { case GGML_TYPE_Q1_0: result = quantize_q1_0 (src + start, (char *) dst + start_row * row_size, nrows, n_per_row, imatrix); break; case GGML_TYPE_Q2_0: result = quantize_q2_0 (src + start, (char *) dst + start_row * row_size, nrows, n_per_row, imatrix); break; + case GGML_TYPE_PQ2_0: result = quantize_pq2_0 (src + start, (char *) dst + start_row * row_size, nrows, n_per_row, imatrix); break; + case GGML_TYPE_PTQ1_0: result = quantize_ptq1_0 (src + start, (char *) dst + start_row * row_size, nrows, n_per_row, imatrix); break; case GGML_TYPE_Q4_0: result = quantize_q4_0 (src + start, (char *) dst + start_row * row_size, nrows, n_per_row, imatrix); break; case GGML_TYPE_Q4_1: result = quantize_q4_1 (src + start, (char *) dst + start_row * row_size, nrows, n_per_row, imatrix); break; case GGML_TYPE_Q5_0: result = quantize_q5_0 (src + start, (char *) dst + start_row * row_size, nrows, n_per_row, imatrix); break; diff --git a/gguf-py/gguf/constants.py b/gguf-py/gguf/constants.py index d86e614d8fd4..b8e511b77774 100644 --- a/gguf-py/gguf/constants.py +++ b/gguf-py/gguf/constants.py @@ -211,6 +211,19 @@ class HyperConnection: COUNT = "{arch}.hyper_connection.count" SINKHORN_ITERATIONS = "{arch}.hyper_connection.sinkhorn_iterations" EPSILON = "{arch}.hyper_connection.epsilon" + # absent means the mix projection is full rank (DeepSeek-V4 behaviour) + LOW_RANK = "{arch}.hyper_connection.low_rank" + + class PerLayerEmbedding: + LAYERS = "{arch}.ple.layers" + NGRAM_SIZE = "{arch}.ple.ngram_size" + HEADS_PER_NGRAM = "{arch}.ple.heads_per_ngram" + CONV_KERNEL = "{arch}.ple.conv_kernel" + LAYER_MULTIPLIERS = "{arch}.ple.layer_multipliers" + HEAD_OFFSETS = "{arch}.ple.head_offsets" + HEAD_VOCAB_SIZES = "{arch}.ple.head_vocab_sizes" + EOS_TOKEN_ID = "{arch}.ple.eos_token_id" + IMAGE_TOKEN_ID = "{arch}.ple.image_token_id" class Rope: DIMENSION_COUNT = "{arch}.rope.dimension_count" @@ -343,6 +356,7 @@ class ClipVision: USE_GELU = "clip.use_gelu" USE_SILU = "clip.use_silu" N_WA_PATTERN = "clip.vision.n_wa_pattern" # used by qwen2.5vl + DECODE_NON_CAUSAL = "clip.vision.decode_non_causal" # used by zaya1-vl WA_LAYER_INDEXES = "clip.vision.wa_layer_indexes" # used by youtuvl WA_PATTERN_MODE = "clip.vision.wa_pattern_mode" # used by mimovl, per-layer -1/0/1 IS_DEEPSTACK_LAYERS = "clip.vision.is_deepstack_layers" @@ -430,6 +444,9 @@ class MODEL_ARCH(IntEnum): GPT2 = auto() GPTJ = auto() GPTNEOX = auto() + OPT = auto() + CODEGEN = auto() + GPTNEO = auto() MPT = auto() STARCODER = auto() REFACT = auto() @@ -454,6 +471,7 @@ class MODEL_ARCH(IntEnum): QWEN3VLMOE = auto() QWEN35 = auto() QWEN35MOE = auto() + QWEN4EXP = auto() PHI2 = auto() PHI3 = auto() PHIMOE = auto() @@ -557,6 +575,7 @@ class MODEL_ARCH(IntEnum): TALKIE = auto() MELLUM = auto() NANBEIGE = auto() + ZAYA = auto() class VISION_PROJECTOR_TYPE(IntEnum): @@ -587,6 +606,9 @@ class MODEL_TENSOR(IntEnum): HC_HEAD_FN = auto() HC_HEAD_BASE = auto() HC_HEAD_SCALE = auto() + HC_HEAD_NORM = auto() # qwen4exp + HC_HEAD_DOWN = auto() # qwen4exp + HC_HEAD_UP = auto() # qwen4exp ROPE_FREQS = auto() ROPE_FACTORS_LONG = auto() ROPE_FACTORS_SHORT = auto() @@ -670,6 +692,35 @@ class MODEL_TENSOR(IntEnum): SSM_BETA = auto() # Kimi Linear qwen3.5 SSM_G_A = auto() # Kimi Linear SSM_G_B = auto() # Kimi Linear + CCA_CONV_GRP = auto() # Zaya + CCA_K_SCALE = auto() # Zaya + CCA_VAL_PROJ1 = auto() # Zaya + CCA_VAL_PROJ2 = auto() # Zaya + RES_SCALE_HS = auto() # Zaya + RES_SCALE_RES = auto() # Zaya + RES_SCALE_HS_MLP = auto() # Zaya + RES_SCALE_RES_MLP = auto() # Zaya + RES_SCALE_HS_FINAL = auto() # Zaya + RES_SCALE_RES_FINAL = auto() # Zaya + INPUT_HIDDEN_STATES_SCALE = auto() # Zaya + ZAYA_ROUTER_MLP2 = auto() # Zaya + ZAYA_ROUTER_MLP4 = auto() # Zaya + ZAYA_ROUTER_BIASES = auto() # Zaya + ZAYA_ROUTER_EDA_SCALE = auto() # Zaya + ZAYA_VLORA_Q_A = auto() # Zaya1-VL + ZAYA_VLORA_Q_B = auto() # Zaya1-VL + ZAYA_VLORA_K_A = auto() # Zaya1-VL + ZAYA_VLORA_K_B = auto() # Zaya1-VL + ZAYA_VLORA_V1_A = auto() # Zaya1-VL + ZAYA_VLORA_V1_B = auto() # Zaya1-VL + ZAYA_VLORA_V2_A = auto() # Zaya1-VL + ZAYA_VLORA_V2_B = auto() # Zaya1-VL + ZAYA_VLORA_O_A = auto() # Zaya1-VL + ZAYA_VLORA_O_B = auto() # Zaya1-VL + ZAYA_VLORA_UP_EXPS_A = auto() # Zaya1-VL + ZAYA_VLORA_UP_EXPS_B = auto() # Zaya1-VL + ZAYA_VLORA_DOWN_EXPS_A = auto() # Zaya1-VL + ZAYA_VLORA_DOWN_EXPS_B = auto() # Zaya1-VL TIME_MIX_W0 = auto() TIME_MIX_W1 = auto() TIME_MIX_W2 = auto() @@ -724,6 +775,20 @@ class MODEL_TENSOR(IntEnum): HC_FFN_FN = auto() HC_FFN_BASE = auto() HC_FFN_SCALE = auto() + HC_ATTN_NORM = auto() # qwen4exp + HC_ATTN_DOWN = auto() # qwen4exp + HC_ATTN_UP = auto() # qwen4exp + HC_ATTN_INJECT = auto() # qwen4exp + HC_FFN_NORM = auto() # qwen4exp + HC_FFN_DOWN = auto() # qwen4exp + HC_FFN_UP = auto() # qwen4exp + HC_FFN_INJECT = auto() # qwen4exp + PLE_KEY = auto() # qwen4exp + PLE_VALUE = auto() # qwen4exp + PLE_NORM_KEY = auto() # qwen4exp + PLE_NORM_QUERY = auto() # qwen4exp + PLE_NORM_CONV = auto() # qwen4exp + PLE_CONV1D = auto() # qwen4exp ATTN_COMPRESSOR_WKV = auto() ATTN_COMPRESSOR_WGATE = auto() ATTN_COMPRESSOR_APE = auto() @@ -987,6 +1052,9 @@ class MODEL_TENSOR(IntEnum): NEXTN_HNORM = auto() NEXTN_SHARED_HEAD_HEAD = auto() NEXTN_SHARED_HEAD_NORM = auto() + NEXTN_HC_HEAD_NORM = auto() + NEXTN_HC_HEAD_DOWN = auto() + NEXTN_HC_HEAD_UP = auto() # eagle3 FC = auto() # feature fusion layer D2T = auto() # draft to target vocabulary mapping @@ -1041,6 +1109,9 @@ class MODEL_TENSOR(IntEnum): MODEL_ARCH.GPT2: "gpt2", MODEL_ARCH.GPTJ: "gptj", MODEL_ARCH.GPTNEOX: "gptneox", + MODEL_ARCH.OPT: "opt", + MODEL_ARCH.CODEGEN: "codegen", + MODEL_ARCH.GPTNEO: "gptneo", MODEL_ARCH.MPT: "mpt", MODEL_ARCH.STARCODER: "starcoder", MODEL_ARCH.REFACT: "refact", @@ -1065,6 +1136,7 @@ class MODEL_TENSOR(IntEnum): MODEL_ARCH.QWEN3VLMOE: "qwen3vlmoe", MODEL_ARCH.QWEN35: "qwen35", MODEL_ARCH.QWEN35MOE: "qwen35moe", + MODEL_ARCH.QWEN4EXP: "qwen4exp", MODEL_ARCH.PHI2: "phi2", MODEL_ARCH.PHI3: "phi3", MODEL_ARCH.PHIMOE: "phimoe", @@ -1169,6 +1241,7 @@ class MODEL_TENSOR(IntEnum): MODEL_ARCH.TALKIE: "talkie", MODEL_ARCH.MELLUM: "mellum", MODEL_ARCH.NANBEIGE: "nanbeige", + MODEL_ARCH.ZAYA: "zaya", } VISION_PROJECTOR_TYPE_NAMES: dict[VISION_PROJECTOR_TYPE, str] = { @@ -1197,6 +1270,9 @@ class MODEL_TENSOR(IntEnum): MODEL_TENSOR.HC_HEAD_FN: "output_hc_fn", MODEL_TENSOR.HC_HEAD_BASE: "output_hc_base", MODEL_TENSOR.HC_HEAD_SCALE: "output_hc_scale", + MODEL_TENSOR.HC_HEAD_NORM: "output_hc_norm", # qwen4exp + MODEL_TENSOR.HC_HEAD_DOWN: "output_hc_down", # qwen4exp + MODEL_TENSOR.HC_HEAD_UP: "output_hc_up", # qwen4exp MODEL_TENSOR.ROPE_FREQS: "rope_freqs", MODEL_TENSOR.ROPE_FACTORS_LONG: "rope_factors_long", MODEL_TENSOR.ROPE_FACTORS_SHORT: "rope_factors_short", @@ -1280,6 +1356,35 @@ class MODEL_TENSOR(IntEnum): MODEL_TENSOR.SSM_BETA: "blk.{bid}.ssm_beta", # Kimi Linear qwen3.5 MODEL_TENSOR.SSM_G_A: "blk.{bid}.ssm_g_a", # Kimi Linear MODEL_TENSOR.SSM_G_B: "blk.{bid}.ssm_g_b", # Kimi Linear + MODEL_TENSOR.CCA_CONV_GRP: "blk.{bid}.cca_conv_grp", # Zaya + MODEL_TENSOR.CCA_K_SCALE: "blk.{bid}.cca_k_scale", # Zaya + MODEL_TENSOR.CCA_VAL_PROJ1: "blk.{bid}.cca_val_proj1", # Zaya + MODEL_TENSOR.CCA_VAL_PROJ2: "blk.{bid}.cca_val_proj2", # Zaya + MODEL_TENSOR.RES_SCALE_HS: "blk.{bid}.res_scale_hs", # Zaya + MODEL_TENSOR.RES_SCALE_RES: "blk.{bid}.res_scale_res", # Zaya + MODEL_TENSOR.RES_SCALE_HS_MLP: "blk.{bid}.res_scale_hs_mlp", # Zaya + MODEL_TENSOR.RES_SCALE_RES_MLP: "blk.{bid}.res_scale_res_mlp", # Zaya + MODEL_TENSOR.RES_SCALE_HS_FINAL: "res_scale_hs", # Zaya + MODEL_TENSOR.RES_SCALE_RES_FINAL: "res_scale_res", # Zaya + MODEL_TENSOR.INPUT_HIDDEN_STATES_SCALE: "input_hidden_states_scale", # Zaya + MODEL_TENSOR.ZAYA_ROUTER_MLP2: "blk.{bid}.zaya_router_mlp2", # Zaya + MODEL_TENSOR.ZAYA_ROUTER_MLP4: "blk.{bid}.zaya_router_mlp4", # Zaya + MODEL_TENSOR.ZAYA_ROUTER_BIASES: "blk.{bid}.zaya_router_biases", # Zaya + MODEL_TENSOR.ZAYA_ROUTER_EDA_SCALE: "blk.{bid}.zaya_router_eda", # Zaya + MODEL_TENSOR.ZAYA_VLORA_Q_A: "blk.{bid}.zaya_vlora_q_a", # Zaya1-VL + MODEL_TENSOR.ZAYA_VLORA_Q_B: "blk.{bid}.zaya_vlora_q_b", # Zaya1-VL + MODEL_TENSOR.ZAYA_VLORA_K_A: "blk.{bid}.zaya_vlora_k_a", # Zaya1-VL + MODEL_TENSOR.ZAYA_VLORA_K_B: "blk.{bid}.zaya_vlora_k_b", # Zaya1-VL + MODEL_TENSOR.ZAYA_VLORA_V1_A: "blk.{bid}.zaya_vlora_v1_a", # Zaya1-VL + MODEL_TENSOR.ZAYA_VLORA_V1_B: "blk.{bid}.zaya_vlora_v1_b", # Zaya1-VL + MODEL_TENSOR.ZAYA_VLORA_V2_A: "blk.{bid}.zaya_vlora_v2_a", # Zaya1-VL + MODEL_TENSOR.ZAYA_VLORA_V2_B: "blk.{bid}.zaya_vlora_v2_b", # Zaya1-VL + MODEL_TENSOR.ZAYA_VLORA_O_A: "blk.{bid}.zaya_vlora_o_a", # Zaya1-VL + MODEL_TENSOR.ZAYA_VLORA_O_B: "blk.{bid}.zaya_vlora_o_b", # Zaya1-VL + MODEL_TENSOR.ZAYA_VLORA_UP_EXPS_A: "blk.{bid}.zaya_vlora_gate_up_exps_a", # Zaya1-VL + MODEL_TENSOR.ZAYA_VLORA_UP_EXPS_B: "blk.{bid}.zaya_vlora_gate_up_exps_b", # Zaya1-VL + MODEL_TENSOR.ZAYA_VLORA_DOWN_EXPS_A: "blk.{bid}.zaya_vlora_down_exps_a", # Zaya1-VL + MODEL_TENSOR.ZAYA_VLORA_DOWN_EXPS_B: "blk.{bid}.zaya_vlora_down_exps_b", # Zaya1-VL MODEL_TENSOR.TIME_MIX_W0: "blk.{bid}.time_mix_w0", MODEL_TENSOR.TIME_MIX_W1: "blk.{bid}.time_mix_w1", MODEL_TENSOR.TIME_MIX_W2: "blk.{bid}.time_mix_w2", @@ -1334,6 +1439,20 @@ class MODEL_TENSOR(IntEnum): MODEL_TENSOR.HC_FFN_FN: "blk.{bid}.hc_ffn_fn", MODEL_TENSOR.HC_FFN_BASE: "blk.{bid}.hc_ffn_base", MODEL_TENSOR.HC_FFN_SCALE: "blk.{bid}.hc_ffn_scale", + MODEL_TENSOR.HC_ATTN_NORM: "blk.{bid}.hc_attn_norm", # qwen4exp + MODEL_TENSOR.HC_ATTN_DOWN: "blk.{bid}.hc_attn_down", # qwen4exp + MODEL_TENSOR.HC_ATTN_UP: "blk.{bid}.hc_attn_up", # qwen4exp + MODEL_TENSOR.HC_ATTN_INJECT: "blk.{bid}.hc_attn_inject", # qwen4exp + MODEL_TENSOR.HC_FFN_NORM: "blk.{bid}.hc_ffn_norm", # qwen4exp + MODEL_TENSOR.HC_FFN_DOWN: "blk.{bid}.hc_ffn_down", # qwen4exp + MODEL_TENSOR.HC_FFN_UP: "blk.{bid}.hc_ffn_up", # qwen4exp + MODEL_TENSOR.HC_FFN_INJECT: "blk.{bid}.hc_ffn_inject", # qwen4exp + MODEL_TENSOR.PLE_KEY: "blk.{bid}.ple_key", # qwen4exp + MODEL_TENSOR.PLE_VALUE: "blk.{bid}.ple_value", # qwen4exp + MODEL_TENSOR.PLE_NORM_KEY: "blk.{bid}.ple_norm_key", # qwen4exp + MODEL_TENSOR.PLE_NORM_QUERY: "blk.{bid}.ple_norm_query", # qwen4exp + MODEL_TENSOR.PLE_NORM_CONV: "blk.{bid}.ple_norm_conv", # qwen4exp + MODEL_TENSOR.PLE_CONV1D: "blk.{bid}.ple_conv1d", # qwen4exp MODEL_TENSOR.ATTN_COMPRESSOR_WKV: "blk.{bid}.attn_compressor_kv", MODEL_TENSOR.ATTN_COMPRESSOR_WGATE: "blk.{bid}.attn_compressor_gate", MODEL_TENSOR.ATTN_COMPRESSOR_APE: "blk.{bid}.attn_compressor_ape", @@ -1630,6 +1749,9 @@ class MODEL_TENSOR(IntEnum): MODEL_TENSOR.NEXTN_HNORM: "blk.{bid}.nextn.hnorm", MODEL_TENSOR.NEXTN_SHARED_HEAD_HEAD: "blk.{bid}.nextn.shared_head_head", MODEL_TENSOR.NEXTN_SHARED_HEAD_NORM: "blk.{bid}.nextn.shared_head_norm", + MODEL_TENSOR.NEXTN_HC_HEAD_NORM: "blk.{bid}.nextn.hc_head_norm", + MODEL_TENSOR.NEXTN_HC_HEAD_DOWN: "blk.{bid}.nextn.hc_head_down", + MODEL_TENSOR.NEXTN_HC_HEAD_UP: "blk.{bid}.nextn.hc_head_up", MODEL_TENSOR.FC: "fc", MODEL_TENSOR.DSPARK_MARKOV_W1: "markov_w1", MODEL_TENSOR.DSPARK_MARKOV_W2: "markov_w2", @@ -2104,6 +2226,44 @@ class MODEL_TENSOR(IntEnum): MODEL_TENSOR.FFN_UP, MODEL_TENSOR.FFN_DOWN, ], + MODEL_ARCH.OPT: [ + MODEL_TENSOR.TOKEN_EMBD, + MODEL_TENSOR.POS_EMBD, + MODEL_TENSOR.OUTPUT_NORM, + MODEL_TENSOR.OUTPUT, + MODEL_TENSOR.ATTN_NORM, + MODEL_TENSOR.ATTN_Q, + MODEL_TENSOR.ATTN_K, + MODEL_TENSOR.ATTN_V, + MODEL_TENSOR.ATTN_OUT, + MODEL_TENSOR.FFN_NORM, + MODEL_TENSOR.FFN_UP, + MODEL_TENSOR.FFN_DOWN, + ], + MODEL_ARCH.CODEGEN: [ + MODEL_TENSOR.TOKEN_EMBD, + MODEL_TENSOR.OUTPUT_NORM, + MODEL_TENSOR.OUTPUT, + MODEL_TENSOR.ATTN_NORM, + MODEL_TENSOR.ATTN_QKV, + MODEL_TENSOR.ATTN_OUT, + MODEL_TENSOR.FFN_UP, + MODEL_TENSOR.FFN_DOWN, + ], + MODEL_ARCH.GPTNEO: [ + MODEL_TENSOR.TOKEN_EMBD, + MODEL_TENSOR.POS_EMBD, + MODEL_TENSOR.OUTPUT_NORM, + MODEL_TENSOR.OUTPUT, + MODEL_TENSOR.ATTN_NORM, + MODEL_TENSOR.ATTN_Q, + MODEL_TENSOR.ATTN_K, + MODEL_TENSOR.ATTN_V, + MODEL_TENSOR.ATTN_OUT, + MODEL_TENSOR.FFN_NORM, + MODEL_TENSOR.FFN_UP, + MODEL_TENSOR.FFN_DOWN, + ], MODEL_ARCH.MPT: [ MODEL_TENSOR.TOKEN_EMBD, MODEL_TENSOR.OUTPUT_NORM, @@ -2435,6 +2595,66 @@ class MODEL_TENSOR(IntEnum): MODEL_TENSOR.NEXTN_SHARED_HEAD_HEAD, MODEL_TENSOR.NEXTN_SHARED_HEAD_NORM, ], + MODEL_ARCH.QWEN4EXP: [ + MODEL_TENSOR.TOKEN_EMBD, + MODEL_TENSOR.OUTPUT, + # no OUTPUT_NORM / ATTN_NORM / ATTN_POST_NORM: hyper-connections replace every layer norm + MODEL_TENSOR.HC_HEAD_NORM, + MODEL_TENSOR.HC_HEAD_DOWN, + MODEL_TENSOR.HC_HEAD_UP, + MODEL_TENSOR.HC_ATTN_NORM, + MODEL_TENSOR.HC_ATTN_DOWN, + MODEL_TENSOR.HC_ATTN_UP, + MODEL_TENSOR.HC_ATTN_INJECT, + MODEL_TENSOR.HC_FFN_NORM, + MODEL_TENSOR.HC_FFN_DOWN, + MODEL_TENSOR.HC_FFN_UP, + MODEL_TENSOR.HC_FFN_INJECT, + # full attention layers: ATTN_Q holds [q|gate] interleaved per head + MODEL_TENSOR.ATTN_Q, + MODEL_TENSOR.ATTN_Q_NORM, + MODEL_TENSOR.ATTN_K, + MODEL_TENSOR.ATTN_K_NORM, + MODEL_TENSOR.ATTN_V, + MODEL_TENSOR.ATTN_OUT, + MODEL_TENSOR.INDEXER_Q_PROJ, + MODEL_TENSOR.INDEXER_K_PROJ, + MODEL_TENSOR.INDEXER_Q_NORM, + MODEL_TENSOR.INDEXER_K_NORM, + MODEL_TENSOR.ATTN_QKV, + MODEL_TENSOR.ATTN_GATE, + MODEL_TENSOR.SSM_A, + MODEL_TENSOR.SSM_CONV1D, + MODEL_TENSOR.SSM_DT, + MODEL_TENSOR.SSM_NORM, + MODEL_TENSOR.SSM_BETA, + MODEL_TENSOR.SSM_ALPHA, + MODEL_TENSOR.SSM_OUT, + MODEL_TENSOR.FFN_GATE_INP, + MODEL_TENSOR.FFN_GATE_INP_SHEXP, + MODEL_TENSOR.FFN_UP_SHEXP, + MODEL_TENSOR.FFN_DOWN_SHEXP, + MODEL_TENSOR.FFN_GATE_SHEXP, + MODEL_TENSOR.FFN_DOWN_EXP, + MODEL_TENSOR.FFN_UP_EXP, + MODEL_TENSOR.FFN_GATE_EXP, + MODEL_TENSOR.FFN_GATE_UP_EXP, + MODEL_TENSOR.PER_LAYER_TOKEN_EMBD, + MODEL_TENSOR.PLE_KEY, + MODEL_TENSOR.PLE_VALUE, + MODEL_TENSOR.PLE_NORM_KEY, + MODEL_TENSOR.PLE_NORM_QUERY, + MODEL_TENSOR.PLE_NORM_CONV, + MODEL_TENSOR.PLE_CONV1D, + MODEL_TENSOR.NEXTN_EH_PROJ, + MODEL_TENSOR.NEXTN_EMBED_TOKENS, + MODEL_TENSOR.NEXTN_ENORM, + MODEL_TENSOR.NEXTN_HNORM, + MODEL_TENSOR.NEXTN_SHARED_HEAD_HEAD, + MODEL_TENSOR.NEXTN_HC_HEAD_NORM, + MODEL_TENSOR.NEXTN_HC_HEAD_DOWN, + MODEL_TENSOR.NEXTN_HC_HEAD_UP, + ], MODEL_ARCH.PLAMO: [ MODEL_TENSOR.TOKEN_EMBD, MODEL_TENSOR.OUTPUT_NORM, @@ -4603,6 +4823,48 @@ class MODEL_TENSOR(IntEnum): MODEL_TENSOR.FFN_DOWN, MODEL_TENSOR.FFN_UP, ], + MODEL_ARCH.ZAYA: [ + MODEL_TENSOR.TOKEN_EMBD, + MODEL_TENSOR.OUTPUT_NORM, + MODEL_TENSOR.INPUT_HIDDEN_STATES_SCALE, + MODEL_TENSOR.ATTN_NORM, + MODEL_TENSOR.ATTN_POST_NORM, + MODEL_TENSOR.ATTN_Q, + MODEL_TENSOR.ATTN_K, + MODEL_TENSOR.ATTN_OUT, + MODEL_TENSOR.SSM_CONV1D, + MODEL_TENSOR.CCA_CONV_GRP, + MODEL_TENSOR.CCA_K_SCALE, + MODEL_TENSOR.CCA_VAL_PROJ1, + MODEL_TENSOR.CCA_VAL_PROJ2, + MODEL_TENSOR.RES_SCALE_HS, + MODEL_TENSOR.RES_SCALE_RES, + MODEL_TENSOR.RES_SCALE_HS_MLP, + MODEL_TENSOR.RES_SCALE_RES_MLP, + MODEL_TENSOR.FFN_GATE_INP, + MODEL_TENSOR.FFN_NORM, + MODEL_TENSOR.FFN_GATE, + MODEL_TENSOR.ZAYA_ROUTER_MLP2, + MODEL_TENSOR.ZAYA_ROUTER_MLP4, + MODEL_TENSOR.ZAYA_ROUTER_BIASES, + MODEL_TENSOR.ZAYA_ROUTER_EDA_SCALE, + MODEL_TENSOR.ZAYA_VLORA_Q_A, + MODEL_TENSOR.ZAYA_VLORA_Q_B, + MODEL_TENSOR.ZAYA_VLORA_K_A, + MODEL_TENSOR.ZAYA_VLORA_K_B, + MODEL_TENSOR.ZAYA_VLORA_V1_A, + MODEL_TENSOR.ZAYA_VLORA_V1_B, + MODEL_TENSOR.ZAYA_VLORA_V2_A, + MODEL_TENSOR.ZAYA_VLORA_V2_B, + MODEL_TENSOR.ZAYA_VLORA_O_A, + MODEL_TENSOR.ZAYA_VLORA_O_B, + MODEL_TENSOR.ZAYA_VLORA_UP_EXPS_A, + MODEL_TENSOR.ZAYA_VLORA_UP_EXPS_B, + MODEL_TENSOR.ZAYA_VLORA_DOWN_EXPS_A, + MODEL_TENSOR.ZAYA_VLORA_DOWN_EXPS_B, + MODEL_TENSOR.FFN_GATE_UP_EXP, + MODEL_TENSOR.FFN_DOWN_EXP, + ], } # tensors that will not be serialized @@ -4740,6 +5002,8 @@ class GGMLQuantizationType(IntEnum): NVFP4 = 40 Q1_0 = 41 Q2_0 = 42 + PQ2_0 = 142 # PrismML group-128 2-bit + PTQ1_0 = 143 # PrismML group-128 ternary class ExpertGatingFuncType(IntEnum): @@ -4796,6 +5060,8 @@ class LlamaFileType(IntEnum): MOSTLY_NVFP4 = 39 # except 1d tensors MOSTLY_Q1_0 = 40 # except 1d tensors MOSTLY_Q2_0 = 41 # except 1d tensors + MOSTLY_PQ2_0 = 141 # except 1d tensors (PrismML) + MOSTLY_PTQ1_0 = 143 # except 1d tensors (PrismML) GUESSED = 1024 # not specified in the model file @@ -4925,6 +5191,8 @@ class VisionProjectorType: GGMLQuantizationType.NVFP4: (64, 4 + 32), GGMLQuantizationType.Q1_0: (128, 2 + 16), GGMLQuantizationType.Q2_0: (64, 2 + 16), + GGMLQuantizationType.PQ2_0: (128, 2 + 32), + GGMLQuantizationType.PTQ1_0: (128, 24 + 2 + 2), } diff --git a/gguf-py/gguf/gguf_writer.py b/gguf-py/gguf/gguf_writer.py index c5905164c356..345cca5750d5 100644 --- a/gguf-py/gguf/gguf_writer.py +++ b/gguf-py/gguf/gguf_writer.py @@ -995,6 +995,40 @@ def add_hyper_connection_sinkhorn_iterations(self, count: int) -> None: def add_hyper_connection_epsilon(self, value: float) -> None: self.add_float32(Keys.HyperConnection.EPSILON.format(arch=self.arch), value) + def add_hyper_connection_low_rank(self, value: int) -> None: + self.add_uint32(Keys.HyperConnection.LOW_RANK.format(arch=self.arch), value) + + def add_ple_layers(self, values: Sequence[int]) -> None: + self.add_array(Keys.PerLayerEmbedding.LAYERS.format(arch=self.arch), values) + + def add_ple_ngram_size(self, value: int) -> None: + self.add_uint32(Keys.PerLayerEmbedding.NGRAM_SIZE.format(arch=self.arch), value) + + def add_ple_heads_per_ngram(self, value: int) -> None: + self.add_uint32(Keys.PerLayerEmbedding.HEADS_PER_NGRAM.format(arch=self.arch), value) + + def add_ple_conv_kernel(self, value: int) -> None: + self.add_uint32(Keys.PerLayerEmbedding.CONV_KERNEL.format(arch=self.arch), value) + + # multipliers reach ~2.4e13; default INT32 inference would truncate them + def _add_u64_array(self, key: str, values: Sequence[int]) -> None: + self.add_key_value(key, list(values), GGUFValueType.ARRAY, GGUFValueType.UINT64) + + def add_ple_layer_multipliers(self, values: Sequence[int]) -> None: + self._add_u64_array(Keys.PerLayerEmbedding.LAYER_MULTIPLIERS.format(arch=self.arch), values) + + def add_ple_head_offsets(self, values: Sequence[int]) -> None: + self._add_u64_array(Keys.PerLayerEmbedding.HEAD_OFFSETS.format(arch=self.arch), values) + + def add_ple_head_vocab_sizes(self, values: Sequence[int]) -> None: + self._add_u64_array(Keys.PerLayerEmbedding.HEAD_VOCAB_SIZES.format(arch=self.arch), values) + + def add_ple_eos_token_id(self, value: int) -> None: + self.add_uint32(Keys.PerLayerEmbedding.EOS_TOKEN_ID.format(arch=self.arch), value) + + def add_ple_image_token_id(self, value: int) -> None: + self.add_uint32(Keys.PerLayerEmbedding.IMAGE_TOKEN_ID.format(arch=self.arch), value) + def add_attention_scale(self, value: float) -> None: self.add_float32(Keys.Attention.SCALE.format(arch=self.arch), value) @@ -1277,6 +1311,9 @@ def add_vision_use_silu(self, value: bool) -> None: def add_vision_projector_scale_factor(self, value: int) -> None: self.add_uint32(Keys.ClipVision.Projector.SCALE_FACTOR, value) + def add_vision_decode_non_causal(self, value: bool) -> None: + self.add_bool(Keys.ClipVision.DECODE_NON_CAUSAL, value) + def add_vision_n_wa_pattern(self, value: int) -> None: """Add window attention pattern interval for vision models. diff --git a/gguf-py/gguf/lazy.py b/gguf-py/gguf/lazy.py index acbc79258a31..6a0aee881107 100644 --- a/gguf-py/gguf/lazy.py +++ b/gguf-py/gguf/lazy.py @@ -226,3 +226,64 @@ def tofile(self, *args, **kwargs): return eager.tofile(*args, **kwargs) # TODO: __array_function__ + + +# Tensor written to file one row-chunk at a time +class LazyChunkedTensor: + + def __init__( + self, chunks: list[Callable[[], np.ndarray]], shape: tuple[int, ...], dtype: DTypeLike, + qtype: Any = None, byteswap: bool = False, + ): + self._chunks = chunks + self._qtype = qtype + self._byteswap = byteswap + self.shape = tuple(shape) + self.dtype = np.dtype(dtype) + + @property + def nbytes(self) -> int: + n = self.dtype.itemsize + for d in self.shape: + n *= d + return n + + def numpy(self) -> LazyChunkedTensor: + return self + + def quantize(self, qtype: Any) -> LazyChunkedTensor: + from .constants import GGMLQuantizationType + from .quants import QuantError, quant_shape_to_byte_shape + + if qtype == GGMLQuantizationType.F32: + shape, dtype = self.shape, np.dtype(np.float32) + elif qtype == GGMLQuantizationType.F16: + shape, dtype = self.shape, np.dtype(np.float16) + else: + try: + shape, dtype = quant_shape_to_byte_shape(self.shape, qtype), np.dtype(np.uint8) + except ValueError as e: + # raised here and not per chunk, so callers can still fall back to F16 + raise QuantError(str(e)) from e + return LazyChunkedTensor(self._chunks, shape, dtype, qtype, self._byteswap) + + def byteswap(self, inplace: bool = False) -> LazyChunkedTensor: + if inplace: + raise NotImplementedError("a chunked tensor cannot be byteswapped in place") + return LazyChunkedTensor(self._chunks, self.shape, self.dtype, self._qtype, not self._byteswap) + + def tofile(self, *args, **kwargs) -> None: + from .quants import quantize + + written = 0 + for load_chunk in self._chunks: + chunk = load_chunk() + if self._qtype is not None: + # exact only because chunks split on rows, and blocks never cross one + chunk = quantize(chunk, self._qtype) + if self._byteswap: + chunk = chunk.byteswap(inplace=False) + chunk.tofile(*args, **kwargs) + written += chunk.nbytes + del chunk + assert written == self.nbytes, f"chunked tensor wrote {written} bytes, expected {self.nbytes}" diff --git a/gguf-py/gguf/tensor_mapping.py b/gguf-py/gguf/tensor_mapping.py index 1e991b873cea..41fce3f6fbdc 100644 --- a/gguf-py/gguf/tensor_mapping.py +++ b/gguf-py/gguf/tensor_mapping.py @@ -9,6 +9,8 @@ class TensorNameMap: mappings_cfg: dict[MODEL_TENSOR, tuple[str, ...]] = { # Token embeddings MODEL_TENSOR.TOKEN_EMBD: ( + "embedding_proj", # pico_decoder + "model.decoder.embed_tokens", # opt "gpt_neox.embed_in", # gptneox "transformer.wte", # gpt2 gpt-j mpt refact qwen dbrx jais exaone "transformer.word_embeddings", # falcon @@ -67,6 +69,7 @@ class TensorNameMap: # Position embeddings MODEL_TENSOR.POS_EMBD: ( + "model.decoder.embed_positions", # opt "transformer.wpe", # gpt2 "embeddings.position_embeddings", # bert "wpe", # gpt2 @@ -75,6 +78,7 @@ class TensorNameMap: # Output MODEL_TENSOR.OUTPUT: ( + "de_embedding_proj", # pico_decoder "embed_out", # gptneox "lm_head", # gpt2 mpt falcon llama-hf baichuan qwen mamba dbrx jais nemotron exaone olmoe olmo2 phimoe plamo2 "output", # llama-pth bloom internlm2 @@ -95,6 +99,8 @@ class TensorNameMap: ), # Output norm MODEL_TENSOR.OUTPUT_NORM: ( + "output_norm", # pico_decoder + "model.decoder.final_layer_norm", # opt "gpt_neox.final_layer_norm", # gptneox "transformer.ln_f", # gpt2 gpt-j falcon jais exaone "model.norm", # llama-hf baichuan internlm2 olmoe olmo2 phimoe plamo2 @@ -184,6 +190,8 @@ class TensorNameMap: block_mappings_cfg: dict[MODEL_TENSOR, tuple[str, ...]] = { # Attention norm MODEL_TENSOR.ATTN_NORM: ( + "layers.{bid}.attention_norm", # pico_decoder + "model.decoder.layers.{bid}.self_attn_layer_norm", # opt "gpt_neox.layers.{bid}.input_layernorm", # gptneox "transformer.h.{bid}.ln_1", # gpt2 gpt-j refact qwen jais exaone "transformer.blocks.{bid}.norm_1", # mpt @@ -229,6 +237,7 @@ class TensorNameMap: # Attention query-key-value MODEL_TENSOR.ATTN_QKV: ( + "transformer.h.{bid}.attn.qkv_proj", # codegen "gpt_neox.layers.{bid}.attention.query_key_value", # gptneox "transformer.h.{bid}.attn.c_attn", # gpt2 qwen jais "transformer.blocks.{bid}.attn.Wqkv", # mpt @@ -253,6 +262,9 @@ class TensorNameMap: # Attention query MODEL_TENSOR.ATTN_Q: ( + "layers.{bid}.attention.q_proj", # pico_decoder + "transformer.h.{bid}.attn.attention.q_proj", # gpt-neo + "model.decoder.layers.{bid}.self_attn.q_proj", # opt "model.layers.{bid}.self_attn.q_proj", # llama-hf nemotron olmoe olmo2 phimoe "layers.{bid}.self_attn.q_proj", # embeddinggemma "model.layers.{bid}.self_attn.q_proj_no_perm", # llama-custom @@ -273,6 +285,9 @@ class TensorNameMap: # Attention key MODEL_TENSOR.ATTN_K: ( + "layers.{bid}.attention.k_proj", # pico_decoder + "transformer.h.{bid}.attn.attention.k_proj", # gpt-neo + "model.decoder.layers.{bid}.self_attn.k_proj", # opt "model.layers.{bid}.self_attn.k_proj", # llama-hf nemotron olmoe olmo2 phimoe "layers.{bid}.self_attn.k_proj", # embeddinggemma "model.layers.{bid}.self_attn.k_proj_no_perm", # llama-custom @@ -294,6 +309,9 @@ class TensorNameMap: # Attention value MODEL_TENSOR.ATTN_V: ( + "layers.{bid}.attention.v_proj", # pico_decoder + "transformer.h.{bid}.attn.attention.v_proj", # gpt-neo + "model.decoder.layers.{bid}.self_attn.v_proj", # opt "model.layers.{bid}.self_attn.v_proj", # llama-hf nemotron olmoe olmo2 phimoe "layers.{bid}.self_attn.v_proj", # embeddinggemma "layers.{bid}.attention.wv", # llama-pth @@ -314,6 +332,9 @@ class TensorNameMap: # Attention output MODEL_TENSOR.ATTN_OUT: ( + "layers.{bid}.attention.o_proj", # pico_decoder + "transformer.h.{bid}.attn.attention.out_proj", # gpt-neo + "model.decoder.layers.{bid}.self_attn.out_proj", # opt "gpt_neox.layers.{bid}.attention.dense", # gptneox "transformer.h.{bid}.attn.c_proj", # gpt2 refact qwen jais "transformer.blocks.{bid}.attn.out_proj", # mpt @@ -389,6 +410,9 @@ class TensorNameMap: # Feed-forward norm MODEL_TENSOR.FFN_NORM: ( + "layers.{bid}.swiglu_norm", # pico_decoder + "transformer.h.{bid}.ln_2", # gpt-neo + "model.decoder.layers.{bid}.final_layer_norm", # opt "gpt_neox.layers.{bid}.post_attention_layernorm", # gptneox "transformer.h.{bid}.ln_2", # gpt2 refact qwen jais exaone "h.{bid}.post_attention_layernorm", # bloom @@ -484,6 +508,9 @@ class TensorNameMap: # Feed-forward up MODEL_TENSOR.FFN_UP: ( + "layers.{bid}.swiglu.w_1", # pico_decoder + "transformer.h.{bid}.mlp.c_fc", # gpt-neo + "model.decoder.layers.{bid}.fc1", # opt "gpt_neox.layers.{bid}.mlp.dense_h_to_4h", # gptneox "transformer.h.{bid}.mlp.c_fc", # gpt2 jais "transformer.blocks.{bid}.ffn.up_proj", # mpt @@ -560,6 +587,7 @@ class TensorNameMap: # Feed-forward gate MODEL_TENSOR.FFN_GATE: ( + "layers.{bid}.swiglu.w_0", # pico_decoder "model.layers.{bid}.mlp.gate_proj", # llama-hf refact olmo2 "layers.{bid}.mlp.gate_proj", # embeddinggemma "layers.{bid}.feed_forward.w1", # llama-pth @@ -619,6 +647,9 @@ class TensorNameMap: # Feed-forward down MODEL_TENSOR.FFN_DOWN: ( + "layers.{bid}.swiglu.w_2", # pico_decoder + "transformer.h.{bid}.mlp.c_proj", # gpt-neo + "model.decoder.layers.{bid}.fc2", # opt "gpt_neox.layers.{bid}.mlp.dense_4h_to_h", # gptneox "transformer.h.{bid}.mlp.c_proj", # gpt2 refact qwen jais "transformer.blocks.{bid}.ffn.down_proj", # mpt @@ -2556,6 +2587,74 @@ class TensorNameMap: "model.layers.{bid}.post_attention_layernorm", ), }, + MODEL_ARCH.QWEN4EXP: { + MODEL_TENSOR.HC_ATTN_NORM: ( + "model.layers.{bid}.attn_hyper_connection.hc_norm", + ), + MODEL_TENSOR.HC_ATTN_DOWN: ( + "model.layers.{bid}.attn_hyper_connection.input_mix_weight_down", + ), + MODEL_TENSOR.HC_ATTN_UP: ( + "model.layers.{bid}.attn_hyper_connection.input_mix_weight_up", + ), + MODEL_TENSOR.HC_ATTN_INJECT: ( + "model.layers.{bid}.attn_hyper_connection.block_inject_weight", + ), + MODEL_TENSOR.HC_FFN_NORM: ( + "model.layers.{bid}.mlp_hyper_connection.hc_norm", + ), + MODEL_TENSOR.HC_FFN_DOWN: ( + "model.layers.{bid}.mlp_hyper_connection.input_mix_weight_down", + ), + MODEL_TENSOR.HC_FFN_UP: ( + "model.layers.{bid}.mlp_hyper_connection.input_mix_weight_up", + ), + MODEL_TENSOR.HC_FFN_INJECT: ( + "model.layers.{bid}.mlp_hyper_connection.block_inject_weight", + ), + MODEL_TENSOR.HC_HEAD_NORM: ( + "model.hyper_connection_mixer.hc_norm", + ), + MODEL_TENSOR.HC_HEAD_DOWN: ( + "model.hyper_connection_mixer.input_mix_weight_down", + ), + MODEL_TENSOR.HC_HEAD_UP: ( + "model.hyper_connection_mixer.input_mix_weight_up", + ), + MODEL_TENSOR.NEXTN_HC_HEAD_NORM: ( + "model.layers.{bid}.hyper_connection_mixer.hc_norm", + ), + MODEL_TENSOR.NEXTN_HC_HEAD_DOWN: ( + "model.layers.{bid}.hyper_connection_mixer.input_mix_weight_down", + ), + MODEL_TENSOR.NEXTN_HC_HEAD_UP: ( + "model.layers.{bid}.hyper_connection_mixer.input_mix_weight_up", + ), + MODEL_TENSOR.INDEXER_Q_NORM: ( + "model.layers.{bid}.self_attn.indexer.q_layernorm", + ), + MODEL_TENSOR.INDEXER_K_NORM: ( + "model.layers.{bid}.self_attn.indexer.k_layernorm", + ), + MODEL_TENSOR.PLE_KEY: ( + "model.layers.{bid}.ple.key_proj", + ), + MODEL_TENSOR.PLE_VALUE: ( + "model.layers.{bid}.ple.value_proj", + ), + MODEL_TENSOR.PLE_NORM_KEY: ( + "model.layers.{bid}.ple.norm_key", + ), + MODEL_TENSOR.PLE_NORM_QUERY: ( + "model.layers.{bid}.ple.norm_query", + ), + MODEL_TENSOR.PLE_NORM_CONV: ( + "model.layers.{bid}.ple.norm_conv", + ), + MODEL_TENSOR.PLE_CONV1D: ( + "model.layers.{bid}.ple.conv1d", + ), + }, } mapping: dict[str, tuple[MODEL_TENSOR, str]] diff --git a/include/llama.h b/include/llama.h index 6e53e2297235..bbc2da9a0018 100644 --- a/include/llama.h +++ b/include/llama.h @@ -43,10 +43,10 @@ #define LLAMA_FILE_MAGIC_GGSQ 0x67677371u // 'ggsq' #define LLAMA_SESSION_MAGIC LLAMA_FILE_MAGIC_GGSN -#define LLAMA_SESSION_VERSION 9 +#define LLAMA_SESSION_VERSION 10 #define LLAMA_STATE_SEQ_MAGIC LLAMA_FILE_MAGIC_GGSQ -#define LLAMA_STATE_SEQ_VERSION 2 +#define LLAMA_STATE_SEQ_VERSION 3 #ifdef __cplusplus extern "C" { @@ -156,6 +156,9 @@ extern "C" { LLAMA_FTYPE_MOSTLY_NVFP4 = 39, // except 1d tensors LLAMA_FTYPE_MOSTLY_Q1_0 = 40, // except 1d tensors LLAMA_FTYPE_MOSTLY_Q2_0 = 41, // except 1d tensors + LLAMA_FTYPE_MOSTLY_PQ2_0 = 141, // except 1d tensors (PrismML group-128 2-bit) + LLAMA_FTYPE_MOSTLY_PQ2_0_LEGACY = 142, // the same, as older PrismML files mark it + LLAMA_FTYPE_MOSTLY_PTQ1_0 = 143, // except 1d tensors (PrismML group-128 ternary) LLAMA_FTYPE_GUESSED = 1024, // not specified in the model file }; @@ -213,6 +216,12 @@ extern "C" { LLAMA_API const char * llama_load_mode_name(enum llama_load_mode load_mode); LLAMA_API enum llama_load_mode llama_load_mode_from_str(const char * str); + enum llama_lazy_mode { + LLAMA_LAZY_MODE_OFF = 0, // always read the whole tensor up front + LLAMA_LAZY_MODE_AUTO = 1, // lazy only for marked tensors larger than 4 GiB (requires mmap) + LLAMA_LAZY_MODE_ON = 2, // read the rows of tensors marked by the arch on demand (requires mmap) + }; + enum llama_context_type { LLAMA_CONTEXT_TYPE_DEFAULT = 0, LLAMA_CONTEXT_TYPE_MTP = 1, @@ -314,6 +323,8 @@ extern "C" { enum llama_split_mode split_mode; // how to split the model across multiple GPUs enum llama_load_mode load_mode; // how to load the model + enum llama_lazy_mode lazy_mode; // on-demand reading of tensors marked by the arch + // the GPU that is used for the entire model when split_mode is LLAMA_SPLIT_MODE_NONE int32_t main_gpu; diff --git a/src/CMakeLists.txt b/src/CMakeLists.txt index 320784c3a8cc..898fe483f6d3 100644 --- a/src/CMakeLists.txt +++ b/src/CMakeLists.txt @@ -19,16 +19,19 @@ add_library(llama llama-cparams.cpp llama-grammar.cpp llama-graph.cpp + llama-hadamard.cpp llama-hparams.cpp llama-impl.cpp llama-io.cpp llama-kv-cache.cpp + llama-kv-share.cpp llama-kv-cache-iswa.cpp llama-kv-cache-dsa.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 diff --git a/src/llama-arch.cpp b/src/llama-arch.cpp index e81ff647eee4..d333e92e7977 100644 --- a/src/llama-arch.cpp +++ b/src/llama-arch.cpp @@ -15,6 +15,9 @@ static const std::map LLM_ARCH_NAMES = { { LLM_ARCH_GPT2, "gpt2" }, { LLM_ARCH_GPTJ, "gptj" }, { LLM_ARCH_GPTNEOX, "gptneox" }, + { LLM_ARCH_OPT, "opt" }, + { LLM_ARCH_CODEGEN, "codegen" }, + { LLM_ARCH_GPTNEO, "gptneo" }, { LLM_ARCH_MPT, "mpt" }, { LLM_ARCH_BAICHUAN, "baichuan" }, { LLM_ARCH_STARCODER, "starcoder" }, @@ -40,6 +43,7 @@ static const std::map LLM_ARCH_NAMES = { { LLM_ARCH_QWEN3VLMOE, "qwen3vlmoe" }, { LLM_ARCH_QWEN35, "qwen35" }, { LLM_ARCH_QWEN35MOE, "qwen35moe" }, + { LLM_ARCH_QWEN4EXP, "qwen4exp" }, { LLM_ARCH_PHI2, "phi2" }, { LLM_ARCH_PHI3, "phi3" }, { LLM_ARCH_PHIMOE, "phimoe" }, @@ -141,6 +145,7 @@ static const std::map LLM_ARCH_NAMES = { { LLM_ARCH_LLAMA_EMBED, "llama-embed" }, { LLM_ARCH_MAINCODER, "maincoder" }, { LLM_ARCH_KIMI_LINEAR, "kimi-linear" }, + { LLM_ARCH_ZAYA, "zaya" }, { LLM_ARCH_TALKIE, "talkie" }, { LLM_ARCH_MELLUM, "mellum" }, { LLM_ARCH_NANBEIGE, "nanbeige" }, @@ -245,6 +250,8 @@ static const std::map LLM_KV_NAMES = { { LLM_KV_ATTENTION_RELATIVE_BUCKETS_COUNT, "%s.attention.relative_buckets_count" }, { LLM_KV_ATTENTION_SLIDING_WINDOW, "%s.attention.sliding_window" }, { LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, "%s.attention.sliding_window_pattern" }, + { LLM_KV_ZAYA_VLORA_RANK_ATTN, "%s.vision_lora.attention_rank" }, + { LLM_KV_ZAYA_VLORA_RANK_FFN, "%s.vision_lora.ffn_rank" }, { LLM_KV_ATTENTION_SCALE, "%s.attention.scale" }, { LLM_KV_ATTENTION_OUTPUT_SCALE, "%s.attention.output_scale" }, { LLM_KV_ATTENTION_VALUE_SCALE, "%s.attention.value_scale" }, @@ -270,6 +277,17 @@ static const std::map LLM_KV_NAMES = { { LLM_KV_HYPER_CONNECTION_COUNT, "%s.hyper_connection.count" }, { LLM_KV_HYPER_CONNECTION_SINKHORN_ITERATIONS, "%s.hyper_connection.sinkhorn_iterations" }, { LLM_KV_HYPER_CONNECTION_EPSILON, "%s.hyper_connection.epsilon" }, + { LLM_KV_HYPER_CONNECTION_LOW_RANK, "%s.hyper_connection.low_rank" }, + + { LLM_KV_PLE_LAYERS, "%s.ple.layers" }, + { LLM_KV_PLE_NGRAM_SIZE, "%s.ple.ngram_size" }, + { LLM_KV_PLE_HEADS_PER_NGRAM, "%s.ple.heads_per_ngram" }, + { LLM_KV_PLE_CONV_KERNEL, "%s.ple.conv_kernel" }, + { LLM_KV_PLE_LAYER_MULTIPLIERS, "%s.ple.layer_multipliers" }, + { LLM_KV_PLE_HEAD_OFFSETS, "%s.ple.head_offsets" }, + { LLM_KV_PLE_HEAD_VOCAB_SIZES, "%s.ple.head_vocab_sizes" }, + { LLM_KV_PLE_EOS_TOKEN_ID, "%s.ple.eos_token_id" }, + { LLM_KV_PLE_IMAGE_TOKEN_ID, "%s.ple.image_token_id" }, { LLM_KV_HASH_LAYER_COUNT, "%s.hash_layer_count" }, @@ -468,12 +486,29 @@ static const std::map LLM_TENSOR_NAMES = { { LLM_TENSOR_HC_HEAD_FN, "output_hc_fn" }, { LLM_TENSOR_HC_HEAD_BASE, "output_hc_base" }, { LLM_TENSOR_HC_HEAD_SCALE, "output_hc_scale" }, + { LLM_TENSOR_HC_HEAD_NORM, "output_hc_norm" }, + { LLM_TENSOR_HC_HEAD_DOWN, "output_hc_down" }, + { LLM_TENSOR_HC_HEAD_UP, "output_hc_up" }, { LLM_TENSOR_HC_ATTN_FN, "blk.%d.hc_attn_fn" }, { LLM_TENSOR_HC_ATTN_BASE, "blk.%d.hc_attn_base" }, { LLM_TENSOR_HC_ATTN_SCALE, "blk.%d.hc_attn_scale" }, { LLM_TENSOR_HC_FFN_FN, "blk.%d.hc_ffn_fn" }, { LLM_TENSOR_HC_FFN_BASE, "blk.%d.hc_ffn_base" }, { LLM_TENSOR_HC_FFN_SCALE, "blk.%d.hc_ffn_scale" }, + { LLM_TENSOR_HC_ATTN_NORM, "blk.%d.hc_attn_norm" }, + { LLM_TENSOR_HC_ATTN_DOWN, "blk.%d.hc_attn_down" }, + { LLM_TENSOR_HC_ATTN_UP, "blk.%d.hc_attn_up" }, + { LLM_TENSOR_HC_ATTN_INJECT, "blk.%d.hc_attn_inject" }, + { LLM_TENSOR_HC_FFN_NORM, "blk.%d.hc_ffn_norm" }, + { LLM_TENSOR_HC_FFN_DOWN, "blk.%d.hc_ffn_down" }, + { LLM_TENSOR_HC_FFN_UP, "blk.%d.hc_ffn_up" }, + { LLM_TENSOR_HC_FFN_INJECT, "blk.%d.hc_ffn_inject" }, + { LLM_TENSOR_PLE_KEY, "blk.%d.ple_key" }, + { LLM_TENSOR_PLE_VALUE, "blk.%d.ple_value" }, + { LLM_TENSOR_PLE_NORM_KEY, "blk.%d.ple_norm_key" }, + { LLM_TENSOR_PLE_NORM_QUERY, "blk.%d.ple_norm_query" }, + { LLM_TENSOR_PLE_NORM_CONV, "blk.%d.ple_norm_conv" }, + { LLM_TENSOR_PLE_CONV1D, "blk.%d.ple_conv1d" }, { LLM_TENSOR_ATTN_COMPRESSOR_WKV, "blk.%d.attn_compressor_kv" }, { LLM_TENSOR_ATTN_COMPRESSOR_WGATE, "blk.%d.attn_compressor_gate" }, { LLM_TENSOR_ATTN_COMPRESSOR_APE, "blk.%d.attn_compressor_ape" }, @@ -507,6 +542,38 @@ static const std::map LLM_TENSOR_NAMES = { { LLM_TENSOR_NEXTN_HNORM, "blk.%d.nextn.hnorm" }, { LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD, "blk.%d.nextn.shared_head_head" }, { LLM_TENSOR_NEXTN_SHARED_HEAD_NORM, "blk.%d.nextn.shared_head_norm" }, + { LLM_TENSOR_CCA_CONV_GRP, "blk.%d.cca_conv_grp" }, + { LLM_TENSOR_CCA_K_SCALE, "blk.%d.cca_k_scale" }, + { LLM_TENSOR_CCA_VAL_PROJ1, "blk.%d.cca_val_proj1" }, + { LLM_TENSOR_CCA_VAL_PROJ2, "blk.%d.cca_val_proj2" }, + { LLM_TENSOR_RES_SCALE_HS, "blk.%d.res_scale_hs" }, + { LLM_TENSOR_RES_SCALE_RES, "blk.%d.res_scale_res" }, + { LLM_TENSOR_RES_SCALE_HS_MLP, "blk.%d.res_scale_hs_mlp" }, + { LLM_TENSOR_RES_SCALE_RES_MLP, "blk.%d.res_scale_res_mlp" }, + { LLM_TENSOR_RES_SCALE_HS_FINAL, "res_scale_hs" }, + { LLM_TENSOR_RES_SCALE_RES_FINAL, "res_scale_res" }, + { LLM_TENSOR_INPUT_HIDDEN_STATES_SCALE, "input_hidden_states_scale" }, + { LLM_TENSOR_ZAYA_ROUTER_MLP2, "blk.%d.zaya_router_mlp2" }, + { LLM_TENSOR_ZAYA_ROUTER_MLP4, "blk.%d.zaya_router_mlp4" }, + { LLM_TENSOR_ZAYA_ROUTER_BIASES, "blk.%d.zaya_router_biases" }, + { LLM_TENSOR_ZAYA_ROUTER_EDA_SCALE, "blk.%d.zaya_router_eda" }, + { LLM_TENSOR_ZAYA_VLORA_Q_A , "blk.%d.zaya_vlora_q_a" }, + { LLM_TENSOR_ZAYA_VLORA_Q_B , "blk.%d.zaya_vlora_q_b" }, + { LLM_TENSOR_ZAYA_VLORA_K_A , "blk.%d.zaya_vlora_k_a" }, + { LLM_TENSOR_ZAYA_VLORA_K_B , "blk.%d.zaya_vlora_k_b" }, + { LLM_TENSOR_ZAYA_VLORA_V1_A , "blk.%d.zaya_vlora_v1_a" }, + { LLM_TENSOR_ZAYA_VLORA_V1_B , "blk.%d.zaya_vlora_v1_b" }, + { LLM_TENSOR_ZAYA_VLORA_V2_A , "blk.%d.zaya_vlora_v2_a" }, + { LLM_TENSOR_ZAYA_VLORA_V2_B , "blk.%d.zaya_vlora_v2_b" }, + { LLM_TENSOR_ZAYA_VLORA_O_A , "blk.%d.zaya_vlora_o_a" }, + { LLM_TENSOR_ZAYA_VLORA_O_B , "blk.%d.zaya_vlora_o_b" }, + { LLM_TENSOR_ZAYA_VLORA_UP_EXPS_A , "blk.%d.zaya_vlora_gate_up_exps_a" }, + { LLM_TENSOR_ZAYA_VLORA_UP_EXPS_B , "blk.%d.zaya_vlora_gate_up_exps_b" }, + { LLM_TENSOR_ZAYA_VLORA_DOWN_EXPS_A , "blk.%d.zaya_vlora_down_exps_a" }, + { LLM_TENSOR_ZAYA_VLORA_DOWN_EXPS_B , "blk.%d.zaya_vlora_down_exps_b" }, + { LLM_TENSOR_NEXTN_HC_HEAD_NORM, "blk.%d.nextn.hc_head_norm" }, + { LLM_TENSOR_NEXTN_HC_HEAD_DOWN, "blk.%d.nextn.hc_head_down" }, + { LLM_TENSOR_NEXTN_HC_HEAD_UP, "blk.%d.nextn.hc_head_up" }, { LLM_TENSOR_ATTN_SUB_NORM, "blk.%d.attn_sub_norm" }, { LLM_TENSOR_FFN_SUB_NORM, "blk.%d.ffn_sub_norm" }, { LLM_TENSOR_DEC_OUTPUT_NORM, "dec.output_norm" }, @@ -672,12 +739,29 @@ static const std::map LLM_TENSOR_INFOS = { {LLM_TENSOR_HC_HEAD_FN, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL_MAT}}, {LLM_TENSOR_HC_HEAD_BASE, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_ADD}}, {LLM_TENSOR_HC_HEAD_SCALE, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL}}, + {LLM_TENSOR_HC_HEAD_NORM, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL}}, + {LLM_TENSOR_HC_HEAD_DOWN, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_HC_HEAD_UP, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL_MAT}}, {LLM_TENSOR_HC_ATTN_FN, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, {LLM_TENSOR_HC_ATTN_BASE, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_ADD}}, {LLM_TENSOR_HC_ATTN_SCALE, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, {LLM_TENSOR_HC_FFN_FN, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, {LLM_TENSOR_HC_FFN_BASE, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_ADD}}, {LLM_TENSOR_HC_FFN_SCALE, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, + {LLM_TENSOR_HC_ATTN_NORM, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, + {LLM_TENSOR_HC_ATTN_DOWN, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_HC_ATTN_UP, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_HC_ATTN_INJECT, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_HC_FFN_NORM, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, + {LLM_TENSOR_HC_FFN_DOWN, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_HC_FFN_UP, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_HC_FFN_INJECT, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_PLE_KEY, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_PLE_VALUE, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_PLE_NORM_KEY, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, + {LLM_TENSOR_PLE_NORM_QUERY, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, + {LLM_TENSOR_PLE_NORM_CONV, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, + {LLM_TENSOR_PLE_CONV1D, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_SSM_CONV}}, {LLM_TENSOR_ATTN_COMPRESSOR_WKV, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, {LLM_TENSOR_ATTN_COMPRESSOR_WGATE, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, {LLM_TENSOR_ATTN_COMPRESSOR_APE, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_GET_ROWS}}, @@ -864,10 +948,43 @@ static const std::map LLM_TENSOR_INFOS = { {LLM_TENSOR_NEXTN_HNORM, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, {LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, {LLM_TENSOR_NEXTN_SHARED_HEAD_NORM, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, + {LLM_TENSOR_NEXTN_HC_HEAD_NORM, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, + {LLM_TENSOR_NEXTN_HC_HEAD_DOWN, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_NEXTN_HC_HEAD_UP, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, // Nemotron 3 Super // latent projections feed ggml_mul_mat, the buft probe must use MUL_MAT to keep them on GPU {LLM_TENSOR_FFN_LATENT_DOWN, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, {LLM_TENSOR_FFN_LATENT_UP, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + // ZAYA + {LLM_TENSOR_CCA_CONV_GRP, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_CCA_K_SCALE, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, + {LLM_TENSOR_CCA_VAL_PROJ1, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_CCA_VAL_PROJ2, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_RES_SCALE_HS, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, + {LLM_TENSOR_RES_SCALE_RES, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, + {LLM_TENSOR_RES_SCALE_HS_MLP, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, + {LLM_TENSOR_RES_SCALE_RES_MLP, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, + {LLM_TENSOR_RES_SCALE_HS_FINAL, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL}}, + {LLM_TENSOR_RES_SCALE_RES_FINAL, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL}}, + {LLM_TENSOR_INPUT_HIDDEN_STATES_SCALE, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL}}, + {LLM_TENSOR_ZAYA_ROUTER_MLP2, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_ZAYA_ROUTER_MLP4, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_ZAYA_ROUTER_BIASES, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_ADD}}, + {LLM_TENSOR_ZAYA_ROUTER_EDA_SCALE, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, + {LLM_TENSOR_ZAYA_VLORA_Q_A , {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_ZAYA_VLORA_Q_B , {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_ZAYA_VLORA_K_A , {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_ZAYA_VLORA_K_B , {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_ZAYA_VLORA_V1_A , {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_ZAYA_VLORA_V1_B , {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_ZAYA_VLORA_V2_A , {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_ZAYA_VLORA_V2_B , {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_ZAYA_VLORA_O_A , {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_ZAYA_VLORA_O_B , {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_ZAYA_VLORA_UP_EXPS_A , {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT_ID}}, + {LLM_TENSOR_ZAYA_VLORA_UP_EXPS_B , {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT_ID}}, + {LLM_TENSOR_ZAYA_VLORA_DOWN_EXPS_A , {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT_ID}}, + {LLM_TENSOR_ZAYA_VLORA_DOWN_EXPS_B , {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT_ID}}, {LLM_TENSOR_MASKED_EMBD_CENTROIDS, {LLM_TENSOR_LAYER_INPUT, GGML_OP_NONE}}, {LLM_TENSOR_MASKED_EMBD_ORDERING, {LLM_TENSOR_LAYER_INPUT, GGML_OP_NONE}}, // eagle3 @@ -967,7 +1084,9 @@ bool llm_arch_is_hybrid(const llm_arch & arch) { case LLM_ARCH_QWEN3NEXT: case LLM_ARCH_KIMI_LINEAR: case LLM_ARCH_QWEN35: + case LLM_ARCH_ZAYA: case LLM_ARCH_QWEN35MOE: + case LLM_ARCH_QWEN4EXP: return true; default: return false; @@ -990,6 +1109,7 @@ bool llm_arch_supports_rs_rollback(const llm_arch & arch) { switch (arch) { case LLM_ARCH_QWEN35: case LLM_ARCH_QWEN35MOE: + case LLM_ARCH_QWEN4EXP: return true; default: return false; @@ -1023,7 +1143,9 @@ bool llm_arch_supports_sm_tensor(const llm_arch & arch) { case LLM_ARCH_MINIMAX_M2: case LLM_ARCH_MINIMAX_M3: case LLM_ARCH_MISTRAL4: + case LLM_ARCH_ZAYA: case LLM_ARCH_KIMI_LINEAR: + case LLM_ARCH_QWEN4EXP: // TODO: fix test-llama-archs return false; default: return true; diff --git a/src/llama-arch.h b/src/llama-arch.h index cbc97085ea79..dbd5fb5df094 100644 --- a/src/llama-arch.h +++ b/src/llama-arch.h @@ -21,6 +21,9 @@ enum llm_arch { LLM_ARCH_GPT2, LLM_ARCH_GPTJ, LLM_ARCH_GPTNEOX, + LLM_ARCH_OPT, + LLM_ARCH_CODEGEN, + LLM_ARCH_GPTNEO, LLM_ARCH_MPT, LLM_ARCH_STARCODER, LLM_ARCH_REFACT, @@ -45,6 +48,7 @@ enum llm_arch { LLM_ARCH_QWEN3VLMOE, LLM_ARCH_QWEN35, LLM_ARCH_QWEN35MOE, + LLM_ARCH_QWEN4EXP, LLM_ARCH_PHI2, LLM_ARCH_PHI3, LLM_ARCH_PHIMOE, @@ -143,6 +147,7 @@ enum llm_arch { LLM_ARCH_LLAMA_EMBED, LLM_ARCH_MAINCODER, LLM_ARCH_KIMI_LINEAR, + LLM_ARCH_ZAYA, LLM_ARCH_TALKIE, LLM_ARCH_MELLUM, LLM_ARCH_EAGLE3, @@ -250,6 +255,8 @@ enum llm_kv { LLM_KV_ATTENTION_RELATIVE_BUCKETS_COUNT, LLM_KV_ATTENTION_SLIDING_WINDOW, LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, + LLM_KV_ZAYA_VLORA_RANK_ATTN, + LLM_KV_ZAYA_VLORA_RANK_FFN, LLM_KV_ATTENTION_SCALE, LLM_KV_ATTENTION_OUTPUT_SCALE, LLM_KV_ATTENTION_VALUE_SCALE, @@ -275,6 +282,17 @@ enum llm_kv { LLM_KV_HYPER_CONNECTION_COUNT, LLM_KV_HYPER_CONNECTION_SINKHORN_ITERATIONS, LLM_KV_HYPER_CONNECTION_EPSILON, + LLM_KV_HYPER_CONNECTION_LOW_RANK, + + LLM_KV_PLE_LAYERS, + LLM_KV_PLE_NGRAM_SIZE, + LLM_KV_PLE_HEADS_PER_NGRAM, + LLM_KV_PLE_CONV_KERNEL, + LLM_KV_PLE_LAYER_MULTIPLIERS, + LLM_KV_PLE_HEAD_OFFSETS, + LLM_KV_PLE_HEAD_VOCAB_SIZES, + LLM_KV_PLE_EOS_TOKEN_ID, + LLM_KV_PLE_IMAGE_TOKEN_ID, LLM_KV_HASH_LAYER_COUNT, @@ -533,12 +551,29 @@ enum llm_tensor { LLM_TENSOR_HC_HEAD_FN, LLM_TENSOR_HC_HEAD_BASE, LLM_TENSOR_HC_HEAD_SCALE, + LLM_TENSOR_HC_HEAD_NORM, // qwen4exp + LLM_TENSOR_HC_HEAD_DOWN, // qwen4exp + LLM_TENSOR_HC_HEAD_UP, // qwen4exp LLM_TENSOR_HC_ATTN_FN, LLM_TENSOR_HC_ATTN_BASE, LLM_TENSOR_HC_ATTN_SCALE, LLM_TENSOR_HC_FFN_FN, LLM_TENSOR_HC_FFN_BASE, LLM_TENSOR_HC_FFN_SCALE, + LLM_TENSOR_HC_ATTN_NORM, // qwen4exp + LLM_TENSOR_HC_ATTN_DOWN, // qwen4exp + LLM_TENSOR_HC_ATTN_UP, // qwen4exp + LLM_TENSOR_HC_ATTN_INJECT, // qwen4exp + LLM_TENSOR_HC_FFN_NORM, // qwen4exp + LLM_TENSOR_HC_FFN_DOWN, // qwen4exp + LLM_TENSOR_HC_FFN_UP, // qwen4exp + LLM_TENSOR_HC_FFN_INJECT, // qwen4exp + LLM_TENSOR_PLE_KEY, // qwen4exp + LLM_TENSOR_PLE_VALUE, // qwen4exp + LLM_TENSOR_PLE_NORM_KEY, // qwen4exp + LLM_TENSOR_PLE_NORM_QUERY, // qwen4exp + LLM_TENSOR_PLE_NORM_CONV, // qwen4exp + LLM_TENSOR_PLE_CONV1D, // qwen4exp LLM_TENSOR_ATTN_COMPRESSOR_WKV, LLM_TENSOR_ATTN_COMPRESSOR_WGATE, LLM_TENSOR_ATTN_COMPRESSOR_APE, @@ -620,6 +655,44 @@ enum llm_tensor { LLM_TENSOR_NEXTN_HNORM, LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD, LLM_TENSOR_NEXTN_SHARED_HEAD_NORM, + + // ZAYA CCA (Compressed Convolutional Attention) + LLM_TENSOR_CCA_CONV_GRP, + LLM_TENSOR_CCA_K_SCALE, + LLM_TENSOR_CCA_VAL_PROJ1, + LLM_TENSOR_CCA_VAL_PROJ2, + // ZAYA residual scaling + LLM_TENSOR_RES_SCALE_HS, + LLM_TENSOR_RES_SCALE_RES, + LLM_TENSOR_RES_SCALE_HS_MLP, + LLM_TENSOR_RES_SCALE_RES_MLP, + LLM_TENSOR_RES_SCALE_HS_FINAL, + LLM_TENSOR_RES_SCALE_RES_FINAL, + // ZAYA input embedding scaling + LLM_TENSOR_INPUT_HIDDEN_STATES_SCALE, + // ZAYA router (MoE gating) + LLM_TENSOR_ZAYA_ROUTER_MLP2, + LLM_TENSOR_ZAYA_ROUTER_MLP4, + LLM_TENSOR_ZAYA_ROUTER_BIASES, + LLM_TENSOR_ZAYA_ROUTER_EDA_SCALE, + // ZAYA1-VL: vision-only LoRA (A = down to the rank, B = back up), used on image tokens + LLM_TENSOR_ZAYA_VLORA_Q_A, + LLM_TENSOR_ZAYA_VLORA_Q_B, + LLM_TENSOR_ZAYA_VLORA_K_A, + LLM_TENSOR_ZAYA_VLORA_K_B, + LLM_TENSOR_ZAYA_VLORA_V1_A, + LLM_TENSOR_ZAYA_VLORA_V1_B, + LLM_TENSOR_ZAYA_VLORA_V2_A, + LLM_TENSOR_ZAYA_VLORA_V2_B, + LLM_TENSOR_ZAYA_VLORA_O_A, + LLM_TENSOR_ZAYA_VLORA_O_B, + LLM_TENSOR_ZAYA_VLORA_UP_EXPS_A, + LLM_TENSOR_ZAYA_VLORA_UP_EXPS_B, + LLM_TENSOR_ZAYA_VLORA_DOWN_EXPS_A, + LLM_TENSOR_ZAYA_VLORA_DOWN_EXPS_B, + LLM_TENSOR_NEXTN_HC_HEAD_NORM, + LLM_TENSOR_NEXTN_HC_HEAD_DOWN, + LLM_TENSOR_NEXTN_HC_HEAD_UP, LLM_TENSOR_MASKED_EMBD_CENTROIDS, LLM_TENSOR_MASKED_EMBD_ORDERING, LLM_TENSOR_FC, diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 5ef7becf6f29..96aa5925204f 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -10,6 +10,7 @@ #include "llama-mmap.h" #include "llama-model.h" #include "llama-ext.h" +#include "llama-kv-share.h" #include "llama.h" #include @@ -149,7 +150,7 @@ llama_context::llama_context( cparams.ctx_other = params.ctx_other; } - if (model.arch == LLM_ARCH_EAGLE3 || model.arch == LLM_ARCH_DFLASH) { + if (model.arch == LLM_ARCH_EAGLE3 || model.arch == LLM_ARCH_DFLASH || model.arch == LLM_ARCH_QWEN4EXP) { if (model.tok_embd == nullptr || model.output == nullptr) { if (params.ctx_other == nullptr) { throw std::runtime_error(model.arch_name() + " requires ctx_other to be set (this warning is normal during memory fitting)"); @@ -387,6 +388,12 @@ llama_context::llama_context( /*.mem_other =*/ llama_get_memory(cparams.ctx_other), }; + // [1bit] llama_kv_share_next (llama-kv-share.cpp) + struct kv_share_scope { + explicit kv_share_scope(bool enabled) { llama_kv_share_begin(enabled); } + ~kv_share_scope() { llama_kv_share_end(); } + } kv_share(llama_kv_share_take_next(hparams.no_alloc)); + memory.reset(model.create_memory(params_mem, cparams)); } @@ -1716,7 +1723,9 @@ int llama_context::decode(const llama_batch & batch_inp) { const auto & hparams = model.hparams; const int64_t n_vocab = vocab.n_tokens(); - const int64_t n_embd = hparams.n_embd_inp(); + // MTP batches carry the target hidden state, which is n_embd_out wide (4 * n_embd for qwen4exp) + const bool mtp_embd = cparams.ctx_type == LLAMA_CONTEXT_TYPE_MTP && batch_inp.embd; + const int64_t n_embd = mtp_embd ? hparams.n_embd_out() : hparams.n_embd_inp(); // when computing embeddings, all tokens are output const bool output_all = cparams.embeddings; @@ -2295,8 +2304,10 @@ void llama_context::output_reorder() { } if (embd_nextn.size > 0) { - for (uint64_t k = 0; k < n_embd; k++) { - std::swap(embd_nextn.data[i0*n_embd + k], embd_nextn.data[i1*n_embd + k]); + // nextn rows are n_embd_out wide (4 * n_embd for qwen4exp) + const uint64_t n_out = model.hparams.n_embd_out(); + for (uint64_t k = 0; k < n_out; k++) { + std::swap(embd_nextn.data[i0*n_out + k], embd_nextn.data[i1*n_out + k]); } } @@ -2350,6 +2361,7 @@ uint32_t llama_context::graph_max_nodes(uint32_t n_tokens) const { model.arch == LLM_ARCH_KIMI_LINEAR || model.arch == LLM_ARCH_QWEN35 || model.arch == LLM_ARCH_QWEN35MOE || + model.arch == LLM_ARCH_QWEN4EXP || model.arch == LLM_ARCH_DEEPSEEK4 || model.arch == LLM_ARCH_NANBEIGE || model.arch == LLM_ARCH_MINIMAX_M3) { @@ -2407,6 +2419,12 @@ ggml_cgraph * llama_context::graph_reserve( auto * gf = model.build_graph(gparams); + // once, on the pristine graph (after scheduling, cross-backend copies hide the producers) + if (gf && !hadamard_verified && !model.hadamard.empty()) { + model.hadamard.verify_graph(gf); + hadamard_verified = true; + } + this->n_outputs = save_n_outputs; // initialize scheduler with the specified graph @@ -2446,6 +2464,7 @@ llm_graph_params llama_context::graph_params( /*.n_outputs =*/ n_outputs, /*.cb =*/ graph_get_cb(), /*.res =*/ res, + /*.hadamard =*/ model.hadamard.empty() ? nullptr : &model.hadamard, }; } diff --git a/src/llama-context.h b/src/llama-context.h index bf91daa8b562..c68b0c062102 100644 --- a/src/llama-context.h +++ b/src/llama-context.h @@ -239,6 +239,9 @@ struct llama_context { // public: + // [1bit] the KV tensors moved to another buffer (llama_kv_share_from): drop graphs that captured the old one + void kv_memory_rebound(); + uint32_t graph_max_nodes(uint32_t n_tokens) const; // can reuse the llm_graph_result instance of the context (for example to update a memory module) @@ -367,6 +370,8 @@ struct llama_context { llm_graph_result_ptr gf_res_prev; llm_graph_result_ptr gf_res_reserve; + bool hadamard_verified = false; // prism.hadamard coverage checked on the first graph + // host buffer for the model output (logits and embeddings) ggml_backend_buffer_ptr buf_output; diff --git a/src/llama-ext.h b/src/llama-ext.h index 348bbae95770..8ed7f101a2a1 100644 --- a/src/llama-ext.h +++ b/src/llama-ext.h @@ -16,6 +16,17 @@ LLAMA_API struct ggml_cgraph * llama_graph_reserve( uint32_t n_seqs, uint32_t n_outputs); +// [1bit] Zero-copy KV sharing between two contexts of one model on two devices of the same GPU (llama-kv-share.cpp). +// The next llama_init_from_model on this thread allocates its KV cache as one exportable region (dma-buf) with a fixed layout. +LLAMA_API void llama_kv_share_next(bool enabled); + +// [1bit] Put the KV tensors of dst into the region of src (created after llama_kv_share_next), mapped with no copy. +// Needs the same model and context params. Frees the KV buffer dst had. +LLAMA_API bool llama_kv_share_from(struct llama_context * dst, struct llama_context * src); + +// [1bit] Copy KV cell metadata (positions, sequences) from src to dst after src ran. Tensor data is shared, not copied. +LLAMA_API bool llama_kv_cells_copy(struct llama_context * dst, const struct llama_context * src); + // Get the default ggml_type for a given ftype. LLAMA_API ggml_type llama_ftype_get_default_type(llama_ftype ftype); diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index 6d1c8f4e42a8..b10fc6af19b3 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -1360,6 +1360,7 @@ llm_graph_context::llm_graph_context(const llm_graph_params & params) : samplers (params.samplers), cb_func (params.cb), res (params.res), + hadamard (params.hadamard), ctx0 (res->get_ctx()), gf (res->get_gf()) { res->set_params(params); @@ -1383,7 +1384,7 @@ ggml_tensor * llm_graph_context::build_lora_mm( ggml_tensor * w, ggml_tensor * cur, ggml_tensor * w_s) const { - ggml_tensor * res = ggml_mul_mat(ctx0, w, cur); + ggml_tensor * res = ggml_mul_mat(ctx0, w, hadamard ? hadamard->forward(ctx0, w, cur, hadamard_memo) : cur); if (w_s) { res = ggml_mul(ctx0, res, w_s); @@ -1415,7 +1416,7 @@ ggml_tensor * llm_graph_context::build_lora_mm_id( ggml_tensor * cur, // ggml_tensor * b ggml_tensor * ids, ggml_tensor * w_s) const { - ggml_tensor * res = ggml_mul_mat_id(ctx0, w, cur, ids); + ggml_tensor * res = ggml_mul_mat_id(ctx0, w, hadamard ? hadamard->forward(ctx0, w, cur, hadamard_memo) : cur, ids); if (w_s) { const int64_t n_expert = w_s->ne[0]; @@ -1867,6 +1868,11 @@ ggml_tensor * llm_graph_context::build_moe_ffn( { probs = logits; // [n_expert, n_tokens] } break; + case LLAMA_EXPERT_GATING_FUNC_TYPE_NONE: + { + GGML_ASSERT(probs_in != nullptr); + probs = logits; // already-normalized probabilities from the caller (zaya) + } break; case LLAMA_EXPERT_GATING_FUNC_TYPE_SQRT_SOFTPLUS: { probs = ggml_sqrt(ctx0, ggml_softplus(ctx0, logits)); // [n_expert, n_tokens] @@ -2185,6 +2191,7 @@ ggml_tensor * llm_graph_context::build_inp_embd(ggml_tensor * tok_embd) const { auto & cur = inps[0]; cur = ggml_get_rows(ctx0, tok_embd, inp->tokens); + cur = hadamard ? hadamard->inverse(ctx0, tok_embd, cur) : cur; // apply lora for embedding tokens if needed for (const auto & lora : *loras) { diff --git a/src/llama-graph.h b/src/llama-graph.h index 7ed490ce6728..6d325c2cf156 100644 --- a/src/llama-graph.h +++ b/src/llama-graph.h @@ -1,6 +1,7 @@ #pragma once #include "llama-arch.h" +#include "llama-hadamard.h" #include "llama-batch.h" #include "llama-hparams.h" #include "llama-adapter.h" @@ -711,6 +712,8 @@ struct llm_graph_params { llm_graph_result * res; + const llama_hadamard * hadamard = nullptr; // prism.hadamard folds (llama-hadamard.h) + // return true if the "other" params would result in a graph with the same topology as with the current params // having the same topology allows us to reuse the graph in some cases bool allow_reuse(const llm_graph_params & other) const { @@ -934,6 +937,9 @@ struct llm_graph_context { llm_graph_result * res; + const llama_hadamard * hadamard; // prism.hadamard folds (llama-hadamard.h) + mutable llama_hadamard_memo hadamard_memo; + ggml_context * ctx0 = nullptr; ggml_cgraph * gf = nullptr; diff --git a/src/llama-hadamard.cpp b/src/llama-hadamard.cpp new file mode 100644 index 000000000000..b562033d78df --- /dev/null +++ b/src/llama-hadamard.cpp @@ -0,0 +1,533 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// Portions follow PrismML-Eng/llama.cpp (https://github.com/PrismML-Eng/llama.cpp), the llama.cpp +// fork of PrismML, who made the Ternary Bonsai models and the prism.hadamard.* format. Thanks to +// PrismML for publishing both. Those portions are used under the MIT License: +// +// MIT License +// +// Copyright (c) 2023-2026 The ggml authors +// +// Permission is hereby granted, free of charge, to any person obtaining a copy +// of this software and associated documentation files (the "Software"), to deal +// in the Software without restriction, including without limitation the rights +// to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +// copies of the Software, and to permit persons to whom the Software is +// furnished to do so, subject to the following conditions: +// +// The above copyright notice and this permission notice shall be included in all +// copies or substantial portions of the Software. +// +// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +// IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +// FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +// AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +// LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +// OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +// SOFTWARE. + +// Hadamard-folded weights (prism.hadamard.*): see llama-hadamard.h. + +#include "llama-hadamard.h" + +#include "llama-impl.h" +#include "llama-model.h" +#include "llama-model-loader.h" + +#include "ggml-backend.h" +#include "gguf.h" + +#include +#include +#include +#include +#include + +// keys straight from the GGUF: llama_model_loader instantiates its string-keyed get_key / get_arr +// for a few types only (bool and array ones fail to link with GCC) +static bool gguf_bool(const gguf_context * ctx, const char * key, bool & out) { + const int64_t id = gguf_find_key(ctx, key); + if (id < 0) { + return false; + } + if (gguf_get_kv_type(ctx, id) != GGUF_TYPE_BOOL) { + throw std::runtime_error(format("%s is not a bool", key)); + } + out = gguf_get_val_bool(ctx, id); + return true; +} + +static bool gguf_strings(const gguf_context * ctx, const char * key, std::vector & out, bool required) { + const int64_t id = gguf_find_key(ctx, key); + if (id < 0) { + if (required) { + throw std::runtime_error(format("key not found in model: %s", key)); + } + return false; + } + if (gguf_get_kv_type(ctx, id) != GGUF_TYPE_ARRAY || gguf_get_arr_type(ctx, id) != GGUF_TYPE_STRING) { + throw std::runtime_error(format("%s is not an array of strings", key)); + } + out.clear(); + for (size_t i = 0; i < gguf_get_arr_n(ctx, id); ++i) { + out.emplace_back(gguf_get_arr_str(ctx, id, i)); + } + return true; +} + +static void gguf_ints(const gguf_context * ctx, const char * key, std::vector & out) { + const int64_t id = gguf_find_key(ctx, key); + if (id < 0) { + throw std::runtime_error(format("key not found in model: %s", key)); + } + if (gguf_get_kv_type(ctx, id) != GGUF_TYPE_ARRAY || gguf_get_arr_type(ctx, id) != GGUF_TYPE_INT32) { + throw std::runtime_error(format("%s is not an array of int32", key)); + } + const auto * data = static_cast(gguf_get_arr_data(ctx, id)); + out.assign(data, data + gguf_get_arr_n(ctx, id)); +} + +static bool gguf_u32(const gguf_context * ctx, const char * key, uint32_t & out, bool required) { + const int64_t id = gguf_find_key(ctx, key); + if (id < 0) { + if (required) { + throw std::runtime_error(format("key not found in model: %s", key)); + } + return false; + } + switch (gguf_get_kv_type(ctx, id)) { + case GGUF_TYPE_UINT32: out = gguf_get_val_u32(ctx, id); return true; + case GGUF_TYPE_INT32: out = (uint32_t) gguf_get_val_i32(ctx, id); return true; + case GGUF_TYPE_UINT16: out = gguf_get_val_u16(ctx, id); return true; + case GGUF_TYPE_UINT8: out = gguf_get_val_u8(ctx, id); return true; + default: throw std::runtime_error(format("%s is not an unsigned integer", key)); + } +} + +static void gguf_str(const gguf_context * ctx, const char * key, std::string & out) { + const int64_t id = gguf_find_key(ctx, key); + if (id < 0 || gguf_get_kv_type(ctx, id) != GGUF_TYPE_STRING) { + throw std::runtime_error(format("%s: missing or not a string", key)); + } + out = gguf_get_val_str(ctx, id); +} + +void llama_hadamard::load_keys(llama_model_loader & ml, llm_arch arch) { + const gguf_context * meta = ml.metadata; + uint32_t version = 0; + gguf_bool(meta, "prism.hadamard.tied_output", tied_output); + // The engine's tools/hadamard_q4_0.py stamps onebit.hadamard_q4_0 = 32 instead: every Q4_0 weight of + // the file is rotated by the normalized 32-point Sylvester Walsh-Hadamard matrix per block along + // its input dimension, without signs, which is the prism.hadamard transform with block_size 32 and + // sign_mode identity. It names no weights: the Q4_0 tensors are the rotated ones. + bool onebit = false; + if (!gguf_u32(meta, "prism.hadamard.version", version, false)) { + if (tied_output) { + throw std::runtime_error("prism.hadamard.tied_output without prism.hadamard.version"); + } + const int64_t kid = gguf_find_key(meta, "onebit.hadamard_q4_0"); + if (kid < 0) { + return; + } + const gguf_type kt = gguf_get_kv_type(meta, kid); + const int64_t stamp = kt == GGUF_TYPE_INT32 ? gguf_get_val_i32(meta, kid) : + kt == GGUF_TYPE_UINT32 ? (int64_t) gguf_get_val_u32(meta, kid) : -1; + if (stamp != 32) { + throw std::runtime_error(format("unsupported onebit.hadamard_q4_0: %lld", (long long) stamp)); + } + onebit = true; + onebit_q4_0 = true; + } else if (version != 1 && version != 2) { + throw std::runtime_error(format("unsupported prism.hadamard.version: %u", version)); + } + if ((version == 2) != tied_output) { + throw std::runtime_error("prism.hadamard version 2 requires tied_output=true; version 1 forbids it"); + } + if (!onebit && tied_output && ml.get_weight("output.weight")) { + throw std::runtime_error("prism.hadamard.tied_output requires output.weight to be absent"); + } + + uint32_t block_size = 0; + std::string transform, axis, sign_mode; + std::vector weight_names; + if (onebit) { + block_size = 32; + transform = "normalized-sylvester-walsh-hadamard"; + axis = "input-last-dimension"; + sign_mode = "identity"; + // the names are collected below, from the Q4_0 tensors on a verified matmul path + } else { + gguf_u32(meta, "prism.hadamard.block_size", block_size, true); + gguf_str(meta, "prism.hadamard.transform", transform); + gguf_str(meta, "prism.hadamard.axis", axis); + gguf_str(meta, "prism.hadamard.sign_mode", sign_mode); + gguf_strings(meta, "prism.hadamard.weight_names", weight_names, true); + } + + if (block_size == 0 || (block_size & (block_size - 1)) != 0) { + throw std::runtime_error(format("invalid prism.hadamard.block_size: %u", block_size)); + } + if (transform != "normalized-sylvester-walsh-hadamard") { + throw std::runtime_error(format("unsupported prism.hadamard.transform: %s", transform.c_str())); + } + if (axis != "input-last-dimension") { + throw std::runtime_error(format("unsupported prism.hadamard.axis: %s", axis.c_str())); + } + if (sign_mode != "identity" && sign_mode != "explicit") { + throw std::runtime_error(format("unsupported prism.hadamard.sign_mode: %s", sign_mode.c_str())); + } + if (!onebit && weight_names.empty()) { + throw std::runtime_error("prism.hadamard.weight_names is empty"); + } + + if (sign_mode == "explicit") { + std::vector widths, values; + gguf_ints(meta, "prism.hadamard.sign_widths", widths); + gguf_ints(meta, "prism.hadamard.sign_values", values); + // explicit with no widths would read as identity later and silently change the model + if (widths.empty()) { + throw std::runtime_error("prism.hadamard.sign_mode is explicit but sign_widths is empty"); + } + size_t off = 0; + for (const int32_t width : widths) { + if (width <= 0 || (uint32_t) width % block_size != 0 || off + width > values.size()) { + throw std::runtime_error(format("invalid prism.hadamard sign width: %d", width)); + } + auto & vec = sign_data[width]; + vec.assign(values.begin() + off, values.begin() + off + width); + for (const int32_t v : vec) { + if (v != 1 && v != -1) { + throw std::runtime_error("prism.hadamard sign values must be +/-1"); + } + } + off += width; + } + if (off != values.size()) { + throw std::runtime_error("prism.hadamard.sign_values length mismatch"); + } + } + + gguf_bool(meta, "prism.hadamard.gdn_v_grouped", gdn_v_grouped); + + // the activation transform is applied in build_lora_mm / build_lora_mm_id only: refuse + // architectures and tensor kinds not verified to route every matmul through them + switch (arch) { + case LLM_ARCH_LLAMA: + case LLM_ARCH_QWEN3: + case LLM_ARCH_QWEN3MOE: + case LLM_ARCH_QWEN35: + case LLM_ARCH_QWEN35MOE: + case LLM_ARCH_QWEN3NEXT: + break; + default: + throw std::runtime_error(format( + "prism.hadamard: arch '%s' is not verified to apply the activation transform to all folded weights", + llm_arch_name(arch))); + } + + const auto foldable = [](const std::string & name) { + static const char * kinds[] = { + "attn_q", "attn_k", "attn_v", "attn_qkv", "attn_gate", "attn_output", + "ffn_gate", "ffn_up", "ffn_down", + "ffn_gate_exps", "ffn_up_exps", "ffn_down_exps", "ffn_gate_up_exps", + "ffn_gate_shexp", "ffn_up_shexp", "ffn_down_shexp", + "ssm_out", "ssm_alpha", "ssm_beta", "ssm_ba", + }; + if (name == "output.weight") { + return true; // the output head goes through build_lora_mm in every arch + } + if (name.compare(0, 4, "blk.") != 0) { + return false; + } + size_t pos = 4; + while (pos < name.size() && isdigit((unsigned char) name[pos])) { + pos++; + } + if (pos == 4 || pos >= name.size() || name[pos] != '.') { + return false; + } + pos++; + for (const char * kind : kinds) { + if (name.compare(pos, std::string::npos, std::string(kind) + ".weight") == 0) { + return true; + } + } + return false; + }; + if (onebit) { + // tools/hadamard_q4_0.py rotates the Q4_0 matmul weights; a Q4_0 lookup table (token_embd) + // is written in the plain basis and needs no transform + // Every other Q4_0 tensor in a stamped file is rotated (the tool checks that before it writes the + // stamp), so a Q4_0 weight this whitelist does not cover would run without its activation transform + // and answer wrongly. Refuse the file instead. + std::vector unfoldable; + for (const auto & [name, w] : ml.weights_map) { + if (w.tensor->type != GGML_TYPE_Q4_0 || name == "token_embd.weight") { + continue; + } + if (foldable(name)) { + weight_names.push_back(name); + } else { + unfoldable.push_back(name); + } + } + if (!unfoldable.empty()) { + std::sort(unfoldable.begin(), unfoldable.end()); + throw std::runtime_error(format( + "onebit.hadamard_q4_0: %zu rotated Q4_0 weights are not on a verified Hadamard-aware matmul path " + "(first: %s), so arch '%s' can't run this file here", + unfoldable.size(), unfoldable.front().c_str(), llm_arch_name(arch))); + } + if (weight_names.empty()) { + throw std::runtime_error("onebit.hadamard_q4_0: the file has no Q4_0 matmul weights"); + } + } + for (const auto & name : weight_names) { + if (!foldable(name)) { + throw std::runtime_error(format("prism.hadamard: weight '%s' is not on a verified Hadamard-aware matmul path", name.c_str())); + } + if (!weight_blocks.emplace(name, block_size).second) { + throw std::runtime_error(format("duplicate prism.hadamard weight: %s", name.c_str())); + } + } + + // tables read by row lookup store rotated rows: the lookup result gets the inverse + std::vector inverse_names; + gguf_strings(meta, "prism.hadamard.inverse_weight_names", inverse_names, false); + for (const auto & name : inverse_names) { + // only the token-embedding lookup applies it; any other table would stay rotated + if (name != "token_embd.weight") { + throw std::runtime_error(format("prism.hadamard: weight '%s' is not a verified inverse-after-lookup table", name.c_str())); + } + if (weight_blocks.count(name) || !inverse_blocks.emplace(name, block_size).second) { + throw std::runtime_error(format("duplicate prism.hadamard inverse weight: %s", name.c_str())); + } + } + + if (tied_output) { + const auto it = inverse_blocks.find("token_embd.weight"); + if (it == inverse_blocks.end()) { + throw std::runtime_error("prism.hadamard.tied_output requires a latent token embedding"); + } + weight_blocks.emplace("token_embd.weight", it->second); + } else if (inverse_blocks.count("token_embd.weight") && !ml.get_weight("output.weight")) { + throw std::runtime_error("a tied Hadamard output requires version 2 and tied_output=true"); + } +} + +// one [n] F32 tensor or [n, n] matrix in its own buffer of type buft, owned by this struct +static ggml_tensor * hadamard_const(std::vector & ctxs, std::vector & bufs, + ggml_backend_buffer_type_t buft, const char * name, int64_t ne0, int64_t ne1, + const std::vector & data) { + ggml_init_params params = { ggml_tensor_overhead(), nullptr, true }; + ggml_context_ptr ctx { ggml_init(params) }; + if (!ctx) { + throw std::runtime_error("failed to create a Hadamard context"); + } + ggml_tensor * t = ne1 > 1 ? ggml_new_tensor_2d(ctx.get(), GGML_TYPE_F32, ne0, ne1) + : ggml_new_tensor_1d(ctx.get(), GGML_TYPE_F32, ne0); + ggml_set_name(t, name); + ggml_backend_buffer_ptr buf { ggml_backend_alloc_ctx_tensors_from_buft(ctx.get(), buft) }; + if (!buf) { + throw std::runtime_error(format("unable to allocate a %s buffer for %s", ggml_backend_buft_name(buft), name)); + } + ggml_backend_buffer_set_usage(buf.get(), GGML_BACKEND_BUFFER_USAGE_WEIGHTS); + ggml_backend_tensor_set(t, data.data(), 0, data.size() * sizeof(float)); + ctxs.emplace_back(std::move(ctx)); + bufs.emplace_back(std::move(buf)); + return t; +} + +void llama_hadamard::setup(const llama_model & model) { + if (weight_blocks.empty() && inverse_blocks.empty()) { + return; + } + std::map, ggml_tensor *> rots, signs; + // the inverse (lookup side) must not inherit a host buffer from a CPU-mapped table, or every + // token's transform crosses devices: it takes the forward rotations' buffer type + ggml_backend_buffer_type_t preferred = nullptr; + + const std::pair *, decltype(rotations) *> groups[] = { + { &weight_blocks, &rotations }, { &inverse_blocks, &inverses }, + }; + for (const auto & [blocks, target] : groups) { + for (const auto & [name, block] : *blocks) { + const ggml_tensor * w = model.get_tensor(name.c_str()); + if (tied_output && name == "token_embd.weight") { + w = target == &rotations ? model.output : model.tok_embd; + if (!w || strcmp(w->name, "token_embd.weight") != 0) { + throw std::runtime_error("prism.hadamard.tied_output is not bound to the token embedding"); + } + } + if (!w && onebit_q4_0) { + continue; // a weight of the file the model does not load (the MTP layer without --mtp) + } + if (!w) { + throw std::runtime_error(format("prism.hadamard weight not found: %s", name.c_str())); + } + if (w->ne[0] % block != 0) { + throw std::runtime_error(format("prism.hadamard block size %u does not divide input dimension %lld for %s", + block, (long long) w->ne[0], name.c_str())); + } + if (!w->buffer) { + throw std::runtime_error(format("prism.hadamard weight has no buffer: %s", name.c_str())); + } + ggml_backend_buffer_type_t buft = ggml_backend_buffer_get_type(w->buffer); + // CPU extra buffer types (CPU_REPACK) only take tensors they can repack + if (ggml_backend_dev_t dev = ggml_backend_buft_get_device(buft)) { + if (ggml_backend_dev_type(dev) == GGML_BACKEND_DEVICE_TYPE_CPU) { + buft = ggml_backend_dev_buffer_type(dev); + } + } + if (target == &rotations) { + preferred = buft; + } else if (preferred) { + buft = preferred; + } + + ggml_tensor *& rot = rots[{ block, buft }]; + if (!rot) { + std::vector h((size_t) block * block); + const float scale = 1.0f / sqrtf((float) block); + for (uint32_t r = 0; r < block; ++r) { + for (uint32_t c = 0; c < block; ++c) { + h[(size_t) r * block + c] = __builtin_parity(r & c) ? -scale : scale; + } + } + char n[GGML_MAX_NAME]; + snprintf(n, sizeof(n), "prism.hadamard.%u", block); + rot = hadamard_const(ctxs, bufs, buft, n, block, block, h); + } + + ggml_tensor * sign = nullptr; + if (!sign_data.empty()) { + const uint32_t width = (uint32_t) w->ne[0]; + const auto sd = sign_data.find(width); + if (sd == sign_data.end()) { + throw std::runtime_error(format("prism.hadamard has no sign vector for width %u (%s)", width, name.c_str())); + } + ggml_tensor *& s = signs[{ width, buft }]; + if (!s) { + char n[GGML_MAX_NAME]; + snprintf(n, sizeof(n), "prism.hadamard.signs.%u", width); + s = hadamard_const(ctxs, bufs, buft, n, width, 1, std::vector(sd->second.begin(), sd->second.end())); + } + sign = s; + } + + llama_hadamard_transform t { rot, sign }; + if (gdn_v_grouped && name.find(".ssm_out.") != std::string::npos) { + const int64_t n_v = model.hparams.ssm_dt_rank; + const int64_t n_k = model.hparams.ssm_n_group; + if (n_k <= 0 || n_v <= 0 || n_v % n_k != 0 || w->ne[0] % n_v != 0) { + throw std::runtime_error(format("prism.hadamard: bad GDN head geometry for %s", name.c_str())); + } + t.perm_hd = w->ne[0] / n_v; + t.perm_nk = n_k; + t.perm_rep = n_v / n_k; + } + target->emplace(w, t); + } + } + LLAMA_LOG_INFO("%s: %zu Hadamard-folded weight(s) (%zu inverse-lookup), %zu rotation(s), %zu sign vector(s)\n", + __func__, rotations.size() + inverses.size(), inverses.size(), rots.size(), signs.size()); +} + +ggml_tensor * llama_hadamard::forward(ggml_context * ctx, const ggml_tensor * w, ggml_tensor * x, llama_hadamard_memo & memo) const { + const auto it = rotations.find(w); + if (it == rotations.end()) { + return x; + } + const auto & t = it->second; + const auto key = std::make_pair((const ggml_tensor *) x, (const ggml_tensor *) t.rot); + if (const auto m = memo.find(key); m != memo.end()) { + return m->second; + } + ggml_tensor * cur = x; + if (t.perm_rep > 1) { + // tiled [hd, nk, rep] -> grouped [hd, rep, nk] feature order + ggml_tensor * c = ggml_is_contiguous(cur) ? cur : ggml_cont(ctx, cur); + const int64_t n = c->ne[1] * c->ne[2] * c->ne[3]; + const int64_t ne1 = c->ne[1], ne2 = c->ne[2], ne3 = c->ne[3]; + c = ggml_reshape_4d(ctx, c, t.perm_hd, t.perm_nk, t.perm_rep, n); + c = ggml_cont(ctx, ggml_permute(ctx, c, 0, 2, 1, 3)); + cur = ggml_reshape_4d(ctx, c, t.perm_hd * t.perm_nk * t.perm_rep, ne1, ne2, ne3); + } + if (t.signs) { + cur = ggml_mul(ctx, cur, t.signs); + } + cur = llama_mul_mat_hadamard(ctx, cur, t.rot); + memo[key] = cur; + return cur; +} + +ggml_tensor * llama_hadamard::inverse(ggml_context * ctx, const ggml_tensor * table, ggml_tensor * rows) const { + const auto it = inverses.find(table); + if (it == inverses.end()) { + return rows; + } + // rows hold z = H (s * h); the normalized Sylvester matrix is symmetric and its own + // inverse, so h = s * (H z) + ggml_tensor * cur = llama_mul_mat_hadamard(ctx, rows, it->second.rot); + return it->second.signs ? ggml_mul(ctx, cur, it->second.signs) : cur; +} + +void llama_hadamard::verify_graph(ggml_cgraph * gf) const { + const auto unwrap = [](const ggml_tensor * t) { + while (t && (t->op == GGML_OP_RESHAPE || t->op == GGML_OP_VIEW)) { + t = t->src[0]; + } + return t; + }; + const auto is_rotation = [](const ggml_tensor * t) { + return t && t->op == GGML_OP_MUL_MAT && ((const int32_t *) t->op_params)[1] == GGML_HINT_SRC0_IS_HADAMARD; + }; + std::map lookups; // lookups of rotated tables -> inverse applied + + for (int i = 0; i < ggml_graph_n_nodes(gf); ++i) { + const ggml_tensor * node = ggml_graph_node(gf, i); + if (node->op == GGML_OP_GET_ROWS && inverses.count(node->src[0])) { + lookups.emplace(node, false); + continue; + } + if (node->op != GGML_OP_MUL_MAT && node->op != GGML_OP_MUL_MAT_ID) { + continue; + } + if (is_rotation(node)) { + if (const auto lk = lookups.find(unwrap(node->src[1])); lk != lookups.end()) { + lk->second = true; + } + continue; + } + const auto it = rotations.find(node->src[0]); + if (it == rotations.end()) { + if (inverses.count(node->src[0])) { + throw std::runtime_error(format("Hadamard-latent table '%s' is used as a head without a forward transform", node->src[0]->name)); + } + continue; + } + const ggml_tensor * src = unwrap(node->src[1]); + if (!(is_rotation(src) && src->src[0] == it->second.rot)) { + throw std::runtime_error(format("Hadamard-folded weight '%s' is consumed without its activation transform; " + "this graph's matmul path does not support prism.hadamard folding", node->src[0]->name)); + } + } + for (const auto & [node, ok] : lookups) { + if (!ok) { + throw std::runtime_error(format("Hadamard-latent table '%s' is read without the inverse transform", node->src[0]->name)); + } + } +} diff --git a/src/llama-hadamard.h b/src/llama-hadamard.h new file mode 100644 index 000000000000..4e938c6cb691 --- /dev/null +++ b/src/llama-hadamard.h @@ -0,0 +1,115 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// Portions follow PrismML-Eng/llama.cpp (https://github.com/PrismML-Eng/llama.cpp), the llama.cpp +// fork of PrismML, who made the Ternary Bonsai models and the prism.hadamard.* format. Thanks to +// PrismML for publishing both. Those portions are used under the MIT License: +// +// MIT License +// +// Copyright (c) 2023-2026 The ggml authors +// +// Permission is hereby granted, free of charge, to any person obtaining a copy +// of this software and associated documentation files (the "Software"), to deal +// in the Software without restriction, including without limitation the rights +// to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +// copies of the Software, and to permit persons to whom the Software is +// furnished to do so, subject to the following conditions: +// +// The above copyright notice and this permission notice shall be included in all +// copies or substantial portions of the Software. +// +// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +// IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +// FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +// AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +// LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +// OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +// SOFTWARE. + +// Hadamard-folded weights (the prism.hadamard.* GGUF keys of PrismML's Ternary Bonsai models). +// Such a file stores most matmul weights in a Walsh-Hadamard-rotated basis: W' = W (S H)^T with +// H the normalized Sylvester-Walsh-Hadamard matrix of prism.hadamard.block_size and S an optional +// +/-1 sign vector per input width. The model is unchanged if every activation x feeding one of +// those weights is rotated the same way first, x' = H (S x) blockwise, and the token-embedding +// rows (stored rotated too) get the inverse after the lookup. Here the rotation is an ordinary +// MUL_MAT against the H matrix, so any backend with an F32 matmul (HRX among them) runs these +// files; with the weights themselves converted to Q4_0 (the engine's tools/ternary_to_q4_0.py, +// exact for ternary weights) no new weight type is needed. +// +// Follows PrismML-Eng/llama.cpp's implementation (MIT, the ggml authors and PrismML): same GGUF +// keys, same checks, same activation transform and the same graph-coverage check. + +#pragma once + +#include "llama-arch.h" + +#include "ggml.h" +#include "ggml-cpp.h" + +#include +#include +#include +#include +#include +#include + +struct llama_model; +struct llama_model_loader; + +struct llama_hadamard_transform { + ggml_tensor * rot = nullptr; // [block, block] F32 + ggml_tensor * signs = nullptr; // [width] F32, nullptr for the identity sign mode + // > 1: the activation arrives with its features in tiled head order [hd, nk, rep] and the + // fold was computed in grouped order [hd, rep, nk] (prism.hadamard.gdn_v_grouped, ssm_out) + int64_t perm_hd = 0, perm_nk = 0, perm_rep = 0; +}; + +// per-graph cache: several weights reading the same activation share one transform +using llama_hadamard_memo = std::map, ggml_tensor *>; + +struct llama_hadamard { + // from the GGUF keys (load_keys) + std::unordered_map weight_blocks; // weight name -> block size + std::unordered_map inverse_blocks; // lookup tables stored rotated + std::map> sign_data; // input width -> signs + bool gdn_v_grouped = false; + bool tied_output = false; + bool onebit_q4_0 = false; // from onebit.hadamard_q4_0: the Q4_0 matmul weights, MTP ones only if loaded + + // bound to the model's tensors (setup) + std::unordered_map rotations; + std::unordered_map inverses; + + bool empty() const { return rotations.empty() && inverses.empty(); } + + // reads prism.hadamard.*; throws on anything this implementation does not verify + void load_keys(llama_model_loader & ml, llm_arch arch); + // makes the rotation and sign tensors next to the weights, after the weights are loaded + void setup(const llama_model & model); + + // x as the folded weight w expects it (x itself when w is not folded) + ggml_tensor * forward(ggml_context * ctx, const ggml_tensor * w, ggml_tensor * x, llama_hadamard_memo & memo) const; + // rows looked up from table: back to the plain basis (rows itself when table is not rotated) + ggml_tensor * inverse(ggml_context * ctx, const ggml_tensor * table, ggml_tensor * rows) const; + + // throws if a folded weight is consumed without its transform (a matmul path that bypasses + // build_lora_mm) or a rotated table is read without the inverse + void verify_graph(ggml_cgraph * gf) const; + + private: + std::vector ctxs; + std::vector bufs; +}; diff --git a/src/llama-hparams.cpp b/src/llama-hparams.cpp index 50af97f358c3..85b91cc3c70a 100644 --- a/src/llama-hparams.cpp +++ b/src/llama-hparams.cpp @@ -211,7 +211,11 @@ uint32_t llama_hparams::n_embd_r() const { // TODO: maybe support other convolution strides than 1 // NOTE: since the first column of the conv_state is shifted out each time, it's not actually needed // Corresponds to Mamba's conv_states size - return (ssm_d_conv > 0 ? ssm_d_conv - 1 : 0) * (ssm_d_inner + 2*ssm_n_group*ssm_d_state); + const uint32_t n_conv = (ssm_d_conv > 0 ? ssm_d_conv - 1 : 0) * (ssm_d_inner + 2*ssm_n_group*ssm_d_state); + + // PLE conv history needs its own row: Meta splits cache_r_l by head, so a history packed behind the first is unaddressable + // it lives in cache_ple_r_l instead, mirrored like the rest of the PLE module + return n_conv; } uint32_t llama_hparams::n_embd_s() const { @@ -239,6 +243,23 @@ bool llama_hparams::is_recr(uint32_t il) const { GGML_ABORT("%s: il (%u) out of bounds (n_layer_all: %u)\n", __func__, il, n_layer_all); } +uint32_t llama_hparams::ple_conv_state() const { + if (ple_n_heads == 0 || ple_conv_kernel == 0) { + return 0; + } + + // dilation equals the n-gram size, matching the reference module + return (ple_conv_kernel - 1) * ple_ngram_size * dsv4_hc_mult * n_embd; +} + +bool llama_hparams::is_ple(uint32_t il) const { + if (il < n_layer_all) { + return is_ple_impl[il]; + } + + GGML_ABORT("%s: il (%u) out of bounds (n_layer_all: %u)\n", __func__, il, n_layer_all); +} + uint32_t llama_hparams::n_pos_per_embd() const { return rope_type == LLAMA_ROPE_TYPE_MROPE || rope_type == LLAMA_ROPE_TYPE_IMROPE ? 4 : 1; } diff --git a/src/llama-hparams.h b/src/llama-hparams.h index fc770bf003e6..3c6c127a9896 100644 --- a/src/llama-hparams.h +++ b/src/llama-hparams.h @@ -3,11 +3,14 @@ #include "llama.h" #include +#include #include // bump if necessary #define LLAMA_MAX_LAYERS 512 #define LLAMA_MAX_EXPERTS 512 // Qwen3 Next +#define LLAMA_MAX_PLE_NGRAM 8 // qwen4exp +#define LLAMA_MAX_PLE_HEADS 64 // qwen4exp enum llama_expert_gating_func_type { LLAMA_EXPERT_GATING_FUNC_TYPE_NONE = 0, @@ -247,6 +250,30 @@ struct llama_hparams { float dsv4_hc_eps = 0.0f; std::array dsv4_compress_ratios; + // 0 = full rank (DeepSeek-V4) + uint32_t hc_low_rank = 0; + + uint32_t ple_ngram_size = 0; + uint32_t ple_heads_per_ngram = 0; + uint32_t ple_conv_kernel = 0; + uint32_t ple_n_heads = 0; // (ngram_size - 1) * heads_per_ngram + uint32_t ple_head_dim = 0; + uint32_t ple_eos_token_id = 0; + // the id the PLE hash stands in at image positions; 0 makes the loader fall back to EOS + uint32_t ple_image_token_id = 0; + // the file lists PLE layer indices, so this is never a per-layer gguf array and can hold one bit per layer + std::bitset is_ple_impl; + // the hash multipliers reach ~2e13 and have to stay 64-bit + std::array ple_layer_multipliers; + // head offsets and vocab sizes are token-space indices; the gather truncates them to int32 anyway + std::array ple_head_offsets; + std::array ple_head_vocab_sizes; + + bool is_ple(uint32_t il) const; + + // PLE conv history rows: (kernel - 1) * ngram_size; 0 without a PLE module + uint32_t ple_conv_state() const; + // qwen3vl deepstack // When parsed from GGUF, this implies the first N layers consume the first // N deepstack embeddings. Use deepstack_mapping_arr if you need a more diff --git a/src/llama-impl.cpp b/src/llama-impl.cpp index b3a94b946d28..2b4a7a3a2b97 100644 --- a/src/llama-impl.cpp +++ b/src/llama-impl.cpp @@ -126,8 +126,8 @@ static std::string gguf_data_to_str(enum gguf_type type, const void * data, int case GGUF_TYPE_INT32: return std::to_string(((const int32_t *)data)[i]); case GGUF_TYPE_UINT64: return std::to_string(((const uint64_t *)data)[i]); case GGUF_TYPE_INT64: return std::to_string(((const int64_t *)data)[i]); - case GGUF_TYPE_FLOAT32: return std::to_string(((const float *)data)[i]); - case GGUF_TYPE_FLOAT64: return std::to_string(((const double *)data)[i]); + case GGUF_TYPE_FLOAT32: return format("%f", (double) ((const float *)data)[i]); // std::to_string format before C++26 + case GGUF_TYPE_FLOAT64: return format("%f", ((const double *)data)[i]); case GGUF_TYPE_BOOL: return ((const int8_t *)data)[i] != 0 ? "true" : "false"; default: return format("unknown type %d", type); } diff --git a/src/llama-kv-cache.cpp b/src/llama-kv-cache.cpp index 44cb1668dacf..5857ec871aee 100644 --- a/src/llama-kv-cache.cpp +++ b/src/llama-kv-cache.cpp @@ -12,6 +12,7 @@ #include #include #include +#include static bool ggml_is_power_of_2(int n) { return (n & (n - 1)) == 0; @@ -77,7 +78,8 @@ llama_kv_cache::llama_kv_cache( llama_memory_t mem_other, const layer_filter_cb & filter, const layer_reuse_cb & reuse, - const layer_share_cb & share) : + const layer_share_cb & share, + const char * name_tag) : model(model), hparams(hparams), v_trans(v_trans), n_seq_max(n_seq_max), n_stream(unified ? 1 : n_seq_max), n_pad(n_pad), n_swa(n_swa), swa_type(swa_type), other(static_cast(mem_other)), @@ -231,8 +233,8 @@ llama_kv_cache::llama_kv_cache( ggml_tensor * k = has_k ? ggml_new_tensor_3d(ctx, type_k, n_embd_k_gqa, kv_size, n_stream) : nullptr; ggml_tensor * v = has_v ? ggml_new_tensor_3d(ctx, type_v, n_embd_v_gqa, kv_size, n_stream) : nullptr; - has_k && ggml_format_name(k, "cache_k_l%d", il); - has_v && ggml_format_name(v, "cache_v_l%d", il); + has_k && ggml_format_name(k, "cache_%sk_l%d", name_tag, il); + has_v && ggml_format_name(v, "cache_%sv_l%d", name_tag, il); std::vector k_stream; std::vector v_stream; @@ -290,7 +292,9 @@ llama_kv_cache::llama_kv_cache( // allocate tensors and initialize the buffers to avoid NaNs in the padding for (auto & [buft, ctx] : ctx_map) { ggml_backend_buffer_t buf; - if (hparams.no_alloc) { + if (llama_kv_share_enabled() && !hparams.no_alloc) { + buf = share_alloc(ctx.get(), buft); // [1bit] exportable region (llama-kv-share.cpp) + } else if (hparams.no_alloc) { buf = ggml_backend_buft_alloc_buffer(buft, /*size =*/ 0); // dummy buffer for (ggml_tensor * t = ggml_get_first_tensor(ctx.get()); t != nullptr; t = ggml_get_next_tensor(ctx.get(), t)) { t->buffer = buf; // set dummy buffer for KV cache so that the backend scheduler won't try to allocate it @@ -1232,11 +1236,24 @@ void llama_kv_cache::apply_ubatch(const slot_info & sinfo, const llama_ubatch & cells.pos_set(idx, ubatch.pos[i]); - if (ubatch.is_pos_2d()) { - llama_kv_cell_ext ext { - /*.x =*/ ubatch.pos[i + ubatch.n_tokens*2], - /*.y =*/ ubatch.pos[i + ubatch.n_tokens], - }; + if (ubatch.is_pos_2d() || ubatch.token || hparams.ple_n_heads > 0) { + llama_kv_cell_ext ext; + + if (ubatch.is_pos_2d()) { + ext.x = ubatch.pos[i + ubatch.n_tokens*2]; + ext.y = ubatch.pos[i + ubatch.n_tokens]; + } + + if (ubatch.token) { + ext.tok = ubatch.token[i]; + } else if (hparams.ple_n_heads > 0) { + // embd batch (multimodal input) has no token ids, need to pad it with the correct ID for PLE layers + // TODO @ngxson : check if we can do the same as gemma 3n / gemma 4 + ext.tok = hparams.ple_image_token_id != 0 + ? (llama_token) hparams.ple_image_token_id + : (llama_token) hparams.ple_eos_token_id; + } + cells.ext_set(idx, ext); } @@ -1337,6 +1354,12 @@ ggml_tensor * llama_kv_cache::get_k_storage(int32_t il) const { return layers[ikv].k; } +const llama_kv_cells & llama_kv_cache::get_cells(llama_seq_id seq_id) const { + GGML_ASSERT(seq_id >= 0 && (size_t) seq_id < seq_to_stream.size()); + + return v_cells[seq_to_stream[seq_id]]; +} + uint32_t llama_kv_cache::get_n_kv(const slot_info & sinfo) const { uint32_t result = 0; @@ -1949,6 +1972,67 @@ void llama_kv_cache::set_input_v_rot(ggml_tensor * dst) const { memcpy(dst->data, attn_rot_hadamard.at(n_rot).data(), ggml_nbytes(dst)); } +bool llama_kv_cache::has_cell_ext() const { + // M-RoPE needs the 2D position, the PLE n-gram hash needs the token id + return hparams.n_pos_per_embd() > 1 || hparams.ple_n_heads > 0; +} + +void llama_kv_cache::get_prev_tokens(const llama_ubatch & ubatch, uint32_t n, std::vector & res) const { + const uint32_t n_tokens = ubatch.n_tokens; + + res.clear(); + res.resize(n_tokens*n, LLAMA_TOKEN_NULL); + + if (n == 0) { + return; + } + + // note: apply_ubatch() has already stored the current ubatch, so the cells cover the tokens + // of this very ubatch as well, which is what we want + // the nearest cell at or before a position also resolves M-RoPE gaps, where multiple tokens + // share the same temporal pos + + // an embd (multimodal) ubatch can repeat one position for a whole image, so positions + // do not encode the token order; resolve its predecessors by ubatch order instead + std::vector ord; // index among the ubatch tokens of the same seq + std::unordered_map> seq_idx; + + if (!ubatch.token) { + ord.resize(n_tokens); + for (uint32_t i = 0; i < n_tokens; ++i) { + auto & v = seq_idx[ubatch.seq_id[i][0]]; + ord[i] = v.size(); + v.push_back(i); + } + } + + for (uint32_t i = 0; i < n_tokens; ++i) { + // TODO: a token that belongs to more than one sequence has an ambiguous history. + // the n-gram architectures have to reject such batches + const llama_seq_id seq_id = ubatch.seq_id[i][0]; + + for (uint32_t j = 0; j < n; ++j) { + const llama_pos d = (llama_pos) (n - j); + + llama_pos p; + if (!ubatch.token) { + const auto & v = seq_idx[seq_id]; + const int64_t k = (int64_t) ord[i] - d; + // k >= 0: an earlier token of this very ubatch; k < 0: before the chunk + p = k >= 0 ? ubatch.pos[v[k]] : ubatch.pos[v[0]] + (llama_pos) k; + } else { + p = ubatch.pos[i] - d; + } + + if (p < 0) { + continue; + } + + res[i*n + j] = v_cells[seq_to_stream[seq_id]].seq_pos_tok_le(seq_id, p); + } + } +} + size_t llama_kv_cache::total_size() const { size_t size = 0; @@ -2189,6 +2273,15 @@ void llama_kv_cache::state_write(llama_io_write_i & io, llama_seq_id seq_id, lla } void llama_kv_cache::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) { + state_read_sinfo(io, seq_id, flags, nullptr, nullptr); +} + +void llama_kv_cache::state_read_sinfo( + llama_io_read_i & io, + llama_seq_id seq_id, + llama_state_seq_flags flags, + slot_info_vec_t * sinfos_out, +const slot_info_vec_t * sinfos_in) { // TODO: refactor [TAG_KV_CACHE_SHARE_CELLS] if (other) { return; @@ -2198,17 +2291,35 @@ void llama_kv_cache::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama GGML_ASSERT(seq_id == -1 || (seq_id >= 0 && (size_t) seq_id < seq_to_stream.size())); + if (sinfos_out) { + sinfos_out->assign(n_stream, slot_info{}); + } + + if (sinfos_in && sinfos_in->size() != n_stream) { + throw std::runtime_error("failed to restore kv cache: mirrored slot layout has the wrong stream count"); + } + uint32_t n_stream_cur; io.read(&n_stream_cur, sizeof(n_stream_cur)); if (n_stream_cur != n_stream) { throw std::runtime_error("n_stream mismatch"); } + // a whole-context restore replaces every stream, so the cache is emptied once here + // clear() resets all streams at once, so doing it per stream below would keep only the last one + if (seq_id == -1) { + clear(true); + } + for (uint32_t s = 0; s < n_stream; ++s) { uint32_t cell_count; io.read(&cell_count, sizeof(cell_count)); if (cell_count == 0) { + // a mirrored cache must be empty here as well, or the two no longer agree cell for cell + if (sinfos_in && !(*sinfos_in)[s].empty()) { + throw std::runtime_error("failed to restore kv cache: mirrored cache holds cells this one does not"); + } continue; } @@ -2217,7 +2328,7 @@ void llama_kv_cache::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama slot_info sinfo; bool res = true; - res = res && state_read_meta(io, strm, cell_count, sinfo, seq_id); + res = res && state_read_meta(io, strm, cell_count, sinfo, seq_id, sinfos_in ? &(*sinfos_in)[s] : nullptr); try { res = res && state_read_data(io, strm, cell_count, sinfo); @@ -2233,6 +2344,10 @@ void llama_kv_cache::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama } throw std::runtime_error("failed to restore kv cache"); } + + if (sinfos_out) { + (*sinfos_out)[s] = sinfo; + } } } @@ -2257,7 +2372,7 @@ void llama_kv_cache::state_write_meta(llama_io_write_i & io, const cell_ranges_t io.write(&pos, sizeof(pos)); io.write(&n_seq_id, sizeof(n_seq_id)); - if (hparams.n_pos_per_embd() > 1) { + if (has_cell_ext()) { const llama_kv_cell_ext ext = cells.ext_get(i); io.write(&ext, sizeof(ext)); } @@ -2398,7 +2513,7 @@ void llama_kv_cache::state_write_data(llama_io_write_i & io, const cell_ranges_t } } -bool llama_kv_cache::state_read_meta(llama_io_read_i & io, uint32_t strm, uint32_t cell_count, slot_info & sinfo, llama_seq_id dest_seq_id) { +bool llama_kv_cache::state_read_meta(llama_io_read_i & io, uint32_t strm, uint32_t cell_count, slot_info & sinfo, llama_seq_id dest_seq_id, const slot_info * sinfo_in) { auto & cells = v_cells[strm]; auto & head = v_heads[strm]; @@ -2412,6 +2527,12 @@ bool llama_kv_cache::state_read_meta(llama_io_read_i & io, uint32_t strm, uint32 ubatch.seq_id_unq[0] = dest_seq_id; + // the ext as it was saved, to put back after apply_ubatch() + std::vector exts; + if (has_cell_ext()) { + exts.resize(cell_count); + } + for (uint32_t i = 0; i < cell_count; ++i) { llama_pos pos; uint32_t n_seq_id; @@ -2424,12 +2545,19 @@ bool llama_kv_cache::state_read_meta(llama_io_read_i & io, uint32_t strm, uint32 return false; } - if (hparams.n_pos_per_embd() > 1) { + if (has_cell_ext()) { llama_kv_cell_ext ext; io.read(&ext, sizeof(ext)); - ubatch.pos[i + ubatch.n_tokens] = ext.y; - ubatch.pos[i + ubatch.n_tokens*2] = ext.x; + if (hparams.n_pos_per_embd() > 1) { + ubatch.pos[i + ubatch.n_tokens] = ext.y; + ubatch.pos[i + ubatch.n_tokens*2] = ext.x; + } + + // apply_ubatch() below restores ext.tok from the ubatch tokens + ubatch.token[i] = ext.tok; + + exts[i] = ext; } // read the sequence id, but directly discard it - we will use dest_seq_id instead @@ -2443,16 +2571,52 @@ bool llama_kv_cache::state_read_meta(llama_io_read_i & io, uint32_t strm, uint32 ubatch.seq_id[i] = &dest_seq_id; } - sinfo = find_slot(ubatch, false); - if (sinfo.empty()) { - LLAMA_LOG_ERROR("%s: failed to find %d available cells in kv cache\n", __func__, cell_count); - return false; + if (sinfo_in) { + // this cache mirrors another one, so it takes that cache's layout instead of searching for its own cells + if (sinfo_in->empty() || sinfo_in->n_stream() != 1 || sinfo_in->idxs[0].size() != cell_count) { + LLAMA_LOG_ERROR("%s: mirrored slot layout holds %d cells, this cache restores %d\n", __func__, + sinfo_in->empty() ? 0 : (int) sinfo_in->idxs[0].size(), cell_count); + return false; + } + + sinfo = *sinfo_in; + + // the layout is cell indices, so it means the same in both caches only while their streams line up + sinfo.s0 = strm; + sinfo.s1 = strm; + sinfo.strm[0] = strm; + + // seq_rm above freed exactly the cells this sequence held + // anything else in the way is a cache that had already drifted, which this restore must not hide + for (uint32_t i = 0; i < cell_count; ++i) { + const uint32_t idx = sinfo.idxs[0][i]; + + if (idx >= cells.size() || !cells.is_empty(idx)) { + LLAMA_LOG_ERROR("%s: cell %u of the mirrored slot layout is not free\n", __func__, idx); + return false; + } + } + } else { + sinfo = find_slot(ubatch, false); + if (sinfo.empty()) { + LLAMA_LOG_ERROR("%s: failed to find %d available cells in kv cache\n", __func__, cell_count); + return false; + } } - // TODO: we cannot yet restore llama_kv_cell_ext as the apply_ubatch() does not support it yet + // note: apply_ubatch() rebuilds llama_kv_cell_ext from the ubatch + // only ext.tok and the M-RoPE 2D position round-trip through it // see: https://github.com/ggml-org/llama.cpp/pull/16825#issuecomment-3460868350 apply_ubatch(sinfo, ubatch); + // apply_ubatch() takes the 2D position from the ubatch, and that ubatch is built with this + // cache's own n_pos_per_embd. a cache that does not use M-RoPE itself but mirrors one that + // does (the qwen4exp QSA indexer) would drop x and y. put the saved ext back instead, which + // is what the whole-context path below already does. + for (uint32_t i = 0; i < (uint32_t) exts.size(); ++i) { + cells.ext_set(sinfo.idxs[0][i], exts[i]); + } + LLAMA_LOG_DEBUG("%s: cell_count = %d, dest_seq_id = %d\n", __func__, cell_count, dest_seq_id); // DEBUG CHECK: verify that all cells were allocated and have correct seq_id and pos values @@ -2471,7 +2635,12 @@ bool llama_kv_cache::state_read_meta(llama_io_read_i & io, uint32_t strm, uint32 return false; } - clear(true); + // the cells go in from 0, so a mirrored cache lands on the same ones as long as it restores the same count. the layout itself carries no more information here + if (sinfo_in && (sinfo_in->empty() || sinfo_in->n_stream() != 1 || sinfo_in->idxs[0].size() != cell_count)) { + LLAMA_LOG_ERROR("%s: mirrored slot layout holds %d cells, this cache restores %d\n", __func__, + sinfo_in->empty() ? 0 : (int) sinfo_in->idxs[0].size(), cell_count); + return false; + } for (uint32_t i = 0; i < cell_count; ++i) { llama_pos pos; @@ -2482,7 +2651,7 @@ bool llama_kv_cache::state_read_meta(llama_io_read_i & io, uint32_t strm, uint32 cells.pos_set(i, pos); - if (hparams.n_pos_per_embd() > 1) { + if (has_cell_ext()) { llama_kv_cell_ext ext; io.read(&ext, sizeof(ext)); cells.ext_set(i, ext); @@ -2903,3 +3072,7 @@ void llama_kv_cache_context::set_input_k_rot(ggml_tensor * dst) const { void llama_kv_cache_context::set_input_v_rot(ggml_tensor * dst) const { kv->set_input_v_rot(dst); } + +void llama_kv_cache_context::get_prev_tokens(const llama_ubatch & ubatch, uint32_t n, std::vector & res) const { + kv->get_prev_tokens(ubatch, n, res); +} diff --git a/src/llama-kv-cache.h b/src/llama-kv-cache.h index d5a92f4405b5..e3ab0c8f9afc 100644 --- a/src/llama-kv-cache.h +++ b/src/llama-kv-cache.h @@ -4,6 +4,7 @@ #include "llama-graph.h" #include "llama-kv-cells.h" #include "llama-memory.h" +#include "llama-kv-share.h" #include #include @@ -112,10 +113,16 @@ class llama_kv_cache : public llama_memory_i { llama_memory_t mem_other, const layer_filter_cb & filter, const layer_reuse_cb & reuse, - const layer_share_cb & share); + const layer_share_cb & share, + // a model can hold more than one cache, so the tensor names have to stay unique + const char * name_tag = ""); ~llama_kv_cache() = default; + // [1bit] zero-copy KV sharing (llama-kv-share.cpp) + bool share_from(const llama_kv_cache & src); + bool cells_copy_from(const llama_kv_cache & src); + // // llama_memory_i // @@ -164,6 +171,19 @@ class llama_kv_cache : public llama_memory_i { std::vector get_layer_ids() const; ggml_tensor * get_k_storage(int32_t il) const; + const llama_kv_cells & get_cells(llama_seq_id seq_id) const; + + // state_read, plus the cells the restored tokens were placed in + // a cache that mirrors another one (the qwen4exp indexer) must not search for its own cells: two searches agree only by luck + // sinfos_out: if set, filled with the layout used; a stream with no cells leaves an empty entry + // sinfos_in : if set, the layout to use instead of searching. one entry per stream, cell count must match the blob + void state_read_sinfo( + llama_io_read_i & io, + llama_seq_id seq_id, + llama_state_seq_flags flags, + slot_info_vec_t * sinfos_out, + const slot_info_vec_t * sinfos_in); + // // graph_build API // @@ -219,6 +239,17 @@ class llama_kv_cache : public llama_memory_i { void set_input_k_rot(ggml_tensor * dst) const; void set_input_v_rot(ggml_tensor * dst) const; + // true if llama_kv_cell_ext holds information that has to survive a state save/restore + bool has_cell_ext() const; + + // for every token of the ubatch, the ids of the n tokens that precede it in its sequence + // example for M-RoPE image case: tokens A B X X X C, where X is a 3-token image at pos 2 spanning positions 2..4: + // tok: A B X X X C + // pos: 0 1 2 2 2 5 + // prev, n=2: A -> [NULL, NULL], B -> [NULL, A], 3rd X -> [X, X], C -> [X, X] + // note: used by n-gram input embeddings + void get_prev_tokens(const llama_ubatch & ubatch, uint32_t n, std::vector & res) const; + private: const llama_model & model; const llama_hparams & hparams; @@ -269,6 +300,10 @@ class llama_kv_cache : public llama_memory_i { // this is the SWA type of the cache - not to be confused with the model SWA type const llama_swa_type swa_type = LLAMA_SWA_TYPE_NONE; + // [1bit] exported regions of this cache's buffers, owner side (llama-kv-share.cpp) + std::vector> shared_; // per buffer: its chunks + ggml_backend_buffer_t share_alloc(ggml_context * ctx, ggml_backend_buffer_type_t buft); + // ggml contexts for the KV cache along with the allocated backend buffers: std::vector> ctxs_bufs; @@ -324,7 +359,8 @@ class llama_kv_cache : public llama_memory_i { void state_write_meta(llama_io_write_i & io, const cell_ranges_t & cr, llama_seq_id seq_id = -1) const; void state_write_data(llama_io_write_i & io, const cell_ranges_t & cr) const; - bool state_read_meta(llama_io_read_i & io, uint32_t strm, uint32_t cell_count, slot_info & sinfo, llama_seq_id dest_seq_id = -1); + // sinfo_in, when set, replaces the find_slot call: the cells are given by the caller + bool state_read_meta(llama_io_read_i & io, uint32_t strm, uint32_t cell_count, slot_info & sinfo, llama_seq_id dest_seq_id = -1, const slot_info * sinfo_in = nullptr); bool state_read_data(llama_io_read_i & io, uint32_t strm, uint32_t cell_count, const slot_info & sinfo); }; @@ -409,6 +445,9 @@ class llama_kv_cache_context : public llama_memory_context_i { void set_input_k_rot(ggml_tensor * dst) const; void set_input_v_rot(ggml_tensor * dst) const; + // see llama_kv_cache::get_prev_tokens() + void get_prev_tokens(const llama_ubatch & ubatch, uint32_t n, std::vector & res) const; + private: llama_memory_status status; diff --git a/src/llama-kv-cells.h b/src/llama-kv-cells.h index fddd31a0b219..5d567a6ed0b8 100644 --- a/src/llama-kv-cells.h +++ b/src/llama-kv-cells.h @@ -6,7 +6,7 @@ #include #include #include -#include +#include #include #include @@ -15,6 +15,10 @@ struct llama_kv_cell_ext { llama_pos x = 0; llama_pos y = 0; + // when tok = LLAMA_TOKEN_NULL when the cell is produced by embedding input (i.e. multimodal) + // use case: n-gram embeddings hash + llama_token tok = LLAMA_TOKEN_NULL; + // return true if the current 2D spatial position is greater than other bool is_2d_gt(llama_pos ox, llama_pos oy) const { return (y > oy) || (y == oy && x > ox); @@ -23,7 +27,7 @@ struct llama_kv_cell_ext { void reset() { static_assert(std::is_trivially_copyable_v); - memset(this, 0, sizeof(*this)); + *this = llama_kv_cell_ext{}; } }; @@ -31,6 +35,8 @@ struct llama_kv_cell_ext { // TODO: add unit tests class llama_kv_cells { public: + using seq_set_t = std::bitset; + void reset() { for (uint32_t i = 0; i < pos.size(); ++i) { pos[i] = -1; @@ -242,7 +248,7 @@ class llama_kv_cells { assert(seq_id >= 0); seq[i].reset(seq_id); - seq_pos_dec(seq_id, pos[i]); + seq_pos_dec(seq_id, i); if (seq[i].none()) { pos[i] = -1; @@ -266,7 +272,7 @@ class llama_kv_cells { seq[i].reset(); seq[i].set(seq_id); - seq_pos_inc(seq_id, pos[i]); + seq_pos_inc(seq_id, i); return false; } @@ -297,6 +303,13 @@ class llama_kv_cells { return seq[i].count(); } + // the full set of sequences this cell is visible to + const seq_set_t & seq_get_all(uint32_t i) const { + assert(i < pos.size()); + + return seq[i]; + } + // check if the cell contains seq_id bool seq_has(uint32_t i, llama_seq_id seq_id) const { assert(i < pos.size()); @@ -305,6 +318,24 @@ class llama_kv_cells { return seq[i].test(seq_id); } + // the token of the cell of sequence seq_id at the largest position <= p + // when several cells share that position, the one with the highest index wins + // return LLAMA_TOKEN_NULL if the sequence has no cell at or before p + // note: used by n-gram input embeddings to recover the tokens preceding a ubatch + llama_token seq_pos_tok_le(llama_seq_id seq_id, llama_pos p) const { + assert(seq_id >= 0); + assert(seq_id < LLAMA_MAX_SEQ); + + const auto & sp = seq_pos[seq_id]; + + auto it = sp.upper_bound({ p, std::numeric_limits::max() }); + if (it == sp.begin()) { + return LLAMA_TOKEN_NULL; + } + + return ext[(--it)->second].tok; + } + // note: call only if the cell is not empty and the seq_id is not in the cell void seq_add(uint32_t i, llama_seq_id seq_id) { assert(i < pos.size()); @@ -312,7 +343,7 @@ class llama_kv_cells { assert(!seq[i].test(seq_id)); seq[i].set(seq_id); - seq_pos_inc(seq_id, pos[i]); + seq_pos_inc(seq_id, i); } // return the sequence id of this cell @@ -339,8 +370,6 @@ class llama_kv_cells { return -1; } - assert(seq_pos[seq_id].begin()->second > 0); - return seq_pos[seq_id].begin()->first; } @@ -354,8 +383,6 @@ class llama_kv_cells { return -1; } - assert(seq_pos[seq_id].rbegin()->second > 0); - return seq_pos[seq_id].rbegin()->first; } @@ -483,41 +510,36 @@ class llama_kv_cells { // std::vector shift; - using seq_set_t = std::bitset; - // the bitset seq[i] tells us which sequences are currently occupying the i-th cell std::vector seq; - // the set seq_pos[s][p] tells us how many times the position p is currently present for sequence s - // if the position p is not present, seq_pos[s][p] is not set + // the set seq_pos[s] holds one (pos, cell) pair per cell that carries sequence s, ordered by position // this way seq_pos[s].begin() and seq_pos[s].rbegin() give us the min/max positions currently in the cache + // and upper_bound() on a position finds the nearest cell of the sequence in logarithmic time // - // note that we cannot a use an std::set because in some cases a position can occur more than once for the same seq: + // the cell index is part of the key because a position can occur more than once for the same seq: // - during performing a cache reuse via (rm + add) // - some vision models have input embeddings with repeating positions // - std::map seq_pos[LLAMA_MAX_SEQ]; + std::set> seq_pos[LLAMA_MAX_SEQ]; // helper functions for updating `seq_pos`, once cell at a time: - void seq_pos_dec(llama_seq_id s, llama_pos p) { - auto it = seq_pos[s].find(p); - assert(it != seq_pos[s].end()); - - if (--it->second == 0) { - seq_pos[s].erase(it); - } + void seq_pos_dec(llama_seq_id s, uint32_t i) { + const auto n = seq_pos[s].erase({ pos[i], i }); + assert(n == 1); + GGML_UNUSED(n); } - void seq_pos_inc(llama_seq_id s, llama_pos p) { - seq_pos[s][p]++; + void seq_pos_inc(llama_seq_id s, uint32_t i) { + seq_pos[s].insert({ pos[i], i }); } // remove cell i void seq_pos_rm(uint32_t i) { for (int s = 0; s < LLAMA_MAX_SEQ; ++s) { if (seq[i].test(s)) { - seq_pos_dec(s, pos[i]); + seq_pos_dec(s, i); } } } @@ -526,7 +548,7 @@ class llama_kv_cells { void seq_pos_add(uint32_t i) { for (int s = 0; s < LLAMA_MAX_SEQ; ++s) { if (seq[i].test(s)) { - seq_pos_inc(s, pos[i]); + seq_pos_inc(s, i); } } } diff --git a/src/llama-kv-share.cpp b/src/llama-kv-share.cpp new file mode 100644 index 000000000000..8237997951f8 --- /dev/null +++ b/src/llama-kv-share.cpp @@ -0,0 +1,252 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "llama-kv-share.h" + +#include "llama-context.h" +#include "llama-ext.h" +#include "llama-impl.h" +#include "llama-kv-cache-iswa.h" +#include "llama-kv-cache.h" + +#include "ggml-backend.h" + +#include +#include +#ifndef _WIN32 +#include +#endif +#include + +// exported by ggml-base (ggml-backend-impl.h): a buffer made of several, freed together +extern "C" { GGML_API ggml_backend_buffer_t ggml_backend_multi_buffer_alloc_buffer(ggml_backend_buffer_t * buffers, size_t n_buffers); } + +namespace { + +thread_local bool g_share_next = false; +thread_local bool g_share = false; + +// every non-view tensor at a 4 KiB boundary, in chunks of at most 1 GiB: both devices compute the same layout +// whatever their own alignment, and each chunk stays below Vulkan's per-buffer limit (reads past 4 GiB go wrong) +constexpr size_t SHARE_ALIGN = 4096; +constexpr size_t SHARE_CHUNK = size_t(1) << 30; + +struct share_chunk { + size_t size = 0; + uint64_t layout = 1469598103934665603ull; // FNV-1a over type and shape of every tensor, in order + std::vector offs; // offset of each of its tensors +}; + +std::vector share_layout(ggml_context * ctx, ggml_backend_buffer_type_t buft) { + std::vector chunks(1); + for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != nullptr; t = ggml_get_next_tensor(ctx, t)) { + if (t->view_src != nullptr) { + continue; + } + const size_t n = GGML_PAD(std::max(ggml_nbytes(t), ggml_backend_buft_get_alloc_size(buft, t)), SHARE_ALIGN); + if (!chunks.back().offs.empty() && chunks.back().size + n > SHARE_CHUNK) { + chunks.emplace_back(); + } + share_chunk & c = chunks.back(); + auto mix = [&](uint64_t v) { c.layout = (c.layout ^ v) * 1099511628211ull; }; + c.offs.push_back(c.size); + c.size += n; + mix(t->type); + for (int d = 0; d < GGML_MAX_DIMS; d++) { + mix((uint64_t) t->ne[d]); + } + } + for (auto & c : chunks) { + c.size = std::max(c.size, SHARE_ALIGN); + } + return chunks; +} + +// place the tensors of ctx into the chunk buffers, in the order share_layout used +void share_place(ggml_context * ctx, const std::vector & chunks, const std::vector & bufs) { + size_t c = 0, i = 0; + for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != nullptr; t = ggml_get_next_tensor(ctx, t)) { + if (t->view_src != nullptr) { + continue; + } + if (i == chunks[c].offs.size()) { + c++; + i = 0; + } + t->buffer = nullptr; + t->data = nullptr; + ggml_backend_tensor_alloc(bufs[c], t, (char *) ggml_backend_buffer_get_base(bufs[c]) + chunks[c].offs[i++]); + } + for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != nullptr; t = ggml_get_next_tensor(ctx, t)) { + if (t->view_src != nullptr) { + t->buffer = nullptr; + t->data = nullptr; + ggml_backend_view_init(t); + } + } +} + +// one buffer for the KV cache's bookkeeping, like ggml_backend_alloc_ctx_tensors_from_buft returns +ggml_backend_buffer_t share_wrap(std::vector & bufs) { + return bufs.size() == 1 ? bufs[0] : ggml_backend_multi_buffer_alloc_buffer(bufs.data(), bufs.size()); +} + +void * share_proc(ggml_backend_buffer_type_t buft, const char * name) { + ggml_backend_dev_t dev = ggml_backend_buft_get_device(buft); + ggml_backend_reg_t reg = dev ? ggml_backend_dev_backend_reg(dev) : nullptr; + return reg ? ggml_backend_reg_get_proc_address(reg, name) : nullptr; +} + +template +bool share_each(llama_context * dst, const llama_context * src, F && f) { + if (auto * d = dynamic_cast(llama_get_memory(dst))) { + auto * s = dynamic_cast(llama_get_memory(src)); + return s != nullptr && f(*d, *s); + } + if (auto * d = dynamic_cast(llama_get_memory(dst))) { + auto * s = dynamic_cast(llama_get_memory(src)); + return s != nullptr && f(*d->get_base(), *s->get_base()) && f(*d->get_swa(), *s->get_swa()); + } + return false; +} + +} // namespace + +llama_kv_shared_region::~llama_kv_shared_region() { +#ifndef _WIN32 + if (fd >= 0) { + close(fd); // a dma-buf, Linux only + } +#endif +} + +bool llama_kv_share_take_next(bool no_alloc) { + const bool take = !no_alloc && g_share_next; + if (!no_alloc) { + g_share_next = false; + } + return take; +} + +void llama_kv_share_begin(bool enabled) { g_share = enabled; } +void llama_kv_share_end() { g_share = false; } +bool llama_kv_share_enabled() { return g_share; } + +ggml_backend_buffer_t llama_kv_cache::share_alloc(ggml_context * ctx, ggml_backend_buffer_type_t buft) { + using export_fn = bool (*)(ggml_backend_buffer_t, int *, size_t *); + auto export_dmabuf = (export_fn) share_proc(buft, "ggml_backend_buffer_export_dmabuf"); + if (export_dmabuf == nullptr) { + throw std::runtime_error(format("shared KV: %s cannot export its memory", ggml_backend_buft_name(buft))); + } + const std::vector chunks = share_layout(ctx, buft); + std::vector bufs; + std::vector regions(chunks.size()); + for (size_t c = 0; c < chunks.size(); ++c) { + ggml_backend_buffer_t buf = ggml_backend_buft_alloc_buffer(buft, chunks[c].size); + if (buf == nullptr || !export_dmabuf(buf, ®ions[c].fd, ®ions[c].offset)) { + if (buf != nullptr) { + ggml_backend_buffer_free(buf); + } + for (auto * b : bufs) { + ggml_backend_buffer_free(b); + } + throw std::runtime_error(format("shared KV: %s allocation or dma-buf export failed", ggml_backend_buft_name(buft))); + } + regions[c].size = chunks[c].size; + regions[c].layout = chunks[c].layout; + bufs.push_back(buf); + } + share_place(ctx, chunks, bufs); + shared_.push_back(std::move(regions)); + return share_wrap(bufs); +} + +bool llama_kv_cache::share_from(const llama_kv_cache & src) { + if (other || src.other || src.shared_.size() != ctxs_bufs.size()) { + LLAMA_LOG_ERROR("%s: the source KV cache has %zu shared buffers, this one has %zu\n", __func__, src.shared_.size(), ctxs_bufs.size()); + return false; + } + using import_fn = ggml_backend_buffer_t (*)(ggml_backend_dev_t, int, size_t, size_t); + for (size_t i = 0; i < ctxs_bufs.size(); ++i) { + ggml_context * ctx = ctxs_bufs[i].first.get(); + ggml_backend_buffer_type_t buft = ggml_backend_buffer_get_type(ctxs_bufs[i].second.get()); + const auto & regions = src.shared_[i]; + const std::vector chunks = share_layout(ctx, buft); + bool same = chunks.size() == regions.size(); + for (size_t c = 0; same && c < chunks.size(); ++c) { + same = chunks[c].size == regions[c].size && chunks[c].layout == regions[c].layout; + } + if (!same) { + LLAMA_LOG_ERROR("%s: KV layout differs from the source (K/V types, shapes or flash attention)\n", __func__); + return false; + } + auto import_dmabuf = (import_fn) share_proc(buft, "ggml_backend_dev_buffer_from_dmabuf"); + std::vector bufs; + size_t total = 0; + for (const auto & region : regions) { + ggml_backend_buffer_t buf = import_dmabuf ? import_dmabuf(ggml_backend_buft_get_device(buft), region.fd, region.offset, region.size) : nullptr; + if (buf == nullptr) { + for (auto * b : bufs) { + ggml_backend_buffer_free(b); + } + LLAMA_LOG_ERROR("%s: %s cannot import the shared KV region\n", __func__, ggml_backend_buft_name(buft)); + return false; + } + bufs.push_back(buf); + total += region.size; + } + share_place(ctx, chunks, bufs); + ctxs_bufs[i].second.reset(share_wrap(bufs)); // frees the buffer this cache had + LLAMA_LOG_INFO("%s: %10s KV buffer now maps the shared region (%.2f MiB in %zu chunks, zero copy)\n", __func__, + ggml_backend_buft_name(buft), total / 1024.0 / 1024.0, regions.size()); + } + return true; +} + +bool llama_kv_cache::cells_copy_from(const llama_kv_cache & src) { + if (other || src.other || v_cells.size() != src.v_cells.size()) { + return false; + } + for (size_t s = 0; s < v_cells.size(); ++s) { + if (v_cells[s].size() != src.v_cells[s].size()) { + return false; + } + } + v_cells = src.v_cells; + v_heads = src.v_heads; + return true; +} + +void llama_context::kv_memory_rebound() { + gf_res_prev->reset(); + ggml_backend_sched_reset(sched.get()); + sched_need_reserve = true; +} + +void llama_kv_share_next(bool enabled) { + g_share_next = enabled; +} + +bool llama_kv_share_from(llama_context * dst, llama_context * src) { + if (!share_each(dst, src, [](llama_kv_cache & d, const llama_kv_cache & s) { return d.share_from(s); })) { + return false; + } + dst->kv_memory_rebound(); + return true; +} + +bool llama_kv_cells_copy(llama_context * dst, const llama_context * src) { + return share_each(dst, src, [](llama_kv_cache & d, const llama_kv_cache & s) { return d.cells_copy_from(s); }); +} diff --git a/src/llama-kv-share.h b/src/llama-kv-share.h new file mode 100644 index 000000000000..cca4752e0708 --- /dev/null +++ b/src/llama-kv-share.h @@ -0,0 +1,42 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include +#include + +// Zero-copy KV sharing between two contexts of one model on two devices of the same GPU (API in llama-ext.h). +// The owner context allocates each KV buffer as one exportable region with a fixed layout; the other maps it. + +struct llama_kv_shared_region { + int fd = -1; + size_t offset = 0; + size_t size = 0; + uint64_t layout = 0; // hash of tensor types and shapes: both sides must lay the cache out the same way + + llama_kv_shared_region() = default; + llama_kv_shared_region(llama_kv_shared_region && o) noexcept : fd(o.fd), offset(o.offset), size(o.size), layout(o.layout) { o.fd = -1; } + llama_kv_shared_region(const llama_kv_shared_region &) = delete; + ~llama_kv_shared_region(); +}; + +// the pending llama_kv_share_next request; a real (not no_alloc) context takes and clears it +bool llama_kv_share_take_next(bool no_alloc); + +// set by llama_context around memory creation, read by llama_kv_cache while it allocates +void llama_kv_share_begin(bool enabled); +void llama_kv_share_end(); +bool llama_kv_share_enabled(); diff --git a/src/llama-memory-hybrid-idx.cpp b/src/llama-memory-hybrid-idx.cpp new file mode 100644 index 000000000000..93b468784a33 --- /dev/null +++ b/src/llama-memory-hybrid-idx.cpp @@ -0,0 +1,679 @@ +#include "llama-memory-hybrid-idx.h" + +#include "llama-impl.h" +#include "llama-batch.h" +#include "llama-io.h" +#include "llama-model.h" + + +#include +#include +#include +#include +#include + +// +// llama_memory_hybrid_idx +// + +llama_memory_hybrid_idx::llama_memory_hybrid_idx( + const llama_model & model, + /* attn */ + ggml_type type_k, + ggml_type type_v, + bool v_trans, + uint32_t kv_size, + uint32_t n_pad, + uint32_t n_swa, + llama_swa_type swa_type, + /* recurrent */ + ggml_type type_r, + ggml_type type_s, + uint32_t rs_size, + /* common */ + uint32_t n_seq_max, + uint32_t n_rs_seq, + bool offload, + bool unified, + /* layer filters */ + const layer_filter_cb & filter_attn, + const layer_filter_cb & filter_recr, + const layer_filter_cb & filter_idx) : + llama_memory_hybrid( + model, + type_k, type_v, v_trans, kv_size, n_pad, n_swa, swa_type, + type_r, type_s, rs_size, + n_seq_max, n_rs_seq, offload, unified, + filter_attn, filter_recr), + hparams_idx(model.hparams), + mem_idx(filter_idx == nullptr ? nullptr : [&] { + // MQA with a single key head of indexer_head_size, as llama_kv_cache_dsa shapes its own + std::fill(hparams_idx.n_head_kv_arr.begin(), hparams_idx.n_head_kv_arr.end(), 1); + hparams_idx.n_embd_head_k_full = model.hparams.indexer_head_size; + + // the cached indexer keys are raw, rotation happens after pooling at read time, so a + // K-shift must not rotate them while the stream copies in the same update still apply + hparams_idx.rope_type = LLAMA_ROPE_TYPE_NONE; + + LLAMA_LOG_INFO("%s: creating indexer KV cache, size = %u cells\n", __func__, kv_size); + + return new llama_kv_cache( + model, hparams_idx, type_k, type_v, v_trans, offload, unified, + kv_size, n_seq_max, n_pad, n_swa, swa_type, + nullptr, filter_idx, nullptr, nullptr, "idx_"); + }()) {} + +llama_memory_context_ptr llama_memory_hybrid_idx::init_batch(llama_batch_allocr & balloc, uint32_t n_ubatch, bool embd_all) { + // note: repeats llama_memory_hybrid::init_batch, as the indexer needs the attention slot infos that the base context hides + do { + balloc.split_reset(); + + // follow the recurrent pattern for creating the ubatch splits + std::vector ubatches; + + while (true) { + llama_ubatch ubatch; + + if (embd_all) { + // if all tokens are output, split by sequence + ubatch = balloc.split_seq(n_ubatch); + } else { + // Use non-sequential split when KV cache is unified (needed for hellaswag/winogrande/multiple-choice) + const bool unified = (get_mem_attn()->get_n_stream() == 1); + + // [TAG_RECURRENT_ROLLBACK_SPLITS] + // the trailing (1 + n_rs_seq) tokens of each seq must stay in the same ubatch + // so that the rollback snapshots remain valid + const uint32_t n_rs_seq = get_mem_recr()->n_rs_seq; + + ubatch = balloc.split_equal(n_ubatch, !unified, n_rs_seq > 0 ? n_rs_seq + 1 : 0); + } + + if (ubatch.n_tokens == 0) { + break; + } + + ubatches.push_back(std::move(ubatch)); // NOLINT + } + + if (balloc.get_n_used() < balloc.get_n_tokens()) { + // failed to find a suitable split + break; + } + + // prepare the recurrent batches first + if (!get_mem_recr()->prepare(ubatches)) { + // TODO: will the recurrent cache be in an undefined context at this point? + LLAMA_LOG_ERROR("%s: failed to prepare recurrent ubatches\n", __func__); + return std::make_unique(LLAMA_MEMORY_STATUS_FAILED_PREPARE); + } + + // prepare the attention cache + auto heads_attn = get_mem_attn()->prepare(ubatches); + if (heads_attn.empty()) { + LLAMA_LOG_ERROR("%s: failed to prepare attention ubatches\n", __func__); + return std::make_unique(LLAMA_MEMORY_STATUS_FAILED_PREPARE); + } + + // the indexer uses the attention cache's slot layout; a separate one can drift from it + llama_kv_cache::slot_info_vec_t heads_idx; + if (mem_idx) { + heads_idx = heads_attn; + } + + return std::make_unique( + this, std::move(heads_attn), std::move(heads_idx), std::move(ubatches)); + } while(false); + + return std::make_unique(LLAMA_MEMORY_STATUS_FAILED_PREPARE); +} + +llama_memory_context_ptr llama_memory_hybrid_idx::init_full() { + return std::make_unique(this); +} + +llama_memory_context_ptr llama_memory_hybrid_idx::init_update(llama_context * lctx, bool optimize) { + return std::make_unique(this, lctx, optimize); +} + +void llama_memory_hybrid_idx::clear(bool data) { + llama_memory_hybrid::clear(data); + + if (mem_idx) { + mem_idx->clear(data); + } +} + +bool llama_memory_hybrid_idx::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) { + // same order as llama_memory_hybrid::seq_rm: the recurrent cache can refuse, so try it first + if (!get_mem_recr()->seq_rm(seq_id, p0, p1)) { + return false; + } + + if (mem_idx) { + mem_idx->seq_rm(seq_id, p0, p1); + } + + return get_mem_attn()->seq_rm(seq_id, p0, p1); +} + +void llama_memory_hybrid_idx::seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) { + llama_memory_hybrid::seq_cp(seq_id_src, seq_id_dst, p0, p1); + + if (mem_idx) { + mem_idx->seq_cp(seq_id_src, seq_id_dst, p0, p1); + } +} + +void llama_memory_hybrid_idx::seq_keep(llama_seq_id seq_id) { + llama_memory_hybrid::seq_keep(seq_id); + + if (mem_idx) { + mem_idx->seq_keep(seq_id); + } +} + +void llama_memory_hybrid_idx::seq_add(llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) { + llama_memory_hybrid::seq_add(seq_id, p0, p1, shift); + + if (mem_idx) { + mem_idx->seq_add(seq_id, p0, p1, shift); + } +} + +void llama_memory_hybrid_idx::seq_div(llama_seq_id seq_id, llama_pos p0, llama_pos p1, int d) { + llama_memory_hybrid::seq_div(seq_id, p0, p1, d); + + if (mem_idx) { + mem_idx->seq_div(seq_id, p0, p1, d); + } +} + +std::map llama_memory_hybrid_idx::memory_breakdown() const { + std::map mb = llama_memory_hybrid::memory_breakdown(); + + if (mem_idx) { + for (const auto & buft_size : mem_idx->memory_breakdown()) { + mb[buft_size.first] += buft_size.second; + } + } + + return mb; +} + +void llama_memory_hybrid_idx::state_write(llama_io_write_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) const { + llama_memory_hybrid::state_write(io, seq_id, flags); + + // [TAG_HYBRID_IDX_STATE] the indexer section goes last, so it is a pure suffix: an old reader stops early instead of misparsing it + // The indexer mirrors the attention cache, so it uses the same PARTIAL_ONLY gate. + if ((flags & LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY) == 0) { + if (mem_idx) { + mem_idx->state_write(io, seq_id, flags); + } + } + +} + +void llama_memory_hybrid_idx::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) { + // note: repeats llama_memory_hybrid::state_read + // the indexer needs the attention cache's cells, and a half-failed restore must leave all three caches alike + + // [TAG_HYBRID_IDX_SINFO] + // the indexer restore adopts the attention cache's layout instead of searching for cells of its own + // two find_slot calls agree only while both caches see the same occupancy, which a restore cannot promise + llama_kv_cache::slot_info_vec_t sinfos_attn; + + try { + if ((flags & LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY) == 0) { + get_mem_attn()->state_read_sinfo(io, seq_id, flags, mem_idx ? &sinfos_attn : nullptr, nullptr); + } + + get_mem_recr()->state_read(io, seq_id, flags); + + // [TAG_HYBRID_IDX_STATE] must mirror the write order in state_write + if ((flags & LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY) == 0) { + if (mem_idx) { + mem_idx->state_read_sinfo(io, seq_id, flags, nullptr, &sinfos_attn); + } + } + + } catch (...) { + // a half-restored context is the one state the indexer cannot fix by itself: attention holds new cells, the indexer old ones + // drop what was being restored from all of them, which is a state they do agree on. + state_drop(seq_id); + + throw; + } +} + +void llama_memory_hybrid_idx::state_drop(llama_seq_id seq_id) { + // dropped directly, not via seq_rm: the recurrent cache may refuse it and then only the other two get cleared + if (seq_id < 0) { + clear(true); + + return; + } + + get_mem_attn()->seq_rm(seq_id, -1, -1); + get_mem_recr()->seq_rm(seq_id, -1, -1); + + if (mem_idx) { + mem_idx->seq_rm(seq_id, -1, -1); + } +} + +llama_kv_cache * llama_memory_hybrid_idx::get_mem_idx() const { + return mem_idx.get(); +} + +void llama_memory_hybrid_idx::set_input_qsa( + ggml_tensor * cell_blk, + ggml_tensor * blk_cells, + ggml_tensor * blk_pos, + ggml_tensor * bias, + const llama_ubatch * ubatch, + uint32_t ratio, + bool blk_bias) const { + GGML_ASSERT(ratio > 0); + GGML_ASSERT(get_mem_idx() != nullptr); + + GGML_ASSERT(ggml_backend_buffer_is_host(cell_blk->buffer)); + + const int64_t n_kv = cell_blk->ne[0]; + const int64_t n_ns = cell_blk->ne[1]; // streams in this ubatch + const int64_t n_blocks = blk_pos->ne[0]/(4*n_ns); + const int64_t n_tokens = ubatch->n_tokens; + const int64_t r = ratio; + + GGML_ASSERT(n_tokens % n_ns == 0); + const int64_t n_tps = n_tokens/n_ns; // tokens per stream + + int32_t * dst_cell_blk = (int32_t *) cell_blk->data; + int32_t * dst_blk_cells = (int32_t *) blk_cells->data; + int32_t * dst_blk_pos = (int32_t *) blk_pos->data; + float * dst_bias = (float *) bias->data; + + // a block is keyed on (sequence set, index bucket): a unified cache counts every sequence + // from zero, so the bucket alone would pool two sequences into one block + GGML_ASSERT(r <= 64); + const uint64_t slots_full = r == 64 ? ~uint64_t(0) : ((uint64_t(1) << r) - 1); + + // TODO: this runs per ubatch and is O(n_kv) per stream, about 865 us at 33k context. the cost + // is the per-cell scan rather than these allocations, so hoisting them buys nothing + std::vector blk_of(n_kv); + std::vector cell_grp(n_kv); + std::vector grp_head(n_blocks); + std::vector grp_next; + std::vector grp_first; + std::vector grp_slot0; + std::vector grp_slots; + std::vector grp_bid; + std::vector bid_idx; + std::vector bid_cell; + std::vector bid_slot0; + + std::vector order; + std::vector rank; + + std::fill(dst_blk_pos, dst_blk_pos + 4*n_blocks*n_ns, 0); + + for (int64_t s = 0; s < n_ns; ++s) { + // ubatch index s*n_tps belongs to this stream; ask which cells array it uses + const llama_seq_id seq_of_stream = ubatch->seq_id[s*n_tps][0]; + const auto & cells = get_mem_idx()->get_cells(seq_of_stream); + + int32_t * cur_cell_blk = dst_cell_blk + s*n_kv; + int32_t * cur_blk_cells = dst_blk_cells + s*(r*n_blocks); + + std::fill(cur_blk_cells, cur_blk_cells + r*n_blocks, 0); + + bid_idx .clear(); + bid_cell .clear(); + bid_slot0.clear(); + + int n_seq_present = 0; + + for (int sq = 0; sq < LLAMA_MAX_SEQ && n_seq_present < 2; ++sq) { + if (cells.seq_pos_min(sq) >= 0) { + n_seq_present++; + } + } + + const bool one_seq = n_seq_present <= 1; + + // a cell no block covers needs its own -inf, which a per-block bias cannot carry + // every cache path keeps the position below the cell window, so this stays false + bool oor = false; + + bool dup = false; + + bool ranked = false; + + auto group_cells = [&]() { + // -1 means no usable block: an incomplete or short group cannot be pooled + std::fill(blk_of.begin(), blk_of.end(), -1); + std::fill(cell_grp.begin(), cell_grp.end(), -1); + std::fill(grp_head.begin(), grp_head.end(), -1); + + grp_next .clear(); + grp_first.clear(); + grp_slot0.clear(); + grp_slots.clear(); + grp_bid .clear(); + + oor = false; + dup = false; + + for (int64_t j = 0; j < n_kv; ++j) { + if (cells.is_empty(j)) { + continue; + } + + const int64_t idx = ranked ? rank[j] : cells.pos_get(j); + const int64_t pb = idx/r; + + if (pb >= n_blocks) { + oor = true; + continue; + } + + int32_t g = -1; + + for (int32_t c = grp_head[pb]; c >= 0; c = grp_next[c]) { + if (one_seq || cells.seq_get_all((uint32_t) grp_first[c]) == cells.seq_get_all((uint32_t) j)) { + g = c; + break; + } + } + + if (g < 0) { + g = (int32_t) grp_first.size(); + + grp_next .push_back(grp_head[pb]); + grp_first.push_back((int32_t) j); + grp_slot0.push_back(-1); + grp_slots.push_back(0); + grp_bid .push_back(-1); + + grp_head[pb] = g; + } + + const uint64_t bit = uint64_t(1) << (idx%r); + + dup |= (grp_slots[g] & bit) != 0; + + cell_grp[j] = g; + grp_slots[g] |= bit; + + if (idx%r == 0) { + grp_slot0[g] = (int32_t) j; + } + } + }; + + group_cells(); + + // mrope repeats one position across an image, so rank cells instead of using the position + if (dup && ubatch->is_pos_2d() && one_seq) { + order.clear(); + order.reserve(n_kv); + + for (int64_t j = 0; j < n_kv; ++j) { + if (!cells.is_empty(j)) { + order.push_back((int32_t) j); + } + } + + // same total order the mrope causal mask uses: pos, then ext.y, then ext.x + std::sort(order.begin(), order.end(), [&cells](int32_t a, int32_t b) { + const llama_pos pa = cells.pos_get(a); + const llama_pos pb = cells.pos_get(b); + + if (pa != pb) { + return pa < pb; + } + + const auto & ea = cells.ext_get(a); + + return cells.ext_get(b).is_2d_gt(ea.x, ea.y); + }); + + rank.assign(n_kv, -1); + + for (int64_t k = 0; k < (int64_t) order.size(); ++k) { + rank[order[k]] = (int32_t) k; + } + + ranked = true; + + group_cells(); + } + + GGML_ASSERT((!blk_bias || !oor) && "qsa: cell position runs past the cell window"); + + int32_t n_bid = 0; + + for (int64_t pb = 0; pb < n_blocks; ++pb) { + for (int32_t g = grp_head[pb]; g >= 0; g = grp_next[g]) { + if (grp_slots[g] != slots_full) { + continue; + } + + grp_bid[g] = n_bid++; + + bid_idx .push_back((int32_t) (pb*r)); + bid_cell .push_back(grp_first[g]); + bid_slot0.push_back(grp_slot0[g]); + } + } + + GGML_ASSERT(n_bid <= n_blocks); + + for (int32_t b = 0; b < n_bid; ++b) { + int32_t sec_pos[4] = { bid_idx[b], bid_idx[b], bid_idx[b], bid_idx[b] }; + + if (ranked) { + const int32_t c = bid_slot0[b]; + const llama_pos p = cells.pos_get(c); + const auto & e = cells.ext_get(c); + + sec_pos[0] = p; + sec_pos[1] = e.y; + sec_pos[2] = e.x; + sec_pos[3] = p; + } + + for (int64_t sec = 0; sec < 4; ++sec) { + dst_blk_pos[sec*(n_blocks*n_ns) + s*n_blocks + b] = sec_pos[sec]; + } + } + + // unpooled cells all point at one spare block. a spare block exists only when some + // cell is unpooled: n_bid == n_blocks means every cell sits in a full block. + const bool have_dead = n_bid < n_blocks; + const int32_t dead_bid = have_dead ? n_bid : n_blocks - 1; + + for (int64_t j = 0; j < n_kv; ++j) { + const int32_t g = cell_grp[j]; + + blk_of[j] = g < 0 ? -1 : grp_bid[g]; + + if (blk_of[j] >= 0) { + const int64_t idx = ranked ? rank[j] : cells.pos_get(j); + + cur_blk_cells[blk_of[j]*r + (idx%r)] = (int32_t) j; + } + + cur_cell_blk[j] = blk_of[j] < 0 ? dead_bid : blk_of[j]; + } + + for (int64_t ii = 0; ii < n_tps; ++ii) { + const int64_t i = s*n_tps + ii; + const llama_seq_id seq_id = ubatch->seq_id[i][0]; + + int64_t q = ubatch->pos[i]; + + if (ranked) { + const llama_pos qt = ubatch->pos[i]; + const llama_pos qy = ubatch->pos[i + n_tokens]; + const llama_pos qx = ubatch->pos[i + n_tokens*2]; + + int64_t lo = 0; + int64_t hi = (int64_t) order.size(); + + while (lo < hi) { + const int64_t mid = (lo + hi)/2; + const int32_t c = order[mid]; + const llama_pos pc = cells.pos_get(c); + + if (pc < qt || (pc == qt && !cells.ext_get(c).is_2d_gt(qx, qy))) { + lo = mid + 1; + } else { + hi = mid; + } + } + + q = lo - 1; + } + + // the tail is an incomplete block and is always visible, as in the reference + const int64_t tail_start = (q + 1)/r*r; + + if (blk_bias) { + // a block sits wholly inside or outside the tail, so one value covers it + // the caller adds the attention mask, which drops empty, foreign and future cells + float * cur_blk_bias = dst_bias + i*n_blocks; + + for (int64_t b = 0; b < n_blocks; ++b) { + if (b >= n_bid || !cells.seq_has((uint32_t) bid_cell[b], seq_id)) { + cur_blk_bias[b] = -INFINITY; + continue; + } + + // finite, so it can never meet a -inf and produce a nan + cur_blk_bias[b] = bid_idx[b] >= tail_start ? 1e9f : 0.0f; + } + + // the spare block holds the unpooled cells, which are the incomplete tail, so + // it gets the tail value. it must stay finite: a sequence with fewer than + // `ratio` cells owns no full block, and a row of -inf only gives a nan. + if (have_dead) { + cur_blk_bias[dead_bid] = 1e9f; + } + + continue; + } + + float * cur_bias = dst_bias + i*n_kv; + + for (int64_t j = 0; j < n_kv; ++j) { + float v = -INFINITY; + + if (!cells.is_empty(j) && cells.seq_has(j, seq_id)) { + const int64_t idx = ranked ? rank[j] : cells.pos_get(j); + + if (idx <= q) { + // finite, so it can never meet a -inf and produce a nan + v = idx >= tail_start ? 1e9f : (blk_of[j] < 0 ? -INFINITY : 0.0f); + } + } + + cur_bias[j] = v; + } + } + } +} + +// +// llama_memory_hybrid_idx_context +// + +// streams in each ubatch's slot info, matching get_k/get_v's `ns` +static std::vector llama_memory_hybrid_idx_ns(const llama_kv_cache::slot_info_vec_t & sinfos) { + std::vector res; + res.reserve(sinfos.size()); + + for (const auto & sinfo : sinfos) { + res.push_back(sinfo.s1 - sinfo.s0 + 1); + } + + return res; +} + +llama_memory_hybrid_idx_context::llama_memory_hybrid_idx_context(llama_memory_status status) : + llama_memory_hybrid_context(status) {} + +llama_memory_hybrid_idx_context::llama_memory_hybrid_idx_context(llama_memory_hybrid_idx * mem) : + llama_memory_hybrid_context(mem), + mem(mem), + // graph reservation walks a full context, and qwen4exp builds the sparse attention only when this is set + // without it the reserved worst case is the dense graph, so ggml-alloc must grow the buffer on the first decode + ns_ubatch(mem->get_mem_idx() == nullptr ? + std::vector() : std::vector{ mem->get_mem_idx()->get_n_stream() }), + ctx_idx(mem->get_mem_idx() == nullptr ? nullptr : + new llama_kv_cache_context(mem->get_mem_idx())) {} + +llama_memory_hybrid_idx_context::llama_memory_hybrid_idx_context( + llama_memory_hybrid_idx * mem, + llama_context * lctx, + bool optimize) : + llama_memory_hybrid_context(mem, lctx, optimize), + mem(mem), + // update() applies a pending cross-stream seq_cp, else the copy keeps stale indexer keys + ctx_idx(mem->get_mem_idx() == nullptr ? nullptr : + mem->get_mem_idx()->init_update(lctx, optimize)) {} + +llama_memory_hybrid_idx_context::llama_memory_hybrid_idx_context( + llama_memory_hybrid_idx * mem, + slot_info_vec_t sinfos_attn, + slot_info_vec_t sinfos_idx, + std::vector ubatches) : + // note: the base copies the ubatches; ctx_idx gets a copy of its own + llama_memory_hybrid_context(mem, std::move(sinfos_attn), ubatches), + mem(mem), + ns_ubatch(llama_memory_hybrid_idx_ns(sinfos_idx)), + ctx_idx(mem->get_mem_idx() == nullptr ? nullptr : + new llama_kv_cache_context(mem->get_mem_idx(), std::move(sinfos_idx), ubatches)) {} + +bool llama_memory_hybrid_idx_context::next() { + if (ctx_idx) { + ctx_idx->next(); + } + + ++i_cur; + + return llama_memory_hybrid_context::next(); +} + +bool llama_memory_hybrid_idx_context::apply() { + bool res = llama_memory_hybrid_context::apply(); + + if (ctx_idx) { + res = res & ctx_idx->apply(); + } + + return res; +} + +const llama_kv_cache_context * llama_memory_hybrid_idx_context::get_idx() const { + return static_cast(ctx_idx.get()); +} + +uint32_t llama_memory_hybrid_idx_context::get_n_stream() const { + GGML_ASSERT(i_cur < ns_ubatch.size()); + + return ns_ubatch[i_cur]; +} + +void llama_memory_hybrid_idx_context::set_input_qsa( + ggml_tensor * cell_blk, + ggml_tensor * blk_cells, + ggml_tensor * blk_pos, + ggml_tensor * bias, + const llama_ubatch * ubatch, + uint32_t ratio, + bool blk_bias) const { + GGML_ASSERT(mem != nullptr); + + mem->set_input_qsa(cell_blk, blk_cells, blk_pos, bias, ubatch, ratio, blk_bias); +} diff --git a/src/llama-memory-hybrid-idx.h b/src/llama-memory-hybrid-idx.h new file mode 100644 index 000000000000..705189e7eb58 --- /dev/null +++ b/src/llama-memory-hybrid-idx.h @@ -0,0 +1,160 @@ +#pragma once + +#include "llama-memory-hybrid.h" + +#include +#include + +// +// llama_memory_hybrid_idx +// + +// llama_memory_hybrid plus a third cache with one indexer key per token, for block-sparse attention (qwen4exp QSA) +// the indexer is a side buffer over the attention cells: same size, padding, streams and slots, so cell j is one token in both + +class llama_memory_hybrid_idx : public llama_memory_hybrid { +public: + llama_memory_hybrid_idx( + const llama_model & model, + /* attn */ + ggml_type type_k, + ggml_type type_v, + bool v_trans, + uint32_t kv_size, + uint32_t n_pad, + uint32_t n_swa, + llama_swa_type swa_type, + /* recurrent */ + ggml_type type_r, + ggml_type type_s, + uint32_t rs_size, + /* common */ + uint32_t n_seq_max, + uint32_t n_rs_seq, + bool offload, + bool unified, + /* layer filters */ + const layer_filter_cb & filter_attn, + const layer_filter_cb & filter_recr, + /* the indexer cache exists only if this is given */ + const layer_filter_cb & filter_idx); + + ~llama_memory_hybrid_idx() = default; + + // + // llama_memory_i + // + + llama_memory_context_ptr init_batch( + llama_batch_allocr & balloc, + uint32_t n_ubatch, + bool embd_all) override; + + llama_memory_context_ptr init_full() override; + + llama_memory_context_ptr init_update(llama_context * lctx, bool optimize) override; + + void clear(bool data) override; + + bool seq_rm (llama_seq_id seq_id, llama_pos p0, llama_pos p1) override; + void seq_cp (llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) override; + void seq_keep(llama_seq_id seq_id) override; + void seq_add (llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) override; + void seq_div (llama_seq_id seq_id, llama_pos p0, llama_pos p1, int d) override; + + std::map memory_breakdown() const override; + + // state write/load + + void state_write(llama_io_write_i & io, llama_seq_id seq_id = -1, llama_state_seq_flags flags = 0) const override; + void state_read (llama_io_read_i & io, llama_seq_id seq_id = -1, llama_state_seq_flags flags = 0) override; + + // + // llama_memory_hybrid_idx specific API + // + + llama_kv_cache * get_mem_idx() const; // nullptr when the model carries no indexer + + // block-compressed sparse attention (qwen4exp QSA) over the cells of the indexer cache. + // Blocks cut the position line, not the cell array, so no caller assumes a contiguous layout: + // cell_blk I32 [n_kv, ns] block each cell belongs to + // blk_cells I32 [ratio*n_blocks, ns] cells making up each block + // blk_pos I32 [4*n_blocks*ns] mrope position rows of each block's first token + // bias F32 [n_kv, n_tokens/ns, ns] -inf where invisible, large where always visible + // blk_bias asks for the bias per block instead: [n_blocks, n_tokens/ns, ns] + // the caller then adds the attention mask, the only part of the bias that varies within a block + void set_input_qsa(ggml_tensor * cell_blk, ggml_tensor * blk_cells, ggml_tensor * blk_pos, + ggml_tensor * bias, const llama_ubatch * ubatch, uint32_t ratio, + bool blk_bias) const; + +private: + // forget seq_id (all of it if seq_id < 0) in every cache at once, so a failed restore cannot leave the caches out of step + // seq_id < 0 drops the whole context, as the caches themselves do on a failed restore + void state_drop(llama_seq_id seq_id); + + // the indexer cache holds one key head per layer, so it needs its own hparams: + // llama_kv_cache keeps a reference to what it is given + llama_hparams hparams_idx; + + const std::unique_ptr mem_idx; +}; + +class llama_memory_hybrid_idx_context : public llama_memory_hybrid_context { +public: + using slot_info_vec_t = llama_kv_cache::slot_info_vec_t; + + // used for errors + explicit llama_memory_hybrid_idx_context(llama_memory_status status); + + // used to create a full-cache context + explicit llama_memory_hybrid_idx_context(llama_memory_hybrid_idx * mem); + + // used to create an update context + llama_memory_hybrid_idx_context( + llama_memory_hybrid_idx * mem, + llama_context * lctx, + bool optimize); + + // used to create a batch processing context from a batch + llama_memory_hybrid_idx_context( + llama_memory_hybrid_idx * mem, + slot_info_vec_t sinfos_attn, + slot_info_vec_t sinfos_idx, + std::vector ubatches); + + ~llama_memory_hybrid_idx_context() = default; + + // + // llama_memory_context_i + // + + bool next() override; + bool apply() override; + + // + // llama_memory_hybrid_idx_context specific API + // + + // nullptr with no indexer + const llama_kv_cache_context * get_idx() const; + + // streams in the current slot info, the `ns` of get_k/get_v; 1 if unified + uint32_t get_n_stream() const; + + void set_input_qsa(ggml_tensor * cell_blk, ggml_tensor * blk_cells, ggml_tensor * blk_pos, + ggml_tensor * bias, const llama_ubatch * ubatch, uint32_t ratio, + bool blk_bias) const; + +private: + const llama_memory_hybrid_idx * mem = nullptr; + + // streams per ubatch, read from the slot infos before ctx_idx takes them + // declared first, so it is initialised while sinfos_idx is still intact + const std::vector ns_ubatch; + + // null unless the model has an indexer + const llama_memory_context_ptr ctx_idx; + + // mirrors the base class's ubatch cursor, which is private there + size_t i_cur = 0; +}; diff --git a/src/llama-memory-recurrent.cpp b/src/llama-memory-recurrent.cpp index ef82eb976ca7..b96f221112b8 100644 --- a/src/llama-memory-recurrent.cpp +++ b/src/llama-memory-recurrent.cpp @@ -51,7 +51,8 @@ llama_memory_recurrent::llama_memory_recurrent( auto it = ctx_map.find(buft); if (it == ctx_map.end()) { ggml_init_params params = { - /*.mem_size =*/ size_t(2u*n_layer*ggml_tensor_overhead()), + // r and s per layer, plus the separate PLE conv row where the model has one + /*.mem_size =*/ size_t((hparams.ple_conv_state() > 0 ? 3u : 2u)*n_layer*ggml_tensor_overhead()), /*.mem_buffer =*/ NULL, /*.no_alloc =*/ true, }; @@ -71,6 +72,7 @@ llama_memory_recurrent::llama_memory_recurrent( r_l.resize(n_layer); s_l.resize(n_layer); + p_l.resize(n_layer); for (int i = 0; i < n_layer; i++) { if (filter && !filter(i)) { @@ -103,6 +105,13 @@ llama_memory_recurrent::llama_memory_recurrent( ggml_format_name(s, "cache_s_l%d", i); r_l[i] = r; s_l[i] = s; + + // the PLE history needs its own row: Meta must mirror it while the delta-net conv state next door stays split + if (hparams.ple_conv_state() > 0 && hparams.is_ple(i)) { + ggml_tensor * p = ggml_new_tensor_2d(ctx, type_r, hparams.ple_conv_state(), n_rows); + ggml_format_name(p, "cache_ple_r_l%d", i); + p_l[i] = p; + } } // allocate tensors and initialize the buffers to avoid NaNs in the padding @@ -119,11 +128,13 @@ llama_memory_recurrent::llama_memory_recurrent( { const size_t memory_size_r = size_r_bytes(); const size_t memory_size_s = size_s_bytes(); + const size_t memory_size_p = size_p_bytes(); - LLAMA_LOG_INFO("%s: size = %7.2f MiB (%6u cells, %3d layers, %2u seqs %2u rs_seq), R (%s): %7.2f MiB, S (%s): %7.2f MiB\n", __func__, - (float)(memory_size_r + memory_size_s) / (1024.0f * 1024.0f), mem_size, n_layer, n_seq_max, n_rs_seq, + LLAMA_LOG_INFO("%s: size = %7.2f MiB (%6u cells, %3d layers, %2u seqs %2u rs_seq), R (%s): %7.2f MiB, S (%s): %7.2f MiB, P (%s): %7.2f MiB\n", __func__, + (float)(memory_size_r + memory_size_s + memory_size_p) / (1024.0f * 1024.0f), mem_size, n_layer, n_seq_max, n_rs_seq, ggml_type_name(type_r), (float)memory_size_r / (1024.0f * 1024.0f), - ggml_type_name(type_s), (float)memory_size_s / (1024.0f * 1024.0f)); + ggml_type_name(type_s), (float)memory_size_s / (1024.0f * 1024.0f), + ggml_type_name(type_r), (float)memory_size_p / (1024.0f * 1024.0f)); } } @@ -730,6 +741,18 @@ size_t llama_memory_recurrent::size_s_bytes() const { return size_s_bytes; } +size_t llama_memory_recurrent::size_p_bytes() const { + size_t size_p_bytes = 0; + + for (const auto & p : p_l) { + if (p != nullptr) { + size_p_bytes += ggml_nbytes(p); + } + } + + return size_p_bytes; +} + void llama_memory_recurrent::state_write(llama_io_write_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) const { GGML_UNUSED(flags); @@ -891,6 +914,17 @@ void llama_memory_recurrent::state_write_data(llama_io_write_i & io, const std:: const size_t buf_size = range_size * r_size_row; io.write_tensor(r_l[il], range.first * r_size_row, buf_size); } + + // the PLE conv history is a second recurrent row, so it has to travel with the first + if (p_l[il] != nullptr) { + const uint64_t p_size_row = ggml_row_size(p_l[il]->type, hparams.ple_conv_state()); + io.write(&p_size_row, sizeof(p_size_row)); + + for (const auto & range : cell_ranges) { + const size_t range_size = range.second - range.first; + io.write_tensor(p_l[il], range.first * p_size_row, range_size * p_size_row); + } + } } if (!s_trans) { @@ -1089,6 +1123,20 @@ bool llama_memory_recurrent::state_read_data(llama_io_read_i & io, uint32_t cell // Read and set the keys for the whole cell range io.read_tensor(r_l[il], head * r_size_row, cell_count * r_size_row); } + + if (p_l[il] != nullptr) { + uint64_t p_size_row_ref; + io.read(&p_size_row_ref, sizeof(p_size_row_ref)); + const size_t p_size_row = ggml_row_size(p_l[il]->type, hparams.ple_conv_state()); + if (p_size_row != p_size_row_ref) { + LLAMA_LOG_ERROR("%s: mismatched ple row size (%zu != %zu, layer %d)\n", __func__, p_size_row, (size_t) p_size_row_ref, il); + return false; + } + + if (cell_count) { + io.read_tensor(p_l[il], head * p_size_row, cell_count * p_size_row); + } + } } if (!s_trans) { @@ -1243,6 +1291,10 @@ ggml_tensor * llama_memory_recurrent_context::get_s_l(int32_t il) const { return mem->s_l[il]; } +ggml_tensor * llama_memory_recurrent_context::get_p_l(int32_t il) const { + return mem->p_l[il]; +} + int32_t llama_memory_recurrent_context::s_copy(int i) const { const uint32_t cell_idx = i + mem->head; const int32_t src0 = mem->cells[cell_idx].src0; diff --git a/src/llama-memory-recurrent.h b/src/llama-memory-recurrent.h index b13b7b748f5e..4abb3f5cf5c0 100644 --- a/src/llama-memory-recurrent.h +++ b/src/llama-memory-recurrent.h @@ -111,6 +111,8 @@ class llama_memory_recurrent : public llama_memory_i { // per layer std::vector r_l; std::vector s_l; + // a second conv history that must stay replicated across devices, so it cannot share the r row + std::vector p_l; private: //const llama_model & model; @@ -125,6 +127,7 @@ class llama_memory_recurrent : public llama_memory_i { size_t size_r_bytes() const; size_t size_s_bytes() const; + size_t size_p_bytes() const; void state_write_meta(llama_io_write_i & io, const std::vector> & cell_ranges, llama_seq_id seq_id = -1) const; void state_write_data(llama_io_write_i & io, const std::vector> & cell_ranges) const; @@ -170,6 +173,7 @@ class llama_memory_recurrent_context : public llama_memory_context_i { ggml_tensor * get_r_l(int32_t il) const; ggml_tensor * get_s_l(int32_t il) const; + ggml_tensor * get_p_l(int32_t il) const; int32_t s_copy(int i) const; diff --git a/src/llama-mmap.cpp b/src/llama-mmap.cpp index ed572da7fb54..4d183cbc9c45 100644 --- a/src/llama-mmap.cpp +++ b/src/llama-mmap.cpp @@ -438,11 +438,34 @@ void llama_file::write_u32(uint32_t val) const { pimpl->write_u32(val); } // llama_mmap +#if defined(_POSIX_MAPPED_FILES) || defined(_WIN32) +// merge `ranges` and return their complement within [0, limit) +static llama_mmap::ranges ranges_complement(llama_mmap::ranges ranges, size_t limit) { + llama_mmap::ranges res; + std::sort(ranges.begin(), ranges.end()); + + size_t pos = 0; + for (const auto & range : ranges) { + const size_t beg = std::min(range.first, limit); + const size_t end = std::min(range.second, limit); + if (beg > pos) { + res.emplace_back(pos, beg); + } + pos = std::max(pos, end); + } + if (pos < limit) { + res.emplace_back(pos, limit); + } + + return res; +} +#endif + struct llama_mmap::impl { #ifdef _POSIX_MAPPED_FILES std::vector> mapped_fragments; - impl(struct llama_file * file, size_t prefetch, bool numa) { + impl(struct llama_file * file, size_t prefetch, bool numa, const llama_mmap::ranges & lazy_ranges) { size = file->size(); int fd = file->file_id(); int flags = MAP_SHARED; @@ -452,19 +475,35 @@ struct llama_mmap::impl { LLAMA_LOG_WARN("warning: posix_fadvise(.., POSIX_FADV_SEQUENTIAL) failed: %s\n", strerror(errno)); } - if (prefetch) { flags |= MAP_POPULATE; } + // MAP_POPULATE would fault in the lazy ranges too + if (prefetch && lazy_ranges.empty()) { flags |= MAP_POPULATE; } #endif addr = mmap(NULL, file->size(), PROT_READ, flags, fd, 0); if (addr == MAP_FAILED) { throw std::runtime_error(format("mmap failed: %s", strerror(errno))); } + // page-aligned madvise over [beg, end), clamped to the file + auto advise = [&](size_t beg, size_t end, int advice, const char * name) { + const size_t page_size = sysconf(_SC_PAGESIZE); + beg = beg & ~(page_size - 1); + end = std::min((end + page_size - 1) & ~(page_size - 1), file->size()); + if (beg >= end) { + return; + } + if (posix_madvise((char *) addr + beg, end - beg, advice)) { + LLAMA_LOG_WARN("warning: posix_madvise(.., %s) failed: %s\n", name, strerror(errno)); + } + }; + if (prefetch > 0) { - if (posix_madvise(addr, std::min(file->size(), prefetch), POSIX_MADV_WILLNEED)) { - LLAMA_LOG_WARN("warning: posix_madvise(.., POSIX_MADV_WILLNEED) failed: %s\n", - strerror(errno)); + for (const auto & range : ranges_complement(lazy_ranges, std::min(file->size(), prefetch))) { + advise(range.first, range.second, POSIX_MADV_WILLNEED, "POSIX_MADV_WILLNEED"); } } + for (const auto & range : lazy_ranges) { + advise(range.first, range.second, POSIX_MADV_RANDOM, "POSIX_MADV_RANDOM"); + } if (numa) { if (posix_madvise(addr, file->size(), POSIX_MADV_RANDOM)) { LLAMA_LOG_WARN("warning: posix_madvise(.., POSIX_MADV_RANDOM) failed: %s\n", @@ -533,7 +572,7 @@ struct llama_mmap::impl { #elif defined(_WIN32) HANDLE hMapping = nullptr; - impl(struct llama_file * file, size_t prefetch, bool numa) { + impl(struct llama_file * file, size_t prefetch, bool numa, const llama_mmap::ranges & lazy_ranges) { GGML_UNUSED(numa); size = file->size(); @@ -563,10 +602,15 @@ struct llama_mmap::impl { pPrefetchVirtualMemory = (decltype(pPrefetchVirtualMemory))(void *) GetProcAddress(hKernel32, "PrefetchVirtualMemory"); if (pPrefetchVirtualMemory) { - WIN32_MEMORY_RANGE_ENTRY range; - range.VirtualAddress = addr; - range.NumberOfBytes = (SIZE_T) std::min(size, prefetch); - if (!pPrefetchVirtualMemory(GetCurrentProcess(), 1, &range, 0)) { + std::vector entries; + for (const auto & range : ranges_complement(lazy_ranges, std::min(size, prefetch))) { + WIN32_MEMORY_RANGE_ENTRY entry; + entry.VirtualAddress = (char *) addr + range.first; + entry.NumberOfBytes = (SIZE_T) (range.second - range.first); + entries.push_back(entry); + } + if (!entries.empty() && + !pPrefetchVirtualMemory(GetCurrentProcess(), (ULONG_PTR) entries.size(), entries.data(), 0)) { LLAMA_LOG_WARN("warning: PrefetchVirtualMemory failed: %s\n", llama_format_win_err(GetLastError()).c_str()); } @@ -597,10 +641,11 @@ struct llama_mmap::impl { } } #else - impl(struct llama_file * file, size_t prefetch, bool numa) { + impl(struct llama_file * file, size_t prefetch, bool numa, const llama_mmap::ranges & lazy_ranges) { GGML_UNUSED(file); GGML_UNUSED(prefetch); GGML_UNUSED(numa); + GGML_UNUSED(lazy_ranges); throw std::runtime_error("mmap not supported"); } @@ -617,7 +662,8 @@ struct llama_mmap::impl { size_t size; }; -llama_mmap::llama_mmap(struct llama_file * file, size_t prefetch, bool numa) : pimpl(std::make_unique(file, prefetch, numa)) {} +llama_mmap::llama_mmap(struct llama_file * file, size_t prefetch, bool numa, + const ranges & lazy_ranges) : pimpl(std::make_unique(file, prefetch, numa, lazy_ranges)) {} llama_mmap::~llama_mmap() = default; size_t llama_mmap::size() const { return pimpl->size; } diff --git a/src/llama-mmap.h b/src/llama-mmap.h index b7d5c61e95ff..cc28c8a73fa5 100644 --- a/src/llama-mmap.h +++ b/src/llama-mmap.h @@ -2,6 +2,7 @@ #include #include +#include #include #include @@ -41,8 +42,12 @@ struct llama_file { }; struct llama_mmap { + // list of [first, last) byte ranges within a file + using ranges = std::vector>; + llama_mmap(const llama_mmap &) = delete; - llama_mmap(struct llama_file * file, size_t prefetch = (size_t) -1, bool numa = false); + llama_mmap(struct llama_file * file, size_t prefetch = (size_t) -1, bool numa = false, + const ranges & lazy_ranges = {}); ~llama_mmap(); size_t size() const; diff --git a/src/llama-model-loader.cpp b/src/llama-model-loader.cpp index df8313e81976..bf0e652886d5 100644 --- a/src/llama-model-loader.cpp +++ b/src/llama-model-loader.cpp @@ -39,6 +39,9 @@ const char * llama_ftype_name(llama_ftype ftype) { case LLAMA_FTYPE_MOSTLY_BF16: name = LLAMA_FTYPE_PREFIX "BF16"; break; case LLAMA_FTYPE_MOSTLY_Q1_0: name = LLAMA_FTYPE_PREFIX "Q1_0"; break; case LLAMA_FTYPE_MOSTLY_Q2_0: name = LLAMA_FTYPE_PREFIX "Q2_0"; break; + case LLAMA_FTYPE_MOSTLY_PQ2_0: + case LLAMA_FTYPE_MOSTLY_PQ2_0_LEGACY: name = LLAMA_FTYPE_PREFIX "PQ2_0 - 2.13 bpw (PrismML, group 128)"; break; + case LLAMA_FTYPE_MOSTLY_PTQ1_0: name = LLAMA_FTYPE_PREFIX "PTQ1_0 - 1.75 bpw ternary (PrismML, group 128)"; break; case LLAMA_FTYPE_MOSTLY_Q4_0: name = LLAMA_FTYPE_PREFIX "Q4_0"; break; case LLAMA_FTYPE_MOSTLY_Q4_1: name = LLAMA_FTYPE_PREFIX "Q4_1"; break; case LLAMA_FTYPE_MOSTLY_Q5_0: name = LLAMA_FTYPE_PREFIX "Q5_0"; break; @@ -322,8 +325,9 @@ namespace GGUFMeta { (std::is_same::value)); break; case GGUF_TYPE_FLOAT32: GGML_ASSERT((std::is_same::value)); break; case GGUF_TYPE_STRING: GGML_ASSERT((std::is_same::value)); break; + case GGUF_TYPE_UINT64: GGML_ASSERT((std::is_same::value)); break; default: - throw std::runtime_error(format("%s is not a string/float32/uint32/int32 array", key.c_str())); + throw std::runtime_error(format("%s is not a string/float32/uint32/int32/uint64 array", key.c_str())); } if constexpr (std::is_same::value) { @@ -364,8 +368,9 @@ namespace GGUFMeta { (std::is_same::value)); break; case GGUF_TYPE_FLOAT32: GGML_ASSERT((std::is_same::value)); break; case GGUF_TYPE_STRING: GGML_ASSERT((std::is_same::value)); break; + case GGUF_TYPE_UINT64: GGML_ASSERT((std::is_same::value)); break; default: - throw std::runtime_error(format("%s is not a string/float32/uint32/int32 array", key.c_str())); + throw std::runtime_error(format("%s is not a string/float32/uint32/int32/uint64 array", key.c_str())); } if (arr_info.length > N_MAX) { @@ -402,6 +407,9 @@ namespace GGUFMeta { template bool llama_model_loader::get_arr>(enum llm_kv kid, std::array & result, bool required); template bool llama_model_loader::get_arr>(enum llm_kv kid, std::vector & result, bool required); template bool llama_model_loader::get_arr>(enum llm_kv kid, std::array & result, bool required); + template bool llama_model_loader::get_arr>(enum llm_kv kid, std::vector & result, bool required); + template bool llama_model_loader::get_arr>(enum llm_kv kid, std::array & result, bool required); + template bool llama_model_loader::get_arr>(enum llm_kv kid, std::array & result, bool required); template bool llama_model_loader::get_key(const std::string & key, T & result, bool required) { @@ -759,6 +767,8 @@ llama_model_loader::llama_model_loader( case GGML_TYPE_NVFP4: ftype = LLAMA_FTYPE_MOSTLY_NVFP4; break; case GGML_TYPE_Q1_0: ftype = LLAMA_FTYPE_MOSTLY_Q1_0; break; case GGML_TYPE_Q2_0: ftype = LLAMA_FTYPE_MOSTLY_Q2_0; break; + case GGML_TYPE_PQ2_0: ftype = LLAMA_FTYPE_MOSTLY_PQ2_0; break; + case GGML_TYPE_PTQ1_0: ftype = LLAMA_FTYPE_MOSTLY_PTQ1_0; break; default: { LLAMA_LOG_WARN("%s: unknown type %s\n", __func__, ggml_type_name(type_max)); @@ -857,7 +867,11 @@ struct ggml_tensor * llama_model_loader::require_tensor_meta(const std::string & return tensor; } -const struct ggml_tensor * llama_model_loader::check_tensor_dims(const std::string & name, const std::vector & ne, bool required) const { +const struct ggml_tensor * llama_model_loader::check_tensor_dims( + const std::string & name, + const std::vector & ne, + bool required, + bool allow_reshape) const { const struct ggml_tensor * cur = get_tensor_meta(name.c_str()); if (cur == NULL) { @@ -867,21 +881,33 @@ const struct ggml_tensor * llama_model_loader::check_tensor_dims(const std::stri throw std::runtime_error(format("%s: tensor '%s' not found", __func__, name.c_str())); } - { - bool is_ok = true; + bool is_ok = true; + + if (allow_reshape) { + // check total number of elements only + const int64_t ncur = ggml_nelements(cur); + int64_t nexp = 1; + for (size_t i = 0; i < ne.size(); ++i) { + nexp *= ne[i]; + } + if (ncur != nexp) { + is_ok = false; + } + } else { for (size_t i = 0; i < GGML_MAX_DIMS; ++i) { if ((i < ne.size() && ne[i] != cur->ne[i]) || (i >= ne.size() && cur->ne[i] != 1)) { is_ok = false; break; } } - if (!is_ok) { - throw std::runtime_error( - format("%s: tensor '%s' has wrong shape; expected %s, got %s", - __func__, name.c_str(), - llama_format_tensor_shape(ne).c_str(), - llama_format_tensor_shape(cur).c_str())); - } + } + + if (!is_ok) { + throw std::runtime_error( + format("%s: tensor '%s' has wrong shape; expected %s, got %s", + __func__, name.c_str(), + llama_format_tensor_shape(ne).c_str(), + llama_format_tensor_shape(cur).c_str())); } return cur; @@ -1040,11 +1066,52 @@ static ggml_backend_buffer_type_t select_weight_buft(const llama_hparams & hpara return nullptr; } +ggml_backend_buffer_type_t llama_model_loader::lazy_read::buft() { + auto * cpu_dev = ggml_backend_dev_by_type(GGML_BACKEND_DEVICE_TYPE_CPU); + if (!cpu_dev) { + throw std::runtime_error("no CPU backend found"); + } + return ggml_backend_dev_buffer_type(cpu_dev); +} + +bool llama_model_loader::lazy_read::add(const std::string & name, const ggml_tensor * t, const llama_tensor_weight * w) { + if (mode == LLAMA_LAZY_MODE_OFF) { + return false; + } + + // do not lazy-read small tensors, it has significant overhead and is not worth it + constexpr size_t auto_min_size = 4ull * 1024 * 1024 * 1024; + if (mode != LLAMA_LAZY_MODE_ON && ggml_nbytes(t) <= auto_min_size) { + return false; + } + + if (!llama_mmap::SUPPORTED) { + LLAMA_LOG_WARN("%s: mmap is not available, so tensor %s (size = %zu MiB) is loaded into RAM in full\n", + __func__, name.c_str(), ggml_nbytes(t)/1024/1024); + return false; + } + + if (w) { + ranges[w->idx].emplace_back(w->offs, w->offs + ggml_nbytes(t)); + tensors.insert(name); + + LLAMA_LOG_INFO("%s: tensor %s (size = %zu MiB) lazy read enabled\n", + __func__, name.c_str(), ggml_nbytes(t)/1024/1024); + } + + return true; +} + struct ggml_tensor * llama_model_loader::create_tensor( const llama_hparams & hparams, const buft_list_t * buft_list_cpu, const buft_list_t * buft_list_input, const buft_list_t * buft_list_output, const buft_list_t * buft_list_layer, const LLM_TN_IMPL & tn, const std::initializer_list & ne, int flags) { + // set below, before buft_for_tensor() runs + bool is_lazy = false; + auto ctx_for_buft = [&](ggml_backend_buffer_type_t buft) -> ggml_context * { - auto it = ctx_map.find(buft); + const ctx_key key { buft, is_lazy }; + + auto it = ctx_map.find(key); if (it == ctx_map.end()) { // one ggml context per buffer type int max_n_tensors = n_tensors; @@ -1066,7 +1133,7 @@ struct ggml_tensor * llama_model_loader::create_tensor( throw std::runtime_error(format("failed to create ggml context")); } - ctx_map.emplace(buft, ctx); + ctx_map.emplace(key, ctx); return ctx; } @@ -1131,6 +1198,10 @@ struct ggml_tensor * llama_model_loader::create_tensor( } } + if (is_lazy) { + return lazy_read::buft(); + } + // select the buffer type for this tensor const buft_list_t * buft_list; switch (info.layer) { @@ -1246,11 +1317,30 @@ struct ggml_tensor * llama_model_loader::create_tensor( return ret; } - ggml_tensor * t_meta = get_tensor_meta(tn.str().c_str()); - ggml_backend_buffer_type_t buft = buft_for_tensor(t_meta); + LLAMA_LOG_DEBUG("%s: loading tensor %s\n", __func__, tn.str().c_str()); + const struct ggml_tensor * cur = check_tensor_dims(tn.str(), ne, !(flags & TENSOR_NOT_REQUIRED), flags & TENSOR_ALLOW_RESHAPE); + if (cur == NULL) { + return NULL; + } + + if (flags & TENSOR_READ_LAZY) { + // the decision must not depend on the load mode, or the memory-fit pass (no_alloc, no mmap) + is_lazy = lazy.add(tn.str(), cur, no_alloc ? nullptr : &require_weight(tn.str().c_str())); + } + + ggml_tensor t_meta = *cur; + if (flags & TENSOR_ALLOW_RESHAPE) { + for (size_t dim = 0; dim < GGML_MAX_DIMS; dim++) { + t_meta.ne[dim] = dim < ne.size() ? ne.begin()[dim] : 1; + t_meta.nb[dim] = dim == 0 ? ggml_type_size(t_meta.type) : t_meta.ne[dim-1]*t_meta.nb[dim-1]; + } + } + + ggml_backend_buffer_type_t buft = buft_for_tensor(&t_meta); if (buft == nullptr) { - return nullptr; // return type is ggml_tensor * + return nullptr; } + ggml_context * ctx = ctx_for_buft(buft); // if duplicated, check if the original tensor was allocated in the same buffer type context and avoid creating a new one @@ -1261,20 +1351,13 @@ struct ggml_tensor * llama_model_loader::create_tensor( } } - LLAMA_LOG_DEBUG("%s: loading tensor %s\n", __func__, tn.str().c_str()); - const struct ggml_tensor * cur = check_tensor_dims(tn.str(), ne, !(flags & TENSOR_NOT_REQUIRED)); - - if (cur == NULL) { - return NULL; - } - const bool duplicated = flags & TENSOR_DUPLICATED; - struct ggml_tensor * tensor = ggml_dup_tensor(ctx, cur); - ggml_set_name(tensor, ggml_get_name(cur)); + struct ggml_tensor * tensor = ggml_dup_tensor(ctx, &t_meta); + ggml_set_name(tensor, ggml_get_name(&t_meta)); if (duplicated) { - size_data += ggml_nbytes(cur); + size_data += ggml_nbytes(&t_meta); } else { n_created++; } @@ -1282,34 +1365,6 @@ struct ggml_tensor * llama_model_loader::create_tensor( return tensor; } -struct ggml_tensor * llama_model_loader::create_tensor_as_view(struct ggml_context * ctx, struct ggml_tensor * base, const std::string & name, const std::initializer_list & ne, size_t offset, bool required) { - const struct ggml_tensor * cur = check_tensor_dims(name, ne, required); - - if (cur == NULL) { - return NULL; - } - - if (cur->type != base->type) { - throw std::runtime_error(format("%s: tensor '%s' has wrong type; expected %s, got %s", __func__, name.c_str(), ggml_type_name(base->type), ggml_type_name(cur->type))); - } - - std::array dims; - for (size_t i = 0; i < GGML_MAX_DIMS; ++i) { - dims[i] = i < ne.size() ? ne.begin()[i] : 1; - } - - struct ggml_tensor * tensor = ggml_view_4d(ctx, base, - dims[0], dims[1], dims[2], dims[3], - cur->nb[1], cur->nb[2], cur->nb[3], - offset); - - ggml_set_name(tensor, name.c_str()); - - n_created++; - - return tensor; -} - void llama_model_loader::done_getting_tensors(bool partial) const { if (n_created > n_tensors) { throw std::runtime_error(format("%s: too many tensors created; expected %d, got %d", __func__, n_tensors, n_created)); @@ -1329,10 +1384,13 @@ void llama_model_loader::done_getting_tensors(bool partial) const { } void llama_model_loader::init_mappings(bool prefetch, llama_mlocks * mlock_mmaps) { - if (use_mmap) { + // note: read_lazy also requires mmap; this condition make sure it's usable even when --load-mode is not set to mmap + if (use_mmap || lazy.any()) { mappings.reserve(files.size()); mmaps_used.reserve(files.size()); - for (const auto & file : files) { + for (uint32_t idx = 0; idx < files.size(); idx++) { + const auto & file = files[idx]; + bool is_numa = false; auto * dev = ggml_backend_dev_by_type(GGML_BACKEND_DEVICE_TYPE_CPU); @@ -1344,7 +1402,10 @@ void llama_model_loader::init_mappings(bool prefetch, llama_mlocks * mlock_mmaps } } - std::unique_ptr mapping = std::make_unique(file.get(), prefetch ? -1 : 0, is_numa); + const size_t prefetch_size = prefetch && use_mmap ? -1 : 0; + + std::unique_ptr mapping = std::make_unique(file.get(), prefetch_size, is_numa, + lazy.for_file(idx)); mmaps_used.emplace_back(mapping->size(), 0); if (mlock_mmaps) { std::unique_ptr mlock_mmap(new llama_mlock()); @@ -1531,7 +1592,9 @@ bool llama_model_loader::load_all_data( size_t n_size = ggml_nbytes(cur); - if (use_mmap) { + const bool from_mapping = use_mmap || lazy.has(cur); + + if (from_mapping) { const auto & mapping = mappings.at(weight->idx); ggml_backend_buffer_t buf_mmap = nullptr; if (bufs.count(weight->idx)) { @@ -1548,7 +1611,9 @@ bool llama_model_loader::load_all_data( GGML_ASSERT(buf_mmap || cur->data); // either we have a buffer to allocate the tensor in, or it is already allocated if (buf_mmap && cur->data == nullptr) { ggml_backend_tensor_alloc(buf_mmap, cur, data); - if (lmlocks) { + + // locking a lazy tensor would fault all of it in, which is what lazy avoids + if (lmlocks && !lazy.has(cur)) { const auto & lmlock = lmlocks->at(weight->idx); lmlock->grow_to(weight->offs + n_size); } diff --git a/src/llama-model-loader.h b/src/llama-model-loader.h index 7ad380782291..8198202dad3e 100644 --- a/src/llama-model-loader.h +++ b/src/llama-model-loader.h @@ -12,6 +12,7 @@ #include #include #include +#include #include #include @@ -67,6 +68,8 @@ struct llama_model_loader { static const int TENSOR_DUPLICATED = 1 << 1; static const int TENSOR_SKIP = 1 << 2; static const int TENSOR_SKIP_IF_VIRTUAL = 1 << 3; + static const int TENSOR_ALLOW_RESHAPE = 1 << 4; + static const int TENSOR_READ_LAZY = 1 << 5; // read rows on demand instead of loading whole tensor; requires mmap for now int n_kv = 0; int n_tensors = 0; @@ -81,6 +84,39 @@ struct llama_model_loader { bool no_alloc; bool load_mtp; + // handle TENSOR_READ_LAZY + // use case: keep PLE / engrams embd tensors on disk, read them on demand + struct lazy_read { + // set by the caller before the create_tensor() calls + enum llama_lazy_mode mode = LLAMA_LAZY_MODE_OFF; + + // decide whether this tensor is read lazily + // pass w to also record it, or nullptr to only ask + bool add(const std::string & name, const ggml_tensor * t, const llama_tensor_weight * w); + + bool any() const { + return !ranges.empty(); + } + + bool has(const ggml_tensor * t) const { + return tensors.count(ggml_get_name(t)) > 0; + } + + const llama_mmap::ranges & for_file(uint32_t idx) const { + static const llama_mmap::ranges none; + + const auto it = ranges.find(idx); + return it == ranges.end() ? none : it->second; + } + + // lazy tensors are gathered on the host, so no offload setting applies to them + static ggml_backend_buffer_type_t buft(); + + private: + std::map ranges; + std::set tensors; + } lazy; + llama_files files; llama_ftype ftype; llama_fver fver; @@ -111,7 +147,22 @@ struct llama_model_loader { } }; - std::map ctx_map; + // lazy tensors need dedicated context + struct ctx_key { + ggml_backend_buffer_type_t buft; + bool lazy; + }; + + struct ctx_key_comparator { + bool operator()(const ctx_key & lhs, const ctx_key & rhs) const { + if (lhs.lazy != rhs.lazy) { + return lhs.lazy < rhs.lazy; + } + return strcmp(ggml_backend_buft_name(lhs.buft), ggml_backend_buft_name(rhs.buft)) < 0; + } + }; + + std::map ctx_map; // track tensors that had to be moved for debugging: size_t n_tensors_moved = 0; @@ -177,14 +228,16 @@ struct llama_model_loader { struct ggml_tensor * require_tensor_meta(const std::string & name) const; - const struct ggml_tensor * check_tensor_dims(const std::string & name, const std::vector & ne, bool required) const; + const struct ggml_tensor * check_tensor_dims( + const std::string & name, + const std::vector & ne, + bool required, + bool allow_reshape) const; struct ggml_tensor * create_tensor( const llama_hparams & hparams, const buft_list_t * buft_list_cpu, const buft_list_t * buft_list_input, const buft_list_t * buft_list_output, const buft_list_t * buft_list_layer, const LLM_TN_IMPL & tn, const std::initializer_list & ne, int flags); - struct ggml_tensor * create_tensor_as_view(struct ggml_context * ctx, struct ggml_tensor * base, const std::string & name, const std::initializer_list & ne, size_t offset, bool required = true); - void done_getting_tensors(bool partial = false) const; void init_mappings(bool prefetch = true, llama_mlocks * mlock_mmaps = nullptr); diff --git a/src/llama-model-saver.cpp b/src/llama-model-saver.cpp index 3812c594e795..c2898b5ba525 100644 --- a/src/llama-model-saver.cpp +++ b/src/llama-model-saver.cpp @@ -57,6 +57,10 @@ void llama_model_saver::add_kv(const enum llm_kv key, const int32_t value) { gguf_set_val_i32(gguf_ctx, llm_kv(key).c_str(), value); } +void llama_model_saver::add_kv(const enum llm_kv key, const uint64_t value) { + gguf_set_val_u64(gguf_ctx, llm_kv(key).c_str(), value); +} + void llama_model_saver::add_kv(const enum llm_kv key, const float value) { gguf_set_val_f32(gguf_ctx, llm_kv(key).c_str(), value); } @@ -110,6 +114,8 @@ void llama_model_saver::add_kv(const enum llm_kv key, const Container & value, c gguf_set_arr_data(gguf_ctx, llm_kv(key).c_str(), GGUF_TYPE_BOOL, value.data(), n_values); } else if (std::is_same::value) { gguf_set_arr_data(gguf_ctx, llm_kv(key).c_str(), GGUF_TYPE_INT32, value.data(), n_values); + } else if (std::is_same::value) { + gguf_set_arr_data(gguf_ctx, llm_kv(key).c_str(), GGUF_TYPE_UINT64, value.data(), n_values); } else if (std::is_same::value) { gguf_set_arr_data(gguf_ctx, llm_kv(key).c_str(), GGUF_TYPE_FLOAT32, value.data(), n_values); } else if (std::is_same::value) { @@ -120,6 +126,8 @@ void llama_model_saver::add_kv(const enum llm_kv key, const Container & value, c } // instantiate for external usage: template void llama_model_saver::add_kv>(const enum llm_kv, const std::vector &, const bool); +template void llama_model_saver::add_kv>(const enum llm_kv, const std::vector &, const bool); +template void llama_model_saver::add_kv>(const enum llm_kv, const std::vector &, const bool); void llama_model_saver::add_kv(const enum llm_kv key, const std::vector & value) { std::vector tmp(value.size()); @@ -285,6 +293,47 @@ void llama_model_saver::add_kv_from_model() { add_kv(LLM_KV_ATTENTION_INDEXER_LOCAL_BLOCKS, hparams.indexer_local_blocks); add_kv(LLM_KV_ATTENTION_INDEXER_TYPES, hparams.is_indexer_full_impl, true); add_kv(LLM_KV_ATTENTION_RECURRENT_LAYERS, hparams.is_recr_impl, true); + add_kv(LLM_KV_ATTENTION_OUTPUT_GROUP_COUNT, hparams.dsv4_o_group_count); + add_kv(LLM_KV_ATTENTION_OUTPUT_LORA_RANK, hparams.dsv4_o_lora_rank); + add_kv(LLM_KV_ATTENTION_COMPRESS_ROPE_FREQ_BASE, hparams.dsv4_compress_rope_base); + if (model->arch == LLM_ARCH_DEEPSEEK4 || hparams.dsv4_hc_mult > 0) { + // the loader requires one compress ratio per layer, including nextn layers + const std::vector compress_ratios( + hparams.dsv4_compress_ratios.begin(), hparams.dsv4_compress_ratios.begin() + hparams.n_layer_all); + add_kv(LLM_KV_ATTENTION_COMPRESS_RATIOS, compress_ratios); + } else { + add_kv(LLM_KV_ATTENTION_COMPRESS_RATIOS, hparams.dsv4_compress_ratios, true); + } + add_kv(LLM_KV_HYPER_CONNECTION_COUNT, hparams.dsv4_hc_mult); + add_kv(LLM_KV_HYPER_CONNECTION_SINKHORN_ITERATIONS, hparams.dsv4_hc_sinkhorn_iters); + add_kv(LLM_KV_HYPER_CONNECTION_EPSILON, hparams.dsv4_hc_eps); + add_kv(LLM_KV_HASH_LAYER_COUNT, hparams.dsv4_hash_layer_count); + add_kv(LLM_KV_HYPER_CONNECTION_LOW_RANK, hparams.hc_low_rank); + + // the PLE group only means anything whole: write all of it or none + if (hparams.ple_n_heads > 0) { + std::vector ple_layers; + for (uint32_t il = 0; il < hparams.n_layer_all; ++il) { + if (hparams.is_ple_impl[il]) { + ple_layers.push_back(il); + } + } + add_kv(LLM_KV_PLE_LAYERS, ple_layers); + add_kv(LLM_KV_PLE_NGRAM_SIZE, hparams.ple_ngram_size); + add_kv(LLM_KV_PLE_HEADS_PER_NGRAM, hparams.ple_heads_per_ngram); + add_kv(LLM_KV_PLE_CONV_KERNEL, hparams.ple_conv_kernel); + add_kv(LLM_KV_PLE_EOS_TOKEN_ID, hparams.ple_eos_token_id); + add_kv(LLM_KV_EMBEDDING_LENGTH_PER_LAYER, hparams.ple_head_dim); + add_kv(LLM_KV_PLE_LAYER_MULTIPLIERS, std::vector( + hparams.ple_layer_multipliers.begin(), + hparams.ple_layer_multipliers.begin() + hparams.ple_ngram_size)); + add_kv(LLM_KV_PLE_HEAD_OFFSETS, std::vector( + hparams.ple_head_offsets.begin(), + hparams.ple_head_offsets.begin() + hparams.ple_n_heads)); + add_kv(LLM_KV_PLE_HEAD_VOCAB_SIZES, std::vector( + hparams.ple_head_vocab_sizes.begin(), + hparams.ple_head_vocab_sizes.begin() + hparams.ple_n_heads)); + } const float rope_scaling_factor = hparams.rope_freq_scale_train == 1.0f ? 0.0f : 1.0f/hparams.rope_freq_scale_train; @@ -407,6 +456,19 @@ void llama_model_saver::add_tensors_from_model() { add_tensor(model->cls_out); add_tensor(model->cls_out_b); add_tensor(model->cls_norm); + add_tensor(model->zaya_input_hs_scale); + add_tensor(model->zaya_input_hs_bias); + add_tensor(model->zaya_res_scale_hs); + add_tensor(model->zaya_res_scale_hs_b); + add_tensor(model->zaya_res_scale_res); + add_tensor(model->zaya_res_scale_res_b); + add_tensor(model->hc_head_fn); + add_tensor(model->hc_head_base); + add_tensor(model->hc_head_scale); + add_tensor(model->per_layer_tok_embd); + add_tensor(model->hc_head_norm); + add_tensor(model->hc_head_down); + add_tensor(model->hc_head_up); for (const struct llama_layer & layer : model->layers) { for (size_t i = 0; i < sizeof(layer)/sizeof(struct ggml_tensor *); ++i) { diff --git a/src/llama-model-saver.h b/src/llama-model-saver.h index 36a715e2b6bf..95e19e666e7f 100644 --- a/src/llama-model-saver.h +++ b/src/llama-model-saver.h @@ -21,6 +21,7 @@ struct llama_model_saver { void add_kv(enum llm_kv key, uint32_t value); void add_kv(enum llm_kv key, int32_t value); + void add_kv(enum llm_kv key, uint64_t value); void add_kv(enum llm_kv key, float value); void add_kv(enum llm_kv key, bool value); void add_kv(enum llm_kv key, const char * value); diff --git a/src/llama-model.cpp b/src/llama-model.cpp index 91c2fc9c008d..b44bd9b1fe8c 100644 --- a/src/llama-model.cpp +++ b/src/llama-model.cpp @@ -14,6 +14,7 @@ #include "llama-kv-cache-dsv4.h" #include "llama-memory-hybrid.h" #include "llama-memory-hybrid-iswa.h" +#include "llama-memory-hybrid-idx.h" #include "llama-memory-recurrent.h" #include "llama.h" @@ -79,8 +80,14 @@ static llama_model * llama_model_mapping(llm_arch arch, const llama_model_params return new llama_model_eurobert(params); case LLM_ARCH_BLOOM: return new llama_model_bloom(params); + case LLM_ARCH_GPTNEO: + return new llama_model_gptneo(params); + case LLM_ARCH_CODEGEN: + return new llama_model_codegen(params); case LLM_ARCH_MPT: return new llama_model_mpt(params); + case LLM_ARCH_OPT: + return new llama_model_opt(params); case LLM_ARCH_STABLELM: return new llama_model_stablelm(params); case LLM_ARCH_MELLUM: @@ -175,6 +182,8 @@ static llama_model * llama_model_mapping(llm_arch arch, const llama_model_params return new llama_model_openelm(params); case LLM_ARCH_GPTNEOX: return new llama_model_gptneox(params); + case LLM_ARCH_GPTJ: + return new llama_model_gptj(params); case LLM_ARCH_ARCTIC: return new llama_model_arctic(params); case LLM_ARCH_DEEPSEEK: @@ -275,6 +284,8 @@ static llama_model * llama_model_mapping(llm_arch arch, const llama_model_params return new llama_model_openai_moe(params); case LLM_ARCH_FALCON_H1: return new llama_model_falcon_h1(params); + case LLM_ARCH_ZAYA: + return new llama_model_zaya(params); case LLM_ARCH_LFM2: return new llama_model_lfm2(params); case LLM_ARCH_LFM2MOE: @@ -299,6 +310,8 @@ static llama_model * llama_model_mapping(llm_arch arch, const llama_model_params return new llama_model_qwen35(params); case LLM_ARCH_QWEN35MOE: return new llama_model_qwen35moe(params); + case LLM_ARCH_QWEN4EXP: + return new llama_model_qwen4exp(params); case LLM_ARCH_MISTRAL3: return new llama_model_mistral3(params); case LLM_ARCH_EAGLE3: @@ -352,6 +365,7 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str static const std::regex pattern_qkv_bias ("blk\\.\\d*\\.attn_qkv.bias"); static const std::regex pattern_qk_norm ("blk\\.\\d*\\.attn_(q|k)_norm\\.weight"); static const std::regex pattern_kv_cache ("cache_(k|v)_l\\d*"); + static const std::regex pattern_idx_cache ("cache_idx_(k|v)_l\\d*"); static const std::regex pattern_attn_sinks ("blk\\.\\d*\\.attn_sinks.weight"); static const std::regex pattern_attn_out_weight ("blk\\.\\d*\\.attn_output.weight"); static const std::regex pattern_attn_out_bias ("blk\\.\\d*\\.attn_output.bias"); @@ -363,6 +377,7 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str static const std::regex pattern_ssm_beta ("blk\\.\\d*\\.ssm_beta.weight"); static const std::regex pattern_ssm_beta_alpha ("blk\\.\\d*\\.ssm_ba.weight"); static const std::regex pattern_r_cache ("cache_r_l\\d*"); + static const std::regex pattern_ple_r_cache ("cache_ple_r_l\\d*"); static const std::regex pattern_s_cache ("cache_s_l\\d*"); static const std::regex pattern_ssm_conv1d ("blk\\.\\d*\\.ssm_conv1d.weight"); static const std::regex pattern_ssm_out_weight ("blk\\.\\d*\\.ssm_out.weight"); @@ -431,6 +446,16 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str }; auto get_tensor_config = [&]() -> tensor_config { + // the qsa indexer has one key head and its projections are mirrored, so its cache cannot be split + if (std::regex_match(tensor_name, pattern_idx_cache)) { + return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_MIRRORED); + } + + // the PLE table is model-level and its conv is mirrored, so every device runs the whole conv and needs the whole history + if (std::regex_match(tensor_name, pattern_ple_r_cache)) { + return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_MIRRORED); + } + // standard attention if (std::regex_match(tensor_name, pattern_q_weight) || std::regex_match(tensor_name, pattern_kv_weight)) { return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_1, "attn_output.weight", "ssm_out.weight"); @@ -512,7 +537,8 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str }; auto get_split_segments = [&](int axis, uint32_t il) -> std::vector> { - if (ud->model->arch == LLM_ARCH_QWEN3NEXT || ud->model->arch == LLM_ARCH_QWEN35 || ud->model->arch == LLM_ARCH_QWEN35MOE) { + if (ud->model->arch == LLM_ARCH_QWEN3NEXT || ud->model->arch == LLM_ARCH_QWEN35 || ud->model->arch == LLM_ARCH_QWEN35MOE || + ud->model->arch == LLM_ARCH_QWEN4EXP) { const int64_t head_k_dim = hparams.ssm_d_state; const int64_t head_v_dim = hparams.ssm_d_state; const int64_t n_k_heads = hparams.ssm_n_group; @@ -625,7 +651,8 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str if (std::regex_match(tensor_name, pattern_q_weight) || std::regex_match(tensor_name, pattern_q_bias)) { GGML_ASSERT(segments.size() == 1); // some models have Q gate tensors, for those cases the granularity needs to be doubled: - if (ud->model->arch == LLM_ARCH_QWEN3NEXT || ud->model->arch == LLM_ARCH_QWEN35 || ud->model->arch == LLM_ARCH_QWEN35MOE) { + if (ud->model->arch == LLM_ARCH_QWEN3NEXT || ud->model->arch == LLM_ARCH_QWEN35 || ud->model->arch == LLM_ARCH_QWEN35MOE || + ud->model->arch == LLM_ARCH_QWEN4EXP) { return {std::lcm(2*n_embd_q, blck_size_perf)}; } return {granularity_q}; @@ -815,6 +842,7 @@ const char * llm_type_name(llm_type type) { case LLM_TYPE_35B_A3B: return "35B.A3B"; case LLM_TYPE_48B_A3B: return "48B.A3B"; case LLM_TYPE_80B_A3B: return "80B.A3B"; + case LLM_TYPE_A3B: return "A3B"; case LLM_TYPE_100B_A6B: return "100B.A6B"; case LLM_TYPE_102B_A12B: return "102B.A12B"; case LLM_TYPE_106B_A12B: return "106B.A12B"; @@ -1074,6 +1102,8 @@ void llama_model_base::load_hparams(llama_model_loader & ml) { gguf_kv.emplace(name, value); } + hadamard.load_keys(ml, arch); + // get general kv ml.get_key(LLM_KV_GENERAL_NAME, name, false); @@ -1537,7 +1567,8 @@ bool llama_model_base::load_tensors(llama_model_loader & ml) { const size_t n_max_backend_buffer = ml.ctx_map.size() * ml.files.size(); pimpl->ctxs_bufs.reserve(n_max_backend_buffer); - for (auto & [buft, ctx_ptr] : ml.ctx_map) { + for (auto & [ctx_key, ctx_ptr] : ml.ctx_map) { + ggml_backend_buffer_type_t buft = ctx_key.buft; ggml_context * ctx = ctx_ptr.get(); // skip contexts without tensors @@ -1563,7 +1594,11 @@ bool llama_model_base::load_tensors(llama_model_loader & ml) { bool is_default_buft = buft == ggml_backend_dev_buffer_type(dev); std::vector bufs; - if (ml.use_mmap && use_mmap_buffer && buffer_from_host_ptr_supported && is_default_buft) { + + // a lazy context is mapped whatever the load mode, but the memory-fit pass maps nothing + const bool is_lazy_mapped = ctx_key.lazy && !ml.no_alloc; + + if ((ml.use_mmap || is_lazy_mapped) && use_mmap_buffer && buffer_from_host_ptr_supported && is_default_buft) { GGML_ASSERT(!ml.no_alloc); for (uint32_t idx = 0; idx < ml.files.size(); idx++) { // only the mmap region containing the tensors in the model is mapped to the backend buffer @@ -1661,6 +1696,8 @@ bool llama_model_base::load_tensors(llama_model_loader & ml) { } } + hadamard.setup(*this); + return true; } @@ -2146,7 +2183,7 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, // attention KV cache for the MTP context instead of the hybrid wrapper. const bool mtp_on_hybrid_qwen35 = params.ctx_type == LLAMA_CONTEXT_TYPE_MTP && - (arch == LLM_ARCH_QWEN35 || arch == LLM_ARCH_QWEN35MOE); + (arch == LLM_ARCH_QWEN35 || arch == LLM_ARCH_QWEN35MOE || arch == LLM_ARCH_QWEN4EXP); if (llm_arch_is_recurrent(arch)) { res = new llama_memory_recurrent( @@ -2163,9 +2200,18 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, // layer filters, so pick the right one here llama_memory_hybrid::layer_filter_cb filter_attn = nullptr; llama_memory_hybrid::layer_filter_cb filter_recr = nullptr; + // only the sparse-attention architectures use llama_memory_hybrid_idx + // a null filter_idx means the GGUF has no indexer tensors + llama_memory_hybrid::layer_filter_cb filter_idx = nullptr; + const bool needs_mem_idx = (arch == LLM_ARCH_QWEN4EXP); if (arch == LLM_ARCH_FALCON_H1) { filter_attn = [&](uint32_t) { return true; }; filter_recr = [&](uint32_t) { return true; }; + } else if (arch == LLM_ARCH_ZAYA) { + // every ZAYA layer runs CCA attention (KV cache) and carries a + // recurrent conv / previous-hidden-state row + filter_attn = [&](uint32_t) { return true; }; + filter_recr = [&](uint32_t) { return true; }; } else if (arch == LLM_ARCH_NEMOTRON_H || arch == LLM_ARCH_NEMOTRON_H_MOE) { filter_attn = [&](uint32_t il) { return !hparams.is_recr(il) && hparams.n_ff(il) == 0; @@ -2173,13 +2219,20 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, filter_recr = [&](uint32_t il) { return hparams.is_recr(il) && hparams.n_ff(il) == 0; }; - } else if (arch == LLM_ARCH_QWEN35 || arch == LLM_ARCH_QWEN35MOE) { + } else if (arch == LLM_ARCH_QWEN35 || arch == LLM_ARCH_QWEN35MOE || arch == LLM_ARCH_QWEN4EXP) { filter_attn = [&](uint32_t il) { return il < hparams.n_layer() && !hparams.is_recr(il); }; filter_recr = [&](uint32_t il) { return il < hparams.n_layer() && hparams.is_recr(il); }; + + if (arch == LLM_ARCH_QWEN4EXP && hparams.indexer_head_size > 0) { + // QSA runs on the dense-attention layers only + filter_idx = [&](uint32_t il) { + return il < hparams.n_layer() && !hparams.is_recr(il); + }; + } } if (hparams.swa_type != LLAMA_SWA_TYPE_NONE) { @@ -2202,6 +2255,27 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, /* unified */ cparams.kv_unified, /* filter_attn */ std::move(filter_attn), /* filter_recr */ std::move(filter_recr)); + } else if (needs_mem_idx) { + // sparse attention over a per-token indexer cache, in its own memory type + res = new llama_memory_hybrid_idx( + /* model */ *this, + /* attn_type_k */ params.type_k, + /* attn_type_v */ params.type_v, + /* attn_v_trans */ !cparams.flash_attn, + /* attn_kv_size */ cparams.n_ctx_seq, + /* attn_n_pad */ 1, + /* attn_n_swa */ hparams.n_swa, + /* attn_swa_type */ hparams.swa_type, + /* recurrent_type_k */ GGML_TYPE_F32, + /* recurrent_type_v */ GGML_TYPE_F32, + /* recurrent_kv_size */ std::max((uint32_t) 1, cparams.n_seq_max), + /* n_seq_max */ cparams.n_seq_max, + /* n_rs_seq */ cparams.n_rs_seq, + /* offload */ cparams.offload_kqv, + /* unified */ cparams.kv_unified, + /* filter_attn */ std::move(filter_attn), + /* filter_recr */ std::move(filter_recr), + /* filter_idx */ std::move(filter_idx)); } else { res = new llama_memory_hybrid( /* model */ *this, @@ -2380,6 +2454,7 @@ llama_model_params llama_model_default_params() { /*.n_gpu_layers =*/ -1, /*.split_mode =*/ LLAMA_SPLIT_MODE_LAYER, /*.load_mode =*/ LLAMA_LOAD_MODE_MMAP, + /*.lazy_mode =*/ LLAMA_LAZY_MODE_AUTO, /*.main_gpu =*/ 0, /*.tensor_split =*/ nullptr, /*.progress_callback =*/ nullptr, @@ -2487,7 +2562,8 @@ llama_rope_type llama_model_rope_type(const llama_model * model) { // these models do not use RoPE case LLM_ARCH_CLIP: case LLM_ARCH_GPT2: - case LLM_ARCH_GPTJ: + case LLM_ARCH_GPTNEO: + case LLM_ARCH_OPT: case LLM_ARCH_MPT: case LLM_ARCH_REFACT: case LLM_ARCH_BLOOM: @@ -2509,6 +2585,8 @@ llama_rope_type llama_model_rope_type(const llama_model * model) { return LLAMA_ROPE_TYPE_NONE; // use what we call a normal RoPE, operating on pairs of consecutive head values + case LLM_ARCH_CODEGEN: + case LLM_ARCH_GPTJ: case LLM_ARCH_LLAMA: case LLM_ARCH_LLADA: case LLM_ARCH_LLAMA4: @@ -2616,6 +2694,7 @@ llama_rope_type llama_model_rope_type(const llama_model * model) { case LLM_ARCH_LAGUNA: case LLM_ARCH_QWEN3NEXT: case LLM_ARCH_MIMO2: + case LLM_ARCH_ZAYA: case LLM_ARCH_STEP35: case LLM_ARCH_TALKIE: case LLM_ARCH_MELLUM: @@ -2629,6 +2708,7 @@ llama_rope_type llama_model_rope_type(const llama_model * model) { case LLM_ARCH_QWEN3VLMOE: case LLM_ARCH_QWEN35: case LLM_ARCH_QWEN35MOE: + case LLM_ARCH_QWEN4EXP: return LLAMA_ROPE_TYPE_IMROPE; case LLM_ARCH_GLM4: @@ -2803,7 +2883,9 @@ llama_model_base::llama_model_base(const struct llama_model_params & params) : l TENSOR_DUPLICATED (llama_model_loader::TENSOR_DUPLICATED), TENSOR_NOT_REQUIRED (llama_model_loader::TENSOR_NOT_REQUIRED), TENSOR_SKIP (llama_model_loader::TENSOR_SKIP), - TENSOR_SKIP_IF_VIRTUAL(llama_model_loader::TENSOR_SKIP_IF_VIRTUAL) {} + TENSOR_SKIP_IF_VIRTUAL(llama_model_loader::TENSOR_SKIP_IF_VIRTUAL), + TENSOR_ALLOW_RESHAPE (llama_model_loader::TENSOR_ALLOW_RESHAPE), + TENSOR_READ_LAZY (llama_model_loader::TENSOR_READ_LAZY) {} ggml_tensor * llama_model_base::create_tensor(const LLM_TN_IMPL & tn, const std::initializer_list & ne, int flags) { GGML_ASSERT(ml != nullptr); diff --git a/src/llama-model.h b/src/llama-model.h index 056a6efa59e8..746992dd88bb 100644 --- a/src/llama-model.h +++ b/src/llama-model.h @@ -3,6 +3,7 @@ #include "llama.h" #include "llama-arch.h" #include "llama-graph.h" +#include "llama-hadamard.h" #include "llama-hparams.h" #include "llama-memory.h" #include "llama-vocab.h" @@ -127,6 +128,7 @@ enum llm_type { LLM_TYPE_35B_A3B, // Qwen3.5 LLM_TYPE_48B_A3B, // Kimi Linear LLM_TYPE_80B_A3B, // Qwen3 Next + LLM_TYPE_A3B, // Qwen3.8 Flash Next LLM_TYPE_100B_A6B, LLM_TYPE_102B_A12B, // Solar-Open LLM_TYPE_106B_A12B, // GLM-4.5-Air @@ -221,6 +223,10 @@ struct llama_layer_nextn { struct ggml_tensor * shared_head_head_s = nullptr; struct ggml_tensor * shared_head_head_in_s = nullptr; struct ggml_tensor * shared_head_norm = nullptr; + + struct ggml_tensor * hc_head_norm = nullptr; + struct ggml_tensor * hc_head_down = nullptr; + struct ggml_tensor * hc_head_up = nullptr; }; struct llama_layer { @@ -320,6 +326,43 @@ struct llama_layer { struct ggml_tensor * ffn_latent_down = nullptr; struct ggml_tensor * ffn_latent_up = nullptr; + // ZAYA CCA (Compressed Convolutional Attention) + struct ggml_tensor * cca_conv_grp = nullptr; + struct ggml_tensor * cca_conv_grp_b = nullptr; + struct ggml_tensor * cca_k_scale = nullptr; + struct ggml_tensor * cca_val_proj1 = nullptr; + struct ggml_tensor * cca_val_proj2 = nullptr; + // ZAYA residual scaling (per layer) + struct ggml_tensor * res_scale_hs = nullptr; + struct ggml_tensor * res_scale_hs_b = nullptr; + struct ggml_tensor * res_scale_res = nullptr; + struct ggml_tensor * res_scale_res_b = nullptr; + struct ggml_tensor * res_scale_hs_mlp = nullptr; + struct ggml_tensor * res_scale_hs_mlp_b = nullptr; + struct ggml_tensor * res_scale_res_mlp = nullptr; + struct ggml_tensor * res_scale_res_mlp_b = nullptr; + // ZAYA router (MoE gating) + struct ggml_tensor * zaya_router_mlp2 = nullptr; + struct ggml_tensor * zaya_router_mlp2_b = nullptr; + struct ggml_tensor * zaya_router_mlp4 = nullptr; + struct ggml_tensor * zaya_router_biases = nullptr; + struct ggml_tensor * zaya_router_eda_scale = nullptr; + // ZAYA1-VL: vision-only LoRA, A down to the rank and B back up (experts stacked) + struct ggml_tensor * zaya_vlora_q_a = nullptr; + struct ggml_tensor * zaya_vlora_q_b = nullptr; + struct ggml_tensor * zaya_vlora_k_a = nullptr; + struct ggml_tensor * zaya_vlora_k_b = nullptr; + struct ggml_tensor * zaya_vlora_v1_a = nullptr; + struct ggml_tensor * zaya_vlora_v1_b = nullptr; + struct ggml_tensor * zaya_vlora_v2_a = nullptr; + struct ggml_tensor * zaya_vlora_v2_b = nullptr; + struct ggml_tensor * zaya_vlora_o_a = nullptr; + struct ggml_tensor * zaya_vlora_o_b = nullptr; + struct ggml_tensor * zaya_vlora_up_exps_a = nullptr; + struct ggml_tensor * zaya_vlora_up_exps_b = nullptr; + struct ggml_tensor * zaya_vlora_down_exps_a = nullptr; + struct ggml_tensor * zaya_vlora_down_exps_b = nullptr; + // ff shared expert (shexp) struct ggml_tensor * ffn_gate_inp_shexp = nullptr; struct ggml_tensor * ffn_gate_shexp = nullptr; @@ -523,6 +566,22 @@ struct llama_layer { struct ggml_tensor * index_q_norm = nullptr; struct ggml_tensor * index_k_norm = nullptr; + struct ggml_tensor * hc_attn_norm = nullptr; + struct ggml_tensor * hc_attn_down = nullptr; + struct ggml_tensor * hc_attn_up = nullptr; + struct ggml_tensor * hc_attn_inject = nullptr; + struct ggml_tensor * hc_ffn_norm = nullptr; + struct ggml_tensor * hc_ffn_down = nullptr; + struct ggml_tensor * hc_ffn_up = nullptr; + struct ggml_tensor * hc_ffn_inject = nullptr; + + struct ggml_tensor * ple_key = nullptr; + struct ggml_tensor * ple_value = nullptr; + struct ggml_tensor * ple_norm_key = nullptr; + struct ggml_tensor * ple_norm_query = nullptr; + struct ggml_tensor * ple_norm_conv = nullptr; + struct ggml_tensor * ple_conv1d = nullptr; + // gemma4 layer output scale, reused for talkie embedding skip scale struct ggml_tensor * out_scale = nullptr; @@ -555,6 +614,9 @@ struct llama_model { std::string name = "n/a"; llama_hparams hparams = {}; + + // prism.hadamard folds (llama-hadamard.h): read in load_hparams, bound in load_tensors + llama_hadamard hadamard; llama_vocab vocab; // for classifier models @@ -572,6 +634,14 @@ struct llama_model { struct ggml_tensor * output_b = nullptr; struct ggml_tensor * output_norm_enc = nullptr; + // ZAYA input embedding scaling + final residual scaling + struct ggml_tensor * zaya_input_hs_scale = nullptr; + struct ggml_tensor * zaya_input_hs_bias = nullptr; + struct ggml_tensor * zaya_res_scale_hs = nullptr; + struct ggml_tensor * zaya_res_scale_hs_b = nullptr; + struct ggml_tensor * zaya_res_scale_res = nullptr; + struct ggml_tensor * zaya_res_scale_res_b = nullptr; + // NVFP4 per-tensor scale2, input_scale for LM head struct ggml_tensor * output_s = nullptr; @@ -600,6 +670,10 @@ struct llama_model { struct ggml_tensor * altup_proj = nullptr; struct ggml_tensor * altup_unembd_proj = nullptr; struct ggml_tensor * per_layer_tok_embd = nullptr; + + struct ggml_tensor * hc_head_norm = nullptr; + struct ggml_tensor * hc_head_down = nullptr; + struct ggml_tensor * hc_head_up = nullptr; struct ggml_tensor * per_layer_model_proj = nullptr; struct ggml_tensor * per_layer_proj_norm = nullptr; @@ -719,6 +793,8 @@ struct llama_model_base : public llama_model { const int TENSOR_NOT_REQUIRED; const int TENSOR_SKIP; const int TENSOR_SKIP_IF_VIRTUAL; + const int TENSOR_ALLOW_RESHAPE; + const int TENSOR_READ_LAZY; explicit llama_model_base(const llama_model_params & params); virtual ~llama_model_base() = default; diff --git a/src/llama-quant.cpp b/src/llama-quant.cpp index fd6e787bd7d2..d1429fa50aa1 100644 --- a/src/llama-quant.cpp +++ b/src/llama-quant.cpp @@ -324,6 +324,7 @@ static bool tensor_allows_quantization(const llama_model_quantize_params * param // do not quantize Mamba/Kimi's small conv1d weights // NOTE: can't use LLM_TN here because the layer number is not known quantize &= name.find("ssm_conv1d") == std::string::npos; + quantize &= name.find("cca_conv_grp") == std::string::npos; // zaya: the grouped conv of q/k stays F16 quantize &= name.find("shortconv.conv.weight") == std::string::npos; // do not quantize MiniMax's indexer projection weights, they are tiny @@ -401,6 +402,12 @@ static ggml_type tensor_type_fallback(quantize_state_impl & qs, const ggml_tenso case GGML_TYPE_Q5_K: return_type = GGML_TYPE_Q5_1; break; case GGML_TYPE_Q6_K: return_type = GGML_TYPE_Q8_0; break; default: + if (qk_k <= 32) { + // the target is already a 32-block type, so there is no smaller block to demote to + // the check below turns it into F16, as a 256-block type does when its fallback does not fit + return_type = target_type; + break; + } throw std::runtime_error(format("no tensor type fallback is defined for type %s", ggml_type_name(target_type))); } @@ -676,7 +683,21 @@ static ggml_type llama_tensor_get_type(quantize_state_impl & qs, const llama_mod return tensor->type; } if (params->token_embedding_type < GGML_TYPE_COUNT && tm.category == tensor_category::TOKEN_EMBD) { - return params->token_embedding_type; + // per_layer_token_embd follows --token-embedding-type by default, but it is a large + // separate table, so let an explicit --tensor-type name it + bool named = false; + if (std::strcmp(tensor->name, "per_layer_token_embd.weight") == 0) { + const std::string tensor_name(tensor->name); + for (const auto & [pattern, qtype] : qs.tensor_type_patterns) { + if (std::regex_search(tensor_name, pattern)) { + named = true; + break; + } + } + } + if (!named) { + return params->token_embedding_type; + } } if (params->output_tensor_type < GGML_TYPE_COUNT && tm.category == tensor_category::OUTPUT) { return params->output_tensor_type; diff --git a/src/llama.cpp b/src/llama.cpp index d6e0bbfefa72..ee7595ae22fc 100644 --- a/src/llama.cpp +++ b/src/llama.cpp @@ -307,6 +307,8 @@ static std::pair llama_model_load(struct gguf_context * meta llama_model_loader ml(metadata, set_tensor_data, set_tensor_data_ud, fname, splits, file, params.load_mode, params.check_tensors, params.no_alloc, params.load_mtp, params.kv_overrides, params.tensor_buft_overrides); + ml.lazy.mode = params.lazy_mode; + ml.print_info(); std::unique_ptr model_ptr(llama_model_create(ml, params)); diff --git a/src/models/codegen.cpp b/src/models/codegen.cpp new file mode 100644 index 000000000000..d02fe7bfaea1 --- /dev/null +++ b/src/models/codegen.cpp @@ -0,0 +1,136 @@ +#include "models.h" + +void llama_model_codegen::load_arch_hparams(llama_model_loader & ml) { + ml.get_key(LLM_KV_ATTENTION_LAYERNORM_EPS, hparams.f_norm_eps); + + // CodeGen runs attention and the MLP in parallel off a single layer norm. + hparams.use_par_res = true; +} + +void llama_model_codegen::load_arch_tensors(llama_model_loader &) { + LLAMA_LOAD_LOCALS; + + tok_embd = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, 0); + + output_norm = create_tensor(tn(LLM_TENSOR_OUTPUT_NORM, "weight"), {n_embd}, 0); + output_norm_b = create_tensor(tn(LLM_TENSOR_OUTPUT_NORM, "bias"), {n_embd}, 0); + output = create_tensor(tn(LLM_TENSOR_OUTPUT, "weight"), {n_embd, n_vocab}, TENSOR_NOT_REQUIRED); + output_b = create_tensor(tn(LLM_TENSOR_OUTPUT, "bias"), {n_vocab}, TENSOR_NOT_REQUIRED); // lm_head has a bias + if (output == NULL) { + output = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, TENSOR_DUPLICATED); + } + + for (int i = 0; i < n_layer; ++i) { + auto & layer = layers[i]; + + layer.attn_norm = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "weight", i), {n_embd}, 0); + layer.attn_norm_b = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "bias", i), {n_embd}, 0); + + // fused, bias-free qkv; bias-free output projection + layer.wqkv = create_tensor(tn(LLM_TENSOR_ATTN_QKV, "weight", i), {n_embd, n_embd + 2*n_embd_gqa}, 0); + layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", i), {n_embd, n_embd}, 0); + + // the MLP carries biases + layer.ffn_up = create_tensor(tn(LLM_TENSOR_FFN_UP, "weight", i), {n_embd, n_ff}, 0); + layer.ffn_up_b = create_tensor(tn(LLM_TENSOR_FFN_UP, "bias", i), {n_ff}, 0); + layer.ffn_down = create_tensor(tn(LLM_TENSOR_FFN_DOWN, "weight", i), {n_ff, n_embd}, 0); + layer.ffn_down_b = create_tensor(tn(LLM_TENSOR_FFN_DOWN, "bias", i), {n_embd}, 0); + } +} + +std::unique_ptr llama_model_codegen::build_arch_graph(const llm_graph_params & params) const { + return std::make_unique(*this, params); +} + +llama_model_codegen::graph::graph(const llama_model & model, const llm_graph_params & params) : llm_graph_context(params) { + const int64_t n_embd_head = hparams.n_embd_head_v(); + + GGML_ASSERT(n_embd_head == hparams.n_embd_head_k()); + + ggml_tensor * cur; + ggml_tensor * inpL; + + inpL = build_inp_embd(model.tok_embd); + + ggml_tensor * inp_pos = build_inp_pos(); + + auto * inp_attn = build_attn_inp_kv(); + + ggml_tensor * inp_out_ids = build_inp_out_ids(); + + for (int il = 0; il < n_layer; ++il) { + cur = build_norm(inpL, + model.layers[il].attn_norm, + model.layers[il].attn_norm_b, + LLM_NORM, il); + cb(cur, "attn_norm", il); + + // x = x + attn(ln_1(x)) + mlp(ln_1(x)) + ggml_tensor * attn_out; + { + auto [Qcur, Kcur, Vcur] = build_qkv(model.layers[il], cur, + n_embd_head, n_head, n_head_kv, il); + + // partial RoPE over the first n_rot head dimensions + Qcur = ggml_rope_ext( + ctx0, Qcur, inp_pos, nullptr, + n_rot, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow + ); + + Kcur = ggml_rope_ext( + ctx0, Kcur, inp_pos, nullptr, + n_rot, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow + ); + + cb(Qcur, "Qcur", il); + cb(Kcur, "Kcur", il); + cb(Vcur, "Vcur", il); + + attn_out = build_attn(inp_attn, + model.layers[il].wo, NULL, model.layers[il].wo_s, + Qcur, Kcur, Vcur, nullptr, nullptr, nullptr, 1.0f/sqrtf(float(n_embd_head)), il); + } + + ggml_tensor * ffn_out = build_ffn(cur, + model.layers[il].ffn_up, model.layers[il].ffn_up_b, NULL, + NULL, NULL, NULL, + model.layers[il].ffn_down, model.layers[il].ffn_down_b, NULL, + NULL, + LLM_FFN_GELU, LLM_FFN_SEQ, il); + cb(ffn_out, "ffn_out", il); + + if (il == n_layer - 1 && inp_out_ids) { + attn_out = ggml_get_rows(ctx0, attn_out, inp_out_ids); + ffn_out = ggml_get_rows(ctx0, ffn_out, inp_out_ids); + inpL = ggml_get_rows(ctx0, inpL, inp_out_ids); + } + + cur = ggml_add(ctx0, inpL, attn_out); + cur = ggml_add(ctx0, cur, ffn_out); + + cur = build_cvec(cur, il); + cb(cur, "l_out", il); + + inpL = cur; + } + + cur = build_norm(inpL, + model.output_norm, + model.output_norm_b, + LLM_NORM, -1); + + cb(cur, "result_norm", -1); + res->t_embd = cur; + + cur = build_lora_mm(model.output, cur, model.output_s); + if (model.output_b) { + cur = ggml_add(ctx0, cur, model.output_b); + } + + cb(cur, "result_output", -1); + res->t_logits = cur; + + ggml_build_forward_expand(gf, cur); +} diff --git a/src/models/gemma4.cpp b/src/models/gemma4.cpp index e44f423bdbc5..aa518c6df504 100644 --- a/src/models/gemma4.cpp +++ b/src/models/gemma4.cpp @@ -50,7 +50,7 @@ void llama_model_gemma4::load_arch_tensors(llama_model_loader &) { tok_embd = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, 0); if (n_embd_per_layer > 0) { - per_layer_tok_embd = create_tensor(tn(LLM_TENSOR_PER_LAYER_TOKEN_EMBD, "weight"), {n_embd_per_layer * n_layer, n_vocab}, 0); + per_layer_tok_embd = create_tensor(tn(LLM_TENSOR_PER_LAYER_TOKEN_EMBD, "weight"), {n_embd_per_layer * n_layer, n_vocab}, TENSOR_READ_LAZY); per_layer_model_proj = create_tensor(tn(LLM_TENSOR_PER_LAYER_MODEL_PROJ, "weight", 0), {n_embd, n_embd_per_layer * n_layer}, 0); per_layer_proj_norm = create_tensor(tn(LLM_TENSOR_PER_LAYER_PROJ_NORM, "weight", 0), {n_embd_per_layer}, 0); } diff --git a/src/models/gptj.cpp b/src/models/gptj.cpp new file mode 100644 index 000000000000..5e8342ca3fde --- /dev/null +++ b/src/models/gptj.cpp @@ -0,0 +1,144 @@ +#include "models.h" + +void llama_model_gptj::load_arch_hparams(llama_model_loader & ml) { + ml.get_key(LLM_KV_ATTENTION_LAYERNORM_EPS, hparams.f_norm_eps); + + // GPT-J runs attention and the MLP in parallel off a single layer norm. + hparams.use_par_res = true; + + switch (hparams.n_layer()) { + case 28: type = LLM_TYPE_6B; break; + default: type = LLM_TYPE_UNKNOWN; + } +} + +void llama_model_gptj::load_arch_tensors(llama_model_loader &) { + LLAMA_LOAD_LOCALS; + + tok_embd = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, 0); + + // output + output_norm = create_tensor(tn(LLM_TENSOR_OUTPUT_NORM, "weight"), {n_embd}, 0); + output_norm_b = create_tensor(tn(LLM_TENSOR_OUTPUT_NORM, "bias"), {n_embd}, 0); + output = create_tensor(tn(LLM_TENSOR_OUTPUT, "weight"), {n_embd, n_vocab}, TENSOR_NOT_REQUIRED); + output_b = create_tensor(tn(LLM_TENSOR_OUTPUT, "bias"), {n_vocab}, TENSOR_NOT_REQUIRED); // lm_head has a bias + if (output == NULL) { + output = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, TENSOR_DUPLICATED); + } + + for (int i = 0; i < n_layer; ++i) { + auto & layer = layers[i]; + + layer.attn_norm = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "weight", i), {n_embd}, 0); + layer.attn_norm_b = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "bias", i), {n_embd}, 0); + + // GPT-J stores q/k/v/out as separate, bias-free projections + layer.wq = create_tensor(tn(LLM_TENSOR_ATTN_Q, "weight", i), {n_embd, n_embd}, 0); + layer.wk = create_tensor(tn(LLM_TENSOR_ATTN_K, "weight", i), {n_embd, n_embd_gqa}, 0); + layer.wv = create_tensor(tn(LLM_TENSOR_ATTN_V, "weight", i), {n_embd, n_embd_gqa}, 0); + layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", i), {n_embd, n_embd}, 0); + + // the MLP does carry biases + layer.ffn_up = create_tensor(tn(LLM_TENSOR_FFN_UP, "weight", i), {n_embd, n_ff}, 0); + layer.ffn_up_b = create_tensor(tn(LLM_TENSOR_FFN_UP, "bias", i), {n_ff}, 0); + layer.ffn_down = create_tensor(tn(LLM_TENSOR_FFN_DOWN, "weight", i), {n_ff, n_embd}, 0); + layer.ffn_down_b = create_tensor(tn(LLM_TENSOR_FFN_DOWN, "bias", i), {n_embd}, 0); + } +} + +std::unique_ptr llama_model_gptj::build_arch_graph(const llm_graph_params & params) const { + return std::make_unique(*this, params); +} + +llama_model_gptj::graph::graph(const llama_model & model, const llm_graph_params & params) : llm_graph_context(params) { + const int64_t n_embd_head = hparams.n_embd_head_v(); + + GGML_ASSERT(n_embd_head == hparams.n_embd_head_k()); + + ggml_tensor * cur; + ggml_tensor * inpL; + + inpL = build_inp_embd(model.tok_embd); + + ggml_tensor * inp_pos = build_inp_pos(); + + auto * inp_attn = build_attn_inp_kv(); + + ggml_tensor * inp_out_ids = build_inp_out_ids(); + + for (int il = 0; il < n_layer; ++il) { + cur = build_norm(inpL, + model.layers[il].attn_norm, + model.layers[il].attn_norm_b, + LLM_NORM, il); + cb(cur, "attn_norm", il); + + // x = x + attn(ln_1(x)) + mlp(ln_1(x)) + ggml_tensor * attn_out; + { + auto [Qcur, Kcur, Vcur] = build_qkv(model.layers[il], cur, + n_embd_head, n_head, n_head_kv, il); + + Qcur = ggml_rope_ext( + ctx0, Qcur, inp_pos, nullptr, + n_rot, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow + ); + + Kcur = ggml_rope_ext( + ctx0, Kcur, inp_pos, nullptr, + n_rot, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow + ); + + cb(Qcur, "Qcur", il); + cb(Kcur, "Kcur", il); + cb(Vcur, "Vcur", il); + + attn_out = build_attn(inp_attn, + model.layers[il].wo, NULL, model.layers[il].wo_s, + Qcur, Kcur, Vcur, nullptr, nullptr, nullptr, 1.0f/sqrtf(float(n_embd_head)), il); + } + + ggml_tensor * ffn_out = build_ffn(cur, + model.layers[il].ffn_up, model.layers[il].ffn_up_b, NULL, + NULL, NULL, NULL, + model.layers[il].ffn_down, model.layers[il].ffn_down_b, NULL, + NULL, + LLM_FFN_GELU, LLM_FFN_SEQ, il); + cb(ffn_out, "ffn_out", il); + + if (il == n_layer - 1 && inp_out_ids) { + attn_out = ggml_get_rows(ctx0, attn_out, inp_out_ids); + ffn_out = ggml_get_rows(ctx0, ffn_out, inp_out_ids); + inpL = ggml_get_rows(ctx0, inpL, inp_out_ids); + } + + cur = ggml_add(ctx0, inpL, attn_out); + cur = ggml_add(ctx0, cur, ffn_out); + + cur = build_cvec(cur, il); + cb(cur, "l_out", il); + + // input for next layer + inpL = cur; + } + + cur = build_norm(inpL, + model.output_norm, + model.output_norm_b, + LLM_NORM, -1); + + cb(cur, "result_norm", -1); + res->t_embd = cur; + + cur = build_lora_mm(model.output, cur, model.output_s); + if (model.output_b) { + cur = ggml_add(ctx0, cur, model.output_b); + } + + cb(cur, "result_output", -1); + res->t_logits = cur; + + ggml_build_forward_expand(gf, cur); +} diff --git a/src/models/gptneo.cpp b/src/models/gptneo.cpp new file mode 100644 index 000000000000..3f35164624db --- /dev/null +++ b/src/models/gptneo.cpp @@ -0,0 +1,139 @@ +#include "models.h" + +void llama_model_gptneo::load_arch_hparams(llama_model_loader & ml) { + ml.get_key(LLM_KV_ATTENTION_LAYERNORM_EPS, hparams.f_norm_eps); + + // GPT-Neo alternates global and local attention, dense layer first, window 256. + hparams.swa_type = LLAMA_SWA_TYPE_STANDARD; + hparams.n_swa = 256; + ml.get_key(LLM_KV_ATTENTION_SLIDING_WINDOW, hparams.n_swa, false); + hparams.set_swa_pattern(2); +} + +void llama_model_gptneo::load_arch_tensors(llama_model_loader &) { + LLAMA_LOAD_LOCALS; + + tok_embd = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, 0); + pos_embd = create_tensor(tn(LLM_TENSOR_POS_EMBD, "weight"), {n_embd, n_ctx_train}, 0); + + output_norm = create_tensor(tn(LLM_TENSOR_OUTPUT_NORM, "weight"), {n_embd}, 0); + output_norm_b = create_tensor(tn(LLM_TENSOR_OUTPUT_NORM, "bias"), {n_embd}, 0); + output = create_tensor(tn(LLM_TENSOR_OUTPUT, "weight"), {n_embd, n_vocab}, TENSOR_NOT_REQUIRED); + if (output == NULL) { + output = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, TENSOR_DUPLICATED); + } + + for (int i = 0; i < n_layer; ++i) { + auto & layer = layers[i]; + + layer.attn_norm = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "weight", i), {n_embd}, 0); + layer.attn_norm_b = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "bias", i), {n_embd}, 0); + + // q/k/v are bias-free; only the output projection carries a bias + layer.wq = create_tensor(tn(LLM_TENSOR_ATTN_Q, "weight", i), {n_embd, n_embd}, 0); + layer.wk = create_tensor(tn(LLM_TENSOR_ATTN_K, "weight", i), {n_embd, n_embd_gqa}, 0); + layer.wv = create_tensor(tn(LLM_TENSOR_ATTN_V, "weight", i), {n_embd, n_embd_gqa}, 0); + layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", i), {n_embd, n_embd}, 0); + layer.wo_b = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "bias", i), {n_embd}, 0); + + layer.ffn_norm = create_tensor(tn(LLM_TENSOR_FFN_NORM, "weight", i), {n_embd}, 0); + layer.ffn_norm_b = create_tensor(tn(LLM_TENSOR_FFN_NORM, "bias", i), {n_embd}, 0); + layer.ffn_up = create_tensor(tn(LLM_TENSOR_FFN_UP, "weight", i), {n_embd, n_ff}, 0); + layer.ffn_up_b = create_tensor(tn(LLM_TENSOR_FFN_UP, "bias", i), {n_ff}, 0); + layer.ffn_down = create_tensor(tn(LLM_TENSOR_FFN_DOWN, "weight", i), {n_ff, n_embd}, 0); + layer.ffn_down_b = create_tensor(tn(LLM_TENSOR_FFN_DOWN, "bias", i), {n_embd}, 0); + } +} + +std::unique_ptr llama_model_gptneo::build_arch_graph(const llm_graph_params & params) const { + return std::make_unique(*this, params); +} + +llama_model_gptneo::graph::graph(const llama_model & model, const llm_graph_params & params) : llm_graph_context(params) { + const int64_t n_embd_head = hparams.n_embd_head_v(); + + GGML_ASSERT(n_embd_head == hparams.n_embd_head_k()); + + ggml_tensor * cur; + ggml_tensor * inpL; + + inpL = build_inp_embd(model.tok_embd); + + ggml_tensor * inp_pos = build_inp_pos(); + + auto * inp_attn = build_attn_inp_kv_iswa(); + + ggml_tensor * inp_out_ids = build_inp_out_ids(); + + cur = ggml_get_rows(ctx0, model.pos_embd, inp_pos); + cb(cur, "pos_embd", -1); + inpL = ggml_add(ctx0, inpL, cur); + cb(inpL, "inpL", -1); + + for (int il = 0; il < n_layer; ++il) { + cur = build_norm(inpL, + model.layers[il].attn_norm, + model.layers[il].attn_norm_b, + LLM_NORM, il); + cb(cur, "attn_norm", il); + + { + auto [Qcur, Kcur, Vcur] = build_qkv(model.layers[il], cur, + n_embd_head, n_head, n_head_kv, il); + + cb(Qcur, "Qcur", il); + cb(Kcur, "Kcur", il); + cb(Vcur, "Vcur", il); + + cur = build_attn(inp_attn, + model.layers[il].wo, model.layers[il].wo_b, model.layers[il].wo_s, + Qcur, Kcur, Vcur, nullptr, nullptr, nullptr, 1.0f, il); + } + + if (il == n_layer - 1 && inp_out_ids) { + cur = ggml_get_rows(ctx0, cur, inp_out_ids); + inpL = ggml_get_rows(ctx0, inpL, inp_out_ids); + } + + ggml_tensor * ffn_inp = ggml_add(ctx0, cur, inpL); + cb(ffn_inp, "ffn_inp", il); + + { + cur = build_norm(ffn_inp, + model.layers[il].ffn_norm, + model.layers[il].ffn_norm_b, + LLM_NORM, il); + cb(cur, "ffn_norm", il); + + cur = build_ffn(cur, + model.layers[il].ffn_up, model.layers[il].ffn_up_b, NULL, + NULL, NULL, NULL, + model.layers[il].ffn_down, model.layers[il].ffn_down_b, NULL, + NULL, + LLM_FFN_GELU, LLM_FFN_SEQ, il); + cb(cur, "ffn_out", il); + } + + cur = ggml_add(ctx0, cur, ffn_inp); + + cur = build_cvec(cur, il); + cb(cur, "l_out", il); + + inpL = cur; + } + + cur = build_norm(inpL, + model.output_norm, + model.output_norm_b, + LLM_NORM, -1); + + cb(cur, "result_norm", -1); + res->t_embd = cur; + + cur = build_lora_mm(model.output, cur, model.output_s); + + cb(cur, "result_output", -1); + res->t_logits = cur; + + ggml_build_forward_expand(gf, cur); +} diff --git a/src/models/models.h b/src/models/models.h index bb372ece81ee..dea1e1366c2f 100644 --- a/src/models/models.h +++ b/src/models/models.h @@ -6,6 +6,16 @@ // note: almost all graphs require at least sqrtf, so include cmath globally #include +#include + +class llama_memory_hybrid_idx_context; + +// ref: https://github.com/ggml-org/llama.cpp/pull/28068 +static inline ggml_tensor * build_gdn_l2_norm(ggml_context * ctx, ggml_tensor * x, float eps) { + const float n = x->ne[0]; + + return ggml_scale(ctx, ggml_rms_norm(ctx, x, eps/n), 1.0f/sqrtf(n)); +} // // base classes @@ -386,6 +396,32 @@ struct llama_model_bloom : public llama_model_base { }; +struct llama_model_codegen : public llama_model_base { + llama_model_codegen(const struct llama_model_params & params) : llama_model_base(params) {} + void load_arch_hparams(llama_model_loader & ml) override; + void load_arch_tensors(llama_model_loader & ml) override; + + struct graph : public llm_graph_context { + graph(const llama_model & model, const llm_graph_params & params); + }; + + std::unique_ptr build_arch_graph(const llm_graph_params & params) const override; +}; + + +struct llama_model_gptneo : public llama_model_base { + llama_model_gptneo(const struct llama_model_params & params) : llama_model_base(params) {} + void load_arch_hparams(llama_model_loader & ml) override; + void load_arch_tensors(llama_model_loader & ml) override; + + struct graph : public llm_graph_context { + graph(const llama_model & model, const llm_graph_params & params); + }; + + std::unique_ptr build_arch_graph(const llm_graph_params & params) const override; +}; + + struct llama_model_mpt : public llama_model_base { llama_model_mpt(const struct llama_model_params & params) : llama_model_base(params) {} void load_arch_hparams(llama_model_loader & ml) override; @@ -398,6 +434,18 @@ struct llama_model_mpt : public llama_model_base { std::unique_ptr build_arch_graph(const llm_graph_params & params) const override; }; +struct llama_model_opt : public llama_model_base { + llama_model_opt(const struct llama_model_params & params) : llama_model_base(params) {} + void load_arch_hparams(llama_model_loader & ml) override; + void load_arch_tensors(llama_model_loader & ml) override; + + struct graph : public llm_graph_context { + graph(const llama_model & model, const llm_graph_params & params); + }; + + std::unique_ptr build_arch_graph(const llm_graph_params & params) const override; +}; + struct llama_model_stablelm : public llama_model_base { llama_model_stablelm(const struct llama_model_params & params) : llama_model_base(params) {} @@ -1049,6 +1097,19 @@ struct llama_model_gptneox : public llama_model_base { }; +struct llama_model_gptj : public llama_model_base { + llama_model_gptj(const struct llama_model_params & params) : llama_model_base(params) {} + void load_arch_hparams(llama_model_loader & ml) override; + void load_arch_tensors(llama_model_loader & ml) override; + + struct graph : public llm_graph_context { + graph(const llama_model & model, const llm_graph_params & params); + }; + + std::unique_ptr build_arch_graph(const llm_graph_params & params) const override; +}; + + struct llama_model_arctic : public llama_model_base { llama_model_arctic(const struct llama_model_params & params) : llama_model_base(params) {} void load_arch_hparams(llama_model_loader & ml) override; @@ -1259,6 +1320,11 @@ struct llama_model_eagle3 : public llama_model_base { std::unique_ptr build_arch_graph(const llm_graph_params & params) const override; }; +// explicit specializations, declared before their first use (make_unique is constexpr since C++23) +template <> llama_model_eagle3::graph::graph(const llama_model & model, const llm_graph_params & params); +template <> llama_model_eagle3::graph::graph(const llama_model & model, const llm_graph_params & params); +template <> ggml_tensor * llama_model_eagle3::graph::build_inp_embd_enc() const; + struct llama_model_dflash : public llama_model_base { llama_model_dflash(const struct llama_model_params & params) : llama_model_base(params) {} @@ -1275,6 +1341,11 @@ struct llama_model_dflash : public llama_model_base { std::unique_ptr build_arch_graph(const llm_graph_params & params) const override; }; +// explicit specializations, declared before their first use (make_unique is constexpr since C++23) +template <> llama_model_dflash::graph::graph(const llama_model & model, const llm_graph_params & params); +template <> llama_model_dflash::graph::graph(const llama_model & model, const llm_graph_params & params); +template <> ggml_tensor * llama_model_dflash::graph::build_inp_embd_enc() const; + struct llama_model_mistral4 : public llama_model_deepseek2 { llama_model_mistral4(const struct llama_model_params & params) : llama_model_deepseek2(params) {} @@ -1351,6 +1422,10 @@ struct llama_model_t5 : public llama_model_base { std::unique_ptr build_arch_graph(const llm_graph_params & params) const override; }; +// explicit specializations, declared before their first use (make_unique is constexpr since C++23) +template <> llama_model_t5::graph::graph(const llama_model & model, const llm_graph_params & params); +template <> llama_model_t5::graph::graph(const llama_model & model, const llm_graph_params & params); + struct llama_model_t5encoder : public llama_model_base { llama_model_t5encoder(const struct llama_model_params & params) : llama_model_base(params) {} @@ -1844,6 +1919,25 @@ struct llama_model_falcon_h1 : public llama_model_base { }; +struct llama_model_zaya : public llama_model_base { + llama_model_zaya(const struct llama_model_params & params) : llama_model_base(params) {} + void load_arch_hparams(llama_model_loader & ml) override; + void load_arch_tensors(llama_model_loader & ml) override; + + // ZAYA1-VL: ranks of the vision-only LoRA (0: none) + uint32_t vlora_rank_attn = 0; + uint32_t vlora_rank_ffn = 0; + + // iswa: sliding-window layers (ZAYA1-74B), on the hybrid iSWA memory + template + struct graph : public llm_graph_context { + graph(const llama_model & model, const llm_graph_params & params); + }; + + std::unique_ptr build_arch_graph(const llm_graph_params & params) const override; +}; + + struct llama_model_lfm2 : public llama_model_base { llama_model_lfm2(const struct llama_model_params & params) : llama_model_base(params) {} void load_arch_hparams(llama_model_loader & ml) override; @@ -2059,6 +2153,119 @@ struct llama_model_qwen35 : public llama_model_base { }; +struct llama_model_qwen4exp : public llama_model_base { + llama_model_qwen4exp(const struct llama_model_params & params) : llama_model_base(params) {} + + class llm_graph_input_qsa; + + void load_arch_hparams(llama_model_loader & ml) override; + void load_arch_tensors(llama_model_loader & ml) override; + + struct graph : public llm_build_delta_net_base { + graph(const llama_model & model, const llm_graph_params & params); + protected: + struct no_build_t {}; + graph(const llama_model & model, const llm_graph_params & params, no_build_t) : + llm_build_delta_net_base(params), model(model) {} + + // HC replaces every layer norm: residual is [n_embd, hc, n_tokens] + ggml_tensor * build_hc_mix( + ggml_tensor * x, + ggml_tensor * w_norm, + ggml_tensor * w_down, + ggml_tensor * w_up, + ggml_tensor * w_inject, + ggml_tensor ** inject, + int il); + + ggml_tensor * build_hc_combine( + ggml_tensor * residual, + ggml_tensor * block_out, + ggml_tensor * inject, + int il); + + ggml_tensor * build_layer_attn( + llm_graph_input_attn_kv * inp_attn, + const llama_memory_hybrid_idx_context * mctx_hyb, + ggml_tensor * cur, + ggml_tensor * inp_pos, + int * sections, + int il); + + // dense self-attention restricted to the cells that top_k names + ggml_tensor * build_attn_qsa( + llm_graph_input_attn_kv * inp, + ggml_tensor * q_cur, + ggml_tensor * k_cur, + ggml_tensor * v_cur, + ggml_tensor * top_k, + float kq_scale, + int il); + + // the QSA cache layout inputs do not depend on the layer, only on its compress ratio, + // so the layers sharing a ratio share one input set + std::map qsa_inps; + + // QSA: token indices this layer's queries may attend to, or nullptr for dense + ggml_tensor * build_qsa_top_k( + const llama_memory_hybrid_idx_context * mctx_hyb, + ggml_tensor * cur, + ggml_tensor * inp_pos, + ggml_tensor * kq_mask, + int * sections, + int il); + + ggml_tensor * build_layer_attn_linear( + llm_graph_input_rs * inp, + ggml_tensor * cur, + int il); + + ggml_tensor * build_layer_ffn( + ggml_tensor * cur, + int il); + + ggml_tensor * build_norm_gated( + ggml_tensor * input, + ggml_tensor * weights, + ggml_tensor * gate, + int layer); + + // build_rs writes the state tensor in place, so one gather per cache tensor is reused + std::map rs_rows; + + // one conv history per cache tensor: delta-net and PLE each have their own + ggml_tensor * build_conv_state_at( + llm_graph_input_rs * inp, + ggml_tensor * conv_states_all, + ggml_tensor * x, + int64_t state_cols, + int64_t channels, + int il); + + ggml_tensor * build_inp_ple( + const llama_memory_hybrid_idx_context * mctx_hyb); + + ggml_tensor * build_ple( + llm_graph_input_rs * inp, + ggml_tensor * emb, + ggml_tensor * hidden, + int il); + + // returns pair of qkv, z + std::pair build_qkvz( + ggml_tensor * input, + int il); + + const llama_model & model; + }; + + struct graph_mtp : public graph { + graph_mtp(const llama_model & model, const llm_graph_params & params); + }; + + std::unique_ptr build_arch_graph(const llm_graph_params & params) const override; +}; + struct llama_model_qwen35moe : public llama_model_base { llama_model_qwen35moe(const struct llama_model_params & params) : llama_model_base(params) {} void load_arch_hparams(llama_model_loader & ml) override; diff --git a/src/models/opt.cpp b/src/models/opt.cpp new file mode 100644 index 000000000000..ce90d91b7d3d --- /dev/null +++ b/src/models/opt.cpp @@ -0,0 +1,139 @@ +#include "models.h" + +void llama_model_opt::load_arch_hparams(llama_model_loader & ml) { + ml.get_key(LLM_KV_ATTENTION_LAYERNORM_EPS, hparams.f_norm_eps); +} + +void llama_model_opt::load_arch_tensors(llama_model_loader &) { + LLAMA_LOAD_LOCALS; + + tok_embd = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, 0); + pos_embd = create_tensor(tn(LLM_TENSOR_POS_EMBD, "weight"), {n_embd, n_ctx_train}, 0); + + // output + output_norm = create_tensor(tn(LLM_TENSOR_OUTPUT_NORM, "weight"), {n_embd}, 0); + output_norm_b = create_tensor(tn(LLM_TENSOR_OUTPUT_NORM, "bias"), {n_embd}, 0); + output = create_tensor(tn(LLM_TENSOR_OUTPUT, "weight"), {n_embd, n_vocab}, TENSOR_NOT_REQUIRED); + if (output == NULL) { + output = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, TENSOR_DUPLICATED); + } + + for (int i = 0; i < n_layer; ++i) { + auto & layer = layers[i]; + + layer.attn_norm = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "weight", i), {n_embd}, 0); + layer.attn_norm_b = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "bias", i), {n_embd}, 0); + + layer.wq = create_tensor(tn(LLM_TENSOR_ATTN_Q, "weight", i), {n_embd, n_embd}, 0); + layer.wq_b = create_tensor(tn(LLM_TENSOR_ATTN_Q, "bias", i), {n_embd}, 0); + layer.wk = create_tensor(tn(LLM_TENSOR_ATTN_K, "weight", i), {n_embd, n_embd_gqa}, 0); + layer.wk_b = create_tensor(tn(LLM_TENSOR_ATTN_K, "bias", i), {n_embd_gqa}, 0); + layer.wv = create_tensor(tn(LLM_TENSOR_ATTN_V, "weight", i), {n_embd, n_embd_gqa}, 0); + layer.wv_b = create_tensor(tn(LLM_TENSOR_ATTN_V, "bias", i), {n_embd_gqa}, 0); + layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", i), {n_embd, n_embd}, 0); + layer.wo_b = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "bias", i), {n_embd}, 0); + + layer.ffn_norm = create_tensor(tn(LLM_TENSOR_FFN_NORM, "weight", i), {n_embd}, 0); + layer.ffn_norm_b = create_tensor(tn(LLM_TENSOR_FFN_NORM, "bias", i), {n_embd}, 0); + layer.ffn_up = create_tensor(tn(LLM_TENSOR_FFN_UP, "weight", i), {n_embd, n_ff}, 0); + layer.ffn_up_b = create_tensor(tn(LLM_TENSOR_FFN_UP, "bias", i), {n_ff}, 0); + layer.ffn_down = create_tensor(tn(LLM_TENSOR_FFN_DOWN, "weight", i), {n_ff, n_embd}, 0); + layer.ffn_down_b = create_tensor(tn(LLM_TENSOR_FFN_DOWN, "bias", i), {n_embd}, 0); + } +} + +std::unique_ptr llama_model_opt::build_arch_graph(const llm_graph_params & params) const { + return std::make_unique(*this, params); +} + +llama_model_opt::graph::graph(const llama_model & model, const llm_graph_params & params) : llm_graph_context(params) { + const int64_t n_embd_head = hparams.n_embd_head_v(); + + GGML_ASSERT(n_embd_head == hparams.n_embd_head_k()); + + ggml_tensor * cur; + ggml_tensor * inpL; + + inpL = build_inp_embd(model.tok_embd); + + ggml_tensor * inp_pos = build_inp_pos(); + + auto * inp_attn = build_attn_inp_kv(); + + ggml_tensor * inp_out_ids = build_inp_out_ids(); + + cur = ggml_get_rows(ctx0, model.pos_embd, inp_pos); + cb(cur, "pos_embd", -1); + inpL = ggml_add(ctx0, inpL, cur); + cb(inpL, "inpL", -1); + + for (int il = 0; il < n_layer; ++il) { + // pre-norm attention + ggml_tensor * residual = inpL; + + cur = build_norm(inpL, + model.layers[il].attn_norm, + model.layers[il].attn_norm_b, + LLM_NORM, il); + cb(cur, "attn_norm", il); + + { + auto [Qcur, Kcur, Vcur] = build_qkv(model.layers[il], cur, + n_embd_head, n_head, n_head_kv, il); + + cb(Qcur, "Qcur", il); + cb(Kcur, "Kcur", il); + cb(Vcur, "Vcur", il); + + cur = build_attn(inp_attn, + model.layers[il].wo, model.layers[il].wo_b, model.layers[il].wo_s, + Qcur, Kcur, Vcur, nullptr, nullptr, nullptr, 1.0f/sqrtf(float(n_embd_head)), il); + } + + cur = ggml_add(ctx0, cur, residual); + cb(cur, "attn_out", il); + + // pre-norm mlp + residual = cur; + + cur = build_norm(cur, + model.layers[il].ffn_norm, + model.layers[il].ffn_norm_b, + LLM_NORM, il); + cb(cur, "ffn_norm", il); + + cur = build_ffn(cur, + model.layers[il].ffn_up, model.layers[il].ffn_up_b, NULL, + NULL, NULL, NULL, + model.layers[il].ffn_down, model.layers[il].ffn_down_b, NULL, + NULL, + LLM_FFN_RELU, LLM_FFN_SEQ, il); + cb(cur, "ffn_out", il); + + cur = ggml_add(ctx0, cur, residual); + + if (il == n_layer - 1 && inp_out_ids) { + cur = ggml_get_rows(ctx0, cur, inp_out_ids); + } + + cur = build_cvec(cur, il); + cb(cur, "l_out", il); + + inpL = cur; + } + + cur = build_norm(inpL, + model.output_norm, + model.output_norm_b, + LLM_NORM, -1); + + cb(cur, "result_norm", -1); + res->t_embd = cur; + + cur = build_lora_mm(model.output, cur, model.output_s); + + cb(cur, "result_output", -1); + res->t_logits = cur; + + ggml_build_forward_expand(gf, cur); +} diff --git a/src/models/qwen35.cpp b/src/models/qwen35.cpp index 309dd432447c..3e5fdf58d065 100644 --- a/src/models/qwen35.cpp +++ b/src/models/qwen35.cpp @@ -522,6 +522,7 @@ llama_model_qwen35::graph_mtp::graph_mtp(const llama_model & model, const llm_gr ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd; tok_embd = ggml_get_rows(ctx0, tok_embd_w, inp->tokens); + tok_embd = hadamard ? hadamard->inverse(ctx0, tok_embd_w, tok_embd) : tok_embd; } else { tok_embd = inp->embd; } diff --git a/src/models/qwen4exp-draft-vocab.cpp b/src/models/qwen4exp-draft-vocab.cpp new file mode 100644 index 000000000000..e76636a071ad --- /dev/null +++ b/src/models/qwen4exp-draft-vocab.cpp @@ -0,0 +1,71 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +#include "qwen4exp-draft-vocab.h" + +#include "llama-impl.h" +#include "llama-mmap.h" + +#include "ggml.h" + +#include +#include +#include +#include + +int64_t qwen4exp_draft_vocab_rows(const llama_model_loader & ml, const std::string & d2t_name, int64_t n_vocab) { + const auto * w = ml.get_weight(d2t_name.c_str()); + if (w == nullptr) { + return 0; + } + const ggml_tensor * d2t = w->tensor; + if (d2t->type != GGML_TYPE_I64 || ggml_n_dims(d2t) != 1) { + throw std::runtime_error(format("QWEN4EXP MTP: d2t must be a 1-D I64 tensor, got %s", ggml_type_name(d2t->type))); + } + if (d2t->ne[0] <= 0 || d2t->ne[0] > n_vocab) { + throw std::runtime_error(format("QWEN4EXP MTP: d2t has %lld rows for a %lld-token vocabulary", + (long long) d2t->ne[0], (long long) n_vocab)); + } + // set_rows scatters each draft row to its d2t id, so every id must be a distinct token: read the K ids from + // the file now (K int64s), before any graph uses them, and refuse the file otherwise. + const int64_t k = d2t->ne[0]; + std::vector ids(k); + llama_file & file = *ml.files.at(w->idx); + file.seek(w->offs, SEEK_SET); + file.read_raw(ids.data(), ids.size() * sizeof(int64_t)); + std::vector seen(n_vocab, false); + for (int64_t i = 0; i < k; ++i) { + if (ids[i] < 0 || ids[i] >= n_vocab) { + throw std::runtime_error(format("QWEN4EXP MTP: d2t[%lld] = %lld is outside the %lld-token vocabulary", + (long long) i, (long long) ids[i], (long long) n_vocab)); + } + if (seen[ids[i]]) { + throw std::runtime_error(format("QWEN4EXP MTP: d2t repeats token id %lld", (long long) ids[i])); + } + seen[ids[i]] = true; + } + return k; +} + +ggml_tensor * qwen4exp_draft_vocab_scatter(ggml_context * ctx, ggml_tensor * logits, ggml_tensor * d2t, int64_t n_vocab) { + const int64_t n_rows = logits->ne[0]; + const int64_t n_outputs = logits->ne[1]; + GGML_ASSERT(d2t->type == GGML_TYPE_I64 && d2t->ne[0] == n_rows); + + ggml_tensor * full = ggml_fill(ctx, ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 1, n_vocab, n_outputs), -INFINITY); + full = ggml_set_rows(ctx, full, + ggml_reshape_3d(ctx, logits, 1, n_rows, n_outputs), + ggml_reshape_3d(ctx, d2t, n_rows, 1, 1)); + return ggml_reshape_2d(ctx, full, n_vocab, n_outputs); +} diff --git a/src/models/qwen4exp-draft-vocab.h b/src/models/qwen4exp-draft-vocab.h new file mode 100644 index 000000000000..2bd9a1082b96 --- /dev/null +++ b/src/models/qwen4exp-draft-vocab.h @@ -0,0 +1,40 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +#pragma once + +// Reduced draft vocabulary for the qwen4exp (Qwen3.8-Flash-Next) MTP head (1bit engine, +// tools/mtp_draft_vocab.py). A draft head GGUF may carry +// blk..nextn.shared_head_head.weight [n_embd, K] the target's output rows for K token ids +// d2t I64 [K] draft row -> target token id +// The draft pass then reads a K-row head instead of the full one (65,536 of 248,320 rows is +// 170 MiB instead of 644 MiB at Q8_0, three times per decode step) and its logits are scattered +// into a full-vocabulary row that is -inf elsewhere. The target verifies every drafted token, +// so only the acceptance rate can change, never the verified output. The loader reads the d2t +// values and refuses a file whose ids are out of range or repeated, since set_rows trusts them. + +#include "llama-model-loader.h" + +#include +#include + +struct ggml_context; +struct ggml_tensor; + +// Rows of the reduced draft head: 0 when the file has no d2t (full head), else d2t's length. +// Throws if d2t is not I64, is longer than the vocabulary, or holds an id outside it or twice. +int64_t qwen4exp_draft_vocab_rows(const llama_model_loader & ml, const std::string & d2t_name, int64_t n_vocab); + +// logits [K, n_outputs] -> [n_vocab, n_outputs], -inf outside the d2t rows. +ggml_tensor * qwen4exp_draft_vocab_scatter(ggml_context * ctx, ggml_tensor * logits, ggml_tensor * d2t, int64_t n_vocab); diff --git a/src/models/qwen4exp.cpp b/src/models/qwen4exp.cpp new file mode 100644 index 000000000000..b04840fdf941 --- /dev/null +++ b/src/models/qwen4exp.cpp @@ -0,0 +1,1579 @@ +#include "models.h" +#include +#include "qwen4exp-draft-vocab.h" +#include "llama-impl.h" +#include "llama-memory-hybrid-idx.h" +#include "llama-memory-recurrent.h" + +#include +#include + +// bad metadata must be catchable: GGML_ASSERT aborts the whole process +static void qwen4exp_require_nonzero(const llama_model_loader & ml, llm_kv kid, uint32_t value) { + if (value == 0) { + throw std::runtime_error(format("%s must be greater than zero, got %u", ml.llm_kv(kid).c_str(), value)); + } +} + +// get_arr() copies a short array as-is, leaving a zero tail the n-gram hash silently drops +static void qwen4exp_require_arr_len(llama_model_loader & ml, llm_kv kid, uint32_t n_min) { + uint32_t n_arr = 0; + ml.get_arr_n(kid, n_arr, true); + if (n_arr < n_min) { + throw std::runtime_error(format("%s has %u entries, but at least %u are required", + ml.llm_kv(kid).c_str(), n_arr, n_min)); + } +} + +static const llama_model & qwen4exp_shared_model(const llama_cparams & cparams, const llama_model & model, const char * name) { + if (cparams.ctx_other == nullptr) { + throw std::runtime_error(format("QWEN4EXP MTP: this draft head has no '%s' of its own; " + "load it as a draft of its target model (-md), not on its own", name)); + } + const llama_model & other = *llama_get_model(cparams.ctx_other); + if (other.hparams.n_embd != model.hparams.n_embd || other.vocab.n_tokens() != model.vocab.n_tokens()) { + throw std::runtime_error(format("QWEN4EXP MTP: draft and target disagree on the shape of '%s'", name)); + } + return other; +} + +void llama_model_qwen4exp::load_arch_hparams(llama_model_loader & ml) { + // this tree reads the MTP block count per architecture (qwen35 does the same) + ml.get_key(LLM_KV_NEXTN_PREDICT_LAYERS, hparams.n_layer_nextn, false); + + // the trunk must keep at least one block: n_layer() == n_layer_all - n_layer_nextn + if (hparams.n_layer_nextn >= hparams.n_layer_all) { + throw std::runtime_error(format("%s must be less than %s, got %u", + ml.llm_kv(LLM_KV_NEXTN_PREDICT_LAYERS).c_str(), + ml.llm_kv(LLM_KV_BLOCK_COUNT).c_str(), hparams.n_layer_nextn)); + } + + ml.get_key(LLM_KV_EXPERT_FEED_FORWARD_LENGTH, hparams.n_ff_exp, false); + ml.get_key(LLM_KV_EXPERT_SHARED_FEED_FORWARD_LENGTH, hparams.n_ff_shexp, false); + ml.get_key(LLM_KV_ATTENTION_LAYERNORM_RMS_EPS, hparams.f_norm_rms_eps); + + ml.get_key_or_arr(LLM_KV_ROPE_DIMENSION_SECTIONS, hparams.rope_sections, 4, true); + + ml.get_key(LLM_KV_SSM_CONV_KERNEL, hparams.ssm_d_conv); + ml.get_key(LLM_KV_SSM_INNER_SIZE, hparams.ssm_d_inner); + ml.get_key(LLM_KV_SSM_STATE_SIZE, hparams.ssm_d_state); + ml.get_key(LLM_KV_SSM_TIME_STEP_RANK, hparams.ssm_dt_rank); + ml.get_key(LLM_KV_SSM_GROUP_COUNT, hparams.ssm_n_group); + qwen4exp_require_nonzero(ml, LLM_KV_SSM_CONV_KERNEL, hparams.ssm_d_conv); + qwen4exp_require_nonzero(ml, LLM_KV_SSM_INNER_SIZE, hparams.ssm_d_inner); + qwen4exp_require_nonzero(ml, LLM_KV_SSM_STATE_SIZE, hparams.ssm_d_state); + qwen4exp_require_nonzero(ml, LLM_KV_SSM_TIME_STEP_RANK, hparams.ssm_dt_rank); + qwen4exp_require_nonzero(ml, LLM_KV_SSM_GROUP_COUNT, hparams.ssm_n_group); + + // HC; low_rank is qwen4exp-specific, DeepSeek-V4 leaves it absent (full rank) + ml.get_key(LLM_KV_HYPER_CONNECTION_COUNT, hparams.dsv4_hc_mult); + ml.get_key(LLM_KV_HYPER_CONNECTION_LOW_RANK, hparams.hc_low_rank); + // a count of 1 has nothing to mix: transformers configuration_qwen4_exp.py:196, vLLM + // config.py:49 and SGLang configs/qwen4_exp.py:38 all raise on hc_count <= 1 + if (hparams.dsv4_hc_mult <= 1) { + throw std::runtime_error(format("%s must be greater than one, got %u", + ml.llm_kv(LLM_KV_HYPER_CONNECTION_COUNT).c_str(), hparams.dsv4_hc_mult)); + } + qwen4exp_require_nonzero(ml, LLM_KV_HYPER_CONNECTION_LOW_RANK, hparams.hc_low_rank); + hparams.n_embd_out_impl = hparams.dsv4_hc_mult * hparams.n_embd; + + ml.get_key(LLM_KV_ATTENTION_INDEXER_HEAD_COUNT, hparams.indexer_n_head); + ml.get_key(LLM_KV_ATTENTION_INDEXER_KEY_LENGTH, hparams.indexer_head_size); + ml.get_key(LLM_KV_ATTENTION_INDEXER_TOP_K, hparams.indexer_top_k); + qwen4exp_require_nonzero(ml, LLM_KV_ATTENTION_INDEXER_HEAD_COUNT, hparams.indexer_n_head); + qwen4exp_require_nonzero(ml, LLM_KV_ATTENTION_INDEXER_KEY_LENGTH, hparams.indexer_head_size); + qwen4exp_require_nonzero(ml, LLM_KV_ATTENTION_INDEXER_TOP_K, hparams.indexer_top_k); + ml.get_key_or_arr(LLM_KV_ATTENTION_COMPRESS_RATIOS, hparams.dsv4_compress_ratios, hparams.n_layer_all, false); + + // PLE n-gram hash embeddings; if the key group is absent every field stays zero + hparams.is_ple_impl.reset(); + hparams.ple_n_heads = 0; + + uint32_t n_ple = 0; + ml.get_arr_n(LLM_KV_PLE_LAYERS, n_ple, false); + if (n_ple > 0) { + std::vector ple_layers; + ml.get_arr(LLM_KV_PLE_LAYERS, ple_layers); + if (n_ple != 1) { + // hparams holds one set of hash constants, so several PLE modules cannot be represented + throw std::runtime_error(format("%s lists %u layers, but only one PLE layer is supported", + ml.llm_kv(LLM_KV_PLE_LAYERS).c_str(), n_ple)); + } + for (uint32_t il : ple_layers) { + if (il >= hparams.n_layer_all) { + throw std::runtime_error(format("PLE layer %u is out of range", il)); + } + hparams.is_ple_impl.set(il); + } + + ml.get_key(LLM_KV_PLE_NGRAM_SIZE, hparams.ple_ngram_size); + ml.get_key(LLM_KV_PLE_HEADS_PER_NGRAM, hparams.ple_heads_per_ngram); + ml.get_key(LLM_KV_PLE_CONV_KERNEL, hparams.ple_conv_kernel); + ml.get_key(LLM_KV_PLE_EOS_TOKEN_ID, hparams.ple_eos_token_id); + // optional: files written before this key fall back to the EOS token + ml.get_key(LLM_KV_PLE_IMAGE_TOKEN_ID, hparams.ple_image_token_id, false); + ml.get_key(LLM_KV_EMBEDDING_LENGTH_PER_LAYER, hparams.n_embd_per_layer); + qwen4exp_require_nonzero(ml, LLM_KV_PLE_CONV_KERNEL, hparams.ple_conv_kernel); + qwen4exp_require_nonzero(ml, LLM_KV_EMBEDDING_LENGTH_PER_LAYER, hparams.n_embd_per_layer); + + hparams.ple_n_heads = (hparams.ple_ngram_size - 1) * hparams.ple_heads_per_ngram; + hparams.ple_head_dim = hparams.n_embd_per_layer; + if (hparams.ple_ngram_size < 2 || hparams.ple_ngram_size > LLAMA_MAX_PLE_NGRAM) { + throw std::runtime_error(format("PLE n-gram size %u is out of range", hparams.ple_ngram_size)); + } + if (hparams.ple_n_heads == 0 || hparams.ple_n_heads > LLAMA_MAX_PLE_HEADS) { + throw std::runtime_error(format("PLE head count %u is out of range", hparams.ple_n_heads)); + } + + qwen4exp_require_arr_len(ml, LLM_KV_PLE_LAYER_MULTIPLIERS, hparams.ple_ngram_size); + qwen4exp_require_arr_len(ml, LLM_KV_PLE_HEAD_OFFSETS, hparams.ple_n_heads); + qwen4exp_require_arr_len(ml, LLM_KV_PLE_HEAD_VOCAB_SIZES, hparams.ple_n_heads); + + ml.get_arr(LLM_KV_PLE_LAYER_MULTIPLIERS, hparams.ple_layer_multipliers); + + // the file stores the head ranges as uint64, so read at that width and narrow to the int32 the gather uses + std::array head_offsets = {}; + std::array head_vocab_sizes = {}; + ml.get_arr(LLM_KV_PLE_HEAD_OFFSETS, head_offsets); + ml.get_arr(LLM_KV_PLE_HEAD_VOCAB_SIZES, head_vocab_sizes); + for (uint32_t h = 0; h < hparams.ple_n_heads; ++h) { + if (head_vocab_sizes[h] == 0 || + head_offsets[h] > INT32_MAX || + head_vocab_sizes[h] > INT32_MAX || + head_offsets[h] + head_vocab_sizes[h] > INT32_MAX) { + throw std::runtime_error(format("PLE head %u range does not fit the int32 row index", h)); + } + hparams.ple_head_offsets[h] = (uint32_t) head_offsets[h]; + hparams.ple_head_vocab_sizes[h] = (uint32_t) head_vocab_sizes[h]; + } + } + + // linear attention everywhere except every full_attention_interval-th layer + if (!ml.get_key_or_arr(LLM_KV_ATTENTION_RECURRENT_LAYERS, hparams.is_recr_impl, hparams.n_layer_all, false)) { + uint32_t full_attn_interval = 4; + ml.get_key(LLM_KV_FULL_ATTENTION_INTERVAL, full_attn_interval, false); + qwen4exp_require_nonzero(ml, LLM_KV_FULL_ATTENTION_INTERVAL, full_attn_interval); + for (uint32_t i = 0; i < hparams.n_layer_all; ++i) { + hparams.is_recr_impl[i] = (i < hparams.n_layer()) && ((i + 1) % full_attn_interval != 0); + } + } + + // the PLE conv history is a row of the recurrent cache, which linear layers alone have + for (uint32_t i = 0; i < hparams.n_layer_all; ++i) { + if (hparams.is_ple(i) && !hparams.is_recr(i)) { + throw std::runtime_error(format("PLE layer %u is not a linear attention layer", i)); + } + } + + switch (hparams.n_layer()) { + case 48: type = LLM_TYPE_A3B; break; + default: type = LLM_TYPE_UNKNOWN; + } +} + +void llama_model_qwen4exp::load_arch_tensors(llama_model_loader & ml) { + LLAMA_LOAD_LOCALS; + + const int64_t hc = hparams.dsv4_hc_mult; + const int64_t hc_dim = hc * n_embd; + const int64_t hc_lr = hparams.hc_low_rank; + + const bool mtp_only = (hparams.n_layer_nextn > 0) && (ml.get_weight("blk.0.hc_attn_norm.weight") == nullptr); + const int trunk_flags = mtp_only ? TENSOR_NOT_REQUIRED : 0; + + tok_embd = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), { n_embd, n_vocab }, trunk_flags); + + hc_head_norm = create_tensor(tn(LLM_TENSOR_HC_HEAD_NORM, "weight"), { n_embd, hc }, trunk_flags | TENSOR_ALLOW_RESHAPE); + hc_head_down = create_tensor(tn(LLM_TENSOR_HC_HEAD_DOWN, "weight"), { hc_dim, hc_lr }, trunk_flags); + hc_head_up = create_tensor(tn(LLM_TENSOR_HC_HEAD_UP, "weight"), { hc_lr, hc_dim }, trunk_flags); + + output = create_tensor(tn(LLM_TENSOR_OUTPUT, "weight"), { n_embd, n_vocab }, TENSOR_NOT_REQUIRED); + // tie_word_embeddings is false here: never tie to a token_embd a borrowing draft lacks. + if (output == NULL && tok_embd != NULL) { + output = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), { n_embd, n_vocab }, TENSOR_DUPLICATED); + } + + // flat [ple_head_dim, n_rows] gather target + if (hparams.ple_n_heads > 0) { + // the head ranges are what the gather indexes, so they set the minimum row count + int64_t ple_rows = 0; + for (uint32_t h = 0; h < hparams.ple_n_heads; ++h) { + ple_rows = std::max(ple_rows, (int64_t) hparams.ple_head_offsets[h] + hparams.ple_head_vocab_sizes[h]); + } + + // the converter pads the table; a model synthesised from metadata has no tensor to ask + const std::string ple_name = tn(LLM_TENSOR_PER_LAYER_TOKEN_EMBD, "weight").str(); + if (const auto * ple_w = ml.get_weight(ple_name.c_str())) { + if (ple_w->tensor->ne[1] < ple_rows) { + throw std::runtime_error(format("%s has %" PRId64 " rows, too few for the PLE head ranges (%" PRId64 ")", + ple_name.c_str(), ple_w->tensor->ne[1], ple_rows)); + } + ple_rows = ple_w->tensor->ne[1]; + } + + per_layer_tok_embd = create_tensor(tn(LLM_TENSOR_PER_LAYER_TOKEN_EMBD, "weight"), + { hparams.ple_head_dim, ple_rows }, TENSOR_READ_LAZY); + } + + const int mtp_flags = !ml.load_mtp ? TENSOR_SKIP : 0; + + for (int il = 0; il < (int) hparams.n_layer_all; ++il) { + auto & layer = layers[il]; + + const int flags = il < n_layer ? trunk_flags : mtp_flags; + + const int64_t n_ff_exp = hparams.n_ff_exp ? hparams.n_ff_exp : n_ff / n_expert_used; + const int64_t n_ff_shexp = hparams.n_ff_shexp ? hparams.n_ff_shexp : n_ff; + + const int64_t head_k_dim = hparams.ssm_d_state; + const int64_t head_v_dim = hparams.ssm_d_state; + const int64_t n_k_heads = hparams.ssm_n_group; + const int64_t n_v_heads = hparams.ssm_dt_rank; + const int64_t key_dim = head_k_dim * n_k_heads; + const int64_t value_dim = head_v_dim * n_v_heads; + const int64_t conv_dim = key_dim * 2 + value_dim; + + // two HC modules per layer: before the token mixer, before the MoE + layer.hc_attn_norm = create_tensor(tn(LLM_TENSOR_HC_ATTN_NORM, "weight", il), { n_embd, hc }, flags | TENSOR_ALLOW_RESHAPE); + layer.hc_attn_down = create_tensor(tn(LLM_TENSOR_HC_ATTN_DOWN, "weight", il), { hc_dim, hc_lr }, flags); + layer.hc_attn_up = create_tensor(tn(LLM_TENSOR_HC_ATTN_UP, "weight", il), { hc_lr, hc_dim }, flags); + layer.hc_attn_inject = create_tensor(tn(LLM_TENSOR_HC_ATTN_INJECT, "weight", il), { hc_dim, hc }, flags); + layer.hc_ffn_norm = create_tensor(tn(LLM_TENSOR_HC_FFN_NORM, "weight", il), { n_embd, hc }, flags | TENSOR_ALLOW_RESHAPE); + layer.hc_ffn_down = create_tensor(tn(LLM_TENSOR_HC_FFN_DOWN, "weight", il), { hc_dim, hc_lr }, flags); + layer.hc_ffn_up = create_tensor(tn(LLM_TENSOR_HC_FFN_UP, "weight", il), { hc_lr, hc_dim }, flags); + layer.hc_ffn_inject = create_tensor(tn(LLM_TENSOR_HC_FFN_INJECT, "weight", il), { hc_dim, hc }, flags); + + if (!hparams.is_recr(il)) { + // full attention: wq holds [q|gate] interleaved per head + create_tensor_qkv(layer, il, n_embd, n_embd_head_k * n_head * 2, n_embd_k_gqa, n_embd_v_gqa, flags); + layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", il), { n_embd_head_k * n_head, n_embd }, flags); + + layer.attn_q_norm = create_tensor(tn(LLM_TENSOR_ATTN_Q_NORM, "weight", il), { n_embd_head_k }, flags); + layer.attn_k_norm = create_tensor(tn(LLM_TENSOR_ATTN_K_NORM, "weight", il), { n_embd_head_k }, flags); + + const int64_t idx_dim = hparams.indexer_head_size; + layer.index_q_proj = create_tensor(tn(LLM_TENSOR_INDEXER_Q_PROJ, "weight", il), { n_embd, hparams.indexer_n_head * idx_dim }, flags); + layer.index_k_proj = create_tensor(tn(LLM_TENSOR_INDEXER_K_PROJ, "weight", il), { n_embd, idx_dim }, flags); + layer.index_q_norm = create_tensor(tn(LLM_TENSOR_INDEXER_Q_NORM, "weight", il), { idx_dim }, flags); + layer.index_k_norm = create_tensor(tn(LLM_TENSOR_INDEXER_K_NORM, "weight", il), { idx_dim }, flags); + } else { + layer.wqkv = create_tensor(tn(LLM_TENSOR_ATTN_QKV, "weight", il), { n_embd, key_dim * 2 + value_dim }, flags); + layer.wqkv_gate = create_tensor(tn(LLM_TENSOR_ATTN_GATE, "weight", il), { n_embd, value_dim }, flags); + layer.ssm_conv1d = create_tensor(tn(LLM_TENSOR_SSM_CONV1D, "weight", il), { hparams.ssm_d_conv, conv_dim }, flags); + layer.ssm_dt = create_tensor(tn(LLM_TENSOR_SSM_DT, "bias", il), { hparams.ssm_dt_rank }, flags); + layer.ssm_a = create_tensor(tn(LLM_TENSOR_SSM_A_NOSCAN, il), { hparams.ssm_dt_rank }, flags); + layer.ssm_beta = create_tensor(tn(LLM_TENSOR_SSM_BETA, "weight", il), { n_embd, n_v_heads }, flags); + layer.ssm_alpha = create_tensor(tn(LLM_TENSOR_SSM_ALPHA, "weight", il), { n_embd, n_v_heads }, flags); + layer.ssm_norm = create_tensor(tn(LLM_TENSOR_SSM_NORM, "weight", il), { head_v_dim }, flags); + layer.ssm_out = create_tensor(tn(LLM_TENSOR_SSM_OUT, "weight", il), { value_dim, n_embd }, flags); + } + + if (hparams.is_ple(il)) { + layer.ple_key = create_tensor(tn(LLM_TENSOR_PLE_KEY, "weight", il), { n_embd, hc_dim }, flags); + layer.ple_value = create_tensor(tn(LLM_TENSOR_PLE_VALUE, "weight", il), { n_embd, n_embd }, flags); + layer.ple_norm_key = create_tensor(tn(LLM_TENSOR_PLE_NORM_KEY, "weight", il), { n_embd, hc }, flags | TENSOR_ALLOW_RESHAPE); + layer.ple_norm_query = create_tensor(tn(LLM_TENSOR_PLE_NORM_QUERY, "weight", il), { n_embd, hc }, flags | TENSOR_ALLOW_RESHAPE); + layer.ple_norm_conv = create_tensor(tn(LLM_TENSOR_PLE_NORM_CONV, "weight", il), { n_embd, hc }, flags | TENSOR_ALLOW_RESHAPE); + layer.ple_conv1d = create_tensor(tn(LLM_TENSOR_PLE_CONV1D, "weight", il), { hparams.ple_conv_kernel, hc_dim }, flags); + } + + layer.ffn_gate_inp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP, "weight", il), { n_embd, n_expert }, flags); + layer.ffn_down_exps = create_tensor(tn(LLM_TENSOR_FFN_DOWN_EXPS, "weight", il), { n_ff_exp, n_embd, n_expert }, flags); + create_tensor_gate_up_exps(layer, il, n_embd, n_ff_exp, n_expert, flags); + + layer.ffn_gate_inp_shexp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP_SHEXP, "weight", il), { n_embd }, flags); + layer.ffn_gate_shexp = create_tensor(tn(LLM_TENSOR_FFN_GATE_SHEXP, "weight", il), { n_embd, n_ff_shexp }, flags); + layer.ffn_up_shexp = create_tensor(tn(LLM_TENSOR_FFN_UP_SHEXP, "weight", il), { n_embd, n_ff_shexp }, flags); + layer.ffn_down_shexp = create_tensor(tn(LLM_TENSOR_FFN_DOWN_SHEXP, "weight", il), { n_ff_shexp, n_embd }, flags); + + if (il < n_layer) { + continue; + } + + layer.nextn.enorm = create_tensor(tn(LLM_TENSOR_NEXTN_ENORM, "weight", il), { n_embd }, flags); + layer.nextn.hnorm = create_tensor(tn(LLM_TENSOR_NEXTN_HNORM, "weight", il), { hc_dim }, flags); + layer.nextn.eh_proj = create_tensor(tn(LLM_TENSOR_NEXTN_EH_PROJ, "weight", il), { 2 * n_embd, n_embd }, flags); + + layer.nextn.hc_head_norm = create_tensor(tn(LLM_TENSOR_NEXTN_HC_HEAD_NORM, "weight", il), { n_embd, hc }, flags | TENSOR_ALLOW_RESHAPE); + layer.nextn.hc_head_down = create_tensor(tn(LLM_TENSOR_NEXTN_HC_HEAD_DOWN, "weight", il), { hc_dim, hc_lr }, flags); + layer.nextn.hc_head_up = create_tensor(tn(LLM_TENSOR_NEXTN_HC_HEAD_UP, "weight", il), { hc_lr, hc_dim }, flags); + + layer.nextn.embed_tokens = create_tensor(tn(LLM_TENSOR_NEXTN_EMBED_TOKENS, "weight", il), { n_embd, n_vocab }, flags | TENSOR_NOT_REQUIRED); + // reduced draft vocabulary (qwen4exp-draft-vocab.h): a K-row head plus d2t + const int64_t n_draft_rows = mtp_flags == 0 ? qwen4exp_draft_vocab_rows(ml, tn(LLM_TENSOR_D2T).str(), n_vocab) : 0; + if (n_draft_rows > 0) { + d2t = create_tensor(tn(LLM_TENSOR_D2T), { n_draft_rows }, 0); + } + layer.nextn.shared_head_head = create_tensor(tn(LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD, "weight", il), + { n_embd, n_draft_rows > 0 ? n_draft_rows : n_vocab }, + flags | (n_draft_rows > 0 ? 0 : TENSOR_NOT_REQUIRED)); + } +} + +std::unique_ptr llama_model_qwen4exp::build_arch_graph(const llm_graph_params & params) const { + if (params.gtype == LLM_GRAPH_TYPE_DECODER_MTP) { + return std::make_unique(*this, params); + } + // without this a self-contained draft loads, then walks the null trunk and segfaults. + if (hc_head_norm == nullptr) { + throw std::runtime_error("this model is an MTP draft head without a trunk; " + "load it as a draft of its target model (-md), not on its own"); + } + return std::make_unique(*this, params); +} + +// Hyper-connections keep hc parallel residual streams [n_embd, hc, T] in place of layer norms. +// Returns the mixed [n_embd, T] stream; `inject` gets the [hc, T] scatter weights. +ggml_tensor * llama_model_qwen4exp::graph::build_hc_mix( + ggml_tensor * x, + ggml_tensor * w_norm, + ggml_tensor * w_down, + ggml_tensor * w_up, + ggml_tensor * w_inject, + ggml_tensor ** inject, + int il) { + const int64_t hc = hparams.dsv4_hc_mult; + const int64_t hc_dim = hc * n_embd; + const int64_t nt = x->ne[2]; + + // grouped RMSNorm: reduce over one stream, then scale all streams with the [n_embd, hc] gamma + // the converter folded each gamma to (1 + w) + ggml_tensor * xn = ggml_mul(ctx0, ggml_rms_norm(ctx0, x, hparams.f_norm_rms_eps), w_norm); + xn = ggml_reshape_2d(ctx0, xn, hc_dim, nt); + cb(xn, "hc_norm", il); + + ggml_tensor * lo = build_lora_mm(w_down, xn); + lo = ggml_silu(ctx0, ggml_scale(ctx0, lo, 1.0f / (float) hc)); + ggml_tensor * gate = build_lora_mm(w_up, lo); + cb(gate, "hc_gate", il); + + ggml_tensor * mixed = nullptr; + if (cparams.fused_dsv4_hc_pre && il >= 0) { + // sigmoid gate and mean over the streams in one op + mixed = ggml_dsv4_hc_pre_gated(ctx0, + ggml_reshape_3d(ctx0, xn, n_embd, hc, nt), + ggml_reshape_3d(ctx0, gate, n_embd, hc, nt), 1.0f / (float) hc); + res->add_fused_node({LLM_FUSED_OP_DSV4_HC_PRE, mixed, il}); + } else { + ggml_tensor * gated = ggml_mul(ctx0, xn, ggml_sigmoid(ctx0, gate)); + gated = ggml_reshape_3d(ctx0, gated, n_embd, hc, nt); + + // collapse the streams by their mean + mixed = ggml_view_2d(ctx0, gated, n_embd, nt, + ggml_row_size(gated->type, n_embd) * hc, 0); + mixed = ggml_cont(ctx0, mixed); + for (int64_t c = 1; c < hc; ++c) { + ggml_tensor * s = ggml_view_2d(ctx0, gated, n_embd, nt, + ggml_row_size(gated->type, n_embd) * hc, + ggml_row_size(gated->type, n_embd) * c); + mixed = ggml_add(ctx0, mixed, s); + } + mixed = ggml_scale(ctx0, mixed, 1.0f / (float) hc); + } + cb(mixed, "hc_mixed", il); + + if (inject) { + *inject = build_lora_mm(w_inject, xn); + cb(*inject, "hc_inject", il); + } + + return mixed; +} + +ggml_tensor * llama_model_qwen4exp::graph::build_hc_combine( + ggml_tensor * residual, + ggml_tensor * block_out, + ggml_tensor * inject, + int il) { + const int64_t hc = hparams.dsv4_hc_mult; + const int64_t nt = residual->ne[2]; + + // 2*sigmoid centres the scatter weights on 1, so a zero injection is a plain residual add + ggml_tensor * w = ggml_sigmoid(ctx0, ggml_scale(ctx0, inject, 1.0f / (float) hc)); + w = ggml_scale(ctx0, w, 2.0f); + + ggml_tensor * cur = nullptr; + if (cparams.fused_dsv4_hc_post && il >= 0) { + // identity comb: every stream adds the same block output, scaled by its own weight + cur = ggml_dsv4_hc_post(ctx0, block_out, residual, w, nullptr); + res->add_fused_node({LLM_FUSED_OP_DSV4_HC_POST, cur, il}); + } else { + w = ggml_reshape_3d(ctx0, w, 1, hc, nt); + + ggml_tensor * b = ggml_reshape_3d(ctx0, block_out, n_embd, 1, nt); + b = ggml_repeat_4d(ctx0, b, n_embd, hc, nt, 1); + + cur = ggml_add(ctx0, residual, ggml_mul(ctx0, b, w)); + } + cb(cur, "hc_combine", il); + + return cur; +} + +llama_model_qwen4exp::graph::graph(const llama_model & model, const llm_graph_params & params) : + llm_build_delta_net_base(params), model(model) { + const int64_t hc = hparams.dsv4_hc_mult; + + GGML_ASSERT(hparams.n_embd_head_v() == hparams.n_embd_head_k()); + + int sections[4]; + std::copy(std::begin(hparams.rope_sections), std::begin(hparams.rope_sections) + 4, sections); + + ggml_tensor * inpL = build_inp_embd(model.tok_embd); + cb(inpL, "model.input_embed", -1); + ggml_build_forward_expand(gf, inpL); + + auto * inp = build_inp_mem_hybrid(); + + // qwen4exp always builds llama_memory_hybrid_idx, so this downcast is safe + // the indexer cache inside it is absent when the GGUF has no indexer tensors + const auto * mctx_hyb = static_cast(inp->mctx); + + const llama_kv_cache_context * mctx_idx = mctx_hyb->get_idx(); + if (mctx_idx) { + GGML_ASSERT(mctx_idx->get_n_kv() == inp->mctx->get_attn()->get_n_kv() && + "the indexer cache must track the attention cache cell for cell"); + } + + ggml_tensor * inp_pos = build_inp_pos(); + ggml_tensor * inp_out_ids = build_inp_out_ids(); + + ggml_tensor * ple_emb = nullptr; + if (hparams.ple_n_heads > 0) { + ple_emb = build_inp_ple(mctx_hyb); + // make sure ple_emb and build_inp_embd are in the same graph split + ggml_build_forward_expand(gf, ple_emb); + } + + // the wide residual starts as hc identical copies of the embedding + ggml_tensor * res_hc = ggml_repeat_4d(ctx0, + ggml_reshape_3d(ctx0, inpL, n_embd, 1, n_tokens), + n_embd, hc, n_tokens, 1); + cb(res_hc, "hc_init", -1); + + for (int il = 0; il < n_layer; ++il) { + res->t_layer_inp[il] = res_hc; + + if (hparams.is_ple(il)) { + res_hc = build_ple(inp->get_recr(), ple_emb, res_hc, il); + } + + ggml_tensor * inject = nullptr; + ggml_tensor * cur = build_hc_mix(res_hc, + model.layers[il].hc_attn_norm, + model.layers[il].hc_attn_down, + model.layers[il].hc_attn_up, + model.layers[il].hc_attn_inject, + &inject, il); + + ggml_build_forward_expand(gf, cur); + + if (hparams.is_recr(il)) { + cur = build_layer_attn_linear(inp->get_recr(), cur, il); + } else { + cur = build_layer_attn(inp->get_attn(), mctx_hyb, cur, inp_pos, sections, il); + } + + const bool gather_now = !cparams.embeddings_nextn || cparams.embeddings_nextn_masked; + + if (il == n_layer - 1 && inp_out_ids && gather_now) { + // everything below is per token, so drop the rows that produce no output + cur = ggml_get_rows(ctx0, cur, inp_out_ids); + inject = ggml_get_rows(ctx0, inject, inp_out_ids); + + res_hc = ggml_reshape_2d(ctx0, res_hc, n_embd*hc, res_hc->ne[2]); + res_hc = ggml_get_rows(ctx0, res_hc, inp_out_ids); + res_hc = ggml_reshape_3d(ctx0, res_hc, n_embd, hc, res_hc->ne[1]); + } + + res_hc = build_hc_combine(res_hc, cur, inject, il); + + cur = build_hc_mix(res_hc, + model.layers[il].hc_ffn_norm, + model.layers[il].hc_ffn_down, + model.layers[il].hc_ffn_up, + model.layers[il].hc_ffn_inject, + &inject, il); + + cur = build_layer_ffn(cur, il); + cb(cur, "ffn_out", il); + + res_hc = build_hc_combine(res_hc, cur, inject, il); + + // "l_last" is the layer output name that build_cvec and imatrix look for + cb(res_hc, "l_last", il); + } + + // export res_hc itself, never a reshape view: a pure view gets no backend assignment to read back. + if (cparams.embeddings_nextn) { + cb(res_hc, "h_nextn", -1); + res->t_h_nextn = res_hc; + + if (!cparams.embeddings_nextn_masked && inp_out_ids) { + res_hc = ggml_reshape_2d(ctx0, res_hc, n_embd*hc, res_hc->ne[2]); + res_hc = ggml_get_rows(ctx0, res_hc, inp_out_ids); + res_hc = ggml_reshape_3d(ctx0, res_hc, n_embd, hc, res_hc->ne[1]); + } + } + + // the final mixer is the output norm: there is no separate one + ggml_tensor * cur = build_hc_mix(res_hc, + model.hc_head_norm, model.hc_head_down, model.hc_head_up, + nullptr, nullptr, -1); + + cb(cur, "result_norm", -1); + res->t_embd = cur; + + cur = build_lora_mm(model.output, cur, model.output_s); + cb(cur, "result_output", -1); + res->t_logits = cur; + + ggml_build_forward_expand(gf, cur); +} + +// TODO: QSA for the draft head; dense is a numerical superset below the 2048-token budget. +llama_model_qwen4exp::graph_mtp::graph_mtp(const llama_model & model, const llm_graph_params & params) : + graph(model, params, no_build_t{}) { + GGML_ASSERT(hparams.n_layer_nextn > 0 && "QWEN4EXP MTP requires n_layer_nextn > 0"); + GGML_ASSERT(hparams.n_layer_nextn == 1 && "QWEN4EXP MTP currently only supports a single MTP block"); + GGML_ASSERT(ubatch.token && "QWEN4EXP MTP requires token input"); + + const int64_t hc = hparams.dsv4_hc_mult; + const int64_t hc_dim = hc * n_embd; + GGML_ASSERT(hparams.n_embd_out() == (uint32_t) hc_dim && "QWEN4EXP MTP hidden width mismatch"); + + const int il = hparams.n_layer() + cparams.nextn_layer_offset; + GGML_ASSERT(cparams.nextn_layer_offset >= 0 && + cparams.nextn_layer_offset < (int) hparams.n_layer_nextn && + "nextn_layer_offset out of range [0, n_layer_nextn)"); + const auto & layer = model.layers[il]; + + GGML_ASSERT(layer.nextn.eh_proj && "MTP block missing nextn.eh_proj"); + GGML_ASSERT(layer.nextn.enorm && "MTP block missing nextn.enorm"); + GGML_ASSERT(layer.nextn.hnorm && "MTP block missing nextn.hnorm"); + GGML_ASSERT(layer.nextn.hc_head_norm && "MTP block missing nextn.hc_head_norm"); + + int sections[4]; + std::copy(std::begin(hparams.rope_sections), std::begin(hparams.rope_sections) + 4, sections); + + auto inp = std::make_unique(hc_dim); + + inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens); + ggml_set_input(inp->tokens); + + inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hc_dim, n_tokens); + ggml_set_input(inp->embd); + + inp->h = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hc_dim, n_tokens); + ggml_set_input(inp->h); + ggml_set_name(inp->h, "mtp_h_input"); + + ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd; + if (tok_embd_w == nullptr) { + tok_embd_w = qwen4exp_shared_model(cparams, model, "token_embd.weight").tok_embd; + GGML_ASSERT(tok_embd_w && "QWEN4EXP MTP: the target model has no token embeddings to borrow"); + } + ggml_tensor * tok_embd = ggml_get_rows(ctx0, tok_embd_w, inp->tokens); + cb(tok_embd, "mtp_tok_embd", il); + + ggml_tensor * h_state = ggml_reshape_3d(ctx0, inp->h, n_embd, hc, n_tokens); + cb(h_state, "mtp_h_state", il); + + res->add_input(std::move(inp)); + + ggml_tensor * inp_pos = build_inp_pos(); + ggml_tensor * inp_out_ids = build_inp_out_ids(); + + auto * inp_attn = build_attn_inp_kv(); + + ggml_tensor * h_norm = ggml_rms_norm(ctx0, h_state, hparams.f_norm_rms_eps); + h_norm = ggml_reshape_2d(ctx0, h_norm, hc_dim, n_tokens); + h_norm = ggml_mul(ctx0, h_norm, layer.nextn.hnorm); + h_norm = ggml_reshape_3d(ctx0, h_norm, n_embd, hc, n_tokens); + cb(h_norm, "mtp_hnorm", il); + + ggml_tensor * e_norm = build_norm(tok_embd, layer.nextn.enorm, nullptr, LLM_NORM_RMS, il); + e_norm = ggml_repeat_4d(ctx0, + ggml_reshape_3d(ctx0, e_norm, n_embd, 1, n_tokens), + n_embd, hc, n_tokens, 1); + cb(e_norm, "mtp_enorm", il); + + // per stream, not pooled: pooling before the projection discards the hyper-connection residual. + ggml_tensor * concat = ggml_concat(ctx0, e_norm, h_norm, /*dim=*/ 0); + cb(concat, "mtp_concat", il); + + ggml_tensor * res_hc = build_lora_mm(layer.nextn.eh_proj, concat, layer.nextn.eh_proj_s); + cb(res_hc, "mtp_eh_proj", il); + + ggml_tensor * inject = nullptr; + ggml_tensor * cur = build_hc_mix(res_hc, + layer.hc_attn_norm, layer.hc_attn_down, layer.hc_attn_up, layer.hc_attn_inject, + &inject, il); + cb(cur, "mtp_hc_attn_pre", il); + + const int64_t n_embd_head = hparams.n_embd_head_v(); + GGML_ASSERT(n_embd_head == hparams.n_embd_head_k()); + + ggml_tensor * Qcur_full = build_lora_mm(layer.wq, cur, layer.wq_s); + cb(Qcur_full, "mtp_Qcur_full", il); + + ggml_tensor * Qcur = ggml_view_3d(ctx0, Qcur_full, n_embd_head, n_head, n_tokens, + ggml_element_size(Qcur_full) * n_embd_head * 2, + ggml_element_size(Qcur_full) * n_embd_head * 2 * n_head, 0); + Qcur = build_norm(Qcur, layer.attn_q_norm, nullptr, LLM_NORM_RMS, il); + cb(Qcur, "mtp_Qcur_normed", il); + + ggml_tensor * gate = ggml_view_3d(ctx0, Qcur_full, n_embd_head, n_head, n_tokens, + ggml_element_size(Qcur_full) * n_embd_head * 2, + ggml_element_size(Qcur_full) * n_embd_head * 2 * n_head, + ggml_element_size(Qcur_full) * n_embd_head); + gate = ggml_cont_2d(ctx0, gate, n_embd_head * n_head, n_tokens); + cb(gate, "mtp_gate", il); + + ggml_tensor * Kcur = build_lora_mm(layer.wk, cur, layer.wk_s); + Kcur = ggml_reshape_3d(ctx0, Kcur, n_embd_head, n_head_kv, n_tokens); + Kcur = build_norm(Kcur, layer.attn_k_norm, nullptr, LLM_NORM_RMS, il); + cb(Kcur, "mtp_Kcur_normed", il); + + ggml_tensor * Vcur = build_lora_mm(layer.wv, cur, layer.wv_s); + Vcur = ggml_reshape_3d(ctx0, Vcur, n_embd_head, n_head_kv, n_tokens); + cb(Vcur, "mtp_Vcur", il); + + Qcur = ggml_rope_multi(ctx0, Qcur, inp_pos, nullptr, + n_rot, sections, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow); + Kcur = ggml_rope_multi(ctx0, Kcur, inp_pos, nullptr, + n_rot, sections, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow); + cb(Qcur, "mtp_Qcur", il); + cb(Kcur, "mtp_Kcur", il); + + const float kq_scale = hparams.f_attention_scale == 0.0f + ? 1.0f / sqrtf(float(n_embd_head)) : hparams.f_attention_scale; + + cur = build_attn(inp_attn, + nullptr, nullptr, nullptr, + Qcur, Kcur, Vcur, nullptr, nullptr, nullptr, kq_scale, il); + cb(cur, "mtp_attn_pregate", il); + + cur = ggml_mul(ctx0, cur, ggml_sigmoid(ctx0, gate)); + cb(cur, "mtp_attn_gated", il); + + cur = build_lora_mm(layer.wo, cur, layer.wo_s); + cb(cur, "mtp_attn_out", il); + + if (inp_out_ids) { + cur = ggml_get_rows(ctx0, cur, inp_out_ids); + inject = ggml_get_rows(ctx0, inject, inp_out_ids); + + res_hc = ggml_reshape_2d(ctx0, res_hc, hc_dim, res_hc->ne[2]); + res_hc = ggml_get_rows(ctx0, res_hc, inp_out_ids); + res_hc = ggml_reshape_3d(ctx0, res_hc, n_embd, hc, res_hc->ne[1]); + } + + res_hc = build_hc_combine(res_hc, cur, inject, il); + cb(res_hc, "mtp_hc_attn_post", il); + + cur = build_hc_mix(res_hc, + layer.hc_ffn_norm, layer.hc_ffn_down, layer.hc_ffn_up, layer.hc_ffn_inject, + &inject, il); + cb(cur, "mtp_hc_ffn_pre", il); + + cur = build_layer_ffn(cur, il); + cb(cur, "mtp_ffn_out", il); + + res_hc = build_hc_combine(res_hc, cur, inject, il); + cb(res_hc, "mtp_hc_ffn_post", il); + + // the next draft step re-enters here, so export the wide stream before it is collapsed. + cb(res_hc, "h_nextn", -1); + res->t_h_nextn = res_hc; + + cur = build_hc_mix(res_hc, + layer.nextn.hc_head_norm, layer.nextn.hc_head_down, layer.nextn.hc_head_up, + nullptr, nullptr, -1); + cb(cur, "mtp_hc_head", -1); + + // no res->t_embd: it is n_embd wide, but the context sizes that buffer by n_embd_out. + + ggml_tensor * head_w = layer.nextn.shared_head_head ? layer.nextn.shared_head_head : model.output; + ggml_tensor * head_s = layer.nextn.shared_head_head ? layer.nextn.shared_head_head_s : model.output_s; + if (head_w == nullptr) { + const llama_model & other = qwen4exp_shared_model(cparams, model, "output.weight"); + head_w = other.output; + head_s = other.output_s; + GGML_ASSERT(head_w && "QWEN4EXP MTP: the target model has no LM head to borrow"); + } + + cur = build_lora_mm(head_w, cur, head_s); + if (model.d2t) { + cur = qwen4exp_draft_vocab_scatter(ctx0, cur, model.d2t, (int64_t) model.vocab.n_tokens()); + } + cb(cur, "result_output", -1); + res->t_logits = cur; + + ggml_build_forward_expand(gf, cur); +} + +std::pair llama_model_qwen4exp::graph::build_qkvz( + ggml_tensor * input, + int il) { + const int64_t n_seqs = ubatch.n_seqs; + const int64_t n_seq_tokens = ubatch.n_seq_tokens; + + ggml_tensor * qkv_mixed = build_lora_mm(model.layers[il].wqkv, input, model.layers[il].wqkv_s); + qkv_mixed = ggml_reshape_3d(ctx0, qkv_mixed, qkv_mixed->ne[0], n_seq_tokens, n_seqs); + cb(qkv_mixed, "linear_attn_qkv_mixed", il); + + ggml_tensor * z = build_lora_mm(model.layers[il].wqkv_gate, input, model.layers[il].wqkv_gate_s); + cb(z, "z", il); + + return { qkv_mixed, z }; +} + +ggml_tensor * llama_model_qwen4exp::graph::build_norm_gated( + ggml_tensor * input, + ggml_tensor * weights, + ggml_tensor * gate, + int layer) { + // the one numerical difference from Qwen3.5's GDN: sigmoid output gate, not silu + ggml_tensor * normalized = build_norm(input, weights, nullptr, LLM_NORM_RMS, layer); + ggml_tensor * gated = ggml_sigmoid(ctx0, gate); + + return ggml_mul(ctx0, normalized, gated); +} + +// QSA attends to a budget of whole blocks of compress_ratio tokens, plus the incomplete tail +// one mean-pooled indexer key scores each block; set_input resolves the cache layout +class llama_model_qwen4exp::llm_graph_input_qsa : public llm_graph_input_i { +public: + llm_graph_input_qsa(const llama_memory_hybrid_idx_context * mctx, uint32_t ratio, bool blk_bias, bool dense = false, int64_t width_max = 0) : + mctx(mctx), ratio(ratio), blk_bias(blk_bias), dense(dense), width_max(width_max) {} + virtual ~llm_graph_input_qsa() = default; + + void set_input(const llama_ubatch * ubatch) override { + mctx->get_idx()->set_input_k_idxs(k_idxs, ubatch); + if (!dense) { + mctx->set_input_qsa(cell_blk, blk_cells, blk_pos, bias, ubatch, ratio, blk_bias); + } + } + + bool can_reuse(const llm_graph_params & params) override { + mctx = static_cast(params.mctx); + + const auto * idx = mctx->get_idx(); + if (idx == nullptr) { + return false; + } + + const int64_t n_kv = idx->get_n_kv(); + const int64_t n_stream = mctx->get_n_stream(); + const int64_t n_blocks = (n_kv + ratio - 1)/ratio; + + bool res = true; + + res &= params.ubatch.n_tokens % n_stream == 0; + + res &= k_idxs->ne[0] == params.ubatch.n_tokens; + if (dense) { + return res && n_kv <= width_max; + } + res &= n_kv > width_max; + res &= cell_blk->ne[0] == n_kv; + res &= cell_blk->ne[1] == n_stream; + res &= blk_cells->ne[0] == (int64_t) ratio*n_blocks; + res &= blk_pos->ne[0] == 4*n_blocks*n_stream; + res &= bias->ne[0] == (blk_bias ? n_blocks : n_kv); + res &= bias->ne[1] == params.ubatch.n_tokens/n_stream; + + return res; + } + + // per stream: a cell index names a different token in each stream + ggml_tensor * k_idxs = nullptr; // I32 [n_tokens] + ggml_tensor * cell_blk = nullptr; // I32 [n_kv, n_stream] + ggml_tensor * blk_cells = nullptr; // I32 [ratio*n_blocks, n_stream] + ggml_tensor * blk_pos = nullptr; // I32 [4*n_blocks*n_stream] + ggml_tensor * bias = nullptr; // F32 [n_blocks or n_kv, n_tokens/n_stream, n_stream] + + const llama_memory_hybrid_idx_context * mctx; + const uint32_t ratio; + + // the per-cell half of the bias is the attention mask, so only the per-block half is uploaded + const bool blk_bias; + + // dense: every cell is selected (n_kv <= indexer_top_k + ratio - 1), so only the indexer keys are written + const bool dense; + const int64_t width_max; +}; + +ggml_tensor * llama_model_qwen4exp::graph::build_qsa_top_k( + const llama_memory_hybrid_idx_context * mctx_hyb, + ggml_tensor * cur, + ggml_tensor * inp_pos, + ggml_tensor * kq_mask, + int * sections, + int il) { + const llama_kv_cache_context * mctx_idx = mctx_hyb->get_idx(); + + const int64_t idx_dim = hparams.indexer_head_size; + const int64_t n_idx_h = hparams.indexer_n_head; + const int64_t r = hparams.dsv4_compress_ratios[il]; + const int64_t n_kv = mctx_idx->get_n_kv(); + + GGML_ASSERT(r > 0); + + const int64_t n_blocks = (n_kv + r - 1)/r; + + // build_attn_qsa and the KQ mask need the tokens to divide evenly across the streams + const int64_t n_stream = mctx_hyb->get_n_stream(); + GGML_ASSERT(n_tokens % n_stream == 0); + const int64_t n_tps = n_tokens/n_stream; + + // only the "which block is visible" half of the bias varies per block + // the rest is the visible/not test the attention mask already carries, so upload the per-block half only: 1/ratio of the cells + // alibi writes distances instead of a mask and non-causal keeps future cells, so both opt out + // the mask also holds an mrope rule for the query's own position, but only 2d image positions can differ there + const bool blk_bias = kq_mask != nullptr && + kq_mask->ne[0] == n_kv && kq_mask->ne[1] == n_tps && kq_mask->ne[3] == n_stream && + cparams.causal_attn && !hparams.use_alibi; + + // top-k keeps min(n_kv, indexer_top_k + r - 1) cells. When that is every cell and the mask covers the same + // cells, the top-k mask is 0 + kq_mask = kq_mask: plain dense attention gives the same result, so only the + // indexer keys are written (later, longer contexts score them). LLAMA_QSA_DENSE=0 keeps the top-k graph. + const int64_t width_max = (int64_t) hparams.indexer_top_k + r - 1; + static const bool dense_ok = [] { const char * e = getenv("LLAMA_QSA_DENSE"); return e == nullptr || atoi(e) != 0; }(); + const bool dense = dense_ok && n_kv <= width_max && kq_mask != nullptr && kq_mask->ne[0] == n_kv; + + // nothing above depends on the layer, so the layers sharing a ratio share one input set + llm_graph_input_qsa * inp = nullptr; + + const auto it = qsa_inps.find((uint32_t) r); + if (it != qsa_inps.end()) { + inp = it->second; + } else if (dense) { + auto qsa = std::make_unique(mctx_hyb, (uint32_t) r, blk_bias, true, width_max); + qsa->k_idxs = mctx_idx->build_input_k_idxs(ctx0, ubatch); + inp = qsa.get(); + res->add_input(std::move(qsa)); + qsa_inps.emplace((uint32_t) r, inp); + } else { + auto qsa = std::make_unique(mctx_hyb, (uint32_t) r, blk_bias, false, width_max); + + qsa->k_idxs = mctx_idx->build_input_k_idxs(ctx0, ubatch); + qsa->cell_blk = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, n_kv, n_stream); + qsa->blk_cells = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, r*n_blocks, n_stream); + qsa->blk_pos = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, 4*n_blocks*n_stream); + qsa->bias = ggml_new_tensor_3d(ctx0, GGML_TYPE_F32, blk_bias ? n_blocks : n_kv, n_tps, n_stream); + + ggml_set_input(qsa->cell_blk); + ggml_set_input(qsa->blk_cells); + ggml_set_input(qsa->blk_pos); + ggml_set_input(qsa->bias); + + inp = qsa.get(); + res->add_input(std::move(qsa)); + qsa_inps.emplace((uint32_t) r, inp); + } + + // cached indexer keys are raw: pooling precedes norm and rotation, so apply neither + ggml_tensor * k_raw = build_lora_mm(model.layers[il].index_k_proj, cur); + k_raw = ggml_reshape_3d(ctx0, k_raw, idx_dim, 1, n_tokens); + cb(k_raw, "indexer_k_raw", il); + + ggml_build_forward_expand(gf, mctx_idx->cpy_k(ctx0, k_raw, inp->k_idxs, il)); + + if (inp->dense) { + return nullptr; + } + + // one key head, so rows are contiguous. get_k gives [idx_dim, n_head_kv, n_kv, n_stream]. + ggml_tensor * k_all = mctx_idx->get_k(ctx0, il); + k_all = ggml_view_3d(ctx0, k_all, idx_dim, n_kv, n_stream, k_all->nb[2], k_all->nb[3], 0); + + // gathers per stream: blk_cells row s indexes stream s's own cells + ggml_tensor * members = ggml_get_rows(ctx0, k_all, inp->blk_cells); + members = ggml_reshape_4d(ctx0, members, idx_dim, r, n_blocks, n_stream); + + // mean over the block members; r is small, so summing slices beats a transpose plus sum_rows + ggml_tensor * pooled = nullptr; + for (int64_t i = 0; i < r; ++i) { + ggml_tensor * slice = ggml_cont(ctx0, + ggml_view_3d(ctx0, members, idx_dim, n_blocks, n_stream, + members->nb[2], members->nb[3], i*members->nb[1])); + pooled = pooled ? ggml_add(ctx0, pooled, slice) : slice; + } + pooled = ggml_scale(ctx0, pooled, 1.0f/(float) r); + cb(pooled, "indexer_k_pooled", il); + + // rope wants [n_dims, n_head, n_tokens]: lay every stream's blocks flat, split after. + // HRX: norm, weight and rope on the same [idx_dim, 1, blocks] shape, so ggml-hrx's fused + // rms_norm + mul + rope dispatch takes them (it has no standalone imrope). + pooled = ggml_reshape_3d(ctx0, pooled, idx_dim, 1, n_blocks*n_stream); + pooled = build_norm(pooled, model.layers[il].index_k_norm, nullptr, LLM_NORM_RMS, il); + pooled = ggml_rope_multi(ctx0, pooled, inp->blk_pos, nullptr, + n_rot, sections, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow); + pooled = ggml_reshape_3d(ctx0, pooled, idx_dim, n_blocks, n_stream); + cb(pooled, "indexer_k", il); + + ggml_tensor * q = build_lora_mm(model.layers[il].index_q_proj, cur); + q = ggml_reshape_3d(ctx0, q, idx_dim, n_idx_h, n_tokens); + q = build_norm(q, model.layers[il].index_q_norm, nullptr, LLM_NORM_RMS, il); + q = ggml_rope_multi(ctx0, q, inp_pos, nullptr, + n_rot, sections, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow); + cb(q, "indexer_q", il); + + // rectify each head dot product before the sum, as in the DeepSeek lightning indexer + // mul_mat matches ne[2], so the queries of stream s only meet the blocks of stream s + ggml_tensor * score = ggml_mul_mat(ctx0, pooled, + ggml_reshape_3d(ctx0, q, idx_dim, n_idx_h*n_tps, n_stream)); + score = ggml_reshape_4d(ctx0, score, n_blocks, n_idx_h, n_tps, n_stream); + score = ggml_relu(ctx0, score); + + // the heads sit side by side on ne[1] and there are only a few of them + ggml_tensor * summed = nullptr; + for (int64_t h = 0; h < n_idx_h; ++h) { + ggml_tensor * slice = ggml_view_3d(ctx0, score, n_blocks, n_tps, n_stream, + score->nb[2], score->nb[3], h*score->nb[1]); + summed = summed ? ggml_add(ctx0, summed, slice) : ggml_cont(ctx0, slice); + } + + score = summed; + cb(score, "indexer_score", il); + + // one value per block, so it is cheaper to bias here than after the cells are expanded + if (blk_bias) { + score = ggml_add(ctx0, score, inp->bias); + } + + // every token of a block gets the block score; the budget is whole blocks, so top-k cuts on a block boundary + ggml_tensor * expanded = ggml_get_rows(ctx0, + ggml_cont(ctx0, ggml_permute(ctx0, score, 1, 0, 2, 3)), inp->cell_blk); + expanded = ggml_cont(ctx0, ggml_permute(ctx0, expanded, 1, 0, 2, 3)); + + if (blk_bias) { + // flash attention keeps the mask in f16; the scores are f32 + ggml_tensor * mask = kq_mask->type == GGML_TYPE_F32 ? kq_mask : ggml_cast(ctx0, kq_mask, GGML_TYPE_F32); + expanded = ggml_add(ctx0, expanded, ggml_reshape_3d(ctx0, mask, n_kv, n_tps, n_stream)); + } else { + expanded = ggml_add(ctx0, expanded, inp->bias); + } + cb(expanded, "indexer_score_tokens", il); + + // the reference returns indexer_top_k + compress_ratio - 1: whole blocks plus the tail + const int64_t width = std::min(n_kv, (int64_t) hparams.indexer_top_k + r - 1); + + ggml_tensor * top_k = ggml_cont(ctx0, ggml_top_k(ctx0, expanded, width)); + + // build_attn_qsa reads [n_top_k, n_batch, 1, n_stream], matching the KQ mask. + top_k = ggml_reshape_4d(ctx0, top_k, width, n_tps, 1, n_stream); + cb(top_k, "indexer_top_k", il); + + return top_k; +} + +// Dense GQA self-attention restricted to the cells that top_k names. +// The mask build below copies the MLA sparse path in llm_graph_context::build_attn. +ggml_tensor * llama_model_qwen4exp::graph::build_attn_qsa( + llm_graph_input_attn_kv * inp, + ggml_tensor * q_cur, + ggml_tensor * k_cur, + ggml_tensor * v_cur, + ggml_tensor * top_k, + float kq_scale, + int il) { + // rotate q/k/v before they reach a quantized cache, as the dense path does. the indexer + // has already scored with its own query in build_qsa_top_k, so top_k is unaffected. + if (inp->self_k_rot) { + q_cur = llama_mul_mat_hadamard(ctx0, q_cur, inp->self_k_rot); + k_cur = llama_mul_mat_hadamard(ctx0, k_cur, inp->self_k_rot); + } + + if (inp->self_v_rot) { + v_cur = llama_mul_mat_hadamard(ctx0, v_cur, inp->self_v_rot); + } + + // these nodes are added to the graph together so that they are not reordered + // by doing so, the number of splits in the graph is reduced + // expand k later to enable rope fusion which directly writes into k-v cache + ggml_build_forward_expand(gf, q_cur); + ggml_build_forward_expand(gf, v_cur); + ggml_build_forward_expand(gf, k_cur); + + const auto * mctx_cur = inp->mctx; + + // store to KV cache + { + const auto & k_idxs = inp->get_k_idxs(); + const auto & v_idxs = inp->get_v_idxs(); + + ggml_build_forward_expand(gf, mctx_cur->cpy_k(ctx0, k_cur, k_idxs, il)); + ggml_build_forward_expand(gf, mctx_cur->cpy_v(ctx0, v_cur, v_idxs, il)); + } + + ggml_tensor * kq_mask = inp->get_kq_mask(); + + // prepare new kq mask - starts filled with -INFINITY + ggml_tensor * kq_mask_all = ggml_fill(ctx0, kq_mask, -INFINITY); + + // reshape KQ mask into tensor with rows of size 1: + // [n_kv, n_batch, 1, n_stream] -> [1, n_kv, n_batch, n_stream] + kq_mask_all = ggml_view_4d(ctx0, kq_mask_all, 1, kq_mask_all->ne[0], kq_mask_all->ne[1], kq_mask_all->ne[3], kq_mask_all->nb[0], kq_mask_all->nb[1], kq_mask_all->nb[2], 0); + + // reshape top_k indices: [n_top_k, n_batch, 1, n_stream] -> [n_top_k, n_batch, n_stream, 1] + ggml_tensor * top_k_3d = ggml_view_4d(ctx0, top_k, top_k->ne[0], top_k->ne[1], top_k->ne[3], 1, top_k->nb[1], top_k->nb[2], top_k->ne[3]*top_k->nb[3], 0); + + // prepare zero-filled tensor with rows of size 1: [1, n_top_k, n_batch, n_stream] + // this will be our source of zero values for unmasking top k mask elements + ggml_tensor * zeros = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, 1, top_k_3d->ne[0], top_k_3d->ne[1], top_k_3d->ne[2]); + zeros = ggml_fill(ctx0, zeros, 0.0f); + + // modify KQ mask by unmasking elements that are in top_k indices + // ggml_set_rows([1, n_kv, n_batch, n_stream], [1, n_top_k, n_batch, n_stream], [n_top_k, n_batch, n_stream, 1]) + ggml_tensor * kq_mask_top_k = ggml_set_rows(ctx0, kq_mask_all, zeros, top_k_3d); + + // reshape to restore the original shape of KQ mask: + // [1, n_kv, n_batch, n_stream] -> [n_kv, n_batch, 1, n_stream] + kq_mask_top_k = ggml_view_4d(ctx0, kq_mask_top_k, kq_mask_top_k->ne[1], kq_mask_top_k->ne[2], 1, kq_mask_top_k->ne[3], kq_mask_top_k->nb[2], kq_mask_top_k->nb[3], kq_mask_top_k->nb[3], 0); + + // combine with the original kq mask + kq_mask_top_k = ggml_add(ctx0, kq_mask_top_k, kq_mask); + + ggml_tensor * q = q_cur; + ggml_tensor * k = mctx_cur->get_k(ctx0, il); + ggml_tensor * v = mctx_cur->get_v(ctx0, il); + + ggml_tensor * cur = build_attn_mha(q, k, v, nullptr, kq_mask_top_k, nullptr, nullptr, kq_scale, il); + cb(cur, "kqv_out", il); + + // the rotation is its own inverse, so undo it on the value side of the output + if (inp->self_v_rot) { + cur = llama_mul_mat_hadamard(ctx0, cur, inp->self_v_rot); + } + + return cur; +} + +ggml_tensor * llama_model_qwen4exp::graph::build_layer_attn( + llm_graph_input_attn_kv * inp, + const llama_memory_hybrid_idx_context * mctx_hyb, + ggml_tensor * cur, + ggml_tensor * inp_pos, + int * sections, + int il) { + const int64_t n_embd_head = hparams.n_embd_head_v(); + GGML_ASSERT(n_embd_head == hparams.n_embd_head_k()); + + // indexer reads the same block input as q/k/v; no cache or no ratio means dense + const bool qsa = mctx_hyb->get_idx() != nullptr && hparams.dsv4_compress_ratios[il] > 0; + + ggml_tensor * top_k = qsa ? build_qsa_top_k(mctx_hyb, cur, inp_pos, inp->get_kq_mask(), sections, il) : nullptr; + + // Qwen3Next uses a single Q projection that outputs query + gate + ggml_tensor * Qcur_full = build_lora_mm(model.layers[il].wq, cur, model.layers[il].wq_s); // [ (n_embd_head * 2) * n_head, n_tokens ] + cb(Qcur_full, "Qcur_full", il); + + ggml_tensor * Qcur = ggml_view_3d(ctx0, Qcur_full, n_embd_head, n_head, n_tokens, + ggml_element_size(Qcur_full) * n_embd_head * 2, + ggml_element_size(Qcur_full) * n_embd_head * 2 * n_head, 0); + cb(Qcur, "Qcur_reshaped", il); + + Qcur = build_norm(Qcur, model.layers[il].attn_q_norm, nullptr, LLM_NORM_RMS, il); + cb(Qcur, "Qcur_normed", il); + + ggml_tensor * Kcur = build_lora_mm(model.layers[il].wk, cur, model.layers[il].wk_s); + cb(Kcur, "Kcur", il); + + ggml_tensor * Vcur = build_lora_mm(model.layers[il].wv, cur, model.layers[il].wv_s); + cb(Vcur, "Vcur", il); + + Kcur = ggml_reshape_3d(ctx0, Kcur, n_embd_head, n_head_kv, n_tokens); + Kcur = build_norm(Kcur, model.layers[il].attn_k_norm, nullptr, LLM_NORM_RMS, il); + cb(Kcur, "Kcur_normed", il); + + ggml_tensor * gate = ggml_view_3d(ctx0, Qcur_full, n_embd_head, n_head, n_tokens, + ggml_element_size(Qcur_full) * n_embd_head * 2, + ggml_element_size(Qcur_full) * n_embd_head * 2 * n_head, + ggml_element_size(Qcur_full) * n_embd_head); + gate = ggml_cont_2d(ctx0, gate, n_embd_head * n_head, n_tokens); + cb(gate, "gate_reshaped", il); + + Vcur = ggml_reshape_3d(ctx0, Vcur, n_embd_head, n_head_kv, n_tokens); + + // Apply IMRoPE + Qcur = ggml_rope_multi( + ctx0, Qcur, inp_pos, nullptr, + n_rot, sections, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow + ); + + Kcur = ggml_rope_multi( + ctx0, Kcur, inp_pos, nullptr, + n_rot, sections, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow + ); + + cb(Qcur, "Qcur", il); + cb(Kcur, "Kcur", il); + cb(Vcur, "Vcur", il); + + const float kq_scale = hparams.f_attention_scale == 0.0f ? 1.0f / sqrtf(float(n_embd_head)) : hparams.f_attention_scale; + + if (top_k) { + cur = build_attn_qsa(inp, Qcur, Kcur, Vcur, top_k, kq_scale, il); + } else { + cur = build_attn(inp, + nullptr, nullptr, nullptr, + Qcur, Kcur, Vcur, nullptr, nullptr, nullptr, kq_scale, il); + } + cb(cur, "attn_pregate", il); + + ggml_tensor * gate_sigmoid = ggml_sigmoid(ctx0, gate); + cb(gate_sigmoid, "gate_sigmoid", il); + + cur = ggml_mul(ctx0, cur, gate_sigmoid); + cb(cur, "attn_gated", il); + + cur = build_lora_mm(model.layers[il].wo, cur, model.layers[il].wo_s); + cb(cur, "attn_output", il); + + return cur; +} + +ggml_tensor * llama_model_qwen4exp::graph::build_layer_attn_linear( + llm_graph_input_rs * inp, + ggml_tensor * cur, + int il) { + const auto * mctx_cur = inp->mctx; + + const int64_t d_inner = hparams.ssm_d_inner; + const int64_t n_seqs = ubatch.n_seqs; + const int64_t head_k_dim = hparams.ssm_d_state; + const int64_t num_k_heads = hparams.ssm_n_group; + const int64_t num_v_heads = hparams.ssm_dt_rank; + const int64_t head_v_dim = hparams.ssm_d_state; + const int64_t n_seq_tokens = ubatch.n_seq_tokens; + + GGML_ASSERT(n_seqs != 0); + GGML_ASSERT(ubatch.equal_seqs()); + GGML_ASSERT(ubatch.n_tokens == n_seq_tokens * n_seqs); + GGML_ASSERT(head_v_dim * num_v_heads == d_inner); + + auto qkvz = build_qkvz(cur, il); + ggml_tensor * qkv_mixed = qkvz.first; + ggml_tensor * z = qkvz.second; + + ggml_tensor * beta = build_lora_mm(model.layers[il].ssm_beta, cur, model.layers[il].ssm_beta_s); + beta = ggml_reshape_4d(ctx0, beta, 1, num_v_heads, n_seq_tokens, n_seqs); + cb(beta, "beta", il); + + beta = ggml_sigmoid(ctx0, beta); + cb(beta, "beta_sigmoid", il); + + ggml_tensor * alpha = build_lora_mm(model.layers[il].ssm_alpha, cur, model.layers[il].ssm_alpha_s); + alpha = ggml_reshape_3d(ctx0, alpha, num_v_heads, n_seq_tokens, n_seqs); + cb(alpha, "alpha", il); + + ggml_tensor * alpha_biased = ggml_add(ctx0, alpha, model.layers[il].ssm_dt); + ggml_tensor * alpha_softplus = ggml_softplus(ctx0, alpha_biased); + cb(alpha_softplus, "a_softplus", il); + + ggml_tensor * gate = ggml_mul(ctx0, alpha_softplus, model.layers[il].ssm_a); // -A_log.exp() * softplus + cb(gate, "gate", il); + + gate = ggml_reshape_4d(ctx0, gate, 1, num_v_heads, n_seq_tokens, n_seqs); + + ggml_tensor * conv_states_all = mctx_cur->get_r_l(il); + ggml_tensor * ssm_states_all = mctx_cur->get_s_l(il); + + ggml_tensor * conv_kernel = model.layers[il].ssm_conv1d; + const int64_t conv_kernel_size = conv_kernel->ne[0]; + + // the channels must match how load_arch_tensors sizes wqkv, not ssm_d_inner + const int64_t conv_channels = head_k_dim * num_k_heads * 2 + head_v_dim * num_v_heads; + + ggml_tensor * conv_input = build_conv_state_at(inp, conv_states_all, qkv_mixed, + conv_kernel_size - 1, conv_channels, il); + + ggml_tensor * state = build_rs(inp, ssm_states_all, hparams.n_embd_s(), n_seqs); + state = ggml_reshape_4d(ctx0, state, head_v_dim, head_v_dim, num_v_heads, n_seqs); + cb(state, "state_predelta", il); + + ggml_tensor * conv_output_proper = ggml_ssm_conv(ctx0, conv_input, conv_kernel); + cb(conv_output_proper, "conv_output_raw", il); + + ggml_tensor * conv_output_silu = ggml_silu(ctx0, conv_output_proper); + cb(conv_output_silu, "conv_output_silu", il); + + ggml_tensor * conv_qkv_mix = conv_output_silu; + + int64_t nb1_qkv = ggml_row_size(conv_qkv_mix->type, conv_channels); + + // Extract the convolved Q, K, V from conv_output + ggml_tensor * q_conv = ggml_view_4d(ctx0, conv_qkv_mix, head_k_dim, num_k_heads, n_seq_tokens, n_seqs, + ggml_row_size(conv_qkv_mix->type, head_k_dim), + nb1_qkv, + nb1_qkv * n_seq_tokens, + 0); + + ggml_tensor * k_conv = ggml_view_4d(ctx0, conv_qkv_mix, head_k_dim, num_k_heads, n_seq_tokens, n_seqs, + ggml_row_size(conv_qkv_mix->type, head_k_dim), + nb1_qkv, + nb1_qkv * n_seq_tokens, + head_k_dim * num_k_heads * ggml_element_size(conv_qkv_mix)); + + ggml_tensor * v_conv = ggml_view_4d(ctx0, conv_qkv_mix, head_v_dim, num_v_heads, n_seq_tokens, n_seqs, + ggml_row_size(conv_qkv_mix->type, head_v_dim), + nb1_qkv, + nb1_qkv * n_seq_tokens, + ggml_row_size(conv_qkv_mix->type, 2 * head_k_dim * num_k_heads)); + + cb(q_conv, "q_conv", il); + cb(k_conv, "k_conv", il); + cb(v_conv, "v_conv", il); + + const float eps_norm = hparams.f_norm_rms_eps; + + // HRX: ggml_l2_norm, which ggml-hrx fuses into its gated delta net dispatch (the qwen35 path); + // the rsqrt form of #28068 (rms_norm + scale on a strided view) has no HRX matcher + q_conv = ggml_l2_norm(ctx0, q_conv, eps_norm); + k_conv = ggml_l2_norm(ctx0, k_conv, eps_norm); + + // repeat to match shapes when head keys != value keys; unneeded with the fused GDN + if (num_k_heads != num_v_heads && (!cparams.fused_gdn_ar || !cparams.fused_gdn_ch)) { + GGML_ASSERT(num_v_heads % num_k_heads == 0); + q_conv = ggml_repeat_4d(ctx0, q_conv, head_k_dim, num_v_heads, n_seq_tokens, n_seqs); + k_conv = ggml_repeat_4d(ctx0, k_conv, head_k_dim, num_v_heads, n_seq_tokens, n_seqs); + } + + cb(q_conv, "q_conv_predelta", il); + cb(k_conv, "k_conv_predelta", il); + cb(v_conv, "v_conv_predelta", il); + + ggml_tensor * output = build_recurrent_attn(inp, ssm_states_all, q_conv, k_conv, v_conv, gate, beta, state, il); + + ggml_tensor * z_2d = ggml_reshape_4d(ctx0, z, head_v_dim, num_v_heads, n_seq_tokens, n_seqs); + + // gated normalization, as self.norm(core_attn_out, z) in the reference + ggml_tensor * attn_out_norm = build_norm_gated(output, model.layers[il].ssm_norm, z_2d, il); + + ggml_tensor * final_output = ggml_reshape_3d(ctx0, attn_out_norm, head_v_dim * num_v_heads, n_seq_tokens, n_seqs); + cb(final_output, "final_output", il); + + cur = build_lora_mm(model.layers[il].ssm_out, final_output, model.layers[il].ssm_out_s); + cb(cur, "linear_attn_out", il); + + cur = ggml_reshape_2d(ctx0, cur, n_embd, n_seq_tokens * n_seqs); + + return cur; +} + +ggml_tensor * llama_model_qwen4exp::graph::build_layer_ffn(ggml_tensor * cur, const int il) { + GGML_ASSERT(model.layers[il].ffn_gate_inp != nullptr); + + ggml_tensor * moe_out = + build_moe_ffn(cur, + model.layers[il].ffn_gate_inp, + model.layers[il].ffn_up_exps, + model.layers[il].ffn_gate_exps, + model.layers[il].ffn_down_exps, + nullptr, + n_expert, n_expert_used, + LLM_FFN_SILU, true, + hparams.expert_weights_scale, + LLAMA_EXPERT_GATING_FUNC_TYPE_SOFTMAX, il, + nullptr, model.layers[il].ffn_gate_up_exps, + model.layers[il].ffn_up_exps_s, + model.layers[il].ffn_gate_exps_s, + model.layers[il].ffn_down_exps_s); + cb(moe_out, "ffn_moe_out", il); + + // shared experts, as in the Qwen3Next reference + if (model.layers[il].ffn_up_shexp != nullptr) { + ggml_tensor * ffn_shexp = + build_ffn(cur, + model.layers[il].ffn_up_shexp, NULL, model.layers[il].ffn_up_shexp_s, + model.layers[il].ffn_gate_shexp, NULL, model.layers[il].ffn_gate_shexp_s, + model.layers[il].ffn_down_shexp, NULL, model.layers[il].ffn_down_shexp_s, + NULL, + LLM_FFN_SILU, LLM_FFN_PAR, il); + cb(ffn_shexp, "ffn_shexp", il); + + // shared expert has its own sigmoided gate (ffn_gate_inp_shexp, one value per token) + ggml_tensor * shared_gate = build_lora_mm(model.layers[il].ffn_gate_inp_shexp, cur); + cb(shared_gate, "shared_expert_gate", il); + + shared_gate = ggml_sigmoid(ctx0, shared_gate); + cb(shared_gate, "shared_expert_gate_sigmoid", il); + + ffn_shexp = ggml_mul(ctx0, ffn_shexp, shared_gate); + cb(ffn_shexp, "ffn_shexp_gated", il); + + cur = ggml_add(ctx0, moe_out, ffn_shexp); + cb(cur, "ffn_out", il); + } else { + cur = moe_out; + } + + return cur; +} + +// PLE n-gram hash embedding: each token gathers ple_n_heads rows of a shared table. +// mixed_n = (t[p]*m[0]) ^ ... ^ (t[p-n+1]*m[n-1]); row = mixed_n % vocab[h] + offset[h] +// The hash runs host-side because ggml has no int64 and no xor. EOS resets the window. + +class llm_graph_input_ple : public llm_graph_input_i { +public: + llm_graph_input_ple(const llama_model_qwen4exp & pmodel, + const llama_kv_cache_context * mctx) : pmodel(pmodel), mctx(mctx) {} + virtual ~llm_graph_input_ple() = default; + + void set_input(const llama_ubatch * ubatch) override; + + bool can_reuse(const llm_graph_params & params) override { + mctx = static_cast(params.mctx)->get_attn(); + return rows->ne[0] == (int64_t) pmodel.hparams.ple_n_heads * params.ubatch.n_tokens; + } + + ggml_tensor * rows = nullptr; // I32 [ple_n_heads * n_tokens] + + const llama_model_qwen4exp & pmodel; + + // the predecessor tokens live in the attention KV cells (ext.tok) + const llama_kv_cache_context * mctx; + + // scratch, reused across set_input() calls + std::vector prev; +}; + +void llm_graph_input_ple::set_input(const llama_ubatch * ubatch) { + const auto & hp = pmodel.hparams; + + // an image arrives as an embd batch, so ubatch->token is null, but every position still needs a row for ggml_get_rows + // stand in the image token id that the reference hashes, or EOS if the file has no such key + // gemma3n and gemma4 do the same with a hardcoded row 0 of per_layer_token_embd. + const llama_token img_tok = hp.ple_image_token_id != 0 + ? (llama_token) hp.ple_image_token_id + : (llama_token) hp.ple_eos_token_id; + auto tok_of = [&](int64_t k) -> llama_token { + return ubatch->token ? ubatch->token[k] : img_tok; + }; + + const int64_t n_tokens = ubatch->n_tokens; + const int64_t n_gram = hp.ple_ngram_size; + const int64_t n_heads = hp.ple_n_heads; + const int64_t per_gram = hp.ple_heads_per_ngram; + const int64_t eos = hp.ple_eos_token_id; + const int64_t n_prev = n_gram - 1; + + std::vector idx(n_heads * n_tokens); + + GGML_ASSERT(mctx != nullptr); + + for (int64_t i = 0; i < n_tokens; ++i) { + // the preceding tokens would be ambiguous, see get_prev_tokens() + GGML_ASSERT(ubatch->n_seq_id[i] == 1 && "PLE n-gram embeddings do not support tokens shared by multiple sequences"); + } + + // predecessors come from the KV cells (ext.tok); apply_ubatch() already stored this ubatch, so its own tokens count too + mctx->get_prev_tokens(*ubatch, n_prev, prev); + + for (int64_t i = 0; i < n_tokens; ++i) { + // an EOS in the window resets everything at or before it + // a missing predecessor (before the sequence start, or no cached cell) reads as EOS + // the EOS of the token itself does not cut its own context, as in the reference + std::vector ctx(n_gram); + ctx[0] = tok_of(i); + bool cut = false; + for (int64_t s = 1; s < n_gram; ++s) { + // predecessor s positions back; prev[] is oldest-first, missing entries are LLAMA_TOKEN_NULL + const llama_token t = cut ? LLAMA_TOKEN_NULL : prev[i*n_prev + (n_prev - s)]; + cut = cut || t < 0 || t == eos; + ctx[s] = cut ? eos : t; + } + + for (int64_t n = 2; n <= n_gram; ++n) { + uint64_t mixed = (uint64_t) ctx[0] * hp.ple_layer_multipliers[0]; + for (int64_t j = 1; j < n; ++j) { + mixed ^= (uint64_t) ctx[j] * hp.ple_layer_multipliers[j]; + } + const int64_t base = (n - 2) * per_gram; + for (int64_t g = 0; g < per_gram; ++g) { + const int64_t h_i = base + g; + idx[i * n_heads + h_i] = + (int32_t) (mixed % hp.ple_head_vocab_sizes[h_i] + hp.ple_head_offsets[h_i]); + } + } + } + + ggml_backend_tensor_set(rows, idx.data(), 0, idx.size()*ggml_element_size(rows)); +} + +// Read a conv history out of its own recurrent row and write the new tail back. +// The shared build_conv_state cannot do this: qwen4exp has two such rows per layer. +ggml_tensor * llama_model_qwen4exp::graph::build_conv_state_at( + llm_graph_input_rs * inp, + ggml_tensor * conv_states_all, + ggml_tensor * x, + int64_t state_cols, + int64_t channels, + int il) { + const auto * mctx_cur = inp->mctx; + + const auto kv_head = mctx_cur->get_head(); + + const int64_t n_seqs = ubatch.n_seqs; + const int64_t row_total = conv_states_all->ne[0]; + + // the row is exactly this convolution's state, so the gather is reused as a whole + GGML_ASSERT(state_cols * channels == row_total); + + auto it = rs_rows.find(conv_states_all); + if (it == rs_rows.end()) { + it = rs_rows.emplace(conv_states_all, build_rs(inp, conv_states_all, row_total, n_seqs)).first; + } + ggml_tensor * rows = it->second; + + ggml_tensor * state = ggml_reshape_3d(ctx0, rows, state_cols, channels, n_seqs); + cb(state, "conv_state_at", il); + + ggml_tensor * conv_input = ggml_concat(ctx0, state, ggml_transpose(ctx0, x), 0); + + // [TAG_RECURRENT_ROLLBACK_SPLITS] keep the last state_cols columns once per rollback slot, + // slot s ending s tokens earlier so a rollback of s tokens reads a history that never saw them + const size_t row_size = ggml_row_size(conv_states_all->type, row_total); + const uint32_t mem_size = mctx_cur->get_size(); + + const int64_t n_slots = (int64_t) cparams.n_rs_seq + 1; + + for (int64_t slot = 0; slot < n_slots; ++slot) { + const int64_t s_idx = std::max(0, conv_input->ne[0] - state_cols - slot); + + ggml_tensor * tail = ggml_view_3d(ctx0, conv_input, + state_cols, channels, n_seqs, + conv_input->nb[1], conv_input->nb[2], + ggml_row_size(conv_input->type, s_idx)); + + ggml_tensor * dst = ggml_view_2d(ctx0, conv_states_all, + state_cols * channels, n_seqs, + conv_states_all->nb[1], + (slot * mem_size + kv_head) * row_size); + + ggml_build_forward_expand(gf, ggml_cpy(ctx0, ggml_cont(ctx0, tail), dst)); + } + + return conv_input; +} + +ggml_tensor * llama_model_qwen4exp::graph::build_inp_ple( + const llama_memory_hybrid_idx_context * mctx_hyb) { + const int64_t n_heads = hparams.ple_n_heads; + + // the attention cells see every ubatch regardless of the layer types + auto ple_inp = std::make_unique( + static_cast(model), mctx_hyb->get_attn()); + + ple_inp->rows = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_heads * n_tokens); + ggml_set_input(ple_inp->rows); + ggml_tensor * rows = ple_inp->rows; + res->add_input(std::move(ple_inp)); + + // gather then flatten the heads: get_rows lays the head dimension out slowest, as the reference does + ggml_tensor * emb = ggml_get_rows(ctx0, model.per_layer_tok_embd, rows); + emb = ggml_reshape_2d(ctx0, emb, hparams.ple_head_dim * n_heads, n_tokens); + cb(emb, "ple_embd", -1); + + return emb; +} + +ggml_tensor * llama_model_qwen4exp::graph::build_ple( + llm_graph_input_rs * inp, + ggml_tensor * emb, + ggml_tensor * hidden, + int il) { + const int64_t hc = hparams.dsv4_hc_mult; + const int64_t hc_dim = hc * n_embd; + + ggml_tensor * key = build_lora_mm(model.layers[il].ple_key, emb); + ggml_tensor * value = build_lora_mm(model.layers[il].ple_value, emb); + + // both norms group over one hc stream, with a [n_embd, hc] weight + auto grouped_norm = [&](ggml_tensor * x, ggml_tensor * w) { + ggml_tensor * t = ggml_reshape_3d(ctx0, x, n_embd, hc, n_tokens); + return ggml_mul(ctx0, ggml_rms_norm(ctx0, t, hparams.f_norm_rms_eps), w); + }; + + key = grouped_norm(key, model.layers[il].ple_norm_key); + ggml_tensor * query = grouped_norm(hidden, model.layers[il].ple_norm_query); + + // per-stream dot product, then a signed square root before the sigmoid + ggml_tensor * s = ggml_sum_rows(ctx0, ggml_mul(ctx0, key, query)); + s = ggml_scale(ctx0, s, 1.0f / sqrtf((float) n_embd)); + + ggml_tensor * mag = ggml_sqrt(ctx0, ggml_clamp(ctx0, ggml_abs(ctx0, s), 1e-6f, INFINITY)); + ggml_tensor * gate = ggml_sigmoid(ctx0, ggml_mul(ctx0, ggml_sgn(ctx0, s), mag)); + cb(gate, "ple_gate", il); + + // [n_embd, 1, T] value broadcast across the hc streams, scaled by the gate + ggml_tensor * v3 = ggml_reshape_3d(ctx0, value, n_embd, 1, n_tokens); + v3 = ggml_repeat_4d(ctx0, v3, n_embd, hc, n_tokens, 1); + + ggml_tensor * gated = ggml_mul(ctx0, v3, gate); + cb(gated, "ple_gated_value", il); + + ggml_tensor * normalized = grouped_norm( + ggml_reshape_2d(ctx0, gated, hc_dim, n_tokens), + model.layers[il].ple_norm_conv); + normalized = ggml_reshape_2d(ctx0, normalized, hc_dim, n_tokens); + + // depthwise causal conv, dilated by the n-gram size, as a sum of shifted copies + // ggml_conv_1d_dw is documented as unreliable: + // out[c, t] = sum_k w[k, c] * x[c, t - (K-1-k)*dilation] + // The history of the earlier ubatches is prepended, so a chunked prefill matches a single-shot one. + const int64_t kern = hparams.ple_conv_kernel; + const int64_t dil = hparams.ple_ngram_size; + const int64_t hist = (kern - 1) * dil; + + // the conv history is per sequence, so the input carries the sequence axis too + const int64_t n_seqs = ubatch.n_seqs; + const int64_t n_seq_tokens = ubatch.n_seq_tokens; + + // [hist + n_seq_tokens, hc_dim, n_seqs], tokens on ne[0] + ggml_tensor * padded = build_conv_state_at(inp, inp->mctx->get_p_l(il), + ggml_reshape_3d(ctx0, normalized, hc_dim, n_seq_tokens, n_seqs), + hist, hc_dim, il); + + ggml_tensor * conv_out = nullptr; + for (int64_t k = 0; k < kern; ++k) { + // tap k reads (kern-1-k)*dilation positions back + const int64_t start = hist - (kern - 1 - k) * dil; + + ggml_tensor * shifted = ggml_cont(ctx0, + ggml_transpose(ctx0, + ggml_view_3d(ctx0, padded, n_seq_tokens, hc_dim, n_seqs, + padded->nb[1], padded->nb[2], + ggml_row_size(padded->type, start)))); + + // column k of the [kern, hc_dim] kernel is one weight per channel + ggml_tensor * wk = ggml_cont(ctx0, + ggml_view_2d(ctx0, model.layers[il].ple_conv1d, 1, hc_dim, + model.layers[il].ple_conv1d->nb[1], + k * model.layers[il].ple_conv1d->nb[0])); + // this kernel keeps the file type, so cast it before it multiplies an f32 activation + wk = ggml_reshape_1d(ctx0, wk, hc_dim); + if (wk->type != GGML_TYPE_F32) { + wk = ggml_cast(ctx0, wk, GGML_TYPE_F32); + } + + ggml_tensor * term = ggml_mul(ctx0, shifted, wk); + conv_out = conv_out ? ggml_add(ctx0, conv_out, term) : term; + } + + conv_out = ggml_silu(ctx0, conv_out); + conv_out = ggml_reshape_3d(ctx0, ggml_cont(ctx0, conv_out), n_embd, hc, n_tokens); + cb(conv_out, "ple_conv_out", il); + + return ggml_add(ctx0, hidden, ggml_add(ctx0, gated, conv_out)); +} diff --git a/src/models/zaya.cpp b/src/models/zaya.cpp new file mode 100644 index 000000000000..15a83354899c --- /dev/null +++ b/src/models/zaya.cpp @@ -0,0 +1,612 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// The CCA graph started from Juste-Leo2's ZAYA1 draft for llama.cpp +// (ggml-org/llama.cpp PR #23112, MIT). + +#include "models.h" + +#include "ggml.h" +#include "llama-memory-recurrent.h" + +#include + +// ZAYA1 (Zyphra). Every layer runs CCA attention and then a MoE, each followed by a +// learned residual scale: +// - CCA: q and k go through a 2-tap depthwise conv (ssm_conv1d) and a grouped conv +// (cca_conv_grp) over time, so each sequence carries two recurrent rows: the +// conv state of q|k (n_embd_r = 2*n_qk) and the previous hidden state +// (n_embd_s = n_embd); v is two projections, of the current and of the +// previous hidden state. +// - MoE router: down_proj -> EDA (adds the previous layer's router state) -> RMSNorm +// -> MLP (GELU) x2 -> softmax -> top-1 over n_expert + 1 slots, the last being a +// skip expert with zero output. Experts are pre-stacked (ffn_gate_up_exps). +// - input_hidden_states_scale/bias apply to the embeddings. + +void llama_model_zaya::load_arch_hparams(llama_model_loader & ml) { + ml.get_key(LLM_KV_ATTENTION_LAYERNORM_RMS_EPS, hparams.f_norm_rms_eps, false); + if (hparams.f_norm_rms_eps == 0.0f) { + hparams.f_norm_rms_eps = 1e-5f; + } + ml.get_key(LLM_KV_SSM_CONV_KERNEL, hparams.ssm_d_conv); + GGML_ASSERT(hparams.ssm_d_conv == 2 && "the CCA graph is written for 2-tap convs"); + ml.get_key(LLM_KV_EXPERT_FEED_FORWARD_LENGTH, hparams.n_ff_exp, false); + hparams.n_rot_full = hparams.n_embd_head_k() / 2; + hparams.n_rot_swa = hparams.n_embd_head_k() / 2; + // two recurrent rows per sequence, written whole (as the gated-delta-net models do): + // r = the 2-tap conv state of q|k, 2*n_qk; s = the previous hidden state, n_embd. + // With d_conv = 2 and d_state = 1, n_embd_r() = d_inner + 2*n_group and n_embd_s() = d_inner. + const uint32_t n_qk = (hparams.n_head() + hparams.n_head_kv()) * hparams.n_embd_head_k(); + GGML_ASSERT(2*n_qk > hparams.n_embd && (2*n_qk - hparams.n_embd) % 2 == 0); + hparams.ssm_d_inner = hparams.n_embd; + hparams.ssm_d_state = 1; + hparams.ssm_n_group = (2*n_qk - hparams.n_embd) / 2; + GGML_ASSERT(hparams.n_embd_r() == 2*n_qk && hparams.n_embd_s() == hparams.n_embd); + std::fill(hparams.is_recr_impl.begin(), hparams.is_recr_impl.end(), true); + + // ZAYA1-74B: sliding-window attention on the layers the pattern marks, with their own rope base + if (ml.get_key(LLM_KV_ATTENTION_SLIDING_WINDOW, hparams.n_swa, false) && hparams.n_swa > 0) { + hparams.swa_type = LLAMA_SWA_TYPE_STANDARD; + // the converter writes a per-layer array; a scalar is a period, as the other SWA archs read it. + // Without the key (llama_model_saver does not write it) default to ZAYA1-74B's: even layers slide. + // Read the scalar first: the array read would take a scalar too, as the raw value in every layer. + uint32_t swa_period = 2; + if (ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false) || + !ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, hparams.is_swa_impl, hparams.n_layer(), false)) { + hparams.set_swa_pattern(swa_period); + } + hparams.rope_freq_base_train_swa = hparams.rope_freq_base_train; + hparams.rope_freq_scale_train_swa = hparams.rope_freq_scale_train; + ml.get_key(LLM_KV_ROPE_FREQ_BASE_SWA, hparams.rope_freq_base_train_swa, false); + } + + ml.get_key(LLM_KV_ZAYA_VLORA_RANK_ATTN, vlora_rank_attn, false); + ml.get_key(LLM_KV_ZAYA_VLORA_RANK_FFN, vlora_rank_ffn, false); + + switch (hparams.n_layer()) { + case 40: type = LLM_TYPE_8B; break; + default: type = LLM_TYPE_UNKNOWN; + } +} + +void llama_model_zaya::load_arch_tensors(llama_model_loader &) { + LLAMA_LOAD_LOCALS; + + const int64_t n_embd_head = hparams.n_embd_head_k(); + const int64_t n_ff_exp = hparams.n_ff_exp; + const int64_t d_conv = hparams.ssm_d_conv; + + tok_embd = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, 0); + + output_norm = create_tensor(tn(LLM_TENSOR_OUTPUT_NORM, "weight"), {n_embd}, 0); + output = create_tensor(tn(LLM_TENSOR_OUTPUT, "weight"), {n_embd, n_vocab}, TENSOR_NOT_REQUIRED); + if (output == NULL) { + output = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, TENSOR_DUPLICATED); + } + + zaya_input_hs_scale = create_tensor(tn(LLM_TENSOR_INPUT_HIDDEN_STATES_SCALE, "weight"), {n_embd}, TENSOR_NOT_REQUIRED); + zaya_input_hs_bias = create_tensor(tn(LLM_TENSOR_INPUT_HIDDEN_STATES_SCALE, "bias"), {n_embd}, TENSOR_NOT_REQUIRED); + + zaya_res_scale_hs = create_tensor(tn(LLM_TENSOR_RES_SCALE_HS_FINAL, "weight"), {n_embd}, TENSOR_NOT_REQUIRED); + zaya_res_scale_hs_b = create_tensor(tn(LLM_TENSOR_RES_SCALE_HS_FINAL, "bias"), {n_embd}, TENSOR_NOT_REQUIRED); + zaya_res_scale_res = create_tensor(tn(LLM_TENSOR_RES_SCALE_RES_FINAL, "weight"), {n_embd}, TENSOR_NOT_REQUIRED); + zaya_res_scale_res_b = create_tensor(tn(LLM_TENSOR_RES_SCALE_RES_FINAL, "bias"), {n_embd}, TENSOR_NOT_REQUIRED); + + for (int i = 0; i < n_layer; ++i) { + auto & layer = layers[i]; + + const int64_t n_head_l = hparams.n_head(i); + const int64_t n_head_kv_l = hparams.n_head_kv(i); + const int64_t n_embd_q = n_head_l * n_embd_head; + const int64_t n_embd_k = n_head_kv_l * n_embd_head; + const int64_t n_qk = n_embd_q + n_embd_k; + const int64_t n_groups = n_head_l + n_head_kv_l; + const int64_t n_ff_l = hparams.n_ff(i); + + layer.attn_norm = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "weight", i), {n_embd}, 0); + layer.attn_post_norm = create_tensor(tn(LLM_TENSOR_ATTN_POST_NORM, "weight", i), {n_embd}, 0); + + // every layer has both CCA attention and the MoE + layer.wq = create_tensor(tn(LLM_TENSOR_ATTN_Q, "weight", i), {n_embd, n_embd_q}, 0); + layer.wk = create_tensor(tn(LLM_TENSOR_ATTN_K, "weight", i), {n_embd, n_embd_k}, 0); + layer.cca_val_proj1 = create_tensor(tn(LLM_TENSOR_CCA_VAL_PROJ1, "weight", i), {n_embd, n_embd_k / 2}, 0); + layer.cca_val_proj2 = create_tensor(tn(LLM_TENSOR_CCA_VAL_PROJ2, "weight", i), {n_embd, n_embd_k / 2}, 0); + layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", i), {n_embd_q, n_embd}, 0); + layer.ssm_conv1d = create_tensor(tn(LLM_TENSOR_SSM_CONV1D, "weight", i), {d_conv, n_qk}, 0); + layer.ssm_conv1d_b = create_tensor(tn(LLM_TENSOR_SSM_CONV1D, "bias", i), {n_qk}, TENSOR_NOT_REQUIRED); + layer.cca_conv_grp = create_tensor(tn(LLM_TENSOR_CCA_CONV_GRP, "weight", i), {n_qk / n_groups, n_qk, d_conv}, 0); // tap-major + layer.cca_conv_grp_b = create_tensor(tn(LLM_TENSOR_CCA_CONV_GRP, "bias", i), {n_qk}, 0); + layer.cca_k_scale = create_tensor(tn(LLM_TENSOR_CCA_K_SCALE, "weight", i), {n_head_kv_l}, 0); + + layer.ffn_gate_inp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP, "weight", i), {n_embd, n_ff_exp}, 0); + layer.ffn_gate_inp_b = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP, "bias", i), {n_ff_exp}, TENSOR_NOT_REQUIRED); + layer.ffn_norm = create_tensor(tn(LLM_TENSOR_FFN_NORM, "weight", i), {n_ff_exp}, 0); + layer.ffn_gate = create_tensor(tn(LLM_TENSOR_FFN_GATE, "weight", i), {n_ff_exp, n_ff_exp}, 0); + layer.ffn_gate_b = create_tensor(tn(LLM_TENSOR_FFN_GATE, "bias", i), {n_ff_exp}, TENSOR_NOT_REQUIRED); + layer.zaya_router_mlp2 = create_tensor(tn(LLM_TENSOR_ZAYA_ROUTER_MLP2, "weight", i), {n_ff_exp, n_ff_exp}, 0); + layer.zaya_router_mlp2_b = create_tensor(tn(LLM_TENSOR_ZAYA_ROUTER_MLP2, "bias", i), {n_ff_exp}, TENSOR_NOT_REQUIRED); + layer.zaya_router_mlp4 = create_tensor(tn(LLM_TENSOR_ZAYA_ROUTER_MLP4, "weight", i), {n_ff_exp, n_expert + 1}, 0); + layer.zaya_router_biases = create_tensor(tn(LLM_TENSOR_ZAYA_ROUTER_BIASES, "weight", i), {n_expert + 1}, TENSOR_NOT_REQUIRED); + layer.zaya_router_eda_scale = create_tensor(tn(LLM_TENSOR_ZAYA_ROUTER_EDA_SCALE, "weight", i), {n_ff_exp}, TENSOR_NOT_REQUIRED); + layer.ffn_gate_up_exps = create_tensor(tn(LLM_TENSOR_FFN_GATE_UP_EXPS, "weight", i), {n_embd, n_ff_l * 2, n_expert}, 0); + layer.ffn_down_exps = create_tensor(tn(LLM_TENSOR_FFN_DOWN_EXPS, "weight", i), {n_ff_l, n_embd, n_expert}, 0); + + layer.res_scale_hs = create_tensor(tn(LLM_TENSOR_RES_SCALE_HS, "weight", i), {n_embd}, TENSOR_NOT_REQUIRED); + layer.res_scale_hs_b = create_tensor(tn(LLM_TENSOR_RES_SCALE_HS, "bias", i), {n_embd}, TENSOR_NOT_REQUIRED); + layer.res_scale_res = create_tensor(tn(LLM_TENSOR_RES_SCALE_RES, "weight", i), {n_embd}, TENSOR_NOT_REQUIRED); + layer.res_scale_res_b = create_tensor(tn(LLM_TENSOR_RES_SCALE_RES, "bias", i), {n_embd}, TENSOR_NOT_REQUIRED); + layer.res_scale_hs_mlp = create_tensor(tn(LLM_TENSOR_RES_SCALE_HS_MLP, "weight", i), {n_embd}, TENSOR_NOT_REQUIRED); + layer.res_scale_hs_mlp_b = create_tensor(tn(LLM_TENSOR_RES_SCALE_HS_MLP, "bias", i), {n_embd}, TENSOR_NOT_REQUIRED); + layer.res_scale_res_mlp = create_tensor(tn(LLM_TENSOR_RES_SCALE_RES_MLP, "weight", i), {n_embd}, TENSOR_NOT_REQUIRED); + layer.res_scale_res_mlp_b = create_tensor(tn(LLM_TENSOR_RES_SCALE_RES_MLP, "bias", i), {n_embd}, TENSOR_NOT_REQUIRED); + + if (vlora_rank_attn > 0) { + const int64_t r = vlora_rank_attn; + layer.zaya_vlora_q_a = create_tensor(tn(LLM_TENSOR_ZAYA_VLORA_Q_A, "weight", i), {n_embd, r}, 0); + layer.zaya_vlora_q_b = create_tensor(tn(LLM_TENSOR_ZAYA_VLORA_Q_B, "weight", i), {r, n_embd_q}, 0); + layer.zaya_vlora_k_a = create_tensor(tn(LLM_TENSOR_ZAYA_VLORA_K_A, "weight", i), {n_embd, r}, 0); + layer.zaya_vlora_k_b = create_tensor(tn(LLM_TENSOR_ZAYA_VLORA_K_B, "weight", i), {r, n_embd_k}, 0); + layer.zaya_vlora_v1_a = create_tensor(tn(LLM_TENSOR_ZAYA_VLORA_V1_A, "weight", i), {n_embd, r}, 0); + layer.zaya_vlora_v1_b = create_tensor(tn(LLM_TENSOR_ZAYA_VLORA_V1_B, "weight", i), {r, n_embd_k / 2}, 0); + layer.zaya_vlora_v2_a = create_tensor(tn(LLM_TENSOR_ZAYA_VLORA_V2_A, "weight", i), {n_embd, r}, 0); + layer.zaya_vlora_v2_b = create_tensor(tn(LLM_TENSOR_ZAYA_VLORA_V2_B, "weight", i), {r, n_embd_k / 2}, 0); + layer.zaya_vlora_o_a = create_tensor(tn(LLM_TENSOR_ZAYA_VLORA_O_A, "weight", i), {n_embd_q, r}, 0); + layer.zaya_vlora_o_b = create_tensor(tn(LLM_TENSOR_ZAYA_VLORA_O_B, "weight", i), {r, n_embd}, 0); + } + if (vlora_rank_ffn > 0) { + const int64_t r = vlora_rank_ffn; + layer.zaya_vlora_up_exps_a = create_tensor(tn(LLM_TENSOR_ZAYA_VLORA_UP_EXPS_A, "weight", i), {n_embd, r, n_expert}, 0); + layer.zaya_vlora_up_exps_b = create_tensor(tn(LLM_TENSOR_ZAYA_VLORA_UP_EXPS_B, "weight", i), {r, n_ff_l * 2, n_expert}, 0); + layer.zaya_vlora_down_exps_a = create_tensor(tn(LLM_TENSOR_ZAYA_VLORA_DOWN_EXPS_A, "weight", i), {n_ff_l, r, n_expert}, 0); + layer.zaya_vlora_down_exps_b = create_tensor(tn(LLM_TENSOR_ZAYA_VLORA_DOWN_EXPS_B, "weight", i), {r, n_embd, n_expert}, 0); + } + } +} + +std::unique_ptr llama_model_zaya::build_arch_graph(const llm_graph_params & params) const { + if (hparams.swa_type == LLAMA_SWA_TYPE_STANDARD) { + return std::make_unique>(*this, params); + } + return std::make_unique>(*this, params); +} + +template +llama_model_zaya::graph::graph(const llama_model & model, const llm_graph_params & params) : llm_graph_context(params) { + const int64_t n_embd_head = hparams.n_embd_head_k(); + const int64_t n_expert = hparams.n_expert; + const int64_t n_seqs = ubatch.n_seqs; + + GGML_ASSERT(n_seqs != 0); + GGML_ASSERT(ubatch.equal_seqs()); + GGML_ASSERT(n_tokens % n_seqs == 0); + + const int64_t n_seq_tokens = n_tokens / n_seqs; + + ggml_tensor * cur; + ggml_tensor * inpL; + + inpL = build_inp_embd(model.tok_embd); + + if (model.zaya_input_hs_scale != nullptr) { + if (model.zaya_input_hs_bias != nullptr) { + inpL = ggml_add(ctx0, inpL, model.zaya_input_hs_bias); + } + inpL = ggml_mul(ctx0, inpL, model.zaya_input_hs_scale); + cb(inpL, "input_hs_scaled", -1); + } + + auto * inp = [&] { + if constexpr (iswa) { + return build_inp_mem_hybrid_iswa(); + } else { + return build_inp_mem_hybrid(); + } + }(); + auto * inp_recr = inp->get_recr(); + + ggml_tensor * inp_pos = build_inp_pos(); + ggml_tensor * inp_out_ids = build_inp_out_ids(); + ggml_tensor * prev_router = nullptr; + + // A CONT of a fresh, already contiguous result (not a view) copies it for nothing; each copy is + // a dispatch per layer on a GPU backend. Views keep their CONT even when contiguous: on HRX the + // consumers of a sliced or permuted view (flash attention, the grouped conv matmul) take a + // much slower strided path (2026-09-29). The MoE router keeps its CONTs as well: dropping them + // made each layer wait ~110 us before the top-k gather. + const auto cont_if_needed = [&](ggml_tensor * t) { + return ggml_is_contiguous(t) && t->view_src == nullptr ? t : ggml_cont(ctx0, t); + }; + + const auto apply_res_scale = [&](ggml_tensor * x, ggml_tensor * scale, ggml_tensor * bias, const char * name, int il) { + if (scale == nullptr) { + return x; + } + if (bias != nullptr) { + x = ggml_add(ctx0, x, bias); + } + x = ggml_mul(ctx0, x, scale); + cb(x, name, il); + return x; + }; + + for (int il = 0; il < n_layer; ++il) { + const auto & layer = model.layers[il]; + + const int64_t n_head = hparams.n_head(il); + const int64_t n_head_kv = hparams.n_head_kv(il); + const int64_t n_embd_q = n_head * n_embd_head; + const int64_t n_embd_k = n_head_kv * n_embd_head; + const int64_t n_qk = n_embd_q + n_embd_k; + const int64_t n_groups = n_head + n_head_kv; + const int64_t n_gqa = n_head / n_head_kv; + + // Zaya 8B (HF ZayaDecoderLayer): EVERY layer runs BOTH blocks. + // residual = h (layer input, fp32 stream) + // cur = input_layernorm(residual) (attn_norm) + // attn = CCA(cur) + // residual = (attn+pa_hsb)*pa_hss + (residual+pa_rsb)*pa_rss (res_scale_hs/res) + // cur = post_attention_layernorm(residual) (post_attn_norm) + // moe = MoE(cur, prev_router) + // h = (moe+pm_hsb)*pm_hss + (residual+pm_rsb)*pm_rss (res_scale_hs_mlp/res_mlp) + ggml_tensor * residual = inpL; + + cur = build_norm(residual, layer.attn_norm, nullptr, LLM_NORM_RMS, il); + cb(cur, "input_norm", il); + + // ===== CCA attention (every layer) ===== + + const int64_t conv_state_size = 2*n_qk; + + ggml_tensor * conv_states_all = inp_recr->mctx->get_r_l(il); + ggml_tensor * conv_state = build_rs(inp_recr, conv_states_all, hparams.n_embd_r(), n_seqs); + conv_state = ggml_reshape_3d(ctx0, conv_state, 2, n_qk, n_seqs); + cb(conv_state, "cca_conv_state", il); + + ggml_tensor * hs_states_all = inp_recr->mctx->get_s_l(il); + ggml_tensor * prev_hs = build_rs(inp_recr, hs_states_all, hparams.n_embd_s(), n_seqs); + cb(prev_hs, "cca_prev_hs", il); + + // ZAYA1-VL: the vision-only LoRA runs on image tokens, which mtmd decodes as their own + // embedding ubatches, so a ubatch of embeddings takes it for every token + const bool vlora = ubatch.embd != nullptr && layer.zaya_vlora_q_a != nullptr; + auto lora = [&](ggml_tensor * a, ggml_tensor * b, ggml_tensor * x) { + return ggml_mul_mat(ctx0, b, ggml_mul_mat(ctx0, a, x)); + }; + + ggml_tensor * Qraw = ggml_mul_mat(ctx0, layer.wq, cur); + ggml_tensor * Kraw = ggml_mul_mat(ctx0, layer.wk, cur); + if (vlora) { + Qraw = ggml_add(ctx0, Qraw, lora(layer.zaya_vlora_q_a, layer.zaya_vlora_q_b, cur)); + Kraw = ggml_add(ctx0, Kraw, lora(layer.zaya_vlora_k_a, layer.zaya_vlora_k_b, cur)); + } + cb(Qraw, "Qraw", il); + cb(Kraw, "Kraw", il); + + ggml_tensor * cur_state_src = cont_if_needed(cur); + ggml_tensor * cur_seq = ggml_reshape_3d(ctx0, cur_state_src, n_embd, n_seq_tokens, n_seqs); + + ggml_tensor * hs_d = ggml_reshape_3d(ctx0, cont_if_needed(prev_hs), n_embd, 1, n_seqs); + if (n_seq_tokens > 1) { + ggml_tensor * cur_shift = ggml_view_3d(ctx0, cur_seq, n_embd, n_seq_tokens - 1, n_seqs, + cur_seq->nb[1], + cur_seq->nb[2], + 0); + hs_d = ggml_concat(ctx0, hs_d, cur_shift, 1); + } + hs_d = ggml_reshape_2d(ctx0, cont_if_needed(hs_d), n_embd, n_tokens); + cb(hs_d, "cca_hs_d", il); + + ggml_tensor * V1 = ggml_mul_mat(ctx0, layer.cca_val_proj1, cur); + ggml_tensor * V2 = ggml_mul_mat(ctx0, layer.cca_val_proj2, hs_d); + if (vlora) { + // V2's LoRA reads the previous token's state but follows the current token's mask + V1 = ggml_add(ctx0, V1, lora(layer.zaya_vlora_v1_a, layer.zaya_vlora_v1_b, cur)); + V2 = ggml_add(ctx0, V2, lora(layer.zaya_vlora_v2_a, layer.zaya_vlora_v2_b, hs_d)); + } + cb(V1, "V1", il); + cb(V2, "V2", il); + ggml_tensor * Vcur = ggml_concat(ctx0, V1, V2, 0); + cb(Vcur, "Vcur", il); + + ggml_tensor * QKraw = ggml_concat(ctx0, Qraw, Kraw, 0); + cb(QKraw, "QKraw", il); + + // Qraw and Kraw are fresh matmul outputs, already contiguous: reshape them in place. A CONT + // here copies for nothing, and with this block entirely on HRX the copied K read back wrong + // in 512-token batches (wikitext perplexity 71.8 instead of 21.6; right with a CPU split + // after the copy, or without the copy). Decode was unaffected. + ggml_tensor * Qpre = ggml_reshape_3d(ctx0, ggml_is_contiguous(Qraw) ? Qraw : cont_if_needed(Qraw), n_embd_head, n_head, n_tokens); + ggml_tensor * Kpre = ggml_reshape_3d(ctx0, ggml_is_contiguous(Kraw) ? Kraw : cont_if_needed(Kraw), n_embd_head, n_head_kv, n_tokens); + + ggml_tensor * Kpre_grouped = ggml_reshape_4d(ctx0, Kpre, n_embd_head, 1, n_head_kv, n_tokens); + Kpre_grouped = ggml_repeat_4d(ctx0, Kpre_grouped, n_embd_head, n_gqa, n_head_kv, n_tokens); + ggml_tensor * Kpre_rep = ggml_reshape_3d(ctx0, Kpre_grouped, n_embd_head, n_head, n_tokens); + ggml_tensor * qk_mean_q = ggml_scale(ctx0, ggml_add(ctx0, Qpre, Kpre_rep), 0.5f); + cb(qk_mean_q, "qk_mean_q", il); + + ggml_tensor * Qgroup = ggml_reshape_4d(ctx0, Qpre, n_embd_head, n_gqa, n_head_kv, n_tokens); + Qgroup = ggml_permute(ctx0, Qgroup, 1, 0, 2, 3); + Qgroup = cont_if_needed(Qgroup); + ggml_tensor * Qmean = ggml_scale(ctx0, ggml_sum_rows(ctx0, Qgroup), 1.0f/n_gqa); // MEAN, as SUM_ROWS + SCALE + Qmean = ggml_reshape_3d(ctx0, Qmean, n_embd_head, n_head_kv, n_tokens); + ggml_tensor * qk_mean_k = ggml_scale(ctx0, ggml_add(ctx0, Qmean, Kpre), 0.5f); + cb(qk_mean_k, "qk_mean_k", il); + + // [n_qk, T, S] -> [T, n_qk, S]: split the sequences before transposing, or with more than + // one sequence in the ubatch a channel's row would run across all of them + ggml_tensor * QKraw_t = ggml_reshape_3d(ctx0, QKraw, n_qk, n_seq_tokens, n_seqs); + // with one token per sequence the transpose moves no data: a reshape does it without a copy + QKraw_t = n_seq_tokens == 1 ? ggml_reshape_3d(ctx0, QKraw_t, 1, n_qk, n_seqs) + : cont_if_needed(ggml_transpose(ctx0, QKraw_t)); + + ggml_tensor * conv_input = ggml_concat(ctx0, conv_state, QKraw_t, 0); + cb(conv_input, "cca_conv_input", il); + + ggml_tensor * last_conv_states = ggml_view_3d(ctx0, conv_input, 2, n_qk, n_seqs, + conv_input->nb[1], + conv_input->nb[2], + n_seq_tokens*conv_input->nb[0]); + cb(last_conv_states, "cca_last_conv_states", il); + + const auto kv_head = inp_recr->mctx->get_head(); + ggml_tensor * conv_state_update_target = ggml_view_2d(ctx0, conv_states_all, conv_state_size, n_seqs, + conv_states_all->nb[1], + kv_head*conv_states_all->nb[1]); + ggml_build_forward_expand(gf, ggml_cpy(ctx0, + ggml_reshape_2d(ctx0, cont_if_needed(last_conv_states), conv_state_size, n_seqs), + conv_state_update_target)); + + ggml_tensor * last_hs = ggml_view_2d(ctx0, cur_seq, n_embd, n_seqs, + cur_seq->nb[2], + (n_seq_tokens - 1)*cur_seq->nb[1]); + ggml_tensor * prev_hs_update_target = ggml_view_2d(ctx0, hs_states_all, n_embd, n_seqs, + hs_states_all->nb[1], + kv_head*hs_states_all->nb[1]); + ggml_build_forward_expand(gf, ggml_cpy(ctx0, cont_if_needed(last_hs), prev_hs_update_target)); + + ggml_tensor * conv_dw = layer.ssm_conv1d; + if (conv_dw->type != GGML_TYPE_F32) { + conv_dw = cont_if_needed(ggml_cast(ctx0, conv_dw, GGML_TYPE_F32)); + } + ggml_tensor * QK = ggml_ssm_conv(ctx0, conv_input, conv_dw); + // Grouped conv (2 taps, no padding) as one batched matmul per tap: the weights are + // stored tap-major, {IC_G, n_qk, 2}, so tap k is the contiguous block + // [IC_G, OC_G, groups], applied to the depthwise output shifted by k steps. + // QK is the depthwise output, [n_qk, T + 1, S]. + if (layer.ssm_conv1d_b) { + QK = ggml_add(ctx0, QK, ggml_reshape_2d(ctx0, layer.ssm_conv1d_b, n_qk, 1)); + } + cb(QK, "QK_dw", il); + const int64_t ic_g = n_qk / n_groups; + ggml_tensor * w_grp = layer.cca_conv_grp; + ggml_tensor * grp = nullptr; + // The sequences fold into the token dimension, so the batched matmul broadcasts the + // weight over dim 2 (the groups) only: backends such as Vulkan cannot broadcast src0 + // over dim 3, and with several sequences reserved the op would fall back to the CPU. + for (int tap = 0; tap < 2; ++tap) { + ggml_tensor * x = ggml_view_4d(ctx0, QK, ic_g, n_groups, n_seq_tokens, n_seqs, + ic_g*ggml_element_size(QK), QK->nb[1], QK->nb[2], tap*QK->nb[1]); + x = cont_if_needed(ggml_permute(ctx0, x, 0, 3, 1, 2)); // [IC_G, T, S, G] + x = ggml_reshape_3d(ctx0, x, ic_g, n_seq_tokens*n_seqs, n_groups); // [IC_G, T*S, G] + ggml_tensor * w = ggml_view_3d(ctx0, w_grp, ic_g, ic_g, n_groups, + w_grp->nb[1], ic_g*w_grp->nb[1], tap*w_grp->nb[2]); // [IC_G, OC_G, G] + ggml_tensor * y = ggml_mul_mat(ctx0, w, x); // [OC_G, T*S, G] + grp = grp ? ggml_add(ctx0, grp, y) : y; + } + grp = ggml_reshape_4d(ctx0, grp, ic_g, n_seq_tokens, n_seqs, n_groups); // [OC_G, T, S, G] + QK = cont_if_needed(ggml_permute(ctx0, grp, 0, 2, 3, 1)); // [OC_G, G, T, S] + QK = ggml_reshape_2d(ctx0, QK, n_qk, n_tokens); + QK = ggml_add(ctx0, QK, layer.cca_conv_grp_b); + cb(QK, "QK_grp", il); + + ggml_tensor * Q_conv = ggml_view_2d(ctx0, QK, n_embd_q, n_tokens, QK->nb[1], 0); + ggml_tensor * K_conv = ggml_view_2d(ctx0, QK, n_embd_k, n_tokens, QK->nb[1], n_embd_q*ggml_element_size(QK)); + + ggml_tensor * Qcur = ggml_reshape_3d(ctx0, cont_if_needed(Q_conv), n_embd_head, n_head, n_tokens); + ggml_tensor * Kcur = ggml_reshape_3d(ctx0, cont_if_needed(K_conv), n_embd_head, n_head_kv, n_tokens); + + Qcur = ggml_add(ctx0, Qcur, qk_mean_q); + Kcur = ggml_add(ctx0, Kcur, qk_mean_k); + + // l2_norm(x) * sqrt(head_dim) is rms_norm(x): one op, and one every backend has. + // eps 1e-24 matches the reference clamping |x| at 1e-12. + Qcur = ggml_rms_norm(ctx0, Qcur, 1e-24f); + Kcur = ggml_rms_norm(ctx0, Kcur, 1e-24f); + Kcur = ggml_mul(ctx0, Kcur, ggml_reshape_3d(ctx0, layer.cca_k_scale, 1, n_head_kv, 1)); + cb(Qcur, "Qcur_pre_rope", il); + cb(Kcur, "Kcur_pre_rope", il); + + ggml_tensor * rope_factors = model.get_rope_factors(cparams, il); + // sliding layers (ZAYA1-74B) have their own base; for other models this is freq_base + const float freq_base_l = model.get_rope_freq_base (cparams, il); + const float freq_scale_l = model.get_rope_freq_scale(cparams, il); + Qcur = ggml_rope_ext(ctx0, Qcur, inp_pos, rope_factors, + n_rot, rope_type, n_ctx_orig, freq_base_l, freq_scale_l, + ext_factor, attn_factor, beta_fast, beta_slow); + Kcur = ggml_rope_ext(ctx0, Kcur, inp_pos, rope_factors, + n_rot, rope_type, n_ctx_orig, freq_base_l, freq_scale_l, + ext_factor, attn_factor, beta_fast, beta_slow); + cb(Qcur, "Qcur", il); + cb(Kcur, "Kcur", il); + + Vcur = ggml_reshape_3d(ctx0, cont_if_needed(Vcur), n_embd_head, n_head_kv, n_tokens); + + cur = build_attn(inp->get_attn(), vlora ? nullptr : layer.wo, nullptr, nullptr, + Qcur, Kcur, Vcur, nullptr, nullptr, nullptr, + 1.0f / sqrtf((float) n_embd_head), il); + if (vlora) { + cur = ggml_add(ctx0, ggml_mul_mat(ctx0, layer.wo, cur), lora(layer.zaya_vlora_o_a, layer.zaya_vlora_o_b, cur)); + } + cb(cur, "attn_out", il); + + // ---- post-attention residual scale ---- + ggml_tensor * hs_scaled = apply_res_scale(cur, layer.res_scale_hs, layer.res_scale_hs_b, "res_scale_hs", il); + ggml_tensor * res_scaled = apply_res_scale(residual, layer.res_scale_res, layer.res_scale_res_b, "res_scale_res", il); + residual = ggml_add(ctx0, hs_scaled, res_scaled); + cb(residual, "residual_post_attn", il); + + // ---- post-attention layernorm ---- + cur = build_norm(residual, layer.attn_post_norm, nullptr, LLM_NORM_RMS, il); + cb(cur, "post_attn_norm", il); + + // ===== MoE (every layer) ===== + + // EDA: the previous layer's router state, scaled. Built before the down projection so + // it is written before HRX's fused matmul + bias + add kernel reads it. + ggml_tensor * eda = nullptr; + if (prev_router != nullptr && layer.zaya_router_eda_scale != nullptr) { + eda = ggml_mul(ctx0, prev_router, layer.zaya_router_eda_scale); + ggml_build_forward_expand(gf, eda); + } + + ggml_tensor * router_h = ggml_mul_mat(ctx0, layer.ffn_gate_inp, cur); + if (layer.ffn_gate_inp_b) { + router_h = ggml_add(ctx0, router_h, layer.ffn_gate_inp_b); + } + cb(router_h, "router_down", il); + + if (eda != nullptr) { + router_h = ggml_add(ctx0, router_h, eda); + cb(router_h, "router_eda", il); + } + + prev_router = router_h; + + router_h = build_norm(router_h, layer.ffn_norm, nullptr, LLM_NORM_RMS, il); + cb(router_h, "router_norm", il); + + router_h = ggml_mul_mat(ctx0, layer.ffn_gate, router_h); + if (layer.ffn_gate_b) { + router_h = ggml_add(ctx0, router_h, layer.ffn_gate_b); + } + router_h = ggml_gelu(ctx0, router_h); + cb(router_h, "router_mlp0", il); + + router_h = ggml_mul_mat(ctx0, layer.zaya_router_mlp2, router_h); + if (layer.zaya_router_mlp2_b) { + router_h = ggml_add(ctx0, router_h, layer.zaya_router_mlp2_b); + } + router_h = ggml_gelu(ctx0, router_h); + cb(router_h, "router_mlp2", il); + + router_h = ggml_mul_mat(ctx0, layer.zaya_router_mlp4, router_h); + cb(router_h, "router_logits", il); + + router_h = ggml_soft_max(ctx0, router_h); + cb(router_h, "router_probs", il); + + ggml_tensor * gate_probs = ggml_cont(ctx0, + ggml_view_2d(ctx0, router_h, n_expert, n_tokens, router_h->nb[1], 0)); + cb(gate_probs, "gate_probs", il); + + ggml_tensor * expert_biases = nullptr; + if (layer.zaya_router_biases != nullptr) { + expert_biases = ggml_view_1d(ctx0, layer.zaya_router_biases, n_expert, 0); + } + + if (ubatch.embd != nullptr && layer.zaya_vlora_up_exps_a != nullptr) { + // the experts with their vision LoRA (fc1 before the SwiGLU, fc2 after it): the same + // top-1 choice and weight as build_moe_ffn below + ggml_tensor * probs = gate_probs; + if (expert_biases != nullptr) { + probs = ggml_add(ctx0, gate_probs, expert_biases); + } + ggml_tensor * sel = ggml_argsort_top_k(ctx0, probs, hparams.n_expert_used); // [k, T] + ggml_tensor * w = ggml_get_rows(ctx0, ggml_reshape_3d(ctx0, gate_probs, 1, n_expert, n_tokens), sel); // [1, k, T] + ggml_tensor * x = ggml_reshape_3d(ctx0, cur, n_embd, 1, n_tokens); + ggml_tensor * up = ggml_mul_mat_id(ctx0, layer.ffn_gate_up_exps, x, sel); // [2 n_ff, k, T] + up = ggml_add(ctx0, up, ggml_mul_mat_id(ctx0, layer.zaya_vlora_up_exps_b, + ggml_mul_mat_id(ctx0, layer.zaya_vlora_up_exps_a, x, sel), sel)); + ggml_tensor * act = ggml_swiglu(ctx0, up); // [n_ff, k, T] + ggml_tensor * dn = ggml_mul_mat_id(ctx0, layer.ffn_down_exps, act, sel); // [n_embd, k, T] + dn = ggml_add(ctx0, dn, ggml_mul_mat_id(ctx0, layer.zaya_vlora_down_exps_b, + ggml_mul_mat_id(ctx0, layer.zaya_vlora_down_exps_a, act, sel), sel)); + dn = ggml_mul(ctx0, dn, w); + cur = ggml_sum_rows(ctx0, ggml_cont(ctx0, ggml_permute(ctx0, dn, 1, 0, 2, 3))); // [1, n_embd, T] + cur = ggml_reshape_2d(ctx0, cur, n_embd, n_tokens); + } else { + cur = build_moe_ffn(cur, + /* gate_inp */ nullptr, + /* gate_inp_b */ nullptr, + /* up_exps */ nullptr, + /* up_exps_b */ nullptr, + /* gate_exps */ nullptr, + /* gate_exps_b */ nullptr, + /* down_exps */ layer.ffn_down_exps, + /* down_exps_b */ nullptr, + /* exp_probs_b */ expert_biases, + /* n_expert */ n_expert, + /* n_expert_used */ hparams.n_expert_used, + /* type_op */ LLM_FFN_SILU, + /* norm_w */ false, + /* w_scale */ 1.0f, + /* gating_op */ LLAMA_EXPERT_GATING_FUNC_TYPE_NONE, + /* il */ il, + /* probs_in */ gate_probs, + /* gate_up_exps */ layer.ffn_gate_up_exps, + /* gate_up_exps_b */ nullptr, + /* up_exps_s */ nullptr, + /* gate_exps_s */ nullptr, + /* down_exps_s */ nullptr); + } + cb(cur, "moe_out", il); + + // The router picks top-1 over n_expert + 1 slots; the last is a skip expert whose + // output is zero (HF ZayaRouter masks it). build_moe_ffn chose the best of the real + // experts; keep its output only where that expert's biased score beats the skip slot's. + if (layer.zaya_router_biases != nullptr) { + ggml_tensor * biased = ggml_add(ctx0, router_h, layer.zaya_router_biases); // [n_expert + 1, T] + ggml_tensor * biased_e = ggml_cont(ctx0, ggml_view_2d(ctx0, biased, n_expert, n_tokens, biased->nb[1], 0)); + ggml_tensor * best = ggml_argsort_top_k(ctx0, biased_e, 1); // [1, T] + ggml_tensor * best_v = ggml_get_rows(ctx0, + ggml_reshape_3d(ctx0, biased_e, 1, n_expert, n_tokens), best); // [1, 1, T] + ggml_tensor * skip_v = ggml_view_2d(ctx0, biased, 1, n_tokens, biased->nb[1], + n_expert*ggml_element_size(biased)); + ggml_tensor * keep = ggml_step(ctx0, ggml_sub(ctx0, + ggml_reshape_2d(ctx0, best_v, 1, n_tokens), ggml_cont(ctx0, skip_v))); // [1, T] + cb(keep, "moe_keep", il); + cur = ggml_mul(ctx0, cur, keep); + } + + // ---- post-MLP residual scale ---- + hs_scaled = apply_res_scale(cur, layer.res_scale_hs_mlp, layer.res_scale_hs_mlp_b, "res_scale_hs_mlp", il); + res_scaled = apply_res_scale(residual, layer.res_scale_res_mlp, layer.res_scale_res_mlp_b, "res_scale_res_mlp", il); + inpL = ggml_add(ctx0, hs_scaled, res_scaled); + cb(inpL, "layer_out", il); + } + + cur = inpL; + + if (inp_out_ids) { + cur = ggml_get_rows(ctx0, cur, inp_out_ids); + } + + cur = build_norm(cur, model.output_norm, nullptr, LLM_NORM_RMS, -1); + cb(cur, "result_norm", -1); + res->t_embd = cur; + + cur = ggml_mul_mat(ctx0, model.output, cur); + cb(cur, "result_output", -1); + + cur = ggml_cont(ctx0, ggml_cast(ctx0, cur, GGML_TYPE_F32)); + cb(cur, "result_output_fp32", -1); + + res->t_logits = cur; + + ggml_build_forward_expand(gf, cur); +} + +template struct llama_model_zaya::graph; +template struct llama_model_zaya::graph; diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 881e55c75a1d..219bbbaacaa2 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -336,3 +336,58 @@ if (TARGET gguf-model-data) target_link_libraries(test-export-graph-ops PRIVATE gguf-model-data) target_compile_definitions(test-export-graph-ops PRIVATE LLAMA_HF_FETCH) endif() + +if (TARGET ggml-hrx) + add_executable(test-hrx-buffer test-hrx-buffer.cpp) + target_link_libraries(test-hrx-buffer PRIVATE ggml-hrx ggml hrx::hrx) + target_include_directories(test-hrx-buffer PRIVATE ../ggml/include ../ggml/src ../ggml/src/ggml-hrx) + add_test(NAME test-hrx-buffer COMMAND test-hrx-buffer) + + add_executable(hrx-backend-test hrx-backend-test.cpp) + target_link_libraries(hrx-backend-test PRIVATE ggml-hrx ggml hrx::hrx loomc::loomc) + target_include_directories(hrx-backend-test PRIVATE ../ggml/include ../ggml/src ../ggml/src/ggml-hrx) + add_test(NAME hrx-backend-test COMMAND hrx-backend-test) + + add_executable(test-hrx-loom-jit test-hrx-loom-jit.cpp) + target_link_libraries(test-hrx-loom-jit PRIVATE ggml-hrx ggml-hrx-kernel-corpus hrx::hrx loomc::loomc) + target_include_directories(test-hrx-loom-jit PRIVATE ../ggml/include ../ggml/src ../ggml/src/ggml-hrx) + add_test(NAME test-hrx-loom-jit COMMAND test-hrx-loom-jit) + + add_executable(test-hrx-ops test-hrx-ops.cpp) + target_link_libraries(test-hrx-ops PRIVATE ggml-hrx ggml hrx::hrx loomc::loomc) + target_include_directories(test-hrx-ops PRIVATE ../ggml/include ../ggml/src ../ggml/src/ggml-hrx) + add_test(NAME test-hrx-ops COMMAND test-hrx-ops) + + add_executable(test-hrx-hadamard test-hrx-hadamard.cpp) + target_link_libraries(test-hrx-hadamard PRIVATE ggml) + add_test(NAME test-hrx-hadamard COMMAND test-hrx-hadamard) + + add_executable(test-hrx-mxfp4 test-hrx-mxfp4.cpp) + target_link_libraries(test-hrx-mxfp4 PRIVATE ggml-hrx ggml hrx::hrx loomc::loomc) + target_include_directories(test-hrx-mxfp4 PRIVATE ../ggml/include ../ggml/src ../ggml/src/ggml-hrx) + add_test(NAME test-hrx-mxfp4 COMMAND test-hrx-mxfp4) + + add_executable(test-hrx-decode-stride test-hrx-decode-stride.cpp) + target_link_libraries(test-hrx-decode-stride PRIVATE ggml-hrx ggml hrx::hrx loomc::loomc) + target_include_directories(test-hrx-decode-stride PRIVATE ../ggml/include ../ggml/src ../ggml/src/ggml-hrx) + add_test(NAME test-hrx-decode-stride COMMAND test-hrx-decode-stride) + + add_executable(test-hrx-mul-mat-id-k32 test-hrx-mul-mat-id-k32.cpp) + target_link_libraries(test-hrx-mul-mat-id-k32 PRIVATE ggml-hrx ggml hrx::hrx loomc::loomc) + target_include_directories(test-hrx-mul-mat-id-k32 PRIVATE ../ggml/include ../ggml/src ../ggml/src/ggml-hrx) + add_test(NAME test-hrx-mul-mat-id-k32 COMMAND test-hrx-mul-mat-id-k32) + + add_executable(test-hrx-moe-split test-hrx-moe-split.cpp) + target_link_libraries(test-hrx-moe-split PRIVATE ggml) + add_test(NAME test-hrx-moe-split COMMAND test-hrx-moe-split) + + add_executable(test-hrx-fa-masked-v test-hrx-fa-masked-v.cpp) + target_link_libraries(test-hrx-fa-masked-v PRIVATE ggml-hrx ggml hrx::hrx loomc::loomc) + target_include_directories(test-hrx-fa-masked-v PRIVATE ../ggml/include ../ggml/src ../ggml/src/ggml-hrx) + add_test(NAME test-hrx-fa-masked-v COMMAND test-hrx-fa-masked-v) + + add_executable(test-hrx-attention-sink test-hrx-attention-sink.cpp) + target_link_libraries(test-hrx-attention-sink PRIVATE ggml-hrx ggml hrx::hrx loomc::loomc) + target_include_directories(test-hrx-attention-sink PRIVATE ../ggml/include ../ggml/src ../ggml/src/ggml-hrx) + add_test(NAME test-hrx-attention-sink COMMAND test-hrx-attention-sink) +endif() diff --git a/tests/hrx-backend-test.cpp b/tests/hrx-backend-test.cpp new file mode 100644 index 000000000000..2c2f930b7f10 --- /dev/null +++ b/tests/hrx-backend-test.cpp @@ -0,0 +1,13833 @@ +#include "backend-buffer-binding.h" +#include "backend-context.h" +#include "dispatch/command-program-bindings.h" +#include "dispatch/command-program-diagnostics.h" +#include "dispatch/command-program-resolver.h" +#include "dispatch/command-program.h" +#include "dispatch/dispatch-scheduler.h" +#include "dispatch_registration/common/dispatch-gather-add.h" +#include "dispatch_registration/common/dispatch-symmetric-i4.h" +#include "dispatch_registration/dispatch-registry.h" +#include "ggml-alloc.h" +#include "ggml-backend.h" +#include "ggml-cpu.h" +#include "ggml-hrx.h" +#include "ggml-impl.h" +#include "ggml.h" +#include "graph/graph-diagnostics.h" +#include "graph/graph-matcher.h" +#include "graph/graph-traversal.h" +#include "graph/graph.h" +#include "hrx-interop-utils.h" +#include "kernel-corpus/kernel-corpus.h" +#include "runtime/command-program-executor.h" +#include "runtime/graph-executor.h" +#include "runtime/graph-program-cache.h" +#include "runtime/graph-replay.h" +#include "runtime/loom-kernel-jit.h" +#include "testing_suite.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#define REQUIRE(condition) \ + do { \ + if (!(condition)) { \ + std::fprintf(stderr, "%s:%d: requirement failed: %s\n", __FILE__, __LINE__, #condition); \ + std::abort(); \ + } \ + } while (false) + +static hrx_buffer_t dummy_hrx_buffer(uintptr_t value) { + return reinterpret_cast(value); +} + +static hrx_stream_t dummy_hrx_stream(uintptr_t value) { + return reinterpret_cast(value); +} + +static hrx_graph_exec_t dummy_hrx_graph_exec(uintptr_t value) { + return reinterpret_cast(value); +} + +static bool contains_value_id(const std::vector & ids, ggml::hrx::ValueId id) { + for (const ggml::hrx::ValueId candidate : ids) { + if (candidate == id) { + return true; + } + } + return false; +} + +static bool command_program_verifies(const ggml::hrx::CommandProgram & program) { + const ggml::hrx::VerificationResult result = + ggml::hrx::verify_command_program(program, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + if (!result.valid()) { + for (const std::string & error : result.status.errors()) { + std::fprintf(stderr, "%s\n", error.c_str()); + } + } + return result.valid(); +} + +static bool async_jit_expected_from_environment() { + const char * value = std::getenv("GGML_HRX_ASYNC_JIT"); + return value == nullptr || + (std::strcmp(value, "0") != 0 && std::strcmp(value, "false") != 0 && std::strcmp(value, "FALSE") != 0 && + std::strcmp(value, "off") != 0 && std::strcmp(value, "OFF") != 0); +} + +static ggml::hrx::CommandProgram copy_command_program_shape(const ggml::hrx::CommandProgram & program) { + ggml::hrx::CommandProgram copy; + copy.initialization_commands = program.initialization_commands; + copy.commands = program.commands; + copy.transients = program.transients; + copy.completion_counters = program.completion_counters; + copy.constant_initializations = program.constant_initializations; + return copy; +} + +static bool status_contains(const ggml::hrx::Status & status, const char * text) { + for (const std::string & message : status.errors()) { + if (message.find(text) != std::string::npos) { + return true; + } + } + return false; +} + +static bool string_contains(const std::string & value, const char * text) { + return value.find(text) != std::string::npos; +} + +static std::string read_text_file(const std::filesystem::path & path) { + std::ifstream input(path, std::ios::binary); + REQUIRE(input.good()); + return { std::istreambuf_iterator(input), std::istreambuf_iterator() }; +} + +static std::vector list_directories(const std::filesystem::path & path) { + std::vector directories; + if (!std::filesystem::exists(path)) { + return directories; + } + for (const std::filesystem::directory_entry & entry : std::filesystem::directory_iterator(path)) { + if (entry.is_directory()) { + directories.push_back(entry.path()); + } + } + std::sort(directories.begin(), directories.end()); + return directories; +} + +struct EnvironmentVariableGuard { + const char * name; + bool had_value = false; + std::string value; + + explicit EnvironmentVariableGuard(const char * name) : name(name) { + const char * current = std::getenv(name); + if (current != nullptr) { + had_value = true; + value = current; + } + } + + ~EnvironmentVariableGuard() { + if (had_value) { + setenv(name, value.c_str(), 1); + } else { + unsetenv(name); + } + } + + void set(const std::filesystem::path & path) const { setenv(name, path.string().c_str(), 1); } + + void unset() const { unsetenv(name); } +}; + +static std::filesystem::path fresh_test_directory(const char * name) { + static uint64_t sequence = 0; + std::filesystem::path path = std::filesystem::temp_directory_path() / + ("hrx-backend-test-" + std::string(name) + "-" + std::to_string(sequence++)); + std::filesystem::remove_all(path); + std::filesystem::create_directories(path); + return path; +} + +static void require_hrx_status(hrx_status_t status) { + if (ggml::hrx::ErrorResult error = ggml::hrx::take_status(status)) { + std::fprintf(stderr, "HRX status failed: %s\n", error->c_str()); + std::abort(); + } +} + +static ggml::hrx::DispatchTarget test_dispatch_target() { + return { "gfx1151" }; +} + +static const ggml::hrx::DispatchRegistry & test_dispatch_registry() { + const ggml::hrx::DispatchRegistry * registry = ggml::hrx::find_dispatch_registry(test_dispatch_target()); + REQUIRE(registry != nullptr); + return *registry; +} + +static std::string kernel_name_for_id(uint64_t kernel_id) { + const ggml::hrx::KernelResolveResult resolved = + ggml::hrx::resolve_kernel_definition(ggml::hrx::get_qwen_kernel_corpus(), "gfx1151", kernel_id); + REQUIRE(resolved.found()); + return ggml::hrx::kernel_definition_name(*resolved.definition); +} + +template static void require_scheduled_command_program(ggml_cgraph * graph, Check && check) { + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + check(imported.graph, scheduler.plan(), commands); +} + +static std::vector scheduled_kernel_names(ggml_cgraph * graph) { + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + + std::vector result; + for (const ggml::hrx::Dispatch & dispatch : scheduler.plan().dispatches) { + result.push_back(kernel_name_for_id(dispatch.kernel.kernel_id)); + } + return result; +} + +static void require_compile_parameter(const ggml::hrx::Dispatch & dispatch, + const char * name, + const std::string & value) { + const auto found = dispatch.kernel.compile_parameters.find(name); + REQUIRE(found != dispatch.kernel.compile_parameters.end()); + if (found->second != value) { + std::fprintf(stderr, "compile parameter %s expected '%s' got '%s'\n", name, value.c_str(), + found->second.c_str()); + std::abort(); + } +} + +static std::string expected_config_value(float value) { + std::ostringstream out; + out.precision(9); + out << value; + return out.str(); +} + +static constexpr int64_t kQwenFlashHeadSize = 128; +static constexpr int64_t kQwenRouterExpertCount = 128; +static constexpr int64_t kQwenRouterRouteCount = 8; +static constexpr int64_t kQwenMoeHiddenSize = 2048; +static constexpr int64_t kQwenMoeIntermediateSize = 768; + +static size_t qwen_expert_table_size(int64_t token_count, int64_t expert_count = kQwenRouterExpertCount) { + return static_cast(expert_count + expert_count * token_count) * sizeof(int32_t); +} + +static size_t qwen_partition_table_size(int64_t token_count, + int64_t route_count = kQwenRouterRouteCount, + int64_t expert_count = kQwenRouterExpertCount) { + const int64_t assignment_count = token_count * route_count; + const int64_t assignment_partition_count = (assignment_count + 31) / 32; + return static_cast(1 + assignment_partition_count + expert_count) * sizeof(int32_t); +} + +static size_t qwen_q8_1_x4_size(int64_t token_count, int64_t hidden_size) { + return static_cast(token_count) * ggml_row_size(GGML_TYPE_Q8_1, hidden_size); +} + +static int64_t matmul_weight_format_config(ggml_type type) { + switch (type) { + case GGML_TYPE_Q1_0: + return 10; + case GGML_TYPE_Q3_K: + return 11; + case GGML_TYPE_Q4_K: + return 4; + case GGML_TYPE_Q5_K: + return 5; + case GGML_TYPE_Q6_K: + return 6; + case GGML_TYPE_Q4_0: + return 40; + case GGML_TYPE_Q4_1: + return 41; + case GGML_TYPE_Q5_0: + return 50; + case GGML_TYPE_Q5_1: + return 51; + case GGML_TYPE_IQ1_S: + return 19; + case GGML_TYPE_IQ1_M: + return 29; + case GGML_TYPE_IQ2_S: + return 22; + case GGML_TYPE_IQ3_S: + return 21; + case GGML_TYPE_IQ4_NL: + return 20; + case GGML_TYPE_IQ4_XS: + return 23; + case GGML_TYPE_Q8_0: + return 80; + case GGML_TYPE_Q8_1: + return 81; + case GGML_TYPE_F16: + return 16; + case GGML_TYPE_BF16: + return 30; + case GGML_TYPE_F32: + return 32; + default: + REQUIRE(false); + return 0; + } +} + +static std::vector make_qwen_route_ids_iota(int64_t token_count, + int64_t route_count, + int64_t route_stride, + int64_t expert_count) { + std::vector route_ids(static_cast(token_count * route_stride), -1); + for (int64_t token = 0; token < token_count; ++token) { + for (int64_t route = 0; route < route_count; ++route) { + route_ids[static_cast(token * route_stride + route)] = + static_cast((token * route_count + route) % expert_count); + } + } + return route_ids; +} + +static std::vector make_qwen_expert_table_reference(const std::vector & route_ids, + int64_t token_count, + int64_t route_count, + int64_t route_stride, + int64_t expert_count) { + std::vector expert_table(qwen_expert_table_size(token_count, expert_count) / sizeof(int32_t), -1); + for (int64_t expert = 0; expert < expert_count; ++expert) { + expert_table[static_cast(expert)] = 0; + } + for (int64_t token = 0; token < token_count; ++token) { + for (int64_t route = 0; route < route_count; ++route) { + const int32_t expert = route_ids[static_cast(token * route_stride + route)]; + REQUIRE(expert >= 0); + REQUIRE(expert < expert_count); + int32_t & count = expert_table[static_cast(expert)]; + const size_t assignment_offset = + static_cast(expert_count + static_cast(expert) * token_count + count); + expert_table[assignment_offset] = static_cast(token * route_count + route); + ++count; + } + } + return expert_table; +} + +static std::vector make_qwen_partition_table_reference(const std::vector & expert_table, + int64_t token_count, + int64_t route_count, + int64_t expert_count) { + std::vector partition_table( + qwen_partition_table_size(token_count, route_count, expert_count) / sizeof(int32_t), -1); + int32_t partition_count = 0; + for (int64_t expert = 0; expert < expert_count; ++expert) { + const int32_t expert_assignment_count = expert_table[static_cast(expert)]; + REQUIRE(expert_assignment_count >= 0); + const int32_t expert_partition_count = (expert_assignment_count + 31) / 32; + for (int32_t partition = 0; partition < expert_partition_count; ++partition) { + int32_t row_count = expert_assignment_count - partition * 32; + if (row_count > 32) { + row_count = 32; + } + const int32_t descriptor = static_cast(expert) | (partition << 7) | ((row_count - 1) << 13); + partition_table[static_cast(1 + partition_count)] = descriptor; + ++partition_count; + } + } + partition_table[0] = partition_count; + return partition_table; +} + +static void require_qwen_expert_table_matches(const std::vector & actual, + const std::vector & expected, + int64_t token_count, + int64_t expert_count) { + REQUIRE(actual.size() == expected.size()); + for (int64_t expert = 0; expert < expert_count; ++expert) { + const size_t count_index = static_cast(expert); + REQUIRE(actual[count_index] == expected[count_index]); + for (int32_t ordinal = 0; ordinal < expected[count_index]; ++ordinal) { + const size_t assignment_index = + static_cast(expert_count + expert * token_count + static_cast(ordinal)); + REQUIRE(actual[assignment_index] == expected[assignment_index]); + } + } +} + +static void require_qwen_partition_table_matches(const std::vector & actual, + const std::vector & expected) { + REQUIRE(actual.size() == expected.size()); + REQUIRE(actual[0] == expected[0]); + for (int32_t i = 0; i < actual[0]; ++i) { + REQUIRE(actual[static_cast(1 + i)] == expected[static_cast(1 + i)]); + } +} + +static size_t qwen_routed_gate_up_f16_output_size(int64_t token_count) { + return static_cast(token_count * kQwenRouterRouteCount * kQwenMoeIntermediateSize) * sizeof(ggml_fp16_t); +} + +static size_t qwen_routed_down_f16_output_size(int64_t token_count) { + return static_cast(token_count * kQwenRouterRouteCount * kQwenMoeHiddenSize) * sizeof(ggml_fp16_t); +} + +static void set_qwen_flash_query_layout(ggml_tensor * tensor, int64_t head_count, int64_t head_size) { + REQUIRE(tensor != nullptr); + tensor->nb[0] = sizeof(float); + tensor->nb[1] = static_cast(head_count * head_size) * sizeof(float); + tensor->nb[2] = static_cast(head_size) * sizeof(float); +} + +static void set_qwen_flash_key_value_layout(ggml_tensor * tensor, int64_t head_count, int64_t head_size) { + REQUIRE(tensor != nullptr); + tensor->nb[0] = sizeof(ggml_fp16_t); + tensor->nb[1] = static_cast(head_count * head_size) * sizeof(ggml_fp16_t); + tensor->nb[2] = static_cast(head_size) * sizeof(ggml_fp16_t); +} + +static ggml_tensor * build_qwen_flash_attention_graph(ggml_context * ctx, + int64_t query_token_count, + int64_t key_value_token_count, + int64_t query_head_count, + int64_t key_value_head_count, + ggml_type query_type = GGML_TYPE_F32, + ggml_type key_value_type = GGML_TYPE_F16, + bool include_mask = true, + bool include_sinks = false, + int64_t head_size = kQwenFlashHeadSize, + float scale = 1.0f / std::sqrt(128.0f), + int64_t value_head_size = -1, + float max_bias = 0.0f, + float logit_softcap = 0.0f) { + const int64_t actual_value_head_size = value_head_size > 0 ? value_head_size : head_size; + ggml_tensor * query = ggml_new_tensor_3d(ctx, query_type, head_size, query_token_count, query_head_count); + ggml_tensor * key = ggml_new_tensor_3d(ctx, key_value_type, head_size, key_value_token_count, key_value_head_count); + ggml_tensor * value = + ggml_new_tensor_3d(ctx, key_value_type, actual_value_head_size, key_value_token_count, key_value_head_count); + REQUIRE(query != nullptr); + REQUIRE(key != nullptr); + REQUIRE(value != nullptr); + if (query_type == GGML_TYPE_F32) { + set_qwen_flash_query_layout(query, query_head_count, head_size); + } + if (key_value_type == GGML_TYPE_F16) { + set_qwen_flash_key_value_layout(key, key_value_head_count, head_size); + set_qwen_flash_key_value_layout(value, key_value_head_count, actual_value_head_size); + } + ggml_tensor * mask = nullptr; + if (include_mask) { + mask = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, key_value_token_count, query_token_count); + REQUIRE(mask != nullptr); + } + ggml_tensor * output = ggml_flash_attn_ext(ctx, query, key, value, mask, scale, max_bias, logit_softcap); + REQUIRE(output != nullptr); + if (include_sinks) { + ggml_tensor * sinks = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, query_head_count); + REQUIRE(sinks != nullptr); + ggml_flash_attn_ext_add_sinks(output, sinks); + } + return output; +} + +static ggml_tensor * build_qwen_router_top8_graph(ggml_context * ctx, + ggml_tensor * logits, + ggml_tensor ** route_ids = nullptr, + ggml_sort_order order = GGML_SORT_ORDER_DESC, + int64_t route_count = kQwenRouterRouteCount, + float clamp_min = 1.0e-7f) { + ggml_tensor * probs = ggml_soft_max(ctx, logits); + REQUIRE(probs != nullptr); + ggml_tensor * probs_reshaped = ggml_reshape_3d(ctx, probs, 1, logits->ne[0], logits->ne[1]); + REQUIRE(probs_reshaped != nullptr); + ggml_tensor * argsort = ggml_argsort(ctx, probs, order); + REQUIRE(argsort != nullptr); + ggml_tensor * topk = ggml_view_2d(ctx, argsort, route_count, logits->ne[1], argsort->nb[1], 0); + REQUIRE(topk != nullptr); + if (route_ids != nullptr) { + *route_ids = topk; + } + ggml_tensor * selected = ggml_get_rows(ctx, probs_reshaped, topk); + REQUIRE(selected != nullptr); + ggml_tensor * selected_reshaped = ggml_reshape_2d(ctx, selected, route_count, logits->ne[1]); + REQUIRE(selected_reshaped != nullptr); + ggml_tensor * sum = ggml_sum_rows(ctx, selected_reshaped); + REQUIRE(sum != nullptr); + ggml_tensor * clamped_sum = ggml_clamp(ctx, sum, clamp_min, std::numeric_limits::infinity()); + REQUIRE(clamped_sum != nullptr); + ggml_tensor * normalized = ggml_div(ctx, selected_reshaped, clamped_sum); + REQUIRE(normalized != nullptr); + ggml_tensor * output = ggml_reshape_3d(ctx, normalized, 1, route_count, logits->ne[1]); + REQUIRE(output != nullptr); + return output; +} + +static std::vector traversal_indices(const ggml::hrx::Graph & graph) { + const ggml::hrx::GraphTraversalOrder order = ggml::hrx::GraphTraversalOrder::build(graph); + std::vector indices; + indices.reserve(order.nodes().size()); + for (const ggml::hrx::GraphNode * node : order.nodes()) { + size_t index = 0; + REQUIRE(node != nullptr); + REQUIRE(graph.index().node_index(node, index)); + indices.push_back(index); + } + return indices; +} + +static size_t find_position(const std::vector & indices, size_t node_index) { + const std::vector::const_iterator it = std::find(indices.begin(), indices.end(), node_index); + REQUIRE(it != indices.end()); + return static_cast(it - indices.begin()); +} + +static size_t producer_index_for_tensor(const ggml::hrx::Graph & graph, const ggml_tensor * tensor) { + const ggml::hrx::Value * value = graph.values().find_tensor(tensor); + REQUIRE(value != nullptr); + const ggml::hrx::GraphNode * producer = graph.index().producer(value->id); + REQUIRE(producer != nullptr); + size_t index = 0; + REQUIRE(graph.index().node_index(producer, index)); + return index; +} + +static bool match_dispatch_at_index(const ggml::hrx::Graph & graph, + const ggml::hrx::CommandPlan & plan, + const std::vector & covered_nodes, + size_t node_index, + ggml::hrx::DispatchMatch & match); + +static void append_match_to_plan(ggml::hrx::CommandPlan & plan, + ggml::hrx::DispatchMatch & match, + std::vector & covered_nodes, + ggml::hrx::Graph * graph = nullptr); + +static void run_status_checks() { + ggml::hrx::Status status; + REQUIRE(status.success()); + REQUIRE(status.errors().empty()); + + status.log("first"); + REQUIRE(!status.success()); + REQUIRE(!status.errors().empty()); + REQUIRE(status.errors().size() == 1); + REQUIRE(status.errors()[0] == "first"); + + status.log("value %d", 7); + REQUIRE(status.errors().size() == 2); + REQUIRE(status.errors()[1] == "value 7"); + + ggml::hrx::Status other; + other.log("third"); + status.append(other); + REQUIRE(status.errors().size() == 3); + REQUIRE(status.errors()[2] == "third"); +} + +static void run_command_plan_metadata_checks() { + const ggml::hrx::MoeRoutingResourceMetadata routing = { + 4, + 8, + 128, + 128, + }; + const ggml::hrx::CommandPlanResourceMetadata metadata = ggml::hrx::make_command_plan_resource_metadata(routing); + + REQUIRE(metadata.kind == ggml::hrx::CommandPlanResourceMetadataKind::MoeRoutingResource); + ggml::hrx::MoeRoutingResourceMetadata decoded; + REQUIRE(metadata.read(decoded)); + REQUIRE(decoded.token_count == routing.token_count); + REQUIRE(decoded.route_count == routing.route_count); + REQUIRE(decoded.route_stride == routing.route_stride); + REQUIRE(decoded.expert_count == routing.expert_count); + + const ggml::hrx::CommandPlanResourceMetadata empty; + REQUIRE(!empty.read(decoded)); + + ggml::hrx::CommandPlanMetadata metadata_plan; + ggml::hrx::Status status; + REQUIRE(metadata_plan.append_alternate_value( + { ggml::hrx::ValueId(1), ggml::hrx::ValueId(2), GGML_TYPE_F16, 16, "alternate" }, status)); + REQUIRE(metadata_plan.append_alternate_value( + { ggml::hrx::ValueId(1), ggml::hrx::ValueId(2), GGML_TYPE_F16, 16, "alternate" }, status)); + REQUIRE(metadata_plan.alternate_values().size() == 1); + REQUIRE(metadata_plan.find_alternate_value(ggml::hrx::ValueId(1), GGML_TYPE_F16, 16) != nullptr); + REQUIRE(metadata_plan.find_alternate_value(ggml::hrx::ValueId(1), GGML_TYPE_F32, 16) == nullptr); + REQUIRE(metadata_plan.find_alternate_value(ggml::hrx::ValueId(1), GGML_TYPE_F16, 32) == nullptr); + REQUIRE(!metadata_plan.append_alternate_value( + { ggml::hrx::ValueId(1), ggml::hrx::ValueId(3), GGML_TYPE_F16, 16, "alternate" }, status)); + REQUIRE(!status.success()); + + ggml::hrx::CommandPlanMetadata bundle_plan; + ggml::hrx::Status bundle_status; + const ggml::hrx::CommandPlanMoeRoutingBundle bundle = { + ggml::hrx::ValueId(10), + ggml::hrx::ValueId(11), + ggml::hrx::ValueId(12), + ggml::hrx::ValueId(13), + 128, + 64, + 4, + 8, + 128, + 128, + }; + REQUIRE(bundle_plan.append_moe_routing_bundle(bundle, bundle_status)); + REQUIRE(bundle_plan.append_moe_routing_bundle(bundle, bundle_status)); + REQUIRE(bundle_plan.moe_routing_bundles().size() == 1); + const ggml::hrx::CommandPlanMoeRoutingBundle * found_bundle = + bundle_plan.find_moe_routing_bundle(ggml::hrx::ValueId(10)); + REQUIRE(found_bundle != nullptr); + REQUIRE(found_bundle->route_weights == ggml::hrx::ValueId(11)); + REQUIRE(found_bundle->expert_table == ggml::hrx::ValueId(12)); + REQUIRE(found_bundle->partition_table == ggml::hrx::ValueId(13)); + ggml::hrx::CommandPlanMoeRoutingBundle conflicting_bundle = bundle; + conflicting_bundle.route_weights = ggml::hrx::ValueId(14); + REQUIRE(!bundle_plan.append_moe_routing_bundle(conflicting_bundle, bundle_status)); + REQUIRE(!bundle_status.success()); +} + +static bool has_dispatch_registration(const std::vector & registrations, + const char * name) { + for (const ggml::hrx::DispatchRegistration & registration : registrations) { + if (std::string(registration.name) == name) { + return true; + } + } + return false; +} + +static bool has_dispatch_registration_kind(const std::vector & registrations, + const char * name, + ggml::hrx::DispatchMatchKind kind) { + for (const ggml::hrx::DispatchRegistration & registration : registrations) { + if (std::string(registration.name) == name && registration.kind == kind) { + return true; + } + } + return false; +} + +static void restore_environment_value(const char * name, bool had_value, const std::string & value) { + if (had_value) { + REQUIRE(setenv(name, value.c_str(), 1) == 0); + } else { + REQUIRE(unsetenv(name) == 0); + } +} + +static void add_binary_f32_exact_dispatch_params(ggml::hrx::Dispatch & dispatch, int64_t element_count) { + dispatch.kernel.integer_parameters.emplace("element_count", element_count); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.ne0", std::to_string(element_count)); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.ne1", "1"); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.ne2", "1"); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src0_stride1", std::to_string(element_count)); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src0_stride2", std::to_string(element_count)); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src0_stride3", std::to_string(element_count)); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src1_stride1", std::to_string(element_count)); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src1_stride2", std::to_string(element_count)); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src1_stride3", std::to_string(element_count)); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src0_span", std::to_string(element_count)); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.src1_span", std::to_string(element_count)); + dispatch.kernel.compile_parameters.emplace("ggml.binary_f32.op", "0"); +} + +static bool match_test_single_dispatch(const ggml::hrx::DispatchMatchContext & context, + ggml::hrx::DispatchMatch & match) { + ggml::hrx::Dispatch dispatch; + dispatch.kernel.integer_parameters.emplace("route", 1); + match.covered_nodes.push_back(context.root_index); + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_test_fused_dispatch(const ggml::hrx::DispatchMatchContext & context, + ggml::hrx::DispatchMatch & match) { + ggml::hrx::Dispatch dispatch; + dispatch.kernel.integer_parameters.emplace("route", 2); + match.covered_nodes.push_back(context.root_index); + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static bool match_test_wrong_root_dispatch(const ggml::hrx::DispatchMatchContext & context, + ggml::hrx::DispatchMatch & match) { + ggml::hrx::Dispatch dispatch; + dispatch.kernel.integer_parameters.emplace("route", 3); + match.covered_nodes.push_back(context.root_index); + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +static void run_dispatch_registry_checks() { + static constexpr const char * kDisableQwenDispatchEnv = "GGML_HRX_DISABLE_QWEN_DISPATCH"; + const char * original_env = std::getenv(kDisableQwenDispatchEnv); + const bool had_original_env = original_env != nullptr; + const std::string original_env_value = had_original_env ? original_env : ""; + REQUIRE(unsetenv(kDisableQwenDispatchEnv) == 0); + + const ggml::hrx::DispatchRegistry & registry = test_dispatch_registry(); + REQUIRE(ggml::hrx::find_dispatch_registry({ "gfx1100" }) != nullptr); + REQUIRE(ggml::hrx::find_dispatch_registry({ "gfx1151" }) != nullptr); + REQUIRE(ggml::hrx::find_dispatch_registry({ "gfx0000" }) == nullptr); + + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_ADD), "common.binary_f32")); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_SUB), "common.binary_f32")); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_MUL), "common.binary_f32")); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_DIV), "common.binary_f32")); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat_swiglu.tiled_pair_f32_f32", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat_swiglu.tiled_pair_postops_f32_f32", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.packed_mul_mat_glu.tiled_pair_f32_f32", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_MUL_MAT), + "llm.attention_qkv_matmul_postprocess.tiled_vector_f32_f32", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat.tiled_f32_f32", ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat.skinny_f32_f32", ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat.skinny_f32_f32_decode", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat_add.skinny_f32_f32_decode", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat_unary.tiled_f32_f32", ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat_unary.skinny_f32_f32", ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat_postops.tiled_f32_f32", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat_postops.skinny_f32_f32", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_MUL_MAT), "qwen.matmul.q6k_q8_1_x4")); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_MUL_MAT), + "llm.moe_router.projection_f32_four_row_wave32")); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_RMS_NORM), + "qwen.rmsnorm_f32_quantize_q8_1_x4")); + REQUIRE( + has_dispatch_registration(registry.registrations_for_root(GGML_OP_RMS_NORM), "common.rmsnorm_binary_q8_1_x4")); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_RMS_NORM), "common.rmsnorm_binary_f32")); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_RMS_NORM), "common.rmsnorm_f32")); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_FLASH_ATTN_EXT), + "common.flash_attention_f32_f16_wmma")); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_FLASH_ATTN_EXT), + "common.flash_attention_decode_split_next_q8")); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_RESHAPE), + "qwen.attention_postprocess_f32_f16")); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_GET_ROWS), "common.get_rows.f32")); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_GET_ROWS), "common.get_rows.f32_next")); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_GET_ROWS), "common.get_rows_scale.f32")); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_GET_ROWS), "common.gather_add_f32")); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_ROPE), "common.rope_set_rows.f32", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_ROPE), "common.rope_concat.f32", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_ROPE), "common.rope.f32", + ggml::hrx::DispatchMatchKind::SingleOp)); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_SET_ROWS), "common.set_rows")); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_SOFT_MAX), "llm.moe_router.top8_f32")); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_MUL_MAT_ID), + "common.mul_mat_id_swiglu.f32_f32_wmma", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_MUL_MAT_ID), + "common.mul_mat_id_postops.f32_f32_wmma", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(registry.registrations_for_root(GGML_OP_MUL_MAT_ID), + "common.mul_mat_id.f32_f32_wmma", ggml::hrx::DispatchMatchKind::SingleOp)); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_MUL_MAT_ID), + "llm.routed_ffn.gate_up_swiglu_f16_wmma")); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_MUL_MAT_ID), + "llm.routed_ffn.down_q4k_f16_wmma_grouped")); + REQUIRE(has_dispatch_registration(registry.registrations_for_root(GGML_OP_MUL_MAT_ID), + "llm.routed_ffn.down_q6k_f16_wmma_grouped")); + REQUIRE(registry.single_op_registrations().size() >= 5); + + REQUIRE(setenv(kDisableQwenDispatchEnv, "0", 1) == 0); + const ggml::hrx::DispatchRegistry & false_env_registry = test_dispatch_registry(); + REQUIRE( + has_dispatch_registration(false_env_registry.registrations_for_root(GGML_OP_GET_ROWS), "common.get_rows.f32")); + + REQUIRE(setenv(kDisableQwenDispatchEnv, "1", 1) == 0); + const ggml::hrx::DispatchRegistry & generic_registry = test_dispatch_registry(); + REQUIRE(has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_ADD), "common.binary_f32")); + REQUIRE(has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_RMS_NORM), + "common.rmsnorm_binary_f32")); + REQUIRE(has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_RMS_NORM), + "common.rmsnorm_binary_q8_1_x4")); + REQUIRE(has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_RMS_NORM), "common.rmsnorm_f32")); + REQUIRE( + has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_GET_ROWS), "common.gather_add_f32")); + REQUIRE( + has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_GET_ROWS), "common.get_rows.f32")); + REQUIRE(has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_GET_ROWS), + "common.get_rows.f32_next")); + REQUIRE(has_dispatch_registration_kind(generic_registry.registrations_for_root(GGML_OP_ROPE), + "common.rope_set_rows.f32", ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(generic_registry.registrations_for_root(GGML_OP_ROPE), "common.rope.f32", + ggml::hrx::DispatchMatchKind::SingleOp)); + REQUIRE(has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_SET_ROWS), "common.set_rows")); + REQUIRE(!has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_RMS_NORM), + "qwen.rmsnorm_f32_quantize_q8_1_x4")); + REQUIRE(has_dispatch_registration_kind(generic_registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat_swiglu.tiled_pair_f32_f32", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(generic_registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat_swiglu.tiled_pair_postops_f32_f32", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(generic_registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.packed_mul_mat_glu.tiled_pair_f32_f32", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(generic_registry.registrations_for_root(GGML_OP_MUL_MAT), + "llm.attention_qkv_matmul_postprocess.tiled_vector_f32_f32", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(generic_registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat.tiled_f32_f32", ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(generic_registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat.skinny_f32_f32", ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(generic_registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat.skinny_f32_f32_decode", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(generic_registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat_add.skinny_f32_f32_decode", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(generic_registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat_unary.tiled_f32_f32", ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(generic_registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat_unary.skinny_f32_f32", ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(generic_registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat_postops.tiled_f32_f32", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(generic_registry.registrations_for_root(GGML_OP_MUL_MAT), + "common.mul_mat_postops.skinny_f32_f32", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(!has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_MUL_MAT), + "qwen.matmul.q6k_q8_1_x4")); + REQUIRE(has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_FLASH_ATTN_EXT), + "common.flash_attention_f32_f16_wmma")); + REQUIRE(has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_FLASH_ATTN_EXT), + "common.flash_attention_decode_split_next_q8")); + REQUIRE(!has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_RESHAPE), + "qwen.attention_postprocess_f32_f16")); + REQUIRE(!has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_GET_ROWS), + "qwen.preamble.token_embedding_q4k")); + REQUIRE(!has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_SOFT_MAX), + "llm.moe_router.top8_f32")); + REQUIRE(has_dispatch_registration_kind(generic_registry.registrations_for_root(GGML_OP_MUL_MAT_ID), + "common.mul_mat_id_swiglu.f32_f32_wmma", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(generic_registry.registrations_for_root(GGML_OP_MUL_MAT_ID), + "common.mul_mat_id_postops.f32_f32_wmma", + ggml::hrx::DispatchMatchKind::Fused)); + REQUIRE(has_dispatch_registration_kind(generic_registry.registrations_for_root(GGML_OP_MUL_MAT_ID), + "common.mul_mat_id.f32_f32_wmma", ggml::hrx::DispatchMatchKind::SingleOp)); + REQUIRE(!has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_MUL_MAT_ID), + "llm.routed_ffn.gate_up_swiglu_f16_wmma")); + REQUIRE(!has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_MUL_MAT_ID), + "llm.routed_ffn.down_q4k_f16_wmma_grouped")); + REQUIRE(!has_dispatch_registration(generic_registry.registrations_for_root(GGML_OP_MUL_MAT_ID), + "llm.routed_ffn.down_q6k_f16_wmma_grouped")); + restore_environment_value(kDisableQwenDispatchEnv, had_original_env, original_env_value); + + ggml::hrx::DispatchRegistryBuilder builder; + builder.add({ + "test.single_add", + GGML_OP_ADD, + ggml::hrx::DispatchMatchKind::SingleOp, + 1000, + ggml::hrx::DispatchSource::Common, + match_test_single_dispatch, + }); + builder.add({ + "test.fused_add", + GGML_OP_ADD, + ggml::hrx::DispatchMatchKind::Fused, + 0, + ggml::hrx::DispatchSource::Common, + match_test_fused_dispatch, + }); + builder.add({ + "test.wrong_root", + GGML_OP_MUL_MAT, + ggml::hrx::DispatchMatchKind::Fused, + 2000, + ggml::hrx::DispatchSource::Common, + match_test_wrong_root_dispatch, + }); + const ggml::hrx::DispatchRegistry ordering_registry = builder.build(); + + ggml::hrx::Graph graph; + graph.add_node(GGML_OP_ADD, ggml::hrx::ValueId(0), {}); + REQUIRE(graph.build_index().success()); + + const std::vector covered_nodes(graph.nodes().size(), false); + const ggml::hrx::CommandPlan plan; + const ggml::hrx::DispatchMatchContext context = { + graph, &graph.nodes().front(), + 0, covered_nodes, + plan, ggml::hrx::ValueId(static_cast(graph.values().size())), + }; + ggml::hrx::DispatchMatch match; + REQUIRE(ordering_registry.match(context, match)); + REQUIRE(match.dispatches.size() == 1); + REQUIRE(match.dispatches.front().kernel.integer_parameters.at("route") == 2); +} + +static void run_graph_import_checks() { + ggml_init_params params = {}; + params.mem_size = 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * out = ggml_add(ctx, a, b); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(out != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 1); + const ggml::hrx::GraphNode * node = &imported.graph.nodes().front(); + REQUIRE(node->op == GGML_OP_ADD); + REQUIRE(node->inputs.size() == 2); + REQUIRE(node->inputs[0] != node->inputs[1]); + REQUIRE(ggml::hrx::DispatchScheduler::supports_node(imported.graph, node, test_dispatch_target())); + + const ggml::hrx::Value * a_value = imported.graph.values().find_tensor(a); + const ggml::hrx::Value * b_value = imported.graph.values().find_tensor(b); + const ggml::hrx::Value * out_value = imported.graph.values().find_tensor(out); + REQUIRE(a_value != nullptr); + REQUIRE(b_value != nullptr); + REQUIRE(out_value != nullptr); + REQUIRE(a_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(b_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(out_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(!a_value->buffer.has_value()); + REQUIRE(!b_value->buffer.has_value()); + REQUIRE(!out_value->buffer.has_value()); + + const std::vector external_ids = imported.graph.values().external_value_ids(); + REQUIRE(external_ids.size() == 3); + REQUIRE(contains_value_id(external_ids, a_value->id)); + REQUIRE(contains_value_id(external_ids, b_value->id)); + REQUIRE(contains_value_id(external_ids, out_value->id)); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + REQUIRE(scheduler.plan().dispatches.front().bindings.size() == 3); + REQUIRE(scheduler.plan().dispatches.front().bindings[0].length == a_value->byte_count); + + ggml::hrx::CommandProgramBindings missing_bindings = + ggml::hrx::CommandProgramBindings::from_value_map(imported.graph.values()); + REQUIRE(!missing_bindings.valid()); + + REQUIRE(imported.graph.values().bind_buffer(a_value->id, { dummy_hrx_buffer(0x1000), 0, a_value->byte_count })); + ggml::hrx::CommandProgramBindings partial_bindings = + ggml::hrx::CommandProgramBindings::from_value_map(imported.graph.values()); + REQUIRE(!partial_bindings.valid()); + + REQUIRE(imported.graph.values().bind_buffer(b_value->id, { dummy_hrx_buffer(0x1800), 0, b_value->byte_count })); + ggml::hrx::CommandProgramBindings missing_output_binding = + ggml::hrx::CommandProgramBindings::from_value_map(imported.graph.values()); + REQUIRE(!missing_output_binding.valid()); + + REQUIRE(imported.graph.values().bind_buffer(out_value->id, { dummy_hrx_buffer(0x2000), 0, 0 })); + ggml::hrx::CommandProgramBindings empty_runtime_binding = + ggml::hrx::CommandProgramBindings::from_value_map(imported.graph.values()); + REQUIRE(!empty_runtime_binding.valid()); + + REQUIRE(imported.graph.values().bind_buffer(out_value->id, { dummy_hrx_buffer(0x2000), 0, out_value->byte_count })); + ggml::hrx::CommandProgramBindings runtime_bindings = + ggml::hrx::CommandProgramBindings::from_value_map(imported.graph.values()); + REQUIRE(runtime_bindings.valid()); + REQUIRE(runtime_bindings.bindings().size() == 3); + const ggml::hrx::CommandProgramBinding * a_binding = runtime_bindings.find(a_value->id); + REQUIRE(a_binding != nullptr); + REQUIRE(a_binding->buffer == dummy_hrx_buffer(0x1000)); + REQUIRE(a_binding->offset == 0); + REQUIRE(a_binding->length == a_value->byte_count); + const ggml::hrx::CommandProgramBinding * b_binding = runtime_bindings.find(b_value->id); + REQUIRE(b_binding != nullptr); + REQUIRE(b_binding->buffer == dummy_hrx_buffer(0x1800)); + REQUIRE(b_binding->offset == 0); + REQUIRE(b_binding->length == b_value->byte_count); + REQUIRE(runtime_bindings.find(ggml::hrx::ValueId(123456)) == nullptr); + + const ggml::hrx::CommandProgramBindingsFingerprint runtime_fingerprint = + ggml::hrx::command_program_bindings_fingerprint(runtime_bindings); + REQUIRE(!runtime_fingerprint.value.empty()); + + ggml::hrx::ValueMap changed_identity = imported.graph.values(); + REQUIRE(changed_identity.bind_buffer( + a_value->id, { dummy_hrx_buffer(0x1000), 0, a_value->byte_count, 1, 0, a_value->byte_count })); + const ggml::hrx::CommandProgramBindings changed_identity_bindings = + ggml::hrx::CommandProgramBindings::from_value_map(changed_identity); + REQUIRE(changed_identity_bindings.valid()); + REQUIRE(ggml::hrx::command_program_bindings_fingerprint(changed_identity_bindings).value != + runtime_fingerprint.value); + + ggml::hrx::ValueMap changed_generation = imported.graph.values(); + REQUIRE(changed_generation.bind_buffer( + a_value->id, { dummy_hrx_buffer(0x1000), 0, a_value->byte_count, 0, 1, a_value->byte_count })); + const ggml::hrx::CommandProgramBindings changed_generation_bindings = + ggml::hrx::CommandProgramBindings::from_value_map(changed_generation); + REQUIRE(changed_generation_bindings.valid()); + REQUIRE(ggml::hrx::command_program_bindings_fingerprint(changed_generation_bindings).value != + runtime_fingerprint.value); + + ggml::hrx::ValueMap changed_capacity = imported.graph.values(); + REQUIRE(changed_capacity.bind_buffer( + a_value->id, { dummy_hrx_buffer(0x1000), 0, a_value->byte_count, 0, 0, a_value->byte_count + 256 })); + const ggml::hrx::CommandProgramBindings changed_capacity_bindings = + ggml::hrx::CommandProgramBindings::from_value_map(changed_capacity); + REQUIRE(changed_capacity_bindings.valid()); + REQUIRE(ggml::hrx::command_program_bindings_fingerprint(changed_capacity_bindings).value != + runtime_fingerprint.value); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 1); + const ggml::hrx::Command & command = commands.commands.front(); + REQUIRE(command.ordinal == 0); + REQUIRE(command.kind == ggml::hrx::CommandKind::Kernel); + REQUIRE(command.kernel.kernel_id != ggml::hrx::kUncatalogedKernelId); + REQUIRE(command.bindings.size() == 3); + REQUIRE(command.bindings[0].name == "lhs"); + REQUIRE(command.bindings[0].origin == ggml::hrx::CommandBindingOrigin::GraphValue); + REQUIRE(command.bindings[0].access == ggml::hrx::ResourceAccess::Read); + REQUIRE(command.bindings[1].name == "rhs"); + REQUIRE(command.bindings[1].origin == ggml::hrx::CommandBindingOrigin::GraphValue); + REQUIRE(command.bindings[1].access == ggml::hrx::ResourceAccess::Read); + REQUIRE(command.bindings[2].name == "output"); + REQUIRE(command.bindings[2].origin == ggml::hrx::CommandBindingOrigin::GraphValue); + REQUIRE(command.bindings[2].access == ggml::hrx::ResourceAccess::ReadWrite); + REQUIRE(command_program_verifies(commands)); + + const ggml::hrx::PreparedCommand default_prepared_command; + REQUIRE(default_prepared_command.kind == ggml::hrx::CommandKind::Invalid); + REQUIRE(ggml::hrx::command_kind_name(ggml::hrx::CommandKind::Invalid) == "Invalid"); + REQUIRE(ggml::hrx::command_kind_name(ggml::hrx::CommandKind::Kernel) == "Kernel"); + REQUIRE(ggml::hrx::command_kind_name(static_cast(255)) == "Unknown(255)"); + REQUIRE(ggml::hrx::command_binding_origin_name(ggml::hrx::CommandBindingOrigin::GraphValue) == "GraphValue"); + REQUIRE(ggml::hrx::command_binding_origin_name(ggml::hrx::CommandBindingOrigin::Transient) == "Transient"); + REQUIRE(ggml::hrx::command_binding_origin_name(ggml::hrx::CommandBindingOrigin::ProgramConstant) == + "ProgramConstant"); + REQUIRE(ggml::hrx::command_binding_origin_name(static_cast(255)) == + "Unknown(255)"); + REQUIRE(ggml::hrx::resource_access_name(ggml::hrx::ResourceAccess::Read) == "Read"); + REQUIRE(ggml::hrx::resource_access_name(ggml::hrx::ResourceAccess::Write) == "Write"); + REQUIRE(ggml::hrx::resource_access_name(ggml::hrx::ResourceAccess::ReadWrite) == "ReadWrite"); + REQUIRE(ggml::hrx::resource_access_name(static_cast(255)) == "Unknown(255)"); + + const std::string binding_text = ggml::hrx::format_command_binding(command.bindings[0]); + REQUIRE(string_contains(binding_text, "binding lhs")); + REQUIRE(string_contains(binding_text, "value=")); + REQUIRE(string_contains(binding_text, "origin=GraphValue")); + REQUIRE(string_contains(binding_text, "access=Read")); + REQUIRE(std::string(ggml::hrx::hrx_graph_replay_event_name(ggml::hrx::HrxGraphReplayEvent::Disabled)) == + "disabled"); + REQUIRE(std::string(ggml::hrx::hrx_graph_replay_event_name(ggml::hrx::HrxGraphReplayEvent::Hit)) == "hit"); + REQUIRE(string_contains(binding_text, "range=[0, ")); + REQUIRE(string_contains(binding_text, std::to_string(a_value->byte_count).c_str())); + + const std::string command_text = ggml::hrx::format_command(command); + REQUIRE(string_contains(command_text, "command 0")); + REQUIRE(string_contains(command_text, "kind=Kernel")); + REQUIRE(string_contains(command_text, "kernel_id=")); + REQUIRE(string_contains(command_text, "bindings=3")); + REQUIRE(string_contains(command_text, "deps=0")); + + const std::string program_text = ggml::hrx::format_command_program(commands); + REQUIRE(string_contains(program_text, "command_program commands=1")); + REQUIRE(string_contains(program_text, "command 0")); + REQUIRE(string_contains(program_text, "binding lhs")); + REQUIRE(string_contains(program_text, "binding rhs")); + REQUIRE(string_contains(program_text, "binding output")); + + ggml::hrx::ResolvedCommandProgram resolved = + ggml::hrx::resolve_command_program_bindings(commands, runtime_bindings); + REQUIRE(resolved.valid()); + REQUIRE(resolved.commands.size() == 1); + REQUIRE(resolved.commands.front().ordinal == command.ordinal); + REQUIRE(resolved.commands.front().kind == command.kind); + REQUIRE(resolved.commands.front().kernel.kernel_id == command.kernel.kernel_id); + REQUIRE(resolved.commands.front().bindings.size() == 3); + REQUIRE(resolved.commands.front().bindings[0].binding.name == "lhs"); + REQUIRE(resolved.commands.front().bindings[0].ref.buffer == dummy_hrx_buffer(0x1000)); + REQUIRE(resolved.commands.front().bindings[0].ref.offset == 0); + REQUIRE(resolved.commands.front().bindings[0].ref.length == a_value->byte_count); + REQUIRE(resolved.commands.front().bindings[1].binding.name == "rhs"); + REQUIRE(resolved.commands.front().bindings[1].ref.buffer == dummy_hrx_buffer(0x1800)); + REQUIRE(resolved.commands.front().bindings[1].ref.offset == 0); + REQUIRE(resolved.commands.front().bindings[1].ref.length == b_value->byte_count); + REQUIRE(resolved.commands.front().bindings[2].binding.name == "output"); + REQUIRE(resolved.commands.front().bindings[2].ref.buffer == dummy_hrx_buffer(0x2000)); + REQUIRE(resolved.commands.front().bindings[2].ref.offset == 0); + REQUIRE(resolved.commands.front().bindings[2].ref.length == out_value->byte_count); + + ggml::hrx::CommandProgram offset_command = copy_command_program_shape(commands); + offset_command.commands.front().bindings[0].offset = 4; + offset_command.commands.front().bindings[0].length = 8; + ggml::hrx::ValueMap offset_values = imported.graph.values(); + REQUIRE(offset_values.bind_buffer(a_value->id, { dummy_hrx_buffer(0x3000), 16, a_value->byte_count })); + REQUIRE(offset_values.bind_buffer(out_value->id, { dummy_hrx_buffer(0x4000), 32, out_value->byte_count })); + const ggml::hrx::CommandProgramBindings offset_bindings = + ggml::hrx::CommandProgramBindings::from_value_map(offset_values); + REQUIRE(offset_bindings.valid()); + resolved = ggml::hrx::resolve_command_program_bindings(offset_command, offset_bindings); + REQUIRE(resolved.valid()); + REQUIRE(resolved.commands.front().bindings[0].ref.buffer == dummy_hrx_buffer(0x3000)); + REQUIRE(resolved.commands.front().bindings[0].ref.offset == 20); + REQUIRE(resolved.commands.front().bindings[0].ref.length == 8); + + std::vector host_storage(a_value->byte_count + 64); + ggml::hrx::ValueBufferBinding host_value_binding; + host_value_binding.host_data = host_storage.data(); + host_value_binding.offset = 16; + host_value_binding.length = a_value->byte_count; + host_value_binding.identity = 42; + host_value_binding.generation = 1; + host_value_binding.capacity = host_storage.size(); + host_value_binding.weight = true; + ggml::hrx::ValueMap host_values = imported.graph.values(); + REQUIRE(host_values.bind_buffer(a_value->id, host_value_binding)); + REQUIRE(host_values.bind_buffer(out_value->id, { dummy_hrx_buffer(0x5000), 0, out_value->byte_count })); + const ggml::hrx::CommandProgramBindings host_bindings = + ggml::hrx::CommandProgramBindings::from_value_map(host_values); + REQUIRE(host_bindings.valid()); + const ggml::hrx::CommandProgramBinding * host_binding = host_bindings.find(a_value->id); + REQUIRE(host_binding != nullptr); + REQUIRE(host_binding->buffer == nullptr); + REQUIRE(host_binding->host_data == host_storage.data()); + REQUIRE(host_binding->offset == 16); + REQUIRE(host_binding->length == a_value->byte_count); + REQUIRE(host_binding->weight); + REQUIRE(ggml::hrx::command_program_bindings_fingerprint(host_bindings).value != runtime_fingerprint.value); + + resolved = ggml::hrx::resolve_command_program_bindings(commands, missing_bindings); + REQUIRE(!resolved.valid()); + REQUIRE(status_contains(resolved.status, "is not bound")); + REQUIRE(status_contains(resolved.status, "binding output")); + REQUIRE(status_contains(resolved.status, "value=")); + + resolved = ggml::hrx::resolve_command_program_bindings(commands, partial_bindings); + REQUIRE(!resolved.valid()); + REQUIRE(status_contains(resolved.status, "is not bound")); + + resolved = ggml::hrx::resolve_command_program_bindings(commands, empty_runtime_binding); + REQUIRE(!resolved.valid()); + REQUIRE(status_contains(resolved.status, "empty binding")); + + ggml::hrx::ValueMap null_values = imported.graph.values(); + REQUIRE(null_values.bind_buffer(a_value->id, { nullptr, 0, a_value->byte_count })); + REQUIRE(null_values.bind_buffer(out_value->id, { dummy_hrx_buffer(0x2000), 0, out_value->byte_count })); + const ggml::hrx::CommandProgramBindings null_bindings = + ggml::hrx::CommandProgramBindings::from_value_map(null_values); + REQUIRE(!null_bindings.valid()); + resolved = ggml::hrx::resolve_command_program_bindings(commands, null_bindings); + REQUIRE(!resolved.valid()); + REQUIRE(status_contains(resolved.status, "null buffer")); + REQUIRE(status_contains(resolved.status, "binding lhs")); + + ggml::hrx::CommandProgram empty_resolve_binding = copy_command_program_shape(commands); + empty_resolve_binding.commands.front().bindings[0].length = 0; + resolved = ggml::hrx::resolve_command_program_bindings(empty_resolve_binding, runtime_bindings); + REQUIRE(!resolved.valid()); + REQUIRE(status_contains(resolved.status, "empty range")); + REQUIRE(status_contains(resolved.status, "range=[0, 0)")); + + ggml::hrx::CommandProgram out_of_range_binding = copy_command_program_shape(commands); + out_of_range_binding.commands.front().bindings[0].offset = a_value->byte_count; + out_of_range_binding.commands.front().bindings[0].length = 4; + resolved = ggml::hrx::resolve_command_program_bindings(out_of_range_binding, runtime_bindings); + REQUIRE(!resolved.valid()); + REQUIRE(status_contains(resolved.status, "outside runtime binding length")); + REQUIRE(status_contains(resolved.status, "binding lhs")); + + ggml::hrx::CommandProgram unsupported_origin = copy_command_program_shape(commands); + unsupported_origin.commands.front().bindings[0].origin = static_cast(255); + resolved = ggml::hrx::resolve_command_program_bindings(unsupported_origin, runtime_bindings); + REQUIRE(!resolved.valid()); + REQUIRE(status_contains(resolved.status, "unsupported binding origin")); + REQUIRE(status_contains(resolved.status, "origin=Unknown(255)")); + ggml::hrx::VerificationResult verification = + ggml::hrx::verify_command_program(unsupported_origin, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(!verification.valid()); + REQUIRE(status_contains(verification.status, "unsupported binding origin")); + REQUIRE(status_contains(verification.status, "origin=Unknown(255)")); + + ggml::hrx::CommandProgram invalid_kernel = copy_command_program_shape(commands); + invalid_kernel.commands.front().kernel.kernel_id = ggml::hrx::kUncatalogedKernelId; + verification = ggml::hrx::verify_command_program(invalid_kernel, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(!verification.valid()); + REQUIRE(status_contains(verification.status, "command 0")); + REQUIRE(status_contains(verification.status, "kernel_id=")); + + ggml::hrx::CommandProgram empty_bindings = copy_command_program_shape(commands); + empty_bindings.commands.front().bindings.clear(); + verification = ggml::hrx::verify_command_program(empty_bindings, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(!verification.valid()); + REQUIRE(status_contains(verification.status, "bindings=0")); + + ggml::hrx::CommandProgram empty_binding = copy_command_program_shape(commands); + empty_binding.commands.front().bindings[0].length = 0; + verification = ggml::hrx::verify_command_program(empty_binding, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(!verification.valid()); + REQUIRE(status_contains(verification.status, "binding lhs")); + REQUIRE(status_contains(verification.status, "range=[0, 0)")); + + ggml::hrx::CommandProgram invalid_value = copy_command_program_shape(commands); + invalid_value.commands.front().bindings[0].value = ggml::hrx::ValueId(-1); + verification = ggml::hrx::verify_command_program(invalid_value, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(!verification.valid()); + REQUIRE(status_contains(verification.status, "value=-1")); + + ggml::hrx::CommandProgram forward_dependency = copy_command_program_shape(commands); + forward_dependency.commands.front().dependencies.push_back(0); + verification = + ggml::hrx::verify_command_program(forward_dependency, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(!verification.valid()); + REQUIRE(status_contains(verification.status, "forward dependency 0")); + + ggml::hrx::CommandProgram wrong_binding_name = copy_command_program_shape(commands); + wrong_binding_name.commands.front().bindings[0].name = "wrong"; + verification = + ggml::hrx::verify_command_program(wrong_binding_name, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(!verification.valid()); + REQUIRE(status_contains(verification.status, "binding wrong")); + + ggml::hrx::CommandProgram wrong_binding_access = copy_command_program_shape(commands); + wrong_binding_access.commands.front().bindings[0].access = ggml::hrx::ResourceAccess::Write; + verification = + ggml::hrx::verify_command_program(wrong_binding_access, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(!verification.valid()); + REQUIRE(status_contains(verification.status, "access=Write")); + + ggml::hrx::CommandProgram transient_flow = copy_command_program_shape(commands); + transient_flow.commands.resize(2); + transient_flow.commands[0].ordinal = 0; + transient_flow.commands[0].dependencies.clear(); + transient_flow.commands[1] = transient_flow.commands[0]; + transient_flow.commands[1].ordinal = 1; + transient_flow.commands[1].dependencies.push_back(0); + const ggml::hrx::ValueId transient_value(100000); + transient_flow.transients.allocations.push_back({ + transient_value, + command.bindings[2].length, + 256, + 0, + }); + transient_flow.transients.arena_size = command.bindings[2].length; + transient_flow.commands[0].bindings[2].origin = ggml::hrx::CommandBindingOrigin::Transient; + transient_flow.commands[0].bindings[2].value = transient_value; + transient_flow.commands[1].bindings[0].origin = ggml::hrx::CommandBindingOrigin::Transient; + transient_flow.commands[1].bindings[0].value = transient_value; + transient_flow.commands[1].bindings[0].length = command.bindings[2].length; + verification = ggml::hrx::verify_command_program(transient_flow, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(verification.valid()); + + ggml::hrx::CommandProgram read_before_write = copy_command_program_shape(transient_flow); + std::swap(read_before_write.commands[0], read_before_write.commands[1]); + read_before_write.commands[0].ordinal = 0; + read_before_write.commands[0].dependencies.clear(); + read_before_write.commands[1].ordinal = 1; + read_before_write.commands[1].dependencies.clear(); + read_before_write.commands[1].dependencies.push_back(0); + verification = ggml::hrx::verify_command_program(read_before_write, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(!verification.valid()); + REQUIRE(status_contains(verification.status, "reads transient value 100000 before write by command 1")); + + const ggml::hrx::KernelCorpus & corpus = ggml::hrx::get_qwen_kernel_corpus(); + const ggml::hrx::CommandProgramExecutionContext prepare_context = { + nullptr, nullptr, "gfx1151", &corpus, nullptr, nullptr, nullptr, nullptr, + }; + + ggml::hrx::PreparedCommandProgram prepared = + ggml::hrx::prepare_command_program(prepare_context, invalid_kernel, runtime_bindings); + REQUIRE(!prepared.valid()); + REQUIRE(status_contains(prepared.status, "kernel_id=")); + + prepared = ggml::hrx::prepare_command_program(prepare_context, commands, missing_bindings); + REQUIRE(!prepared.valid()); + REQUIRE(status_contains(prepared.status, "is not bound")); + + prepared = ggml::hrx::prepare_command_program(prepare_context, commands, runtime_bindings); + REQUIRE(!prepared.valid()); + REQUIRE(status_contains(prepared.status, "missing HRX device")); + + const ggml::hrx::CommandProgramExecutionContext missing_kernel_cache_context = { + reinterpret_cast(uintptr_t(1)), nullptr, "gfx1151", &corpus, nullptr, nullptr, nullptr, nullptr, + }; + prepared = ggml::hrx::prepare_command_program(missing_kernel_cache_context, commands, runtime_bindings); + REQUIRE(!prepared.valid()); + REQUIRE(status_contains(prepared.status, "missing HRX kernel executable cache")); + + ggml_free(ctx); +} + +static void run_graph_import_mixed_backend_boundary_checks() { + ggml_init_params params = {}; + params.mem_size = 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 8, 4); + ggml_tensor * cpu_sin = ggml_sin(ctx, input); + ggml_tensor * hrx_norm = ggml_rms_norm(ctx, cpu_sin, 1.0e-6f); + REQUIRE(input != nullptr); + REQUIRE(cpu_sin != nullptr); + REQUIRE(hrx_norm != nullptr); + + ggml_cgraph * cpu_to_hrx = ggml_new_graph(ctx); + REQUIRE(cpu_to_hrx != nullptr); + ggml_graph_add_node(cpu_to_hrx, hrx_norm); + + ggml::hrx::GraphImportResult imported_cpu_to_hrx = ggml::hrx::import_ggml_graph(*cpu_to_hrx); + REQUIRE(imported_cpu_to_hrx.valid()); + REQUIRE(imported_cpu_to_hrx.graph.nodes().size() == 1); + REQUIRE(imported_cpu_to_hrx.graph.nodes()[0].op == GGML_OP_RMS_NORM); + + const ggml::hrx::Value * cpu_sin_value = imported_cpu_to_hrx.graph.values().find_tensor(cpu_sin); + const ggml::hrx::Value * hrx_norm_value = imported_cpu_to_hrx.graph.values().find_tensor(hrx_norm); + REQUIRE(cpu_sin_value != nullptr); + REQUIRE(hrx_norm_value != nullptr); + REQUIRE(cpu_sin_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(hrx_norm_value->kind == ggml::hrx::ValueKind::External); + + ggml_tensor * hrx_norm_for_cpu = ggml_rms_norm(ctx, input, 1.0e-6f); + ggml_tensor * cpu_sin_after = ggml_sin(ctx, hrx_norm_for_cpu); + REQUIRE(hrx_norm_for_cpu != nullptr); + REQUIRE(cpu_sin_after != nullptr); + + ggml_cgraph * hrx_to_cpu = ggml_new_graph(ctx); + REQUIRE(hrx_to_cpu != nullptr); + ggml_graph_add_node(hrx_to_cpu, hrx_norm_for_cpu); + + ggml::hrx::GraphImportResult imported_hrx_to_cpu = ggml::hrx::import_ggml_graph(*hrx_to_cpu); + REQUIRE(imported_hrx_to_cpu.valid()); + REQUIRE(imported_hrx_to_cpu.graph.nodes().size() == 1); + REQUIRE(imported_hrx_to_cpu.graph.nodes()[0].op == GGML_OP_RMS_NORM); + + const ggml::hrx::Value * input_value = imported_hrx_to_cpu.graph.values().find_tensor(input); + const ggml::hrx::Value * output_value = imported_hrx_to_cpu.graph.values().find_tensor(hrx_norm_for_cpu); + REQUIRE(input_value != nullptr); + REQUIRE(output_value != nullptr); + REQUIRE(input_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(output_value->kind == ggml::hrx::ValueKind::External); + + ggml_tensor * cpu_sin_for_view = ggml_sin(ctx, input); + ggml_tensor * cpu_sin_view = ggml_reshape_2d(ctx, cpu_sin_for_view, 8, 4); + ggml_tensor * hrx_norm_from_view = ggml_rms_norm(ctx, cpu_sin_view, 1.0e-6f); + REQUIRE(cpu_sin_for_view != nullptr); + REQUIRE(cpu_sin_view != nullptr); + REQUIRE(hrx_norm_from_view != nullptr); + + ggml_cgraph * cpu_view_to_hrx = ggml_new_graph(ctx); + REQUIRE(cpu_view_to_hrx != nullptr); + ggml_graph_add_node(cpu_view_to_hrx, cpu_sin_view); + ggml_graph_add_node(cpu_view_to_hrx, hrx_norm_from_view); + + ggml::hrx::GraphImportResult imported_cpu_view_to_hrx = ggml::hrx::import_ggml_graph(*cpu_view_to_hrx); + REQUIRE(imported_cpu_view_to_hrx.valid()); + REQUIRE(imported_cpu_view_to_hrx.graph.nodes().size() == 2); + + const ggml::hrx::Value * cpu_sin_for_view_value = + imported_cpu_view_to_hrx.graph.values().find_tensor(cpu_sin_for_view); + const ggml::hrx::Value * cpu_sin_view_value = imported_cpu_view_to_hrx.graph.values().find_tensor(cpu_sin_view); + REQUIRE(cpu_sin_for_view_value != nullptr); + REQUIRE(cpu_sin_view_value != nullptr); + REQUIRE(cpu_sin_for_view_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(cpu_sin_view_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(imported_cpu_view_to_hrx.graph.values().same_storage(cpu_sin_for_view_value->id, cpu_sin_view_value->id)); + + ggml_tensor * hrx_norm_for_view = ggml_rms_norm(ctx, input, 1.0e-6f); + ggml_tensor * hrx_norm_view = ggml_reshape_2d(ctx, hrx_norm_for_view, 8, 4); + ggml_tensor * cpu_sin_view_out = ggml_sin(ctx, hrx_norm_view); + REQUIRE(hrx_norm_for_view != nullptr); + REQUIRE(hrx_norm_view != nullptr); + REQUIRE(cpu_sin_view_out != nullptr); + + ggml_cgraph * hrx_view_to_cpu = ggml_new_graph(ctx); + REQUIRE(hrx_view_to_cpu != nullptr); + ggml_build_forward_expand(hrx_view_to_cpu, hrx_norm_view); + + ggml::hrx::GraphImportResult imported_hrx_view_to_cpu = ggml::hrx::import_ggml_graph(*hrx_view_to_cpu); + REQUIRE(imported_hrx_view_to_cpu.valid()); + + const ggml::hrx::Value * hrx_norm_for_view_value = + imported_hrx_view_to_cpu.graph.values().find_tensor(hrx_norm_for_view); + const ggml::hrx::Value * hrx_norm_view_value = imported_hrx_view_to_cpu.graph.values().find_tensor(hrx_norm_view); + REQUIRE(hrx_norm_for_view_value != nullptr); + REQUIRE(hrx_norm_view_value != nullptr); + REQUIRE(hrx_norm_for_view_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(hrx_norm_view_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(imported_hrx_view_to_cpu.graph.values().same_storage(hrx_norm_for_view_value->id, hrx_norm_view_value->id)); + + ggml_free(ctx); +} + +static void run_graph_view_external_use_checks() { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * c = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * shared = ggml_add(ctx, a, b); + ggml_tensor * left = ggml_mul(ctx, shared, c); + ggml_tensor * right = ggml_sqr(ctx, shared); + ggml_tensor * output = ggml_add(ctx, left, right); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(c != nullptr); + REQUIRE(shared != nullptr); + REQUIRE(left != nullptr); + REQUIRE(right != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + REQUIRE(graph->n_nodes == 4); + REQUIRE(graph->nodes[0] == shared); + REQUIRE(graph->nodes[1] == left); + + ggml_cgraph graph_view = ggml_graph_view(graph, 0, 2); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(graph_view); + REQUIRE(imported.valid()); + + const ggml::hrx::Value * shared_value = imported.graph.values().find_tensor(shared); + const ggml::hrx::Value * left_value = imported.graph.values().find_tensor(left); + REQUIRE(shared_value != nullptr); + REQUIRE(left_value != nullptr); + REQUIRE(shared_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(left_value->kind == ggml::hrx::ValueKind::External); + + ggml_free(ctx); +} + +static ggml::hrx::Dispatch schedule_single_dispatch_for_tensor(ggml_context * ctx, ggml_tensor * output) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + return scheduler.plan().dispatches.front(); +} + +static bool can_schedule_tensor(ggml_context * ctx, ggml_tensor * output) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + return ggml::hrx::DispatchScheduler::can_schedule_graph(imported.graph, test_dispatch_target()); +} + +static void require_compile_param(const ggml::hrx::Dispatch & dispatch, const char * name, const char * value) { + REQUIRE(dispatch.kernel.compile_parameters.at(name) == value); +} + +static void run_scale_f32_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 3840, 18); + ggml_tensor * output = ggml_scale(ctx, input, 1.75f); + REQUIRE(input != nullptr); + REQUIRE(output != nullptr); + + const ggml::hrx::Dispatch dispatch = schedule_single_dispatch_for_tensor(ctx, output); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_scale_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("element_count") == 69120); + require_compile_parameter(dispatch, "ggml.scale_f32.scale", expected_config_value(1.75f)); + require_compile_parameter(dispatch, "ggml.scale_f32.bias", expected_config_value(0.0f)); + + ggml_tensor * biased = ggml_scale_bias(ctx, input, -0.5f, 0.25f); + REQUIRE(biased != nullptr); + + const ggml::hrx::Dispatch bias_dispatch = schedule_single_dispatch_for_tensor(ctx, biased); + REQUIRE(kernel_name_for_id(bias_dispatch.kernel.kernel_id) == "loom_libs:ggml_scale_f32"); + REQUIRE(bias_dispatch.kernel.integer_parameters.at("element_count") == 69120); + require_compile_parameter(bias_dispatch, "ggml.scale_f32.scale", expected_config_value(-0.5f)); + require_compile_parameter(bias_dispatch, "ggml.scale_f32.bias", expected_config_value(0.25f)); + + ggml_tensor * inplace = ggml_scale_bias_inplace(ctx, input, 2.0f, -1.0f); + REQUIRE(inplace != nullptr); + + const ggml::hrx::Dispatch inplace_dispatch = schedule_single_dispatch_for_tensor(ctx, inplace); + REQUIRE(kernel_name_for_id(inplace_dispatch.kernel.kernel_id) == "loom_libs:ggml_scale_bias_f32"); + REQUIRE(inplace_dispatch.kernel.integer_parameters.at("element_count") == 69120); + require_compile_parameter(inplace_dispatch, "ggml.scale.scale", expected_config_value(2.0f)); + require_compile_parameter(inplace_dispatch, "ggml.scale.bias", expected_config_value(-1.0f)); + + ggml_free(ctx); +} + +static void require_scale_add_dispatch(ggml_context * ctx, + ggml_tensor * input, + ggml_tensor * residual, + int64_t expected_element_count, + int64_t expected_input_stride1, + int64_t expected_residual_stride1, + ggml::hrx::BinaryKind binary_kind = ggml::hrx::BinaryKind::Add, + bool scaled_lhs = true) { + ggml_tensor * scaled = ggml_scale_bias(ctx, input, 0.177800179f, 0.25f); + ggml_tensor * output = nullptr; + switch (binary_kind) { + case ggml::hrx::BinaryKind::Add: + output = ggml_add(ctx, scaled, residual); + break; + case ggml::hrx::BinaryKind::Sub: + output = scaled_lhs ? ggml_sub(ctx, scaled, residual) : ggml_sub(ctx, residual, scaled); + break; + case ggml::hrx::BinaryKind::Mul: + output = ggml_mul(ctx, scaled, residual); + break; + case ggml::hrx::BinaryKind::Div: + output = scaled_lhs ? ggml_div(ctx, scaled, residual) : ggml_div(ctx, residual, scaled); + break; + default: + REQUIRE(false); + } + REQUIRE(scaled != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + require_scheduled_command_program(graph, [&](const ggml::hrx::Graph &, const ggml::hrx::CommandPlan & plan, + const ggml::hrx::CommandProgram & commands) { + REQUIRE(plan.dispatches.size() == 1); + const ggml::hrx::Dispatch & dispatch = plan.dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_scale_add_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("element_count") == expected_element_count); + require_compile_parameter(dispatch, "ggml.scale_add_f32.scale", expected_config_value(0.177800179f)); + require_compile_parameter(dispatch, "ggml.scale_add_f32.bias", expected_config_value(0.25f)); + require_compile_parameter(dispatch, "ggml.scale_add_f32.binary_op", + std::to_string(ggml::hrx::binary_kind_config_value(binary_kind))); + require_compile_parameter(dispatch, "ggml.scale_add_f32.scaled_lhs", scaled_lhs ? "1" : "0"); + require_compile_parameter(dispatch, "ggml.scale_add_f32.input_stride1", std::to_string(expected_input_stride1)); + require_compile_parameter(dispatch, "ggml.scale_add_f32.residual_stride1", + std::to_string(expected_residual_stride1)); + REQUIRE(dispatch.bindings.size() == 3); + REQUIRE(commands.commands.size() == 1); + REQUIRE(commands.commands.front().bindings[0].name == "input"); + REQUIRE(commands.commands.front().bindings[1].name == "residual"); + REQUIRE(commands.commands.front().bindings[2].name == "output"); + }); +} + +static void run_scale_add_f32_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 2 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * packed_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2560, 64); + ggml_tensor * packed_residual = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2560, 64); + REQUIRE(packed_input != nullptr); + REQUIRE(packed_residual != nullptr); + require_scale_add_dispatch(ctx, packed_input, packed_residual, 163840, 2560, 2560); + require_scale_add_dispatch(ctx, packed_input, packed_residual, 163840, 2560, 2560, ggml::hrx::BinaryKind::Sub, + false); + require_scale_add_dispatch(ctx, packed_input, packed_residual, 163840, 2560, 2560, ggml::hrx::BinaryKind::Mul, + true); + require_scale_add_dispatch(ctx, packed_input, packed_residual, 163840, 2560, 2560, ggml::hrx::BinaryKind::Div, + true); + + ggml_tensor * input_storage = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 320, 3); + ggml_tensor * residual_storage = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 384, 3); + ggml_tensor * strided_input = ggml_view_2d(ctx, input_storage, 257, 3, input_storage->nb[1], 0); + ggml_tensor * strided_residual = ggml_view_2d(ctx, residual_storage, 257, 3, residual_storage->nb[1], 0); + REQUIRE(strided_input != nullptr); + REQUIRE(strided_residual != nullptr); + require_scale_add_dispatch(ctx, strided_input, strided_residual, 771, 320, 384); + + ggml_tensor * scaled = ggml_scale(ctx, packed_input, 0.5f); + ggml_tensor * alias_output = ggml_add_inplace(ctx, scaled, packed_residual); + REQUIRE(scaled != nullptr); + REQUIRE(alias_output != nullptr); + REQUIRE(!can_schedule_tensor(ctx, alias_output)); + + ggml_tensor * alias_scaled = ggml_scale_inplace(ctx, packed_input, 0.5f); + ggml_tensor * output_from_alias = ggml_add(ctx, alias_scaled, packed_residual); + REQUIRE(alias_scaled != nullptr); + REQUIRE(output_from_alias != nullptr); + ggml_cgraph * alias_graph = ggml_new_graph(ctx); + REQUIRE(alias_graph != nullptr); + ggml_build_forward_expand(alias_graph, output_from_alias); + ggml::hrx::GraphImportResult alias_imported = ggml::hrx::import_ggml_graph(*alias_graph); + REQUIRE(alias_imported.valid()); + ggml::hrx::DispatchScheduler alias_scheduler; + REQUIRE(alias_scheduler.schedule_graph(alias_imported.graph, test_dispatch_target())); + REQUIRE(alias_scheduler.plan().dispatches.size() == 2); + REQUIRE(kernel_name_for_id(alias_scheduler.plan().dispatches[0].kernel.kernel_id) == + "loom_libs:ggml_scale_bias_f32"); + REQUIRE(kernel_name_for_id(alias_scheduler.plan().dispatches[1].kernel.kernel_id) == "loom_libs:ggml_binary_f32"); + + ggml_free(ctx); +} + +static ggml_tensor * make_test_binary(ggml_context * ctx, + ggml_tensor * lhs, + ggml_tensor * rhs, + ggml::hrx::BinaryKind kind) { + switch (kind) { + case ggml::hrx::BinaryKind::Add: + return ggml_add(ctx, lhs, rhs); + case ggml::hrx::BinaryKind::Sub: + return ggml_sub(ctx, lhs, rhs); + case ggml::hrx::BinaryKind::Mul: + return ggml_mul(ctx, lhs, rhs); + case ggml::hrx::BinaryKind::Div: + return ggml_div(ctx, lhs, rhs); + default: + REQUIRE(false); + return nullptr; + } +} + +static void run_rmsnorm_two_binary_dispatch_checks() { + const struct { + ggml::hrx::BinaryKind normalized_op; + ggml::hrx::BinaryKind output_op; + bool binary_lhs; + } cases[] = { + { ggml::hrx::BinaryKind::Add, ggml::hrx::BinaryKind::Mul, true }, + { ggml::hrx::BinaryKind::Sub, ggml::hrx::BinaryKind::Sub, true }, + { ggml::hrx::BinaryKind::Mul, ggml::hrx::BinaryKind::Div, true }, + { ggml::hrx::BinaryKind::Div, ggml::hrx::BinaryKind::Sub, false }, + }; + + for (const auto & test : cases) { + ggml_init_params params = {}; + params.mem_size = 512 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 128, 4); + ggml_tensor * weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 128); + ggml_tensor * residual = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 128, 4); + ggml_tensor * rms = ggml_rms_norm(ctx, input, 1.0e-6f); + ggml_tensor * scaled = make_test_binary(ctx, rms, weight, test.normalized_op); + ggml_tensor * output = test.binary_lhs ? make_test_binary(ctx, scaled, residual, test.output_op) : + make_test_binary(ctx, residual, scaled, test.output_op); + REQUIRE(output != nullptr); + + const ggml::hrx::Dispatch dispatch = schedule_single_dispatch_for_tensor(ctx, output); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rmsnorm_mul_add_f32"); + REQUIRE(dispatch.bindings.size() == 4); + require_compile_parameter(dispatch, "ggml.rmsnorm_mul_add_f32.normalized_op", + std::to_string(ggml::hrx::binary_kind_config_value(test.normalized_op))); + require_compile_parameter(dispatch, "ggml.rmsnorm_mul_add_f32.output_op", + std::to_string(ggml::hrx::binary_kind_config_value(test.output_op))); + require_compile_parameter(dispatch, "ggml.rmsnorm_mul_add_f32.binary_lhs", test.binary_lhs ? "1" : "0"); + ggml_free(ctx); + } +} + +static void run_cont_f32_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 512 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t width = 64; + constexpr int64_t row_count = 40; + constexpr int64_t token_count = 23; + constexpr int64_t token_stride = 4096; + constexpr size_t view_offset = 8 * sizeof(float); + const size_t storage_elements = + view_offset / sizeof(float) + static_cast((token_count - 1) * token_stride + width * row_count); + + ggml_tensor * storage = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, storage_elements); + ggml_tensor * view = ggml_view_3d(ctx, storage, width, row_count, token_count, width * sizeof(float), + token_stride * sizeof(float), view_offset); + ggml_tensor * output = ggml_cont(ctx, view); + REQUIRE(storage != nullptr); + REQUIRE(view != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 2); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_VIEW); + REQUIRE(imported.graph.nodes()[1].op == GGML_OP_CONT); + + const ggml::hrx::Value * view_value = imported.graph.values().find(imported.graph.nodes()[1].inputs[0]); + REQUIRE(view_value != nullptr); + const ggml::hrx::Value * storage_value = imported.graph.values().find(view_value->storage_root); + REQUIRE(storage_value != nullptr); + REQUIRE(view_value->storage_offset == view_offset); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches[0]; + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_copy_strided_source_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("element_count") == width * row_count * token_count); + REQUIRE(dispatch.kernel.integer_parameters.at("source_span") == + static_cast((token_count - 1) * token_stride + width * row_count)); + require_compile_parameter(dispatch, "ggml.copy_strided_source_f32.ne0", std::to_string(width)); + require_compile_parameter(dispatch, "ggml.copy_strided_source_f32.ne1", std::to_string(row_count)); + require_compile_parameter(dispatch, "ggml.copy_strided_source_f32.ne2", std::to_string(token_count)); + require_compile_parameter(dispatch, "ggml.copy_strided_source_f32.stride0", "1"); + require_compile_parameter(dispatch, "ggml.copy_strided_source_f32.stride1", std::to_string(width)); + require_compile_parameter(dispatch, "ggml.copy_strided_source_f32.stride2", std::to_string(token_stride)); + REQUIRE(dispatch.bindings.size() == 2); + REQUIRE(dispatch.bindings[0].value == storage_value->id); + REQUIRE(dispatch.bindings[0].offset == view_offset); + REQUIRE(dispatch.bindings[0].length == + static_cast((token_count - 1) * token_stride + width * row_count) * sizeof(float)); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + + ggml_free(ctx); +} + +static void run_binary_f32_broadcast_dispatch_checks() { + { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * lhs = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 16, 3); + ggml_tensor * rhs = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 16, 3); + ggml_tensor * output = ggml_add(ctx, lhs, rhs); + REQUIRE(lhs != nullptr); + REQUIRE(rhs != nullptr); + REQUIRE(output != nullptr); + + const ggml::hrx::Dispatch dispatch = schedule_single_dispatch_for_tensor(ctx, output); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_binary_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("element_count") == 48); + require_compile_param(dispatch, "ggml.binary_f32.ne0", "16"); + require_compile_param(dispatch, "ggml.binary_f32.src0_stride1", "16"); + require_compile_param(dispatch, "ggml.binary_f32.src1_stride1", "16"); + REQUIRE(dispatch.kernel.integer_parameters.count("src0_element_count") == 0); + require_compile_param(dispatch, "ggml.binary_f32.op", "0"); + REQUIRE(dispatch.kernel.compile_parameters.count("ggml.binary_bc_f32.src0_broadcast_dim0") == 0); + + ggml_free(ctx); + } + + { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * source = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 2048, 3, 2); + REQUIRE(source != nullptr); + ggml_tensor * lhs = ggml_view_2d(ctx, source, 2048, 2, source->nb[2], 0); + ggml_tensor * rhs = ggml_view_2d(ctx, source, 2048, 2, source->nb[2], 2 * source->nb[1]); + REQUIRE(lhs != nullptr); + REQUIRE(rhs != nullptr); + ggml_tensor * output = ggml_mul(ctx, lhs, rhs); + REQUIRE(output != nullptr); + + const ggml::hrx::Dispatch dispatch = schedule_single_dispatch_for_tensor(ctx, output); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_binary_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("element_count") == 4096); + require_compile_param(dispatch, "ggml.binary_f32.src0_stride1", "6144"); + require_compile_param(dispatch, "ggml.binary_f32.src1_stride1", "6144"); + require_compile_param(dispatch, "ggml.binary_f32.src0_span", "8192"); + require_compile_param(dispatch, "ggml.binary_f32.src1_span", "8192"); + require_compile_param(dispatch, "ggml.binary_f32.op", "2"); + + ggml_free(ctx); + } + + { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 16, 3); + ggml_tensor * weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 16); + ggml_tensor * output = ggml_mul(ctx, input, weight); + REQUIRE(input != nullptr); + REQUIRE(weight != nullptr); + REQUIRE(output != nullptr); + + const ggml::hrx::Dispatch dispatch = schedule_single_dispatch_for_tensor(ctx, output); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_binary_bc_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("element_count") == 48); + REQUIRE(dispatch.kernel.integer_parameters.at("ne0") == 16); + REQUIRE(dispatch.kernel.integer_parameters.at("ne1") == 3); + REQUIRE(dispatch.kernel.integer_parameters.at("src0_element_count") == 48); + REQUIRE(dispatch.kernel.integer_parameters.at("src1_element_count") == 16); + require_compile_param(dispatch, "ggml.binary_bc_f32.op", "2"); + require_compile_param(dispatch, "ggml.binary_bc_f32.src0_broadcast_dim0", "0"); + require_compile_param(dispatch, "ggml.binary_bc_f32.src0_broadcast_dim1", "0"); + require_compile_param(dispatch, "ggml.binary_bc_f32.src1_broadcast_dim0", "0"); + require_compile_param(dispatch, "ggml.binary_bc_f32.src1_broadcast_dim1", "1"); + + ggml_free(ctx); + } + + { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * numerator = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 8, 4); + ggml_tensor * denominator = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 1, 4); + ggml_tensor * output = ggml_div(ctx, numerator, denominator); + REQUIRE(numerator != nullptr); + REQUIRE(denominator != nullptr); + REQUIRE(output != nullptr); + + const ggml::hrx::Dispatch dispatch = schedule_single_dispatch_for_tensor(ctx, output); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_binary_bc_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("element_count") == 32); + REQUIRE(dispatch.kernel.integer_parameters.at("ne0") == 8); + REQUIRE(dispatch.kernel.integer_parameters.at("ne1") == 4); + REQUIRE(dispatch.kernel.integer_parameters.at("src0_element_count") == 32); + REQUIRE(dispatch.kernel.integer_parameters.at("src1_element_count") == 4); + require_compile_param(dispatch, "ggml.binary_bc_f32.op", "3"); + require_compile_param(dispatch, "ggml.binary_bc_f32.src0_broadcast_dim0", "0"); + require_compile_param(dispatch, "ggml.binary_bc_f32.src1_broadcast_dim0", "1"); + require_compile_param(dispatch, "ggml.binary_bc_f32.src1_broadcast_dim1", "0"); + + ggml_free(ctx); + } + + { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 5, 6, 7); + ggml_tensor * weight = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 1, 6, 7); + ggml_tensor * output = ggml_mul(ctx, input, weight); + REQUIRE(input != nullptr); + REQUIRE(weight != nullptr); + REQUIRE(output != nullptr); + + const ggml::hrx::Dispatch dispatch = schedule_single_dispatch_for_tensor(ctx, output); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_binary_bc_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("element_count") == 210); + REQUIRE(dispatch.kernel.integer_parameters.at("ne0") == 5); + REQUIRE(dispatch.kernel.integer_parameters.at("ne1") == 6); + REQUIRE(dispatch.kernel.integer_parameters.at("ne2") == 7); + REQUIRE(dispatch.kernel.integer_parameters.at("src0_element_count") == 210); + REQUIRE(dispatch.kernel.integer_parameters.at("src1_element_count") == 42); + require_compile_param(dispatch, "ggml.binary_bc_f32.op", "2"); + require_compile_param(dispatch, "ggml.binary_bc_f32.src1_broadcast_dim0", "1"); + require_compile_param(dispatch, "ggml.binary_bc_f32.src1_broadcast_dim1", "0"); + require_compile_param(dispatch, "ggml.binary_bc_f32.src1_broadcast_dim2", "0"); + + ggml_free(ctx); + } + + { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 8, 4); + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2, 4); + ggml_tensor * output = ggml_mul(ctx, input, weight); + REQUIRE(input != nullptr); + REQUIRE(weight != nullptr); + REQUIRE(output != nullptr); + REQUIRE(!can_schedule_tensor(ctx, output)); + + ggml_free(ctx); + } +} + +static void run_binary_q8_publication_checks() { + const struct { + int64_t hidden; + int64_t tokens; + bool broadcast; + bool multiply; + bool side_use; + int projections; + const char * expected_kernel; + } cases[] = { + { 2048, 2, false, true, false, 1, "loom_libs:ggml_binary_f32_publish_q8_1_x4" }, + { 2048, 4, false, false, true, 2, "loom_libs:ggml_binary_f32_publish_q8_1_x4" }, + { 256, 2, true, true, false, 1, "loom_libs:ggml_binary_bc_f32_publish_q8_1_x4" }, + { 2048, 6, false, true, false, 1, "loom_libs:ggml_binary_f32" }, + { 2048, 2, false, true, false, 0, "loom_libs:ggml_binary_f32" }, + }; + + for (const auto & test : cases) { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * lhs = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, test.hidden, test.tokens); + ggml_tensor * rhs = test.broadcast ? ggml_new_tensor_1d(ctx, GGML_TYPE_F32, test.hidden) : + ggml_new_tensor_2d(ctx, GGML_TYPE_F32, test.hidden, test.tokens); + ggml_tensor * binary = test.multiply ? ggml_mul(ctx, lhs, rhs) : ggml_add(ctx, lhs, rhs); + REQUIRE(lhs != nullptr); + REQUIRE(rhs != nullptr); + REQUIRE(binary != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + for (int i = 0; i < test.projections; ++i) { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, test.hidden, 512); + ggml_build_forward_expand(graph, ggml_mul_mat(ctx, weight, binary)); + } + if (test.projections == 0 || test.side_use) { + ggml_build_forward_expand(graph, ggml_scale(ctx, binary, 0.5f)); + } + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const ggml::hrx::CommandPlan & plan = scheduler.plan(); + REQUIRE(plan.valid()); + const ggml::hrx::Value * binary_value = imported.graph.values().find_tensor(binary); + REQUIRE(binary_value != nullptr); + const auto producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return dispatch.bindings.size() >= 3 && dispatch.bindings[2].value == binary_value->id; + }); + REQUIRE(producer != plan.dispatches.end()); + REQUIRE(kernel_name_for_id(producer->kernel.kernel_id) == test.expected_kernel); + + const bool publishes = std::string_view(test.expected_kernel).find("publish_q8_1_x4") != std::string_view::npos; + const size_t q8_bytes = qwen_q8_1_x4_size(test.tokens, test.hidden); + const auto * q8 = plan.metadata.find_alternate_value(binary_value->id, GGML_TYPE_Q8_1, q8_bytes); + REQUIRE((q8 != nullptr) == publishes); + if (publishes) { + REQUIRE(producer->bindings.size() == 4); + REQUIRE(producer->bindings.back().value == q8->alternate_value); + REQUIRE(std::count_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "qwen3_moe:ggml_quantize_q8_1_x4_f32"; + }) == 0); + REQUIRE(std::count_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return &dispatch != &*producer && !dispatch.bindings.empty() && + dispatch.bindings.front().value == q8->alternate_value; + }) == test.projections); + } + + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + ggml_free(ctx); + } +} + +static void run_graph_snapshot_diagnostics_checks() { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * out = ggml_add(ctx, a, b); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(out != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + const std::string snapshot_json = ggml::hrx::serialize_graph_snapshot_json(imported.graph, "gfx1151", 42); + REQUIRE(string_contains(snapshot_json, "ggml-hrx-graph-snapshot-v1")); + REQUIRE(string_contains(snapshot_json, "ADD")); + + ggml::hrx::GraphSnapshotLoadResult loaded = ggml::hrx::load_graph_snapshot_json(snapshot_json); + REQUIRE(loaded.valid()); + REQUIRE(loaded.uid == 42); + REQUIRE(loaded.target == "gfx1151"); + REQUIRE(loaded.graph.nodes().size() == imported.graph.nodes().size()); + REQUIRE(loaded.graph.values().size() == imported.graph.values().size()); + + ggml::hrx::DispatchScheduler scheduler; + ggml::hrx::DispatchScheduleDiagnostics diagnostics; + REQUIRE(scheduler.schedule_graph(loaded.graph, { loaded.target }, &diagnostics)); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + + ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + loaded.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), loaded.target); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + + ggml_free(ctx); +} + +static void run_unmatched_graph_diagnostics_checks() { + ggml_init_params params = {}; + params.mem_size = 512 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 512, 512); + ggml_tensor * rows = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 512, 2); + ggml_tensor * indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 2); + REQUIRE(cache != nullptr); + REQUIRE(rows != nullptr); + REQUIRE(indices != nullptr); + ggml_tensor * out = ggml_set_rows(ctx, cache, rows, indices); + REQUIRE(out != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 1); + REQUIRE(imported.graph.nodes().front().op == GGML_OP_SET_ROWS); + + ggml::hrx::DispatchScheduler scheduler; + ggml::hrx::DispatchScheduleDiagnostics diagnostics; + REQUIRE(!scheduler.schedule_graph(imported.graph, test_dispatch_target(), &diagnostics)); + REQUIRE(diagnostics.unsupported_node != nullptr); + REQUIRE(diagnostics.unsupported_node_index == 0); + REQUIRE(diagnostics.unsupported_node->op == GGML_OP_SET_ROWS); + REQUIRE(!diagnostics.match.attempts.empty()); + REQUIRE(status_contains(scheduler.plan().status, "unsupported HRX node 0: SET_ROWS")); + + const std::string diagnostics_text = + ggml::hrx::format_schedule_diagnostics_text(imported.graph, scheduler.plan(), diagnostics); + REQUIRE(string_contains(diagnostics_text, "unsupported_node=0:SET_ROWS")); + REQUIRE(string_contains(diagnostics_text, "matcher_attempts=")); + const std::string diagnostics_json = + ggml::hrx::serialize_schedule_diagnostics_json(imported.graph, scheduler.plan(), diagnostics); + REQUIRE(string_contains(diagnostics_json, "SET_ROWS")); + REQUIRE(string_contains(diagnostics_json, "matcher_attempts")); + + ggml_free(ctx); +} + +static void run_completion_counter_plan_checks() { + ggml::hrx::Graph graph; + ggml::hrx::CommandPlan plan; + const ggml::hrx::ValueId first_counter(1000); + const ggml::hrx::ValueId second_counter(1001); + plan.completion_counter_requests.push_back({ first_counter, "first_completion_counter", 1 }); + plan.completion_counter_requests.push_back({ second_counter, "second_completion_counters", 2 }); + + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.empty()); + REQUIRE(commands.constant_initializations.empty()); + REQUIRE(commands.completion_counters.count == 3); + REQUIRE(commands.completion_counters.arena_offset == 0); + REQUIRE(commands.completion_counters.byte_count == 24); + REQUIRE(commands.transients.allocations.size() == 2); + REQUIRE(commands.transients.arena_size == 256); + REQUIRE(command_program_verifies(commands)); + + const ggml::hrx::TransientAllocation * first_allocation = + ggml::hrx::find_transient_allocation(commands.transients, first_counter); + const ggml::hrx::TransientAllocation * second_allocation = + ggml::hrx::find_transient_allocation(commands.transients, second_counter); + REQUIRE(first_allocation != nullptr); + REQUIRE(second_allocation != nullptr); + REQUIRE(first_allocation->size == sizeof(int32_t)); + REQUIRE(first_allocation->alignment == 16); + REQUIRE(first_allocation->arena_offset == commands.completion_counters.arena_offset); + REQUIRE(second_allocation->size == 2 * sizeof(int32_t)); + REQUIRE(second_allocation->alignment == 16); + REQUIRE(second_allocation->arena_offset == 16); + + ggml::hrx::CommandPlan bound_plan; + bound_plan.completion_counter_requests.push_back({ first_counter, "first_bound_completion_counter", 1 }); + bound_plan.completion_counter_requests.push_back({ second_counter, "second_bound_completion_counter", 1 }); + ggml::hrx::Dispatch first_dispatch; + first_dispatch.kernel = + ggml::hrx::make_kernel_specialization(ggml::hrx::kernel_catalog_ref("loom_libs", "ggml_binary_f32")); + add_binary_f32_exact_dispatch_params(first_dispatch, 1); + first_dispatch.bindings.push_back({ first_counter, 0, sizeof(int32_t) }); + first_dispatch.bindings.push_back({ second_counter, 0, sizeof(int32_t) }); + first_dispatch.bindings.push_back({ first_counter, 0, sizeof(int32_t) }); + bound_plan.dispatches.push_back(std::move(first_dispatch)); + ggml::hrx::Dispatch second_dispatch; + second_dispatch.kernel = + ggml::hrx::make_kernel_specialization(ggml::hrx::kernel_catalog_ref("loom_libs", "ggml_binary_f32")); + add_binary_f32_exact_dispatch_params(second_dispatch, 1); + second_dispatch.bindings.push_back({ second_counter, 0, sizeof(int32_t) }); + second_dispatch.bindings.push_back({ first_counter, 0, sizeof(int32_t) }); + second_dispatch.bindings.push_back({ second_counter, 0, sizeof(int32_t) }); + bound_plan.dispatches.push_back(std::move(second_dispatch)); + + const ggml::hrx::CommandProgram bound_commands = + ggml::hrx::build_command_program(graph, bound_plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(bound_commands.valid()); + REQUIRE(bound_commands.commands.size() == 2); + REQUIRE(bound_commands.completion_counters.count == 2); + REQUIRE(bound_commands.completion_counters.arena_offset == 0); + REQUIRE(bound_commands.completion_counters.byte_count == 20); + REQUIRE(command_program_verifies(bound_commands)); + + const ggml::hrx::TransientArenaAllocationRef transient_arena = { + dummy_hrx_buffer(0x8000), + bound_commands.transients.arena_size, + 7, + }; + ggml::hrx::PreparedCommandProgram prepared_shape; + for (const ggml::hrx::Command & prepared_source : bound_commands.commands) { + ggml::hrx::PreparedCommand prepared_command; + prepared_command.ordinal = prepared_source.ordinal; + prepared_command.kind = prepared_source.kind; + prepared_command.kernel.specialization = prepared_source.kernel; + for (const ggml::hrx::CommandBinding & binding : prepared_source.bindings) { + prepared_command.kernel.bindings.push_back({ + binding, { dummy_hrx_buffer(0x4000), 123, binding.length } + }); + } + prepared_shape.commands.push_back(std::move(prepared_command)); + } + prepared_shape.bound_transient_arena_allocation_id = 1; + + REQUIRE(ggml::hrx::bind_prepared_command_program_transients(bound_commands, transient_arena, prepared_shape)); + REQUIRE(prepared_shape.commands[0].kernel.bindings[0].ref.buffer == dummy_hrx_buffer(0x8000)); + REQUIRE(prepared_shape.commands[0].kernel.bindings[0].ref.offset == 0); + REQUIRE(prepared_shape.commands[0].kernel.bindings[1].ref.buffer == dummy_hrx_buffer(0x8000)); + REQUIRE(prepared_shape.commands[0].kernel.bindings[1].ref.offset == 16); + REQUIRE(prepared_shape.commands[1].kernel.bindings[0].ref.offset == 16); + REQUIRE(prepared_shape.commands[1].kernel.bindings[1].ref.offset == 0); +} + +static void run_graph_index_checks() { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 256, 1); + ggml_tensor * weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 256); + REQUIRE(input != nullptr); + REQUIRE(weight != nullptr); + ggml_tensor * rms = ggml_rms_norm(ctx, input, 0.000001f); + REQUIRE(rms != nullptr); + ggml_tensor * out = ggml_mul(ctx, rms, weight); + REQUIRE(out != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.has_index()); + REQUIRE(imported.graph.nodes().size() == 2); + const ggml::hrx::GraphNode * rms_node = &imported.graph.nodes()[0]; + const ggml::hrx::GraphNode * mul_node = &imported.graph.nodes()[1]; + REQUIRE(rms_node->op == GGML_OP_RMS_NORM); + REQUIRE(mul_node->op == GGML_OP_MUL); + const ggml::hrx::RmsNormParams * rms_params = ggml::hrx::op_params_as(rms_node->params); + REQUIRE(rms_params != nullptr); + REQUIRE(rms_params->eps == 0.000001f); + REQUIRE(imported.graph.index().producer(rms_node->output) == rms_node); + REQUIRE(imported.graph.index().producer(mul_node->output) == mul_node); + REQUIRE(imported.graph.index().has_single_consumer(rms_node->output)); + const std::vector & consumers = imported.graph.index().consumers(rms_node->output); + REQUIRE(consumers.size() == 1); + REQUIRE(consumers.front() == mul_node); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + + ggml_free(ctx); + + params.mem_size = 2 * 1024 * 1024; + params.no_alloc = true; + ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 256, 4); + REQUIRE(input != nullptr); + rms = ggml_rms_norm(ctx, input, 0.00001f); + REQUIRE(rms != nullptr); + + graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, rms); + + imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 1); + + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches[0]; + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rmsnorm_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == 4); + require_compile_parameter(dispatch, "ggml.rmsnorm_f32.hidden_size", "256"); + require_compile_parameter(dispatch, "ggml.rmsnorm_f32.input_stride", "256"); + REQUIRE(dispatch.kernel.compile_parameters.count("ggml.rmsnorm_f32.rms_epsilon") == 1); + REQUIRE(dispatch.bindings.size() == 2); + + { + constexpr int64_t view_hidden_size = 256; + constexpr int64_t view_tokens = 23; + constexpr int64_t view_stride = 512; + constexpr size_t view_offset = 4 * sizeof(float); + ggml_tensor * storage = ggml_new_tensor_1d( + ctx, GGML_TYPE_F32, view_offset / sizeof(float) + (view_tokens - 1) * view_stride + view_hidden_size); + ggml_tensor * view = + ggml_view_2d(ctx, storage, view_hidden_size, view_tokens, view_stride * sizeof(float), view_offset); + ggml_tensor * norm = ggml_rms_norm(ctx, view, 0.00001f); + REQUIRE(storage != nullptr); + REQUIRE(view != nullptr); + REQUIRE(norm != nullptr); + + ggml_cgraph * view_graph = ggml_new_graph(ctx); + REQUIRE(view_graph != nullptr); + ggml_build_forward_expand(view_graph, norm); + + ggml::hrx::GraphImportResult view_imported = ggml::hrx::import_ggml_graph(*view_graph); + REQUIRE(view_imported.valid()); + REQUIRE(view_imported.graph.nodes().size() == 2); + REQUIRE(view_imported.graph.nodes()[0].op == GGML_OP_VIEW); + REQUIRE(view_imported.graph.nodes()[1].op == GGML_OP_RMS_NORM); + + const ggml::hrx::Value * view_value = + view_imported.graph.values().find(view_imported.graph.nodes()[1].inputs[0]); + REQUIRE(view_value != nullptr); + const ggml::hrx::Value * storage_value = view_imported.graph.values().find(view_value->storage_root); + REQUIRE(storage_value != nullptr); + REQUIRE(view_value->storage_offset == view_offset); + + REQUIRE(scheduler.schedule_graph(view_imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + const ggml::hrx::Dispatch & view_dispatch = scheduler.plan().dispatches[0]; + REQUIRE(kernel_name_for_id(view_dispatch.kernel.kernel_id) == "loom_libs:ggml_rmsnorm_f32"); + REQUIRE(view_dispatch.kernel.integer_parameters.at("token_count") == view_tokens); + require_compile_parameter(view_dispatch, "ggml.rmsnorm_f32.hidden_size", std::to_string(view_hidden_size)); + require_compile_parameter(view_dispatch, "ggml.rmsnorm_f32.input_stride", std::to_string(view_stride)); + REQUIRE(view_dispatch.bindings.size() == 2); + REQUIRE(view_dispatch.bindings[0].value == storage_value->id); + REQUIRE(view_dispatch.bindings[0].offset == view_offset); + REQUIRE(view_dispatch.bindings[0].length == static_cast(view_tokens - 1) * view->nb[1] + + static_cast(view_hidden_size) * sizeof(float)); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + view_imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + + { + constexpr int64_t hidden_size = 256; + constexpr int64_t token_count = 64; + constexpr int64_t input_stride = 288; + const size_t storage_elements = static_cast(token_count - 1) * input_stride + hidden_size; + ggml_tensor * storage = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, storage_elements); + ggml_tensor * view = ggml_view_2d(ctx, storage, hidden_size, token_count, input_stride * sizeof(float), 0); + ggml_tensor * weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, hidden_size); + ggml_tensor * norm = ggml_rms_norm(ctx, view, 0.00001f); + ggml_tensor * scaled = ggml_mul(ctx, norm, weight); + REQUIRE(storage != nullptr); + REQUIRE(view != nullptr); + REQUIRE(weight != nullptr); + REQUIRE(norm != nullptr); + REQUIRE(scaled != nullptr); + + ggml_cgraph * strided_graph = ggml_new_graph(ctx); + REQUIRE(strided_graph != nullptr); + ggml_build_forward_expand(strided_graph, scaled); + ggml::hrx::GraphImportResult strided_imported = ggml::hrx::import_ggml_graph(*strided_graph); + REQUIRE(strided_imported.valid()); + REQUIRE(scheduler.schedule_graph(strided_imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + const ggml::hrx::Dispatch & strided_dispatch = scheduler.plan().dispatches[0]; + REQUIRE(kernel_name_for_id(strided_dispatch.kernel.kernel_id) == "loom_libs:ggml_rmsnorm_binary_strided_f32"); + require_compile_parameter(strided_dispatch, "ggml.rmsnorm_binary_f32.hidden_size", "256"); + require_compile_parameter(strided_dispatch, "ggml.rmsnorm_binary_f32.input_stride", "288"); + require_compile_parameter(strided_dispatch, "ggml.rmsnorm_binary_f32.op", "2"); + REQUIRE(strided_dispatch.bindings.size() == 3); + REQUIRE(strided_dispatch.bindings[0].length == + static_cast(token_count - 1) * view->nb[1] + hidden_size * sizeof(float)); + + ggml_tensor * full_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden_size, token_count); + ggml_tensor * rejected_norm = ggml_rms_norm(ctx, view, 0.00001f); + ggml_tensor * rejected_scaled = ggml_mul(ctx, rejected_norm, full_weight); + REQUIRE(full_weight != nullptr); + REQUIRE(rejected_scaled != nullptr); + ggml_cgraph * rejected_graph = ggml_new_graph(ctx); + REQUIRE(rejected_graph != nullptr); + ggml_build_forward_expand(rejected_graph, rejected_scaled); + ggml::hrx::GraphImportResult rejected_imported = ggml::hrx::import_ggml_graph(*rejected_graph); + REQUIRE(rejected_imported.valid()); + REQUIRE(scheduler.schedule_graph(rejected_imported.graph, test_dispatch_target())); + for (const ggml::hrx::Dispatch & rejected_dispatch : scheduler.plan().dispatches) { + REQUIRE(kernel_name_for_id(rejected_dispatch.kernel.kernel_id) != + "loom_libs:ggml_rmsnorm_binary_strided_f32"); + } + + ggml_tensor * second_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, hidden_size); + ggml_tensor * shared_norm = ggml_rms_norm(ctx, view, 0.00001f); + ggml_tensor * scaled_a = ggml_mul(ctx, shared_norm, weight); + ggml_tensor * scaled_b = ggml_add(ctx, shared_norm, second_weight); + ggml_tensor * joined = ggml_add(ctx, scaled_a, scaled_b); + REQUIRE(second_weight != nullptr); + REQUIRE(joined != nullptr); + ggml_cgraph * fanout_graph = ggml_new_graph(ctx); + REQUIRE(fanout_graph != nullptr); + ggml_build_forward_expand(fanout_graph, joined); + ggml::hrx::GraphImportResult fanout_imported = ggml::hrx::import_ggml_graph(*fanout_graph); + REQUIRE(fanout_imported.valid()); + REQUIRE(scheduler.schedule_graph(fanout_imported.graph, test_dispatch_target())); + for (const ggml::hrx::Dispatch & fanout_dispatch : scheduler.plan().dispatches) { + REQUIRE(kernel_name_for_id(fanout_dispatch.kernel.kernel_id) != + "loom_libs:ggml_rmsnorm_binary_strided_f32"); + } + } + + ggml_free(ctx); +} + +static void run_rope_set_rows_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 2 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t head_size = 8; + constexpr int64_t head_count = 2; + constexpr int64_t token_count = 3; + constexpr int64_t cache_rows = 8; + constexpr int64_t hidden_size = head_size * head_count; + + const auto check_rope_concat = [&](int64_t concat_head_count, int64_t concat_token_count, int64_t n_dims) { + constexpr int64_t prefix_size = 64; + constexpr int64_t rope_size = 32; + constexpr int64_t source_size = prefix_size + rope_size; + const size_t source_elements = static_cast(source_size) * concat_head_count * concat_token_count; + ggml_tensor * storage = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, source_elements); + ggml_tensor * prefix = + ggml_view_3d(ctx, storage, prefix_size, concat_head_count, concat_token_count, source_size * sizeof(float), + source_size * concat_head_count * sizeof(float), 0); + ggml_tensor * rope_input = + ggml_view_3d(ctx, storage, rope_size, concat_head_count, concat_token_count, source_size * sizeof(float), + source_size * concat_head_count * sizeof(float), prefix_size * sizeof(float)); + ggml_tensor * positions = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, concat_token_count); + ggml_tensor * freqs = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, n_dims / 2); + ggml_tensor * rope = ggml_rope_ext(ctx, rope_input, positions, freqs, n_dims, GGML_ROPE_TYPE_NEOX, 0, 10000.0f, + 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + ggml_tensor * output = ggml_concat(ctx, prefix, rope, 0); + REQUIRE(storage != nullptr); + REQUIRE(prefix != nullptr); + REQUIRE(rope_input != nullptr); + REQUIRE(positions != nullptr); + REQUIRE(freqs != nullptr); + REQUIRE(rope != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rope_concat_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == concat_token_count); + require_compile_parameter(dispatch, "ggml.rope_concat_f32.prefix_size", std::to_string(prefix_size)); + require_compile_parameter(dispatch, "ggml.rope_concat_f32.rope_size", std::to_string(rope_size)); + require_compile_parameter(dispatch, "ggml.rope_concat_f32.n_dims", std::to_string(n_dims)); + require_compile_parameter(dispatch, "ggml.rope_concat_f32.head_count", std::to_string(concat_head_count)); + require_compile_parameter(dispatch, "ggml.rope_concat_f32.token_capacity", std::to_string(concat_token_count)); + require_compile_parameter(dispatch, "ggml.rope_concat_f32.input_stride1", std::to_string(source_size)); + require_compile_parameter(dispatch, "ggml.rope_concat_f32.input_stride2", + std::to_string(source_size * concat_head_count)); + require_compile_parameter(dispatch, "ggml.rope_concat_f32.mode", "2"); + REQUIRE(dispatch.bindings.size() == 5); + REQUIRE(dispatch.bindings[1].length == source_elements * sizeof(float)); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + }; + + check_rope_concat(40, 64, 32); + check_rope_concat(3, 2, 16); + + { + ggml_tensor * input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * pos = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, token_count); + ggml_tensor * freqs = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, head_size / 2); + ggml_tensor * output = ggml_rope_ext(ctx, input, pos, freqs, head_size, GGML_ROPE_TYPE_NEOX, 0, 10000.0f, 1.0f, + 0.0f, 1.0f, 0.0f, 0.0f); + REQUIRE(input != nullptr); + REQUIRE(pos != nullptr); + REQUIRE(freqs != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 1); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_ROPE); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rope_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == token_count); + require_compile_parameter(dispatch, "ggml.rope_f32.head_size", "8"); + require_compile_parameter(dispatch, "ggml.rope_f32.n_dims", "8"); + require_compile_parameter(dispatch, "ggml.rope_f32.head_count", "2"); + require_compile_parameter(dispatch, "ggml.rope_f32.token_capacity", "3"); + require_compile_parameter(dispatch, "ggml.rope_f32.mode", "2"); + REQUIRE(dispatch.bindings.size() == 5); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + + { + ggml_tensor * input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * pos = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, token_count); + ggml_tensor * freqs = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, head_size / 2); + ggml_tensor * output = ggml_rope_ext(ctx, input, pos, freqs, head_size, GGML_ROPE_TYPE_NORMAL, 0, 10000.0f, + 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + REQUIRE(input != nullptr); + REQUIRE(pos != nullptr); + REQUIRE(freqs != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 1); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_ROPE); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rope_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == token_count); + require_compile_parameter(dispatch, "ggml.rope_f32.head_size", "8"); + require_compile_parameter(dispatch, "ggml.rope_f32.n_dims", "8"); + require_compile_parameter(dispatch, "ggml.rope_f32.head_count", "2"); + require_compile_parameter(dispatch, "ggml.rope_f32.token_capacity", "3"); + require_compile_parameter(dispatch, "ggml.rope_f32.mode", "0"); + REQUIRE(dispatch.bindings.size() == 5); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + + { + ggml_tensor * input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * pos = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, token_count); + ggml_tensor * output = ggml_rope(ctx, input, pos, head_size, GGML_ROPE_TYPE_NEOX); + REQUIRE(input != nullptr); + REQUIRE(pos != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + REQUIRE(scheduler.plan().transients.size() == 2); + REQUIRE(scheduler.plan().constant_initializations.size() == 2); + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rope_f32"); + require_compile_parameter(dispatch, "ggml.rope_f32.mode", "2"); + require_compile_parameter(dispatch, "ggml.rope_f32.n_dims", "8"); + } + + { + constexpr int64_t partial_head_size = 128; + constexpr int64_t partial_n_dims = 96; + constexpr int64_t partial_head_count = 24; + constexpr int64_t partial_tokens = 2; + ggml_tensor * input = + ggml_new_tensor_3d(ctx, GGML_TYPE_F32, partial_head_size, partial_head_count, partial_tokens); + ggml_tensor * pos = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, partial_tokens); + ggml_tensor * freqs = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, partial_n_dims / 2); + ggml_tensor * output = ggml_rope_ext(ctx, input, pos, freqs, partial_n_dims, GGML_ROPE_TYPE_NEOX, 0, 10000.0f, + 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + REQUIRE(input != nullptr); + REQUIRE(pos != nullptr); + REQUIRE(freqs != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rope_f32"); + require_compile_parameter(dispatch, "ggml.rope_f32.head_size", std::to_string(partial_head_size)); + require_compile_parameter(dispatch, "ggml.rope_f32.n_dims", std::to_string(partial_n_dims)); + require_compile_parameter(dispatch, "ggml.rope_f32.head_count", std::to_string(partial_head_count)); + require_compile_parameter(dispatch, "ggml.rope_f32.token_capacity", std::to_string(partial_tokens)); + require_compile_parameter(dispatch, "ggml.rope_f32.mode", "2"); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + + { + constexpr int64_t partial_head_size = 128; + constexpr int64_t partial_n_dims = 96; + constexpr int64_t partial_head_count = 24; + constexpr int64_t partial_tokens = 2; + constexpr int64_t partial_token_stride = 5120; + constexpr size_t input_offset = 64; + constexpr size_t input_elements = input_offset / sizeof(float) + + static_cast(partial_tokens - 1) * partial_token_stride + + partial_head_size * partial_head_count; + + ggml_tensor * storage = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, input_elements); + ggml_tensor * input = + ggml_view_3d(ctx, storage, partial_head_size, partial_head_count, partial_tokens, + partial_head_size * sizeof(float), partial_token_stride * sizeof(float), input_offset); + ggml_tensor * pos = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, partial_tokens); + ggml_tensor * freqs = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, partial_n_dims / 2); + ggml_tensor * output = ggml_rope_ext(ctx, input, pos, freqs, partial_n_dims, GGML_ROPE_TYPE_NEOX, 0, 10000.0f, + 1.0f, 0.0f, 1.19024f, 32.0f, 1.0f); + REQUIRE(storage != nullptr); + REQUIRE(input != nullptr); + REQUIRE(pos != nullptr); + REQUIRE(freqs != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rope_f32"); + require_compile_parameter(dispatch, "ggml.rope_f32.head_size", std::to_string(partial_head_size)); + require_compile_parameter(dispatch, "ggml.rope_f32.n_dims", std::to_string(partial_n_dims)); + require_compile_parameter(dispatch, "ggml.rope_f32.head_count", std::to_string(partial_head_count)); + require_compile_parameter(dispatch, "ggml.rope_f32.token_capacity", std::to_string(partial_tokens)); + require_compile_parameter(dispatch, "ggml.rope_f32.input_stride1", std::to_string(partial_head_size)); + require_compile_parameter(dispatch, "ggml.rope_f32.input_stride2", std::to_string(partial_token_stride)); + require_compile_parameter(dispatch, "ggml.rope_f32.mscale", "1.19024003"); + require_compile_parameter(dispatch, "ggml.rope_f32.mode", "2"); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + + { + ggml_tensor * input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * pos = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, token_count); + ggml_tensor * raw_scales = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 1, 1, token_count); + ggml_tensor * scales = ggml_scale(ctx, raw_scales, 0.5f); + ggml_tensor * rope = ggml_rope(ctx, input, pos, head_size, GGML_ROPE_TYPE_NORMAL); + ggml_tensor * output = ggml_mul(ctx, rope, scales); + REQUIRE(input != nullptr); + REQUIRE(pos != nullptr); + REQUIRE(raw_scales != nullptr); + REQUIRE(scales != nullptr); + REQUIRE(rope != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + const size_t rope_index = producer_index_for_tensor(imported.graph, rope); + const size_t scale_index = producer_index_for_tensor(imported.graph, scales); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan; + ggml::hrx::DispatchMatch match; + + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, rope_index, match)); + REQUIRE(match.dispatches.size() == 1); + REQUIRE(kernel_name_for_id(match.dispatches.front().kernel.kernel_id) == "loom_libs:ggml_rope_f32"); + + covered_nodes[scale_index] = true; + match = {}; + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, rope_index, match)); + REQUIRE(match.dispatches.size() == 1); + REQUIRE(kernel_name_for_id(match.dispatches.front().kernel.kernel_id) == "loom_libs:ggml_rope_token_scale_f32"); + } + + { + ggml_tensor * input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * pos = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, token_count); + ggml_tensor * rope = ggml_rope(ctx, input, pos, head_size, GGML_ROPE_TYPE_NEOX); + ggml_tensor * output = ggml_scale(ctx, rope, 0.125f); + REQUIRE(input != nullptr); + REQUIRE(pos != nullptr); + REQUIRE(rope != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rope_scale_f32"); + require_compile_parameter(dispatch, "ggml.rope_scale_f32.scale", "0.125"); + REQUIRE(dispatch.bindings.size() == 5); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + + { + ggml_tensor * input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * pos = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, token_count); + ggml_tensor * scales = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 1, 1, token_count); + ggml_tensor * rope = ggml_rope(ctx, input, pos, head_size, GGML_ROPE_TYPE_NORMAL); + ggml_tensor * output = ggml_mul(ctx, rope, scales); + REQUIRE(input != nullptr); + REQUIRE(pos != nullptr); + REQUIRE(scales != nullptr); + REQUIRE(rope != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rope_token_scale_f32"); + REQUIRE(dispatch.bindings.size() == 6); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + + { + ggml_tensor * input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * pos = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, token_count); + ggml_tensor * rope = ggml_rope(ctx, input, pos, head_size, GGML_ROPE_TYPE_NEOX); + ggml_tensor * output = ggml_scale(ctx, rope, 0.125f); + ggml_tensor * side = ggml_scale(ctx, rope, 0.25f); + REQUIRE(input != nullptr); + REQUIRE(pos != nullptr); + REQUIRE(rope != nullptr); + REQUIRE(output != nullptr); + REQUIRE(side != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + ggml_build_forward_expand(graph, side); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 3); + REQUIRE(kernel_name_for_id(scheduler.plan().dispatches.front().kernel.kernel_id) == "loom_libs:ggml_rope_f32"); + } + + { + ggml_tensor * cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, hidden_size, cache_rows); + ggml_tensor * rows = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden_size, token_count); + ggml_tensor * indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, token_count); + ggml_tensor * output = ggml_set_rows(ctx, cache, rows, indices); + REQUIRE(cache != nullptr); + REQUIRE(rows != nullptr); + REQUIRE(indices != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 1); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_SET_ROWS); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_set_rows"); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == token_count); + REQUIRE(dispatch.kernel.integer_parameters.at("cache_row_count") == cache_rows); + REQUIRE(dispatch.kernel.integer_parameters.at("hidden_size") == hidden_size); + require_compile_parameter(dispatch, "ggml.set_rows.input_format", "32"); + require_compile_parameter(dispatch, "ggml.set_rows.output_format", "16"); + REQUIRE(dispatch.bindings.size() == 3); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + + { + constexpr int64_t row_stride = hidden_size * 5; + constexpr size_t row_offset = 4 * sizeof(float); + ggml_tensor * cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, hidden_size, cache_rows); + ggml_tensor * storage = ggml_new_tensor_1d( + ctx, GGML_TYPE_F32, row_offset / sizeof(float) + (token_count - 1) * row_stride + hidden_size); + ggml_tensor * rows = + ggml_view_2d(ctx, storage, hidden_size, token_count, row_stride * sizeof(float), row_offset); + ggml_tensor * indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, token_count); + ggml_tensor * output = ggml_set_rows(ctx, cache, rows, indices); + REQUIRE(cache != nullptr); + REQUIRE(storage != nullptr); + REQUIRE(rows != nullptr); + REQUIRE(indices != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 2); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_VIEW); + REQUIRE(imported.graph.nodes()[1].op == GGML_OP_SET_ROWS); + + const ggml::hrx::Value * rows_value = imported.graph.values().find(imported.graph.nodes()[1].inputs[0]); + REQUIRE(rows_value != nullptr); + const ggml::hrx::Value * storage_value = imported.graph.values().find(rows_value->storage_root); + REQUIRE(storage_value != nullptr); + REQUIRE(rows_value->storage_offset == row_offset); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_set_rows"); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == token_count); + REQUIRE(dispatch.kernel.integer_parameters.at("cache_row_count") == cache_rows); + REQUIRE(dispatch.kernel.integer_parameters.at("hidden_size") == hidden_size); + require_compile_parameter(dispatch, "ggml.set_rows.input_format", "32"); + require_compile_parameter(dispatch, "ggml.set_rows.output_format", "16"); + require_compile_parameter(dispatch, "ggml.set_rows.input_stride", std::to_string(row_stride)); + REQUIRE(dispatch.bindings.size() == 3); + REQUIRE(dispatch.bindings[0].value == storage_value->id); + REQUIRE(dispatch.bindings[0].offset == row_offset); + REQUIRE(dispatch.bindings[0].length == + static_cast(token_count - 1) * rows->nb[1] + static_cast(hidden_size) * sizeof(float)); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + + { + ggml_tensor * cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, hidden_size, cache_rows); + ggml_tensor * rows = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, hidden_size, token_count); + ggml_tensor * indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, token_count); + ggml_tensor * output = ggml_set_rows(ctx, cache, rows, indices); + REQUIRE(cache != nullptr); + REQUIRE(rows != nullptr); + REQUIRE(indices != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_set_rows"); + require_compile_parameter(dispatch, "ggml.set_rows.input_format", "16"); + require_compile_parameter(dispatch, "ggml.set_rows.output_format", "16"); + } + + { + ggml_tensor * input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * pos = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, token_count); + ggml_tensor * freqs = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, head_size / 2); + ggml_tensor * rope = ggml_rope_ext(ctx, input, pos, freqs, head_size, GGML_ROPE_TYPE_NEOX, 0, 10000.0f, 1.0f, + 0.0f, 1.0f, 0.0f, 0.0f); + ggml_tensor * rows = ggml_reshape_2d(ctx, rope, hidden_size, token_count); + ggml_tensor * cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, hidden_size, cache_rows); + ggml_tensor * indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, token_count); + ggml_tensor * output = ggml_set_rows(ctx, cache, rows, indices); + REQUIRE(input != nullptr); + REQUIRE(pos != nullptr); + REQUIRE(freqs != nullptr); + REQUIRE(rope != nullptr); + REQUIRE(rows != nullptr); + REQUIRE(cache != nullptr); + REQUIRE(indices != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 3); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_ROPE); + REQUIRE(imported.graph.nodes()[1].op == GGML_OP_RESHAPE); + REQUIRE(imported.graph.nodes()[2].op == GGML_OP_SET_ROWS); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rope_set_rows_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == token_count); + REQUIRE(dispatch.kernel.integer_parameters.at("cache_row_count") == cache_rows); + require_compile_parameter(dispatch, "ggml.rope_set_rows_f32.output_format", "16"); + require_compile_parameter(dispatch, "ggml.rope_set_rows_f32.n_dims", "8"); + require_compile_parameter(dispatch, "ggml.rope_set_rows_f32.mode", "2"); + REQUIRE(dispatch.bindings.size() == 6); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + + { + ggml_tensor * input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * pos = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, token_count); + ggml_tensor * freqs = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, head_size / 2); + ggml_tensor * rope = ggml_rope_ext(ctx, input, pos, freqs, head_size, GGML_ROPE_TYPE_NORMAL, 0, 10000.0f, 1.0f, + 0.0f, 1.0f, 0.0f, 0.0f); + ggml_tensor * rows = ggml_reshape_2d(ctx, rope, hidden_size, token_count); + ggml_tensor * cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, hidden_size, cache_rows); + ggml_tensor * indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, token_count); + ggml_tensor * output = ggml_set_rows(ctx, cache, rows, indices); + REQUIRE(input != nullptr); + REQUIRE(pos != nullptr); + REQUIRE(freqs != nullptr); + REQUIRE(rope != nullptr); + REQUIRE(rows != nullptr); + REQUIRE(cache != nullptr); + REQUIRE(indices != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 3); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_ROPE); + REQUIRE(imported.graph.nodes()[1].op == GGML_OP_RESHAPE); + REQUIRE(imported.graph.nodes()[2].op == GGML_OP_SET_ROWS); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rope_set_rows_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == token_count); + REQUIRE(dispatch.kernel.integer_parameters.at("cache_row_count") == cache_rows); + require_compile_parameter(dispatch, "ggml.rope_set_rows_f32.output_format", "16"); + require_compile_parameter(dispatch, "ggml.rope_set_rows_f32.n_dims", "8"); + require_compile_parameter(dispatch, "ggml.rope_set_rows_f32.mode", "0"); + REQUIRE(dispatch.bindings.size() == 6); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + + ggml_free(ctx); +} + +static void run_graph_traversal_checks() { + { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * c = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * d = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * out0 = ggml_add(ctx, a, b); + ggml_tensor * out1 = ggml_add(ctx, c, d); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(c != nullptr); + REQUIRE(d != nullptr); + REQUIRE(out0 != nullptr); + REQUIRE(out1 != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out0); + ggml_build_forward_expand(graph, out1); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 2); + REQUIRE(imported.graph.nodes()[0].output == imported.graph.values().find_tensor(out0)->id); + REQUIRE(imported.graph.nodes()[1].output == imported.graph.values().find_tensor(out1)->id); + + const std::vector order = traversal_indices(imported.graph); + REQUIRE(order.size() == 2); + REQUIRE(order[0] == 0); + REQUIRE(order[1] == 1); + + ggml_free(ctx); + } + + { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * add_out = ggml_add(ctx, a, b); + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 4, 3); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 4, 2); + ggml_tensor * matmul = ggml_mul_mat(ctx, weight, input); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(add_out != nullptr); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + REQUIRE(matmul != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, add_out); + ggml_build_forward_expand(graph, matmul); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 2); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_ADD); + REQUIRE(imported.graph.nodes()[1].op == GGML_OP_MUL_MAT); + + const std::vector order = traversal_indices(imported.graph); + REQUIRE(order.size() == 2); + REQUIRE(order[0] == 1); + REQUIRE(order[1] == 0); + + ggml_free(ctx); + } + + { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 256, 1); + ggml_tensor * weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 256); + ggml_tensor * rms = ggml_rms_norm(ctx, input, 0.000001f); + ggml_tensor * out = ggml_mul(ctx, rms, weight); + REQUIRE(input != nullptr); + REQUIRE(weight != nullptr); + REQUIRE(rms != nullptr); + REQUIRE(out != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 2); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_RMS_NORM); + REQUIRE(imported.graph.nodes()[1].op == GGML_OP_MUL); + + const std::vector order = traversal_indices(imported.graph); + REQUIRE(order.size() == 2); + REQUIRE(order[0] == 0); + REQUIRE(order[1] == 1); + + ggml_free(ctx); + } + + { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * c = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * left = ggml_add(ctx, a, b); + ggml_tensor * right = ggml_mul(ctx, a, c); + ggml_tensor * join = ggml_add(ctx, left, right); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(c != nullptr); + REQUIRE(left != nullptr); + REQUIRE(right != nullptr); + REQUIRE(join != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, join); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 3); + + const size_t left_index = producer_index_for_tensor(imported.graph, left); + const size_t right_index = producer_index_for_tensor(imported.graph, right); + const size_t join_index = producer_index_for_tensor(imported.graph, join); + + const std::vector order = traversal_indices(imported.graph); + REQUIRE(order.size() == 3); + REQUIRE(find_position(order, join_index) > find_position(order, left_index)); + REQUIRE(find_position(order, join_index) > find_position(order, right_index)); + + ggml_free(ctx); + } +} + +static bool matmul_graph_is_supported(ggml_context * ctx, ggml_tensor * output); + +static void schedule_single_matmul_command(ggml_context * ctx, + ggml_tensor * output, + const char * expected_kernel_name, + int64_t expected_token_count, + int64_t expected_input_size, + int64_t expected_output_size) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 1); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_MUL_MAT); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + const ggml_type weight_type = output->src[0]->type; + REQUIRE(!scheduler.plan().dispatches.empty()); + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.back(); + const std::string kernel_name = kernel_name_for_id(dispatch.kernel.kernel_id); + const bool vector_kernel = string_contains(kernel_name, "ggml_mul_mat_vector"); + const bool q6_vector_kernel = kernel_name == "loom_libs:ggml_mul_mat_vector_q6_f32_f32"; + const bool aligned_dynamic_kernel = kernel_name == "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32_aligned"; + const bool codebook_kernel = string_contains(kernel_name, "_iq1_") || string_contains(kernel_name, "_iq2_s_") || + string_contains(kernel_name, "_iq3_s_") || + string_contains(kernel_name, "_iq3_xxs_"); + const std::string expected_kernel(expected_kernel_name); + const bool q8_capable_kernel = expected_kernel == "loom_libs:ggml_mul_mat_f32_f32_wmma" || + expected_kernel == "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32"; + const bool q8_input = !vector_kernel && expected_token_count <= 5 && expected_output_size % 64 == 0 && + (weight_type == GGML_TYPE_Q4_K || weight_type == GGML_TYPE_Q6_K) && q8_capable_kernel; + const size_t command_count = q8_input ? 2 : 1; + REQUIRE(scheduler.plan().dispatches.size() == command_count); + if (q8_input) { + REQUIRE(kernel_name_for_id(scheduler.plan().dispatches.front().kernel.kernel_id) == + "qwen3_moe:ggml_quantize_q8_1_x4_f32"); + } + + const bool accepts_vector_token1 = + expected_token_count == 1 && (expected_kernel == "loom_libs:ggml_mul_mat_f32_f32_wmma" || + expected_kernel == "loom_libs:ggml_mul_mat_f32_f32_decode_wave64" || + expected_kernel == "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma" || + expected_kernel == "qwen3_moe:qwen3_moe_router_projection_f32_four_row_wave32"); + if (vector_kernel && accepts_vector_token1) { + REQUIRE(string_contains(kernel_name, "ggml_mul_mat_vector")); + } else { + REQUIRE(kernel_name == expected_kernel_name); + } + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == expected_token_count); + if (vector_kernel) { + REQUIRE(dispatch.bindings.size() == (codebook_kernel ? 5 : q6_vector_kernel ? 3 : 4)); + if (!q6_vector_kernel) { + require_compile_parameter(dispatch, "ggml.workload.token_capacity", std::to_string(expected_token_count)); + } + require_compile_parameter(dispatch, "ggml.matmul.vector.input_size", std::to_string(expected_input_size)); + require_compile_parameter(dispatch, "ggml.matmul.vector.output_size", std::to_string(expected_output_size)); + require_compile_parameter(dispatch, "ggml.matmul.vector.output_accumulation", "0"); + if (q6_vector_kernel) { + require_compile_parameter(dispatch, "ggml.matmul.vector.q6_storage", "0"); + } + } else if (string_contains(kernel_name, "ggml_mul_mat_q6_k_packed_token1_f16_wmma")) { + REQUIRE(dispatch.bindings.size() == 3); + require_compile_parameter(dispatch, "ggml.mul_mat_q6_k_packed.input_size", std::to_string(expected_input_size)); + require_compile_parameter(dispatch, "ggml.mul_mat_q6_k_packed.output_size", + std::to_string(expected_output_size)); + require_compile_parameter(dispatch, "ggml.mul_mat_q6_k_packed.output_accumulation", "0"); + } else if (string_contains(kernel_name, "_decode")) { + REQUIRE(dispatch.bindings.size() == 3); + REQUIRE(dispatch.kernel.integer_parameters.at("input_size") == expected_input_size); + REQUIRE(dispatch.kernel.integer_parameters.at("output_size") == expected_output_size); + require_compile_parameter(dispatch, "ggml.mul_mat_f32_f32_decode.token_capacity", + std::to_string(expected_token_count)); + require_compile_parameter(dispatch, "ggml.mul_mat_f32_f32_decode.output_capacity", + std::to_string(expected_output_size)); + const ggml::hrx::Value * weight = imported.graph.values().find(dispatch.bindings[1].value); + REQUIRE(weight != nullptr); + require_compile_parameter(dispatch, "ggml.mul_mat_f32_f32_decode.weight_format", + std::to_string(matmul_weight_format_config(weight->type))); + } else if (string_contains(kernel_name, "ggml_mul_mat")) { + REQUIRE(dispatch.bindings.size() == (codebook_kernel ? 4 : 3)); + if (aligned_dynamic_kernel) { + REQUIRE(dispatch.kernel.workload_specialization == ggml::hrx::WorkloadSpecialization::Dynamic); + REQUIRE(dispatch.kernel.compile_parameters.count("ggml.workload.token_capacity") == 0); + } else { + require_compile_parameter(dispatch, "ggml.workload.token_capacity", std::to_string(expected_token_count)); + } + require_compile_parameter(dispatch, "ggml.mul_mat.input_size", std::to_string(expected_input_size)); + require_compile_parameter(dispatch, "ggml.mul_mat.output_size", std::to_string(expected_output_size)); + require_compile_parameter(dispatch, "ggml.mul_mat.output_accumulation", "0"); + require_compile_parameter(dispatch, "ggml.mul_mat.output_unary_op", + std::to_string(ggml::hrx::unary_kind_config_value(ggml::hrx::UnaryKind::Identity))); + const ggml::hrx::Value * weight = imported.graph.values().find(dispatch.bindings[1].value); + REQUIRE(weight != nullptr); + require_compile_parameter(dispatch, "ggml.mul_mat.weight_format", + std::to_string(q8_input ? (weight_type == GGML_TYPE_Q4_K ? 44 : 46) : + matmul_weight_format_config(weight->type))); + } else if (string_contains(kernel_name, "dense_linear")) { + REQUIRE(dispatch.bindings.size() == 3); + require_compile_parameter(dispatch, "qwen3_moe.workload.token_capacity", std::to_string(expected_token_count)); + require_compile_parameter(dispatch, "qwen3_moe.dense_quantized.input_size", + std::to_string(expected_input_size)); + require_compile_parameter(dispatch, "qwen3_moe.dense_quantized.output_size", + std::to_string(expected_output_size)); + require_compile_parameter(dispatch, "qwen3_moe.dense_quantized.output_accumulation", "0"); + } else { + REQUIRE(dispatch.bindings.size() == 3); + require_compile_parameter(dispatch, "qwen3_moe.workload.token_capacity", std::to_string(expected_token_count)); + require_compile_parameter(dispatch, "qwen3_moe.model.hidden_size", std::to_string(expected_input_size)); + require_compile_parameter(dispatch, "qwen3_moe.router.expert_count", std::to_string(expected_output_size)); + } + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == command_count); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.back().bindings.size() == + (codebook_kernel ? (vector_kernel ? 5 : 4) : vector_kernel && !q6_vector_kernel ? 4 : 3)); + REQUIRE(commands.commands.back().bindings[0].name == "input"); + REQUIRE(commands.commands.back().bindings[1].name == "weight"); + const size_t output_binding = codebook_kernel ? 3 : 2; + if (codebook_kernel) { + REQUIRE(string_contains(commands.commands.back().bindings[2].name, "grid")); + } + REQUIRE(commands.commands.back().bindings[output_binding].name == "output"); + if (vector_kernel && !q6_vector_kernel) { + REQUIRE(commands.commands.back().bindings[output_binding + 1].name == "next_output"); + } +} + +static void run_iq1_codebook_matmul_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + auto check_single = [&](ggml_type type, int64_t token_count, const char * expected_kernel) { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, type, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, token_count); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + REQUIRE(kernel_name_for_id(scheduler.plan().dispatches[0].kernel.kernel_id) == expected_kernel); + REQUIRE(scheduler.plan().transients.size() == 1); + REQUIRE(scheduler.plan().transients[0].name == "common.iq1.grid.v1"); + REQUIRE(scheduler.plan().transients[0].size == 16384); + REQUIRE(scheduler.plan().constant_initializations.size() == 1); + REQUIRE(scheduler.plan().constant_initializations[0].name == "common.iq1.grid.v1"); + REQUIRE(scheduler.plan().constant_initializations[0].data.size() == 16384); + const std::array expected_grid_1542 = { 0, 1, 0, 0xff, 1, 0, 0xff, 1 }; + REQUIRE(std::equal(expected_grid_1542.begin(), expected_grid_1542.end(), + scheduler.plan().constant_initializations[0].data.begin() + 1542 * 8)); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 1); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands[0].bindings.size() == (token_count == 1 ? 5 : 4)); + REQUIRE(commands.commands[0].bindings[0].name == "input"); + REQUIRE(commands.commands[0].bindings[1].name == "weight"); + REQUIRE(commands.commands[0].bindings[2].name == "iq1_grid"); + REQUIRE(commands.commands[0].bindings[3].name == "output"); + }; + + check_single(GGML_TYPE_IQ1_S, 1, "loom_libs:ggml_mul_mat_vector_iq1_s_f32_f32"); + check_single(GGML_TYPE_IQ1_S, 64, "loom_libs:ggml_mul_mat_tiled_input_f32_iq1_s_publish_f32"); + check_single(GGML_TYPE_IQ1_M, 1, "loom_libs:ggml_mul_mat_vector_iq1_m_f32_f32"); + check_single(GGML_TYPE_IQ1_M, 64, "loom_libs:ggml_mul_mat_tiled_input_f32_iq1_m_publish_f32"); + + for (const auto [type, disable_env] : + { std::pair{ GGML_TYPE_IQ1_S, "GGML_HRX_DISABLE_IQ1_S_CODEBOOK_MATMUL" }, + std::pair{ GGML_TYPE_IQ1_M, "GGML_HRX_DISABLE_IQ1_M_CODEBOOK_MATMUL" } }) { + EnvironmentVariableGuard guard(disable_env); + REQUIRE(setenv(disable_env, "1", 1) == 0); + ggml_tensor * weight = ggml_new_tensor_2d(ctx, type, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + REQUIRE(!matmul_graph_is_supported(ctx, output)); + } + + ggml_tensor * iq1_s_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ1_S, 2048, 128); + ggml_tensor * iq1_m_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ1_M, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + ggml_tensor * iq1_s_output = ggml_mul_mat(ctx, iq1_s_weight, input); + ggml_tensor * iq1_m_output = ggml_mul_mat(ctx, iq1_m_weight, input); + REQUIRE(iq1_s_output != nullptr); + REQUIRE(iq1_m_output != nullptr); + ggml_cgraph * graph = ggml_new_graph_custom(ctx, GGML_DEFAULT_GRAPH_SIZE, false); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, iq1_s_output); + ggml_build_forward_expand(graph, iq1_m_output); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 2); + REQUIRE(scheduler.plan().transients.size() == 1); + REQUIRE(scheduler.plan().transients[0].name == "common.iq1.grid.v1"); + REQUIRE(scheduler.plan().constant_initializations.size() == 1); + REQUIRE(scheduler.plan().dispatches[0].bindings[2].value == scheduler.plan().transients[0].value); + REQUIRE(scheduler.plan().dispatches[1].bindings[2].value == scheduler.plan().transients[0].value); + + ggml_free(ctx); +} + +static void schedule_fused_matmul_unary_command(ggml_context * ctx, + ggml_tensor * output, + const char * expected_kernel_name, + ggml::hrx::UnaryKind expected_unary_op, + int64_t expected_token_count, + int64_t expected_input_size, + int64_t expected_output_size) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 2); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_MUL_MAT); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + const ggml_type weight_type = output->src[0]->src[0]->type; + const bool q8_input = expected_token_count <= 5 && expected_output_size % 64 == 0 && + (weight_type == GGML_TYPE_Q4_K || weight_type == GGML_TYPE_Q6_K); + const size_t command_count = q8_input ? 2 : 1; + REQUIRE(scheduler.plan().dispatches.size() == command_count); + if (q8_input) { + REQUIRE(kernel_name_for_id(scheduler.plan().dispatches.front().kernel.kernel_id) == + "qwen3_moe:ggml_quantize_q8_1_x4_f32"); + } + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.back(); + const std::string kernel_name = kernel_name_for_id(dispatch.kernel.kernel_id); + REQUIRE(kernel_name == expected_kernel_name); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == expected_token_count); + REQUIRE(dispatch.bindings.size() == 3); + require_compile_parameter(dispatch, "ggml.workload.token_capacity", std::to_string(expected_token_count)); + require_compile_parameter(dispatch, "ggml.mul_mat.input_size", std::to_string(expected_input_size)); + require_compile_parameter(dispatch, "ggml.mul_mat.output_size", std::to_string(expected_output_size)); + require_compile_parameter(dispatch, "ggml.mul_mat.output_accumulation", "0"); + require_compile_parameter(dispatch, "ggml.mul_mat.output_unary_op", + std::to_string(ggml::hrx::unary_kind_config_value(expected_unary_op))); + const ggml::hrx::Value * weight = imported.graph.values().find(dispatch.bindings[1].value); + REQUIRE(weight != nullptr); + require_compile_parameter(dispatch, "ggml.mul_mat.weight_format", + std::to_string(q8_input ? (weight_type == GGML_TYPE_Q4_K ? 44 : 46) : + matmul_weight_format_config(weight->type))); + REQUIRE(dispatch.bindings[2].value == imported.graph.nodes()[1].output); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == command_count); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.back().bindings.size() == 3); + REQUIRE(commands.commands.back().bindings[0].name == "input"); + REQUIRE(commands.commands.back().bindings[1].name == "weight"); + REQUIRE(commands.commands.back().bindings[2].name == "output"); +} + +static void schedule_fused_matmul_swiglu_command( + ggml_context * ctx, + ggml_tensor * output, + ggml_type expected_gate_type, + ggml_type expected_up_type, + int64_t expected_token_count, + int64_t expected_input_size, + int64_t expected_output_size, + ggml::hrx::BinaryKind expected_binary_op = ggml::hrx::BinaryKind::SwiGLU, + ggml_op expected_output_op = GGML_OP_GLU) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 3); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_MUL_MAT); + REQUIRE(imported.graph.nodes()[1].op == GGML_OP_MUL_MAT); + REQUIRE(imported.graph.nodes()[2].op == expected_output_op); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + const bool low_token = expected_token_count <= 5 && expected_input_size % 256 == 0; + const bool q8_input = low_token && expected_gate_type == GGML_TYPE_Q4_K && expected_up_type == GGML_TYPE_Q4_K && + expected_output_size % 64 == 0; + const size_t command_count = q8_input ? 2 : 1; + REQUIRE(scheduler.plan().dispatches.size() == command_count); + if (q8_input) { + REQUIRE(kernel_name_for_id(scheduler.plan().dispatches.front().kernel.kernel_id) == + "qwen3_moe:ggml_quantize_q8_1_x4_f32"); + } + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.back(); + const std::string kernel_name = kernel_name_for_id(dispatch.kernel.kernel_id); + REQUIRE(kernel_name == (q8_input ? "loom_libs:ggml_mul_mat_swiglu_q4_q8_1_x4_lowtoken_dot" : + low_token ? "loom_libs:ggml_mul_mat_swiglu_f32_f32_lowtoken_dot" : + "loom_libs:ggml_mul_mat_tiled_pair_input_f32_binary_publish_f32")); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == expected_token_count); + REQUIRE(dispatch.bindings.size() == 4); + require_compile_parameter(dispatch, "ggml.workload.token_capacity", std::to_string(expected_token_count)); + require_compile_parameter(dispatch, "ggml.mul_mat_swiglu.input_size", std::to_string(expected_input_size)); + require_compile_parameter(dispatch, "ggml.mul_mat_swiglu.output_size", std::to_string(expected_output_size)); + require_compile_parameter(dispatch, "ggml.mul_mat_swiglu.gate_weight_format", + std::to_string(q8_input ? 44 : matmul_weight_format_config(expected_gate_type))); + require_compile_parameter(dispatch, "ggml.mul_mat_swiglu.up_weight_format", + std::to_string(q8_input ? 44 : matmul_weight_format_config(expected_up_type))); + require_compile_parameter(dispatch, "ggml.mul_mat_swiglu.op", + std::to_string(ggml::hrx::binary_kind_config_value(expected_binary_op))); + REQUIRE(dispatch.bindings[3].value == imported.graph.nodes()[2].output); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == command_count); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.back().bindings.size() == 4); + REQUIRE(commands.commands.back().bindings[0].name == "input"); + REQUIRE(commands.commands.back().bindings[1].name == "gate_weight"); + REQUIRE(commands.commands.back().bindings[2].name == "up_weight"); + REQUIRE(commands.commands.back().bindings[3].name == "output"); +} + +static void schedule_fused_matmul_swiglu_postops_command(ggml_context * ctx, + ggml_tensor * output, + const char * expected_kernel_name, + size_t expected_node_count, + const std::vector & expected_ops, + const std::vector expected_binding_names) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == expected_node_count); + REQUIRE(expected_ops.size() == expected_node_count); + for (size_t i = 0; i < expected_ops.size(); ++i) { + REQUIRE(imported.graph.nodes()[i].op == expected_ops[i]); + } + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == expected_kernel_name); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == 8); + REQUIRE(dispatch.bindings.size() == expected_binding_names.size()); + require_compile_parameter(dispatch, "ggml.workload.token_capacity", "8"); + require_compile_parameter(dispatch, "ggml.mul_mat_swiglu.input_size", "640"); + require_compile_parameter(dispatch, "ggml.mul_mat_swiglu.output_size", "256"); + require_compile_parameter(dispatch, "ggml.mul_mat_swiglu.gate_weight_format", + std::to_string(matmul_weight_format_config(GGML_TYPE_IQ4_NL))); + require_compile_parameter(dispatch, "ggml.mul_mat_swiglu.up_weight_format", + std::to_string(matmul_weight_format_config(GGML_TYPE_IQ4_NL))); + require_compile_parameter(dispatch, "ggml.mul_mat_swiglu.op", + std::to_string(ggml::hrx::binary_kind_config_value(ggml::hrx::BinaryKind::SwiGLU))); + REQUIRE(dispatch.bindings.back().value == imported.graph.nodes().back().output); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 1); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.front().bindings.size() == expected_binding_names.size()); + for (size_t i = 0; i < expected_binding_names.size(); ++i) { + REQUIRE(commands.commands.front().bindings[i].name == expected_binding_names[i]); + } +} + +static void schedule_packed_glu_command(ggml_context * ctx, + ggml_tensor * output, + ggml::hrx::BinaryKind expected_binary_op, + int64_t expected_hidden_size, + int64_t expected_token_count, + bool expected_swapped) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 1); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_GLU); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + const std::string kernel_name = kernel_name_for_id(dispatch.kernel.kernel_id); + REQUIRE(kernel_name == "loom_libs:ggml_binary_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("element_count") == expected_hidden_size * expected_token_count); + require_compile_parameter(dispatch, "ggml.binary_f32.op", + std::to_string(ggml::hrx::binary_kind_config_value(expected_binary_op))); + require_compile_parameter(dispatch, "ggml.binary_f32.ne0", std::to_string(expected_hidden_size)); + require_compile_parameter(dispatch, "ggml.binary_f32.ne1", std::to_string(expected_token_count)); + REQUIRE(dispatch.bindings.size() == 3); + + const size_t half_offset = static_cast(expected_hidden_size) * sizeof(float); + REQUIRE(dispatch.bindings[0].offset == (expected_swapped ? half_offset : 0)); + REQUIRE(dispatch.bindings[1].offset == (expected_swapped ? 0 : half_offset)); + REQUIRE(dispatch.bindings[2].value == imported.graph.nodes()[0].output); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 1); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.front().bindings.size() == 3); + REQUIRE(commands.commands.front().bindings[0].name == "lhs"); + REQUIRE(commands.commands.front().bindings[1].name == "rhs"); + REQUIRE(commands.commands.front().bindings[2].name == "output"); +} + +static void schedule_packed_matmul_glu_command(ggml_context * ctx, + ggml_tensor * output, + ggml_type expected_weight_type, + int64_t expected_token_count, + int64_t expected_input_size, + int64_t expected_output_size, + ggml::hrx::BinaryKind expected_binary_op, + bool expected_swapped) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 2); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_MUL_MAT); + REQUIRE(imported.graph.nodes()[1].op == GGML_OP_GLU); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + const std::string kernel_name = kernel_name_for_id(dispatch.kernel.kernel_id); + REQUIRE(kernel_name == "loom_libs:ggml_mul_mat_tiled_pair_input_f32_binary_publish_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == expected_token_count); + REQUIRE(dispatch.bindings.size() == 4); + require_compile_parameter(dispatch, "ggml.workload.token_capacity", std::to_string(expected_token_count)); + require_compile_parameter(dispatch, "ggml.mul_mat_swiglu.input_size", std::to_string(expected_input_size)); + require_compile_parameter(dispatch, "ggml.mul_mat_swiglu.output_size", std::to_string(expected_output_size)); + require_compile_parameter(dispatch, "ggml.mul_mat_swiglu.gate_weight_format", + std::to_string(matmul_weight_format_config(expected_weight_type))); + require_compile_parameter(dispatch, "ggml.mul_mat_swiglu.up_weight_format", + std::to_string(matmul_weight_format_config(expected_weight_type))); + require_compile_parameter(dispatch, "ggml.mul_mat_swiglu.op", + std::to_string(ggml::hrx::binary_kind_config_value(expected_binary_op))); + + const size_t half_offset = + ggml_row_size(expected_weight_type, expected_input_size) * static_cast(expected_output_size); + REQUIRE(dispatch.bindings[1].offset == (expected_swapped ? half_offset : 0)); + REQUIRE(dispatch.bindings[2].offset == (expected_swapped ? 0 : half_offset)); + REQUIRE(dispatch.bindings[3].value == imported.graph.nodes()[1].output); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 1); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.front().bindings.size() == 4); + REQUIRE(commands.commands.front().bindings[0].name == "input"); + REQUIRE(commands.commands.front().bindings[1].name == "gate_weight"); + REQUIRE(commands.commands.front().bindings[2].name == "up_weight"); + REQUIRE(commands.commands.front().bindings[3].name == "output"); +} + +static void schedule_fused_matmul_postops_command(ggml_context * ctx, + ggml_tensor * output, + const char * expected_kernel_name, + ggml_type expected_weight_type, + int64_t expected_token_count, + int64_t expected_input_size, + int64_t expected_output_size, + size_t expected_node_count, + const std::vector & expected_ops, + const std::vector expected_binding_names, + bool expected_rmsnorm) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == expected_node_count); + REQUIRE(expected_ops.size() == expected_node_count); + for (size_t i = 0; i < expected_ops.size(); ++i) { + REQUIRE(imported.graph.nodes()[i].op == expected_ops[i]); + } + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == (expected_rmsnorm ? 2 : 1)); + REQUIRE(scheduler.plan().completion_counter_requests.empty()); + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + const std::string kernel_name = kernel_name_for_id(dispatch.kernel.kernel_id); + REQUIRE(kernel_name == expected_kernel_name); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == expected_token_count); + REQUIRE(dispatch.bindings.size() == expected_binding_names.size()); + require_compile_parameter(dispatch, "ggml.workload.token_capacity", std::to_string(expected_token_count)); + require_compile_parameter(dispatch, "ggml.mul_mat_postops.input_size", std::to_string(expected_input_size)); + require_compile_parameter(dispatch, "ggml.mul_mat_postops.output_size", std::to_string(expected_output_size)); + require_compile_parameter(dispatch, "ggml.mul_mat_postops.weight_format", + std::to_string(matmul_weight_format_config(expected_weight_type))); + REQUIRE(dispatch.kernel.compile_parameters.count("ggml.mul_mat_postops.rms_epsilon") == 0); + if (expected_rmsnorm) { + const ggml::hrx::Dispatch & rmsnorm_dispatch = scheduler.plan().dispatches[1]; + REQUIRE(kernel_name_for_id(rmsnorm_dispatch.kernel.kernel_id) == "loom_libs:ggml_rmsnorm_binary_f32"); + } + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == (expected_rmsnorm ? 2 : 1)); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.front().bindings.size() == expected_binding_names.size()); + for (size_t i = 0; i < expected_binding_names.size(); ++i) { + REQUIRE(commands.commands.front().bindings[i].name == expected_binding_names[i]); + } +} + +static void schedule_fused_llama_attention_matmul_command(ggml_context * ctx, + ggml_tensor * output, + const char * expected_kernel_name, + ggml_type expected_weight_type, + int64_t expected_token_count, + int64_t expected_input_size, + int64_t expected_output_size, + int64_t expected_head_size, + int64_t expected_head_count, + int64_t expected_cache_row_count, + int64_t expected_cache_format, + size_t expected_node_count, + const std::vector & expected_ops, + const std::vector expected_binding_names) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == expected_node_count); + REQUIRE(expected_ops.size() == expected_node_count); + for (size_t i = 0; i < expected_ops.size(); ++i) { + REQUIRE(imported.graph.nodes()[i].op == expected_ops[i]); + } + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + const std::string kernel_name = kernel_name_for_id(dispatch.kernel.kernel_id); + REQUIRE(kernel_name == expected_kernel_name); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == expected_token_count); + REQUIRE(dispatch.bindings.size() == expected_binding_names.size()); + require_compile_parameter(dispatch, "ggml.workload.token_capacity", std::to_string(expected_token_count)); + require_compile_parameter(dispatch, "llm.attention_qkv.input_size", std::to_string(expected_input_size)); + require_compile_parameter(dispatch, "llm.attention_qkv.output_size", std::to_string(expected_output_size)); + require_compile_parameter(dispatch, "llm.attention_qkv.weight_format", + std::to_string(matmul_weight_format_config(expected_weight_type))); + if (expected_head_size > 0) { + require_compile_parameter(dispatch, "llm.attention_qkv.head_size", std::to_string(expected_head_size)); + require_compile_parameter(dispatch, "llm.attention_qkv.head_count", std::to_string(expected_head_count)); + } + if (expected_cache_row_count > 0) { + require_compile_parameter(dispatch, "llm.attention_qkv.cache_row_count", + std::to_string(expected_cache_row_count)); + require_compile_parameter(dispatch, "llm.attention_qkv.cache_output_format", + std::to_string(expected_cache_format)); + } + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 1); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.front().bindings.size() == expected_binding_names.size()); + for (size_t i = 0; i < expected_binding_names.size(); ++i) { + REQUIRE(commands.commands.front().bindings[i].name == expected_binding_names[i]); + } +} + +static void require_matmul_swiglu_falls_back(ggml_context * ctx, ggml_tensor * output) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() > 1); + for (const ggml::hrx::Dispatch & dispatch : scheduler.plan().dispatches) { + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) != "loom_libs:ggml_mul_mat_swiglu_f32_f32_wmma"); + } +} + +static void require_matmul_swiglu_falls_back_to_compilable_plan(ggml_context * ctx, ggml_tensor * output) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() > 1); + for (const ggml::hrx::Dispatch & dispatch : scheduler.plan().dispatches) { + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) != "loom_libs:ggml_mul_mat_swiglu_f32_f32_wmma"); + } + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); +} + +static void require_matmul_root_matches_identity_single(ggml_context * ctx, ggml_tensor * output) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(!imported.graph.nodes().empty()); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_MUL_MAT); + + const std::vector covered_nodes(imported.graph.nodes().size(), false); + const ggml::hrx::CommandPlan plan; + const ggml::hrx::DispatchMatchContext context = { + imported.graph, + &imported.graph.nodes()[0], + 0, + covered_nodes, + plan, + ggml::hrx::ValueId(static_cast(imported.graph.values().size())), + }; + ggml::hrx::DispatchMatch match; + REQUIRE(test_dispatch_registry().match(context, match)); + REQUIRE(match.covered_nodes.size() == 1); + REQUIRE(match.covered_nodes.front() == 0); + const ggml::hrx::Value * weight = imported.graph.values().find(imported.graph.nodes()[0].inputs[0]); + const ggml::hrx::Value * input = imported.graph.values().find(imported.graph.nodes()[0].inputs[1]); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + const bool q8_input = input->ne[1] <= 5 && weight->ne[1] % 64 == 0 && + (weight->type == GGML_TYPE_Q4_K || weight->type == GGML_TYPE_Q6_K); + REQUIRE(match.dispatches.size() == (q8_input ? 2 : 1)); + require_compile_parameter(match.dispatches.back(), "ggml.mul_mat.output_unary_op", + std::to_string(ggml::hrx::unary_kind_config_value(ggml::hrx::UnaryKind::Identity))); +} + +static bool matmul_graph_is_supported(ggml_context * ctx, ggml_tensor * output) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + return ggml::hrx::DispatchScheduler::can_schedule_graph(imported.graph, test_dispatch_target()); +} + +static void schedule_qwen_terminal_q6k_q8_command(int64_t token_count) { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, token_count); + ggml_tensor * norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 2048); + ggml_tensor * vocab_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, 2048, 151936); + REQUIRE(input != nullptr); + REQUIRE(norm_weight != nullptr); + REQUIRE(vocab_weight != nullptr); + ggml_tensor * rms = ggml_rms_norm(ctx, input, 0.000001f); + ggml_tensor * normalized = ggml_mul(ctx, rms, norm_weight); + ggml_tensor * logits = ggml_mul_mat(ctx, vocab_weight, normalized); + REQUIRE(rms != nullptr); + REQUIRE(normalized != nullptr); + REQUIRE(logits != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, logits); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 3); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 2); + REQUIRE(scheduler.plan().transients.size() == 1); + + const ggml::hrx::Dispatch & rms_dispatch = scheduler.plan().dispatches[0]; + REQUIRE(kernel_name_for_id(rms_dispatch.kernel.kernel_id) == "qwen3_moe:qwen3_moe_rmsnorm_f32_quantize_q8_1_x4"); + REQUIRE(rms_dispatch.kernel.integer_parameters.at("token_count") == token_count); + require_compile_parameter(rms_dispatch, "qwen3_moe.model.hidden_size", "2048"); + require_compile_parameter(rms_dispatch, "qwen3_moe.workload.token_capacity", std::to_string(token_count)); + REQUIRE(rms_dispatch.bindings.size() == 4); + + const ggml::hrx::CommandPlanTransient & q8_transient = scheduler.plan().transients.front(); + REQUIRE(q8_transient.size == qwen_q8_1_x4_size(token_count, 2048)); + REQUIRE(rms_dispatch.bindings[3].value == q8_transient.value); + REQUIRE(rms_dispatch.bindings[3].length == q8_transient.size); + + const ggml::hrx::Dispatch & vocab_dispatch = scheduler.plan().dispatches[1]; + REQUIRE(kernel_name_for_id(vocab_dispatch.kernel.kernel_id) == "qwen3_moe:ggml_linear_q6k_q8_1_x4"); + REQUIRE(vocab_dispatch.kernel.integer_parameters.at("token_count") == token_count); + REQUIRE(vocab_dispatch.kernel.integer_parameters.at("input_size") == 2048); + REQUIRE(vocab_dispatch.kernel.integer_parameters.at("output_size") == 151936); + require_compile_parameter(vocab_dispatch, "ggml.linear_q6k_q8_1_x4.token_capacity", std::to_string(token_count)); + require_compile_parameter(vocab_dispatch, "ggml.linear_q6k_q8_1_x4.output_capacity", "151936"); + REQUIRE(vocab_dispatch.bindings.size() == 3); + REQUIRE(vocab_dispatch.bindings[0].value == q8_transient.value); + REQUIRE(vocab_dispatch.bindings[0].length == q8_transient.size); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 2); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands[0].bindings.size() == 4); + REQUIRE(commands.commands[0].bindings[3].name == "q8_output"); + REQUIRE(commands.commands[0].bindings[3].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands[1].bindings.size() == 3); + REQUIRE(commands.commands[1].bindings[0].name == "q8_input"); + REQUIRE(commands.commands[1].bindings[0].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands[1].bindings[0].value == q8_transient.value); + + ggml_free(ctx); +} + +static void run_qwen_decode_rmsnorm_publication_checks() { + const struct { + bool projection; + bool side_use; + bool q8; + } cases[] = { + { true, false, true }, + { true, true, true }, + { false, false, false }, + }; + for (const auto & test : cases) { + ggml_init_params params = {}; + params.mem_size = 8 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + ggml_tensor * norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 2048); + ggml_tensor * normalized = ggml_mul(ctx, ggml_rms_norm(ctx, input, 0.000001f), norm_weight); + ggml_tensor * projection = nullptr; + if (test.projection) { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 512); + projection = ggml_mul_mat(ctx, weight, normalized); + } + ggml_tensor * side = test.side_use ? ggml_scale(ctx, normalized, 0.5f) : nullptr; + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, projection != nullptr ? projection : normalized); + if (side != nullptr) { + ggml_build_forward_expand(graph, side); + } + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const ggml::hrx::CommandPlan & plan = scheduler.plan(); + REQUIRE(plan.valid()); + + const auto producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == + "qwen3_moe:qwen3_moe_rmsnorm_f32_quantize_q8_1_x4"; + }); + REQUIRE((producer != plan.dispatches.end()) == test.q8); + const ggml::hrx::Value * output = imported.graph.values().find_tensor(normalized); + REQUIRE(output != nullptr); + const auto * q8 = plan.metadata.find_alternate_value( + output->id, GGML_TYPE_Q8_1, qwen_q8_1_x4_size(1, 2048)); + REQUIRE((q8 != nullptr) == test.q8); + if (test.q8) { + REQUIRE(producer->bindings.size() == 4); + REQUIRE(producer->bindings.back().value == q8->alternate_value); + REQUIRE(std::any_of(std::next(producer), plan.dispatches.end(), [&](const auto & dispatch) { + return !dispatch.bindings.empty() && dispatch.bindings.front().value == q8->alternate_value; + })); + } + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + ggml_free(ctx); + } +} + +static void run_qwen_decode_publication_consumer_qualification_checks() { + enum class ConsumerCase { + Dense, + UnsupportedWeight, + WrongOperand, + IncompatibleGeometry, + Routed, + RoutedQ6, + RoutedShape, + RoutedRoute, + AliasedDense, + }; + const struct { + ConsumerCase consumer; + bool q8; + } cases[] = { + { ConsumerCase::Dense, true }, + { ConsumerCase::UnsupportedWeight, false }, + { ConsumerCase::WrongOperand, false }, + { ConsumerCase::IncompatibleGeometry, false }, + { ConsumerCase::Routed, true }, + { ConsumerCase::RoutedQ6, true }, + { ConsumerCase::RoutedShape, false }, + { ConsumerCase::RoutedRoute, false }, + { ConsumerCase::AliasedDense, true }, + }; + + for (const auto & test : cases) { + ggml_init_params params = {}; + params.mem_size = 32 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + ggml_tensor * norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 2048); + ggml_tensor * rms = ggml_rms_norm(ctx, input, 0.000001f); + ggml_tensor * normalized = ggml_mul(ctx, rms, norm_weight); + REQUIRE(rms != nullptr); + REQUIRE(normalized != nullptr); + + ggml_tensor * consumer_input = normalized; + if (test.consumer == ConsumerCase::AliasedDense) { + consumer_input = ggml_reshape_2d(ctx, normalized, 2048, 1); + REQUIRE(consumer_input != nullptr); + } + + ggml_tensor * terminal = nullptr; + if (test.consumer == ConsumerCase::Dense || test.consumer == ConsumerCase::AliasedDense) { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 512); + terminal = ggml_mul_mat(ctx, weight, consumer_input); + } else if (test.consumer == ConsumerCase::UnsupportedWeight) { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_BF16, 2048, 512); + terminal = ggml_mul_mat(ctx, weight, normalized); + } else if (test.consumer == ConsumerCase::WrongOperand) { + ggml_tensor * activation = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + terminal = ggml_mul_mat(ctx, normalized, activation); + } else if (test.consumer == ConsumerCase::IncompatibleGeometry) { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 63); + terminal = ggml_mul_mat(ctx, weight, normalized); + } else { + const int64_t route_count = test.consumer == ConsumerCase::RoutedShape ? 4 : 8; + ggml_tensor * route_ids = ggml_new_tensor_2d(ctx, GGML_TYPE_I32, route_count, 1); + ggml_tensor * gate_weight = ggml_new_tensor_3d(ctx, GGML_TYPE_Q4_K, 2048, 768, 128); + ggml_tensor * up_weight = ggml_new_tensor_3d(ctx, GGML_TYPE_Q4_K, 2048, 768, 128); + const ggml_type down_weight_type = test.consumer == ConsumerCase::RoutedQ6 ? GGML_TYPE_Q6_K : GGML_TYPE_Q4_K; + ggml_tensor * down_weight = ggml_new_tensor_3d(ctx, down_weight_type, 768, 2048, 128); + ggml_tensor * gate = ggml_mul_mat_id(ctx, gate_weight, normalized, route_ids); + ggml_tensor * up = ggml_mul_mat_id(ctx, up_weight, normalized, route_ids); + ggml_tensor * glu = ggml_glu_split(ctx, gate, up, GGML_GLU_OP_SWIGLU); + ggml_tensor * down_route_ids = test.consumer == ConsumerCase::RoutedRoute ? + ggml_new_tensor_2d(ctx, GGML_TYPE_I32, route_count, 1) : + route_ids; + ggml_tensor * down = ggml_mul_mat_id(ctx, down_weight, glu, down_route_ids); + ggml_tensor * route_weights = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 1, route_count, 1); + terminal = ggml_mul(ctx, down, route_weights); + } + REQUIRE(terminal != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, terminal); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + const ggml::hrx::Value * rms_value = imported.graph.values().find_tensor(rms); + const ggml::hrx::Value * normalized_value = imported.graph.values().find_tensor(normalized); + REQUIRE(rms_value != nullptr); + REQUIRE(normalized_value != nullptr); + const ggml::hrx::GraphNode * rms_node = imported.graph.index().producer(rms_value->id); + size_t rms_index = 0; + REQUIRE(rms_node != nullptr); + REQUIRE(imported.graph.index().node_index(rms_node, rms_index)); + + const ggml::hrx::CommandPlan empty_plan; + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::DispatchMatch match; + REQUIRE(match_dispatch_at_index(imported.graph, empty_plan, covered_nodes, rms_index, match)); + const auto * q8 = match.metadata.find_alternate_value( + normalized_value->id, GGML_TYPE_Q8_1, qwen_q8_1_x4_size(1, 2048)); + REQUIRE((q8 != nullptr) == test.q8); + if (test.q8) { + REQUIRE(kernel_name_for_id(match.dispatches.back().kernel.kernel_id) == + "qwen3_moe:qwen3_moe_rmsnorm_f32_quantize_q8_1_x4"); + } + ggml_free(ctx); + } +} + +static void schedule_qwen_terminal_q6k_vector_command() { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + ggml_tensor * vocab_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, 2048, 151936); + REQUIRE(input != nullptr); + REQUIRE(vocab_weight != nullptr); + ggml_tensor * logits = ggml_mul_mat(ctx, vocab_weight, input); + REQUIRE(logits != nullptr); + + schedule_single_matmul_command(ctx, logits, "loom_libs:ggml_mul_mat_vector_q6_f32_f32", 1, 2048, 151936); + ggml_free(ctx); +} + +static void schedule_get_rows_q8_1_alternate_command(ggml_type embedding_weight_type) { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * embedding_weight = ggml_new_tensor_2d(ctx, embedding_weight_type, 2048, 151936); + ggml_tensor * token_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 1); + ggml_tensor * vocab_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, 2048, 151936); + REQUIRE(embedding_weight != nullptr); + REQUIRE(token_ids != nullptr); + REQUIRE(vocab_weight != nullptr); + + ggml_tensor * embedding = ggml_get_rows(ctx, embedding_weight, token_ids); + ggml_tensor * logits = ggml_mul_mat(ctx, vocab_weight, embedding); + REQUIRE(embedding != nullptr); + REQUIRE(logits != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, logits); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 2); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_GET_ROWS); + REQUIRE(imported.graph.nodes()[1].op == GGML_OP_MUL_MAT); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 2); + REQUIRE(scheduler.plan().transients.size() == 1); + + const size_t q8_byte_count = qwen_q8_1_x4_size(1, 2048); + + const ggml::hrx::Dispatch & get_rows_dispatch = scheduler.plan().dispatches[0]; + REQUIRE(kernel_name_for_id(get_rows_dispatch.kernel.kernel_id) == "loom_libs:ggml_get_rows_f32_next"); + REQUIRE(get_rows_dispatch.kernel.integer_parameters.at("token_count") == 1); + REQUIRE(get_rows_dispatch.kernel.integer_parameters.at("row_count") == 151936); + REQUIRE(get_rows_dispatch.kernel.integer_parameters.at("hidden_size") == 2048); + require_compile_parameter(get_rows_dispatch, "ggml.get_rows_f32.weight_format", + std::to_string(matmul_weight_format_config(embedding_weight_type))); + require_compile_parameter(get_rows_dispatch, "ggml.get_rows_f32.next_format", "81"); + REQUIRE(get_rows_dispatch.bindings.size() == 4); + + const ggml::hrx::CommandPlanTransient & q8_transient = scheduler.plan().transients.front(); + REQUIRE(q8_transient.size == q8_byte_count); + REQUIRE(get_rows_dispatch.bindings[3].value == q8_transient.value); + REQUIRE(get_rows_dispatch.bindings[3].length == q8_transient.size); + + const ggml::hrx::CommandPlanAlternateValue * alternate = + scheduler.plan().metadata.find_alternate_value(imported.graph.nodes()[0].output, GGML_TYPE_Q8_1, q8_byte_count); + REQUIRE(alternate != nullptr); + REQUIRE(alternate->alternate_value == q8_transient.value); + + const ggml::hrx::Dispatch & vocab_dispatch = scheduler.plan().dispatches[1]; + const std::string vocab_kernel_name = kernel_name_for_id(vocab_dispatch.kernel.kernel_id); + REQUIRE(vocab_dispatch.bindings.size() == 3); + if (vocab_kernel_name == "qwen3_moe:ggml_linear_q6k_q8_1_x4") { + REQUIRE(vocab_dispatch.bindings[0].value == alternate->alternate_value); + REQUIRE(vocab_dispatch.bindings[0].length == alternate->byte_count); + } else { + REQUIRE(vocab_kernel_name == "loom_libs:ggml_mul_mat_vector_q6_f32_f32"); + REQUIRE(vocab_dispatch.bindings[0].value == imported.graph.nodes()[0].output); + require_compile_parameter(vocab_dispatch, "ggml.matmul.vector.input_size", "2048"); + require_compile_parameter(vocab_dispatch, "ggml.matmul.vector.output_size", "151936"); + require_compile_parameter(vocab_dispatch, "ggml.matmul.vector.q6_storage", "0"); + } + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 2); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands[0].bindings.size() == 4); + REQUIRE(commands.commands[0].bindings[3].name == "next_output"); + REQUIRE(commands.commands[0].bindings[3].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands[1].bindings.size() == 3); + REQUIRE(commands.commands[1].bindings[0].name == + (vocab_kernel_name == "qwen3_moe:ggml_linear_q6k_q8_1_x4" ? "q8_input" : "input")); + + ggml_free(ctx); +} + +static bool graph_is_supported(ggml_context * ctx, ggml_tensor * output) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + return ggml::hrx::DispatchScheduler::can_schedule_graph(imported.graph, test_dispatch_target()); +} + +struct ManualQwenRouterTop8Graph { + ggml::hrx::Graph graph; + ggml::hrx::ValueId route_ids; +}; + +static ggml::hrx::ValueId add_manual_graph_tensor(ggml::hrx::Graph & graph, + ggml_tensor * tensor, + ggml::hrx::ValueKind kind) { + REQUIRE(tensor != nullptr); + return graph.values().get_or_add_tensor_value(tensor, kind); +} + +static ManualQwenRouterTop8Graph build_manual_qwen_router_top8_graph(ggml_context * ctx, + int64_t expert_count, + int64_t route_count, + int64_t token_count) { + ManualQwenRouterTop8Graph manual; + ggml::hrx::Graph & graph = manual.graph; + + const ggml::hrx::ValueId logits = add_manual_graph_tensor( + graph, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, expert_count, token_count), ggml::hrx::ValueKind::External); + const ggml::hrx::ValueId probs = add_manual_graph_tensor( + graph, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, expert_count, token_count), ggml::hrx::ValueKind::Transient); + const ggml::hrx::ValueId probs_reshaped = add_manual_graph_tensor( + graph, ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 1, expert_count, token_count), ggml::hrx::ValueKind::Transient); + const ggml::hrx::ValueId argsort = add_manual_graph_tensor( + graph, ggml_new_tensor_2d(ctx, GGML_TYPE_I32, expert_count, token_count), ggml::hrx::ValueKind::Transient); + manual.route_ids = add_manual_graph_tensor(graph, ggml_new_tensor_2d(ctx, GGML_TYPE_I32, route_count, token_count), + ggml::hrx::ValueKind::Transient); + const ggml::hrx::ValueId selected = add_manual_graph_tensor( + graph, ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 1, route_count, token_count), ggml::hrx::ValueKind::Transient); + const ggml::hrx::ValueId weights_flat = add_manual_graph_tensor( + graph, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, route_count, token_count), ggml::hrx::ValueKind::Transient); + const ggml::hrx::ValueId sum = add_manual_graph_tensor( + graph, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 1, token_count), ggml::hrx::ValueKind::Transient); + const ggml::hrx::ValueId clamped = add_manual_graph_tensor( + graph, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 1, token_count), ggml::hrx::ValueKind::Transient); + const ggml::hrx::ValueId normalized = add_manual_graph_tensor( + graph, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, route_count, token_count), ggml::hrx::ValueKind::Transient); + const ggml::hrx::ValueId route_weights = add_manual_graph_tensor( + graph, ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 1, route_count, token_count), ggml::hrx::ValueKind::External); + + ggml::hrx::GraphNode & softmax = graph.add_node(GGML_OP_SOFT_MAX, probs, { logits }); + softmax.params = ggml::hrx::SoftMaxParams{ 1.0f, 0.0f }; + graph.add_node(GGML_OP_RESHAPE, probs_reshaped, { probs }); + ggml::hrx::GraphNode & argsort_node = graph.add_node(GGML_OP_ARGSORT, argsort, { probs }); + argsort_node.params = ggml::hrx::ArgsortParams{ GGML_SORT_ORDER_DESC }; + graph.add_node(GGML_OP_VIEW, manual.route_ids, { argsort }); + graph.add_node(GGML_OP_GET_ROWS, selected, { probs_reshaped, manual.route_ids }); + graph.add_node(GGML_OP_RESHAPE, weights_flat, { selected }); + graph.add_node(GGML_OP_SUM_ROWS, sum, { weights_flat }); + ggml::hrx::GraphNode & clamp = graph.add_node(GGML_OP_CLAMP, clamped, { sum }); + clamp.params = ggml::hrx::ClampParams{ 0.00006103515625f, std::numeric_limits::infinity() }; + graph.add_node(GGML_OP_DIV, normalized, { weights_flat, clamped }); + graph.add_node(GGML_OP_RESHAPE, route_weights, { normalized }); + REQUIRE(graph.build_index().success()); + return manual; +} + +static ggml::hrx::Graph build_manual_token_embedding_graph(ggml_tensor * weight, + ggml_tensor * token_ids, + ggml_tensor * output) { + ggml::hrx::Graph graph; + ggml::hrx::ValueId weight_value = graph.values().get_or_add_tensor_value(weight, ggml::hrx::ValueKind::External); + ggml::hrx::ValueId token_ids_value = + graph.values().get_or_add_tensor_value(token_ids, ggml::hrx::ValueKind::External); + ggml::hrx::ValueId output_value = graph.values().get_or_add_tensor_value(output, ggml::hrx::ValueKind::External); + graph.add_node(GGML_OP_GET_ROWS, output_value, { weight_value, token_ids_value }); + REQUIRE(graph.build_index().success()); + return graph; +} + +static bool manual_token_embedding_graph_is_supported(ggml_context * ctx, + ggml_type weight_type, + ggml_type token_ids_type, + ggml_type output_type, + int64_t hidden_size, + int64_t vocabulary_count, + int64_t token_count, + int64_t output_hidden_size = -1, + int64_t output_token_count = -1) { + if (output_hidden_size < 0) { + output_hidden_size = hidden_size; + } + if (output_token_count < 0) { + output_token_count = token_count; + } + ggml_tensor * weight = ggml_new_tensor_2d(ctx, weight_type, hidden_size, vocabulary_count); + ggml_tensor * token_ids = ggml_new_tensor_1d(ctx, token_ids_type, token_count); + ggml_tensor * output = ggml_new_tensor_2d(ctx, output_type, output_hidden_size, output_token_count); + REQUIRE(weight != nullptr); + REQUIRE(token_ids != nullptr); + REQUIRE(output != nullptr); + + ggml::hrx::Graph graph = build_manual_token_embedding_graph(weight, token_ids, output); + return ggml::hrx::DispatchScheduler::can_schedule_graph(graph, test_dispatch_target()); +} + +static void schedule_token_embedding_command(ggml_context * ctx, + ggml_tensor * output, + int64_t expected_token_count, + int64_t expected_vocabulary_count, + int64_t expected_hidden_size, + const char * expected_weight_format) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 1); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_GET_ROWS); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + const std::string kernel_name = kernel_name_for_id(dispatch.kernel.kernel_id); + REQUIRE(kernel_name == "loom_libs:ggml_get_rows_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == expected_token_count); + REQUIRE(dispatch.kernel.integer_parameters.at("row_count") == expected_vocabulary_count); + REQUIRE(dispatch.kernel.integer_parameters.at("hidden_size") == expected_hidden_size); + require_compile_parameter(dispatch, "ggml.get_rows_f32.weight_format", expected_weight_format); + REQUIRE(dispatch.bindings.size() == 3); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 1); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.front().bindings.size() == 3); + REQUIRE(commands.commands.front().bindings[0].name == "token_ids"); + REQUIRE(commands.commands.front().bindings[1].name == "weight"); + REQUIRE(commands.commands.front().bindings[2].name == "output"); +} + +static void run_qwen_token_embedding_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 151936); + ggml_tensor * token_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 1); + REQUIRE(weight != nullptr); + REQUIRE(token_ids != nullptr); + ggml_tensor * output = ggml_get_rows(ctx, weight, token_ids); + REQUIRE(output != nullptr); + schedule_token_embedding_command(ctx, output, 1, 151936, 2048, "4"); + } + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 151936); + ggml_tensor * token_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 13); + REQUIRE(weight != nullptr); + REQUIRE(token_ids != nullptr); + ggml_tensor * output = ggml_get_rows(ctx, weight, token_ids); + REQUIRE(output != nullptr); + schedule_token_embedding_command(ctx, output, 13, 151936, 2048, "4"); + } + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_BF16, 2048, 151936); + ggml_tensor * token_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 5); + REQUIRE(weight != nullptr); + REQUIRE(token_ids != nullptr); + ggml_tensor * output = ggml_get_rows(ctx, weight, token_ids); + REQUIRE(output != nullptr); + schedule_token_embedding_command(ctx, output, 5, 151936, 2048, "30"); + } + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q1_0, 2048, 151936); + ggml_tensor * token_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 1); + REQUIRE(weight != nullptr); + REQUIRE(token_ids != nullptr); + ggml_tensor * output = ggml_get_rows(ctx, weight, token_ids); + REQUIRE(output != nullptr); + schedule_token_embedding_command(ctx, output, 1, 151936, 2048, "10"); + } + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q5_1, 640, 262144); + ggml_tensor * token_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 1); + REQUIRE(weight != nullptr); + REQUIRE(token_ids != nullptr); + ggml_tensor * output = ggml_get_rows(ctx, weight, token_ids); + REQUIRE(output != nullptr); + schedule_token_embedding_command(ctx, output, 1, 262144, 640, "51"); + } + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q8_0, 3840, 262208); + ggml_tensor * token_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 64); + REQUIRE(weight != nullptr); + REQUIRE(token_ids != nullptr); + ggml_tensor * output = ggml_get_rows(ctx, weight, token_ids); + REQUIRE(output != nullptr); + schedule_token_embedding_command(ctx, output, 64, 262208, 3840, "80"); + } + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, 3840, 262208); + ggml_tensor * token_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 64); + REQUIRE(weight != nullptr); + REQUIRE(token_ids != nullptr); + ggml_tensor * output = ggml_get_rows(ctx, weight, token_ids); + REQUIRE(output != nullptr); + schedule_token_embedding_command(ctx, output, 64, 262208, 3840, "6"); + } + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_BF16, 3840, 262208); + ggml_tensor * token_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 64); + REQUIRE(weight != nullptr); + REQUIRE(token_ids != nullptr); + ggml_tensor * output = ggml_get_rows(ctx, weight, token_ids); + REQUIRE(output != nullptr); + schedule_token_embedding_command(ctx, output, 64, 262208, 3840, "30"); + } + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 262144, 4); + ggml_tensor * token_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 1); + REQUIRE(weight != nullptr); + REQUIRE(token_ids != nullptr); + ggml_tensor * output = ggml_get_rows(ctx, weight, token_ids); + REQUIRE(output != nullptr); + schedule_token_embedding_command(ctx, output, 1, 4, 262144, "32"); + } + + REQUIRE( + manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_F32, GGML_TYPE_I32, GGML_TYPE_F32, 2048, 151936, 1)); + REQUIRE( + manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_Q6_K, GGML_TYPE_I32, GGML_TYPE_F32, 2048, 151936, 1)); + REQUIRE( + manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_Q1_0, GGML_TYPE_I32, GGML_TYPE_F32, 2048, 151936, 1)); + REQUIRE( + manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_Q5_1, GGML_TYPE_I32, GGML_TYPE_F32, 640, 262144, 1)); + REQUIRE(manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_IQ4_XS, GGML_TYPE_I32, GGML_TYPE_F32, 2048, 151936, + 1)); + REQUIRE(!manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_IQ3_S, GGML_TYPE_I32, GGML_TYPE_F32, 2048, 151936, + 1)); + REQUIRE(!manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_IQ4_NL, GGML_TYPE_I32, GGML_TYPE_F32, 2048, + 151936, 1)); + REQUIRE( + manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_Q8_0, GGML_TYPE_I32, GGML_TYPE_F32, 2048, 151936, 1)); + REQUIRE( + manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_Q8_1, GGML_TYPE_I32, GGML_TYPE_F32, 2048, 151936, 1)); + REQUIRE( + manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_F16, GGML_TYPE_I32, GGML_TYPE_F32, 2048, 151936, 1)); + REQUIRE( + manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_BF16, GGML_TYPE_I32, GGML_TYPE_F32, 2048, 151936, 1)); + REQUIRE( + manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_BF16, GGML_TYPE_I32, GGML_TYPE_F32, 3840, 262208, 64)); + REQUIRE( + !manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_Q4_K, GGML_TYPE_I64, GGML_TYPE_F32, 2048, 151936, 1)); + REQUIRE( + !manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_Q4_K, GGML_TYPE_I32, GGML_TYPE_F16, 2048, 151936, 1)); + REQUIRE( + manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_F32, GGML_TYPE_I32, GGML_TYPE_F32, 128, 151936, 1)); + REQUIRE(manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_F32, GGML_TYPE_I32, GGML_TYPE_F32, 262144, 4, 1)); + REQUIRE( + manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_F32, GGML_TYPE_I32, GGML_TYPE_F32, 33024, 151936, 1)); + REQUIRE(!manual_token_embedding_graph_is_supported(ctx, GGML_TYPE_Q4_K, GGML_TYPE_I32, GGML_TYPE_F32, 2048, 151936, + 1, 2048, 2)); + + ggml_free(ctx); +} + +static void run_get_rows_rmsnorm_binary_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * embedding_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, 2048, 128256); + ggml_tensor * token_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 512); + ggml_tensor * norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 2048); + REQUIRE(embedding_weight != nullptr); + REQUIRE(token_ids != nullptr); + REQUIRE(norm_weight != nullptr); + ggml_tensor * embedding = ggml_get_rows(ctx, embedding_weight, token_ids); + ggml_tensor * rms = ggml_rms_norm(ctx, embedding, 0.00001f); + ggml_tensor * output = ggml_mul(ctx, rms, norm_weight); + REQUIRE(embedding != nullptr); + REQUIRE(rms != nullptr); + REQUIRE(output != nullptr); + ggml_set_output(embedding); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 3); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const ggml::hrx::CommandPlan & plan = scheduler.plan(); + REQUIRE(plan.valid()); + REQUIRE(plan.dispatches.size() == 1); + REQUIRE(plan.transients.size() == 2); + const ggml::hrx::Dispatch & dispatch = plan.dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_get_rows_rmsnorm_binary_q8_1_x4_f16"); + REQUIRE(dispatch.bindings.size() == 7); + require_compile_parameter(dispatch, "ggml.get_rows_f32.weight_format", "6"); + require_compile_parameter(dispatch, "ggml.get_rows_rmsnorm.rms_epsilon", "9.99999975e-06"); + + const ggml::hrx::Value * output_value = imported.graph.values().find_tensor(output); + REQUIRE(output_value != nullptr); + const size_t q8_bytes = qwen_q8_1_x4_size(512, 2048); + const size_t f16_bytes = output_value->byte_count / 2; + const auto * q8 = plan.metadata.find_alternate_value(output_value->id, GGML_TYPE_Q8_1, q8_bytes); + const auto * f16 = plan.metadata.find_alternate_value(output_value->id, GGML_TYPE_F16, f16_bytes); + REQUIRE(q8 == nullptr); + REQUIRE(f16 == nullptr); + REQUIRE(dispatch.bindings[5].value == plan.transients[0].value); + REQUIRE(dispatch.bindings[6].value == plan.transients[1].value); + REQUIRE(plan.transients[0].size == q8_bytes); + REQUIRE(plan.transients[1].size == f16_bytes); + + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 1); + REQUIRE(command_program_verifies(commands)); + + ggml_free(ctx); +} + +static void run_get_rows_scale_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2560, 64); + ggml_tensor * ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 5); + REQUIRE(weight != nullptr); + REQUIRE(ids != nullptr); + ggml_tensor * rows = ggml_get_rows(ctx, weight, ids); + ggml_tensor * output = ggml_scale(ctx, rows, 0.177800179f); + REQUIRE(rows != nullptr); + REQUIRE(output != nullptr); + ggml_set_output(rows); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + require_scheduled_command_program( + graph, [&](const ggml::hrx::Graph & imported_graph, const ggml::hrx::CommandPlan & plan, + const ggml::hrx::CommandProgram & commands) { + REQUIRE(imported_graph.nodes().size() == 2); + REQUIRE(plan.dispatches.size() == 1); + const ggml::hrx::Dispatch & dispatch = plan.dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_get_rows_scale_f32"); + REQUIRE(dispatch.bindings.size() == 4); + require_compile_parameter(dispatch, "ggml.get_rows_f32.weight_format", "32"); + require_compile_parameter(dispatch, "ggml.get_rows_scale_f32.scale", expected_config_value(0.177800179f)); + require_compile_parameter(dispatch, "ggml.get_rows_scale_f32.bias", expected_config_value(0.0f)); + REQUIRE(commands.commands.size() == 1); + REQUIRE(commands.commands.front().bindings[0].name == "token_ids"); + REQUIRE(commands.commands.front().bindings[1].name == "weight"); + REQUIRE(commands.commands.front().bindings[2].name == "raw_output"); + REQUIRE(commands.commands.front().bindings[3].name == "output"); + }); + + ggml_tensor * q5_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q5_1, 640, 262144); + ggml_tensor * q5_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 64); + REQUIRE(q5_weight != nullptr); + REQUIRE(q5_ids != nullptr); + ggml_tensor * q5_rows = ggml_get_rows(ctx, q5_weight, q5_ids); + ggml_tensor * q5_output = ggml_scale(ctx, q5_rows, 25.2982216f); + REQUIRE(q5_rows != nullptr); + REQUIRE(q5_output != nullptr); + ggml_set_output(q5_rows); + + ggml_cgraph * q5_graph = ggml_new_graph(ctx); + REQUIRE(q5_graph != nullptr); + ggml_build_forward_expand(q5_graph, q5_output); + require_scheduled_command_program( + q5_graph, [&](const ggml::hrx::Graph & imported_graph, const ggml::hrx::CommandPlan & plan, + const ggml::hrx::CommandProgram & commands) { + REQUIRE(imported_graph.nodes().size() == 2); + REQUIRE(plan.dispatches.size() == 1); + const ggml::hrx::Dispatch & dispatch = plan.dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_get_rows_scale_f32"); + REQUIRE(dispatch.bindings.size() == 4); + require_compile_parameter(dispatch, "ggml.get_rows_f32.weight_format", "51"); + require_compile_parameter(dispatch, "ggml.get_rows_scale_f32.scale", expected_config_value(25.2982216f)); + require_compile_parameter(dispatch, "ggml.get_rows_scale_f32.bias", expected_config_value(0.0f)); + REQUIRE(commands.commands.size() == 1); + }); + + ggml_tensor * q5_decode_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 1); + REQUIRE(q5_decode_ids != nullptr); + ggml_tensor * q5_decode_rows = ggml_get_rows(ctx, q5_weight, q5_decode_ids); + ggml_tensor * q5_decode_output = ggml_scale(ctx, q5_decode_rows, 25.2982216f); + REQUIRE(q5_decode_rows != nullptr); + REQUIRE(q5_decode_output != nullptr); + ggml_cgraph * q5_decode_graph = ggml_new_graph(ctx); + REQUIRE(q5_decode_graph != nullptr); + ggml_build_forward_expand(q5_decode_graph, q5_decode_output); + const std::vector q5_decode_kernels = scheduled_kernel_names(q5_decode_graph); + REQUIRE(std::find(q5_decode_kernels.begin(), q5_decode_kernels.end(), "loom_libs:ggml_get_rows_scale_f32") == + q5_decode_kernels.end()); + REQUIRE(std::find(q5_decode_kernels.begin(), q5_decode_kernels.end(), "loom_libs:ggml_get_rows_f32") != + q5_decode_kernels.end()); + REQUIRE(std::find(q5_decode_kernels.begin(), q5_decode_kernels.end(), "loom_libs:ggml_scale_f32") != + q5_decode_kernels.end()); + + ggml_tensor * q4_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_0, 640, 64); + ggml_tensor * q4_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 5); + REQUIRE(q4_weight != nullptr); + REQUIRE(q4_ids != nullptr); + ggml_tensor * q4_rows = ggml_get_rows(ctx, q4_weight, q4_ids); + ggml_tensor * q4_output = ggml_scale(ctx, q4_rows, 0.5f); + REQUIRE(q4_rows != nullptr); + REQUIRE(q4_output != nullptr); + ggml_cgraph * q4_graph = ggml_new_graph(ctx); + REQUIRE(q4_graph != nullptr); + ggml_build_forward_expand(q4_graph, q4_output); + const std::vector q4_kernels = scheduled_kernel_names(q4_graph); + REQUIRE(std::find(q4_kernels.begin(), q4_kernels.end(), "loom_libs:ggml_get_rows_scale_f32") == q4_kernels.end()); + REQUIRE(std::find(q4_kernels.begin(), q4_kernels.end(), "loom_libs:ggml_get_rows_f32") != q4_kernels.end()); + + ggml_tensor * side_use = ggml_add(ctx, rows, rows); + REQUIRE(side_use != nullptr); + ggml_cgraph * fallback_graph = ggml_new_graph(ctx); + REQUIRE(fallback_graph != nullptr); + ggml_build_forward_expand(fallback_graph, output); + ggml_build_forward_expand(fallback_graph, side_use); + const std::vector fallback_kernels = scheduled_kernel_names(fallback_graph); + REQUIRE(std::find(fallback_kernels.begin(), fallback_kernels.end(), "loom_libs:ggml_get_rows_scale_f32") == + fallback_kernels.end()); + REQUIRE(std::find(fallback_kernels.begin(), fallback_kernels.end(), "loom_libs:ggml_get_rows_f32") != + fallback_kernels.end()); + + ggml_tensor * inplace_rows = ggml_get_rows(ctx, weight, ids); + ggml_tensor * inplace_output = ggml_scale_inplace(ctx, inplace_rows, 0.5f); + REQUIRE(inplace_rows != nullptr); + REQUIRE(inplace_output != nullptr); + ggml_cgraph * inplace_graph = ggml_new_graph(ctx); + REQUIRE(inplace_graph != nullptr); + ggml_build_forward_expand(inplace_graph, inplace_output); + const std::vector inplace_kernels = scheduled_kernel_names(inplace_graph); + REQUIRE(std::find(inplace_kernels.begin(), inplace_kernels.end(), "loom_libs:ggml_get_rows_scale_f32") == + inplace_kernels.end()); + REQUIRE(std::find(inplace_kernels.begin(), inplace_kernels.end(), "loom_libs:ggml_get_rows_f32") != + inplace_kernels.end()); + + ggml_free(ctx); +} + +static ggml::hrx::Graph build_manual_gather_add_graph(ggml_context * ctx, + int64_t hidden_size, + int64_t source_token_count, + int64_t output_token_count, + bool shared_row_ids, + int64_t second_source_hidden_size = -1) { + if (second_source_hidden_size < 0) { + second_source_hidden_size = hidden_size; + } + + ggml::hrx::Graph graph; + + ggml_tensor * attention = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden_size, source_token_count); + ggml_tensor * residual = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, second_source_hidden_size, source_token_count); + ggml_tensor * row_ids0 = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, output_token_count); + ggml_tensor * row_ids1 = shared_row_ids ? row_ids0 : ggml_new_tensor_1d(ctx, GGML_TYPE_I32, output_token_count); + ggml_tensor * selected0 = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden_size, output_token_count); + ggml_tensor * selected1 = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden_size, output_token_count); + ggml_tensor * output = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden_size, output_token_count); + REQUIRE(attention != nullptr); + REQUIRE(residual != nullptr); + REQUIRE(row_ids0 != nullptr); + REQUIRE(row_ids1 != nullptr); + REQUIRE(selected0 != nullptr); + REQUIRE(selected1 != nullptr); + REQUIRE(output != nullptr); + + const ggml::hrx::ValueId attention_value = + graph.values().get_or_add_tensor_value(attention, ggml::hrx::ValueKind::External); + const ggml::hrx::ValueId residual_value = + graph.values().get_or_add_tensor_value(residual, ggml::hrx::ValueKind::External); + const ggml::hrx::ValueId row_ids0_value = + graph.values().get_or_add_tensor_value(row_ids0, ggml::hrx::ValueKind::External); + const ggml::hrx::ValueId row_ids1_value = + graph.values().get_or_add_tensor_value(row_ids1, ggml::hrx::ValueKind::External); + const ggml::hrx::ValueId selected0_value = + graph.values().get_or_add_tensor_value(selected0, ggml::hrx::ValueKind::Transient); + const ggml::hrx::ValueId selected1_value = + graph.values().get_or_add_tensor_value(selected1, ggml::hrx::ValueKind::Transient); + const ggml::hrx::ValueId output_value = + graph.values().get_or_add_tensor_value(output, ggml::hrx::ValueKind::External); + + graph.add_node(GGML_OP_GET_ROWS, selected0_value, { attention_value, row_ids0_value }); + graph.add_node(GGML_OP_GET_ROWS, selected1_value, { residual_value, row_ids1_value }); + graph.add_node(GGML_OP_ADD, output_value, { selected0_value, selected1_value }); + REQUIRE(graph.build_index().success()); + return graph; +} + +struct GatherAddRmsNormGraph { + ggml_tensor * attention = nullptr; + ggml_tensor * residual = nullptr; + ggml_tensor * row_ids = nullptr; + ggml_tensor * weight = nullptr; + ggml_tensor * selected = nullptr; + ggml_tensor * raw_output = nullptr; + ggml_tensor * normalized_output = nullptr; + ggml_cgraph * graph = nullptr; +}; + +static GatherAddRmsNormGraph build_gather_add_rmsnorm_graph( + ggml_context * ctx, + int64_t hidden_size, + int64_t source_token_count, + int64_t output_token_count, + bool produced_row_ids = false, + bool produced_weight = false, + ggml::hrx::BinaryKind gather_op = ggml::hrx::BinaryKind::Add, + bool first_is_lhs = true) { + GatherAddRmsNormGraph result; + result.attention = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden_size, source_token_count); + result.residual = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden_size, source_token_count); + ggml_tensor * raw_row_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, output_token_count); + ggml_tensor * raw_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, hidden_size); + result.row_ids = produced_row_ids ? ggml_dup(ctx, raw_row_ids) : raw_row_ids; + result.weight = produced_weight ? ggml_scale(ctx, raw_weight, 0.5f) : raw_weight; + ggml_tensor * selected_attention = ggml_get_rows(ctx, result.attention, result.row_ids); + ggml_tensor * selected_residual = ggml_get_rows(ctx, result.residual, result.row_ids); + result.selected = selected_attention; + switch (gather_op) { + case ggml::hrx::BinaryKind::Add: + result.raw_output = ggml_add(ctx, selected_attention, selected_residual); + break; + case ggml::hrx::BinaryKind::Sub: + result.raw_output = first_is_lhs ? ggml_sub(ctx, selected_attention, selected_residual) : + ggml_sub(ctx, selected_residual, selected_attention); + break; + case ggml::hrx::BinaryKind::Mul: + result.raw_output = ggml_mul(ctx, selected_attention, selected_residual); + break; + case ggml::hrx::BinaryKind::Div: + result.raw_output = first_is_lhs ? ggml_div(ctx, selected_attention, selected_residual) : + ggml_div(ctx, selected_residual, selected_attention); + break; + default: + REQUIRE(false); + } + ggml_tensor * rms = ggml_rms_norm(ctx, result.raw_output, 0.000001f); + result.normalized_output = ggml_mul(ctx, rms, result.weight); + REQUIRE(result.attention != nullptr); + REQUIRE(result.residual != nullptr); + REQUIRE(raw_row_ids != nullptr); + REQUIRE(raw_weight != nullptr); + REQUIRE(result.row_ids != nullptr); + REQUIRE(result.weight != nullptr); + REQUIRE(selected_attention != nullptr); + REQUIRE(selected_residual != nullptr); + REQUIRE(result.raw_output != nullptr); + REQUIRE(rms != nullptr); + REQUIRE(result.normalized_output != nullptr); + + ggml_set_output(result.raw_output); + result.graph = ggml_new_graph(ctx); + REQUIRE(result.graph != nullptr); + ggml_build_forward_expand(result.graph, result.normalized_output); + return result; +} + +static ggml::hrx::DispatchRegistry gather_add_test_registry() { + ggml::hrx::DispatchRegistryBuilder builder; + ggml::hrx::register_gather_add_dispatch(builder); + return builder.build(); +} + +static bool match_gather_add_at(const ggml::hrx::DispatchRegistry & registry, + const ggml::hrx::Graph & graph, + size_t root_index, + const std::vector & covered_nodes, + ggml::hrx::DispatchMatch & match) { + ggml::hrx::CommandPlan plan; + const ggml::hrx::DispatchMatchContext context = { + graph, &graph.nodes()[root_index], + root_index, covered_nodes, + plan, ggml::hrx::ValueId(static_cast(graph.values().size())), + }; + return registry.match(context, match); +} + +static void require_gather_add_rmsnorm_route(ggml_context * ctx, + ggml::hrx::BinaryKind gather_op = ggml::hrx::BinaryKind::Add, + bool first_is_lhs = true) { + GatherAddRmsNormGraph tensors = + build_gather_add_rmsnorm_graph(ctx, 128, 4, 5, false, false, gather_op, first_is_lhs); + require_scheduled_command_program(tensors.graph, [&](const ggml::hrx::Graph &, const ggml::hrx::CommandPlan & plan, + const ggml::hrx::CommandProgram & commands) { + REQUIRE(plan.dispatches.size() == 1); + const ggml::hrx::Dispatch & dispatch = plan.dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_gather_add_rmsnorm_binary_f32"); + REQUIRE(dispatch.bindings.size() == 6); + require_compile_parameter(dispatch, "ggml.gather_add_rmsnorm_binary_f32.hidden_size", "128"); + require_compile_parameter(dispatch, "ggml.gather_add_rmsnorm_binary_f32.rms_epsilon", "9.99999997e-07"); + require_compile_parameter(dispatch, "ggml.gather_add_rmsnorm_binary_f32.gather_op", + std::to_string(ggml::hrx::binary_kind_config_value(gather_op))); + REQUIRE(commands.commands.size() == 1); + REQUIRE(commands.commands.front().bindings.size() == 6); + REQUIRE(commands.commands.front().bindings[0].name == "attention"); + REQUIRE(commands.commands.front().bindings[1].name == "residual"); + REQUIRE(commands.commands.front().bindings[2].name == "output_ids"); + REQUIRE(commands.commands.front().bindings[3].name == "raw_output"); + REQUIRE(commands.commands.front().bindings[4].name == "weight"); + REQUIRE(commands.commands.front().bindings[5].name == "normalized_output"); + }); +} + +static void require_gather_add_rmsnorm_availability(ggml_context * ctx) { + GatherAddRmsNormGraph tensors = build_gather_add_rmsnorm_graph(ctx, 128, 4, 5, true, true); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*tensors.graph); + REQUIRE(imported.valid()); + const ggml::hrx::DispatchRegistry registry = gather_add_test_registry(); + const size_t root_index = producer_index_for_tensor(imported.graph, tensors.selected); + const size_t row_index = producer_index_for_tensor(imported.graph, tensors.row_ids); + const size_t weight_index = producer_index_for_tensor(imported.graph, tensors.weight); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::DispatchMatch match; + + REQUIRE(!match_gather_add_at(registry, imported.graph, root_index, covered_nodes, match)); + + covered_nodes[row_index] = true; + match = {}; + REQUIRE(match_gather_add_at(registry, imported.graph, root_index, covered_nodes, match)); + REQUIRE(match.dispatches.size() == 1); + REQUIRE(kernel_name_for_id(match.dispatches.front().kernel.kernel_id) == "hrx:ggml_gather_add_f32"); + + covered_nodes[weight_index] = true; + match = {}; + REQUIRE(match_gather_add_at(registry, imported.graph, root_index, covered_nodes, match)); + REQUIRE(match.dispatches.size() == 1); + REQUIRE(kernel_name_for_id(match.dispatches.front().kernel.kernel_id) == + "loom_libs:ggml_gather_add_rmsnorm_binary_f32"); +} + +static void require_gather_add_rmsnorm_row_ids_noalias(ggml_context * ctx) { + GatherAddRmsNormGraph tensors = build_gather_add_rmsnorm_graph(ctx, 128, 4, 5); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*tensors.graph); + REQUIRE(imported.valid()); + const ggml::hrx::Value * row_ids = imported.graph.values().find_tensor(tensors.row_ids); + const ggml::hrx::Value * attention = imported.graph.values().find_tensor(tensors.attention); + REQUIRE(row_ids != nullptr); + REQUIRE(attention != nullptr); + + ggml::hrx::Graph aliased_graph; + for (const ggml::hrx::ValueStorage & storage : imported.graph.values().storages()) { + REQUIRE(aliased_graph.values().add_snapshot_storage(storage).success()); + } + for (ggml::hrx::Value value : imported.graph.values().values()) { + if (value.id == row_ids->id) { + value.storage = attention->storage; + value.storage_root = attention->storage_root; + value.alias_source = attention->id; + value.storage_offset = attention->storage_offset; + value.storage_byte_count = attention->storage_byte_count; + } + REQUIRE(aliased_graph.values().add_snapshot_value(std::move(value)).success()); + } + for (const ggml::hrx::GraphNode & node : imported.graph.nodes()) { + ggml::hrx::GraphNode & copied = aliased_graph.add_node(node.op, node.output, node.inputs); + copied.params = node.params; + } + REQUIRE(aliased_graph.build_index().success()); + + const ggml::hrx::DispatchRegistry registry = gather_add_test_registry(); + size_t root_index = aliased_graph.nodes().size(); + for (size_t i = 0; i < aliased_graph.nodes().size(); ++i) { + if (aliased_graph.nodes()[i].op == GGML_OP_GET_ROWS) { + root_index = i; + break; + } + } + REQUIRE(root_index < aliased_graph.nodes().size()); + std::vector covered_nodes(aliased_graph.nodes().size(), false); + ggml::hrx::DispatchMatch match; + REQUIRE(match_gather_add_at(registry, aliased_graph, root_index, covered_nodes, match)); + REQUIRE(match.dispatches.size() == 1); + REQUIRE(kernel_name_for_id(match.dispatches.front().kernel.kernel_id) == "hrx:ggml_gather_add_f32"); +} + +static bool manual_gather_add_graph_is_supported(ggml_context * ctx, + int64_t hidden_size, + int64_t source_token_count, + int64_t output_token_count, + bool shared_row_ids, + int64_t second_source_hidden_size = -1) { + ggml::hrx::Graph graph = build_manual_gather_add_graph(ctx, hidden_size, source_token_count, output_token_count, + shared_row_ids, second_source_hidden_size); + return ggml::hrx::DispatchScheduler::can_schedule_graph(graph, test_dispatch_target()); +} + +static bool partial_gather_add_graph_is_supported(ggml_context * ctx) { + ggml::hrx::Graph graph; + + ggml_tensor * attention = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 13); + ggml_tensor * row_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 1); + ggml_tensor * selected = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + ggml_tensor * add_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + ggml_tensor * add_output = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(attention != nullptr); + REQUIRE(row_ids != nullptr); + REQUIRE(selected != nullptr); + REQUIRE(add_input != nullptr); + REQUIRE(add_output != nullptr); + + const ggml::hrx::ValueId attention_value = + graph.values().get_or_add_tensor_value(attention, ggml::hrx::ValueKind::External); + const ggml::hrx::ValueId row_ids_value = + graph.values().get_or_add_tensor_value(row_ids, ggml::hrx::ValueKind::External); + const ggml::hrx::ValueId selected_value = + graph.values().get_or_add_tensor_value(selected, ggml::hrx::ValueKind::Transient); + const ggml::hrx::ValueId add_input_value = + graph.values().get_or_add_tensor_value(add_input, ggml::hrx::ValueKind::External); + const ggml::hrx::ValueId add_output_value = + graph.values().get_or_add_tensor_value(add_output, ggml::hrx::ValueKind::External); + + graph.add_node(GGML_OP_GET_ROWS, selected_value, { attention_value, row_ids_value }); + graph.add_node(GGML_OP_ADD, add_output_value, { selected_value, add_input_value }); + REQUIRE(graph.build_index().success()); + return ggml::hrx::DispatchScheduler::can_schedule_graph(graph, test_dispatch_target()); +} + +static void require_gather_add_rejects_unavailable_sources(ggml_context * ctx) { + ggml::hrx::Graph graph; + + ggml_tensor * raw_attention = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 13); + ggml_tensor * attention = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 13); + ggml_tensor * residual = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 13); + ggml_tensor * row_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 1); + ggml_tensor * selected0 = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + ggml_tensor * selected1 = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + ggml_tensor * output = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(raw_attention != nullptr); + REQUIRE(attention != nullptr); + REQUIRE(residual != nullptr); + REQUIRE(row_ids != nullptr); + REQUIRE(selected0 != nullptr); + REQUIRE(selected1 != nullptr); + REQUIRE(output != nullptr); + + const ggml::hrx::ValueId raw_attention_value = + graph.values().get_or_add_tensor_value(raw_attention, ggml::hrx::ValueKind::External); + const ggml::hrx::ValueId attention_value = + graph.values().get_or_add_tensor_value(attention, ggml::hrx::ValueKind::Transient); + const ggml::hrx::ValueId residual_value = + graph.values().get_or_add_tensor_value(residual, ggml::hrx::ValueKind::External); + const ggml::hrx::ValueId row_ids_value = + graph.values().get_or_add_tensor_value(row_ids, ggml::hrx::ValueKind::External); + const ggml::hrx::ValueId selected0_value = + graph.values().get_or_add_tensor_value(selected0, ggml::hrx::ValueKind::Transient); + const ggml::hrx::ValueId selected1_value = + graph.values().get_or_add_tensor_value(selected1, ggml::hrx::ValueKind::Transient); + const ggml::hrx::ValueId output_value = + graph.values().get_or_add_tensor_value(output, ggml::hrx::ValueKind::External); + + graph.add_node(GGML_OP_SCALE, attention_value, { raw_attention_value }); + graph.add_node(GGML_OP_GET_ROWS, selected0_value, { attention_value, row_ids_value }); + graph.add_node(GGML_OP_GET_ROWS, selected1_value, { residual_value, row_ids_value }); + graph.add_node(GGML_OP_ADD, output_value, { selected0_value, selected1_value }); + REQUIRE(graph.build_index().success()); + + ggml::hrx::DispatchRegistryBuilder builder; + ggml::hrx::register_gather_add_dispatch(builder); + const ggml::hrx::DispatchRegistry gather_add_registry = builder.build(); + + const size_t producer_index = 0; + const size_t get_rows_index = 1; + ggml::hrx::CommandPlan plan; + ggml::hrx::DispatchMatch match; + std::vector covered_nodes(graph.nodes().size(), false); + ggml::hrx::DispatchMatchContext context = { + graph, &graph.nodes()[get_rows_index], + get_rows_index, covered_nodes, + plan, ggml::hrx::ValueId(static_cast(graph.values().size())), + }; + + REQUIRE(!gather_add_registry.match(context, match)); + + covered_nodes[producer_index] = true; + const ggml::hrx::DispatchMatchContext available_context = { + graph, &graph.nodes()[get_rows_index], + get_rows_index, covered_nodes, + plan, ggml::hrx::ValueId(static_cast(graph.values().size())), + }; + match = {}; + REQUIRE(gather_add_registry.match(available_context, match)); + REQUIRE(match.dispatches.size() == 1); + REQUIRE(kernel_name_for_id(match.dispatches.front().kernel.kernel_id) == "hrx:ggml_gather_add_f32"); +} + +static void schedule_gather_add_command(ggml::hrx::Graph & graph, + int64_t expected_hidden_size, + int64_t expected_source_token_count, + int64_t expected_output_token_count) { + ggml::hrx::DispatchScheduler scheduler; + if (!scheduler.schedule_graph(graph, test_dispatch_target())) { + for (const std::string & error : scheduler.plan().status.errors()) { + std::fprintf(stderr, "%s\n", error.c_str()); + } + REQUIRE(false); + } + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "hrx:ggml_gather_add_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("hidden_size") == expected_hidden_size); + REQUIRE(dispatch.kernel.integer_parameters.at("source_token_count") == expected_source_token_count); + REQUIRE(dispatch.kernel.integer_parameters.at("output_token_count") == expected_output_token_count); + require_compile_parameter(dispatch, "ggml.gather_add_f32.binary_op", "0"); + REQUIRE(dispatch.bindings.size() == 4); + + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 1); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.front().bindings.size() == 4); + REQUIRE(commands.commands.front().bindings[0].name == "attention"); + REQUIRE(commands.commands.front().bindings[1].name == "residual"); + REQUIRE(commands.commands.front().bindings[2].name == "output_ids"); + REQUIRE(commands.commands.front().bindings[3].name == "output"); +} + +static void run_gather_add_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + { + ggml::hrx::Graph graph = build_manual_gather_add_graph(ctx, 2048, 13, 1, true); + schedule_gather_add_command(graph, 2048, 13, 1); + } + { + ggml::hrx::Graph graph = build_manual_gather_add_graph(ctx, 2048, 128, 8, true); + schedule_gather_add_command(graph, 2048, 128, 8); + } + + REQUIRE(!manual_gather_add_graph_is_supported(ctx, 2048, 13, 1, false)); + REQUIRE(!manual_gather_add_graph_is_supported(ctx, 2048, 13, 1, true, 1024)); + REQUIRE(!manual_gather_add_graph_is_supported(ctx, 96, 13, 1, true)); + REQUIRE(!partial_gather_add_graph_is_supported(ctx)); + require_gather_add_rejects_unavailable_sources(ctx); + require_gather_add_rmsnorm_route(ctx); + require_gather_add_rmsnorm_route(ctx, ggml::hrx::BinaryKind::Sub, false); + require_gather_add_rmsnorm_route(ctx, ggml::hrx::BinaryKind::Mul); + require_gather_add_rmsnorm_route(ctx, ggml::hrx::BinaryKind::Div); + require_gather_add_rmsnorm_availability(ctx); + require_gather_add_rmsnorm_row_ids_noalias(ctx); + + ggml_free(ctx); +} + +static void schedule_flash_attention_command(ggml_context * ctx, + ggml_tensor * output, + const char * expected_kernel_name, + const char * config_prefix, + int64_t query_token_count, + int64_t key_value_token_count, + int64_t query_head_count, + int64_t key_value_head_count, + int64_t qk_head_size = kQwenFlashHeadSize, + int64_t value_head_size = kQwenFlashHeadSize, + float scale = 1.0f / std::sqrt(128.0f), + bool transposed_value = false) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 1); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_FLASH_ATTN_EXT); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == (transposed_value ? 2 : 1)); + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.back(); + const std::string kernel_name = kernel_name_for_id(dispatch.kernel.kernel_id); + REQUIRE(kernel_name == expected_kernel_name); + REQUIRE(dispatch.kernel.integer_parameters.at("query_token_count") == query_token_count); + REQUIRE(dispatch.kernel.integer_parameters.at("key_value_token_count") == key_value_token_count); + REQUIRE(dispatch.bindings.size() == 6); + require_compile_parameter(dispatch, (std::string(config_prefix) + "query_head_count").c_str(), + std::to_string(query_head_count)); + require_compile_parameter(dispatch, (std::string(config_prefix) + "key_value_head_count").c_str(), + std::to_string(key_value_head_count)); + if (std::strcmp(config_prefix, "ggml.flash_attention.") == 0) { + require_compile_parameter(dispatch, "ggml.flash_attention.qk_head_size", std::to_string(qk_head_size)); + require_compile_parameter(dispatch, "ggml.flash_attention.value_head_size", std::to_string(value_head_size)); + require_compile_parameter(dispatch, "ggml.flash_attention.attention_scale", expected_config_value(scale)); + require_compile_parameter(dispatch, "ggml.flash_attention.apply_gate", "0"); + require_compile_parameter(dispatch, "ggml.flash_attention.gate_stride_head", "1"); + require_compile_parameter(dispatch, "ggml.flash_attention.gate_stride_token", "1"); + } else { + require_compile_parameter(dispatch, "qwen3_moe.workload.token_capacity", std::to_string(query_token_count)); + } + + if (transposed_value) { + REQUIRE(scheduler.plan().transients.size() == 1); + const auto & copy = scheduler.plan().dispatches.front(); + REQUIRE(kernel_name_for_id(copy.kernel.kernel_id) == "loom_libs:ggml_copy_transpose_f16"); + REQUIRE(copy.kernel.workload_specialization == ggml::hrx::WorkloadSpecialization::Dynamic); + REQUIRE(copy.kernel.integer_parameters.at("row_count") == key_value_token_count); + REQUIRE(copy.kernel.integer_parameters.at("column_count") == key_value_head_count * value_head_size); + REQUIRE(copy.bindings.size() == 2); + REQUIRE(copy.bindings[1].value == dispatch.bindings[2].value); + REQUIRE(copy.bindings[0].value != copy.bindings[1].value); + REQUIRE(copy.bindings[1].length == + static_cast(key_value_token_count * key_value_head_count * value_head_size) * + sizeof(ggml_fp16_t)); + require_compile_parameter(dispatch, "ggml.flash_attention.value_layout", "1"); + } else { + REQUIRE(dispatch.kernel.compile_parameters.count("ggml.flash_attention.value_layout") == 0); + } + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == (transposed_value ? 2 : 1)); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.back().bindings.size() == 6); + REQUIRE(commands.commands.back().bindings[0].name == "query"); + REQUIRE(commands.commands.back().bindings[1].name == "key"); + REQUIRE(commands.commands.back().bindings[2].name == "value"); + REQUIRE(commands.commands.back().bindings[3].name == "mask"); + REQUIRE(commands.commands.back().bindings[4].name == "gate"); + REQUIRE(commands.commands.back().bindings[5].name == "output"); +} + +static void schedule_qwen_flash_attention_command(ggml_context * ctx, ggml_tensor * output) { + schedule_flash_attention_command(ctx, output, "loom_libs:ggml_flash_attention_f32_f16_wmma", + "ggml.flash_attention.", 16, 16, 4, 2); +} + +static void schedule_common_flash_attention_command(ggml_context * ctx, + ggml_tensor * output, + int64_t query_token_count, + int64_t key_value_token_count, + int64_t query_head_count, + int64_t key_value_head_count, + int64_t qk_head_size = kQwenFlashHeadSize, + int64_t value_head_size = kQwenFlashHeadSize, + float scale = 1.0f / std::sqrt(128.0f)) { + schedule_flash_attention_command(ctx, output, "loom_libs:ggml_flash_attention_f32_f16_wmma", + "ggml.flash_attention.", query_token_count, key_value_token_count, + query_head_count, key_value_head_count, qk_head_size, value_head_size, scale); +} + +static void schedule_common_flash_attention_decode_fallback_command(ggml_context * ctx, + ggml_tensor * output, + int64_t query_token_count, + int64_t key_value_token_count, + int64_t query_head_count, + int64_t key_value_head_count, + int64_t qk_head_size = kQwenFlashHeadSize, + int64_t value_head_size = kQwenFlashHeadSize, + float scale = 1.0f / std::sqrt(128.0f)) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 1); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_FLASH_ATTN_EXT); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == static_cast(query_token_count)); + REQUIRE(scheduler.plan().transients.size() == 4); + REQUIRE(scheduler.plan().completion_counter_requests.size() == 1); + REQUIRE(scheduler.plan().completion_counter_requests.front().count == key_value_head_count); + REQUIRE(scheduler.plan().metadata.alternate_values().empty()); + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == + "loom_libs:ggml_flash_attention_decode_split_f32_f16_wmma_next_q8"); + REQUIRE(dispatch.kernel.integer_parameters.at("key_value_token_count") == key_value_token_count); + REQUIRE(dispatch.bindings.size() == 10); + require_compile_parameter(dispatch, "ggml.flash_attention.query_head_count", std::to_string(query_head_count)); + require_compile_parameter(dispatch, "ggml.flash_attention.key_value_head_count", + std::to_string(key_value_head_count)); + require_compile_parameter(dispatch, "ggml.flash_attention.qk_head_size", std::to_string(qk_head_size)); + require_compile_parameter(dispatch, "ggml.flash_attention.value_head_size", std::to_string(value_head_size)); + require_compile_parameter(dispatch, "ggml.flash_attention.attention_scale", expected_config_value(scale)); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == static_cast(query_token_count)); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.front().bindings.size() == 10); + REQUIRE(commands.commands.front().bindings[0].name == "query"); + REQUIRE(commands.commands.front().bindings[1].name == "key"); + REQUIRE(commands.commands.front().bindings[2].name == "value"); + REQUIRE(commands.commands.front().bindings[3].name == "mask"); + REQUIRE(commands.commands.front().bindings[8].name == "output"); + REQUIRE(commands.commands.front().bindings[9].name == "next_q8_output"); +} + +static void schedule_common_flash_attention_decode_q8_command(ggml_context * ctx, ggml_type weight_type) { + constexpr int64_t query_token_count = 1; + constexpr int64_t key_value_token_count = 512; + constexpr int64_t query_head_count = 4; + constexpr int64_t key_value_head_count = 2; + constexpr int64_t hidden_size = query_head_count * kQwenFlashHeadSize; + constexpr int64_t output_size = 256; + const size_t q8_bytes = ggml_row_size(GGML_TYPE_Q8_1, hidden_size); + + ggml_tensor * output = build_qwen_flash_attention_graph( + ctx, query_token_count, key_value_token_count, query_head_count, key_value_head_count); + ggml_tensor * reshaped = ggml_reshape_2d(ctx, output, hidden_size, query_token_count); + ggml_tensor * weight = ggml_new_tensor_2d(ctx, weight_type, hidden_size, output_size); + ggml_tensor * projected = ggml_mul_mat(ctx, weight, reshaped); + REQUIRE(reshaped != nullptr); + REQUIRE(weight != nullptr); + REQUIRE(projected != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, projected); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const ggml::hrx::CommandPlan & plan = scheduler.plan(); + REQUIRE(plan.valid()); + REQUIRE(plan.dispatches.size() == 2); + REQUIRE(plan.transients.size() == 4); + REQUIRE(plan.completion_counter_requests.size() == 1); + + const ggml::hrx::Value * native_value = imported.graph.values().find_tensor(output); + const ggml::hrx::Value * reshape_value = imported.graph.values().find_tensor(reshaped); + REQUIRE(native_value != nullptr); + REQUIRE(reshape_value != nullptr); + const ggml::hrx::CommandPlanAlternateValue * q8 = + ggml::hrx::find_alternate_value(plan, reshape_value->id, GGML_TYPE_Q8_1, q8_bytes); + REQUIRE(q8 != nullptr); + REQUIRE(ggml::hrx::find_alternate_value(plan, native_value->id, GGML_TYPE_Q8_1, q8_bytes) == nullptr); + + const ggml::hrx::Dispatch & producer = plan.dispatches.front(); + const ggml::hrx::Dispatch & consumer = plan.dispatches.back(); + REQUIRE(kernel_name_for_id(producer.kernel.kernel_id) == + "loom_libs:ggml_flash_attention_decode_split_f32_f16_wmma_next_q8"); + REQUIRE(producer.bindings.size() == 10); + REQUIRE(producer.bindings[9].value == q8->alternate_value); + REQUIRE(consumer.bindings[0].value == q8->alternate_value); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 2); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.front().bindings.size() == 10); + REQUIRE(commands.commands.front().bindings[8].name == "output"); + REQUIRE(commands.commands.front().bindings[9].name == "next_q8_output"); +} + +static void schedule_common_flash_attention_decode_unqualified_command(ggml_context * ctx) { + constexpr int64_t hidden_size = 4 * kQwenFlashHeadSize; + ggml_tensor * output = build_qwen_flash_attention_graph(ctx, 1, 512, 4, 2); + ggml_tensor * reshaped = ggml_reshape_2d(ctx, output, hidden_size, 1); + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden_size, 256); + ggml_tensor * projected = ggml_mul_mat(ctx, weight, reshaped); + REQUIRE(reshaped != nullptr); + REQUIRE(weight != nullptr); + REQUIRE(projected != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, projected); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const ggml::hrx::CommandPlan & plan = scheduler.plan(); + REQUIRE(plan.valid()); + REQUIRE(plan.metadata.alternate_values().empty()); + REQUIRE(plan.dispatches.size() == 2); + REQUIRE(plan.transients.size() == 4); + REQUIRE(plan.completion_counter_requests.size() == 1); + REQUIRE(kernel_name_for_id(plan.dispatches.front().kernel.kernel_id) == + "loom_libs:ggml_flash_attention_decode_split_f32_f16_wmma_next_q8"); + REQUIRE(plan.dispatches.front().bindings.size() == 10); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.size() == 2); + REQUIRE(commands.commands.front().bindings.size() == 10); +} + +static void run_qwen_flash_attention_dispatch_checks() { + static constexpr const char * kDisableQwenDispatchEnv = "GGML_HRX_DISABLE_QWEN_DISPATCH"; + const char * original_env = std::getenv(kDisableQwenDispatchEnv); + const bool had_original_env = original_env != nullptr; + const std::string original_env_value = had_original_env ? original_env : ""; + REQUIRE(unsetenv(kDisableQwenDispatchEnv) == 0); + + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + for (const auto & shape : { + std::array{ 256, 512, 64, 1 }, + std::array{ 512, 512, 256, 1 }, + std::array{ 512, 2048, 320, 1 }, + std::array{ 255, 512, 256, 0 }, + std::array{ 512, 513, 256, 0 } + }) { + const float scale = 1.0f / std::sqrt(static_cast(shape[2])); + ggml_tensor * output = build_qwen_flash_attention_graph(ctx, shape[0], shape[1], 4, 2, GGML_TYPE_F32, + GGML_TYPE_F16, true, false, shape[2], scale); + schedule_flash_attention_command(ctx, output, "loom_libs:ggml_flash_attention_f32_f16_wmma", + "ggml.flash_attention.", shape[0], shape[1], 4, 2, shape[2], shape[2], scale, + shape[3] != 0); + } + + { + ggml_tensor * output = build_qwen_flash_attention_graph(ctx, 4, 8, 4, 2); + schedule_common_flash_attention_decode_fallback_command(ctx, output, 4, 8, 4, 2); + } + { + ggml_tensor * output = build_qwen_flash_attention_graph(ctx, 16, 16, 4, 2); + schedule_qwen_flash_attention_command(ctx, output); + } + { + REQUIRE(setenv(kDisableQwenDispatchEnv, "1", 1) == 0); + ggml_tensor * output = build_qwen_flash_attention_graph(ctx, 16, 16, 4, 2); + schedule_common_flash_attention_command(ctx, output, 16, 16, 4, 2); + REQUIRE(unsetenv(kDisableQwenDispatchEnv) == 0); + } + { + REQUIRE(setenv(kDisableQwenDispatchEnv, "1", 1) == 0); + constexpr int64_t head_size = 256; + ggml_tensor * output = + build_qwen_flash_attention_graph(ctx, 16, 16, 4, 2, GGML_TYPE_F32, GGML_TYPE_F16, true, false, head_size, + 1.0f / std::sqrt(static_cast(head_size))); + schedule_common_flash_attention_command(ctx, output, 16, 16, 4, 2, head_size, head_size, + 1.0f / std::sqrt(static_cast(head_size))); + REQUIRE(unsetenv(kDisableQwenDispatchEnv) == 0); + } + { + REQUIRE(setenv(kDisableQwenDispatchEnv, "1", 1) == 0); + constexpr int64_t qk_head_size = 96; + constexpr int64_t value_head_size = 64; + ggml_tensor * output = build_qwen_flash_attention_graph( + ctx, 23, 256, 40, 40, GGML_TYPE_F32, GGML_TYPE_F16, true, false, qk_head_size, + 1.0f / std::sqrt(static_cast(qk_head_size)), value_head_size); + schedule_common_flash_attention_command(ctx, output, 23, 256, 40, 40, qk_head_size, value_head_size, + 1.0f / std::sqrt(static_cast(qk_head_size))); + REQUIRE(unsetenv(kDisableQwenDispatchEnv) == 0); + } + { + ggml_tensor * output = build_qwen_flash_attention_graph(ctx, 1, 8, 4, 2); + schedule_common_flash_attention_decode_fallback_command(ctx, output, 1, 8, 4, 2); + } + { + REQUIRE(setenv(kDisableQwenDispatchEnv, "1", 1) == 0); + constexpr int64_t qk_head_size = 96; + constexpr int64_t value_head_size = 64; + ggml_tensor * output = build_qwen_flash_attention_graph( + ctx, 1, 512, 32, 4, GGML_TYPE_F32, GGML_TYPE_F16, true, false, qk_head_size, + 1.0f / std::sqrt(static_cast(qk_head_size)), value_head_size); + schedule_common_flash_attention_decode_fallback_command(ctx, output, 1, 512, 32, 4, qk_head_size, + value_head_size, + 1.0f / std::sqrt(static_cast(qk_head_size))); + REQUIRE(unsetenv(kDisableQwenDispatchEnv) == 0); + } + { + REQUIRE(setenv(kDisableQwenDispatchEnv, "1", 1) == 0); + schedule_common_flash_attention_decode_q8_command(ctx, GGML_TYPE_Q4_K); + schedule_common_flash_attention_decode_q8_command(ctx, GGML_TYPE_Q6_K); + schedule_common_flash_attention_decode_unqualified_command(ctx); + REQUIRE(unsetenv(kDisableQwenDispatchEnv) == 0); + } + { + REQUIRE(setenv(kDisableQwenDispatchEnv, "1", 1) == 0); + constexpr int64_t tokens = 512; + constexpr int64_t hidden = 512; + ggml_tensor * output = build_qwen_flash_attention_graph(ctx, tokens, tokens, 4, 2); + ggml_tensor * reshaped = ggml_reshape_2d(ctx, output, hidden, tokens); + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, hidden, hidden); + ggml_tensor * projected = ggml_mul_mat(ctx, weight, reshaped); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, projected); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const ggml::hrx::CommandPlan & plan = scheduler.plan(); + REQUIRE(plan.valid()); + const auto producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == + "loom_libs:ggml_flash_attention_f32_f16_wmma_publish_f16"; + }); + REQUIRE(producer != plan.dispatches.end()); + REQUIRE(producer->bindings.size() == 7); + const ggml::hrx::Value * output_value = imported.graph.values().find_tensor(output); + REQUIRE(output_value != nullptr); + const auto * alternate = + plan.metadata.find_alternate_value(output_value->id, GGML_TYPE_F16, output_value->byte_count / 2); + REQUIRE(alternate != nullptr); + REQUIRE(producer->bindings.back().value == alternate->alternate_value); + const auto consumer = std::next(producer); + REQUIRE(consumer != plan.dispatches.end()); + REQUIRE(consumer->bindings.front().value == alternate->alternate_value); + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + REQUIRE(unsetenv(kDisableQwenDispatchEnv) == 0); + } + { + ggml_tensor * output = build_qwen_flash_attention_graph(ctx, 4, 8, 4, 2, GGML_TYPE_F32, GGML_TYPE_F32); + REQUIRE(!graph_is_supported(ctx, output)); + } + { + ggml_tensor * output = build_qwen_flash_attention_graph(ctx, 4, 8, 4, 2, GGML_TYPE_F32, GGML_TYPE_F16, false); + REQUIRE(!graph_is_supported(ctx, output)); + } + { + ggml_tensor * output = + build_qwen_flash_attention_graph(ctx, 4, 8, 4, 2, GGML_TYPE_F32, GGML_TYPE_F16, true, true); + REQUIRE(!graph_is_supported(ctx, output)); + } + { + ggml_tensor * output = build_qwen_flash_attention_graph(ctx, 16, 16, 4, 2, GGML_TYPE_F32, GGML_TYPE_F16, true, + false, kQwenFlashHeadSize, 1.0f); + schedule_common_flash_attention_command(ctx, output, 16, 16, 4, 2, kQwenFlashHeadSize, kQwenFlashHeadSize, + 1.0f); + } + { + ggml_tensor * output = + build_qwen_flash_attention_graph(ctx, 4, 8, 4, 2, GGML_TYPE_F32, GGML_TYPE_F16, true, false, + kQwenFlashHeadSize, 1.0f / std::sqrt(128.0f), -1, 1.0f); + REQUIRE(!graph_is_supported(ctx, output)); + } + + ggml_free(ctx); + restore_environment_value(kDisableQwenDispatchEnv, had_original_env, original_env_value); +} + +struct QwenAttentionPostprocessTensors { + ggml_tensor * query_raw = nullptr; + ggml_tensor * key_raw = nullptr; + ggml_tensor * value_raw = nullptr; + ggml_tensor * query_reshape = nullptr; + ggml_tensor * key_reshape = nullptr; + ggml_tensor * value_reshape = nullptr; + ggml_tensor * query_output = nullptr; + ggml_tensor * key_cache = nullptr; + ggml_tensor * value_cache = nullptr; + ggml_tensor * key_output = nullptr; + ggml_tensor * value_output = nullptr; + ggml_tensor * positions = nullptr; + ggml_tensor * key_cache_indices = nullptr; + ggml_tensor * value_cache_indices = nullptr; + ggml_tensor * mask = nullptr; + ggml_tensor * flash_output = nullptr; +}; + +static QwenAttentionPostprocessTensors build_qwen_attention_postprocess_graph(ggml_context * ctx, + int64_t token_count, + int64_t query_head_count, + int64_t key_value_head_count, + int64_t cache_row_count, + float rms_epsilon = 0.000001f, + bool include_inverse_frequencies = true, + int rope_n_ctx_orig = 0, + float rope_freq_base = 10000.0f, + float rope_freq_scale = 1.0f, + float rope_ext_factor = 0.0f, + float rope_beta_fast = 0.0f, + float rope_beta_slow = 0.0f) { + QwenAttentionPostprocessTensors tensors; + const int64_t query_size = query_head_count * kQwenFlashHeadSize; + const int64_t key_value_size = key_value_head_count * kQwenFlashHeadSize; + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenMoeHiddenSize, token_count); + ggml_tensor * query_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, kQwenMoeHiddenSize, query_size); + ggml_tensor * key_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, kQwenMoeHiddenSize, key_value_size); + ggml_tensor * value_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, kQwenMoeHiddenSize, key_value_size); + REQUIRE(input != nullptr); + REQUIRE(query_weight != nullptr); + REQUIRE(key_weight != nullptr); + REQUIRE(value_weight != nullptr); + + tensors.query_raw = ggml_mul_mat(ctx, query_weight, input); + tensors.key_raw = ggml_mul_mat(ctx, key_weight, input); + tensors.value_raw = ggml_mul_mat(ctx, value_weight, input); + REQUIRE(tensors.query_raw != nullptr); + REQUIRE(tensors.key_raw != nullptr); + REQUIRE(tensors.value_raw != nullptr); + + tensors.query_reshape = ggml_reshape_3d(ctx, tensors.query_raw, kQwenFlashHeadSize, query_head_count, token_count); + tensors.key_reshape = ggml_reshape_3d(ctx, tensors.key_raw, kQwenFlashHeadSize, key_value_head_count, token_count); + tensors.value_reshape = + ggml_reshape_3d(ctx, tensors.value_raw, kQwenFlashHeadSize, key_value_head_count, token_count); + REQUIRE(tensors.query_reshape != nullptr); + REQUIRE(tensors.key_reshape != nullptr); + REQUIRE(tensors.value_reshape != nullptr); + + ggml_tensor * query_norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, kQwenFlashHeadSize); + ggml_tensor * key_norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, kQwenFlashHeadSize); + tensors.positions = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, token_count); + ggml_tensor * inverse_frequencies = + include_inverse_frequencies ? ggml_new_tensor_1d(ctx, GGML_TYPE_F32, kQwenFlashHeadSize / 2) : nullptr; + REQUIRE(query_norm_weight != nullptr); + REQUIRE(key_norm_weight != nullptr); + REQUIRE(tensors.positions != nullptr); + REQUIRE(include_inverse_frequencies == (inverse_frequencies != nullptr)); + + ggml_tensor * query_norm = ggml_rms_norm(ctx, tensors.query_reshape, rms_epsilon); + ggml_tensor * query_mul = ggml_mul(ctx, query_norm, query_norm_weight); + REQUIRE(query_norm != nullptr); + REQUIRE(query_mul != nullptr); + tensors.query_output = ggml_rope_ext(ctx, query_mul, tensors.positions, inverse_frequencies, kQwenFlashHeadSize, + GGML_ROPE_TYPE_NEOX, rope_n_ctx_orig, rope_freq_base, rope_freq_scale, + rope_ext_factor, 1.0f, rope_beta_fast, rope_beta_slow); + REQUIRE(tensors.query_output != nullptr); + + ggml_tensor * key_norm = ggml_rms_norm(ctx, tensors.key_reshape, rms_epsilon); + ggml_tensor * key_mul = ggml_mul(ctx, key_norm, key_norm_weight); + REQUIRE(key_norm != nullptr); + REQUIRE(key_mul != nullptr); + ggml_tensor * key_rope = ggml_rope_ext(ctx, key_mul, tensors.positions, inverse_frequencies, kQwenFlashHeadSize, + GGML_ROPE_TYPE_NEOX, rope_n_ctx_orig, rope_freq_base, rope_freq_scale, + rope_ext_factor, 1.0f, rope_beta_fast, rope_beta_slow); + REQUIRE(key_rope != nullptr); + + ggml_tensor * key_cache_rows = + ggml_reshape_2d(ctx, key_rope, kQwenFlashHeadSize * key_value_head_count, token_count); + ggml_tensor * value_cache_rows = + ggml_reshape_2d(ctx, tensors.value_reshape, kQwenFlashHeadSize * key_value_head_count, token_count); + REQUIRE(key_cache_rows != nullptr); + REQUIRE(value_cache_rows != nullptr); + + tensors.key_cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, key_value_size, cache_row_count); + tensors.value_cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, key_value_size, cache_row_count); + tensors.key_cache_indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, token_count); + tensors.value_cache_indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, token_count); + REQUIRE(tensors.key_cache != nullptr); + REQUIRE(tensors.value_cache != nullptr); + REQUIRE(tensors.key_cache_indices != nullptr); + REQUIRE(tensors.value_cache_indices != nullptr); + + tensors.key_output = ggml_set_rows(ctx, tensors.key_cache, key_cache_rows, tensors.key_cache_indices); + tensors.value_output = ggml_set_rows(ctx, tensors.value_cache, value_cache_rows, tensors.value_cache_indices); + REQUIRE(tensors.key_output != nullptr); + REQUIRE(tensors.value_output != nullptr); + return tensors; +} + +static ggml::hrx::GraphImportResult import_qwen_attention_postprocess_graph( + ggml_context * ctx, + const QwenAttentionPostprocessTensors & tensors) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, tensors.query_output); + ggml_build_forward_expand(graph, tensors.key_output); + ggml_build_forward_expand(graph, tensors.value_output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + return imported; +} + +static ggml_tensor * append_qwen_flash_attention_consumer(ggml_context * ctx, + QwenAttentionPostprocessTensors & tensors, + int64_t token_count, + int64_t query_head_count, + int64_t key_value_head_count, + int64_t cache_row_count) { + ggml_tensor * query_layout = + ggml_reshape_3d(ctx, tensors.query_output, kQwenFlashHeadSize, query_head_count, token_count); + ggml_tensor * query_permute = ggml_permute(ctx, query_layout, 0, 2, 1, 3); + ggml_tensor * key_cache_layout = + ggml_reshape_3d(ctx, tensors.key_cache, kQwenFlashHeadSize, key_value_head_count, cache_row_count); + ggml_tensor * key_permute = ggml_permute(ctx, key_cache_layout, 0, 2, 1, 3); + ggml_tensor * value_cache_layout = + ggml_reshape_3d(ctx, tensors.value_cache, kQwenFlashHeadSize, key_value_head_count, cache_row_count); + ggml_tensor * value_permute = ggml_permute(ctx, value_cache_layout, 0, 2, 1, 3); + tensors.mask = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, cache_row_count, token_count); + REQUIRE(query_layout != nullptr); + REQUIRE(query_permute != nullptr); + REQUIRE(key_cache_layout != nullptr); + REQUIRE(key_permute != nullptr); + REQUIRE(value_cache_layout != nullptr); + REQUIRE(value_permute != nullptr); + REQUIRE(tensors.mask != nullptr); + + tensors.flash_output = ggml_flash_attn_ext(ctx, query_permute, key_permute, value_permute, tensors.mask, + 1.0f / std::sqrt(static_cast(kQwenFlashHeadSize)), 0.0f, 0.0f); + REQUIRE(tensors.flash_output != nullptr); + return tensors.flash_output; +} + +static void schedule_qwen_attention_postprocess_command(ggml_context * ctx, + const QwenAttentionPostprocessTensors & tensors, + int64_t token_count, + int64_t query_head_count, + int64_t key_value_head_count, + int64_t cache_row_count, + bool expect_synthetic_inverse_frequencies = false, + const std::string & expected_rope_mscale = "1") { + ggml::hrx::GraphImportResult imported = import_qwen_attention_postprocess_graph(ctx, tensors); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().initialization_dispatches.empty()); + const bool q8_input = token_count <= 5; + REQUIRE(scheduler.plan().dispatches.size() == (q8_input ? 5 : 4)); + if (q8_input) { + REQUIRE(kernel_name_for_id(scheduler.plan().dispatches.front().kernel.kernel_id) == + "qwen3_moe:ggml_quantize_q8_1_x4_f32"); + } + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches.back(); + const std::string kernel_name = kernel_name_for_id(dispatch.kernel.kernel_id); + REQUIRE(kernel_name == "qwen3_moe:qwen3_moe_attention_postprocess_f32_f16"); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == token_count); + REQUIRE(dispatch.kernel.integer_parameters.at("cache_row_count") == cache_row_count); + REQUIRE(dispatch.bindings.size() == 12); + require_compile_parameter(dispatch, "qwen3_moe.model.rms_epsilon", "0.000001"); + require_compile_parameter(dispatch, "qwen3_moe.attention.head_size", std::to_string(kQwenFlashHeadSize)); + require_compile_parameter(dispatch, "qwen3_moe.attention.rope_mscale", expected_rope_mscale); + require_compile_parameter(dispatch, "qwen3_moe.attention.query_size", + std::to_string(query_head_count * kQwenFlashHeadSize)); + require_compile_parameter(dispatch, "qwen3_moe.attention.key_value_size", + std::to_string(key_value_head_count * kQwenFlashHeadSize)); + require_compile_parameter(dispatch, "qwen3_moe.workload.token_capacity", std::to_string(token_count)); + + const ggml::hrx::Value * positions_value = imported.graph.values().find_tensor(tensors.positions); + const ggml::hrx::Value * key_indices_value = imported.graph.values().find_tensor(tensors.key_cache_indices); + const ggml::hrx::Value * value_indices_value = imported.graph.values().find_tensor(tensors.value_cache_indices); + const ggml::hrx::Value * query_raw_value = imported.graph.values().find_tensor(tensors.query_raw); + const ggml::hrx::Value * key_raw_value = imported.graph.values().find_tensor(tensors.key_raw); + const ggml::hrx::Value * value_raw_value = imported.graph.values().find_tensor(tensors.value_raw); + const ggml::hrx::Value * query_output_value = imported.graph.values().find_tensor(tensors.query_output); + const ggml::hrx::Value * key_cache_value = imported.graph.values().find_tensor(tensors.key_cache); + const ggml::hrx::Value * value_cache_value = imported.graph.values().find_tensor(tensors.value_cache); + REQUIRE(positions_value != nullptr); + REQUIRE(key_indices_value != nullptr); + REQUIRE(value_indices_value != nullptr); + REQUIRE(query_raw_value != nullptr); + REQUIRE(key_raw_value != nullptr); + REQUIRE(value_raw_value != nullptr); + REQUIRE(query_output_value != nullptr); + REQUIRE(key_cache_value != nullptr); + REQUIRE(value_cache_value != nullptr); + REQUIRE(dispatch.bindings[0].value == positions_value->id); + REQUIRE(dispatch.bindings[1].value == key_indices_value->id); + REQUIRE(dispatch.bindings[2].value == value_indices_value->id); + REQUIRE(dispatch.bindings[3].value == query_raw_value->id); + REQUIRE(dispatch.bindings[4].value == key_raw_value->id); + REQUIRE(dispatch.bindings[5].value == value_raw_value->id); + REQUIRE(dispatch.bindings[9].value == query_output_value->id); + REQUIRE(dispatch.bindings[10].value == key_cache_value->id); + REQUIRE(dispatch.bindings[11].value == value_cache_value->id); + if (expect_synthetic_inverse_frequencies) { + REQUIRE(scheduler.plan().transients.size() == (q8_input ? 2 : 1)); + REQUIRE(scheduler.plan().constant_initializations.size() == 1); + REQUIRE(dispatch.bindings[8].value == scheduler.plan().transients.back().value); + REQUIRE(scheduler.plan().constant_initializations[0].value == dispatch.bindings[8].value); + REQUIRE(scheduler.plan().constant_initializations[0].data.size() == + static_cast(kQwenFlashHeadSize / 2) * sizeof(float)); + } else { + REQUIRE(scheduler.plan().transients.size() == (q8_input ? 1 : 0)); + REQUIRE(scheduler.plan().constant_initializations.empty()); + } + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.initialization_commands.empty()); + REQUIRE(commands.commands.size() == (q8_input ? 5 : 4)); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.back().bindings.size() == 12); + REQUIRE(commands.commands.back().bindings[0].name == "positions"); + REQUIRE(commands.commands.back().bindings[1].name == "key_cache_indices"); + REQUIRE(commands.commands.back().bindings[2].name == "value_cache_indices"); + REQUIRE(commands.commands.back().bindings[3].name == "query_input"); + REQUIRE(commands.commands.back().bindings[4].name == "key_input"); + REQUIRE(commands.commands.back().bindings[5].name == "value_input"); + REQUIRE(commands.commands.back().bindings[6].name == "query_norm_weight"); + REQUIRE(commands.commands.back().bindings[7].name == "key_norm_weight"); + REQUIRE(commands.commands.back().bindings[8].name == "inverse_frequencies"); + REQUIRE(commands.commands.back().bindings[9].name == "query_output"); + REQUIRE(commands.commands.back().bindings[10].name == "key_cache"); + REQUIRE(commands.commands.back().bindings[11].name == "value_cache"); + REQUIRE(commands.constant_initializations.size() == (expect_synthetic_inverse_frequencies ? 1 : 0)); +} + +static void run_qwen_attention_postprocess_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 8 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + { + constexpr int64_t token_count = 4; + constexpr int64_t query_head_count = 4; + constexpr int64_t key_value_head_count = 2; + constexpr int64_t cache_row_count = 16; + const QwenAttentionPostprocessTensors tensors = build_qwen_attention_postprocess_graph( + ctx, token_count, query_head_count, key_value_head_count, cache_row_count); + schedule_qwen_attention_postprocess_command(ctx, tensors, token_count, query_head_count, key_value_head_count, + cache_row_count); + } + + { + const QwenAttentionPostprocessTensors tensors = + build_qwen_attention_postprocess_graph(ctx, 4, 4, 2, 16, 0.00001f); + ggml::hrx::GraphImportResult imported = import_qwen_attention_postprocess_graph(ctx, tensors); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + for (const ggml::hrx::Dispatch & dispatch : scheduler.plan().dispatches) { + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) != + "qwen3_moe:qwen3_moe_attention_postprocess_f32_f16"); + } + } + + { + const QwenAttentionPostprocessTensors tensors = + build_qwen_attention_postprocess_graph(ctx, 4, 4, 2, 16, 0.000001f, false); + schedule_qwen_attention_postprocess_command(ctx, tensors, 4, 4, 2, 16, true); + } + + { + const QwenAttentionPostprocessTensors tensors = build_qwen_attention_postprocess_graph( + ctx, 4, 4, 2, 16, 0.000001f, false, 8192, 1000000.0f, 0.25f, 1.0f, 32.0f, 1.0f); + schedule_qwen_attention_postprocess_command(ctx, tensors, 4, 4, 2, 16, true, "1.13862944"); + } + + { + constexpr int64_t token_count = 13; + constexpr int64_t query_head_count = 32; + constexpr int64_t key_value_head_count = 4; + constexpr int64_t cache_row_count = 512; + const QwenAttentionPostprocessTensors tensors = build_qwen_attention_postprocess_graph( + ctx, token_count, query_head_count, key_value_head_count, cache_row_count, 0.000001f, false); + schedule_qwen_attention_postprocess_command(ctx, tensors, token_count, query_head_count, key_value_head_count, + cache_row_count, true); + } + + { + constexpr int64_t token_count = 13; + constexpr int64_t query_head_count = 32; + constexpr int64_t key_value_head_count = 4; + constexpr int64_t cache_row_count = 512; + + const QwenAttentionPostprocessTensors first = build_qwen_attention_postprocess_graph( + ctx, token_count, query_head_count, key_value_head_count, cache_row_count, 0.000001f, false); + const QwenAttentionPostprocessTensors second = build_qwen_attention_postprocess_graph( + ctx, token_count, query_head_count, key_value_head_count, cache_row_count, 0.000001f, false); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, first.query_output); + ggml_build_forward_expand(graph, first.key_output); + ggml_build_forward_expand(graph, first.value_output); + ggml_build_forward_expand(graph, second.query_output); + ggml_build_forward_expand(graph, second.key_output); + ggml_build_forward_expand(graph, second.value_output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 8); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 8); + REQUIRE(command_program_verifies(commands)); + } + + { + constexpr int64_t token_count = 4; + constexpr int64_t query_head_count = 4; + constexpr int64_t key_value_head_count = 2; + constexpr int64_t cache_row_count = 16; + QwenAttentionPostprocessTensors tensors = build_qwen_attention_postprocess_graph( + ctx, token_count, query_head_count, key_value_head_count, cache_row_count); + ggml_tensor * flash_output = append_qwen_flash_attention_consumer(ctx, tensors, token_count, query_head_count, + key_value_head_count, cache_row_count); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, tensors.key_output); + ggml_build_forward_expand(graph, tensors.value_output); + ggml_build_forward_expand(graph, flash_output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan; + ggml::hrx::DispatchMatch match; + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, + producer_index_for_tensor(imported.graph, tensors.query_reshape), match)); + append_match_to_plan(plan, match, covered_nodes); + REQUIRE(plan.valid()); + REQUIRE(plan.initialization_dispatches.size() == 2); + + const ggml::hrx::Dispatch & context_capture = plan.initialization_dispatches[0]; + const ggml::hrx::Dispatch & metadata = plan.initialization_dispatches[1]; + REQUIRE(kernel_name_for_id(context_capture.kernel.kernel_id) == "qwen:qwen_attention_context_base_capture"); + REQUIRE(kernel_name_for_id(metadata.kernel.kernel_id) == "qwen3_moe:qwen_attention_metadata"); + REQUIRE(context_capture.bindings.size() == 2); + REQUIRE(metadata.bindings.size() == 5); + REQUIRE(context_capture.bindings[1].value == metadata.bindings[0].value); + REQUIRE(metadata.kernel.integer_parameters.at("token_count") == token_count); + REQUIRE(metadata.kernel.integer_parameters.at("context_capacity") == cache_row_count); + + const ggml::hrx::Value * positions_value = imported.graph.values().find_tensor(tensors.positions); + const ggml::hrx::Value * key_indices_value = imported.graph.values().find_tensor(tensors.key_cache_indices); + const ggml::hrx::Value * value_indices_value = imported.graph.values().find_tensor(tensors.value_cache_indices); + const ggml::hrx::Value * mask_value = imported.graph.values().find_tensor(tensors.mask); + REQUIRE(positions_value != nullptr); + REQUIRE(key_indices_value != nullptr); + REQUIRE(value_indices_value != nullptr); + REQUIRE(mask_value != nullptr); + REQUIRE(context_capture.bindings[0].value == positions_value->id); + REQUIRE(metadata.bindings[1].value == positions_value->id); + REQUIRE(metadata.bindings[2].value == key_indices_value->id); + REQUIRE(metadata.bindings[3].value == value_indices_value->id); + REQUIRE(metadata.bindings[4].value == mask_value->id); + + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.initialization_commands.size() == 2); + REQUIRE(commands.initialization_commands[0].bindings[0].name == "positions"); + REQUIRE(commands.initialization_commands[0].bindings[1].name == "control"); + REQUIRE(commands.initialization_commands[1].bindings[0].name == "control"); + REQUIRE(commands.initialization_commands[1].bindings[4].name == "attention_mask"); + REQUIRE(ggml::hrx::find_transient_allocation(commands.transients, context_capture.bindings[1].value) != + nullptr); + + ggml::hrx::CommandProgram initialization_commands = copy_command_program_shape(commands); + initialization_commands.commands.clear(); + REQUIRE(command_program_verifies(initialization_commands)); + } + + ggml_free(ctx); +} + +static void run_packed_f16_producer_consumer_checks() { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + const struct { + ggml_type type; + int64_t tokens; + int64_t hidden; + int64_t outputs; + bool side_use; + bool packed; + bool alternate; + bool packed_copy = false; + } cases[] = { + { GGML_TYPE_Q6_K, 512, 4096, 2048, false, true, false }, + { GGML_TYPE_Q6_K, 1024, 4096, 2048, false, true, false }, + { GGML_TYPE_Q6_K, 256, 4096, 2048, false, false, true }, + { GGML_TYPE_Q6_K, 512, 4096, 1024, false, false, true }, + { GGML_TYPE_Q4_K, 512, 4096, 4096, false, true, false }, + { GGML_TYPE_Q4_K, 512, 8192, 2048, false, false, true }, + { GGML_TYPE_Q4_K, 512,12288, 4096, false, true, false }, + { GGML_TYPE_Q4_K, 512,12288, 4096, true, false, true }, + { GGML_TYPE_Q4_K, 512,32768, 4096, false, false, false }, + { GGML_TYPE_Q4_K, 1024, 4096, 4096, false, false, true }, + { GGML_TYPE_Q6_K, 512, 4096, 2048, true, false, true, true }, + }; + + for (const auto & test : cases) { + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 256, test.tokens); + ggml_tensor * gate_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 256, test.hidden); + ggml_tensor * up_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 256, test.hidden); + ggml_tensor * gate = ggml_mul_mat(ctx, gate_weight, input); + ggml_tensor * up = ggml_mul_mat(ctx, up_weight, input); + ggml_tensor * glu = ggml_swiglu_split(ctx, gate, up); + ggml_tensor * weight = ggml_new_tensor_2d(ctx, test.type, test.hidden, test.outputs); + ggml_tensor * output = ggml_mul_mat(ctx, weight, glu); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, output); + if (test.side_use) { + ggml_build_forward_expand(graph, ggml_scale(ctx, glu, 0.5f)); + } + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const auto & plan = scheduler.plan(); + REQUIRE(plan.valid()); + const auto * value = imported.graph.values().find_tensor(glu); + REQUIRE(value != nullptr); + const auto * generated = + plan.metadata.find_generated_resource(value->id, ggml::hrx::GeneratedResourceRole::F16K16Major); + const auto * alternate = plan.metadata.find_alternate_value(value->id, GGML_TYPE_F16, value->byte_count / 2); + REQUIRE((generated != nullptr) == (test.packed || test.packed_copy)); + REQUIRE((alternate != nullptr) == test.alternate); + const auto producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == + "loom_libs:ggml_mul_mat_swiglu_q4_k_f16_wmma_prefill_wave32"; + }); + REQUIRE(producer != plan.dispatches.end()); + require_compile_parameter(*producer, "ggml.mul_mat_swiglu.f16_output_layout", test.packed ? "1" : "0"); + if (test.packed) { + REQUIRE(generated->byte_count == value->byte_count / 2); + REQUIRE(producer->bindings.back().value == generated->generated_value); + const auto consumer = std::next(producer); + REQUIRE(consumer != plan.dispatches.end()); + REQUIRE(consumer->bindings.front().value == generated->generated_value); + if (test.type == GGML_TYPE_Q4_K) { + require_compile_parameter(*consumer, "ggml.mul_mat.f16_input_layout", "1"); + } + } else if (test.packed_copy) { + const auto copy = std::next(producer); + REQUIRE(copy != plan.dispatches.end()); + REQUIRE(kernel_name_for_id(copy->kernel.kernel_id) == "loom_libs:ggml_copy_f16_k16_major"); + REQUIRE(copy->bindings.front().value == alternate->alternate_value); + REQUIRE(copy->bindings.back().value == generated->generated_value); + REQUIRE(std::next(copy) != plan.dispatches.end()); + REQUIRE(std::next(copy)->bindings.front().value == generated->generated_value); + } + const auto commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + + ggml_free(ctx); +} + +static void run_generic_swiglu_k16_publish_checks() { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t input_size = 2048; + constexpr int64_t hidden_size = 8192; + constexpr int64_t output_size = 2048; + constexpr int64_t token_count = 512; + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * gate_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, input_size, hidden_size); + ggml_tensor * up_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, input_size, hidden_size); + ggml_tensor * gate = ggml_mul_mat(ctx, gate_weight, input); + ggml_tensor * up = ggml_mul_mat(ctx, up_weight, input); + ggml_tensor * glu = ggml_swiglu_split(ctx, gate, up); + ggml_tensor * projection_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, hidden_size, output_size); + ggml_tensor * output = ggml_mul_mat(ctx, projection_weight, glu); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const ggml::hrx::CommandPlan & plan = scheduler.plan(); + REQUIRE(plan.valid()); + + const auto producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == + "loom_libs:ggml_mul_mat_tiled_pair_input_f32_binary_publish_f32_k16"; + }); + REQUIRE(producer != plan.dispatches.end()); + REQUIRE(producer->bindings.size() == 5); + const ggml::hrx::Value * glu_value = imported.graph.values().find_tensor(glu); + REQUIRE(glu_value != nullptr); + const auto * packed = + plan.metadata.find_generated_resource(glu_value->id, ggml::hrx::GeneratedResourceRole::F16K16Major); + REQUIRE(packed != nullptr); + REQUIRE(packed->byte_count == glu_value->byte_count / 2); + REQUIRE(producer->bindings.back().value == packed->generated_value); + const auto consumer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return &dispatch != &*producer && !dispatch.bindings.empty() && + dispatch.bindings.front().value == packed->generated_value; + }); + REQUIRE(consumer != plan.dispatches.end()); + REQUIRE(std::none_of(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_copy_f16_k16_major"; + })); + + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + ggml_free(ctx); +} + +static void run_generic_swiglu_k16_publication_edge_checks() { + { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t input_size = 2048; + constexpr int64_t hidden_size = 8192; + constexpr int64_t token_count = 512; + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * gate_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, input_size, hidden_size); + ggml_tensor * up_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, input_size, hidden_size); + ggml_tensor * gate = ggml_mul_mat(ctx, gate_weight, input); + ggml_tensor * up = ggml_mul_mat(ctx, up_weight, input); + ggml_tensor * glu = ggml_swiglu_split(ctx, gate, up); + REQUIRE(glu != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, glu); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const ggml::hrx::CommandPlan & plan = scheduler.plan(); + REQUIRE(plan.valid()); + + const ggml::hrx::Value * glu_value = imported.graph.values().find_tensor(glu); + REQUIRE(glu_value != nullptr); + REQUIRE(plan.metadata.find_generated_resource(glu_value->id, ggml::hrx::GeneratedResourceRole::F16K16Major) == + nullptr); + const auto producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == + "loom_libs:ggml_mul_mat_tiled_pair_input_f32_binary_publish_f32"; + }); + REQUIRE(producer != plan.dispatches.end()); + REQUIRE(producer->bindings.size() == 4); + REQUIRE(producer->bindings.back().value == glu_value->id); + + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + ggml_free(ctx); + } + + { + ggml_init_params params = {}; + params.mem_size = 8 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t input_size = 2048; + constexpr int64_t hidden_size = 8192; + constexpr int64_t token_count = 512; + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * gate_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, input_size, hidden_size); + ggml_tensor * up_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, input_size, hidden_size); + ggml_tensor * gate = ggml_mul_mat(ctx, gate_weight, input); + ggml_tensor * up = ggml_mul_mat(ctx, up_weight, input); + ggml_tensor * glu = ggml_swiglu_split(ctx, gate, up); + ggml_tensor * reshaped = ggml_reshape_2d(ctx, glu, 4096, 1024); + ggml_tensor * projection_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, 4096, 2048); + ggml_tensor * output = ggml_mul_mat(ctx, projection_weight, reshaped); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const ggml::hrx::CommandPlan & plan = scheduler.plan(); + REQUIRE(plan.valid()); + + const ggml::hrx::Value * glu_value = imported.graph.values().find_tensor(glu); + const ggml::hrx::Value * reshaped_value = imported.graph.values().find_tensor(reshaped); + REQUIRE(glu_value != nullptr); + REQUIRE(reshaped_value != nullptr); + REQUIRE(plan.metadata.find_generated_resource(glu_value->id, ggml::hrx::GeneratedResourceRole::F16K16Major) == + nullptr); + const auto * packed = + plan.metadata.find_generated_resource(reshaped_value->id, ggml::hrx::GeneratedResourceRole::F16K16Major); + REQUIRE(packed != nullptr); + + const auto producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == + "loom_libs:ggml_mul_mat_tiled_pair_input_f32_binary_publish_f32"; + }); + REQUIRE(producer != plan.dispatches.end()); + REQUIRE(producer->bindings.size() == 4); + const auto copy = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_copy_f16_k16_major" && + dispatch.bindings.front().value == reshaped_value->id; + }); + REQUIRE(copy != plan.dispatches.end()); + REQUIRE(copy->bindings.back().value == packed->generated_value); + require_compile_parameter(*copy, "ggml.copy_f16_k16_major.input_is_f16", "0"); + const auto consumer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return &dispatch != &*copy && !dispatch.bindings.empty() && + dispatch.bindings.front().value == packed->generated_value; + }); + REQUIRE(consumer != plan.dispatches.end()); + + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + ggml_free(ctx); + } +} + +static void run_tiled_matmul_alternate_publish_checks() { + struct Case { + ggml_type consumer_type; + const char * producer_kernel; + const char * consumer_kernel; + ggml_type alternate_type; + }; + + const Case cases[] = { + { GGML_TYPE_Q6_K, "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32_f16_alternate", + "loom_libs:ggml_mul_mat_q6_k_f16_wmma_prefill_wave32", GGML_TYPE_F16 }, + { GGML_TYPE_Q5_K, "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32", + "loom_libs:ggml_mul_mat_q5_k_iq4_xs_q8_1_x4_wmma_token256", GGML_TYPE_Q8_1 }, + }; + + for (const Case & test : cases) { + ggml_init_params params = {}; + params.mem_size = 8 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 256); + ggml_tensor * producer_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 2048, 2048); + ggml_tensor * projection = ggml_mul_mat(ctx, producer_weight, input); + ggml_tensor * consumer_weight = ggml_new_tensor_2d(ctx, test.consumer_type, 2048, 1024); + ggml_tensor * output = ggml_mul_mat(ctx, consumer_weight, projection); + + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, output); + auto imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const auto & plan = scheduler.plan(); + REQUIRE(plan.valid()); + + const auto * projection_value = imported.graph.values().find_tensor(projection); + REQUIRE(projection_value != nullptr); + const size_t alternate_bytes = test.alternate_type == GGML_TYPE_F16 ? + projection_value->byte_count / 2 : + static_cast(256) * ggml_row_size(GGML_TYPE_Q8_1, 2048); + const auto * alternate = + plan.metadata.find_alternate_value(projection_value->id, test.alternate_type, alternate_bytes); + REQUIRE(alternate != nullptr); + + const auto producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == test.producer_kernel; + }); + REQUIRE(producer != plan.dispatches.end()); + REQUIRE(producer->bindings[2].value == projection_value->id); + if (test.alternate_type == GGML_TYPE_F16) { + REQUIRE(producer->bindings.size() == 4); + REQUIRE(producer->bindings[3].value == alternate->alternate_value); + } else { + const auto quantize = + std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "qwen3_moe:ggml_quantize_q8_1_x4_f32" && + dispatch.bindings.front().value == projection_value->id; + }); + REQUIRE(quantize != plan.dispatches.end()); + REQUIRE(quantize->bindings.back().value == alternate->alternate_value); + } + + const auto consumer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == test.consumer_kernel; + }); + REQUIRE(consumer != plan.dispatches.end()); + REQUIRE(consumer->bindings.front().value == alternate->alternate_value); + + const auto commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + ggml_free(ctx); + } + + { + ggml_init_params params = {}; + params.mem_size = 12 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 256); + ggml_tensor * producer_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 2048, 2048); + ggml_tensor * projection = ggml_mul_mat(ctx, producer_weight, input); + ggml_tensor * q5_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q5_K, 2048, 1024); + ggml_tensor * q6_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, 2048, 1024); + ggml_tensor * q5_output = ggml_mul_mat(ctx, q5_weight, projection); + ggml_tensor * q6_output = ggml_mul_mat(ctx, q6_weight, projection); + + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, q5_output); + ggml_build_forward_expand(graph, q6_output); + auto imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const auto & plan = scheduler.plan(); + REQUIRE(plan.valid()); + + const auto * projection_value = imported.graph.values().find_tensor(projection); + REQUIRE(projection_value != nullptr); + const size_t q8_bytes = static_cast(256) * ggml_row_size(GGML_TYPE_Q8_1, 2048); + const auto * q8 = plan.metadata.find_alternate_value(projection_value->id, GGML_TYPE_Q8_1, q8_bytes); + REQUIRE(q8 != nullptr); + const auto publication_count = std::count_if( + plan.metadata.activation_publication_diagnostics().begin(), + plan.metadata.activation_publication_diagnostics().end(), [&](const auto & diagnostic) { + return diagnostic.source_value == projection_value->id; + }); + REQUIRE(publication_count == 1); + const auto publication = std::find_if( + plan.metadata.activation_publication_diagnostics().begin(), + plan.metadata.activation_publication_diagnostics().end(), [&](const auto & diagnostic) { + return diagnostic.source_value == projection_value->id; + }); + REQUIRE(publication != plan.metadata.activation_publication_diagnostics().end()); + REQUIRE(publication->requested_format == "q8-1-x4"); + + const auto producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32" && + dispatch.bindings[2].value == projection_value->id; + }); + REQUIRE(producer != plan.dispatches.end()); + REQUIRE(producer->bindings.size() == 3); + + const auto q5_consumer = + std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == + "loom_libs:ggml_mul_mat_q5_k_iq4_xs_q8_1_x4_wmma_token256" && + dispatch.bindings.front().value == q8->alternate_value; + }); + REQUIRE(q5_consumer != plan.dispatches.end()); + const auto q6_consumer = + std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == + "loom_libs:ggml_mul_mat_q6_k_f16_wmma_prefill_wave32"; + }); + REQUIRE(q6_consumer != plan.dispatches.end()); + const auto f16_copy = + std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_copy_f32_f16" && + dispatch.bindings.front().value == projection_value->id; + }); + REQUIRE(f16_copy != plan.dispatches.end()); + REQUIRE(q6_consumer->bindings.front().value == f16_copy->bindings.back().value); + + const size_t q8_quantizer_count = std::count_if( + plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "qwen3_moe:ggml_quantize_q8_1_x4_f32" && + dispatch.bindings.front().value == projection_value->id; + }); + REQUIRE(q8_quantizer_count == 1); + + const auto commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + ggml_free(ctx); + } +} + +static void run_tiled_pair_matmul_postops_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 16 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + auto build_glu = [&]() { + ggml_tensor * gate_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ4_NL, 640, 256); + ggml_tensor * up_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ4_NL, 640, 256); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 640, 8); + REQUIRE(gate_weight != nullptr); + REQUIRE(up_weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * gate = ggml_mul_mat(ctx, gate_weight, input); + ggml_tensor * up = ggml_mul_mat(ctx, up_weight, input); + REQUIRE(gate != nullptr); + REQUIRE(up != nullptr); + ggml_tensor * glu = ggml_glu_split(ctx, gate, up, GGML_GLU_OP_SWIGLU); + REQUIRE(glu != nullptr); + return glu; + }; + + { + ggml_tensor * bias = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 256); + REQUIRE(bias != nullptr); + ggml_tensor * output = ggml_add(ctx, build_glu(), bias); + REQUIRE(output != nullptr); + schedule_fused_matmul_swiglu_postops_command( + ctx, output, "loom_libs:ggml_mul_mat_tiled_pair_input_f32_binary_bias_residual_publish_f32", 4, + { GGML_OP_MUL_MAT, GGML_OP_MUL_MAT, GGML_OP_GLU, GGML_OP_ADD }, + { "input", "gate_weight", "up_weight", "bias", "residual_input", "residual_output" }); + } + + { + ggml_tensor * residual = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 256, 8); + REQUIRE(residual != nullptr); + ggml_tensor * output = ggml_add(ctx, build_glu(), residual); + REQUIRE(output != nullptr); + schedule_fused_matmul_swiglu_postops_command( + ctx, output, "loom_libs:ggml_mul_mat_tiled_pair_input_f32_binary_bias_residual_publish_f32", 4, + { GGML_OP_MUL_MAT, GGML_OP_MUL_MAT, GGML_OP_GLU, GGML_OP_ADD }, + { "input", "gate_weight", "up_weight", "bias", "residual_input", "residual_output" }); + } + + { + ggml_tensor * bias = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 256); + ggml_tensor * residual = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 256, 8); + REQUIRE(bias != nullptr); + REQUIRE(residual != nullptr); + ggml_tensor * biased = ggml_add(ctx, build_glu(), bias); + REQUIRE(biased != nullptr); + ggml_tensor * output = ggml_add(ctx, biased, residual); + REQUIRE(output != nullptr); + schedule_fused_matmul_swiglu_postops_command( + ctx, output, "loom_libs:ggml_mul_mat_tiled_pair_input_f32_binary_bias_residual_publish_f32", 5, + { GGML_OP_MUL_MAT, GGML_OP_MUL_MAT, GGML_OP_GLU, GGML_OP_ADD, GGML_OP_ADD }, + { "input", "gate_weight", "up_weight", "bias", "residual_input", "residual_output" }); + } + + ggml_free(ctx); +} + +static void run_swiglu_q8_output_checks() { + const struct { + int64_t tokens; + int64_t inputs; + int64_t outputs; + ggml_type down_type; + int projections; + bool terminal; + bool side_use; + bool observed; + bool geglu; + bool weight_alias; + bool packed; + } cases[] = { + { 2, 5120, 17408, GGML_TYPE_Q4_K, 1, false, false, false, false, false, true }, + { 3, 5120, 17408, GGML_TYPE_Q6_K, 1, false, false, false, false, false, true }, + { 4, 5120, 17408, GGML_TYPE_Q4_K, 1, false, false, false, false, false, true }, + { 5, 5120, 17408, GGML_TYPE_Q6_K, 2, false, true, true, false, false, true }, + { 5, 4096, 16384, GGML_TYPE_Q4_K, 1, true, false, false, false, false, true }, + { 5, 5120, 32768, GGML_TYPE_Q6_K, 1, false, false, false, false, false, true }, + { 1, 5120, 17408, GGML_TYPE_Q4_K, 1, false, false, false, false, false, true }, + { 6, 5120, 17408, GGML_TYPE_Q4_K, 1, false, false, false, false, false, false }, + { 5, 2048, 17408, GGML_TYPE_Q4_K, 1, false, false, false, false, false, true }, + { 5, 6144, 17408, GGML_TYPE_Q4_K, 1, false, false, false, false, false, true }, + { 5, 5120, 8192, GGML_TYPE_Q4_K, 1, false, false, false, false, false, true }, + { 5, 5120, 33024, GGML_TYPE_Q4_K, 1, false, false, false, false, false, false }, + { 5, 5120, 17408, GGML_TYPE_F32, 1, false, true, false, false, false, false }, + { 5, 5120, 17408, GGML_TYPE_Q6_K, 1, true, false, false, false, false, true }, + { 5, 5120, 17408, GGML_TYPE_Q4_K, 0, false, true, true, false, false, false }, + { 5, 5120, 17408, GGML_TYPE_Q4_K, 1, false, false, false, true, false, false }, + { 5, 5120, 17408, GGML_TYPE_Q4_K, 1, false, false, false, false, true, false }, + }; + + for (const auto & test : cases) { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, test.inputs, test.tokens); + ggml_tensor * gate_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, test.inputs, test.outputs); + ggml_tensor * up_weight = test.weight_alias ? + ggml_view_tensor(ctx, gate_weight) : + ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, test.inputs, test.outputs); + ggml_tensor * gate = ggml_mul_mat(ctx, gate_weight, input); + ggml_tensor * up = ggml_mul_mat(ctx, up_weight, input); + ggml_tensor * glu = ggml_glu_split(ctx, gate, up, test.geglu ? GGML_GLU_OP_GEGLU : GGML_GLU_OP_SWIGLU); + if (test.observed) { + ggml_set_output(glu); + } + ggml_cgraph * graph = ggml_new_graph(ctx); + for (int i = 0; i < test.projections; ++i) { + ggml_tensor * down_weight = ggml_new_tensor_2d(ctx, test.down_type, test.outputs, 5120); + ggml_tensor * down = ggml_mul_mat(ctx, down_weight, glu); + ggml_build_forward_expand(graph, test.terminal ? down : ggml_add(ctx, down, ggml_dup_tensor(ctx, down))); + } + ggml_tensor * side = test.side_use ? ggml_scale(ctx, glu, 0.5f) : nullptr; + if (side != nullptr) { + ggml_build_forward_expand(graph, side); + } + auto imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + const bool scheduled = scheduler.schedule_graph(imported.graph, test_dispatch_target()); + const auto & plan = scheduler.plan(); + const auto producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == + "loom_libs:ggml_mul_mat_swiglu_q4_q8_1_x4_lowtoken_dot_q8_output"; + }); + REQUIRE((producer != plan.dispatches.end()) == test.packed); + if (test.outputs > 32768) { + REQUIRE(!scheduled); + REQUIRE(status_contains(plan.status, "unsupported HRX node")); + ggml_free(ctx); + continue; + } + REQUIRE(scheduled); + REQUIRE(plan.valid()); + const auto * output = imported.graph.values().find_tensor(glu); + REQUIRE(output != nullptr); + if (test.packed) { + const size_t bytes = static_cast(test.tokens) * ggml_row_size(GGML_TYPE_Q8_1, test.outputs); + const auto * alternate = plan.metadata.find_alternate_value(output->id, GGML_TYPE_Q8_1, bytes); + REQUIRE(alternate != nullptr); + REQUIRE(producer->bindings.size() == 5); + REQUIRE(producer->bindings[3].value == output->id); + REQUIRE(producer->bindings[4].value == alternate->alternate_value); + REQUIRE(std::count_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "qwen3_moe:ggml_quantize_q8_1_x4_f32" && + dispatch.bindings.front().value == output->id; + }) == 0); + REQUIRE(std::count_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return &dispatch != &*producer && !dispatch.bindings.empty() && + dispatch.bindings.front().value == alternate->alternate_value; + }) == test.projections); + } + if (side != nullptr) { + const auto * side_value = imported.graph.values().find_tensor(side); + REQUIRE(side_value != nullptr); + const auto consumer = + std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return !dispatch.bindings.empty() && dispatch.bindings.back().value == side_value->id; + }); + REQUIRE(consumer != plan.dispatches.end()); + REQUIRE(consumer->bindings.front().value == output->id); + } + const auto commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + ggml_free(ctx); + } +} + +static void run_k16_major_preparation_reuse_checks() { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + const struct { + ggml_type first_type; + ggml_type second_type; + bool shared; + bool reshape; + } cases[] = { + { GGML_TYPE_Q4_K, GGML_TYPE_Q4_K, true, false }, + { GGML_TYPE_Q4_K, GGML_TYPE_Q6_K, true, false }, + { GGML_TYPE_Q6_K, GGML_TYPE_Q4_K, true, false }, + { GGML_TYPE_Q6_K, GGML_TYPE_Q6_K, true, false }, + { GGML_TYPE_Q4_K, GGML_TYPE_Q4_K, false, false }, + { GGML_TYPE_Q6_K, GGML_TYPE_Q6_K, false, false }, + { GGML_TYPE_Q4_K, GGML_TYPE_Q6_K, false, true }, + { GGML_TYPE_Q6_K, GGML_TYPE_Q6_K, false, true }, + }; + + for (const auto & test : cases) { + ggml_tensor * first = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 4096, 512); + ggml_tensor * second = test.reshape ? ggml_reshape_2d(ctx, first, 2048, 1024) : + test.shared ? first : + ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 4096, 512); + ggml_tensor * inputs[] = { first, second }; + ggml_tensor * outputs[2]; + ggml_cgraph * graph = ggml_new_graph(ctx); + for (int i = 0; i < 2; ++i) { + const ggml_type type = i == 0 ? test.first_type : test.second_type; + ggml_tensor * weight = ggml_new_tensor_2d(ctx, type, inputs[i]->ne[0], 4096); + outputs[i] = ggml_mul_mat(ctx, weight, inputs[i]); + ggml_build_forward_expand(graph, outputs[i]); + } + auto imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const auto & plan = scheduler.plan(); + REQUIRE(plan.valid()); + const auto copy_count = + std::count_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_copy_f16_k16_major"; + }); + REQUIRE(copy_count == (test.shared ? 1 : 2)); + ggml::hrx::ValueId packed[2]; + for (int i = 0; i < 2; ++i) { + const auto * input = imported.graph.values().find_tensor(inputs[i]); + const auto * output = imported.graph.values().find_tensor(outputs[i]); + REQUIRE(input != nullptr); + REQUIRE(output != nullptr); + const auto * generated = + plan.metadata.find_generated_resource(input->id, ggml::hrx::GeneratedResourceRole::F16K16Major); + REQUIRE(generated != nullptr); + REQUIRE(generated->byte_count == input->byte_count / 2); + packed[i] = generated->generated_value; + const auto consumer = + std::find_if(plan.dispatches.begin(), plan.dispatches.end(), + [&](const auto & dispatch) { return dispatch.bindings.back().value == output->id; }); + REQUIRE(consumer != plan.dispatches.end()); + REQUIRE(consumer->bindings.front().value == packed[i]); + } + REQUIRE((packed[0] == packed[1]) == test.shared); + const auto commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + ggml_free(ctx); +} + +static void run_rmsnorm_binary_k16_output_checks() { + const struct { + ggml_type type; + int64_t hidden; + int64_t tokens; + int projections; + bool side_use; + bool observed; + bool inplace; + bool multiply; + bool packed; + } cases[] = { + { GGML_TYPE_Q4_K, 2048, 512, 1, false, false, false, true, true }, + { GGML_TYPE_Q4_K, 3072, 512, 1, false, false, false, true, true }, + { GGML_TYPE_Q4_K, 4096, 512, 1, false, false, false, true, true }, + { GGML_TYPE_Q4_K, 5120, 512, 2, false, false, false, true, true }, + { GGML_TYPE_Q6_K, 8192, 512, 2, true, false, false, true, true }, + { GGML_TYPE_Q6_K, 5120, 1024, 1, false, false, false, true, true }, + { GGML_TYPE_Q4_K, 5120, 512, 2, true, true, false, true, true }, + { GGML_TYPE_Q4_K, 6144, 512, 1, false, false, false, true, true }, + { GGML_TYPE_Q4_K, 7168, 512, 1, false, false, false, true, false }, + { GGML_TYPE_Q4_K, 4096, 256, 1, false, false, false, true, false }, + { GGML_TYPE_Q4_K, 5120, 5, 1, false, false, false, true, false }, + { GGML_TYPE_Q4_K, 4096, 1024, 1, false, false, false, true, false }, + { GGML_TYPE_Q4_K, 5120, 512, 0, true, true, false, true, false }, + { GGML_TYPE_F32, 5120, 512, 1, true, false, false, true, false }, + { GGML_TYPE_Q4_K, 5120, 512, 1, false, false, true, true, false }, + { GGML_TYPE_Q4_K, 5120, 512, 1, false, false, false, false, false }, + }; + + for (const auto & test : cases) { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, test.hidden, test.tokens); + ggml_tensor * weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, test.hidden); + ggml_tensor * rms = + test.inplace ? ggml_rms_norm_inplace(ctx, input, 1.0e-6f) : ggml_rms_norm(ctx, input, 1.0e-6f); + ggml_tensor * norm = test.inplace ? ggml_mul_inplace(ctx, rms, weight) : + test.multiply ? ggml_mul(ctx, rms, weight) : + ggml_add(ctx, rms, weight); + if (test.observed) { + ggml_set_output(norm); + } + ggml_cgraph * graph = ggml_new_graph(ctx); + for (int i = 0; i < test.projections; ++i) { + ggml_tensor * projection = ggml_new_tensor_2d(ctx, test.type, test.hidden, 4096); + ggml_build_forward_expand(graph, ggml_mul_mat(ctx, projection, norm)); + } + ggml_tensor * side = test.side_use ? ggml_scale(ctx, norm, 0.5f) : nullptr; + if (side != nullptr) { + ggml_build_forward_expand(graph, side); + } + auto imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const auto & plan = scheduler.plan(); + REQUIRE(plan.valid()); + const auto producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rmsnorm_binary_f32_k16"; + }); + REQUIRE((producer != plan.dispatches.end()) == test.packed); + const auto * output = imported.graph.values().find_tensor(norm); + REQUIRE(output != nullptr); + if (test.packed) { + const auto * generated = + plan.metadata.find_generated_resource(output->id, ggml::hrx::GeneratedResourceRole::F16K16Major); + REQUIRE(generated != nullptr); + REQUIRE(generated->byte_count == output->byte_count / 2); + REQUIRE(producer->bindings.size() == 4); + REQUIRE(producer->bindings[2].value == output->id); + REQUIRE(producer->bindings[3].value == generated->generated_value); + REQUIRE(plan.metadata.find_alternate_value(output->id, GGML_TYPE_F16, output->byte_count / 2) == nullptr); + REQUIRE(std::count_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_copy_f16_k16_major"; + }) == 0); + REQUIRE(std::count_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return &dispatch != &*producer && !dispatch.bindings.empty() && + dispatch.bindings.front().value == generated->generated_value; + }) == test.projections); + } else if (test.hidden == 7168 && test.tokens == 512) { + const auto * f16 = + plan.metadata.find_alternate_value(output->id, GGML_TYPE_F16, output->byte_count / 2); + REQUIRE(f16 != nullptr); + REQUIRE(std::count_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_copy_f16_k16_major"; + }) == 1); + } else if (test.projections == 0) { + REQUIRE(plan.metadata.find_generated_resource( + output->id, ggml::hrx::GeneratedResourceRole::F16K16Major) == nullptr); + REQUIRE(plan.metadata.find_alternate_value(output->id) == nullptr); + } + if (side != nullptr) { + const auto * side_value = imported.graph.values().find_tensor(side); + REQUIRE(side_value != nullptr); + const auto consumer = + std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return !dispatch.bindings.empty() && dispatch.bindings.back().value == side_value->id; + }); + REQUIRE(consumer != plan.dispatches.end()); + REQUIRE(consumer->bindings.front().value == output->id); + } + const auto commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + ggml_free(ctx); + } +} + +static void run_rmsnorm_binary_f16_output_checks() { + { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 7168, 512); + ggml_tensor * norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 7168); + ggml_tensor * rms = ggml_rms_norm(ctx, input, 1.0e-6f); + ggml_tensor * norm = ggml_mul(ctx, rms, norm_weight); + ggml_tensor * projection_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 7168, 2048); + ggml_tensor * output = ggml_mul_mat(ctx, projection_weight, norm); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const ggml::hrx::CommandPlan & plan = scheduler.plan(); + REQUIRE(plan.valid()); + const auto producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rmsnorm_binary_f32_f16"; + }); + REQUIRE(producer != plan.dispatches.end()); + REQUIRE(producer->bindings.size() == 4); + const ggml::hrx::Value * norm_value = imported.graph.values().find_tensor(norm); + REQUIRE(norm_value != nullptr); + const auto * f16 = + plan.metadata.find_alternate_value(norm_value->id, GGML_TYPE_F16, norm_value->byte_count / 2); + REQUIRE(f16 != nullptr); + REQUIRE(producer->bindings.back().value == f16->alternate_value); + const auto consumer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return &dispatch != &*producer && !dispatch.bindings.empty() && + dispatch.bindings.front().value == f16->alternate_value; + }); + REQUIRE(consumer != plan.dispatches.end()); + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + ggml_free(ctx); + } + + { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 512); + ggml_tensor * norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 2048); + ggml_tensor * rms = ggml_rms_norm(ctx, input, 1.0e-6f); + ggml_tensor * norm = ggml_mul(ctx, rms, norm_weight); + ggml_tensor * q8_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q5_K, 2048, 1024); + ggml_tensor * f16_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 2048); + ggml_tensor * q8_output = ggml_mul_mat(ctx, q8_weight, norm); + ggml_tensor * f16_output = ggml_mul_mat(ctx, f16_weight, norm); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, q8_output); + ggml_build_forward_expand(graph, f16_output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const ggml::hrx::CommandPlan & plan = scheduler.plan(); + REQUIRE(plan.valid()); + const auto producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rmsnorm_binary_q8_1_x4_publish"; + }); + REQUIRE(producer != plan.dispatches.end()); + require_compile_parameter(*producer, "ggml.rmsnorm_binary_q8_1_x4.publish_f16", "1"); + REQUIRE(producer->bindings.size() == 5); + const ggml::hrx::Value * norm_value = imported.graph.values().find_tensor(norm); + REQUIRE(norm_value != nullptr); + const auto * q8 = + plan.metadata.find_alternate_value(norm_value->id, GGML_TYPE_Q8_1, qwen_q8_1_x4_size(512, 2048)); + const auto * f16 = + plan.metadata.find_alternate_value(norm_value->id, GGML_TYPE_F16, norm_value->byte_count / 2); + REQUIRE(q8 != nullptr); + REQUIRE(f16 != nullptr); + REQUIRE(producer->bindings[3].value == q8->alternate_value); + REQUIRE(producer->bindings[4].value == f16->alternate_value); + REQUIRE(std::any_of(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return &dispatch != &*producer && !dispatch.bindings.empty() && + dispatch.bindings.front().value == q8->alternate_value; + })); + REQUIRE(std::any_of(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return &dispatch != &*producer && !dispatch.bindings.empty() && + dispatch.bindings.front().value == f16->alternate_value; + })); + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + ggml_free(ctx); + } + + { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t hidden = 2048; + constexpr int64_t tokens = 2; + ggml_tensor * input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, hidden, 1, tokens); + ggml_tensor * norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, hidden); + ggml_tensor * norm = ggml_mul(ctx, ggml_rms_norm(ctx, input, 1.0e-5f), norm_weight); + ggml_tensor * reshaped = ggml_reshape_2d(ctx, norm, hidden, tokens); + ggml_tensor * projection = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, hidden, 6144); + ggml_tensor * output = ggml_silu(ctx, ggml_mul_mat(ctx, projection, reshaped)); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const ggml::hrx::CommandPlan & plan = scheduler.plan(); + REQUIRE(plan.valid()); + const auto producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rmsnorm_binary_q8_1_x4_publish"; + }); + REQUIRE(producer != plan.dispatches.end()); + require_compile_parameter(*producer, "ggml.rmsnorm_binary_q8_1_x4.publish_f16", "1"); + REQUIRE(producer->bindings.size() == 5); + const ggml::hrx::Value * norm_value = imported.graph.values().find_tensor(norm); + REQUIRE(norm_value != nullptr); + const size_t q8_bytes = qwen_q8_1_x4_size(tokens, hidden); + const auto * q8 = plan.metadata.find_alternate_value(norm_value->id, GGML_TYPE_Q8_1, q8_bytes); + REQUIRE(q8 != nullptr); + REQUIRE(producer->bindings[3].value == q8->alternate_value); + REQUIRE(std::count_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "qwen3_moe:ggml_quantize_q8_1_x4_f32"; + }) == 0); + REQUIRE(std::any_of(std::next(producer), plan.dispatches.end(), [&](const auto & dispatch) { + return !dispatch.bindings.empty() && dispatch.bindings.front().value == q8->alternate_value; + })); + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + ggml_free(ctx); + } +} + +static void run_rmsnorm_gate_packed_output_checks() { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + const struct { + ggml_type type; + int64_t hidden; + int64_t channels; + int64_t tokens; + bool side_use; + bool packed; + int copies; + } cases[] = { + { GGML_TYPE_Q4_K, 128, 6144, 512, false, true, 0 }, + { GGML_TYPE_Q6_K, 256, 4096, 1024, false, true, 0 }, + { GGML_TYPE_Q6_K, 1024, 4096, 512, false, true, 0 }, + { GGML_TYPE_Q4_K, 128, 6144, 512, true, false, 1 }, + { GGML_TYPE_Q4_K, 128, 4096, 256, false, false, 0 }, + { GGML_TYPE_Q4_K, 256, 256, 512, false, true, 0 }, + { GGML_TYPE_Q5_K, 128, 6144, 19, false, false, 0 }, + { GGML_TYPE_IQ4_XS, 128, 6144, 19, false, false, 0 }, + }; + + for (const auto & test : cases) { + const int64_t heads = test.channels / test.hidden; + ggml_tensor * input = heads == 1 ? ggml_new_tensor_2d(ctx, GGML_TYPE_F32, test.hidden, test.tokens) : + ggml_new_tensor_3d(ctx, GGML_TYPE_F32, test.hidden, heads, test.tokens); + ggml_tensor * weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, test.hidden); + ggml_tensor * gate = ggml_dup_tensor(ctx, input); + ggml_tensor * norm = ggml_mul(ctx, ggml_rms_norm(ctx, input, 1.0e-6f), weight); + ggml_tensor * gated = ggml_mul(ctx, norm, ggml_silu(ctx, gate)); + ggml_tensor * reshaped = heads == 1 ? gated : ggml_reshape_2d(ctx, gated, test.channels, test.tokens); + ggml_tensor * projection = ggml_new_tensor_2d(ctx, test.type, test.channels, 4096); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, ggml_mul_mat(ctx, projection, reshaped)); + if (test.side_use) { + ggml_build_forward_expand(graph, ggml_scale(ctx, gated, 0.5f)); + } + auto imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const auto & plan = scheduler.plan(); + REQUIRE(plan.valid()); + const auto producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rmsnorm_gate_f32_publish"; + }); + REQUIRE(producer != plan.dispatches.end()); + require_compile_parameter(*producer, "ggml.rmsnorm_gate_f32.publish_q8", "0"); + require_compile_parameter(*producer, "ggml.rmsnorm_gate_f32.f16_output_row_width", + std::to_string(test.packed ? test.channels : 0)); + const auto * output = imported.graph.values().find_tensor(gated); + const auto * matmul_input = imported.graph.values().find_tensor(reshaped); + REQUIRE(output != nullptr && matmul_input != nullptr); + const auto * alternate = plan.metadata.find_alternate_value(output->id, GGML_TYPE_F16, output->byte_count / 2); + REQUIRE((alternate == nullptr) == test.packed); + if (test.packed) { + const auto * generated = + plan.metadata.find_generated_resource(matmul_input->id, ggml::hrx::GeneratedResourceRole::F16K16Major); + REQUIRE(generated != nullptr); + REQUIRE(producer->bindings.back().value == generated->generated_value); + REQUIRE(std::next(producer) != plan.dispatches.end()); + REQUIRE(std::next(producer)->bindings.front().value == generated->generated_value); + } + REQUIRE(std::count_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_copy_f16_k16_major"; + }) == test.copies); + const auto commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + + { + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 128, 2); + ggml_tensor * weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 128); + ggml_tensor * gate = ggml_dup_tensor(ctx, input); + ggml_tensor * norm = ggml_mul(ctx, ggml_rms_norm(ctx, input, 1.0e-6f), weight); + ggml_tensor * gated = ggml_mul(ctx, norm, ggml_silu(ctx, gate)); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, ggml_scale(ctx, gated, 0.5f)); + + auto imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const auto & plan = scheduler.plan(); + REQUIRE(plan.valid()); + const auto producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rmsnorm_gate_f32_publish"; + }); + REQUIRE(producer != plan.dispatches.end()); + require_compile_parameter(*producer, "ggml.rmsnorm_gate_f32.publish_q8", "0"); + REQUIRE(producer->bindings.size() == 5); + const auto * output = imported.graph.values().find_tensor(gated); + REQUIRE(output != nullptr); + REQUIRE(plan.metadata.find_alternate_value(output->id) == nullptr); + REQUIRE(plan.metadata.find_generated_resource( + output->id, ggml::hrx::GeneratedResourceRole::F16K16Major) == nullptr); + REQUIRE(std::any_of(plan.transients.begin(), plan.transients.end(), [&](const auto & transient) { + return transient.value == producer->bindings.back().value && + transient.size == output->byte_count / 2; + })); + const auto commands = ggml::hrx::build_command_program( + imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + ggml_free(ctx); +} + +static void run_rmsnorm_gate_q8_output_checks() { + ggml_init_params params = {}; + params.mem_size = 8 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + const struct { + ggml_type type; + int64_t hidden; + int64_t channels; + int64_t tokens; + int64_t outputs; + bool side_use; + bool terminal; + bool q8; + } cases[] = { + { GGML_TYPE_Q4_K, 128, 6144, 1, 5120, false, false, true }, + { GGML_TYPE_Q4_K, 128, 6144, 2, 5120, false, false, true }, + { GGML_TYPE_Q4_K, 128, 6144, 3, 5120, false, false, true }, + { GGML_TYPE_Q4_K, 128, 6144, 4, 5120, false, false, true }, + { GGML_TYPE_Q4_K, 128, 6144, 5, 5120, false, false, true }, + { GGML_TYPE_Q6_K, 256, 4096, 3, 256, false, false, true }, + { GGML_TYPE_Q4_K, 512, 1536, 5, 256, false, false, true }, + { GGML_TYPE_Q4_K, 1024, 1024, 4, 256, false, false, true }, + { GGML_TYPE_Q4_K, 256, 256, 1, 64, false, false, true }, + { GGML_TYPE_Q4_K, 128, 6144, 6, 5120, false, false, false }, + { GGML_TYPE_Q4_K, 128, 6144, 256, 5120, false, false, false }, + { GGML_TYPE_Q4_K, 128, 6144, 512, 5120, false, false, false }, + { GGML_TYPE_Q4_K, 128, 6144, 5, 5120, true, false, false }, + { GGML_TYPE_Q4_K, 128, 6144, 5, 48, false, false, false }, + { GGML_TYPE_Q6_K, 128, 6144, 5, 256, false, true, false }, + }; + + for (const auto & test : cases) { + const int64_t heads = test.channels / test.hidden; + ggml_tensor * input = heads == 1 ? ggml_new_tensor_2d(ctx, GGML_TYPE_F32, test.hidden, test.tokens) : + ggml_new_tensor_3d(ctx, GGML_TYPE_F32, test.hidden, heads, test.tokens); + ggml_tensor * weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, test.hidden); + ggml_tensor * gate = ggml_dup_tensor(ctx, input); + ggml_tensor * norm = ggml_mul(ctx, ggml_rms_norm(ctx, input, 1.0e-6f), weight); + ggml_tensor * gated = ggml_mul(ctx, norm, ggml_silu(ctx, gate)); + ggml_tensor * reshaped = heads == 1 ? gated : ggml_reshape_2d(ctx, gated, test.channels, test.tokens); + ggml_tensor * projection = ggml_new_tensor_2d(ctx, test.type, test.channels, test.outputs); + ggml_tensor * projected = ggml_mul_mat(ctx, projection, reshaped); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, test.terminal ? projected : ggml_scale(ctx, projected, 0.5f)); + if (test.side_use) { + ggml_build_forward_expand(graph, ggml_scale(ctx, gated, 0.25f)); + } + auto imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const auto & plan = scheduler.plan(); + REQUIRE(plan.valid()); + const auto producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + const auto publish_q8 = dispatch.kernel.compile_parameters.find("ggml.rmsnorm_gate_f32.publish_q8"); + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rmsnorm_gate_f32_publish" && + publish_q8 != dispatch.kernel.compile_parameters.end() && publish_q8->second == "1"; + }); + REQUIRE((producer != plan.dispatches.end()) == test.q8); + if (test.q8) { + require_compile_parameter(*producer, "ggml.rmsnorm_gate_f32.publish_q8", "1"); + const auto * output = imported.graph.values().find_tensor(gated); + REQUIRE(output != nullptr); + const size_t bytes = static_cast(test.tokens) * ggml_row_size(GGML_TYPE_Q8_1, test.channels); + const auto * alternate = plan.metadata.find_alternate_value(output->id, GGML_TYPE_Q8_1, bytes); + REQUIRE(alternate != nullptr); + REQUIRE(producer->bindings.back().value == alternate->alternate_value); + REQUIRE(producer->bindings.back().length == bytes); + REQUIRE(std::count_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "qwen3_moe:ggml_quantize_q8_1_x4_f32"; + }) == 0); + REQUIRE(std::any_of(std::next(producer), plan.dispatches.end(), [&](const auto & dispatch) { + return std::any_of(dispatch.bindings.begin(), dispatch.bindings.end(), + [&](const auto & binding) { return binding.value == alternate->alternate_value; }); + })); + } + const auto commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + ggml_free(ctx); +} + +static void run_symmetric_i4_consumer_qualification_checks() { + const struct { + bool direct_consumer; + bool reshape_consumer; + bool strided_consumer; + bool q8_publication; + } cases[] = { + { true, false, false, true }, + { false, true, false, true }, + { false, false, true, false }, + { false, false, false, false }, + }; + for (const auto & test : cases) { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t hidden = 2048; + constexpr int64_t tokens = 2; + const int64_t producer_hidden = test.strided_consumer ? hidden + 128 : hidden; + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, producer_hidden, tokens); + ggml_tensor * weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, producer_hidden); + ggml_tensor * normalized = ggml_mul(ctx, ggml_rms_norm(ctx, input, 1.0e-6f), weight); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_tensor * consumer_input = normalized; + if (test.reshape_consumer) { + consumer_input = ggml_reshape_2d(ctx, normalized, hidden, tokens); + } else if (test.strided_consumer) { + consumer_input = ggml_view_2d( + ctx, normalized, hidden, tokens, static_cast(producer_hidden) * sizeof(float), 0); + } + if (test.direct_consumer || test.reshape_consumer || test.strided_consumer) { + ggml_tensor * i4_projection = ggml_new_tensor_2d(ctx, GGML_TYPE_Q5_K, hidden, hidden); + ggml_tensor * i4_projected = ggml_mul_mat(ctx, i4_projection, consumer_input); + ggml_build_forward_expand(graph, ggml_scale(ctx, i4_projected, 0.5f)); + if (!test.strided_consumer) { + ggml_tensor * q8_projection = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, hidden, 512); + ggml_tensor * q8_projected = ggml_mul_mat(ctx, q8_projection, consumer_input); + ggml_build_forward_expand(graph, ggml_scale(ctx, q8_projected, 0.5f)); + } + } else { + ggml_build_forward_expand(graph, ggml_scale(ctx, normalized, 0.5f)); + } + + auto imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + const ggml::hrx::Value * normalized_value = imported.graph.values().find_tensor(normalized); + const ggml::hrx::Value * qualified_input = imported.graph.values().find_tensor(consumer_input); + REQUIRE(normalized_value != nullptr); + REQUIRE(qualified_input != nullptr); + for (const ggml::hrx::GraphNode * consumer : imported.graph.index().consumers(qualified_input->id)) { + REQUIRE(consumer == nullptr || + !ggml::hrx::common_symmetric_i4_lowrow_mul_mat_eligible( + imported.graph, *consumer, *qualified_input)); + } + if (test.strided_consumer) { + ggml_free(ctx); + continue; + } + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const auto & plan = scheduler.plan(); + REQUIRE(plan.valid()); + const auto symmetric_producer = + std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == + "loom_libs:ggml_rmsnorm_binary_symmetric_i4_k32"; + }); + REQUIRE(symmetric_producer == plan.dispatches.end()); + const auto q8_producer = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rmsnorm_binary_q8_1_x4_publish"; + }); + REQUIRE((q8_producer != plan.dispatches.end()) == test.q8_publication); + if (test.q8_publication) { + require_compile_parameter(*q8_producer, "ggml.rmsnorm_binary_q8_1_x4.publish_f16", "1"); + const auto * q8 = plan.metadata.find_alternate_value( + normalized_value->id, GGML_TYPE_Q8_1, qwen_q8_1_x4_size(tokens, hidden)); + REQUIRE(q8 != nullptr); + REQUIRE(q8_producer->bindings.size() == 5); + REQUIRE(q8_producer->bindings[3].value == q8->alternate_value); + REQUIRE(std::count_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "qwen3_moe:ggml_quantize_q8_1_x4_f32"; + }) == 0); + REQUIRE(std::any_of(std::next(q8_producer), plan.dispatches.end(), [&](const auto & dispatch) { + return !dispatch.bindings.empty() && dispatch.bindings.front().value == q8->alternate_value; + })); + } else if (!test.strided_consumer) { + REQUIRE(plan.metadata.alternate_values().empty()); + } + const auto commands = ggml::hrx::build_command_program( + imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + ggml_free(ctx); + } +} + +static void run_quantized_conv4_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 8 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + const struct { + ggml_type type; + int64_t inputs; + int64_t channels; + int64_t tokens; + bool side_use; + bool exposed; + bool fused; + bool gathered_state = false; + } cases[] = { + { GGML_TYPE_Q4_K, 5120, 10240, 512, false, false, true }, + { GGML_TYPE_Q6_K, 5120, 10240, 512, false, false, true }, + { GGML_TYPE_Q4_K, 4096, 8192, 512, false, false, true }, + { GGML_TYPE_Q6_K, 4096, 8192, 512, false, false, true }, + { GGML_TYPE_Q4_K, 5120, 10240, 512, true, false, false }, + { GGML_TYPE_Q4_K, 5120, 10240, 512, false, true, false }, + { GGML_TYPE_Q4_K, 5120, 10240, 256, false, false, false }, + { GGML_TYPE_Q4_K, 5120, 10240, 512, false, false, true, true }, + { GGML_TYPE_Q6_K, 5120, 10240, 512, false, false, true, true }, + { GGML_TYPE_Q4_K, 4096, 8192, 512, false, false, true, true }, + { GGML_TYPE_Q6_K, 4096, 8192, 512, false, false, true, true }, + { GGML_TYPE_Q4_K, 5120, 10240, 512, true, false, false, true }, + { GGML_TYPE_Q4_K, 5120, 10240, 512, false, true, false, true }, + }; + + for (const auto & test : cases) { + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, test.inputs, test.tokens); + ggml_tensor * weight = ggml_new_tensor_2d(ctx, test.type, test.inputs, test.channels); + ggml_tensor * raw = ggml_mul_mat(ctx, weight, input); + if (test.exposed) { + ggml_set_output(raw); + } + ggml_tensor * x = ggml_reshape_3d(ctx, raw, test.channels, test.tokens, 1); + ggml_tensor * state = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 3, test.channels, 1); + if (test.gathered_state) { + ggml_tensor * indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 1); + state = ggml_get_rows(ctx, ggml_reshape_2d(ctx, state, 3 * test.channels, 1), indices); + state = ggml_reshape_3d(ctx, state, 3, test.channels, 1); + } + ggml_tensor * filter = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 4, test.channels); + ggml_tensor * window = ggml_concat(ctx, state, ggml_transpose(ctx, x), 0); + ggml_tensor * tail = + ggml_view_3d(ctx, window, 3, test.channels, 1, window->nb[1], window->nb[2], test.tokens * sizeof(float)); + ggml_tensor * cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 3 * test.channels, 1); + ggml_tensor * target = ggml_view_2d(ctx, cache, 3 * test.channels, 1, cache->nb[1], 0); + ggml_tensor * output = ggml_silu(ctx, ggml_ssm_conv(ctx, window, filter)); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, ggml_cpy(ctx, tail, target)); + ggml_build_forward_expand(graph, output); + if (test.side_use) { + ggml_build_forward_expand(graph, ggml_scale(ctx, raw, 0.5f)); + } + auto imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const auto & plan = scheduler.plan(); + REQUIRE(plan.valid()); + const auto fused = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [&](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == + (test.gathered_state ? "loom_libs:ggml_mul_mat_quantized_f16_wmma_prefill_conv4_interior" : + "loom_libs:ggml_mul_mat_quantized_f16_wmma_prefill_conv4"); + }); + REQUIRE((fused != plan.dispatches.end()) == test.fused); + if (test.fused) { + REQUIRE(fused->bindings.size() == (test.gathered_state ? 5 : 6)); + require_compile_parameter(*fused, "ggml.mul_mat.weight_format", test.type == GGML_TYPE_Q4_K ? "4" : "6"); + if (test.gathered_state) { + const auto finish = + std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == + "loom_libs:llm_ssm_conv_dconv4_silu_prefill_finish_f32"; + }); + const auto gather = + std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id).find("get_rows") != std::string::npos; + }); + REQUIRE(finish != plan.dispatches.end()); + REQUIRE(gather != plan.dispatches.end()); + REQUIRE(fused < gather && gather < finish); + REQUIRE(finish->bindings[2].value == fused->bindings[4].value); + REQUIRE(finish->bindings[3].value == fused->bindings[3].value); + REQUIRE(finish->bindings[2].length == static_cast(6 * test.channels) * sizeof(float)); + } + REQUIRE(std::none_of(plan.transients.begin(), plan.transients.end(), [](const auto & transient) { + return transient.name == "llm.ssm_conv.window_snapshot"; + })); + } + const auto commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + ggml_free(ctx); +} + +static void run_generic_ssm_conv_binary_dispatch_checks() { + const struct { + ggml::hrx::BinaryKind kind; + bool conv_is_lhs; + } cases[] = { + { ggml::hrx::BinaryKind::Add, true }, + { ggml::hrx::BinaryKind::Sub, true }, + { ggml::hrx::BinaryKind::Sub, false }, + { ggml::hrx::BinaryKind::Mul, true }, + { ggml::hrx::BinaryKind::Div, true }, + { ggml::hrx::BinaryKind::Div, false }, + }; + + for (const auto & test : cases) { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t d_conv = 4; + constexpr int64_t d_inner = 64; + constexpr int64_t n_t = 8; + ggml_tensor * window = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, d_conv + n_t - 1, d_inner); + ggml_tensor * filter = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, d_conv, d_inner); + ggml_tensor * operand = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, d_inner, n_t); + ggml_tensor * conv = ggml_ssm_conv(ctx, window, filter); + ggml_tensor * output = nullptr; + switch (test.kind) { + case ggml::hrx::BinaryKind::Add: + output = ggml_add(ctx, conv, operand); + break; + case ggml::hrx::BinaryKind::Sub: + output = test.conv_is_lhs ? ggml_sub(ctx, conv, operand) : ggml_sub(ctx, operand, conv); + break; + case ggml::hrx::BinaryKind::Mul: + output = ggml_mul(ctx, conv, operand); + break; + case ggml::hrx::BinaryKind::Div: + output = test.conv_is_lhs ? ggml_div(ctx, conv, operand) : ggml_div(ctx, operand, conv); + break; + default: + REQUIRE(false); + } + + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, output); + auto imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const auto & plan = scheduler.plan(); + REQUIRE(plan.valid()); + REQUIRE(plan.dispatches.size() == 1); + const auto & dispatch = plan.dispatches.front(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:llm_ssm_conv_binary_f32"); + REQUIRE(dispatch.kernel.workload_specialization == ggml::hrx::WorkloadSpecialization::Dynamic); + REQUIRE(dispatch.kernel.integer_parameters.at("n_t") == n_t); + REQUIRE(dispatch.kernel.integer_parameters.at("n_s") == 1); + REQUIRE(dispatch.kernel.compile_parameters.count("llm.ssm_conv.generic.n_t") == 0); + REQUIRE(dispatch.kernel.compile_parameters.count("llm.ssm_conv.generic.n_s") == 0); + REQUIRE(dispatch.bindings.size() == 4); + require_compile_parameter(dispatch, "llm.ssm_conv.generic.binary_op", + std::to_string(ggml::hrx::binary_kind_config_value(test.kind))); + require_compile_parameter(dispatch, "llm.ssm_conv.generic.binary_lhs", test.conv_is_lhs ? "1" : "0"); + REQUIRE(dispatch.bindings[2].value == imported.graph.values().find_tensor(operand)->id); + REQUIRE(dispatch.bindings[3].value == imported.graph.values().find_tensor(output)->id); + + const auto commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + ggml_free(ctx); + } +} + +static void run_lfm_dconv3_dispatch_checks() { + const struct { + int64_t tokens; + int64_t hidden; + int64_t state_rows; + int64_t filter_rows; + bool disabled; + bool fused; + } cases[] = { + {1, 2048, 2, 3, false, true}, + {64, 2048, 2, 3, false, true}, + {2, 2048, 2, 3, false, false}, + {64, 2016, 2, 3, false, false}, + {64, 2048, 3, 3, false, false}, + {64, 2048, 2, 4, false, false}, + {64, 2048, 2, 3, true, false}, + }; + + for (const auto & test : cases) { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + ggml_tensor * x = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, test.hidden, test.tokens, 1); + ggml_tensor * state = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, test.state_rows, test.hidden, 1); + ggml_tensor * filter = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, test.filter_rows, test.hidden); + ggml_tensor * window = ggml_concat(ctx, state, ggml_transpose(ctx, x), 0); + ggml_tensor * tail = ggml_view_3d(ctx, window, test.state_rows, test.hidden, 1, window->nb[1], window->nb[2], + test.tokens * sizeof(float)); + ggml_tensor * cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, test.state_rows * test.hidden, 1); + ggml_tensor * target = ggml_view_2d(ctx, cache, test.state_rows * test.hidden, 1, cache->nb[1], 0); + ggml_tensor * output = ggml_ssm_conv(ctx, window, filter); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, ggml_cpy(ctx, tail, target)); + ggml_build_forward_expand(graph, output); + if (test.disabled) { + REQUIRE(setenv("GGML_HRX_DISABLE_LFM_DCONV3_FUSION", "1", 1) == 0); + } + auto imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const auto & plan = scheduler.plan(); + REQUIRE(plan.valid()); + const auto fused = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:llm_ssm_conv_lfm_dconv3_state_f32"; + }); + REQUIRE((fused != plan.dispatches.end()) == test.fused); + if (test.fused) { + REQUIRE(fused->bindings.size() == 5); + require_compile_parameter(*fused, "llm.ssm_conv.lfm_dconv3.n_t", std::to_string(test.tokens)); + } + const auto commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + if (test.disabled) { + REQUIRE(unsetenv("GGML_HRX_DISABLE_LFM_DCONV3_FUSION") == 0); + } + ggml_free(ctx); + } +} + +static void run_gdn_native_projection_pair_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 8 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + const struct { + int64_t tokens; + int64_t inputs; + int64_t heads; + bool different_input; + bool mutate_input; + bool aliased_weight; + ggml_type weight_type; + bool fused; + } cases[] = { + {1, 5120, 48, false, false, false, GGML_TYPE_Q4_K, true}, + {1, 4096, 24, false, false, false, GGML_TYPE_Q4_K, true}, + {1, 8192, 48, false, false, false, GGML_TYPE_Q4_K, true}, + {1, 32768, 48, false, false, false, GGML_TYPE_Q4_K, true}, + {2, 5120, 48, false, false, false, GGML_TYPE_Q4_K, false}, + {5, 5120, 48, false, false, false, GGML_TYPE_Q4_K, false}, + {16, 5120, 48, false, false, false, GGML_TYPE_Q4_K, false}, + {17, 5120, 48, false, false, false, GGML_TYPE_Q4_K, false}, + {1, 3072, 48, false, false, false, GGML_TYPE_Q4_K, false}, + {1, 4352, 48, false, false, false, GGML_TYPE_Q4_K, false}, + {1, 5120, 16, false, false, false, GGML_TYPE_Q4_K, false}, + {1, 5120, 48, true, false, false, GGML_TYPE_Q4_K, false}, + {1, 5120, 48, false, true, false, GGML_TYPE_Q4_K, false}, + {1, 5120, 48, false, false, true, GGML_TYPE_Q4_K, false}, + {1, 5120, 48, false, false, false, GGML_TYPE_Q6_K, false}, + }; + + for (const auto & test : cases) { + const int64_t width = 128; + const int64_t qheads = test.heads / 4; + const int64_t hidden = width * (2 * qheads + test.heads); + const int64_t state_elements = width * width * test.heads; + const size_t state_bytes = state_elements * sizeof(float); + ggml_tensor * input = ggml_scale(ctx, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, test.inputs, test.tokens), 0.5f); + ggml_tensor * first_weight = ggml_new_tensor_2d(ctx, test.weight_type, test.inputs, test.heads); + ggml_tensor * second_weight = ggml_dup_tensor(ctx, first_weight); + if (test.aliased_weight) { + second_weight = ggml_view_2d(ctx, second_weight, test.inputs, test.heads, second_weight->nb[1], 0); + } + ggml_tensor * alpha = ggml_mul_mat(ctx, first_weight, input); + ggml_tensor * second_input = test.different_input ? ggml_dup_tensor(ctx, input) : input; + ggml_tensor * beta_raw = ggml_mul_mat(ctx, second_weight, second_input); + ggml_tensor * bias = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, test.heads); + ggml_tensor * scale = ggml_dup_tensor(ctx, bias); + ggml_tensor * gate = ggml_mul( + ctx, ggml_softplus(ctx, ggml_add(ctx, ggml_reshape_3d(ctx, alpha, test.heads, test.tokens, 1), bias)), + scale); + gate = ggml_reshape_4d(ctx, gate, 1, test.heads, test.tokens, 1); + ggml_tensor * beta = ggml_sigmoid(ctx, ggml_reshape_4d(ctx, beta_raw, 1, test.heads, test.tokens, 1)); + ggml_tensor * raw = ggml_scale(ctx, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden, test.tokens), 0.5f); + const auto view = [&](int64_t heads, size_t offset) { + return ggml_view_4d(ctx, raw, width, heads, test.tokens, 1, width * sizeof(float), hidden * sizeof(float), + hidden * test.tokens * sizeof(float), offset); + }; + ggml_tensor * q = ggml_l2_norm(ctx, view(qheads, 0), 1.e-6f); + ggml_tensor * k = ggml_l2_norm(ctx, view(qheads, width * qheads * sizeof(float)), 1.e-6f); + ggml_tensor * v = view(test.heads, 2 * width * qheads * sizeof(float)); + ggml_tensor * cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, state_elements, 20); + ggml_tensor * ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 1); + ggml_tensor * gathered = ggml_get_rows(ctx, cache, ids); + ggml_tensor * state = ggml_reshape_4d(ctx, gathered, width, width, test.heads, 1); + ggml_tensor * gdn = ggml_gated_delta_net(ctx, q, k, v, gate, beta, state, 5); + const int64_t count = std::min(test.tokens, 5); + const size_t attention_bytes = width * test.heads * test.tokens * sizeof(float); + ggml_tensor * attention = ggml_view_4d(ctx, gdn, width, test.heads, test.tokens, 1, width * sizeof(float), + width * test.heads * sizeof(float), attention_bytes, 0); + ggml_tensor * snapshots = + ggml_view_3d(ctx, gdn, state_elements, 1, count, state_bytes, state_bytes, attention_bytes); + ggml_tensor * target = ggml_view_3d(ctx, cache, state_elements, 1, count, state_bytes, 4 * state_bytes, 0); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, raw); + ggml_build_forward_expand(graph, alpha); + if (test.mutate_input) { + ggml_build_forward_expand(graph, ggml_cpy(ctx, ggml_dup_tensor(ctx, input), input)); + } + ggml_build_forward_expand(graph, beta_raw); + ggml_build_forward_expand(graph, gathered); + ggml_build_forward_expand(graph, ggml_cpy(ctx, snapshots, target)); + ggml_build_forward_expand(graph, ggml_scale(ctx, attention, 0.5f)); + auto imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const auto & plan = scheduler.plan(); + REQUIRE(plan.valid()); + const auto fused = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_mul_mat_dual_q4_f32_decode"; + }); + REQUIRE((fused != plan.dispatches.end()) == test.fused); + if (test.fused) { + REQUIRE(fused->bindings.size() == 5); + require_compile_parameter(*fused, "ggml.mul_mat_dual_q4_f32_c1.input_size", std::to_string(test.inputs)); + REQUIRE(fused->bindings[0].value == imported.graph.values().find_tensor(input)->id); + REQUIRE(fused->bindings[1].value == imported.graph.values().find_tensor(first_weight)->id); + REQUIRE(fused->bindings[2].value == imported.graph.values().find_tensor(second_weight)->id); + REQUIRE(fused->bindings[3].value == imported.graph.values().find_tensor(alpha)->id); + REQUIRE(fused->bindings[4].value == imported.graph.values().find_tensor(beta_raw)->id); + } + const auto epilogue = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == + "loom_libs:llm_gated_delta_net_projection_epilogue_f32"; + }); + REQUIRE((epilogue != plan.dispatches.end()) == (test.tokens > 16)); + if (epilogue != plan.dispatches.end()) { + REQUIRE(epilogue->kernel.workload_specialization == ggml::hrx::WorkloadSpecialization::Dynamic); + REQUIRE(epilogue->kernel.integer_parameters.at("element_count") == test.heads * test.tokens); + REQUIRE(epilogue->kernel.compile_parameters.count("llm.gated_delta_net.epilogue_element_count") == 0); + } + const auto commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + ggml_free(ctx); +} + +static void run_gdn_selected_snapshot_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 8 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + const struct { + int64_t tokens; + int64_t sequences; + bool ready; + bool side_use; + bool cache_read; + bool observed; + bool aligned; + bool fused; + } cases[] = { + {2,1,true,false,false,false,true,true}, + {3,1,true,false,false,false,true,true}, + {4,1,true,false,false,false,true,true}, + {5,1,true,false,false,false,true,true}, + {1,1,true,false,false,false,true,false}, + {5,2,true,false,false,false,true,false}, + {5,1,false,false,false,false,true,false}, + {5,1,true,true,false,false,true,false}, + {5,1,true,false,true,false,true,false}, + {5,1,true,false,false,true,true,false}, + {5,1,true,false,false,false,false,false}, + }; + for (const auto & test : cases) { + const int64_t width=128, heads=6, qheads=2, hidden=1280; + const int64_t state_elements=width*width*heads; + const size_t state_bytes=state_elements*sizeof(float); + ggml_tensor * raw=ggml_scale(ctx, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden, test.tokens*test.sequences), 0.5f); + ggml_tensor * alpha=ggml_scale(ctx, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, heads, test.tokens*test.sequences), 0.5f); + ggml_tensor * beta_raw=ggml_scale(ctx, ggml_dup_tensor(ctx, alpha), 0.5f); + ggml_tensor * bias=ggml_new_tensor_1d(ctx, GGML_TYPE_F32, heads); + ggml_tensor * scale=ggml_dup_tensor(ctx, bias); + ggml_tensor * gate=ggml_mul(ctx, ggml_softplus(ctx, ggml_add(ctx, ggml_reshape_3d(ctx, alpha, heads, test.tokens, test.sequences), bias)), scale); + gate=ggml_reshape_4d(ctx, gate, 1, heads, test.tokens, test.sequences); + ggml_tensor * beta=ggml_sigmoid(ctx, ggml_reshape_4d(ctx, beta_raw, 1, heads, test.tokens, test.sequences)); + ggml_tensor * cache=ggml_new_tensor_2d(ctx, GGML_TYPE_F32, state_elements, 20); + ggml_tensor * ids=ggml_new_tensor_1d(ctx, GGML_TYPE_I32, test.sequences); + ggml_tensor * gathered=ggml_get_rows(ctx, cache, ids); + if (test.observed) { + ggml_set_output(gathered); + } + ggml_tensor * state=ggml_reshape_4d(ctx, gathered, width, width, heads, test.sequences); + const auto view=[&](int64_t nheads, size_t offset) { + return ggml_view_4d(ctx, raw, width, nheads, test.tokens, test.sequences, + width*sizeof(float), hidden*sizeof(float), hidden*test.tokens*sizeof(float), offset); + }; + ggml_tensor * q=ggml_l2_norm(ctx, view(qheads,0), 1.e-6f); + ggml_tensor * k=ggml_l2_norm(ctx, view(qheads,width*qheads*sizeof(float)), 1.e-6f); + ggml_tensor * v=view(heads,2*width*qheads*sizeof(float)); + ggml_tensor * gdn=ggml_gated_delta_net(ctx,q,k,v,gate,beta,state,5); + const size_t attention_bytes=width*heads*test.tokens*test.sequences*sizeof(float); + ggml_tensor * attention=ggml_view_4d(ctx,gdn,width,heads,test.tokens,test.sequences, + width*sizeof(float),width*heads*sizeof(float), + width*heads*test.tokens*sizeof(float),0); + ggml_tensor * snapshots=ggml_view_3d(ctx,gdn,state_elements,test.sequences,test.tokens, + state_bytes,state_bytes*test.sequences,attention_bytes); + ggml_tensor * target=ggml_view_3d(ctx,cache,state_elements,test.sequences,test.tokens, + state_bytes,4*state_bytes,test.aligned ? state_bytes : sizeof(float)); + ggml_cgraph * graph=ggml_new_graph(ctx); + if (test.ready) { + ggml_build_forward_expand(graph,raw); + } + ggml_build_forward_expand(graph,alpha); + ggml_build_forward_expand(graph,beta_raw); + ggml_build_forward_expand(graph,gathered); + if (test.cache_read) { + ggml_build_forward_expand(graph,ggml_scale(ctx,cache,0.25f)); + } + ggml_build_forward_expand(graph,ggml_cpy(ctx,snapshots,target)); + ggml_build_forward_expand(graph,ggml_scale(ctx,attention,0.5f)); + if (test.side_use) { + ggml_build_forward_expand(graph,ggml_scale(ctx,gathered,0.5f)); + } + auto imported=ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph,test_dispatch_target())); + const auto & plan=scheduler.plan(); + REQUIRE(plan.valid()); + const auto fused=std::find_if(plan.dispatches.begin(),plan.dispatches.end(),[](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id)== + "loom_libs:llm_gated_delta_net_f32_wmma_head128_selected_snapshot_projection_epilogue"; + }); + REQUIRE((fused!=plan.dispatches.end())==test.fused); + if (test.fused) { + REQUIRE(fused->bindings.size()==11); + require_compile_parameter(*fused,"llm.gated_delta_net.state_row_count","20"); + REQUIRE(fused->bindings[7].value==imported.graph.values().find_tensor(cache)->id); + REQUIRE(fused->bindings[10].value==imported.graph.values().find_tensor(ids)->id); + REQUIRE(std::none_of(plan.dispatches.begin(),plan.dispatches.end(),[](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id)=="loom_libs:ggml_get_rows_f32"; + })); + } + const auto commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + ggml_free(ctx); +} + +static void run_gdn_rmsnorm_gate_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + const struct { + ggml_type type; + int64_t tokens; + bool side_use; + bool attention_output; + bool gdn_output; + bool fused; + } cases[] = { + { GGML_TYPE_Q4_K, 512, false, false, false, true }, + { GGML_TYPE_Q6_K, 512, false, false, false, true }, + { GGML_TYPE_Q4_K, 512, true, false, false, false }, + { GGML_TYPE_Q4_K, 256, false, false, false, false }, + { GGML_TYPE_Q4_K, 512, false, true, false, false }, + { GGML_TYPE_Q4_K, 512, false, false, true, false }, + }; + + for (const auto & test : cases) { + const int64_t width = 128; + const int64_t heads = 2; + const int64_t channels = width * heads; + const int64_t state_elements = width * channels; + ggml_tensor * q = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, width, 1, test.tokens, 1); + ggml_tensor * k = ggml_dup_tensor(ctx, q); + ggml_tensor * v = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, width, heads, test.tokens, 1); + ggml_tensor * g = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, 1, heads, test.tokens, 1); + ggml_tensor * beta = ggml_dup_tensor(ctx, g); + ggml_tensor * state = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, width, width, heads, 1); + ggml_tensor * gdn = ggml_gated_delta_net(ctx, q, k, v, g, beta, state, 1); + ggml_tensor * attention = ggml_view_4d(ctx, gdn, width, heads, test.tokens, 1, width * sizeof(float), + channels * sizeof(float), channels * test.tokens * sizeof(float), 0); + if (test.attention_output) { + ggml_set_output(attention); + } + if (test.gdn_output) { + ggml_set_output(gdn); + } + ggml_tensor * new_state = ggml_view_2d(ctx, gdn, state_elements, 1, state_elements * sizeof(float), + channels * test.tokens * sizeof(float)); + ggml_tensor * cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, state_elements, 1); + ggml_tensor * weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, width); + ggml_tensor * gate = ggml_dup_tensor(ctx, attention); + ggml_tensor * norm = ggml_mul(ctx, ggml_rms_norm(ctx, attention, 2.0e-6f), weight); + ggml_tensor * gated = ggml_mul(ctx, norm, ggml_silu(ctx, gate)); + ggml_tensor * input = ggml_reshape_2d(ctx, gated, channels, test.tokens); + ggml_tensor * projection = ggml_new_tensor_2d(ctx, test.type, channels, 4096); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, ggml_cpy(ctx, new_state, cache)); + ggml_build_forward_expand(graph, ggml_mul_mat(ctx, projection, input)); + if (test.side_use) { + ggml_build_forward_expand(graph, ggml_scale(ctx, attention, 0.5f)); + } + auto imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const auto & plan = scheduler.plan(); + REQUIRE(plan.valid()); + const auto fused = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == + "loom_libs:llm_gated_delta_net_f32_wmma_head128_rmsnorm_gate"; + }); + REQUIRE((fused != plan.dispatches.end()) == test.fused); + REQUIRE(std::count_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_rmsnorm_gate_f32_publish"; + }) == (test.fused ? 0 : 1)); + if (test.fused) { + REQUIRE(fused->bindings.size() == 11); + const auto * projected = imported.graph.values().find_tensor(input); + REQUIRE(projected != nullptr); + const auto * packed = + plan.metadata.find_generated_resource(projected->id, ggml::hrx::GeneratedResourceRole::F16K16Major); + REQUIRE(packed != nullptr && fused->bindings.back().value == packed->generated_value); + require_compile_parameter(*fused, "ggml.rmsnorm_gate_f32.rms_epsilon", "1.99999999e-06"); + } + const auto commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + ggml_free(ctx); +} + +static void run_qwen_matmul_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 12 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, 2048, + 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_f32_f32_wmma", 1, 2048, 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 2); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 2, 2048, + 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 3); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 3, 2048, + 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 5); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 5, 2048, + 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 6); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32", 6, 2048, 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q5_0, 640, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 640, 2); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32", 2, 640, 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma", 1, 2048, 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ4_XS, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, 2048, + 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ4_XS, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", 1, 2048, 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q8_0, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, 2048, + 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q8_0, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", 1, 2048, 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q8_1, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, 2048, + 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q8_1, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", 1, 2048, 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, 2048, + 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", 1, 2048, 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_BF16, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, 2048, + 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_BF16, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", 1, 2048, 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 256); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, 2048, + 256); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 256); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", 1, 2048, 256); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 2048, 256); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 33); + ggml_tensor * bias = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 256); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + REQUIRE(bias != nullptr); + ggml_tensor * projection = ggml_mul_mat(ctx, weight, input); + REQUIRE(projection != nullptr); + ggml_tensor * output = ggml_add(ctx, projection, bias); + REQUIRE(output != nullptr); + schedule_fused_matmul_postops_command( + ctx, output, "loom_libs:ggml_mul_mat_tiled_input_f32_bias_residual_publish_f32", GGML_TYPE_F16, 33, 2048, + 256, 2, { GGML_OP_MUL_MAT, GGML_OP_ADD }, + { "input", "weight", "bias", "residual_input", "residual_output" }, false); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_BF16, 2048, 256); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 33); + ggml_tensor * bias = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 256); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + REQUIRE(bias != nullptr); + ggml_tensor * projection = ggml_mul_mat(ctx, weight, input); + REQUIRE(projection != nullptr); + ggml_tensor * output = ggml_add(ctx, projection, bias); + REQUIRE(output != nullptr); + schedule_fused_matmul_postops_command( + ctx, output, "loom_libs:ggml_mul_mat_tiled_input_f32_bias_residual_publish_f32", GGML_TYPE_BF16, 33, 2048, + 256, 2, { GGML_OP_MUL_MAT, GGML_OP_ADD }, + { "input", "weight", "bias", "residual_input", "residual_output" }, false); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 2048, 256); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 33); + ggml_tensor * residual = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 256, 33); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + REQUIRE(residual != nullptr); + ggml_tensor * projection = ggml_mul_mat(ctx, weight, input); + REQUIRE(projection != nullptr); + ggml_tensor * output = ggml_add(ctx, projection, residual); + REQUIRE(output != nullptr); + schedule_fused_matmul_postops_command( + ctx, output, "loom_libs:ggml_mul_mat_tiled_input_f32_bias_residual_publish_f32", GGML_TYPE_F16, 33, 2048, + 256, 2, { GGML_OP_MUL_MAT, GGML_OP_ADD }, + { "input", "weight", "bias", "residual_input", "residual_output" }, false); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 2048, 256); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 33); + ggml_tensor * bias = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 256); + ggml_tensor * residual = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 256, 33); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + REQUIRE(bias != nullptr); + REQUIRE(residual != nullptr); + ggml_tensor * projection = ggml_mul_mat(ctx, weight, input); + REQUIRE(projection != nullptr); + ggml_tensor * biased = ggml_add(ctx, projection, bias); + REQUIRE(biased != nullptr); + ggml_tensor * output = ggml_add(ctx, biased, residual); + REQUIRE(output != nullptr); + schedule_fused_matmul_postops_command( + ctx, output, "loom_libs:ggml_mul_mat_tiled_input_f32_bias_residual_publish_f32", GGML_TYPE_F16, 33, 2048, + 256, 3, { GGML_OP_MUL_MAT, GGML_OP_ADD, GGML_OP_ADD }, + { "input", "weight", "bias", "residual_input", "residual_output" }, false); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 2048, 256); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 33); + ggml_tensor * residual = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 256, 33); + ggml_tensor * norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 256); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + REQUIRE(residual != nullptr); + REQUIRE(norm_weight != nullptr); + ggml_tensor * projection = ggml_mul_mat(ctx, weight, input); + REQUIRE(projection != nullptr); + ggml_tensor * added = ggml_add(ctx, projection, residual); + REQUIRE(added != nullptr); + ggml_tensor * rms = ggml_rms_norm(ctx, added, 0.000001f); + REQUIRE(rms != nullptr); + ggml_tensor * output = ggml_mul(ctx, rms, norm_weight); + REQUIRE(output != nullptr); + schedule_fused_matmul_postops_command( + ctx, output, "loom_libs:ggml_mul_mat_tiled_input_f32_bias_residual_publish_f32", GGML_TYPE_F16, 33, 2048, + 256, 4, { GGML_OP_MUL_MAT, GGML_OP_ADD, GGML_OP_RMS_NORM, GGML_OP_MUL }, + { "input", "weight", "bias", "residual_input", "residual_output" }, true); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 2048, 256); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 33); + ggml_tensor * bias = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 256); + ggml_tensor * residual = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 256, 33); + ggml_tensor * norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 256); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + REQUIRE(bias != nullptr); + REQUIRE(residual != nullptr); + REQUIRE(norm_weight != nullptr); + ggml_tensor * projection = ggml_mul_mat(ctx, weight, input); + REQUIRE(projection != nullptr); + ggml_tensor * biased = ggml_add(ctx, projection, bias); + REQUIRE(biased != nullptr); + ggml_tensor * added = ggml_add(ctx, biased, residual); + REQUIRE(added != nullptr); + ggml_tensor * rms = ggml_rms_norm(ctx, added, 0.000001f); + REQUIRE(rms != nullptr); + ggml_tensor * output = ggml_mul(ctx, rms, norm_weight); + REQUIRE(output != nullptr); + schedule_fused_matmul_postops_command( + ctx, output, "loom_libs:ggml_mul_mat_tiled_input_f32_bias_residual_publish_f32", GGML_TYPE_F16, 33, 2048, + 256, 5, { GGML_OP_MUL_MAT, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_RMS_NORM, GGML_OP_MUL }, + { "input", "weight", "bias", "residual_input", "residual_output" }, true); + } + + { + struct FormatCase { + ggml_type type; + int64_t output_size; + }; + + const FormatCase cases[] = { + { GGML_TYPE_Q3_K, 128 }, + { GGML_TYPE_Q4_K, 128 }, + { GGML_TYPE_Q6_K, 128 }, + { GGML_TYPE_IQ3_S, 128 }, + { GGML_TYPE_IQ4_NL, 128 }, + { GGML_TYPE_IQ4_XS, 128 }, + { GGML_TYPE_Q8_0, 128 }, + { GGML_TYPE_Q8_1, 128 }, + { GGML_TYPE_F16, 128 }, + { GGML_TYPE_BF16, 128 }, + { GGML_TYPE_F32, 256 }, + }; + + for (const FormatCase & c : cases) { + ggml_tensor * gate_weight = ggml_new_tensor_2d(ctx, c.type, 2048, c.output_size); + ggml_tensor * up_weight = ggml_new_tensor_2d(ctx, c.type, 2048, c.output_size); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(gate_weight != nullptr); + REQUIRE(up_weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * gate = ggml_mul_mat(ctx, gate_weight, input); + ggml_tensor * up = ggml_mul_mat(ctx, up_weight, input); + REQUIRE(gate != nullptr); + REQUIRE(up != nullptr); + ggml_tensor * output = ggml_glu_split(ctx, gate, up, GGML_GLU_OP_SWIGLU); + REQUIRE(output != nullptr); + schedule_fused_matmul_swiglu_command(ctx, output, c.type, c.type, 4, 2048, c.output_size); + } + } + + { + ggml_tensor * gate_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ4_NL, 640, 2048); + ggml_tensor * up_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ4_NL, 640, 2048); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 640, 2); + REQUIRE(gate_weight != nullptr); + REQUIRE(up_weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * gate = ggml_mul_mat(ctx, gate_weight, input); + ggml_tensor * up = ggml_mul_mat(ctx, up_weight, input); + REQUIRE(gate != nullptr); + REQUIRE(up != nullptr); + ggml_tensor * output = ggml_glu_split(ctx, gate, up, GGML_GLU_OP_GEGLU); + REQUIRE(output != nullptr); + schedule_fused_matmul_swiglu_command(ctx, output, GGML_TYPE_IQ4_NL, GGML_TYPE_IQ4_NL, 2, 640, 2048, + ggml::hrx::BinaryKind::GeGLU); + } + + { + ggml_tensor * gate_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ4_NL, 640, 2048); + ggml_tensor * up_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ4_NL, 640, 2048); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 640, 1); + REQUIRE(gate_weight != nullptr); + REQUIRE(up_weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * gate = ggml_mul_mat(ctx, gate_weight, input); + ggml_tensor * up = ggml_mul_mat(ctx, up_weight, input); + REQUIRE(gate != nullptr); + REQUIRE(up != nullptr); + ggml_tensor * output = ggml_glu_split(ctx, gate, up, GGML_GLU_OP_SWIGLU); + REQUIRE(output != nullptr); + require_matmul_swiglu_falls_back_to_compilable_plan(ctx, output); + } + + { + ggml_tensor * packed = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 256, 3); + REQUIRE(packed != nullptr); + ggml_tensor * output = ggml_glu(ctx, packed, GGML_GLU_OP_SWIGLU, false); + REQUIRE(output != nullptr); + schedule_packed_glu_command(ctx, output, ggml::hrx::BinaryKind::SwiGLU, 128, 3, false); + } + + { + ggml_tensor * packed = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 256, 3); + REQUIRE(packed != nullptr); + ggml_tensor * output = ggml_glu(ctx, packed, GGML_GLU_OP_GEGLU, true); + REQUIRE(output != nullptr); + schedule_packed_glu_command(ctx, output, ggml::hrx::BinaryKind::GeGLU, 128, 3, true); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 256, 256); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 256, 2); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * packed = ggml_mul_mat(ctx, weight, input); + REQUIRE(packed != nullptr); + ggml_tensor * output = ggml_glu(ctx, packed, GGML_GLU_OP_SWIGLU, false); + REQUIRE(output != nullptr); + schedule_packed_matmul_glu_command(ctx, output, GGML_TYPE_Q4_K, 2, 256, 128, ggml::hrx::BinaryKind::SwiGLU, + false); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 256, 256); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 256, 2); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * packed = ggml_mul_mat(ctx, weight, input); + REQUIRE(packed != nullptr); + ggml_tensor * output = ggml_glu(ctx, packed, GGML_GLU_OP_SWIGLU, true); + REQUIRE(output != nullptr); + schedule_packed_matmul_glu_command(ctx, output, GGML_TYPE_Q4_K, 2, 256, 128, ggml::hrx::BinaryKind::SwiGLU, + true); + } + + { + ggml_tensor * gate_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ4_NL, 640, 256); + ggml_tensor * up_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q5_0, 640, 256); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 640, 2); + REQUIRE(gate_weight != nullptr); + REQUIRE(up_weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * gate = ggml_mul_mat(ctx, gate_weight, input); + ggml_tensor * up = ggml_mul_mat(ctx, up_weight, input); + REQUIRE(gate != nullptr); + REQUIRE(up != nullptr); + ggml_tensor * output = ggml_glu_split(ctx, gate, up, GGML_GLU_OP_SWIGLU); + REQUIRE(output != nullptr); + schedule_fused_matmul_swiglu_command(ctx, output, GGML_TYPE_IQ4_NL, GGML_TYPE_Q5_0, 2, 640, 256); + } + + { + ggml_tensor * gate_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * up_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(gate_weight != nullptr); + REQUIRE(up_weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * gate = ggml_mul_mat(ctx, gate_weight, input); + ggml_tensor * up = ggml_mul_mat(ctx, up_weight, input); + REQUIRE(gate != nullptr); + REQUIRE(up != nullptr); + ggml_tensor * output = ggml_glu_split(ctx, gate, up, GGML_GLU_OP_SWIGLU); + REQUIRE(output != nullptr); + schedule_fused_matmul_swiglu_command(ctx, output, GGML_TYPE_Q4_K, GGML_TYPE_F16, 4, 2048, 128); + } + + { + ggml_tensor * gate_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 128); + ggml_tensor * up_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q8_1, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(gate_weight != nullptr); + REQUIRE(up_weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * gate = ggml_mul_mat(ctx, gate_weight, input); + ggml_tensor * up = ggml_mul_mat(ctx, up_weight, input); + REQUIRE(gate != nullptr); + REQUIRE(up != nullptr); + ggml_tensor * output = ggml_glu_split(ctx, gate, up, GGML_GLU_OP_SWIGLU); + REQUIRE(output != nullptr); + schedule_fused_matmul_swiglu_command(ctx, output, GGML_TYPE_F32, GGML_TYPE_Q8_1, 4, 2048, 128); + } + + { + ggml_tensor * gate_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * up_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * gate_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + ggml_tensor * up_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(gate_weight != nullptr); + REQUIRE(up_weight != nullptr); + REQUIRE(gate_input != nullptr); + REQUIRE(up_input != nullptr); + ggml_tensor * gate = ggml_mul_mat(ctx, gate_weight, gate_input); + ggml_tensor * up = ggml_mul_mat(ctx, up_weight, up_input); + REQUIRE(gate != nullptr); + REQUIRE(up != nullptr); + ggml_tensor * output = ggml_glu_split(ctx, gate, up, GGML_GLU_OP_SWIGLU); + REQUIRE(output != nullptr); + require_matmul_swiglu_falls_back(ctx, output); + } + + for (const auto & variant : { + std::pair{ GGML_GLU_OP_GEGLU, ggml::hrx::BinaryKind::GeGLU }, + std::pair{ GGML_GLU_OP_REGLU, ggml::hrx::BinaryKind::RegLU }, + std::pair{ GGML_GLU_OP_GEGLU_ERF, ggml::hrx::BinaryKind::GeGLUErf }, + std::pair{ GGML_GLU_OP_GEGLU_QUICK, + ggml::hrx::BinaryKind::GeGLUQuick }, + }) { + ggml_tensor * gate_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * up_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(gate_weight != nullptr); + REQUIRE(up_weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * gate = ggml_mul_mat(ctx, gate_weight, input); + ggml_tensor * up = ggml_mul_mat(ctx, up_weight, input); + REQUIRE(gate != nullptr); + REQUIRE(up != nullptr); + ggml_tensor * output = ggml_glu_split(ctx, gate, up, variant.first); + REQUIRE(output != nullptr); + schedule_fused_matmul_swiglu_command(ctx, output, GGML_TYPE_Q4_K, GGML_TYPE_Q4_K, 4, 2048, 128, variant.second, + GGML_OP_GLU); + } + + { + ggml_tensor * lhs_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * rhs_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(lhs_weight != nullptr); + REQUIRE(rhs_weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * lhs = ggml_mul_mat(ctx, lhs_weight, input); + ggml_tensor * rhs = ggml_mul_mat(ctx, rhs_weight, input); + REQUIRE(lhs != nullptr); + REQUIRE(rhs != nullptr); + ggml_tensor * output = ggml_sub(ctx, lhs, rhs); + REQUIRE(output != nullptr); + schedule_fused_matmul_swiglu_command(ctx, output, GGML_TYPE_Q4_K, GGML_TYPE_Q4_K, 4, 2048, 128, + ggml::hrx::BinaryKind::Sub, GGML_OP_SUB); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * matmul = ggml_mul_mat(ctx, weight, input); + REQUIRE(matmul != nullptr); + ggml_tensor * output = ggml_silu(ctx, matmul); + REQUIRE(output != nullptr); + schedule_fused_matmul_unary_command(ctx, output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", + ggml::hrx::UnaryKind::Silu, 4, 2048, 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q8_0, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * matmul = ggml_mul_mat(ctx, weight, input); + REQUIRE(matmul != nullptr); + ggml_tensor * output = ggml_relu(ctx, matmul); + REQUIRE(output != nullptr); + schedule_fused_matmul_unary_command(ctx, output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", + ggml::hrx::UnaryKind::Relu, 4, 2048, 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * matmul = ggml_mul_mat(ctx, weight, input); + REQUIRE(matmul != nullptr); + ggml_tensor * output = ggml_gelu(ctx, matmul); + REQUIRE(output != nullptr); + schedule_fused_matmul_unary_command(ctx, output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", + ggml::hrx::UnaryKind::Gelu, 4, 2048, 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q5_0, 640, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 640, 2); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * matmul = ggml_mul_mat(ctx, weight, input); + REQUIRE(matmul != nullptr); + ggml_tensor * output = ggml_gelu(ctx, matmul); + REQUIRE(output != nullptr); + schedule_fused_matmul_unary_command(ctx, output, "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32", + ggml::hrx::UnaryKind::Gelu, 2, 640, 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * matmul = ggml_mul_mat(ctx, weight, input); + REQUIRE(matmul != nullptr); + ggml_tensor * relu = ggml_relu(ctx, matmul); + ggml_tensor * sqr = ggml_sqr(ctx, matmul); + REQUIRE(relu != nullptr); + REQUIRE(sqr != nullptr); + ggml_tensor * output = ggml_add(ctx, relu, sqr); + REQUIRE(output != nullptr); + require_matmul_root_matches_identity_single(ctx, output); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * matmul = ggml_mul_mat(ctx, weight, input); + REQUIRE(matmul != nullptr); + ggml_tensor * output = ggml_softplus(ctx, matmul); + REQUIRE(output != nullptr); + require_matmul_root_matches_identity_single(ctx, output); + } + + { + static constexpr const char * kDisableQwenDispatchEnv = "GGML_HRX_DISABLE_QWEN_DISPATCH"; + const char * original_env = std::getenv(kDisableQwenDispatchEnv); + const bool had_original_env = original_env != nullptr; + const std::string original_env_value = had_original_env ? original_env : ""; + REQUIRE(setenv(kDisableQwenDispatchEnv, "1", 1) == 0); + + ggml_tensor * q4_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * q4_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(q4_weight != nullptr); + REQUIRE(q4_input != nullptr); + ggml_tensor * q4_output = ggml_mul_mat(ctx, q4_weight, q4_input); + REQUIRE(q4_output != nullptr); + schedule_single_matmul_command(ctx, q4_output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, 2048, + 128); + + ggml_tensor * q4_decode_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * q4_decode_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(q4_decode_weight != nullptr); + REQUIRE(q4_decode_input != nullptr); + ggml_tensor * q4_decode_output = ggml_mul_mat(ctx, q4_decode_weight, q4_decode_input); + REQUIRE(q4_decode_output != nullptr); + schedule_single_matmul_command(ctx, q4_decode_output, "loom_libs:ggml_mul_mat_f32_f32_wmma", 1, 2048, 128); + + for (ggml_type legacy_type : + { GGML_TYPE_Q1_0, GGML_TYPE_Q4_0, GGML_TYPE_Q4_1, GGML_TYPE_Q5_0, GGML_TYPE_Q5_1 }) { + ggml_tensor * legacy_weight = ggml_new_tensor_2d(ctx, legacy_type, 2048, 128); + ggml_tensor * legacy_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(legacy_weight != nullptr); + REQUIRE(legacy_input != nullptr); + ggml_tensor * legacy_output = ggml_mul_mat(ctx, legacy_weight, legacy_input); + REQUIRE(legacy_output != nullptr); + schedule_single_matmul_command(ctx, legacy_output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, + 2048, 128); + + ggml_tensor * legacy_decode_weight = ggml_new_tensor_2d(ctx, legacy_type, 2048, 128); + ggml_tensor * legacy_decode_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(legacy_decode_weight != nullptr); + REQUIRE(legacy_decode_input != nullptr); + ggml_tensor * legacy_decode_output = ggml_mul_mat(ctx, legacy_decode_weight, legacy_decode_input); + REQUIRE(legacy_decode_output != nullptr); + schedule_single_matmul_command(ctx, legacy_decode_output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", 1, + 2048, 128); + } + + ggml_tensor * q3_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q3_K, 2048, 128); + ggml_tensor * q3_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(q3_weight != nullptr); + REQUIRE(q3_input != nullptr); + ggml_tensor * q3_output = ggml_mul_mat(ctx, q3_weight, q3_input); + REQUIRE(q3_output != nullptr); + schedule_single_matmul_command(ctx, q3_output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, 2048, + 128); + + ggml_tensor * q3_decode_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q3_K, 2048, 128); + ggml_tensor * q3_decode_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(q3_decode_weight != nullptr); + REQUIRE(q3_decode_input != nullptr); + ggml_tensor * q3_decode_output = ggml_mul_mat(ctx, q3_decode_weight, q3_decode_input); + REQUIRE(q3_decode_output != nullptr); + schedule_single_matmul_command(ctx, q3_decode_output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", 1, 2048, + 128); + + ggml_tensor * iq2_s_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ2_S, 2048, 640); + ggml_tensor * iq2_s_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 2); + REQUIRE(iq2_s_weight != nullptr); + REQUIRE(iq2_s_input != nullptr); + ggml_tensor * iq2_s_output = ggml_mul_mat(ctx, iq2_s_weight, iq2_s_input); + REQUIRE(iq2_s_output != nullptr); + schedule_single_matmul_command(ctx, iq2_s_output, + "loom_libs:ggml_mul_mat_tiled_input_f32_iq2_s_publish_f32", 2, 2048, 640); + + ggml_tensor * iq2_s_decode_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ2_S, 2048, 640); + ggml_tensor * iq2_s_decode_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(iq2_s_decode_weight != nullptr); + REQUIRE(iq2_s_decode_input != nullptr); + ggml_tensor * iq2_s_decode_output = ggml_mul_mat(ctx, iq2_s_decode_weight, iq2_s_decode_input); + REQUIRE(iq2_s_decode_output != nullptr); + schedule_single_matmul_command(ctx, iq2_s_decode_output, "loom_libs:ggml_mul_mat_vector_iq2_s_f32_f32", 1, + 2048, 640); + + ggml_tensor * iq3_s_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ3_S, 2048, 128); + ggml_tensor * iq3_s_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(iq3_s_weight != nullptr); + REQUIRE(iq3_s_input != nullptr); + ggml_tensor * iq3_s_output = ggml_mul_mat(ctx, iq3_s_weight, iq3_s_input); + REQUIRE(iq3_s_output != nullptr); + schedule_single_matmul_command(ctx, iq3_s_output, + "loom_libs:ggml_mul_mat_tiled_input_f32_iq3_s_publish_f32", 4, 2048, 128); + + ggml_tensor * iq3_s_decode_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ3_S, 2048, 128); + ggml_tensor * iq3_s_decode_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(iq3_s_decode_weight != nullptr); + REQUIRE(iq3_s_decode_input != nullptr); + ggml_tensor * iq3_s_decode_output = ggml_mul_mat(ctx, iq3_s_decode_weight, iq3_s_decode_input); + REQUIRE(iq3_s_decode_output != nullptr); + schedule_single_matmul_command(ctx, iq3_s_decode_output, "loom_libs:ggml_mul_mat_vector_iq3_s_f32_f32", 1, + 2048, 128); + + ggml_tensor * iq4_nl_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ4_NL, 2048, 128); + ggml_tensor * iq4_nl_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(iq4_nl_weight != nullptr); + REQUIRE(iq4_nl_input != nullptr); + ggml_tensor * iq4_nl_output = ggml_mul_mat(ctx, iq4_nl_weight, iq4_nl_input); + REQUIRE(iq4_nl_output != nullptr); + schedule_single_matmul_command(ctx, iq4_nl_output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, + 2048, 128); + + ggml_tensor * iq4_nl_decode_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ4_NL, 2048, 128); + ggml_tensor * iq4_nl_decode_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(iq4_nl_decode_weight != nullptr); + REQUIRE(iq4_nl_decode_input != nullptr); + ggml_tensor * iq4_nl_decode_output = ggml_mul_mat(ctx, iq4_nl_decode_weight, iq4_nl_decode_input); + REQUIRE(iq4_nl_decode_output != nullptr); + schedule_single_matmul_command(ctx, iq4_nl_decode_output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", 1, + 2048, 128); + + ggml_tensor * iq4_nl_small_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ4_NL, 640, 1024); + ggml_tensor * iq4_nl_small_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 640, 2); + REQUIRE(iq4_nl_small_weight != nullptr); + REQUIRE(iq4_nl_small_input != nullptr); + ggml_tensor * iq4_nl_small_output = ggml_mul_mat(ctx, iq4_nl_small_weight, iq4_nl_small_input); + REQUIRE(iq4_nl_small_output != nullptr); + schedule_single_matmul_command(ctx, iq4_nl_small_output, "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32", + 2, 640, 1024); + + ggml_tensor * iq4_nl_small_decode_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ4_NL, 640, 1024); + ggml_tensor * iq4_nl_small_decode_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 640, 1); + REQUIRE(iq4_nl_small_decode_weight != nullptr); + REQUIRE(iq4_nl_small_decode_input != nullptr); + ggml_tensor * iq4_nl_small_decode_output = + ggml_mul_mat(ctx, iq4_nl_small_decode_weight, iq4_nl_small_decode_input); + REQUIRE(iq4_nl_small_decode_output != nullptr); + schedule_single_matmul_command(ctx, iq4_nl_small_decode_output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", + 1, 640, 1024); + + ggml_tensor * q5_0_small_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q5_0, 640, 256); + ggml_tensor * q5_0_small_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 640, 2); + REQUIRE(q5_0_small_weight != nullptr); + REQUIRE(q5_0_small_input != nullptr); + ggml_tensor * q5_0_small_output = ggml_mul_mat(ctx, q5_0_small_weight, q5_0_small_input); + REQUIRE(q5_0_small_output != nullptr); + schedule_single_matmul_command(ctx, q5_0_small_output, "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32", 2, + 640, 256); + + ggml_tensor * q5_0_small_decode_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q5_0, 640, 256); + ggml_tensor * q5_0_small_decode_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 640, 1); + REQUIRE(q5_0_small_decode_weight != nullptr); + REQUIRE(q5_0_small_decode_input != nullptr); + ggml_tensor * q5_0_small_decode_output = ggml_mul_mat(ctx, q5_0_small_decode_weight, q5_0_small_decode_input); + REQUIRE(q5_0_small_decode_output != nullptr); + schedule_single_matmul_command(ctx, q5_0_small_decode_output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", 1, + 640, 256); + + ggml_tensor * q6_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, 2048, 128); + ggml_tensor * q6_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 2); + REQUIRE(q6_weight != nullptr); + REQUIRE(q6_input != nullptr); + ggml_tensor * q6_output = ggml_mul_mat(ctx, q6_weight, q6_input); + REQUIRE(q6_output != nullptr); + schedule_single_matmul_command(ctx, q6_output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 2, 2048, + 128); + + ggml_tensor * q6_decode_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, 2048, 128); + ggml_tensor * q6_decode_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(q6_decode_weight != nullptr); + REQUIRE(q6_decode_input != nullptr); + ggml_tensor * q6_decode_output = ggml_mul_mat(ctx, q6_decode_weight, q6_decode_input); + REQUIRE(q6_decode_output != nullptr); + schedule_single_matmul_command(ctx, q6_decode_output, "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma", 1, + 2048, 128); + + ggml_tensor * q6_gemma_decode_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, 3840, 262208); + ggml_tensor * q6_gemma_decode_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 3840, 1); + REQUIRE(q6_gemma_decode_weight != nullptr); + REQUIRE(q6_gemma_decode_input != nullptr); + ggml_tensor * q6_gemma_decode_output = ggml_mul_mat(ctx, q6_gemma_decode_weight, q6_gemma_decode_input); + REQUIRE(q6_gemma_decode_output != nullptr); + schedule_single_matmul_command(ctx, q6_gemma_decode_output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", 1, + 3840, 262208); + + ggml_tensor * iq4_xs_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ4_XS, 2048, 128); + ggml_tensor * iq4_xs_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(iq4_xs_weight != nullptr); + REQUIRE(iq4_xs_input != nullptr); + ggml_tensor * iq4_xs_output = ggml_mul_mat(ctx, iq4_xs_weight, iq4_xs_input); + REQUIRE(iq4_xs_output != nullptr); + schedule_single_matmul_command(ctx, iq4_xs_output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, + 2048, 128); + + ggml_tensor * iq4_xs_decode_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ4_XS, 2048, 128); + ggml_tensor * iq4_xs_decode_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(iq4_xs_decode_weight != nullptr); + REQUIRE(iq4_xs_decode_input != nullptr); + ggml_tensor * iq4_xs_decode_output = ggml_mul_mat(ctx, iq4_xs_decode_weight, iq4_xs_decode_input); + REQUIRE(iq4_xs_decode_output != nullptr); + schedule_single_matmul_command(ctx, iq4_xs_decode_output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", 1, + 2048, 128); + + ggml_tensor * q8_0_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q8_0, 2048, 128); + ggml_tensor * q8_0_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(q8_0_weight != nullptr); + REQUIRE(q8_0_input != nullptr); + ggml_tensor * q8_0_output = ggml_mul_mat(ctx, q8_0_weight, q8_0_input); + REQUIRE(q8_0_output != nullptr); + schedule_single_matmul_command(ctx, q8_0_output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, 2048, + 128); + + ggml_tensor * q8_0_decode_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q8_0, 2048, 128); + ggml_tensor * q8_0_decode_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(q8_0_decode_weight != nullptr); + REQUIRE(q8_0_decode_input != nullptr); + ggml_tensor * q8_0_decode_output = ggml_mul_mat(ctx, q8_0_decode_weight, q8_0_decode_input); + REQUIRE(q8_0_decode_output != nullptr); + schedule_single_matmul_command(ctx, q8_0_decode_output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", 1, 2048, + 128); + + ggml_tensor * q8_0_gemma_decode_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q8_0, 3840, 262208); + ggml_tensor * q8_0_gemma_decode_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 3840, 1); + REQUIRE(q8_0_gemma_decode_weight != nullptr); + REQUIRE(q8_0_gemma_decode_input != nullptr); + ggml_tensor * q8_0_gemma_decode_output = ggml_mul_mat(ctx, q8_0_gemma_decode_weight, q8_0_gemma_decode_input); + REQUIRE(q8_0_gemma_decode_output != nullptr); + schedule_single_matmul_command(ctx, q8_0_gemma_decode_output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", 1, + 3840, 262208); + + ggml_tensor * q8_0_unsafe_decode_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q8_0, 32768, 262144); + ggml_tensor * q8_0_unsafe_decode_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 32768, 1); + REQUIRE(q8_0_unsafe_decode_weight != nullptr); + REQUIRE(q8_0_unsafe_decode_input != nullptr); + ggml_tensor * q8_0_unsafe_decode_output = + ggml_mul_mat(ctx, q8_0_unsafe_decode_weight, q8_0_unsafe_decode_input); + REQUIRE(q8_0_unsafe_decode_output != nullptr); + REQUIRE(!matmul_graph_is_supported(ctx, q8_0_unsafe_decode_output)); + + ggml_tensor * q8_1_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q8_1, 2048, 128); + ggml_tensor * q8_1_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(q8_1_weight != nullptr); + REQUIRE(q8_1_input != nullptr); + ggml_tensor * q8_1_output = ggml_mul_mat(ctx, q8_1_weight, q8_1_input); + REQUIRE(q8_1_output != nullptr); + schedule_single_matmul_command(ctx, q8_1_output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, 2048, + 128); + + ggml_tensor * q8_1_decode_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q8_1, 2048, 128); + ggml_tensor * q8_1_decode_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(q8_1_decode_weight != nullptr); + REQUIRE(q8_1_decode_input != nullptr); + ggml_tensor * q8_1_decode_output = ggml_mul_mat(ctx, q8_1_decode_weight, q8_1_decode_input); + REQUIRE(q8_1_decode_output != nullptr); + schedule_single_matmul_command(ctx, q8_1_decode_output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", 1, 2048, + 128); + + ggml_tensor * f16_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 2048, 128); + ggml_tensor * f16_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(f16_weight != nullptr); + REQUIRE(f16_input != nullptr); + ggml_tensor * f16_output = ggml_mul_mat(ctx, f16_weight, f16_input); + REQUIRE(f16_output != nullptr); + schedule_single_matmul_command(ctx, f16_output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, 2048, + 128); + + ggml_tensor * f16_decode_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 2048, 128); + ggml_tensor * f16_decode_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(f16_decode_weight != nullptr); + REQUIRE(f16_decode_input != nullptr); + ggml_tensor * f16_decode_output = ggml_mul_mat(ctx, f16_decode_weight, f16_decode_input); + REQUIRE(f16_decode_output != nullptr); + schedule_single_matmul_command(ctx, f16_decode_output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", 1, 2048, + 128); + + ggml_tensor * f32_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 128); + ggml_tensor * f32_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(f32_weight != nullptr); + REQUIRE(f32_input != nullptr); + ggml_tensor * f32_output = ggml_mul_mat(ctx, f32_weight, f32_input); + REQUIRE(f32_output != nullptr); + schedule_single_matmul_command(ctx, f32_output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, 2048, + 128); + + ggml_tensor * f32_decode_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 128); + ggml_tensor * f32_decode_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(f32_decode_weight != nullptr); + REQUIRE(f32_decode_input != nullptr); + ggml_tensor * f32_decode_output = ggml_mul_mat(ctx, f32_decode_weight, f32_decode_input); + REQUIRE(f32_decode_output != nullptr); + schedule_single_matmul_command(ctx, f32_decode_output, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64", 1, 2048, + 128); + + restore_environment_value(kDisableQwenDispatchEnv, had_original_env, original_env_value); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, 2048, 151936); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_q6_k_packed_token1_f16_wmma", 1, 2048, + 151936); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "qwen3_moe:qwen3_moe_router_projection_f32_four_row_wave32", 4, + 2048, 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "qwen3_moe:qwen3_moe_router_projection_f32_four_row_wave32", 1, + 2048, 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 4); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, 2048, + 128); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 2048, 4); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + REQUIRE(!matmul_graph_is_supported(ctx, output)); + } + + { + ggml_tensor * weight = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 2048, 128, 2); + ggml_tensor * input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 2048, 4, 2); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + REQUIRE(!matmul_graph_is_supported(ctx, output)); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 128, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 128, 4); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + REQUIRE(!matmul_graph_is_supported(ctx, output)); + } + + ggml_free(ctx); +} + +static void run_iq3_xxs_codebook_matmul_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + auto check_plan = [&](int64_t token_count, const char * expected_kernel, size_t expected_binding_count, + const char * weight_format_parameter) { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ3_XXS, 2048, 128); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, token_count); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const ggml::hrx::CommandPlan & plan = scheduler.plan(); + REQUIRE(plan.dispatches.size() == 1); + REQUIRE(kernel_name_for_id(plan.dispatches[0].kernel.kernel_id) == expected_kernel); + REQUIRE(plan.dispatches[0].bindings.size() == expected_binding_count); + REQUIRE(plan.dispatches[0].bindings[2].length == 1024); + REQUIRE(plan.transients.size() == 1); + REQUIRE(plan.constant_initializations.size() == 1); + REQUIRE(plan.constant_initializations[0].value == plan.dispatches[0].bindings[2].value); + REQUIRE(plan.constant_initializations[0].data.size() == 1024); + require_compile_parameter(plan.dispatches[0], weight_format_parameter, "18"); + }; + + check_plan(1, "loom_libs:ggml_mul_mat_vector_iq3_xxs_f32_f32", 5, "ggml.matmul.vector.weight_format"); + check_plan(64, "loom_libs:ggml_mul_mat_tiled_input_f32_iq3_xxs_publish_f32", 4, + "ggml.mul_mat.weight_format"); + + const char * disable_name = "GGML_HRX_DISABLE_IQ3_XXS_CODEBOOK_MATMUL"; + REQUIRE(setenv(disable_name, "1", 1) == 0); + ggml_tensor * disabled_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_IQ3_XXS, 2048, 128); + ggml_tensor * disabled_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + REQUIRE(disabled_weight != nullptr); + REQUIRE(disabled_input != nullptr); + ggml_tensor * disabled_output = ggml_mul_mat(ctx, disabled_weight, disabled_input); + REQUIRE(disabled_output != nullptr); + ggml_cgraph * disabled_graph = ggml_new_graph(ctx); + REQUIRE(disabled_graph != nullptr); + ggml_build_forward_expand(disabled_graph, disabled_output); + ggml::hrx::GraphImportResult disabled_imported = ggml::hrx::import_ggml_graph(*disabled_graph); + REQUIRE(disabled_imported.valid()); + ggml::hrx::DispatchScheduler disabled_scheduler; + REQUIRE(!disabled_scheduler.schedule_graph(disabled_imported.graph, test_dispatch_target())); + REQUIRE(unsetenv(disable_name) == 0); + + ggml_free(ctx); +} + +static void run_iq2_codebook_resource_reuse_checks() { + for (const ggml_type type : { GGML_TYPE_IQ2_XXS, GGML_TYPE_IQ2_XS }) { + ggml_init_params params = {}; + params.mem_size = 8 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * vector_weight = ggml_new_tensor_2d(ctx, type, 2048, 128); + ggml_tensor * vector_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 1); + ggml_tensor * tiled_weight = ggml_new_tensor_2d(ctx, type, 2048, 128); + ggml_tensor * tiled_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 2048, 64); + REQUIRE(vector_weight != nullptr); + REQUIRE(vector_input != nullptr); + REQUIRE(tiled_weight != nullptr); + REQUIRE(tiled_input != nullptr); + ggml_tensor * vector_output = ggml_mul_mat(ctx, vector_weight, vector_input); + ggml_tensor * tiled_output = ggml_mul_mat(ctx, tiled_weight, tiled_input); + REQUIRE(vector_output != nullptr); + REQUIRE(tiled_output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, vector_output); + ggml_build_forward_expand(graph, tiled_output); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 2); + REQUIRE(scheduler.plan().constant_initializations.size() == 1); + REQUIRE(scheduler.plan().transients.size() == 1); + + const char * expected_key = + type == GGML_TYPE_IQ2_XXS ? "common.iq2_xxs.grid.v1" : "common.iq2_xs.grid.v1"; + const ggml::hrx::CommandPlanConstantInitialization & initialization = + scheduler.plan().constant_initializations.front(); + REQUIRE(initialization.name == expected_key); + REQUIRE(scheduler.plan().transients.front().name == expected_key); + REQUIRE(scheduler.plan().transients.front().value == initialization.value); + REQUIRE(std::count_if(scheduler.plan().dispatches.begin(), scheduler.plan().dispatches.end(), + [&](const ggml::hrx::Dispatch & dispatch) { + return std::any_of(dispatch.bindings.begin(), dispatch.bindings.end(), + [&](const ggml::hrx::DispatchBinding & binding) { + return binding.value == initialization.value; + }); + }) == 2); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.constant_initializations.size() == 1); + ggml_free(ctx); + } +} + +static void run_iq4_matmul_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 2 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + const auto check = [ctx](ggml_type weight_type, int64_t input_size, int64_t output_size, + int64_t token_count, const char * expected_kernel) { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, weight_type, input_size, output_size); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, input_size, token_count); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + schedule_single_matmul_command(ctx, output, expected_kernel, token_count, input_size, output_size); + }; + + check(GGML_TYPE_IQ4_NL, 640, 1024, 1, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64"); + check(GGML_TYPE_IQ4_NL, 640, 1024, 64, "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32_aligned"); + check(GGML_TYPE_IQ4_XS, 2048, 128, 1, "loom_libs:ggml_mul_mat_f32_f32_decode_wave64"); + check(GGML_TYPE_IQ4_XS, 2048, 128, 64, "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32_aligned"); + + ggml_free(ctx); +} + +static void run_quantized_value_projection_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 16 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + for (ggml_type weight_type : { GGML_TYPE_Q4_K, GGML_TYPE_Q6_K }) { + for (ggml_type cache_type : { GGML_TYPE_F16, GGML_TYPE_F32 }) { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, weight_type, 768, 256); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 768, 1); + ggml_tensor * cache = ggml_new_tensor_2d(ctx, cache_type, 256, 17); + ggml_tensor * indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, 1); + ggml_tensor * output = ggml_set_rows(ctx, cache, ggml_mul_mat(ctx, weight, input), indices); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, output); + auto imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const auto & plan = scheduler.plan(); + REQUIRE(plan.valid()); + REQUIRE(plan.dispatches.size() == 1 || plan.dispatches.size() == 2); + const auto & dispatch = plan.dispatches.back(); + const std::string kernel_name = kernel_name_for_id(dispatch.kernel.kernel_id); + if (plan.dispatches.size() == 1) { + REQUIRE(kernel_name == "loom_libs:llm_attention_v_matmul_set_rows_vector_f32_f32"); + require_compile_parameter(dispatch, "llm.attention_qkv.weight_format", + std::to_string(matmul_weight_format_config(weight_type))); + } else { + REQUIRE(kernel_name_for_id(plan.dispatches.front().kernel.kernel_id) == + "qwen3_moe:ggml_quantize_q8_1_x4_f32"); + REQUIRE(kernel_name == "loom_libs:llm_attention_v_matmul_set_rows_legacy_f32_f32"); + require_compile_parameter(dispatch, "llm.attention_qkv.weight_format", + weight_type == GGML_TYPE_Q4_K ? "44" : "46"); + require_compile_parameter(dispatch, "ggml.mul_mat.activation_format", std::to_string(GGML_TYPE_Q8_1)); + REQUIRE(!dispatch.bindings[1].layout.empty()); + } + require_compile_parameter(dispatch, "llm.attention_qkv.cache_output_format", + cache_type == GGML_TYPE_F16 ? "16" : "32"); + const auto commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + } + ggml_free(ctx); +} + +static void run_prefill_value_projection_cache_dispatch_checks() { + static constexpr const char * kDisableFusionEnv = "GGML_HRX_DISABLE_PREFILL_V_CACHE_FUSION"; + EnvironmentVariableGuard disable_fusion(kDisableFusionEnv); + disable_fusion.unset(); + + auto check_case = [&](ggml_type weight_type, int64_t input_size, int64_t output_size, int64_t token_count, + ggml_type cache_type, int64_t cache_row_count, bool publish_f16_alternate, + bool set_rows_available, bool disable, bool expect_fusion, bool observe_projection = false, + bool alias_weight = false) { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * weight_storage = ggml_new_tensor_2d(ctx, weight_type, input_size, output_size); + REQUIRE(weight_storage != nullptr); + ggml_tensor * weight = + alias_weight ? ggml_view_2d(ctx, weight_storage, input_size, output_size, weight_storage->nb[1], 0) : + weight_storage; + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * cache = ggml_new_tensor_2d(ctx, cache_type, output_size, cache_row_count); + ggml_tensor * indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, token_count); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + REQUIRE(cache != nullptr); + REQUIRE(indices != nullptr); + ggml_tensor * projection = ggml_mul_mat(ctx, weight, input); + ggml_tensor * value = ggml_reshape_3d(ctx, projection, 128, output_size / 128, token_count); + ggml_tensor * rows = ggml_reshape_2d(ctx, value, output_size, token_count); + ggml_tensor * output = ggml_set_rows(ctx, cache, rows, indices); + REQUIRE(projection != nullptr); + REQUIRE(value != nullptr); + REQUIRE(rows != nullptr); + REQUIRE(output != nullptr); + if (observe_projection) { + ggml_set_output(projection); + } + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::CommandPlan plan; + const ggml::hrx::Value * input_value = imported.graph.values().find_tensor(input); + REQUIRE(input_value != nullptr); + if (publish_f16_alternate) { + const size_t bytes = static_cast(input_size * token_count) * sizeof(ggml_fp16_t); + const ggml::hrx::ValueId alternate(static_cast(imported.graph.values().size())); + plan.transients.push_back({ alternate, "test.prefill_v_cache.f16", bytes, 256 }); + plan.constant_initializations.push_back( + { alternate, "test.prefill_v_cache.f16", 0, std::vector(bytes) }); + ggml::hrx::Status status; + REQUIRE(plan.metadata.append_alternate_value( + { input_value->id, alternate, GGML_TYPE_F16, bytes, "test.prefill_v_cache.f16" }, status)); + } + + std::vector covered_nodes(imported.graph.nodes().size(), false); + if (!set_rows_available) { + covered_nodes[producer_index_for_tensor(imported.graph, output)] = true; + } + if (disable) { + disable_fusion.set("1"); + } else { + disable_fusion.unset(); + } + + ggml::hrx::DispatchMatch match; + const size_t projection_index = producer_index_for_tensor(imported.graph, projection); + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, projection_index, match)); + REQUIRE(!match.dispatches.empty()); + const ggml::hrx::Dispatch & dispatch = match.dispatches.back(); + const bool fused = + kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:llm_attention_v_matmul_set_rows_tiled_f32_f32"; + REQUIRE(fused == expect_fusion); + if (expect_fusion) { + REQUIRE(match.dispatches.size() == 1); + REQUIRE(match.covered_nodes.size() == 4); + REQUIRE(dispatch.bindings.size() == 4); + REQUIRE(dispatch.bindings[0].value == plan.transients.front().value); + REQUIRE(dispatch.bindings[0].length == plan.transients.front().size); + require_compile_parameter(dispatch, "ggml.mul_mat.activation_format", std::to_string(GGML_TYPE_F16)); + require_compile_parameter(dispatch, "llm.attention_qkv.weight_format", "6"); + require_compile_parameter(dispatch, "llm.attention_qkv.input_size", "1536"); + require_compile_parameter(dispatch, "llm.attention_qkv.output_size", "256"); + require_compile_parameter(dispatch, "llm.attention_qkv.cache_row_count", "512"); + require_compile_parameter(dispatch, "llm.attention_qkv.cache_output_format", "16"); + + append_match_to_plan(plan, match, covered_nodes, &imported.graph); + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.size() == 1); + REQUIRE(commands.commands.front().bindings.size() == 4); + REQUIRE(commands.commands.front().bindings[0].name == "input"); + REQUIRE(commands.commands.front().bindings[1].name == "weight"); + REQUIRE(commands.commands.front().bindings[2].name == "indices"); + REQUIRE(commands.commands.front().bindings[3].name == "cache"); + } + ggml_free(ctx); + }; + + check_case(GGML_TYPE_Q6_K, 1536, 256, 64, GGML_TYPE_F16, 512, true, true, false, true); + check_case(GGML_TYPE_Q6_K, 1536, 256, 64, GGML_TYPE_F16, 512, true, true, true, false); + check_case(GGML_TYPE_Q6_K, 1536, 256, 64, GGML_TYPE_F16, 512, false, true, false, false); + check_case(GGML_TYPE_Q4_K, 1536, 256, 64, GGML_TYPE_F16, 512, true, true, false, false); + check_case(GGML_TYPE_Q6_K, 2048, 256, 64, GGML_TYPE_F16, 512, true, true, false, false); + check_case(GGML_TYPE_Q6_K, 1536, 512, 64, GGML_TYPE_F16, 512, true, true, false, false); + check_case(GGML_TYPE_Q6_K, 1536, 256, 32, GGML_TYPE_F16, 512, true, true, false, false); + check_case(GGML_TYPE_Q6_K, 1536, 256, 64, GGML_TYPE_F32, 512, true, true, false, false); + check_case(GGML_TYPE_Q6_K, 1536, 256, 64, GGML_TYPE_F16, 1024, true, true, false, false); + check_case(GGML_TYPE_Q6_K, 1536, 256, 64, GGML_TYPE_F16, 512, true, false, false, false); + check_case(GGML_TYPE_Q6_K, 1536, 256, 64, GGML_TYPE_F16, 512, true, true, false, false, true); + check_case(GGML_TYPE_Q6_K, 1536, 256, 64, GGML_TYPE_F16, 512, true, true, false, false, false, true); +} + +static void run_llama_attention_matmul_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 16 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t input_size = 2048; + constexpr int64_t head_size = 128; + constexpr int64_t head_count = 2; + constexpr int64_t output_size = head_size * head_count; + constexpr int64_t token_count = 7; + constexpr int64_t cache_row_count = 16; + + const ggml_type weight_types[] = { GGML_TYPE_F16, GGML_TYPE_BF16 }; + for (ggml_type weight_type : weight_types) { + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, weight_type, input_size, output_size); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * positions = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, token_count); + ggml_tensor * freqs = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, head_size / 2); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + REQUIRE(positions != nullptr); + REQUIRE(freqs != nullptr); + ggml_tensor * projection = ggml_mul_mat(ctx, weight, input); + ggml_tensor * query = ggml_reshape_3d(ctx, projection, head_size, head_count, token_count); + ggml_tensor * output = ggml_rope_ext(ctx, query, positions, freqs, head_size, GGML_ROPE_TYPE_NORMAL, 0, + 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + REQUIRE(projection != nullptr); + REQUIRE(query != nullptr); + REQUIRE(output != nullptr); + schedule_fused_llama_attention_matmul_command( + ctx, output, "loom_libs:llm_attention_q_matmul_rope_tiled_f32_f32", weight_type, token_count, + input_size, output_size, head_size, head_count, 0, 0, 3, + { GGML_OP_MUL_MAT, GGML_OP_RESHAPE, GGML_OP_ROPE }, + { "input", "weight", "positions", "theta", "freq_factors", "output" }); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, weight_type, input_size, output_size); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * positions = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, token_count); + ggml_tensor * freqs = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, head_size / 2); + ggml_tensor * cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, output_size, cache_row_count); + ggml_tensor * indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, token_count); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + REQUIRE(positions != nullptr); + REQUIRE(freqs != nullptr); + REQUIRE(cache != nullptr); + REQUIRE(indices != nullptr); + ggml_tensor * projection = ggml_mul_mat(ctx, weight, input); + ggml_tensor * key = ggml_reshape_3d(ctx, projection, head_size, head_count, token_count); + ggml_tensor * rope = ggml_rope_ext(ctx, key, positions, freqs, head_size, GGML_ROPE_TYPE_NORMAL, 0, + 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + ggml_tensor * rows = ggml_reshape_2d(ctx, rope, output_size, token_count); + ggml_tensor * output = ggml_set_rows(ctx, cache, rows, indices); + REQUIRE(projection != nullptr); + REQUIRE(key != nullptr); + REQUIRE(rope != nullptr); + REQUIRE(rows != nullptr); + REQUIRE(output != nullptr); + schedule_fused_llama_attention_matmul_command( + ctx, output, "loom_libs:llm_attention_k_matmul_rope_set_rows_tiled_f32_f32", weight_type, token_count, + input_size, output_size, head_size, head_count, cache_row_count, 16, 5, + { GGML_OP_MUL_MAT, GGML_OP_RESHAPE, GGML_OP_ROPE, GGML_OP_RESHAPE, GGML_OP_SET_ROWS }, + { "input", "weight", "positions", "indices", "theta", "freq_factors", "cache" }); + } + + { + ggml_tensor * weight = ggml_new_tensor_2d(ctx, weight_type, input_size, output_size); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, output_size, cache_row_count); + ggml_tensor * indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, token_count); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + REQUIRE(cache != nullptr); + REQUIRE(indices != nullptr); + ggml_tensor * projection = ggml_mul_mat(ctx, weight, input); + ggml_tensor * value = ggml_reshape_3d(ctx, projection, head_size, head_count, token_count); + ggml_tensor * rows = ggml_reshape_2d(ctx, value, output_size, token_count); + ggml_tensor * output = ggml_set_rows(ctx, cache, rows, indices); + REQUIRE(projection != nullptr); + REQUIRE(value != nullptr); + REQUIRE(rows != nullptr); + REQUIRE(output != nullptr); + schedule_fused_llama_attention_matmul_command( + ctx, output, "loom_libs:llm_attention_v_matmul_set_rows_tiled_f32_f32", weight_type, token_count, + input_size, output_size, 0, 0, cache_row_count, 16, 4, + { GGML_OP_MUL_MAT, GGML_OP_RESHAPE, GGML_OP_RESHAPE, GGML_OP_SET_ROWS }, + { "input", "weight", "indices", "cache" }); + } + } + + ggml_free(ctx); +} + +static void schedule_qwen_router_top8_command(ggml_context * ctx, + ggml_tensor * output, + ggml_tensor * route_ids, + int64_t expected_expert_count = kQwenRouterExpertCount, + int64_t expected_route_count = kQwenRouterRouteCount) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 10); + + const ggml::hrx::Value * route_ids_value = imported.graph.values().find_tensor(route_ids); + const ggml::hrx::Value * output_value = imported.graph.values().find_tensor(output); + REQUIRE(route_ids_value != nullptr); + REQUIRE(output_value != nullptr); + REQUIRE(route_ids_value->kind == ggml::hrx::ValueKind::Transient); + REQUIRE(output_value->kind == ggml::hrx::ValueKind::External); + + const ggml::hrx::GraphNode * softmax_node = nullptr; + const ggml::hrx::GraphNode * get_rows_node = nullptr; + for (const ggml::hrx::GraphNode & node : imported.graph.nodes()) { + if (node.op == GGML_OP_SOFT_MAX) { + softmax_node = &node; + } else if (node.op == GGML_OP_GET_ROWS) { + get_rows_node = &node; + } + } + REQUIRE(softmax_node != nullptr); + REQUIRE(get_rows_node != nullptr); + const ggml::hrx::GraphNode * probs_reshape = + ggml::hrx::find_single_layout_alias_consumer_with_op(imported.graph, softmax_node->output, GGML_OP_RESHAPE); + REQUIRE(probs_reshape != nullptr); + REQUIRE(ggml::hrx::is_layout_alias_node(imported.graph, *probs_reshape)); + REQUIRE(ggml::hrx::find_single_consumer_with_op_through_layout_aliases(imported.graph, softmax_node->output, + GGML_OP_GET_ROWS) == get_rows_node); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + + const int64_t token_count = output->ne[2]; + const size_t route_id_length = static_cast(token_count) * route_ids->nb[1]; + const size_t expert_table_bytes = qwen_expert_table_size(token_count, expected_expert_count); + const size_t partition_table_bytes = + qwen_partition_table_size(token_count, expected_route_count, expected_expert_count); + const bool uses_fused_prefill_expert_table_partition = + token_count == 512 && expected_route_count == 8 && + route_ids->nb[1] / sizeof(int32_t) == static_cast(expected_route_count) && expected_expert_count == 128; + REQUIRE(scheduler.plan().dispatches.size() == (uses_fused_prefill_expert_table_partition ? 2 : 3)); + REQUIRE(scheduler.plan().transients.size() == 2); + REQUIRE(scheduler.plan().constant_initializations.empty()); + REQUIRE(scheduler.plan().completion_counter_requests.size() == (uses_fused_prefill_expert_table_partition ? 1 : 0)); + + const ggml::hrx::CommandPlanTransient & expert_table_transient = scheduler.plan().transients[0]; + const ggml::hrx::CommandPlanTransient & partition_table_transient = scheduler.plan().transients[1]; + REQUIRE(expert_table_transient.value.value == static_cast(imported.graph.values().size())); + REQUIRE(expert_table_transient.name == "qwen.router.expert_table"); + REQUIRE(expert_table_transient.size == expert_table_bytes); + REQUIRE(partition_table_transient.value.value == expert_table_transient.value.value + 1); + REQUIRE(partition_table_transient.name == "qwen.router.partition_table"); + REQUIRE(partition_table_transient.size == partition_table_bytes); + if (uses_fused_prefill_expert_table_partition) { + const ggml::hrx::CommandPlanCompletionCounterRequest & completion_counter_request = + scheduler.plan().completion_counter_requests[0]; + REQUIRE(completion_counter_request.value.value == partition_table_transient.value.value + 1); + REQUIRE(completion_counter_request.name == "qwen.router.prefill_expert_table_partition_completion_counter"); + REQUIRE(completion_counter_request.count == 1); + } + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches[0]; + const std::string kernel_name = kernel_name_for_id(dispatch.kernel.kernel_id); + REQUIRE(kernel_name == "qwen3_moe:qwen3_moe_router_top8_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == token_count); + REQUIRE(dispatch.kernel.integer_parameters.at("route_id_stride") == route_ids->nb[1] / sizeof(int32_t)); + REQUIRE(dispatch.bindings.size() == 3); + REQUIRE(dispatch.bindings[1].value == route_ids_value->id); + REQUIRE(dispatch.bindings[1].length == route_id_length); + REQUIRE(dispatch.bindings[2].value == output_value->id); + require_compile_parameter(dispatch, "qwen3_moe.router.expert_count", std::to_string(expected_expert_count)); + require_compile_parameter(dispatch, "qwen3_moe.router.route_count", std::to_string(expected_route_count)); + require_compile_parameter(dispatch, "qwen3_moe.workload.token_capacity", std::to_string(token_count)); + + if (uses_fused_prefill_expert_table_partition) { + const ggml::hrx::CommandPlanCompletionCounterRequest & completion_counter_request = + scheduler.plan().completion_counter_requests[0]; + const ggml::hrx::Dispatch & expert_table_partition_dispatch = scheduler.plan().dispatches[1]; + REQUIRE(kernel_name_for_id(expert_table_partition_dispatch.kernel.kernel_id) == + "qwen3_moe:qwen3_moe_build_expert_table_partition_prefill_512"); + REQUIRE(expert_table_partition_dispatch.kernel.integer_parameters.at("token_count") == token_count); + REQUIRE(expert_table_partition_dispatch.kernel.integer_parameters.at("route_count") == expected_route_count); + REQUIRE(expert_table_partition_dispatch.kernel.integer_parameters.at("route_stride") == + route_ids->nb[1] / sizeof(int32_t)); + REQUIRE(expert_table_partition_dispatch.kernel.integer_parameters.at("expert_count") == expected_expert_count); + REQUIRE(expert_table_partition_dispatch.bindings.size() == 4); + REQUIRE(expert_table_partition_dispatch.bindings[0].value == route_ids_value->id); + REQUIRE(expert_table_partition_dispatch.bindings[0].length == route_id_length); + REQUIRE(expert_table_partition_dispatch.bindings[1].value == expert_table_transient.value); + REQUIRE(expert_table_partition_dispatch.bindings[1].length == expert_table_bytes); + REQUIRE(expert_table_partition_dispatch.bindings[2].value == partition_table_transient.value); + REQUIRE(expert_table_partition_dispatch.bindings[2].length == partition_table_bytes); + REQUIRE(expert_table_partition_dispatch.bindings[3].value == completion_counter_request.value); + REQUIRE(expert_table_partition_dispatch.bindings[3].length == sizeof(int32_t)); + } else { + const ggml::hrx::Dispatch & expert_table_dispatch = scheduler.plan().dispatches[1]; + REQUIRE(kernel_name_for_id(expert_table_dispatch.kernel.kernel_id) == "loom_libs:ggml_moe_build_expert_table"); + REQUIRE(expert_table_dispatch.kernel.integer_parameters.at("token_count") == token_count); + REQUIRE(expert_table_dispatch.kernel.integer_parameters.at("route_count") == expected_route_count); + REQUIRE(expert_table_dispatch.kernel.integer_parameters.at("route_stride") == + route_ids->nb[1] / sizeof(int32_t)); + REQUIRE(expert_table_dispatch.kernel.integer_parameters.at("expert_count") == expected_expert_count); + REQUIRE(expert_table_dispatch.bindings.size() == 2); + REQUIRE(expert_table_dispatch.bindings[0].value == route_ids_value->id); + REQUIRE(expert_table_dispatch.bindings[0].length == route_id_length); + REQUIRE(expert_table_dispatch.bindings[1].value == expert_table_transient.value); + REQUIRE(expert_table_dispatch.bindings[1].length == expert_table_bytes); + require_compile_parameter(expert_table_dispatch, "ggml.moe_routing.expert_count", + std::to_string(expected_expert_count)); + require_compile_parameter(expert_table_dispatch, "ggml.moe_routing.route_count", + std::to_string(expected_route_count)); + require_compile_parameter(expert_table_dispatch, "ggml.workload.token_capacity", std::to_string(token_count)); + + const ggml::hrx::Dispatch & partition_table_dispatch = scheduler.plan().dispatches[2]; + REQUIRE(kernel_name_for_id(partition_table_dispatch.kernel.kernel_id) == + "loom_libs:ggml_moe_build_expert_partition_table"); + REQUIRE(partition_table_dispatch.kernel.integer_parameters.at("token_count") == token_count); + REQUIRE(partition_table_dispatch.kernel.integer_parameters.at("route_count") == expected_route_count); + REQUIRE(partition_table_dispatch.kernel.integer_parameters.at("expert_count") == expected_expert_count); + REQUIRE(partition_table_dispatch.bindings.size() == 2); + REQUIRE(partition_table_dispatch.bindings[0].value == expert_table_transient.value); + REQUIRE(partition_table_dispatch.bindings[0].length == expert_table_bytes); + REQUIRE(partition_table_dispatch.bindings[1].value == partition_table_transient.value); + REQUIRE(partition_table_dispatch.bindings[1].length == partition_table_bytes); + require_compile_parameter(partition_table_dispatch, "ggml.moe_routing.expert_count", + std::to_string(expected_expert_count)); + require_compile_parameter(partition_table_dispatch, "ggml.moe_routing.route_count", + std::to_string(expected_route_count)); + require_compile_parameter(partition_table_dispatch, "ggml.workload.token_capacity", + std::to_string(token_count)); + require_compile_parameter(partition_table_dispatch, "ggml.moe_routing.descriptor_expert_mask", "127"); + require_compile_parameter(partition_table_dispatch, "ggml.moe_routing.descriptor_partition_shift", "7"); + require_compile_parameter(partition_table_dispatch, "ggml.moe_routing.descriptor_row_count_shift", "13"); + require_compile_parameter(partition_table_dispatch, "ggml.moe_routing.partition_workgroup_size", "128"); + } + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == (uses_fused_prefill_expert_table_partition ? 2 : 3)); + REQUIRE(commands.constant_initializations.empty()); + REQUIRE(commands.completion_counters.count == (uses_fused_prefill_expert_table_partition ? 1 : 0)); + REQUIRE(commands.completion_counters.byte_count == + (uses_fused_prefill_expert_table_partition ? sizeof(int32_t) : 0)); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands[0].bindings.size() == 3); + REQUIRE(commands.commands[0].bindings[0].name == "logits"); + REQUIRE(commands.commands[0].bindings[1].name == "route_ids"); + REQUIRE(commands.commands[0].bindings[1].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands[0].bindings[1].value == route_ids_value->storage_root); + REQUIRE(commands.commands[0].bindings[1].offset == route_ids_value->storage_offset); + REQUIRE(commands.commands[0].bindings[1].length == dispatch.bindings[1].length); + REQUIRE(commands.commands[0].bindings[2].name == "route_weights"); + REQUIRE(commands.commands[1].dependencies.size() == 1); + REQUIRE(commands.commands[1].dependencies[0] == 0); + REQUIRE(commands.commands[1].bindings.size() == (uses_fused_prefill_expert_table_partition ? 4 : 2)); + REQUIRE(commands.commands[1].bindings[0].name == "route_ids"); + REQUIRE(commands.commands[1].bindings[0].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands[1].bindings[0].value == route_ids_value->storage_root); + REQUIRE(commands.commands[1].bindings[0].offset == route_ids_value->storage_offset); + REQUIRE(commands.commands[1].bindings[1].name == "expert_table"); + REQUIRE(commands.commands[1].bindings[1].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands[1].bindings[1].length == expert_table_bytes); + if (uses_fused_prefill_expert_table_partition) { + REQUIRE(commands.commands[1].bindings[2].name == "partition_table"); + REQUIRE(commands.commands[1].bindings[2].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands[1].bindings[2].length == partition_table_bytes); + REQUIRE(commands.commands[1].bindings[3].name == "completion_counter"); + REQUIRE(commands.commands[1].bindings[3].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands[1].bindings[3].length == sizeof(int32_t)); + } else { + REQUIRE(commands.commands[2].dependencies.size() == 1); + REQUIRE(commands.commands[2].dependencies[0] == 1); + REQUIRE(commands.commands[2].bindings.size() == 2); + REQUIRE(commands.commands[2].bindings[0].name == "expert_table"); + REQUIRE(commands.commands[2].bindings[0].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands[2].bindings[1].name == "partition_table"); + REQUIRE(commands.commands[2].bindings[1].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands[2].bindings[1].length == partition_table_bytes); + } + const ggml::hrx::TransientAllocation * route_ids_allocation = + ggml::hrx::find_transient_allocation(commands.transients, route_ids_value->storage_root); + REQUIRE(route_ids_allocation != nullptr); + REQUIRE(route_ids_allocation->size == dispatch.bindings[1].length); + const ggml::hrx::TransientAllocation * expert_table_allocation = + ggml::hrx::find_transient_allocation(commands.transients, expert_table_transient.value); + REQUIRE(expert_table_allocation != nullptr); + REQUIRE(expert_table_allocation->size == expert_table_bytes); + const ggml::hrx::TransientAllocation * partition_table_allocation = + ggml::hrx::find_transient_allocation(commands.transients, partition_table_transient.value); + REQUIRE(partition_table_allocation != nullptr); + REQUIRE(partition_table_allocation->size == partition_table_bytes); + if (uses_fused_prefill_expert_table_partition) { + const ggml::hrx::CommandPlanCompletionCounterRequest & completion_counter_request = + scheduler.plan().completion_counter_requests[0]; + const ggml::hrx::TransientAllocation * completion_counter_allocation = + ggml::hrx::find_transient_allocation(commands.transients, completion_counter_request.value); + REQUIRE(completion_counter_allocation != nullptr); + REQUIRE(completion_counter_allocation->size == sizeof(int32_t)); + REQUIRE(completion_counter_allocation->arena_offset == commands.completion_counters.arena_offset); + } +} + +static void schedule_manual_qwen_router_top8_command(ggml::hrx::Graph & graph, + ggml::hrx::ValueId route_ids, + int64_t token_count, + int64_t expert_count, + int64_t route_count) { + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 3); + REQUIRE(scheduler.plan().transients.size() == 2); + + const size_t route_id_length = static_cast(token_count * route_count) * sizeof(int32_t); + const size_t expert_table_bytes = qwen_expert_table_size(token_count, expert_count); + const size_t partition_table_bytes = qwen_partition_table_size(token_count, route_count, expert_count); + + const ggml::hrx::Dispatch & dispatch = scheduler.plan().dispatches[0]; + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "qwen3_moe:qwen3_moe_router_top8_f32"); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == token_count); + REQUIRE(dispatch.kernel.integer_parameters.at("route_id_stride") == route_count); + REQUIRE(dispatch.bindings.size() == 3); + REQUIRE(dispatch.bindings[1].value == route_ids); + REQUIRE(dispatch.bindings[1].length == route_id_length); + require_compile_parameter(dispatch, "qwen3_moe.router.expert_count", std::to_string(expert_count)); + require_compile_parameter(dispatch, "qwen3_moe.router.route_count", std::to_string(route_count)); + + const ggml::hrx::CommandPlanTransient & expert_table_transient = scheduler.plan().transients[0]; + const ggml::hrx::CommandPlanTransient & partition_table_transient = scheduler.plan().transients[1]; + REQUIRE(expert_table_transient.size == expert_table_bytes); + REQUIRE(partition_table_transient.size == partition_table_bytes); + + const ggml::hrx::CommandPlanMoeRoutingBundle * bundle = + scheduler.plan().metadata.find_moe_routing_bundle(route_ids); + REQUIRE(bundle != nullptr); + REQUIRE(bundle->token_count == token_count); + REQUIRE(bundle->route_count == route_count); + REQUIRE(bundle->route_stride == route_count); + REQUIRE(bundle->expert_count == expert_count); + REQUIRE(bundle->expert_table_byte_count == expert_table_bytes); + REQUIRE(bundle->partition_table_byte_count == partition_table_bytes); + + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 3); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands[0].bindings[1].value == route_ids); + REQUIRE(commands.commands[0].bindings[1].length == route_id_length); +} + +struct QwenRoutedGateUpTensors { + ggml_tensor * route_ids = nullptr; + ggml_tensor * route_weights = nullptr; + ggml_tensor * input = nullptr; + ggml_tensor * gate_weight = nullptr; + ggml_tensor * up_weight = nullptr; + ggml_tensor * down_weight = nullptr; + ggml_tensor * gate = nullptr; + ggml_tensor * up = nullptr; + ggml_tensor * glu = nullptr; + ggml_tensor * output = nullptr; + ggml_tensor * weighted = nullptr; + ggml_tensor * hidden_state = nullptr; + ggml_tensor * residual = nullptr; + ggml_tensor * next_rms = nullptr; + ggml_tensor * next_output = nullptr; + ggml_tensor * hidden_use = nullptr; + std::vector route_views; +}; + +static QwenRoutedGateUpTensors build_qwen_routed_gate_up_graph(ggml_context * ctx, + int64_t token_count, + ggml_glu_op glu_op = GGML_GLU_OP_SWIGLU, + ggml_type up_weight_type = GGML_TYPE_Q4_K, + bool include_down = true, + ggml_type down_weight_type = GGML_TYPE_Q6_K) { + QwenRoutedGateUpTensors tensors; + ggml_tensor * logits = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenRouterExpertCount, token_count); + REQUIRE(logits != nullptr); + tensors.route_weights = build_qwen_router_top8_graph(ctx, logits, &tensors.route_ids); + REQUIRE(tensors.route_weights != nullptr); + REQUIRE(tensors.route_ids != nullptr); + + tensors.input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, kQwenMoeHiddenSize, 1, token_count); + tensors.gate_weight = + ggml_new_tensor_3d(ctx, GGML_TYPE_Q4_K, kQwenMoeHiddenSize, kQwenMoeIntermediateSize, kQwenRouterExpertCount); + tensors.up_weight = + ggml_new_tensor_3d(ctx, up_weight_type, kQwenMoeHiddenSize, kQwenMoeIntermediateSize, kQwenRouterExpertCount); + REQUIRE(tensors.input != nullptr); + REQUIRE(tensors.gate_weight != nullptr); + REQUIRE(tensors.up_weight != nullptr); + + tensors.gate = ggml_mul_mat_id(ctx, tensors.gate_weight, tensors.input, tensors.route_ids); + tensors.up = ggml_mul_mat_id(ctx, tensors.up_weight, tensors.input, tensors.route_ids); + REQUIRE(tensors.gate != nullptr); + REQUIRE(tensors.up != nullptr); + tensors.glu = ggml_glu_split(ctx, tensors.gate, tensors.up, glu_op); + REQUIRE(tensors.glu != nullptr); + + if (include_down) { + tensors.down_weight = ggml_new_tensor_3d(ctx, down_weight_type, kQwenMoeIntermediateSize, kQwenMoeHiddenSize, + kQwenRouterExpertCount); + REQUIRE(tensors.down_weight != nullptr); + tensors.output = ggml_mul_mat_id(ctx, tensors.down_weight, tensors.glu, tensors.route_ids); + REQUIRE(tensors.output != nullptr); + } else { + tensors.output = tensors.glu; + } + return tensors; +} + +static void append_qwen_weighted_reduce_tail(ggml_context * ctx, + QwenRoutedGateUpTensors & tensors, + bool include_next_rmsnorm = false, + ggml_tensor * routed_input = nullptr) { + REQUIRE(tensors.output != nullptr); + REQUIRE(tensors.route_weights != nullptr); + tensors.weighted = ggml_mul(ctx, routed_input != nullptr ? routed_input : tensors.output, tensors.route_weights); + REQUIRE(tensors.weighted != nullptr); + tensors.route_views.clear(); + for (int64_t route = 0; route < kQwenRouterRouteCount; ++route) { + ggml_tensor * view = + ggml_view_2d(ctx, tensors.weighted, kQwenMoeHiddenSize, tensors.weighted->ne[2], tensors.weighted->nb[2], + static_cast(route) * tensors.weighted->nb[1]); + REQUIRE(view != nullptr); + tensors.route_views.push_back(view); + } + + ggml_tensor * reduced = tensors.route_views.front(); + for (size_t i = 1; i < tensors.route_views.size(); ++i) { + reduced = ggml_add(ctx, reduced, tensors.route_views[i]); + REQUIRE(reduced != nullptr); + } + + tensors.hidden_state = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenMoeHiddenSize, tensors.output->ne[2]); + REQUIRE(tensors.hidden_state != nullptr); + tensors.residual = ggml_add(ctx, tensors.hidden_state, reduced); + REQUIRE(tensors.residual != nullptr); + + if (include_next_rmsnorm) { + ggml_tensor * next_norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, kQwenMoeHiddenSize); + REQUIRE(next_norm_weight != nullptr); + tensors.next_rms = ggml_rms_norm(ctx, tensors.residual, 0.000001f); + REQUIRE(tensors.next_rms != nullptr); + tensors.next_output = ggml_mul(ctx, tensors.next_rms, next_norm_weight); + REQUIRE(tensors.next_output != nullptr); + } +} + +static ggml::hrx::GraphImportResult import_qwen_routed_gate_up_graph(ggml_context * ctx, + const QwenRoutedGateUpTensors & tensors) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, tensors.route_weights); + ggml_build_forward_expand(graph, tensors.output); + if (tensors.residual != nullptr) { + ggml_build_forward_expand(graph, tensors.residual); + } + if (tensors.next_output != nullptr) { + ggml_build_forward_expand(graph, tensors.next_output); + } + if (tensors.hidden_use != nullptr) { + ggml_build_forward_expand(graph, tensors.hidden_use); + } + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + return imported; +} + +static bool match_dispatch_at_index(const ggml::hrx::Graph & graph, + const ggml::hrx::CommandPlan & plan, + const std::vector & covered_nodes, + size_t node_index, + ggml::hrx::DispatchMatch & match) { + REQUIRE(node_index < graph.nodes().size()); + const ggml::hrx::DispatchMatchContext context = { + graph, + &graph.nodes()[node_index], + node_index, + covered_nodes, + plan, + ggml::hrx::ValueId(static_cast(graph.values().size() + plan.transients.size() + + plan.completion_counter_requests.size())), + &test_dispatch_registry(), + }; + return test_dispatch_registry().match(context, match); +} + +static void append_match_to_plan(ggml::hrx::CommandPlan & plan, + ggml::hrx::DispatchMatch & match, + std::vector & covered_nodes, + ggml::hrx::Graph * graph) { + if (graph != nullptr) { + for (const ggml::hrx::DispatchValueAliasRequest & alias : match.value_aliases) { + ggml::hrx::Status status = graph->values().alias_storage(alias.target_value, alias.source_value); + REQUIRE(status.success()); + } + } + for (ggml::hrx::Dispatch & dispatch : match.initialization_dispatches) { + plan.initialization_dispatches.push_back(std::move(dispatch)); + } + for (ggml::hrx::Dispatch & dispatch : match.dispatches) { + plan.dispatches.push_back(std::move(dispatch)); + } + for (ggml::hrx::CommandPlanTransient & transient : match.transients) { + plan.transients.push_back(std::move(transient)); + } + for (ggml::hrx::CommandPlanConstantInitialization & initialization : match.constant_initializations) { + plan.constant_initializations.push_back(std::move(initialization)); + } + for (ggml::hrx::CommandPlanCompletionCounterRequest & request : match.completion_counter_requests) { + plan.completion_counter_requests.push_back(std::move(request)); + } + REQUIRE(plan.metadata.append(std::move(match.metadata), plan.status)); + for (const size_t covered_node : match.covered_nodes) { + REQUIRE(covered_node < covered_nodes.size()); + REQUIRE(!covered_nodes[covered_node]); + covered_nodes[covered_node] = true; + } +} + +static ggml::hrx::CommandPlan build_qwen_router_plan_for_graph(ggml::hrx::Graph & graph, + std::vector & covered_nodes) { + ggml::hrx::CommandPlan plan; + size_t softmax_index = graph.nodes().size(); + for (size_t i = 0; i < graph.nodes().size(); ++i) { + if (graph.nodes()[i].op == GGML_OP_SOFT_MAX) { + softmax_index = i; + break; + } + } + REQUIRE(softmax_index < graph.nodes().size()); + ggml::hrx::DispatchMatch router_match; + REQUIRE(match_dispatch_at_index(graph, plan, covered_nodes, softmax_index, router_match)); + append_match_to_plan(plan, router_match, covered_nodes, &graph); + return plan; +} + +static void append_qwen_routed_gate_up_for_graph(ggml::hrx::Graph & graph, + const QwenRoutedGateUpTensors & tensors, + std::vector & covered_nodes, + ggml::hrx::CommandPlan & plan) { + const size_t gate_index = producer_index_for_tensor(graph, tensors.gate); + ggml::hrx::DispatchMatch gate_up_match; + REQUIRE(match_dispatch_at_index(graph, plan, covered_nodes, gate_index, gate_up_match)); + append_match_to_plan(plan, gate_up_match, covered_nodes, &graph); +} + +static void append_qwen_routed_down_for_graph(ggml::hrx::Graph & graph, + const QwenRoutedGateUpTensors & tensors, + std::vector & covered_nodes, + ggml::hrx::CommandPlan & plan, + const char * expected_kernel_name) { + const size_t down_index = producer_index_for_tensor(graph, tensors.output); + ggml::hrx::DispatchMatch down_match; + REQUIRE(match_dispatch_at_index(graph, plan, covered_nodes, down_index, down_match)); + append_match_to_plan(plan, down_match, covered_nodes, &graph); + + REQUIRE(plan.dispatches.size() >= 1); + REQUIRE(plan.transients.size() >= 1); + const ggml::hrx::Value * route_ids_value = graph.values().find_tensor(tensors.route_ids); + const ggml::hrx::Value * glu_value = graph.values().find_tensor(tensors.glu); + const ggml::hrx::Value * output_value = graph.values().find_tensor(tensors.output); + REQUIRE(route_ids_value != nullptr); + REQUIRE(glu_value != nullptr); + REQUIRE(output_value != nullptr); + const ggml::hrx::CommandPlanMoeRoutingBundle * routing_bundle = + plan.metadata.find_moe_routing_bundle(route_ids_value->id); + const ggml::hrx::CommandPlanAlternateValue * gate_up_alternate = plan.metadata.find_alternate_value( + glu_value->id, GGML_TYPE_F16, qwen_routed_gate_up_f16_output_size(tensors.output->ne[2])); + const ggml::hrx::CommandPlanAlternateValue * routed_down_alternate = plan.metadata.find_alternate_value( + output_value->id, GGML_TYPE_F16, qwen_routed_down_f16_output_size(tensors.output->ne[2])); + REQUIRE(routing_bundle != nullptr); + REQUIRE(gate_up_alternate != nullptr); + REQUIRE(routed_down_alternate != nullptr); + const ggml::hrx::Dispatch & dispatch = plan.dispatches.back(); + const ggml::hrx::CommandPlanTransient & routed_down_transient = plan.transients.back(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == expected_kernel_name); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == tensors.output->ne[2]); + REQUIRE(dispatch.bindings.size() == 4); + REQUIRE(routed_down_transient.name == "qwen.moe.routed_down_f16"); + REQUIRE(routed_down_transient.size == qwen_routed_down_f16_output_size(tensors.output->ne[2])); + REQUIRE(dispatch.bindings[0].value == gate_up_alternate->alternate_value); + REQUIRE(dispatch.bindings[0].length == qwen_routed_gate_up_f16_output_size(tensors.output->ne[2])); + REQUIRE(dispatch.bindings[1].value == routing_bundle->expert_table); + REQUIRE(dispatch.bindings[1].length == routing_bundle->expert_table_byte_count); + REQUIRE(routed_down_alternate->alternate_value == routed_down_transient.value); + REQUIRE(routed_down_alternate->byte_count == routed_down_transient.size); + REQUIRE(dispatch.bindings[3].value == routed_down_transient.value); + REQUIRE(dispatch.bindings[3].length == routed_down_transient.size); + require_compile_parameter(dispatch, "ggml.mul_mat_id_f16_f16.input_size", "768"); + require_compile_parameter(dispatch, "ggml.mul_mat_id_f16_f16.route_count", "8"); + require_compile_parameter(dispatch, "ggml.mul_mat_id_f16_f16.expert_count", "128"); + require_compile_parameter(dispatch, "ggml.mul_mat_id_f16_f16.output_size", "2048"); + require_compile_parameter(dispatch, "ggml.mul_mat_id_f16_f16.weight_format", + tensors.down_weight->type == GGML_TYPE_Q4_K ? "4" : "6"); + require_compile_parameter(dispatch, "ggml.workload.token_capacity", std::to_string(tensors.output->ne[2])); +} + +static std::string common_mul_mat_weight_format(ggml_type type) { + switch (type) { + case GGML_TYPE_Q3_K: + return "11"; + case GGML_TYPE_Q4_K: + return "4"; + case GGML_TYPE_Q5_K: + return "5"; + case GGML_TYPE_Q6_K: + return "6"; + case GGML_TYPE_IQ3_S: + return "21"; + case GGML_TYPE_IQ4_NL: + return "20"; + case GGML_TYPE_IQ4_XS: + return "23"; + case GGML_TYPE_Q8_0: + return "80"; + case GGML_TYPE_Q8_1: + return "81"; + case GGML_TYPE_F16: + return "16"; + case GGML_TYPE_BF16: + return "30"; + case GGML_TYPE_F32: + return "32"; + default: + return "0"; + } +} + +static const char * common_mul_mat_id_f32_kernel_name(int64_t token_count) { + return token_count <= 5 ? "loom_libs:ggml_mul_mat_id_skinny_input_f32_publish_f32" : + "loom_libs:ggml_mul_mat_id_tiled_input_f32_publish_f32"; +} + +static const char * common_mul_mat_id_postops_f32_kernel_name(int64_t token_count, bool has_rmsnorm) { + if (has_rmsnorm) { + return token_count <= 5 ? "loom_libs:ggml_mul_mat_id_skinny_input_f32_postops_next_rmsnorm_publish_f32" : + "loom_libs:ggml_mul_mat_id_tiled_input_f32_postops_next_rmsnorm_publish_f32"; + } + return token_count <= 5 ? "loom_libs:ggml_mul_mat_id_skinny_input_f32_postops_publish_f32" : + "loom_libs:ggml_mul_mat_id_tiled_input_f32_postops_publish_f32"; +} + +static const char * common_mul_mat_id_swiglu_f32_kernel_name(int64_t token_count) { + return token_count <= 5 ? "loom_libs:ggml_mul_mat_id_skinny_pair_input_f32_swiglu_publish_f32" : + "loom_libs:ggml_mul_mat_id_tiled_pair_input_f32_swiglu_publish_f32"; +} + +static ggml::hrx::DispatchMatch require_common_mul_mat_id_for_graph(ggml::hrx::Graph & graph, + ggml::hrx::CommandPlan & plan, + std::vector & covered_nodes, + ggml_tensor * output, + ggml_tensor * input, + ggml_tensor * weight, + ggml_tensor * route_ids) { + const ggml::hrx::Value * output_value = graph.values().find_tensor(output); + const ggml::hrx::Value * input_value = graph.values().find_tensor(input); + const ggml::hrx::Value * weight_value = graph.values().find_tensor(weight); + const ggml::hrx::Value * route_ids_value = graph.values().find_tensor(route_ids); + REQUIRE(output_value != nullptr); + REQUIRE(input_value != nullptr); + REQUIRE(weight_value != nullptr); + REQUIRE(route_ids_value != nullptr); + const ggml::hrx::CommandPlanMoeRoutingBundle * routing_bundle = + plan.metadata.find_moe_routing_bundle(route_ids_value->id); + REQUIRE(routing_bundle != nullptr); + + const size_t output_index = producer_index_for_tensor(graph, output); + ggml::hrx::DispatchMatch match; + REQUIRE(match_dispatch_at_index(graph, plan, covered_nodes, output_index, match)); + REQUIRE(match.dispatches.size() == 1); + const ggml::hrx::Dispatch & dispatch = match.dispatches[0]; + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == common_mul_mat_id_f32_kernel_name(output->ne[2])); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == output->ne[2]); + REQUIRE(dispatch.bindings.size() == 5); + REQUIRE(dispatch.bindings[0].value == input_value->id); + REQUIRE(dispatch.bindings[0].length == input_value->byte_count); + REQUIRE(dispatch.bindings[1].value == routing_bundle->expert_table); + REQUIRE(dispatch.bindings[1].length == qwen_expert_table_size(output->ne[2], weight->ne[2])); + REQUIRE(dispatch.bindings[2].value == routing_bundle->partition_table); + REQUIRE(dispatch.bindings[2].length == qwen_partition_table_size(output->ne[2], route_ids->ne[0], weight->ne[2])); + REQUIRE(dispatch.bindings[3].value == weight_value->id); + REQUIRE(dispatch.bindings[3].length == weight_value->byte_count); + REQUIRE(dispatch.bindings[4].value == output_value->id); + REQUIRE(dispatch.bindings[4].length == output_value->byte_count); + require_compile_parameter(dispatch, "ggml.mul_mat_id.input_size", std::to_string(weight->ne[0])); + require_compile_parameter(dispatch, "ggml.mul_mat_id.output_size", std::to_string(weight->ne[1])); + require_compile_parameter(dispatch, "ggml.mul_mat_id.expert_count", std::to_string(weight->ne[2])); + require_compile_parameter(dispatch, "ggml.mul_mat_id.route_count", std::to_string(route_ids->ne[0])); + require_compile_parameter(dispatch, "ggml.mul_mat_id.input_route_count", std::to_string(input->ne[1])); + require_compile_parameter(dispatch, "ggml.mul_mat_id.weight_format", common_mul_mat_weight_format(weight->type)); + require_compile_parameter(dispatch, "ggml.workload.token_capacity", std::to_string(output->ne[2])); + return match; +} + +struct CommonMulMatIdSwiGLUTensors { + ggml_tensor * input = nullptr; + ggml_tensor * route_ids = nullptr; + ggml_tensor * gate_weight = nullptr; + ggml_tensor * up_weight = nullptr; + ggml_tensor * gate = nullptr; + ggml_tensor * up = nullptr; + ggml_tensor * output = nullptr; +}; + +static CommonMulMatIdSwiGLUTensors build_common_mul_mat_id_swiglu_graph(ggml_context * ctx, + ggml_type gate_weight_type, + ggml_type up_weight_type, + ggml_glu_op glu_op = GGML_GLU_OP_SWIGLU, + int64_t input_size = 256, + int64_t output_size = 128, + int64_t token_count = 4) { + CommonMulMatIdSwiGLUTensors tensors; + constexpr int64_t route_count = 8; + constexpr int64_t input_route_count = 1; + constexpr int64_t expert_count = 16; + tensors.route_ids = ggml_new_tensor_2d(ctx, GGML_TYPE_I32, route_count, token_count); + tensors.input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, input_size, input_route_count, token_count); + tensors.gate_weight = ggml_new_tensor_3d(ctx, gate_weight_type, input_size, output_size, expert_count); + tensors.up_weight = ggml_new_tensor_3d(ctx, up_weight_type, input_size, output_size, expert_count); + REQUIRE(tensors.route_ids != nullptr); + REQUIRE(tensors.input != nullptr); + REQUIRE(tensors.gate_weight != nullptr); + REQUIRE(tensors.up_weight != nullptr); + + tensors.gate = ggml_mul_mat_id(ctx, tensors.gate_weight, tensors.input, tensors.route_ids); + tensors.up = ggml_mul_mat_id(ctx, tensors.up_weight, tensors.input, tensors.route_ids); + REQUIRE(tensors.gate != nullptr); + REQUIRE(tensors.up != nullptr); + tensors.output = ggml_glu_split(ctx, tensors.gate, tensors.up, glu_op); + REQUIRE(tensors.output != nullptr); + return tensors; +} + +static ggml::hrx::GraphImportResult import_common_mul_mat_id_swiglu_graph(ggml_context * ctx, + const CommonMulMatIdSwiGLUTensors & tensors) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, tensors.output); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + return imported; +} + +static void require_common_mul_mat_id_swiglu_match(ggml_context * ctx, + const CommonMulMatIdSwiGLUTensors & tensors, + ggml_tensor * root_projection) { + ggml::hrx::GraphImportResult imported = import_common_mul_mat_id_swiglu_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan; + const size_t root_index = producer_index_for_tensor(imported.graph, root_projection); + ggml::hrx::DispatchMatch match; + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, root_index, match)); + REQUIRE(match.dispatches.size() == 3); + REQUIRE(match.transients.size() == 2); + REQUIRE(match.metadata.generated_resources().size() == 2); + REQUIRE(match.metadata.moe_routing_bundles().size() == 1); + + const ggml::hrx::Dispatch & dispatch = match.dispatches.back(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == + common_mul_mat_id_swiglu_f32_kernel_name(tensors.output->ne[2])); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == tensors.output->ne[2]); + REQUIRE(dispatch.bindings.size() == 6); + require_compile_parameter(dispatch, "ggml.mul_mat_id_swiglu.input_size", + std::to_string(tensors.gate_weight->ne[0])); + require_compile_parameter(dispatch, "ggml.mul_mat_id_swiglu.output_size", + std::to_string(tensors.gate_weight->ne[1])); + require_compile_parameter(dispatch, "ggml.mul_mat_id_swiglu.expert_count", + std::to_string(tensors.gate_weight->ne[2])); + require_compile_parameter(dispatch, "ggml.mul_mat_id_swiglu.route_count", std::to_string(tensors.route_ids->ne[0])); + require_compile_parameter(dispatch, "ggml.mul_mat_id_swiglu.input_route_count", + std::to_string(tensors.input->ne[1])); + require_compile_parameter(dispatch, "ggml.mul_mat_id_swiglu.gate_weight_format", + common_mul_mat_weight_format(tensors.gate_weight->type)); + require_compile_parameter(dispatch, "ggml.mul_mat_id_swiglu.up_weight_format", + common_mul_mat_weight_format(tensors.up_weight->type)); + require_compile_parameter(dispatch, "ggml.workload.token_capacity", std::to_string(tensors.output->ne[2])); + + const ggml::hrx::Value * input_value = imported.graph.values().find_tensor(tensors.input); + const ggml::hrx::Value * route_ids_value = imported.graph.values().find_tensor(tensors.route_ids); + const ggml::hrx::Value * gate_weight_value = imported.graph.values().find_tensor(tensors.gate_weight); + const ggml::hrx::Value * up_weight_value = imported.graph.values().find_tensor(tensors.up_weight); + const ggml::hrx::Value * output_value = imported.graph.values().find_tensor(tensors.output); + REQUIRE(input_value != nullptr); + REQUIRE(route_ids_value != nullptr); + REQUIRE(gate_weight_value != nullptr); + REQUIRE(up_weight_value != nullptr); + REQUIRE(output_value != nullptr); + const ggml::hrx::CommandPlanMoeRoutingBundle * routing_bundle = + match.metadata.find_moe_routing_bundle(route_ids_value->id); + REQUIRE(routing_bundle != nullptr); + + REQUIRE(dispatch.bindings[0].value == input_value->id); + REQUIRE(dispatch.bindings[1].value == routing_bundle->expert_table); + REQUIRE(dispatch.bindings[2].value == routing_bundle->partition_table); + REQUIRE(dispatch.bindings[3].value == gate_weight_value->id); + REQUIRE(dispatch.bindings[4].value == up_weight_value->id); + REQUIRE(dispatch.bindings[5].value == output_value->id); + + append_match_to_plan(plan, match, covered_nodes, &imported.graph); + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 3); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands[2].bindings.size() == 6); + REQUIRE(commands.commands[2].bindings[0].name == "input"); + REQUIRE(commands.commands[2].bindings[1].name == "expert_table"); + REQUIRE(commands.commands[2].bindings[2].name == "partition_table"); + REQUIRE(commands.commands[2].bindings[3].name == "gate_weight"); + REQUIRE(commands.commands[2].bindings[4].name == "up_weight"); + REQUIRE(commands.commands[2].bindings[5].name == "output"); +} + +static void run_common_mul_mat_id_swiglu_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + { + const CommonMulMatIdSwiGLUTensors tensors = + build_common_mul_mat_id_swiglu_graph(ctx, GGML_TYPE_Q4_K, GGML_TYPE_F16); + require_common_mul_mat_id_swiglu_match(ctx, tensors, tensors.gate); + } + { + const CommonMulMatIdSwiGLUTensors tensors = + build_common_mul_mat_id_swiglu_graph(ctx, GGML_TYPE_F32, GGML_TYPE_Q8_1); + require_common_mul_mat_id_swiglu_match(ctx, tensors, tensors.up); + } + { + const CommonMulMatIdSwiGLUTensors tensors = + build_common_mul_mat_id_swiglu_graph(ctx, GGML_TYPE_BF16, GGML_TYPE_BF16); + require_common_mul_mat_id_swiglu_match(ctx, tensors, tensors.gate); + } + { + const CommonMulMatIdSwiGLUTensors tensors = + build_common_mul_mat_id_swiglu_graph(ctx, GGML_TYPE_IQ4_XS, GGML_TYPE_IQ4_XS); + require_common_mul_mat_id_swiglu_match(ctx, tensors, tensors.gate); + } + { + const CommonMulMatIdSwiGLUTensors tensors = + build_common_mul_mat_id_swiglu_graph(ctx, GGML_TYPE_Q3_K, GGML_TYPE_Q3_K); + require_common_mul_mat_id_swiglu_match(ctx, tensors, tensors.gate); + } + { + const CommonMulMatIdSwiGLUTensors tensors = + build_common_mul_mat_id_swiglu_graph(ctx, GGML_TYPE_IQ3_S, GGML_TYPE_IQ3_S); + require_common_mul_mat_id_swiglu_match(ctx, tensors, tensors.gate); + } + { + const CommonMulMatIdSwiGLUTensors tensors = + build_common_mul_mat_id_swiglu_graph(ctx, GGML_TYPE_IQ4_NL, GGML_TYPE_IQ4_NL); + require_common_mul_mat_id_swiglu_match(ctx, tensors, tensors.gate); + } + { + const CommonMulMatIdSwiGLUTensors tensors = + build_common_mul_mat_id_swiglu_graph(ctx, GGML_TYPE_IQ4_NL, GGML_TYPE_IQ4_NL, GGML_GLU_OP_SWIGLU, 640, 128); + require_common_mul_mat_id_swiglu_match(ctx, tensors, tensors.gate); + } + { + const CommonMulMatIdSwiGLUTensors tensors = + build_common_mul_mat_id_swiglu_graph(ctx, GGML_TYPE_Q4_K, GGML_TYPE_Q4_K, GGML_GLU_OP_GEGLU); + ggml::hrx::GraphImportResult imported = import_common_mul_mat_id_swiglu_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan; + const size_t gate_index = producer_index_for_tensor(imported.graph, tensors.gate); + ggml::hrx::DispatchMatch match; + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, gate_index, match)); + REQUIRE(kernel_name_for_id(match.dispatches.back().kernel.kernel_id) == + common_mul_mat_id_f32_kernel_name(tensors.output->ne[2])); + } + + ggml_free(ctx); +} + +struct CommonMulMatIdPostOpsTensors { + ggml_tensor * input = nullptr; + ggml_tensor * weight = nullptr; + ggml_tensor * route_ids = nullptr; + ggml_tensor * projection = nullptr; + ggml_tensor * bias = nullptr; + ggml_tensor * residual_input = nullptr; + ggml_tensor * residual_output = nullptr; + ggml_tensor * rms_output = nullptr; + ggml_tensor * norm_weight = nullptr; + ggml_tensor * normalized_output = nullptr; +}; + +static int64_t common_mul_mat_id_partition_descriptor_capacity(const CommonMulMatIdPostOpsTensors & tensors) { + const int64_t token_count = tensors.projection->ne[2]; + const int64_t route_count = tensors.route_ids->ne[0]; + const int64_t expert_count = tensors.weight->ne[2]; + const int64_t assignment_count = token_count * route_count; + const int64_t assignment_partition_count = (assignment_count + 31) / 32; + return assignment_partition_count + expert_count; +} + +static CommonMulMatIdPostOpsTensors build_common_mul_mat_id_postops_graph(ggml_context * ctx, + bool include_bias, + bool include_residual, + bool include_rmsnorm, + ggml_type weight_type = GGML_TYPE_Q4_K, + int64_t token_count = 4) { + CommonMulMatIdPostOpsTensors tensors; + constexpr int64_t route_count = 8; + constexpr int64_t input_size = 256; + constexpr int64_t output_size = 128; + constexpr int64_t expert_count = 16; + tensors.route_ids = ggml_new_tensor_2d(ctx, GGML_TYPE_I32, route_count, token_count); + tensors.input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, input_size, 1, token_count); + tensors.weight = ggml_new_tensor_3d(ctx, weight_type, input_size, output_size, expert_count); + REQUIRE(tensors.route_ids != nullptr); + REQUIRE(tensors.input != nullptr); + REQUIRE(tensors.weight != nullptr); + + tensors.projection = ggml_mul_mat_id(ctx, tensors.weight, tensors.input, tensors.route_ids); + REQUIRE(tensors.projection != nullptr); + + ggml_tensor * current = tensors.projection; + if (include_bias) { + tensors.bias = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, output_size); + REQUIRE(tensors.bias != nullptr); + current = ggml_add(ctx, current, tensors.bias); + REQUIRE(current != nullptr); + } + if (include_residual) { + tensors.residual_input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, output_size, route_count, token_count); + REQUIRE(tensors.residual_input != nullptr); + current = ggml_add(ctx, current, tensors.residual_input); + REQUIRE(current != nullptr); + } + tensors.residual_output = current; + + if (include_rmsnorm) { + REQUIRE(include_residual); + tensors.norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, output_size); + REQUIRE(tensors.norm_weight != nullptr); + tensors.rms_output = ggml_rms_norm(ctx, tensors.residual_output, 0.000001f); + REQUIRE(tensors.rms_output != nullptr); + tensors.normalized_output = ggml_mul(ctx, tensors.rms_output, tensors.norm_weight); + REQUIRE(tensors.normalized_output != nullptr); + } else { + tensors.normalized_output = tensors.residual_output; + } + return tensors; +} + +static ggml::hrx::GraphImportResult import_common_mul_mat_id_postops_graph( + ggml_context * ctx, + const CommonMulMatIdPostOpsTensors & tensors) { + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, tensors.normalized_output); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + return imported; +} + +static void require_common_mul_mat_id_postops_match(ggml_context * ctx, + const CommonMulMatIdPostOpsTensors & tensors, + const char * expected_kernel_name, + bool expect_bias, + bool expect_residual, + bool expect_rmsnorm) { + ggml::hrx::GraphImportResult imported = import_common_mul_mat_id_postops_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan; + const size_t projection_index = producer_index_for_tensor(imported.graph, tensors.projection); + ggml::hrx::DispatchMatch match; + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, projection_index, match)); + REQUIRE(match.dispatches.size() == 3); + REQUIRE(match.transients.size() == 2); + + const ggml::hrx::Dispatch & dispatch = match.dispatches.back(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == expected_kernel_name); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == tensors.projection->ne[2]); + REQUIRE(dispatch.bindings.size() == (expect_rmsnorm ? 10 : 7)); + require_compile_parameter(dispatch, "ggml.mul_mat_id_postops.input_size", std::to_string(tensors.weight->ne[0])); + require_compile_parameter(dispatch, "ggml.mul_mat_id_postops.output_size", std::to_string(tensors.weight->ne[1])); + require_compile_parameter(dispatch, "ggml.mul_mat_id_postops.expert_count", std::to_string(tensors.weight->ne[2])); + require_compile_parameter(dispatch, "ggml.mul_mat_id_postops.route_count", + std::to_string(tensors.route_ids->ne[0])); + require_compile_parameter(dispatch, "ggml.mul_mat_id_postops.input_route_count", + std::to_string(tensors.input->ne[1])); + require_compile_parameter(dispatch, "ggml.mul_mat_id_postops.weight_format", + common_mul_mat_weight_format(tensors.weight->type)); + require_compile_parameter(dispatch, "ggml.mul_mat_id_postops.has_bias", expect_bias ? "1" : "0"); + require_compile_parameter(dispatch, "ggml.mul_mat_id_postops.has_residual", expect_residual ? "1" : "0"); + + const ggml::hrx::Value * input_value = imported.graph.values().find_tensor(tensors.input); + const ggml::hrx::Value * weight_value = imported.graph.values().find_tensor(tensors.weight); + const ggml::hrx::Value * route_ids_value = imported.graph.values().find_tensor(tensors.route_ids); + const ggml::hrx::Value * residual_output_value = imported.graph.values().find_tensor(tensors.residual_output); + const ggml::hrx::Value * normalized_output_value = imported.graph.values().find_tensor(tensors.normalized_output); + REQUIRE(input_value != nullptr); + REQUIRE(weight_value != nullptr); + REQUIRE(route_ids_value != nullptr); + REQUIRE(residual_output_value != nullptr); + REQUIRE(normalized_output_value != nullptr); + const ggml::hrx::CommandPlanMoeRoutingBundle * routing_bundle = + match.metadata.find_moe_routing_bundle(route_ids_value->id); + REQUIRE(routing_bundle != nullptr); + + REQUIRE(dispatch.bindings[0].value == input_value->id); + REQUIRE(dispatch.bindings[1].value == routing_bundle->expert_table); + REQUIRE(dispatch.bindings[2].value == routing_bundle->partition_table); + REQUIRE(dispatch.bindings[3].value == weight_value->id); + if (expect_bias) { + const ggml::hrx::Value * bias_value = imported.graph.values().find_tensor(tensors.bias); + REQUIRE(bias_value != nullptr); + REQUIRE(dispatch.bindings[4].value == bias_value->id); + } + if (expect_residual) { + const ggml::hrx::Value * residual_input_value = imported.graph.values().find_tensor(tensors.residual_input); + REQUIRE(residual_input_value != nullptr); + REQUIRE(dispatch.bindings[5].value == residual_input_value->id); + } + REQUIRE(dispatch.bindings[6].value == residual_output_value->id); + + if (expect_rmsnorm) { + const ggml::hrx::Value * norm_weight_value = imported.graph.values().find_tensor(tensors.norm_weight); + REQUIRE(norm_weight_value != nullptr); + REQUIRE(dispatch.bindings[7].value == norm_weight_value->id); + REQUIRE(dispatch.bindings[8].value == normalized_output_value->id); + REQUIRE(match.completion_counter_requests.size() == 1); + REQUIRE(match.completion_counter_requests.front().count == + common_mul_mat_id_partition_descriptor_capacity(tensors)); + REQUIRE(dispatch.bindings[9].value == match.completion_counter_requests.front().value); + require_compile_parameter(dispatch, "ggml.mul_mat_id_postops.rms_epsilon", "9.99999997e-07"); + } else { + REQUIRE(match.completion_counter_requests.empty()); + } + + append_match_to_plan(plan, match, covered_nodes, &imported.graph); + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 3); +} + +static void run_common_mul_mat_id_postops_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + { + const CommonMulMatIdPostOpsTensors tensors = build_common_mul_mat_id_postops_graph(ctx, true, false, false); + require_common_mul_mat_id_postops_match( + ctx, tensors, common_mul_mat_id_postops_f32_kernel_name(tensors.projection->ne[2], false), true, false, + false); + } + { + const CommonMulMatIdPostOpsTensors tensors = build_common_mul_mat_id_postops_graph(ctx, false, true, false); + require_common_mul_mat_id_postops_match( + ctx, tensors, common_mul_mat_id_postops_f32_kernel_name(tensors.projection->ne[2], false), false, true, + false); + } + { + const CommonMulMatIdPostOpsTensors tensors = build_common_mul_mat_id_postops_graph(ctx, true, true, false); + require_common_mul_mat_id_postops_match( + ctx, tensors, common_mul_mat_id_postops_f32_kernel_name(tensors.projection->ne[2], false), true, true, + false); + } + { + const CommonMulMatIdPostOpsTensors tensors = + build_common_mul_mat_id_postops_graph(ctx, true, true, false, GGML_TYPE_BF16); + require_common_mul_mat_id_postops_match( + ctx, tensors, common_mul_mat_id_postops_f32_kernel_name(tensors.projection->ne[2], false), true, true, + false); + } + { + const CommonMulMatIdPostOpsTensors tensors = build_common_mul_mat_id_postops_graph(ctx, false, true, true); + require_common_mul_mat_id_postops_match( + ctx, tensors, common_mul_mat_id_postops_f32_kernel_name(tensors.projection->ne[2], true), false, true, + true); + } + { + const CommonMulMatIdPostOpsTensors tensors = build_common_mul_mat_id_postops_graph(ctx, true, true, true); + require_common_mul_mat_id_postops_match( + ctx, tensors, common_mul_mat_id_postops_f32_kernel_name(tensors.projection->ne[2], true), true, true, true); + } + + ggml_free(ctx); +} + +static void append_qwen_weighted_reduce_for_graph(ggml::hrx::Graph & graph, + const QwenRoutedGateUpTensors & tensors, + std::vector & covered_nodes, + ggml::hrx::CommandPlan & plan, + const char * expected_kernel_name) { + REQUIRE(tensors.weighted != nullptr); + REQUIRE(tensors.residual != nullptr); + const size_t weighted_index = producer_index_for_tensor(graph, tensors.weighted); + ggml::hrx::DispatchMatch weighted_match; + REQUIRE(match_dispatch_at_index(graph, plan, covered_nodes, weighted_index, weighted_match)); + append_match_to_plan(plan, weighted_match, covered_nodes, &graph); + + const ggml::hrx::Value * route_weights_value = graph.values().find_tensor(tensors.route_weights); + const ggml::hrx::Value * routed_output_value = graph.values().find_tensor(tensors.output); + const ggml::hrx::Value * hidden_state_value = graph.values().find_tensor(tensors.hidden_state); + const ggml::hrx::Value * residual_value = graph.values().find_tensor(tensors.residual); + REQUIRE(route_weights_value != nullptr); + REQUIRE(routed_output_value != nullptr); + REQUIRE(hidden_state_value != nullptr); + REQUIRE(residual_value != nullptr); + + const ggml::hrx::Dispatch & dispatch = plan.dispatches.back(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == expected_kernel_name); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == tensors.output->ne[2]); + REQUIRE(dispatch.bindings.size() == (tensors.next_output != nullptr ? 5 : 4)); + REQUIRE(dispatch.bindings[0].value == route_weights_value->id); + REQUIRE(dispatch.bindings[0].length == route_weights_value->byte_count); + REQUIRE(dispatch.bindings[1].value == plan.metadata.alternate_values().back().alternate_value); + REQUIRE(dispatch.bindings[1].length == qwen_routed_down_f16_output_size(tensors.output->ne[2])); + require_compile_parameter(dispatch, "qwen3_moe.routed_down.route_count", "8"); + require_compile_parameter(dispatch, "qwen3_moe.routed_down.output_size", "2048"); + require_compile_parameter(dispatch, "qwen3_moe.workload.token_capacity", std::to_string(tensors.output->ne[2])); + + if (tensors.next_output != nullptr) { + REQUIRE(dispatch.bindings[2].value == residual_value->id); + REQUIRE(dispatch.bindings[2].length == residual_value->byte_count); + const ggml::hrx::Value * next_output_value = graph.values().find_tensor(tensors.next_output); + REQUIRE(next_output_value != nullptr); + REQUIRE(dispatch.bindings[4].value == next_output_value->id); + REQUIRE(dispatch.bindings[4].length == next_output_value->byte_count); + require_compile_parameter(dispatch, "qwen3_moe.model.hidden_size", "2048"); + require_compile_parameter(dispatch, "qwen3_moe.model.rms_epsilon", "0.000001"); + } else { + REQUIRE(dispatch.bindings[2].value == hidden_state_value->id); + REQUIRE(dispatch.bindings[2].length == hidden_state_value->byte_count); + REQUIRE(dispatch.bindings[3].value == residual_value->id); + REQUIRE(dispatch.bindings[3].length == residual_value->byte_count); + } +} + +static void run_qwen_routed_gate_up_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + { + constexpr int64_t token_count = 4; + QwenRoutedGateUpTensors tensors = build_qwen_routed_gate_up_graph(ctx, token_count, GGML_GLU_OP_GEGLU); + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + ggml::hrx::DispatchMatch gate_match = require_common_mul_mat_id_for_graph( + imported.graph, plan, covered_nodes, tensors.gate, tensors.input, tensors.gate_weight, tensors.route_ids); + append_match_to_plan(plan, gate_match, covered_nodes, &imported.graph); + + REQUIRE(plan.dispatches.size() == 4); + REQUIRE(plan.transients.size() == 2); + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 4); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands[3].bindings.size() == 5); + REQUIRE(commands.commands[3].bindings[0].name == "input"); + REQUIRE(commands.commands[3].bindings[1].name == "expert_table"); + REQUIRE(commands.commands[3].bindings[1].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands[3].bindings[2].name == "partition_table"); + REQUIRE(commands.commands[3].bindings[2].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands[3].bindings[3].name == "weight"); + REQUIRE(commands.commands[3].bindings[4].name == "output"); + } + + { + constexpr int64_t token_count = 4; + const QwenRoutedGateUpTensors tensors = build_qwen_routed_gate_up_graph(ctx, token_count); + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + ggml::hrx::DispatchMatch down_match = require_common_mul_mat_id_for_graph( + imported.graph, plan, covered_nodes, tensors.output, tensors.glu, tensors.down_weight, tensors.route_ids); + append_match_to_plan(plan, down_match, covered_nodes, &imported.graph); + + REQUIRE(plan.dispatches.size() == 4); + REQUIRE(plan.transients.size() == 2); + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 4); + REQUIRE(commands.commands[3].bindings[0].name == "input"); + REQUIRE(commands.commands[3].bindings[2].name == "partition_table"); + REQUIRE(commands.commands[3].bindings[3].name == "weight"); + REQUIRE(commands.commands[3].bindings[4].name == "output"); + } + + { + constexpr int64_t token_count = 4; + const QwenRoutedGateUpTensors tensors = + build_qwen_routed_gate_up_graph(ctx, token_count, GGML_GLU_OP_SWIGLU, GGML_TYPE_Q4_K, true, GGML_TYPE_BF16); + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + ggml::hrx::DispatchMatch down_match = require_common_mul_mat_id_for_graph( + imported.graph, plan, covered_nodes, tensors.output, tensors.glu, tensors.down_weight, tensors.route_ids); + append_match_to_plan(plan, down_match, covered_nodes, &imported.graph); + + REQUIRE(plan.dispatches.size() == 4); + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 4); + } + + { + constexpr int64_t token_count = 4; + QwenRoutedGateUpTensors tensors = build_qwen_routed_gate_up_graph(ctx, token_count); + append_qwen_weighted_reduce_tail(ctx, tensors); + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + + const ggml::hrx::Value * route_ids_value = imported.graph.values().find_tensor(tensors.route_ids); + const ggml::hrx::Value * route_weights_value = imported.graph.values().find_tensor(tensors.route_weights); + const ggml::hrx::Value * glu_value = imported.graph.values().find_tensor(tensors.glu); + REQUIRE(route_ids_value != nullptr); + REQUIRE(route_weights_value != nullptr); + REQUIRE(glu_value != nullptr); + REQUIRE(route_ids_value->kind == ggml::hrx::ValueKind::Transient); + REQUIRE(route_weights_value->kind == ggml::hrx::ValueKind::Transient); + REQUIRE(glu_value->kind == ggml::hrx::ValueKind::Transient); + + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + const ggml::hrx::CommandPlanGeneratedResource * expert_table_resource = plan.metadata.find_generated_resource( + route_ids_value->id, ggml::hrx::GeneratedResourceRole::MoeExpertTable); + const ggml::hrx::CommandPlanGeneratedResource * partition_table_resource = + plan.metadata.find_generated_resource(route_ids_value->id, + ggml::hrx::GeneratedResourceRole::MoePartitionTable); + const ggml::hrx::CommandPlanMoeRoutingBundle * routing_bundle = + plan.metadata.find_moe_routing_bundle(route_ids_value->id); + REQUIRE(expert_table_resource != nullptr); + REQUIRE(partition_table_resource != nullptr); + REQUIRE(routing_bundle != nullptr); + REQUIRE(routing_bundle->route_ids == route_ids_value->id); + REQUIRE(routing_bundle->route_weights == route_weights_value->id); + REQUIRE(routing_bundle->expert_table == expert_table_resource->generated_value); + REQUIRE(routing_bundle->partition_table == partition_table_resource->generated_value); + REQUIRE(routing_bundle->expert_table_byte_count == qwen_expert_table_size(token_count)); + REQUIRE(routing_bundle->partition_table_byte_count == qwen_partition_table_size(token_count)); + REQUIRE(routing_bundle->route_count == kQwenRouterRouteCount); + REQUIRE(routing_bundle->expert_count == kQwenRouterExpertCount); + ggml::hrx::MoeRoutingResourceMetadata expert_metadata; + ggml::hrx::MoeRoutingResourceMetadata partition_metadata; + REQUIRE(expert_table_resource->metadata.read(expert_metadata)); + REQUIRE(partition_table_resource->metadata.read(partition_metadata)); + REQUIRE(expert_metadata.token_count == token_count); + REQUIRE(expert_metadata.route_count == kQwenRouterRouteCount); + REQUIRE(expert_metadata.expert_count == kQwenRouterExpertCount); + REQUIRE(partition_metadata.route_stride == expert_metadata.route_stride); + REQUIRE(expert_table_resource->byte_count == qwen_expert_table_size(token_count)); + REQUIRE(partition_table_resource->byte_count == qwen_partition_table_size(token_count)); + + const size_t gate_index = producer_index_for_tensor(imported.graph, tensors.gate); + ggml::hrx::DispatchMatch gate_up_match; + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, gate_index, gate_up_match)); + append_match_to_plan(plan, gate_up_match, covered_nodes); + + REQUIRE(plan.dispatches.size() == 4); + REQUIRE(plan.transients.size() == 3); + REQUIRE(plan.metadata.alternate_values().size() == 1); + const ggml::hrx::CommandPlanTransient & f16_output_transient = plan.transients.back(); + REQUIRE(f16_output_transient.name == "qwen.moe.gate_up_swiglu_f16"); + REQUIRE(f16_output_transient.size == qwen_routed_gate_up_f16_output_size(token_count)); + REQUIRE(plan.metadata.alternate_values().front().graph_value == glu_value->id); + REQUIRE(plan.metadata.alternate_values().front().alternate_value == f16_output_transient.value); + REQUIRE(plan.metadata.alternate_values().front().type == GGML_TYPE_F16); + REQUIRE(plan.metadata.alternate_values().front().byte_count == f16_output_transient.size); + + const ggml::hrx::Dispatch & dispatch = plan.dispatches.back(); + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_mul_mat_id_swiglu_f16_f16_wmma"); + REQUIRE(dispatch.kernel.integer_parameters.at("token_count") == token_count); + REQUIRE(dispatch.bindings.size() == 6); + REQUIRE(dispatch.bindings[1].value == routing_bundle->expert_table); + REQUIRE(dispatch.bindings[1].length == qwen_expert_table_size(token_count)); + REQUIRE(dispatch.bindings[2].value == routing_bundle->partition_table); + REQUIRE(dispatch.bindings[2].length == qwen_partition_table_size(token_count)); + REQUIRE(dispatch.bindings[5].value == f16_output_transient.value); + REQUIRE(dispatch.bindings[5].length == f16_output_transient.size); + require_compile_parameter(dispatch, "ggml.mul_mat_id_swiglu_f16_f16.input_size", "2048"); + require_compile_parameter(dispatch, "ggml.mul_mat_id_swiglu_f16_f16.expert_count", "128"); + require_compile_parameter(dispatch, "ggml.mul_mat_id_swiglu_f16_f16.route_count", "8"); + require_compile_parameter(dispatch, "ggml.mul_mat_id_swiglu_f16_f16.output_size", "768"); + require_compile_parameter(dispatch, "ggml.mul_mat_id_swiglu_f16_f16.gate_weight_format", "4"); + require_compile_parameter(dispatch, "ggml.mul_mat_id_swiglu_f16_f16.up_weight_format", "4"); + + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 4); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands[3].bindings.size() == 6); + REQUIRE(commands.commands[3].bindings[0].name == "input"); + REQUIRE(commands.commands[3].bindings[0].origin == ggml::hrx::CommandBindingOrigin::GraphValue); + REQUIRE(commands.commands[3].bindings[1].name == "expert_table"); + REQUIRE(commands.commands[3].bindings[1].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands[3].bindings[2].name == "partition_table"); + REQUIRE(commands.commands[3].bindings[2].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands[3].bindings[3].name == "gate_weight"); + REQUIRE(commands.commands[3].bindings[3].origin == ggml::hrx::CommandBindingOrigin::GraphValue); + REQUIRE(commands.commands[3].bindings[4].name == "up_weight"); + REQUIRE(commands.commands[3].bindings[4].origin == ggml::hrx::CommandBindingOrigin::GraphValue); + REQUIRE(commands.commands[3].bindings[5].name == "output"); + REQUIRE(commands.commands[3].bindings[5].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(ggml::hrx::find_transient_allocation(commands.transients, f16_output_transient.value) != nullptr); + } + + { + constexpr int64_t token_count = 1; + QwenRoutedGateUpTensors tensors = build_qwen_routed_gate_up_graph(ctx, token_count); + append_qwen_weighted_reduce_tail(ctx, tensors); + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + const size_t gate_index = producer_index_for_tensor(imported.graph, tensors.gate); + ggml::hrx::DispatchMatch gate_up_match; + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, gate_index, gate_up_match)); + append_match_to_plan(plan, gate_up_match, covered_nodes); + + REQUIRE(plan.dispatches.size() == 4); + REQUIRE(plan.transients.size() == 3); + REQUIRE(kernel_name_for_id(plan.dispatches.back().kernel.kernel_id) == + "loom_libs:ggml_mul_mat_id_swiglu_f16_f16_wmma"); + REQUIRE(plan.dispatches.back().kernel.integer_parameters.at("token_count") == 1); + + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + + { + constexpr int64_t token_count = 4; + ggml_tensor * first_logits = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenRouterExpertCount, token_count); + ggml_tensor * first_route_ids = nullptr; + REQUIRE(first_logits != nullptr); + ggml_tensor * first_route_weights = build_qwen_router_top8_graph(ctx, first_logits, &first_route_ids); + REQUIRE(first_route_ids != nullptr); + REQUIRE(first_route_weights != nullptr); + + QwenRoutedGateUpTensors tensors = build_qwen_routed_gate_up_graph(ctx, token_count); + append_qwen_weighted_reduce_tail(ctx, tensors); + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, first_route_weights); + ggml_build_forward_expand(graph, tensors.route_weights); + ggml_build_forward_expand(graph, tensors.residual); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + const ggml::hrx::Value * first_route_ids_value = imported.graph.values().find_tensor(first_route_ids); + const ggml::hrx::Value * second_route_ids_value = imported.graph.values().find_tensor(tensors.route_ids); + REQUIRE(first_route_ids_value != nullptr); + REQUIRE(second_route_ids_value != nullptr); + REQUIRE(first_route_ids_value->id != second_route_ids_value->id); + + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan; + size_t router_matches = 0; + for (size_t i = 0; i < imported.graph.nodes().size(); ++i) { + if (imported.graph.nodes()[i].op != GGML_OP_SOFT_MAX) { + continue; + } + ggml::hrx::DispatchMatch router_match; + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, i, router_match)); + append_match_to_plan(plan, router_match, covered_nodes); + ++router_matches; + } + REQUIRE(router_matches == 2); + REQUIRE(plan.metadata.generated_resources().size() == 4); + REQUIRE(plan.metadata.moe_routing_bundles().size() == 2); + + const ggml::hrx::CommandPlanGeneratedResource * first_expert_table = plan.metadata.find_generated_resource( + first_route_ids_value->id, ggml::hrx::GeneratedResourceRole::MoeExpertTable); + const ggml::hrx::CommandPlanGeneratedResource * second_expert_table = plan.metadata.find_generated_resource( + second_route_ids_value->id, ggml::hrx::GeneratedResourceRole::MoeExpertTable); + const ggml::hrx::CommandPlanGeneratedResource * second_partition_table = plan.metadata.find_generated_resource( + second_route_ids_value->id, ggml::hrx::GeneratedResourceRole::MoePartitionTable); + const ggml::hrx::CommandPlanMoeRoutingBundle * second_routing_bundle = + plan.metadata.find_moe_routing_bundle(second_route_ids_value->id); + REQUIRE(first_expert_table != nullptr); + REQUIRE(second_expert_table != nullptr); + REQUIRE(second_partition_table != nullptr); + REQUIRE(second_routing_bundle != nullptr); + REQUIRE(second_routing_bundle->expert_table == second_expert_table->generated_value); + REQUIRE(second_routing_bundle->partition_table == second_partition_table->generated_value); + REQUIRE(first_expert_table->generated_value != second_expert_table->generated_value); + + const size_t gate_index = producer_index_for_tensor(imported.graph, tensors.gate); + ggml::hrx::DispatchMatch gate_up_match; + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, gate_index, gate_up_match)); + append_match_to_plan(plan, gate_up_match, covered_nodes); + + REQUIRE(plan.dispatches.size() == 7); + REQUIRE(plan.transients.size() == 5); + const ggml::hrx::Dispatch & dispatch = plan.dispatches.back(); + REQUIRE(dispatch.bindings.size() == 6); + REQUIRE(dispatch.bindings[1].value == second_routing_bundle->expert_table); + REQUIRE(dispatch.bindings[1].value != first_expert_table->generated_value); + REQUIRE(dispatch.bindings[2].value == second_routing_bundle->partition_table); + + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 7); + REQUIRE(command_program_verifies(commands)); + } + + { + constexpr int64_t token_count = 4; + const QwenRoutedGateUpTensors tensors = build_qwen_routed_gate_up_graph(ctx, token_count); + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + append_qwen_routed_gate_up_for_graph(imported.graph, tensors, covered_nodes, plan); + REQUIRE(kernel_name_for_id(plan.dispatches.back().kernel.kernel_id) == + common_mul_mat_id_swiglu_f32_kernel_name(token_count)); + ggml::hrx::DispatchMatch down_match = require_common_mul_mat_id_for_graph( + imported.graph, plan, covered_nodes, tensors.output, tensors.glu, tensors.down_weight, tensors.route_ids); + append_match_to_plan(plan, down_match, covered_nodes, &imported.graph); + + REQUIRE(plan.dispatches.size() == 5); + REQUIRE(plan.transients.size() == 2); + REQUIRE(plan.metadata.alternate_values().empty()); + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 5); + REQUIRE(command_program_verifies(commands)); + } + + { + constexpr int64_t token_count = 4; + QwenRoutedGateUpTensors tensors = build_qwen_routed_gate_up_graph(ctx, token_count); + append_qwen_weighted_reduce_tail(ctx, tensors); + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + append_qwen_routed_gate_up_for_graph(imported.graph, tensors, covered_nodes, plan); + append_qwen_routed_down_for_graph(imported.graph, tensors, covered_nodes, plan, + "loom_libs:ggml_mul_mat_id_f16_f16_wmma"); + append_qwen_weighted_reduce_for_graph(imported.graph, tensors, covered_nodes, plan, + "qwen3_moe:qwen3_moe_routed_down_weighted_reduce_f16_f32"); + + REQUIRE(plan.dispatches.size() == 6); + REQUIRE(plan.transients.size() == 4); + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 6); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.back().bindings.size() == 4); + REQUIRE(commands.commands.back().bindings[0].name == "route_weights"); + REQUIRE(commands.commands.back().bindings[1].name == "routed_output"); + REQUIRE(commands.commands.back().bindings[1].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands.back().bindings[2].name == "residual_input"); + REQUIRE(commands.commands.back().bindings[2].access == ggml::hrx::ResourceAccess::Read); + REQUIRE(commands.commands.back().bindings[3].name == "output"); + REQUIRE(commands.commands.back().bindings[3].access == ggml::hrx::ResourceAccess::Write); + } + + { + constexpr int64_t token_count = 1; + QwenRoutedGateUpTensors tensors = build_qwen_routed_gate_up_graph(ctx, token_count); + append_qwen_weighted_reduce_tail(ctx, tensors); + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + append_qwen_routed_gate_up_for_graph(imported.graph, tensors, covered_nodes, plan); + append_qwen_routed_down_for_graph(imported.graph, tensors, covered_nodes, plan, + "loom_libs:ggml_mul_mat_id_f16_f16_wmma"); + append_qwen_weighted_reduce_for_graph(imported.graph, tensors, covered_nodes, plan, + "qwen3_moe:qwen3_moe_routed_down_weighted_reduce_f16_f32"); + + REQUIRE(plan.dispatches.size() == 6); + REQUIRE(plan.dispatches.back().kernel.integer_parameters.at("token_count") == 1); + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + + { + constexpr int64_t token_count = 4; + QwenRoutedGateUpTensors tensors = build_qwen_routed_gate_up_graph(ctx, token_count); + ggml_tensor * glu_reshape = ggml_reshape_4d(ctx, tensors.glu, tensors.glu->ne[0], tensors.glu->ne[1], + tensors.glu->ne[2], tensors.glu->ne[3]); + REQUIRE(glu_reshape != nullptr); + tensors.output = ggml_mul_mat_id(ctx, tensors.down_weight, glu_reshape, tensors.route_ids); + REQUIRE(tensors.output != nullptr); + ggml_tensor * routed_output_view = + ggml_view_4d(ctx, tensors.output, tensors.output->ne[0], tensors.output->ne[1], tensors.output->ne[2], + tensors.output->ne[3], tensors.output->nb[1], tensors.output->nb[2], tensors.output->nb[3], 0); + REQUIRE(routed_output_view != nullptr); + append_qwen_weighted_reduce_tail(ctx, tensors, false, routed_output_view); + + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + const ggml::hrx::Value * glu_reshape_value = imported.graph.values().find_tensor(glu_reshape); + const ggml::hrx::Value * routed_output_view_value = imported.graph.values().find_tensor(routed_output_view); + REQUIRE(glu_reshape_value != nullptr); + REQUIRE(routed_output_view_value != nullptr); + const ggml::hrx::GraphNode * glu_reshape_node = imported.graph.index().producer(glu_reshape_value->id); + const ggml::hrx::GraphNode * routed_output_view_node = + imported.graph.index().producer(routed_output_view_value->id); + REQUIRE(glu_reshape_node != nullptr); + REQUIRE(routed_output_view_node != nullptr); + REQUIRE(glu_reshape_node->op == GGML_OP_RESHAPE); + REQUIRE(routed_output_view_node->op == GGML_OP_VIEW); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + append_qwen_routed_gate_up_for_graph(imported.graph, tensors, covered_nodes, plan); + const ggml::hrx::Value * glu_value = imported.graph.values().find_tensor(tensors.glu); + REQUIRE(glu_value != nullptr); + const ggml::hrx::CommandPlanAlternateValue * gate_up_alternate = ggml::hrx::find_alternate_value( + imported.graph, plan, glu_value->id, GGML_TYPE_F16, qwen_routed_gate_up_f16_output_size(token_count)); + REQUIRE(gate_up_alternate != nullptr); + const ggml::hrx::ValueId gate_up_alternate_value = gate_up_alternate->alternate_value; + + append_qwen_routed_down_for_graph(imported.graph, tensors, covered_nodes, plan, + "loom_libs:ggml_mul_mat_id_f16_f16_wmma"); + REQUIRE(plan.dispatches.back().bindings[0].value == gate_up_alternate_value); + const ggml::hrx::Value * routed_output_value = imported.graph.values().find_tensor(tensors.output); + REQUIRE(routed_output_value != nullptr); + const ggml::hrx::CommandPlanAlternateValue * routed_down_alternate = ggml::hrx::find_alternate_value( + imported.graph, plan, routed_output_value->id, GGML_TYPE_F16, + qwen_routed_down_f16_output_size(token_count)); + REQUIRE(routed_down_alternate != nullptr); + const ggml::hrx::ValueId routed_down_alternate_value = routed_down_alternate->alternate_value; + + append_qwen_weighted_reduce_for_graph(imported.graph, tensors, covered_nodes, plan, + "qwen3_moe:qwen3_moe_routed_down_weighted_reduce_f16_f32"); + REQUIRE(plan.dispatches.back().bindings[1].value == routed_down_alternate_value); + for (const ggml::hrx::Dispatch & dispatch : plan.dispatches) { + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) != "loom_libs:ggml_copy_f32_f16"); + } + + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(command_program_verifies(commands)); + } + + { + constexpr int64_t token_count = 4; + QwenRoutedGateUpTensors tensors = build_qwen_routed_gate_up_graph(ctx, token_count); + append_qwen_weighted_reduce_tail(ctx, tensors, true); + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + append_qwen_routed_gate_up_for_graph(imported.graph, tensors, covered_nodes, plan); + append_qwen_routed_down_for_graph(imported.graph, tensors, covered_nodes, plan, + "loom_libs:ggml_mul_mat_id_f16_f16_wmma"); + append_qwen_weighted_reduce_for_graph(imported.graph, tensors, covered_nodes, plan, + "qwen3_moe:qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32"); + + REQUIRE(plan.dispatches.size() == 6); + REQUIRE(plan.transients.size() == 4); + const ggml::hrx::Value * hidden_state_value = imported.graph.values().find_tensor(tensors.hidden_state); + const ggml::hrx::Value * residual_value = imported.graph.values().find_tensor(tensors.residual); + REQUIRE(hidden_state_value != nullptr); + REQUIRE(residual_value != nullptr); + REQUIRE(residual_value->alias_source == hidden_state_value->id); + REQUIRE(imported.graph.values().same_storage(hidden_state_value->id, residual_value->id)); + REQUIRE(residual_value->storage_root == hidden_state_value->storage_root); + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 6); + REQUIRE(command_program_verifies(commands)); + REQUIRE(commands.commands.back().bindings.size() == 5); + REQUIRE(commands.commands.back().bindings[0].name == "route_weights"); + REQUIRE(commands.commands.back().bindings[1].name == "routed_output"); + REQUIRE(commands.commands.back().bindings[1].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands.back().bindings[2].name == "hidden_state"); + REQUIRE(commands.commands.back().bindings[2].value == hidden_state_value->storage_root); + REQUIRE(commands.commands.back().bindings[2].access == ggml::hrx::ResourceAccess::ReadWrite); + REQUIRE(commands.commands.back().bindings[3].name == "next_norm_weight"); + REQUIRE(commands.commands.back().bindings[4].name == "next_projection_input"); + REQUIRE(commands.commands.back().bindings[4].access == ggml::hrx::ResourceAccess::Write); + } + + { + constexpr int64_t token_count = 4; + QwenRoutedGateUpTensors tensors = build_qwen_routed_gate_up_graph(ctx, token_count); + append_qwen_weighted_reduce_tail(ctx, tensors, true); + ggml_tensor * hidden_bias = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenMoeHiddenSize, token_count); + REQUIRE(hidden_bias != nullptr); + tensors.hidden_use = ggml_add(ctx, tensors.hidden_state, hidden_bias); + REQUIRE(tensors.hidden_use != nullptr); + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + append_qwen_routed_gate_up_for_graph(imported.graph, tensors, covered_nodes, plan); + append_qwen_routed_down_for_graph(imported.graph, tensors, covered_nodes, plan, + "loom_libs:ggml_mul_mat_id_f16_f16_wmma"); + + const size_t weighted_index = producer_index_for_tensor(imported.graph, tensors.weighted); + ggml::hrx::DispatchMatch weighted_match; + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, weighted_index, weighted_match)); + append_match_to_plan(plan, weighted_match, covered_nodes, &imported.graph); + + REQUIRE(plan.dispatches.size() == 6); + REQUIRE(kernel_name_for_id(plan.dispatches.back().kernel.kernel_id) == + "qwen3_moe:qwen3_moe_routed_down_weighted_reduce_f16_f32"); + REQUIRE(weighted_match.value_aliases.empty()); + REQUIRE(plan.dispatches.back().bindings.size() == 4); + const ggml::hrx::Value * hidden_state_value = imported.graph.values().find_tensor(tensors.hidden_state); + const ggml::hrx::Value * residual_value = imported.graph.values().find_tensor(tensors.residual); + REQUIRE(hidden_state_value != nullptr); + REQUIRE(residual_value != nullptr); + REQUIRE(residual_value->alias_source != hidden_state_value->id); + REQUIRE(!imported.graph.values().same_storage(hidden_state_value->id, residual_value->id)); + } + + { + constexpr int64_t token_count = 4; + QwenRoutedGateUpTensors tensors = + build_qwen_routed_gate_up_graph(ctx, token_count, GGML_GLU_OP_SWIGLU, GGML_TYPE_Q4_K, true, GGML_TYPE_Q4_K); + append_qwen_weighted_reduce_tail(ctx, tensors); + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + append_qwen_routed_gate_up_for_graph(imported.graph, tensors, covered_nodes, plan); + append_qwen_routed_down_for_graph(imported.graph, tensors, covered_nodes, plan, + "loom_libs:ggml_mul_mat_id_f16_f16_wmma"); + + REQUIRE(plan.dispatches.size() == 5); + REQUIRE(plan.transients.size() == 4); + const ggml::hrx::CommandProgram commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 5); + REQUIRE(command_program_verifies(commands)); + } + + { + const QwenRoutedGateUpTensors tensors = build_qwen_routed_gate_up_graph(ctx, 4); + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + const size_t gate_index = producer_index_for_tensor(imported.graph, tensors.gate); + std::vector covered_nodes(imported.graph.nodes().size(), false); + const ggml::hrx::CommandPlan empty_plan; + ggml::hrx::DispatchMatch gate_up_match; + REQUIRE(match_dispatch_at_index(imported.graph, empty_plan, covered_nodes, gate_index, gate_up_match)); + REQUIRE(gate_up_match.dispatches.size() == 3); + REQUIRE(gate_up_match.transients.size() == 2); + REQUIRE(kernel_name_for_id(gate_up_match.dispatches[0].kernel.kernel_id) == + "loom_libs:ggml_moe_build_expert_table"); + REQUIRE(kernel_name_for_id(gate_up_match.dispatches[1].kernel.kernel_id) == + "loom_libs:ggml_moe_build_expert_partition_table"); + REQUIRE(kernel_name_for_id(gate_up_match.dispatches[2].kernel.kernel_id) == + common_mul_mat_id_swiglu_f32_kernel_name(tensors.output->ne[2])); + REQUIRE(gate_up_match.dispatches[2].bindings.size() == 6); + REQUIRE(gate_up_match.dispatches[2].bindings[1].value == gate_up_match.transients[0].value); + REQUIRE(gate_up_match.dispatches[2].bindings[2].value == gate_up_match.transients[1].value); + REQUIRE(gate_up_match.metadata.generated_resources().size() == 2); + REQUIRE(gate_up_match.metadata.moe_routing_bundles().size() == 1); + } + + { + const QwenRoutedGateUpTensors tensors = build_qwen_routed_gate_up_graph(ctx, 4, GGML_GLU_OP_GEGLU); + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + const size_t gate_index = producer_index_for_tensor(imported.graph, tensors.gate); + ggml::hrx::DispatchMatch gate_up_match; + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, gate_index, gate_up_match)); + REQUIRE(gate_up_match.dispatches.size() == 1); + REQUIRE(kernel_name_for_id(gate_up_match.dispatches[0].kernel.kernel_id) == + common_mul_mat_id_f32_kernel_name(tensors.gate->ne[2])); + } + + { + const QwenRoutedGateUpTensors tensors = + build_qwen_routed_gate_up_graph(ctx, 4, GGML_GLU_OP_SWIGLU, GGML_TYPE_Q6_K); + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + const size_t gate_index = producer_index_for_tensor(imported.graph, tensors.gate); + ggml::hrx::DispatchMatch gate_up_match; + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, gate_index, gate_up_match)); + REQUIRE(gate_up_match.dispatches.size() == 1); + REQUIRE(kernel_name_for_id(gate_up_match.dispatches[0].kernel.kernel_id) == + common_mul_mat_id_swiglu_f32_kernel_name(tensors.output->ne[2])); + REQUIRE(gate_up_match.dispatches[0].bindings.size() == 6); + } + + { + const QwenRoutedGateUpTensors tensors = + build_qwen_routed_gate_up_graph(ctx, 4, GGML_GLU_OP_SWIGLU, GGML_TYPE_Q4_K, false); + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + const size_t gate_index = producer_index_for_tensor(imported.graph, tensors.gate); + ggml::hrx::DispatchMatch gate_up_match; + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, gate_index, gate_up_match)); + REQUIRE(gate_up_match.dispatches.size() == 1); + REQUIRE(kernel_name_for_id(gate_up_match.dispatches[0].kernel.kernel_id) == + common_mul_mat_id_swiglu_f32_kernel_name(tensors.output->ne[2])); + REQUIRE(gate_up_match.dispatches[0].bindings.size() == 6); + } + + { + const QwenRoutedGateUpTensors tensors = build_qwen_routed_gate_up_graph(ctx, 4); + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + const size_t down_index = producer_index_for_tensor(imported.graph, tensors.output); + ggml::hrx::DispatchMatch down_match; + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, down_index, down_match)); + REQUIRE(down_match.dispatches.size() == 1); + REQUIRE(kernel_name_for_id(down_match.dispatches[0].kernel.kernel_id) == + common_mul_mat_id_f32_kernel_name(tensors.output->ne[2])); + } + + { + const QwenRoutedGateUpTensors tensors = + build_qwen_routed_gate_up_graph(ctx, 4, GGML_GLU_OP_SWIGLU, GGML_TYPE_Q4_K, true, GGML_TYPE_Q5_K); + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + append_qwen_routed_gate_up_for_graph(imported.graph, tensors, covered_nodes, plan); + const size_t down_index = producer_index_for_tensor(imported.graph, tensors.output); + ggml::hrx::DispatchMatch down_match; + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, down_index, down_match)); + REQUIRE(down_match.dispatches.size() == 1); + REQUIRE(kernel_name_for_id(down_match.dispatches[0].kernel.kernel_id) == + common_mul_mat_id_f32_kernel_name(tensors.output->ne[2])); + } + + ggml_free(ctx); +} + +static void run_qwen_weighted_reduce_next_rmsnorm_q8_publication_checks() { + enum class ConsumerCase { + Direct, + Alias, + UnsupportedWeight, + WrongOperand, + IncompatibleGeometry, + }; + const struct { + ConsumerCase consumer; + bool publishes_q8; + } cases[] = { + { ConsumerCase::Direct, true }, + { ConsumerCase::Alias, true }, + { ConsumerCase::UnsupportedWeight, false }, + { ConsumerCase::WrongOperand, false }, + { ConsumerCase::IncompatibleGeometry, false }, + }; + + for (const auto & test : cases) { + ggml_init_params params = {}; + params.mem_size = 128 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t token_count = 2; + QwenRoutedGateUpTensors tensors = build_qwen_routed_gate_up_graph(ctx, token_count); + append_qwen_weighted_reduce_tail(ctx, tensors, true); + ggml_tensor * consumer_input = tensors.next_output; + if (test.consumer == ConsumerCase::Alias) { + consumer_input = ggml_reshape_2d(ctx, tensors.next_output, kQwenMoeHiddenSize, token_count); + REQUIRE(consumer_input != nullptr); + } + + if (test.consumer == ConsumerCase::WrongOperand) { + ggml_tensor * activation = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenMoeHiddenSize, 64); + REQUIRE(activation != nullptr); + tensors.hidden_use = ggml_mul_mat(ctx, tensors.next_output, activation); + } else { + const ggml_type weight_type = + test.consumer == ConsumerCase::UnsupportedWeight ? GGML_TYPE_BF16 : GGML_TYPE_Q4_K; + const int64_t output_size = test.consumer == ConsumerCase::IncompatibleGeometry ? 63 : 512; + ggml_tensor * weight = ggml_new_tensor_2d(ctx, weight_type, kQwenMoeHiddenSize, output_size); + REQUIRE(weight != nullptr); + tensors.hidden_use = ggml_mul_mat(ctx, weight, consumer_input); + } + REQUIRE(tensors.hidden_use != nullptr); + + ggml::hrx::GraphImportResult imported = import_qwen_routed_gate_up_graph(ctx, tensors); + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan = build_qwen_router_plan_for_graph(imported.graph, covered_nodes); + append_qwen_routed_gate_up_for_graph(imported.graph, tensors, covered_nodes, plan); + append_qwen_routed_down_for_graph(imported.graph, tensors, covered_nodes, plan, + "loom_libs:ggml_mul_mat_id_f16_f16_wmma"); + + const size_t weighted_index = producer_index_for_tensor(imported.graph, tensors.weighted); + ggml::hrx::DispatchMatch weighted_match; + REQUIRE(match_dispatch_at_index(imported.graph, plan, covered_nodes, weighted_index, weighted_match)); + const char * expected_kernel = + test.publishes_q8 ? + "qwen3_moe:qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32_publish_q8" : + "qwen3_moe:qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32"; + REQUIRE(kernel_name_for_id(weighted_match.dispatches.back().kernel.kernel_id) == expected_kernel); + REQUIRE(weighted_match.dispatches.back().bindings.size() == (test.publishes_q8 ? 6 : 5)); + + const ggml::hrx::Value * next_output_value = imported.graph.values().find_tensor(tensors.next_output); + REQUIRE(next_output_value != nullptr); + const ggml::hrx::CommandPlanAlternateValue * q8 = weighted_match.metadata.find_alternate_value( + next_output_value->id, GGML_TYPE_Q8_1, qwen_q8_1_x4_size(token_count, kQwenMoeHiddenSize)); + REQUIRE((q8 != nullptr) == test.publishes_q8); + if (test.publishes_q8) { + REQUIRE(weighted_match.metadata.activation_publication_diagnostics().size() == 1); + const auto & diagnostic = weighted_match.metadata.activation_publication_diagnostics().front(); + REQUIRE(diagnostic.source_value == next_output_value->id); + REQUIRE(diagnostic.alternate_value == q8->alternate_value); + REQUIRE(diagnostic.publication_stage == "published"); + REQUIRE(diagnostic.requested_format == "q8-1-x4"); + } + ggml_free(ctx); + } +} + +static void run_qwen_router_top8_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 2 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + { + ggml_tensor * logits = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenRouterExpertCount, 4); + ggml_tensor * route_ids = nullptr; + REQUIRE(logits != nullptr); + ggml_tensor * output = build_qwen_router_top8_graph(ctx, logits, &route_ids); + schedule_qwen_router_top8_command(ctx, output, route_ids); + } + { + ggml_tensor * logits = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenRouterExpertCount, 1); + ggml_tensor * route_ids = nullptr; + REQUIRE(logits != nullptr); + ggml_tensor * output = build_qwen_router_top8_graph(ctx, logits, &route_ids); + schedule_qwen_router_top8_command(ctx, output, route_ids); + } + { + ggml_tensor * logits = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenRouterExpertCount, 13); + ggml_tensor * route_ids = nullptr; + REQUIRE(logits != nullptr); + ggml_tensor * output = build_qwen_router_top8_graph(ctx, logits, &route_ids, GGML_SORT_ORDER_DESC, + kQwenRouterRouteCount, 0.00006103515625f); + schedule_qwen_router_top8_command(ctx, output, route_ids); + } + { + ggml_tensor * logits = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenRouterExpertCount, 512); + ggml_tensor * route_ids = nullptr; + REQUIRE(logits != nullptr); + ggml_tensor * output = build_qwen_router_top8_graph(ctx, logits, &route_ids); + schedule_qwen_router_top8_command(ctx, output, route_ids); + } + { + ggml_tensor * logits = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 64, 4); + ggml_tensor * route_ids = nullptr; + REQUIRE(logits != nullptr); + ggml_tensor * output = build_qwen_router_top8_graph(ctx, logits, &route_ids); + schedule_qwen_router_top8_command(ctx, output, route_ids, 64, kQwenRouterRouteCount); + } + { + ggml_tensor * logits = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenRouterExpertCount, 4); + ggml_tensor * route_ids = nullptr; + REQUIRE(logits != nullptr); + ggml_tensor * output = build_qwen_router_top8_graph(ctx, logits, &route_ids, GGML_SORT_ORDER_DESC, 4); + schedule_qwen_router_top8_command(ctx, output, route_ids, kQwenRouterExpertCount, 4); + } + { + ManualQwenRouterTop8Graph manual = build_manual_qwen_router_top8_graph(ctx, 64, 4, 5); + schedule_manual_qwen_router_top8_command(manual.graph, manual.route_ids, 5, 64, 4); + } + { + ggml_tensor * logits = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenRouterExpertCount, 4); + REQUIRE(logits != nullptr); + ggml_tensor * output = build_qwen_router_top8_graph(ctx, logits, nullptr, GGML_SORT_ORDER_ASC); + REQUIRE(!graph_is_supported(ctx, output)); + } + { + ggml_tensor * logits = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenRouterExpertCount, 4); + REQUIRE(logits != nullptr); + ggml_tensor * output = build_qwen_router_top8_graph(ctx, logits, nullptr, GGML_SORT_ORDER_DESC, 33); + REQUIRE(!graph_is_supported(ctx, output)); + } + { + ggml_tensor * logits = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 16, 4); + REQUIRE(logits != nullptr); + ggml_tensor * output = build_qwen_router_top8_graph(ctx, logits); + REQUIRE(!graph_is_supported(ctx, output)); + } + { + ggml_tensor * logits = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, kQwenRouterExpertCount, 4); + REQUIRE(logits != nullptr); + ggml_tensor * output = build_qwen_router_top8_graph(ctx, logits); + REQUIRE(!graph_is_supported(ctx, output)); + } + + ggml_free(ctx); +} + +static void bind_external_values(ggml::hrx::ValueMap & values) { + uintptr_t buffer = 0x1000; + for (const ggml::hrx::ValueId id : values.external_value_ids()) { + const ggml::hrx::Value * value = values.find(id); + REQUIRE(value != nullptr); + REQUIRE(values.bind_buffer(id, { dummy_hrx_buffer(buffer), 0, value->byte_count })); + buffer += 0x1000; + } +} + +static void run_alias_value_import_checks() { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * source = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * view = ggml_view_1d(ctx, source, 4, 2 * sizeof(float)); + REQUIRE(source != nullptr); + REQUIRE(view != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, view); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 1); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_VIEW); + + const ggml::hrx::Value * source_value = imported.graph.values().find_tensor(source); + const ggml::hrx::Value * view_value = imported.graph.values().find_tensor(view); + REQUIRE(source_value != nullptr); + REQUIRE(view_value != nullptr); + REQUIRE(source_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(view_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(imported.graph.values().same_storage(source_value->id, view_value->id)); + REQUIRE(view_value->storage_root == source_value->id); + REQUIRE(view_value->alias_source == source_value->id); + REQUIRE(view_value->storage_offset == 2 * sizeof(float)); + REQUIRE(view_value->byte_count == 4 * sizeof(float)); + REQUIRE(ggml::hrx::is_layout_alias_node(imported.graph, imported.graph.nodes()[0])); + + REQUIRE(imported.graph.values().bind_buffer(source_value->id, + { dummy_hrx_buffer(0x4000), 128, source_value->byte_count })); + const ggml::hrx::CommandProgramBindings bindings = + ggml::hrx::CommandProgramBindings::from_value_map(imported.graph.values()); + REQUIRE(bindings.valid()); + const ggml::hrx::CommandProgramBinding * view_binding = bindings.find(view_value->id); + REQUIRE(view_binding != nullptr); + REQUIRE(view_binding->buffer == dummy_hrx_buffer(0x4000)); + REQUIRE(view_binding->offset == 128 + 2 * sizeof(float)); + REQUIRE(view_binding->length == view_value->byte_count); + + ggml_tensor * internal_source = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 8, 2); + ggml_tensor * internal_bias = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 8, 1); + ggml_tensor * internal_view = + ggml_view_2d(ctx, internal_source, 8, 1, internal_source->nb[1], internal_source->nb[1]); + ggml_tensor * internal_out = ggml_add(ctx, internal_view, internal_bias); + REQUIRE(internal_source != nullptr); + REQUIRE(internal_bias != nullptr); + REQUIRE(internal_view != nullptr); + REQUIRE(internal_out != nullptr); + + ggml_cgraph * internal_graph = ggml_new_graph(ctx); + REQUIRE(internal_graph != nullptr); + ggml_build_forward_expand(internal_graph, internal_out); + + ggml::hrx::GraphImportResult internal_imported = ggml::hrx::import_ggml_graph(*internal_graph); + REQUIRE(internal_imported.valid()); + REQUIRE(internal_imported.graph.nodes().size() == 2); + REQUIRE(internal_imported.graph.nodes()[0].op == GGML_OP_VIEW); + REQUIRE(internal_imported.graph.nodes()[1].op == GGML_OP_ADD); + + const ggml::hrx::Value * internal_source_value = internal_imported.graph.values().find_tensor(internal_source); + const ggml::hrx::Value * internal_view_value = internal_imported.graph.values().find_tensor(internal_view); + const ggml::hrx::Value * internal_bias_value = internal_imported.graph.values().find_tensor(internal_bias); + const ggml::hrx::Value * internal_out_value = internal_imported.graph.values().find_tensor(internal_out); + REQUIRE(internal_source_value != nullptr); + REQUIRE(internal_view_value != nullptr); + REQUIRE(internal_bias_value != nullptr); + REQUIRE(internal_out_value != nullptr); + REQUIRE(internal_source_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(internal_view_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(internal_view_value->storage_root == internal_source_value->id); + REQUIRE(internal_view_value->storage_offset == internal_source->nb[1]); + + ggml::hrx::DispatchScheduler internal_scheduler; + REQUIRE(internal_scheduler.schedule_graph(internal_imported.graph, test_dispatch_target())); + + REQUIRE(internal_imported.graph.values().bind_buffer( + internal_source_value->id, { dummy_hrx_buffer(0x5000), 256, internal_source_value->byte_count })); + REQUIRE(internal_imported.graph.values().bind_buffer( + internal_bias_value->id, { dummy_hrx_buffer(0x6000), 0, internal_bias_value->byte_count })); + REQUIRE(internal_imported.graph.values().bind_buffer( + internal_out_value->id, { dummy_hrx_buffer(0x7000), 0, internal_out_value->byte_count })); + const ggml::hrx::CommandProgramBindings internal_bindings = + ggml::hrx::CommandProgramBindings::from_value_map(internal_imported.graph.values()); + REQUIRE(internal_bindings.valid()); + const ggml::hrx::CommandProgramBinding * internal_view_binding = internal_bindings.find(internal_view_value->id); + REQUIRE(internal_view_binding != nullptr); + REQUIRE(internal_view_binding->buffer == dummy_hrx_buffer(0x5000)); + REQUIRE(internal_view_binding->offset == 256 + internal_source->nb[1]); + REQUIRE(internal_view_binding->length == internal_view_value->byte_count); + + ggml_tensor * cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 8, 4); + ggml_tensor * rows = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 8, 2); + ggml_tensor * row_indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, 2); + ggml_tensor * updated = ggml_set_rows(ctx, cache, rows, row_indices); + REQUIRE(cache != nullptr); + REQUIRE(rows != nullptr); + REQUIRE(row_indices != nullptr); + REQUIRE(updated != nullptr); + ggml_tensor * updated_view = ggml_view_2d(ctx, updated, 8, 2, updated->nb[1], 0); + ggml_tensor * updated_bias = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 8, 2); + ggml_tensor * updated_out = ggml_add(ctx, updated_view, updated_bias); + REQUIRE(updated_view != nullptr); + REQUIRE(updated_bias != nullptr); + REQUIRE(updated_out != nullptr); + + ggml_cgraph * set_rows_graph = ggml_new_graph(ctx); + REQUIRE(set_rows_graph != nullptr); + ggml_build_forward_expand(set_rows_graph, updated_out); + + ggml::hrx::GraphImportResult set_rows_imported = ggml::hrx::import_ggml_graph(*set_rows_graph); + REQUIRE(set_rows_imported.valid()); + REQUIRE(set_rows_imported.graph.nodes().size() == 3); + REQUIRE(set_rows_imported.graph.nodes()[0].op == GGML_OP_SET_ROWS); + REQUIRE(set_rows_imported.graph.nodes()[1].op == GGML_OP_VIEW); + REQUIRE(set_rows_imported.graph.nodes()[2].op == GGML_OP_ADD); + + const ggml::hrx::Value * cache_value = set_rows_imported.graph.values().find_tensor(cache); + const ggml::hrx::Value * updated_value = set_rows_imported.graph.values().find_tensor(updated); + const ggml::hrx::Value * updated_view_value = set_rows_imported.graph.values().find_tensor(updated_view); + REQUIRE(cache_value != nullptr); + REQUIRE(updated_value != nullptr); + REQUIRE(updated_view_value != nullptr); + REQUIRE(cache_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(updated_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(updated_view_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(updated_value->storage_root == cache_value->id); + REQUIRE(updated_view_value->storage_root == cache_value->id); + REQUIRE(updated_view_value->alias_source == cache_value->id); + + ggml::hrx::ValueMap value_map; + const std::array ne = { 8, 1, 1, 1 }; + const std::array nb = { sizeof(float), 8 * sizeof(float), 8 * sizeof(float), + 8 * sizeof(float) }; + REQUIRE(value_map.add_snapshot_storage({ ggml::hrx::ValueStorageId(0), ggml::hrx::ValueId(0), 8 * sizeof(float) }) + .success()); + REQUIRE( + value_map + .add_snapshot_value({ ggml::hrx::ValueId(0), ggml::hrx::ValueKind::External, ggml::hrx::ValueStorageId(0), + ggml::hrx::ValueId(0), ggml::hrx::ValueId(), 0, 8 * sizeof(float), GGML_TYPE_F32, ne, + nb, 8, 8 * sizeof(float), true, nullptr, std::nullopt }) + .success()); + REQUIRE(value_map.add_snapshot_storage({ ggml::hrx::ValueStorageId(1), ggml::hrx::ValueId(1), 8 * sizeof(float) }) + .success()); + REQUIRE( + value_map + .add_snapshot_value({ ggml::hrx::ValueId(1), ggml::hrx::ValueKind::Transient, ggml::hrx::ValueStorageId(1), + ggml::hrx::ValueId(1), ggml::hrx::ValueId(), 0, 8 * sizeof(float), GGML_TYPE_F32, ne, + nb, 8, 8 * sizeof(float), true, nullptr, std::nullopt }) + .success()); + REQUIRE(value_map.alias_storage(ggml::hrx::ValueId(1), ggml::hrx::ValueId(0)).success()); + const ggml::hrx::Value * aliased_value = value_map.find(ggml::hrx::ValueId(1)); + REQUIRE(aliased_value != nullptr); + REQUIRE(aliased_value->kind == ggml::hrx::ValueKind::Transient); + REQUIRE(aliased_value->alias_source == ggml::hrx::ValueId(0)); + REQUIRE(aliased_value->storage_root == ggml::hrx::ValueId(0)); + REQUIRE(value_map.same_storage(ggml::hrx::ValueId(0), ggml::hrx::ValueId(1))); + REQUIRE(value_map.bind_buffer(ggml::hrx::ValueId(0), { dummy_hrx_buffer(0x8000), 64, 8 * sizeof(float) })); + const std::optional aliased_binding = + value_map.resolve_buffer_binding(ggml::hrx::ValueId(1)); + REQUIRE(aliased_binding.has_value()); + REQUIRE(aliased_binding->buffer == dummy_hrx_buffer(0x8000)); + REQUIRE(aliased_binding->offset == 64); + REQUIRE(aliased_binding->length == 8 * sizeof(float)); + + ggml_free(ctx); +} + +static void run_multi_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * c = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * d = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * out0 = ggml_add(ctx, a, b); + ggml_tensor * out1 = ggml_add(ctx, c, d); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(c != nullptr); + REQUIRE(d != nullptr); + REQUIRE(out0 != nullptr); + REQUIRE(out1 != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out0); + ggml_build_forward_expand(graph, out1); + graph->uid = 1001; + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 2); + const ggml::hrx::Value * out0_value = imported.graph.values().find_tensor(out0); + const ggml::hrx::Value * out1_value = imported.graph.values().find_tensor(out1); + REQUIRE(out0_value != nullptr); + REQUIRE(out1_value != nullptr); + REQUIRE(imported.graph.nodes()[0].output == out0_value->id); + REQUIRE(imported.graph.nodes()[1].output == out1_value->id); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 2); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 2); + REQUIRE(commands.commands[0].ordinal == 0); + REQUIRE(commands.commands[0].dependencies.empty()); + REQUIRE(commands.commands[1].ordinal == 1); + REQUIRE(commands.commands[1].dependencies.size() == 1); + REQUIRE(commands.commands[1].dependencies[0] == 0); + REQUIRE(command_program_verifies(commands)); + + bind_external_values(imported.graph.values()); + const ggml::hrx::CommandProgramBindings bindings = + ggml::hrx::CommandProgramBindings::from_value_map(imported.graph.values()); + REQUIRE(bindings.valid()); + ggml::hrx::ResolvedCommandProgram resolved = ggml::hrx::resolve_command_program_bindings(commands, bindings); + REQUIRE(resolved.valid()); + REQUIRE(resolved.commands.size() == 2); + + ggml_free(ctx); +} + +static void run_layout_alias_scheduler_elision_checks() { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * sum = ggml_add(ctx, a, b); + ggml_tensor * reshaped = ggml_reshape_2d(ctx, sum, 4, 2); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(sum != nullptr); + REQUIRE(reshaped != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, reshaped); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 2); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_ADD); + REQUIRE(imported.graph.nodes()[1].op == GGML_OP_RESHAPE); + + const ggml::hrx::Value * sum_value = imported.graph.values().find_tensor(sum); + const ggml::hrx::Value * reshaped_value = imported.graph.values().find_tensor(reshaped); + REQUIRE(sum_value != nullptr); + REQUIRE(reshaped_value != nullptr); + REQUIRE(sum_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(reshaped_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(imported.graph.values().same_storage(sum_value->id, reshaped_value->id)); + REQUIRE(reshaped_value->storage_root == sum_value->id); + REQUIRE(reshaped_value->alias_source == sum_value->id); + REQUIRE(ggml::hrx::is_layout_alias_node(imported.graph, imported.graph.nodes()[1])); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 1); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 1); + REQUIRE(commands.commands[0].bindings.size() == 3); + REQUIRE(commands.commands[0].bindings[2].value == sum_value->id); + REQUIRE(commands.commands[0].bindings[2].origin == ggml::hrx::CommandBindingOrigin::GraphValue); + REQUIRE(commands.transients.allocations.empty()); + REQUIRE(ggml::hrx::find_transient_allocation(commands.transients, sum_value->id) == nullptr); + REQUIRE(ggml::hrx::find_transient_allocation(commands.transients, reshaped_value->id) == nullptr); + REQUIRE(command_program_verifies(commands)); + + ggml_free(ctx); +} + +static void run_zero_output_scheduler_elision_checks() { + ggml_init_params params = {}; + params.mem_size = 512 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + { + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 3840, 1); + ggml_tensor * ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 0); + REQUIRE(input != nullptr); + REQUIRE(ids != nullptr); + ggml_tensor * rows = ggml_get_rows(ctx, input, ids); + REQUIRE(rows != nullptr); + REQUIRE(rows->ne[0] == 3840); + REQUIRE(rows->ne[1] == 0); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, rows); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 1); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_GET_ROWS); + const ggml::hrx::Value * rows_value = imported.graph.values().find_tensor(rows); + REQUIRE(rows_value != nullptr); + REQUIRE(rows_value->element_count == 0); + REQUIRE(rows_value->byte_count == 0); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.empty()); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.empty()); + REQUIRE(command_program_verifies(commands)); + + ggml::hrx::GraphProgramCache cache; + ggml::hrx::GraphProgramLookup lookup = + cache.get_or_build(*graph, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(lookup.valid()); + REQUIRE(lookup.match.external_bindings.size() == 1); + REQUIRE(lookup.match.external_bindings[0].tensor == input); + + const auto * ids_value = lookup.program->graph().values().find_tensor(ids); + REQUIRE(ids_value != nullptr); + ggml::hrx::Command referenced; + referenced.bindings.push_back({ "ids", ids_value->id }); + for (auto * list : + { &lookup.program->commands().initialization_commands, &lookup.program->commands().commands }) { + list->push_back(referenced); + const auto match = lookup.program->match_current_graph(*graph); + REQUIRE(match.valid()); + REQUIRE(match.external_bindings.size() == 2); + REQUIRE(std::any_of(match.external_bindings.begin(), match.external_bindings.end(), + [&](const auto & binding) { return binding.tensor == ids; })); + list->clear(); + const auto unreferenced = lookup.program->match_current_graph(*graph); + REQUIRE(unreferenced.valid()); + REQUIRE(unreferenced.external_bindings.size() == 1); + for (const auto origin : + { ggml::hrx::CommandBindingOrigin::Transient, ggml::hrx::CommandBindingOrigin::ProgramConstant }) { + ggml::hrx::Command unrelated; + unrelated.bindings.push_back({ "negative", ggml::hrx::ValueId(-1) }); + unrelated.bindings.push_back({ "large", ggml::hrx::ValueId(INT32_MAX) }); + unrelated.bindings.push_back({ "ids", ids_value->id, origin }); + list->push_back(unrelated); + const auto ignored = lookup.program->match_current_graph(*graph); + REQUIRE(ignored.valid()); + REQUIRE(ignored.external_bindings.size() == 1); + REQUIRE(ignored.external_bindings[0].tensor == input); + list->clear(); + } + } + } + + { + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 3840, 1); + ggml_tensor * ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 0); + REQUIRE(input != nullptr); + REQUIRE(ids != nullptr); + ggml_tensor * rows = ggml_get_rows(ctx, input, ids); + ggml_tensor * norm = ggml_rms_norm(ctx, rows, 0.000001f); + REQUIRE(rows != nullptr); + REQUIRE(norm != nullptr); + REQUIRE(norm->ne[0] == 3840); + REQUIRE(norm->ne[1] == 0); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, norm); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 2); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_GET_ROWS); + REQUIRE(imported.graph.nodes()[1].op == GGML_OP_RMS_NORM); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.empty()); + + ggml::hrx::GraphProgramCache cache; + ggml::hrx::GraphProgramLookup lookup = + cache.get_or_build(*graph, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(lookup.valid()); + REQUIRE(lookup.match.external_bindings.size() == 1); + REQUIRE(lookup.match.external_bindings[0].tensor == input); + } + + { + ggml_tensor * lhs = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 3840, 0); + ggml_tensor * rhs = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 3840, 0); + REQUIRE(lhs != nullptr); + REQUIRE(rhs != nullptr); + ggml_tensor * sum = ggml_add(ctx, lhs, rhs); + REQUIRE(sum != nullptr); + REQUIRE(sum->ne[0] == 3840); + REQUIRE(sum->ne[1] == 0); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, sum); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 1); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_ADD); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.empty()); + + ggml::hrx::GraphProgramCache cache; + ggml::hrx::GraphProgramLookup lookup = + cache.get_or_build(*graph, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(lookup.valid()); + REQUIRE(lookup.match.external_bindings.empty()); + } + + { + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 3840, 0); + REQUIRE(input != nullptr); + ggml_tensor * scaled = ggml_scale(ctx, input, 2.0f); + REQUIRE(scaled != nullptr); + REQUIRE(scaled->ne[0] == 3840); + REQUIRE(scaled->ne[1] == 0); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, scaled); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 1); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_SCALE); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.empty()); + + ggml::hrx::GraphProgramCache cache; + ggml::hrx::GraphProgramLookup lookup = + cache.get_or_build(*graph, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(lookup.valid()); + REQUIRE(lookup.match.external_bindings.empty()); + } + + { + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 3840, 1); + REQUIRE(input != nullptr); + ggml_tensor * sin = ggml_sin(ctx, input); + REQUIRE(sin != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, sin); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(!scheduler.schedule_graph(imported.graph, test_dispatch_target())); + } + + ggml_free(ctx); +} + +static void run_zero_output_device_support_checks() { + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 512 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * lhs = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 3840, 0); + ggml_tensor * rhs = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 3840, 0); + REQUIRE(lhs != nullptr); + REQUIRE(rhs != nullptr); + + ggml_tensor * sum = ggml_add(ctx, lhs, rhs); + ggml_tensor * prod = ggml_mul(ctx, lhs, rhs); + ggml_tensor * scaled = ggml_scale(ctx, lhs, 2.0f); + REQUIRE(sum != nullptr); + REQUIRE(prod != nullptr); + REQUIRE(scaled != nullptr); + REQUIRE(ggml_backend_supports_op(backend, sum)); + REQUIRE(ggml_backend_supports_op(backend, prod)); + REQUIRE(ggml_backend_supports_op(backend, scaled)); + + ggml_tensor * acc = ggml_acc(ctx, lhs, rhs, lhs->nb[1], lhs->nb[2], lhs->nb[3], 0); + ggml_tensor * set = ggml_set(ctx, lhs, rhs, lhs->nb[1], lhs->nb[2], lhs->nb[3], 0); + ggml_tensor * cpy = ggml_cpy(ctx, lhs, rhs); + REQUIRE(acc != nullptr); + REQUIRE(set != nullptr); + REQUIRE(cpy != nullptr); + REQUIRE(!ggml_backend_supports_op(backend, acc)); + REQUIRE(!ggml_backend_supports_op(backend, set)); + REQUIRE(ggml_backend_supports_op(backend, cpy)); + + ggml_free(ctx); + ggml_backend_free(backend); +} + +static void run_scale_f32_device_support_checks() { + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 512 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 32); + ggml_tensor * out = ggml_scale_bias(ctx, input, 1.5f, 0.5f); + ggml_tensor * inplace = ggml_scale_bias_inplace(ctx, input, -2.0f, 1.0f); + ggml_tensor * view = ggml_view_1d(ctx, input, 16, 8 * sizeof(float)); + ggml_tensor * view_out = ggml_scale_bias(ctx, view, 0.25f, -0.5f); + ggml_tensor * view_inpl = ggml_scale_bias_inplace(ctx, view, 0.5f, 0.25f); + REQUIRE(input != nullptr); + REQUIRE(out != nullptr); + REQUIRE(inplace != nullptr); + REQUIRE(view != nullptr); + REQUIRE(view_out != nullptr); + REQUIRE(view_inpl != nullptr); + + REQUIRE(ggml_backend_supports_op(backend, out)); + REQUIRE(ggml_backend_supports_op(backend, inplace)); + REQUIRE(ggml_backend_supports_op(backend, view_out)); + REQUIRE(ggml_backend_supports_op(backend, view_inpl)); + + ggml_free(ctx); + ggml_backend_free(backend); +} + +static void run_transient_import_checks() { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * sum = ggml_add(ctx, a, b); + ggml_tensor * out = ggml_sin(ctx, sum); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(sum != nullptr); + REQUIRE(out != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 2); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_ADD); + REQUIRE(imported.graph.nodes()[1].op == GGML_OP_SIN); + + const ggml::hrx::Value * a_value = imported.graph.values().find_tensor(a); + const ggml::hrx::Value * sum_value = imported.graph.values().find_tensor(sum); + const ggml::hrx::Value * out_value = imported.graph.values().find_tensor(out); + REQUIRE(a_value != nullptr); + REQUIRE(sum_value != nullptr); + REQUIRE(out_value != nullptr); + REQUIRE(a_value->kind == ggml::hrx::ValueKind::External); + REQUIRE(sum_value->kind == ggml::hrx::ValueKind::Transient); + REQUIRE(out_value->kind == ggml::hrx::ValueKind::External); + REQUIRE( + !imported.graph.values().bind_buffer(sum_value->id, { dummy_hrx_buffer(0x3000), 0, sum_value->byte_count })); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(!scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(!scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.empty()); + REQUIRE(status_contains(scheduler.plan().status, "unsupported HRX node 1")); + + ggml_free(ctx); +} + +static void run_graph_view_preserves_shared_output_storage_checks() { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * c = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * d = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * sum = ggml_add(ctx, a, b); + ggml_tensor * out0 = ggml_add(ctx, sum, c); + ggml_tensor * out1 = ggml_add(ctx, sum, d); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(c != nullptr); + REQUIRE(d != nullptr); + REQUIRE(sum != nullptr); + REQUIRE(out0 != nullptr); + REQUIRE(out1 != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out0); + ggml_build_forward_expand(graph, out1); + REQUIRE(graph->n_nodes == 3); + REQUIRE(graph->nodes[0] == sum); + REQUIRE(graph->nodes[1] == out0); + REQUIRE(graph->nodes[2] == out1); + REQUIRE(ggml_node_get_use_count(graph, 0) == 2); + + ggml_cgraph graph_view = ggml_graph_view(graph, 0, 2); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(graph_view); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 2); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_ADD); + REQUIRE(imported.graph.nodes()[1].op == GGML_OP_ADD); + + const ggml::hrx::Value * sum_value = imported.graph.values().find_tensor(sum); + REQUIRE(sum_value != nullptr); + REQUIRE(sum_value->kind == ggml::hrx::ValueKind::External); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 2); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 2); + REQUIRE(commands.commands[0].bindings[2].value == sum_value->id); + REQUIRE(commands.commands[0].bindings[2].origin == ggml::hrx::CommandBindingOrigin::GraphValue); + REQUIRE(commands.commands[1].bindings[0].value == sum_value->id); + REQUIRE(commands.commands[1].bindings[0].origin == ggml::hrx::CommandBindingOrigin::GraphValue); + REQUIRE(ggml::hrx::find_transient_allocation(commands.transients, sum_value->id) == nullptr); + REQUIRE(command_program_verifies(commands)); + + ggml_free(ctx); +} + +static void run_chained_dispatch_requires_transients() { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * c = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * sum = ggml_add(ctx, a, b); + ggml_tensor * out = ggml_add(ctx, sum, c); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(c != nullptr); + REQUIRE(sum != nullptr); + REQUIRE(out != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 2); + REQUIRE(imported.graph.nodes()[0].op == GGML_OP_ADD); + REQUIRE(imported.graph.nodes()[1].op == GGML_OP_ADD); + + const ggml::hrx::Value * sum_value = imported.graph.values().find_tensor(sum); + REQUIRE(sum_value != nullptr); + REQUIRE(sum_value->kind == ggml::hrx::ValueKind::Transient); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 2); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 2); + REQUIRE(commands.commands[1].dependencies.size() == 1); + REQUIRE(commands.commands[1].dependencies[0] == 0); + REQUIRE(commands.commands[0].bindings.size() == 3); + REQUIRE(commands.commands[1].bindings.size() == 3); + REQUIRE(commands.commands[0].bindings[0].origin == ggml::hrx::CommandBindingOrigin::GraphValue); + REQUIRE(commands.commands[0].bindings[1].origin == ggml::hrx::CommandBindingOrigin::GraphValue); + REQUIRE(commands.commands[0].bindings[2].value == sum_value->id); + REQUIRE(commands.commands[0].bindings[2].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands[1].bindings[0].value == sum_value->id); + REQUIRE(commands.commands[1].bindings[0].origin == ggml::hrx::CommandBindingOrigin::Transient); + REQUIRE(commands.commands[1].bindings[1].origin == ggml::hrx::CommandBindingOrigin::GraphValue); + REQUIRE(commands.commands[1].bindings[2].origin == ggml::hrx::CommandBindingOrigin::GraphValue); + REQUIRE(commands.transients.allocations.size() == 1); + const ggml::hrx::TransientAllocation * sum_allocation = + ggml::hrx::find_transient_allocation(commands.transients, sum_value->id); + REQUIRE(sum_allocation != nullptr); + REQUIRE(sum_allocation->value == sum_value->id); + REQUIRE(sum_allocation->size == sum_value->byte_count); + REQUIRE(sum_allocation->alignment == 256); + REQUIRE(sum_allocation->arena_offset == 0); + REQUIRE(commands.transients.arena_size == 256); + REQUIRE(command_program_verifies(commands)); + + const std::string transient_binding_text = ggml::hrx::format_command_binding(commands.commands[1].bindings[0]); + REQUIRE(string_contains(transient_binding_text, "origin=Transient")); + + bind_external_values(imported.graph.values()); + const ggml::hrx::CommandProgramBindings bindings = + ggml::hrx::CommandProgramBindings::from_value_map(imported.graph.values()); + REQUIRE(bindings.valid()); + REQUIRE(bindings.find(sum_value->id) == nullptr); + + const ggml::hrx::ResolvedCommandProgram resolved = ggml::hrx::resolve_command_program_bindings(commands, bindings); + REQUIRE(!resolved.valid()); + REQUIRE(status_contains(resolved.status, "no transient arena")); + REQUIRE(status_contains(resolved.status, "origin=Transient")); + REQUIRE(status_contains(resolved.status, "value=")); + + const ggml::hrx::TransientArenaAllocationRef transient_arena = { + dummy_hrx_buffer(0x8000), + commands.transients.arena_size, + 7, + }; + const ggml::hrx::ResolvedCommandProgram resolved_with_transients = + ggml::hrx::resolve_command_program_bindings(commands, bindings, &transient_arena); + REQUIRE(resolved_with_transients.valid()); + REQUIRE(resolved_with_transients.commands.size() == 2); + REQUIRE(resolved_with_transients.commands[0].bindings[2].ref.buffer == dummy_hrx_buffer(0x8000)); + REQUIRE(resolved_with_transients.commands[0].bindings[2].ref.offset == 0); + REQUIRE(resolved_with_transients.commands[0].bindings[2].ref.length == sum_value->byte_count); + REQUIRE(resolved_with_transients.commands[1].bindings[0].ref.buffer == dummy_hrx_buffer(0x8000)); + REQUIRE(resolved_with_transients.commands[1].bindings[0].ref.offset == 0); + REQUIRE(resolved_with_transients.commands[1].bindings[0].ref.length == sum_value->byte_count); + + ggml::hrx::PreparedCommandProgram prepared_shape; + for (const ggml::hrx::Command & prepared_source : commands.commands) { + ggml::hrx::PreparedCommand prepared_command; + prepared_command.ordinal = prepared_source.ordinal; + prepared_command.kind = prepared_source.kind; + prepared_command.kernel.specialization = prepared_source.kernel; + for (const ggml::hrx::CommandBinding & binding : prepared_source.bindings) { + prepared_command.kernel.bindings.push_back({ + binding, { dummy_hrx_buffer(0x4000), 123, binding.length } + }); + } + prepared_shape.commands.push_back(prepared_command); + } + prepared_shape.bound_transient_arena_allocation_id = 1; + + REQUIRE(ggml::hrx::bind_prepared_command_program_transients(commands, transient_arena, prepared_shape)); + REQUIRE(prepared_shape.bound_transient_arena_allocation_id == transient_arena.allocation_id); + REQUIRE(prepared_shape.commands[0].kernel.bindings[2].ref.buffer == dummy_hrx_buffer(0x8000)); + REQUIRE(prepared_shape.commands[0].kernel.bindings[2].ref.offset == 0); + REQUIRE(prepared_shape.commands[1].kernel.bindings[0].ref.buffer == dummy_hrx_buffer(0x8000)); + REQUIRE(prepared_shape.commands[1].kernel.bindings[0].ref.offset == 0); + + const ggml::hrx::TransientArenaAllocationRef grown_transient_arena = { + dummy_hrx_buffer(0x9000), + commands.transients.arena_size + 256, + 8, + }; + REQUIRE(ggml::hrx::bind_prepared_command_program_transients(commands, grown_transient_arena, prepared_shape)); + REQUIRE(prepared_shape.bound_transient_arena_allocation_id == grown_transient_arena.allocation_id); + REQUIRE(prepared_shape.commands[0].kernel.bindings[2].ref.buffer == dummy_hrx_buffer(0x9000)); + REQUIRE(prepared_shape.commands[1].kernel.bindings[0].ref.buffer == dummy_hrx_buffer(0x9000)); + + prepared_shape.commands[1].kernel.bindings[0].binding.origin = ggml::hrx::CommandBindingOrigin::ProgramConstant; + prepared_shape.commands[1].kernel.bindings[0].ref = { dummy_hrx_buffer(0xb000), 32, sum_value->byte_count }; + const ggml::hrx::TransientArenaAllocationRef rebinding_transient_arena = { + dummy_hrx_buffer(0xc000), + commands.transients.arena_size + 512, + 9, + }; + REQUIRE(ggml::hrx::bind_prepared_command_program_transients(commands, rebinding_transient_arena, prepared_shape)); + REQUIRE(prepared_shape.commands[0].kernel.bindings[2].ref.buffer == dummy_hrx_buffer(0xc000)); + REQUIRE(prepared_shape.commands[1].kernel.bindings[0].ref.buffer == dummy_hrx_buffer(0xb000)); + REQUIRE(prepared_shape.commands[1].kernel.bindings[0].ref.offset == 32); + + const ggml::hrx::TransientArenaAllocationRef invalid_transient_arena = { + dummy_hrx_buffer(0xa000), + commands.transients.arena_size, + ggml::hrx::kInvalidTransientArenaAllocationId, + }; + const ggml::hrx::ResolvedCommandProgram invalid_transient_resolved = + ggml::hrx::resolve_command_program_bindings(commands, bindings, &invalid_transient_arena); + REQUIRE(!invalid_transient_resolved.valid()); + REQUIRE(status_contains(invalid_transient_resolved.status, "no transient arena allocation id")); + + ggml::hrx::CommandProgram missing_allocation = copy_command_program_shape(commands); + missing_allocation.transients.allocations.clear(); + ggml::hrx::VerificationResult verification = + ggml::hrx::verify_command_program(missing_allocation, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(!verification.valid()); + REQUIRE(status_contains(verification.status, "no transient allocation")); + + ggml::hrx::CommandProgram out_of_range = copy_command_program_shape(commands); + out_of_range.commands[0].bindings[2].length = sum_value->byte_count + 1; + verification = ggml::hrx::verify_command_program(out_of_range, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(!verification.valid()); + REQUIRE(status_contains(verification.status, "outside transient allocation length")); + + ggml_free(ctx); +} + +static void run_graph_replay_host_staging_is_not_ineligible() { + ggml::hrx::CommandProgram commands; + + std::vector source0(64, 1); + std::vector source1(64, 2); + + ggml::hrx::PreparedCommandProgram prepared; + ggml::hrx::HostStagingBuffer staging; + staging.buffer = dummy_hrx_buffer(0x1000); + staging.host_data = source0.data(); + staging.value = 7; + staging.length = source0.size(); + staging.upload = true; + prepared.host_staging.push_back(std::move(staging)); + + ggml::hrx::CommandProgramBinding live_binding; + live_binding.value = ggml::hrx::ValueId(7); + live_binding.length = source1.size(); + live_binding.capacity = source1.size(); + live_binding.host_data = source1.data(); + const ggml::hrx::CommandProgramBindings bindings = + ggml::hrx::CommandProgramBindings::from_bindings({ live_binding }); + REQUIRE(bindings.valid()); + + ggml::hrx::RecordedCommandGraph recorded; + recorded.exec = dummy_hrx_graph_exec(0x2000); + recorded.bound_transient_arena_allocation_id = ggml::hrx::kInvalidTransientArenaAllocationId; + + ggml::hrx::CommandProgramExecutionContext context; + context.stream = dummy_hrx_stream(0x3000); + + const ggml::hrx::RecordedCommandGraphExecutionResult result = + ggml::hrx::bind_and_launch_recorded_command_graph(context, commands, bindings, prepared, recorded); + recorded.exec = nullptr; + + REQUIRE(!result.success); + REQUIRE(result.event == ggml::hrx::HrxGraphReplayEvent::LaunchFailed); + REQUIRE(result.ineligible_reason.empty()); + REQUIRE(status_contains(result.status, "missing HRX host transfer manager")); + REQUIRE(prepared.host_staging.size() == 1); + REQUIRE(prepared.host_staging[0].buffer == dummy_hrx_buffer(0x1000)); + REQUIRE(prepared.host_staging[0].host_data == source1.data()); + prepared.host_staging[0].buffer = nullptr; +} + +static void run_multiple_transient_plan_checks() { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * c = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * d = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * sum0 = ggml_add(ctx, a, b); + ggml_tensor * sum1 = ggml_add(ctx, c, d); + ggml_tensor * out = ggml_add(ctx, sum0, sum1); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(c != nullptr); + REQUIRE(d != nullptr); + REQUIRE(sum0 != nullptr); + REQUIRE(sum1 != nullptr); + REQUIRE(out != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 3); + + const ggml::hrx::Value * sum0_value = imported.graph.values().find_tensor(sum0); + const ggml::hrx::Value * sum1_value = imported.graph.values().find_tensor(sum1); + REQUIRE(sum0_value != nullptr); + REQUIRE(sum1_value != nullptr); + REQUIRE(sum0_value->kind == ggml::hrx::ValueKind::Transient); + REQUIRE(sum1_value->kind == ggml::hrx::ValueKind::Transient); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 3); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.transients.allocations.size() == 2); + REQUIRE(commands.transients.arena_size == 512); + const ggml::hrx::TransientAllocation * sum0_allocation = + ggml::hrx::find_transient_allocation(commands.transients, sum0_value->id); + const ggml::hrx::TransientAllocation * sum1_allocation = + ggml::hrx::find_transient_allocation(commands.transients, sum1_value->id); + REQUIRE(sum0_allocation != nullptr); + REQUIRE(sum1_allocation != nullptr); + REQUIRE(sum0_allocation->arena_offset != sum1_allocation->arena_offset); + REQUIRE(sum0_allocation->arena_offset % 256 == 0); + REQUIRE(sum1_allocation->arena_offset % 256 == 0); + + ggml::hrx::CommandProgram overlapping_live_transients = copy_command_program_shape(commands); + ggml::hrx::TransientAllocation * overlapping_sum0 = nullptr; + ggml::hrx::TransientAllocation * overlapping_sum1 = nullptr; + for (ggml::hrx::TransientAllocation & allocation : overlapping_live_transients.transients.allocations) { + if (allocation.value == sum0_value->id) { + overlapping_sum0 = &allocation; + } else if (allocation.value == sum1_value->id) { + overlapping_sum1 = &allocation; + } + } + REQUIRE(overlapping_sum0 != nullptr); + REQUIRE(overlapping_sum1 != nullptr); + overlapping_sum1->arena_offset = overlapping_sum0->arena_offset; + ggml::hrx::VerificationResult verification = + ggml::hrx::verify_command_program(overlapping_live_transients, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(!verification.valid()); + REQUIRE(status_contains(verification.status, "transient allocations overlap")); + + ggml_free(ctx); +} + +static void run_disjoint_transient_plan_packing_checks() { + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * c = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * d = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * e = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * f = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * sum0 = ggml_add(ctx, a, b); + ggml_tensor * out0 = ggml_add(ctx, sum0, c); + ggml_tensor * sum1 = ggml_add(ctx, d, e); + ggml_tensor * out1 = ggml_add(ctx, sum1, f); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(c != nullptr); + REQUIRE(d != nullptr); + REQUIRE(e != nullptr); + REQUIRE(f != nullptr); + REQUIRE(sum0 != nullptr); + REQUIRE(out0 != nullptr); + REQUIRE(sum1 != nullptr); + REQUIRE(out1 != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out0); + ggml_build_forward_expand(graph, out1); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 4); + + const ggml::hrx::Value * sum0_value = imported.graph.values().find_tensor(sum0); + const ggml::hrx::Value * sum1_value = imported.graph.values().find_tensor(sum1); + REQUIRE(sum0_value != nullptr); + REQUIRE(sum1_value != nullptr); + REQUIRE(sum0_value->kind == ggml::hrx::ValueKind::Transient); + REQUIRE(sum1_value->kind == ggml::hrx::ValueKind::Transient); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + REQUIRE(scheduler.plan().valid()); + REQUIRE(scheduler.plan().dispatches.size() == 4); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + REQUIRE(commands.transients.allocations.size() == 2); + REQUIRE(commands.transients.arena_size == 256); + const ggml::hrx::TransientAllocation * sum0_allocation = + ggml::hrx::find_transient_allocation(commands.transients, sum0_value->id); + const ggml::hrx::TransientAllocation * sum1_allocation = + ggml::hrx::find_transient_allocation(commands.transients, sum1_value->id); + REQUIRE(sum0_allocation != nullptr); + REQUIRE(sum1_allocation != nullptr); + REQUIRE(sum0_allocation->arena_offset == sum1_allocation->arena_offset); + REQUIRE(sum0_allocation->arena_offset % 256 == 0); + REQUIRE(command_program_verifies(commands)); + + ggml_free(ctx); +} + +static void run_command_shape_hash_checks() { + REQUIRE(ggml::hrx::command_program_shape_hash("") == UINT64_C(1469598103934665603)); + std::string bytes; + for (int i = 0; i < 256; ++i) { + bytes.push_back(static_cast(i)); + } + bytes.append("\0key\xff", 5); + REQUIRE(ggml::hrx::command_program_shape_hash(bytes) == UINT64_C(6643458358697029213)); + ggml::hrx::GraphProgram program(1, "gfx1151", nullptr, nullptr, bytes); + bytes.clear(); + REQUIRE(program.command_shape_hash() == UINT64_C(6643458358697029213)); + REQUIRE(program.command_shape_hash() == ggml::hrx::command_program_shape_hash(program.command_shape())); +} + +static void run_graph_program_cache_uid_mismatch_checks() { + ggml::hrx::GraphProgramCache cache; + const ggml::hrx::KernelCorpus & corpus = ggml::hrx::get_qwen_kernel_corpus(); + + ggml_init_params params = {}; + params.mem_size = 512 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * c = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * d = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * out0 = ggml_add(ctx, a, b); + ggml_tensor * out1 = ggml_add(ctx, c, d); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(c != nullptr); + REQUIRE(d != nullptr); + REQUIRE(out0 != nullptr); + REQUIRE(out1 != nullptr); + + ggml_cgraph * graph0 = ggml_new_graph(ctx); + REQUIRE(graph0 != nullptr); + ggml_build_forward_expand(graph0, out0); + graph0->uid = 3001; + + ggml::hrx::GraphProgramLookup lookup = cache.get_or_build(*graph0, corpus, "gfx1151"); + REQUIRE(lookup.valid()); + const uint64_t first_shape_hash = lookup.program->command_shape_hash(); + REQUIRE(first_shape_hash == ggml::hrx::command_program_shape_hash(lookup.program->command_shape())); + ggml::hrx::GraphProgramCacheStats stats = cache.stats(); + REQUIRE(stats.builds == 1); + REQUIRE(stats.hits == 0); + + lookup = cache.get_or_build(*graph0, corpus, "gfx1151"); + REQUIRE(lookup.valid()); + stats = cache.stats(); + REQUIRE(stats.builds == 1); + REQUIRE(stats.hits == 1); + REQUIRE(lookup.program->command_shape_hash() == first_shape_hash); + + ggml_cgraph * graph1 = ggml_new_graph(ctx); + REQUIRE(graph1 != nullptr); + ggml_build_forward_expand(graph1, out0); + ggml_build_forward_expand(graph1, out1); + graph1->uid = 3001; + + lookup = cache.get_or_build(*graph1, corpus, "gfx1151"); + REQUIRE(lookup.valid()); + stats = cache.stats(); + REQUIRE(stats.builds == 2); + REQUIRE(stats.hits == 1); + REQUIRE(lookup.program->command_shape_hash() != first_shape_hash); + REQUIRE(lookup.program->command_shape_hash() == + ggml::hrx::command_program_shape_hash(lookup.program->command_shape())); + + ggml_tensor * unsupported = ggml_sin(ctx, a); + REQUIRE(unsupported != nullptr); + ggml_cgraph * graph2 = ggml_new_graph(ctx); + REQUIRE(graph2 != nullptr); + ggml_build_forward_expand(graph2, unsupported); + graph2->uid = 3001; + + lookup = cache.get_or_build(*graph2, corpus, "gfx1151"); + REQUIRE(!lookup.valid()); + stats = cache.stats(); + REQUIRE(stats.builds == 2); + REQUIRE(stats.hits == 1); + + ggml::hrx::GraphProgramCache alias_cache; + ggml_tensor * alias_sum = ggml_add(ctx, a, b); + ggml_tensor * view0 = ggml_view_1d(ctx, alias_sum, 4, 0); + ggml_tensor * view1 = ggml_view_1d(ctx, alias_sum, 4, sizeof(float)); + REQUIRE(alias_sum != nullptr); + REQUIRE(view0 != nullptr); + REQUIRE(view1 != nullptr); + + ggml_cgraph * alias_graph0 = ggml_new_graph(ctx); + REQUIRE(alias_graph0 != nullptr); + ggml_build_forward_expand(alias_graph0, view0); + alias_graph0->uid = 3002; + lookup = alias_cache.get_or_build(*alias_graph0, corpus, "gfx1151"); + REQUIRE(lookup.valid()); + stats = alias_cache.stats(); + REQUIRE(stats.builds == 1); + REQUIRE(stats.hits == 0); + + ggml_cgraph * alias_graph1 = ggml_new_graph(ctx); + REQUIRE(alias_graph1 != nullptr); + ggml_build_forward_expand(alias_graph1, view1); + alias_graph1->uid = 3002; + lookup = alias_cache.get_or_build(*alias_graph1, corpus, "gfx1151"); + REQUIRE(lookup.valid()); + stats = alias_cache.stats(); + REQUIRE(stats.builds == 1); + REQUIRE(stats.hits == 1); + + ggml_free(ctx); +} + +static void run_graph_match_hash_collision_checks() { + ggml_context * ctx = ggml_init({ 2 * 1024 * 1024, nullptr, true }); + REQUIRE(ctx != nullptr); + const size_t buckets = ggml_hash_size(10); + std::vector seen(buckets, nullptr); + ggml_tensor * a = nullptr; + ggml_tensor * b = nullptr; + for (size_t i = 0; i <= buckets && b == nullptr; ++i) { + ggml_tensor * tensor = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + const size_t bucket = ggml_hash(tensor) % buckets; + if (seen[bucket] != nullptr) { + a = seen[bucket]; + b = tensor; + } else { + seen[bucket] = tensor; + } + } + REQUIRE(a != nullptr && b != nullptr && a != b); + ggml_tensor * c = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * sum = ggml_add(ctx, a, b); + ggml_tensor * out = ggml_add(ctx, sum, c); + ggml_cgraph * graph = ggml_new_graph_custom(ctx, 256, false); + ggml_build_forward_expand(graph, out); + ggml::hrx::GraphProgramCache cache; + auto lookup = cache.get_or_build(*graph, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(lookup.valid()); + REQUIRE(lookup.program->graph().values().size() == 5); + const auto match = lookup.program->match_current_graph(*graph); + REQUIRE(match.valid()); + for (const ggml_tensor * tensor : { a, b, c, out }) { + REQUIRE(std::any_of(match.external_bindings.begin(), match.external_bindings.end(), + [&](const auto & binding) { return binding.tensor == tensor; })); + } + for (int i = 0; i < 128; ++i) { + out = ggml_add(ctx, out, i % 2 == 0 ? a : b); + } + ggml_build_forward_expand(graph, out); + lookup = cache.get_or_build(*graph, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(lookup.valid()); + REQUIRE(lookup.program->graph().values().size() == 133); + REQUIRE(lookup.program->match_current_graph(*graph).valid()); + ggml_free(ctx); +} + +static void run_graph_match_bijection_checks() { + ggml_context * ctx = ggml_init({ 512 * 1024, nullptr, true }); + REQUIRE(ctx != nullptr); + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * c = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * sum = ggml_add(ctx, a, b); + ggml_tensor * out = ggml_add(ctx, sum, b); + ggml_cgraph * graph = ggml_new_graph_custom(ctx, 16, false); + ggml_build_forward_expand(graph, out); + ggml::hrx::GraphProgramCache cache; + auto lookup = cache.get_or_build(*graph, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(lookup.valid()); + REQUIRE(lookup.program->match_current_graph(*graph).valid()); + + sum->src[1] = a; + auto duplicate_tensor = lookup.program->match_current_graph(*graph); + REQUIRE(!duplicate_tensor.valid()); + REQUIRE(status_contains(duplicate_tensor.status, "tensor maps to cached values")); + sum->src[1] = b; + out->src[1] = c; + auto duplicate_value = lookup.program->match_current_graph(*graph); + REQUIRE(!duplicate_value.valid()); + REQUIRE(status_contains(duplicate_value.status, "maps to multiple current tensors")); + out->src[1] = b; + REQUIRE(lookup.program->match_current_graph(*graph).valid()); + ggml_free(ctx); +} + +static void run_validated_graph_uid_match_checks() { + ggml::hrx::GraphProgramCache cache; + const auto & corpus = ggml::hrx::get_qwen_kernel_corpus(); + ggml_context * ctx = ggml_init({ 2 * 1024 * 1024, nullptr, true }); + REQUIRE(ctx != nullptr); + auto make_graph = [&](uint64_t uid, int64_t elements, bool two_nodes) { + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, elements); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, elements); + ggml_tensor * out = ggml_add(ctx, a, b); + if (two_nodes) { + out = ggml_add(ctx, out, b); + } + ggml_cgraph * graph = ggml_new_graph_custom(ctx, 16, false); + ggml_build_forward_expand(graph, out); + graph->uid = uid; + return std::make_pair(graph, a); + }; + auto first = make_graph(4001, 8, false); + auto alias = make_graph(4002, 8, false); + auto original = cache.get_or_build(*first.first, corpus, "gfx1151"); + REQUIRE(original.valid()); + for (int i = 0; i < 2; ++i) { + auto match = cache.get_or_build(*alias.first, corpus, "gfx1151"); + REQUIRE(match.valid()); + REQUIRE(match.program == original.program); + REQUIRE(std::any_of(match.match.external_bindings.begin(), match.match.external_bindings.end(), + [&](const auto & binding) { return binding.tensor == alias.second; })); + } + REQUIRE(cache.stats().builds == 1); + + auto new_nodes = make_graph(4002, 16, false); + auto changed = cache.get_or_build(*new_nodes.first, corpus, "gfx1151"); + REQUIRE(changed.valid()); + REQUIRE(changed.program->uid() == 4002); + REQUIRE(cache.stats().builds == 2); + + auto second_alias = make_graph(4003, 8, false); + REQUIRE(cache.get_or_build(*second_alias.first, corpus, "gfx1151").program == original.program); + auto replacement = make_graph(4001, 8, true); + REQUIRE(cache.get_or_build(*replacement.first, corpus, "gfx1151").valid()); + REQUIRE(cache.stats().builds == 3); + auto after_replacement = cache.get_or_build(*second_alias.first, corpus, "gfx1151"); + REQUIRE(after_replacement.valid()); + REQUIRE(after_replacement.program->uid() == 4003); + REQUIRE(cache.stats().builds == 4); + + auto checked_alias = make_graph(4004, 8, false); + REQUIRE(cache.get_or_build(*checked_alias.first, corpus, "gfx1151").valid()); + const char * validate_name = "GGML_HRX_VALIDATE_GRAPH_UID_CACHE"; + const char * validate_value = std::getenv(validate_name); + const bool had_validate = validate_value != nullptr; + const std::string saved_validate = had_validate ? validate_value : ""; + REQUIRE(setenv(validate_name, "1", 1) == 0); + checked_alias.second->ne[0] = 12; + REQUIRE(!cache.get_or_build(*checked_alias.first, corpus, "gfx1151").valid()); + checked_alias.second->ne[0] = 8; + REQUIRE(cache.get_or_build(*checked_alias.first, corpus, "gfx1151").valid()); + restore_environment_value(validate_name, had_validate, saved_validate); + + for (uint64_t uid = 4100; uid < 4230; ++uid) { + checked_alias.first->uid = uid; + auto match = cache.get_or_build(*checked_alias.first, corpus, "gfx1151"); + REQUIRE(match.valid()); + REQUIRE(match.program->uid() == 4003); + } + REQUIRE(cache.stats().builds == 4); + cache.clear(); + auto after_clear = cache.get_or_build(*second_alias.first, corpus, "gfx1151"); + REQUIRE(after_clear.valid()); + REQUIRE(after_clear.program->uid() == 4003); + REQUIRE(cache.stats().builds == 5); + ggml_free(ctx); +} + +static void run_graph_executor_contract_checks() { + ggml_backend_hrx_device_context device_context = {}; + ggml_backend_hrx_context backend_context = {}; + device_context.architecture = "gfx1151"; + backend_context.device = &device_context; + const ggml::hrx::GraphExecutor executor(backend_context); + + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * add_out = ggml_add(ctx, a, b); + ggml_tensor * sin_out = ggml_sin(ctx, a); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(add_out != nullptr); + REQUIRE(sin_out != nullptr); + + ggml_cgraph * add_graph = ggml_new_graph(ctx); + REQUIRE(add_graph != nullptr); + ggml_build_forward_expand(add_graph, add_out); + const ggml::hrx::GraphSupportResult add_support = executor.can_execute(*add_graph); + REQUIRE(add_support.supported); + REQUIRE(add_support.status.success()); + + const ggml::hrx::GraphExecutionResult missing_binding = executor.execute(*add_graph); + REQUIRE(!missing_binding.success()); + REQUIRE(missing_binding.code == GGML_STATUS_FAILED); + REQUIRE(status_contains(missing_binding.status, "external value")); + REQUIRE(status_contains(missing_binding.status, "not bound")); + + ggml_cgraph * sin_graph = ggml_new_graph(ctx); + REQUIRE(sin_graph != nullptr); + ggml_build_forward_expand(sin_graph, sin_out); + const ggml::hrx::GraphSupportResult sin_support = executor.can_execute(*sin_graph); + REQUIRE(!sin_support.supported); + REQUIRE(status_contains(sin_support.status, "unsupported HRX node 0")); + REQUIRE(status_contains(sin_support.status, "SIN")); + + ggml_free(ctx); +} + +static void build_add_graph_program(ggml::hrx::GraphProgramCache & cache, uint64_t uid, bool two_outputs) { + ggml_init_params params = {}; + params.mem_size = 512 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * out0 = ggml_add(ctx, a, b); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(out0 != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out0); + if (two_outputs) { + ggml_tensor * c = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * d = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * out1 = ggml_add(ctx, c, d); + REQUIRE(c != nullptr); + REQUIRE(d != nullptr); + REQUIRE(out1 != nullptr); + ggml_build_forward_expand(graph, out1); + } + graph->uid = uid; + + const ggml::hrx::KernelCorpus & corpus = ggml::hrx::get_qwen_kernel_corpus(); + ggml::hrx::GraphProgramLookup lookup = cache.get_or_build(*graph, corpus, "gfx1151"); + REQUIRE(lookup.valid()); + ggml_free(ctx); +} + +static void run_command_program_kernel_dump_checks() { + static constexpr const char * kDumpEnv = "GGML_HRX_DUMP_COMMAND_PROGRAM_DIR"; + EnvironmentVariableGuard env(kDumpEnv); + + { + const std::filesystem::path dump_dir = fresh_test_directory("kernel-dump-disabled"); + env.unset(); + ggml::hrx::GraphProgramCache cache; + build_add_graph_program(cache, 5001, false); + REQUIRE(list_directories(dump_dir).empty()); + std::filesystem::remove_all(dump_dir); + } + + const std::filesystem::path dump_dir = fresh_test_directory("kernel-dump-enabled"); + env.set(dump_dir); + + ggml::hrx::GraphProgramCache cache; + build_add_graph_program(cache, 5002, false); + std::vector dumps = list_directories(dump_dir); + REQUIRE(dumps.size() == 1); + REQUIRE(std::filesystem::exists(dumps[0] / "kernels.txt")); + REQUIRE(std::filesystem::exists(dumps[0] / "kernels.dot")); + + const std::string add_kernels = read_text_file(dumps[0] / "kernels.txt"); + const std::string add_dot = read_text_file(dumps[0] / "kernels.dot"); + REQUIRE(string_contains(add_kernels, "main 0 loom_libs:ggml_binary_f32")); + REQUIRE(!string_contains(add_kernels, "binding")); + REQUIRE(!string_contains(add_kernels, "transient")); + REQUIRE(string_contains(add_dot, "digraph hrx_kernel_invocations")); + REQUIRE(string_contains(add_dot, "main 0\\nloom_libs:ggml_binary_f32")); + + build_add_graph_program(cache, 5003, false); + dumps = list_directories(dump_dir); + REQUIRE(dumps.size() == 1); + + const std::filesystem::path second_dump_dir = fresh_test_directory("kernel-dump-second-directory"); + env.set(second_dump_dir); + build_add_graph_program(cache, 5005, false); + std::vector second_dumps = list_directories(second_dump_dir); + REQUIRE(second_dumps.size() == 1); + REQUIRE(std::filesystem::exists(second_dumps[0] / "kernels.txt")); + std::filesystem::remove_all(second_dump_dir); + env.set(dump_dir); + + build_add_graph_program(cache, 5004, true); + dumps = list_directories(dump_dir); + REQUIRE(dumps.size() == 2); + + const std::string two_add_kernels = read_text_file(dumps[1] / "kernels.txt"); + const std::string two_add_dot = read_text_file(dumps[1] / "kernels.dot"); + REQUIRE(string_contains(two_add_kernels, "main 0 loom_libs:ggml_binary_f32")); + REQUIRE(string_contains(two_add_kernels, "main 1 loom_libs:ggml_binary_f32 deps=0")); + REQUIRE(string_contains(two_add_dot, "main0 -> main1")); + REQUIRE(!std::filesystem::exists(dumps[1] / "commands.txt")); + REQUIRE(!std::filesystem::exists(dumps[1] / "commands.json")); + REQUIRE(!std::filesystem::exists(dumps[1] / "manifest.json")); + + std::filesystem::remove_all(dump_dir); +} + +static std::vector read_transient_i32(ggml_backend_hrx_context * context, + ggml::hrx::TransientArenaAllocationRef arena, + const ggml::hrx::TransientAllocation & allocation) { + REQUIRE(context != nullptr); + REQUIRE(arena.buffer != nullptr); + REQUIRE(allocation.size % sizeof(int32_t) == 0); + std::vector data(allocation.size / sizeof(int32_t)); + require_hrx_status(hrx_synchronous_d2h(context->device->device, arena.buffer, allocation.arena_offset, data.data(), + allocation.size)); + return data; +} + +static void write_transient_i32_value(ggml_backend_hrx_context * context, + ggml::hrx::TransientArenaAllocationRef arena, + const ggml::hrx::TransientAllocation & allocation, + int32_t value) { + REQUIRE(context != nullptr); + REQUIRE(arena.buffer != nullptr); + REQUIRE(allocation.size >= sizeof(value)); + require_hrx_status( + hrx_synchronous_h2d(context->device->device, &value, arena.buffer, allocation.arena_offset, sizeof(value))); +} + +static std::vector read_transient_bytes(ggml_backend_hrx_context * context, + ggml::hrx::TransientArenaAllocationRef arena, + const ggml::hrx::TransientAllocation & allocation) { + REQUIRE(context != nullptr); + REQUIRE(arena.buffer != nullptr); + std::vector data(allocation.size); + require_hrx_status(hrx_synchronous_d2h(context->device->device, arena.buffer, allocation.arena_offset, data.data(), + allocation.size)); + return data; +} + +static void run_binary_q8_publication_numerics() { + const struct { + int64_t hidden; + int64_t tokens; + bool broadcast; + int op; + } cases[] = { + { 2048, 3, false, 0 }, + { 2176, 4, false, 2 }, + { 2048, 4, true, 0 }, + { 2176, 5, true, 2 }, + }; + + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + auto * backend_context = static_cast(backend->context); + REQUIRE(backend_context != nullptr); + const ggml::hrx::KernelCorpus & corpus = ggml::hrx::get_qwen_kernel_corpus(); + const char * target = backend_context->device->architecture.c_str(); + + for (const auto & test : cases) { + ggml_init_params params = {}; + params.mem_size = 512 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + const int64_t element_count = test.hidden * test.tokens; + const int64_t rhs_count = test.broadcast ? test.hidden : element_count; + ggml_tensor * lhs = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, test.hidden, test.tokens); + ggml_tensor * rhs = test.broadcast ? ggml_new_tensor_1d(ctx, GGML_TYPE_F32, test.hidden) : + ggml_new_tensor_2d(ctx, GGML_TYPE_F32, test.hidden, test.tokens); + REQUIRE(lhs != nullptr && rhs != nullptr); + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + + std::vector lhs_data(element_count); + std::vector rhs_data(rhs_count); + std::vector expected(element_count); + for (int64_t i = 0; i < element_count; ++i) { + lhs_data[i] = static_cast((i * 17 + 3) % 251 - 125) / 31.0f; + } + for (int64_t i = 0; i < rhs_count; ++i) { + rhs_data[i] = static_cast((i * 29 + 11) % 239 - 119) / 37.0f; + } + for (int64_t token = 0; token < test.tokens; ++token) { + for (int64_t channel = 0; channel < test.hidden; ++channel) { + const int64_t i = token * test.hidden + channel; + const float rhs_v = rhs_data[test.broadcast ? channel : i]; + expected[i] = test.op == 0 ? lhs_data[i] + rhs_v : lhs_data[i] * rhs_v; + } + } + ggml_backend_tensor_set(lhs, lhs_data.data(), 0, lhs_data.size() * sizeof(float)); + ggml_backend_tensor_set(rhs, rhs_data.data(), 0, rhs_data.size() * sizeof(float)); + ggml_backend_synchronize(backend); + + ggml::hrx::Graph graph; + const ggml::hrx::ValueId lhs_value = + graph.values().get_or_add_tensor_value(lhs, ggml::hrx::ValueKind::External); + const ggml::hrx::ValueId rhs_value = + graph.values().get_or_add_tensor_value(rhs, ggml::hrx::ValueKind::External); + ggml::hrx::ValueBufferBinding lhs_binding; + ggml::hrx::ValueBufferBinding rhs_binding; + REQUIRE(ggml_backend_hrx_resolve_value_buffer(lhs, lhs_binding)); + REQUIRE(ggml_backend_hrx_resolve_value_buffer(rhs, rhs_binding)); + REQUIRE(graph.values().bind_buffer(lhs_value, lhs_binding)); + REQUIRE(graph.values().bind_buffer(rhs_value, rhs_binding)); + + const ggml::hrx::ValueId reference_f32(static_cast(graph.values().size())); + const ggml::hrx::ValueId reference_q8(reference_f32.value + 1); + const ggml::hrx::ValueId published_f32(reference_f32.value + 2); + const ggml::hrx::ValueId published_q8(reference_f32.value + 3); + const size_t f32_bytes = static_cast(element_count) * sizeof(float); + const size_t q8_bytes = qwen_q8_1_x4_size(test.tokens, test.hidden); + + ggml::hrx::CommandPlan plan; + plan.transients.push_back({ reference_f32, "test.binary.reference_f32", f32_bytes, 256 }); + plan.transients.push_back({ reference_q8, "test.binary.reference_q8", q8_bytes, 256 }); + plan.transients.push_back({ published_f32, "test.binary.published_f32", f32_bytes, 256 }); + plan.transients.push_back({ published_q8, "test.binary.published_q8", q8_bytes, 256 }); + + ggml::hrx::Dispatch legacy; + legacy.kernel = ggml::hrx::make_kernel_specialization(ggml::hrx::kernel_catalog_ref( + "loom_libs", test.broadcast ? "ggml_binary_bc_f32" : "ggml_binary_f32")); + legacy.kernel.integer_parameters.emplace("element_count", element_count); + legacy.kernel.compile_parameters.emplace( + test.broadcast ? "ggml.binary_bc_f32.op" : "ggml.binary_f32.op", std::to_string(test.op)); + if (test.broadcast) { + legacy.kernel.integer_parameters.emplace("ne0", test.hidden); + legacy.kernel.integer_parameters.emplace("ne1", test.tokens); + legacy.kernel.integer_parameters.emplace("ne2", 1); + legacy.kernel.integer_parameters.emplace("ne3", 1); + legacy.kernel.integer_parameters.emplace("src0_element_count", element_count); + legacy.kernel.integer_parameters.emplace("src1_element_count", rhs_count); + } else { + legacy.kernel.compile_parameters.emplace("ggml.binary_f32.ne0", std::to_string(test.hidden)); + legacy.kernel.compile_parameters.emplace("ggml.binary_f32.ne1", std::to_string(test.tokens)); + legacy.kernel.compile_parameters.emplace("ggml.binary_f32.ne2", "1"); + legacy.kernel.compile_parameters.emplace("ggml.binary_f32.src0_stride1", std::to_string(test.hidden)); + legacy.kernel.compile_parameters.emplace("ggml.binary_f32.src0_stride2", std::to_string(element_count)); + legacy.kernel.compile_parameters.emplace("ggml.binary_f32.src0_stride3", std::to_string(element_count)); + legacy.kernel.compile_parameters.emplace("ggml.binary_f32.src1_stride1", std::to_string(test.hidden)); + legacy.kernel.compile_parameters.emplace("ggml.binary_f32.src1_stride2", std::to_string(element_count)); + legacy.kernel.compile_parameters.emplace("ggml.binary_f32.src1_stride3", std::to_string(element_count)); + legacy.kernel.compile_parameters.emplace("ggml.binary_f32.src0_span", std::to_string(element_count)); + legacy.kernel.compile_parameters.emplace("ggml.binary_f32.src1_span", std::to_string(element_count)); + } + if (test.broadcast) { + for (int dim = 0; dim < GGML_MAX_DIMS; ++dim) { + legacy.kernel.compile_parameters.emplace("ggml.binary_bc_f32.src0_broadcast_dim" + + std::to_string(dim), + "0"); + legacy.kernel.compile_parameters.emplace("ggml.binary_bc_f32.src1_broadcast_dim" + + std::to_string(dim), + dim == 1 ? "1" : "0"); + } + } + legacy.bindings.push_back({ lhs_value, 0, ggml_nbytes(lhs) }); + legacy.bindings.push_back({ rhs_value, 0, ggml_nbytes(rhs) }); + legacy.bindings.push_back({ reference_f32, 0, f32_bytes }); + plan.dispatches.push_back(std::move(legacy)); + + ggml::hrx::Dispatch quantize; + quantize.kernel = ggml::hrx::make_kernel_specialization( + ggml::hrx::kernel_catalog_ref("qwen3_moe", "ggml_quantize_q8_1_x4_f32")); + quantize.kernel.integer_parameters.emplace("token_count", test.tokens); + quantize.kernel.integer_parameters.emplace("input_size", test.hidden); + quantize.kernel.compile_parameters.emplace("ggml.quantize_q8_1_x4.group_capacity", + std::to_string(element_count / 128)); + quantize.bindings.push_back({ reference_f32, 0, f32_bytes }); + quantize.bindings.push_back({ reference_q8, 0, q8_bytes }); + plan.dispatches.push_back(std::move(quantize)); + + ggml::hrx::Dispatch publish; + publish.kernel = ggml::hrx::make_kernel_specialization(ggml::hrx::kernel_catalog_ref( + "loom_libs", test.broadcast ? "ggml_binary_bc_f32_publish_q8_1_x4" : + "ggml_binary_f32_publish_q8_1_x4")); + publish.kernel.integer_parameters.emplace("token_count", test.tokens); + if (test.broadcast) { + publish.kernel.integer_parameters.emplace("hidden_size", test.hidden); + publish.kernel.integer_parameters.emplace("src0_element_count", element_count); + publish.kernel.integer_parameters.emplace("src1_element_count", rhs_count); + publish.kernel.compile_parameters = plan.dispatches.front().kernel.compile_parameters; + } else { + publish.kernel.compile_parameters = plan.dispatches.front().kernel.compile_parameters; + } + publish.bindings.push_back({ lhs_value, 0, ggml_nbytes(lhs) }); + publish.bindings.push_back({ rhs_value, 0, ggml_nbytes(rhs) }); + publish.bindings.push_back({ published_f32, 0, f32_bytes }); + publish.bindings.push_back({ published_q8, 0, q8_bytes }); + plan.dispatches.push_back(std::move(publish)); + + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program(graph, plan, corpus, target); + REQUIRE(commands.valid()); + REQUIRE(ggml::hrx::verify_command_program(commands, corpus, target).valid()); + const ggml::hrx::CommandProgramBindings bindings = + ggml::hrx::CommandProgramBindings::from_value_map(graph.values()); + REQUIRE(bindings.valid()); + const ggml::hrx::CommandProgramExecutionContext execution_context = { + backend_context->device->device, + backend_context->stream, + target, + &corpus, + &backend_context->kernel_executables, + &backend_context->transient_arena, + &backend_context->host_transfers, + &backend_context->host_weights, + }; + REQUIRE(ggml::hrx::execute_command_program(execution_context, commands, bindings)); + ggml_backend_synchronize(backend); + + const auto arena = backend_context->transient_arena.current_allocation(); + const auto * reference_f32_allocation = + ggml::hrx::find_transient_allocation(commands.transients, reference_f32); + const auto * reference_q8_allocation = + ggml::hrx::find_transient_allocation(commands.transients, reference_q8); + const auto * published_f32_allocation = + ggml::hrx::find_transient_allocation(commands.transients, published_f32); + const auto * published_q8_allocation = + ggml::hrx::find_transient_allocation(commands.transients, published_q8); + REQUIRE(reference_f32_allocation != nullptr && reference_q8_allocation != nullptr); + REQUIRE(published_f32_allocation != nullptr && published_q8_allocation != nullptr); + const auto reference_f32_bytes = read_transient_bytes(backend_context, arena, *reference_f32_allocation); + const auto reference_q8_bytes = read_transient_bytes(backend_context, arena, *reference_q8_allocation); + const auto published_f32_bytes = read_transient_bytes(backend_context, arena, *published_f32_allocation); + const auto published_q8_bytes = read_transient_bytes(backend_context, arena, *published_q8_allocation); + REQUIRE(reference_f32_bytes == published_f32_bytes); + REQUIRE(reference_q8_bytes == published_q8_bytes); + REQUIRE(reference_f32_bytes.size() == expected.size() * sizeof(float)); + REQUIRE(std::memcmp(reference_f32_bytes.data(), expected.data(), reference_f32_bytes.size()) == 0); + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + } + ggml_backend_free(backend); +} + +static void run_qwen_expert_table_partition_prefill_512_execution() { + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + auto * backend_context = static_cast(backend->context); + REQUIRE(backend_context != nullptr); + + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t token_count = 512; + constexpr int64_t route_count = 8; + constexpr int64_t route_stride = 8; + constexpr int64_t expert_count = 128; + + ggml_tensor * route_ids_tensor = ggml_new_tensor_2d(ctx, GGML_TYPE_I32, route_stride, token_count); + REQUIRE(route_ids_tensor != nullptr); + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + + const std::vector route_ids = + make_qwen_route_ids_iota(token_count, route_count, route_stride, expert_count); + const std::vector expected_expert_table = + make_qwen_expert_table_reference(route_ids, token_count, route_count, route_stride, expert_count); + const std::vector expected_partition_table = + make_qwen_partition_table_reference(expected_expert_table, token_count, route_count, expert_count); + ggml_backend_tensor_set(route_ids_tensor, route_ids.data(), 0, route_ids.size() * sizeof(int32_t)); + ggml_backend_synchronize(backend); + + ggml::hrx::Graph graph; + const ggml::hrx::ValueId route_ids_value = + graph.values().get_or_add_tensor_value(route_ids_tensor, ggml::hrx::ValueKind::External); + ggml::hrx::ValueBufferBinding route_ids_binding; + REQUIRE(ggml_backend_hrx_resolve_value_buffer(route_ids_tensor, route_ids_binding)); + REQUIRE(graph.values().bind_buffer(route_ids_value, route_ids_binding)); + + const ggml::hrx::ValueId expert_table_value(static_cast(graph.values().size())); + const ggml::hrx::ValueId partition_table_value(expert_table_value.value + 1); + const ggml::hrx::ValueId completion_counter_value(expert_table_value.value + 2); + ggml::hrx::CommandPlan plan; + plan.transients.push_back( + { expert_table_value, "qwen.test.expert_table", qwen_expert_table_size(token_count, expert_count), 256 }); + plan.transients.push_back({ partition_table_value, "qwen.test.partition_table", + qwen_partition_table_size(token_count, route_count, expert_count), 256 }); + plan.completion_counter_requests.push_back({ completion_counter_value, "qwen.test.completion_counter", 1 }); + ggml::hrx::Dispatch dispatch; + dispatch.kernel = ggml::hrx::make_kernel_specialization( + ggml::hrx::kernel_catalog_ref("qwen3_moe", "qwen3_moe_build_expert_table_partition_prefill_512")); + dispatch.kernel.integer_parameters.emplace("token_count", token_count); + dispatch.kernel.integer_parameters.emplace("route_count", route_count); + dispatch.kernel.integer_parameters.emplace("route_stride", route_stride); + dispatch.kernel.integer_parameters.emplace("expert_count", expert_count); + dispatch.bindings.push_back({ route_ids_value, 0, route_ids.size() * sizeof(int32_t) }); + dispatch.bindings.push_back({ expert_table_value, 0, qwen_expert_table_size(token_count, expert_count) }); + dispatch.bindings.push_back( + { partition_table_value, 0, qwen_partition_table_size(token_count, route_count, expert_count) }); + dispatch.bindings.push_back({ completion_counter_value, 0, sizeof(int32_t) }); + plan.dispatches.push_back(std::move(dispatch)); + + const ggml::hrx::KernelCorpus & corpus = ggml::hrx::get_qwen_kernel_corpus(); + const char * target = backend_context->device->architecture.c_str(); + const ggml::hrx::CommandProgram commands = ggml::hrx::build_command_program(graph, plan, corpus, target); + REQUIRE(commands.valid()); + REQUIRE(commands.commands.size() == 1); + REQUIRE(commands.completion_counters.count == 1); + REQUIRE(commands.completion_counters.byte_count == sizeof(int32_t)); + REQUIRE(commands.transients.allocations.size() == 3); + REQUIRE(ggml::hrx::verify_command_program(commands, corpus, target).valid()); + + const ggml::hrx::TransientAllocation * expert_table_allocation = + ggml::hrx::find_transient_allocation(commands.transients, expert_table_value); + const ggml::hrx::TransientAllocation * partition_table_allocation = + ggml::hrx::find_transient_allocation(commands.transients, partition_table_value); + const ggml::hrx::TransientAllocation * completion_counter_allocation = + ggml::hrx::find_transient_allocation(commands.transients, completion_counter_value); + REQUIRE(expert_table_allocation != nullptr); + REQUIRE(partition_table_allocation != nullptr); + REQUIRE(completion_counter_allocation != nullptr); + REQUIRE(completion_counter_allocation->arena_offset == commands.completion_counters.arena_offset); + + const ggml::hrx::CommandProgramBindings bindings = + ggml::hrx::CommandProgramBindings::from_value_map(graph.values()); + REQUIRE(bindings.valid()); + const ggml::hrx::CommandProgramExecutionContext execution_context = { + backend_context->device->device, + backend_context->stream, + target, + &corpus, + &backend_context->kernel_executables, + &backend_context->transient_arena, + &backend_context->host_transfers, + &backend_context->host_weights, + }; + + REQUIRE(ggml::hrx::execute_command_program(execution_context, commands, bindings)); + ggml_backend_synchronize(backend); + ggml::hrx::TransientArenaAllocationRef arena = backend_context->transient_arena.current_allocation(); + require_qwen_expert_table_matches(read_transient_i32(backend_context, arena, *expert_table_allocation), + expected_expert_table, token_count, expert_count); + require_qwen_partition_table_matches(read_transient_i32(backend_context, arena, *partition_table_allocation), + expected_partition_table); + REQUIRE(read_transient_i32(backend_context, arena, *completion_counter_allocation)[0] == 0); + + write_transient_i32_value(backend_context, arena, *completion_counter_allocation, 17); + REQUIRE(ggml::hrx::execute_command_program(execution_context, commands, bindings)); + ggml_backend_synchronize(backend); + arena = backend_context->transient_arena.current_allocation(); + require_qwen_expert_table_matches(read_transient_i32(backend_context, arena, *expert_table_allocation), + expected_expert_table, token_count, expert_count); + require_qwen_partition_table_matches(read_transient_i32(backend_context, arena, *partition_table_allocation), + expected_partition_table); + REQUIRE(read_transient_i32(backend_context, arena, *completion_counter_allocation)[0] == 0); + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(backend); +} + +static void run_add_f32() { + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t element_count = 1024; + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, element_count); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, element_count); + ggml_tensor * out = ggml_add(ctx, a, b); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(out != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out); + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + + std::vector a_data(element_count); + std::vector b_data(element_count); + std::vector expected(element_count); + for (int64_t i = 0; i < element_count; ++i) { + a_data[i] = static_cast(i % 17) * 0.25f - 2.0f; + b_data[i] = static_cast(i % 13) * -0.5f + 3.0f; + expected[i] = a_data[i] + b_data[i]; + } + + ggml_backend_tensor_set(a, a_data.data(), 0, a_data.size() * sizeof(float)); + ggml_backend_tensor_set(b, b_data.data(), 0, b_data.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + + std::vector actual(element_count); + ggml_backend_tensor_get(out, actual.data(), 0, actual.size() * sizeof(float)); + for (int64_t i = 0; i < element_count; ++i) { + REQUIRE(actual[i] == expected[i]); + } + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(backend); +} + +static std::vector run_rope_scale_numerical_case(ggml_backend_t backend, bool per_token, int mode) { + ggml_init_params params = {}; + params.mem_size = 512 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t head_size = 128; + constexpr int64_t n_dims = 96; + constexpr int64_t head_count = 2; + constexpr int64_t token_count = 3; + constexpr int64_t elements = head_size * head_count * token_count; + + ggml_tensor * input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * pos = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, token_count); + ggml_tensor * freqs = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, n_dims / 2); + ggml_tensor * rope = + ggml_rope_ext(ctx, input, pos, freqs, n_dims, mode, 0, 10000.0f, 1.0f, 0.0f, 1.19024f, 32.0f, 1.0f); + ggml_tensor * scales = per_token ? ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 1, 1, token_count) : nullptr; + ggml_tensor * output = per_token ? ggml_mul(ctx, rope, scales) : ggml_scale(ctx, rope, 0.0883883461f); + REQUIRE(input != nullptr); + REQUIRE(pos != nullptr); + REQUIRE(freqs != nullptr); + REQUIRE(rope != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + std::vector input_data(elements); + std::vector freq_data(n_dims / 2); + const std::vector positions = { 3, 7, 11 }; + const std::vector scale_data = { 0.125f, -0.25f, 0.5f }; + for (int64_t i = 0; i < elements; ++i) { + input_data[i] = static_cast((i * 17) % 113) * 0.03125f - 1.5f; + } + for (int64_t i = 0; i < n_dims / 2; ++i) { + freq_data[i] = 0.75f + static_cast(i % 7) * 0.125f; + } + ggml_backend_tensor_set(input, input_data.data(), 0, input_data.size() * sizeof(float)); + ggml_backend_tensor_set(pos, positions.data(), 0, positions.size() * sizeof(int32_t)); + ggml_backend_tensor_set(freqs, freq_data.data(), 0, freq_data.size() * sizeof(float)); + if (per_token) { + ggml_backend_tensor_set(scales, scale_data.data(), 0, scale_data.size() * sizeof(float)); + } + + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + std::vector actual(elements); + ggml_backend_tensor_get(output, actual.data(), 0, actual.size() * sizeof(float)); + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + return actual; +} + +static void run_rope_scale_f32_numerics() { + ggml_backend_t hrx = ggml_backend_hrx_init(0); + ggml_backend_t cpu = ggml_backend_cpu_init(); + REQUIRE(hrx != nullptr); + REQUIRE(cpu != nullptr); + + for (const bool per_token : { false, true }) { + for (const int mode : { static_cast(GGML_ROPE_TYPE_NORMAL), static_cast(GGML_ROPE_TYPE_NEOX) }) { + const std::vector expected = run_rope_scale_numerical_case(cpu, per_token, mode); + const std::vector actual = run_rope_scale_numerical_case(hrx, per_token, mode); + REQUIRE(actual.size() == expected.size()); + for (size_t i = 0; i < actual.size(); ++i) { + const float tolerance = 2.0e-4f + 2.0e-4f * std::fabs(expected[i]); + REQUIRE(std::fabs(actual[i] - expected[i]) <= tolerance); + } + } + } + + ggml_backend_free(cpu); + ggml_backend_free(hrx); +} + +struct GatherAddRmsNormOutputs { + std::vector raw; + std::vector normalized; +}; + +static GatherAddRmsNormOutputs run_gather_add_rmsnorm_numerical_case(ggml_backend_t backend) { + ggml_init_params params = {}; + params.mem_size = 512 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t hidden_size = 128; + constexpr int64_t source_token_count = 4; + constexpr int64_t output_token_count = 5; + GatherAddRmsNormGraph tensors = + build_gather_add_rmsnorm_graph(ctx, hidden_size, source_token_count, output_token_count); + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + + std::vector attention(hidden_size * source_token_count); + std::vector residual(hidden_size * source_token_count); + std::vector weight(hidden_size); + const std::vector row_ids = { 3, 1, 3, 0, 2 }; + for (size_t i = 0; i < attention.size(); ++i) { + attention[i] = static_cast((i * 13) % 47) * 0.03125f - 0.5f; + residual[i] = static_cast((i * 7) % 31) * -0.015625f + 0.25f; + } + for (size_t i = 0; i < weight.size(); ++i) { + weight[i] = 0.5f + static_cast(i % 11) * 0.0625f; + } + ggml_backend_tensor_set(tensors.attention, attention.data(), 0, attention.size() * sizeof(float)); + ggml_backend_tensor_set(tensors.residual, residual.data(), 0, residual.size() * sizeof(float)); + ggml_backend_tensor_set(tensors.row_ids, row_ids.data(), 0, row_ids.size() * sizeof(int32_t)); + ggml_backend_tensor_set(tensors.weight, weight.data(), 0, weight.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(backend, tensors.graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + + GatherAddRmsNormOutputs outputs; + outputs.raw.resize(hidden_size * output_token_count); + outputs.normalized.resize(hidden_size * output_token_count); + ggml_backend_tensor_get(tensors.raw_output, outputs.raw.data(), 0, outputs.raw.size() * sizeof(float)); + ggml_backend_tensor_get(tensors.normalized_output, outputs.normalized.data(), 0, + outputs.normalized.size() * sizeof(float)); + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + return outputs; +} + +static void run_gather_add_rmsnorm_f32_numerics() { + ggml_backend_t hrx = ggml_backend_hrx_init(0); + ggml_backend_t cpu = ggml_backend_cpu_init(); + REQUIRE(hrx != nullptr); + REQUIRE(cpu != nullptr); + + const GatherAddRmsNormOutputs expected = run_gather_add_rmsnorm_numerical_case(cpu); + const GatherAddRmsNormOutputs actual = run_gather_add_rmsnorm_numerical_case(hrx); + REQUIRE(actual.raw.size() == expected.raw.size()); + REQUIRE(actual.normalized.size() == expected.normalized.size()); + for (size_t i = 0; i < actual.raw.size(); ++i) { + REQUIRE(actual.raw[i] == expected.raw[i]); + const float tolerance = 3.0e-4f + 3.0e-4f * std::fabs(expected.normalized[i]); + REQUIRE(std::fabs(actual.normalized[i] - expected.normalized[i]) <= tolerance); + } + + ggml_backend_free(cpu); + ggml_backend_free(hrx); +} + +static void run_scale_f32() { + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 512 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t element_count = 1024; + ggml_tensor * input = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, element_count); + ggml_tensor * out = ggml_scale_bias(ctx, input, -1.5f, 0.25f); + REQUIRE(input != nullptr); + REQUIRE(out != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out); + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + + std::vector input_data(element_count); + std::vector expected(element_count); + for (int64_t i = 0; i < element_count; ++i) { + input_data[i] = static_cast(i % 23) * 0.125f - 1.0f; + expected[i] = input_data[i] * -1.5f + 0.25f; + } + + ggml_backend_tensor_set(input, input_data.data(), 0, input_data.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + + std::vector actual(element_count); + ggml_backend_tensor_get(out, actual.data(), 0, actual.size() * sizeof(float)); + for (int64_t i = 0; i < element_count; ++i) { + REQUIRE(std::fabs(actual[i] - expected[i]) <= 1.0e-6f); + } + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(backend); +} + +static void run_scale_f32_inplace() { + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 512 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t element_count = 1024; + ggml_tensor * input = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, element_count); + ggml_tensor * out = ggml_scale_bias_inplace(ctx, input, -1.5f, 0.25f); + REQUIRE(input != nullptr); + REQUIRE(out != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out); + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + + std::vector input_data(element_count); + std::vector expected(element_count); + for (int64_t i = 0; i < element_count; ++i) { + input_data[i] = static_cast(i % 23) * 0.125f - 1.0f; + expected[i] = input_data[i] * -1.5f + 0.25f; + } + + ggml_backend_tensor_set(input, input_data.data(), 0, input_data.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + + std::vector actual(element_count); + ggml_backend_tensor_get(out, actual.data(), 0, actual.size() * sizeof(float)); + for (int64_t i = 0; i < element_count; ++i) { + REQUIRE(std::fabs(actual[i] - expected[i]) <= 1.0e-6f); + } + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(backend); +} + +static void run_two_independent_add_f32() { + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t element_count = 1024; + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, element_count); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, element_count); + ggml_tensor * c = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, element_count); + ggml_tensor * d = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, element_count); + ggml_tensor * out0 = ggml_add(ctx, a, b); + ggml_tensor * out1 = ggml_add(ctx, c, d); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(c != nullptr); + REQUIRE(d != nullptr); + REQUIRE(out0 != nullptr); + REQUIRE(out1 != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out0); + ggml_build_forward_expand(graph, out1); + graph->uid = 1002; + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + + std::vector a_data(element_count); + std::vector b_data(element_count); + std::vector c_data(element_count); + std::vector d_data(element_count); + std::vector expected0(element_count); + std::vector expected1(element_count); + for (int64_t i = 0; i < element_count; ++i) { + a_data[i] = static_cast(i % 17) * 0.25f - 2.0f; + b_data[i] = static_cast(i % 13) * -0.5f + 3.0f; + c_data[i] = static_cast(i % 19) * 0.125f + 1.0f; + d_data[i] = static_cast(i % 11) * 0.75f - 4.0f; + expected0[i] = a_data[i] + b_data[i]; + expected1[i] = c_data[i] + d_data[i]; + } + + ggml_backend_tensor_set(a, a_data.data(), 0, a_data.size() * sizeof(float)); + ggml_backend_tensor_set(b, b_data.data(), 0, b_data.size() * sizeof(float)); + ggml_backend_tensor_set(c, c_data.data(), 0, c_data.size() * sizeof(float)); + ggml_backend_tensor_set(d, d_data.data(), 0, d_data.size() * sizeof(float)); + + ggml_backend_hrx_cache_stats cache_stats = {}; + REQUIRE(ggml_backend_hrx_get_cache_stats(backend, &cache_stats)); + REQUIRE(cache_stats.graph_program_builds == 0); + REQUIRE(cache_stats.graph_program_hits == 0); + REQUIRE(cache_stats.prepared_program_builds == 0); + REQUIRE(cache_stats.prepared_program_hits == 0); + + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + + REQUIRE(ggml_backend_hrx_get_cache_stats(backend, &cache_stats)); + REQUIRE(cache_stats.graph_program_builds == 1); + REQUIRE(cache_stats.graph_program_hits == 0); + REQUIRE(cache_stats.prepared_program_builds == 1); + REQUIRE(cache_stats.prepared_program_hits == 0); + + std::vector actual0(element_count); + std::vector actual1(element_count); + ggml_backend_tensor_get(out0, actual0.data(), 0, actual0.size() * sizeof(float)); + ggml_backend_tensor_get(out1, actual1.data(), 0, actual1.size() * sizeof(float)); + for (int64_t i = 0; i < element_count; ++i) { + REQUIRE(actual0[i] == expected0[i]); + REQUIRE(actual1[i] == expected1[i]); + } + + for (int64_t i = 0; i < element_count; ++i) { + a_data[i] = static_cast(i % 23) * -0.25f + 5.0f; + b_data[i] = static_cast(i % 7) * 0.5f - 1.0f; + c_data[i] = static_cast(i % 5) * -0.125f + 2.0f; + d_data[i] = static_cast(i % 29) * 0.75f - 6.0f; + expected0[i] = a_data[i] + b_data[i]; + expected1[i] = c_data[i] + d_data[i]; + } + ggml_backend_tensor_set(a, a_data.data(), 0, a_data.size() * sizeof(float)); + ggml_backend_tensor_set(b, b_data.data(), 0, b_data.size() * sizeof(float)); + ggml_backend_tensor_set(c, c_data.data(), 0, c_data.size() * sizeof(float)); + ggml_backend_tensor_set(d, d_data.data(), 0, d_data.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + + REQUIRE(ggml_backend_hrx_get_cache_stats(backend, &cache_stats)); + REQUIRE(cache_stats.graph_program_builds == 1); + REQUIRE(cache_stats.graph_program_hits == 1); + REQUIRE(cache_stats.prepared_program_builds == 1); + REQUIRE(cache_stats.prepared_program_hits == 1); + + ggml_backend_tensor_get(out0, actual0.data(), 0, actual0.size() * sizeof(float)); + ggml_backend_tensor_get(out1, actual1.data(), 0, actual1.size() * sizeof(float)); + for (int64_t i = 0; i < element_count; ++i) { + REQUIRE(actual0[i] == expected0[i]); + REQUIRE(actual1[i] == expected1[i]); + } + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(backend); +} + +static void run_chained_add_f32() { + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t element_count = 1024; + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, element_count); + ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, element_count); + ggml_tensor * c = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, element_count); + ggml_tensor * sum = ggml_add(ctx, a, b); + ggml_tensor * out = ggml_add(ctx, sum, c); + REQUIRE(a != nullptr); + REQUIRE(b != nullptr); + REQUIRE(c != nullptr); + REQUIRE(sum != nullptr); + REQUIRE(out != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out); + graph->uid = 1004; + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + + std::vector a_data(element_count); + std::vector b_data(element_count); + std::vector c_data(element_count); + std::vector expected(element_count); + for (int64_t i = 0; i < element_count; ++i) { + a_data[i] = static_cast(i % 17) * 0.25f - 2.0f; + b_data[i] = static_cast(i % 13) * -0.5f + 3.0f; + c_data[i] = static_cast(i % 7) * 0.125f + 1.0f; + expected[i] = a_data[i] + b_data[i] + c_data[i]; + } + + ggml_backend_tensor_set(a, a_data.data(), 0, a_data.size() * sizeof(float)); + ggml_backend_tensor_set(b, b_data.data(), 0, b_data.size() * sizeof(float)); + ggml_backend_tensor_set(c, c_data.data(), 0, c_data.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + + std::vector actual(element_count); + ggml_backend_tensor_get(out, actual.data(), 0, actual.size() * sizeof(float)); + for (int64_t i = 0; i < element_count; ++i) { + REQUIRE(actual[i] == expected[i]); + } + + ggml_backend_hrx_cache_stats cache_stats = {}; + REQUIRE(ggml_backend_hrx_get_cache_stats(backend, &cache_stats)); + REQUIRE(cache_stats.graph_program_builds == 1); + REQUIRE(cache_stats.prepared_program_builds == 1); + + for (int64_t i = 0; i < element_count; ++i) { + a_data[i] = static_cast(i % 11) * -0.25f + 5.0f; + b_data[i] = static_cast(i % 5) * 0.5f - 1.0f; + c_data[i] = static_cast(i % 19) * 0.75f - 6.0f; + expected[i] = a_data[i] + b_data[i] + c_data[i]; + } + ggml_backend_tensor_set(a, a_data.data(), 0, a_data.size() * sizeof(float)); + ggml_backend_tensor_set(b, b_data.data(), 0, b_data.size() * sizeof(float)); + ggml_backend_tensor_set(c, c_data.data(), 0, c_data.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + + REQUIRE(ggml_backend_hrx_get_cache_stats(backend, &cache_stats)); + REQUIRE(cache_stats.graph_program_hits == 1); + REQUIRE(cache_stats.prepared_program_hits == 1); + + ggml_backend_tensor_get(out, actual.data(), 0, actual.size() * sizeof(float)); + for (int64_t i = 0; i < element_count; ++i) { + REQUIRE(actual[i] == expected[i]); + } + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(backend); +} + +static void run_same_uid_distinct_graph_reuses_graph_program(bool distinct_uid = false) { + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + + constexpr int64_t element_count = 1024; + std::vector a_data(element_count); + std::vector b_data(element_count); + std::vector expected(element_count); + + ggml_init_params params0 = {}; + params0.mem_size = 256 * 1024; + params0.no_alloc = true; + ggml_context * ctx0 = ggml_init(params0); + REQUIRE(ctx0 != nullptr); + + ggml_tensor * a0 = ggml_new_tensor_1d(ctx0, GGML_TYPE_F32, element_count); + ggml_tensor * b0 = ggml_new_tensor_1d(ctx0, GGML_TYPE_F32, element_count); + ggml_tensor * out0 = ggml_add(ctx0, a0, b0); + REQUIRE(a0 != nullptr); + REQUIRE(b0 != nullptr); + REQUIRE(out0 != nullptr); + + ggml_cgraph * graph0 = ggml_new_graph(ctx0); + REQUIRE(graph0 != nullptr); + ggml_build_forward_expand(graph0, out0); + graph0->uid = 1003; + + ggml_backend_buffer_t buffer0 = ggml_backend_alloc_ctx_tensors(ctx0, backend); + REQUIRE(buffer0 != nullptr); + + for (int64_t i = 0; i < element_count; ++i) { + a_data[i] = static_cast(i % 17) * 0.25f - 2.0f; + b_data[i] = static_cast(i % 13) * -0.5f + 3.0f; + expected[i] = a_data[i] + b_data[i]; + } + ggml_backend_tensor_set(a0, a_data.data(), 0, a_data.size() * sizeof(float)); + ggml_backend_tensor_set(b0, b_data.data(), 0, b_data.size() * sizeof(float)); + REQUIRE(ggml_backend_graph_compute(backend, graph0) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + + ggml_backend_hrx_cache_stats cache_stats = {}; + REQUIRE(ggml_backend_hrx_get_cache_stats(backend, &cache_stats)); + REQUIRE(cache_stats.graph_program_builds == 1); + REQUIRE(cache_stats.graph_program_hits == 0); + REQUIRE(cache_stats.prepared_program_builds == 1); + REQUIRE(cache_stats.prepared_program_hits == 0); + + ggml_init_params params1 = {}; + params1.mem_size = 256 * 1024; + params1.no_alloc = true; + ggml_context * ctx1 = ggml_init(params1); + REQUIRE(ctx1 != nullptr); + + ggml_tensor * a1 = ggml_new_tensor_1d(ctx1, GGML_TYPE_F32, element_count); + ggml_tensor * b1 = ggml_new_tensor_1d(ctx1, GGML_TYPE_F32, element_count); + ggml_tensor * out1 = ggml_add(ctx1, a1, b1); + REQUIRE(a1 != nullptr); + REQUIRE(b1 != nullptr); + REQUIRE(out1 != nullptr); + + ggml_cgraph * graph1 = ggml_new_graph(ctx1); + REQUIRE(graph1 != nullptr); + ggml_build_forward_expand(graph1, out1); + graph1->uid = distinct_uid ? 1004 : 1003; + + ggml_backend_buffer_t buffer1 = ggml_backend_alloc_ctx_tensors(ctx1, backend); + REQUIRE(buffer1 != nullptr); + + for (int64_t i = 0; i < element_count; ++i) { + a_data[i] = static_cast(i % 23) * -0.25f + 5.0f; + b_data[i] = static_cast(i % 7) * 0.5f - 1.0f; + expected[i] = a_data[i] + b_data[i]; + } + ggml_backend_tensor_set(a1, a_data.data(), 0, a_data.size() * sizeof(float)); + ggml_backend_tensor_set(b1, b_data.data(), 0, b_data.size() * sizeof(float)); + REQUIRE(ggml_backend_graph_compute(backend, graph1) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + + REQUIRE(ggml_backend_hrx_get_cache_stats(backend, &cache_stats)); + REQUIRE(cache_stats.graph_program_builds == 1); + REQUIRE(cache_stats.graph_program_hits == 1); + REQUIRE(cache_stats.prepared_program_builds == 2); + REQUIRE(cache_stats.prepared_program_hits == 0); + + std::vector actual(element_count); + ggml_backend_tensor_get(out1, actual.data(), 0, actual.size() * sizeof(float)); + for (int64_t i = 0; i < element_count; ++i) { + REQUIRE(actual[i] == expected[i]); + } + + if (distinct_uid) { + for (int64_t i = 0; i < element_count; ++i) { + a_data[i] += 1.0f; + expected[i] = a_data[i] + b_data[i]; + } + ggml_backend_tensor_set(a1, a_data.data(), 0, a_data.size() * sizeof(float)); + REQUIRE(ggml_backend_graph_compute(backend, graph1) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + ggml_backend_tensor_get(out1, actual.data(), 0, actual.size() * sizeof(float)); + REQUIRE(actual == expected); + REQUIRE(ggml_backend_hrx_get_cache_stats(backend, &cache_stats)); + REQUIRE(cache_stats.graph_program_builds == 1); + REQUIRE(cache_stats.graph_program_hits == 2); + REQUIRE(cache_stats.prepared_program_builds == 2); + REQUIRE(cache_stats.prepared_program_hits == 1); + } + + ggml_backend_buffer_free(buffer1); + ggml_free(ctx1); + ggml_backend_buffer_free(buffer0); + ggml_free(ctx0); + ggml_backend_free(backend); +} + +static void run_unsupported_op_fails() { + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * a = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 8); + ggml_tensor * out = ggml_sin(ctx, a); + REQUIRE(a != nullptr); + REQUIRE(out != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, out); + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + + std::vector input(8, 2.0f); + ggml_backend_tensor_set(a, input.data(), 0, input.size() * sizeof(float)); + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_FAILED); + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(backend); +} + +static void run_lowtoken_residual_dispatch_checks() { + struct Case { + ggml_type type; + int64_t k, n, tokens; + int layout; + int side_use; + bool alias_weight; + bool second_add; + bool fused; + }; + + const Case cases[] = { + { GGML_TYPE_Q4_K, 6144, 5120, 1, 0, 0, false, false, true }, + { GGML_TYPE_Q4_K, 17408, 5120, 1, 0, 0, false, false, true }, + { GGML_TYPE_Q6_K, 17408, 5120, 1, 0, 0, false, false, true }, + { GGML_TYPE_Q4_K, 4096, 4096, 1, 0, 0, false, false, true }, + { GGML_TYPE_Q4_K, 32768, 4096, 1, 0, 0, false, false, true }, + { GGML_TYPE_Q4_K, 6144, 5120, 1, 1, 0, false, false, true }, + { GGML_TYPE_Q4_K, 6144, 5120, 2, 1, 0, false, false, true }, + { GGML_TYPE_Q4_K, 6144, 5120, 3, 1, 0, false, false, true }, + { GGML_TYPE_Q4_K, 6144, 5120, 4, 1, 0, false, false, true }, + { GGML_TYPE_Q4_K, 6144, 5120, 5, 1, 0, false, false, true }, + { GGML_TYPE_Q6_K, 17408, 5120, 5, 1, 0, false, false, true }, + { GGML_TYPE_Q4_K, 6144, 5120, 1, 0, 0, false, true, true }, + { GGML_TYPE_Q4_K, 6144, 5120, 5, 1, 0, false, true, true }, + { GGML_TYPE_Q4_K, 6144, 5120, 1, 0, 1, false, false, false }, + { GGML_TYPE_Q4_K, 6144, 5120, 1, 1, 1, false, false, false }, + { GGML_TYPE_Q4_K, 6144, 5120, 5, 1, 1, false, false, false }, + { GGML_TYPE_Q4_K, 6144, 5120, 1, 1, 2, false, false, false }, + { GGML_TYPE_Q4_K, 6144, 5120, 5, 1, 2, false, false, false }, + { GGML_TYPE_Q4_K, 6144, 5120, 1, 2, 0, false, false, false }, + { GGML_TYPE_Q4_K, 6144, 5120, 5, 2, 0, false, false, false }, + { GGML_TYPE_Q4_K, 6144, 5120, 1, 3, 0, false, false, false }, + { GGML_TYPE_Q4_K, 6144, 5120, 5, 3, 0, false, false, false }, + { GGML_TYPE_Q4_K, 6144, 5120, 1, 0, 0, true, false, true }, + { GGML_TYPE_Q4_K, 6144, 5120, 5, 1, 0, true, false, false }, + { GGML_TYPE_Q4_K, 3840, 4096, 1, 0, 0, false, false, true }, + { GGML_TYPE_Q4_K, 6144, 4032, 1, 0, 0, false, false, true }, + { GGML_TYPE_Q4_K, 4096, 8192, 1, 0, 0, false, false, true }, + { GGML_TYPE_F16, 6144, 5120, 1, 0, 0, false, false, true }, + { GGML_TYPE_Q4_K, 6144, 5120, 6, 1, 0, false, false, false }, + { GGML_TYPE_Q4_K, 6144, 5120, 256, 1, 0, false, false, false }, + }; + for (const Case & c : cases) { + ggml_context * ctx = ggml_init({ 16 * 1024 * 1024, nullptr, true }); + REQUIRE(ctx != nullptr); + ggml_tensor * weight = ggml_new_tensor_2d(ctx, c.type, c.k, c.n + (c.alias_weight ? 64 : 0)); + if (c.alias_weight) { + weight = ggml_view_2d(ctx, weight, c.k, c.n, weight->nb[1], 0); + } + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, c.k, c.tokens); + ggml_tensor * projection = ggml_mul_mat(ctx, weight, input); + ggml_tensor * laid_out = projection; + if (c.layout == 1) { + laid_out = ggml_reshape_2d(ctx, projection, c.n, c.tokens); + } else if (c.layout == 2) { + laid_out = ggml_reshape_2d(ctx, projection, c.n / 2, c.tokens * 2); + } else if (c.layout == 3) { + laid_out = ggml_view_2d(ctx, projection, c.n - 64, c.tokens, projection->nb[1], 0); + } + ggml_tensor * residual = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, laid_out->ne[0], laid_out->ne[1]); + ggml_tensor * output = ggml_add(ctx, laid_out, residual); + if (c.second_add) { + output = ggml_add(ctx, output, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, laid_out->ne[0], laid_out->ne[1])); + } + ggml_cgraph * graph = ggml_new_graph_custom(ctx, 64, false); + ggml_build_forward_expand(graph, output); + if (c.side_use != 0) { + ggml_build_forward_expand(graph, ggml_scale(ctx, c.side_use == 1 ? projection : laid_out, 0.5f)); + } + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const ggml::hrx::CommandProgram program = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(program.valid()); + REQUIRE(command_program_verifies(program)); + size_t fused = 0, binary = 0; + bool fused_second_add = false; + for (const auto & command : program.commands) { + const std::string name = kernel_name_for_id(command.kernel.kernel_id); + binary += name == "loom_libs:ggml_binary_f32"; + const bool tiled_postops = name == "loom_libs:ggml_mul_mat_tiled_input_f32_bias_residual_publish_f32"; + const bool skinny_postops = name == "loom_libs:ggml_mul_mat_skinny_input_f32_bias_residual_publish_f32"; + const bool vector_postops = name == "loom_libs:ggml_mul_mat_vector_bias_residual_f32_f32"; + if (!tiled_postops && !skinny_postops && !vector_postops) { + continue; + } + ++fused; + if (vector_postops) { + fused_second_add = + fused_second_add || + (command.kernel.compile_parameters.at("ggml.matmul.vector_postops.apply_bias") == "1" && + command.kernel.compile_parameters.at("ggml.matmul.vector_postops.apply_residual") == "1"); + } else { + fused_second_add = + fused_second_add || command.kernel.compile_parameters.at("ggml.mul_mat_postops.epilogue") == "3"; + } + REQUIRE(command.bindings.size() >= 4); + REQUIRE(command.kernel.integer_parameters.at("token_count") == c.tokens); + if (vector_postops) { + REQUIRE(command.kernel.compile_parameters.count("ggml.matmul.vector_postops.weight_format") == 1); + } else { + REQUIRE(command.kernel.compile_parameters.count("ggml.mul_mat_postops.weight_format") == 1); + } + } + REQUIRE(fused == (c.fused ? 1 : 0)); + const size_t expected_binary = c.fused ? (c.second_add && !fused_second_add ? 1 : 0) : 1; + REQUIRE(binary == expected_binary); + ggml_free(ctx); + } +} + +static void run_lowtoken_bias_dispatch_checks() { + ggml_context * ctx = ggml_init({ 16 * 1024 * 1024, nullptr, true }); + REQUIRE(ctx != nullptr); + + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, 3584, 1024); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 3584, 2); + ggml_tensor * projection = ggml_mul_mat(ctx, weight, input); + ggml_tensor * bias = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 1024); + ggml_tensor * output = ggml_add(ctx, projection, bias); + + ggml_cgraph * graph = ggml_new_graph_custom(ctx, 64, false); + ggml_build_forward_expand(graph, output); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const ggml::hrx::CommandProgram program = ggml::hrx::build_command_program( + imported.graph, scheduler.plan(), ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(program.valid()); + REQUIRE(command_program_verifies(program)); + + size_t skinny_bias = 0; + size_t tiled_bias = 0; + for (const auto & command : program.commands) { + const std::string name = kernel_name_for_id(command.kernel.kernel_id); + skinny_bias += name == "loom_libs:ggml_mul_mat_skinny_input_f32_bias_residual_publish_f32"; + tiled_bias += name == "loom_libs:ggml_mul_mat_tiled_input_f32_bias_residual_publish_f32"; + } + REQUIRE(skinny_bias == 1); + REQUIRE(tiled_bias == 0); + + ggml_free(ctx); +} + +static void run_gdn_selected_rms_q8_dispatch_checks() { + ggml_init_params params = {}; + params.mem_size = 32 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + enum { + UnreadyInput = 1, + UnreadyGate = 2, + GatherSide = 4, + AttentionSide = 8, + GdnObserved = 16, + GatherObserved = 32, + Unaligned = 64, + CacheRead = 128, + OutputSide = 256, + CachedGate = 512, + TanhGate = 1024, + AttentionReshape = 2048, + AttentionObserved = 4096, + }; + + const struct { + int64_t tokens, heads, sequences; + int flags; + ggml_type type; + bool fused; + } cases[] = { + {1,48,1,0,GGML_TYPE_Q4_K,true}, + {2,48,1,0,GGML_TYPE_Q4_K,true}, + {3,48,1,0,GGML_TYPE_Q4_K,true}, + {4,48,1,0,GGML_TYPE_Q4_K,true}, + {5,48,1,0,GGML_TYPE_Q4_K,true}, + {1,24,1,0,GGML_TYPE_Q4_K,true}, + {5,24,1,0,GGML_TYPE_Q4_K,true}, + {1,48,1,0,GGML_TYPE_Q6_K,true}, + {5,48,1,0,GGML_TYPE_Q6_K,true}, + {1,48,1,AttentionReshape,GGML_TYPE_Q4_K,true}, + {5,48,1,AttentionReshape,GGML_TYPE_Q4_K,true}, + {1,48,1,UnreadyInput,GGML_TYPE_Q4_K,false}, + {1,48,1,UnreadyGate,GGML_TYPE_Q4_K,false}, + {5,48,1,UnreadyGate,GGML_TYPE_Q4_K,false}, + {1,48,1,GatherSide,GGML_TYPE_Q4_K,false}, + {5,48,1,AttentionSide,GGML_TYPE_Q4_K,false}, + {1,48,1,GdnObserved,GGML_TYPE_Q4_K,false}, + {1,48,1,GatherObserved,GGML_TYPE_Q4_K,false}, + {1,48,1,AttentionObserved,GGML_TYPE_Q4_K,false}, + {1,48,1,Unaligned,GGML_TYPE_Q4_K,false}, + {5,48,1,CacheRead,GGML_TYPE_Q4_K,false}, + {1,48,1,OutputSide,GGML_TYPE_Q4_K,false}, + {1,48,1,CachedGate,GGML_TYPE_Q4_K,false}, + {5,48,1,TanhGate,GGML_TYPE_Q4_K,false}, + {1,48,1,0,GGML_TYPE_F16,false}, + {6,48,1,0,GGML_TYPE_Q4_K,false}, + {1,48,2,0,GGML_TYPE_Q4_K,false}, + {1,6,1,0,GGML_TYPE_Q4_K,false}, + }; + for (const auto & test : cases) { + const int64_t width = 128, heads = test.heads, qheads = heads / 3; + const int64_t hidden = width * (heads + 2 * qheads), channels = width * heads; + const int64_t state_elements = width * channels; + const int64_t written = std::min(test.tokens, int64_t{5}); + const size_t state_bytes = state_elements * sizeof(float); + ggml_tensor * raw = ggml_scale(ctx, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden, test.tokens * test.sequences), 0.5f); + ggml_tensor * alpha = ggml_scale(ctx, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, heads, test.tokens * test.sequences), 0.5f); + ggml_tensor * beta_raw = ggml_scale(ctx, ggml_dup_tensor(ctx, alpha), 0.5f); + ggml_tensor * bias = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, heads); + ggml_tensor * scale = ggml_dup_tensor(ctx, bias); + ggml_tensor * gate = ggml_mul(ctx, ggml_softplus(ctx, ggml_add(ctx, ggml_reshape_3d(ctx, alpha, heads, test.tokens, test.sequences), bias)), scale); + gate = ggml_reshape_4d(ctx, gate, 1, heads, test.tokens, test.sequences); + ggml_tensor * beta = ggml_sigmoid(ctx, ggml_reshape_4d(ctx, beta_raw, 1, heads, test.tokens, test.sequences)); + ggml_tensor * cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, state_elements, 20); + ggml_tensor * ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, test.sequences); + ggml_tensor * gathered = ggml_get_rows(ctx, cache, ids); + if (test.flags & GatherObserved) { + ggml_set_output(gathered); + } + ggml_tensor * state = ggml_reshape_4d(ctx, gathered, width, width, heads, test.sequences); + const auto view = [&](int64_t nheads, size_t offset) { + return ggml_view_4d(ctx, raw, width, nheads, test.tokens, test.sequences, width * sizeof(float), + hidden * sizeof(float), hidden * test.tokens * sizeof(float), offset); + }; + ggml_tensor * q = ggml_l2_norm(ctx, view(qheads, 0), 1.e-6f); + ggml_tensor * k = ggml_l2_norm(ctx, view(qheads, width * qheads * sizeof(float)), 1.e-6f); + ggml_tensor * v = view(heads, 2 * width * qheads * sizeof(float)); + ggml_tensor * gdn = ggml_gated_delta_net(ctx, q, k, v, gate, beta, state, 5); + if (test.flags & GdnObserved) { + ggml_set_output(gdn); + } + const size_t attention_bytes = channels * test.tokens * test.sequences * sizeof(float); + ggml_tensor * attention = + ggml_view_4d(ctx, gdn, width, heads, test.tokens, test.sequences, width * sizeof(float), + channels * sizeof(float), channels * test.tokens * sizeof(float), 0); + if (test.flags & AttentionObserved) { + ggml_set_output(attention); + } + ggml_tensor * snapshots = ggml_view_3d(ctx, gdn, state_elements, test.sequences, written, state_bytes, + state_bytes * test.sequences, attention_bytes); + ggml_tensor * target = ggml_view_3d(ctx, cache, state_elements, test.sequences, written, state_bytes, + 4 * state_bytes, (test.flags & Unaligned) ? sizeof(float) : state_bytes); + ggml_tensor * norm_input = (test.flags & AttentionReshape) ? + ggml_reshape_2d(ctx, attention, width, heads * test.tokens * test.sequences) : + attention; + ggml_tensor * rms_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, width); + ggml_tensor * raw_gate = + (test.flags & CachedGate) ? + ggml_view_4d(ctx, cache, width, heads, test.tokens, test.sequences, width * sizeof(float), + channels * sizeof(float), channels * test.tokens * sizeof(float), 0) : + ggml_scale(ctx, ggml_new_tensor_4d(ctx, GGML_TYPE_F32, width, heads, test.tokens, test.sequences), + 0.5f); + ggml_tensor * shaped_gate = (test.flags & AttentionReshape) ? + ggml_reshape_2d(ctx, raw_gate, width, heads * test.tokens * test.sequences) : + raw_gate; + ggml_tensor * norm = ggml_mul(ctx, ggml_rms_norm(ctx, norm_input, 1.e-6f), rms_weight); + ggml_tensor * gated = + ggml_mul(ctx, norm, (test.flags & TanhGate) ? ggml_tanh(ctx, shaped_gate) : ggml_silu(ctx, shaped_gate)); + ggml_tensor * projection_input = ggml_reshape_2d(ctx, gated, channels, test.tokens * test.sequences); + ggml_tensor * projection_weight = ggml_new_tensor_2d(ctx, test.type, channels, 4096); + ggml_tensor * projection = ggml_mul_mat(ctx, projection_weight, projection_input); + ggml_cgraph * graph = ggml_new_graph(ctx); + if (!(test.flags & UnreadyInput)) { + ggml_build_forward_expand(graph, raw); + } + ggml_build_forward_expand(graph, alpha); + ggml_build_forward_expand(graph, beta_raw); + if (!(test.flags & UnreadyGate)) { + ggml_build_forward_expand(graph, raw_gate); + } + ggml_build_forward_expand(graph, gathered); + if (test.flags & CacheRead) { + ggml_build_forward_expand(graph, ggml_scale(ctx, cache, 0.25f)); + } + ggml_build_forward_expand(graph, ggml_cpy(ctx, snapshots, target)); + ggml_build_forward_expand(graph, ggml_scale(ctx, projection, 0.5f)); + if (test.flags & GatherSide) { + ggml_build_forward_expand(graph, ggml_scale(ctx, gathered, 0.5f)); + } + if (test.flags & AttentionSide) { + ggml_build_forward_expand(graph, ggml_scale(ctx, attention, 0.5f)); + } + if (test.flags & OutputSide) { + ggml_build_forward_expand(graph, ggml_scale(ctx, gated, 0.5f)); + } + auto imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, test_dispatch_target())); + const auto & plan = scheduler.plan(); + REQUIRE(plan.valid()); + const auto fused = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == + "loom_libs:llm_gated_delta_net_f32_wmma_head128_selected_snapshot_projection_rms_gate_q8"; + }); + REQUIRE((fused != plan.dispatches.end()) == test.fused); + if (test.fused) { + REQUIRE(fused->bindings.size() == 14); + require_compile_parameter(*fused, "llm.gated_delta_net.state_row_count", "20"); + REQUIRE(fused->bindings[7].value == imported.graph.values().find_tensor(cache)->id); + REQUIRE(fused->bindings[10].value == imported.graph.values().find_tensor(ids)->id); + const auto * output = imported.graph.values().find_tensor(gated); + const size_t bytes = static_cast(test.tokens * heads) * ggml_row_size(GGML_TYPE_Q8_1, width); + const auto * alternate = plan.metadata.find_alternate_value(output->id, GGML_TYPE_Q8_1, bytes); + REQUIRE(alternate != nullptr && alternate->alternate_value == fused->bindings.back().value); + REQUIRE(std::none_of(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + const auto name = kernel_name_for_id(dispatch.kernel.kernel_id); + return name == "loom_libs:ggml_get_rows_f32" || name == "loom_libs:ggml_copy_f32" || + name == "loom_libs:ggml_rmsnorm_gate_f32_publish"; + })); + } + const auto commands = + ggml::hrx::build_command_program(imported.graph, plan, ggml::hrx::get_qwen_kernel_corpus(), "gfx1151"); + REQUIRE(commands.valid()); + if (!(test.flags & UnreadyGate)) { + REQUIRE(command_program_verifies(commands)); + } + } + ggml_free(ctx); +} + +static void register_hrx_backend_host_cases(test_runner::Suite & suite) { + suite.host_case("status", [] { run_status_checks(); }); + suite.host_case("command_plan_metadata", [] { run_command_plan_metadata_checks(); }); + suite.host_case("dispatch_registry", [] { run_dispatch_registry_checks(); }); + suite.host_case("graph_import", [] { run_graph_import_checks(); }); + suite.host_case("graph_import_mixed_backend_boundary", [] { run_graph_import_mixed_backend_boundary_checks(); }); + suite.host_case("graph_view_external_use", [] { run_graph_view_external_use_checks(); }); + suite.host_case("scale_f32_dispatch", [] { run_scale_f32_dispatch_checks(); }); + suite.host_case("scale_add_f32_dispatch", [] { run_scale_add_f32_dispatch_checks(); }); + suite.host_case("rmsnorm_two_binary_dispatch", [] { run_rmsnorm_two_binary_dispatch_checks(); }); + suite.host_case("cont_f32_dispatch", [] { run_cont_f32_dispatch_checks(); }); + suite.host_case("binary_f32_broadcast_dispatch", [] { run_binary_f32_broadcast_dispatch_checks(); }); + suite.host_case("binary_q8_publication", [] { run_binary_q8_publication_checks(); }); + suite.host_case("graph_snapshot_diagnostics", [] { run_graph_snapshot_diagnostics_checks(); }); + suite.host_case("unmatched_graph_diagnostics", [] { run_unmatched_graph_diagnostics_checks(); }); + suite.host_case("completion_counter_plan", [] { run_completion_counter_plan_checks(); }); + suite.host_case("graph_index", [] { run_graph_index_checks(); }); + suite.host_case("rope_set_rows_dispatch", [] { run_rope_set_rows_dispatch_checks(); }); + suite.host_case("graph_traversal", [] { run_graph_traversal_checks(); }); + suite.host_case("qwen_token_embedding_dispatch", [] { run_qwen_token_embedding_dispatch_checks(); }); + suite.host_case("get_rows_rmsnorm_binary_dispatch", [] { run_get_rows_rmsnorm_binary_dispatch_checks(); }); + suite.host_case("get_rows_scale_dispatch", [] { run_get_rows_scale_dispatch_checks(); }); + suite.host_case("gather_add_dispatch", [] { run_gather_add_dispatch_checks(); }); + suite.host_case("qwen_flash_attention_dispatch", [] { run_qwen_flash_attention_dispatch_checks(); }); + suite.host_case("tiled_pair_matmul_postops_dispatch", [] { run_tiled_pair_matmul_postops_dispatch_checks(); }); + suite.host_case("qwen_attention_postprocess_dispatch", [] { run_qwen_attention_postprocess_dispatch_checks(); }); + suite.host_case("qwen_matmul_dispatch", [] { run_qwen_matmul_dispatch_checks(); }); + suite.host_case("iq1_codebook_matmul_dispatch", [] { run_iq1_codebook_matmul_dispatch_checks(); }); + suite.host_case("iq3_xxs_codebook_matmul_dispatch", [] { run_iq3_xxs_codebook_matmul_dispatch_checks(); }); + suite.host_case("iq2_codebook_resource_reuse", [] { run_iq2_codebook_resource_reuse_checks(); }); + suite.host_case("iq4_matmul_dispatch", [] { run_iq4_matmul_dispatch_checks(); }); + suite.host_case("packed_f16_producer_consumer", [] { run_packed_f16_producer_consumer_checks(); }); + suite.host_case("generic_swiglu_k16_publish", [] { run_generic_swiglu_k16_publish_checks(); }); + suite.host_case("generic_swiglu_k16_publication_edges", [] { run_generic_swiglu_k16_publication_edge_checks(); }); + suite.host_case("tiled_matmul_alternate_publish", [] { run_tiled_matmul_alternate_publish_checks(); }); + suite.host_case("k16_major_preparation_reuse", [] { run_k16_major_preparation_reuse_checks(); }); + suite.host_case("rmsnorm_binary_k16_output", [] { run_rmsnorm_binary_k16_output_checks(); }); + suite.host_case("rmsnorm_binary_f16_output", [] { run_rmsnorm_binary_f16_output_checks(); }); + suite.host_case("swiglu_q8_output", [] { run_swiglu_q8_output_checks(); }); + suite.host_case("rmsnorm_gate_packed_output", [] { run_rmsnorm_gate_packed_output_checks(); }); + suite.host_case("rmsnorm_gate_q8_output", [] { run_rmsnorm_gate_q8_output_checks(); }); + suite.host_case("symmetric_i4_consumer_qualification", + [] { run_symmetric_i4_consumer_qualification_checks(); }); + suite.host_case("lowtoken_residual_dispatch", [] { run_lowtoken_residual_dispatch_checks(); }); + suite.host_case("lowtoken_bias_dispatch", [] { run_lowtoken_bias_dispatch_checks(); }); + suite.host_case("generic_ssm_conv_binary_dispatch", [] { run_generic_ssm_conv_binary_dispatch_checks(); }); + suite.host_case("quantized_conv4_dispatch", [] { run_quantized_conv4_dispatch_checks(); }); + suite.host_case("lfm_dconv3_dispatch", [] { run_lfm_dconv3_dispatch_checks(); }); + suite.host_case("gdn_rmsnorm_gate_dispatch", [] { run_gdn_rmsnorm_gate_dispatch_checks(); }); + suite.host_case("gdn_native_projection_pair_dispatch", [] { run_gdn_native_projection_pair_dispatch_checks(); }); + suite.host_case("gdn_selected_snapshot_dispatch", [] { run_gdn_selected_snapshot_dispatch_checks(); }); + suite.host_case("gdn_selected_rms_q8_dispatch", [] { run_gdn_selected_rms_q8_dispatch_checks(); }); + suite.host_case("llama_attention_matmul_dispatch", [] { run_llama_attention_matmul_dispatch_checks(); }); + suite.host_case("quantized_value_projection_dispatch", [] { run_quantized_value_projection_dispatch_checks(); }); + suite.host_case("prefill_value_projection_cache_dispatch", + [] { run_prefill_value_projection_cache_dispatch_checks(); }); + suite.host_case("qwen_terminal_q6k_q8_command.tokens1", [] { schedule_qwen_terminal_q6k_q8_command(1); }); + suite.host_case("qwen_terminal_q6k_q8_command.tokens18", [] { schedule_qwen_terminal_q6k_q8_command(18); }); + suite.host_case("qwen_terminal_q6k_vector_command", [] { schedule_qwen_terminal_q6k_vector_command(); }); + suite.host_case("qwen_decode_rmsnorm_publication", [] { run_qwen_decode_rmsnorm_publication_checks(); }); + suite.host_case("qwen_decode_publication_consumer_qualification", [] { + run_qwen_decode_publication_consumer_qualification_checks(); + }); + suite.host_case("get_rows_q8_1_alternate_command.q4_k", + [] { schedule_get_rows_q8_1_alternate_command(GGML_TYPE_Q4_K); }); + suite.host_case("get_rows_q8_1_alternate_command.bf16", + [] { schedule_get_rows_q8_1_alternate_command(GGML_TYPE_BF16); }); + suite.host_case("qwen_router_top8_dispatch", [] { run_qwen_router_top8_dispatch_checks(); }); + suite.host_case("common_mul_mat_id_swiglu_dispatch", [] { run_common_mul_mat_id_swiglu_dispatch_checks(); }); + suite.host_case("qwen_routed_gate_up_dispatch", [] { run_qwen_routed_gate_up_dispatch_checks(); }); + suite.host_case("qwen_weighted_reduce_next_rmsnorm_q8_publication", [] { + run_qwen_weighted_reduce_next_rmsnorm_q8_publication_checks(); + }); + suite.host_case("common_mul_mat_id_postops_dispatch", [] { run_common_mul_mat_id_postops_dispatch_checks(); }); + suite.host_case("alias_value_import", [] { run_alias_value_import_checks(); }); + suite.host_case("multi_dispatch", [] { run_multi_dispatch_checks(); }); + suite.host_case("layout_alias_scheduler_elision", [] { run_layout_alias_scheduler_elision_checks(); }); + suite.host_case("zero_output_scheduler_elision", [] { run_zero_output_scheduler_elision_checks(); }); + suite.host_case("transient_import", [] { run_transient_import_checks(); }); + suite.host_case("graph_view_preserves_shared_output_storage", + [] { run_graph_view_preserves_shared_output_storage_checks(); }); + suite.host_case("chained_dispatch_requires_transients", [] { run_chained_dispatch_requires_transients(); }); + suite.host_case("graph_replay_host_staging_is_not_ineligible", + [] { run_graph_replay_host_staging_is_not_ineligible(); }); + suite.host_case("multiple_transient_plan", [] { run_multiple_transient_plan_checks(); }); + suite.host_case("disjoint_transient_plan_packing", [] { run_disjoint_transient_plan_packing_checks(); }); + suite.host_case("command_program_kernel_dump", [] { run_command_program_kernel_dump_checks(); }); + suite.host_case("command_shape_hash", [] { run_command_shape_hash_checks(); }); + suite.host_case("graph_program_cache_uid_mismatch", [] { run_graph_program_cache_uid_mismatch_checks(); }); + suite.host_case("graph_match_hash_collision", [] { run_graph_match_hash_collision_checks(); }); + suite.host_case("graph_match_bijection", [] { run_graph_match_bijection_checks(); }); + suite.host_case("validated_graph_uid_match", [] { run_validated_graph_uid_match_checks(); }); + suite.host_case("graph_executor_contract", [] { run_graph_executor_contract_checks(); }); + suite.host_case("loom_async_jit_environment", [] { + REQUIRE(ggml::hrx::loom_async_jit_enabled_from_environment() == async_jit_expected_from_environment()); + }); +} + +static void register_hrx_backend_device_cases(test_runner::Suite & suite) { + suite.device_case("zero_output_device_support", [] { run_zero_output_device_support_checks(); }); + suite.device_case("scale_f32_device_support", [] { run_scale_f32_device_support_checks(); }); + suite.device_case("binary_q8_publication_numerics", [] { run_binary_q8_publication_numerics(); }); + suite.device_case("qwen_expert_table_partition_prefill_512_execution", + [] { run_qwen_expert_table_partition_prefill_512_execution(); }); + suite.device_case("add_f32", [] { run_add_f32(); }); + suite.device_case("rope_scale_f32_numerics", [] { run_rope_scale_f32_numerics(); }); + suite.device_case("gather_add_rmsnorm_f32_numerics", [] { run_gather_add_rmsnorm_f32_numerics(); }); + suite.device_case("scale_f32", [] { run_scale_f32(); }); + suite.device_case("scale_f32_inplace", [] { run_scale_f32_inplace(); }); + suite.device_case("two_independent_add_f32", [] { run_two_independent_add_f32(); }); + suite.device_case("chained_add_f32", [] { run_chained_add_f32(); }); + suite.device_case("same_uid_distinct_graph_reuses_graph_program.default", + [] { run_same_uid_distinct_graph_reuses_graph_program(); }); + suite.device_case("same_uid_distinct_graph_reuses_graph_program.distinct_uid", + [] { run_same_uid_distinct_graph_reuses_graph_program(true); }); + suite.device_case("unsupported_op_fails", [] { run_unsupported_op_fails(); }); +} + +static void register_hrx_backend_cases(test_runner::Suite & suite) { + register_hrx_backend_host_cases(suite); + register_hrx_backend_device_cases(suite); +} + +int main(int argc, char ** argv) { + test_runner::Suite suite( + test_runner::Config::with_prefix("HRX backend test", "hrx-backend", "GGML_HRX_BACKEND_TEST")); + register_hrx_backend_cases(suite); + + const bool has_device = ggml_backend_hrx_get_device_count() != 0; + return suite.run(argc, argv, has_device); +} diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 4098acaaf91a..7b5a2cada8ef 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -351,9 +351,20 @@ static std::string var_to_str(const std::string & x) { return x; } +// "%f": the std::to_string format for floating point before C++26 +static std::string fp_to_str(double x) { + char buf[512]; + snprintf(buf, sizeof(buf), "%f", x); + return buf; +} + template static std::string var_to_str(const T & x) { - return std::to_string(x); + if constexpr (std::is_floating_point_v) { + return fp_to_str(x); + } else { + return std::to_string(x); + } } template @@ -406,10 +417,10 @@ static std::string var_to_str(ggml_scale_mode mode) { case GGML_SCALE_MODE_BICUBIC: str = "bicubic"; break; default: str = std::to_string(mode); break; } - if (mode & GGML_SCALE_FLAG_ALIGN_CORNERS) { + if (mode & (int) GGML_SCALE_FLAG_ALIGN_CORNERS) { str += "|align_corners"; } - if (mode & GGML_SCALE_FLAG_ANTIALIAS) { + if (mode & (int) GGML_SCALE_FLAG_ANTIALIAS) { str += "|antialias"; } return str; @@ -621,9 +632,9 @@ struct test_result { std::to_string(supported), std::to_string(passed), error_message, - std::to_string(time_us), - std::to_string(flops), - std::to_string(bandwidth_gb_s), + fp_to_str(time_us), + fp_to_str(flops), + fp_to_str(bandwidth_gb_s), std::to_string(memory_kb), std::to_string(n_runs), device_description, @@ -3794,6 +3805,9 @@ struct test_dsv4_hc : public test_case { if (name == "post") { lo = 0.0f; hi = 2.0f; return true; } + if (name == "gate") { + lo = -4.0f; hi = 4.0f; return true; + } if (name == "x" || name == "residual") { lo = -1.0f; hi = 1.0f; return true; } @@ -3858,6 +3872,7 @@ struct test_dsv4_hc_comb : public test_dsv4_hc { struct test_dsv4_hc_pre : public test_dsv4_hc { const int64_t n_embd; const int64_t n_tokens; + const bool gated; std::string op_desc(ggml_tensor * t) override { GGML_UNUSED(t); @@ -3865,20 +3880,27 @@ struct test_dsv4_hc_pre : public test_dsv4_hc { } std::string vars() override { - return VARS_TO_STR2(n_embd, n_tokens); + return VARS_TO_STR3(n_embd, n_tokens, gated); } - test_dsv4_hc_pre(int64_t n_embd = 31, int64_t n_tokens = 17) - : n_embd(n_embd), n_tokens(n_tokens) {} + test_dsv4_hc_pre(int64_t n_embd = 31, int64_t n_tokens = 17, bool gated = false) + : n_embd(n_embd), n_tokens(n_tokens), gated(gated) {} ggml_tensor * build_graph(ggml_context * ctx) override { ggml_tensor * x = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd, hc, n_tokens); ggml_set_name(x, "x"); - ggml_tensor * weights = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hc, n_tokens); - ggml_set_name(weights, "weights"); + if (gated) { + ggml_tensor * gate = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd, hc, n_tokens); + ggml_set_name(gate, "gate"); - out = ggml_dsv4_hc_pre(ctx, x, weights); + out = ggml_dsv4_hc_pre_gated(ctx, x, gate, 1.0f/hc); + } else { + ggml_tensor * weights = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hc, n_tokens); + ggml_set_name(weights, "weights"); + + out = ggml_dsv4_hc_pre(ctx, x, weights); + } ggml_set_name(out, "out"); return out; } @@ -3887,6 +3909,7 @@ struct test_dsv4_hc_pre : public test_dsv4_hc { struct test_dsv4_hc_post : public test_dsv4_hc { const int64_t n_embd; const int64_t n_tokens; + const bool identity; std::string op_desc(ggml_tensor * t) override { GGML_UNUSED(t); @@ -3894,11 +3917,11 @@ struct test_dsv4_hc_post : public test_dsv4_hc { } std::string vars() override { - return VARS_TO_STR2(n_embd, n_tokens); + return VARS_TO_STR3(n_embd, n_tokens, identity); } - test_dsv4_hc_post(int64_t n_embd = 31, int64_t n_tokens = 17) - : n_embd(n_embd), n_tokens(n_tokens) {} + test_dsv4_hc_post(int64_t n_embd = 31, int64_t n_tokens = 17, bool identity = false) + : n_embd(n_embd), n_tokens(n_tokens), identity(identity) {} ggml_tensor * build_graph(ggml_context * ctx) override { ggml_tensor * x = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, n_embd, n_tokens); @@ -3910,8 +3933,11 @@ struct test_dsv4_hc_post : public test_dsv4_hc { ggml_tensor * post = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hc, n_tokens); ggml_set_name(post, "post"); - ggml_tensor * comb = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, hc, hc, n_tokens); - ggml_set_name(comb, "comb"); + ggml_tensor * comb = nullptr; + if (!identity) { + comb = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, hc, hc, n_tokens); + ggml_set_name(comb, "comb"); + } out = ggml_dsv4_hc_post(ctx, x, residual, post, comb); ggml_set_name(out, "out"); @@ -3943,6 +3969,24 @@ struct test_ssm_conv : public test_case { } }; +// GGML_OP_SSM_CONV + GGML_OP_MUL +struct test_ssm_conv_mul : public test_ssm_conv { + using test_ssm_conv::test_ssm_conv; + + std::string op_desc(ggml_tensor * t) override { + GGML_UNUSED(t); + return "SSM_CONV_MUL"; + } + + ggml_tensor * build_graph(ggml_context * ctx) override { + ggml_tensor * a = ggml_new_tensor(ctx, type, 4, ne_a.data()); + ggml_tensor * b = ggml_new_tensor(ctx, type, 4, ne_b.data()); + ggml_tensor * out = ggml_ssm_conv(ctx, a, b); + ggml_tensor * rhs = ggml_new_tensor_3d(ctx, type, out->ne[0], out->ne[1], out->ne[2]); + return ggml_mul(ctx, out, rhs); + } +}; + // GGML_OP_SSM_CONV + GGML_OP_ADD (channel-wise bias, optional) + GGML_OP_UNARY(SILU) (fused operation) struct test_ssm_conv_bias_silu : public test_case { const ggml_type type; @@ -7986,6 +8030,7 @@ static const ggml_type all_types[] = { GGML_TYPE_Q8_0, GGML_TYPE_Q1_0, GGML_TYPE_Q2_0, + GGML_TYPE_PQ2_0, GGML_TYPE_PTQ1_0, // PrismML group-128 types GGML_TYPE_MXFP4, GGML_TYPE_NVFP4, GGML_TYPE_Q2_K, GGML_TYPE_Q3_K, GGML_TYPE_Q4_K, GGML_TYPE_Q5_K, @@ -8014,6 +8059,7 @@ static const ggml_type other_types[] = { GGML_TYPE_Q8_0, GGML_TYPE_Q1_0, GGML_TYPE_Q2_0, + GGML_TYPE_PQ2_0, GGML_TYPE_PTQ1_0, // PrismML group-128 types GGML_TYPE_Q2_K, GGML_TYPE_Q3_K, GGML_TYPE_Q5_K, GGML_TYPE_Q6_K, @@ -8074,10 +8120,15 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_dsv4_hc_pre(31, 17)); test_cases.emplace_back(new test_dsv4_hc_pre(128, 257)); test_cases.emplace_back(new test_dsv4_hc_pre(4096, 21)); + test_cases.emplace_back(new test_dsv4_hc_pre(31, 17, true)); + test_cases.emplace_back(new test_dsv4_hc_pre(4096, 21, true)); test_cases.emplace_back(new test_dsv4_hc_post(1, 1)); test_cases.emplace_back(new test_dsv4_hc_post(31, 17)); test_cases.emplace_back(new test_dsv4_hc_post(128, 257)); + test_cases.emplace_back(new test_dsv4_hc_post(4096, 21)); + test_cases.emplace_back(new test_dsv4_hc_post(31, 17, true)); + test_cases.emplace_back(new test_dsv4_hc_post(4096, 21, true)); // glu ops for (ggml_type type : {GGML_TYPE_F16, GGML_TYPE_F32}) { @@ -8723,6 +8774,8 @@ static std::vector> make_test_cases_eval() { // in-place tests test_cases.emplace_back(new test_rms_norm(GGML_TYPE_F32, {64, 5, 4, 3}, false, 1e-6f, true)); + test_cases.emplace_back(new test_rms_norm(GGML_TYPE_F32, {64, 32, 2, 1}, false, 1e-6f)); + test_cases.emplace_back(new test_rms_norm(GGML_TYPE_F32, {64, 32, 22, 1}, false, 1e-6f)); for (float eps : { 0.0f, 1e-6f, 1e-4f, 1e-1f, 1.0f }) { for (uint32_t n : { 64, 1025 }) { @@ -8766,6 +8819,10 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_ssm_conv(GGML_TYPE_F32, {d_conv - 1 + 64, d_inner, 4, 1}, {d_conv, d_inner, 1, 1})); } } + test_cases.emplace_back(new test_ssm_conv(GGML_TYPE_F32, {4, 2048, 1, 1}, {3, 2048, 1, 1})); + test_cases.emplace_back(new test_ssm_conv(GGML_TYPE_F32, {24, 2048, 1, 1}, {3, 2048, 1, 1})); + test_cases.emplace_back(new test_ssm_conv_mul(GGML_TYPE_F32, {4, 2048, 1, 1}, {3, 2048, 1, 1})); + test_cases.emplace_back(new test_ssm_conv_mul(GGML_TYPE_F32, {24, 2048, 1, 1}, {3, 2048, 1, 1})); // fused ssm_conv + (optional) bias_add + silu. The bias-only graph (no silu) is intentionally // not tested since there's no fusion for that pattern in ggml_cuda_can_fuse. @@ -9315,6 +9372,10 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_concat(GGML_TYPE_I64, {11, 12, 13, 14}, 7, dim, v)); } } + for (int v : { 0, 1, 2, 3 }) { + test_cases.emplace_back(new test_concat(GGML_TYPE_F32, {64, 40, 2, 1}, 32, 0, v)); + test_cases.emplace_back(new test_concat(GGML_TYPE_F32, {2, 2048, 1, 1}, 22, 0, v)); + } for (ggml_type type_a : { GGML_TYPE_Q4_0, GGML_TYPE_Q4_1, GGML_TYPE_Q5_0, GGML_TYPE_Q5_1, GGML_TYPE_Q8_0 }) { for (int v : { 0, 4, 8, 12 }) { @@ -9373,16 +9434,16 @@ static std::vector> make_test_cases_eval() { // test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, {i, 2, 1, 3}, rand() % i + 1)); //} - for (ggml_scale_mode mode : {GGML_SCALE_MODE_NEAREST, GGML_SCALE_MODE_BILINEAR, GGML_SCALE_MODE_BICUBIC, ggml_scale_mode(GGML_SCALE_MODE_BILINEAR | GGML_SCALE_FLAG_ANTIALIAS)}) { + for (ggml_scale_mode mode : {GGML_SCALE_MODE_NEAREST, GGML_SCALE_MODE_BILINEAR, GGML_SCALE_MODE_BICUBIC, ggml_scale_mode(GGML_SCALE_MODE_BILINEAR | (int) GGML_SCALE_FLAG_ANTIALIAS)}) { test_cases.emplace_back(new test_upscale(GGML_TYPE_F32, {512, 512, 3, 2}, 2, mode)); test_cases.emplace_back(new test_upscale(GGML_TYPE_F32, {512, 512, 3, 2}, 2, mode, true)); test_cases.emplace_back(new test_interpolate(GGML_TYPE_F32, {2, 5, 7, 11}, {5, 7, 11, 13}, mode)); test_cases.emplace_back(new test_interpolate(GGML_TYPE_F32, {5, 7, 11, 13}, {2, 5, 7, 11}, mode)); } for (ggml_scale_mode mode : {GGML_SCALE_MODE_BILINEAR, GGML_SCALE_MODE_BICUBIC}) { - test_cases.emplace_back(new test_interpolate(GGML_TYPE_F32, {2, 5, 7, 11}, {5, 7, 11, 13}, (ggml_scale_mode)(mode | GGML_SCALE_FLAG_ALIGN_CORNERS))); - test_cases.emplace_back(new test_interpolate(GGML_TYPE_F32, {1, 4, 3, 2}, {2, 8, 3, 2}, (ggml_scale_mode)(mode | GGML_SCALE_FLAG_ALIGN_CORNERS))); - test_cases.emplace_back(new test_interpolate(GGML_TYPE_F32, {4, 1, 3, 2}, {1, 1, 3, 2}, (ggml_scale_mode)(mode | GGML_SCALE_FLAG_ALIGN_CORNERS))); + test_cases.emplace_back(new test_interpolate(GGML_TYPE_F32, {2, 5, 7, 11}, {5, 7, 11, 13}, (ggml_scale_mode)(mode | (int) GGML_SCALE_FLAG_ALIGN_CORNERS))); + test_cases.emplace_back(new test_interpolate(GGML_TYPE_F32, {1, 4, 3, 2}, {2, 8, 3, 2}, (ggml_scale_mode)(mode | (int) GGML_SCALE_FLAG_ALIGN_CORNERS))); + test_cases.emplace_back(new test_interpolate(GGML_TYPE_F32, {4, 1, 3, 2}, {1, 1, 3, 2}, (ggml_scale_mode)(mode | (int) GGML_SCALE_FLAG_ALIGN_CORNERS))); } test_cases.emplace_back(new test_sum()); @@ -9643,6 +9704,9 @@ static std::vector> make_test_cases_eval() { } } + // one-token Q4_K/Q5_K gate/up + SwiGLU at the HRX HIP kernel sizes (output % 128 == 0) + for (ggml_type t : {GGML_TYPE_Q4_K, GGML_TYPE_Q5_K}) { for (int64_t n : {128, 384}) { for (int64_t k : {256, 1024, 2560}) { test_cases.emplace_back(new test_mul_mat_vec_fusion(t, GGML_GLU_OP_SWIGLU, 1, n, k, false, 1, 1, false, false, true, false, {1, 1})); } } } + for (auto gate : {GATING_FUNC_SOFTMAX, GATING_FUNC_SIGMOID, GATING_FUNC_SOFTMAX_WEIGHT, GATING_FUNC_SQRT_SOFTPLUS}) { for (bool with_norm : {false, true}) { for (bool bias_probs : {false, true}) { diff --git a/tests/test-hrx-attention-sink.cpp b/tests/test-hrx-attention-sink.cpp new file mode 100644 index 000000000000..702dea88672c --- /dev/null +++ b/tests/test-hrx-attention-sink.cpp @@ -0,0 +1,75 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// The attention-sink dispatch rewrites FlashAttention's output in place, so it must refuse a node whose output +// shares storage with one of its inputs, including a different value (view) of the same allocation. + +#include "dispatch_registration/common/dispatch-attention-sink.h" +#include "ggml.h" +#include "graph/graph.h" +#include "graph/op-params.h" + +#include +#include +#include + +#define REQUIRE(condition) \ + do { \ + if (!(condition)) { \ + std::fprintf(stderr, "%s:%d: requirement failed: %s\n", __FILE__, __LINE__, #condition); \ + std::abort(); \ + } \ + } while (false) + +static bool append_for(bool alias_output_onto_query) { + ggml_init_params params = { 16 * ggml_tensor_overhead(), nullptr, true }; + ggml_context * ctx = ggml_init(params); + const int64_t d = 64, tokens = 8, heads = 8, kv_heads = 2, keys = 32; // tokens == heads: output and query share a layout + ggml_tensor * q = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, d, tokens, heads); + ggml_tensor * k = ggml_new_tensor_3d(ctx, GGML_TYPE_F16, d, keys, kv_heads); + ggml_tensor * v = ggml_new_tensor_3d(ctx, GGML_TYPE_F16, d, keys, kv_heads); + ggml_tensor * mask = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, keys, tokens); + ggml_tensor * sinks = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, heads); + ggml_tensor * fa = ggml_flash_attn_ext(ctx, q, k, v, mask, 0.125f, 0.0f, 0.0f); + ggml_flash_attn_ext_add_sinks(fa, sinks); + + ggml::hrx::Graph graph; + std::vector inputs; + for (ggml_tensor * source : fa->src) { + if (source != nullptr) { + inputs.push_back(graph.values().get_or_add_tensor_value(source, ggml::hrx::ValueKind::External)); + } + } + const ggml::hrx::ValueId output = graph.values().get_or_add_tensor_value(fa, ggml::hrx::ValueKind::Transient); + if (alias_output_onto_query) { + REQUIRE(graph.values().alias_storage(output, inputs[0]).success()); + REQUIRE(graph.values().same_storage(output, inputs[0])); + REQUIRE(output.value != inputs[0].value); + } + ggml::hrx::GraphNode & node = graph.add_node(fa->op, output, inputs); + node.params = ggml::hrx::import_op_params(*fa); + REQUIRE(ggml::hrx::attention_sinks_supported(graph, node)); + ggml::hrx::DispatchMatch match; + const bool appended = ggml::hrx::append_attention_sink_dispatch(graph, node, match); + ggml_free(ctx); + return appended; +} + +int main() { + REQUIRE(append_for(false)); + REQUIRE(!append_for(true)); + std::printf("test-hrx-attention-sink: disjoint output accepted, output aliasing the query refused\n"); + return 0; +} diff --git a/tests/test-hrx-buffer.cpp b/tests/test-hrx-buffer.cpp new file mode 100644 index 000000000000..e281a3199f09 --- /dev/null +++ b/tests/test-hrx-buffer.cpp @@ -0,0 +1,400 @@ +#include "backend-buffer-binding.h" +#include "backend-context.h" +#include "ggml-backend-impl.h" +#include "ggml-backend.h" +#include "ggml-hrx.h" +#include "ggml.h" +#include "hrx-interop-utils.h" +#include "runtime/host-memory.h" +#include "testing_suite.h" + +#include +#include +#include +#include +#include +#include + +#define REQUIRE(condition) \ + do { \ + if (!(condition)) { \ + std::fprintf(stderr, "%s:%d: requirement failed: %s\n", __FILE__, __LINE__, #condition); \ + std::abort(); \ + } \ + } while (false) + +static void require_hrx_status(hrx_status_t status) { + if (ggml::hrx::ErrorResult error = ggml::hrx::take_status(status)) { + std::fprintf(stderr, "HRX status failed: %s\n", error->c_str()); + std::abort(); + } +} + +static ggml_backend_hrx_context * backend_context(ggml_backend_t backend) { + auto * context = static_cast(backend->context); + REQUIRE(context != nullptr); + REQUIRE(context->device != nullptr); + REQUIRE(context->device->device != nullptr); + REQUIRE(context->stream != nullptr); + return context; +} + +static void run_backend_buffer_checks(ggml_backend_t backend) { + ggml_backend_hrx_context * hrx = backend_context(backend); + + ggml_init_params params = {}; + params.mem_size = 16 * 1024; + params.no_alloc = true; + ggml_context * context = ggml_init(params); + REQUIRE(context != nullptr); + ggml_tensor * tensor = ggml_new_tensor_1d(context, GGML_TYPE_I32, 64); + ggml_tensor * copy = ggml_new_tensor_1d(context, GGML_TYPE_I32, 64); + ggml_backend_buffer_t buffer = ggml_backend_alloc_buffer(backend, 4096); + ggml_backend_buffer_t copy_buffer = ggml_backend_alloc_buffer(backend, 4096); + REQUIRE(buffer != nullptr); + REQUIRE(copy_buffer != nullptr); + tensor->buffer = buffer; + tensor->data = ggml_backend_buffer_get_base(buffer); + copy->buffer = copy_buffer; + copy->data = ggml_backend_buffer_get_base(copy_buffer); + REQUIRE(ggml_backend_buffer_init_tensor(buffer, tensor) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_buffer_init_tensor(copy_buffer, copy) == GGML_STATUS_SUCCESS); + + std::array input = {}; + for (size_t i = 0; i < input.size(); ++i) { + input[i] = static_cast(i * 17 + 3); + } + ggml_backend_tensor_set(tensor, input.data(), 0, sizeof(input)); + std::array output = {}; + ggml_backend_tensor_get(tensor, output.data(), 0, sizeof(output)); + REQUIRE(output == input); + ggml_backend_tensor_copy(tensor, copy); + output.fill(0); + ggml_backend_tensor_get(copy, output.data(), 0, sizeof(output)); + REQUIRE(output == input); + + input[0] = 0x12345678; + ggml_backend_tensor_set_async(backend, tensor, input.data(), 0, sizeof(input)); + ggml_backend_synchronize(backend); + output.fill(0); + ggml_backend_tensor_get_async(backend, tensor, output.data(), 0, sizeof(output)); + ggml_backend_synchronize(backend); + REQUIRE(output == input); + REQUIRE(hrx->device->synchronous_upload_fallbacks.load(std::memory_order_relaxed) == 1); + REQUIRE(hrx->device->synchronous_download_fallbacks.load(std::memory_order_relaxed) == 1); + + ggml_backend_tensor_memset(tensor, 0x5a, 16, 32); + ggml_backend_tensor_get(tensor, output.data(), 0, sizeof(output)); + const uint8_t * bytes = reinterpret_cast(output.data()); + for (size_t i = 16; i < 48; ++i) { + REQUIRE(bytes[i] == 0x5a); + } + + ggml_backend_buffer_clear(buffer, 0); + ggml_backend_tensor_get(tensor, output.data(), 0, sizeof(output)); + for (uint32_t value : output) { + REQUIRE(value == 0); + } + + ggml_backend_buffer_free(buffer); + ggml_backend_buffer_free(copy_buffer); + ggml_free(context); + ggml_backend_synchronize(backend); + REQUIRE(hrx->device->device != nullptr); +} + +static void run_host_buffer_checks(ggml_backend_t backend) { + ggml_backend_hrx_context * context = backend_context(backend); + ggml_backend_buffer_type_t buft = ggml_backend_dev_host_buffer_type(ggml_backend_get_device(backend)); + REQUIRE(buft != nullptr); + REQUIRE(ggml_backend_buft_is_host(buft)); + + const bool original_direct_host_bindings = context->device->use_direct_host_bindings; + context->device->use_direct_host_bindings = false; + ggml_backend_buffer_t buffer = ggml_backend_buft_alloc_buffer(buft, 4096); + context->device->use_direct_host_bindings = original_direct_host_bindings; + REQUIRE(buffer != nullptr); + REQUIRE(ggml_backend_buffer_is_host(buffer)); + auto * buffer_context = ggml_backend_hrx_buffer_context_from_buffer(buffer); + REQUIRE(buffer_context != nullptr); + REQUIRE(buffer_context->buffer != nullptr); + REQUIRE(buffer_context->base == ggml_backend_buffer_get_base(buffer)); + REQUIRE(!buffer_context->direct_host_binding); + + const uint32_t pattern = 0x12345678; + require_hrx_status( + hrx_stream_fill_buffer(context->stream, buffer_context->buffer, 0, 4096, &pattern, sizeof(pattern))); + require_hrx_status(hrx_stream_synchronize(context->stream)); + const auto * words = static_cast(ggml_backend_buffer_get_base(buffer)); + for (size_t i = 0; i < 4096 / sizeof(uint32_t); ++i) { + REQUIRE(words[i] == pattern); + } + + ggml_init_params params = {}; + params.mem_size = 4096; + params.no_alloc = true; + ggml_context * ggml = ggml_init(params); + REQUIRE(ggml != nullptr); + ggml_tensor * host_tensor = ggml_new_tensor_1d(ggml, GGML_TYPE_I32, 64); + host_tensor->buffer = buffer; + host_tensor->data = ggml_backend_buffer_get_base(buffer); + REQUIRE(ggml_backend_buffer_init_tensor(buffer, host_tensor) == GGML_STATUS_SUCCESS); + ggml::hrx::ValueBufferBinding staged_binding; + REQUIRE(ggml_backend_hrx_resolve_value_buffer(host_tensor, staged_binding)); + REQUIRE(staged_binding.buffer == nullptr); + REQUIRE(staged_binding.host_data == ggml_backend_buffer_get_base(buffer)); + REQUIRE(staged_binding.offset == 0); + REQUIRE(staged_binding.length == ggml_nbytes(host_tensor)); + + context->device->use_direct_host_bindings = true; + ggml_backend_buffer_t direct_buffer = ggml_backend_buft_alloc_buffer(buft, 4096); + context->device->use_direct_host_bindings = original_direct_host_bindings; + REQUIRE(direct_buffer != nullptr); + auto * direct_buffer_context = ggml_backend_hrx_buffer_context_from_buffer(direct_buffer); + REQUIRE(direct_buffer_context->direct_host_binding); + ggml_tensor * direct_tensor = ggml_new_tensor_1d(ggml, GGML_TYPE_I32, 64); + direct_tensor->buffer = direct_buffer; + direct_tensor->data = ggml_backend_buffer_get_base(direct_buffer); + REQUIRE(ggml_backend_buffer_init_tensor(direct_buffer, direct_tensor) == GGML_STATUS_SUCCESS); + ggml::hrx::ValueBufferBinding direct_binding; + REQUIRE(ggml_backend_hrx_resolve_value_buffer(direct_tensor, direct_binding)); + REQUIRE(direct_binding.buffer == direct_buffer_context->buffer); + REQUIRE(direct_binding.host_data == nullptr); + REQUIRE(direct_binding.offset == 0); + REQUIRE(direct_binding.length == ggml_nbytes(direct_tensor)); + + ggml_tensor * tensor = ggml_new_tensor_1d(ggml, GGML_TYPE_I32, 64); + ggml_backend_buffer_t local = ggml_backend_alloc_buffer(backend, 4096); + REQUIRE(local != nullptr); + tensor->buffer = local; + tensor->data = ggml_backend_buffer_get_base(local); + REQUIRE(ggml_backend_buffer_init_tensor(local, tensor) == GGML_STATUS_SUCCESS); + + const uint64_t upload_fallbacks = context->device->synchronous_upload_fallbacks.load(std::memory_order_relaxed); + const uint64_t download_fallbacks = context->device->synchronous_download_fallbacks.load(std::memory_order_relaxed); + auto * host_words = static_cast(ggml_backend_buffer_get_base(buffer)); + for (size_t i = 0; i < 64; ++i) { + host_words[i] = static_cast(i * 13 + 7); + } + ggml_backend_tensor_set_async(backend, tensor, host_words, 0, 64 * sizeof(uint32_t)); + ggml_backend_synchronize(backend); + std::memset(host_words, 0, 64 * sizeof(uint32_t)); + ggml_backend_tensor_get_async(backend, tensor, host_words, 0, 64 * sizeof(uint32_t)); + ggml_backend_synchronize(backend); + for (size_t i = 0; i < 64; ++i) { + REQUIRE(host_words[i] == static_cast(i * 13 + 7)); + } + REQUIRE(context->device->synchronous_upload_fallbacks.load(std::memory_order_relaxed) == upload_fallbacks); + REQUIRE(context->device->synchronous_download_fallbacks.load(std::memory_order_relaxed) == download_fallbacks); + + ggml_backend_buffer_free(local); + ggml_backend_buffer_free(direct_buffer); + ggml_backend_buffer_free(buffer); + ggml_free(ggml); +} + +static void run_host_transfer_checks(ggml_backend_hrx_context * context) { + ggml::hrx::HostTransferManager transfers; + ggml::hrx::HostStagingBuffer staging; + REQUIRE(ggml::hrx::allocate_host_staging_buffer(context->device->device, 64, staging).success()); + + const std::array zero = {}; + require_hrx_status(hrx_synchronous_h2d(context->device->device, zero.data(), staging.buffer, 0, zero.size())); + + std::array host = {}; + for (size_t i = 0; i < host.size(); ++i) { + host[i] = static_cast(i + 1); + } + + REQUIRE(transfers.upload_synchronous(context->stream, host.data(), staging.buffer, 0, 0).success()); + REQUIRE(transfers.upload_synchronous(context->stream, host.data() + 8, staging.buffer, 16, 24).success()); + ggml::hrx::HostTransferStats stats = transfers.stats(); + REQUIRE(stats.uploads == 1); + REQUIRE(stats.upload_bytes == 24); + require_hrx_status(hrx_stream_synchronize(context->stream)); + + std::array upload_result = {}; + require_hrx_status( + hrx_synchronous_d2h(context->device->device, staging.buffer, 0, upload_result.data(), upload_result.size())); + for (size_t i = 0; i < upload_result.size(); ++i) { + const uint8_t expected = i >= 16 && i < 40 ? host[i - 8] : 0; + REQUIRE(upload_result[i] == expected); + } + + std::array device_values = {}; + for (size_t i = 0; i < device_values.size(); ++i) { + device_values[i] = static_cast(0xa0 + i); + } + require_hrx_status( + hrx_synchronous_h2d(context->device->device, device_values.data(), staging.buffer, 0, device_values.size())); + + std::array download_result = {}; + REQUIRE( + transfers.download_synchronous(context->stream, staging.buffer, 12, download_result.data() + 4, 20).success()); + stats = transfers.stats(); + REQUIRE(stats.downloads == 1); + REQUIRE(stats.download_bytes == 20); + require_hrx_status(hrx_stream_synchronize(context->stream)); + for (size_t i = 0; i < download_result.size(); ++i) { + const uint8_t expected = i >= 4 && i < 24 ? device_values[i + 8] : 0; + REQUIRE(download_result[i] == expected); + } + + REQUIRE(!transfers.upload_synchronous(nullptr, host.data(), staging.buffer, 0, 4).success()); + REQUIRE(!transfers.upload_synchronous(context->stream, nullptr, staging.buffer, 0, 4).success()); + REQUIRE(!transfers.upload_synchronous(context->stream, host.data(), nullptr, 0, 4).success()); + REQUIRE(!transfers.download_synchronous(nullptr, staging.buffer, 0, download_result.data(), 4).success()); + REQUIRE(!transfers.download_synchronous(context->stream, nullptr, 0, download_result.data(), 4).success()); + REQUIRE(!transfers.download_synchronous(context->stream, staging.buffer, 0, nullptr, 4).success()); + + transfers.clear(); + stats = transfers.stats(); + REQUIRE(stats.uploads == 0); + REQUIRE(stats.downloads == 0); + REQUIRE(stats.upload_bytes == 0); + REQUIRE(stats.download_bytes == 0); +} + +static void run_host_staging_checks(ggml_backend_hrx_context * context) { + ggml::hrx::HostStagingBuffer staging; + REQUIRE(ggml::hrx::allocate_host_staging_buffer(context->device->device, 32, staging).success()); + REQUIRE(staging.buffer != nullptr); + REQUIRE(staging.length == 32); + + hrx_buffer_t original = staging.buffer; + ggml::hrx::HostStagingBuffer moved(std::move(staging)); + REQUIRE(moved.buffer == original); + REQUIRE(moved.length == 32); + REQUIRE(staging.buffer == nullptr); + REQUIRE(staging.length == 0); + + ggml::hrx::HostStagingBuffer assigned; + assigned = std::move(moved); + REQUIRE(assigned.buffer == original); + REQUIRE(assigned.length == 32); + REQUIRE(moved.buffer == nullptr); + REQUIRE(moved.length == 0); + + assigned.clear(); + REQUIRE(assigned.buffer == nullptr); + REQUIRE(assigned.length == 0); + assigned.clear(); + REQUIRE(assigned.buffer == nullptr); +} + +static void run_host_weight_cache_checks(ggml_backend_hrx_context * context) { + ggml::hrx::HostTransferManager transfers; + ggml::hrx::HostWeightCache weights; + + std::array host = {}; + for (size_t i = 0; i < host.size(); ++i) { + host[i] = static_cast(i); + } + + ggml::hrx::HostWeightSource source; + source.host_data = host.data(); + source.identity = 0x1234; + source.generation = 1; + source.capacity = host.size(); + source.offset = 16; + source.length = 32; + source.layout = "ggml-native"; + + ggml::hrx::HostWeightAcquireResult first = + weights.acquire(context->device->device, context->stream, transfers, source); + REQUIRE(first.valid()); + + std::array first_bytes = {}; + require_hrx_status( + hrx_synchronous_d2h(context->device->device, first.lease.buffer(), 0, first_bytes.data(), first_bytes.size())); + for (size_t i = 0; i < first_bytes.size(); ++i) { + REQUIRE(first_bytes[i] == host[source.offset + i]); + } + + ggml::hrx::HostWeightAcquireResult second = + weights.acquire(context->device->device, context->stream, transfers, source); + REQUIRE(second.valid()); + REQUIRE(second.lease.buffer() == first.lease.buffer()); + + source.offset = 32; + ggml::hrx::HostWeightAcquireResult slice = + weights.acquire(context->device->device, context->stream, transfers, source); + REQUIRE(slice.valid()); + REQUIRE(slice.lease.buffer() != first.lease.buffer()); + + source.offset = 16; + source.generation = 2; + ggml::hrx::HostWeightAcquireResult next_generation = + weights.acquire(context->device->device, context->stream, transfers, source); + REQUIRE(next_generation.valid()); + REQUIRE(next_generation.lease.buffer() != first.lease.buffer()); + + source.generation = 1; + source.layout = "alternate-layout"; + ggml::hrx::HostWeightAcquireResult conflict = + weights.acquire(context->device->device, context->stream, transfers, source); + REQUIRE(!conflict.valid()); + + ggml::hrx::HostWeightCacheStats weight_stats = weights.stats(); + REQUIRE(weight_stats.hits == 1); + REQUIRE(weight_stats.misses == 3); + REQUIRE(weight_stats.allocation_count == 3); + REQUIRE(weight_stats.resident_bytes == 96); + + const ggml::hrx::HostTransferStats transfer_stats = transfers.stats(); + REQUIRE(transfer_stats.uploads == 3); + REQUIRE(transfer_stats.upload_bytes == 96); + + weights.clear(); + weight_stats = weights.stats(); + REQUIRE(weight_stats.hits == 0); + REQUIRE(weight_stats.misses == 0); + REQUIRE(weight_stats.allocation_count == 0); + REQUIRE(weight_stats.resident_bytes == 0); +} + +template +static void run_with_hrx_backend(F && f) { + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + ggml_backend_hrx_context * context = backend_context(backend); + f(backend, context); + ggml_backend_free(backend); +} + +static void register_hrx_buffer_cases(test_runner::Suite & suite) { + suite.device_case("backend_buffer", [] { + run_with_hrx_backend([](ggml_backend_t backend, ggml_backend_hrx_context *) { + run_backend_buffer_checks(backend); + }); + }); + suite.device_case("host_buffer", [] { + run_with_hrx_backend([](ggml_backend_t backend, ggml_backend_hrx_context *) { + run_host_buffer_checks(backend); + }); + }); + suite.device_case("host_transfer", [] { + run_with_hrx_backend([](ggml_backend_t, ggml_backend_hrx_context * context) { + run_host_transfer_checks(context); + }); + }); + suite.device_case("host_staging", [] { + run_with_hrx_backend([](ggml_backend_t, ggml_backend_hrx_context * context) { + run_host_staging_checks(context); + }); + }); + suite.device_case("host_weight_cache", [] { + run_with_hrx_backend([](ggml_backend_t, ggml_backend_hrx_context * context) { + run_host_weight_cache_checks(context); + }); + }); +} + +int main(int argc, char ** argv) { + test_runner::Suite suite(test_runner::Config::with_prefix( + "HRX buffer test", "hrx-buffer", "GGML_HRX_BUFFER_TEST", 1)); + register_hrx_buffer_cases(suite); + + const bool has_device = ggml_backend_hrx_get_device_count() != 0; + return suite.run(argc, argv, has_device); +} diff --git a/tests/test-hrx-decode-stride.cpp b/tests/test-hrx-decode-stride.cpp new file mode 100644 index 000000000000..6cd8ec9b25ef --- /dev/null +++ b/tests/test-hrx-decode-stride.cpp @@ -0,0 +1,227 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Decode-kernel input rows for input sizes that are a multiple of 32 but not of 256 (gpt-oss 2880, BlackMamba 1152; +// 2048 as the control): every token / input row after the first must be read input_size values after the previous +// one. MUL_MAT with 2..8 tokens and MUL_MAT_ID with 1 or 4 input rows per token (the down projection reads one row +// per route) on MXFP4, Q8_0 and Q4_0 weights, against the CPU backend (normalized MSE, as test-backend-ops). Each +// graph runs on the HRX device itself (no scheduler) and must be planned on HRX. + +#include "dispatch/dispatch-scheduler.h" +#include "ggml-alloc.h" +#include "ggml-backend.h" +#include "ggml-cpu.h" +#include "ggml.h" +#include "graph/graph.h" +#include "kernel-corpus/kernel-corpus.h" + +#include +#include +#include +#include +#include +#include +#include + +#define REQUIRE(condition) \ + do { \ + if (!(condition)) { \ + std::fprintf(stderr, "%s:%d: requirement failed: %s\n", __FILE__, __LINE__, #condition); \ + std::abort(); \ + } \ + } while (false) + +static constexpr int64_t kOutputSize = 96; +static constexpr int64_t kExpertCount = 32; +static constexpr int64_t kRouteCount = 4; +static constexpr double kMaxNmse = 5e-4; + +static std::string kernel_name_for_id(uint64_t kernel_id) { + const ggml::hrx::KernelResolveResult resolved = + ggml::hrx::resolve_kernel_definition(ggml::hrx::get_qwen_kernel_corpus(), "gfx1151", kernel_id); + REQUIRE(resolved.found()); + return ggml::hrx::kernel_definition_name(*resolved.definition); +} + +// The HRX kernels planned for the graph, or "" when HRX does not take it. +static std::string hrx_plan(ggml_cgraph * graph) { + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + ggml::hrx::DispatchScheduleDiagnostics diagnostics; + if (!scheduler.schedule_graph(imported.graph, { "gfx1151" }, &diagnostics) || !scheduler.plan().valid()) { + return ""; + } + std::string names; + for (const ggml::hrx::Dispatch & dispatch : scheduler.plan().dispatches) { + names += (names.empty() ? "" : "+") + kernel_name_for_id(dispatch.kernel.kernel_id); + } + return names; +} + +struct Case { + ggml_type type; + int64_t input_size; + int64_t tokens; + int64_t input_rows; // 0: MUL_MAT; 1 or kRouteCount: MUL_MAT_ID +}; + +struct Data { + std::vector weights; + std::vector x; + std::vector ids; +}; + +static Data make_data(const Case & c, uint32_t seed) { + std::mt19937 rng(seed); + std::uniform_real_distribution uniform(-1.0f, 1.0f); + const int64_t experts = c.input_rows == 0 ? 1 : kExpertCount; + const int64_t rows = kOutputSize * experts; + std::vector w(static_cast(rows * c.input_size)); + for (float & v : w) { + v = uniform(rng); + } + Data d; + d.weights.resize(ggml_row_size(c.type, c.input_size) * rows); + ggml_quantize_chunk(c.type, w.data(), d.weights.data(), 0, rows, c.input_size, nullptr); + const int64_t x_rows = c.input_rows == 0 ? 1 : c.input_rows; + d.x.resize(static_cast(c.input_size * x_rows * c.tokens)); + for (float & v : d.x) { + v = uniform(rng); + } + if (c.input_rows != 0) { + d.ids.resize(static_cast(kExpertCount * c.tokens)); + for (int64_t t = 0; t < c.tokens; ++t) { + std::vector order(kExpertCount); + for (int64_t e = 0; e < kExpertCount; ++e) { + order[e] = static_cast(e); + } + std::shuffle(order.begin(), order.end(), rng); + std::copy(order.begin(), order.end(), d.ids.begin() + t * kExpertCount); + } + } + return d; +} + +// Runs the case on a backend; returns false when HRX (plan != nullptr) does not take the graph. +static bool run(ggml_backend_t backend, const Case & c, const Data & d, std::vector & result, std::string * plan) { + ggml_init_params params = { 32 * ggml_tensor_overhead() + ggml_graph_overhead(), nullptr, true }; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + ggml_tensor * weights = nullptr; + ggml_tensor * x = nullptr; + ggml_tensor * ids_all = nullptr; + ggml_tensor * out = nullptr; + if (c.input_rows == 0) { + weights = ggml_new_tensor_2d(ctx, c.type, c.input_size, kOutputSize); + x = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, c.input_size, c.tokens); + out = ggml_mul_mat(ctx, weights, x); + } else { + weights = ggml_new_tensor_3d(ctx, c.type, c.input_size, kOutputSize, kExpertCount); + x = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, c.input_size, c.input_rows, c.tokens); + ids_all = ggml_new_tensor_2d(ctx, GGML_TYPE_I32, kExpertCount, c.tokens); + ggml_tensor * ids = ggml_view_2d(ctx, ids_all, kRouteCount, c.tokens, ids_all->nb[1], 0); + out = ggml_mul_mat_id(ctx, weights, x, ids); + } + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, out); + if (plan != nullptr) { + *plan = ggml_backend_supports_op(backend, out) ? hrx_plan(graph) : ""; + if (plan->empty()) { + ggml_free(ctx); + return false; + } + } + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + ggml_backend_tensor_set(weights, d.weights.data(), 0, d.weights.size()); + ggml_backend_tensor_set(x, d.x.data(), 0, d.x.size() * sizeof(float)); + if (ids_all != nullptr) { + ggml_backend_tensor_set(ids_all, d.ids.data(), 0, d.ids.size() * sizeof(int32_t)); + } + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + result.resize(static_cast(ggml_nelements(out))); + ggml_backend_tensor_get(out, result.data(), 0, result.size() * sizeof(float)); + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + return true; +} + +static double nmse(const std::vector & got, const std::vector & expected) { + double err = 0.0; + double ref = 0.0; + for (size_t i = 0; i < got.size(); ++i) { + const double diff = static_cast(got[i]) - expected[i]; + err += diff * diff; + ref += static_cast(expected[i]) * expected[i]; + } + return ref > 0.0 ? err / ref : err; +} + +int main() { + ggml_backend_dev_t device = ggml_backend_dev_by_name("HRX0"); + if (device == nullptr) { + ggml_backend_load_all(); + device = ggml_backend_dev_by_name("HRX0"); + } + if (device == nullptr) { + std::printf("test-hrx-decode-stride: no HRX0 device, skipped\n"); + return 0; + } + ggml_backend_t hrx = ggml_backend_dev_init(device, nullptr); + ggml_backend_t cpu = ggml_backend_cpu_init(); + REQUIRE(hrx != nullptr && cpu != nullptr); + + std::vector cases; + for (ggml_type type : { GGML_TYPE_MXFP4, GGML_TYPE_Q8_0, GGML_TYPE_Q4_0 }) { + for (int64_t input_size : { 2880, 1152, 2048 }) { + for (int64_t tokens : { 2, 3, 4, 8 }) { + cases.push_back({ type, input_size, tokens, 0 }); + } + for (int64_t input_rows : { int64_t(1), kRouteCount }) { + for (int64_t tokens : { 1, 2 }) { + cases.push_back({ type, input_size, tokens, input_rows }); + } + } + } + } + int failures = 0; + uint32_t seed = 1; + for (const Case & c : cases) { + const Data d = make_data(c, seed++); + std::string plan; + std::vector got; + std::vector expected; + const char * op = c.input_rows == 0 ? "MUL_MAT" : "MUL_MAT_ID"; + if (!run(hrx, c, d, got, &plan)) { + std::printf("%-10s %-5s k=%4lld tokens=%lld rows=%lld not admitted by HRX\n", op, ggml_type_name(c.type), + (long long) c.input_size, (long long) c.tokens, (long long) c.input_rows); + ++failures; + continue; + } + REQUIRE(run(cpu, c, d, expected, nullptr)); + const double e = nmse(got, expected); + const bool ok = e <= kMaxNmse; + std::printf("%-10s %-5s k=%4lld tokens=%lld rows=%lld %-44s nmse=%.3g %s\n", op, ggml_type_name(c.type), + (long long) c.input_size, (long long) c.tokens, (long long) c.input_rows, plan.c_str(), e, + ok ? "OK" : "FAIL"); + failures += ok ? 0 : 1; + } + ggml_backend_free(cpu); + ggml_backend_free(hrx); + std::printf("test-hrx-decode-stride: %zu cases, %d failures\n", cases.size(), failures); + REQUIRE(failures == 0); + return 0; +} diff --git a/tests/test-hrx-fa-masked-v.cpp b/tests/test-hrx-fa-masked-v.cpp new file mode 100644 index 000000000000..96b7f4e24606 --- /dev/null +++ b/tests/test-hrx-fa-masked-v.cpp @@ -0,0 +1,255 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// FLASH_ATTN_EXT on HRX must not depend on K/V rows that the mask hides. A KV cache keeps whatever an earlier request +// (or a rejected draft) wrote past the current sequence end, so a dependence there makes identical requests give +// different logits. Each case runs the same attention several times on the HRX device, changing only the masked K/V +// rows (random values, +x, -x, +0, -0), and requires bitwise-identical outputs; one run is also checked against the +// CPU backend (normalized MSE). Cases cover the decode-split kernel (1..15 query rows) and the prefill kernel +// (16+ rows), with the sequence end inside a 64-key block and in the final partial block. + +#include "dispatch/dispatch-scheduler.h" +#include "ggml-alloc.h" +#include "ggml-backend.h" +#include "ggml-cpu.h" +#include "ggml.h" +#include "graph/graph.h" +#include "kernel-corpus/kernel-corpus.h" + +#include +#include +#include +#include +#include +#include +#include +#include + +#define REQUIRE(condition) \ + do { \ + if (!(condition)) { \ + std::fprintf(stderr, "%s:%d: requirement failed: %s\n", __FILE__, __LINE__, #condition); \ + std::abort(); \ + } \ + } while (false) + +static constexpr int64_t kHeadSize = 128; +static constexpr int64_t kQueryHeadCount = 16; +static constexpr int64_t kKeyValueHeadCount = 8; +static constexpr double kMaxNmse = 5e-4; + +static std::string kernel_name_for_id(uint64_t kernel_id) { + const ggml::hrx::KernelResolveResult resolved = + ggml::hrx::resolve_kernel_definition(ggml::hrx::get_qwen_kernel_corpus(), "gfx1151", kernel_id); + REQUIRE(resolved.found()); + return ggml::hrx::kernel_definition_name(*resolved.definition); +} + +// The HRX kernels planned for the graph, or "" when HRX does not take it. +static std::string hrx_plan(ggml_cgraph * graph) { + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + ggml::hrx::DispatchScheduleDiagnostics diagnostics; + if (!scheduler.schedule_graph(imported.graph, { "gfx1151" }, &diagnostics) || !scheduler.plan().valid()) { + return ""; + } + std::string names; + for (const ggml::hrx::Dispatch & dispatch : scheduler.plan().dispatches) { + const std::string name = kernel_name_for_id(dispatch.kernel.kernel_id); + if (names.find(name) == std::string::npos) { + names += (names.empty() ? "" : "+") + name; + } + } + return names; +} + +struct Case { + int64_t queries; // query rows in the batch + int64_t visible; // cells 0 .. visible-1 hold the sequence; the last query row sees all of them + int64_t cells; // KV cells in the graph (the mask width) +}; + +// Graph like llama.cpp builds it: K/V cache views [head, cell, kv head], Q permuted to [head, token, head]. +struct Graph { + ggml_context * ctx = nullptr; + ggml_cgraph * graph = nullptr; + ggml_tensor * q = nullptr; + ggml_tensor * k = nullptr; + ggml_tensor * v = nullptr; + ggml_tensor * mask = nullptr; + ggml_tensor * out = nullptr; + ggml_backend_buffer_t buffer = nullptr; +}; + +static Graph build(ggml_backend_t backend, const Case & c) { + Graph g; + ggml_init_params params = { 32 * ggml_tensor_overhead() + ggml_graph_overhead(), nullptr, true }; + g.ctx = ggml_init(params); + REQUIRE(g.ctx != nullptr); + g.q = ggml_new_tensor_3d(g.ctx, GGML_TYPE_F32, kHeadSize, kQueryHeadCount, c.queries); + g.k = ggml_new_tensor_3d(g.ctx, GGML_TYPE_F16, kHeadSize, kKeyValueHeadCount, c.cells); + g.v = ggml_new_tensor_3d(g.ctx, GGML_TYPE_F16, kHeadSize, kKeyValueHeadCount, c.cells); + g.mask = ggml_new_tensor_2d(g.ctx, GGML_TYPE_F16, c.cells, c.queries); + ggml_tensor * k_view = ggml_view_3d(g.ctx, g.k, kHeadSize, c.cells, kKeyValueHeadCount, g.k->nb[2], g.k->nb[1], 0); + ggml_tensor * v_view = ggml_view_3d(g.ctx, g.v, kHeadSize, c.cells, kKeyValueHeadCount, g.v->nb[2], g.v->nb[1], 0); + ggml_tensor * q_permuted = ggml_permute(g.ctx, g.q, 0, 2, 1, 3); + ggml_tensor * attention = ggml_flash_attn_ext(g.ctx, q_permuted, k_view, v_view, g.mask, + 1.0f / std::sqrt(static_cast(kHeadSize)), 0.0f, 0.0f); + ggml_flash_attn_ext_set_prec(attention, GGML_PREC_F32); + g.out = ggml_reshape_2d(g.ctx, attention, kHeadSize * kQueryHeadCount, c.queries); + g.graph = ggml_new_graph(g.ctx); + ggml_build_forward_expand(g.graph, g.out); + g.buffer = ggml_backend_alloc_ctx_tensors(g.ctx, backend); + REQUIRE(g.buffer != nullptr); + return g; +} + +static void release(Graph & g) { + ggml_backend_buffer_free(g.buffer); + ggml_free(g.ctx); +} + +static std::vector compute(ggml_backend_t backend, Graph & g, const std::vector & q, + const std::vector & k, const std::vector & v, + const std::vector & mask) { + ggml_backend_tensor_set(g.q, q.data(), 0, ggml_nbytes(g.q)); + ggml_backend_tensor_set(g.k, k.data(), 0, ggml_nbytes(g.k)); + ggml_backend_tensor_set(g.v, v.data(), 0, ggml_nbytes(g.v)); + ggml_backend_tensor_set(g.mask, mask.data(), 0, ggml_nbytes(g.mask)); + REQUIRE(ggml_backend_graph_compute(backend, g.graph) == GGML_STATUS_SUCCESS); + std::vector result(static_cast(ggml_nelements(g.out))); + ggml_backend_tensor_get(g.out, result.data(), 0, result.size() * sizeof(float)); + return result; +} + +static double nmse(const std::vector & got, const std::vector & expected) { + double err = 0.0; + double ref = 0.0; + for (size_t i = 0; i < got.size(); ++i) { + const double diff = static_cast(got[i]) - expected[i]; + err += diff * diff; + ref += static_cast(expected[i]) * expected[i]; + } + return ref > 0.0 ? err / ref : err; +} + +int main() { + ggml_backend_dev_t device = ggml_backend_dev_by_name("HRX0"); + if (device == nullptr) { + ggml_backend_load_all(); + device = ggml_backend_dev_by_name("HRX0"); + } + if (device == nullptr) { + std::printf("test-hrx-fa-masked-v: no HRX0 device, skipped\n"); + return 0; + } + ggml_backend_t hrx = ggml_backend_dev_init(device, nullptr); + ggml_backend_t cpu = ggml_backend_cpu_init(); + REQUIRE(hrx != nullptr && cpu != nullptr); + + std::vector cases; + for (int64_t seed_case = 0; seed_case < 8; ++seed_case) { + cases.push_back({ 1 + seed_case % 15, 6 + 7 * seed_case, 256 }); // decode split, end in block 0 + cases.push_back({ 1 + (3 * seed_case) % 15, 70 + 11 * seed_case, 256 }); // decode split, end in block 1 + cases.push_back({ 16 + 9 * seed_case, 30 + 13 * seed_case, 256 }); // prefill, end inside a block + } + cases.push_back({ 1, 150, 200 }); // decode split, partial last block + cases.push_back({ 24, 150, 200 }); // prefill, partial last block (tail path) + cases.push_back({ 40, 40, 512 }); // prefill, fresh prompt in a long cache + + const size_t kv_row = static_cast(kHeadSize * kKeyValueHeadCount); + int failures = 0; + for (size_t ci = 0; ci < cases.size(); ++ci) { + const Case & c = cases[ci]; + REQUIRE(c.visible >= c.queries && c.visible <= c.cells); + std::mt19937 rng(static_cast(1000 + ci)); + std::normal_distribution normal(0.0f, 1.0f); + std::vector q(static_cast(kHeadSize * kQueryHeadCount * c.queries)); + for (float & x : q) { + x = 2.0f * normal(rng); + } + std::vector k(kv_row * c.cells); + std::vector v(kv_row * c.cells); + for (size_t i = 0; i < k.size(); ++i) { + k[i] = ggml_fp32_to_fp16(normal(rng)); + v[i] = ggml_fp32_to_fp16(normal(rng)); + } + // Causal mask: query row r sits at position visible - queries + r. + std::vector mask(static_cast(c.cells * c.queries)); + for (int64_t r = 0; r < c.queries; ++r) { + const int64_t position = c.visible - c.queries + r; + for (int64_t cell = 0; cell < c.cells; ++cell) { + mask[r * c.cells + cell] = ggml_fp32_to_fp16(cell <= position ? 0.0f : -INFINITY); + } + } + + Graph g = build(hrx, c); + const std::string plan = ggml_backend_supports_op(hrx, g.out->src[0]) ? hrx_plan(g.graph) : ""; + if (plan.empty()) { + std::printf("FAIL case %zu (q=%lld visible=%lld cells=%lld): not planned on HRX\n", ci, + static_cast(c.queries), static_cast(c.visible), + static_cast(c.cells)); + ++failures; + release(g); + continue; + } + const std::vector reference = compute(hrx, g, q, k, v, mask); + + // Rewrite only the rows past the sequence end, five ways. + int mismatches = 0; + for (int variant = 0; variant < 5; ++variant) { + std::vector k2 = k; + std::vector v2 = v; + std::mt19937 stale_rng(static_cast(77 + variant)); + for (size_t i = kv_row * c.visible; i < k2.size(); ++i) { + float kv = 0.0f; + float vv = 0.0f; + switch (variant) { + case 0: kv = 3.0f * normal(stale_rng); vv = 3.0f * normal(stale_rng); break; + case 1: kv = 30.0f; vv = 1.0f; break; + case 2: kv = -30.0f; vv = -1.0f; break; + case 3: kv = 0.0f; vv = 0.0f; break; + default: kv = -0.0f; vv = -0.0f; break; + } + k2[i] = ggml_fp32_to_fp16(kv); + v2[i] = ggml_fp32_to_fp16(vv); + } + const std::vector got = compute(hrx, g, q, k2, v2, mask); + if (std::memcmp(got.data(), reference.data(), got.size() * sizeof(float)) != 0) { + ++mismatches; + } + } + release(g); + + Graph gc = build(cpu, c); + const std::vector expected = compute(cpu, gc, q, k, v, mask); + release(gc); + const double error = nmse(reference, expected); + const bool ok = mismatches == 0 && error <= kMaxNmse; + failures += ok ? 0 : 1; + std::printf("%s case %zu q=%lld visible=%lld cells=%lld (%s): masked-row variants differing %d/5, nmse %.3g\n", + ok ? "ok " : "FAIL", ci, static_cast(c.queries), static_cast(c.visible), + static_cast(c.cells), plan.c_str(), mismatches, error); + } + ggml_backend_free(hrx); + ggml_backend_free(cpu); + if (failures != 0) { + std::printf("test-hrx-fa-masked-v: %d of %zu cases failed\n", failures, cases.size()); + return 1; + } + std::printf("test-hrx-fa-masked-v: %zu cases passed\n", cases.size()); + return 0; +} diff --git a/tests/test-hrx-hadamard.cpp b/tests/test-hrx-hadamard.cpp new file mode 100644 index 000000000000..52df4985d335 --- /dev/null +++ b/tests/test-hrx-hadamard.cpp @@ -0,0 +1,148 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Known-answer test for a MUL_MAT hinted GGML_HINT_SRC0_IS_HADAMARD on HRX0 (ops/hadamard_f32.loom +// through dispatch-hadamard.cpp). Each shape is computed on HRX0 directly (no scheduler, so it +// cannot fall back to the CPU) and on the CPU backend (its fwht path for the same hint), and for +// small blocks also against a double-precision dense product. Every shape runs four times and the +// HRX results must be bitwise equal across runs (races hide in single runs). + +#include "ggml-alloc.h" +#include "ggml-backend.h" +#include "ggml-cpu.h" +#include "ggml.h" + +#include +#include +#include +#include +#include +#include + +namespace { + +struct Shape { + int64_t n; + int64_t rows; +}; + +std::vector sylvester(int64_t n) { + std::vector h(n * n); + const float scale = 1.0f / std::sqrt((float) n); + for (int64_t r = 0; r < n; ++r) { + for (int64_t c = 0; c < n; ++c) { + h[r * n + c] = __builtin_parityll(r & c) ? -scale : scale; + } + } + return h; +} + +bool run(ggml_backend_t backend, const Shape & s, const std::vector & rot, const std::vector & x, + std::vector & out) { + ggml_init_params params = { 3 * ggml_tensor_overhead() + ggml_graph_overhead(), nullptr, true }; + ggml_context * ctx = ggml_init(params); + ggml_tensor * a = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, s.n, s.n); + ggml_tensor * b = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, s.n, s.rows); + ggml_tensor * y = ggml_mul_mat(ctx, a, b); + ggml_mul_mat_set_hint(y, GGML_HINT_SRC0_IS_HADAMARD); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, y); + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + bool ok = buffer != nullptr; + if (ok) { + ggml_backend_tensor_set(a, rot.data(), 0, rot.size() * sizeof(float)); + ggml_backend_tensor_set(b, x.data(), 0, x.size() * sizeof(float)); + ok = ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS; + if (ok) { + out.resize(s.n * s.rows); + ggml_backend_tensor_get(y, out.data(), 0, out.size() * sizeof(float)); + } + ggml_backend_buffer_free(buffer); + } + ggml_free(ctx); + return ok; +} + +double max_rel_error(const std::vector & got, const std::vector & want) { + double worst = 0.0; + for (size_t i = 0; i < got.size(); ++i) { + worst = std::max(worst, std::fabs(got[i] - want[i]) / std::max(1.0, std::fabs(want[i]))); + } + return worst; +} + +} // namespace + +int main() { + ggml_backend_dev_t dev = ggml_backend_dev_by_name("HRX0"); + if (dev == nullptr) { + printf("test-hrx-hadamard: no HRX0 device, skipped\n"); + return 0; + } + ggml_backend_t hrx = ggml_backend_dev_init(dev, nullptr); + ggml_backend_t cpu = ggml_backend_cpu_init(); + // decode-sized and prompt-sized rows, including Bonsai's 1024-block rotation at 512 tokens + const Shape shapes[] = { { 64, 1 }, { 64, 7 }, { 1024, 1 }, { 1024, 17 }, { 1024, 8704 }, { 4096, 3 } }; + std::mt19937 rng(20261002); + std::uniform_real_distribution dist(-1.0f, 1.0f); + int failures = 0; + for (const Shape & s : shapes) { + const std::vector rot = sylvester(s.n); + std::vector x(s.n * s.rows); + for (float & v : x) { + v = dist(rng); + } + std::vector want_cpu; + if (!run(cpu, s, rot, x, want_cpu)) { + printf("n=%lld rows=%lld: CPU compute failed\n", (long long) s.n, (long long) s.rows); + ++failures; + continue; + } + std::vector want(want_cpu.begin(), want_cpu.end()); + if (s.n <= 1024 && s.rows <= 17) { // dense double-precision product + for (int64_t r = 0; r < s.rows; ++r) { + for (int64_t j = 0; j < s.n; ++j) { + double sum = 0.0; + for (int64_t i = 0; i < s.n; ++i) { + sum += (double) rot[j * s.n + i] * x[r * s.n + i]; + } + want[r * s.n + j] = sum; + } + } + } + std::vector first; + for (int rep = 0; rep < 4; ++rep) { + std::vector got; + if (!run(hrx, s, rot, x, got)) { + printf("n=%lld rows=%lld rep %d: HRX0 compute failed\n", (long long) s.n, (long long) s.rows, rep); + ++failures; + break; + } + const double err = max_rel_error(got, want); + const bool same = rep == 0 || std::memcmp(got.data(), first.data(), got.size() * sizeof(float)) == 0; + if (rep == 0) { + first = got; + } + const bool ok = err <= 1e-5 && same; + printf("n=%-5lld rows=%-5lld rep %d: max rel error %.2e%s %s\n", (long long) s.n, (long long) s.rows, rep, + err, same ? "" : " (differs from rep 0)", ok ? "ok" : "FAIL"); + failures += ok ? 0 : 1; + } + } + ggml_backend_free(hrx); + ggml_backend_free(cpu); + printf(failures == 0 ? "test-hrx-hadamard: PASS\n" : "test-hrx-hadamard: FAIL (%d)\n", failures); + return failures == 0 ? 0 : 1; +} diff --git a/tests/test-hrx-loom-jit.cpp b/tests/test-hrx-loom-jit.cpp new file mode 100644 index 000000000000..ee5272eaef95 --- /dev/null +++ b/tests/test-hrx-loom-jit.cpp @@ -0,0 +1,1562 @@ +#include "dispatch/dispatch.h" +#include "hrx-interop-utils.h" +#include "kernel-corpus/kernel-corpus.h" +#include "runtime/kernel-executable-cache.h" +#include "runtime/loom-kernel-jit.h" +#include "testing_suite.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#define REQUIRE(condition) \ + do { \ + if (!(condition)) { \ + std::fprintf(stderr, "%s:%d: requirement failed: %s\n", __FILE__, __LINE__, #condition); \ + std::abort(); \ + } \ + } while (false) + +namespace { + +class HrxTestDevice { + public: + ~HrxTestDevice() { + if (device != nullptr) { + hrx_device_release(device); + } + } + + bool open() { + hrx_status_t init_status = hrx_gpu_initialize(0); + if (!hrx_status_is_ok(init_status)) { + if (hrx_status_code(init_status) == HRX_STATUS_ALREADY_EXISTS) { + hrx_status_ignore(init_status); + } else { + hrx_status_ignore(init_status); + return false; + } + } + + int count = 0; + if (ggml::hrx::ErrorResult error = ggml::hrx::take_status(hrx_gpu_device_count(&count))) { + std::fprintf(stderr, "skip executable cache materialization: device count failed: %s\n", error->c_str()); + return false; + } + for (int i = 0; i < count; ++i) { + hrx_device_t candidate = nullptr; + if (ggml::hrx::ErrorResult error = ggml::hrx::take_status(hrx_gpu_device_get(i, &candidate))) { + std::fprintf(stderr, "skip HRX device %d: %s\n", i, error->c_str()); + continue; + } + if (candidate == nullptr) { + continue; + } + hrx_device_retain(candidate); + std::optional candidate_architecture = + read_string_property(candidate, HRX_DEVICE_PROPERTY_ARCHITECTURE); + if (!candidate_architecture) { + hrx_device_release(candidate); + continue; + } + device = candidate; + architecture = *candidate_architecture; + return true; + } + return false; + } + + hrx_device_t device = nullptr; + std::string architecture; + + private: + static std::optional read_string_property(hrx_device_t device, hrx_device_property_t property) { + std::vector buffer(64); + while (buffer.size() <= 4096) { + hrx_status_t status = hrx_device_get_property(device, property, buffer.data(), buffer.size()); + if (hrx_status_is_ok(status)) { + return std::string(buffer.data()); + } + if (hrx_status_code(status) != HRX_STATUS_OUT_OF_RANGE) { + if (ggml::hrx::ErrorResult error = ggml::hrx::take_status(status)) { + std::fprintf(stderr, "HRX property query failed: %s\n", error->c_str()); + } + return std::nullopt; + } + hrx_status_ignore(status); + buffer.resize(buffer.size() * 2); + } + return std::nullopt; + } +}; + +static ggml_hrx_loom_jit_source_format to_jit_source_format(ggml::hrx::KernelSourceFormat format) { + switch (format) { + case ggml::hrx::KERNEL_SOURCE_FORMAT_TEXT: + return GGML_HRX_LOOM_JIT_SOURCE_FORMAT_TEXT; + case ggml::hrx::KERNEL_SOURCE_FORMAT_BINARY: + return GGML_HRX_LOOM_JIT_SOURCE_FORMAT_BYTECODE; + } + return GGML_HRX_LOOM_JIT_SOURCE_FORMAT_TEXT; +} + +static const ggml::hrx::KernelDefinition & find_kernel(const char * name) { + const ggml::hrx::KernelCorpus & corpus = ggml::hrx::get_qwen_kernel_corpus(); + for (const ggml::hrx::KernelDefinition & kernel : corpus.kernels) { + if (std::strcmp(kernel.name, name) == 0) { + return kernel; + } + } + std::fprintf(stderr, "missing test kernel: %s\n", name); + std::abort(); +} + +static const ggml::hrx::KernelDefinition * find_targeted_kernel(const char * name, const char * target) { + const ggml::hrx::KernelCorpus & corpus = ggml::hrx::get_qwen_kernel_corpus(); + for (const ggml::hrx::KernelDefinition & kernel : corpus.kernels) { + if (std::strcmp(kernel.name, name) == 0 && std::strcmp(kernel.target_selector, target) == 0) { + return &kernel; + } + } + return nullptr; +} + +static std::map binary_f32_exact_config(const char * op, int64_t element_count) { + const std::string count = std::to_string(element_count); + return { + { "ggml.binary_f32.op", op }, + { "ggml.binary_f32.ne0", count }, + { "ggml.binary_f32.ne1", "1" }, + { "ggml.binary_f32.ne2", "1" }, + { "ggml.binary_f32.src0_stride1", count }, + { "ggml.binary_f32.src0_stride2", count }, + { "ggml.binary_f32.src0_stride3", count }, + { "ggml.binary_f32.src1_stride1", count }, + { "ggml.binary_f32.src1_stride2", count }, + { "ggml.binary_f32.src1_stride3", count }, + { "ggml.binary_f32.src0_span", count }, + { "ggml.binary_f32.src1_span", count }, + }; +} + +static ggml::hrx::Dispatch make_binary_add_dispatch(const ggml::hrx::KernelDefinition & definition, + int64_t element_count) { + ggml::hrx::Dispatch dispatch; + dispatch.kernel.kernel_id = definition.id; + dispatch.kernel.integer_parameters.emplace("element_count", element_count); + dispatch.kernel.compile_parameters = binary_f32_exact_config("0", element_count); + dispatch.bindings.reserve(definition.bindings.size()); + for (size_t i = 0; i < definition.bindings.size(); ++i) { + dispatch.bindings.push_back({ + ggml::hrx::ValueId(static_cast(i + 1)), + 0, + static_cast(element_count) * sizeof(float), + }); + } + return dispatch; +} + +static std::map binary_f32_exact_workload(int64_t element_count) { + return { + { "element_count", element_count }, + }; +} + +static std::map binary_bc_f32_workload() { + return { + { "element_count", 48 }, + { "ne0", 16 }, + { "ne1", 3 }, + { "ne2", 1 }, + { "ne3", 1 }, + { "src0_element_count", 48 }, + { "src1_element_count", 16 }, + }; +} + +static std::map binary_bc_f32_config(const char * op) { + return { + { "ggml.binary_bc_f32.op", op }, + { "ggml.binary_bc_f32.src0_broadcast_dim0", "0" }, + { "ggml.binary_bc_f32.src0_broadcast_dim1", "0" }, + { "ggml.binary_bc_f32.src0_broadcast_dim2", "0" }, + { "ggml.binary_bc_f32.src0_broadcast_dim3", "0" }, + { "ggml.binary_bc_f32.src1_broadcast_dim0", "0" }, + { "ggml.binary_bc_f32.src1_broadcast_dim1", "1" }, + { "ggml.binary_bc_f32.src1_broadcast_dim2", "0" }, + { "ggml.binary_bc_f32.src1_broadcast_dim3", "0" }, + }; +} + +static std::map token_count_workload(int64_t token_count) { + return { + { "token_count", token_count }, + }; +} + +static std::map moe_routing_config() { + return { + { "ggml.moe_routing.route_count", "8" }, + { "ggml.moe_routing.expert_count", "128" }, + { "ggml.moe_routing.descriptor_expert_mask", "127" }, + { "ggml.moe_routing.descriptor_partition_shift", "7" }, + { "ggml.moe_routing.descriptor_row_count_shift", "13" }, + { "ggml.moe_routing.partition_workgroup_size", "128" }, + { "ggml.workload.token_capacity", "1" }, + }; +} + +static std::map flash_attention_workload(int64_t query_token_count, + int64_t key_value_token_count) { + return { + { "query_token_count", query_token_count }, + { "key_value_token_count", key_value_token_count }, + }; +} + +static std::map flash_attention_decode_workload(int64_t key_value_token_count) { + return { + { "key_value_token_count", key_value_token_count }, + }; +} + +static std::map flash_attention_config(const char * query_head_count, + const char * key_value_head_count, + const char * head_size, + const char * attention_scale) { + return { + { "ggml.flash_attention.query_head_count", query_head_count }, + { "ggml.flash_attention.key_value_head_count", key_value_head_count }, + { "ggml.flash_attention.qk_head_size", head_size }, + { "ggml.flash_attention.value_head_size", head_size }, + { "ggml.flash_attention.attention_scale", attention_scale }, + { "ggml.flash_attention.apply_gate", "0" }, + { "ggml.flash_attention.gate_stride_head", "1" }, + { "ggml.flash_attention.gate_stride_token", "1" }, + }; +} + +static std::map flash_attention_decode_base_config(const char * query_head_count, + const char * key_value_head_count, + const char * head_size, + const char * attention_scale) { + return { + { "ggml.flash_attention.query_head_count", query_head_count }, + { "ggml.flash_attention.key_value_head_count", key_value_head_count }, + { "ggml.flash_attention.qk_head_size", head_size }, + { "ggml.flash_attention.value_head_size", head_size }, + { "ggml.flash_attention.attention_scale", attention_scale }, + { "ggml.flash_attention.apply_gate", "0" }, + { "ggml.flash_attention.gate_stride_head", head_size }, + { "ggml.flash_attention.gate_stride_token", head_size }, + }; +} + +static std::map flash_attention_decode_config(const char * query_head_count, + const char * key_value_head_count, + const char * head_size, + const char * attention_scale, + const char * key_value_token_capacity) { + std::map config = + flash_attention_decode_base_config(query_head_count, key_value_head_count, head_size, attention_scale); + config["ggml.flash_attention.decode.key_value_token_capacity"] = key_value_token_capacity; + return config; +} + +static std::map mul_mat_id_f16_f16_config(const char * weight_format) { + return { + { "ggml.mul_mat_id_f16_f16.input_size", "768" }, + { "ggml.mul_mat_id_f16_f16.route_count", "8" }, + { "ggml.mul_mat_id_f16_f16.expert_count", "128" }, + { "ggml.mul_mat_id_f16_f16.output_size", "2048" }, + { "ggml.mul_mat_id_f16_f16.weight_format", weight_format }, + { "ggml.workload.token_capacity", "4" }, + }; +} + +static std::map mul_mat_id_swiglu_f16_f16_config(const char * gate_format, + const char * up_format) { + return { + { "ggml.mul_mat_id_swiglu_f16_f16.input_size", "2048" }, + { "ggml.mul_mat_id_swiglu_f16_f16.output_size", "768" }, + { "ggml.mul_mat_id_swiglu_f16_f16.expert_count", "128" }, + { "ggml.mul_mat_id_swiglu_f16_f16.route_count", "8" }, + { "ggml.mul_mat_id_swiglu_f16_f16.gate_weight_format", gate_format }, + { "ggml.mul_mat_id_swiglu_f16_f16.up_weight_format", up_format }, + { "ggml.mul_mat_id_swiglu_f16_f16.descriptor_expert_mask", "127" }, + { "ggml.mul_mat_id_swiglu_f16_f16.descriptor_partition_shift", "7" }, + { "ggml.mul_mat_id_swiglu_f16_f16.descriptor_row_count_shift", "13" }, + }; +} + +static ggml::hrx::LoomKernelCompileRequest make_compile_request( + const ggml::hrx::KernelDefinition & definition, + const std::map & workload, + const std::map & compile_config = {}) { + ggml::hrx::LoomKernelCompileRequest request; + REQUIRE(!definition.compile_recipe.primary_sources.empty()); + + const ggml::hrx::KernelSourceRef & primary_source = definition.compile_recipe.primary_sources.front(); + REQUIRE(primary_source.contents != nullptr); + request.source_data = primary_source.contents->source.data; + request.source_size = primary_source.contents->source.length; + request.source_format = to_jit_source_format(primary_source.contents->source.format); + request.source_identifier = primary_source.path != nullptr ? primary_source.path : ""; + request.symbol = definition.symbol != nullptr ? definition.symbol : ""; + request.launch_config_symbol = definition.name != nullptr ? definition.name : ""; + + request.dependencies.reserve(definition.compile_recipe.library_sources.size()); + for (const ggml::hrx::KernelSourceRef & dependency_ref : definition.compile_recipe.library_sources) { + REQUIRE(dependency_ref.contents != nullptr); + request.dependencies.push_back({ + dependency_ref.contents->source.data, + dependency_ref.contents->source.length, + to_jit_source_format(dependency_ref.contents->source.format), + dependency_ref.path, + }); + } + + std::map merged_config; + for (const ggml::hrx::KernelCompileConfig & config : definition.compile_config) { + merged_config[config.key != nullptr ? config.key : ""] = config.value != nullptr ? config.value : ""; + } + for (const auto & item : compile_config) { + merged_config[item.first] = item.second; + } + request.config_storage.reserve(merged_config.size()); + for (const auto & item : merged_config) { + request.config_storage.push_back(item); + } + + request.workload.reserve(definition.workload_parameters.size()); + for (const ggml::hrx::KernelScalarDefinition & parameter : definition.workload_parameters) { + REQUIRE(parameter.type != nullptr); + REQUIRE(std::strcmp(parameter.type, "index") == 0); + const auto found = workload.find(parameter.name != nullptr ? parameter.name : ""); + REQUIRE(found != workload.end()); + request.workload.push_back(found->second); + } + return request; +} + +static ggml::hrx::LoomCompiledKernelRef compile_kernel(ggml::hrx::LoomJit & jit, + const std::string & key, + const ggml::hrx::KernelDefinition & definition, + const std::map & workload, + const std::map & compile_config = {}) { + return jit.compile(key, make_compile_request(definition, workload, compile_config)); +} + +static bool resolve_with_timeout(const ggml::hrx::LoomCompiledKernelRef & ref, std::chrono::seconds timeout) { + std::atomic done = false; + std::atomic success = false; + std::thread resolver([&] { + success.store(ref->resolve(), std::memory_order_release); + done.store(true, std::memory_order_release); + }); + + const auto deadline = std::chrono::steady_clock::now() + timeout; + while (!done.load(std::memory_order_acquire)) { + if (std::chrono::steady_clock::now() >= deadline) { + std::fprintf(stderr, "timed out waiting for async Loom compile: %s\n", ref->key().c_str()); + std::abort(); + } + std::this_thread::sleep_for(std::chrono::milliseconds(10)); + } + resolver.join(); + return success.load(std::memory_order_acquire); +} + +static void require_compiled_kernel(const ggml::hrx::LoomCompiledKernelRef & ref) { + REQUIRE(ref != nullptr); + REQUIRE(resolve_with_timeout(ref, std::chrono::seconds(120))); + ggml_hrx_loom_jit_compile_result compiled = ref->take_result(); + REQUIRE(compiled.hsaco_data != nullptr); + REQUIRE(compiled.hsaco_size > 0); + REQUIRE(compiled.launch_config.workgroup_count[0] > 0); + REQUIRE(compiled.launch_config.workgroup_size[0] > 0); + compiled.reset(); +} + +static int64_t elapsed_us(std::chrono::steady_clock::time_point begin, std::chrono::steady_clock::time_point end) { + return std::chrono::duration_cast(end - begin).count(); +} + +static void run_cache_materialize_case(HrxTestDevice & device, + ggml::hrx::LoomJitMode mode, + const ggml::hrx::KernelDefinition & binary) { + ggml::hrx::KernelExecutablePrepareContext context = {}; + context.device = device.device; + context.target = device.architecture.c_str(); + + ggml::hrx::Dispatch dispatch = make_binary_add_dispatch(binary, 64); + std::vector constants; + ggml_hrx_loom_jit_launch_config launch; + ggml::hrx::KernelExecutableCache cache(mode); + + ggml::hrx::KernelExecutableRef ref = cache.get_or_compile(context, binary, dispatch); + REQUIRE(ref.valid()); + + std::shared_ptr executable = + cache.materialize(context, ref, dispatch.kernel, constants, launch); + REQUIRE(executable != nullptr); + REQUIRE(executable->executable != nullptr); + REQUIRE(executable->export_info.binding_count == dispatch.bindings.size()); + REQUIRE(executable->export_info.constant_byte_length == constants.size()); + REQUIRE(launch.workgroup_count[0] > 0); + REQUIRE(launch.workgroup_size[0] > 0); + + std::shared_ptr loaded_hit = + cache.materialize(context, ref, dispatch.kernel, constants, launch); + REQUIRE(loaded_hit == executable); + + std::vector constants_for_prepare; + ggml_hrx_loom_jit_launch_config launch_for_prepare; + std::shared_ptr prepared = + cache.prepare(context, binary, dispatch, constants_for_prepare, launch_for_prepare); + REQUIRE(prepared == executable); + + const char * mode_name = mode == ggml::hrx::LoomJitMode::Async ? "async" : "sync"; + std::printf("%s KernelExecutableCache materialized ggml_binary_f32 for %s\n", mode_name, + device.architecture.c_str()); +} + +static void run_dynamic_cache_materialize_case(HrxTestDevice & device, + ggml::hrx::LoomJitMode mode, + const ggml::hrx::KernelDefinition & binary) { + ggml::hrx::KernelExecutablePrepareContext context = {}; + context.device = device.device; + context.target = device.architecture.c_str(); + + ggml::hrx::Dispatch first = make_binary_add_dispatch(binary, 64); + ggml::hrx::Dispatch second = make_binary_add_dispatch(binary, 257); + first.kernel.compile_parameters = binary_f32_exact_config("0", 512); + second.kernel.compile_parameters = first.kernel.compile_parameters; + first.kernel.workload_specialization = ggml::hrx::WorkloadSpecialization::Dynamic; + second.kernel.workload_specialization = ggml::hrx::WorkloadSpecialization::Dynamic; + + ggml::hrx::KernelExecutableCache cache(mode); + ggml::hrx::KernelExecutableRef first_ref = cache.get_or_compile(context, binary, first); + ggml::hrx::KernelExecutableRef second_ref = cache.get_or_compile(context, binary, second); + REQUIRE(first_ref.valid()); + REQUIRE(second_ref.valid()); + REQUIRE(first_ref.entry == second_ref.entry); + + std::vector first_constants; + std::vector second_constants; + ggml_hrx_loom_jit_launch_config first_launch; + ggml_hrx_loom_jit_launch_config second_launch; + std::shared_ptr first_executable = + cache.materialize(context, first_ref, first.kernel, first_constants, first_launch); + std::shared_ptr second_executable = + cache.materialize(context, second_ref, second.kernel, second_constants, second_launch); + REQUIRE(first_executable != nullptr); + REQUIRE(second_executable == first_executable); + REQUIRE(first_launch.workgroup_count[0] == 1); + REQUIRE(second_launch.workgroup_count[0] == 2); + REQUIRE(first_constants.size() == second_constants.size()); + REQUIRE(first_constants != second_constants); + + ggml::hrx::Dispatch exact_second = second; + exact_second.kernel.workload_specialization = ggml::hrx::WorkloadSpecialization::Exact; + ggml::hrx::KernelExecutableRef exact_ref = cache.get_or_compile(context, binary, exact_second); + REQUIRE(exact_ref.valid()); + REQUIRE(exact_ref.entry != first_ref.entry); +} + +static void run_dynamic_export_reuse_case( + HrxTestDevice & device, + const char * kernel_name, + const std::map & first_workload, + const std::map & second_workload, + const std::map & compile_config = {}) { + const ggml::hrx::KernelDefinition & definition = find_kernel(kernel_name); + ggml::hrx::KernelExecutablePrepareContext context = {}; + context.device = device.device; + context.target = device.architecture.c_str(); + + auto make_dispatch = [&](const std::map & workload) { + ggml::hrx::Dispatch dispatch; + dispatch.kernel.kernel_id = definition.id; + dispatch.kernel.integer_parameters = workload; + dispatch.kernel.compile_parameters = compile_config; + dispatch.kernel.workload_specialization = ggml::hrx::WorkloadSpecialization::Dynamic; + dispatch.bindings.resize(definition.bindings.size()); + return dispatch; + }; + + ggml::hrx::Dispatch first = make_dispatch(first_workload); + ggml::hrx::Dispatch second = make_dispatch(second_workload); + ggml::hrx::KernelExecutableCache cache(ggml::hrx::LoomJitMode::Sync); + ggml::hrx::KernelExecutableRef first_ref = cache.get_or_compile(context, definition, first); + ggml::hrx::KernelExecutableRef second_ref = cache.get_or_compile(context, definition, second); + REQUIRE(first_ref.valid()); + REQUIRE(second_ref.valid()); + REQUIRE(first_ref.entry == second_ref.entry); + + std::vector first_constants; + std::vector second_constants; + ggml_hrx_loom_jit_launch_config first_launch; + ggml_hrx_loom_jit_launch_config second_launch; + std::shared_ptr first_executable = + cache.materialize(context, first_ref, first.kernel, first_constants, first_launch); + std::shared_ptr second_executable = + cache.materialize(context, second_ref, second.kernel, second_constants, second_launch); + REQUIRE(first_executable != nullptr); + REQUIRE(second_executable == first_executable); + REQUIRE(first_launch.workgroup_count[0] > 0); + REQUIRE(second_launch.workgroup_count[0] > 0); + REQUIRE(first_constants.size() == second_constants.size()); + REQUIRE(first_constants != second_constants); + std::printf("dynamic executable reuse materialized %s for two workloads\n", kernel_name); +} + +static void run_targeted_export_materialize_case(HrxTestDevice & device) { + const ggml::hrx::KernelDefinition * definition = + find_targeted_kernel("ggml_linear_q6k_q8_1_x4", device.architecture.c_str()); + if (definition == nullptr) { + return; + } + + ggml::hrx::KernelExecutablePrepareContext context = {}; + context.device = device.device; + context.target = device.architecture.c_str(); + + ggml::hrx::Dispatch dispatch; + dispatch.kernel.kernel_id = definition->id; + dispatch.kernel.integer_parameters.emplace("token_count", 1); + dispatch.kernel.integer_parameters.emplace("input_size", 2048); + dispatch.kernel.integer_parameters.emplace("output_size", 151936); + dispatch.kernel.compile_parameters.emplace("ggml.linear_q6k_q8_1_x4.token_capacity", "1"); + dispatch.kernel.compile_parameters.emplace("ggml.linear_q6k_q8_1_x4.output_capacity", "151936"); + dispatch.bindings.resize(3); + + ggml::hrx::KernelExecutableCache cache(ggml::hrx::LoomJitMode::Sync); + std::vector constants; + ggml_hrx_loom_jit_launch_config launch; + std::shared_ptr executable = + cache.prepare(context, *definition, dispatch, constants, launch); + REQUIRE(executable != nullptr); + REQUIRE(executable->executable != nullptr); + REQUIRE(launch.workgroup_count[0] == 151936); + std::printf("KernelExecutableCache materialized targeted %s as export %s for %s\n", definition->symbol, + definition->name, device.architecture.c_str()); +} + +} // namespace + +static void run_loom_jit_compile_checks() { + static constexpr const char * kTarget = "gfx1100"; + + const ggml::hrx::KernelDefinition & binary = find_kernel("ggml_binary_f32"); + const ggml::hrx::KernelDefinition & binary_bc = find_kernel("ggml_binary_bc_f32"); + const ggml::hrx::KernelDefinition & gather_add = find_kernel("ggml_gather_add_f32"); + const ggml::hrx::KernelDefinition & rmsnorm = find_kernel("ggml_rmsnorm_f32"); + const ggml::hrx::KernelDefinition & router_top8 = find_kernel("qwen3_moe_router_top8_f32"); + const ggml::hrx::KernelDefinition & expert_table = find_kernel("ggml_moe_build_expert_table"); + const ggml::hrx::KernelDefinition & partition_table = find_kernel("ggml_moe_build_expert_partition_table"); + const ggml::hrx::KernelDefinition & mul_mat_id_f16 = find_kernel("ggml_mul_mat_id_f16_f16_wmma"); + const ggml::hrx::KernelDefinition & swiglu_f16 = find_kernel("ggml_mul_mat_id_swiglu_f16_f16_wmma"); + const ggml::hrx::KernelDefinition & flash_prefill = find_kernel("ggml_flash_attention_f32_f16_wmma"); + const ggml::hrx::KernelDefinition & decode_mul_mat = find_kernel("ggml_mul_mat_f32_f32_decode_wave64"); + const ggml::hrx::KernelDefinition & flash_decode = + find_kernel("ggml_flash_attention_decode_split_f32_f16_wmma_next_q8"); + + std::string sync_error; + std::unique_ptr sync_jit = + ggml::hrx::create_loom_jit(kTarget, ggml::hrx::LoomJitMode::Sync, sync_error); + REQUIRE(sync_jit != nullptr); + REQUIRE(!sync_jit->async_enabled()); + + const auto sync_begin = std::chrono::steady_clock::now(); + ggml::hrx::LoomCompiledKernelRef sync_ref = compile_kernel( + *sync_jit, "sync-binary-add-64", binary, binary_f32_exact_workload(64), binary_f32_exact_config("0", 64)); + require_compiled_kernel(sync_ref); + const auto sync_end = std::chrono::steady_clock::now(); + const int64_t sync_compile_us = elapsed_us(sync_begin, sync_end); + std::printf("sync Loom compile completed in %ld us\n", static_cast(sync_compile_us)); + + std::string async_error; + std::unique_ptr async_jit = + ggml::hrx::create_loom_jit(kTarget, ggml::hrx::LoomJitMode::Async, async_error); + REQUIRE(async_jit != nullptr); + REQUIRE(async_jit->async_enabled()); + + std::vector refs; + refs.reserve(31); + const auto enqueue_begin = std::chrono::steady_clock::now(); + refs.push_back(compile_kernel(*async_jit, "async-binary-add-64", binary, binary_f32_exact_workload(64), + binary_f32_exact_config("0", 64))); + refs.push_back(compile_kernel(*async_jit, "async-binary-add-128", binary, binary_f32_exact_workload(128), + binary_f32_exact_config("0", 128))); + refs.push_back(compile_kernel(*async_jit, "async-binary-add-256", binary, binary_f32_exact_workload(256), + binary_f32_exact_config("0", 256))); + refs.push_back(compile_kernel(*async_jit, "async-binary-add-512", binary, binary_f32_exact_workload(512), + binary_f32_exact_config("0", 512))); + refs.push_back(compile_kernel(*async_jit, "async-binary-geglu-256", binary, binary_f32_exact_workload(256), + binary_f32_exact_config("5", 256))); + refs.push_back(compile_kernel(*async_jit, "async-binary-reglu-256", binary, binary_f32_exact_workload(256), + binary_f32_exact_config("6", 256))); + refs.push_back(compile_kernel(*async_jit, "async-binary-geglu-erf-256", binary, binary_f32_exact_workload(256), + binary_f32_exact_config("7", 256))); + refs.push_back(compile_kernel(*async_jit, "async-binary-geglu-quick-256", binary, binary_f32_exact_workload(256), + binary_f32_exact_config("8", 256))); + refs.push_back(compile_kernel(*async_jit, "async-binary-bc-mul-48", binary_bc, binary_bc_f32_workload(), + binary_bc_f32_config("2"))); + refs.push_back(compile_kernel(*async_jit, "async-rmsnorm-1", rmsnorm, + { + { "token_count", 1 } + }, + { + { "ggml.rmsnorm_f32.hidden_size", "2048" }, + { "ggml.rmsnorm_f32.input_stride", "2048" }, + { "ggml.rmsnorm_f32.rms_epsilon", "0.000001" }, + })); + refs.push_back(compile_kernel(*async_jit, "async-router-top8-1", router_top8, + { + { "token_count", 1 }, + { "route_id_stride", 8 }, + }, + { + { "qwen3_moe.router.expert_count", "128" }, + { "qwen3_moe.router.route_count", "8" }, + { "qwen3_moe.workload.token_capacity", "1" }, + })); + refs.push_back(compile_kernel(*async_jit, "async-expert-table-1", expert_table, + { + { "token_count", 1 }, + { "route_count", 8 }, + { "route_stride", 8 }, + { "expert_count", 128 }, + }, + moe_routing_config())); + refs.push_back(compile_kernel(*async_jit, "async-partition-table-1", partition_table, + { + { "token_count", 1 }, + { "route_count", 8 }, + { "expert_count", 128 }, + }, + moe_routing_config())); + // Exercise gather-add coverage in the async JIT path. + refs.push_back(compile_kernel(*async_jit, "async-gather-add-2-to-1", gather_add, + { + { "source_token_count", 2 }, + { "output_token_count", 1 }, + { "hidden_size", 2048 }, + }, + { + { "ggml.gather_add_f32.binary_op", "3" }, + })); + refs.push_back(compile_kernel(*async_jit, "async-mul-mat-id-f16-q4", mul_mat_id_f16, token_count_workload(4), + mul_mat_id_f16_f16_config("4"))); + refs.push_back(compile_kernel(*async_jit, "async-mul-mat-id-f16-q6", mul_mat_id_f16, token_count_workload(4), + mul_mat_id_f16_f16_config("6"))); + refs.push_back(compile_kernel(*async_jit, "async-mul-mat-id-swiglu-f16-q4k", swiglu_f16, token_count_workload(4), + mul_mat_id_swiglu_f16_f16_config("4", "4"))); + refs.push_back(compile_kernel(*async_jit, "async-mul-mat-id-swiglu-f16-q6k", swiglu_f16, token_count_workload(4), + mul_mat_id_swiglu_f16_f16_config("6", "6"))); + refs.push_back(compile_kernel(*async_jit, "async-mul-mat-id-swiglu-f16-iq4-xs", swiglu_f16, token_count_workload(4), + mul_mat_id_swiglu_f16_f16_config("23", "23"))); + refs.push_back(compile_kernel(*async_jit, "async-mul-mat-id-swiglu-f16-q3k", swiglu_f16, token_count_workload(4), + mul_mat_id_swiglu_f16_f16_config("11", "11"))); + refs.push_back(compile_kernel(*async_jit, "async-mul-mat-id-swiglu-f16-iq4-nl", swiglu_f16, token_count_workload(4), + mul_mat_id_swiglu_f16_f16_config("20", "20"))); + refs.push_back(compile_kernel(*async_jit, "async-mul-mat-id-swiglu-f16-q8-0", swiglu_f16, token_count_workload(4), + mul_mat_id_swiglu_f16_f16_config("80", "80"))); + refs.push_back(compile_kernel(*async_jit, "async-mul-mat-id-swiglu-f16-q8-1", swiglu_f16, token_count_workload(4), + mul_mat_id_swiglu_f16_f16_config("81", "81"))); + refs.push_back(compile_kernel(*async_jit, "async-mul-mat-id-swiglu-f16-f16", swiglu_f16, token_count_workload(4), + mul_mat_id_swiglu_f16_f16_config("16", "16"))); + refs.push_back(compile_kernel(*async_jit, "async-mul-mat-decode-q8-0-gemma-vocab", decode_mul_mat, + { + { "token_count", 1 }, + { "input_size", 3840 }, + { "output_size", 262208 }, + }, + { + { "ggml.mul_mat_f32_f32_decode.token_capacity", "1" }, + { "ggml.mul_mat_f32_f32_decode.output_capacity", "262208" }, + { "ggml.mul_mat_f32_f32_decode.weight_format", "80" }, + })); + refs.push_back(compile_kernel(*async_jit, "async-flash-prefill-128", flash_prefill, + flash_attention_workload(128, 128), + flash_attention_config("32", "4", "128", "0.0883883461"))); + refs.push_back(compile_kernel(*async_jit, "async-flash-prefill-256", flash_prefill, + flash_attention_workload(128, 128), + flash_attention_config("32", "4", "256", "0.0625"))); + refs.push_back(compile_kernel(*async_jit, "async-flash-decode-128", flash_decode, + flash_attention_decode_workload(64), + flash_attention_decode_config("32", "4", "128", "0.0883883461", "64"))); + const auto enqueue_end = std::chrono::steady_clock::now(); + const int64_t enqueue_us = elapsed_us(enqueue_begin, enqueue_end); + std::printf("async Loom enqueue completed in %ld us\n", static_cast(enqueue_us)); + + const int64_t max_expected_enqueue_us = std::max(100000, sync_compile_us / 2); + REQUIRE(enqueue_us < max_expected_enqueue_us); + + for (const ggml::hrx::LoomCompiledKernelRef & ref : refs) { + require_compiled_kernel(ref); + } + + std::printf("async Loom JIT compiled %zu kernels\n", refs.size()); + + const auto & q6_prefill = find_kernel("ggml_mul_mat_q6_k_f16_wmma_prefill_wave32"); + + const struct { + int64_t tokens; + int64_t inputs; + int64_t outputs; + int64_t accumulate; + int64_t workgroup_size; + int64_t token_tiles; + int64_t output_tiles; + } q6_launch_cases[] = { + { 128, 256, 4096, 0, 256, 1, 32 }, + { 256, 256, 8192, 0, 512, 1, 64 }, + { 512, 256, 1024, 0, 512, 2, 8 }, + { 512, 256, 3840, 0, 512, 2, 30 }, + { 512, 256, 4096, 0, 1024, 1, 32 }, + { 512, 256, 4224, 0, 1024, 1, 33 }, + { 512, 256, 4096, 1, 256, 4, 32 }, + {1024, 256, 2048, 0, 1024, 2, 16 }, + { 512, 17408, 5120, 0, 512, 1, 80 }, + { 512, 5120, 10240, 0, 1024, 1, 80 }, + { 512, 5120, 1024, 0, 512, 2, 8 }, + { 512, 17408, 4096, 0, 512, 1, 64 }, + { 512, 5120, 1984, 0, 256, 4, 16 }, + { 512, 5120, 2048, 0, 512, 1, 32 }, + {1024, 5120, 1024, 0, 512, 2, 16 }, + {2048, 5120, 1024, 0, 512, 4, 16 }, + }; + + for (const auto & test : q6_launch_cases) { + const std::string key = "q6-prefill-" + std::to_string(test.tokens) + "-" + std::to_string(test.inputs) + "-" + + std::to_string(test.outputs) + "-" + std::to_string(test.accumulate); + auto ref = + compile_kernel(*sync_jit, key, q6_prefill, token_count_workload(test.tokens), + { + { "ggml.mul_mat_q6_k_packed.input_size", std::to_string(test.inputs) }, + { "ggml.mul_mat_q6_k_packed.output_size", std::to_string(test.outputs) }, + { "ggml.mul_mat_q6_k_packed.output_accumulation", std::to_string(test.accumulate) }, + { "ggml.mul_mat_q6_k_packed.weight_offset", "0" }, + { "ggml.mul_mat_q6_k_packed.token_capacity", std::to_string(test.tokens) }, + }); + REQUIRE(ref != nullptr); + REQUIRE(resolve_with_timeout(ref, std::chrono::seconds(120))); + auto compiled = ref->take_result(); + REQUIRE(compiled.hsaco_data != nullptr); + REQUIRE(compiled.launch_config.workgroup_size[0] == test.workgroup_size); + REQUIRE(compiled.launch_config.workgroup_count[0] == test.token_tiles); + REQUIRE(compiled.launch_config.workgroup_count[1] == test.output_tiles); + compiled.reset(); + } + + const auto & q4_prefill = find_kernel("ggml_mul_mat_q4_k_f16_wmma_prefill_wave32"); + + const struct { + int64_t tokens; + int64_t inputs; + int64_t outputs; + int64_t workgroup_size; + int64_t token_tiles; + int64_t output_tiles; + int64_t input_layout; + } q4_launch_cases[] = { + { 256, 8192, 8192, 512, 1, 64, 0 }, + { 512, 8192, 1024, 512, 2, 8, 0 }, + { 512, 8192, 3840, 512, 2, 30, 0 }, + { 512, 8192, 4096, 512, 1, 64, 1 }, + { 512, 8192, 4224, 512, 1, 66, 1 }, + { 512, 8192, 12288, 512, 1, 192, 1 }, + { 512, 8192, 16384, 512, 1, 256, 1 }, + {1024, 8192, 2048, 1024, 4, 8, 0 }, + { 512, 17408, 5120, 1024, 2, 20, 0 }, + { 512, 5120, 6144, 512, 1, 96, 1 }, + { 512, 6144, 5120, 512, 1, 80, 1 }, + { 512, 5120, 10240, 512, 1, 160, 1 }, + { 512, 5120, 8192, 512, 1, 128, 1 }, + {1024, 5120, 4096, 1024, 4, 16, 0 }, + { 512, 17408, 5120, 512, 1, 80, 1 }, + }; + + for (const auto & test : q4_launch_cases) { + const std::string key = "q4-prefill-" + std::to_string(test.tokens) + "-" + std::to_string(test.inputs) + "-" + + std::to_string(test.outputs) + "-" + std::to_string(test.input_layout); + auto ref = compile_kernel(*sync_jit, key, q4_prefill, token_count_workload(test.tokens), + { + { "ggml.mul_mat.input_size", std::to_string(test.inputs) }, + { "ggml.mul_mat.output_size", std::to_string(test.outputs) }, + { "ggml.workload.token_capacity", std::to_string(test.tokens) }, + { "ggml.mul_mat.f16_input_layout", std::to_string(test.input_layout) }, + }); + REQUIRE(ref != nullptr); + REQUIRE(resolve_with_timeout(ref, std::chrono::seconds(120))); + auto compiled = ref->take_result(); + REQUIRE(compiled.hsaco_data != nullptr); + REQUIRE(compiled.launch_config.workgroup_size[0] == test.workgroup_size); + REQUIRE(compiled.launch_config.workgroup_count[0] == test.token_tiles); + REQUIRE(compiled.launch_config.workgroup_count[1] == test.output_tiles); + compiled.reset(); + } + + const auto & packed_f16 = find_kernel("ggml_copy_f16_k16_major"); + for (const auto * input_is_f16 : { "0", "1" }) { + auto ref = + compile_kernel(*sync_jit, std::string("packed-f16-") + input_is_f16, packed_f16, token_count_workload(512), + { + { "ggml.copy_f16_k16_major.input_size", "256" }, + { "ggml.copy_f16_k16_major.token_count", "512" }, + { "ggml.copy_f16_k16_major.input_is_f16", input_is_f16 }, + }); + require_compiled_kernel(ref); + } + + const auto & transpose_f16 = find_kernel("ggml_copy_transpose_f16"); + auto transpose_ref = compile_kernel(*sync_jit, "transpose-f16-8192x1024", transpose_f16, + { + { "row_count", 8192 }, + { "column_count", 1024 }, + }, + { + }); + require_compiled_kernel(transpose_ref); + auto transposed_attention_config = flash_attention_config("24", "4", "256", "0.0625"); + transposed_attention_config.emplace("ggml.flash_attention.value_layout", "1"); + auto transposed_attention_ref = + compile_kernel(*sync_jit, "flash-transposed-value", flash_prefill, + { + { "query_token_count", 512 }, + { "key_value_token_count", 8192 } + }, + transposed_attention_config); + require_compiled_kernel(transposed_attention_ref); + + const auto & paired_q4 = find_kernel("ggml_mul_mat_swiglu_q4_k_f16_wmma_prefill_wave32"); + + const struct { + int64_t tokens; + int64_t outputs; + int64_t threads; + int64_t token_tiles; + } paired_q4_cases[] = { + { 256, 4096, 512, 1 }, + { 512, 3840, 512, 2 }, + { 512, 4096, 512, 1 }, + { 768, 8192, 512, 3 }, + {1024, 2048, 512, 2 }, + {2048,17408, 512, 4 }, + }; + + for (const auto & test : paired_q4_cases) { + const std::string key = "paired-q4-" + std::to_string(test.tokens) + "-" + std::to_string(test.outputs); + auto ref = compile_kernel(*sync_jit, key, paired_q4, token_count_workload(test.tokens), + { + { "ggml.mul_mat_swiglu.input_size", "5120" }, + { "ggml.mul_mat_swiglu.output_size", std::to_string(test.outputs) }, + { "ggml.mul_mat_swiglu.op", "4" }, + { "ggml.mul_mat_swiglu.f16_output_layout", "0" }, + { "ggml.workload.token_capacity", std::to_string(test.tokens) }, + }); + REQUIRE(ref != nullptr); + REQUIRE(resolve_with_timeout(ref, std::chrono::seconds(120))); + auto compiled = ref->take_result(); + REQUIRE(compiled.hsaco_data != nullptr); + REQUIRE(compiled.launch_config.workgroup_size[0] == test.threads); + REQUIRE(compiled.launch_config.workgroup_count[0] == test.token_tiles); + const bool packed = test.tokens % 512 == 0 && (test.tokens / 512) * (test.outputs / 64) >= 64; + REQUIRE(compiled.launch_config.workgroup_count[1] == test.outputs / (packed ? 32 : 64)); + compiled.reset(); + } + + for (const int64_t tokens : { 256, 512, 1024, 2048 }) { + auto ref = compile_kernel(*sync_jit, "paired-q4-packed-output-" + std::to_string(tokens), paired_q4, + token_count_workload(tokens), + { + { "ggml.mul_mat_swiglu.input_size", "5120" }, + { "ggml.mul_mat_swiglu.output_size", "17408" }, + { "ggml.mul_mat_swiglu.op", "4" }, + { "ggml.mul_mat_swiglu.f16_output_layout", "1" }, + { "ggml.workload.token_capacity", std::to_string(tokens) }, + }); + require_compiled_kernel(ref); + } + + const auto & paired_q4_decode = find_kernel("ggml_mul_mat_swiglu_q4_q8_1_x4_lowtoken_dot"); + + const struct { + int64_t tokens; + int64_t inputs; + int64_t outputs; + int64_t threads; + int64_t row_tiles; + } paired_q4_decode_cases[] = { + { 1, 5120, 17408, 128, 2176 }, + { 1, 4096, 8192, 128, 1024 }, + { 1, 4096, 16384, 128, 2048 }, + { 1, 8192, 8192, 128, 1024 }, + { 1, 5376, 17408, 128, 2176 }, + { 2, 5120, 17408, 512, 1088 }, + { 3, 5120, 17408, 512, 1088 }, + { 4, 5120, 17408, 512, 1088 }, + { 5, 5120, 17408, 512, 1088 }, + { 5, 4096, 8192, 512, 256 }, + { 5, 4096, 16384, 512, 1024 }, + { 3, 8192, 8192, 512, 512 }, + { 5, 5632, 12288, 512, 768 }, + { 5, 5376, 17408, 512, 544 }, + { 5, 6144, 17408, 128, 2176 }, + { 5, 256, 4032, 128, 504 }, + { 5, 256, 4096, 512, 128 }, + { 5, 5632, 4096, 512, 128 }, + { 5, 5888, 4096, 128, 512 }, + }; + + for (const auto & test : paired_q4_decode_cases) { + const std::string key = "paired-q4-decode-" + std::to_string(test.tokens) + "-" + std::to_string(test.inputs) + + "-" + std::to_string(test.outputs); + auto ref = compile_kernel(*sync_jit, key, paired_q4_decode, token_count_workload(test.tokens), + { + { "ggml.mul_mat_swiglu.input_size", std::to_string(test.inputs) }, + { "ggml.mul_mat_swiglu.output_size", std::to_string(test.outputs) }, + { "ggml.mul_mat_swiglu.op", "4" }, + { "ggml.workload.token_capacity", std::to_string(test.tokens) }, + }); + REQUIRE(ref != nullptr); + REQUIRE(resolve_with_timeout(ref, std::chrono::seconds(120))); + auto compiled = ref->take_result(); + REQUIRE(compiled.hsaco_data != nullptr); + REQUIRE(compiled.launch_config.workgroup_size[0] == test.threads); + REQUIRE(compiled.launch_config.workgroup_count[0] == test.row_tiles); + REQUIRE(compiled.launch_config.workgroup_count[1] == 1); + compiled.reset(); + } + + const auto & quant_decode = find_kernel("ggml_mul_mat_f32_f32_wmma"); + + const struct { + int64_t tokens; + int64_t inputs; + int64_t outputs; + int64_t format; + int64_t threads; + int64_t row_tiles; + } quant_decode_cases[] = { + { 1, 5120, 48, 4, 256, 12 }, + { 2, 5120, 48, 4, 1024, 12 }, + { 3, 5120, 48, 4, 1024, 12 }, + { 4, 5120, 48, 4, 1024, 12 }, + { 5, 5120, 48, 4, 1024, 12 }, + { 5, 5120, 48, 5, 1024, 12 }, + { 5, 5120, 48, 6, 1024, 12 }, + { 5, 4096, 1, 4, 1024, 1 }, + { 5, 4096, 16, 4, 1024, 4 }, + { 5, 8192, 63, 4, 1024, 16 }, + { 5, 32768, 48, 4, 1024, 12 }, + { 5, 3840, 48, 4, 256, 12 }, + { 5, 4352, 48, 4, 256, 12 }, + { 5, 4096, 64, 4, 128, 8 }, + { 6, 5120, 48, 4, 128, 1 }, + { 1, 17408, 5120, 46, 128, 640 }, + { 4, 17408, 5120, 46, 128, 640 }, + { 5, 17408, 5120, 46, 512, 160 }, + { 5, 8192, 4096, 46, 512, 128 }, + { 5, 8448, 4160, 46, 512, 130 }, + { 5, 7936, 4096, 46, 128, 512 }, + { 5, 8192, 4032, 46, 128, 504 }, + { 5, 8192, 8256, 46, 128, 1032 }, + { 1, 17408, 5120, 44, 128, 640 }, + { 4, 17408, 5120, 44, 128, 640 }, + { 5, 17408, 5120, 44, 1024, 80 }, + { 5, 8192, 4096, 44, 128, 1024 }, + { 1, 6144, 5120, 44, 128, 640 }, + { 4, 6144, 5120, 44, 128, 640 }, + { 5, 6144, 5120, 44, 128, 1280 }, + { 5, 6400, 5120, 44, 128, 640 }, + { 5, 6656, 5120, 44, 128, 1280 }, + { 5, 6144, 4096, 44, 128, 512 }, + { 5, 6144, 6144, 44, 128, 768 }, + { 5, 5120, 6144, 44, 128, 768 }, + { 5, 8448, 4160, 44, 128, 520 }, + { 5, 9216, 9216, 44, 128, 1152 }, + { 5, 12288, 8192, 44, 128, 1024 }, + { 5, 16384, 4096, 44, 128, 512 }, + { 5, 16384, 5056, 44, 128, 632 }, + { 5, 16384, 5120, 44, 1024, 80 }, + { 5, 20480, 4096, 44, 1024, 64 }, + { 5, 20736, 4160, 44, 1024, 65 }, + { 5, 7936, 4096, 44, 128, 512 }, + { 5, 8192, 4032, 44, 128, 504 }, + { 5, 8192, 8256, 44, 128, 1032 }, + }; + + for (const auto & test : quant_decode_cases) { + const std::string key = "quant-decode-" + std::to_string(test.tokens) + "-" + std::to_string(test.inputs) + + "-" + std::to_string(test.outputs) + "-" + std::to_string(test.format); + auto ref = + compile_kernel(*sync_jit, key, quant_decode, token_count_workload(test.tokens), + { + { "ggml.mul_mat.input_size", std::to_string(test.inputs) }, + { "ggml.mul_mat.output_size", std::to_string(test.outputs) }, + { "ggml.mul_mat.weight_format", std::to_string(test.format) }, + { "ggml.mul_mat.activation_format", test.format == 44 || test.format == 46 ? "9" : "0" }, + { "ggml.mul_mat.output_accumulation", "0" }, + { "ggml.mul_mat.output_unary_op", "23" }, + { "ggml.workload.token_capacity", std::to_string(test.tokens) }, + }); + REQUIRE(ref != nullptr); + REQUIRE(resolve_with_timeout(ref, std::chrono::seconds(120))); + auto compiled = ref->take_result(); + REQUIRE(compiled.hsaco_data != nullptr); + REQUIRE(compiled.launch_config.workgroup_size[0] == test.threads); + REQUIRE(compiled.launch_config.workgroup_count[0] == test.row_tiles); + REQUIRE(compiled.launch_config.workgroup_count[1] == 1); + compiled.reset(); + } + + + for (const char * name : { "ggml_mul_mat_quantized_f16_wmma_prefill_conv4", + "ggml_mul_mat_quantized_f16_wmma_prefill_conv4_interior" }) { + const auto & quantized_conv4 = find_kernel(name); + REQUIRE(quantized_conv4.bindings.size() == (std::string(name).find("interior") != std::string::npos ? 5 : 6)); + for (const int format : {4, 6}) { + for (const int64_t channels : {2112, 8192, 10240}) { + auto ref = compile_kernel(*sync_jit, std::string(name) + "-" + std::to_string(format) + "-" + + std::to_string(channels), quantized_conv4, token_count_workload(512), { + {"ggml.mul_mat.input_size", "5120"}, + {"ggml.mul_mat.output_size", std::to_string(channels)}, + {"ggml.mul_mat.weight_format", std::to_string(format)}, + }); + REQUIRE(ref != nullptr); + REQUIRE(resolve_with_timeout(ref, std::chrono::seconds(120))); + auto compiled = ref->take_result(); + REQUIRE(compiled.hsaco_data != nullptr); + REQUIRE(compiled.launch_config.workgroup_size[0] == 512); + REQUIRE(compiled.launch_config.workgroup_count[0] == 1); + REQUIRE(compiled.launch_config.workgroup_count[1] == channels / 64); + compiled.reset(); + } + } + } + + const auto & conv_finish = find_kernel("llm_ssm_conv_dconv4_silu_prefill_finish_f32"); + REQUIRE(conv_finish.bindings.size() == 5); + for (const int64_t channels : { 2112, 8192, 8256, 10240 }) { + auto ref = compile_kernel(*sync_jit, "conv-finish-" + std::to_string(channels), conv_finish, + { + }, + { + { "llm.ssm_conv.generic.d_inner", std::to_string(channels) }, + }); + REQUIRE(ref != nullptr); + REQUIRE(resolve_with_timeout(ref, std::chrono::seconds(120))); + auto compiled = ref->take_result(); + REQUIRE(compiled.hsaco_data != nullptr); + REQUIRE(compiled.launch_config.workgroup_size[0] == 256); + REQUIRE(compiled.launch_config.workgroup_count[0] == (channels + 255) / 256); + } + + const auto & ssm_prefill = find_kernel("llm_ssm_conv_dconv4_silu_prefill_512_wg1024"); + for (const int64_t channels : { 8192, 8224, 8256, 10208, 10240 }) { + auto ref = compile_kernel(*sync_jit, "ssm-prefill-" + std::to_string(channels), ssm_prefill, + { + }, + { + { "llm.ssm_conv.prefill.d_conv", "4" }, + { "llm.ssm_conv.prefill.d_inner", std::to_string(channels) }, + { "llm.ssm_conv.prefill.n_t", "512" }, + { "llm.ssm_conv.prefill.n_s", "1" }, + { "llm.ssm_conv.prefill.state_row_stride", std::to_string(channels) }, + { "llm.ssm_conv.prefill.x_row_stride", std::to_string(channels) }, + }); + REQUIRE(ref != nullptr); + REQUIRE(resolve_with_timeout(ref, std::chrono::seconds(120))); + auto compiled = ref->take_result(); + REQUIRE(compiled.hsaco_data != nullptr); + REQUIRE(compiled.launch_config.workgroup_size[0] == 1024); + REQUIRE(compiled.launch_config.workgroup_count[0] == channels / (channels % 64 == 0 ? 64 : 32)); + REQUIRE(compiled.launch_config.workgroup_count[1] == 1); + compiled.reset(); + } + + const struct { + const char * name; + int64_t tokens; + size_t bindings; + } gdn_cases[] = { + { "llm_gated_delta_net_f32_wmma_head128", 17, 7 }, + { "llm_gated_delta_net_f32_wmma_head128_rmsnorm_gate", 17, 11 }, + { "llm_gated_delta_net_f32_wmma_head128_projection_epilogue", 16, 9 }, + { "llm_gated_delta_net_f32_wmma_head128_inplace", 1, 7 }, + { "llm_gated_delta_net_f32_wmma_head128_inplace_projection_epilogue", 1, 9 }, + { "llm_gated_delta_net_f32_wmma_head128_snapshot", 5, 8 }, + { "llm_gated_delta_net_f32_wmma_head128_snapshot_projection_epilogue", 5, 10 }, + }; + + for (const auto & test : gdn_cases) { + const auto & definition = find_kernel(test.name); + REQUIRE(definition.bindings.size() == test.bindings); + auto ref = compile_kernel(*sync_jit, test.name, definition, + { + }, + { + { "llm.gated_delta_net.head_width", "128" }, + { "llm.gated_delta_net.head_count", "3" }, + { "llm.gated_delta_net.token_count", std::to_string(test.tokens) }, + { "llm.gated_delta_net.sequence_count", "1" }, + { "llm.gated_delta_net.qk_stride1", "128" }, + { "llm.gated_delta_net.qk_stride2", "640" }, + { "llm.gated_delta_net.qk_stride3", std::to_string(test.tokens * 640) }, + { "llm.gated_delta_net.value_stride1", "128" }, + { "llm.gated_delta_net.value_stride2", "640" }, + { "llm.gated_delta_net.value_stride3", std::to_string(test.tokens * 640) }, + { "llm.gated_delta_net.scalar_stride1", "1" }, + { "llm.gated_delta_net.scalar_stride2", "3" }, + { "llm.gated_delta_net.scalar_stride3", std::to_string(test.tokens * 3) }, + { "llm.gated_delta_net.query_head_count", "1" }, + { "llm.gated_delta_net.query_sequence_ratio", "1" }, + { "llm.gated_delta_net.snapshot_stride", "49216" }, + { "llm.gated_delta_net.l2_epsilon", "0.000001" }, + { "ggml.rmsnorm_gate_f32.rms_epsilon", "0.000001" }, + { "ggml.rmsnorm_gate_f32.gate_op", "15" }, + }); + REQUIRE(ref != nullptr); + REQUIRE(resolve_with_timeout(ref, std::chrono::seconds(120))); + auto compiled = ref->take_result(); + REQUIRE(compiled.hsaco_data != nullptr); + REQUIRE(compiled.launch_config.workgroup_size[0] == 256); + REQUIRE(compiled.launch_config.workgroup_count[0] == 3); + REQUIRE(compiled.launch_config.workgroup_count[1] == 1); + compiled.reset(); + } +} + +static void run_moe_routing_compile_checks() { + static constexpr const char * kTarget = "gfx1100"; + + std::string error; + std::unique_ptr jit = ggml::hrx::create_loom_jit(kTarget, ggml::hrx::LoomJitMode::Sync, error); + REQUIRE(jit != nullptr); + + const std::map expert_table_workload = { + { "token_count", 1 }, + { "route_count", 8 }, + { "route_stride", 8 }, + { "expert_count", 128 }, + }; + const auto & expert_table = find_kernel("ggml_moe_build_expert_table"); + require_compiled_kernel( + compile_kernel(*jit, "moe-expert-table", expert_table, expert_table_workload, moe_routing_config())); + + const std::map partition_table_workload = { + { "token_count", 1 }, + { "route_count", 8 }, + { "expert_count", 128 }, + }; + const auto & partition_table = find_kernel("ggml_moe_build_expert_partition_table"); + require_compiled_kernel(compile_kernel(*jit, "moe-expert-partition-table", partition_table, + partition_table_workload, moe_routing_config())); + + const auto & weighted_reduce_publish_q8 = + find_kernel("qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32_publish_q8"); + REQUIRE(weighted_reduce_publish_q8.bindings.size() == 6); + require_compiled_kernel(compile_kernel(*jit, "moe-weighted-reduce-next-rmsnorm-publish-q8", + weighted_reduce_publish_q8, token_count_workload(2), { + { "qwen3_moe.model.hidden_size", "2048" }, + { "qwen3_moe.model.rms_epsilon", "0.000001" }, + { "qwen3_moe.routed_down.output_size", "2048" }, + { "qwen3_moe.routed_down.route_count", "8" }, + { "qwen3_moe.workload.token_capacity", "2" }, + { "ggml.quantize_q8_1_x4.group_capacity", "32" }, + })); +} + +static void run_prefill_fusion_compile_checks() { + static constexpr const char * kTarget = "gfx1100"; + + std::string error; + std::unique_ptr jit = ggml::hrx::create_loom_jit(kTarget, ggml::hrx::LoomJitMode::Sync, error); + REQUIRE(jit != nullptr); + + const auto & rmsnorm_f32_f16 = find_kernel("ggml_rmsnorm_binary_f32_f16"); + require_compiled_kernel(compile_kernel(*jit, "prefill-rmsnorm-f32-f16", rmsnorm_f32_f16, token_count_workload(512), + { + { "ggml.rmsnorm_binary_f32.hidden_size", "2048" }, + { "ggml.rmsnorm_binary_f32.rms_epsilon", "0.00001" }, + { "ggml.rmsnorm_binary_f32.op", "2" }, + })); + + const auto & rmsnorm_strided = find_kernel("ggml_rmsnorm_binary_strided_f32"); + require_compiled_kernel(compile_kernel(*jit, "prefill-rmsnorm-strided", rmsnorm_strided, token_count_workload(64), + { + { "ggml.rmsnorm_binary_f32.hidden_size", "256" }, + { "ggml.rmsnorm_binary_f32.input_stride", "288" }, + { "ggml.rmsnorm_binary_f32.rms_epsilon", "0.00001" }, + { "ggml.rmsnorm_binary_f32.op", "2" }, + })); + + const auto & rmsnorm_q8_f16 = find_kernel("ggml_rmsnorm_binary_q8_1_x4_publish"); + require_compiled_kernel(compile_kernel(*jit, "prefill-rmsnorm-q8-f16", rmsnorm_q8_f16, token_count_workload(512), + { + { "ggml.rmsnorm_binary_q8_1_x4.hidden_size", "2048" }, + { "ggml.rmsnorm_binary_q8_1_x4.rms_epsilon", "0.00001" }, + { "ggml.rmsnorm_binary_q8_1_x4.op", "2" }, + { "ggml.rmsnorm_binary_q8_1_x4.publish_f16", "1" }, + })); + + const auto & flash_f16 = find_kernel("ggml_flash_attention_f32_f16_wmma_publish_f16"); + require_compiled_kernel(compile_kernel(*jit, "prefill-flash-f16", flash_f16, flash_attention_workload(512, 512), + flash_attention_config("32", "8", "64", "0.125"))); + + const auto & swiglu_k16 = find_kernel("ggml_mul_mat_tiled_pair_input_f32_binary_publish_f32_k16"); + require_compiled_kernel(compile_kernel(*jit, "prefill-swiglu-k16", swiglu_k16, token_count_workload(512), + { + { "ggml.mul_mat_swiglu.input_size", "2048" }, + { "ggml.mul_mat_swiglu.output_size", "8192" }, + { "ggml.mul_mat_swiglu.gate_weight_format", "4" }, + { "ggml.mul_mat_swiglu.up_weight_format", "4" }, + { "ggml.mul_mat_swiglu.op", "4" }, + { "ggml.workload.token_capacity", "512" }, + })); + + const auto & get_rows_rmsnorm = find_kernel("ggml_get_rows_rmsnorm_binary_q8_1_x4_f16"); + require_compiled_kernel(compile_kernel(*jit, "prefill-get-rows-rmsnorm", get_rows_rmsnorm, + { + { "token_count", 512 }, + { "row_count", 128256 }, + { "hidden_size", 2048 }, + }, + { + { "ggml.get_rows_f32.token_capacity", "512" }, + { "ggml.get_rows_f32.hidden_capacity", "2048" }, + { "ggml.get_rows_f32.weight_format", "6" }, + { "ggml.get_rows_rmsnorm.rms_epsilon", "0.00001" }, + })); + + const auto & get_rows_scale = find_kernel("ggml_get_rows_scale_f32"); + require_compiled_kernel(compile_kernel(*jit, "prefill-get-rows-scale", get_rows_scale, + { + { "token_count", 64 }, + { "row_count", 122753 }, + { "hidden_size", 2560 }, + }, + { + { "ggml.get_rows_f32.token_capacity", "64" }, + { "ggml.get_rows_f32.hidden_capacity", "2560" }, + { "ggml.get_rows_f32.weight_format", "32" }, + { "ggml.get_rows_scale_f32.scale", "0.177800179" }, + { "ggml.get_rows_scale_f32.bias", "0" }, + })); + require_compiled_kernel(compile_kernel(*jit, "gemma-prefill-get-rows-scale-q5_1", get_rows_scale, + { + { "token_count", 64 }, + { "row_count", 262144 }, + { "hidden_size", 640 }, + }, + { + { "ggml.get_rows_f32.token_capacity", "64" }, + { "ggml.get_rows_f32.hidden_capacity", "640" }, + { "ggml.get_rows_f32.weight_format", "51" }, + { "ggml.get_rows_scale_f32.scale", "25.2982216" }, + { "ggml.get_rows_scale_f32.bias", "0" }, + })); + require_compiled_kernel(compile_kernel(*jit, "gemma-decode-get-rows-scale-q5_1", get_rows_scale, + { + { "token_count", 1 }, + { "row_count", 262144 }, + { "hidden_size", 640 }, + }, + { + { "ggml.get_rows_f32.token_capacity", "1" }, + { "ggml.get_rows_f32.hidden_capacity", "640" }, + { "ggml.get_rows_f32.weight_format", "51" }, + { "ggml.get_rows_scale_f32.scale", "25.2982216" }, + { "ggml.get_rows_scale_f32.bias", "0" }, + })); + + const auto & lfm_dconv3 = find_kernel("llm_ssm_conv_lfm_dconv3_state_f32"); + REQUIRE(lfm_dconv3.bindings.size() == 5); + for (const int64_t tokens : {1, 64}) { + auto ref = compile_kernel(*jit, "lfm-dconv3-" + std::to_string(tokens), lfm_dconv3, {}, { + {"llm.ssm_conv.lfm_dconv3.n_t", std::to_string(tokens)}, + }); + REQUIRE(ref != nullptr); + REQUIRE(resolve_with_timeout(ref, std::chrono::seconds(120))); + auto compiled = ref->take_result(); + REQUIRE(compiled.hsaco_data != nullptr); + REQUIRE(compiled.launch_config.workgroup_size[0] == 256); + REQUIRE(compiled.launch_config.workgroup_count[0] == static_cast(8 * tokens)); + } +} + +static void run_binary_fusion_compile_checks() { + std::string error; + std::unique_ptr jit = + ggml::hrx::create_loom_jit("gfx1100", ggml::hrx::LoomJitMode::Sync, error); + REQUIRE(jit != nullptr); + + const auto & gather = find_kernel("ggml_gather_add_f32"); + require_compiled_kernel(compile_kernel(*jit, "binary-fusion-gather-div", gather, + { + { "source_token_count", 2 }, + { "output_token_count", 1 }, + { "hidden_size", 2048 }, + }, + { + { "ggml.gather_add_f32.binary_op", "3" }, + })); + + const auto & scale = find_kernel("ggml_scale_add_f32"); + require_compiled_kernel(compile_kernel(*jit, "binary-fusion-scale-sub", scale, + { + { "element_count", 128 }, + }, + { + { "ggml.scale_add_f32.ne0", "128" }, + { "ggml.scale_add_f32.ne1", "1" }, + { "ggml.scale_add_f32.ne2", "1" }, + { "ggml.scale_add_f32.input_stride1", "128" }, + { "ggml.scale_add_f32.input_stride2", "128" }, + { "ggml.scale_add_f32.input_stride3", "128" }, + { "ggml.scale_add_f32.input_span", "128" }, + { "ggml.scale_add_f32.residual_stride1", "128" }, + { "ggml.scale_add_f32.residual_stride2", "128" }, + { "ggml.scale_add_f32.residual_stride3", "128" }, + { "ggml.scale_add_f32.residual_span", "128" }, + { "ggml.scale_add_f32.scale", "0.5" }, + { "ggml.scale_add_f32.bias", "0" }, + { "ggml.scale_add_f32.binary_op", "1" }, + { "ggml.scale_add_f32.scaled_lhs", "0" }, + })); + + const auto & rms = find_kernel("ggml_rmsnorm_mul_add_f32"); + require_compiled_kernel(compile_kernel(*jit, "binary-fusion-rms-div-sub", rms, + { + { "token_count", 1 }, + }, + { + { "ggml.rmsnorm_mul_add_f32.hidden_size", "128" }, + { "ggml.rmsnorm_mul_add_f32.rms_epsilon", "0.000001" }, + { "ggml.rmsnorm_mul_add_f32.normalized_op", "3" }, + { "ggml.rmsnorm_mul_add_f32.output_op", "1" }, + { "ggml.rmsnorm_mul_add_f32.binary_lhs", "0" }, + })); + + const auto & gather_rms = find_kernel("ggml_gather_add_rmsnorm_binary_f32"); + require_compiled_kernel(compile_kernel(*jit, "binary-fusion-gather-rms-mul", gather_rms, + { + { "source_token_count", 2 }, + { "output_token_count", 1 }, + }, + { + { "ggml.gather_add_rmsnorm_binary_f32.hidden_size", "128" }, + { "ggml.gather_add_rmsnorm_binary_f32.rms_epsilon", "0.000001" }, + { "ggml.gather_add_rmsnorm_binary_f32.gather_op", "2" }, + })); + + const auto & ssm = find_kernel("llm_ssm_conv_binary_f32"); + require_compiled_kernel(compile_kernel(*jit, "binary-fusion-ssm-rhs-div", ssm, + { + { "n_t", 8 }, + { "n_s", 1 }, + }, + { + { "llm.ssm_conv.generic.d_conv", "4" }, + { "llm.ssm_conv.generic.d_inner", "64" }, + { "llm.ssm_conv.generic.unary_op", "23" }, + { "llm.ssm_conv.generic.binary_op", "3" }, + { "llm.ssm_conv.generic.binary_lhs", "0" }, + { "llm.ssm_conv.generic.workgroup_size", "256" }, + })); + + const auto & vector_postops = find_kernel("ggml_mul_mat_vector_bias_residual_f32_f32"); + for (const int64_t epilogue : { int64_t{ 1 }, int64_t{ 2 }, int64_t{ 3 } }) { + require_compiled_kernel(compile_kernel( + *jit, "binary-fusion-vector-postops-" + std::to_string(epilogue), vector_postops, token_count_workload(1), + { + { "ggml.matmul.vector_postops.input_size", "2048" }, + { "ggml.matmul.vector_postops.output_size", "2048" }, + { "ggml.matmul.vector_postops.weight_format", "4" }, + { "ggml.matmul.vector_postops.apply_bias", epilogue & 1 ? "1" : "0" }, + { "ggml.matmul.vector_postops.apply_residual", epilogue & 2 ? "1" : "0" }, + { "ggml.workload.token_capacity", "1" }, + })); + } + + const auto compile_postops = [&](const char * kernel_name, const char * key_prefix, int64_t token_count, + int64_t token_capacity) { + const auto & kernel = find_kernel(kernel_name); + for (const int64_t epilogue : { int64_t{ 1 }, int64_t{ 2 }, int64_t{ 3 } }) { + require_compiled_kernel(compile_kernel( + *jit, std::string(key_prefix) + std::to_string(epilogue), kernel, token_count_workload(token_count), + { + { "ggml.mul_mat_postops.input_size", "2048" }, + { "ggml.mul_mat_postops.output_size", "2048" }, + { "ggml.mul_mat_postops.weight_format", "4" }, + { "ggml.mul_mat_postops.epilogue", std::to_string(epilogue) }, + { "ggml.workload.token_capacity", std::to_string(token_capacity) }, + })); + } + }; + compile_postops("ggml_mul_mat_skinny_input_f32_bias_residual_publish_f32", "binary-fusion-skinny-postops-", 2, 2); + compile_postops("ggml_mul_mat_tiled_input_f32_bias_residual_publish_f32", "binary-fusion-tiled-postops-", 32, 32); + + const auto & pair_postops = find_kernel("ggml_mul_mat_tiled_pair_input_f32_binary_bias_residual_publish_f32"); + for (const int64_t epilogue : { int64_t{ 1 }, int64_t{ 2 }, int64_t{ 3 } }) { + require_compiled_kernel(compile_kernel(*jit, "binary-fusion-pair-postops-" + std::to_string(epilogue), + pair_postops, token_count_workload(32), + { + { "ggml.mul_mat_swiglu.input_size", "2048" }, + { "ggml.mul_mat_swiglu.output_size", "2048" }, + { "ggml.mul_mat_swiglu.gate_weight_format", "4" }, + { "ggml.mul_mat_swiglu.up_weight_format", "4" }, + { "ggml.mul_mat_swiglu.op", "4" }, + { "ggml.mul_mat_swiglu.epilogue", std::to_string(epilogue) }, + { "ggml.workload.token_capacity", "32" }, + })); + } +} + +static void run_loom_jit_materialize_checks() { + const ggml::hrx::KernelDefinition & binary = find_kernel("ggml_binary_f32"); + + HrxTestDevice device; + REQUIRE(device.open()); + run_cache_materialize_case(device, ggml::hrx::LoomJitMode::Sync, binary); + run_cache_materialize_case(device, ggml::hrx::LoomJitMode::Async, binary); + run_dynamic_cache_materialize_case(device, ggml::hrx::LoomJitMode::Sync, binary); + run_dynamic_cache_materialize_case(device, ggml::hrx::LoomJitMode::Async, binary); + run_dynamic_export_reuse_case(device, "ggml_unary_f32", { { "element_count", 64 } }, + { { "element_count", 257 } }, { { "ggml.unary_f32.op", "4" } }); + run_dynamic_export_reuse_case(device, "ggml_copy_f32", { { "element_count", 64 } }, + { { "element_count", 257 } }); + run_dynamic_export_reuse_case( + device, "ggml_concat_dim0_f32", + { { "row_count", 3 }, { "lhs_span", 192 }, { "rhs_span", 96 } }, + { { "row_count", 7 }, { "lhs_span", 448 }, { "rhs_span", 224 } }, + { { "ggml.concat_dim0_f32.lhs_width", "64" }, { "ggml.concat_dim0_f32.rhs_width", "32" } }); + run_dynamic_export_reuse_case( + device, "ggml_rmsnorm_f32", { { "token_count", 1 } }, { { "token_count", 4 } }, + { { "ggml.rmsnorm_f32.hidden_size", "256" }, { "ggml.rmsnorm_f32.input_stride", "256" }, + { "ggml.rmsnorm_f32.rms_epsilon", "1e-06" } }); + run_dynamic_export_reuse_case( + device, "ggml_get_rows_f32", + { { "token_count", 3 }, { "row_count", 64 }, { "hidden_size", 256 } }, + { { "token_count", 7 }, { "row_count", 64 }, { "hidden_size", 256 } }, + { { "ggml.get_rows_f32.token_capacity", "8" }, { "ggml.get_rows_f32.hidden_capacity", "256" }, + { "ggml.get_rows_f32.weight_format", "32" } }); + run_dynamic_export_reuse_case( + device, "ggml_rope_f32", { { "token_count", 3 }, { "input_span", 768 } }, + { { "token_count", 7 }, { "input_span", 1792 } }, + { { "ggml.rope_f32.head_size", "64" }, { "ggml.rope_f32.n_dims", "64" }, + { "ggml.rope_f32.head_count", "4" }, { "ggml.rope_f32.token_capacity", "8" }, + { "ggml.rope_f32.input_stride1", "64" }, { "ggml.rope_f32.input_stride2", "256" }, + { "ggml.rope_f32.mscale", "1" }, { "ggml.rope_f32.mode", "0" } }); + run_dynamic_export_reuse_case( + device, "ggml_mul_mat_tiled_input_f32_publish_f32_aligned", { { "token_count", 32 } }, + { { "token_count", 64 } }, + { { "ggml.mul_mat.input_size", "256" }, { "ggml.mul_mat.output_size", "64" }, + { "ggml.mul_mat.output_accumulation", "0" }, { "ggml.mul_mat.output_unary_op", "23" }, + { "ggml.mul_mat.weight_format", "4" } }); + run_dynamic_export_reuse_case( + device, "ggml_mul_mat_tiled_pair_input_f32_binary_publish_f32_aligned", { { "token_count", 32 } }, + { { "token_count", 64 } }, + { { "ggml.mul_mat_swiglu.input_size", "256" }, { "ggml.mul_mat_swiglu.output_size", "64" }, + { "ggml.mul_mat_swiglu.gate_weight_format", "4" }, + { "ggml.mul_mat_swiglu.up_weight_format", "4" }, { "ggml.mul_mat_swiglu.op", "4" } }); + run_dynamic_export_reuse_case( + device, "llm_ssm_conv_f32", { { "n_t", 8 }, { "n_s", 1 } }, { { "n_t", 16 }, { "n_s", 2 } }, + { { "llm.ssm_conv.generic.d_conv", "4" }, { "llm.ssm_conv.generic.d_inner", "64" }, + { "llm.ssm_conv.generic.unary_op", "23" }, { "llm.ssm_conv.generic.workgroup_size", "256" } }); + run_dynamic_export_reuse_case( + device, "llm_ssm_conv_binary_f32", { { "n_t", 8 }, { "n_s", 1 } }, + { { "n_t", 16 }, { "n_s", 2 } }, + { { "llm.ssm_conv.generic.d_conv", "4" }, { "llm.ssm_conv.generic.d_inner", "64" }, + { "llm.ssm_conv.generic.unary_op", "23" }, { "llm.ssm_conv.generic.binary_op", "3" }, + { "llm.ssm_conv.generic.binary_lhs", "0" }, + { "llm.ssm_conv.generic.workgroup_size", "256" } }); + run_dynamic_export_reuse_case(device, "ggml_copy_transpose_f16", + { { "row_count", 512 }, { "column_count", 256 } }, + { { "row_count", 1024 }, { "column_count", 512 } }); + run_dynamic_export_reuse_case( + device, "ggml_flash_attention_f32_f16_wmma", flash_attention_workload(16, 64), + flash_attention_workload(17, 65), flash_attention_config("32", "4", "128", "0.0883883461")); + run_dynamic_export_reuse_case( + device, "ggml_flash_attention_f32_f16_wmma_publish_f16", flash_attention_workload(16, 64), + flash_attention_workload(17, 65), flash_attention_config("32", "4", "128", "0.0883883461")); + run_dynamic_export_reuse_case( + device, "ggml_flash_attention_decode_split_f32_f16_wmma_next_q8", + flash_attention_decode_workload(65), flash_attention_decode_workload(127), + flash_attention_decode_config("32", "4", "128", "0.0883883461", "128")); + run_dynamic_export_reuse_case( + device, "llm_gated_delta_net_projection_epilogue_f32", { { "element_count", 48 } }, + { { "element_count", 192 } }, + { { "llm.gated_delta_net.epilogue_head_count", "48" }, + { "llm.gated_delta_net.epilogue_workgroup_size", "256" } }); + run_targeted_export_materialize_case(device); +} + +static void run_binary_publication_compile_checks() { + std::string error; + std::unique_ptr jit = + ggml::hrx::create_loom_jit("gfx1100", ggml::hrx::LoomJitMode::Sync, error); + REQUIRE(jit != nullptr); + + const auto & binary_q8 = find_kernel("ggml_binary_f32_publish_q8_1_x4"); + require_compiled_kernel(compile_kernel(*jit, "binary-mul-q8", binary_q8, { + { "token_count", 2 }, + }, { + { "ggml.binary_f32.op", "2" }, + { "ggml.binary_f32.ne0", "2048" }, + { "ggml.binary_f32.src0_stride1", "2048" }, + { "ggml.binary_f32.src1_stride1", "2048" }, + { "ggml.binary_f32.src0_span", "4096" }, + { "ggml.binary_f32.src1_span", "4096" }, + })); + + const auto & binary_bc_q8 = find_kernel("ggml_binary_bc_f32_publish_q8_1_x4"); + require_compiled_kernel(compile_kernel(*jit, "binary-bc-mul-q8", binary_bc_q8, { + { "token_count", 2 }, + { "hidden_size", 256 }, + { "src0_element_count", 512 }, + { "src1_element_count", 256 }, + }, { + { "ggml.binary_bc_f32.op", "2" }, + { "ggml.binary_bc_f32.src0_broadcast_dim0", "0" }, + { "ggml.binary_bc_f32.src0_broadcast_dim1", "0" }, + { "ggml.binary_bc_f32.src1_broadcast_dim0", "0" }, + { "ggml.binary_bc_f32.src1_broadcast_dim1", "1" }, + })); +} + +static bool has_hrx_test_device() { + HrxTestDevice device; + return device.open(); +} + +static void register_hrx_loom_jit_cases(test_runner::Suite & suite) { + suite.host_case("compile", [] { run_loom_jit_compile_checks(); }); + suite.host_case("binary-publication", [] { run_binary_publication_compile_checks(); }); + suite.host_case("moe-routing", [] { run_moe_routing_compile_checks(); }); + suite.host_case("prefill-fusions", [] { run_prefill_fusion_compile_checks(); }); + suite.host_case("binary-fusions", [] { run_binary_fusion_compile_checks(); }); + suite.device_case("materialize", [] { run_loom_jit_materialize_checks(); }); +} + +int main(int argc, char ** argv) { + test_runner::Suite suite( + test_runner::Config::with_prefix("HRX Loom JIT test", "hrx-loom-jit", "GGML_HRX_LOOM_JIT_TEST", 2)); + register_hrx_loom_jit_cases(suite); + + const bool has_device = has_hrx_test_device(); + return suite.run(argc, argv, has_device); +} diff --git a/tests/test-hrx-moe-split.cpp b/tests/test-hrx-moe-split.cpp new file mode 100644 index 000000000000..20a821769004 --- /dev/null +++ b/tests/test-hrx-moe-split.cpp @@ -0,0 +1,240 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// gpt-oss's MoE block with its MXFP4 experts in CPU memory and its biases and routing in HRX memory, as llama.cpp +// lays out gpt-oss when HRX declines the expert MUL_MAT_ID at load time: +// gate = add_id(mul_mat_id(W_gate, x, ids), b_gate, ids) up = add_id(mul_mat_id(W_up, x, ids), b_up, ids) +// h = swiglu_oai(gate, up, 1.702, 7) out = add_id(mul_mat_id(W_down, h, ids), b_down, ids) +// ids is a [n_used of n_expert] view of a [n_expert, tokens] tensor in HRX memory, as llama.cpp passes the +// argsort and gpt-oss computes it on HRX. Two passes: experts in a CPU-only buffer (CPU_REPACK, the layout above, +// where every MUL_MAT_ID stays on the CPU), and experts in a plain host buffer (HRX may read them). Checks the +// result against the same graph on the CPU alone, and the placement guard (moe-placement-guard.cpp): an ADD_ID or +// SWIGLU_OAI never runs on HRX while the MUL_MAT_ID it follows runs on the CPU (engine #286). + +#include "ggml-alloc.h" +#include "ggml-backend.h" +#include "ggml-cpu.h" +#include "ggml.h" + +#include +#include +#include +#include +#include +#include +#include + +#define REQUIRE(condition) \ + do { \ + if (!(condition)) { \ + std::fprintf(stderr, "%s:%d: requirement failed: %s\n", __FILE__, __LINE__, #condition); \ + std::abort(); \ + } \ + } while (false) + +static constexpr int64_t kHidden = 256; +static constexpr int64_t kFfn = 256; +static constexpr int64_t kExperts = 8; +static constexpr int64_t kUsed = 4; + +struct Model { + std::vector w_gate, w_up, w_down; // MXFP4 + std::vector b_gate, b_up, b_down, x; + std::vector ids; +}; + +static Model make_model(int64_t tokens) { + std::mt19937 rng(1234); + std::uniform_real_distribution u(-0.5f, 0.5f); + auto fill = [&](std::vector & v, size_t n) { + v.resize(n); + for (float & f : v) { + f = u(rng); + } + }; + auto experts = [&](std::vector & out, int64_t cols, int64_t rows) { + std::vector f; + fill(f, cols * rows * kExperts); + out.resize(ggml_row_size(GGML_TYPE_MXFP4, cols) * rows * kExperts); + ggml_quantize_chunk(GGML_TYPE_MXFP4, f.data(), out.data(), 0, rows * kExperts, cols, nullptr); + }; + Model m; + experts(m.w_gate, kHidden, kFfn); + experts(m.w_up, kHidden, kFfn); + experts(m.w_down, kFfn, kHidden); + fill(m.b_gate, kFfn * kExperts); + fill(m.b_up, kFfn * kExperts); + fill(m.b_down, kHidden * kExperts); + fill(m.x, kHidden * tokens); + m.ids.resize(kExperts * tokens); + for (int64_t t = 0; t < tokens; ++t) { + std::vector order(kExperts); + for (int64_t e = 0; e < kExperts; ++e) { + order[e] = (int32_t) e; + } + std::shuffle(order.begin(), order.end(), rng); + std::copy(order.begin(), order.end(), m.ids.begin() + t * kExperts); + } + return m; +} + +// The CPU's CPU_REPACK buffer type, or nullptr when this build has none. +static ggml_backend_buffer_type_t cpu_repack_buft() { + ggml_backend_dev_t dev = ggml_backend_dev_by_type(GGML_BACKEND_DEVICE_TYPE_CPU); + if (dev == nullptr) { + return nullptr; + } + auto get_extra = (ggml_backend_dev_get_extra_bufts_t) ggml_backend_reg_get_proc_address( + ggml_backend_dev_backend_reg(dev), "ggml_backend_dev_get_extra_bufts"); + if (get_extra == nullptr) { + return nullptr; + } + for (ggml_backend_buffer_type_t * b = get_extra(dev); b != nullptr && *b != nullptr; ++b) { + if (std::strcmp(ggml_backend_buft_name(*b), "CPU_REPACK") == 0) { + return *b; + } + } + return nullptr; +} + +// Runs the block; with hrx != nullptr the experts go to expert_buft and the biases and routing to HRX memory. +// Returns false in *guard_ok when an ADD_ID / SWIGLU_OAI ran on HRX while its MUL_MAT_ID ran elsewhere. +static std::vector run(ggml_backend_t hrx, ggml_backend_t cpu, ggml_backend_buffer_type_t expert_buft, + const Model & m, int64_t tokens, bool * guard_ok) { + ggml_init_params wp = { 16 * ggml_tensor_overhead(), nullptr, true }; + ggml_context * ectx = ggml_init(wp); // experts + ggml_context * xctx = ggml_init(wp); // the activations, in CPU memory + ggml_context * bctx = ggml_init(wp); // biases and routing (HRX-resident when hrx != nullptr) + ggml_tensor * w_gate = ggml_new_tensor_3d(ectx, GGML_TYPE_MXFP4, kHidden, kFfn, kExperts); + ggml_tensor * w_up = ggml_new_tensor_3d(ectx, GGML_TYPE_MXFP4, kHidden, kFfn, kExperts); + ggml_tensor * w_down = ggml_new_tensor_3d(ectx, GGML_TYPE_MXFP4, kFfn, kHidden, kExperts); + ggml_tensor * x = ggml_new_tensor_3d(xctx, GGML_TYPE_F32, kHidden, 1, tokens); + ggml_tensor * b_gate = ggml_new_tensor_2d(bctx, GGML_TYPE_F32, kFfn, kExperts); + ggml_tensor * b_up = ggml_new_tensor_2d(bctx, GGML_TYPE_F32, kFfn, kExperts); + ggml_tensor * b_down = ggml_new_tensor_2d(bctx, GGML_TYPE_F32, kHidden, kExperts); + ggml_tensor * ids_all = ggml_new_tensor_2d(bctx, GGML_TYPE_I32, kExperts, tokens); + ggml_set_input(x); + ggml_set_input(ids_all); + ggml_backend_buffer_t ebuf = ggml_backend_alloc_ctx_tensors_from_buft( + ectx, hrx != nullptr && expert_buft != nullptr ? expert_buft : ggml_backend_get_default_buffer_type(cpu)); + ggml_backend_buffer_t xbuf = ggml_backend_alloc_ctx_tensors(xctx, cpu); + ggml_backend_buffer_t bbuf = ggml_backend_alloc_ctx_tensors(bctx, hrx != nullptr ? hrx : cpu); + REQUIRE(ebuf != nullptr && xbuf != nullptr && bbuf != nullptr); + ggml_backend_buffer_set_usage(ebuf, GGML_BACKEND_BUFFER_USAGE_WEIGHTS); + ggml_backend_buffer_set_usage(bbuf, GGML_BACKEND_BUFFER_USAGE_WEIGHTS); + ggml_backend_tensor_set(w_gate, m.w_gate.data(), 0, ggml_nbytes(w_gate)); + ggml_backend_tensor_set(w_up, m.w_up.data(), 0, ggml_nbytes(w_up)); + ggml_backend_tensor_set(w_down, m.w_down.data(), 0, ggml_nbytes(w_down)); + ggml_backend_tensor_set(b_gate, m.b_gate.data(), 0, ggml_nbytes(b_gate)); + ggml_backend_tensor_set(b_up, m.b_up.data(), 0, ggml_nbytes(b_up)); + ggml_backend_tensor_set(b_down, m.b_down.data(), 0, ggml_nbytes(b_down)); + ggml_backend_tensor_set(x, m.x.data(), 0, ggml_nbytes(x)); + ggml_backend_tensor_set(ids_all, m.ids.data(), 0, ggml_nbytes(ids_all)); + + ggml_init_params gp = { 64 * ggml_tensor_overhead() + ggml_graph_overhead(), nullptr, true }; + ggml_context * gctx = ggml_init(gp); + ggml_tensor * ids = ggml_view_2d(gctx, ids_all, kUsed, tokens, ids_all->nb[1], 0); + ggml_tensor * mm_gate = ggml_mul_mat_id(gctx, w_gate, x, ids); + ggml_tensor * mm_up = ggml_mul_mat_id(gctx, w_up, x, ids); + ggml_tensor * gate = ggml_add_id(gctx, mm_gate, b_gate, ids); + ggml_tensor * up = ggml_add_id(gctx, mm_up, b_up, ids); + ggml_tensor * h = ggml_swiglu_oai(gctx, gate, up, 1.702f, 7.0f); + ggml_tensor * mm_down = ggml_mul_mat_id(gctx, w_down, h, ids); + ggml_tensor * out = ggml_add_id(gctx, mm_down, b_down, ids); + ggml_set_output(out); + ggml_cgraph * graph = ggml_new_graph(gctx); + ggml_build_forward_expand(graph, out); + + std::vector backends; + if (hrx != nullptr) { + backends.push_back(hrx); + } + backends.push_back(cpu); + ggml_backend_sched_t sched = + ggml_backend_sched_new(backends.data(), nullptr, (int) backends.size(), 1024, false, true); + REQUIRE(sched != nullptr); + REQUIRE(ggml_backend_sched_graph_compute(sched, graph) == GGML_STATUS_SUCCESS); + *guard_ok = true; + if (hrx != nullptr) { + auto on_hrx = [&](ggml_tensor * t) { return ggml_backend_sched_get_tensor_backend(sched, t) == hrx; }; + std::printf(" placement:"); + for (ggml_tensor * node : { mm_gate, gate, mm_up, up, h, mm_down, out }) { + std::printf(" %s=%s", ggml_op_desc(node), on_hrx(node) ? "HRX" : "CPU"); + } + std::printf("\n"); + const bool tail_ok = (!on_hrx(gate) || on_hrx(mm_gate)) && (!on_hrx(up) || on_hrx(mm_up)) && + (!on_hrx(h) || (on_hrx(mm_gate) && on_hrx(mm_up))) && (!on_hrx(out) || on_hrx(mm_down)); + *guard_ok = tail_ok; + } + std::vector result(ggml_nelements(out)); + ggml_backend_tensor_get(out, result.data(), 0, ggml_nbytes(out)); + ggml_backend_sched_free(sched); + ggml_free(gctx); + ggml_backend_buffer_free(ebuf); + ggml_backend_buffer_free(xbuf); + ggml_backend_buffer_free(bbuf); + ggml_free(ectx); + ggml_free(xctx); + ggml_free(bctx); + return result; +} + +int main() { + std::setvbuf(stdout, nullptr, _IONBF, 0); // keep the lines printed before a failed REQUIRE + ggml_backend_dev_t device = ggml_backend_dev_by_name("HRX0"); + if (device == nullptr) { + ggml_backend_load_all(); + device = ggml_backend_dev_by_name("HRX0"); + } + if (device == nullptr) { + std::printf("test-hrx-moe-split: no HRX0 device, skipped\n"); + return 0; + } + ggml_backend_t hrx = ggml_backend_dev_init(device, nullptr); + ggml_backend_t cpu = ggml_backend_cpu_init(); + REQUIRE(hrx != nullptr && cpu != nullptr); + ggml_backend_buffer_type_t repack = cpu_repack_buft(); + if (repack == nullptr) { + std::printf("test-hrx-moe-split: no CPU_REPACK buffer type, the CPU-only expert pass is skipped\n"); + } + int failures = 0; + for (ggml_backend_buffer_type_t buft : { repack, ggml_backend_get_default_buffer_type(cpu) }) { + if (buft == nullptr) { + continue; + } + for (int64_t tokens : { 1, 3, 12, 40 }) { + const Model m = make_model(tokens); + bool ref_ok = true, guard_ok = true; + std::vector expected = run(nullptr, cpu, nullptr, m, tokens, &ref_ok); + std::vector got = run(hrx, cpu, buft, m, tokens, &guard_ok); + double err = 0.0, ref = 0.0; + for (size_t i = 0; i < got.size(); ++i) { + const double d = (double) got[i] - expected[i]; + err += d * d; + ref += (double) expected[i] * expected[i]; + } + const double nmse = ref > 0.0 ? err / ref : err; + const bool ok = nmse <= 5e-4 && guard_ok; + std::printf("experts=%-10s tokens=%2lld nmse=%.3g guard=%s %s\n", ggml_backend_buft_name(buft), + (long long) tokens, nmse, guard_ok ? "ok" : "VIOLATED", ok ? "OK" : "FAIL"); + failures += ok ? 0 : 1; + } + } + ggml_backend_free(cpu); + ggml_backend_free(hrx); + REQUIRE(failures == 0); + std::printf("test-hrx-moe-split: all cases OK\n"); + return 0; +} diff --git a/tests/test-hrx-mul-mat-id-k32.cpp b/tests/test-hrx-mul-mat-id-k32.cpp new file mode 100644 index 000000000000..a18fbba5c204 --- /dev/null +++ b/tests/test-hrx-mul-mat-id-k32.cpp @@ -0,0 +1,194 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// MUL_MAT_ID on HRX with input sizes that are a multiple of 32 but not of 256: gpt-oss-20b's experts (2880, the last +// 256-value tile 64 values long) and BlackMamba's (1152, last tile 128 long), with 2816 (no tail) as the control. +// The 32-value block formats MXFP4, Q8_0 and Q4_0; 1 token (decode kernel) and 7 / 40 tokens (WMMA kernels). Each +// graph runs on the HRX device itself (no scheduler, no CPU fallback), its HRX plan must contain a mul_mat_id +// kernel, and the result is compared with the CPU backend (normalized MSE, as test-backend-ops). + +#include "dispatch/dispatch-scheduler.h" +#include "ggml-alloc.h" +#include "ggml-backend.h" +#include "ggml-cpu.h" +#include "ggml.h" +#include "graph/graph.h" +#include "kernel-corpus/kernel-corpus.h" + +#include +#include +#include +#include +#include +#include +#include + +#define REQUIRE(condition) \ + do { \ + if (!(condition)) { \ + std::fprintf(stderr, "%s:%d: requirement failed: %s\n", __FILE__, __LINE__, #condition); \ + std::abort(); \ + } \ + } while (false) + +static constexpr int64_t kOutputSize = 96; // one full 64-row tile and a partial one +static constexpr int64_t kExpertCount = 8; +static constexpr int64_t kRouteCount = 4; +static constexpr double kMaxNmse = 5e-4; + +static std::string kernel_name_for_id(uint64_t kernel_id) { + const ggml::hrx::KernelResolveResult resolved = + ggml::hrx::resolve_kernel_definition(ggml::hrx::get_qwen_kernel_corpus(), "gfx1151", kernel_id); + REQUIRE(resolved.found()); + return ggml::hrx::kernel_definition_name(*resolved.definition); +} + +static std::string hrx_mul_mat_id_kernel(ggml_cgraph * graph) { + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + ggml::hrx::DispatchScheduleDiagnostics diagnostics; + if (!scheduler.schedule_graph(imported.graph, { "gfx1151" }, &diagnostics)) { + std::fprintf(stderr, "unsupported: %s\n", diagnostics.unsupported_message.c_str()); + std::abort(); + } + REQUIRE(scheduler.plan().valid()); + std::string found; + for (const ggml::hrx::Dispatch & dispatch : scheduler.plan().dispatches) { + const std::string name = kernel_name_for_id(dispatch.kernel.kernel_id); + if (name.find("mul_mat_id") != std::string::npos) { + found = name; + } + } + return found; +} + +struct Inputs { + std::vector weights; + std::vector activations; + std::vector ids; +}; + +static Inputs make_inputs(ggml_type type, int64_t input_size, int64_t token_count, uint32_t seed) { + std::mt19937 rng(seed); + std::uniform_real_distribution uniform(-1.0f, 1.0f); + Inputs in; + const int64_t rows = kOutputSize * kExpertCount; + std::vector w(static_cast(rows * input_size)); + for (float & v : w) { + v = uniform(rng); + } + in.weights.resize(ggml_row_size(type, input_size) * rows); + ggml_quantize_chunk(type, w.data(), in.weights.data(), 0, rows, input_size, nullptr); + in.activations.resize(static_cast(input_size * token_count)); + for (float & v : in.activations) { + v = uniform(rng); + } + // ids [kExpertCount, token_count] (an argsort-like layout); MUL_MAT_ID reads the first kRouteCount of each column + in.ids.resize(static_cast(kExpertCount * token_count)); + for (int64_t t = 0; t < token_count; ++t) { + std::vector order(kExpertCount); + for (int64_t e = 0; e < kExpertCount; ++e) { + order[e] = static_cast(e); + } + std::shuffle(order.begin(), order.end(), rng); + for (int64_t e = 0; e < kExpertCount; ++e) { + in.ids[t * kExpertCount + e] = order[e]; + } + } + return in; +} + +static std::vector run(ggml_backend_t backend, ggml_type type, int64_t input_size, int64_t token_count, + const Inputs & in, std::string * hrx_kernel) { + ggml_init_params params = { 32 * ggml_tensor_overhead() + ggml_graph_overhead(), nullptr, true }; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + ggml_tensor * weights = ggml_new_tensor_3d(ctx, type, input_size, kOutputSize, kExpertCount); + ggml_tensor * x = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, input_size, 1, token_count); + ggml_tensor * ids_all = ggml_new_tensor_2d(ctx, GGML_TYPE_I32, kExpertCount, token_count); + ggml_tensor * ids = ggml_view_2d(ctx, ids_all, kRouteCount, token_count, ids_all->nb[1], 0); + ggml_tensor * out = ggml_mul_mat_id(ctx, weights, x, ids); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, out); + REQUIRE(ggml_backend_supports_op(backend, out)); + if (hrx_kernel != nullptr) { + *hrx_kernel = hrx_mul_mat_id_kernel(graph); + } + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + ggml_backend_tensor_set(weights, in.weights.data(), 0, in.weights.size()); + ggml_backend_tensor_set(x, in.activations.data(), 0, in.activations.size() * sizeof(float)); + ggml_backend_tensor_set(ids_all, in.ids.data(), 0, in.ids.size() * sizeof(int32_t)); + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + std::vector result(static_cast(ggml_nelements(out))); + ggml_backend_tensor_get(out, result.data(), 0, result.size() * sizeof(float)); + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + return result; +} + +static double nmse(const std::vector & got, const std::vector & expected) { + double err = 0.0; + double ref = 0.0; + for (size_t i = 0; i < got.size(); ++i) { + const double d = static_cast(got[i]) - expected[i]; + err += d * d; + ref += static_cast(expected[i]) * expected[i]; + } + return ref > 0.0 ? err / ref : err; +} + +int main() { + ggml_backend_dev_t device = ggml_backend_dev_by_name("HRX0"); + if (device == nullptr) { + ggml_backend_load_all(); + device = ggml_backend_dev_by_name("HRX0"); + } + if (device == nullptr) { + std::printf("test-hrx-mul-mat-id-k32: no HRX0 device, skipped\n"); + return 0; + } + ggml_backend_t hrx = ggml_backend_dev_init(device, nullptr); + ggml_backend_t cpu = ggml_backend_cpu_init(); + REQUIRE(hrx != nullptr && cpu != nullptr); + + const ggml_type types[] = { GGML_TYPE_MXFP4, GGML_TYPE_Q8_0, GGML_TYPE_Q4_0 }; + const int64_t input_sizes[] = { 2880, 1152, 2816 }; + const int64_t token_counts[] = { 1, 7, 40 }; + int failures = 0; + uint32_t seed = 1; + for (ggml_type type : types) { + for (int64_t input_size : input_sizes) { + for (int64_t token_count : token_counts) { + const Inputs in = make_inputs(type, input_size, token_count, seed++); + std::string kernel; + std::vector got = run(hrx, type, input_size, token_count, in, &kernel); + std::vector expected = run(cpu, type, input_size, token_count, in, nullptr); + const double e = nmse(got, expected); + const bool ok = !kernel.empty() && e <= kMaxNmse; + std::printf("%-5s k=%5lld tokens=%2lld %-48s nmse=%.3g %s\n", ggml_type_name(type), + static_cast(input_size), static_cast(token_count), + kernel.empty() ? "(no HRX mul_mat_id kernel)" : kernel.c_str(), e, ok ? "OK" : "FAIL"); + failures += ok ? 0 : 1; + } + } + } + ggml_backend_free(cpu); + ggml_backend_free(hrx); + REQUIRE(failures == 0); + std::printf("test-hrx-mul-mat-id-k32: all cases OK\n"); + return 0; +} diff --git a/tests/test-hrx-mxfp4.cpp b/tests/test-hrx-mxfp4.cpp new file mode 100644 index 000000000000..d815fa129e53 --- /dev/null +++ b/tests/test-hrx-mxfp4.cpp @@ -0,0 +1,165 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Known-answer check for MXFP4 weights on HRX: GET_ROWS of MXFP4 rows must equal ggml's dequantize_row_mxfp4 +// bit for bit (the E8M0 half scale is a power of two and every E2M1 value is exact). Rows use the exponents +// 120, 127, 134 and the edges 0, 1, 2, 254; every block holds all 16 codes in both nibbles. Exponents 0 and 1 +// have f32 subnormal scales (2^-128 and 2^-127), which the GPU kernels flush to zero, so those rows' values +// (at most 12 * 2^-127) may come back as zeros of the same sign; every other value must match exactly. The graph runs on the HRX device itself (no scheduler, so no CPU fallback), and the HRX +// dispatch plan for it must contain the get_rows kernel. + +#include "dispatch/dispatch-scheduler.h" +#include "ggml-alloc.h" +#include "ggml-backend.h" +#include "ggml.h" +#include "graph/graph.h" +#include "kernel-corpus/kernel-corpus.h" + +#include +#include +#include +#include +#include +#include +#include + +#define REQUIRE(condition) \ + do { \ + if (!(condition)) { \ + std::fprintf(stderr, "%s:%d: requirement failed: %s\n", __FILE__, __LINE__, #condition); \ + std::abort(); \ + } \ + } while (false) + +static constexpr int kExponents[] = { 120, 127, 134, 0, 1, 2, 254 }; +static constexpr int kRows = sizeof(kExponents) / sizeof(kExponents[0]); +static constexpr int kValues = 256; // eight 32-value blocks per row +static constexpr int kBlocks = kValues / 32; +static constexpr int kBlockBytes = 17; // e (E8M0), qs[16] +static constexpr int kRowBytes = kBlocks * kBlockBytes; + +static std::vector make_weights() { + std::vector weights(static_cast(kRowBytes) * kRows); + for (int r = 0; r < kRows; ++r) { + for (int b = 0; b < kBlocks; ++b) { + uint8_t * block = weights.data() + static_cast(r) * kRowBytes + b * kBlockBytes; + block[0] = static_cast(kExponents[r]); + for (int j = 0; j < 16; ++j) { + const int low = (j + b) & 15; + const int high = (15 - j + 3 * b) & 15; + block[1 + j] = static_cast(low | (high << 4)); + } + } + } + return weights; +} + +static std::string kernel_name_for_id(uint64_t kernel_id) { + const ggml::hrx::KernelResolveResult resolved = + ggml::hrx::resolve_kernel_definition(ggml::hrx::get_qwen_kernel_corpus(), "gfx1151", kernel_id); + REQUIRE(resolved.found()); + return ggml::hrx::kernel_definition_name(*resolved.definition); +} + +// The HRX dispatch plan for the graph: every node must be covered, and a get_rows kernel must run. +static void require_hrx_get_rows_plan(ggml_cgraph * graph) { + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + ggml::hrx::DispatchScheduleDiagnostics diagnostics; + if (!scheduler.schedule_graph(imported.graph, { "gfx1151" }, &diagnostics)) { + std::fprintf(stderr, "unsupported: %s\n", diagnostics.unsupported_message.c_str()); + std::abort(); + } + REQUIRE(scheduler.plan().valid()); + bool found = false; + for (const ggml::hrx::Dispatch & dispatch : scheduler.plan().dispatches) { + const std::string name = kernel_name_for_id(dispatch.kernel.kernel_id); + std::printf("dispatch: %s\n", name.c_str()); + found = found || name.find("get_rows") != std::string::npos; + } + REQUIRE(found); +} + +int main() { + ggml_backend_dev_t device = ggml_backend_dev_by_name("HRX0"); + if (device == nullptr) { + ggml_backend_load_all(); + device = ggml_backend_dev_by_name("HRX0"); + } + if (device == nullptr) { + std::printf("test-hrx-mxfp4: no HRX0 device, skipped\n"); + return 0; + } + ggml_backend_t backend = ggml_backend_dev_init(device, nullptr); + REQUIRE(backend != nullptr); + + ggml_init_params params = { 16 * ggml_tensor_overhead() + ggml_graph_overhead(), nullptr, true }; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + ggml_tensor * weights = ggml_new_tensor_2d(ctx, GGML_TYPE_MXFP4, kValues, kRows); + ggml_tensor * ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, kRows); + ggml_tensor * rows = ggml_get_rows(ctx, weights, ids); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, rows); + REQUIRE(ggml_backend_supports_op(backend, rows)); + require_hrx_get_rows_plan(graph); + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + const std::vector host_weights = make_weights(); + std::vector host_ids(kRows); + for (int r = 0; r < kRows; ++r) { + host_ids[r] = kRows - 1 - r; + } + ggml_backend_tensor_set(weights, host_weights.data(), 0, host_weights.size()); + ggml_backend_tensor_set(ids, host_ids.data(), 0, host_ids.size() * sizeof(int32_t)); + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + + std::vector got(static_cast(kValues) * kRows); + std::vector expected(kValues); + ggml_backend_tensor_get(rows, got.data(), 0, got.size() * sizeof(float)); + const ggml_type_traits * traits = ggml_get_type_traits(GGML_TYPE_MXFP4); + REQUIRE(traits->to_float != nullptr); + int mismatches = 0; + int flushed = 0; + for (int r = 0; r < kRows; ++r) { + const int source = host_ids[r]; + traits->to_float(host_weights.data() + static_cast(source) * kRowBytes, expected.data(), kValues); + for (int k = 0; k < kValues; ++k) { + const float value = got[static_cast(r) * kValues + k]; + if (std::memcmp(&value, &expected[k], sizeof(float)) == 0) { + continue; + } + if (kExponents[source] < 2 && value == 0.0f && std::signbit(value) == std::signbit(expected[k])) { + ++flushed; + continue; + } + if (mismatches < 8) { + std::fprintf(stderr, "e=%d value %d: got %.9g, expected %.9g\n", kExponents[source], k, value, + expected[k]); + } + ++mismatches; + } + } + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(backend); + REQUIRE(mismatches == 0); + std::printf("test-hrx-mxfp4: %d rows x %d values bit-exact (%d values with a subnormal scale flushed to zero)\n", + kRows, kValues, flushed); + return 0; +} diff --git a/tests/test-hrx-ops.cpp b/tests/test-hrx-ops.cpp new file mode 100644 index 000000000000..12064f3f63c2 --- /dev/null +++ b/tests/test-hrx-ops.cpp @@ -0,0 +1,6515 @@ +#include "backend-context.h" +#include "dispatch/dispatch-scheduler.h" +#include "dispatch_registration/common/dispatch-activation-publication.h" +#include "ggml-alloc.h" +#include "ggml-backend.h" +#include "ggml-hrx.h" +#include "ggml-quants.h" +#include "ggml.h" +#include "graph/graph.h" +#include "kernel-corpus/kernel-corpus.h" +#include "runtime/graph-executor.h" +#include "testing_suite.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#define REQUIRE(condition) \ + do { \ + if (!(condition)) { \ + std::fprintf(stderr, "%s:%d: requirement failed: %s\n", __FILE__, __LINE__, #condition); \ + std::abort(); \ + } \ + } while (false) + +static void restore_environment_value(const char * name, bool had_value, const std::string & value) { + if (had_value) { + REQUIRE(setenv(name, value.c_str(), 1) == 0); + } else { + REQUIRE(unsetenv(name) == 0); + } +} + +static std::string expected_config_value(float value) { + std::ostringstream out; + out.precision(9); + out << value; + return out.str(); +} + +static constexpr float kQwenRmsNormEps = 0.000001f; +static constexpr int64_t kQwenFlashHeadSize = 128; +static constexpr int64_t kQwenRouterExpertCount = 128; +static constexpr int64_t kQwenRouterRouteCount = 8; +static constexpr int64_t kQwenHiddenSize = 2048; +static constexpr int64_t kQwenMoeIntermediate = 768; +static constexpr int64_t kQwenVocabularyCount = 151936; +static constexpr int64_t kGemmaHiddenSize = 3840; +static constexpr int64_t kGemmaPromptTokenCount = 18; +static constexpr float kGemmaRmsNormEps = 0.000001f; + +static ggml::hrx::Value make_test_value(ggml::hrx::ValueId id, + ggml::hrx::ValueStorageId storage, + ggml::hrx::ValueId storage_root, + ggml::hrx::ValueId alias_source, + size_t storage_offset, + size_t storage_byte_count, + ggml_type type, + int64_t element_count) { + ggml::hrx::Value value = {}; + value.id = id; + value.kind = ggml::hrx::ValueKind::Transient; + value.storage = storage; + value.storage_root = storage_root; + value.alias_source = alias_source; + value.storage_offset = storage_offset; + value.storage_byte_count = storage_byte_count; + value.type = type; + value.ne = { element_count, 1, 1, 1 }; + value.nb = { ggml_type_size(type), ggml_type_size(type) * static_cast(element_count), + ggml_type_size(type) * static_cast(element_count), + ggml_type_size(type) * static_cast(element_count) }; + value.element_count = element_count; + value.byte_count = ggml_row_size(type, element_count); + value.contiguous = true; + return value; +} + +static std::vector make_input(int64_t hidden_size, int64_t token_count) { + std::vector data(hidden_size * token_count); + for (int64_t i = 0; i < static_cast(data.size()); ++i) { + data[i] = static_cast((i % 29) - 14) * 0.125f; + } + return data; +} + +static std::vector make_weight(int64_t hidden_size) { + std::vector data(hidden_size); + for (int64_t i = 0; i < hidden_size; ++i) { + data[i] = 0.5f + static_cast(i % 17) * 0.03125f; + } + return data; +} + +static std::vector make_router_input(int64_t hidden_size, int64_t token_count) { + std::vector data(hidden_size * token_count); + for (int64_t i = 0; i < static_cast(data.size()); ++i) { + data[i] = static_cast((i % 41) - 20) * 0.01f; + } + return data; +} + +static std::vector make_router_weight(int64_t hidden_size, int64_t expert_count) { + std::vector data(hidden_size * expert_count); + for (int64_t expert = 0; expert < expert_count; ++expert) { + for (int64_t column = 0; column < hidden_size; ++column) { + data[expert * hidden_size + column] = static_cast(((expert + column) % 31) - 15) * 0.0025f; + } + } + return data; +} + +static std::vector make_router_logits(int64_t token_count) { + std::vector data(kQwenRouterExpertCount * token_count); + for (int64_t i = 0; i < static_cast(data.size()); ++i) { + data[i] = static_cast(i % kQwenRouterRouteCount); + } + return data; +} + +static std::vector make_flash_query(int64_t token_count, int64_t head_size = kQwenFlashHeadSize) { + std::vector data(static_cast(head_size) * static_cast(token_count)); + for (int64_t i = 0; i < static_cast(data.size()); ++i) { + data[i] = static_cast((i % 37) - 18) * 0.01f; + } + return data; +} + +static std::vector make_flash_key_value(int64_t token_count, + int offset, + int64_t head_size = kQwenFlashHeadSize) { + std::vector data(static_cast(head_size) * static_cast(token_count)); + for (int64_t i = 0; i < static_cast(data.size()); ++i) { + const float value = static_cast(((i + offset) % 31) - 15) * 0.015f; + data[i] = ggml_fp32_to_fp16(value); + } + return data; +} + +static std::vector make_flash_mask(int64_t query_token_count, int64_t key_value_token_count) { + std::vector data(query_token_count * key_value_token_count); + for (int64_t query = 0; query < query_token_count; ++query) { + for (int64_t key = 0; key < key_value_token_count; ++key) { + const float value = key <= query + 1 ? 0.0f : -10000.0f; + data[query * key_value_token_count + key] = ggml_fp32_to_fp16(value); + } + } + return data; +} + +static std::vector rmsnorm_mul_reference(const std::vector & input, + const std::vector & weight, + int64_t hidden_size, + int64_t token_count) { + std::vector output(input.size()); + for (int64_t token = 0; token < token_count; ++token) { + float sum_squares = 0.0f; + for (int64_t column = 0; column < hidden_size; ++column) { + const float value = input[token * hidden_size + column]; + sum_squares += value * value; + } + const float scale = 1.0f / std::sqrt(sum_squares / static_cast(hidden_size) + kQwenRmsNormEps); + for (int64_t column = 0; column < hidden_size; ++column) { + output[token * hidden_size + column] = input[token * hidden_size + column] * scale * weight[column]; + } + } + return output; +} + +static std::vector rmsnorm_reference(const std::vector & input, + int64_t hidden_size, + int64_t token_count, + float epsilon) { + std::vector output(input.size()); + for (int64_t token = 0; token < token_count; ++token) { + float sum_squares = 0.0f; + for (int64_t column = 0; column < hidden_size; ++column) { + const float value = input[token * hidden_size + column]; + sum_squares += value * value; + } + const float scale = 1.0f / std::sqrt(sum_squares / static_cast(hidden_size) + epsilon); + for (int64_t column = 0; column < hidden_size; ++column) { + output[token * hidden_size + column] = input[token * hidden_size + column] * scale; + } + } + return output; +} + +static std::vector router_projection_reference(const std::vector & input, + const std::vector & weight, + int64_t hidden_size, + int64_t expert_count, + int64_t token_count) { + std::vector output(expert_count * token_count); + for (int64_t token = 0; token < token_count; ++token) { + for (int64_t expert = 0; expert < expert_count; ++expert) { + float sum = 0.0f; + for (int64_t column = 0; column < hidden_size; ++column) { + sum += input[token * hidden_size + column] * weight[expert * hidden_size + column]; + } + output[token * expert_count + expert] = sum; + } + } + return output; +} + +static std::vector router_top8_weights_reference(const std::vector & logits, int64_t token_count) { + std::vector output(kQwenRouterRouteCount * token_count); + for (int64_t token = 0; token < token_count; ++token) { + bool used[kQwenRouterExpertCount] = {}; + int64_t selected[kQwenRouterRouteCount] = {}; + for (int64_t route = 0; route < kQwenRouterRouteCount; ++route) { + int64_t best_expert = -1; + float best_value = -std::numeric_limits::infinity(); + for (int64_t expert = 0; expert < kQwenRouterExpertCount; ++expert) { + const float value = logits[token * kQwenRouterExpertCount + expert]; + if (!used[expert] && + (best_expert < 0 || value > best_value || (value == best_value && expert < best_expert))) { + best_value = value; + best_expert = expert; + } + } + selected[route] = best_expert; + used[best_expert] = true; + } + + float max_selected = -std::numeric_limits::infinity(); + for (const int64_t expert : selected) { + max_selected = std::max(max_selected, logits[token * kQwenRouterExpertCount + expert]); + } + float sum = 0.0f; + for (int64_t route = 0; route < kQwenRouterRouteCount; ++route) { + const float value = std::exp(logits[token * kQwenRouterExpertCount + selected[route]] - max_selected); + output[token * kQwenRouterRouteCount + route] = value; + sum += value; + } + for (int64_t route = 0; route < kQwenRouterRouteCount; ++route) { + output[token * kQwenRouterRouteCount + route] /= sum; + } + } + return output; +} + +static std::vector flash_attention_reference(const std::vector & query, + const std::vector & key, + const std::vector & value, + const std::vector & mask, + int64_t query_token_count, + int64_t key_value_token_count, + int64_t qk_head_size = kQwenFlashHeadSize, + int64_t value_head_size = -1, + float scale = 0.0f) { + const int64_t actual_value_head_size = value_head_size > 0 ? value_head_size : qk_head_size; + std::vector output(static_cast(query_token_count) * static_cast(actual_value_head_size)); + const float actual_scale = scale == 0.0f ? 1.0f / std::sqrt(static_cast(qk_head_size)) : scale; + for (int64_t query_token = 0; query_token < query_token_count; ++query_token) { + std::vector scores(key_value_token_count); + float max_score = -std::numeric_limits::infinity(); + for (int64_t key_token = 0; key_token < key_value_token_count; ++key_token) { + float dot = 0.0f; + for (int64_t channel = 0; channel < qk_head_size; ++channel) { + dot += query[query_token * qk_head_size + channel] * + ggml_fp16_to_fp32(key[key_token * qk_head_size + channel]); + } + const float score = + dot * actual_scale + ggml_fp16_to_fp32(mask[query_token * key_value_token_count + key_token]); + scores[key_token] = score; + max_score = std::max(max_score, score); + } + + float sum = 0.0f; + for (float & score : scores) { + score = std::exp(score - max_score); + sum += score; + } + for (int64_t channel = 0; channel < actual_value_head_size; ++channel) { + float weighted_sum = 0.0f; + for (int64_t key_token = 0; key_token < key_value_token_count; ++key_token) { + const float probability = scores[key_token] / sum; + weighted_sum += probability * ggml_fp16_to_fp32(value[key_token * actual_value_head_size + channel]); + } + output[query_token * actual_value_head_size + channel] = weighted_sum; + } + } + return output; +} + +static ggml_tensor * build_rmsnorm_mul_graph(ggml_context * ctx, + ggml_tensor * input, + ggml_tensor * weight, + float eps = kQwenRmsNormEps) { + ggml_tensor * rms = ggml_rms_norm(ctx, input, eps); + REQUIRE(rms != nullptr); + ggml_tensor * output = ggml_mul(ctx, rms, weight); + REQUIRE(output != nullptr); + return output; +} + +static ggml_tensor * build_qwen_flash_attention_graph(ggml_context * ctx, + ggml_tensor * query, + ggml_tensor * key, + ggml_tensor * value, + ggml_tensor * mask, + int64_t head_size = kQwenFlashHeadSize, + float scale = 0.0f) { + const float actual_scale = scale == 0.0f ? 1.0f / std::sqrt(static_cast(head_size)) : scale; + ggml_tensor * output = ggml_flash_attn_ext(ctx, query, key, value, mask, actual_scale, 0.0f, 0.0f); + REQUIRE(output != nullptr); + return output; +} + +static ggml_tensor * build_qwen_router_top8_graph(ggml_context * ctx, + ggml_tensor * logits, + ggml_tensor ** route_ids = nullptr) { + ggml_tensor * probs = ggml_soft_max(ctx, logits); + REQUIRE(probs != nullptr); + ggml_tensor * probs_reshaped = ggml_reshape_3d(ctx, probs, 1, kQwenRouterExpertCount, logits->ne[1]); + REQUIRE(probs_reshaped != nullptr); + ggml_tensor * argsort = ggml_argsort(ctx, probs, GGML_SORT_ORDER_DESC); + REQUIRE(argsort != nullptr); + ggml_tensor * topk = ggml_view_2d(ctx, argsort, kQwenRouterRouteCount, logits->ne[1], argsort->nb[1], 0); + REQUIRE(topk != nullptr); + if (route_ids != nullptr) { + *route_ids = topk; + } + ggml_tensor * selected = ggml_get_rows(ctx, probs_reshaped, topk); + REQUIRE(selected != nullptr); + ggml_tensor * selected_reshaped = ggml_reshape_2d(ctx, selected, kQwenRouterRouteCount, logits->ne[1]); + REQUIRE(selected_reshaped != nullptr); + ggml_tensor * sum = ggml_sum_rows(ctx, selected_reshaped); + REQUIRE(sum != nullptr); + ggml_tensor * clamped_sum = ggml_clamp(ctx, sum, 1.0e-7f, std::numeric_limits::infinity()); + REQUIRE(clamped_sum != nullptr); + ggml_tensor * normalized = ggml_div(ctx, selected_reshaped, clamped_sum); + REQUIRE(normalized != nullptr); + ggml_tensor * output = ggml_reshape_3d(ctx, normalized, 1, kQwenRouterRouteCount, logits->ne[1]); + REQUIRE(output != nullptr); + return output; +} + +static std::string kernel_name_for_id(uint64_t kernel_id) { + const ggml::hrx::KernelResolveResult resolved = + ggml::hrx::resolve_kernel_definition(ggml::hrx::get_qwen_kernel_corpus(), "gfx1151", kernel_id); + REQUIRE(resolved.found()); + return ggml::hrx::kernel_definition_name(*resolved.definition); +} + +static std::vector scheduled_kernel_sequence(ggml_cgraph * graph) { + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + ggml::hrx::DispatchScheduleDiagnostics diagnostics; + if (!scheduler.schedule_graph(imported.graph, { "gfx1151" }, &diagnostics)) { + for (const std::string & error : scheduler.plan().status.errors()) { + std::fprintf(stderr, "scheduler error: %s\n", error.c_str()); + } + std::fprintf(stderr, "unsupported: %s\n", diagnostics.unsupported_message.c_str()); + for (const ggml::hrx::DispatchRegistrationAttempt & attempt : diagnostics.match.attempts) { + std::fprintf(stderr, " attempt %s matched=%d\n", attempt.name.c_str(), attempt.matched ? 1 : 0); + if (!attempt.covered_nodes.empty()) { + std::fprintf(stderr, " covered:"); + for (size_t node : attempt.covered_nodes) { + std::fprintf(stderr, " %zu", node); + } + std::fprintf(stderr, "\n"); + } + for (const std::string & error : attempt.errors) { + std::fprintf(stderr, " %s\n", error.c_str()); + } + } + std::abort(); + } + REQUIRE(scheduler.plan().valid()); + + std::vector names; + names.reserve(scheduler.plan().dispatches.size()); + for (const ggml::hrx::Dispatch & dispatch : scheduler.plan().dispatches) { + names.push_back(kernel_name_for_id(dispatch.kernel.kernel_id)); + } + return names; +} + +static ggml::hrx::KernelSpecialization scheduled_kernel_specialization(ggml_cgraph * graph, + const char * kernel_name) { + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, { "gfx1151" })); + REQUIRE(scheduler.plan().valid()); + for (const ggml::hrx::Dispatch & dispatch : scheduler.plan().dispatches) { + if (kernel_name_for_id(dispatch.kernel.kernel_id) == kernel_name) { + return dispatch.kernel; + } + } + REQUIRE(false); + return {}; +} + +static size_t producer_index_for_tensor(const ggml::hrx::Graph & graph, const ggml_tensor * tensor) { + const ggml::hrx::Value * value = graph.values().find_tensor(tensor); + REQUIRE(value != nullptr); + const ggml::hrx::GraphNode * producer = graph.index().producer(value->id); + REQUIRE(producer != nullptr); + size_t index = 0; + REQUIRE(graph.index().node_index(producer, index)); + return index; +} + +static ggml::hrx::ValueId next_plan_value(const ggml::hrx::Graph & graph, const ggml::hrx::CommandPlan & plan) { + return ggml::hrx::ValueId( + static_cast(graph.values().size() + plan.transients.size() + plan.completion_counter_requests.size())); +} + +static void append_match_to_plan(ggml::hrx::CommandPlan & plan, + ggml::hrx::DispatchMatch & match, + std::vector & covered_nodes) { + for (ggml::hrx::Dispatch & dispatch : match.initialization_dispatches) { + plan.initialization_dispatches.push_back(std::move(dispatch)); + } + for (ggml::hrx::Dispatch & dispatch : match.dispatches) { + plan.dispatches.push_back(std::move(dispatch)); + } + for (ggml::hrx::CommandPlanTransient & transient : match.transients) { + plan.transients.push_back(std::move(transient)); + } + for (ggml::hrx::CommandPlanConstantInitialization & initialization : match.constant_initializations) { + plan.constant_initializations.push_back(std::move(initialization)); + } + for (ggml::hrx::CommandPlanCompletionCounterRequest & request : match.completion_counter_requests) { + plan.completion_counter_requests.push_back(std::move(request)); + } + REQUIRE(plan.metadata.append(std::move(match.metadata), plan.status)); + plan.status.append(match.status); + for (size_t covered_node : match.covered_nodes) { + REQUIRE(covered_node < covered_nodes.size()); + REQUIRE(!covered_nodes[covered_node]); + covered_nodes[covered_node] = true; + } +} + +static void match_dispatch_at_index(const ggml::hrx::Graph & graph, + const ggml::hrx::DispatchRegistry & registry, + ggml::hrx::CommandPlan & plan, + std::vector & covered_nodes, + size_t root_index, + ggml::hrx::DispatchMatch & match) { + REQUIRE(root_index < graph.nodes().size()); + const ggml::hrx::DispatchMatchContext context = { + graph, &graph.nodes()[root_index], root_index, covered_nodes, plan, next_plan_value(graph, plan), ®istry, + }; + ggml::hrx::DispatchMatchDiagnostics diagnostics; + if (!registry.match(context, match, &diagnostics)) { + std::fprintf(stderr, "manual matcher failed for node %zu %s\n", root_index, + ggml_op_name(graph.nodes()[root_index].op)); + const ggml::hrx::GraphNode & node = graph.nodes()[root_index]; + for (size_t input_index = 0; input_index < node.inputs.size(); ++input_index) { + const ggml::hrx::Value * input = graph.values().find(node.inputs[input_index]); + if (input != nullptr) { + std::fprintf(stderr, + " input %zu value=%d type=%d ne=[%" PRId64 ",%" PRId64 ",%" PRId64 ",%" PRId64 + "] nb=[%zu,%zu,%zu,%zu]\n", + input_index, input->id.value, static_cast(input->type), input->ne[0], input->ne[1], + input->ne[2], input->ne[3], input->nb[0], input->nb[1], input->nb[2], input->nb[3]); + } + } + const ggml::hrx::Value * output = graph.values().find(node.output); + if (output != nullptr) { + std::fprintf(stderr, + " output value=%d type=%d ne=[%" PRId64 ",%" PRId64 ",%" PRId64 ",%" PRId64 + "] nb=[%zu,%zu,%zu,%zu]\n", + output->id.value, static_cast(output->type), output->ne[0], output->ne[1], output->ne[2], + output->ne[3], output->nb[0], output->nb[1], output->nb[2], output->nb[3]); + } + for (const ggml::hrx::CommandPlanAlternateValue & alternate : plan.metadata.alternate_values()) { + std::fprintf(stderr, " alternate graph=%d value=%d type=%d bytes=%zu name=%s\n", + alternate.graph_value.value, alternate.alternate_value.value, static_cast(alternate.type), + alternate.byte_count, alternate.name.c_str()); + } + for (const ggml::hrx::DispatchRegistrationAttempt & attempt : diagnostics.attempts) { + std::fprintf(stderr, " attempt %s matched=%d\n", attempt.name.c_str(), attempt.matched ? 1 : 0); + for (const std::string & error : attempt.errors) { + std::fprintf(stderr, " %s\n", error.c_str()); + } + } + std::abort(); + } + append_match_to_plan(plan, match, covered_nodes); + REQUIRE(plan.valid()); +} + +static void require_kernel_subsequence(const std::vector & sequence, + const std::vector & expected) { + size_t sequence_index = 0; + for (const std::string & name : expected) { + while (sequence_index < sequence.size() && sequence[sequence_index] != name) { + ++sequence_index; + } + if (sequence_index >= sequence.size()) { + std::fprintf(stderr, "missing expected kernel: %s\nscheduled kernels:\n", name.c_str()); + for (const std::string & scheduled : sequence) { + std::fprintf(stderr, " %s\n", scheduled.c_str()); + } + std::abort(); + } + ++sequence_index; + } +} + +static void run_alternate_value_alias_lookup_checks() { + constexpr int64_t element_count = 2048; + const size_t full_bytes = ggml_row_size(GGML_TYPE_F32, element_count); + const size_t f16_bytes = ggml_row_size(GGML_TYPE_F16, element_count); + const size_t q8_bytes = ggml_row_size(GGML_TYPE_Q8_1, element_count); + const size_t symmetric_bytes = 1536; + + ggml::hrx::Graph graph; + ggml::hrx::Status status; + ggml::hrx::CommandPlan plan; + + const ggml::hrx::ValueId root(0); + const ggml::hrx::ValueId full_alias(1); + const ggml::hrx::ValueId partial_alias(2); + const ggml::hrx::ValueId reordered_alias(3); + const ggml::hrx::ValueId unrelated_same_storage(4); + const ggml::hrx::ValueId different_shape_alias(5); + const ggml::hrx::ValueId incompatible_q8_shape_alias(6); + const ggml::hrx::ValueId f16_alternate(99); + const ggml::hrx::ValueId q8_alternate(100); + const ggml::hrx::ValueStorageId storage(0); + + status = graph.values().add_snapshot_storage({ storage, root, full_bytes }); + REQUIRE(status.success()); + status = graph.values().add_snapshot_value( + make_test_value(root, storage, root, ggml::hrx::ValueId(), 0, full_bytes, GGML_TYPE_F32, element_count)); + REQUIRE(status.success()); + status = graph.values().add_snapshot_value( + make_test_value(full_alias, storage, root, root, 0, full_bytes, GGML_TYPE_F32, element_count)); + REQUIRE(status.success()); + status = graph.values().add_snapshot_value( + make_test_value(partial_alias, storage, root, root, 0, full_bytes, GGML_TYPE_F32, element_count / 2)); + REQUIRE(status.success()); + ggml::hrx::Value reordered = + make_test_value(reordered_alias, storage, root, root, 0, full_bytes, GGML_TYPE_F32, element_count); + reordered.contiguous = false; + status = graph.values().add_snapshot_value(std::move(reordered)); + REQUIRE(status.success()); + status = graph.values().add_snapshot_value( + make_test_value(unrelated_same_storage, storage, root, root, 0, full_bytes, GGML_TYPE_F32, element_count)); + REQUIRE(status.success()); + ggml::hrx::Value different_shape = + make_test_value(different_shape_alias, storage, root, root, 0, full_bytes, GGML_TYPE_F32, element_count); + different_shape.ne = { element_count / 2, 2, 1, 1 }; + different_shape.nb = { sizeof(float), static_cast(element_count / 2) * sizeof(float), full_bytes, + full_bytes }; + status = graph.values().add_snapshot_value(std::move(different_shape)); + REQUIRE(status.success()); + ggml::hrx::Value incompatible_q8_shape = + make_test_value(incompatible_q8_shape_alias, storage, root, root, 0, full_bytes, GGML_TYPE_F32, element_count); + incompatible_q8_shape.ne = { 64, element_count / 64, 1, 1 }; + incompatible_q8_shape.nb = { sizeof(float), 64 * sizeof(float), full_bytes, full_bytes }; + status = graph.values().add_snapshot_value(std::move(incompatible_q8_shape)); + REQUIRE(status.success()); + + graph.add_node(GGML_OP_RESHAPE, full_alias, { root }); + graph.add_node(GGML_OP_VIEW, partial_alias, { root }); + graph.add_node(GGML_OP_TRANSPOSE, reordered_alias, { root }); + graph.add_node(GGML_OP_CPY, unrelated_same_storage, { root }); + graph.add_node(GGML_OP_RESHAPE, different_shape_alias, { root }); + graph.add_node(GGML_OP_RESHAPE, incompatible_q8_shape_alias, { root }); + status = graph.build_index(); + REQUIRE(status.success()); + + REQUIRE(plan.metadata.append_alternate_value({ root, q8_alternate, GGML_TYPE_Q8_1, q8_bytes, "q8" }, status)); + REQUIRE(plan.metadata.append_alternate_value({ root, f16_alternate, GGML_TYPE_F16, f16_bytes, "f16" }, status)); + REQUIRE(plan.metadata.append_alternate_value( + { root, ggml::hrx::ValueId(104), GGML_TYPE_COUNT, symmetric_bytes, "symmetric-i4" }, status)); + + const ggml::hrx::CommandPlanAlternateValue * exact = + ggml::hrx::find_alternate_value(graph, plan, root, GGML_TYPE_Q8_1, q8_bytes); + REQUIRE(exact != nullptr); + REQUIRE(exact->alternate_value == q8_alternate); + + const ggml::hrx::CommandPlanAlternateValue * through_full_alias = + ggml::hrx::find_alternate_value(graph, plan, full_alias, GGML_TYPE_Q8_1, q8_bytes); + REQUIRE(through_full_alias != nullptr); + REQUIRE(through_full_alias->alternate_value == q8_alternate); + + const ggml::hrx::CommandPlanAlternateValue * through_partial_alias = + ggml::hrx::find_alternate_value(graph, plan, partial_alias, GGML_TYPE_Q8_1, q8_bytes); + REQUIRE(through_partial_alias == nullptr); + + const ggml::hrx::CommandPlanAlternateValue * through_reordered_alias = + ggml::hrx::find_alternate_value(graph, plan, reordered_alias, GGML_TYPE_Q8_1, q8_bytes); + REQUIRE(through_reordered_alias == nullptr); + + REQUIRE(ggml::hrx::find_alternate_value(graph, plan, unrelated_same_storage, GGML_TYPE_Q8_1, q8_bytes) == nullptr); + REQUIRE(ggml::hrx::find_alternate_value(graph, plan, different_shape_alias, GGML_TYPE_Q8_1, q8_bytes) != nullptr); + REQUIRE(ggml::hrx::find_alternate_value(graph, plan, incompatible_q8_shape_alias, GGML_TYPE_Q8_1, q8_bytes) == + nullptr); + + const ggml::hrx::CommandPlanAlternateValue * f16_exact = + ggml::hrx::find_alternate_value(graph, plan, root, GGML_TYPE_F16, f16_bytes); + REQUIRE(f16_exact != nullptr); + REQUIRE(f16_exact->alternate_value == f16_alternate); + REQUIRE(ggml::hrx::find_alternate_value(graph, plan, different_shape_alias, GGML_TYPE_F16, f16_bytes) != nullptr); + REQUIRE(ggml::hrx::find_alternate_value(graph, plan, full_alias, GGML_TYPE_COUNT, symmetric_bytes) != nullptr); + REQUIRE(ggml::hrx::find_alternate_value(graph, plan, different_shape_alias, GGML_TYPE_COUNT, symmetric_bytes) == + nullptr); + + ggml::hrx::Status conflict_status; + REQUIRE(!plan.metadata.append_alternate_value( + { root, ggml::hrx::ValueId(101), GGML_TYPE_Q8_1, q8_bytes, "q8-conflict" }, conflict_status)); + + ggml::hrx::Status generated_status; + const ggml::hrx::ValueId k16_value(102); + REQUIRE(plan.metadata.append_generated_resource( + { root, ggml::hrx::GeneratedResourceRole::F16K16Major, k16_value, f16_bytes, {} }, generated_status)); + const ggml::hrx::CommandPlanGeneratedResource * generated_through_alias = ggml::hrx::find_generated_resource( + graph, plan, full_alias, ggml::hrx::GeneratedResourceRole::F16K16Major, f16_bytes); + REQUIRE(generated_through_alias != nullptr); + REQUIRE(generated_through_alias->generated_value == k16_value); + REQUIRE(ggml::hrx::find_generated_resource(graph, plan, partial_alias, + ggml::hrx::GeneratedResourceRole::F16K16Major, f16_bytes) == nullptr); + REQUIRE(ggml::hrx::find_generated_resource(graph, plan, reordered_alias, + ggml::hrx::GeneratedResourceRole::F16K16Major, f16_bytes) == nullptr); + REQUIRE(ggml::hrx::find_generated_resource(graph, plan, different_shape_alias, + ggml::hrx::GeneratedResourceRole::F16K16Major, f16_bytes) == nullptr); +} + +static bool publication_accepts_add(const ggml::hrx::DispatchMatchContext &, + const ggml::hrx::GraphNode & consumer, + const ggml::hrx::Value & input) { + return consumer.op == GGML_OP_ADD && + std::find(consumer.inputs.begin(), consumer.inputs.end(), input.id) != consumer.inputs.end(); +} + +static bool publication_accepts_mul(const ggml::hrx::DispatchMatchContext &, + const ggml::hrx::GraphNode & consumer, + const ggml::hrx::Value & input) { + return consumer.op == GGML_OP_MUL && + std::find(consumer.inputs.begin(), consumer.inputs.end(), input.id) != consumer.inputs.end(); +} + +static bool publication_accepts_mul_output_13(const ggml::hrx::DispatchMatchContext & context, + const ggml::hrx::GraphNode & consumer, + const ggml::hrx::Value & input) { + return consumer.output == ggml::hrx::ValueId(13) && publication_accepts_mul(context, consumer, input); +} + +static bool publication_accepts_mul_output_15(const ggml::hrx::DispatchMatchContext & context, + const ggml::hrx::GraphNode & consumer, + const ggml::hrx::Value & input) { + return consumer.output == ggml::hrx::ValueId(15) && publication_accepts_mul(context, consumer, input); +} + +static void run_activation_publication_contract_checks() { + constexpr int64_t row_size = 2048; + constexpr int64_t row_count = 2; + constexpr int64_t elements = row_size * row_count; + const size_t f32_bytes = ggml_row_size(GGML_TYPE_F32, elements); + + auto make_activation = [&](ggml::hrx::ValueId id, ggml::hrx::ValueStorageId storage, ggml::hrx::ValueId root, + ggml::hrx::ValueId alias, int64_t element_count = elements) { + ggml::hrx::Value value = make_test_value(id, storage, root, alias, 0, f32_bytes, GGML_TYPE_F32, element_count); + if (element_count == elements) { + value.ne = { row_size, row_count, 1, 1 }; + value.nb = { sizeof(float), static_cast(row_size) * sizeof(float), f32_bytes, f32_bytes }; + } + return value; + }; + + ggml::hrx::Graph graph; + ggml::hrx::Status status; + const ggml::hrx::ValueId root(0); + const ggml::hrx::ValueId reshaped(1); + const ggml::hrx::ValueId partial(2); + const ggml::hrx::ValueId transposed(3); + const ggml::hrx::ValueId different_shape(8); + const ggml::hrx::ValueId incompatible_q8_shape(10); + const ggml::hrx::ValueId shape_back(12); + const ggml::hrx::ValueId mismatched_alias_source(14); + status = graph.values().add_snapshot_storage({ ggml::hrx::ValueStorageId(0), root, f32_bytes }); + REQUIRE(status.success()); + status = graph.values().add_snapshot_value( + make_activation(root, ggml::hrx::ValueStorageId(0), root, ggml::hrx::ValueId())); + REQUIRE(status.success()); + status = graph.values().add_snapshot_value(make_activation(reshaped, ggml::hrx::ValueStorageId(0), root, root)); + REQUIRE(status.success()); + status = graph.values().add_snapshot_value( + make_activation(partial, ggml::hrx::ValueStorageId(0), root, root, elements / 2)); + REQUIRE(status.success()); + ggml::hrx::Value reordered = make_activation(transposed, ggml::hrx::ValueStorageId(0), root, root); + reordered.contiguous = false; + status = graph.values().add_snapshot_value(std::move(reordered)); + REQUIRE(status.success()); + + for (int32_t value = 4; value < 8; ++value) { + const ggml::hrx::ValueStorageId storage(value - 3); + const ggml::hrx::ValueId id(value); + status = graph.values().add_snapshot_storage({ storage, id, f32_bytes }); + REQUIRE(status.success()); + status = graph.values().add_snapshot_value(make_activation(id, storage, id, ggml::hrx::ValueId())); + REQUIRE(status.success()); + } + ggml::hrx::Value different = make_activation(different_shape, ggml::hrx::ValueStorageId(0), root, root); + different.ne = { row_size / 2, row_count * 2, 1, 1 }; + different.nb = { sizeof(float), static_cast(row_size / 2) * sizeof(float), f32_bytes, f32_bytes }; + status = graph.values().add_snapshot_value(std::move(different)); + REQUIRE(status.success()); + status = graph.values().add_snapshot_storage({ ggml::hrx::ValueStorageId(5), ggml::hrx::ValueId(9), f32_bytes }); + REQUIRE(status.success()); + status = graph.values().add_snapshot_value(make_activation(ggml::hrx::ValueId(9), ggml::hrx::ValueStorageId(5), + ggml::hrx::ValueId(9), ggml::hrx::ValueId())); + REQUIRE(status.success()); + ggml::hrx::Value incompatible = make_activation(incompatible_q8_shape, ggml::hrx::ValueStorageId(0), root, root); + incompatible.ne = { 64, elements / 64, 1, 1 }; + incompatible.nb = { sizeof(float), 64 * sizeof(float), f32_bytes, f32_bytes }; + status = graph.values().add_snapshot_value(std::move(incompatible)); + REQUIRE(status.success()); + status = graph.values().add_snapshot_storage({ ggml::hrx::ValueStorageId(6), ggml::hrx::ValueId(11), f32_bytes }); + REQUIRE(status.success()); + status = graph.values().add_snapshot_value(make_activation(ggml::hrx::ValueId(11), ggml::hrx::ValueStorageId(6), + ggml::hrx::ValueId(11), ggml::hrx::ValueId())); + REQUIRE(status.success()); + status = graph.values().add_snapshot_value( + make_activation(shape_back, ggml::hrx::ValueStorageId(0), root, different_shape)); + REQUIRE(status.success()); + status = graph.values().add_snapshot_storage({ ggml::hrx::ValueStorageId(7), ggml::hrx::ValueId(13), f32_bytes }); + REQUIRE(status.success()); + status = graph.values().add_snapshot_value(make_activation(ggml::hrx::ValueId(13), ggml::hrx::ValueStorageId(7), + ggml::hrx::ValueId(13), ggml::hrx::ValueId())); + REQUIRE(status.success()); + ggml::hrx::Value mismatched = make_activation( + mismatched_alias_source, ggml::hrx::ValueStorageId(0), root, root); + mismatched.ne = { row_size / 2, row_count * 2, 1, 1 }; + mismatched.nb = { sizeof(float), static_cast(row_size / 2) * sizeof(float), f32_bytes, f32_bytes }; + status = graph.values().add_snapshot_value(std::move(mismatched)); + REQUIRE(status.success()); + status = graph.values().add_snapshot_storage({ ggml::hrx::ValueStorageId(8), ggml::hrx::ValueId(15), f32_bytes }); + REQUIRE(status.success()); + status = graph.values().add_snapshot_value(make_activation(ggml::hrx::ValueId(15), ggml::hrx::ValueStorageId(8), + ggml::hrx::ValueId(15), ggml::hrx::ValueId())); + REQUIRE(status.success()); + + graph.add_node(GGML_OP_RESHAPE, reshaped, { root }); + graph.add_node(GGML_OP_VIEW, partial, { root }); + graph.add_node(GGML_OP_TRANSPOSE, transposed, { root }); + graph.add_node(GGML_OP_ADD, ggml::hrx::ValueId(4), { root, root }); + graph.add_node(GGML_OP_MUL, ggml::hrx::ValueId(5), { reshaped, reshaped }); + graph.add_node(GGML_OP_MUL, ggml::hrx::ValueId(6), { partial, partial }); + graph.add_node(GGML_OP_MUL, ggml::hrx::ValueId(7), { transposed, transposed }); + graph.add_node(GGML_OP_RESHAPE, different_shape, { root }); + graph.add_node(GGML_OP_MUL, ggml::hrx::ValueId(9), { different_shape, different_shape }); + graph.add_node(GGML_OP_RESHAPE, incompatible_q8_shape, { root }); + graph.add_node(GGML_OP_MUL, ggml::hrx::ValueId(11), { incompatible_q8_shape, incompatible_q8_shape }); + graph.add_node(GGML_OP_RESHAPE, shape_back, { different_shape }); + graph.add_node(GGML_OP_MUL, ggml::hrx::ValueId(13), { shape_back, shape_back }); + graph.add_node(GGML_OP_RESHAPE, mismatched_alias_source, { different_shape }); + graph.add_node(GGML_OP_MUL, ggml::hrx::ValueId(15), { mismatched_alias_source, mismatched_alias_source }); + status = graph.build_index(); + REQUIRE(status.success()); + + ggml::hrx::DispatchRegistryBuilder publication_registry_builder; + publication_registry_builder.add_activation_consumer({ + "test.q8.mul", + GGML_OP_MUL, + ggml::hrx::DispatchActivationInputQ8_1X4, + ggml::hrx::DispatchActivationConsumerUse::BandwidthLimited, + ggml::hrx::DispatchActivationProducerRequirementNone, + ggml::hrx::DispatchSource::Common, + publication_accepts_mul, + }); + publication_registry_builder.add_activation_consumer({ + "test.f16-row.add", + GGML_OP_ADD, + ggml::hrx::DispatchActivationInputF16Row, + ggml::hrx::DispatchActivationConsumerUse::Row, + ggml::hrx::DispatchActivationProducerRequirementNone, + ggml::hrx::DispatchSource::Common, + publication_accepts_add, + }); + const ggml::hrx::DispatchRegistry publication_registry = publication_registry_builder.build(); + const std::vector covered(graph.nodes().size(), false); + ggml::hrx::CommandPlan plan; + const ggml::hrx::DispatchMatchContext context = { + graph, &graph.nodes().front(), 0, covered, plan, ggml::hrx::ValueId(100), &publication_registry, + }; + const ggml::hrx::Value * produced = graph.values().find(root); + REQUIRE(produced != nullptr); + + const ggml::hrx::CommonActivationPublicationCapabilities direct_only = { + ggml::hrx::DispatchActivationInputQ8_1X4, + }; + const ggml::hrx::CommonActivationPublicationPlan direct_only_plan = + ggml::hrx::common_select_activation_publication_plan(context, *produced, direct_only); + REQUIRE(!direct_only_plan.matched()); + REQUIRE(direct_only_plan.fallback_reason == + ggml::hrx::CommonActivationPublicationFallbackReason::NoQualifiedConsumer); + + const ggml::hrx::CommonActivationPublicationCapabilities q8_capabilities = { + ggml::hrx::DispatchActivationInputQ8_1X4, + ggml::hrx::DispatchActivationInputQ8_1X4, + }; + const ggml::hrx::CommonActivationPublicationPlan q8_plan = + ggml::hrx::common_select_activation_publication_plan(context, *produced, q8_capabilities); + REQUIRE(q8_plan.matched()); + const ggml::hrx::CommonActivationPublicationDemand q8_demand = q8_plan.preferred(); + REQUIRE(q8_demand.format == ggml::hrx::CommonActivationPublicationFormat::Q8_1X4); + REQUIRE(q8_demand.consumer_value->id == reshaped); + REQUIRE(q8_demand.alias_depth == 1); + REQUIRE(q8_demand.mixed_fanout); + + const ggml::hrx::CommonActivationPublicationCapabilities f16_capabilities = { + ggml::hrx::DispatchActivationInputF16Row, + }; + const ggml::hrx::CommonActivationPublicationPlan f16_plan = + ggml::hrx::common_select_activation_publication_plan(context, *produced, f16_capabilities); + REQUIRE(f16_plan.matched()); + const ggml::hrx::CommonActivationPublicationDemand f16_demand = f16_plan.preferred(); + REQUIRE(f16_demand.format == ggml::hrx::CommonActivationPublicationFormat::F16Row); + REQUIRE(f16_demand.consumer_value->id == root); + REQUIRE(f16_demand.alias_depth == 0); + REQUIRE(f16_demand.mixed_fanout); + + ggml::hrx::DispatchRegistryBuilder direct_k16_registry_builder; + direct_k16_registry_builder.add_activation_consumer({ + "test.k16.mul", + GGML_OP_MUL, + ggml::hrx::DispatchActivationInputF16K16Major, + ggml::hrx::DispatchActivationConsumerUse::Tiled, + ggml::hrx::DispatchActivationProducerRequirementNone, + ggml::hrx::DispatchSource::Common, + publication_accepts_mul, + }); + const ggml::hrx::DispatchRegistry direct_k16_registry = direct_k16_registry_builder.build(); + const ggml::hrx::DispatchMatchContext direct_k16_context = { + graph, &graph.nodes().front(), 0, covered, plan, ggml::hrx::ValueId(100), &direct_k16_registry, + }; + const ggml::hrx::CommonActivationPublicationCapabilities mixed_alias_capabilities = { + ggml::hrx::DispatchActivationInputQ8_1X4 | ggml::hrx::DispatchActivationInputF16K16Major, + ggml::hrx::DispatchActivationInputQ8_1X4, + ggml::hrx::DispatchActivationProducerRequirementNone, + }; + const ggml::hrx::CommonActivationPublicationDemand direct_k16_through_q8_alias = + ggml::hrx::common_select_activation_publication_plan( + direct_k16_context, *produced, mixed_alias_capabilities).publications.front(); + REQUIRE(!direct_k16_through_q8_alias.matched()); + + ggml::hrx::DispatchRegistryBuilder shape_back_registry_builder; + shape_back_registry_builder.add_activation_consumer({ + "test.k16.shape-back", + GGML_OP_MUL, + ggml::hrx::DispatchActivationInputF16K16Major, + ggml::hrx::DispatchActivationConsumerUse::Tiled, + ggml::hrx::DispatchActivationProducerRequirementNone, + ggml::hrx::DispatchSource::Common, + publication_accepts_mul_output_13, + }); + const ggml::hrx::DispatchRegistry shape_back_registry = shape_back_registry_builder.build(); + const ggml::hrx::DispatchMatchContext shape_back_context = { + graph, &graph.nodes().front(), 0, covered, plan, ggml::hrx::ValueId(100), &shape_back_registry, + }; + REQUIRE(!ggml::hrx::common_select_activation_publication_plan( + shape_back_context, *produced, mixed_alias_capabilities).matched()); + + ggml::hrx::CommonActivationPublicationCapabilities multiple_capabilities = { + ggml::hrx::DispatchActivationInputQ8_1X4 | ggml::hrx::DispatchActivationInputF16Row, + ggml::hrx::DispatchActivationInputQ8_1X4, + ggml::hrx::DispatchActivationProducerRequirementNone, + ggml::hrx::DispatchActivationInputQ8_1X4 | ggml::hrx::DispatchActivationInputF16Row, + }; + ggml::hrx::DispatchRegistryBuilder mismatched_alias_registry_builder; + mismatched_alias_registry_builder.add_activation_consumer({ + "test.q8.mismatched-alias-source", + GGML_OP_MUL, + ggml::hrx::DispatchActivationInputQ8_1X4, + ggml::hrx::DispatchActivationConsumerUse::BandwidthLimited, + ggml::hrx::DispatchActivationProducerRequirementNone, + ggml::hrx::DispatchSource::Common, + publication_accepts_mul_output_15, + }); + const ggml::hrx::DispatchRegistry mismatched_alias_registry = mismatched_alias_registry_builder.build(); + const ggml::hrx::DispatchMatchContext mismatched_alias_context = { + graph, &graph.nodes().front(), 0, covered, plan, ggml::hrx::ValueId(100), &mismatched_alias_registry, + }; + REQUIRE(!ggml::hrx::common_select_activation_publication_plan( + mismatched_alias_context, *produced, multiple_capabilities).matched()); + + const ggml::hrx::CommonActivationPublicationPlan multiple_plan = + ggml::hrx::common_select_activation_publication_plan(context, *produced, multiple_capabilities); + REQUIRE(multiple_plan.publication_count == 2); + REQUIRE(multiple_plan.publications[0].format == ggml::hrx::CommonActivationPublicationFormat::Q8_1X4); + REQUIRE(multiple_plan.publications[1].format == ggml::hrx::CommonActivationPublicationFormat::F16Row); + ggml::hrx::DispatchRegistryBuilder q8_only_registry_builder; + q8_only_registry_builder.add_activation_consumer({ + "test.q8.only", + GGML_OP_MUL, + ggml::hrx::DispatchActivationInputQ8_1X4, + ggml::hrx::DispatchActivationConsumerUse::BandwidthLimited, + ggml::hrx::DispatchActivationProducerRequirementNone, + ggml::hrx::DispatchSource::Common, + publication_accepts_mul, + }); + const ggml::hrx::DispatchRegistry q8_only_registry = q8_only_registry_builder.build(); + const ggml::hrx::DispatchMatchContext q8_only_context = { + graph, &graph.nodes().front(), 0, covered, plan, ggml::hrx::ValueId(100), &q8_only_registry, + }; + const ggml::hrx::CommonActivationPublicationPlan companion_plan = + ggml::hrx::common_select_activation_publication_plan(q8_only_context, *produced, multiple_capabilities); + REQUIRE(companion_plan.publication_count == 2); + REQUIRE(companion_plan.publications[0].format == ggml::hrx::CommonActivationPublicationFormat::Q8_1X4); + REQUIRE(companion_plan.publications[1].format == ggml::hrx::CommonActivationPublicationFormat::F16Row); + REQUIRE(companion_plan.publications[1].consumer == companion_plan.publications[0].consumer); + multiple_capabilities.co_publication_formats = ggml::hrx::DispatchActivationInputNone; + REQUIRE(ggml::hrx::common_select_activation_publication_plan(context, *produced, multiple_capabilities) + .publication_count == 1); + + const ggml::hrx::CommonActivationPublicationCapabilities disabled_capabilities; + const ggml::hrx::CommonActivationPublicationPlan disabled_plan = + ggml::hrx::common_select_activation_publication_plan(context, *produced, disabled_capabilities); + REQUIRE(!disabled_plan.matched()); + REQUIRE(disabled_plan.fallback_reason == + ggml::hrx::CommonActivationPublicationFallbackReason::NoEnabledCandidates); + + ggml::hrx::DispatchMatch match; + match.completion_counter_requests.push_back({ ggml::hrx::ValueId(100), "existing.counter", 1 }); + const ggml::hrx::CommonActivationPublicationDemand symmetric_demand = { + ggml::hrx::CommonActivationPublicationFormat::SymmetricI4K32, + produced, + &graph.nodes()[3], + }; + ggml::hrx::CommonActivationPublication symmetric; + ggml::hrx::CommonActivationPublication q8; + ggml::hrx::CommonActivationPublication f16; + ggml::hrx::CommonActivationPublication k16; + REQUIRE(ggml::hrx::common_reserve_activation_publication(context, match, *produced, symmetric_demand, + "test.symmetric_i4", symmetric)); + REQUIRE(ggml::hrx::common_append_activation_publication(match, symmetric, "test.symmetric_i4")); + REQUIRE(ggml::hrx::common_reserve_activation_publication(context, match, *produced, q8_demand, "test.q8", q8)); + REQUIRE(ggml::hrx::common_append_activation_publication(match, q8, "test.q8")); + REQUIRE(ggml::hrx::common_reserve_activation_publication(context, match, *produced, f16_demand, "test.f16", f16)); + REQUIRE(ggml::hrx::common_append_activation_publication(match, f16, "test.f16")); + const ggml::hrx::CommonActivationPublicationDemand k16_demand = { + ggml::hrx::CommonActivationPublicationFormat::F16K16Major, + q8_demand.consumer_value, + q8_demand.consumer, + }; + REQUIRE(ggml::hrx::common_reserve_activation_publication(context, match, *produced, k16_demand, "test.k16", k16)); + REQUIRE(ggml::hrx::common_append_activation_publication(match, k16, "test.k16")); + + ggml::hrx::DispatchMatch invalid_id_match; + const ggml::hrx::DispatchMatchContext invalid_id_context = { + graph, &graph.nodes().front(), 0, covered, plan, ggml::hrx::ValueId(-1), + }; + ggml::hrx::CommonActivationPublication invalid_id_publication; + REQUIRE(!ggml::hrx::common_reserve_activation_publication(invalid_id_context, invalid_id_match, *produced, + f16_demand, "test.invalid_id", invalid_id_publication)); + REQUIRE(invalid_id_match.transients.empty()); + + ggml::hrx::DispatchMatch conflict_match; + REQUIRE(ggml::hrx::common_append_activation_publication(conflict_match, q8, "test.q8")); + ggml::hrx::CommonActivationPublication conflicting_q8 = q8; + conflicting_q8.alternate_value = ggml::hrx::ValueId(200); + REQUIRE(!ggml::hrx::common_append_activation_publication(conflict_match, conflicting_q8, "test.q8.conflict")); + REQUIRE(!conflict_match.status.success()); + + REQUIRE(match.status.success()); + REQUIRE(match.transients.size() == 4); + REQUIRE(symmetric.alternate_value == ggml::hrx::ValueId(101)); + REQUIRE(q8.alternate_value == ggml::hrx::ValueId(102)); + REQUIRE(f16.alternate_value == ggml::hrx::ValueId(103)); + REQUIRE(k16.alternate_value == ggml::hrx::ValueId(104)); + REQUIRE(symmetric.payload_bytes == static_cast(elements) / 2); + REQUIRE(symmetric.scales_offset % 256 == 0); + REQUIRE(symmetric.sums_offset % 256 == 0); + REQUIRE(q8.byte_count == static_cast(row_count) * ggml_row_size(GGML_TYPE_Q8_1, row_size)); + REQUIRE(f16.byte_count == static_cast(elements) * sizeof(ggml_fp16_t)); + REQUIRE(symmetric.binding(symmetric.scales_offset, symmetric.scales_bytes).value == symmetric.alternate_value); + + ggml::hrx::CommandPlan recorded; + recorded.metadata = std::move(match.metadata); + const std::vector & diagnostics = + recorded.metadata.activation_publication_diagnostics(); + REQUIRE(diagnostics.size() == 4); + REQUIRE(diagnostics[1].source_value == root); + REQUIRE(diagnostics[1].alternate_value == q8.alternate_value); + REQUIRE(diagnostics[1].consumer_value == reshaped); + REQUIRE(diagnostics[1].requested_format == "q8-1-x4"); + REQUIRE(diagnostics[1].publication_name == "test.q8"); + REQUIRE(diagnostics[1].publication_stage == "published"); + REQUIRE(diagnostics[1].fallback_reason == "none"); + REQUIRE(diagnostics[1].alias_depth == 1); + REQUIRE(diagnostics[1].mixed_fanout); + REQUIRE(ggml::hrx::find_alternate_value(graph, recorded, reshaped, GGML_TYPE_Q8_1, q8.byte_count) != nullptr); + REQUIRE(ggml::hrx::find_alternate_value(graph, recorded, reshaped, GGML_TYPE_F16, f16.byte_count) != nullptr); + REQUIRE(ggml::hrx::find_generated_resource(graph, recorded, reshaped, ggml::hrx::GeneratedResourceRole::F16K16Major, + k16.byte_count) != nullptr); +} + +static std::vector make_pattern_f32(size_t element_count, int seed, float scale = 0.01f) { + std::vector data(element_count); + for (size_t i = 0; i < element_count; ++i) { + const int value = static_cast((i * 17 + static_cast(seed) * 29) % 97) - 48; + data[i] = static_cast(value) * scale; + } + return data; +} + +static std::vector make_gemma_scaled_embedding_input(int64_t hidden_size, int64_t token_count) { + std::vector data = make_pattern_f32(static_cast(hidden_size * token_count), 29, 0.0005f); + const float embedding_scale = std::sqrt(static_cast(hidden_size)); + for (float & value : data) { + value *= embedding_scale; + } + return data; +} + +static std::vector make_i32_mod_data(size_t element_count, int32_t modulo) { + std::vector data(element_count); + for (size_t i = 0; i < element_count; ++i) { + data[i] = static_cast(i % static_cast(modulo)); + } + return data; +} + +static std::vector make_i64_mod_data(size_t element_count, int64_t modulo) { + std::vector data(element_count); + for (size_t i = 0; i < element_count; ++i) { + data[i] = static_cast(i % static_cast(modulo)); + } + return data; +} + +static std::vector make_pattern_f16(size_t element_count, int seed, float scale = 0.01f) { + const std::vector f32 = make_pattern_f32(element_count, seed, scale); + std::vector data(element_count); + for (size_t i = 0; i < element_count; ++i) { + data[i] = ggml_fp32_to_fp16(f32[i]); + } + return data; +} + +static std::vector make_quantized_rows(ggml_type type, int64_t row_length, int64_t row_count, int seed) { + const ggml_type_traits * traits = ggml_get_type_traits(type); + REQUIRE(traits != nullptr); + ggml_quantize_init(type); + const size_t row_size = ggml_row_size(type, row_length); + std::vector data(static_cast(row_count) * row_size); + if (type == GGML_TYPE_IQ2_XXS || type == GGML_TYPE_IQ2_XS) { + std::vector source(static_cast(row_length * row_count)); + std::vector importance(static_cast(row_length), 1.0f); + for (int64_t r = 0; r < row_count; ++r) { + for (int64_t c = 0; c < row_length; ++c) { + const int value = static_cast((r * 13 + c * 7 + seed * 31) % 101) - 50; + source[static_cast(r * row_length + c)] = static_cast(value) * 0.005f; + } + } + const size_t bytes_written = type == GGML_TYPE_IQ2_XXS ? + quantize_iq2_xxs(source.data(), data.data(), row_count, row_length, + importance.data()) : + quantize_iq2_xs(source.data(), data.data(), row_count, row_length, + importance.data()); + REQUIRE(bytes_written == data.size()); + return data; + } + REQUIRE(traits->from_float_ref != nullptr); + std::vector row(static_cast(row_length)); + for (int64_t r = 0; r < row_count; ++r) { + for (int64_t c = 0; c < row_length; ++c) { + const int value = static_cast((r * 13 + c * 7 + seed * 31) % 101) - 50; + row[static_cast(c)] = static_cast(value) * 0.005f; + } + traits->from_float_ref(row.data(), data.data() + static_cast(r) * row_size, row_length); + } + return data; +} + +static std::vector make_iq1_rows(ggml_type type, int64_t row_length, int64_t row_count, int seed) { + REQUIRE(type == GGML_TYPE_IQ1_S || type == GGML_TYPE_IQ1_M); + REQUIRE(row_length % 256 == 0); + const size_t row_size = ggml_row_size(type, row_length); + const size_t block_size = type == GGML_TYPE_IQ1_S ? 50 : 56; + REQUIRE(row_size == static_cast(row_length / 256) * block_size); + std::vector data(static_cast(row_count) * row_size); + for (int64_t row = 0; row < row_count; ++row) { + uint8_t * row_data = data.data() + static_cast(row) * row_size; + for (int64_t block = 0; block < row_length / 256; ++block) { + uint8_t * dst = row_data + static_cast(block) * block_size; + if (type == GGML_TYPE_IQ1_S) { + const ggml_fp16_t d = ggml_fp32_to_fp16(0.015625f); + std::memcpy(dst, &d, sizeof(d)); + for (int i = 0; i < 32; ++i) { + dst[2 + i] = static_cast((row * 13 + block * 29 + i * 37 + seed) & 0xff); + } + for (int group = 0; group < 8; ++group) { + uint16_t qh = 0; + for (int entry = 0; entry < 4; ++entry) { + qh |= static_cast(((row + block + group + entry + seed) & 7) << (3 * entry)); + } + qh |= static_cast(((row + block + group) & 7) << 12); + qh |= static_cast(((row + group) & 1) << 15); + std::memcpy(dst + 34 + 2 * group, &qh, sizeof(qh)); + } + } else { + for (int i = 0; i < 32; ++i) { + dst[i] = static_cast((row * 17 + block * 31 + i * 41 + seed) & 0xff); + } + for (int i = 0; i < 16; ++i) { + const uint8_t low = static_cast((row + block + i + seed) & 7); + const uint8_t high = static_cast((row + block + i + seed + 3) & 7); + dst[32 + i] = static_cast(low | ((i & 1) ? 0x08 : 0) | (high << 4) | + ((i & 2) ? 0x80 : 0)); + } + const uint16_t base_bits = ggml_fp32_to_fp16(0.015625f); + for (int word = 0; word < 4; ++word) { + const uint16_t local_scales = static_cast(1 | (2 << 3) | (3 << 6) | (4 << 9)); + const uint16_t base_nibble = static_cast((base_bits >> (4 * word)) & 0xf); + const uint16_t scales = static_cast(local_scales | (base_nibble << 12)); + std::memcpy(dst + 48 + 2 * word, &scales, sizeof(scales)); + } + } + } + } + return data; +} + +static std::vector make_matmul_weight_bytes(ggml_type type, int64_t row_length, int64_t row_count, int seed) { + if (type == GGML_TYPE_IQ1_S || type == GGML_TYPE_IQ1_M) { + return make_iq1_rows(type, row_length, row_count, seed); + } + if (type == GGML_TYPE_F32) { + (void) seed; + const std::vector weights(static_cast(row_length * row_count), 0.00390625f); + std::vector bytes(weights.size() * sizeof(float)); + std::memcpy(bytes.data(), weights.data(), bytes.size()); + return bytes; + } + + if (type == GGML_TYPE_F16) { + (void) seed; + const std::vector weights(static_cast(row_length * row_count), + ggml_fp32_to_fp16(0.00390625f)); + std::vector bytes(weights.size() * sizeof(ggml_fp16_t)); + std::memcpy(bytes.data(), weights.data(), bytes.size()); + return bytes; + } + + if (type == GGML_TYPE_BF16) { + (void) seed; + const std::vector weights(static_cast(row_length * row_count), + ggml_fp32_to_bf16(0.00390625f)); + std::vector bytes(weights.size() * sizeof(ggml_bf16_t)); + std::memcpy(bytes.data(), weights.data(), bytes.size()); + return bytes; + } + + return make_quantized_rows(type, row_length, row_count, seed); +} + +static void set_tensor_bytes(ggml_backend_t backend, ggml_tensor * tensor, const void * data, size_t byte_count) { + REQUIRE(tensor != nullptr); + REQUIRE(ggml_nbytes(tensor) == byte_count); + ggml_backend_tensor_set(tensor, data, 0, byte_count); + ggml_backend_synchronize(backend); +} + +static void set_tensor_pair_bytes(ggml_backend_t cpu_backend, + ggml_tensor * cpu_tensor, + ggml_backend_t hrx_backend, + ggml_tensor * hrx_tensor, + const void * data, + size_t byte_count) { + set_tensor_bytes(cpu_backend, cpu_tensor, data, byte_count); + set_tensor_bytes(hrx_backend, hrx_tensor, data, byte_count); +} + +static std::vector get_f32_tensor(ggml_backend_t backend, ggml_tensor * tensor) { + REQUIRE(tensor != nullptr); + const size_t element_count = static_cast(ggml_nelements(tensor)); + std::vector data(element_count); + if (tensor->type == GGML_TYPE_F32) { + ggml_backend_tensor_get(tensor, data.data(), 0, data.size() * sizeof(float)); + } else if (tensor->type == GGML_TYPE_F16) { + std::vector f16(element_count); + ggml_backend_tensor_get(tensor, f16.data(), 0, f16.size() * sizeof(ggml_fp16_t)); + for (size_t i = 0; i < element_count; ++i) { + data[i] = ggml_fp16_to_fp32(f16[i]); + } + } else { + REQUIRE(false); + } + ggml_backend_synchronize(backend); + return data; +} + +static void require_close(const std::vector & actual, + const std::vector & expected, + float abs_tolerance, + float rel_tolerance = 0.0f) { + REQUIRE(actual.size() == expected.size()); + for (size_t i = 0; i < actual.size(); ++i) { + const float diff = std::fabs(actual[i] - expected[i]); + const float allowed = abs_tolerance + rel_tolerance * std::fabs(expected[i]); + if (!std::isfinite(actual[i]) || !std::isfinite(expected[i]) || diff > allowed) { + std::fprintf(stderr, "value mismatch at %zu: actual=%g expected=%g diff=%g allowed=%g\n", i, actual[i], + expected[i], diff, allowed); + std::abort(); + } + } +} + +static ggml_backend_t init_cpu_backend() { + ggml_backend_t backend = ggml_backend_init_by_type(GGML_BACKEND_DEVICE_TYPE_CPU, nullptr); + REQUIRE(backend != nullptr); + return backend; +} + +struct AttentionPostprocessGraph { + ggml_tensor * input = nullptr; + ggml_tensor * query_weight = nullptr; + ggml_tensor * key_weight = nullptr; + ggml_tensor * value_weight = nullptr; + ggml_tensor * query_norm_weight = nullptr; + ggml_tensor * key_norm_weight = nullptr; + ggml_tensor * positions = nullptr; + ggml_tensor * inverse_frequencies = nullptr; + ggml_tensor * key_cache = nullptr; + ggml_tensor * value_cache = nullptr; + ggml_tensor * key_cache_indices = nullptr; + ggml_tensor * value_cache_indices = nullptr; + ggml_tensor * attention_mask = nullptr; + ggml_tensor * query_reshape = nullptr; + ggml_tensor * query_output = nullptr; + ggml_tensor * key_output = nullptr; + ggml_tensor * value_output = nullptr; +}; + +static AttentionPostprocessGraph build_attention_postprocess_graph(ggml_context * ctx, + int64_t token_count, + int64_t query_head_count, + int64_t key_value_head_count, + int64_t cache_row_count) { + AttentionPostprocessGraph graph; + const int64_t query_size = query_head_count * kQwenFlashHeadSize; + const int64_t key_value_size = key_value_head_count * kQwenFlashHeadSize; + + graph.input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenHiddenSize, token_count); + graph.query_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, kQwenHiddenSize, query_size); + graph.key_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, kQwenHiddenSize, key_value_size); + graph.value_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, kQwenHiddenSize, key_value_size); + REQUIRE(graph.input != nullptr); + REQUIRE(graph.query_weight != nullptr); + REQUIRE(graph.key_weight != nullptr); + REQUIRE(graph.value_weight != nullptr); + + ggml_tensor * query_raw = ggml_mul_mat(ctx, graph.query_weight, graph.input); + ggml_tensor * key_raw = ggml_mul_mat(ctx, graph.key_weight, graph.input); + ggml_tensor * value_raw = ggml_mul_mat(ctx, graph.value_weight, graph.input); + REQUIRE(query_raw != nullptr); + REQUIRE(key_raw != nullptr); + REQUIRE(value_raw != nullptr); + + ggml_tensor * query_reshape = ggml_reshape_3d(ctx, query_raw, kQwenFlashHeadSize, query_head_count, token_count); + ggml_tensor * key_reshape = ggml_reshape_3d(ctx, key_raw, kQwenFlashHeadSize, key_value_head_count, token_count); + ggml_tensor * value_reshape = + ggml_reshape_3d(ctx, value_raw, kQwenFlashHeadSize, key_value_head_count, token_count); + REQUIRE(query_reshape != nullptr); + REQUIRE(key_reshape != nullptr); + REQUIRE(value_reshape != nullptr); + graph.query_reshape = query_reshape; + + graph.query_norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, kQwenFlashHeadSize); + graph.key_norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, kQwenFlashHeadSize); + graph.positions = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, token_count); + graph.inverse_frequencies = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, kQwenFlashHeadSize / 2); + REQUIRE(graph.query_norm_weight != nullptr); + REQUIRE(graph.key_norm_weight != nullptr); + REQUIRE(graph.positions != nullptr); + REQUIRE(graph.inverse_frequencies != nullptr); + + ggml_tensor * query_norm = ggml_rms_norm(ctx, query_reshape, kQwenRmsNormEps); + ggml_tensor * query_mul = ggml_mul(ctx, query_norm, graph.query_norm_weight); + REQUIRE(query_norm != nullptr); + REQUIRE(query_mul != nullptr); + graph.query_output = ggml_rope_ext(ctx, query_mul, graph.positions, graph.inverse_frequencies, kQwenFlashHeadSize, + GGML_ROPE_TYPE_NEOX, 0, 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + REQUIRE(graph.query_output != nullptr); + + ggml_tensor * key_norm = ggml_rms_norm(ctx, key_reshape, kQwenRmsNormEps); + ggml_tensor * key_mul = ggml_mul(ctx, key_norm, graph.key_norm_weight); + REQUIRE(key_norm != nullptr); + REQUIRE(key_mul != nullptr); + ggml_tensor * key_rope = ggml_rope_ext(ctx, key_mul, graph.positions, graph.inverse_frequencies, kQwenFlashHeadSize, + GGML_ROPE_TYPE_NEOX, 0, 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + REQUIRE(key_rope != nullptr); + + ggml_tensor * key_cache_rows = ggml_reshape_2d(ctx, key_rope, key_value_size, token_count); + ggml_tensor * value_cache_rows = ggml_reshape_2d(ctx, value_reshape, key_value_size, token_count); + REQUIRE(key_cache_rows != nullptr); + REQUIRE(value_cache_rows != nullptr); + + graph.key_cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, key_value_size, cache_row_count); + graph.value_cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, key_value_size, cache_row_count); + graph.key_cache_indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, token_count); + graph.value_cache_indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, token_count); + REQUIRE(graph.key_cache != nullptr); + REQUIRE(graph.value_cache != nullptr); + REQUIRE(graph.key_cache_indices != nullptr); + REQUIRE(graph.value_cache_indices != nullptr); + + graph.key_output = ggml_set_rows(ctx, graph.key_cache, key_cache_rows, graph.key_cache_indices); + graph.value_output = ggml_set_rows(ctx, graph.value_cache, value_cache_rows, graph.value_cache_indices); + REQUIRE(graph.key_output != nullptr); + REQUIRE(graph.value_output != nullptr); + return graph; +} + +static ggml_tensor * append_qwen_full_cache_flash_attention_consumer(ggml_context * ctx, + AttentionPostprocessGraph & graph, + int64_t token_count, + int64_t query_head_count, + int64_t key_value_head_count, + int64_t cache_row_count) { + ggml_tensor * query_layout = + ggml_reshape_3d(ctx, graph.query_output, kQwenFlashHeadSize, query_head_count, token_count); + ggml_tensor * query_permute = ggml_permute(ctx, query_layout, 0, 2, 1, 3); + ggml_tensor * key_cache_layout = + ggml_reshape_3d(ctx, graph.key_cache, kQwenFlashHeadSize, key_value_head_count, cache_row_count); + ggml_tensor * key_permute = ggml_permute(ctx, key_cache_layout, 0, 2, 1, 3); + ggml_tensor * value_cache_layout = + ggml_reshape_3d(ctx, graph.value_cache, kQwenFlashHeadSize, key_value_head_count, cache_row_count); + ggml_tensor * value_permute = ggml_permute(ctx, value_cache_layout, 0, 2, 1, 3); + graph.attention_mask = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, cache_row_count, token_count); + REQUIRE(query_layout != nullptr); + REQUIRE(query_permute != nullptr); + REQUIRE(key_cache_layout != nullptr); + REQUIRE(key_permute != nullptr); + REQUIRE(value_cache_layout != nullptr); + REQUIRE(value_permute != nullptr); + REQUIRE(graph.attention_mask != nullptr); + return build_qwen_flash_attention_graph(ctx, query_permute, key_permute, value_permute, graph.attention_mask); +} + +struct QwenFlashAttentionLayoutGraph { + ggml_tensor * query = nullptr; + ggml_tensor * key = nullptr; + ggml_tensor * value = nullptr; + ggml_tensor * mask = nullptr; + ggml_tensor * output = nullptr; +}; + +static QwenFlashAttentionLayoutGraph build_qwen_flash_attention_layout_graph( + ggml_context * ctx, + int64_t query_token_count, + int64_t key_value_token_count, + int64_t qk_head_size = kQwenFlashHeadSize, + int64_t value_head_size = kQwenFlashHeadSize, + float scale = 0.0f) { + QwenFlashAttentionLayoutGraph graph; + constexpr int64_t query_head_count = 32; + constexpr int64_t key_value_head_count = 4; + + ggml_tensor * query_storage = + ggml_new_tensor_3d(ctx, GGML_TYPE_F32, qk_head_size, query_head_count, query_token_count); + ggml_tensor * key_storage = + ggml_new_tensor_3d(ctx, GGML_TYPE_F16, qk_head_size, key_value_head_count, key_value_token_count); + ggml_tensor * value_storage = + ggml_new_tensor_3d(ctx, GGML_TYPE_F16, value_head_size, key_value_head_count, key_value_token_count); + REQUIRE(query_storage != nullptr); + REQUIRE(key_storage != nullptr); + REQUIRE(value_storage != nullptr); + + graph.query = ggml_permute(ctx, query_storage, 0, 2, 1, 3); + graph.key = ggml_permute(ctx, key_storage, 0, 2, 1, 3); + graph.value = ggml_permute(ctx, value_storage, 0, 2, 1, 3); + graph.mask = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, key_value_token_count, query_token_count); + REQUIRE(graph.query != nullptr); + REQUIRE(graph.key != nullptr); + REQUIRE(graph.value != nullptr); + REQUIRE(graph.mask != nullptr); + + graph.output = + build_qwen_flash_attention_graph(ctx, graph.query, graph.key, graph.value, graph.mask, qk_head_size, scale); + return graph; +} + +static AttentionPostprocessGraph build_decode_attention_qkv_graph(ggml_context * ctx) { + constexpr int64_t token_count = 1; + constexpr int64_t query_head_count = 32; + constexpr int64_t key_value_head_count = 4; + constexpr int64_t cache_row_count = 1024; + AttentionPostprocessGraph graph = + build_attention_postprocess_graph(ctx, token_count, query_head_count, key_value_head_count, cache_row_count); + + ggml_tensor * hidden_state = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenHiddenSize, token_count); + ggml_tensor * attention_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, kQwenHiddenSize); + REQUIRE(hidden_state != nullptr); + REQUIRE(attention_weight != nullptr); + ggml_tensor * attention_rms = ggml_rms_norm(ctx, hidden_state, kQwenRmsNormEps); + REQUIRE(attention_rms != nullptr); + graph.input = ggml_mul(ctx, attention_rms, attention_weight); + REQUIRE(graph.input != nullptr); + + const int64_t key_value_size = key_value_head_count * kQwenFlashHeadSize; + ggml_tensor * query_raw = ggml_mul_mat(ctx, graph.query_weight, graph.input); + ggml_tensor * key_raw = ggml_mul_mat(ctx, graph.key_weight, graph.input); + ggml_tensor * value_raw = ggml_mul_mat(ctx, graph.value_weight, graph.input); + REQUIRE(query_raw != nullptr); + REQUIRE(key_raw != nullptr); + REQUIRE(value_raw != nullptr); + + ggml_tensor * query_reshape = ggml_reshape_3d(ctx, query_raw, kQwenFlashHeadSize, query_head_count, token_count); + ggml_tensor * key_reshape = ggml_reshape_3d(ctx, key_raw, kQwenFlashHeadSize, key_value_head_count, token_count); + ggml_tensor * value_reshape = + ggml_reshape_3d(ctx, value_raw, kQwenFlashHeadSize, key_value_head_count, token_count); + REQUIRE(query_reshape != nullptr); + REQUIRE(key_reshape != nullptr); + REQUIRE(value_reshape != nullptr); + + ggml_tensor * query_norm = ggml_rms_norm(ctx, query_reshape, kQwenRmsNormEps); + ggml_tensor * query_mul = ggml_mul(ctx, query_norm, graph.query_norm_weight); + REQUIRE(query_norm != nullptr); + REQUIRE(query_mul != nullptr); + graph.query_output = ggml_rope_ext(ctx, query_mul, graph.positions, graph.inverse_frequencies, kQwenFlashHeadSize, + GGML_ROPE_TYPE_NEOX, 0, 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + REQUIRE(graph.query_output != nullptr); + + ggml_tensor * key_norm = ggml_rms_norm(ctx, key_reshape, kQwenRmsNormEps); + ggml_tensor * key_mul = ggml_mul(ctx, key_norm, graph.key_norm_weight); + REQUIRE(key_norm != nullptr); + REQUIRE(key_mul != nullptr); + ggml_tensor * key_rope = ggml_rope_ext(ctx, key_mul, graph.positions, graph.inverse_frequencies, kQwenFlashHeadSize, + GGML_ROPE_TYPE_NEOX, 0, 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + REQUIRE(key_rope != nullptr); + + ggml_tensor * key_cache_rows = ggml_reshape_2d(ctx, key_rope, key_value_size, token_count); + ggml_tensor * value_cache_rows = ggml_reshape_2d(ctx, value_reshape, key_value_size, token_count); + REQUIRE(key_cache_rows != nullptr); + REQUIRE(value_cache_rows != nullptr); + graph.key_output = ggml_set_rows(ctx, graph.key_cache, key_cache_rows, graph.key_cache_indices); + graph.value_output = ggml_set_rows(ctx, graph.value_cache, value_cache_rows, graph.value_cache_indices); + REQUIRE(graph.key_output != nullptr); + REQUIRE(graph.value_output != nullptr); + return graph; +} + +struct RoutedMoeGraph { + ggml_tensor * logits = nullptr; + ggml_tensor * input = nullptr; + ggml_tensor * gate_weight = nullptr; + ggml_tensor * up_weight = nullptr; + ggml_tensor * down_weight = nullptr; + ggml_tensor * hidden_state = nullptr; + ggml_tensor * norm_weight = nullptr; + ggml_tensor * output = nullptr; +}; + +static RoutedMoeGraph build_routed_moe_graph(ggml_context * ctx, + ggml_type down_weight_type, + bool include_next_rmsnorm) { + RoutedMoeGraph graph; + constexpr int64_t token_count = 1; + graph.logits = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenRouterExpertCount, token_count); + REQUIRE(graph.logits != nullptr); + ggml_tensor * route_ids = nullptr; + ggml_tensor * route_weights = build_qwen_router_top8_graph(ctx, graph.logits, &route_ids); + REQUIRE(route_weights != nullptr); + REQUIRE(route_ids != nullptr); + + graph.input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, kQwenHiddenSize, 1, token_count); + graph.gate_weight = + ggml_new_tensor_3d(ctx, GGML_TYPE_Q4_K, kQwenHiddenSize, kQwenMoeIntermediate, kQwenRouterExpertCount); + graph.up_weight = + ggml_new_tensor_3d(ctx, GGML_TYPE_Q4_K, kQwenHiddenSize, kQwenMoeIntermediate, kQwenRouterExpertCount); + graph.down_weight = + ggml_new_tensor_3d(ctx, down_weight_type, kQwenMoeIntermediate, kQwenHiddenSize, kQwenRouterExpertCount); + graph.hidden_state = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenHiddenSize, token_count); + REQUIRE(graph.input != nullptr); + REQUIRE(graph.gate_weight != nullptr); + REQUIRE(graph.up_weight != nullptr); + REQUIRE(graph.down_weight != nullptr); + REQUIRE(graph.hidden_state != nullptr); + + ggml_tensor * gate = ggml_mul_mat_id(ctx, graph.gate_weight, graph.input, route_ids); + ggml_tensor * up = ggml_mul_mat_id(ctx, graph.up_weight, graph.input, route_ids); + REQUIRE(gate != nullptr); + REQUIRE(up != nullptr); + ggml_tensor * glu = ggml_glu_split(ctx, gate, up, GGML_GLU_OP_SWIGLU); + REQUIRE(glu != nullptr); + ggml_tensor * down = ggml_mul_mat_id(ctx, graph.down_weight, glu, route_ids); + REQUIRE(down != nullptr); + ggml_tensor * weighted = ggml_mul(ctx, down, route_weights); + REQUIRE(weighted != nullptr); + + std::vector route_views; + route_views.reserve(kQwenRouterRouteCount); + for (int64_t route = 0; route < kQwenRouterRouteCount; ++route) { + ggml_tensor * view = ggml_view_2d(ctx, weighted, kQwenHiddenSize, token_count, weighted->nb[2], + static_cast(route) * weighted->nb[1]); + REQUIRE(view != nullptr); + route_views.push_back(view); + } + + ggml_tensor * reduced = route_views.front(); + for (size_t i = 1; i < route_views.size(); ++i) { + reduced = ggml_add(ctx, reduced, route_views[i]); + REQUIRE(reduced != nullptr); + } + + ggml_tensor * residual = ggml_add(ctx, graph.hidden_state, reduced); + REQUIRE(residual != nullptr); + graph.output = residual; + + if (include_next_rmsnorm) { + graph.norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, kQwenHiddenSize); + REQUIRE(graph.norm_weight != nullptr); + ggml_tensor * rms = ggml_rms_norm(ctx, residual, kQwenRmsNormEps); + REQUIRE(rms != nullptr); + graph.output = ggml_mul(ctx, rms, graph.norm_weight); + REQUIRE(graph.output != nullptr); + } + + return graph; +} + +static void run_rmsnorm_support_checks() { + ggml_backend_hrx_device_context device_context = {}; + ggml_backend_hrx_context backend_context = {}; + device_context.architecture = "gfx1151"; + backend_context.device = &device_context; + const ggml::hrx::GraphExecutor executor(backend_context); + + ggml_init_params params = {}; + params.mem_size = 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 256, 1); + ggml_tensor * weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 256); + REQUIRE(input != nullptr); + REQUIRE(weight != nullptr); + ggml_tensor * output = build_rmsnorm_mul_graph(ctx, input, weight); + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + const ggml::hrx::GraphSupportResult support = executor.can_execute(*graph); + REQUIRE(support.supported); + REQUIRE(support.status.success()); + + ggml_tensor * wrong_eps_output = build_rmsnorm_mul_graph(ctx, input, weight, 1.0e-5f); + ggml_cgraph * wrong_eps_graph = ggml_new_graph(ctx); + REQUIRE(wrong_eps_graph != nullptr); + ggml_build_forward_expand(wrong_eps_graph, wrong_eps_output); + const ggml::hrx::GraphSupportResult wrong_eps_support = executor.can_execute(*wrong_eps_graph); + REQUIRE(wrong_eps_support.supported); + REQUIRE(wrong_eps_support.status.success()); + require_kernel_subsequence(scheduled_kernel_sequence(wrong_eps_graph), { "loom_libs:ggml_rmsnorm_binary_f32" }); + + ggml_tensor * standalone_output = ggml_rms_norm(ctx, input, 1.0e-5f); + REQUIRE(standalone_output != nullptr); + ggml_cgraph * standalone_graph = ggml_new_graph(ctx); + REQUIRE(standalone_graph != nullptr); + ggml_build_forward_expand(standalone_graph, standalone_output); + const ggml::hrx::GraphSupportResult standalone_support = executor.can_execute(*standalone_graph); + REQUIRE(standalone_support.supported); + REQUIRE(standalone_support.status.success()); + require_kernel_subsequence(scheduled_kernel_sequence(standalone_graph), { "loom_libs:ggml_rmsnorm_f32" }); + + ggml_tensor * q8_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenHiddenSize, 1); + ggml_tensor * q8_rhs = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, kQwenHiddenSize); + ggml_tensor * q8_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, kQwenHiddenSize, kQwenVocabularyCount); + REQUIRE(q8_input != nullptr); + REQUIRE(q8_rhs != nullptr); + REQUIRE(q8_weight != nullptr); + ggml_tensor * q8_rms = ggml_rms_norm(ctx, q8_input, 1.0e-5f); + ggml_tensor * q8_binary = ggml_add(ctx, q8_rms, q8_rhs); + ggml_tensor * q8_output = ggml_mul_mat(ctx, q8_weight, q8_binary); + REQUIRE(q8_output != nullptr); + ggml_cgraph * q8_graph = ggml_new_graph(ctx); + REQUIRE(q8_graph != nullptr); + ggml_build_forward_expand(q8_graph, q8_output); + require_kernel_subsequence(scheduled_kernel_sequence(q8_graph), + { "loom_libs:ggml_rmsnorm_binary_f32", "loom_libs:ggml_mul_mat_vector_q6_f32_f32" }); + + ggml_tensor * wrong_type_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 256, 1); + ggml_tensor * wrong_type_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F16, 256); + REQUIRE(wrong_type_input != nullptr); + REQUIRE(wrong_type_weight != nullptr); + ggml_tensor * wrong_type_output = build_rmsnorm_mul_graph(ctx, wrong_type_input, wrong_type_weight); + ggml_cgraph * wrong_type_graph = ggml_new_graph(ctx); + REQUIRE(wrong_type_graph != nullptr); + ggml_build_forward_expand(wrong_type_graph, wrong_type_output); + const ggml::hrx::GraphSupportResult wrong_type_support = executor.can_execute(*wrong_type_graph); + REQUIRE(!wrong_type_support.supported); + + ggml_tensor * wrong_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 128); + REQUIRE(wrong_weight != nullptr); + ggml_tensor * wrong_weight_output = build_rmsnorm_mul_graph(ctx, input, wrong_weight); + ggml_cgraph * wrong_weight_graph = ggml_new_graph(ctx); + REQUIRE(wrong_weight_graph != nullptr); + ggml_build_forward_expand(wrong_weight_graph, wrong_weight_output); + const ggml::hrx::GraphSupportResult wrong_weight_support = executor.can_execute(*wrong_weight_graph); + REQUIRE(!wrong_weight_support.supported); + + ggml_free(ctx); +} + +static void run_rmsnorm_mul_case(int64_t hidden_size, int64_t token_count) { + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(hidden_size * token_count * sizeof(float) * 8 + 1024 * 1024); + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden_size, token_count); + ggml_tensor * weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, hidden_size); + REQUIRE(input != nullptr); + REQUIRE(weight != nullptr); + ggml_tensor * output = build_rmsnorm_mul_graph(ctx, input, weight); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + require_kernel_subsequence(scheduled_kernel_sequence(graph), { "loom_libs:ggml_rmsnorm_binary_f32" }); + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + + const std::vector input_data = make_input(hidden_size, token_count); + const std::vector weight_data = make_weight(hidden_size); + const std::vector expected = rmsnorm_mul_reference(input_data, weight_data, hidden_size, token_count); + + ggml_backend_tensor_set(input, input_data.data(), 0, input_data.size() * sizeof(float)); + ggml_backend_tensor_set(weight, weight_data.data(), 0, weight_data.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + + std::vector actual(expected.size()); + ggml_backend_tensor_get(output, actual.data(), 0, actual.size() * sizeof(float)); + for (size_t i = 0; i < actual.size(); ++i) { + const float diff = std::fabs(actual[i] - expected[i]); + REQUIRE(diff <= 5.0e-4f); + } + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(backend); +} + +static void run_rmsnorm_binary_add_cpu_reference_case() { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + constexpr int64_t hidden_size = 256; + constexpr int64_t token_count = 4; + constexpr float epsilon = 1.0e-5f; + + ggml_init_params params = {}; + params.mem_size = 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_input = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, hidden_size, token_count); + ggml_tensor * cpu_rhs = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, hidden_size); + ggml_tensor * hrx_input = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, hidden_size, token_count); + ggml_tensor * hrx_rhs = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, hidden_size); + REQUIRE(cpu_input != nullptr); + REQUIRE(cpu_rhs != nullptr); + REQUIRE(hrx_input != nullptr); + REQUIRE(hrx_rhs != nullptr); + + ggml_tensor * cpu_rms = ggml_rms_norm(cpu_ctx, cpu_input, epsilon); + ggml_tensor * hrx_rms = ggml_rms_norm(hrx_ctx, hrx_input, epsilon); + ggml_tensor * cpu_output = ggml_add(cpu_ctx, cpu_rms, cpu_rhs); + ggml_tensor * hrx_output = ggml_add(hrx_ctx, hrx_rms, hrx_rhs); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_rmsnorm_binary_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector input = make_pattern_f32(hidden_size * token_count, 17, 0.02f); + const std::vector rhs = make_pattern_f32(hidden_size, 18, 0.03f); + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), input.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_rhs, hrx_backend, hrx_rhs, rhs.data(), rhs.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 5.0e-4f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_gemma_post_proj_cpu_reference_case(int64_t token_count) { + constexpr int64_t hidden_size = 640; + constexpr float epsilon = 1.0e-6f; + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + struct GraphTensors { + ggml_tensor * input; + ggml_tensor * weight; + ggml_tensor * residual; + ggml_tensor * raw; + ggml_tensor * raw_side; + ggml_tensor * output; + }; + + auto build = [&](ggml_context * ctx) { + GraphTensors g = {}; + g.input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden_size, token_count); + g.weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, hidden_size); + g.residual = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden_size, token_count); + g.raw = ggml_add(ctx, g.input, g.residual); + ggml_tensor * normalized = ggml_rms_norm(ctx, g.raw, epsilon); + ggml_tensor * scaled = ggml_mul(ctx, normalized, g.weight); + g.output = ggml_add(ctx, scaled, g.residual); + g.raw_side = ggml_mul(ctx, g.raw, g.residual); + REQUIRE(g.output != nullptr); + REQUIRE(g.raw_side != nullptr); + return g; + }; + GraphTensors cpu = build(cpu_ctx); + GraphTensors hrx = build(hrx_ctx); + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu.output); + ggml_build_forward_expand(cpu_graph, cpu.raw_side); + ggml_build_forward_expand(hrx_graph, hrx.output); + ggml_build_forward_expand(hrx_graph, hrx.raw_side); + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_rmsnorm_mul_add_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + const std::vector input = make_pattern_f32(hidden_size * token_count, 37, 0.02f); + const std::vector residual = make_pattern_f32(hidden_size * token_count, 43, 0.04f); + const std::vector weight = make_weight(hidden_size); + set_tensor_pair_bytes(cpu_backend, cpu.input, hrx_backend, hrx.input, input.data(), input.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu.residual, hrx_backend, hrx.residual, residual.data(), + residual.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu.weight, hrx_backend, hrx.weight, weight.data(), + weight.size() * sizeof(float)); + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx.output), get_f32_tensor(cpu_backend, cpu.output), 5.0e-4f); + require_close(get_f32_tensor(hrx_backend, hrx.raw_side), get_f32_tensor(cpu_backend, cpu.raw_side), 1.0e-5f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_gemma_post_proj_matcher_negatives() { + constexpr int64_t hidden_size = 640; + ggml_init_params params = {}; + params.mem_size = 1024 * 1024; + params.no_alloc = true; + const ggml::hrx::DispatchRegistry * registry = ggml::hrx::find_dispatch_registry({ "gfx1151" }); + REQUIRE(registry != nullptr); + + for (const int variant : { 0, 1, 2, 3, 4 }) { + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden_size, 3); + ggml_tensor * weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, hidden_size); + ggml_tensor * rhs = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden_size, 3); + ggml_tensor * raw = ggml_add(ctx, input, rhs); + ggml_tensor * residual = variant == 2 ? ggml_view_2d(ctx, raw, hidden_size, 3, raw->nb[1], 0) : + variant == 3 ? ggml_mul(ctx, raw, rhs) : + rhs; + ggml_tensor * rms = ggml_rms_norm(ctx, raw, 1.0e-6f); + ggml_tensor * scaled = variant == 4 ? ggml_add(ctx, rms, weight) : ggml_mul(ctx, rms, weight); + ggml_tensor * output = ggml_add(ctx, scaled, residual); + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + if (variant == 1) { + ggml_build_forward_expand(graph, ggml_mul(ctx, scaled, rhs)); + } + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + const size_t rms_index = producer_index_for_tensor(imported.graph, rms); + const size_t raw_index = producer_index_for_tensor(imported.graph, raw); + ggml::hrx::CommandPlan plan; + std::vector covered(imported.graph.nodes().size(), false); + auto kernel = [&]() { + ggml::hrx::DispatchMatch match; + const ggml::hrx::DispatchMatchContext context = { + imported.graph, &imported.graph.nodes()[rms_index], rms_index, covered, + plan, next_plan_value(imported.graph, plan), + }; + REQUIRE(registry->match(context, match)); + REQUIRE(match.dispatches.size() == 1); + return kernel_name_for_id(match.dispatches[0].kernel.kernel_id); + }; + REQUIRE(kernel() == "loom_libs:ggml_rmsnorm_binary_f32"); + covered[raw_index] = true; + if (variant == 3) { + REQUIRE(kernel() == "loom_libs:ggml_rmsnorm_binary_f32"); + covered[producer_index_for_tensor(imported.graph, residual)] = true; + } + if (variant == 2) { + const size_t view_index = producer_index_for_tensor(imported.graph, residual); + covered[view_index] = true; + const ggml::hrx::Value * raw_value = imported.graph.values().find_tensor(raw); + const ggml::hrx::Value * residual_value = imported.graph.values().find_tensor(residual); + REQUIRE(raw_value != nullptr); + REQUIRE(residual_value != nullptr); + REQUIRE(raw_value->storage_root == residual_value->storage_root); + } + REQUIRE(kernel() == (variant == 0 || variant == 3 ? "loom_libs:ggml_rmsnorm_mul_add_f32" : + "loom_libs:ggml_rmsnorm_binary_f32")); + ggml_free(ctx); + } +} + +static void run_rmsnorm_gate_cpu_reference_case(int64_t hidden_size, + int64_t head_count, + int64_t token_count, + ggml_unary_op gate_op) { + ggml_backend_t cpu = init_cpu_backend(); + ggml_backend_t hrx = ggml_backend_hrx_init(0); + REQUIRE(hrx != nullptr); + const int64_t elements = hidden_size * head_count * token_count; + ggml_init_params params = {}; + params.mem_size = static_cast(elements * sizeof(float) * 8 + 1024 * 1024); + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr && hrx_ctx != nullptr); + + struct Graph { + ggml_tensor * input; + ggml_tensor * weight; + ggml_tensor * gate; + ggml_tensor * output; + ggml_cgraph * graph; + }; + + auto build = [&](ggml_context * ctx) { + Graph g; + g.input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, hidden_size, head_count, token_count); + g.weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, hidden_size); + g.gate = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hidden_size * head_count, token_count); + ggml_tensor * normalized = build_rmsnorm_mul_graph(ctx, g.input, g.weight); + ggml_tensor * reshaped_gate = ggml_reshape_3d(ctx, g.gate, hidden_size, head_count, token_count); + g.output = ggml_mul(ctx, normalized, ggml_unary(ctx, reshaped_gate, gate_op)); + g.graph = ggml_new_graph(ctx); + ggml_build_forward_expand(g.graph, g.output); + return g; + }; + Graph a = build(cpu_ctx); + Graph b = build(hrx_ctx); + require_kernel_subsequence(scheduled_kernel_sequence(b.graph), { "loom_libs:ggml_rmsnorm_gate_f32_publish" }); + ggml_backend_buffer_t a_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu); + ggml_backend_buffer_t b_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx); + REQUIRE(a_buffer != nullptr && b_buffer != nullptr); + const auto input = make_pattern_f32(elements, 23, 0.05f); + const auto weight = make_pattern_f32(hidden_size, 11, 0.075f); + const auto gate = make_pattern_f32(elements, 29, 0.125f); + set_tensor_pair_bytes(cpu, a.input, hrx, b.input, input.data(), input.size() * sizeof(float)); + set_tensor_pair_bytes(cpu, a.weight, hrx, b.weight, weight.data(), weight.size() * sizeof(float)); + set_tensor_pair_bytes(cpu, a.gate, hrx, b.gate, gate.data(), gate.size() * sizeof(float)); + REQUIRE(ggml_backend_graph_compute(cpu, a.graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx, b.graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu); + ggml_backend_synchronize(hrx); + const auto actual = get_f32_tensor(hrx, b.output); + const auto cpu_output = get_f32_tensor(cpu, a.output); + if (gate_op == GGML_UNARY_OP_GELU) { + // The CPU GELU route uses an F16 lookup table. Keep the tighter HRX + // check against the defining function instead of inheriting its rounding. + auto expected = rmsnorm_mul_reference(input, weight, hidden_size, head_count * token_count); + for (size_t i = 0; i < expected.size(); ++i) { + const double x = gate[i]; + const double gelu = 0.5 * x * (1.0 + std::tanh(0.7978845608028654 * x * (1.0 + 0.044715 * x * x))); + expected[i] *= static_cast(gelu); + } + require_close(actual, expected, 2.0e-5f); + require_close(cpu_output, expected, 2.0e-5f, 5.0e-4f); + } else { + require_close(actual, cpu_output, 2.0e-5f); + } + + ggml_backend_buffer_free(a_buffer); + ggml_backend_buffer_free(b_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu); + ggml_backend_free(hrx); +} + +static void run_rmsnorm_cpu_reference_case(int64_t hidden_size = 256, + int64_t token_count = 4, + float epsilon = 1.0e-5f) { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(hidden_size * token_count * sizeof(float) * 8 + 1024 * 1024); + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_input = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, hidden_size, token_count); + ggml_tensor * hrx_input = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, hidden_size, token_count); + REQUIRE(cpu_input != nullptr); + REQUIRE(hrx_input != nullptr); + + ggml_tensor * cpu_output = ggml_rms_norm(cpu_ctx, cpu_input, epsilon); + ggml_tensor * hrx_output = ggml_rms_norm(hrx_ctx, hrx_input, epsilon); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_rmsnorm_f32" }); + REQUIRE(scheduled_kernel_specialization(hrx_graph, "loom_libs:ggml_rmsnorm_f32").workload_specialization == + ggml::hrx::WorkloadSpecialization::Dynamic); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector input = make_pattern_f32(hidden_size * token_count, 19, 0.02f); + const std::vector expected = rmsnorm_reference(input, hidden_size, token_count, epsilon); + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), input.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(cpu_backend, cpu_output), expected, 5.0e-4f); + require_close(get_f32_tensor(hrx_backend, hrx_output), expected, 5.0e-4f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_strided_rmsnorm_cpu_reference_case(int64_t token_count) { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + constexpr int64_t hidden_size = 256; + constexpr int64_t input_stride = 512; + constexpr float epsilon = 1.0e-5f; + constexpr size_t input_offset = 4 * sizeof(float); + const size_t storage_elements = + input_offset / sizeof(float) + static_cast((token_count - 1) * input_stride + hidden_size); + + ggml_init_params params = {}; + params.mem_size = 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_storage = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, storage_elements); + ggml_tensor * cpu_input = + ggml_view_2d(cpu_ctx, cpu_storage, hidden_size, token_count, input_stride * sizeof(float), input_offset); + ggml_tensor * cpu_output = ggml_rms_norm(cpu_ctx, cpu_input, epsilon); + ggml_tensor * hrx_storage = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, storage_elements); + ggml_tensor * hrx_input = + ggml_view_2d(hrx_ctx, hrx_storage, hidden_size, token_count, input_stride * sizeof(float), input_offset); + ggml_tensor * hrx_output = ggml_rms_norm(hrx_ctx, hrx_input, epsilon); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_rmsnorm_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector rows = make_pattern_f32(static_cast(hidden_size * token_count), 53, 0.015f); + std::vector storage(storage_elements, -19.0f); + for (int64_t token = 0; token < token_count; ++token) { + std::memcpy(storage.data() + input_offset / sizeof(float) + static_cast(token * input_stride), + rows.data() + static_cast(token * hidden_size), + static_cast(hidden_size) * sizeof(float)); + } + const std::vector expected = rmsnorm_reference(rows, hidden_size, token_count, epsilon); + set_tensor_pair_bytes(cpu_backend, cpu_storage, hrx_backend, hrx_storage, storage.data(), + storage.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(cpu_backend, cpu_output), expected, 5.0e-4f); + require_close(get_f32_tensor(hrx_backend, hrx_output), expected, 5.0e-4f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_strided_rmsnorm_binary_cpu_reference_case() { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + constexpr int64_t hidden_size = 256; + constexpr int64_t token_count = 64; + constexpr int64_t input_stride = 288; + constexpr float epsilon = 1.0e-5f; + const size_t storage_elements = static_cast((token_count - 1) * input_stride + hidden_size); + + ggml_init_params params = {}; + params.mem_size = 2 * 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_storage = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, storage_elements); + ggml_tensor * cpu_input = + ggml_view_2d(cpu_ctx, cpu_storage, hidden_size, token_count, input_stride * sizeof(float), 0); + ggml_tensor * cpu_scale = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, hidden_size); + ggml_tensor * cpu_output = ggml_mul(cpu_ctx, ggml_rms_norm(cpu_ctx, cpu_input, epsilon), cpu_scale); + ggml_tensor * hrx_storage = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, storage_elements); + ggml_tensor * hrx_input = + ggml_view_2d(hrx_ctx, hrx_storage, hidden_size, token_count, input_stride * sizeof(float), 0); + ggml_tensor * hrx_scale = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, hidden_size); + ggml_tensor * hrx_output = ggml_mul(hrx_ctx, ggml_rms_norm(hrx_ctx, hrx_input, epsilon), hrx_scale); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_rmsnorm_binary_strided_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector rows = make_pattern_f32(hidden_size * token_count, 54, 0.015f); + const std::vector scale = make_pattern_f32(hidden_size, 55, 0.03f); + std::vector storage(storage_elements, -19.0f); + for (int64_t token = 0; token < token_count; ++token) { + std::memcpy(storage.data() + static_cast(token * input_stride), + rows.data() + static_cast(token * hidden_size), + static_cast(hidden_size) * sizeof(float)); + } + set_tensor_pair_bytes(cpu_backend, cpu_storage, hrx_backend, hrx_storage, storage.data(), + storage.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_scale, hrx_backend, hrx_scale, scale.data(), scale.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 5.0e-4f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_strided_cont_cpu_reference_case(int64_t token_count) { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + constexpr int64_t width = 64; + constexpr int64_t row_count = 40; + constexpr int64_t token_stride = 4096; + constexpr size_t view_offset = 8 * sizeof(float); + const size_t storage_elements = + view_offset / sizeof(float) + static_cast((token_count - 1) * token_stride + width * row_count); + + ggml_init_params params = {}; + params.mem_size = 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_storage = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, storage_elements); + ggml_tensor * cpu_view = ggml_view_3d(cpu_ctx, cpu_storage, width, row_count, token_count, width * sizeof(float), + token_stride * sizeof(float), view_offset); + ggml_tensor * cpu_output = ggml_cont(cpu_ctx, cpu_view); + ggml_tensor * hrx_storage = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, storage_elements); + ggml_tensor * hrx_view = ggml_view_3d(hrx_ctx, hrx_storage, width, row_count, token_count, width * sizeof(float), + token_stride * sizeof(float), view_offset); + ggml_tensor * hrx_output = ggml_cont(hrx_ctx, hrx_view); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_copy_strided_source_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector rows = make_pattern_f32(static_cast(width * row_count * token_count), 59, 0.01f); + std::vector storage(storage_elements, -37.0f); + for (int64_t token = 0; token < token_count; ++token) { + std::memcpy(storage.data() + view_offset / sizeof(float) + static_cast(token * token_stride), + rows.data() + static_cast(token * width * row_count), + static_cast(width * row_count) * sizeof(float)); + } + set_tensor_pair_bytes(cpu_backend, cpu_storage, hrx_backend, hrx_storage, storage.data(), + storage.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(cpu_backend, cpu_output), rows, 0.0f); + require_close(get_f32_tensor(hrx_backend, hrx_output), rows, 0.0f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_rmsnorm_mul_cpu_reference_case(int64_t hidden_size, int64_t token_count, float epsilon) { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(hidden_size * token_count * sizeof(float) * 8 + 1024 * 1024); + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_input = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, hidden_size, token_count); + ggml_tensor * cpu_weight = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, hidden_size); + ggml_tensor * hrx_input = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, hidden_size, token_count); + ggml_tensor * hrx_weight = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, hidden_size); + REQUIRE(cpu_input != nullptr); + REQUIRE(cpu_weight != nullptr); + REQUIRE(hrx_input != nullptr); + REQUIRE(hrx_weight != nullptr); + + ggml_tensor * cpu_output = build_rmsnorm_mul_graph(cpu_ctx, cpu_input, cpu_weight, epsilon); + ggml_tensor * hrx_output = build_rmsnorm_mul_graph(hrx_ctx, hrx_input, hrx_weight, epsilon); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_rmsnorm_binary_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector input = make_pattern_f32(hidden_size * token_count, 23, 0.02f); + const std::vector weight = make_weight(hidden_size); + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), input.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_weight, hrx_backend, hrx_weight, weight.data(), + weight.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 5.0e-4f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_gemma_scaled_rmsnorm_cpu_reference_case() { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(kGemmaHiddenSize * kGemmaPromptTokenCount * sizeof(float) * 8 + 1024 * 1024); + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_input = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, kGemmaHiddenSize, kGemmaPromptTokenCount); + ggml_tensor * hrx_input = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, kGemmaHiddenSize, kGemmaPromptTokenCount); + REQUIRE(cpu_input != nullptr); + REQUIRE(hrx_input != nullptr); + + ggml_tensor * cpu_output = ggml_rms_norm(cpu_ctx, cpu_input, kGemmaRmsNormEps); + ggml_tensor * hrx_output = ggml_rms_norm(hrx_ctx, hrx_input, kGemmaRmsNormEps); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_rmsnorm_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector input = make_gemma_scaled_embedding_input(kGemmaHiddenSize, kGemmaPromptTokenCount); + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), input.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 5.0e-4f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_gemma_scaled_rmsnorm_mul_cpu_reference_case() { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(kGemmaHiddenSize * kGemmaPromptTokenCount * sizeof(float) * 8 + 1024 * 1024); + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_input = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, kGemmaHiddenSize, kGemmaPromptTokenCount); + ggml_tensor * cpu_weight = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, kGemmaHiddenSize); + ggml_tensor * hrx_input = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, kGemmaHiddenSize, kGemmaPromptTokenCount); + ggml_tensor * hrx_weight = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, kGemmaHiddenSize); + REQUIRE(cpu_input != nullptr); + REQUIRE(cpu_weight != nullptr); + REQUIRE(hrx_input != nullptr); + REQUIRE(hrx_weight != nullptr); + + ggml_tensor * cpu_output = build_rmsnorm_mul_graph(cpu_ctx, cpu_input, cpu_weight, kGemmaRmsNormEps); + ggml_tensor * hrx_output = build_rmsnorm_mul_graph(hrx_ctx, hrx_input, hrx_weight, kGemmaRmsNormEps); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_rmsnorm_binary_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector input = make_gemma_scaled_embedding_input(kGemmaHiddenSize, kGemmaPromptTokenCount); + const std::vector weight = make_weight(kGemmaHiddenSize); + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), input.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_weight, hrx_backend, hrx_weight, weight.data(), + weight.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 5.0e-4f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_scheduled_hrx_scale_rmsnorm_case() { + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + ggml_backend_t cpu_backend = init_cpu_backend(); + REQUIRE(hrx_backend != nullptr); + + ggml_backend_t backends[] = { hrx_backend, cpu_backend }; + ggml_backend_sched_t sched = ggml_backend_sched_new(backends, nullptr, 2, GGML_DEFAULT_GRAPH_SIZE, false, true); + REQUIRE(sched != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(kGemmaHiddenSize * kGemmaPromptTokenCount * sizeof(float) * 8 + 1024 * 1024); + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kGemmaHiddenSize, kGemmaPromptTokenCount); + REQUIRE(input != nullptr); + ggml_backend_sched_set_tensor_backend(sched, input, hrx_backend); + + ggml_tensor * scaled = ggml_scale(ctx, input, std::sqrt(static_cast(kGemmaHiddenSize))); + ggml_tensor * output = ggml_rms_norm(ctx, scaled, kGemmaRmsNormEps); + REQUIRE(scaled != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + const ggml::hrx::KernelSpecialization scale_specialization = + scheduled_kernel_specialization(graph, "loom_libs:ggml_scale_f32"); + REQUIRE(scale_specialization.workload_specialization == ggml::hrx::WorkloadSpecialization::Dynamic); + + REQUIRE(ggml_backend_sched_alloc_graph(sched, graph)); + REQUIRE(ggml_backend_sched_get_tensor_backend(sched, scaled) == hrx_backend); + REQUIRE(ggml_backend_sched_get_tensor_backend(sched, output) == hrx_backend); + + const std::vector input_data = + make_pattern_f32(static_cast(kGemmaHiddenSize * kGemmaPromptTokenCount), 31, 0.0005f); + std::vector scaled_input = input_data; + const float embedding_scale = std::sqrt(static_cast(kGemmaHiddenSize)); + for (float & value : scaled_input) { + value *= embedding_scale; + } + const std::vector expected = + rmsnorm_reference(scaled_input, kGemmaHiddenSize, kGemmaPromptTokenCount, kGemmaRmsNormEps); + + ggml_backend_tensor_set(input, input_data.data(), 0, input_data.size() * sizeof(float)); + REQUIRE(ggml_backend_sched_graph_compute(sched, graph) == GGML_STATUS_SUCCESS); + ggml_backend_sched_synchronize(sched); + + std::vector actual(expected.size()); + ggml_backend_tensor_get(output, actual.data(), 0, actual.size() * sizeof(float)); + require_close(actual, expected, 5.0e-4f); + + ggml_free(ctx); + ggml_backend_sched_free(sched); + ggml_backend_free(hrx_backend); + ggml_backend_free(cpu_backend); +} + +static void run_scheduled_hrx_rmsnorm_scale_case() { + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + ggml_backend_t cpu_backend = init_cpu_backend(); + REQUIRE(hrx_backend != nullptr); + + ggml_backend_t backends[] = { hrx_backend, cpu_backend }; + ggml_backend_sched_t sched = ggml_backend_sched_new(backends, nullptr, 2, GGML_DEFAULT_GRAPH_SIZE, false, true); + REQUIRE(sched != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(kGemmaHiddenSize * kGemmaPromptTokenCount * sizeof(float) * 8 + 1024 * 1024); + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kGemmaHiddenSize, kGemmaPromptTokenCount); + REQUIRE(input != nullptr); + ggml_backend_sched_set_tensor_backend(sched, input, hrx_backend); + + ggml_tensor * normalized = ggml_rms_norm(ctx, input, kGemmaRmsNormEps); + ggml_tensor * output = ggml_scale(ctx, normalized, 0.5f); + REQUIRE(normalized != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + REQUIRE(ggml_backend_sched_alloc_graph(sched, graph)); + REQUIRE(ggml_backend_sched_get_tensor_backend(sched, normalized) == hrx_backend); + REQUIRE(ggml_backend_sched_get_tensor_backend(sched, output) == hrx_backend); + + const std::vector input_data = + make_pattern_f32(static_cast(kGemmaHiddenSize * kGemmaPromptTokenCount), 37, 0.0005f); + std::vector expected = + rmsnorm_reference(input_data, kGemmaHiddenSize, kGemmaPromptTokenCount, kGemmaRmsNormEps); + for (float & value : expected) { + value *= 0.5f; + } + + ggml_backend_tensor_set(input, input_data.data(), 0, input_data.size() * sizeof(float)); + REQUIRE(ggml_backend_sched_graph_compute(sched, graph) == GGML_STATUS_SUCCESS); + ggml_backend_sched_synchronize(sched); + + std::vector actual(expected.size()); + ggml_backend_tensor_get(output, actual.data(), 0, actual.size() * sizeof(float)); + require_close(actual, expected, 5.0e-4f); + + ggml_free(ctx); + ggml_backend_sched_free(sched); + ggml_backend_free(hrx_backend); + ggml_backend_free(cpu_backend); +} + +static void run_scheduled_hrx_rmsnorm_view_scale_case() { + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + ggml_backend_t cpu_backend = init_cpu_backend(); + REQUIRE(hrx_backend != nullptr); + + ggml_backend_t backends[] = { hrx_backend, cpu_backend }; + ggml_backend_sched_t sched = ggml_backend_sched_new(backends, nullptr, 2, GGML_DEFAULT_GRAPH_SIZE, false, true); + REQUIRE(sched != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(kGemmaHiddenSize * kGemmaPromptTokenCount * sizeof(float) * 8 + 1024 * 1024); + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kGemmaHiddenSize, kGemmaPromptTokenCount); + REQUIRE(input != nullptr); + ggml_backend_sched_set_tensor_backend(sched, input, hrx_backend); + + ggml_tensor * normalized = ggml_rms_norm(ctx, input, kGemmaRmsNormEps); + ggml_tensor * view = ggml_reshape_2d(ctx, normalized, kGemmaHiddenSize, kGemmaPromptTokenCount); + ggml_tensor * output = ggml_scale(ctx, view, 0.25f); + REQUIRE(normalized != nullptr); + REQUIRE(view != nullptr); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + REQUIRE(ggml_backend_sched_alloc_graph(sched, graph)); + REQUIRE(ggml_backend_sched_get_tensor_backend(sched, normalized) == hrx_backend); + REQUIRE(ggml_backend_sched_get_tensor_backend(sched, output) == hrx_backend); + + const std::vector input_data = + make_pattern_f32(static_cast(kGemmaHiddenSize * kGemmaPromptTokenCount), 41, 0.0005f); + std::vector expected = + rmsnorm_reference(input_data, kGemmaHiddenSize, kGemmaPromptTokenCount, kGemmaRmsNormEps); + for (float & value : expected) { + value *= 0.25f; + } + + ggml_backend_tensor_set(input, input_data.data(), 0, input_data.size() * sizeof(float)); + REQUIRE(ggml_backend_sched_graph_compute(sched, graph) == GGML_STATUS_SUCCESS); + ggml_backend_sched_synchronize(sched); + + std::vector actual(expected.size()); + ggml_backend_tensor_get(output, actual.data(), 0, actual.size() * sizeof(float)); + require_close(actual, expected, 5.0e-4f); + + ggml_free(ctx); + ggml_backend_sched_free(sched); + ggml_backend_free(hrx_backend); + ggml_backend_free(cpu_backend); +} + +static void run_external_view_input_rmsnorm_case() { + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + static constexpr int64_t kHeadSize = 256; + static constexpr int64_t kHeadCount = 16; + const int64_t row_count = kHeadCount * kGemmaPromptTokenCount; + + ggml_init_params params = {}; + params.mem_size = static_cast(kHeadSize * row_count * sizeof(float) * 4 + 1024 * 1024); + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * root = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, kHeadSize, kHeadCount, kGemmaPromptTokenCount, 1); + REQUIRE(root != nullptr); + ggml_tensor * input_view = ggml_view_4d(ctx, root, kHeadSize, kHeadCount, kGemmaPromptTokenCount, 1, root->nb[1], + root->nb[2], root->nb[3], 0); + REQUIRE(input_view != nullptr); + ggml_tensor * output = ggml_rms_norm(ctx, input_view, kGemmaRmsNormEps); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_graph_add_node(graph, output); + + require_kernel_subsequence(scheduled_kernel_sequence(graph), { "loom_libs:ggml_rmsnorm_f32" }); + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, hrx_backend); + REQUIRE(buffer != nullptr); + + const std::vector input = make_pattern_f32(static_cast(kHeadSize * row_count), 29, 0.0025f); + const std::vector expected = rmsnorm_reference(input, kHeadSize, row_count, kGemmaRmsNormEps); + + ggml_backend_tensor_set(root, input.data(), 0, input.size() * sizeof(float)); + REQUIRE(ggml_backend_graph_compute(hrx_backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, output), expected, 5.0e-4f); + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(hrx_backend); +} + +static void run_split_local_view_alias_import_case() { + static constexpr int64_t kHeadSize = 256; + static constexpr int64_t kKeyValueHeads = 8; + static constexpr int64_t kKeyValueHidden = kHeadSize * kKeyValueHeads; + static constexpr int64_t kRootHeads = kKeyValueHeads + 1; + const size_t split_offset = static_cast(kHeadSize * sizeof(float)); + + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = + static_cast(kHeadSize * kRootHeads * kGemmaPromptTokenCount * sizeof(float) * 4 + 1024 * 1024); + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * root = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, kHeadSize, kRootHeads, kGemmaPromptTokenCount, 1); + REQUIRE(root != nullptr); + ggml_tensor * split_input = ggml_view_4d(ctx, root, kHeadSize, kKeyValueHeads, kGemmaPromptTokenCount, 1, + root->nb[1], root->nb[2], root->nb[3], split_offset); + REQUIRE(split_input != nullptr); + ggml_tensor * flattened = + ggml_view_2d(ctx, split_input, kKeyValueHidden, kGemmaPromptTokenCount, split_input->nb[2], 0); + REQUIRE(flattened != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_graph_add_node(graph, flattened); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + REQUIRE(imported.graph.nodes().size() == 1); + REQUIRE(ggml::hrx::is_layout_alias_node(imported.graph, imported.graph.nodes().front())); + + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, { "gfx1151" })); + REQUIRE(scheduler.plan().dispatches.empty()); + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, hrx_backend); + REQUIRE(buffer != nullptr); + + const size_t root_element_count = static_cast(kHeadSize * kRootHeads * kGemmaPromptTokenCount); + const std::vector root_data = make_pattern_f32(root_element_count, 37, 0.001f); + ggml_backend_tensor_set(root, root_data.data(), 0, root_data.size() * sizeof(float)); + REQUIRE(ggml_backend_graph_compute(hrx_backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(hrx_backend); + + std::vector actual(static_cast(kKeyValueHidden * kGemmaPromptTokenCount)); + ggml_backend_tensor_get(flattened, actual.data(), 0, actual.size() * sizeof(float)); + const float * expected = root_data.data() + split_offset / sizeof(float); + require_close(actual, std::vector(expected, expected + actual.size()), 0.0f); + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(hrx_backend); +} + +static void run_router_projection_case(int64_t token_count) { + static constexpr int64_t kHiddenSize = 2048; + static constexpr int64_t kExpertCount = 128; + + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast( + (kHiddenSize * token_count + kHiddenSize * kExpertCount + kExpertCount * token_count) * sizeof(float) * 4 + + 1024 * 1024); + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kHiddenSize, kExpertCount); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kHiddenSize, token_count); + REQUIRE(weight != nullptr); + REQUIRE(input != nullptr); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + require_kernel_subsequence(scheduled_kernel_sequence(graph), + { "qwen3_moe:qwen3_moe_router_projection_f32_four_row_wave32" }); + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + + const std::vector input_data = make_router_input(kHiddenSize, token_count); + const std::vector weight_data = make_router_weight(kHiddenSize, kExpertCount); + const std::vector expected = + router_projection_reference(input_data, weight_data, kHiddenSize, kExpertCount, token_count); + + ggml_backend_tensor_set(input, input_data.data(), 0, input_data.size() * sizeof(float)); + ggml_backend_tensor_set(weight, weight_data.data(), 0, weight_data.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + + std::vector actual(expected.size()); + ggml_backend_tensor_get(output, actual.data(), 0, actual.size() * sizeof(float)); + for (size_t i = 0; i < actual.size(); ++i) { + const float diff = std::fabs(actual[i] - expected[i]); + REQUIRE(diff <= 1.0e-2f); + } + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(backend); +} + +static void run_router_top8_case(int64_t token_count) { + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(kQwenRouterExpertCount * token_count * sizeof(float) * 16 + 1024 * 1024); + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * logits = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenRouterExpertCount, token_count); + REQUIRE(logits != nullptr); + ggml_tensor * output = build_qwen_router_top8_graph(ctx, logits); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + require_kernel_subsequence(scheduled_kernel_sequence(graph), { "qwen3_moe:qwen3_moe_router_top8_f32" }); + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + + const std::vector logits_data = make_router_logits(token_count); + const std::vector expected = router_top8_weights_reference(logits_data, token_count); + + ggml_backend_tensor_set(logits, logits_data.data(), 0, logits_data.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + + std::vector actual(expected.size()); + ggml_backend_tensor_get(output, actual.data(), 0, actual.size() * sizeof(float)); + for (size_t i = 0; i < actual.size(); ++i) { + const float diff = std::fabs(actual[i] - expected[i]); + REQUIRE(diff <= 1.0e-5f); + } + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(backend); +} + +static void run_qwen_flash_attention_case() { + static constexpr int64_t kQueryTokenCount = 2; + static constexpr int64_t kKeyValueTokenCount = 4; + + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 2 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * query = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, kQwenFlashHeadSize, kQueryTokenCount, 1); + ggml_tensor * key = ggml_new_tensor_3d(ctx, GGML_TYPE_F16, kQwenFlashHeadSize, kKeyValueTokenCount, 1); + ggml_tensor * value = ggml_new_tensor_3d(ctx, GGML_TYPE_F16, kQwenFlashHeadSize, kKeyValueTokenCount, 1); + ggml_tensor * mask = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, kKeyValueTokenCount, kQueryTokenCount); + REQUIRE(query != nullptr); + REQUIRE(key != nullptr); + REQUIRE(value != nullptr); + REQUIRE(mask != nullptr); + ggml_tensor * output = build_qwen_flash_attention_graph(ctx, query, key, value, mask); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + require_kernel_subsequence(scheduled_kernel_sequence(graph), + { "loom_libs:ggml_flash_attention_decode_split_f32_f16_wmma_next_q8" }); + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + REQUIRE(scheduler.schedule_graph(imported.graph, { "gfx1151" })); + const ggml::hrx::Value * output_value = imported.graph.values().find_tensor(output); + REQUIRE(output_value != nullptr); + REQUIRE(scheduler.plan().metadata.find_alternate_value( + output_value->id, GGML_TYPE_Q8_1, + static_cast(kQueryTokenCount) * ggml_row_size(GGML_TYPE_Q8_1, kQwenFlashHeadSize)) == nullptr); + static constexpr const char * kDisableQwenDispatchEnv = "GGML_HRX_DISABLE_QWEN_DISPATCH"; + const char * original_env = std::getenv(kDisableQwenDispatchEnv); + const bool had_original_env = original_env != nullptr; + const std::string original_env_value = had_original_env ? original_env : ""; + REQUIRE(setenv(kDisableQwenDispatchEnv, "1", 1) == 0); + ggml_tensor * common_query = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, kQwenFlashHeadSize, 16, 1); + ggml_tensor * common_key = ggml_new_tensor_3d(ctx, GGML_TYPE_F16, kQwenFlashHeadSize, 16, 1); + ggml_tensor * common_value = ggml_new_tensor_3d(ctx, GGML_TYPE_F16, kQwenFlashHeadSize, 16, 1); + ggml_tensor * common_mask = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, 16, 16); + REQUIRE(common_query != nullptr); + REQUIRE(common_key != nullptr); + REQUIRE(common_value != nullptr); + REQUIRE(common_mask != nullptr); + ggml_tensor * common_output = + build_qwen_flash_attention_graph(ctx, common_query, common_key, common_value, common_mask); + ggml_cgraph * common_graph = ggml_new_graph(ctx); + REQUIRE(common_graph != nullptr); + ggml_build_forward_expand(common_graph, common_output); + require_kernel_subsequence(scheduled_kernel_sequence(common_graph), + { "loom_libs:ggml_flash_attention_f32_f16_wmma" }); + restore_environment_value(kDisableQwenDispatchEnv, had_original_env, original_env_value); + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + + const std::vector query_data = make_flash_query(kQueryTokenCount); + const std::vector key_data = make_flash_key_value(kKeyValueTokenCount, 3); + const std::vector value_data = make_flash_key_value(kKeyValueTokenCount, 11); + const std::vector mask_data = make_flash_mask(kQueryTokenCount, kKeyValueTokenCount); + const std::vector expected = + flash_attention_reference(query_data, key_data, value_data, mask_data, kQueryTokenCount, kKeyValueTokenCount); + + ggml_backend_tensor_set(query, query_data.data(), 0, query_data.size() * sizeof(float)); + ggml_backend_tensor_set(key, key_data.data(), 0, key_data.size() * sizeof(ggml_fp16_t)); + ggml_backend_tensor_set(value, value_data.data(), 0, value_data.size() * sizeof(ggml_fp16_t)); + ggml_backend_tensor_set(mask, mask_data.data(), 0, mask_data.size() * sizeof(ggml_fp16_t)); + + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + + std::vector actual(expected.size()); + ggml_backend_tensor_get(output, actual.data(), 0, actual.size() * sizeof(float)); + for (size_t i = 0; i < actual.size(); ++i) { + const float diff = std::fabs(actual[i] - expected[i]); + REQUIRE(diff <= 5.0e-2f); + } + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(backend); +} + +static void run_common_flash_attention_cpu_reference_case(int64_t head_size, + int64_t query_token_count, + int64_t key_value_token_count, + float scale, + const char * expected_kernel, + int64_t value_head_size = -1) { + if (value_head_size < 0) { + value_head_size = head_size; + } + static constexpr const char * kDisableQwenDispatchEnv = "GGML_HRX_DISABLE_QWEN_DISPATCH"; + const char * original_env = std::getenv(kDisableQwenDispatchEnv); + const bool had_original_env = original_env != nullptr; + const std::string original_env_value = had_original_env ? original_env : ""; + REQUIRE(setenv(kDisableQwenDispatchEnv, "1", 1) == 0); + + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * query = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, head_size, query_token_count, 1); + ggml_tensor * key = ggml_new_tensor_3d(ctx, GGML_TYPE_F16, head_size, key_value_token_count, 1); + ggml_tensor * value = ggml_new_tensor_3d(ctx, GGML_TYPE_F16, value_head_size, key_value_token_count, 1); + ggml_tensor * mask = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, key_value_token_count, query_token_count); + REQUIRE(query != nullptr); + REQUIRE(key != nullptr); + REQUIRE(value != nullptr); + REQUIRE(mask != nullptr); + ggml_tensor * output = build_qwen_flash_attention_graph(ctx, query, key, value, mask, head_size, scale); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + require_kernel_subsequence(scheduled_kernel_sequence(graph), { expected_kernel }); + const ggml::hrx::WorkloadSpecialization expected_specialization = + query_token_count == 256 && key_value_token_count == 256 ? ggml::hrx::WorkloadSpecialization::Exact : + ggml::hrx::WorkloadSpecialization::Dynamic; + REQUIRE(scheduled_kernel_specialization(graph, expected_kernel).workload_specialization == + expected_specialization); + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + + const std::vector query_data = make_flash_query(query_token_count, head_size); + const std::vector key_data = make_flash_key_value(key_value_token_count, 3, head_size); + const std::vector value_data = make_flash_key_value(key_value_token_count, 11, value_head_size); + const std::vector mask_data = make_flash_mask(query_token_count, key_value_token_count); + const std::vector expected = + flash_attention_reference(query_data, key_data, value_data, mask_data, query_token_count, key_value_token_count, + head_size, value_head_size, scale); + + ggml_backend_tensor_set(query, query_data.data(), 0, query_data.size() * sizeof(float)); + ggml_backend_tensor_set(key, key_data.data(), 0, key_data.size() * sizeof(ggml_fp16_t)); + ggml_backend_tensor_set(value, value_data.data(), 0, value_data.size() * sizeof(ggml_fp16_t)); + ggml_backend_tensor_set(mask, mask_data.data(), 0, mask_data.size() * sizeof(ggml_fp16_t)); + + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + + std::vector actual(expected.size()); + ggml_backend_tensor_get(output, actual.data(), 0, actual.size() * sizeof(float)); + for (size_t i = 0; i < actual.size(); ++i) { + const float diff = std::fabs(actual[i] - expected[i]); + REQUIRE(diff <= 5.0e-2f); + } + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(backend); + restore_environment_value(kDisableQwenDispatchEnv, had_original_env, original_env_value); +} + +static void run_asymmetric_flash_attention_cpu_reference_case() { + static constexpr const char * kDisableQwenDispatchEnv = "GGML_HRX_DISABLE_QWEN_DISPATCH"; + const char * original_env = std::getenv(kDisableQwenDispatchEnv); + const bool had_original_env = original_env != nullptr; + const std::string original_env_value = had_original_env ? original_env : ""; + REQUIRE(setenv(kDisableQwenDispatchEnv, "1", 1) == 0); + + constexpr int64_t query_token_count = 23; + constexpr int64_t key_value_token_count = 256; + constexpr int64_t qk_head_size = 96; + constexpr int64_t value_head_size = 64; + const float scale = 1.0f / std::sqrt(static_cast(qk_head_size)); + + ggml_backend_t backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * query = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, qk_head_size, query_token_count, 1); + ggml_tensor * key = ggml_new_tensor_3d(ctx, GGML_TYPE_F16, qk_head_size, key_value_token_count, 1); + ggml_tensor * value = ggml_new_tensor_3d(ctx, GGML_TYPE_F16, value_head_size, key_value_token_count, 1); + ggml_tensor * mask = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, key_value_token_count, query_token_count); + REQUIRE(query != nullptr); + REQUIRE(key != nullptr); + REQUIRE(value != nullptr); + REQUIRE(mask != nullptr); + ggml_tensor * output = build_qwen_flash_attention_graph(ctx, query, key, value, mask, qk_head_size, scale); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + require_kernel_subsequence(scheduled_kernel_sequence(graph), { "loom_libs:ggml_flash_attention_f32_f16_wmma" }); + REQUIRE(scheduled_kernel_specialization(graph, "loom_libs:ggml_flash_attention_f32_f16_wmma") + .workload_specialization == ggml::hrx::WorkloadSpecialization::Dynamic); + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + + const std::vector query_data = make_flash_query(query_token_count, qk_head_size); + const std::vector key_data = make_flash_key_value(key_value_token_count, 3, qk_head_size); + const std::vector value_data = make_flash_key_value(key_value_token_count, 11, value_head_size); + const std::vector mask_data = make_flash_mask(query_token_count, key_value_token_count); + const std::vector expected = + flash_attention_reference(query_data, key_data, value_data, mask_data, query_token_count, key_value_token_count, + qk_head_size, value_head_size, scale); + + ggml_backend_tensor_set(query, query_data.data(), 0, query_data.size() * sizeof(float)); + ggml_backend_tensor_set(key, key_data.data(), 0, key_data.size() * sizeof(ggml_fp16_t)); + ggml_backend_tensor_set(value, value_data.data(), 0, value_data.size() * sizeof(ggml_fp16_t)); + ggml_backend_tensor_set(mask, mask_data.data(), 0, mask_data.size() * sizeof(ggml_fp16_t)); + + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + + std::vector actual(expected.size()); + ggml_backend_tensor_get(output, actual.data(), 0, actual.size() * sizeof(float)); + for (size_t i = 0; i < actual.size(); ++i) { + const float diff = std::fabs(actual[i] - expected[i]); + REQUIRE(diff <= 5.0e-2f); + } + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(backend); + restore_environment_value(kDisableQwenDispatchEnv, had_original_env, original_env_value); +} + +static void run_qwen_full_cache_prefill_flash_attention_scheduling_case(int64_t token_count) { + constexpr int64_t kQueryHeadCount = 32; + constexpr int64_t kKeyValueHeadCount = 4; + constexpr int64_t kFullCacheRowCount = 40960; + const size_t active_mask_byte_count = + static_cast(token_count) * static_cast(token_count) * sizeof(ggml_fp16_t); + + ggml_init_params params = {}; + params.mem_size = 128 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + AttentionPostprocessGraph attention = + build_attention_postprocess_graph(ctx, token_count, kQueryHeadCount, kKeyValueHeadCount, kFullCacheRowCount); + ggml_tensor * flash_output = append_qwen_full_cache_flash_attention_consumer( + ctx, attention, token_count, kQueryHeadCount, kKeyValueHeadCount, kFullCacheRowCount); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, attention.key_output); + ggml_build_forward_expand(graph, attention.value_output); + ggml_build_forward_expand(graph, flash_output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + const ggml::hrx::DispatchRegistry * registry = ggml::hrx::find_dispatch_registry({ "gfx1151" }); + REQUIRE(registry != nullptr); + + std::vector covered_nodes(imported.graph.nodes().size(), false); + ggml::hrx::CommandPlan plan; + ggml::hrx::DispatchMatch postprocess_match; + match_dispatch_at_index(imported.graph, *registry, plan, covered_nodes, + producer_index_for_tensor(imported.graph, attention.query_reshape), postprocess_match); + ggml::hrx::DispatchMatch flash_match; + match_dispatch_at_index(imported.graph, *registry, plan, covered_nodes, + producer_index_for_tensor(imported.graph, flash_output), flash_match); + + const ggml::hrx::Value * mask_value = imported.graph.values().find_tensor(attention.attention_mask); + REQUIRE(mask_value != nullptr); + const ggml::hrx::CommandPlanAlternateValue * compact_mask = + ggml::hrx::find_alternate_value(plan, mask_value->id, GGML_TYPE_F16, active_mask_byte_count); + REQUIRE(compact_mask != nullptr); + + const ggml::hrx::Dispatch * metadata_dispatch = nullptr; + for (const ggml::hrx::Dispatch & dispatch : plan.initialization_dispatches) { + if (kernel_name_for_id(dispatch.kernel.kernel_id) == "qwen3_moe:qwen_attention_metadata") { + metadata_dispatch = &dispatch; + } + } + REQUIRE(metadata_dispatch != nullptr); + REQUIRE(metadata_dispatch->kernel.integer_parameters.at("token_count") == token_count); + REQUIRE(metadata_dispatch->kernel.integer_parameters.at("context_capacity") == token_count); + REQUIRE(metadata_dispatch->bindings.size() == 5); + REQUIRE(metadata_dispatch->bindings[4].value == compact_mask->alternate_value); + REQUIRE(metadata_dispatch->bindings[4].length == active_mask_byte_count); + + const ggml::hrx::Dispatch * flash_dispatch = nullptr; + for (const ggml::hrx::Dispatch & dispatch : plan.dispatches) { + if (kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_flash_attention_f32_f16_wmma") { + flash_dispatch = &dispatch; + } + } + REQUIRE(flash_dispatch != nullptr); + REQUIRE(flash_dispatch->kernel.integer_parameters.at("query_token_count") == token_count); + REQUIRE(flash_dispatch->kernel.integer_parameters.at("key_value_token_count") == token_count); + REQUIRE(flash_dispatch->bindings.size() == 6); + REQUIRE(flash_dispatch->kernel.compile_parameters.at("ggml.flash_attention.apply_gate") == "0"); + REQUIRE(flash_dispatch->bindings[3].value == compact_mask->alternate_value); + REQUIRE(flash_dispatch->bindings[3].length == active_mask_byte_count); + + ggml_free(ctx); +} + +static void run_qwen_decode_split_flash_attention_scheduling_case(int64_t query_token_count, + int64_t key_value_token_count, + bool use_common_dispatch = false, + int64_t qk_head_size = kQwenFlashHeadSize, + int64_t value_head_size = kQwenFlashHeadSize, + float scale = 0.0f) { + static constexpr const char * kDisableQwenDispatchEnv = "GGML_HRX_DISABLE_QWEN_DISPATCH"; + const char * original_env = std::getenv(kDisableQwenDispatchEnv); + const bool had_original_env = original_env != nullptr; + const std::string original_env_value = had_original_env ? original_env : ""; + if (use_common_dispatch) { + REQUIRE(setenv(kDisableQwenDispatchEnv, "1", 1) == 0); + } else { + REQUIRE(unsetenv(kDisableQwenDispatchEnv) == 0); + } + + constexpr int64_t query_head_count = 32; + constexpr int64_t key_value_head_count = 4; + const int64_t query_hidden_size = query_head_count * qk_head_size; + const int64_t output_hidden_size = query_head_count * value_head_size; + const size_t q8_row_bytes = ggml_row_size(GGML_TYPE_Q8_1, output_hidden_size); + const size_t q8_output_bytes = static_cast(query_token_count) * q8_row_bytes; + int64_t key_value_capacity = 64; + while (key_value_capacity < key_value_token_count) { + key_value_capacity *= 2; + } + const int64_t key_value_blocks = key_value_capacity / 64; + const size_t partial_scalar_bytes = + static_cast(key_value_head_count * key_value_blocks * 16) * sizeof(float); + const size_t partial_output_bytes = + static_cast(key_value_head_count * key_value_blocks * 16 * value_head_size) * sizeof(ggml_fp16_t); + + ggml_init_params params = {}; + params.mem_size = 128 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + QwenFlashAttentionLayoutGraph attention = build_qwen_flash_attention_layout_graph( + ctx, query_token_count, key_value_token_count, qk_head_size, value_head_size, scale); + ggml_tensor * attention_output = + ggml_reshape_2d(ctx, attention.output, output_hidden_size, query_token_count); + ggml_tensor * consumer_weight = + ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, output_hidden_size, kQwenHiddenSize); + ggml_tensor * consumer_output = ggml_mul_mat(ctx, consumer_weight, attention_output); + REQUIRE(attention_output != nullptr); + REQUIRE(consumer_weight != nullptr); + REQUIRE(consumer_output != nullptr); + + ggml_cgraph * cgraph = ggml_new_graph(ctx); + REQUIRE(cgraph != nullptr); + ggml_build_forward_expand(cgraph, consumer_output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*cgraph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + ggml::hrx::DispatchScheduleDiagnostics diagnostics; + REQUIRE(scheduler.schedule_graph(imported.graph, { "gfx1151" }, &diagnostics)); + const ggml::hrx::CommandPlan & plan = scheduler.plan(); + REQUIRE(plan.valid()); + REQUIRE(plan.dispatches.size() == static_cast(query_token_count + 1)); + REQUIRE(plan.transients.size() == 4); + REQUIRE(plan.completion_counter_requests.size() == 1); + + REQUIRE(plan.transients[0].size == partial_scalar_bytes); + REQUIRE(plan.transients[1].size == partial_scalar_bytes); + REQUIRE(plan.transients[2].size == partial_output_bytes); + REQUIRE(plan.transients[3].size == q8_output_bytes); + REQUIRE(plan.completion_counter_requests[0].count == key_value_head_count); + + const ggml::hrx::Value * query_value = imported.graph.values().find_tensor(attention.query); + const ggml::hrx::Value * mask_value = imported.graph.values().find_tensor(attention.mask); + const ggml::hrx::Value * output_value = imported.graph.values().find_tensor(attention.output); + const ggml::hrx::Value * publication_value = imported.graph.values().find_tensor(attention_output); + REQUIRE(query_value != nullptr); + REQUIRE(mask_value != nullptr); + REQUIRE(output_value != nullptr); + REQUIRE(publication_value != nullptr); + const ggml::hrx::CommandPlanAlternateValue * q8_alternate = + plan.metadata.find_alternate_value(publication_value->id, GGML_TYPE_Q8_1, q8_output_bytes); + REQUIRE(q8_alternate != nullptr); + REQUIRE(q8_alternate->alternate_value == plan.transients[3].value); + + const size_t query_row_bytes = static_cast(query_hidden_size) * sizeof(float); + const size_t mask_row_bytes = static_cast(key_value_token_count) * sizeof(ggml_fp16_t); + const size_t output_row_bytes = static_cast(output_hidden_size) * sizeof(float); + for (int64_t row = 0; row < query_token_count; ++row) { + const ggml::hrx::Dispatch & dispatch = plan.dispatches[static_cast(row)]; + const char * expected_kernel = "loom_libs:ggml_flash_attention_decode_split_f32_f16_wmma_next_q8"; + REQUIRE(kernel_name_for_id(dispatch.kernel.kernel_id) == expected_kernel); + REQUIRE(dispatch.kernel.workload_specialization == ggml::hrx::WorkloadSpecialization::Dynamic); + REQUIRE(dispatch.kernel.integer_parameters.at("key_value_token_count") == key_value_token_count); + REQUIRE(dispatch.kernel.compile_parameters.at("ggml.flash_attention.query_head_count") == + std::to_string(query_head_count)); + REQUIRE(dispatch.kernel.compile_parameters.at("ggml.flash_attention.key_value_head_count") == + std::to_string(key_value_head_count)); + REQUIRE(dispatch.kernel.compile_parameters.at("ggml.flash_attention.qk_head_size") == + std::to_string(qk_head_size)); + REQUIRE(dispatch.kernel.compile_parameters.at("ggml.flash_attention.value_head_size") == + std::to_string(value_head_size)); + REQUIRE(dispatch.kernel.compile_parameters.at("ggml.flash_attention.attention_scale") == + expected_config_value(scale == 0.0f ? 1.0f / std::sqrt(static_cast(qk_head_size)) : scale)); + REQUIRE(dispatch.kernel.compile_parameters.at("ggml.flash_attention.decode.key_value_token_capacity") == + std::to_string(key_value_capacity)); + REQUIRE(dispatch.bindings.size() == 10); + REQUIRE(dispatch.bindings[0].value == query_value->id); + REQUIRE(dispatch.bindings[0].offset == static_cast(row) * query_value->nb[1]); + REQUIRE(dispatch.bindings[0].length == query_row_bytes); + REQUIRE(dispatch.bindings[3].value == mask_value->id); + REQUIRE(dispatch.bindings[3].offset == static_cast(row) * mask_value->nb[1]); + REQUIRE(dispatch.bindings[3].length == mask_row_bytes); + REQUIRE(dispatch.bindings[8].value == output_value->id); + REQUIRE(dispatch.bindings[8].offset == static_cast(row) * output_value->nb[2]); + REQUIRE(dispatch.bindings[8].length == output_row_bytes); + REQUIRE(dispatch.bindings[9].value == q8_alternate->alternate_value); + REQUIRE(dispatch.bindings[9].offset == static_cast(row) * q8_row_bytes); + REQUIRE(dispatch.bindings[9].length == q8_row_bytes); + } + + ggml_free(ctx); + restore_environment_value(kDisableQwenDispatchEnv, had_original_env, original_env_value); +} + +static void run_qwen_decode_attention_output_next_q8_scheduling_case(bool include_get_rows_selectors) { + constexpr int64_t query_token_count = 1; + constexpr int64_t key_value_token_count = 512; + constexpr int64_t attention_hidden_size = 4096; + const size_t attention_q8_bytes = ggml_row_size(GGML_TYPE_Q8_1, attention_hidden_size); + const size_t next_q8_bytes = ggml_row_size(GGML_TYPE_Q8_1, kQwenHiddenSize); + + ggml_init_params params = {}; + params.mem_size = 128 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + QwenFlashAttentionLayoutGraph attention = + build_qwen_flash_attention_layout_graph(ctx, query_token_count, key_value_token_count); + ggml_tensor * attention_output = ggml_reshape_2d(ctx, attention.output, attention_hidden_size, query_token_count); + ggml_tensor * output_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, attention_hidden_size, kQwenHiddenSize); + ggml_tensor * projection = ggml_mul_mat(ctx, output_weight, attention_output); + ggml_tensor * residual_input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenHiddenSize, query_token_count); + ggml_tensor * selected_projection = projection; + ggml_tensor * selected_residual = residual_input; + ggml_tensor * row_indices = nullptr; + if (include_get_rows_selectors) { + row_indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, 1); + selected_projection = ggml_get_rows(ctx, projection, row_indices); + selected_residual = ggml_get_rows(ctx, residual_input, row_indices); + REQUIRE(row_indices != nullptr); + REQUIRE(selected_projection != nullptr); + REQUIRE(selected_residual != nullptr); + } + ggml_tensor * residual = ggml_add(ctx, selected_projection, selected_residual); + ggml_tensor * norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, kQwenHiddenSize); + ggml_tensor * rms = ggml_rms_norm(ctx, residual, kQwenRmsNormEps); + ggml_tensor * normalized = ggml_mul(ctx, rms, norm_weight); + ggml_tensor * next_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, kQwenHiddenSize, kQwenHiddenSize); + ggml_tensor * next_output = ggml_mul_mat(ctx, next_weight, normalized); + REQUIRE(attention_output != nullptr); + REQUIRE(output_weight != nullptr); + REQUIRE(projection != nullptr); + REQUIRE(residual_input != nullptr); + REQUIRE(residual != nullptr); + REQUIRE(norm_weight != nullptr); + REQUIRE(rms != nullptr); + REQUIRE(normalized != nullptr); + REQUIRE(next_weight != nullptr); + REQUIRE(next_output != nullptr); + + ggml_cgraph * cgraph = ggml_new_graph(ctx); + REQUIRE(cgraph != nullptr); + ggml_build_forward_expand(cgraph, next_output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*cgraph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + ggml::hrx::DispatchScheduleDiagnostics diagnostics; + REQUIRE(scheduler.schedule_graph(imported.graph, { "gfx1151" }, &diagnostics)); + const ggml::hrx::CommandPlan & plan = scheduler.plan(); + REQUIRE(plan.valid()); + REQUIRE(plan.dispatches.size() == 3); + REQUIRE(plan.completion_counter_requests.size() == 2); + + const auto projection_it = std::find_if(plan.dispatches.begin(), plan.dispatches.end(), [](const auto & dispatch) { + return kernel_name_for_id(dispatch.kernel.kernel_id) == + "qwen3_moe:qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8"; + }); + REQUIRE(projection_it != plan.dispatches.end()); + const ggml::hrx::Dispatch & projection_dispatch = *projection_it; + REQUIRE(kernel_name_for_id(projection_dispatch.kernel.kernel_id) == + "qwen3_moe:qwen3_moe_dense_linear_q4k_q8_1_x4_next_q8"); + REQUIRE(projection_dispatch.kernel.integer_parameters.at("token_count") == query_token_count); + REQUIRE(projection_dispatch.kernel.compile_parameters.at("qwen3_moe.dense_quantized.input_size") == + std::to_string(attention_hidden_size)); + REQUIRE(projection_dispatch.kernel.compile_parameters.at("qwen3_moe.dense_quantized.output_size") == + std::to_string(kQwenHiddenSize)); + REQUIRE(projection_dispatch.kernel.compile_parameters.at("qwen3_moe.dense_quantized.output_accumulation") == "1"); + REQUIRE(projection_dispatch.bindings.size() == 7); + + const ggml::hrx::Value * flash_output_value = imported.graph.values().find_tensor(attention_output); + const ggml::hrx::Value * residual_input_value = imported.graph.values().find_tensor(residual_input); + const ggml::hrx::Value * residual_value = imported.graph.values().find_tensor(residual); + const ggml::hrx::Value * normalized_value = imported.graph.values().find_tensor(normalized); + REQUIRE(flash_output_value != nullptr); + REQUIRE(residual_input_value != nullptr); + REQUIRE(residual_value != nullptr); + REQUIRE(normalized_value != nullptr); + const ggml::hrx::CommandPlanAlternateValue * attention_q8 = + ggml::hrx::find_alternate_value(imported.graph, plan, flash_output_value->id, GGML_TYPE_Q8_1, + attention_q8_bytes); + const ggml::hrx::CommandPlanAlternateValue * next_q8 = + ggml::hrx::find_alternate_value(imported.graph, plan, normalized_value->id, GGML_TYPE_Q8_1, next_q8_bytes); + REQUIRE(attention_q8 != nullptr); + REQUIRE(next_q8 != nullptr); + REQUIRE(projection_dispatch.bindings[0].value == attention_q8->alternate_value); + REQUIRE(projection_dispatch.bindings[2].value == residual_value->id); + REQUIRE(projection_dispatch.bindings[4].value == normalized_value->id); + REQUIRE(projection_dispatch.bindings[6].value == next_q8->alternate_value); + + const ggml::hrx::Value * aliased_residual = imported.graph.values().find(residual_value->id); + REQUIRE(aliased_residual != nullptr); + REQUIRE(aliased_residual->alias_source == residual_input_value->id); + + ggml_free(ctx); +} + +static void run_add_f32_cpu_reference_case() { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params cpu_params = {}; + cpu_params.mem_size = 256 * 1024; + cpu_params.no_alloc = true; + ggml_init_params hrx_params = cpu_params; + ggml_context * cpu_ctx = ggml_init(cpu_params); + ggml_context * hrx_ctx = ggml_init(hrx_params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + constexpr int64_t element_count = 257; + ggml_tensor * cpu_a = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, element_count); + ggml_tensor * cpu_b = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, element_count); + ggml_tensor * cpu_output = ggml_add(cpu_ctx, cpu_a, cpu_b); + ggml_tensor * hrx_a = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, element_count); + ggml_tensor * hrx_b = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, element_count); + ggml_tensor * hrx_output = ggml_add(hrx_ctx, hrx_a, hrx_b); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_binary_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector a = make_pattern_f32(element_count, 1, 0.125f); + const std::vector b = make_pattern_f32(element_count, 2, 0.25f); + set_tensor_pair_bytes(cpu_backend, cpu_a, hrx_backend, hrx_a, a.data(), a.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_b, hrx_backend, hrx_b, b.data(), b.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 0.0f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_dynamic_generic_ops_cpu_reference_case() { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 512 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + constexpr int64_t element_count = 257; + ggml_tensor * cpu_input = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, element_count); + ggml_tensor * hrx_input = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, element_count); + ggml_tensor * cpu_unary = ggml_sqr(cpu_ctx, cpu_input); + ggml_tensor * hrx_unary = ggml_sqr(hrx_ctx, hrx_input); + ggml_tensor * cpu_copy_target = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, element_count); + ggml_tensor * hrx_copy_target = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, element_count); + ggml_tensor * cpu_copy = ggml_cpy(cpu_ctx, cpu_input, cpu_copy_target); + ggml_tensor * hrx_copy = ggml_cpy(hrx_ctx, hrx_input, hrx_copy_target); + ggml_tensor * cpu_concat_lhs = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, 64, 7); + ggml_tensor * cpu_concat_rhs = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, 32, 7); + ggml_tensor * hrx_concat_lhs = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, 64, 7); + ggml_tensor * hrx_concat_rhs = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, 32, 7); + ggml_tensor * cpu_concat = ggml_concat(cpu_ctx, cpu_concat_lhs, cpu_concat_rhs, 0); + ggml_tensor * hrx_concat = ggml_concat(hrx_ctx, hrx_concat_lhs, hrx_concat_rhs, 0); + REQUIRE(cpu_unary != nullptr && hrx_unary != nullptr); + REQUIRE(cpu_copy != nullptr && hrx_copy != nullptr); + REQUIRE(cpu_concat != nullptr && hrx_concat != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr && hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_unary); + ggml_build_forward_expand(cpu_graph, cpu_copy); + ggml_build_forward_expand(cpu_graph, cpu_concat); + ggml_build_forward_expand(hrx_graph, hrx_unary); + ggml_build_forward_expand(hrx_graph, hrx_copy); + ggml_build_forward_expand(hrx_graph, hrx_concat); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), + { "loom_libs:ggml_unary_f32", "loom_libs:ggml_copy_f32", + "loom_libs:ggml_concat_dim0_f32" }); + for (const char * kernel_name : + { "loom_libs:ggml_unary_f32", "loom_libs:ggml_copy_f32", "loom_libs:ggml_concat_dim0_f32" }) { + const ggml::hrx::KernelSpecialization specialization = + scheduled_kernel_specialization(hrx_graph, kernel_name); + REQUIRE(specialization.workload_specialization == ggml::hrx::WorkloadSpecialization::Dynamic); + } + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr && hrx_buffer != nullptr); + const std::vector input = make_pattern_f32(element_count, 73, 0.03125f); + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), input.size() * sizeof(float)); + const std::vector concat_lhs = make_pattern_f32(64 * 7, 79, 0.015625f); + const std::vector concat_rhs = make_pattern_f32(32 * 7, 83, 0.0234375f); + set_tensor_pair_bytes(cpu_backend, cpu_concat_lhs, hrx_backend, hrx_concat_lhs, concat_lhs.data(), + concat_lhs.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_concat_rhs, hrx_backend, hrx_concat_rhs, concat_rhs.data(), + concat_rhs.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_unary), get_f32_tensor(cpu_backend, cpu_unary), 0.0f); + require_close(get_f32_tensor(hrx_backend, hrx_copy), get_f32_tensor(cpu_backend, cpu_copy), 0.0f); + require_close(get_f32_tensor(hrx_backend, hrx_concat), get_f32_tensor(cpu_backend, cpu_concat), 0.0f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_view_mul_f32_cpu_reference_case() { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 512 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + constexpr int64_t hidden_size = 2048; + constexpr int64_t token_count = 2; + ggml_tensor * cpu_source = ggml_new_tensor_3d(cpu_ctx, GGML_TYPE_F32, hidden_size, 3, token_count); + ggml_tensor * hrx_source = ggml_new_tensor_3d(hrx_ctx, GGML_TYPE_F32, hidden_size, 3, token_count); + REQUIRE(cpu_source != nullptr); + REQUIRE(hrx_source != nullptr); + ggml_tensor * cpu_lhs = ggml_view_2d(cpu_ctx, cpu_source, hidden_size, token_count, cpu_source->nb[2], 0); + ggml_tensor * cpu_rhs = + ggml_view_2d(cpu_ctx, cpu_source, hidden_size, token_count, cpu_source->nb[2], 2 * cpu_source->nb[1]); + ggml_tensor * hrx_lhs = ggml_view_2d(hrx_ctx, hrx_source, hidden_size, token_count, hrx_source->nb[2], 0); + ggml_tensor * hrx_rhs = + ggml_view_2d(hrx_ctx, hrx_source, hidden_size, token_count, hrx_source->nb[2], 2 * hrx_source->nb[1]); + REQUIRE(cpu_lhs != nullptr); + REQUIRE(cpu_rhs != nullptr); + REQUIRE(hrx_lhs != nullptr); + REQUIRE(hrx_rhs != nullptr); + ggml_tensor * cpu_output = ggml_mul(cpu_ctx, cpu_lhs, cpu_rhs); + ggml_tensor * hrx_output = ggml_mul(hrx_ctx, hrx_lhs, hrx_rhs); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_binary_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector source = make_pattern_f32(hidden_size * 3 * token_count, 19, 0.015625f); + set_tensor_pair_bytes(cpu_backend, cpu_source, hrx_backend, hrx_source, source.data(), + source.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 0.0f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_glu_split_f32_cpu_reference_case(ggml_glu_op glu_op) { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 256 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + constexpr int64_t hidden_size = 256; + constexpr int64_t token_count = 3; + ggml_tensor * cpu_gate = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, hidden_size, token_count); + ggml_tensor * cpu_up = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, hidden_size, token_count); + ggml_tensor * cpu_output = ggml_glu_split(cpu_ctx, cpu_gate, cpu_up, glu_op); + ggml_tensor * hrx_gate = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, hidden_size, token_count); + ggml_tensor * hrx_up = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, hidden_size, token_count); + ggml_tensor * hrx_output = ggml_glu_split(hrx_ctx, hrx_gate, hrx_up, glu_op); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_binary_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector gate = make_pattern_f32(hidden_size * token_count, 21, 0.02f); + const std::vector up = make_pattern_f32(hidden_size * token_count, 22, 0.03f); + set_tensor_pair_bytes(cpu_backend, cpu_gate, hrx_backend, hrx_gate, gate.data(), gate.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_up, hrx_backend, hrx_up, up.data(), up.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 5.0e-4f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_glu_packed_f32_cpu_reference_case(ggml_glu_op glu_op, + bool swapped, + int64_t hidden_size = 256, + int64_t token_count = 3) { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(16 * 1024 * 1024); + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_input = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, 2 * hidden_size, token_count); + ggml_tensor * cpu_output = ggml_glu(cpu_ctx, cpu_input, glu_op, swapped); + ggml_tensor * hrx_input = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, 2 * hidden_size, token_count); + ggml_tensor * hrx_output = ggml_glu(hrx_ctx, hrx_input, glu_op, swapped); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_binary_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector input = + make_pattern_f32(static_cast(2 * hidden_size * token_count), swapped ? 24 : 23, 0.02f); + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), input.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 5.0e-4f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_gather_add_f32_cpu_reference_case() { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 512 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + constexpr int64_t hidden_size = 256; + constexpr int64_t source_token_count = 11; + constexpr int64_t output_token_count = 5; + ggml_tensor * cpu_a = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, hidden_size, source_token_count); + ggml_tensor * cpu_b = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, hidden_size, source_token_count); + ggml_tensor * cpu_ids = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_I32, output_token_count); + ggml_tensor * cpu_rows_a = ggml_get_rows(cpu_ctx, cpu_a, cpu_ids); + ggml_tensor * cpu_rows_b = ggml_get_rows(cpu_ctx, cpu_b, cpu_ids); + ggml_tensor * cpu_output = ggml_add(cpu_ctx, cpu_rows_a, cpu_rows_b); + + ggml_tensor * hrx_a = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, hidden_size, source_token_count); + ggml_tensor * hrx_b = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, hidden_size, source_token_count); + ggml_tensor * hrx_ids = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_I32, output_token_count); + ggml_tensor * hrx_rows_a = ggml_get_rows(hrx_ctx, hrx_a, hrx_ids); + ggml_tensor * hrx_rows_b = ggml_get_rows(hrx_ctx, hrx_b, hrx_ids); + ggml_tensor * hrx_output = ggml_add(hrx_ctx, hrx_rows_a, hrx_rows_b); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "hrx:ggml_gather_add_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector a = make_pattern_f32(hidden_size * source_token_count, 3, 0.05f); + const std::vector b = make_pattern_f32(hidden_size * source_token_count, 4, 0.075f); + const std::vector ids = { 9, 3, 7, 1, 5 }; + set_tensor_pair_bytes(cpu_backend, cpu_a, hrx_backend, hrx_a, a.data(), a.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_b, hrx_backend, hrx_b, b.data(), b.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_ids, hrx_backend, hrx_ids, ids.data(), ids.size() * sizeof(int32_t)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 0.0f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_scale_add_f32_cpu_reference_case(bool strided_residual) { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 512 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + constexpr int64_t width = 257; + constexpr int64_t token_count = 3; + constexpr int64_t residual_width = 320; + ggml_tensor * cpu_input = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, width, token_count); + ggml_tensor * hrx_input = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, width, token_count); + ggml_tensor * cpu_residual_storage = + ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, strided_residual ? residual_width : width, token_count); + ggml_tensor * hrx_residual_storage = + ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, strided_residual ? residual_width : width, token_count); + ggml_tensor * cpu_residual = strided_residual ? ggml_view_2d(cpu_ctx, cpu_residual_storage, width, token_count, + cpu_residual_storage->nb[1], 0) : + cpu_residual_storage; + ggml_tensor * hrx_residual = strided_residual ? ggml_view_2d(hrx_ctx, hrx_residual_storage, width, token_count, + hrx_residual_storage->nb[1], 0) : + hrx_residual_storage; + ggml_tensor * cpu_scaled = ggml_scale_bias(cpu_ctx, cpu_input, 0.177800179f, -0.125f); + ggml_tensor * hrx_scaled = ggml_scale_bias(hrx_ctx, hrx_input, 0.177800179f, -0.125f); + ggml_tensor * cpu_output = ggml_add(cpu_ctx, cpu_scaled, cpu_residual); + ggml_tensor * hrx_output = ggml_add(hrx_ctx, hrx_scaled, hrx_residual); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_scale_add_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector input = make_pattern_f32(width * token_count, 31, 0.03125f); + const std::vector residual = + make_pattern_f32((strided_residual ? residual_width : width) * token_count, 37, 0.046875f); + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), input.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_residual_storage, hrx_backend, hrx_residual_storage, residual.data(), + residual.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 0.0f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_get_rows_f32_cpu_reference_case(ggml_type weight_type, + int64_t hidden_size = kQwenHiddenSize, + int64_t vocabulary_count = 64, + int64_t token_count = 7) { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_weight = ggml_new_tensor_2d(cpu_ctx, weight_type, hidden_size, vocabulary_count); + ggml_tensor * cpu_ids = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_I32, token_count); + ggml_tensor * cpu_output = ggml_get_rows(cpu_ctx, cpu_weight, cpu_ids); + ggml_tensor * hrx_weight = ggml_new_tensor_2d(hrx_ctx, weight_type, hidden_size, vocabulary_count); + ggml_tensor * hrx_ids = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_I32, token_count); + ggml_tensor * hrx_output = ggml_get_rows(hrx_ctx, hrx_weight, hrx_ids); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_get_rows_f32" }); + const ggml::hrx::KernelSpecialization specialization = + scheduled_kernel_specialization(hrx_graph, "loom_libs:ggml_get_rows_f32"); + REQUIRE(specialization.workload_specialization == ggml::hrx::WorkloadSpecialization::Dynamic); + int64_t expected_token_capacity = 1; + while (expected_token_capacity < token_count) { + expected_token_capacity *= 2; + } + REQUIRE(specialization.compile_parameters.at("ggml.get_rows_f32.token_capacity") == + std::to_string(expected_token_capacity)); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector weight = make_matmul_weight_bytes(weight_type, hidden_size, vocabulary_count, 5); + std::vector ids(static_cast(token_count)); + for (int64_t i = 0; i < token_count; ++i) { + ids[static_cast(i)] = static_cast((i * 17 + 3) % vocabulary_count); + } + set_tensor_pair_bytes(cpu_backend, cpu_weight, hrx_backend, hrx_weight, weight.data(), weight.size()); + set_tensor_pair_bytes(cpu_backend, cpu_ids, hrx_backend, hrx_ids, ids.data(), ids.size() * sizeof(int32_t)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 5.0e-4f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_get_rows_q8_1_zero_weight_case() { + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(hrx_ctx != nullptr); + + constexpr int64_t vocabulary_count = 64; + constexpr int64_t hidden_size = kQwenHiddenSize; + constexpr int64_t token_count = 7; + ggml_tensor * weight = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_Q8_1, hidden_size, vocabulary_count); + ggml_tensor * ids = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_I32, token_count); + ggml_tensor * output = ggml_get_rows(hrx_ctx, weight, ids); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(hrx_ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + require_kernel_subsequence(scheduled_kernel_sequence(graph), { "loom_libs:ggml_get_rows_f32" }); + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(buffer != nullptr); + + const std::vector weight_bytes(ggml_nbytes(weight), 0); + const std::vector id_values = { 3, 17, 29, 41, 53, 7, 19 }; + set_tensor_bytes(hrx_backend, weight, weight_bytes.data(), weight_bytes.size()); + set_tensor_bytes(hrx_backend, ids, id_values.data(), id_values.size() * sizeof(int32_t)); + + REQUIRE(ggml_backend_graph_compute(hrx_backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, output), + std::vector(static_cast(hidden_size * token_count), 0.0f), 0.0f, 0.0f); + + ggml_backend_buffer_free(buffer); + ggml_free(hrx_ctx); + ggml_backend_free(hrx_backend); +} + +static void run_get_rows_scale_f32_cpu_reference_case(ggml_type weight_type = GGML_TYPE_F32, + int64_t hidden_size = 256, + int64_t row_count = 8, + int64_t token_count = 6) { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 4 * 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + const float scale = weight_type == GGML_TYPE_Q5_1 ? 25.2982216f : 0.177800179f; + ggml_tensor * cpu_weight = ggml_new_tensor_2d(cpu_ctx, weight_type, hidden_size, row_count); + ggml_tensor * cpu_ids = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_I32, token_count); + ggml_tensor * cpu_rows = ggml_get_rows(cpu_ctx, cpu_weight, cpu_ids); + ggml_tensor * cpu_output = ggml_scale(cpu_ctx, cpu_rows, scale); + ggml_tensor * hrx_weight = ggml_new_tensor_2d(hrx_ctx, weight_type, hidden_size, row_count); + ggml_tensor * hrx_ids = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_I32, token_count); + ggml_tensor * hrx_rows = ggml_get_rows(hrx_ctx, hrx_weight, hrx_ids); + ggml_tensor * hrx_output = ggml_scale(hrx_ctx, hrx_rows, scale); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + ggml_set_output(cpu_rows); + ggml_set_output(hrx_rows); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_get_rows_scale_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector weight = make_matmul_weight_bytes(weight_type, hidden_size, row_count, 9); + const std::vector ids = { 0, static_cast(row_count - 1), 2, + 2, static_cast(row_count - 2), 1 }; + set_tensor_pair_bytes(cpu_backend, cpu_weight, hrx_backend, hrx_weight, weight.data(), weight.size()); + set_tensor_pair_bytes(cpu_backend, cpu_ids, hrx_backend, hrx_ids, ids.data(), ids.size() * sizeof(int32_t)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + const float tolerance = weight_type == GGML_TYPE_F32 ? 0.0f : 5.0e-4f; + require_close(get_f32_tensor(hrx_backend, hrx_rows), get_f32_tensor(cpu_backend, cpu_rows), tolerance); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), tolerance); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static auto allocate_test_weight(ggml_backend_t backend, ggml_tensor * weight) { + const auto buft = ggml_backend_get_default_buffer_type(backend); + const size_t alignment = ggml_backend_buft_get_alignment(buft); + const size_t size = GGML_PAD(ggml_backend_buft_get_alloc_size(buft, weight), alignment); + auto buffer = std::unique_ptr( + ggml_backend_buft_alloc_buffer(buft, size), ggml_backend_buffer_free); + REQUIRE(buffer != nullptr); + ggml_tallocr allocator = ggml_tallocr_new(buffer.get()); + REQUIRE(ggml_tallocr_alloc(&allocator, weight) == GGML_STATUS_SUCCESS); + ggml_backend_buffer_set_usage(buffer.get(), GGML_BACKEND_BUFFER_USAGE_WEIGHTS); + return buffer; +} + +static void run_dense_matmul_cpu_reference_case(ggml_type weight_type, + const char * expected_kernel, + int64_t token_count, + int64_t output_size, + int64_t input_size = kQwenHiddenSize, + bool normalize_input = false, + bool square_output = false, + ggml_type cache_type = GGML_TYPE_COUNT) { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(32 * 1024 * 1024); + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_weight = ggml_new_tensor_2d(cpu_ctx, weight_type, input_size, output_size); + ggml_tensor * cpu_input = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * hrx_weight = ggml_new_tensor_2d(hrx_ctx, weight_type, input_size, output_size); + ggml_tensor * hrx_input = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * cpu_norm_w = normalize_input ? ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, input_size) : nullptr; + ggml_tensor * hrx_norm_w = normalize_input ? ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, input_size) : nullptr; + ggml_tensor * cpu_rhs = normalize_input ? build_rmsnorm_mul_graph(cpu_ctx, cpu_input, cpu_norm_w) : cpu_input; + ggml_tensor * hrx_rhs = normalize_input ? build_rmsnorm_mul_graph(hrx_ctx, hrx_input, hrx_norm_w) : hrx_input; + ggml_tensor * cpu_output = ggml_mul_mat(cpu_ctx, cpu_weight, cpu_rhs); + ggml_tensor * hrx_output = ggml_mul_mat(hrx_ctx, hrx_weight, hrx_rhs); + if (square_output) { + cpu_output = ggml_sqr(cpu_ctx, cpu_output); + hrx_output = ggml_sqr(hrx_ctx, hrx_output); + } + ggml_tensor * cpu_indices = nullptr; + ggml_tensor * hrx_indices = nullptr; + ggml_tensor * cpu_cache = nullptr; + ggml_tensor * hrx_cache = nullptr; + if (cache_type != GGML_TYPE_COUNT) { + REQUIRE(token_count == 1); + cpu_indices = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_I64, token_count); + hrx_indices = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_I64, token_count); + cpu_cache = ggml_new_tensor_2d(cpu_ctx, cache_type, output_size, 17); + hrx_cache = ggml_new_tensor_2d(hrx_ctx, cache_type, output_size, 17); + cpu_output = ggml_set_rows(cpu_ctx, cpu_cache, cpu_output, cpu_indices); + hrx_output = ggml_set_rows(hrx_ctx, hrx_cache, hrx_output, hrx_indices); + } + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + auto hrx_weight_buffer = allocate_test_weight(hrx_backend, hrx_weight); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + const auto sequence = scheduled_kernel_sequence(hrx_graph); + require_kernel_subsequence(sequence, { expected_kernel }); + const bool packed_q4_input = weight_type == GGML_TYPE_Q4_K; + const bool iq4_xs_weight = weight_type == GGML_TYPE_IQ4_XS; + const bool prefill_tokens = token_count >= 256; + const bool packed_iq4_xs_input = iq4_xs_weight && prefill_tokens; + const bool packed_input = packed_q4_input || packed_iq4_xs_input || square_output; + const bool normalized_packed_input = normalize_input && packed_input; + if (normalized_packed_input) { + require_kernel_subsequence(sequence, { "loom_libs:ggml_rmsnorm_binary_q8_1_x4_publish", expected_kernel }); + REQUIRE(std::find(sequence.begin(), sequence.end(), "qwen3_moe:ggml_quantize_q8_1_x4_f32") == sequence.end()); + } else if (normalize_input) { + REQUIRE(std::find(sequence.begin(), sequence.end(), "loom_libs:ggml_rmsnorm_binary_q8_1_x4_publish") == sequence.end()); + } + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector weight = make_matmul_weight_bytes(weight_type, input_size, output_size, 6); + const bool weight_is_f16 = weight_type == GGML_TYPE_F16; + const bool weight_is_bf16 = weight_type == GGML_TYPE_BF16; + const bool weight_is_f32 = weight_type == GGML_TYPE_F32; + const bool weight_is_dense_float = weight_is_f16 || weight_is_bf16 || weight_is_f32; + std::vector input = weight_is_dense_float ? + std::vector(static_cast(input_size * token_count), 0.00390625f) : + make_pattern_f32(input_size * token_count, 7, 0.01f); + if (cache_type != GGML_TYPE_COUNT) { + for (size_t i = 0; i < input.size(); ++i) { + const int value = i % 32 == 0 ? 127 : static_cast((i * 29) % 255) - 127; + input[i] = static_cast(value) * 0.0078125f; + } + } + set_tensor_pair_bytes(cpu_backend, cpu_weight, hrx_backend, hrx_weight, weight.data(), weight.size()); + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), input.size() * sizeof(float)); + if (normalize_input) { + const std::vector norm_w = make_weight(input_size); + set_tensor_pair_bytes(cpu_backend, cpu_norm_w, hrx_backend, hrx_norm_w, norm_w.data(), + norm_w.size() * sizeof(float)); + } + if (cache_type != GGML_TYPE_COUNT) { + const int64_t index = 7; + set_tensor_pair_bytes(cpu_backend, cpu_indices, hrx_backend, hrx_indices, &index, sizeof(index)); + const std::vector cache(ggml_nbytes(cpu_cache), 0); + set_tensor_pair_bytes(cpu_backend, cpu_cache, hrx_backend, hrx_cache, cache.data(), cache.size()); + } + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + const float abs_tolerance = cache_type != GGML_TYPE_COUNT ? 2.0e-3f : 1.0f; + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), abs_tolerance, + cache_type != GGML_TYPE_COUNT ? 1.0e-4f : 3.0e-2f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_prefill_value_cache_repeated_indices_case() { + static constexpr const char * disable_env = "GGML_HRX_DISABLE_PREFILL_V_CACHE_FUSION"; + const char * old_env = std::getenv(disable_env); + const bool had_env = old_env != nullptr; + const std::string old_value = had_env ? old_env : ""; + REQUIRE(unsetenv(disable_env) == 0); + + constexpr int64_t input_size = 1536; + constexpr int64_t output_size = 256; + constexpr int64_t token_count = 64; + constexpr int64_t vocabulary_count = 128; + constexpr int64_t cache_row_count = 512; + + ggml_backend_t fused_backend = ggml_backend_hrx_init(0); + ggml_backend_t split_backend = ggml_backend_hrx_init(0); + REQUIRE(fused_backend != nullptr); + REQUIRE(split_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(32 * 1024 * 1024); + params.no_alloc = true; + ggml_context * fused_ctx = ggml_init(params); + ggml_context * split_ctx = ggml_init(params); + REQUIRE(fused_ctx != nullptr); + REQUIRE(split_ctx != nullptr); + + auto build_graph = [](ggml_context * ctx) { + ggml_tensor * embedding = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, input_size, vocabulary_count); + ggml_tensor * tokens = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, token_count); + ggml_tensor * norm = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, input_size); + ggml_tensor * weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q6_K, input_size, output_size); + ggml_tensor * indices = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, token_count); + ggml_tensor * cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, output_size, cache_row_count); + ggml_tensor * rows = ggml_get_rows(ctx, embedding, tokens); + ggml_tensor * input = build_rmsnorm_mul_graph(ctx, rows, norm); + ggml_tensor * projected = ggml_mul_mat(ctx, weight, input); + ggml_tensor * output = ggml_set_rows(ctx, cache, projected, indices); + return std::array{ embedding, tokens, norm, weight, indices, cache, output }; + }; + + const auto fused = build_graph(fused_ctx); + const auto split = build_graph(split_ctx); + auto fused_embedding_buffer = allocate_test_weight(fused_backend, fused[0]); + auto fused_projection_buffer = allocate_test_weight(fused_backend, fused[3]); + auto split_embedding_buffer = allocate_test_weight(split_backend, split[0]); + auto split_projection_buffer = allocate_test_weight(split_backend, split[3]); + + ggml_cgraph * fused_graph = ggml_new_graph(fused_ctx); + ggml_cgraph * split_graph = ggml_new_graph(split_ctx); + REQUIRE(fused_graph != nullptr); + REQUIRE(split_graph != nullptr); + ggml_build_forward_expand(fused_graph, fused[6]); + ggml_build_forward_expand(split_graph, split[6]); + require_kernel_subsequence(scheduled_kernel_sequence(fused_graph), + { "loom_libs:ggml_get_rows_rmsnorm_binary_q8_1_x4_f16", + "loom_libs:llm_attention_v_matmul_set_rows_tiled_f32_f32" }); + REQUIRE(setenv(disable_env, "1", 1) == 0); + require_kernel_subsequence(scheduled_kernel_sequence(split_graph), + { "loom_libs:ggml_get_rows_rmsnorm_binary_q8_1_x4_f16", + "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32_aligned", + "loom_libs:ggml_set_rows" }); + + ggml_backend_buffer_t fused_buffer = ggml_backend_alloc_ctx_tensors(fused_ctx, fused_backend); + ggml_backend_buffer_t split_buffer = ggml_backend_alloc_ctx_tensors(split_ctx, split_backend); + REQUIRE(fused_buffer != nullptr); + REQUIRE(split_buffer != nullptr); + + const std::vector embedding = make_matmul_weight_bytes(GGML_TYPE_Q4_K, input_size, vocabulary_count, 5); + const std::vector projection = make_matmul_weight_bytes(GGML_TYPE_Q6_K, input_size, output_size, 11); + const std::vector norm = make_weight(input_size); + const std::vector tokens(static_cast(token_count), 3); + std::vector indices(static_cast(token_count)); + for (int64_t token = 0; token < token_count; ++token) { + indices[static_cast(token)] = (token * 37 + 11) % 31; + } + const std::vector cache(static_cast(output_size * cache_row_count), ggml_fp32_to_fp16(-2.5f)); + + set_tensor_pair_bytes(fused_backend, fused[0], split_backend, split[0], embedding.data(), embedding.size()); + set_tensor_pair_bytes(fused_backend, fused[1], split_backend, split[1], tokens.data(), + tokens.size() * sizeof(int32_t)); + set_tensor_pair_bytes(fused_backend, fused[2], split_backend, split[2], norm.data(), norm.size() * sizeof(float)); + set_tensor_pair_bytes(fused_backend, fused[3], split_backend, split[3], projection.data(), projection.size()); + set_tensor_pair_bytes(fused_backend, fused[4], split_backend, split[4], indices.data(), + indices.size() * sizeof(int64_t)); + set_tensor_pair_bytes(fused_backend, fused[5], split_backend, split[5], cache.data(), + cache.size() * sizeof(ggml_fp16_t)); + + REQUIRE(unsetenv(disable_env) == 0); + REQUIRE(ggml_backend_graph_compute(fused_backend, fused_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(fused_backend); + REQUIRE(setenv(disable_env, "1", 1) == 0); + REQUIRE(ggml_backend_graph_compute(split_backend, split_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(split_backend); + require_close(get_f32_tensor(fused_backend, fused[6]), get_f32_tensor(split_backend, split[6]), 0.0f); + + ggml_backend_buffer_free(fused_buffer); + ggml_backend_buffer_free(split_buffer); + ggml_free(fused_ctx); + ggml_free(split_ctx); + ggml_backend_free(fused_backend); + ggml_backend_free(split_backend); + restore_environment_value(disable_env, had_env, old_value); +} + +static void run_dense_matmul_unary_cpu_reference_case(ggml_type weight_type, + const char * expected_kernel, + int64_t token_count, + int64_t output_size, + int64_t input_size = kQwenHiddenSize) { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(32 * 1024 * 1024); + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_weight = ggml_new_tensor_2d(cpu_ctx, weight_type, input_size, output_size); + ggml_tensor * cpu_input = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * cpu_matmul = ggml_mul_mat(cpu_ctx, cpu_weight, cpu_input); + ggml_tensor * cpu_output = ggml_sqr(cpu_ctx, cpu_matmul); + ggml_tensor * hrx_weight = ggml_new_tensor_2d(hrx_ctx, weight_type, input_size, output_size); + ggml_tensor * hrx_input = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * hrx_matmul = ggml_mul_mat(hrx_ctx, hrx_weight, hrx_input); + ggml_tensor * hrx_output = ggml_sqr(hrx_ctx, hrx_matmul); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { expected_kernel }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector weight(ggml_nbytes(cpu_weight), 0); + const std::vector input = make_pattern_f32(input_size * token_count, 7, 0.01f); + set_tensor_pair_bytes(cpu_backend, cpu_weight, hrx_backend, hrx_weight, weight.data(), weight.size()); + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), input.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 1.0e-1f, 3.0e-2f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_dense_matmul_swiglu_cpu_reference_case(ggml_type gate_weight_type, + ggml_type up_weight_type, + const char * expected_kernel, + int64_t token_count, + int64_t output_size, + ggml_glu_op glu_op = GGML_GLU_OP_SWIGLU, + int64_t input_size = kQwenHiddenSize) { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(48 * 1024 * 1024); + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_gate_weight = ggml_new_tensor_2d(cpu_ctx, gate_weight_type, input_size, output_size); + ggml_tensor * cpu_up_weight = ggml_new_tensor_2d(cpu_ctx, up_weight_type, input_size, output_size); + ggml_tensor * cpu_input = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * cpu_gate = ggml_mul_mat(cpu_ctx, cpu_gate_weight, cpu_input); + ggml_tensor * cpu_up = ggml_mul_mat(cpu_ctx, cpu_up_weight, cpu_input); + ggml_tensor * cpu_output = ggml_glu_split(cpu_ctx, cpu_gate, cpu_up, glu_op); + ggml_tensor * hrx_gate_weight = ggml_new_tensor_2d(hrx_ctx, gate_weight_type, input_size, output_size); + ggml_tensor * hrx_up_weight = ggml_new_tensor_2d(hrx_ctx, up_weight_type, input_size, output_size); + ggml_tensor * hrx_input = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * hrx_gate = ggml_mul_mat(hrx_ctx, hrx_gate_weight, hrx_input); + ggml_tensor * hrx_up = ggml_mul_mat(hrx_ctx, hrx_up_weight, hrx_input); + ggml_tensor * hrx_output = ggml_glu_split(hrx_ctx, hrx_gate, hrx_up, glu_op); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + auto hrx_gate_buffer = allocate_test_weight(hrx_backend, hrx_gate_weight); + auto hrx_up_buffer = allocate_test_weight(hrx_backend, hrx_up_weight); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { expected_kernel }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector gate_weight = make_matmul_weight_bytes(gate_weight_type, input_size, output_size, 6); + const std::vector up_weight = make_matmul_weight_bytes(up_weight_type, input_size, output_size, 11); + const std::vector input(static_cast(input_size * token_count), 0.00390625f); + set_tensor_pair_bytes(cpu_backend, cpu_gate_weight, hrx_backend, hrx_gate_weight, gate_weight.data(), + gate_weight.size()); + set_tensor_pair_bytes(cpu_backend, cpu_up_weight, hrx_backend, hrx_up_weight, up_weight.data(), up_weight.size()); + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), input.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 1.0e-4f, 1.0e-3f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_dense_matmul_packed_glu_cpu_reference_case(ggml_type weight_type, + const char * expected_kernel, + int64_t token_count, + int64_t output_size, + ggml_glu_op glu_op = GGML_GLU_OP_SWIGLU, + bool swapped = false, + int64_t input_size = 256) { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(64 * 1024 * 1024); + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_weight = ggml_new_tensor_2d(cpu_ctx, weight_type, input_size, 2 * output_size); + ggml_tensor * cpu_input = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * cpu_packed = ggml_mul_mat(cpu_ctx, cpu_weight, cpu_input); + ggml_tensor * cpu_output = ggml_glu(cpu_ctx, cpu_packed, glu_op, swapped); + ggml_tensor * hrx_weight = ggml_new_tensor_2d(hrx_ctx, weight_type, input_size, 2 * output_size); + ggml_tensor * hrx_input = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * hrx_packed = ggml_mul_mat(hrx_ctx, hrx_weight, hrx_input); + ggml_tensor * hrx_output = ggml_glu(hrx_ctx, hrx_packed, glu_op, swapped); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { expected_kernel }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector weight = make_matmul_weight_bytes(weight_type, input_size, 2 * output_size, 15); + const std::vector input(static_cast(input_size * token_count), 0.00390625f); + set_tensor_pair_bytes(cpu_backend, cpu_weight, hrx_backend, hrx_weight, weight.data(), weight.size()); + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), input.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 1.0f, 3.0e-2f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_dense_matmul_binary_cpu_reference_case(ggml_type lhs_weight_type, + ggml_type rhs_weight_type, + const char * expected_kernel, + int64_t token_count, + int64_t output_size) { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(48 * 1024 * 1024); + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + constexpr int64_t input_size = kQwenHiddenSize; + ggml_tensor * cpu_lhs_weight = ggml_new_tensor_2d(cpu_ctx, lhs_weight_type, input_size, output_size); + ggml_tensor * cpu_rhs_weight = ggml_new_tensor_2d(cpu_ctx, rhs_weight_type, input_size, output_size); + ggml_tensor * cpu_input = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * cpu_lhs = ggml_mul_mat(cpu_ctx, cpu_lhs_weight, cpu_input); + ggml_tensor * cpu_rhs = ggml_mul_mat(cpu_ctx, cpu_rhs_weight, cpu_input); + ggml_tensor * cpu_output = ggml_sub(cpu_ctx, cpu_lhs, cpu_rhs); + ggml_tensor * hrx_lhs_weight = ggml_new_tensor_2d(hrx_ctx, lhs_weight_type, input_size, output_size); + ggml_tensor * hrx_rhs_weight = ggml_new_tensor_2d(hrx_ctx, rhs_weight_type, input_size, output_size); + ggml_tensor * hrx_input = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * hrx_lhs = ggml_mul_mat(hrx_ctx, hrx_lhs_weight, hrx_input); + ggml_tensor * hrx_rhs = ggml_mul_mat(hrx_ctx, hrx_rhs_weight, hrx_input); + ggml_tensor * hrx_output = ggml_sub(hrx_ctx, hrx_lhs, hrx_rhs); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + auto hrx_lhs_buffer = allocate_test_weight(hrx_backend, hrx_lhs_weight); + auto hrx_rhs_buffer = allocate_test_weight(hrx_backend, hrx_rhs_weight); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { expected_kernel }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector lhs_weight = make_matmul_weight_bytes(lhs_weight_type, input_size, output_size, 6); + const std::vector rhs_weight = make_matmul_weight_bytes(rhs_weight_type, input_size, output_size, 11); + const std::vector input(static_cast(input_size * token_count), 0.00390625f); + set_tensor_pair_bytes(cpu_backend, cpu_lhs_weight, hrx_backend, hrx_lhs_weight, lhs_weight.data(), + lhs_weight.size()); + set_tensor_pair_bytes(cpu_backend, cpu_rhs_weight, hrx_backend, hrx_rhs_weight, rhs_weight.data(), + rhs_weight.size()); + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), input.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 1.0e-4f, 1.0e-3f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_vector_q8_publish_cpu_reference_case() { + constexpr int64_t input_size = 256; + constexpr int64_t intermediate_size = 256; + constexpr int64_t output_size = 128; + constexpr int64_t token_count = 1; + + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(32 * 1024 * 1024); + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_first_weight = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, input_size, intermediate_size); + ggml_tensor * cpu_second_weight = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_Q4_K, intermediate_size, output_size); + ggml_tensor * cpu_input = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * cpu_hidden = ggml_mul_mat(cpu_ctx, cpu_first_weight, cpu_input); + ggml_tensor * cpu_output = ggml_mul_mat(cpu_ctx, cpu_second_weight, cpu_hidden); + ggml_tensor * hrx_first_weight = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, input_size, intermediate_size); + ggml_tensor * hrx_second_weight = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_Q4_K, intermediate_size, output_size); + ggml_tensor * hrx_input = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * hrx_hidden = ggml_mul_mat(hrx_ctx, hrx_first_weight, hrx_input); + ggml_tensor * hrx_output = ggml_mul_mat(hrx_ctx, hrx_second_weight, hrx_hidden); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + auto hrx_first_weight_buffer = allocate_test_weight(hrx_backend, hrx_first_weight); + auto hrx_second_weight_buffer = allocate_test_weight(hrx_backend, hrx_second_weight); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), + { "loom_libs:ggml_mul_mat_vector_f32_f32", "qwen3_moe:ggml_quantize_q8_1_x4_f32", + "loom_libs:ggml_mul_mat_vector_f32_f32" }); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*hrx_graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + ggml::hrx::DispatchScheduleDiagnostics diagnostics; + REQUIRE(scheduler.schedule_graph(imported.graph, { "gfx1151" }, &diagnostics)); + const ggml::hrx::CommandPlan & plan = scheduler.plan(); + REQUIRE(plan.valid()); + + const ggml::hrx::Value * hidden_value = imported.graph.values().find_tensor(hrx_hidden); + REQUIRE(hidden_value != nullptr); + const size_t q8_byte_count = static_cast(token_count) * ggml_row_size(GGML_TYPE_Q8_1, intermediate_size); + const ggml::hrx::CommandPlanAlternateValue * q8_alternate = + plan.metadata.find_alternate_value(hidden_value->id, GGML_TYPE_Q8_1, q8_byte_count); + REQUIRE(q8_alternate != nullptr); + + const ggml::hrx::Dispatch * second_vector = nullptr; + int vector_count = 0; + for (const ggml::hrx::Dispatch & dispatch : plan.dispatches) { + if (kernel_name_for_id(dispatch.kernel.kernel_id) == "loom_libs:ggml_mul_mat_vector_f32_f32") { + ++vector_count; + if (vector_count == 2) { + second_vector = &dispatch; + break; + } + } + } + REQUIRE(second_vector != nullptr); + REQUIRE(second_vector->bindings.size() == 4); + REQUIRE(second_vector->bindings[0].value == q8_alternate->alternate_value); + REQUIRE(second_vector->kernel.compile_parameters.at("ggml.matmul.vector.activation_format") == + std::to_string(GGML_TYPE_Q8_1)); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector first_weight = make_matmul_weight_bytes(GGML_TYPE_F32, input_size, intermediate_size, 6); + const std::vector second_weight = + make_matmul_weight_bytes(GGML_TYPE_Q4_K, intermediate_size, output_size, 11); + const std::vector input = make_pattern_f32(input_size * token_count, 7, 0.01f); + set_tensor_pair_bytes(cpu_backend, cpu_first_weight, hrx_backend, hrx_first_weight, first_weight.data(), + first_weight.size()); + set_tensor_pair_bytes(cpu_backend, cpu_second_weight, hrx_backend, hrx_second_weight, second_weight.data(), + second_weight.size()); + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), input.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 1.0f, 5.0e-2f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_dense_matmul_postops_cpu_reference_case(ggml_type weight_type, + const char * expected_kernel, + int64_t token_count, + int64_t output_size, + bool include_bias, + bool include_residual, + bool include_next_rmsnorm, + int64_t input_size = kQwenHiddenSize) { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(96 * 1024 * 1024); + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_weight = ggml_new_tensor_2d(cpu_ctx, weight_type, input_size, output_size); + ggml_tensor * cpu_input = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * cpu_bias = include_bias ? ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, output_size) : nullptr; + ggml_tensor * cpu_residual = + include_residual ? ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, output_size, token_count) : nullptr; + ggml_tensor * cpu_norm_weight = + include_next_rmsnorm ? ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, output_size) : nullptr; + ggml_tensor * cpu_projection = ggml_mul_mat(cpu_ctx, cpu_weight, cpu_input); + ggml_tensor * cpu_postops = cpu_projection; + if (include_bias) { + cpu_postops = ggml_add(cpu_ctx, cpu_postops, cpu_bias); + } + if (include_residual) { + cpu_postops = ggml_add(cpu_ctx, cpu_postops, cpu_residual); + } + ggml_tensor * cpu_output = cpu_postops; + if (include_next_rmsnorm) { + ggml_tensor * cpu_rms = ggml_rms_norm(cpu_ctx, cpu_postops, kQwenRmsNormEps); + ggml_tensor * cpu_normalized = ggml_mul(cpu_ctx, cpu_rms, cpu_norm_weight); + cpu_output = ggml_add(cpu_ctx, cpu_normalized, cpu_postops); + } + + ggml_tensor * hrx_weight = ggml_new_tensor_2d(hrx_ctx, weight_type, input_size, output_size); + ggml_tensor * hrx_input = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * hrx_bias = include_bias ? ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, output_size) : nullptr; + ggml_tensor * hrx_residual = + include_residual ? ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, output_size, token_count) : nullptr; + ggml_tensor * hrx_norm_weight = + include_next_rmsnorm ? ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, output_size) : nullptr; + ggml_tensor * hrx_projection = ggml_mul_mat(hrx_ctx, hrx_weight, hrx_input); + ggml_tensor * hrx_postops = hrx_projection; + if (include_bias) { + hrx_postops = ggml_add(hrx_ctx, hrx_postops, hrx_bias); + } + if (include_residual) { + hrx_postops = ggml_add(hrx_ctx, hrx_postops, hrx_residual); + } + ggml_tensor * hrx_output = hrx_postops; + if (include_next_rmsnorm) { + ggml_tensor * hrx_rms = ggml_rms_norm(hrx_ctx, hrx_postops, kQwenRmsNormEps); + ggml_tensor * hrx_normalized = ggml_mul(hrx_ctx, hrx_rms, hrx_norm_weight); + hrx_output = ggml_add(hrx_ctx, hrx_normalized, hrx_postops); + } + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + if (include_next_rmsnorm) { + require_kernel_subsequence( + scheduled_kernel_sequence(hrx_graph), + { expected_kernel, "loom_libs:ggml_rmsnorm_binary_f32", "loom_libs:ggml_binary_f32" }); + } else { + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { expected_kernel }); + } + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector weight(ggml_nbytes(cpu_weight), 0); + const std::vector input(static_cast(input_size * token_count), 0.00390625f); + const std::vector bias = make_pattern_f32(static_cast(output_size), 13, 0.02f); + const std::vector residual = make_pattern_f32(static_cast(output_size * token_count), 23, 0.02f); + const std::vector norm_weight = make_weight(output_size); + set_tensor_pair_bytes(cpu_backend, cpu_weight, hrx_backend, hrx_weight, weight.data(), weight.size()); + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), input.size() * sizeof(float)); + if (include_bias) { + set_tensor_pair_bytes(cpu_backend, cpu_bias, hrx_backend, hrx_bias, bias.data(), bias.size() * sizeof(float)); + } + if (include_residual) { + set_tensor_pair_bytes(cpu_backend, cpu_residual, hrx_backend, hrx_residual, residual.data(), + residual.size() * sizeof(float)); + } + if (include_next_rmsnorm) { + set_tensor_pair_bytes(cpu_backend, cpu_norm_weight, hrx_backend, hrx_norm_weight, norm_weight.data(), + norm_weight.size() * sizeof(float)); + } + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + if (!include_next_rmsnorm) { + require_close(get_f32_tensor(hrx_backend, hrx_postops), get_f32_tensor(cpu_backend, cpu_postops), 1.0f, + 3.0e-2f); + } + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 1.0f, 5.0e-2f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_dense_matmul_zero_weight_case(ggml_type weight_type, + const char * expected_kernel, + int64_t token_count = 2) { + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = static_cast(32 * 1024 * 1024); + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + constexpr int64_t input_size = kQwenHiddenSize; + constexpr int64_t output_size = 128; + ggml_tensor * weight = ggml_new_tensor_2d(ctx, weight_type, input_size, output_size); + ggml_tensor * input = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, input_size, token_count); + ggml_tensor * output = ggml_mul_mat(ctx, weight, input); + REQUIRE(output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + require_kernel_subsequence(scheduled_kernel_sequence(graph), { expected_kernel }); + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, hrx_backend); + REQUIRE(buffer != nullptr); + + const std::vector weight_bytes(ggml_nbytes(weight), 0); + const std::vector input_data = make_pattern_f32(input_size * token_count, 7, 0.01f); + const std::vector expected(static_cast(output_size * token_count), 0.0f); + set_tensor_bytes(hrx_backend, weight, weight_bytes.data(), weight_bytes.size()); + set_tensor_bytes(hrx_backend, input, input_data.data(), input_data.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(hrx_backend, graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, output), expected, 0.0f, 0.0f); + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(hrx_backend); +} + +static void run_endpoint_rmsnorm_q6k_q8_cpu_reference_case() { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 32 * 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + constexpr int64_t token_count = 1; + ggml_tensor * cpu_input = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, kQwenHiddenSize, token_count); + ggml_tensor * cpu_norm_w = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, kQwenHiddenSize); + ggml_tensor * cpu_weight = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_Q6_K, kQwenHiddenSize, kQwenVocabularyCount); + ggml_tensor * cpu_norm = build_rmsnorm_mul_graph(cpu_ctx, cpu_input, cpu_norm_w); + ggml_tensor * cpu_output = ggml_mul_mat(cpu_ctx, cpu_weight, cpu_norm); + + ggml_tensor * hrx_input = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, kQwenHiddenSize, token_count); + ggml_tensor * hrx_norm_w = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, kQwenHiddenSize); + ggml_tensor * hrx_weight = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_Q6_K, kQwenHiddenSize, kQwenVocabularyCount); + ggml_tensor * hrx_norm = build_rmsnorm_mul_graph(hrx_ctx, hrx_input, hrx_norm_w); + ggml_tensor * hrx_output = ggml_mul_mat(hrx_ctx, hrx_weight, hrx_norm); + REQUIRE(cpu_output != nullptr); + REQUIRE(hrx_output != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_output); + ggml_build_forward_expand(hrx_graph, hrx_output); + + require_kernel_subsequence( + scheduled_kernel_sequence(hrx_graph), + { "qwen3_moe:qwen3_moe_rmsnorm_f32_quantize_q8_1_x4", "qwen3_moe:ggml_linear_q6k_q8_1_x4" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector input = make_pattern_f32(kQwenHiddenSize * token_count, 8, 0.01f); + const std::vector norm_w = make_weight(kQwenHiddenSize); + const std::vector weight = make_quantized_rows(GGML_TYPE_Q6_K, kQwenHiddenSize, kQwenVocabularyCount, 9); + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), input.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_norm_w, hrx_backend, hrx_norm_w, norm_w.data(), + norm_w.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_weight, hrx_backend, hrx_weight, weight.data(), weight.size()); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_output), get_f32_tensor(cpu_backend, cpu_output), 6.0e-1f, 3.0e-2f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_lfm_head_rope_cpu_reference_case(int64_t head_count, int64_t token_count, bool cache_output) { + constexpr int64_t head_size = 64; + constexpr float epsilon = 1.0e-5f; + const int64_t row_size = head_size * head_count; + const int64_t cache_rows = token_count + 64; + + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + struct HeadGraph { + ggml_tensor * input; + ggml_tensor * weight; + ggml_tensor * positions; + ggml_tensor * cache; + ggml_tensor * row_ids; + ggml_tensor * output; + }; + + auto build = [&](ggml_context * ctx) { + HeadGraph g = {}; + g.input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, head_size, head_count, token_count); + g.weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, head_size); + g.positions = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, token_count); + ggml_tensor * normalized = ggml_rms_norm(ctx, g.input, epsilon); + ggml_tensor * scaled = ggml_mul(ctx, normalized, g.weight); + ggml_tensor * rope = ggml_rope_ext(ctx, scaled, g.positions, nullptr, head_size, GGML_ROPE_TYPE_NEOX, 0, + 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + g.output = rope; + if (cache_output) { + g.cache = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, row_size, cache_rows); + g.row_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, token_count); + ggml_tensor * rows = ggml_reshape_2d(ctx, rope, row_size, token_count); + ggml_tensor * view = ggml_view_2d(ctx, rows, row_size, token_count, rows->nb[1], 0); + g.output = ggml_set_rows(ctx, g.cache, view, g.row_ids); + } + REQUIRE(g.output != nullptr); + return g; + }; + + HeadGraph cpu = build(cpu_ctx); + HeadGraph hrx = build(hrx_ctx); + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu.output); + ggml_build_forward_expand(hrx_graph, hrx.output); + const std::vector kernels = scheduled_kernel_sequence(hrx_graph); + if (cache_output) { + REQUIRE(std::find(kernels.begin(), kernels.end(), "loom_libs:ggml_rmsnorm_mul_rope_f32") == kernels.end()); + require_kernel_subsequence(kernels, { "loom_libs:ggml_rope_f32", "loom_libs:ggml_set_rows" }); + } else { + require_kernel_subsequence(kernels, { "loom_libs:ggml_rmsnorm_mul_rope_f32" }); + } + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector input = make_pattern_f32(static_cast(row_size * token_count), 47, 0.03f); + const std::vector weight = make_weight(head_size); + std::vector positions(static_cast(token_count)); + std::vector row_ids(static_cast(token_count)); + for (int64_t i = 0; i < token_count; ++i) { + positions[static_cast(i)] = static_cast(3 + 2 * i); + row_ids[static_cast(i)] = (17 * i + 5) % cache_rows; + } + set_tensor_pair_bytes(cpu_backend, cpu.input, hrx_backend, hrx.input, input.data(), input.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu.weight, hrx_backend, hrx.weight, weight.data(), + weight.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu.positions, hrx_backend, hrx.positions, positions.data(), + positions.size() * sizeof(int32_t)); + if (cache_output) { + const std::vector cache(static_cast(row_size * cache_rows), ggml_fp32_to_fp16(-3.0f)); + set_tensor_pair_bytes(cpu_backend, cpu.cache, hrx_backend, hrx.cache, cache.data(), + cache.size() * sizeof(ggml_fp16_t)); + set_tensor_pair_bytes(cpu_backend, cpu.row_ids, hrx_backend, hrx.row_ids, row_ids.data(), + row_ids.size() * sizeof(int64_t)); + } + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + const float tolerance = cache_output ? 1.0e-3f : 1.0e-4f; + require_close(get_f32_tensor(hrx_backend, hrx.output), get_f32_tensor(cpu_backend, cpu.output), tolerance, + tolerance); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_lfm_head_rope_availability_checks() { + ggml_init_params params = {}; + params.mem_size = 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 64, 8, 1); + ggml_tensor * weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 64); + ggml_tensor * pos_f32 = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 1); + ggml_tensor * positions = ggml_cast(ctx, pos_f32, GGML_TYPE_I32); + ggml_tensor * rms = ggml_rms_norm(ctx, input, 1.0e-5f); + ggml_tensor * scaled = ggml_mul(ctx, rms, weight); + ggml_tensor * output = ggml_rope_ext(ctx, scaled, positions, nullptr, 64, GGML_ROPE_TYPE_NEOX, 0, 10000.0f, 1.0f, + 0.0f, 1.0f, 0.0f, 0.0f); + REQUIRE(output != nullptr); + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, output); + + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + const ggml::hrx::DispatchRegistry * registry = ggml::hrx::find_dispatch_registry({ "gfx1151" }); + REQUIRE(registry != nullptr); + const size_t rms_index = producer_index_for_tensor(imported.graph, rms); + const size_t pos_index = producer_index_for_tensor(imported.graph, positions); + ggml::hrx::CommandPlan plan; + std::vector covered(imported.graph.nodes().size(), false); + auto matched_kernel = [&]() { + ggml::hrx::DispatchMatch match; + const ggml::hrx::DispatchMatchContext context = { + imported.graph, &imported.graph.nodes()[rms_index], rms_index, covered, + plan, next_plan_value(imported.graph, plan), + }; + REQUIRE(registry->match(context, match)); + REQUIRE(match.dispatches.size() == 1); + return kernel_name_for_id(match.dispatches[0].kernel.kernel_id); + }; + REQUIRE(matched_kernel() == "loom_libs:ggml_rmsnorm_f32"); + covered[pos_index] = true; + REQUIRE(matched_kernel() == "loom_libs:ggml_rmsnorm_mul_rope_f32"); + + ggml_free(ctx); +} + +static void run_rope_set_rows_cpu_reference_case() { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + constexpr int64_t head_size = 8; + constexpr int64_t head_count = 2; + constexpr int64_t token_count = 3; + constexpr int64_t cache_row_count = 8; + constexpr int64_t hidden_size = head_size * head_count; + + { + ggml_init_params params = {}; + params.mem_size = 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_input = ggml_new_tensor_3d(cpu_ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * cpu_pos = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_I32, token_count); + ggml_tensor * cpu_freq = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, head_size / 2); + ggml_tensor * cpu_out = ggml_rope_ext(cpu_ctx, cpu_input, cpu_pos, cpu_freq, head_size, GGML_ROPE_TYPE_NEOX, 0, + 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + ggml_tensor * hrx_input = ggml_new_tensor_3d(hrx_ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * hrx_pos = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_I32, token_count); + ggml_tensor * hrx_freq = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, head_size / 2); + ggml_tensor * hrx_out = ggml_rope_ext(hrx_ctx, hrx_input, hrx_pos, hrx_freq, head_size, GGML_ROPE_TYPE_NEOX, 0, + 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + REQUIRE(cpu_out != nullptr); + REQUIRE(hrx_out != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_out); + ggml_build_forward_expand(hrx_graph, hrx_out); + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_rope_f32" }); + const ggml::hrx::KernelSpecialization specialization = + scheduled_kernel_specialization(hrx_graph, "loom_libs:ggml_rope_f32"); + REQUIRE(specialization.workload_specialization == ggml::hrx::WorkloadSpecialization::Dynamic); + REQUIRE(specialization.compile_parameters.at("ggml.rope_f32.token_capacity") == "4"); + REQUIRE(specialization.compile_parameters.count("ggml.rope_f32.input_span") == 0); + REQUIRE(specialization.integer_parameters.at("input_span") == hidden_size * token_count); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector input = make_pattern_f32(static_cast(hidden_size * token_count), 41, 0.03f); + const std::vector pos = { 0, 1, 7 }; + const std::vector freq = { 1.0f, 0.5f, 0.25f, 0.125f }; + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), + input.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_pos, hrx_backend, hrx_pos, pos.data(), pos.size() * sizeof(int32_t)); + set_tensor_pair_bytes(cpu_backend, cpu_freq, hrx_backend, hrx_freq, freq.data(), freq.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_out), get_f32_tensor(cpu_backend, cpu_out), 1.0e-4f, 1.0e-4f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + } + + { + ggml_init_params params = {}; + params.mem_size = 8 * 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + constexpr int64_t partial_head_size = 128; + constexpr int64_t partial_n_dims = 96; + constexpr int64_t partial_head_count = 3; + constexpr int64_t partial_tokens = 4; + ggml_tensor * cpu_input = + ggml_new_tensor_3d(cpu_ctx, GGML_TYPE_F32, partial_head_size, partial_head_count, partial_tokens); + ggml_tensor * cpu_pos = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_I32, partial_tokens); + ggml_tensor * cpu_freq = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, partial_n_dims / 2); + ggml_tensor * cpu_out = ggml_rope_ext(cpu_ctx, cpu_input, cpu_pos, cpu_freq, partial_n_dims, + GGML_ROPE_TYPE_NEOX, 0, 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + ggml_tensor * hrx_input = + ggml_new_tensor_3d(hrx_ctx, GGML_TYPE_F32, partial_head_size, partial_head_count, partial_tokens); + ggml_tensor * hrx_pos = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_I32, partial_tokens); + ggml_tensor * hrx_freq = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, partial_n_dims / 2); + ggml_tensor * hrx_out = ggml_rope_ext(hrx_ctx, hrx_input, hrx_pos, hrx_freq, partial_n_dims, + GGML_ROPE_TYPE_NEOX, 0, 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + REQUIRE(cpu_out != nullptr); + REQUIRE(hrx_out != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_out); + ggml_build_forward_expand(hrx_graph, hrx_out); + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_rope_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector input = + make_pattern_f32(static_cast(partial_head_size * partial_head_count * partial_tokens), 47, 0.01f); + const std::vector pos = { 0, 1, 13, 29 }; + std::vector freq(static_cast(partial_n_dims / 2)); + for (size_t i = 0; i < freq.size(); ++i) { + freq[i] = 1.0f + 0.01f * static_cast(i); + } + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), + input.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_pos, hrx_backend, hrx_pos, pos.data(), pos.size() * sizeof(int32_t)); + set_tensor_pair_bytes(cpu_backend, cpu_freq, hrx_backend, hrx_freq, freq.data(), freq.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_out), get_f32_tensor(cpu_backend, cpu_out), 1.0e-4f, 1.0e-4f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + } + + auto run_view_set_rows_case = [&](int64_t view_token_count) { + ggml_init_params params = {}; + params.mem_size = 16 * 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + constexpr int64_t phi_hidden_size = 1024; + constexpr int64_t phi_cache_row_count = 256; + constexpr int64_t phi_row_stride = 5120; + constexpr size_t row_offset = 8 * sizeof(float); + const size_t storage_elements = + row_offset / sizeof(float) + static_cast((view_token_count - 1) * phi_row_stride + phi_hidden_size); + + ggml_tensor * cpu_cache = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F16, phi_hidden_size, phi_cache_row_count); + ggml_tensor * cpu_store = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, storage_elements); + ggml_tensor * cpu_rows = ggml_view_2d(cpu_ctx, cpu_store, phi_hidden_size, view_token_count, + phi_row_stride * sizeof(float), row_offset); + ggml_tensor * cpu_ids = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_I64, view_token_count); + ggml_tensor * cpu_out = ggml_set_rows(cpu_ctx, cpu_cache, cpu_rows, cpu_ids); + ggml_tensor * hrx_cache = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F16, phi_hidden_size, phi_cache_row_count); + ggml_tensor * hrx_store = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, storage_elements); + ggml_tensor * hrx_rows = ggml_view_2d(hrx_ctx, hrx_store, phi_hidden_size, view_token_count, + phi_row_stride * sizeof(float), row_offset); + ggml_tensor * hrx_ids = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_I64, view_token_count); + ggml_tensor * hrx_out = ggml_set_rows(hrx_ctx, hrx_cache, hrx_rows, hrx_ids); + REQUIRE(cpu_out != nullptr); + REQUIRE(hrx_out != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_out); + ggml_build_forward_expand(hrx_graph, hrx_out); + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_set_rows" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector cache(static_cast(phi_hidden_size * phi_cache_row_count), + ggml_fp32_to_fp16(-2.5f)); + std::vector storage(storage_elements, -17.0f); + const std::vector rows = + make_pattern_f32(static_cast(phi_hidden_size * view_token_count), 52, 0.004f); + for (int64_t token = 0; token < view_token_count; ++token) { + std::memcpy(storage.data() + row_offset / sizeof(float) + static_cast(token * phi_row_stride), + rows.data() + static_cast(token * phi_hidden_size), + static_cast(phi_hidden_size) * sizeof(float)); + } + std::vector ids(static_cast(view_token_count)); + for (int64_t i = 0; i < view_token_count; ++i) { + ids[static_cast(i)] = (i * 17 + 3) % phi_cache_row_count; + } + + set_tensor_pair_bytes(cpu_backend, cpu_cache, hrx_backend, hrx_cache, cache.data(), + cache.size() * sizeof(ggml_fp16_t)); + set_tensor_pair_bytes(cpu_backend, cpu_store, hrx_backend, hrx_store, storage.data(), + storage.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_ids, hrx_backend, hrx_ids, ids.data(), ids.size() * sizeof(int64_t)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_out), get_f32_tensor(cpu_backend, cpu_out), 1.0e-3f, 1.0e-3f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + }; + + run_view_set_rows_case(2); + run_view_set_rows_case(14); + + { + ggml_init_params params = {}; + params.mem_size = 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_input = ggml_new_tensor_3d(cpu_ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * cpu_pos = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_I32, token_count); + ggml_tensor * cpu_freq = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, head_size / 2); + ggml_tensor * cpu_out = ggml_rope_ext(cpu_ctx, cpu_input, cpu_pos, cpu_freq, head_size, GGML_ROPE_TYPE_NORMAL, + 0, 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + ggml_tensor * hrx_input = ggml_new_tensor_3d(hrx_ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * hrx_pos = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_I32, token_count); + ggml_tensor * hrx_freq = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, head_size / 2); + ggml_tensor * hrx_out = ggml_rope_ext(hrx_ctx, hrx_input, hrx_pos, hrx_freq, head_size, GGML_ROPE_TYPE_NORMAL, + 0, 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + REQUIRE(cpu_out != nullptr); + REQUIRE(hrx_out != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_out); + ggml_build_forward_expand(hrx_graph, hrx_out); + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_rope_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector input = make_pattern_f32(static_cast(hidden_size * token_count), 45, 0.03f); + const std::vector pos = { 0, 1, 7 }; + const std::vector freq = { 1.0f, 0.5f, 0.25f, 0.125f }; + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), + input.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_pos, hrx_backend, hrx_pos, pos.data(), pos.size() * sizeof(int32_t)); + set_tensor_pair_bytes(cpu_backend, cpu_freq, hrx_backend, hrx_freq, freq.data(), freq.size() * sizeof(float)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_out), get_f32_tensor(cpu_backend, cpu_out), 1.0e-4f, 1.0e-4f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + } + + { + ggml_init_params params = {}; + params.mem_size = 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_cache = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F16, hidden_size, cache_row_count); + ggml_tensor * cpu_rows = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F16, hidden_size, token_count); + ggml_tensor * cpu_ids = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_I64, token_count); + ggml_tensor * cpu_out = ggml_set_rows(cpu_ctx, cpu_cache, cpu_rows, cpu_ids); + ggml_tensor * hrx_cache = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F16, hidden_size, cache_row_count); + ggml_tensor * hrx_rows = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F16, hidden_size, token_count); + ggml_tensor * hrx_ids = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_I64, token_count); + ggml_tensor * hrx_out = ggml_set_rows(hrx_ctx, hrx_cache, hrx_rows, hrx_ids); + REQUIRE(cpu_out != nullptr); + REQUIRE(hrx_out != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_out); + ggml_build_forward_expand(hrx_graph, hrx_out); + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_set_rows" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector cache(static_cast(hidden_size * cache_row_count), + ggml_fp32_to_fp16(-4.0f)); + const std::vector rows = + make_pattern_f16(static_cast(hidden_size * token_count), 44, 0.04f); + const std::vector ids = { 7, 0, 5 }; + set_tensor_pair_bytes(cpu_backend, cpu_cache, hrx_backend, hrx_cache, cache.data(), + cache.size() * sizeof(ggml_fp16_t)); + set_tensor_pair_bytes(cpu_backend, cpu_rows, hrx_backend, hrx_rows, rows.data(), + rows.size() * sizeof(ggml_fp16_t)); + set_tensor_pair_bytes(cpu_backend, cpu_ids, hrx_backend, hrx_ids, ids.data(), ids.size() * sizeof(int64_t)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_out), get_f32_tensor(cpu_backend, cpu_out), 1.0e-3f, 1.0e-3f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + } + + { + ggml_init_params params = {}; + params.mem_size = 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_cache = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F16, hidden_size, cache_row_count); + ggml_tensor * cpu_rows = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F32, hidden_size, token_count); + ggml_tensor * cpu_ids = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_I64, token_count); + ggml_tensor * cpu_out = ggml_set_rows(cpu_ctx, cpu_cache, cpu_rows, cpu_ids); + ggml_tensor * hrx_cache = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F16, hidden_size, cache_row_count); + ggml_tensor * hrx_rows = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F32, hidden_size, token_count); + ggml_tensor * hrx_ids = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_I64, token_count); + ggml_tensor * hrx_out = ggml_set_rows(hrx_ctx, hrx_cache, hrx_rows, hrx_ids); + REQUIRE(cpu_out != nullptr); + REQUIRE(hrx_out != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_out); + ggml_build_forward_expand(hrx_graph, hrx_out); + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_set_rows" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector cache(static_cast(hidden_size * cache_row_count), + ggml_fp32_to_fp16(-2.0f)); + const std::vector rows = make_pattern_f32(static_cast(hidden_size * token_count), 42, 0.04f); + const std::vector ids = { 6, 2, 4 }; + set_tensor_pair_bytes(cpu_backend, cpu_cache, hrx_backend, hrx_cache, cache.data(), + cache.size() * sizeof(ggml_fp16_t)); + set_tensor_pair_bytes(cpu_backend, cpu_rows, hrx_backend, hrx_rows, rows.data(), rows.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_ids, hrx_backend, hrx_ids, ids.data(), ids.size() * sizeof(int64_t)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_out), get_f32_tensor(cpu_backend, cpu_out), 1.0e-3f, 1.0e-3f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + } + + { + ggml_init_params params = {}; + params.mem_size = 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_input = ggml_new_tensor_3d(cpu_ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * cpu_pos = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_I32, token_count); + ggml_tensor * cpu_freq = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, head_size / 2); + ggml_tensor * cpu_rope = ggml_rope_ext(cpu_ctx, cpu_input, cpu_pos, cpu_freq, head_size, GGML_ROPE_TYPE_NEOX, 0, + 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + ggml_tensor * cpu_rows = ggml_reshape_2d(cpu_ctx, cpu_rope, hidden_size, token_count); + ggml_tensor * cpu_cache = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F16, hidden_size, cache_row_count); + ggml_tensor * cpu_ids = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_I64, token_count); + ggml_tensor * cpu_out = ggml_set_rows(cpu_ctx, cpu_cache, cpu_rows, cpu_ids); + ggml_tensor * hrx_input = ggml_new_tensor_3d(hrx_ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * hrx_pos = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_I32, token_count); + ggml_tensor * hrx_freq = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, head_size / 2); + ggml_tensor * hrx_rope = ggml_rope_ext(hrx_ctx, hrx_input, hrx_pos, hrx_freq, head_size, GGML_ROPE_TYPE_NEOX, 0, + 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + ggml_tensor * hrx_rows = ggml_reshape_2d(hrx_ctx, hrx_rope, hidden_size, token_count); + ggml_tensor * hrx_cache = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F16, hidden_size, cache_row_count); + ggml_tensor * hrx_ids = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_I64, token_count); + ggml_tensor * hrx_out = ggml_set_rows(hrx_ctx, hrx_cache, hrx_rows, hrx_ids); + REQUIRE(cpu_out != nullptr); + REQUIRE(hrx_out != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_out); + ggml_build_forward_expand(hrx_graph, hrx_out); + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_rope_set_rows_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector input = make_pattern_f32(static_cast(hidden_size * token_count), 43, 0.03f); + const std::vector pos = { 0, 1, 7 }; + const std::vector freq = { 1.0f, 0.5f, 0.25f, 0.125f }; + const std::vector cache(static_cast(hidden_size * cache_row_count), + ggml_fp32_to_fp16(-3.0f)); + const std::vector ids = { 5, 1, 3 }; + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), + input.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_pos, hrx_backend, hrx_pos, pos.data(), pos.size() * sizeof(int32_t)); + set_tensor_pair_bytes(cpu_backend, cpu_freq, hrx_backend, hrx_freq, freq.data(), freq.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_cache, hrx_backend, hrx_cache, cache.data(), + cache.size() * sizeof(ggml_fp16_t)); + set_tensor_pair_bytes(cpu_backend, cpu_ids, hrx_backend, hrx_ids, ids.data(), ids.size() * sizeof(int64_t)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_out), get_f32_tensor(cpu_backend, cpu_out), 1.0e-3f, 1.0e-3f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + } + + { + ggml_init_params params = {}; + params.mem_size = 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + ggml_tensor * cpu_input = ggml_new_tensor_3d(cpu_ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * cpu_pos = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_I32, token_count); + ggml_tensor * cpu_freq = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_F32, head_size / 2); + ggml_tensor * cpu_rope = ggml_rope_ext(cpu_ctx, cpu_input, cpu_pos, cpu_freq, head_size, GGML_ROPE_TYPE_NORMAL, + 0, 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + ggml_tensor * cpu_rows = ggml_reshape_2d(cpu_ctx, cpu_rope, hidden_size, token_count); + ggml_tensor * cpu_cache = ggml_new_tensor_2d(cpu_ctx, GGML_TYPE_F16, hidden_size, cache_row_count); + ggml_tensor * cpu_ids = ggml_new_tensor_1d(cpu_ctx, GGML_TYPE_I64, token_count); + ggml_tensor * cpu_out = ggml_set_rows(cpu_ctx, cpu_cache, cpu_rows, cpu_ids); + ggml_tensor * hrx_input = ggml_new_tensor_3d(hrx_ctx, GGML_TYPE_F32, head_size, head_count, token_count); + ggml_tensor * hrx_pos = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_I32, token_count); + ggml_tensor * hrx_freq = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_F32, head_size / 2); + ggml_tensor * hrx_rope = ggml_rope_ext(hrx_ctx, hrx_input, hrx_pos, hrx_freq, head_size, GGML_ROPE_TYPE_NORMAL, + 0, 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + ggml_tensor * hrx_rows = ggml_reshape_2d(hrx_ctx, hrx_rope, hidden_size, token_count); + ggml_tensor * hrx_cache = ggml_new_tensor_2d(hrx_ctx, GGML_TYPE_F16, hidden_size, cache_row_count); + ggml_tensor * hrx_ids = ggml_new_tensor_1d(hrx_ctx, GGML_TYPE_I64, token_count); + ggml_tensor * hrx_out = ggml_set_rows(hrx_ctx, hrx_cache, hrx_rows, hrx_ids); + REQUIRE(cpu_out != nullptr); + REQUIRE(hrx_out != nullptr); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu_out); + ggml_build_forward_expand(hrx_graph, hrx_out); + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), { "loom_libs:ggml_rope_set_rows_f32" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector input = make_pattern_f32(static_cast(hidden_size * token_count), 46, 0.03f); + const std::vector pos = { 0, 1, 7 }; + const std::vector freq = { 1.0f, 0.5f, 0.25f, 0.125f }; + const std::vector cache(static_cast(hidden_size * cache_row_count), + ggml_fp32_to_fp16(-5.0f)); + const std::vector ids = { 4, 2, 6 }; + set_tensor_pair_bytes(cpu_backend, cpu_input, hrx_backend, hrx_input, input.data(), + input.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_pos, hrx_backend, hrx_pos, pos.data(), pos.size() * sizeof(int32_t)); + set_tensor_pair_bytes(cpu_backend, cpu_freq, hrx_backend, hrx_freq, freq.data(), freq.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu_cache, hrx_backend, hrx_cache, cache.data(), + cache.size() * sizeof(ggml_fp16_t)); + set_tensor_pair_bytes(cpu_backend, cpu_ids, hrx_backend, hrx_ids, ids.data(), ids.size() * sizeof(int64_t)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx_out), get_f32_tensor(cpu_backend, cpu_out), 1.0e-3f, 1.0e-3f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + } + + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_attention_postprocess_cpu_reference_case() { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 32 * 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + constexpr int64_t token_count = 2; + constexpr int64_t query_head_count = 4; + constexpr int64_t key_value_head_count = 2; + constexpr int64_t cache_row_count = 8; + const int64_t query_size = query_head_count * kQwenFlashHeadSize; + const int64_t key_value_size = key_value_head_count * kQwenFlashHeadSize; + AttentionPostprocessGraph cpu = build_attention_postprocess_graph(cpu_ctx, token_count, query_head_count, + key_value_head_count, cache_row_count); + AttentionPostprocessGraph hrx = build_attention_postprocess_graph(hrx_ctx, token_count, query_head_count, + key_value_head_count, cache_row_count); + + auto hrx_query_buffer = allocate_test_weight(hrx_backend, hrx.query_weight); + auto hrx_key_buffer = allocate_test_weight(hrx_backend, hrx.key_weight); + auto hrx_value_buffer = allocate_test_weight(hrx_backend, hrx.value_weight); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu.query_output); + ggml_build_forward_expand(cpu_graph, cpu.key_output); + ggml_build_forward_expand(cpu_graph, cpu.value_output); + ggml_build_forward_expand(hrx_graph, hrx.query_output); + ggml_build_forward_expand(hrx_graph, hrx.key_output); + ggml_build_forward_expand(hrx_graph, hrx.value_output); + + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), + { "qwen3_moe:qwen3_moe_attention_postprocess_f32_f16" }); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector input = make_pattern_f32(kQwenHiddenSize * token_count, 10, 0.01f); + const std::vector query_w = make_quantized_rows(GGML_TYPE_Q4_K, kQwenHiddenSize, query_size, 11); + const std::vector key_w = make_quantized_rows(GGML_TYPE_Q4_K, kQwenHiddenSize, key_value_size, 12); + const std::vector value_w = make_quantized_rows(GGML_TYPE_Q6_K, kQwenHiddenSize, key_value_size, 13); + const std::vector query_nw = make_weight(kQwenFlashHeadSize); + const std::vector key_nw = make_pattern_f32(kQwenFlashHeadSize, 14, 0.02f); + const std::vector positions = make_i32_mod_data(token_count, 1024); + std::vector inv_freq(static_cast(kQwenFlashHeadSize / 2)); + for (size_t i = 0; i < inv_freq.size(); ++i) { + inv_freq[i] = 1.0f / std::pow(10000.0f, static_cast(2 * i) / static_cast(kQwenFlashHeadSize)); + } + const std::vector key_cache(static_cast(key_value_size * cache_row_count), + ggml_fp32_to_fp16(0.0f)); + const std::vector value_cache(static_cast(key_value_size * cache_row_count), + ggml_fp32_to_fp16(0.0f)); + const std::vector cache_indices = make_i64_mod_data(token_count, cache_row_count); + + set_tensor_pair_bytes(cpu_backend, cpu.input, hrx_backend, hrx.input, input.data(), input.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu.query_weight, hrx_backend, hrx.query_weight, query_w.data(), query_w.size()); + set_tensor_pair_bytes(cpu_backend, cpu.key_weight, hrx_backend, hrx.key_weight, key_w.data(), key_w.size()); + set_tensor_pair_bytes(cpu_backend, cpu.value_weight, hrx_backend, hrx.value_weight, value_w.data(), value_w.size()); + set_tensor_pair_bytes(cpu_backend, cpu.query_norm_weight, hrx_backend, hrx.query_norm_weight, query_nw.data(), + query_nw.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu.key_norm_weight, hrx_backend, hrx.key_norm_weight, key_nw.data(), + key_nw.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu.positions, hrx_backend, hrx.positions, positions.data(), + positions.size() * sizeof(int32_t)); + set_tensor_pair_bytes(cpu_backend, cpu.inverse_frequencies, hrx_backend, hrx.inverse_frequencies, inv_freq.data(), + inv_freq.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu.key_cache, hrx_backend, hrx.key_cache, key_cache.data(), + key_cache.size() * sizeof(ggml_fp16_t)); + set_tensor_pair_bytes(cpu_backend, cpu.value_cache, hrx_backend, hrx.value_cache, value_cache.data(), + value_cache.size() * sizeof(ggml_fp16_t)); + set_tensor_pair_bytes(cpu_backend, cpu.key_cache_indices, hrx_backend, hrx.key_cache_indices, cache_indices.data(), + cache_indices.size() * sizeof(int64_t)); + set_tensor_pair_bytes(cpu_backend, cpu.value_cache_indices, hrx_backend, hrx.value_cache_indices, + cache_indices.data(), cache_indices.size() * sizeof(int64_t)); + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx.query_output), get_f32_tensor(cpu_backend, cpu.query_output), 2.0f, + 5.0e-2f); + require_close(get_f32_tensor(hrx_backend, hrx.key_output), get_f32_tensor(cpu_backend, cpu.key_output), 2.0f, + 5.0e-2f); + require_close(get_f32_tensor(hrx_backend, hrx.value_output), get_f32_tensor(cpu_backend, cpu.value_output), 2.0f, + 5.0e-2f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_routed_moe_cpu_reference_case(ggml_type down_weight_type, bool include_next_rmsnorm) { + ggml_backend_t cpu_backend = init_cpu_backend(); + ggml_backend_t hrx_backend = ggml_backend_hrx_init(0); + REQUIRE(hrx_backend != nullptr); + + ggml_init_params params = {}; + params.mem_size = 64 * 1024 * 1024; + params.no_alloc = true; + ggml_context * cpu_ctx = ggml_init(params); + ggml_context * hrx_ctx = ggml_init(params); + REQUIRE(cpu_ctx != nullptr); + REQUIRE(hrx_ctx != nullptr); + + RoutedMoeGraph cpu = build_routed_moe_graph(cpu_ctx, down_weight_type, include_next_rmsnorm); + RoutedMoeGraph hrx = build_routed_moe_graph(hrx_ctx, down_weight_type, include_next_rmsnorm); + + ggml_cgraph * cpu_graph = ggml_new_graph(cpu_ctx); + ggml_cgraph * hrx_graph = ggml_new_graph(hrx_ctx); + REQUIRE(cpu_graph != nullptr); + REQUIRE(hrx_graph != nullptr); + ggml_build_forward_expand(cpu_graph, cpu.output); + ggml_build_forward_expand(hrx_graph, hrx.output); + + std::vector expected = { + "qwen3_moe:qwen3_moe_router_top8_f32", + "loom_libs:ggml_moe_build_expert_table", + "loom_libs:ggml_moe_build_expert_partition_table", + "loom_libs:ggml_mul_mat_id_swiglu_f16_f16_wmma", + "loom_libs:ggml_mul_mat_id_f16_f16_wmma", + include_next_rmsnorm ? "qwen3_moe:qwen3_moe_routed_down_weighted_reduce_next_rmsnorm_f32" : + "qwen3_moe:qwen3_moe_routed_down_weighted_reduce_f16_f32", + }; + require_kernel_subsequence(scheduled_kernel_sequence(hrx_graph), expected); + + ggml_backend_buffer_t cpu_buffer = ggml_backend_alloc_ctx_tensors(cpu_ctx, cpu_backend); + ggml_backend_buffer_t hrx_buffer = ggml_backend_alloc_ctx_tensors(hrx_ctx, hrx_backend); + REQUIRE(cpu_buffer != nullptr); + REQUIRE(hrx_buffer != nullptr); + + const std::vector logits = make_router_logits(1); + const std::vector input = make_pattern_f32(kQwenHiddenSize, 15, 0.01f); + const std::vector hidden = make_pattern_f32(kQwenHiddenSize, 16, 0.02f); + set_tensor_pair_bytes(cpu_backend, cpu.logits, hrx_backend, hrx.logits, logits.data(), + logits.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu.input, hrx_backend, hrx.input, input.data(), input.size() * sizeof(float)); + set_tensor_pair_bytes(cpu_backend, cpu.hidden_state, hrx_backend, hrx.hidden_state, hidden.data(), + hidden.size() * sizeof(float)); + if (include_next_rmsnorm) { + const std::vector norm = make_weight(kQwenHiddenSize); + set_tensor_pair_bytes(cpu_backend, cpu.norm_weight, hrx_backend, hrx.norm_weight, norm.data(), + norm.size() * sizeof(float)); + } + + { + const std::vector gate = + make_quantized_rows(GGML_TYPE_Q4_K, kQwenHiddenSize, kQwenMoeIntermediate * kQwenRouterExpertCount, 17); + set_tensor_pair_bytes(cpu_backend, cpu.gate_weight, hrx_backend, hrx.gate_weight, gate.data(), gate.size()); + } + { + const std::vector up = + make_quantized_rows(GGML_TYPE_Q4_K, kQwenHiddenSize, kQwenMoeIntermediate * kQwenRouterExpertCount, 18); + set_tensor_pair_bytes(cpu_backend, cpu.up_weight, hrx_backend, hrx.up_weight, up.data(), up.size()); + } + { + const std::vector down = + make_quantized_rows(down_weight_type, kQwenMoeIntermediate, kQwenHiddenSize * kQwenRouterExpertCount, 19); + set_tensor_pair_bytes(cpu_backend, cpu.down_weight, hrx_backend, hrx.down_weight, down.data(), down.size()); + } + + REQUIRE(ggml_backend_graph_compute(cpu_backend, cpu_graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(hrx_backend, hrx_graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(cpu_backend); + ggml_backend_synchronize(hrx_backend); + require_close(get_f32_tensor(hrx_backend, hrx.output), get_f32_tensor(cpu_backend, cpu.output), 2.5f, 1.0e-1f); + + ggml_backend_buffer_free(cpu_buffer); + ggml_backend_buffer_free(hrx_buffer); + ggml_free(cpu_ctx); + ggml_free(hrx_ctx); + ggml_backend_free(cpu_backend); + ggml_backend_free(hrx_backend); +} + +static void run_decode_routed_moe_scheduling_case(ggml_type down_weight_type, bool alias_gate_input = false) { + ggml_init_params params = {}; + params.mem_size = 128 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + ggml_tensor * hidden_state = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenHiddenSize, 1); + ggml_tensor * attention_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, kQwenHiddenSize); + REQUIRE(hidden_state != nullptr); + REQUIRE(attention_weight != nullptr); + ggml_tensor * attention_rms = ggml_rms_norm(ctx, hidden_state, kQwenRmsNormEps); + REQUIRE(attention_rms != nullptr); + ggml_tensor * attention_prepared = ggml_mul(ctx, attention_rms, attention_weight); + REQUIRE(attention_prepared != nullptr); + ggml_tensor * moe_input = attention_prepared; + if (alias_gate_input) { + moe_input = ggml_reshape_2d(ctx, attention_prepared, kQwenHiddenSize, 1); + REQUIRE(moe_input != nullptr); + } + + ggml_tensor * router_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, kQwenHiddenSize, kQwenRouterExpertCount); + REQUIRE(router_weight != nullptr); + ggml_tensor * logits = ggml_mul_mat(ctx, router_weight, attention_prepared); + REQUIRE(logits != nullptr); + ggml_tensor * route_ids = nullptr; + ggml_tensor * route_weights = build_qwen_router_top8_graph(ctx, logits, &route_ids); + REQUIRE(route_ids != nullptr); + REQUIRE(route_weights != nullptr); + + ggml_tensor * gate_weight = + ggml_new_tensor_3d(ctx, GGML_TYPE_Q4_K, kQwenHiddenSize, kQwenMoeIntermediate, kQwenRouterExpertCount); + ggml_tensor * up_weight = + ggml_new_tensor_3d(ctx, GGML_TYPE_Q4_K, kQwenHiddenSize, kQwenMoeIntermediate, kQwenRouterExpertCount); + ggml_tensor * down_weight = + ggml_new_tensor_3d(ctx, down_weight_type, kQwenMoeIntermediate, kQwenHiddenSize, kQwenRouterExpertCount); + REQUIRE(gate_weight != nullptr); + REQUIRE(up_weight != nullptr); + REQUIRE(down_weight != nullptr); + + ggml_tensor * gate = ggml_mul_mat_id(ctx, gate_weight, moe_input, route_ids); + ggml_tensor * up = ggml_mul_mat_id(ctx, up_weight, moe_input, route_ids); + REQUIRE(gate != nullptr); + REQUIRE(up != nullptr); + ggml_tensor * glu = ggml_glu_split(ctx, gate, up, GGML_GLU_OP_SWIGLU); + REQUIRE(glu != nullptr); + ggml_tensor * down = ggml_mul_mat_id(ctx, down_weight, glu, route_ids); + REQUIRE(down != nullptr); + ggml_tensor * weighted = ggml_mul(ctx, down, route_weights); + REQUIRE(weighted != nullptr); + + std::vector route_views; + route_views.reserve(kQwenRouterRouteCount); + for (int64_t route = 0; route < kQwenRouterRouteCount; ++route) { + ggml_tensor * view = ggml_view_2d(ctx, weighted, kQwenHiddenSize, 1, weighted->nb[2], + static_cast(route) * weighted->nb[1]); + REQUIRE(view != nullptr); + route_views.push_back(view); + } + + ggml_tensor * reduced = route_views.front(); + for (size_t i = 1; i < route_views.size(); ++i) { + reduced = ggml_add(ctx, reduced, route_views[i]); + REQUIRE(reduced != nullptr); + } + ggml_tensor * residual = ggml_add(ctx, hidden_state, reduced); + REQUIRE(residual != nullptr); + ggml_tensor * next_norm_weight = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, kQwenHiddenSize); + REQUIRE(next_norm_weight != nullptr); + ggml_tensor * next_rms = ggml_rms_norm(ctx, residual, kQwenRmsNormEps); + REQUIRE(next_rms != nullptr); + ggml_tensor * output = ggml_mul(ctx, next_rms, next_norm_weight); + REQUIRE(output != nullptr); + ggml_tensor * next_weight = ggml_new_tensor_2d(ctx, GGML_TYPE_Q4_K, kQwenHiddenSize, kQwenHiddenSize); + ggml_tensor * next_output = ggml_mul_mat(ctx, next_weight, output); + REQUIRE(next_weight != nullptr); + REQUIRE(next_output != nullptr); + + ggml_cgraph * graph = ggml_new_graph(ctx); + REQUIRE(graph != nullptr); + ggml_build_forward_expand(graph, next_output); + + std::vector expected = { + "qwen3_moe:qwen3_moe_rmsnorm_f32_quantize_q8_1_x4", + "qwen3_moe:qwen3_moe_router_projection_top8_fused_decode_f32", + down_weight_type == GGML_TYPE_Q4_K ? "loom_libs:ggml_mul_mat_id_vector_pair_binary_publish_q8" : + "loom_libs:ggml_mul_mat_id_vector_pair_binary_publish_f32", + down_weight_type == GGML_TYPE_Q4_K ? "loom_libs:ggml_mul_mat_id_vector_weighted_wave32_publish_q8" : + "loom_libs:ggml_mul_mat_id_vector_weighted_wave64_publish_q8", + }; + require_kernel_subsequence(scheduled_kernel_sequence(graph), expected); + ggml_free(ctx); +} + +static void run_decode_attention_qkv_scheduling_case() { + ggml_init_params params = {}; + params.mem_size = 128 * 1024 * 1024; + params.no_alloc = true; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + + AttentionPostprocessGraph graph = build_decode_attention_qkv_graph(ctx); + ggml_cgraph * cgraph = ggml_new_graph(ctx); + REQUIRE(cgraph != nullptr); + ggml_build_forward_expand(cgraph, graph.query_output); + ggml_build_forward_expand(cgraph, graph.key_output); + ggml_build_forward_expand(cgraph, graph.value_output); + + require_kernel_subsequence(scheduled_kernel_sequence(cgraph), + { "qwen3_moe:qwen3_moe_rmsnorm_f32_quantize_q8_1_x4", + "qwen3_moe:qwen3_moe_attention_qkv_postprocess_fused_decode" }); + ggml_free(ctx); +} + +struct SsmConvPrefillGraph { + ggml_context * ctx = nullptr; + ggml_tensor * input = nullptr; + ggml_tensor * weight = nullptr; + ggml_tensor * pool = nullptr; + ggml_tensor * initial = nullptr; + ggml_tensor * ids = nullptr; + ggml_tensor * filter = nullptr; + ggml_tensor * output = nullptr; + ggml_tensor * cache = nullptr; + ggml_cgraph * graph = nullptr; +}; + +static SsmConvPrefillGraph make_ssm_conv_prefill_graph(ggml_type type, + int64_t k, + int64_t n, + bool expose, + bool gathered, + bool clear_state) { + ggml_init_params params = {}; + params.mem_size = 32 * 1024 * 1024; + params.no_alloc = true; + SsmConvPrefillGraph c; + c.ctx = ggml_init(params); + REQUIRE(c.ctx != nullptr); + c.input = ggml_new_tensor_2d(c.ctx, GGML_TYPE_F32, k, 512); + c.weight = ggml_new_tensor_2d(c.ctx, type, k, n); + c.pool = ggml_new_tensor_2d(c.ctx, GGML_TYPE_F32, 3 * n, 2); + c.initial = ggml_new_tensor_3d(c.ctx, GGML_TYPE_F32, 3, n, 1); + c.ids = ggml_new_tensor_1d(c.ctx, GGML_TYPE_I32, 1); + c.filter = ggml_new_tensor_2d(c.ctx, GGML_TYPE_F32, 4, n); + c.graph = ggml_new_graph(c.ctx); + auto state = c.initial; + if (gathered) { + if (clear_state) { + auto row = ggml_view_2d(c.ctx, c.pool, 3 * n, 1, c.pool->nb[1], 0); + ggml_build_forward_expand(c.graph, ggml_scale_inplace(c.ctx, row, 0.0f)); + } + state = ggml_get_rows(c.ctx, c.pool, c.ids); + ggml_build_forward_expand(c.graph, state); + state = ggml_reshape_3d(c.ctx, state, 3, n, 1); + } + auto raw = ggml_mul_mat(c.ctx, c.weight, c.input); + auto x = ggml_reshape_3d(c.ctx, raw, n, 512, 1); + auto window = ggml_concat(c.ctx, state, ggml_transpose(c.ctx, x), 0); + auto tail = ggml_view_3d(c.ctx, window, 3, n, 1, window->nb[1], window->nb[2], 512 * sizeof(float)); + c.cache = ggml_view_2d(c.ctx, c.pool, 3 * n, 1, c.pool->nb[1], 0); + ggml_build_forward_expand(c.graph, ggml_cpy(c.ctx, tail, c.cache)); + c.output = ggml_silu(c.ctx, ggml_ssm_conv(c.ctx, window, c.filter)); + ggml_set_output(c.output); + ggml_build_forward_expand(c.graph, c.output); + if (expose) { + ggml_set_output(raw); + ggml_build_forward_expand(c.graph, ggml_scale(c.ctx, raw, 0.5f)); + } + const auto sequence = scheduled_kernel_sequence(c.graph); + const char * fused = gathered ? "loom_libs:ggml_mul_mat_quantized_f16_wmma_prefill_conv4_interior" : + "loom_libs:ggml_mul_mat_quantized_f16_wmma_prefill_conv4"; + REQUIRE((std::find(sequence.begin(), sequence.end(), fused) != sequence.end()) == !expose); + if (gathered && !expose) { + REQUIRE(std::find(sequence.begin(), sequence.end(), "loom_libs:llm_ssm_conv_dconv4_silu_prefill_finish_f32") != + sequence.end()); + } + return c; +} + +static void run_ssm_conv_prefill_recurrent_case(ggml_type type, bool gathered, bool clear_state) { + const int64_t k = 256; + const int64_t n = 8192; + auto backend = ggml_backend_hrx_init(0); + REQUIRE(backend != nullptr); + auto a = make_ssm_conv_prefill_graph(type, k, n, false, gathered, clear_state); + auto b = make_ssm_conv_prefill_graph(type, k, n, true, gathered, clear_state); + auto aw = allocate_test_weight(backend, a.weight); + auto bw = allocate_test_weight(backend, b.weight); + auto ab = ggml_backend_alloc_ctx_tensors(a.ctx, backend); + auto bb = ggml_backend_alloc_ctx_tensors(b.ctx, backend); + REQUIRE(ab != nullptr && bb != nullptr); + const auto weights = make_matmul_weight_bytes(type, k, n, 6); + const auto filters = make_pattern_f32(4 * n, 11, 0.0625f); + const auto states = make_pattern_f32(6 * n, 17, 0.03125f); + set_tensor_pair_bytes(backend, a.weight, backend, b.weight, weights.data(), weights.size()); + set_tensor_pair_bytes(backend, a.filter, backend, b.filter, filters.data(), filters.size() * sizeof(float)); + set_tensor_pair_bytes(backend, a.pool, backend, b.pool, states.data(), states.size() * sizeof(float)); + set_tensor_pair_bytes(backend, a.initial, backend, b.initial, states.data(), 3 * n * sizeof(float)); + for (int round = 0; round < 5; ++round) { + const int32_t id = round % 2; + set_tensor_pair_bytes(backend, a.ids, backend, b.ids, &id, sizeof(id)); + const auto input = make_pattern_f32(k * 512, 7 + round, 0.0625f); + set_tensor_pair_bytes(backend, a.input, backend, b.input, input.data(), input.size() * sizeof(float)); + REQUIRE(ggml_backend_graph_compute(backend, a.graph) == GGML_STATUS_SUCCESS); + REQUIRE(ggml_backend_graph_compute(backend, b.graph) == GGML_STATUS_SUCCESS); + ggml_backend_synchronize(backend); + const auto actual = get_f32_tensor(backend, a.output); + const auto expected = get_f32_tensor(backend, b.output); + require_close(get_f32_tensor(backend, a.cache), get_f32_tensor(backend, b.cache), 0.0f, 0.0f); + require_close(actual, expected, 0.0f, 0.0f); + } + ggml_backend_buffer_free(ab); + ggml_backend_buffer_free(bb); + aw.reset(); + bw.reset(); + ggml_free(a.ctx); + ggml_free(b.ctx); + ggml_backend_free(backend); +} + +static void run_vector_output_capacity_experiment() { + constexpr int64_t kSmallVectorInputSize = 256; + constexpr int64_t kAboveOldOutputSize = 262208; + constexpr int64_t kLiftedOutputSize = 524288; + + run_dense_matmul_cpu_reference_case(GGML_TYPE_F32, "loom_libs:ggml_mul_mat_vector_f32_f32", 1, kAboveOldOutputSize, + kSmallVectorInputSize); + run_dense_matmul_unary_cpu_reference_case(GGML_TYPE_F32, "loom_libs:ggml_mul_mat_vector_f32_f32", 1, + kAboveOldOutputSize, kSmallVectorInputSize); + run_dense_matmul_postops_cpu_reference_case(GGML_TYPE_F16, "loom_libs:ggml_mul_mat_vector_bias_residual_f32_f32", 1, + kAboveOldOutputSize, true, false, false, kSmallVectorInputSize); + run_dense_matmul_cpu_reference_case(GGML_TYPE_F16, "loom_libs:ggml_mul_mat_vector_f32_f32", 1, kLiftedOutputSize, + kSmallVectorInputSize); +} + +static void run_vector_postops_cpu_reference_case() { + constexpr int64_t kVectorInputSize = 256; + constexpr int64_t kVectorOutputSize = 256; + + run_dense_matmul_postops_cpu_reference_case(GGML_TYPE_F16, "loom_libs:ggml_mul_mat_vector_bias_residual_f32_f32", 1, + kVectorOutputSize, true, false, false, kVectorInputSize); + run_dense_matmul_postops_cpu_reference_case(GGML_TYPE_F16, "loom_libs:ggml_mul_mat_vector_bias_residual_f32_f32", 1, + kVectorOutputSize, true, true, false, kVectorInputSize); +} + +static const char * glu_op_name(ggml_glu_op op) { + switch (op) { + case GGML_GLU_OP_REGLU: + return "reglu"; + case GGML_GLU_OP_GEGLU: + return "geglu"; + case GGML_GLU_OP_SWIGLU: + return "swiglu"; + case GGML_GLU_OP_SWIGLU_OAI: + return "swiglu_oai"; + case GGML_GLU_OP_GEGLU_ERF: + return "geglu_erf"; + case GGML_GLU_OP_GEGLU_QUICK: + return "geglu_quick"; + case GGML_GLU_OP_COUNT: + break; + } + return "unknown_glu"; +} + +static std::string type_name(ggml_type type) { + return ggml_type_name(type); +} + +using test_runner::Suite; + +static void register_basic_ops_cases(Suite & suite) { + suite.host_case("support.rmsnorm", [] { run_rmsnorm_support_checks(); }); + suite.host_case("support.alternate_value_alias_lookup", [] { run_alternate_value_alias_lookup_checks(); }); + suite.host_case("support.activation_publication_contract", [] { run_activation_publication_contract_checks(); }); + suite.device_case("add_f32.cpu_reference", [] { run_add_f32_cpu_reference_case(); }); + suite.device_case("dynamic_generic_ops.cpu_reference", [] { run_dynamic_generic_ops_cpu_reference_case(); }); + for (const ggml_type type : { GGML_TYPE_Q4_K, GGML_TYPE_Q6_K }) { + suite.device_case("ssm_conv_prefill_recurrent." + type_name(type) + ".plain", + [type] { run_ssm_conv_prefill_recurrent_case(type, false, false); }); + suite.device_case("ssm_conv_prefill_recurrent." + type_name(type) + ".gathered", + [type] { run_ssm_conv_prefill_recurrent_case(type, true, false); }); + suite.device_case("ssm_conv_prefill_recurrent." + type_name(type) + ".gathered_clear_state", + [type] { run_ssm_conv_prefill_recurrent_case(type, true, true); }); + } + suite.device_case("view_mul_f32.cpu_reference", [] { run_view_mul_f32_cpu_reference_case(); }); + for (const int64_t token_count : { 2, 23 }) { + suite.device_case("strided_cont.tokens" + std::to_string(token_count), + [token_count] { run_strided_cont_cpu_reference_case(token_count); }); + } + for (const ggml_glu_op op : + { GGML_GLU_OP_REGLU, GGML_GLU_OP_SWIGLU, GGML_GLU_OP_GEGLU, GGML_GLU_OP_GEGLU_ERF, GGML_GLU_OP_GEGLU_QUICK }) { + suite.device_case(std::string("glu_split_f32.") + glu_op_name(op), + [op] { run_glu_split_f32_cpu_reference_case(op); }); + suite.device_case(std::string("glu_packed_f32.") + glu_op_name(op) + ".normal", + [op] { run_glu_packed_f32_cpu_reference_case(op, false); }); + } + suite.device_case("glu_packed_f32.swiglu.gated", + [] { run_glu_packed_f32_cpu_reference_case(GGML_GLU_OP_SWIGLU, true); }); + suite.device_case("glu_packed_f32.swiglu.outputs8192.tokens2", + [] { run_glu_packed_f32_cpu_reference_case(GGML_GLU_OP_SWIGLU, false, 8192, 2); }); + suite.device_case("glu_packed_f32.swiglu.outputs8192.tokens14", + [] { run_glu_packed_f32_cpu_reference_case(GGML_GLU_OP_SWIGLU, false, 8192, 14); }); + suite.device_case("gather_add_f32.cpu_reference", [] { run_gather_add_f32_cpu_reference_case(); }); + suite.device_case("scale_add_f32.cpu_reference.packed", [] { run_scale_add_f32_cpu_reference_case(false); }); + suite.device_case("scale_add_f32.cpu_reference.strided", [] { run_scale_add_f32_cpu_reference_case(true); }); + for (const ggml_type type : { GGML_TYPE_Q4_K, GGML_TYPE_Q6_K, GGML_TYPE_Q8_0, GGML_TYPE_IQ4_XS, GGML_TYPE_F16, + GGML_TYPE_BF16, GGML_TYPE_F32 }) { + suite.device_case("get_rows_f32." + type_name(type), [type] { run_get_rows_f32_cpu_reference_case(type); }); + } + suite.device_case("get_rows_f32.q1_0.rows2048", [] { run_get_rows_f32_cpu_reference_case(GGML_TYPE_Q1_0, 2048); }); + suite.device_case("get_rows_f32.q5_1.rows640", [] { run_get_rows_f32_cpu_reference_case(GGML_TYPE_Q5_1, 640); }); + suite.device_case("get_rows_q8_1.zero_weight", [] { run_get_rows_q8_1_zero_weight_case(); }); + suite.device_case("get_rows_scale_f32.cpu_reference", [] { run_get_rows_scale_f32_cpu_reference_case(); }); + suite.device_case("get_rows_scale_f32.q5_1.cpu_reference", + [] { run_get_rows_scale_f32_cpu_reference_case(GGML_TYPE_Q5_1, 640); }); +} + +static void register_dense_matmul_cases(Suite & suite) { + for (const ggml_type type : { GGML_TYPE_IQ1_S, GGML_TYPE_IQ1_M }) { + const char * type_suffix = type == GGML_TYPE_IQ1_S ? "iq1_s" : "iq1_m"; + const char * vector_kernel = type == GGML_TYPE_IQ1_S ? + "loom_libs:ggml_mul_mat_vector_iq1_s_f32_f32" : + "loom_libs:ggml_mul_mat_vector_iq1_m_f32_f32"; + const char * tiled_kernel = type == GGML_TYPE_IQ1_S ? + "loom_libs:ggml_mul_mat_tiled_input_f32_iq1_s_publish_f32" : + "loom_libs:ggml_mul_mat_tiled_input_f32_iq1_m_publish_f32"; + suite.device_case(std::string("dense_matmul.") + type_suffix + ".vector.tokens1.outputs128", [=] { + run_dense_matmul_cpu_reference_case(type, vector_kernel, 1, 128, 2048); + }); + suite.device_case(std::string("dense_matmul.") + type_suffix + ".tiled.tokens64.outputs128", [=] { + run_dense_matmul_cpu_reference_case(type, tiled_kernel, 64, 128, 2048); + }); + } + for (const int64_t tokens : { 1, 3, 5 }) { + for (const int64_t outputs : { 1, 4, 47, 48, 63 }) { + suite.device_case( + "dense_matmul.q4_k.tokens" + std::to_string(tokens) + ".outputs" + std::to_string(outputs), + [tokens, outputs] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_Q4_K, + tokens == 1 ? + "loom_libs:ggml_mul_mat_vector_f32_f32" : + "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", + tokens, outputs); + }); + } + } + for (const ggml_type type : + { GGML_TYPE_Q5_K, GGML_TYPE_Q6_K, GGML_TYPE_Q8_0, GGML_TYPE_F16, GGML_TYPE_BF16, GGML_TYPE_F32 }) { + suite.device_case("dense_matmul." + type_name(type) + ".tokens3.outputs47", [type] { + run_dense_matmul_cpu_reference_case(type, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 3, 47); + }); + } + for (const ggml_type type : { GGML_TYPE_Q4_K, GGML_TYPE_Q6_K }) { + suite.device_case("dense_matmul." + type_name(type) + ".tiled.tokens33.outputs65.input512", [type] { + run_dense_matmul_cpu_reference_case(type, "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32", 33, 65, + 512); + }); + } + suite.device_case("dense_matmul_unary.f32.tokens3.outputs47", [] { + run_dense_matmul_unary_cpu_reference_case(GGML_TYPE_F32, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", + 3, 47); + }); + suite.device_case("get_rows_f32.f32.rows262144.tokens4.outputs1", + [] { run_get_rows_f32_cpu_reference_case(GGML_TYPE_F32, 262144, 4, 1); }); + suite.device_case("dense_matmul.q4_k.skinny.tokens2.outputs128", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_Q4_K, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 2, + 128); + }); + suite.device_case("dense_matmul.q4_k.vector.tokens1.outputs128", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_Q4_K, "loom_libs:ggml_mul_mat_vector_f32_f32", 1, 128); + }); + for (const ggml_type type : { GGML_TYPE_IQ2_XXS, GGML_TYPE_IQ2_XS }) { + suite.device_case("dense_matmul." + type_name(type) + ".codebook.vector.tokens1.outputs128", [type] { + run_dense_matmul_cpu_reference_case(type, "loom_libs:ggml_mul_mat_vector_iq2_s_f32_f32", 1, 128, + 2048); + }); + suite.device_case("dense_matmul." + type_name(type) + ".codebook.tiled.tokens64.outputs128", [type] { + run_dense_matmul_cpu_reference_case( + type, "loom_libs:ggml_mul_mat_tiled_input_f32_iq2_s_publish_f32", 64, 128, 2048); + }); + } + suite.device_case("dense_matmul.iq3_xxs.codebook.vector.tokens1.outputs128", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_IQ3_XXS, + "loom_libs:ggml_mul_mat_vector_iq3_xxs_f32_f32", 1, 128, 2048); + }); + suite.device_case("dense_matmul.iq3_xxs.codebook.tiled.tokens64.outputs128", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_IQ3_XXS, + "loom_libs:ggml_mul_mat_tiled_input_f32_iq3_xxs_publish_f32", 64, + 128, 2048); + }); + for (const ggml_type type : { GGML_TYPE_IQ2_S, GGML_TYPE_IQ3_S }) { + const char * vector_kernel = type == GGML_TYPE_IQ2_S ? + "loom_libs:ggml_mul_mat_vector_iq2_s_f32_f32" : + "loom_libs:ggml_mul_mat_vector_iq3_s_f32_f32"; + const char * tiled_kernel = type == GGML_TYPE_IQ2_S ? + "loom_libs:ggml_mul_mat_tiled_input_f32_iq2_s_publish_f32" : + "loom_libs:ggml_mul_mat_tiled_input_f32_iq3_s_publish_f32"; + suite.device_case("dense_matmul." + type_name(type) + ".codebook_control.vector.tokens1.outputs128", + [type, vector_kernel] { + run_dense_matmul_cpu_reference_case(type, vector_kernel, 1, 128, 2048); + }); + suite.device_case("dense_matmul." + type_name(type) + ".codebook_control.tiled.tokens64.outputs128", + [type, tiled_kernel] { + run_dense_matmul_cpu_reference_case(type, tiled_kernel, 64, 128, 2048); + }); + } + suite.device_case("dense_matmul.iq4_xs.vector.tokens1.outputs128", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_IQ4_XS, "loom_libs:ggml_mul_mat_vector_f32_f32", 1, 128); + }); + suite.device_case("dense_matmul.iq4_xs.tiled.tokens64.outputs128", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_IQ4_XS, + "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32_aligned", 64, 128); + }); + suite.device_case("dense_matmul.iq4_xs.prefill.tokens256.outputs128", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_IQ4_XS, + "loom_libs:ggml_mul_mat_q5_k_iq4_xs_q8_1_x4_wmma_token256", 256, 128, + kQwenHiddenSize, true); + }); + // Exercise the production K dimension and a partial output tile. + suite.device_case("dense_matmul.iq4_xs.prefill.tokens512.outputs64", [] { + run_dense_matmul_cpu_reference_case( + GGML_TYPE_IQ4_XS, "loom_libs:ggml_mul_mat_q5_k_iq4_xs_q8_1_x4_wmma_token256", 512, 64, 3072, true); + }); + for (const ggml_type type : { GGML_TYPE_Q4_K, GGML_TYPE_Q6_K }) { + for (const ggml_type cache_type : { GGML_TYPE_F16, GGML_TYPE_F32 }) { + suite.device_case( + "attention_v_matmul." + type_name(type) + ".cache_" + type_name(cache_type), [type, cache_type] { + run_dense_matmul_cpu_reference_case(type, + "loom_libs:llm_attention_v_matmul_set_rows_vector_f32_f32", 1, + 1024, 2048, false, false, cache_type); + }); + } + } + suite.device_case("attention_v_matmul.q6_k.prefill_f16.repeated_indices", + [] { run_prefill_value_cache_repeated_indices_case(); }); + for (const ggml_type type : + { GGML_TYPE_Q1_0, GGML_TYPE_Q4_0, GGML_TYPE_Q4_1, GGML_TYPE_Q5_0, GGML_TYPE_Q5_1, GGML_TYPE_IQ2_S, + GGML_TYPE_Q6_K, GGML_TYPE_Q8_0, GGML_TYPE_IQ4_NL, GGML_TYPE_F16, GGML_TYPE_BF16 }) { + const char * vector_kernel = type == GGML_TYPE_Q6_K ? "loom_libs:ggml_mul_mat_vector_q6_f32_f32" : + "loom_libs:ggml_mul_mat_vector_f32_f32"; + const int64_t outputs = type == GGML_TYPE_IQ2_S ? 640 : (type == GGML_TYPE_IQ4_NL ? 1024 : 128); + const int64_t input = type == GGML_TYPE_IQ4_NL ? 640 : kQwenHiddenSize; + const char * route = type == GGML_TYPE_IQ4_NL ? "tiled" : "skinny"; + const char * kernel = type == GGML_TYPE_IQ4_NL ? "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32" : + "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32"; + suite.device_case( + "dense_matmul." + type_name(type) + "." + route + ".tokens2.outputs" + std::to_string(outputs), + [type, outputs, input, kernel] { run_dense_matmul_cpu_reference_case(type, kernel, 2, outputs, input); }); + suite.device_case("dense_matmul." + type_name(type) + ".vector.tokens1.outputs" + std::to_string(outputs), + [type, vector_kernel, outputs, input] { + run_dense_matmul_cpu_reference_case(type, vector_kernel, 1, outputs, input); + }); + } + suite.device_case("dense_matmul.iq4_nl.tiled.tokens64.outputs1024", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_IQ4_NL, + "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32_aligned", 64, 1024, + 640); + }); + suite.device_case("dense_matmul.q5_0.tiled.tokens2.outputs256.input640", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_Q5_0, "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32", 2, + 256, 640); + }); + suite.device_case("dense_matmul.q5_0.tiled.tokens2.outputs128.input544", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_Q5_0, "loom_libs:ggml_mul_mat_tiled_input_f32_publish_f32", 2, + 128, 544); + }); + suite.device_case("dense_matmul.q5_0.vector.tokens1.outputs256.input640", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_Q5_0, "loom_libs:ggml_mul_mat_vector_f32_f32", 1, 256, 640); + }); + suite.device_case("dense_matmul.q8_1.zero_weight.skinny", [] { + run_dense_matmul_zero_weight_case(GGML_TYPE_Q8_1, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32"); + }); + suite.device_case("dense_matmul.q8_1.zero_weight.vector", [] { + run_dense_matmul_zero_weight_case(GGML_TYPE_Q8_1, "loom_libs:ggml_mul_mat_vector_f32_f32", 1); + }); + suite.device_case("dense_matmul.f32.skinny.tokens2.outputs256", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_F32, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 2, + 256); + }); + suite.device_case("dense_matmul.f32.vector.tokens1.outputs256", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_F32, "loom_libs:ggml_mul_mat_vector_f32_f32", 1, 256); + }); + suite.device_case("dense_matmul_unary.f32.tokens2.outputs128", [] { + run_dense_matmul_unary_cpu_reference_case(GGML_TYPE_F32, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", + 2, 128); + }); + suite.device_case("dense_matmul_swiglu.q4_k.f16.swiglu.tokens2.outputs128", [] { + run_dense_matmul_swiglu_cpu_reference_case(GGML_TYPE_Q4_K, GGML_TYPE_F16, + "loom_libs:ggml_mul_mat_swiglu_f32_f32_lowtoken_dot", 2, 128); + }); + suite.device_case("dense_matmul_swiglu.q4_k.f16.geglu.tokens2.outputs128", [] { + run_dense_matmul_swiglu_cpu_reference_case(GGML_TYPE_Q4_K, GGML_TYPE_F16, + "loom_libs:ggml_mul_mat_swiglu_f32_f32_lowtoken_dot", 2, 128, + GGML_GLU_OP_GEGLU); + }); + suite.device_case("dense_matmul_swiglu.iq4_nl.iq4_nl.geglu.tokens2.outputs2048.input640", [] { + run_dense_matmul_swiglu_cpu_reference_case(GGML_TYPE_IQ4_NL, GGML_TYPE_IQ4_NL, + "loom_libs:ggml_mul_mat_tiled_pair_input_f32_binary_publish_f32", 2, + 2048, GGML_GLU_OP_GEGLU, 640); + }); + suite.device_case("dense_matmul_swiglu.iq4_nl.iq4_nl.geglu.tokens64.outputs128.input640", [] { + run_dense_matmul_swiglu_cpu_reference_case( + GGML_TYPE_IQ4_NL, GGML_TYPE_IQ4_NL, + "loom_libs:ggml_mul_mat_tiled_pair_input_f32_binary_publish_f32_aligned", + 64, 128, GGML_GLU_OP_GEGLU, 640); + }); + for (const ggml_glu_op op : { GGML_GLU_OP_REGLU, GGML_GLU_OP_GEGLU_ERF, GGML_GLU_OP_GEGLU_QUICK }) { + suite.device_case(std::string("dense_matmul_swiglu.q4_k.f16.") + glu_op_name(op) + ".tokens2.outputs128", [op] { + run_dense_matmul_swiglu_cpu_reference_case( + GGML_TYPE_Q4_K, GGML_TYPE_F16, "loom_libs:ggml_mul_mat_swiglu_f32_f32_lowtoken_dot", 2, 128, op); + }); + } + suite.device_case("dense_matmul_packed_glu.q4_k.tokens2.outputs128", [] { + run_dense_matmul_packed_glu_cpu_reference_case( + GGML_TYPE_Q4_K, "loom_libs:ggml_mul_mat_tiled_pair_input_f32_binary_publish_f32", 2, 128); + }); + suite.device_case("dense_matmul_packed_glu.q4_k.swiglu.tokens2.outputs128.gated", [] { + run_dense_matmul_packed_glu_cpu_reference_case(GGML_TYPE_Q4_K, + "loom_libs:ggml_mul_mat_tiled_pair_input_f32_binary_publish_f32", + 2, 128, GGML_GLU_OP_SWIGLU, true); + }); + suite.device_case("dense_matmul_binary.q4_k.f16.tokens2.outputs128", [] { + run_dense_matmul_binary_cpu_reference_case(GGML_TYPE_Q4_K, GGML_TYPE_F16, + "loom_libs:ggml_mul_mat_swiglu_f32_f32_lowtoken_dot", 2, 128); + }); + suite.device_case("vector_q8_publish.cpu_reference", [] { run_vector_q8_publish_cpu_reference_case(); }); + for (const ggml_glu_op op : + { GGML_GLU_OP_SWIGLU, GGML_GLU_OP_GEGLU, GGML_GLU_OP_REGLU, GGML_GLU_OP_GEGLU_ERF, GGML_GLU_OP_GEGLU_QUICK }) { + for (const int64_t outputs : { 128, 4096 }) { + suite.device_case(std::string("dense_matmul_swiglu.q4_k.q4_k.lowtoken.") + glu_op_name(op) + + ".tokens4.outputs" + std::to_string(outputs), + [op, outputs] { + run_dense_matmul_swiglu_cpu_reference_case( + GGML_TYPE_Q4_K, GGML_TYPE_Q4_K, + "loom_libs:ggml_mul_mat_swiglu_q4_q8_1_x4_lowtoken_dot", 4, outputs, op); + }); + } + for (const int64_t tokens : { 256, 512 }) { + if (tokens == 512) { + suite.device_case( + std::string("dense_matmul_swiglu.q4_k.q4_k.prefill.") + glu_op_name(op) + ".tokens512.outputs4096", + [op] { + run_dense_matmul_swiglu_cpu_reference_case( + GGML_TYPE_Q4_K, GGML_TYPE_Q4_K, + "loom_libs:ggml_mul_mat_swiglu_q4_k_f16_wmma_prefill_wave32", 512, 4096, op); + }); + } else { + for (const int64_t outputs : { 128, 4096 }) { + suite.device_case(std::string("dense_matmul_swiglu.q4_k.q4_k.prefill.") + glu_op_name(op) + + ".tokens256.outputs" + std::to_string(outputs), + [op, outputs] { + run_dense_matmul_swiglu_cpu_reference_case( + GGML_TYPE_Q4_K, GGML_TYPE_Q4_K, + "loom_libs:ggml_mul_mat_swiglu_q4_k_f16_wmma_prefill_wave32", 256, + outputs, op); + }); + } + } + } + } + for (const int64_t tokens : { 1, 2, 3, 5 }) { + suite.device_case( + "dense_matmul_swiglu.q4_k.q4_k.lowtoken.tokens" + std::to_string(tokens) + ".outputs4160", [tokens] { + run_dense_matmul_swiglu_cpu_reference_case(GGML_TYPE_Q4_K, GGML_TYPE_Q4_K, + "loom_libs:ggml_mul_mat_swiglu_q4_q8_1_x4_lowtoken_dot", + tokens, 4160); + }); + } + for (const int64_t tokens : { 1, 2, 5 }) { + suite.device_case( + "dense_matmul_swiglu.q4_k.q4_k.swiglu.tokens" + std::to_string(tokens) + ".outputs16384.input4096", + [tokens] { + run_dense_matmul_swiglu_cpu_reference_case(GGML_TYPE_Q4_K, GGML_TYPE_Q4_K, + "loom_libs:ggml_mul_mat_swiglu_q4_q8_1_x4_lowtoken_dot", + tokens, 16384, GGML_GLU_OP_SWIGLU, 4096); + }); + } + suite.device_case("dense_matmul_swiglu.q4_k.q4_k.reglu.tokens3.outputs16384.input4352", [] { + run_dense_matmul_swiglu_cpu_reference_case(GGML_TYPE_Q4_K, GGML_TYPE_Q4_K, + "loom_libs:ggml_mul_mat_swiglu_q4_q8_1_x4_lowtoken_dot", 3, 16384, + GGML_GLU_OP_REGLU, 4352); + }); + suite.device_case("dense_matmul_binary.q4_k.q4_k.lowtoken.tokens4.outputs128", [] { + run_dense_matmul_binary_cpu_reference_case(GGML_TYPE_Q4_K, GGML_TYPE_Q4_K, + "loom_libs:ggml_mul_mat_swiglu_q4_q8_1_x4_lowtoken_dot", 4, 128); + }); + suite.device_case("dense_matmul_binary.q4_k.q4_k.prefill.tokens256.outputs128", [] { + run_dense_matmul_binary_cpu_reference_case( + GGML_TYPE_Q4_K, GGML_TYPE_Q4_K, "loom_libs:ggml_mul_mat_swiglu_q4_k_f16_wmma_prefill_wave32", 256, 128); + }); + suite.device_case("dense_matmul_binary.q4_k.q4_k.prefill.tokens256.outputs4096", [] { + run_dense_matmul_binary_cpu_reference_case( + GGML_TYPE_Q4_K, GGML_TYPE_Q4_K, "loom_libs:ggml_mul_mat_swiglu_q4_k_f16_wmma_prefill_wave32", 256, 4096); + }); + suite.device_case("dense_matmul_binary.q4_k.q4_k.prefill.tokens1024.outputs2048", [] { + run_dense_matmul_binary_cpu_reference_case( + GGML_TYPE_Q4_K, GGML_TYPE_Q4_K, "loom_libs:ggml_mul_mat_swiglu_q4_k_f16_wmma_prefill_wave32", 1024, 2048); + }); + suite.device_case("dense_matmul_postops.f16.bias.tokens33.outputs256", [] { + run_dense_matmul_postops_cpu_reference_case(GGML_TYPE_F16, + "loom_libs:ggml_mul_mat_tiled_input_f32_bias_residual_publish_f32", + 33, 256, true, false, false); + }); + suite.device_case("dense_matmul_postops.f16.residual.tokens33.outputs256", [] { + run_dense_matmul_postops_cpu_reference_case(GGML_TYPE_F16, + "loom_libs:ggml_mul_mat_tiled_input_f32_bias_residual_publish_f32", + 33, 256, false, true, false); + }); + suite.device_case("dense_matmul_postops.f16.bias_residual.tokens33.outputs256", [] { + run_dense_matmul_postops_cpu_reference_case(GGML_TYPE_F16, + "loom_libs:ggml_mul_mat_tiled_input_f32_bias_residual_publish_f32", + 33, 256, true, true, false); + }); + suite.device_case("dense_matmul_postops.f16.residual.alias.tokens33.outputs256", [] { + run_dense_matmul_postops_cpu_reference_case(GGML_TYPE_F16, + "loom_libs:ggml_mul_mat_tiled_input_f32_bias_residual_publish_f32", + 33, 256, false, true, true); + }); + suite.device_case("dense_matmul_postops.f16.bias_residual.alias.tokens33.outputs256", [] { + run_dense_matmul_postops_cpu_reference_case(GGML_TYPE_F16, + "loom_libs:ggml_mul_mat_tiled_input_f32_bias_residual_publish_f32", + 33, 256, true, true, true); + }); + for (const ggml_type type : { GGML_TYPE_Q4_K, GGML_TYPE_Q6_K }) { + const int64_t input_size = type == GGML_TYPE_Q4_K ? 20480 : 8192; + suite.device_case( + "dense_matmul." + type_name(type) + ".skinny.tokens5.outputs4096.input" + std::to_string(input_size), + [type, input_size] { + run_dense_matmul_cpu_reference_case(type, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 5, + 4096, input_size); + }); + suite.device_case("dense_matmul_postops." + type_name(type) + ".add.tokens5.outputs4160", [type, input_size] { + run_dense_matmul_postops_cpu_reference_case( + type, "loom_libs:ggml_mul_mat_skinny_input_f32_bias_residual_publish_f32", 5, 4160, false, true, false, + input_size + 256); + }); + suite.device_case("dense_matmul_postops." + type_name(type) + ".add.alias.tokens5.outputs4224", + [type, input_size] { + run_dense_matmul_postops_cpu_reference_case( + type, "loom_libs:ggml_mul_mat_skinny_input_f32_bias_residual_publish_f32", 5, 4224, + false, true, true, input_size + 256); + }); + } + suite.device_case("dense_matmul.q4_k.skinny.tokens5.outputs4096.input8192", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_Q4_K, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 5, + 4096, 8192); + }); + for (const int64_t input_size : { 6144, 6400, 6656 }) { + suite.device_case( + "dense_matmul.q4_k.skinny.tokens5.outputs5120.input" + std::to_string(input_size), [input_size] { + run_dense_matmul_cpu_reference_case( + GGML_TYPE_Q4_K, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 5, 5120, input_size); + }); + suite.device_case("dense_matmul_postops.q4_k.add.alias.tokens5.outputs5120.input" + std::to_string(input_size), + [input_size] { + run_dense_matmul_postops_cpu_reference_case( + GGML_TYPE_Q4_K, "loom_libs:ggml_mul_mat_skinny_input_f32_bias_residual_publish_f32", + 5, 5120, false, true, true, input_size); + }); + } + suite.device_case("dense_matmul.q4_k.skinny.tokens4.outputs4096.input4096", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_Q4_K, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 4, + 4096, 4096); + }); + suite.device_case("dense_matmul_postops.q4_k.add.tokens4.outputs4160.input6400", [] { + run_dense_matmul_postops_cpu_reference_case(GGML_TYPE_Q4_K, + "loom_libs:ggml_mul_mat_skinny_input_f32_bias_residual_publish_f32", + 4, 4160, false, true, false, 6400); + }); + suite.device_case("dense_matmul_postops.q4_k.add.alias.tokens4.outputs4224.input6656", [] { + run_dense_matmul_postops_cpu_reference_case(GGML_TYPE_Q4_K, + "loom_libs:ggml_mul_mat_skinny_input_f32_bias_residual_publish_f32", + 4, 4224, false, true, true, 6656); + }); + suite.device_case("endpoint_rmsnorm_q6k_q8.cpu_reference", + [] { run_endpoint_rmsnorm_q6k_q8_cpu_reference_case(); }); + suite.device_case("dense_matmul.q6_k.skinny.tokens5.outputs4096.input16384", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_Q6_K, "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", 5, + 4096, 16384); + }); + suite.device_case("dense_matmul_postops.q6_k.add.tokens5.outputs4160.input16640", [] { + run_dense_matmul_postops_cpu_reference_case(GGML_TYPE_Q6_K, + "loom_libs:ggml_mul_mat_skinny_input_f32_bias_residual_publish_f32", + 5, 4160, false, true, false, 16640); + }); + suite.device_case("dense_matmul_postops.q6_k.add.alias.tokens5.outputs4224.input16640", [] { + run_dense_matmul_postops_cpu_reference_case(GGML_TYPE_Q6_K, + "loom_libs:ggml_mul_mat_skinny_input_f32_bias_residual_publish_f32", + 5, 4224, false, true, true, 16640); + }); + for (const int64_t tokens : { 1, 3, 5 }) { + suite.device_case( + "dense_matmul.q4_k.contiguous.tokens" + std::to_string(tokens) + ".outputs256.input5120", [tokens] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_Q4_K, + tokens == 1 ? "loom_libs:ggml_mul_mat_vector_f32_f32" : + "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", + tokens, 256, 5120, true); + }); + suite.device_case( + "dense_matmul.q6_k.contiguous.tokens" + std::to_string(tokens) + ".outputs256.input5120", [tokens] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_Q6_K, + tokens == 1 ? "loom_libs:ggml_mul_mat_vector_f32_f32" : + "loom_libs:ggml_mul_mat_skinny_input_f32_publish_f32", + tokens, 256, 5120, true, true); + }); + } + suite.device_case("dense_matmul.q6_k.vector_q6.tokens1.outputs256.input5120", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_Q6_K, "loom_libs:ggml_mul_mat_vector_q6_f32_f32", 1, 256, 5120, + true); + }); + suite.device_case("dense_matmul.q6_k.vector_q6.tokens1.outputs320.input512", [] { + run_dense_matmul_cpu_reference_case(GGML_TYPE_Q6_K, "loom_libs:ggml_mul_mat_vector_q6_f32_f32", 1, 320, 512, + true); + }); +} + +static void register_rmsnorm_and_scheduling_cases(Suite & suite) { + suite.host_case("support.lfm_head_rope.availability", [] { run_lfm_head_rope_availability_checks(); }); + suite.host_case("gemma_post_proj.matcher_negatives", [] { run_gemma_post_proj_matcher_negatives(); }); + for (const int64_t token_count : { 1, 3, 64 }) { + suite.device_case("gemma_post_proj.tokens" + std::to_string(token_count), + [token_count] { run_gemma_post_proj_cpu_reference_case(token_count); }); + } + suite.device_case("rmsnorm.default", [] { run_rmsnorm_cpu_reference_case(); }); + suite.device_case("rmsnorm.hidden3840.tokens18", [] { run_rmsnorm_cpu_reference_case(3840, 18, 1.0e-6f); }); + for (const int64_t token_count : { 2, 23 }) { + suite.device_case("strided_rmsnorm.tokens" + std::to_string(token_count), + [token_count] { run_strided_rmsnorm_cpu_reference_case(token_count); }); + } + suite.device_case("strided_rmsnorm_binary.minicpm_pp512_witness", + [] { run_strided_rmsnorm_binary_cpu_reference_case(); }); + suite.device_case("rmsnorm_binary_add.cpu_reference", [] { run_rmsnorm_binary_add_cpu_reference_case(); }); + suite.device_case("rmsnorm_gate.silu.hidden128.tokens1.stride7", + [] { run_rmsnorm_gate_cpu_reference_case(128, 1, 7, GGML_UNARY_OP_SILU); }); + suite.device_case("rmsnorm_gate.gelu.hidden256.tokens3.stride3", + [] { run_rmsnorm_gate_cpu_reference_case(256, 3, 3, GGML_UNARY_OP_GELU); }); + suite.device_case("rmsnorm_gate.relu.hidden512.tokens3.stride5", + [] { run_rmsnorm_gate_cpu_reference_case(512, 3, 5, GGML_UNARY_OP_RELU); }); + suite.device_case("rmsnorm_gate.gelu_erf.hidden1024.tokens1.stride4", + [] { run_rmsnorm_gate_cpu_reference_case(1024, 1, 4, GGML_UNARY_OP_GELU_ERF); }); + suite.device_case("rmsnorm_gate.silu.hidden128.tokens48.stride512", + [] { run_rmsnorm_gate_cpu_reference_case(128, 48, 512, GGML_UNARY_OP_SILU); }); + suite.device_case("rmsnorm_mul.hidden3840.tokens18", [] { run_rmsnorm_mul_cpu_reference_case(3840, 18, 1.0e-6f); }); + suite.device_case("gemma_scaled_rmsnorm.cpu_reference", [] { run_gemma_scaled_rmsnorm_cpu_reference_case(); }); + suite.device_case("gemma_scaled_rmsnorm_mul.cpu_reference", + [] { run_gemma_scaled_rmsnorm_mul_cpu_reference_case(); }); + suite.device_case("scheduled.scale_rmsnorm", [] { run_scheduled_hrx_scale_rmsnorm_case(); }); + suite.device_case("scheduled.rmsnorm_scale", [] { run_scheduled_hrx_rmsnorm_scale_case(); }); + suite.device_case("scheduled.rmsnorm_view_scale", [] { run_scheduled_hrx_rmsnorm_view_scale_case(); }); + suite.device_case("external_view_input_rmsnorm", [] { run_external_view_input_rmsnorm_case(); }); + suite.device_case("split_local_view_alias_import", [] { run_split_local_view_alias_import_case(); }); + suite.device_case("rope_set_rows.cpu_reference", [] { run_rope_set_rows_cpu_reference_case(); }); + for (const int64_t head_count : { 8, 32 }) { + for (const int64_t token_count : { 1, 64 }) { + suite.device_case( + "lfm_head_rope.query.heads" + std::to_string(head_count) + ".tokens" + std::to_string(token_count), + [head_count, token_count] { run_lfm_head_rope_cpu_reference_case(head_count, token_count, false); }); + } + } + suite.device_case("lfm_head_rope.query.heads7.tokens3", [] { run_lfm_head_rope_cpu_reference_case(7, 3, false); }); + for (const int64_t token_count : { 1, 64 }) { + suite.device_case("lfm_head_rope.cache.heads8.tokens" + std::to_string(token_count), + [token_count] { run_lfm_head_rope_cpu_reference_case(8, token_count, true); }); + } + suite.device_case("attention_postprocess.cpu_reference", [] { run_attention_postprocess_cpu_reference_case(); }); + suite.device_case("routed_moe.q4_k.next_rmsnorm", [] { run_routed_moe_cpu_reference_case(GGML_TYPE_Q4_K, true); }); + suite.device_case("flash_attention.asymmetric", [] { run_asymmetric_flash_attention_cpu_reference_case(); }); + suite.device_case("routed_moe.q6_k", [] { run_routed_moe_cpu_reference_case(GGML_TYPE_Q6_K, false); }); + suite.device_case("decode_attention_qkv.scheduling", [] { run_decode_attention_qkv_scheduling_case(); }); + suite.device_case("decode_routed_moe.q4_k", [] { run_decode_routed_moe_scheduling_case(GGML_TYPE_Q4_K); }); + suite.device_case("decode_routed_moe.q6_k", [] { run_decode_routed_moe_scheduling_case(GGML_TYPE_Q6_K); }); + suite.device_case("decode_routed_moe.q6_k.alias_gate_input", + [] { run_decode_routed_moe_scheduling_case(GGML_TYPE_Q6_K, true); }); + suite.device_case("rmsnorm_mul.hidden256.tokens1", [] { run_rmsnorm_mul_case(256, 1); }); + suite.device_case("rmsnorm_mul.hidden256.tokens4", [] { run_rmsnorm_mul_case(256, 4); }); + suite.device_case("rmsnorm_mul.hidden2048.tokens1", [] { run_rmsnorm_mul_case(2048, 1); }); + suite.device_case("router_projection.tokens4", [] { run_router_projection_case(4); }); + suite.device_case("router_top8.tokens4", [] { run_router_top8_case(4); }); + suite.device_case("qwen_flash_attention", [] { run_qwen_flash_attention_case(); }); +} + +static void register_flash_attention_cases(Suite & suite) { + suite.device_case("common_flash_attention.prefill.head128.query256.kv256.exact", [] { + run_common_flash_attention_cpu_reference_case(128, 256, 256, 1.0f / std::sqrt(128.0f), + "loom_libs:ggml_flash_attention_f32_f16_wmma"); + }); + for (int64_t head_size = 64; head_size <= 512; head_size += 64) { + suite.device_case("common_flash_attention.decode.head" + std::to_string(head_size) + ".query2.kv65", + [head_size] { + run_common_flash_attention_cpu_reference_case( + head_size, 2, 65, 1.0f / std::sqrt(static_cast(head_size)), + "loom_libs:ggml_flash_attention_decode_split_f32_f16_wmma_next_q8"); + }); + suite.device_case( + "common_flash_attention.prefill.head" + std::to_string(head_size) + ".query16.kv64", [head_size] { + run_common_flash_attention_cpu_reference_case(head_size, 16, 64, + 1.0f / std::sqrt(static_cast(head_size)), + "loom_libs:ggml_flash_attention_f32_f16_wmma"); + }); + suite.device_case( + "common_flash_attention.prefill.head" + std::to_string(head_size) + ".query17.kv65", [head_size] { + run_common_flash_attention_cpu_reference_case(head_size, 17, 65, + 1.0f / std::sqrt(static_cast(head_size)), + "loom_libs:ggml_flash_attention_f32_f16_wmma"); + }); + } + suite.device_case("common_flash_attention.decode.head96.query2.kv4.capacity320", [] { + run_common_flash_attention_cpu_reference_case( + 96, 2, 4, 1.0f / std::sqrt(96.0f), "loom_libs:ggml_flash_attention_decode_split_f32_f16_wmma_next_q8", 320); + }); + suite.device_case("common_flash_attention.decode.head320.query2.kv64.capacity64", [] { + run_common_flash_attention_cpu_reference_case( + 320, 2, 64, 1.0f / std::sqrt(320.0f), "loom_libs:ggml_flash_attention_decode_split_f32_f16_wmma_next_q8", + 64); + }); + suite.device_case("qwen_decode_split_flash_attention.query1.kv512", + [] { run_qwen_decode_split_flash_attention_scheduling_case(1, 512); }); + suite.device_case("qwen_decode_split_flash_attention.query4.kv513", + [] { run_qwen_decode_split_flash_attention_scheduling_case(4, 513); }); + suite.device_case("qwen_decode_split_flash_attention.query1.kv512.masked", + [] { run_qwen_decode_split_flash_attention_scheduling_case(1, 512, true); }); + suite.device_case("qwen_decode_split_flash_attention.query1.kv512.masked.head96.value64", [] { + run_qwen_decode_split_flash_attention_scheduling_case(1, 512, true, 96, 64, 1.0f / std::sqrt(96.0f)); + }); + suite.device_case("qwen_decode_split_flash_attention.query4.kv513.masked.head96.value64", [] { + run_qwen_decode_split_flash_attention_scheduling_case(4, 513, true, 96, 64, 1.0f / std::sqrt(96.0f)); + }); + suite.device_case("qwen_decode_attention_output_next_q8.no_selectors", + [] { run_qwen_decode_attention_output_next_q8_scheduling_case(false); }); + suite.device_case("qwen_decode_attention_output_next_q8.selectors", + [] { run_qwen_decode_attention_output_next_q8_scheduling_case(true); }); + suite.device_case("qwen_full_cache_prefill_flash_attention.tokens512", + [] { run_qwen_full_cache_prefill_flash_attention_scheduling_case(512); }); +} + +static void register_hrx_ops_cases(Suite & suite) { + register_basic_ops_cases(suite); + register_dense_matmul_cases(suite); + register_rmsnorm_and_scheduling_cases(suite); + register_flash_attention_cases(suite); +} + +int main(int argc, char ** argv) { + Suite suite(test_runner::Config::with_prefix("HRX test", "hrx", "GGML_HRX_TEST")); + register_hrx_ops_cases(suite); + + test_runner::Options options; + if (!suite.parse_options(argc, argv, options)) { + suite.print_usage(argc > 0 ? argv[0] : "test-hrx-ops"); + return 2; + } + if (options.help) { + return suite.run(options, false); + } + + const bool run_vector_cap_experiment = std::getenv("GGML_HRX_RUN_VECTOR_CAP_EXPERIMENT") != nullptr; + const bool run_vector_publish_test = std::getenv("GGML_HRX_RUN_VECTOR_PUBLISH_TEST") != nullptr; + const bool run_vector_postops_test = std::getenv("GGML_HRX_RUN_VECTOR_POSTOPS_TEST") != nullptr; + + const bool has_device = ggml_backend_hrx_get_device_count() != 0; + if (!options.runner_control && run_vector_cap_experiment) { + run_vector_output_capacity_experiment(); + return 0; + } + if (!options.runner_control && run_vector_publish_test) { + run_vector_q8_publish_cpu_reference_case(); + return 0; + } + if (!options.runner_control && run_vector_postops_test) { + run_vector_postops_cpu_reference_case(); + return 0; + } + + return suite.run(options, has_device); +} diff --git a/tests/test-llama-archs.cpp b/tests/test-llama-archs.cpp index 4336e4e13d4f..8e3642c51455 100644 --- a/tests/test-llama-archs.cpp +++ b/tests/test-llama-archs.cpp @@ -203,8 +203,40 @@ static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe) { // MSA requires one indexer head per GQA (KV) head, unlike the DSA archs where the // indexer head count is independent of the main attention head count. + if (arch == LLM_ARCH_QWEN4EXP) { + ms.add_kv(LLM_KV_HYPER_CONNECTION_COUNT, uint32_t(4)); + ms.add_kv(LLM_KV_HYPER_CONNECTION_LOW_RANK, uint32_t(8)); + // without this the QSA layers fall back to dense and go uncovered + ms.add_kv(LLM_KV_ATTENTION_COMPRESS_RATIOS, std::vector(n_layer, 4)); + + // has_cell_ext() needs ple_n_heads here: the indexer cache serializes no ext without it + const uint32_t ple_ngram_size = 3; + const uint32_t ple_heads_per_ngram = 2; + const uint32_t ple_n_heads = (ple_ngram_size - 1)*ple_heads_per_ngram; + GGML_ASSERT(n_embd % ple_n_heads == 0); + const uint32_t ple_head_dim = n_embd/ple_n_heads; + + std::vector ple_head_offsets(ple_n_heads); + std::vector ple_head_vocab_sizes(ple_n_heads, n_vocab); + for (uint32_t h = 0; h < ple_n_heads; h++) { + ple_head_offsets[h] = uint64_t(h)*n_vocab; + } + + // the PLE history lives in the recurrent cache, so it must sit on a linear attention layer + ms.add_kv(LLM_KV_PLE_LAYERS, std::vector({ 0 })); + ms.add_kv(LLM_KV_PLE_NGRAM_SIZE, ple_ngram_size); + ms.add_kv(LLM_KV_PLE_HEADS_PER_NGRAM, ple_heads_per_ngram); + ms.add_kv(LLM_KV_PLE_CONV_KERNEL, uint32_t(4)); + ms.add_kv(LLM_KV_PLE_EOS_TOKEN_ID, uint32_t(0)); + ms.add_kv(LLM_KV_EMBEDDING_LENGTH_PER_LAYER, ple_head_dim); + ms.add_kv(LLM_KV_PLE_LAYER_MULTIPLIERS, std::vector({ 1, 3, 5 })); + ms.add_kv(LLM_KV_PLE_HEAD_OFFSETS, ple_head_offsets); + ms.add_kv(LLM_KV_PLE_HEAD_VOCAB_SIZES, ple_head_vocab_sizes); + } + ms.add_kv(LLM_KV_ATTENTION_INDEXER_HEAD_COUNT, arch == LLM_ARCH_MINIMAX_M3 ? n_head : uint32_t(1)); - ms.add_kv(LLM_KV_ATTENTION_INDEXER_KEY_LENGTH, uint32_t(64)); + ms.add_kv(LLM_KV_ATTENTION_INDEXER_KEY_LENGTH, + arch == LLM_ARCH_QWEN4EXP ? n_embd_head : uint32_t(64)); ms.add_kv(LLM_KV_ATTENTION_INDEXER_TOP_K, uint32_t(8)); ms.add_kv(LLM_KV_ATTENTION_INDEXER_BLOCK_SIZE, uint32_t(4)); ms.add_kv(LLM_KV_ATTENTION_INDEXER_LOCAL_BLOCKS, uint32_t(1)); @@ -232,8 +264,8 @@ static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe) { ms.add_kv(LLM_KV_XIELU_ALPHA_P, 1.0f); ms.add_kv(LLM_KV_XIELU_BETA, 1.0f); ms.add_kv(LLM_KV_XIELU_EPS, 1.0e-7f); - ms.add_kv(LLM_KV_SSM_INNER_SIZE, arch == LLM_ARCH_QWEN3NEXT || arch == LLM_ARCH_QWEN35 || arch == LLM_ARCH_QWEN35MOE ? 256 : 2*n_embd); - ms.add_kv(LLM_KV_SSM_CONV_KERNEL, uint32_t(4)); + ms.add_kv(LLM_KV_SSM_INNER_SIZE, arch == LLM_ARCH_QWEN3NEXT || arch == LLM_ARCH_QWEN35 || arch == LLM_ARCH_QWEN35MOE || arch == LLM_ARCH_QWEN4EXP ? 256 : 2*n_embd); + ms.add_kv(LLM_KV_SSM_CONV_KERNEL, uint32_t(arch == LLM_ARCH_ZAYA ? 2 : 4)); // zaya: 2-tap CCA convs ms.add_kv(LLM_KV_SSM_STATE_SIZE, uint32_t(128)); ms.add_kv(LLM_KV_SSM_TIME_STEP_RANK, n_head); ms.add_kv(LLM_KV_SSM_GROUP_COUNT, arch == LLM_ARCH_PLAMO2 ? 0 : uint32_t(2)); @@ -338,6 +370,7 @@ static bool moe_mandatory(const llm_arch arch) { case LLM_ARCH_QWEN3NEXT: case LLM_ARCH_QWEN3VLMOE: case LLM_ARCH_QWEN35MOE: + case LLM_ARCH_QWEN4EXP: case LLM_ARCH_PHIMOE: case LLM_ARCH_DBRX: case LLM_ARCH_OLMOE: @@ -371,6 +404,7 @@ static bool moe_mandatory(const llm_arch arch) { case LLM_ARCH_MISTRAL4: case LLM_ARCH_MELLUM: case LLM_ARCH_LAGUNA: + case LLM_ARCH_ZAYA: return true; default: return false; @@ -396,7 +430,7 @@ static bool moe_implemented(const llm_arch arch) { } static bool arch_supported(const llm_arch arch) { - if (arch == LLM_ARCH_CLIP || arch == LLM_ARCH_GPTJ || arch == LLM_ARCH_UNKNOWN) { + if (arch == LLM_ARCH_CLIP || arch == LLM_ARCH_UNKNOWN) { return false; // These models don't have usable implementations. } if (arch == LLM_ARCH_CHAMELEON) { @@ -430,7 +464,7 @@ static bool arch_supported(const llm_arch arch) { // FIXME: these hit scheduler/view-backed-output issues with WebGPU on CI. #ifdef GGML_USE_WEBGPU - if (arch == LLM_ARCH_DEEPSEEK32 || arch == LLM_ARCH_GLM_DSA || arch == LLM_ARCH_MINIMAX_M3) { + if (arch == LLM_ARCH_DEEPSEEK32 || arch == LLM_ARCH_GLM_DSA || arch == LLM_ARCH_MINIMAX_M3 || arch == LLM_ARCH_QWEN4EXP) { return false; } #endif // GGML_USE_WEBGPU @@ -599,6 +633,7 @@ static int test_backends(const llm_arch target_arch, const size_t seed, const gg std::string status_nmse = "\033[1;33mSKIP\033[0m"; std::string status_roundtrip = "\033[1;33mSKIP\033[0m"; char nmse_str[12] = {0}; + bool skip = !arch_supported(arch) || (dc.split_mode == LLAMA_SPLIT_MODE_TENSOR && dc.devs.empty()); if (!skip) { if (logits_cpu.empty()) { diff --git a/tests/test-quantize-fns.cpp b/tests/test-quantize-fns.cpp index 9510ac14ce00..badca1bd5a77 100644 --- a/tests/test-quantize-fns.cpp +++ b/tests/test-quantize-fns.cpp @@ -159,6 +159,8 @@ static int test_vec_dot_q(bool verbose) { type == GGML_TYPE_TQ1_0 ? MAX_QUANTIZATION_TOTAL_ERROR_TERNARY : type == GGML_TYPE_TQ2_0 ? MAX_QUANTIZATION_TOTAL_ERROR_TERNARY : type == GGML_TYPE_Q2_0 ? MAX_QUANTIZATION_TOTAL_ERROR_TERNARY : + type == GGML_TYPE_PQ2_0 ? MAX_QUANTIZATION_TOTAL_ERROR_TERNARY : + type == GGML_TYPE_PTQ1_0 ? MAX_QUANTIZATION_TOTAL_ERROR_TERNARY : type == GGML_TYPE_Q2_K ? MAX_QUANTIZATION_TOTAL_ERROR_2BITS : type == GGML_TYPE_IQ2_S ? MAX_QUANTIZATION_TOTAL_ERROR_2BITS : type == GGML_TYPE_Q3_K ? MAX_QUANTIZATION_TOTAL_ERROR_3BITS : @@ -184,7 +186,8 @@ static int test_vec_dot_q(bool verbose) { ? MAX_DOT_PRODUCT_ERROR_LOWBIT : type == GGML_TYPE_Q1_0 ? MAX_DOT_PRODUCT_ERROR_BINARY - : type == GGML_TYPE_TQ1_0 || type == GGML_TYPE_TQ2_0 || type == GGML_TYPE_Q2_0 + : type == GGML_TYPE_TQ1_0 || type == GGML_TYPE_TQ2_0 || type == GGML_TYPE_Q2_0 || + type == GGML_TYPE_PQ2_0 || type == GGML_TYPE_PTQ1_0 ? MAX_DOT_PRODUCT_ERROR_TERNARY : type == GGML_TYPE_NVFP4 ? MAX_DOT_PRODUCT_ERROR_FP4 diff --git a/tests/testing_suite.h b/tests/testing_suite.h new file mode 100644 index 000000000000..ffa5fa0e66cb --- /dev/null +++ b/tests/testing_suite.h @@ -0,0 +1,548 @@ +#pragma once + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#if !defined(_WIN32) +#include +#include +#include +#endif + +namespace test_runner { + +struct Case { + std::string name; + std::function run; + bool requires_device = true; +}; + +enum class OutputMode { + Failed, + All, + None, +}; + +struct Config { + std::string suite_name = "test"; + std::string list_flag = "--test-list"; + std::string case_flag = "--test-case"; + std::string filter_flag = "--test-filter"; + std::string jobs_flag = "--test-jobs"; + std::string output_flag = "--test-output"; + std::string env_jobs; + std::string env_filter; + std::string env_output; + size_t max_default_jobs = 4; + + static Config with_prefix(std::string suite_name, + const std::string & flag_prefix, + const std::string & env_prefix, + size_t max_default_jobs = 4) { + Config config; + config.suite_name = std::move(suite_name); + config.list_flag = "--" + flag_prefix + "-test-list"; + config.case_flag = "--" + flag_prefix + "-test-case"; + config.filter_flag = "--" + flag_prefix + "-test-filter"; + config.jobs_flag = "--" + flag_prefix + "-test-jobs"; + config.output_flag = "--" + flag_prefix + "-test-output"; + config.env_jobs = env_prefix + "_JOBS"; + config.env_filter = env_prefix + "_FILTER"; + config.env_output = env_prefix + "_OUTPUT"; + config.max_default_jobs = max_default_jobs; + return config; + } +}; + +struct Options { + std::string executable; + std::string single_case; + std::string filter; + size_t jobs = 0; + bool help = false; + bool list = false; + bool runner_control = false; + OutputMode output = OutputMode::Failed; +}; + +struct Result { + std::string name; + std::string output; + int exit_code = -1; + int signal_code = 0; + int64_t elapsed_ms = 0; + bool passed = false; +}; + +struct Selection { + std::vector cases; + size_t matched = 0; + size_t skipped_device = 0; +}; + +inline void add_case(std::vector & cases, + std::string name, + std::function run, + bool requires_device = true) { + cases.push_back({ std::move(name), std::move(run), requires_device }); +} + +inline size_t default_jobs(const Config & config) { + const unsigned int hardware_jobs = std::thread::hardware_concurrency(); + if (hardware_jobs == 0) { + return 1; + } + return std::min(config.max_default_jobs, hardware_jobs); +} + +inline bool parse_positive_size(const char * value, size_t & out) { + if (value == nullptr || *value == '\0') { + return false; + } + char * end = nullptr; + errno = 0; + const unsigned long parsed = std::strtoul(value, &end, 10); + if (errno != 0 || end == value || *end != '\0' || parsed == 0) { + return false; + } + out = static_cast(parsed); + return true; +} + +inline bool parse_output_mode(const char * value, OutputMode & out) { + if (value == nullptr) { + return true; + } + const std::string mode = value; + if (mode == "failed" || mode == "fail") { + out = OutputMode::Failed; + return true; + } + if (mode == "all") { + out = OutputMode::All; + return true; + } + if (mode == "none") { + out = OutputMode::None; + return true; + } + return false; +} + +inline void print_usage(const Config & config, const char * executable) { + std::fprintf(stderr, + "usage: %s [%s N] [%s REGEX] [%s MODE] [%s]\n" + " %s %s NAME\n" + "\n" + "environment: %s, %s, %s\n" + "output modes: failed, all, none\n", + executable, config.jobs_flag.c_str(), config.filter_flag.c_str(), config.output_flag.c_str(), + config.list_flag.c_str(), executable, config.case_flag.c_str(), config.env_jobs.c_str(), + config.env_filter.c_str(), config.env_output.c_str()); +} + +inline bool parse_options(int argc, char ** argv, const Config & config, Options & options) { + options.executable = argc > 0 ? argv[0] : config.suite_name.c_str(); + if (!config.env_filter.empty()) { + if (const char * env_filter = std::getenv(config.env_filter.c_str())) { + options.filter = env_filter; + options.runner_control = true; + } + } + if (!config.env_jobs.empty()) { + if (const char * env_jobs = std::getenv(config.env_jobs.c_str())) { + if (!parse_positive_size(env_jobs, options.jobs)) { + std::fprintf(stderr, "invalid %s value: %s\n", config.env_jobs.c_str(), env_jobs); + return false; + } + options.runner_control = true; + } + } + if (!config.env_output.empty()) { + if (const char * env_output = std::getenv(config.env_output.c_str())) { + if (!parse_output_mode(env_output, options.output)) { + std::fprintf(stderr, "invalid %s value: %s\n", config.env_output.c_str(), env_output); + return false; + } + options.runner_control = true; + } + } + + for (int i = 1; i < argc; ++i) { + const std::string arg = argv[i]; + if (arg == config.list_flag) { + options.list = true; + options.runner_control = true; + } else if (arg == config.case_flag) { + if (++i >= argc) { + std::fprintf(stderr, "%s requires a case name\n", config.case_flag.c_str()); + return false; + } + options.single_case = argv[i]; + options.runner_control = true; + } else if (arg == config.filter_flag) { + if (++i >= argc) { + std::fprintf(stderr, "%s requires a regex\n", config.filter_flag.c_str()); + return false; + } + options.filter = argv[i]; + options.runner_control = true; + } else if (arg == config.jobs_flag) { + if (++i >= argc || !parse_positive_size(argv[i], options.jobs)) { + std::fprintf(stderr, "%s requires a positive integer\n", config.jobs_flag.c_str()); + return false; + } + options.runner_control = true; + } else if (arg == config.output_flag) { + if (++i >= argc || !parse_output_mode(argv[i], options.output)) { + std::fprintf(stderr, "%s requires failed, all, or none\n", config.output_flag.c_str()); + return false; + } + options.runner_control = true; + } else if (arg == "--help" || arg == "-h") { + options.help = true; + options.runner_control = true; + } else { + std::fprintf(stderr, "unknown argument: %s\n", arg.c_str()); + return false; + } + } + + if (options.jobs == 0) { + options.jobs = default_jobs(config); + } + return true; +} + +inline const Case * find_case(const std::vector & cases, const std::string & name) { + for (const Case & test_case : cases) { + if (test_case.name == name) { + return &test_case; + } + } + return nullptr; +} + +inline int run_single_case_in_process(const Config & config, + const std::vector & cases, + const std::string & name, + bool has_device) { + const Case * test_case = find_case(cases, name); + if (test_case == nullptr) { + std::fprintf(stderr, "unknown test case: %s\n", name.c_str()); + return 2; + } + if (test_case->requires_device && !has_device) { + std::fprintf(stderr, "test skipped: no device available for %s\n", config.suite_name.c_str()); + return 0; + } + test_case->run(); + return 0; +} + +inline Result run_single_case_child(const Config & config, const Options & options, const Case & test_case) { + Result result; + result.name = test_case.name; + + const auto start = std::chrono::steady_clock::now(); + +#if defined(_WIN32) + result.output = "process-isolated test runner is not supported on this platform\n"; + result.exit_code = 127; + result.elapsed_ms = + std::chrono::duration_cast(std::chrono::steady_clock::now() - start).count(); + return result; +#else + int pipe_fds[2] = { -1, -1 }; + if (pipe(pipe_fds) != 0) { + result.output = "failed to create output pipe\n"; + result.exit_code = 127; + return result; + } + + std::fflush(nullptr); + const pid_t pid = fork(); + if (pid == 0) { + close(pipe_fds[0]); + dup2(pipe_fds[1], STDOUT_FILENO); + dup2(pipe_fds[1], STDERR_FILENO); + close(pipe_fds[1]); + + std::vector child_argv; + child_argv.push_back(const_cast(options.executable.c_str())); + child_argv.push_back(const_cast(config.case_flag.c_str())); + child_argv.push_back(const_cast(test_case.name.c_str())); + child_argv.push_back(nullptr); + execvp(child_argv[0], child_argv.data()); + std::fprintf(stderr, "failed to exec %s: %s\n", child_argv[0], std::strerror(errno)); + _exit(127); + } + + close(pipe_fds[1]); + if (pid < 0) { + close(pipe_fds[0]); + result.output = "failed to fork child test process\n"; + result.exit_code = 127; + return result; + } + + char buffer[4096]; + for (;;) { + const ssize_t bytes = read(pipe_fds[0], buffer, sizeof(buffer)); + if (bytes > 0) { + result.output.append(buffer, static_cast(bytes)); + continue; + } + if (bytes == 0) { + break; + } + if (errno == EINTR) { + continue; + } + result.output += "failed to read child output\n"; + break; + } + close(pipe_fds[0]); + + int status = 0; + while (waitpid(pid, &status, 0) < 0) { + if (errno == EINTR) { + continue; + } + result.output += "failed to wait for child test process\n"; + result.exit_code = 127; + return result; + } + + if (WIFEXITED(status)) { + result.exit_code = WEXITSTATUS(status); + result.passed = result.exit_code == 0; + } else if (WIFSIGNALED(status)) { + result.signal_code = WTERMSIG(status); + result.passed = false; + } + + result.elapsed_ms = + std::chrono::duration_cast(std::chrono::steady_clock::now() - start).count(); + return result; +#endif +} + +inline bool should_print_output(OutputMode mode, const Result & result) { + if (mode == OutputMode::All) { + return true; + } + if (mode == OutputMode::None) { + return false; + } + return !result.passed; +} + +inline int run_cases_parallel(const Config & config, + const Options & options, + const std::vector & selected_cases) { + if (selected_cases.empty()) { + std::printf("No %s cases selected\n", config.suite_name.c_str()); + return 0; + } + + const size_t jobs = std::min(options.jobs, selected_cases.size()); + std::vector results(selected_cases.size()); + std::atomic next_index(0); + std::atomic finished(0); + std::mutex print_mutex; + const auto suite_start = std::chrono::steady_clock::now(); + + std::printf("Running %zu %s case(s) with %zu worker process(es)\n", selected_cases.size(), + config.suite_name.c_str(), jobs); + + auto worker = [&] { + for (;;) { + const size_t index = next_index.fetch_add(1); + if (index >= selected_cases.size()) { + return; + } + Result result = run_single_case_child(config, options, selected_cases[index]); + const size_t done = finished.fetch_add(1) + 1; + { + std::lock_guard lock(print_mutex); + std::printf("%s %4zu/%zu %-86s %8lld ms\n", result.passed ? "[PASS]" : "[FAIL]", done, + selected_cases.size(), result.name.c_str(), static_cast(result.elapsed_ms)); + if (!result.passed) { + if (result.signal_code != 0) { + std::printf(" terminated by signal %d\n", result.signal_code); + } else { + std::printf(" exit code %d\n", result.exit_code); + } + } + if (should_print_output(options.output, result) && !result.output.empty()) { + std::printf("----- output: %s -----\n%s", result.name.c_str(), result.output.c_str()); + if (result.output.back() != '\n') { + std::printf("\n"); + } + std::printf("----- end output: %s -----\n", result.name.c_str()); + } + std::fflush(stdout); + } + results[index] = std::move(result); + } + }; + + std::vector workers; + workers.reserve(jobs); + for (size_t i = 0; i < jobs; ++i) { + workers.emplace_back(worker); + } + for (std::thread & thread : workers) { + thread.join(); + } + + size_t passed = 0; + for (const Result & result : results) { + if (result.passed) { + ++passed; + } + } + const size_t failed = results.size() - passed; + const int64_t elapsed_ms = + std::chrono::duration_cast(std::chrono::steady_clock::now() - suite_start).count(); + + std::printf("\n%s summary: %zu passed, %zu failed, %zu total, %lld ms\n", config.suite_name.c_str(), passed, failed, + results.size(), static_cast(elapsed_ms)); + if (failed != 0) { + std::printf("Failed cases:\n"); + for (const Result & result : results) { + if (!result.passed) { + std::printf(" %s\n", result.name.c_str()); + } + } + } + return failed == 0 ? 0 : 1; +} + +inline bool select_cases(const std::vector & cases, + const std::string & filter_text, + bool has_device, + bool include_unavailable_device_cases, + Selection & selection) { + try { + std::regex filter; + const bool has_filter = !filter_text.empty(); + if (has_filter) { + filter = std::regex(filter_text); + } + for (const Case & test_case : cases) { + if (has_filter && !std::regex_search(test_case.name, filter)) { + continue; + } + ++selection.matched; + if (!has_device && test_case.requires_device && !include_unavailable_device_cases) { + ++selection.skipped_device; + continue; + } + selection.cases.push_back(test_case); + } + } catch (const std::regex_error & error) { + std::fprintf(stderr, "invalid test filter regex: %s\n", error.what()); + return false; + } + return true; +} + +inline void print_case_list(const std::vector & cases) { + for (const Case & test_case : cases) { + std::printf("[%s] %s\n", test_case.requires_device ? "device" : "host", test_case.name.c_str()); + } +} + +inline bool validate_case_names(const std::vector & cases) { + for (size_t i = 0; i < cases.size(); ++i) { + for (size_t j = i + 1; j < cases.size(); ++j) { + if (cases[i].name == cases[j].name) { + std::fprintf(stderr, "duplicate test case: %s\n", cases[i].name.c_str()); + return false; + } + } + } + return true; +} + +class Suite { + public: + explicit Suite(Config config) : config_(std::move(config)) {} + + template void device_case(std::string name, F && run) { + add_case(cases_, std::move(name), std::function(std::forward(run)), true); + } + + template void host_case(std::string name, F && run) { + add_case(cases_, std::move(name), std::function(std::forward(run)), false); + } + + bool parse_options(int argc, char ** argv, Options & options) const { + return test_runner::parse_options(argc, argv, config_, options); + } + + void print_usage(const char * executable) const { test_runner::print_usage(config_, executable); } + + int run(int argc, char ** argv, bool has_device) const { + Options options; + if (!parse_options(argc, argv, options)) { + print_usage(argc > 0 ? argv[0] : config_.suite_name.c_str()); + return 2; + } + return run(options, has_device); + } + + int run(const Options & options, bool has_device) const { + if (options.help) { + print_usage(options.executable.c_str()); + return 0; + } + if (!validate_case_names(cases_)) { + return 2; + } + if (!options.single_case.empty()) { + return run_single_case_in_process(config_, cases_, options.single_case, has_device); + } + + Selection selection; + if (!select_cases(cases_, options.filter, has_device, options.list, selection)) { + return 2; + } + + if (options.list) { + print_case_list(selection.cases); + return 0; + } + + if (selection.skipped_device != 0) { + std::fprintf(stderr, "test skipped: no device available for %s (%zu device case(s))\n", + config_.suite_name.c_str(), selection.skipped_device); + if (selection.cases.empty()) { + return 0; + } + } + + return run_cases_parallel(config_, options, selection.cases); + } + + private: + Config config_; + std::vector cases_; +}; + +} // namespace test_runner diff --git a/tools/cli/README.md b/tools/cli/README.md index bcddd05702bb..4f276b8e72e3 100644 --- a/tools/cli/README.md +++ b/tools/cli/README.md @@ -59,12 +59,14 @@ | `--mmap, --no-mmap` | DEPRECATED in favor of `--load-mode`: whether to memory-map model. (if mmap disabled, slower load but may reduce pageouts if not using mlock)
(env: LLAMA_ARG_MMAP) | | `-dio, --direct-io, -ndio, --no-direct-io` | DEPRECATED in favor of `--load-mode`: use DirectIO if available
(env: LLAMA_ARG_DIO) | | `-lm, --load-mode MODE` | model loading mode (default: mmap)
- none: no special loading mode
- mmap: memory-map model (if mmap disabled, slower load but may reduce pageouts if not using mlock)
- mlock: force system to keep model in RAM rather than swapping or compressing
- mmap+mlock: mmap + force system to keep model in RAM rather than swapping or compressing
- dio: use DirectIO if available

(env: LLAMA_ARG_LOAD_MODE) | +| `-lzm, --lazy-mode MODE` | on-demand reading of certain tensors, for example per-layer embeddings (default: auto)
- on: read the rows of such tensors from disk on demand instead of keeping them resident (requires mmap)
- auto: on, but only for tensors larger than 4 GiB
- off: always keep them resident
(env: LLAMA_ARG_LAZY_MODE) | | `--numa TYPE` | attempt optimizations that help on some NUMA systems
- distribute: spread execution evenly over all nodes
- isolate: only spawn threads on CPUs on the node that execution started on
- numactl: use the CPU map provided by numactl
if run without this previously, it is recommended to drop the system page cache before using this
see https://github.com/ggml-org/llama.cpp/issues/1437
(env: LLAMA_ARG_NUMA) | | `-dev, --device ` | comma-separated list of devices to use for offloading (none = don't offload)
use --list-devices to see a list of available devices
(env: LLAMA_ARG_DEVICE) | | `--list-devices` | print list of available devices and exit | | `-ot, --override-tensor =,...` | override tensor buffer type
(env: LLAMA_ARG_OVERRIDE_TENSOR) | | `-cmoe, --cpu-moe` | keep all Mixture of Experts (MoE) weights in the CPU
(env: LLAMA_ARG_CPU_MOE) | | `-ncmoe, --n-cpu-moe N` | keep the Mixture of Experts (MoE) weights of the first N layers in the CPU
(env: LLAMA_ARG_N_CPU_MOE) | +| `-ncffn, --n-cpu-ffn N` | keep the dense FFN weights of the first N layers in the CPU
(dense models; for MoE expert weights use --n-cpu-moe)
(env: LLAMA_ARG_N_CPU_FFN) | | `-ngl, --gpu-layers, --n-gpu-layers N` | max. number of layers to store in VRAM, either an exact number, 'auto', or 'all' (default: auto)
(env: LLAMA_ARG_N_GPU_LAYERS) | | `-sm, --split-mode {none,layer,row,tensor}` | how to split the model across multiple GPUs, one of:
- none: use one GPU only
- layer (default): split layers and KV across GPUs (pipelined)
- row: split weight across GPUs by rows (parallelized)
- tensor: split weights and KV across GPUs (parallelized, EXPERIMENTAL)
(env: LLAMA_ARG_SPLIT_MODE) | | `-ts, --tensor-split N0,N1,N2,...` | fraction of the model to offload to each GPU, comma-separated list of proportions, e.g. 3,1
(env: LLAMA_ARG_TENSOR_SPLIT) | @@ -156,7 +158,6 @@ | `-sysf, --system-prompt-file FNAME` | a file containing the system prompt (default: none) | | `-r, --reverse-prompt PROMPT` | halt generation at PROMPT, return control in interactive mode | | `-sp, --special` | special tokens output enabled (default: false) | -| `-cnv, --conversation, -no-cnv, --no-conversation` | whether to run in conversation mode:
- does not print special tokens and suffix/prefix
- interactive mode is also enabled
(default: auto enabled if chat template is available) | | `-st, --single-turn` | run conversation for a single turn only, then exit when done
will not be interactive if first turn is predefined with --prompt
(default: false) | | `-mli, --multiline-input` | allows you to write or paste multiple lines without ending each in '\' | | `--warmup, --no-warmup` | whether to perform warmup with an empty run (default: enabled) | diff --git a/tools/completion/README.md b/tools/completion/README.md index bce71d68d949..d9b672de5371 100644 --- a/tools/completion/README.md +++ b/tools/completion/README.md @@ -142,12 +142,14 @@ llama-completion.exe -m models\gemma-1.1-7b-it.Q4_K_M.gguf --ignore-eos -n -1 | `--mmap, --no-mmap` | DEPRECATED in favor of `--load-mode`: whether to memory-map model. (if mmap disabled, slower load but may reduce pageouts if not using mlock)
(env: LLAMA_ARG_MMAP) | | `-dio, --direct-io, -ndio, --no-direct-io` | DEPRECATED in favor of `--load-mode`: use DirectIO if available
(env: LLAMA_ARG_DIO) | | `-lm, --load-mode MODE` | model loading mode (default: mmap)
- none: no special loading mode
- mmap: memory-map model (if mmap disabled, slower load but may reduce pageouts if not using mlock)
- mlock: force system to keep model in RAM rather than swapping or compressing
- mmap+mlock: mmap + force system to keep model in RAM rather than swapping or compressing
- dio: use DirectIO if available

(env: LLAMA_ARG_LOAD_MODE) | +| `-lzm, --lazy-mode MODE` | on-demand reading of certain tensors, for example per-layer embeddings (default: auto)
- on: read the rows of such tensors from disk on demand instead of keeping them resident (requires mmap)
- auto: on, but only for tensors larger than 4 GiB
- off: always keep them resident
(env: LLAMA_ARG_LAZY_MODE) | | `--numa TYPE` | attempt optimizations that help on some NUMA systems
- distribute: spread execution evenly over all nodes
- isolate: only spawn threads on CPUs on the node that execution started on
- numactl: use the CPU map provided by numactl
if run without this previously, it is recommended to drop the system page cache before using this
see https://github.com/ggml-org/llama.cpp/issues/1437
(env: LLAMA_ARG_NUMA) | | `-dev, --device ` | comma-separated list of devices to use for offloading (none = don't offload)
use --list-devices to see a list of available devices
(env: LLAMA_ARG_DEVICE) | | `--list-devices` | print list of available devices and exit | | `-ot, --override-tensor =,...` | override tensor buffer type
(env: LLAMA_ARG_OVERRIDE_TENSOR) | | `-cmoe, --cpu-moe` | keep all Mixture of Experts (MoE) weights in the CPU
(env: LLAMA_ARG_CPU_MOE) | | `-ncmoe, --n-cpu-moe N` | keep the Mixture of Experts (MoE) weights of the first N layers in the CPU
(env: LLAMA_ARG_N_CPU_MOE) | +| `-ncffn, --n-cpu-ffn N` | keep the dense FFN weights of the first N layers in the CPU
(dense models; for MoE expert weights use --n-cpu-moe)
(env: LLAMA_ARG_N_CPU_FFN) | | `-ngl, --gpu-layers, --n-gpu-layers N` | max. number of layers to store in VRAM, either an exact number, 'auto', or 'all' (default: auto)
(env: LLAMA_ARG_N_GPU_LAYERS) | | `-sm, --split-mode {none,layer,row,tensor}` | how to split the model across multiple GPUs, one of:
- none: use one GPU only
- layer (default): split layers and KV across GPUs (pipelined)
- row: split weight across GPUs by rows (parallelized)
- tensor: split weights and KV across GPUs (parallelized, EXPERIMENTAL)
(env: LLAMA_ARG_SPLIT_MODE) | | `-ts, --tensor-split N0,N1,N2,...` | fraction of the model to offload to each GPU, comma-separated list of proportions, e.g. 3,1
(env: LLAMA_ARG_TENSOR_SPLIT) | diff --git a/tools/llama-bench/README.md b/tools/llama-bench/README.md index d53978548a16..6431e7d6d4aa 100644 --- a/tools/llama-bench/README.md +++ b/tools/llama-bench/README.md @@ -67,6 +67,7 @@ test parameters: -nkvo, --no-kv-offload <0|1> (default: 0) -fa, --flash-attn (default: auto) -dev, --device (default: auto) + -lzm, --lazy-mode (default: auto) -mmp, --mmap <0|1> (default: 1) -dio, --direct-io <0|1> (default: 0) -embd, --embeddings <0|1> (default: 0) diff --git a/tools/llama-bench/llama-bench.cpp b/tools/llama-bench/llama-bench.cpp index c17a27b54019..72868aff73ff 100644 --- a/tools/llama-bench/llama-bench.cpp +++ b/tools/llama-bench/llama-bench.cpp @@ -271,6 +271,19 @@ static const char * split_mode_str(llama_split_mode mode) { } } +static const char * lazy_mode_str(llama_lazy_mode mode) { + switch (mode) { + case LLAMA_LAZY_MODE_OFF: + return "off"; + case LLAMA_LAZY_MODE_AUTO: + return "auto"; + case LLAMA_LAZY_MODE_ON: + return "on"; + default: + GGML_ABORT("invalid lazy mode"); + } +} + static std::string pair_str(const std::pair & p) { static char buf[32]; snprintf(buf, sizeof(buf), "%d,%d", p.first, p.second); @@ -341,6 +354,7 @@ struct cmd_params { std::vector n_cpu_moe; std::vector split_mode; std::vector load_mode; + std::vector lazy_mode; std::vector main_gpu; std::vector no_kv_offload; std::vector flash_attn; @@ -385,6 +399,7 @@ static const cmd_params cmd_params_defaults = { /* n_cpu_moe */ { 0 }, /* split_mode */ { LLAMA_SPLIT_MODE_LAYER }, /* load_mode */ { LLAMA_LOAD_MODE_MMAP }, + /* lazy_mode */ { LLAMA_LAZY_MODE_AUTO }, /* main_gpu */ { 0 }, /* no_kv_offload */ { false }, /* flash_attn */ { LLAMA_FLASH_ATTN_TYPE_AUTO }, @@ -460,6 +475,7 @@ static void print_usage(int /* argc */, char ** argv) { printf(" -fa, --flash-attn (default: %s)\n", join(transform_to_str(cmd_params_defaults.flash_attn, llama_flash_attn_type_name), ",").c_str()); printf(" -dev, --device (default: auto)\n"); printf(" -lm, --load-mode (default: %s)\n", join(transform_to_str(cmd_params_defaults.load_mode, llama_load_mode_name), ",").c_str()); + printf(" -lzm, --lazy-mode (default: %s)\n", join(transform_to_str(cmd_params_defaults.lazy_mode, lazy_mode_str), ",").c_str()); printf(" -mmp, --mmap <0|1> (DEPRECATED IN FAVOUR OF --load-mode)\n"); printf(" -dio, --direct-io <0|1> (DEPRECATED IN FAVOUR OF --load-mode)\n"); printf(" -embd, --embeddings <0|1> (default: %s)\n", join(cmd_params_defaults.embeddings, ",").c_str()); @@ -784,6 +800,32 @@ static cmd_params parse_cmd_params(int argc, char ** argv) { break; } params.load_mode.insert(params.load_mode.end(), modes.begin(), modes.end()); + } else if (arg == "-lzm" || arg == "--lazy-mode") { + if (++i >= argc) { + invalid_param = true; + break; + } + auto p = string_split(argv[i], split_delim); + + std::vector modes; + for (const auto & m : p) { + llama_lazy_mode mode; + if (m == "on") { + mode = LLAMA_LAZY_MODE_ON; + } else if (m == "auto") { + mode = LLAMA_LAZY_MODE_AUTO; + } else if (m == "off") { + mode = LLAMA_LAZY_MODE_OFF; + } else { + invalid_param = true; + break; + } + modes.push_back(mode); + } + if (invalid_param) { + break; + } + params.lazy_mode.insert(params.lazy_mode.end(), modes.begin(), modes.end()); } else if (arg == "-mg" || arg == "--main-gpu") { if (++i >= argc) { invalid_param = true; @@ -1135,6 +1177,9 @@ static cmd_params parse_cmd_params(int argc, char ** argv) { if (params.load_mode.empty()) { params.load_mode = cmd_params_defaults.load_mode; } + if (params.lazy_mode.empty()) { + params.lazy_mode = cmd_params_defaults.lazy_mode; + } if (params.main_gpu.empty()) { params.main_gpu = cmd_params_defaults.main_gpu; } @@ -1201,6 +1246,7 @@ struct cmd_params_instance { int n_cpu_moe; llama_split_mode split_mode; llama_load_mode load_mode; + llama_lazy_mode lazy_mode; int main_gpu; bool no_kv_offload; llama_flash_attn_type flash_attn; @@ -1222,9 +1268,14 @@ struct cmd_params_instance { } mparams.split_mode = split_mode; mparams.load_mode = load_mode; + mparams.lazy_mode = lazy_mode; mparams.main_gpu = main_gpu; mparams.tensor_split = tensor_split.data(); mparams.no_host = no_host; + // LLAMA_BENCH_NO_REPACK=1: no CPU repack copies (big file-backed tensors stay mapped) + if (const char * e = getenv("LLAMA_BENCH_NO_REPACK"); e != nullptr && atoi(e) != 0) { + mparams.use_extra_bufts = false; + } if (n_cpu_moe <= 0) { if (tensor_buft_overrides.empty()) { @@ -1269,7 +1320,8 @@ struct cmd_params_instance { return model == other.model && n_gpu_layers == other.n_gpu_layers && n_cpu_moe == other.n_cpu_moe && split_mode == other.split_mode && main_gpu == other.main_gpu && tensor_split == other.tensor_split && - load_mode == other.load_mode && devices == other.devices && no_host == other.no_host && + load_mode == other.load_mode && lazy_mode == other.lazy_mode && + devices == other.devices && no_host == other.no_host && vec_tensor_buft_override_equal(tensor_buft_overrides, other.tensor_buft_overrides); } @@ -1303,6 +1355,7 @@ static std::vector get_cmd_params_instances(const cmd_param for (const auto & ncmoe : params.n_cpu_moe) for (const auto & sm : params.split_mode) for (const auto & lm : params.load_mode) + for (const auto & lzm : params.lazy_mode) for (const auto & mg : params.main_gpu) for (const auto & devs : params.devices) for (const auto & ts : params.tensor_split) @@ -1342,6 +1395,7 @@ static std::vector get_cmd_params_instances(const cmd_param /* .n_cpu_moe = */ ncmoe, /* .split_mode = */ sm, /* .load_mode = */ lm, + /* .lazy_mode = */ lzm, /* .main_gpu = */ mg, /* .no_kv_offload = */ nkvo, /* .flash_attn = */ fa, @@ -1378,6 +1432,7 @@ static std::vector get_cmd_params_instances(const cmd_param /* .n_cpu_moe = */ ncmoe, /* .split_mode = */ sm, /* .load_mode = */ lm, + /* .lazy_mode = */ lzm, /* .main_gpu = */ mg, /* .no_kv_offload = */ nkvo, /* .flash_attn = */ fa, @@ -1414,6 +1469,7 @@ static std::vector get_cmd_params_instances(const cmd_param /* .n_cpu_moe = */ ncmoe, /* .split_mode = */ sm, /* .load_mode = */ lm, + /* .lazy_mode = */ lzm, /* .main_gpu = */ mg, /* .no_kv_offload = */ nkvo, /* .flash_attn = */ fa, @@ -1455,6 +1511,7 @@ struct test { int n_cpu_moe; llama_split_mode split_mode; llama_load_mode load_mode; + llama_lazy_mode lazy_mode; int main_gpu; bool no_kv_offload; llama_flash_attn_type flash_attn; @@ -1494,6 +1551,7 @@ struct test { n_cpu_moe = inst.n_cpu_moe; split_mode = inst.split_mode; load_mode = inst.load_mode; + lazy_mode = inst.lazy_mode; main_gpu = inst.main_gpu; no_kv_offload = inst.no_kv_offload; flash_attn = inst.flash_attn; @@ -1561,7 +1619,8 @@ struct test { "n_ubatch", "n_threads", "cpu_mask", "cpu_strict", "poll", "type_k", "type_v", "n_gpu_layers", "n_cpu_moe", "split_mode", "main_gpu", "no_kv_offload", "flash_attn", "devices", "tensor_split", - "tensor_buft_overrides", "load_mode", "embeddings", + "tensor_buft_overrides", "load_mode", "lazy_mode", + "embeddings", "no_op_offload", "no_host", "fit_target", "fit_min_ctx", "n_prompt", "n_gen", "n_depth", "test_time", "avg_ns", "stddev_ns", "avg_ts", "stddev_ts" @@ -1586,7 +1645,7 @@ struct test { if (field == "avg_ts" || field == "stddev_ts") { return FLOAT; } - if (field == "load_mode") { + if (field == "load_mode" || field == "lazy_mode") { return STRING; } return STRING; @@ -1656,6 +1715,7 @@ struct test { tensor_split_str, tensor_buft_overrides_str, llama_load_mode_name(load_mode), + lazy_mode_str(lazy_mode), std::to_string(embeddings), std::to_string(no_op_offload), std::to_string(no_host), @@ -1667,8 +1727,8 @@ struct test { test_time, std::to_string(avg_ns()), std::to_string(stdev_ns()), - std::to_string(avg_ts()), - std::to_string(stdev_ts()) }; + string_format("%f", avg_ts()), // std::to_string format before C++26 + string_format("%f", stdev_ts()) }; return values; } @@ -1970,6 +2030,9 @@ struct markdown_printer : public printer { if (params.load_mode.size() > 1 || params.load_mode != cmd_params_defaults.load_mode) { fields.emplace_back("load_mode"); } + if (params.lazy_mode.size() > 1 || params.lazy_mode != cmd_params_defaults.lazy_mode) { + fields.emplace_back("lazy_mode"); + } if (params.embeddings.size() > 1 || params.embeddings != cmd_params_defaults.embeddings) { fields.emplace_back("embeddings"); } diff --git a/tools/mtmd/clip-graph.h b/tools/mtmd/clip-graph.h index 29352abb4c0b..de4d0e741fc3 100644 --- a/tools/mtmd/clip-graph.h +++ b/tools/mtmd/clip-graph.h @@ -9,7 +9,7 @@ #include #include -#define DEFAULT_INTERPOLATION_MODE (GGML_SCALE_MODE_BILINEAR | GGML_SCALE_FLAG_ANTIALIAS) +#define DEFAULT_INTERPOLATION_MODE (GGML_SCALE_MODE_BILINEAR | (uint32_t) GGML_SCALE_FLAG_ANTIALIAS) struct build_vit_opts { ggml_tensor * attn_mask = nullptr; diff --git a/tools/mtmd/clip-impl.h b/tools/mtmd/clip-impl.h index d42b38222c27..483e3ee37618 100644 --- a/tools/mtmd/clip-impl.h +++ b/tools/mtmd/clip-impl.h @@ -65,6 +65,7 @@ #define KEY_MM_PATCH_MERGE_TYPE "clip.vision.mm_patch_merge_type" #define KEY_IMAGE_GRID_PINPOINTS "clip.vision.image_grid_pinpoints" #define KEY_WIN_ATTN_PATTERN "clip.vision.n_wa_pattern" +#define KEY_DECODE_NON_CAUSAL "clip.vision.decode_non_causal" // the text model attends to an image bidirectionally #define KEY_WIN_ATTN_LAYER_INDEXES "clip.vision.wa_layer_indexes" #define KEY_WA_PATTERN_MODE "clip.vision.wa_pattern_mode" #define KEY_ATTN_WINDOW_SIZE "clip.vision.window_size" @@ -797,8 +798,8 @@ static std::string gguf_data_to_str(enum gguf_type type, const void * data, int case GGUF_TYPE_INT32: return std::to_string(((const int32_t *)data)[i]); case GGUF_TYPE_UINT64: return std::to_string(((const uint64_t *)data)[i]); case GGUF_TYPE_INT64: return std::to_string(((const int64_t *)data)[i]); - case GGUF_TYPE_FLOAT32: return std::to_string(((const float *)data)[i]); - case GGUF_TYPE_FLOAT64: return std::to_string(((const double *)data)[i]); + case GGUF_TYPE_FLOAT32: return string_format("%f", (double) ((const float *)data)[i]); // std::to_string format before C++26 + case GGUF_TYPE_FLOAT64: return string_format("%f", ((const double *)data)[i]); case GGUF_TYPE_BOOL: return ((const int8_t *)data)[i] != 0 ? "true" : "false"; default: return string_format("unknown type %d", type); } diff --git a/tools/mtmd/clip-model.h b/tools/mtmd/clip-model.h index 8b9db5101d2c..eea008ed27bc 100644 --- a/tools/mtmd/clip-model.h +++ b/tools/mtmd/clip-model.h @@ -96,6 +96,7 @@ struct clip_hparams { std::vector feature_layers; int32_t attn_window_size = 0; int32_t n_wa_pattern = 0; + bool decode_non_causal = false; // KEY_DECODE_NON_CAUSAL std::unordered_set wa_layer_indexes; // explicit layer indexes that use full attention (for irregular patterns like YoutuVL) std::vector wa_pattern_mode; // mimovl: per-layer window-attention mode diff --git a/tools/mtmd/clip.cpp b/tools/mtmd/clip.cpp index c1870813fb93..03e6b0c82940 100644 --- a/tools/mtmd/clip.cpp +++ b/tools/mtmd/clip.cpp @@ -1250,6 +1250,8 @@ struct clip_model_loader { // default warmup value hparams.warmup_image_size = hparams.image_size; + get_bool(KEY_DECODE_NON_CAUSAL, hparams.decode_non_causal, false); + { bool use_gelu = false; bool use_silu = false; @@ -5073,6 +5075,10 @@ bool clip_support_batch(const struct clip_ctx * ctx) { // TODO @ngxson : this is no longer correct with mtmd_batch API // this was only meant to be used by qwen-vl-based models, to fuse 2 input images into one (qwen-vl video support) // this logic should be refactored in near future to distinctly handle "merge frames" and "batching" +bool clip_decode_non_causal(const struct clip_ctx * ctx) { + return ctx->model.hparams.decode_non_causal; +} + int clip_model_n_temporal_merge(const struct clip_ctx * ctx) { switch (ctx->proj_type()) { case PROJECTOR_TYPE_QWEN2VL: diff --git a/tools/mtmd/clip.h b/tools/mtmd/clip.h index 967093a812d6..c3c584db954d 100644 --- a/tools/mtmd/clip.h +++ b/tools/mtmd/clip.h @@ -95,6 +95,9 @@ bool clip_support_batch(const struct clip_ctx * ctx); int clip_model_n_temporal_merge(const struct clip_ctx * ctx); // TODO @ngxson : remove, refactor this +// the mmproj asks for the text model to attend to its images bidirectionally (clip.vision.decode_non_causal) +bool clip_decode_non_causal(const struct clip_ctx * ctx); + std::map clip_get_mem_usage(const struct clip_ctx * ctx); struct clip_cap { diff --git a/tools/mtmd/models/qwen3vl.cpp b/tools/mtmd/models/qwen3vl.cpp index 48626b221fbb..2c707686626a 100644 --- a/tools/mtmd/models/qwen3vl.cpp +++ b/tools/mtmd/models/qwen3vl.cpp @@ -37,7 +37,7 @@ ggml_cgraph * clip_graph_qwen3vl::build() { } // calculate absolute position embedding and apply - ggml_tensor * learned_pos_embd = resize_position_embeddings(GGML_SCALE_MODE_BILINEAR | GGML_SCALE_FLAG_ALIGN_CORNERS); + ggml_tensor * learned_pos_embd = resize_position_embeddings(GGML_SCALE_MODE_BILINEAR | (uint32_t) GGML_SCALE_FLAG_ALIGN_CORNERS); learned_pos_embd = ggml_cont_4d( ctx0, learned_pos_embd, n_embd * 2, n_patches_x / 2, n_patches_y, batch_size); diff --git a/tools/mtmd/mtmd.cpp b/tools/mtmd/mtmd.cpp index d3899f5c853d..56e3967d6d5f 100644 --- a/tools/mtmd/mtmd.cpp +++ b/tools/mtmd/mtmd.cpp @@ -1692,7 +1692,8 @@ bool mtmd_decode_use_non_causal(const mtmd_context * ctx, const mtmd_input_chunk case PROJECTOR_TYPE_GEMMA4UV: return true; default: - return false; + // any vision projector can ask for it (ZAYA1-VL: a Qwen2.5-VL tower) + return chunk && chunk->type == MTMD_INPUT_CHUNK_TYPE_IMAGE && ctx->ctx_v && clip_decode_non_causal(ctx->ctx_v); } } diff --git a/tools/server/CMakeLists.txt b/tools/server/CMakeLists.txt index 280bd9e19dca..e3e202379cdc 100644 --- a/tools/server/CMakeLists.txt +++ b/tools/server/CMakeLists.txt @@ -15,6 +15,8 @@ add_library(${TARGET} STATIC server-common.h server-context.cpp server-context.h + server-prefill.cpp + server-prefill.h server-stream.cpp server-stream.h server-tools.cpp diff --git a/tools/server/README.md b/tools/server/README.md index f45c018972d2..a384e79db429 100644 --- a/tools/server/README.md +++ b/tools/server/README.md @@ -76,12 +76,14 @@ For the full list of features, please refer to [server's changelog](https://gith | `--mmap, --no-mmap` | DEPRECATED in favor of `--load-mode`: whether to memory-map model. (if mmap disabled, slower load but may reduce pageouts if not using mlock)
(env: LLAMA_ARG_MMAP) | | `-dio, --direct-io, -ndio, --no-direct-io` | DEPRECATED in favor of `--load-mode`: use DirectIO if available
(env: LLAMA_ARG_DIO) | | `-lm, --load-mode MODE` | model loading mode (default: mmap)
- none: no special loading mode
- mmap: memory-map model (if mmap disabled, slower load but may reduce pageouts if not using mlock)
- mlock: force system to keep model in RAM rather than swapping or compressing
- mmap+mlock: mmap + force system to keep model in RAM rather than swapping or compressing
- dio: use DirectIO if available

(env: LLAMA_ARG_LOAD_MODE) | +| `-lzm, --lazy-mode MODE` | on-demand reading of certain tensors, for example per-layer embeddings (default: auto)
- on: read the rows of such tensors from disk on demand instead of keeping them resident (requires mmap)
- auto: on, but only for tensors larger than 4 GiB
- off: always keep them resident
(env: LLAMA_ARG_LAZY_MODE) | | `--numa TYPE` | attempt optimizations that help on some NUMA systems
- distribute: spread execution evenly over all nodes
- isolate: only spawn threads on CPUs on the node that execution started on
- numactl: use the CPU map provided by numactl
if run without this previously, it is recommended to drop the system page cache before using this
see https://github.com/ggml-org/llama.cpp/issues/1437
(env: LLAMA_ARG_NUMA) | | `-dev, --device ` | comma-separated list of devices to use for offloading (none = don't offload)
use --list-devices to see a list of available devices
(env: LLAMA_ARG_DEVICE) | | `--list-devices` | print list of available devices and exit | | `-ot, --override-tensor =,...` | override tensor buffer type
(env: LLAMA_ARG_OVERRIDE_TENSOR) | | `-cmoe, --cpu-moe` | keep all Mixture of Experts (MoE) weights in the CPU
(env: LLAMA_ARG_CPU_MOE) | | `-ncmoe, --n-cpu-moe N` | keep the Mixture of Experts (MoE) weights of the first N layers in the CPU
(env: LLAMA_ARG_N_CPU_MOE) | +| `-ncffn, --n-cpu-ffn N` | keep the dense FFN weights of the first N layers in the CPU
(dense models; for MoE expert weights use --n-cpu-moe)
(env: LLAMA_ARG_N_CPU_FFN) | | `-ngl, --gpu-layers, --n-gpu-layers N` | max. number of layers to store in VRAM, either an exact number, 'auto', or 'all' (default: auto)
(env: LLAMA_ARG_N_GPU_LAYERS) | | `-sm, --split-mode {none,layer,row,tensor}` | how to split the model across multiple GPUs, one of:
- none: use one GPU only
- layer (default): split layers and KV across GPUs (pipelined)
- row: split weight across GPUs by rows (parallelized)
- tensor: split weights and KV across GPUs (parallelized, EXPERIMENTAL)
(env: LLAMA_ARG_SPLIT_MODE) | | `-ts, --tensor-split N0,N1,N2,...` | fraction of the model to offload to each GPU, comma-separated list of proportions, e.g. 3,1
(env: LLAMA_ARG_TENSOR_SPLIT) | diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index 4655b518e21f..fe77d13aaa78 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -1,4 +1,5 @@ #include "server-context.h" +#include "server-prefill.h" #include "server-chat.h" #include "server-common.h" #include "server-http.h" @@ -931,6 +932,8 @@ struct server_context_impl { llama_context * ctx_tgt = nullptr; + server_prefill_device prefill_dev; // [1bit] ONEBIT_PREFILL_DEVICE (server-prefill.h) + server_batch batch; llama_model * model_dft = nullptr; @@ -976,6 +979,8 @@ struct server_context_impl { int64_t t_last_load_progress_ms = 0; void destroy() { + prefill_dev.reset(); + spec.reset(); spec_init.reset(); @@ -1199,6 +1204,10 @@ struct server_context_impl { return false; } + if (const char * pf = getenv("ONEBIT_PREFILL_DEVICE"); pf != nullptr && *pf != '\0') { + prefill_dev.load(pf, params_base, ctx_tgt); + } + vocab = llama_model_get_vocab(model_tgt); n_ctx = llama_n_ctx(ctx_tgt); @@ -2856,6 +2865,11 @@ struct server_context_impl { // TODO @ngxson : maybe handle n_batch == 1 here instead of inside decode() batch_view = batch.get_view(off, n_tokens); + // [1bit] the prompt prefix runs on the prefill device, the rest of the view on ctx_tgt + if (const int32_t k = spec || mctx != nullptr ? 0 : prefill_dev.prefill(ctx_tgt, batch_view); k > 0) { + off_next = off + k; + continue; + } bool ok = decode(n_batch, off, batch_view); #ifdef DEBUG_TIMINGS llama_synchronize(ctx_tgt); diff --git a/tools/server/server-prefill.cpp b/tools/server/server-prefill.cpp new file mode 100644 index 000000000000..7da560ce83fa --- /dev/null +++ b/tools/server/server-prefill.cpp @@ -0,0 +1,93 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "server-prefill.h" + +#include "../../src/llama-ext.h" +#include "log.h" + +#include +#include +#include +#include + +bool server_prefill_device::load(const char * device_name, common_params & params, llama_context * ctx_tgt) { + ggml_backend_dev_t dev = ggml_backend_dev_by_name(device_name); + if (dev == nullptr) { + LOG_WRN("prefill device '%s' not found, prefill stays on the main device\n", device_name); + return false; + } + if (!params.lora_adapters.empty()) { + LOG_WRN("%s", "prefill device is not used with LoRA adapters\n"); + return false; + } + if (const char * v = getenv("ONEBIT_PREFILL_MIN_TOKENS")) { + min_tokens = std::max(1, atoi(v)); + } + ggml_backend_dev_t devs[2] = { dev, nullptr }; + llama_model_params mparams = common_model_params_to_llama(params); + mparams.devices = devs; + mparams.progress_callback = nullptr; + model.reset(llama_model_load_from_file(params.model.path.c_str(), mparams)); + if (!model) { + LOG_WRN("failed to load the model on prefill device '%s'\n", device_name); + return false; + } + llama_context_params cparams = common_context_params_to_llama(params); + cparams.n_ctx = llama_n_ctx(ctx_tgt); + cparams.n_batch = llama_n_batch(ctx_tgt); + cparams.n_ubatch = llama_n_ubatch(ctx_tgt); + cparams.n_seq_max = llama_n_seq_max(ctx_tgt); + llama_kv_share_next(true); // this context owns the KV region; ctx_tgt maps it + try { + ctx.reset(llama_init_from_model(model.get(), cparams)); + } catch (const std::exception & e) { + LOG_WRN("prefill context: %s\n", e.what()); + } + if (!ctx || !llama_kv_share_from(ctx_tgt, ctx.get()) || !llama_kv_cells_copy(ctx.get(), ctx_tgt)) { + LOG_WRN("prefill device '%s' cannot share the KV cache (set -fa on explicitly so both devices lay it out the same way)\n", device_name); + reset(); + return false; + } + LOG_INF("prompt prefill on %s (from %d tokens), decode on the main device; KV cache shared, zero copy\n", device_name, min_tokens); + return true; +} + +int32_t server_prefill_device::prefill(llama_context * ctx_tgt, const llama_batch & view) { + if (!ctx || view.embd != nullptr || view.logits == nullptr) { + return 0; + } + int32_t k = 0; + while (k < view.n_tokens && !view.logits[k]) { + k++; + } + // HRX prefill halves on a partial last ubatch (Qwen2.5-7B: 1085 tok/s at 2047 tokens, 2246 at 2048): whole + // ubatches only, the remainder runs on the main device with the token that needs logits + k -= k % (int32_t) llama_n_ubatch(ctx.get()); + if (k < min_tokens || !llama_kv_cells_copy(ctx.get(), ctx_tgt)) { + return 0; + } + llama_batch prefix = view; + prefix.n_tokens = k; + if (llama_decode(ctx.get(), prefix) != 0) { + LOG_WRN("prefill device decode failed (n = %d), using the main device\n", k); + return 0; + } + llama_synchronize(ctx.get()); + if (!llama_kv_cells_copy(ctx_tgt, ctx.get())) { + throw std::runtime_error("prefill device: KV cells handoff failed"); + } + return k; +} diff --git a/tools/server/server-prefill.h b/tools/server/server-prefill.h new file mode 100644 index 000000000000..11a6f7b62b17 --- /dev/null +++ b/tools/server/server-prefill.h @@ -0,0 +1,43 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include "common.h" +#include "llama-cpp.h" +#include "llama.h" + +#include + +// ONEBIT_PREFILL_DEVICE=: prompt prefill runs on a second device of the same GPU (for example HRX0 while the +// server decodes on Vulkan0). Both contexts use one KV cache (dma-buf, zero copy); only the cell metadata moves. +// ONEBIT_PREFILL_MIN_TOKENS sets the shortest prompt prefix sent there (default 1024: below it the split does not pay). +struct server_prefill_device { + llama_model_ptr model; + llama_context_ptr ctx; + int32_t min_tokens = 1024; + + // loads a second copy of the model on device_name; its KV region becomes ctx_tgt's KV cache + bool load(const char * device_name, common_params & params, llama_context * ctx_tgt); + + void reset() { + ctx.reset(); + model.reset(); + } + + // runs the prompt prefix of view (the tokens before the first one that needs logits) on this device; + // returns how many tokens it ran, 0 when view stays on ctx_tgt + int32_t prefill(llama_context * ctx_tgt, const llama_batch & view); +}; diff --git a/tools/server/server-schema.cpp b/tools/server/server-schema.cpp index 674d3ba337bc..ca2ce70bf9ca 100644 --- a/tools/server/server-schema.cpp +++ b/tools/server/server-schema.cpp @@ -588,6 +588,16 @@ static void handle_with_catch(const char * name, std::function func) { } } +// "%f" for floating point: the std::to_string format before C++26 +template +static std::string num_to_str(T v) { + if constexpr (std::is_floating_point_v) { + return string_format("%f", (double) v); + } else { + return std::to_string(v); + } +} + // treat a null value as absent so clients can send null to request the server default static bool has_value(const json & data, const char * n) { auto it = data.find(n); @@ -606,7 +616,7 @@ void field_num::eval(field_eval_context & ctx, const json & data) { } else { T tmp = data.at(n).template get(); if (tmp < min || tmp > max) { - throw std::invalid_argument(std::string("Value must be between ") + std::to_string(min) + " <= value <= " + std::to_string(max) + ", but got " + std::to_string(tmp)); + throw std::invalid_argument(std::string("Value must be between ") + num_to_str(min) + " <= value <= " + num_to_str(max) + ", but got " + num_to_str(tmp)); } val = tmp; } diff --git a/vendor/nlohmann/json.hpp b/vendor/nlohmann/json.hpp index 82d69f7c5d04..bc750be26f18 100644 --- a/vendor/nlohmann/json.hpp +++ b/vendor/nlohmann/json.hpp @@ -17554,7 +17554,7 @@ class binary_writer static CharType to_char_type(std::uint8_t x) noexcept { static_assert(sizeof(std::uint8_t) == sizeof(CharType), "size of CharType must be equal to std::uint8_t"); - static_assert(std::is_trivial::value, "CharType must be trivial"); + static_assert(std::is_trivially_default_constructible::value && std::is_trivially_copyable::value, "CharType must be trivial"); CharType result; std::memcpy(&result, &x, sizeof(x)); return result;